#!/usr/bin/env python3
import argparse
import csv
import json
import sqlite3
from collections import Counter, defaultdict
from pathlib import Path


ROOT = Path(__file__).resolve().parents[1]
DEFAULT_CASES = ROOT / "data" / "support_split_cases.csv"
DEFAULT_FEATURES = ROOT / "data" / "feature_history.csv"
DEFAULT_ASSIGNMENTS = ROOT / "output" / "dataset_split_assignments.csv"
DEFAULT_POLICY = ROOT / "contracts" / "split_policy.json"
DEFAULT_OUTPUT = ROOT / "output"


def read_csv(path):
    with path.open(newline="", encoding="utf-8") as handle:
        return list(csv.DictReader(handle))


def read_json(path):
    with path.open(encoding="utf-8") as handle:
        return json.load(handle)


def write_json(path, payload):
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text(json.dumps(payload, indent=2, ensure_ascii=False) + "\n", encoding="utf-8")


def create_table(connection, name, rows):
    if not rows:
        raise ValueError(f"{name} no contiene filas")
    columns = list(rows[0])
    quoted_columns = ", ".join(f"{column} TEXT" for column in columns)
    connection.execute(f"CREATE TABLE {name} ({quoted_columns})")
    placeholders = ", ".join("?" for _ in columns)
    connection.executemany(
        f"INSERT INTO {name} ({', '.join(columns)}) VALUES ({placeholders})",
        [[row[column] for column in columns] for row in rows],
    )


def query_dicts(connection, sql, params=()):
    cursor = connection.execute(sql, params)
    columns = [description[0] for description in cursor.description]
    return [dict(zip(columns, row)) for row in cursor.fetchall()]


def build_connection(cases, features, assignments):
    connection = sqlite3.connect(":memory:")
    create_table(connection, "cases", cases)
    create_table(connection, "feature_history", features)
    create_table(connection, "split_assignments", assignments)
    return connection


def as_of_preview(connection):
    rows = []
    cases = query_dicts(
        connection,
        """
        SELECT case_id, student_id, created_at
        FROM cases
        ORDER BY case_id
        """,
    )
    for case in cases:
        selected = query_dicts(
            connection,
            """
            SELECT feature_available_at, open_cases_30d, last_payment_status
            FROM feature_history
            WHERE student_id = ?
              AND feature_available_at <= ?
            ORDER BY feature_available_at DESC
            LIMIT 1
            """,
            (case["student_id"], case["created_at"]),
        )
        item = dict(case)
        if selected:
            item.update(selected[0])
            item["missing_asof_feature"] = False
        else:
            item.update(
                {
                    "feature_available_at": None,
                    "open_cases_30d": None,
                    "last_payment_status": None,
                    "missing_asof_feature": True,
                }
            )
        rows.append(item)
    return rows


def test_label_coverage(connection, required_labels):
    rows = query_dicts(
        connection,
        """
        SELECT DISTINCT c.label
        FROM cases c
        JOIN split_assignments a USING(case_id)
        WHERE a.split = 'test'
        """,
    )
    labels = sorted(row["label"] for row in rows)
    return {
        "observed_test_labels": labels,
        "missing_required_test_labels": sorted(set(required_labels) - set(labels)),
    }


def build_report(cases, features, assignments, policy):
    connection = build_connection(cases, features, assignments)
    duplicate_assignments = query_dicts(
        connection,
        """
        SELECT case_id, COUNT(*) AS assignment_count
        FROM split_assignments
        GROUP BY case_id
        HAVING COUNT(*) <> 1
        ORDER BY case_id
        """,
    )
    missing_assignments = query_dicts(
        connection,
        """
        SELECT c.case_id
        FROM cases c
        LEFT JOIN split_assignments a USING(case_id)
        WHERE a.case_id IS NULL
        ORDER BY c.case_id
        """,
    )
    group_overlap = query_dicts(
        connection,
        """
        SELECT c.student_id,
               GROUP_CONCAT(DISTINCT a.split) AS splits,
               GROUP_CONCAT(c.case_id) AS case_ids
        FROM cases c
        JOIN split_assignments a USING(case_id)
        GROUP BY c.student_id
        HAVING SUM(CASE WHEN a.split = 'train' THEN 1 ELSE 0 END) > 0
           AND SUM(CASE WHEN a.split = 'test' THEN 1 ELSE 0 END) > 0
        ORDER BY c.student_id
        """,
    )
    source_overlap = query_dicts(
        connection,
        """
        SELECT c.source_id,
               GROUP_CONCAT(DISTINCT a.split) AS splits,
               GROUP_CONCAT(c.case_id) AS case_ids
        FROM cases c
        JOIN split_assignments a USING(case_id)
        GROUP BY c.source_id
        HAVING SUM(CASE WHEN a.split = 'train' THEN 1 ELSE 0 END) > 0
           AND SUM(CASE WHEN a.split = 'test' THEN 1 ELSE 0 END) > 0
        ORDER BY c.source_id
        """,
    )
    future_feature_candidates = query_dicts(
        connection,
        """
        SELECT c.case_id,
               c.student_id,
               c.created_at,
               f.feature_available_at,
               f.open_cases_30d,
               f.last_payment_status
        FROM cases c
        JOIN feature_history f USING(student_id)
        WHERE f.feature_available_at > c.created_at
        ORDER BY c.case_id, f.feature_available_at
        """,
    )
    asof_rows = as_of_preview(connection)
    missing_asof = [row["case_id"] for row in asof_rows if row["missing_asof_feature"]]
    coverage = test_label_coverage(connection, policy["required_test_labels"])

    blocking = []
    if duplicate_assignments:
        blocking.append("duplicate_case_assignment")
    if missing_assignments:
        blocking.append("case_missing_from_split")
    if group_overlap:
        blocking.append("student_train_test_overlap")
    if source_overlap:
        blocking.append("source_train_test_overlap")
    if missing_asof:
        blocking.append("missing_asof_features")

    review = []
    if future_feature_candidates:
        review.append("naive_join_would_see_future_features")
    if coverage["missing_required_test_labels"]:
        review.append("missing_required_test_labels")

    decision = "block" if blocking else ("review" if review else "pass")
    counts = Counter(item["split"] for item in assignments)

    return {
        "policy_id": policy["policy_id"],
        "split_version": assignments[0]["split_version"] if assignments else None,
        "split_counts": dict(sorted(counts.items())),
        "checks": {
            "duplicate_assignments": duplicate_assignments,
            "missing_assignments": missing_assignments,
            "student_train_test_overlap": group_overlap,
            "source_train_test_overlap": source_overlap,
            "future_feature_candidates": future_feature_candidates,
            "as_of_join_preview": asof_rows,
            "test_label_coverage": coverage,
        },
        "blocking_failures": blocking,
        "review_failures": review,
        "decision": decision,
    }


def write_decision(path, report):
    future_count = len(report["checks"]["future_feature_candidates"])
    missing_asof = [
        row["case_id"]
        for row in report["checks"]["as_of_join_preview"]
        if row["missing_asof_feature"]
    ]
    lines = [
        "# Decisión de auditoría SQL del split",
        "",
        f"Política: `{report['policy_id']}`.",
        f"Versión de split: `{report['split_version']}`.",
        f"Decisión: `{report['decision']}`.",
        "",
        "## Qué se ha comprobado",
        "",
        "| Check | Resultado | Lectura |",
        "|---|---:|---|",
        f"| Asignaciones duplicadas | {len(report['checks']['duplicate_assignments'])} | Debe ser 0: cada caso pertenece a un único split. |",
        f"| Casos sin asignación | {len(report['checks']['missing_assignments'])} | Debe ser 0: todo el snapshot queda cubierto. |",
        f"| Estudiantes en train y test | {len(report['checks']['student_train_test_overlap'])} | Debe ser 0 si medimos entidades no vistas. |",
        f"| Fuentes en train y test | {len(report['checks']['source_train_test_overlap'])} | Debe ser 0 si medimos fuentes no vistas. |",
        f"| Features futuras candidatas | {future_count} | Riesgo de leakage si se hace un join ingenuo. |",
        f"| Casos sin feature as-of | {len(missing_asof)} | Debe ser 0 si la feature es obligatoria. |",
        "",
        "## Lectura técnica",
        "",
    ]
    if future_count:
        lines.extend(
            [
                "El historial contiene features posteriores a algunos casos. Eso es normal en una tabla histórica viva, pero obliga a usar un as-of join: para cada caso, selecciona el último valor con `feature_available_at <= created_at`.",
                "",
                "Si alguien hiciera un join por `student_id` y eligiera el valor más reciente de toda la tabla, la evaluación miraría futuro. Por eso este reporte queda en `review`: no bloquea el split, pero obliga a revisar el código de feature engineering.",
            ]
        )
    else:
        lines.append("No se han detectado features futuras candidatas. Aun así, conserva la condición temporal en el join.")

    if missing_asof:
        lines.extend(
            [
                "",
                "Casos sin feature as-of:",
                "",
                ", ".join(f"`{case_id}`" for case_id in missing_asof),
            ]
        )

    lines.extend(
        [
            "",
            "## Archivos que sostienen esta decisión",
            "",
            "- `output/dataset_split_assignments.csv`: tabla materializada de split.",
            "- `data/feature_history.csv`: historial de features con fecha de disponibilidad.",
            "- `sql/split_audit_duckdb.sql`: versión SQL que puedes portar a DuckDB o a tu almacén.",
            "- `output/sql_split_audit_report.json`: reporte estructurado de checks.",
            "",
        ]
    )
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text("\n".join(lines), encoding="utf-8")


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--cases", type=Path, default=DEFAULT_CASES)
    parser.add_argument("--features", type=Path, default=DEFAULT_FEATURES)
    parser.add_argument("--assignments", type=Path, default=DEFAULT_ASSIGNMENTS)
    parser.add_argument("--policy", type=Path, default=DEFAULT_POLICY)
    parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT)
    parser.add_argument("--write", action="store_true")
    args = parser.parse_args()

    cases = read_csv(args.cases)
    features = read_csv(args.features)
    assignments = read_csv(args.assignments)
    policy = read_json(args.policy)
    report = build_report(cases, features, assignments, policy)

    if args.write:
        write_json(args.output_dir / "sql_split_audit_report.json", report)
        write_decision(args.output_dir / "sql_split_audit_decision.md", report)

    print(json.dumps(report, indent=2, ensure_ascii=False))


if __name__ == "__main__":
    main()
