#!/usr/bin/env python3
"""Score the final F2 rules against the hand-checked releases (judgments made by reading each release).

Inputs (f2/handcheck/): sample.csv + judgments.csv (425 stratified releases), judgments_new_scope.csv (30),
recall samples (recall2/recall3/recall_final url lists are judged in RECALL_JUDGED below),
defendant_truth.csv + defendant_truth_fresh.csv. Text comes from f2/work/broad.jsonl.gz.
Writes f2/handcheck/validation.json and prints a summary.
Usage: validate_handcheck.py [<raw_doj_dir> <topup.jsonl.gz>] [--amounts]  (the paths are needed only outside the work folder)
"""
import csv, gzip, json, os, sys
from collections import Counter, defaultdict

HERE = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, HERE)
import doj_rules as dr

H = os.path.join(HERE, "handcheck")
by = {}
_broad = os.path.join(HERE, "work", "broad.jsonl.gz")
if os.path.exists(_broad):
    for line in gzip.open(_broad, "rt", encoding="utf-8"):
        r = json.loads(line); by[r["url"]] = r
else:  # packaged use: validate_handcheck.py <raw_doj_dir> <topup.jsonl.gz> [--amounts]
    import load_corpus
    for r in load_corpus.load(sys.argv[1], sys.argv[2]):
        by[r["url"]] = r

S = {r["hid"]: r for r in csv.DictReader(open(os.path.join(H, "sample.csv"), encoding="utf-8"))}
J = [(S[j["hid"]]["url"], j["scope"], j["true_class"], S[j["hid"]]["stratum"], S[j["hid"]]["date"][:4])
     for j in csv.DictReader(open(os.path.join(H, "judgments.csv"), encoding="utf-8"))]
J += [(j["url"], j["scope"], j["true_class"], "NEWSCOPE", by[j["url"]]["date"][:4])
      for j in csv.DictReader(open(os.path.join(H, "judgments_new_scope.csv"), encoding="utf-8"))]


def auto(url):
    r = by[url]
    if dr.is_translation(r["title"]) or not dr.scope(r["title"], r["body"]):
        return "OUT"
    return dr.stage(r["title"], r["body"])[0]


prec = defaultdict(Counter); byyear = defaultdict(Counter); rows = []
for url, sc, tc, stratum, yr in J:
    a = auto(url); rows.append([url, stratum, yr, a, sc, tc])
    p = prec[a]; p["n"] += 1
    if a == "OUT":
        p["in_scope_missed"] += sc == "Y"; continue
    p["stage_right"] += tc == a
    p["strict"] += (sc == "Y" and tc == a)
    p["lenient"] += (sc in ("Y", "P") and tc == a)
    if a in ("CHARGED", "CIVIL_RESOLUTION") and yr in ("2021", "2025"):
        q = byyear[(a, yr)]; q["n"] += 1; q["strict"] += (sc == "Y" and tc == a); q["lenient"] += (sc in ("Y", "P") and tc == a)
# recall samples drawn from out-of-scope releases that mention fraud and COVID/pandemic/CARES Act
RECALL_JUDGED = {}
for fn in ("recall_final_judgments.csv", "recall2_judgments.csv", "recall3_judgments.csv"):
    rr = list(csv.DictReader(open(os.path.join(H, fn), encoding="utf-8")))
    RECALL_JUDGED[fn] = {"n": len(rr), "clear_misses": sum(r["judgment"] == "MISS" for r in rr),
                         "borderline": sum(r["judgment"] == "BORDERLINE" for r in rr),
                         "rules": "final" if fn.startswith("recall_final") else "earlier (misses fixed afterwards)"}
out = {"precision": {k: dict(v) for k, v in prec.items()}, "precision_2021_2025": {f"{k[0]} {k[1]}": dict(v) for k, v in byyear.items()},
       "recall": RECALL_JUDGED}
# defendant estimator
dv = {}
for fn, tag in (("defendant_truth.csv", "tuning"), ("defendant_truth_fresh.csv", "fresh")):
    for r in csv.DictReader(open(os.path.join(H, fn), encoding="utf-8")):
        if not r["true_defendants"]:
            continue
        x = by[r["url"]]; e = dr.defendants_estimate(x["title"], x["body"]); t = int(r["true_defendants"])
        c = dv.setdefault(f"{tag} {r['year']}", Counter()); c["releases"] += 1; c["true"] += t; c["estimate"] += e; c["exact"] += e == t
out["defendants"] = {k: dict(v) for k, v in dv.items()}
with open(os.path.join(H, "precision_rows_final.csv"), "w", newline="", encoding="utf-8") as fh:
    w = csv.writer(fh); w.writerow(["url", "stratum", "year", "auto_class_final", "scope_judged", "true_class_judged"]); w.writerows(rows)
json.dump(out, open(os.path.join(H, "validation.json"), "w"), indent=1, sort_keys=True)
for k in ("CHARGED", "PLEA", "CONVICTED", "SENTENCED", "CIVIL_RESOLUTION", "CIVIL_ACTION", "FORFEITURE", "OTHER", "OUT"):
    print(k, dict(prec[k]))
for k, v in sorted(byyear.items()):
    print(k, dict(v))
for k, v in sorted(dv.items()):
    print("defendants", k, dict(v))
if "--amounts" in sys.argv:
    for url, sc, tc, stratum, yr in J:
        if auto(url) == "CIVIL_RESOLUTION":
            x = by[url]; amt, basis = dr.settlement_usd(x["title"], x["body"])
            print(f"AMT|{yr}|{'' if amt is None else f'{amt:,.2f}'}|{basis}|{x['title'][:110]}")
