from __future__ import annotations

import csv
from dataclasses import dataclass, field
from pathlib import Path

from .csv_dialect import detect_dialect
from .row_hash import compute_row_hash


@dataclass(frozen=True)
class DryRunResult:
    rows_total: int
    rows_valid: int
    rows_failed: int
    row_hashes: list[str]
    errors: list[str]
    rows_new: int = 0
    rows_existing: int = 0
    rows_duplicate: int = 0
    row_statuses: list[str] = field(default_factory=list)


def read_csv_rows(path: str | Path, required_columns: set[str]) -> tuple[list[dict[str, str]], list[str]]:
    path = Path(path)
    sample = path.read_text(encoding="utf-8")[:4096]
    dialect = detect_dialect(sample)
    errors: list[str] = []
    with path.open(newline="", encoding="utf-8") as f:
        reader = csv.DictReader(f, dialect=dialect)
        fieldnames = set(reader.fieldnames or [])
        missing = required_columns - fieldnames
        if missing:
            return [], [f"Missing columns: {sorted(missing)}"]
        rows = [dict(row) for row in reader]
    return rows, errors


def dry_run_csv(
    path: str | Path,
    required_columns: set[str],
    *,
    existing_hashes: set[str] | None = None,
) -> DryRunResult:
    rows, errors = read_csv_rows(path, required_columns)
    if errors:
        return DryRunResult(0, 0, 0, [], errors)
    existing_hashes = existing_hashes or set()
    seen: set[str] = set()
    row_hashes: list[str] = []
    row_statuses: list[str] = []
    rows_new = rows_existing = rows_duplicate = 0
    for index, row in enumerate(rows, start=2):
        row_hash = compute_row_hash(row)
        row_hashes.append(row_hash)
        if row_hash in seen:
            rows_duplicate += 1
            row_statuses.append("duplicate")
            errors.append(f"Duplicate row hash at CSV line {index}")
        elif row_hash in existing_hashes:
            rows_existing += 1
            row_statuses.append("existing")
        else:
            rows_new += 1
            row_statuses.append("new")
        seen.add(row_hash)
    return DryRunResult(
        rows_total=len(rows),
        rows_valid=len(rows) - len(errors),
        rows_failed=len(errors),
        row_hashes=row_hashes,
        errors=errors,
        rows_new=rows_new,
        rows_existing=rows_existing,
        rows_duplicate=rows_duplicate,
        row_statuses=row_statuses,
    )
