#!/usr/bin/env python3
"""
csv-roundtrip / transform.py  (version 1.0, 2026-09-17)

Export -> transform -> validate -> import, with the transform written down in a
rules file instead of done by hand in a spreadsheet.

    python3 transform.py --in export.csv --rules rules.json --out changed.csv --dry-run
    python3 transform.py --in export.csv --rules rules.json --out changed.csv
    python3 transform.py --validate --in export.csv --out changed.csv --rules rules.json

--dry-run   applies the rules and prints what WOULD change; writes nothing.
--validate  compares an existing --out file against --in and the rules and exits
            non-zero if anything other than the intended change happened.

Standard library only. Tested with Python 3.11 and 3.12 on the bundled samples.
License: MIT. Built by the Distru team for the No Bullshit AI Course, module 05.
No credentials, no customer data: the samples use fake ids and fake names.
"""
from __future__ import annotations

import argparse
import csv
import json
import sys
from collections import Counter
from pathlib import Path


# ----------------------------------------------------------------------------- rules


def load_rules(path: Path) -> dict:
    rules = json.loads(path.read_text(encoding="utf-8"))
    for required in ("key", "changes"):
        if required not in rules:
            sys.exit(f"rules: missing required field '{required}'")
    if not isinstance(rules["changes"], list) or not rules["changes"]:
        sys.exit("rules: 'changes' must be a non-empty list")
    for i, change in enumerate(rules["changes"]):
        if "column" not in change:
            sys.exit(f"rules: changes[{i}] has no 'column'")
        kinds = [k for k in ("set", "map") if k in change]
        if len(kinds) != 1:
            sys.exit(f"rules: changes[{i}] needs exactly one of 'set' or 'map'")
    return rules


def row_matches(row: dict, where: dict | None) -> bool:
    """where = {"column": ["allowed", "values"]} ; all columns must match (AND)."""
    if not where:
        return True
    for col, allowed in where.items():
        if col not in row:
            sys.exit(f"rules: 'where' names a column that is not in the file: {col}")
        if row[col] not in allowed:
            return False
    return True


def apply_rules(rows: list[dict], rules: dict) -> tuple[list[dict], Counter, list[str]]:
    """Return (new_rows, per-column change counts, unmapped values seen)."""
    out: list[dict] = []
    changed: Counter = Counter()
    unmapped: list[str] = []
    for row in rows:
        new = dict(row)
        for change in rules["changes"]:
            col = change["column"]
            if col not in new:
                sys.exit(f"rules: column '{col}' is not in the input file")
            if not row_matches(row, change.get("where")):
                continue
            if "set" in change:
                target = change["set"]
            else:  # map
                mapping = change["map"]
                if row[col] in mapping:
                    target = mapping[row[col]]
                elif change.get("keep_unmapped", True):
                    target = row[col]
                    if row[col] not in unmapped:
                        unmapped.append(row[col])
                else:
                    sys.exit(f"rules: value '{row[col]}' in column '{col}' has no mapping")
            if new[col] != target:
                new[col] = target
                changed[col] += 1
        out.append(new)
    return out, changed, unmapped


# ------------------------------------------------------------------------- validation


def validate(original: list[dict], changed: list[dict], header_in: list[str],
             header_out: list[str], rules: dict) -> list[str]:
    """Every check that must pass before you upload. Returns a list of problems."""
    problems: list[str] = []
    key = rules["key"]
    allowed_cols = {c["column"] for c in rules["changes"]}
    required = rules.get("required_columns", [])

    if header_in != header_out:
        problems.append(f"header changed: {header_in} -> {header_out}")
    for col in [key, *required]:
        if col not in header_out:
            problems.append(f"required column missing from output: {col}")
    if len(original) != len(changed):
        problems.append(f"row count changed: {len(original)} -> {len(changed)}")
        return problems  # nothing below is meaningful after this

    keys_in = [r[key] for r in original]
    dup = [k for k, n in Counter(keys_in).items() if n > 1]
    if dup:
        problems.append(f"duplicate keys in input ({len(dup)}): {dup[:5]}")
    blank = sum(1 for k in keys_in if not k.strip())
    if blank:
        problems.append(f"{blank} row(s) have a blank key")

    for i, (a, b) in enumerate(zip(original, changed), start=2):  # line 1 is the header
        if a[key] != b[key]:
            problems.append(f"line {i}: key changed {a[key]!r} -> {b[key]!r}")
        for col in header_in:
            if col in allowed_cols:
                continue
            if a.get(col) != b.get(col):
                problems.append(f"line {i}: column '{col}' changed but is not in the rules "
                                f"({a.get(col)!r} -> {b.get(col)!r})")
    for change in rules["changes"]:
        allowed_values = change.get("allowed_values")
        if allowed_values:
            bad = sorted({r[change['column']] for r in changed} - set(allowed_values))
            if bad:
                problems.append(f"column '{change['column']}' has values outside allowed_values: {bad}")
    return problems


# --------------------------------------------------------------------------------- io


def read_csv(path: Path) -> tuple[list[str], list[dict]]:
    with path.open(newline="", encoding="utf-8-sig") as f:
        reader = csv.DictReader(f)
        if reader.fieldnames is None:
            sys.exit(f"{path}: empty file")
        rows = list(reader)
    for r in rows:
        if None in r:
            sys.exit(f"{path}: a row has more cells than the header (stray comma?)")
    return list(reader.fieldnames), rows


def write_csv(path: Path, header: list[str], rows: list[dict]) -> None:
    with path.open("w", newline="", encoding="utf-8") as f:
        w = csv.DictWriter(f, fieldnames=header, lineterminator="\n")
        w.writeheader()
        w.writerows(rows)


def print_diff(original: list[dict], changed: list[dict], key: str, limit: int = 10) -> None:
    shown = 0
    for a, b in zip(original, changed):
        diffs = {c: (a[c], b[c]) for c in a if a[c] != b.get(c)}
        if not diffs:
            continue
        if shown < limit:
            parts = ", ".join(f"{c}: {x!r} -> {y!r}" for c, (x, y) in diffs.items())
            print(f"  {a[key]}: {parts}")
        shown += 1
    if shown > limit:
        print(f"  ... and {shown - limit} more row(s)")


# ------------------------------------------------------------------------------- main


def main() -> int:
    ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument("--in", dest="inp", required=True, type=Path, help="the export you downloaded")
    ap.add_argument("--rules", required=True, type=Path, help="rules.json (see rules.example.json)")
    ap.add_argument("--out", required=True, type=Path, help="the file you will import")
    ap.add_argument("--dry-run", action="store_true", help="show changes, write nothing")
    ap.add_argument("--validate", action="store_true", help="check an existing --out against --in")
    args = ap.parse_args()

    rules = load_rules(args.rules)
    header_in, original = read_csv(args.inp)
    key = rules["key"]
    if key not in header_in:
        sys.exit(f"key column '{key}' not in {args.inp}")

    if args.validate:
        if not args.out.exists():
            sys.exit(f"--validate: {args.out} does not exist")
        header_out, changed = read_csv(args.out)
        problems = validate(original, changed, header_in, header_out, rules)
        expected, counts, _ = apply_rules(original, rules)
        for i, (want, got) in enumerate(zip(expected, changed), start=2):
            if want != got:
                problems.append(f"line {i}: output differs from what the rules produce")
                break
        return report(problems, counts, len(original))

    changed, counts, unmapped = apply_rules(original, rules)
    problems = validate(original, changed, header_in, header_in, rules)
    if unmapped:
        print(f"note: {len(unmapped)} value(s) had no mapping and were left as-is: {unmapped[:10]}")

    print(f"{'DRY RUN: ' if args.dry_run else ''}{len(original)} rows in, {len(changed)} rows out")
    for col, n in counts.items():
        print(f"  {n} row(s) change in column '{col}'")
    if not counts:
        print("  nothing would change (check your 'where' values and case)")
    print_diff(original, changed, key)

    if problems:
        return report(problems, counts, len(original))
    if args.dry_run:
        print("OK (dry run). Re-run without --dry-run to write the file.")
        return 0
    write_csv(args.out, header_in, changed)
    print(f"wrote {args.out}. Now run --validate, then import.")
    return 0


def report(problems: list[str], counts: Counter, n_rows: int) -> int:
    if problems:
        print(f"STOP: {len(problems)} problem(s). Do not import this file.")
        for p in problems:
            print(f"  - {p}")
        return 1
    print(f"VALID: {n_rows} rows, only the intended column(s) changed: {dict(counts) or 'none'}")
    return 0


if __name__ == "__main__":
    sys.exit(main())
