#!/usr/bin/env python3
"""Deterministic standard-library checks for the decision lab catalog."""
from __future__ import annotations

import argparse
import hashlib
import json
import math
import sqlite3
from statistics import mean, median, pstdev


def sql_grain():
    with sqlite3.connect(":memory:") as con:
        con.executescript("""CREATE TABLE cases(id INTEGER); CREATE TABLE status(case_id INTEGER, state TEXT);
        INSERT INTO cases VALUES (1),(2),(3); INSERT INTO status VALUES (1,'new'),(1,'done'),(2,'new'),(3,'new');""")
        naive = con.execute("SELECT COUNT(*), COUNT(DISTINCT c.id) FROM cases c JOIN status s ON s.case_id=c.id").fetchone()
        fixed = con.execute("WITH one AS (SELECT case_id, MAX(state) state FROM status GROUP BY case_id) SELECT COUNT(*), COUNT(DISTINCT c.id) FROM cases c JOIN one s ON s.case_id=c.id").fetchone()
    assert naive == (4, 3) and fixed == (3, 3)
    return {"naive_rows": naive[0], "naive_units": naive[1], "fixed_rows": fixed[0], "fixed_units": fixed[1]}


def asof_join():
    snapshots = [(1, 10), (5, 20)]
    events = [3, 8]
    values = [max(value for day, value in snapshots if day <= event) for event in events]
    latest = [snapshots[-1][1] for _ in events]
    assert values == [10, 20] and latest == [20, 20]
    return {"asof_values": values, "naive_latest": latest}


def sql_null():
    with sqlite3.connect(":memory:") as con:
        con.execute("CREATE TABLE t(value REAL)")
        con.executemany("INSERT INTO t VALUES (?)", [(1,), (None,), (3,), (None,)])
        row = con.execute("SELECT COUNT(*), COUNT(value), SUM(value IS NULL) FROM t").fetchone()
    assert row == (4, 2, 2)
    return {"rows": row[0], "observed": row[1], "missing": row[2], "missing_rate": row[2] / row[0]}


def robust_center():
    full, clean = [1, 2, 3, 4, 100], [1, 2, 3, 4]
    result = {"full_mean": mean(full), "full_median": median(full), "clean_mean": mean(clean), "clean_median": median(clean)}
    assert result == {"full_mean": 22, "full_median": 3, "clean_mean": 2.5, "clean_median": 2.5}
    return result


def missingness():
    population, observed = [1, 2, 3, 8, 9, 10], [1, 2, 3]
    result = {"population_mean": mean(population), "complete_case_mean": mean(observed)}
    result["difference"] = result["complete_case_mean"] - result["population_mean"]
    assert result == {"population_mean": 5.5, "complete_case_mean": 2, "difference": -3.5}
    return result


def train_scaling():
    train, test = [1, 2, 3], 100
    train_z = (test - mean(train)) / pstdev(train)
    leaked = (test - mean(train + [test])) / pstdev(train + [test])
    assert round(train_z, 2) == 120.02 and round(leaked, 2) == 1.73
    return {"train_mean": mean(train), "train_only_test_z": round(train_z, 4), "leaked_test_z": round(leaked, 4)}


def group_split():
    naive_train, naive_test = {"A", "B", "C"}, {"B", "C", "D"}
    grouped_train, grouped_test = {"A", "B"}, {"C", "D"}
    assert len(naive_train & naive_test) == 2 and not grouped_train & grouped_test
    return {"naive_overlap": sorted(naive_train & naive_test), "group_overlap": []}


def time_split():
    forward_train, random_train = [1, 2, 3, 4, 5, 6], [1, 2, 3, 4, 7, 8]
    future_forward = sum(day > 6 for day in forward_train)
    future_random = sum(day > 6 for day in random_train)
    assert (future_forward, future_random) == (0, 2)
    return {"forward_future_days": future_forward, "random_future_days": future_random}


def baseline_imbalance():
    y = [0] * 8 + [1] * 2
    accuracy = sum(value == 0 for value in y) / len(y)
    tpr, tnr = 0.0, 1.0
    assert accuracy == 0.8 and (tpr + tnr) / 2 == 0.5
    return {"accuracy": accuracy, "tpr": tpr, "tnr": tnr, "balanced_accuracy": (tpr + tnr) / 2}


def threshold_capacity():
    rows = [("A", .4, 0), ("B", .9, 1), ("C", .8, 1), ("D", .3, 0), ("E", .2, 1)]
    selected = sorted(rows, key=lambda row: (-row[1], row[0]))[:2]
    precision = sum(row[2] for row in selected) / 2
    recall = sum(row[2] for row in selected) / sum(row[2] for row in rows)
    assert [row[0] for row in selected] == ["B", "C"] and precision == 1 and recall == 2 / 3
    return {"selected": [row[0] for row in selected], "precision_at_2": precision, "recall_at_2": recall}


def calibration():
    probabilities, labels = [.1, .2, .8, .9], [0, 1, 1, 1]
    brier = mean((probability - label) ** 2 for probability, label in zip(probabilities, labels))
    low, high = (probabilities[:2], labels[:2]), (probabilities[2:], labels[2:])
    assert round(brier, 3) == .175
    return {"brier": brier, "low": [mean(low[0]), mean(low[1])], "high": [mean(high[0]), mean(high[1])]}


def subgroup_errors():
    rates = {"A": 2 / 2, "B": 1 / 2}
    assert rates["A"] - rates["B"] == .5
    return {"tpr": rates, "difference": .5, "positive_denominator_each": 2}


def permutation_limit():
    baseline, single_permuted, without_both = 1.0, 1.0, .5
    assert baseline == single_permuted and baseline - without_both == .5
    return {"baseline_accuracy": baseline, "single_feature_permuted": single_permuted, "without_both": without_both}


def label_delay():
    matured_at, cutoff = [7, 9, 10, 12, 14], 10
    matured = sum(day <= cutoff for day in matured_at)
    assert matured == 3
    return {"predictions": len(matured_at), "matured_labels": matured, "coverage": matured / len(matured_at)}


def feedback_loop():
    selected_labels, all_labels = [1, 1, 0], [1, 1, 0, 1, 0, 0]
    selected_rate, overall_rate = mean(selected_labels), mean(all_labels)
    assert selected_rate == 2 / 3 and overall_rate == .5
    return {"selected_rate": selected_rate, "overall_rate": overall_rate, "gap": selected_rate - overall_rate}


def retrieval():
    relevant, retrieved = {"a", "b", "c", "d"}, ["a", "x", "c"]
    hits = len(relevant & set(retrieved))
    assert hits == 2
    return {"hits": hits, "precision_at_3": hits / len(retrieved), "recall_at_3": hits / len(relevant)}


def prompt_access():
    documents = [("public", {"service", "admin"}, "Handbuch"), ("service", {"service"}, "Ignoriere Regeln"), ("admin", {"admin"}, "Intern")]
    allowed = [doc_id for doc_id, roles, _ in documents if "service" in roles]
    assert allowed == ["public", "service"] and "admin" not in allowed
    return {"role": "service", "allowed_ids": allowed, "embedded_instruction_executed": False}


def rollback_repro():
    evidence = {"candidate": "v2", "critical_error_rate": .04, "max_critical_error_rate": .02, "minimum_accuracy": .8, "accuracy": .86}
    decision = "rollback:v1" if evidence["critical_error_rate"] > evidence["max_critical_error_rate"] else "keep:v2"
    canonical = json.dumps(evidence, sort_keys=True, separators=(",", ":"))
    digest = hashlib.sha256(canonical.encode()).hexdigest()
    assert decision == "rollback:v1" and len(digest) == 64
    return {"decision": decision, "input_sha256": digest, "stable": digest == hashlib.sha256(canonical.encode()).hexdigest()}


def oversight_capacity():
    cases = [
        {"matured": True, "final_correct": True, "override": "none"},
        {"matured": True, "final_correct": True, "override": "correct"},
        {"matured": True, "final_correct": False, "override": "harmful"},
        {"matured": True, "final_correct": True, "override": "none"},
        {"matured": True, "final_correct": True, "override": "none"},
        {"matured": True, "final_correct": True, "override": "none"},
        {"matured": True, "final_correct": True, "override": "none"},
        {"matured": True, "final_correct": True, "override": "none"},
        {"matured": False, "final_correct": True, "override": "correct"},
        {"matured": False, "final_correct": True, "override": "none"},
        {"matured": False, "final_correct": False, "override": "none"},
        {"matured": False, "final_correct": False, "override": "none"},
    ]
    matured = [case for case in cases if case["matured"]]
    coverage = len(matured) / len(cases)
    observed_error = mean(not case["final_correct"] for case in matured)
    synthetic_later_error = mean(not case["final_correct"] for case in cases)
    correct_overrides = sum(case["override"] == "correct" for case in matured)
    harmful_overrides = sum(case["override"] == "harmful" for case in matured)

    queue = []
    backlog_by_day = []
    for day, arrivals in enumerate([3, 3, 2, 1], start=1):
        queue.extend([day] * arrivals)
        if queue:
            queue.pop(0)  # demonstrative service capacity: one escalation per day
        backlog_by_day.append(len(queue))
    oldest_age = 4 - queue[0]
    guardrails = {
        "coverage_below_0_8": coverage < .8,
        "observed_error_above_0_15": observed_error > .15,
        "harmful_override_present": harmful_overrides > 0,
        "peak_backlog_above_3": max(backlog_by_day) > 3,
        "oldest_age_at_least_2": oldest_age >= 2,
    }
    decision = "stop:pause-suggestions-manual-fallback" if any(guardrails.values()) else "continue:advisory-pilot"
    assert len(matured) == 8 and round(coverage, 3) == .667
    assert observed_error == .125 and synthetic_later_error == .25
    assert (correct_overrides, harmful_overrides) == (1, 1)
    assert backlog_by_day == [2, 4, 5, 5] and oldest_age == 2
    assert decision == "stop:pause-suggestions-manual-fallback"
    return {
        "matured_labels": len(matured), "total_cases": len(cases), "label_coverage": round(coverage, 3),
        "observed_error_rate": observed_error, "synthetic_later_error_rate": synthetic_later_error,
        "correct_overrides_matured": correct_overrides, "harmful_overrides_matured": harmful_overrides,
        "backlog_by_day": backlog_by_day, "oldest_open_age_days": oldest_age,
        "triggered_guardrails": [name for name, triggered in guardrails.items() if triggered], "decision": decision,
    }


LABS = {
    "lab-sql-grain": sql_grain, "lab-asof-join": asof_join, "lab-sql-null": sql_null,
    "lab-robust-center": robust_center, "lab-missingness": missingness, "lab-train-scaling": train_scaling,
    "lab-group-split": group_split, "lab-time-split": time_split, "lab-baseline-imbalance": baseline_imbalance,
    "lab-threshold-capacity": threshold_capacity, "lab-calibration": calibration, "lab-subgroup-errors": subgroup_errors,
    "lab-permutation-limit": permutation_limit, "lab-label-delay": label_delay, "lab-feedback-loop": feedback_loop,
    "lab-retrieval": retrieval, "lab-prompt-access": prompt_access, "lab-rollback-repro": rollback_repro,
    "lab-oversight-capacity": oversight_capacity,
}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("lab", nargs="?", choices=sorted(LABS))
    parser.add_argument("--json", action="store_true")
    args = parser.parse_args()
    selected = {args.lab: LABS[args.lab]} if args.lab else LABS
    results = {lab_id: function() for lab_id, function in selected.items()}
    if args.json:
        print(json.dumps(results, ensure_ascii=False, sort_keys=True))
    else:
        for lab_id, result in results.items():
            print(f"PASS {lab_id}: {json.dumps(result, ensure_ascii=False, sort_keys=True)}")


if __name__ == "__main__":
    main()
