#!/usr/bin/env python3
"""Reproduce every F2 number: DOJ pandemic-relief enforcement press releases by year and stage.

Input: the site's DOJ press-release corpus (datasets/raw/doj-press-releases: the 2026-07-04 harvest and the
2026-09-15 sweep) plus the API top-up (f2/doj_topup_slim.jsonl.gz, created-sorted, fetched 2026-09-23).
Rules: doj_rules.py (frozen copy in handcheck/doj_rules_FROZEN.py). Stdlib only.

Outputs in <out_dir>:
  doj-pandemic-enforcement-releases.csv  one row per in-scope release
  doj-pandemic-enforcement-by-year.csv   year x stage: releases, distinct announcements, named defendants, settlement USD
  numbers.json                           every number used on the page
Usage: reproduce_enforcement.py <raw_doj_dir> <topup.jsonl.gz> <out_dir>
"""
import csv, html, json, os, re, statistics, sys
from collections import Counter, defaultdict

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

# Exact prefilter: every scope() path needs one of these lowercase substrings (program names, 'unemployment',
# or a COVID/pandemic/coronavirus/CARES Act context term), so skipping releases without any of them changes nothing.
PREKEYS = ("paycheck", "ppp", "eidl", "economic injury", "unemployment", "pua", "fpuc", "lost wages", "employee retention",
           "ertc", "restaurant revitalization", "shuttered venue", "provider relief", "uninsured", "coronavirus", "covid",
           "pandemic", "cares act", "sick and family leave", "sick leave", "family leave", "ffcra", "american rescue plan",
           "esser", "elementary and secondary school emergency", "fiscal recovery", "payroll support program", "cfap",
           "feeding our future", "economic impact payment", "rental assistance", "tenant", "rent relief", "rent assistance")
STAGES = ["CHARGED", "PLEA", "CONVICTED", "SENTENCED", "CIVIL_RESOLUTION", "CIVIL_ACTION", "FORFEITURE", "OTHER"]


def norm_title(t):
    t = html.unescape(t or "").lower()
    t = re.sub(r"\$\s?([\d.,]+)\s*m\b", r"$\1 million", t)
    t = re.sub(r"\$\s?([\d.,]+)\s*b\b", r"$\1 billion", t)
    t = t.replace("payment protection", "paycheck protection").replace("&#039;", "'")
    t = re.sub(r"[^a-z0-9$. ]", " ", t)
    return set(re.sub(r"\s+", " ", t).split())


def day(d):
    y, m, dd = d.split("-")
    import datetime
    return datetime.date(int(y), int(m), int(dd)).toordinal()


def main(raw, topup, out):
    os.makedirs(out, exist_ok=True)
    corpus = load_corpus.load(raw, topup)
    all_by_year = Counter(r["date"][:4] for r in corpus)
    rows = []
    for r in corpus:
        if not r["date"] or dr.is_translation(r["title"]):
            continue
        low = (r["title"] + " " + r["body"]).lower()
        if not any(k in low for k in PREKEYS):
            continue
        progs = dr.scope(r["title"], r["body"])
        if not progs:
            continue
        st, basis = dr.stage(r["title"], r["body"])
        crim = st in ("CHARGED", "PLEA", "CONVICTED", "SENTENCED")
        names = dr.defendants(r["body"]) if crim else []
        dest = dr.defendants_estimate(r["title"], r["body"]) if crim else ""
        amt, abasis = dr.settlement_usd(r["title"], r["body"]) if st == "CIVIL_RESOLUTION" else (None, "")
        rows.append({"date": r["date"], "year": r["date"][:4], "stage": st, "stage_basis": basis,
                     "programs": "|".join(progs), "title": html.unescape(r["title"]), "url": r["url"],
                     "component": r["component"], "defendants_named": len(names),
                     "defendant_names": "; ".join(names), "defendants_estimate": dest, "settlement_usd": "" if amt is None else round(amt, 2),
                     "settlement_basis": abasis, "_tok": norm_title(r["title"]), "_lead": re.sub(r"\s+", " ", r["body"] or "")[:300]})
    rows.sort(key=lambda x: (x["date"], x["url"]))
    # near-duplicate clustering: same stage, within 30 days, and (title token Jaccard >= 0.8, or identical
    # opening 300 characters of body, or, for civil resolutions, the same settlement amount)
    cid = 0
    by_stage = defaultdict(list)
    for x in rows:
        by_stage[x["stage"]].append(x)
    for st, xs in by_stage.items():
        heads = []
        for x in xs:
            dx = day(x["date"]); hit = None
            for h in reversed(heads):
                if dx - day(h["date"]) > 30:
                    break
                j = len(x["_tok"] & h["_tok"]) / max(1, len(x["_tok"] | h["_tok"]))
                same_amt = False
                if st == "CIVIL_RESOLUTION" and x["settlement_usd"] != "" and h["settlement_usd"] != "":
                    gap, big = dx - day(h["date"]), max(x["settlement_usd"], h["settlement_usd"])
                    diff = abs(x["settlement_usd"] - h["settlement_usd"])
                    same_amt = (gap <= 3 and diff <= 0.001 * big) or (gap <= 7 and diff <= 0.02 * big and j >= 0.25)
                if j >= 0.8 or (x["_lead"] and x["_lead"] == h["_lead"]) or same_amt:
                    hit = h; break
            if hit:
                x["duplicate_of"] = hit["url"]; x["cluster"] = hit["cluster"]
            else:
                cid += 1; x["duplicate_of"] = ""; x["cluster"] = cid; heads.append(x)
    cols = ["date", "year", "stage", "stage_basis", "programs", "title", "url", "component", "defendants_named",
            "defendant_names", "defendants_estimate", "settlement_usd", "settlement_basis", "duplicate_of"]
    with open(os.path.join(out, "doj-pandemic-enforcement-releases.csv"), "w", newline="", encoding="utf-8") as fh:
        w = csv.DictWriter(fh, fieldnames=cols, extrasaction="ignore"); w.writeheader(); w.writerows(rows)
    years = sorted(set(x["year"] for x in rows))
    table = []
    for y in years:
        for st in STAGES:
            xs = [x for x in rows if x["year"] == y and x["stage"] == st]
            dist = [x for x in xs if not x["duplicate_of"]]
            amts = [x["settlement_usd"] for x in dist if x["settlement_usd"] != ""]
            table.append({"year": y, "stage": st, "releases": len(xs), "distinct_announcements": len(dist),
                          "defendants_named": sum(x["defendants_named"] for x in dist),
                          "defendants_estimate": sum(x["defendants_estimate"] for x in dist) if dist and dist[0]["defendants_estimate"] != "" else "",
                          "defendants_estimate_capped20": sum(min(x["defendants_estimate"], 20) for x in dist) if dist and dist[0]["defendants_estimate"] != "" else "",
                          "releases_naming_no_defendant": sum(1 for x in dist if x["defendants_named"] == 0),
                          "settlement_usd_sum": round(sum(amts), 2) if st == "CIVIL_RESOLUTION" else "",
                          "settlement_usd_median": round(statistics.median(amts), 2) if st == "CIVIL_RESOLUTION" and amts else "",
                          "settlements_with_amount": len(amts) if st == "CIVIL_RESOLUTION" else "",
                          "all_doj_releases_that_year": all_by_year[y]})
    with open(os.path.join(out, "doj-pandemic-enforcement-by-year.csv"), "w", newline="", encoding="utf-8") as fh:
        w = csv.DictWriter(fh, fieldnames=list(table[0].keys())); w.writeheader(); w.writerows(table)
    T = {(t["year"], t["stage"]): t for t in table}
    N = {"corpus_releases": len(corpus), "corpus_first_date": min(r["date"] for r in corpus if r["date"]),
         "corpus_last_date": max(r["date"] for r in corpus if r["date"]), "in_scope_releases": len(rows),
         "in_scope_distinct": sum(1 for x in rows if not x["duplicate_of"]),
         "all_doj_by_year": {y: all_by_year[y] for y in sorted(all_by_year)}, "by_year": {}}
    for y in years:
        d = {st: {k: T[(y, st)][k] for k in ("releases", "distinct_announcements", "defendants_named", "defendants_estimate", "defendants_estimate_capped20",
                                               "releases_naming_no_defendant", "settlement_usd_sum",
                                               "settlement_usd_median", "settlements_with_amount")} for st in STAGES}
        d["_all_doj"] = all_by_year[y]
        N["by_year"][y] = d
    top = sorted((x for x in rows if x["stage"] == "CIVIL_RESOLUTION" and not x["duplicate_of"] and x["settlement_usd"] != ""),
                 key=lambda x: -x["settlement_usd"])[:10]
    N["largest_settlements"] = [[x["date"], x["settlement_usd"], x["title"][:120], x["url"]] for x in top]
    N["last_in_scope_date"] = rows[-1]["date"]
    json.dump(N, open(os.path.join(out, "numbers.json"), "w"), indent=1, sort_keys=True)
    open(os.path.join(out, "F2.DONE"), "w").write("ok\n")
    for y in years:
        c, s = T[(y, "CHARGED")], T[(y, "CIVIL_RESOLUTION")]
        print(y, "charged", c["releases"], c["distinct_announcements"], c["defendants_named"], c["defendants_estimate"], "| civil", s["releases"],
              s["distinct_announcements"], s["settlement_usd_sum"], "| all", all_by_year[y])


if __name__ == "__main__":
    main(sys.argv[1], sys.argv[2], sys.argv[3])
