#!/usr/bin/env python3
from __future__ import annotations

import argparse
import json
from pathlib import Path
from datetime import datetime
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt


def log(msg: str) -> None:
    print(msg, flush=True)


def parse_dates(df: pd.DataFrame, cols: list[str]) -> pd.DataFrame:
    for c in cols:
        if c in df.columns:
            df[c] = pd.to_datetime(df[c], errors="coerce")
    return df


def load_inputs(input_dir: Path, input_xlsx: Path | None):
    if input_xlsx is not None and input_xlsx.exists():
        log(f"[read] {input_xlsx}")
        xls = pd.ExcelFile(input_xlsx)
        tfo = pd.read_excel(input_xlsx, sheet_name="r1_tracking_font_objects")
        if "r1_tracked_docs_summary" in xls.sheet_names:
            docs = pd.read_excel(input_xlsx, sheet_name="r1_tracked_docs_summary")
        else:
            docs = None
    else:
        tfo_path = input_dir / "10r1_tracking_font_objects__long.csv"
        docs_path = input_dir / "10r1_tracked_docs__summary.csv"
        if not tfo_path.exists():
            raise FileNotFoundError(f"Missing required file: {tfo_path}")
        log(f"[read] {tfo_path}")
        tfo = pd.read_csv(tfo_path)
        if docs_path.exists():
            log(f"[read] {docs_path}")
            docs = pd.read_csv(docs_path)
        else:
            docs = None

    tfo = parse_dates(tfo, ["filing_date", "efile_date"])
    tfo["cluster_id_num"] = pd.to_numeric(tfo.get("cluster_id"), errors="coerce")
    tfo = tfo.dropna(subset=["doc_sha256", "filing_date", "object_sha256"]).copy()

    if docs is None:
        # derive minimal doc table if tracked_docs summary is unavailable
        docs = (
            tfo.sort_values("filing_date")
            .drop_duplicates("doc_sha256")
            .rename(columns={"cluster_id_num": "cluster_id"})
        )
        keep = [c for c in ["doc_sha256", "case_id", "cluster_id", "cluster_name", "filing_date", "efile_date"] if c in docs.columns]
        docs = docs[keep].copy()
    else:
        docs = parse_dates(docs, ["filing_date", "efile_date"])
        docs["cluster_id_num"] = pd.to_numeric(docs.get("cluster_id"), errors="coerce")
        docs = docs.dropna(subset=["doc_sha256", "filing_date"]).copy()

    # unique docs
    docs_u = docs.sort_values("filing_date").drop_duplicates("doc_sha256").copy()
    if "cluster_id_num" not in docs_u.columns:
        docs_u["cluster_id_num"] = pd.to_numeric(docs_u.get("cluster_id"), errors="coerce")

    return tfo, docs_u


def pct_rank(series: pd.Series) -> pd.Series:
    s = pd.to_numeric(series, errors="coerce")
    if s.notna().sum() == 0:
        return pd.Series(np.nan, index=s.index)
    return s.rank(pct=True, method="average")


def compute_entity_temporal(
    tfo: pd.DataFrame,
    entity_col: str,
    focal_cluster_id: int,
    global_focal_start: pd.Timestamp,
) -> pd.DataFrame:
    rows = []
    for ent, sub in tfo.groupby(entity_col, dropna=False):
        sub = sub.copy()
        sub["is_focal"] = sub["cluster_id_num"].eq(focal_cluster_id)
        sf = sub[sub["is_focal"]]
        sn = sub[~sub["is_focal"]]

        n_docs_all = sub["doc_sha256"].nunique()
        n_docs_focal = sf["doc_sha256"].nunique()
        n_docs_nonfocal = sn["doc_sha256"].nunique()

        first_all = sub["filing_date"].min()
        last_all = sub["filing_date"].max()
        span_all_days = (last_all - first_all).days if pd.notna(first_all) and pd.notna(last_all) else np.nan

        first_focal = sf["filing_date"].min() if len(sf) else pd.NaT
        last_focal = sf["filing_date"].max() if len(sf) else pd.NaT
        first_nonfocal = sn["filing_date"].min() if len(sn) else pd.NaT
        last_nonfocal = sn["filing_date"].max() if len(sn) else pd.NaT

        gap_days = np.nan
        pre_entity_docs = np.nan
        pre_entity_pct = np.nan
        if pd.notna(first_focal) and pd.notna(first_nonfocal):
            gap_days = (first_focal - first_nonfocal).days
            pre_entity_docs = sn.loc[sn["filing_date"] < first_focal, "doc_sha256"].nunique()
            pre_entity_pct = (pre_entity_docs / n_docs_nonfocal * 100.0) if n_docs_nonfocal else np.nan

        # relative to global focal start
        pre_global_docs = (
            sn.loc[sn["filing_date"] < global_focal_start, "doc_sha256"].nunique()
            if pd.notna(global_focal_start)
            else np.nan
        )
        pre_global_pct = (pre_global_docs / n_docs_nonfocal * 100.0) if n_docs_nonfocal and pd.notna(pre_global_docs) else np.nan

        row = {
            entity_col: ent,
            "n_rows": len(sub),
            "n_docs_all": n_docs_all,
            "n_docs_focal": n_docs_focal,
            "n_docs_nonfocal": n_docs_nonfocal,
            "n_clusters_all": sub["cluster_id_num"].nunique(),
            "n_cases_all": sub["case_id"].nunique() if "case_id" in sub.columns else np.nan,
            "first_all": first_all,
            "last_all": last_all,
            "span_all_days": span_all_days,
            "first_focal": first_focal,
            "last_focal": last_focal,
            "first_nonfocal": first_nonfocal,
            "last_nonfocal": last_nonfocal,
            "gap_nonfocal_before_focal_days": gap_days,
            "pre_entity_nonfocal_docs": pre_entity_docs,
            "pre_entity_nonfocal_pct": pre_entity_pct,
            "pre_global_nonfocal_docs": pre_global_docs,
            "pre_global_nonfocal_pct": pre_global_pct,
            "focal_doc_share_pct": (n_docs_focal / n_docs_all * 100.0) if n_docs_all else np.nan,
        }

        if "object_description" in sub.columns:
            v = sub["object_description"].dropna()
            row["object_description"] = v.iloc[0] if len(v) else np.nan
        if "object_pdf_group" in sub.columns:
            v = sub["object_pdf_group"].dropna()
            row["object_pdf_group"] = v.mode().iloc[0] if len(v) else np.nan

        rows.append(row)

    out = pd.DataFrame(rows)

    # Hypothesis-conditioned prioritization score (NOT proof)
    for c in ["gap_nonfocal_before_focal_days", "pre_entity_nonfocal_pct", "span_all_days"]:
        out[c + "_pct_rank"] = pct_rank(out[c])

    out["temporal_pressure_score"] = (
        out["gap_nonfocal_before_focal_days_pct_rank"].fillna(0) * 0.5
        + out["pre_entity_nonfocal_pct_pct_rank"].fillna(0) * 0.3
        + out["span_all_days_pct_rank"].fillna(0) * 0.2
    ) * 100.0

    return out.sort_values("temporal_pressure_score", ascending=False)


def compute_cluster_ranges(docs_u: pd.DataFrame, tfo: pd.DataFrame, focal_cluster_id: int) -> pd.DataFrame:
    if "cluster_name" not in docs_u.columns:
        docs_u["cluster_name"] = np.nan
    if "case_id" not in docs_u.columns:
        docs_u["case_id"] = np.nan

    cl = docs_u.groupby(["cluster_id_num", "cluster_name"], dropna=False).agg(
        n_docs=("doc_sha256", "nunique"),
        n_cases=("case_id", "nunique"),
        first_filing=("filing_date", "min"),
        last_filing=("filing_date", "max"),
    ).reset_index()

    cl["span_days"] = (cl["last_filing"] - cl["first_filing"]).dt.days
    cl["is_focal_cluster"] = cl["cluster_id_num"].eq(focal_cluster_id)

    hcount = tfo.groupby("cluster_id_num")["object_sha256"].nunique().rename("n_hashes")
    gcount = tfo.groupby("cluster_id_num")["object_pdf_group"].nunique().rename("n_pdf_groups")
    cl = cl.merge(hcount, on="cluster_id_num", how="left")
    cl = cl.merge(gcount, on="cluster_id_num", how="left")

    return cl.sort_values(["is_focal_cluster", "span_days"], ascending=[False, False])


def save_csv(df: pd.DataFrame, path: Path) -> None:
    out = df.copy()
    for c in out.columns:
        if np.issubdtype(out[c].dtype, np.datetime64):
            out[c] = out[c].dt.strftime("%Y-%m-%d")
    out.to_csv(path, index=False)
    log(f"[ok] {path}")


def plot_group_ranges(group_df: pd.DataFrame, out_png: Path, focal_start: pd.Timestamp, top_n: int = 12):
    d = group_df.sort_values("span_all_days", ascending=False).head(top_n).copy()
    d = d.sort_values("first_all")
    y = np.arange(len(d))

    fig, ax = plt.subplots(figsize=(12, 7), dpi=160)
    for i, (_, r) in enumerate(d.iterrows()):
        ax.hlines(y=i, xmin=r["first_all"], xmax=r["last_all"], linewidth=3)
        ax.plot(r["first_all"], i, marker="o")
        ax.plot(r["last_all"], i, marker="o")
    if pd.notna(focal_start):
        ax.axvline(focal_start, linestyle="--", linewidth=1)

    ax.set_yticks(y)
    ax.set_yticklabels(d["object_pdf_group"].astype(str))
    ax.set_xlabel("Filing date")
    ax.set_ylabel("Tracking font group")
    ax.set_title("Run 2 temporal ranges by tracking font group")
    ax.grid(True, axis="x", alpha=0.3)
    fig.tight_layout()
    fig.savefig(out_png, bbox_inches="tight")
    plt.close(fig)
    log(f"[ok] {out_png}")


def plot_hash_gaps(hash_df: pd.DataFrame, out_png: Path):
    d = hash_df.sort_values("gap_nonfocal_before_focal_days", ascending=False).copy()
    labels = d["object_pdf_group"].astype(str) + " | " + d["object_sha256"].astype(str).str.slice(0, 10)
    vals = pd.to_numeric(d["gap_nonfocal_before_focal_days"], errors="coerce").fillna(0)

    fig, ax = plt.subplots(figsize=(12, 7), dpi=160)
    ax.barh(np.arange(len(d)), vals)
    ax.set_yticks(np.arange(len(d)))
    ax.set_yticklabels(labels)
    ax.invert_yaxis()
    ax.set_xlabel("Days non-focal starts before focal start (entity-level)")
    ax.set_ylabel("Hash")
    ax.set_title("Run 2 temporal lead-in gap by tracking font hash")
    ax.grid(True, axis="x", alpha=0.3)
    fig.tight_layout()
    fig.savefig(out_png, bbox_inches="tight")
    plt.close(fig)
    log(f"[ok] {out_png}")


def plot_yearly(yearly: pd.DataFrame, out_png: Path):
    fig, ax = plt.subplots(figsize=(11, 5), dpi=160)
    ax.plot(yearly["year"], yearly["nonfocal_docs"], marker="o", label="non-focal docs")
    ax.plot(yearly["year"], yearly["focal_docs"], marker="o", label="focal docs")
    ax.set_xlabel("Filing year")
    ax.set_ylabel("Unique docs")
    ax.set_title("Run 2 yearly doc volume: focal vs non-focal")
    ax.grid(True, alpha=0.3)
    ax.legend()
    fig.tight_layout()
    fig.savefig(out_png, bbox_inches="tight")
    plt.close(fig)
    log(f"[ok] {out_png}")


def make_markdown_report(
    out_md: Path,
    focal_cluster_id: int,
    docs_u: pd.DataFrame,
    tfo: pd.DataFrame,
    group_df: pd.DataFrame,
    hash_df: pd.DataFrame,
    cluster_df: pd.DataFrame,
    yearly: pd.DataFrame,
    global_focal_start: pd.Timestamp,
):
    focal_docs = docs_u[docs_u["cluster_id_num"].eq(focal_cluster_id)]["doc_sha256"].nunique()
    nonfocal_docs = docs_u[~docs_u["cluster_id_num"].eq(focal_cluster_id)]["doc_sha256"].nunique()
    focal_first = docs_u.loc[docs_u["cluster_id_num"].eq(focal_cluster_id), "filing_date"].min()
    focal_last = docs_u.loc[docs_u["cluster_id_num"].eq(focal_cluster_id), "filing_date"].max()
    pre_non = docs_u.loc[
        (~docs_u["cluster_id_num"].eq(focal_cluster_id)) & (docs_u["filing_date"] < global_focal_start),
        "doc_sha256"
    ].nunique()

    delta_all_zero = False
    if "efile_date" in docs_u.columns and docs_u["efile_date"].notna().any():
        d = (docs_u["efile_date"] - docs_u["filing_date"]).dt.days
        delta_all_zero = (d.fillna(0) == 0).all()

    top_groups = group_df[[
        "object_pdf_group", "n_docs_all", "n_docs_focal", "n_docs_nonfocal",
        "first_nonfocal", "first_focal", "gap_nonfocal_before_focal_days",
        "pre_entity_nonfocal_pct", "span_all_days", "temporal_pressure_score"
    ]].sort_values("temporal_pressure_score", ascending=False).head(8)

    top_hashes = hash_df[[
        "object_pdf_group", "object_sha256", "object_description", "n_docs_all",
        "n_docs_focal", "n_docs_nonfocal", "gap_nonfocal_before_focal_days",
        "pre_entity_nonfocal_pct", "span_all_days", "temporal_pressure_score"
    ]].sort_values("temporal_pressure_score", ascending=False).head(8)

    longest_clusters = cluster_df[[
        "cluster_id_num", "cluster_name", "n_docs", "n_hashes", "n_pdf_groups",
        "first_filing", "last_filing", "span_days", "is_focal_cluster"
    ]].sort_values("span_days", ascending=False).head(10)

    md = []
    md.append("# Run 2 Temporal Analysis (Hypothesis-Conditioned)")
    md.append("")
    md.append("This report is intentionally **hypothesis-conditioned**:")
    md.append("it models temporal signatures under the assumption that non-focal groups could include synthetic/back-dated filings.")
    md.append("It does **not** conclude that this is true; it surfaces patterns that would be notable if the hypothesis were true.")
    md.append("")
    md.append("## Dataset Scope")
    md.append(f"- Unique docs: **{docs_u['doc_sha256'].nunique():,}**")
    md.append(f"- Tracking-font object rows: **{len(tfo):,}**")
    md.append(f"- Clusters: **{docs_u['cluster_id_num'].nunique():,}**")
    md.append(f"- Cases: **{docs_u['case_id'].nunique():,}**")
    md.append(f"- Tracking hashes: **{tfo['object_sha256'].nunique():,}**")
    md.append(f"- Tracking groups: **{tfo['object_pdf_group'].nunique():,}**")
    md.append("")
    md.append("## Focal Anchor")
    md.append(f"- Focal cluster ID: **{focal_cluster_id}**")
    md.append(f"- Focal docs: **{focal_docs:,}**")
    md.append(f"- Focal filing window: **{focal_first.date() if pd.notna(focal_first) else 'NA'} → {focal_last.date() if pd.notna(focal_last) else 'NA'}**")
    md.append(f"- Non-focal docs: **{nonfocal_docs:,}**")
    if pd.notna(global_focal_start):
        md.append(f"- Non-focal docs dated earlier than focal start ({global_focal_start.date()}): **{pre_non:,}**")
    md.append("")
    md.append("## Top Temporal Signals by Group")
    md.append("Columns: docs, focal share, first non-focal vs first focal gap, pre-focal non-focal %, span days, composite score.")
    md.append("```")
    md.append(top_groups.to_string(index=False))
    md.append("```")
    md.append("")
    md.append("## Top Temporal Signals by Hash")
    md.append("```")
    md.append(top_hashes.to_string(index=False))
    md.append("```")
    md.append("")
    md.append("## Longest Cluster Time Ranges")
    md.append("```")
    md.append(longest_clusters.to_string(index=False))
    md.append("```")
    md.append("")
    md.append("## Yearly Volume (focal vs non-focal)")
    md.append("```")
    md.append(yearly.to_string(index=False))
    md.append("```")
    md.append("")
    md.append("## QA / Limits")
    if delta_all_zero:
        md.append("- In this run's tables, `filing_date` and `efile_date` are identical for all docs, so this run cannot independently validate backdating from those two fields alone.")
    else:
        md.append("- `filing_date` and `efile_date` are not always identical; inspect delta columns in exports.")
    md.append("- Temporal pressure score is a prioritization heuristic (rank-based), not a proof metric.")
    md.append("- Strong signals indicate where to inspect provenance metadata, signature chains, and raw ingestion timestamps next.")
    md.append("")

    out_md.write_text("\n".join(md), encoding="utf-8")
    log(f"[ok] {out_md}")


def main():
    ap = argparse.ArgumentParser(description="Run 2 temporal one-shot report for 10_font_tracking")
    ap.add_argument("--input_dir", type=Path, default=Path("reports/10_font_tracking"))
    ap.add_argument("--input_xlsx", type=Path, default=None,
                    help="Optional workbook path (reads sheets r1_tracking_font_objects + r1_tracked_docs_summary).")
    ap.add_argument("--out_dir", type=Path, default=Path("reports/10_font_tracking"))
    ap.add_argument("--focal_cluster_id", type=int, default=1570)
    ap.add_argument("--top_n_groups_plot", type=int, default=12)
    args = ap.parse_args()

    args.out_dir.mkdir(parents=True, exist_ok=True)

    tfo, docs_u = load_inputs(args.input_dir, args.input_xlsx)

    docs_u["cluster_id_num"] = pd.to_numeric(docs_u.get("cluster_id_num"), errors="coerce")
    tfo["cluster_id_num"] = pd.to_numeric(tfo.get("cluster_id_num"), errors="coerce")

    global_focal_start = docs_u.loc[docs_u["cluster_id_num"].eq(args.focal_cluster_id), "filing_date"].min()

    group_df = compute_entity_temporal(tfo, "object_pdf_group", args.focal_cluster_id, global_focal_start)
    hash_df = compute_entity_temporal(tfo, "object_sha256", args.focal_cluster_id, global_focal_start)

    cluster_df = compute_cluster_ranges(docs_u, tfo, args.focal_cluster_id)

    yearly = (
        docs_u.assign(is_focal=docs_u["cluster_id_num"].eq(args.focal_cluster_id), year=docs_u["filing_date"].dt.year)
        .groupby(["year", "is_focal"])["doc_sha256"].nunique()
        .unstack(fill_value=0)
        .rename(columns={False: "nonfocal_docs", True: "focal_docs"})
        .reset_index()
    )
    yearly["total_docs"] = yearly["nonfocal_docs"] + yearly["focal_docs"]

    out_group = args.out_dir / "10r2_temporal_group_ranges.csv"
    out_hash = args.out_dir / "10r2_temporal_hash_ranges.csv"
    out_cluster = args.out_dir / "10r2_temporal_cluster_ranges.csv"
    out_yearly = args.out_dir / "10r2_temporal_yearly_focal_vs_nonfocal.csv"
    out_top = args.out_dir / "10r2_temporal_top_signals.csv"

    save_csv(group_df, out_group)
    save_csv(hash_df, out_hash)
    save_csv(cluster_df, out_cluster)
    save_csv(yearly, out_yearly)

    top_signals = pd.concat([
        group_df.assign(entity_type="group", entity_value=group_df["object_pdf_group"].astype(str))[[
            "entity_type", "entity_value", "temporal_pressure_score", "n_docs_all",
            "n_docs_focal", "n_docs_nonfocal", "gap_nonfocal_before_focal_days",
            "pre_entity_nonfocal_pct", "span_all_days"
        ]],
        hash_df.assign(entity_type="hash", entity_value=hash_df["object_sha256"].astype(str))[[
            "entity_type", "entity_value", "temporal_pressure_score", "n_docs_all",
            "n_docs_focal", "n_docs_nonfocal", "gap_nonfocal_before_focal_days",
            "pre_entity_nonfocal_pct", "span_all_days"
        ]],
    ], ignore_index=True).sort_values("temporal_pressure_score", ascending=False)
    save_csv(top_signals, out_top)

    plot_group_ranges(
        group_df,
        args.out_dir / "10r2_plot_group_timespan_ranges.png",
        global_focal_start,
        top_n=args.top_n_groups_plot
    )
    plot_hash_gaps(hash_df, args.out_dir / "10r2_plot_hash_gap_days.png")
    plot_yearly(yearly, args.out_dir / "10r2_plot_yearly_docs_focal_vs_nonfocal.png")

    out_md = args.out_dir / "10r2_temporal_report.md"
    make_markdown_report(
        out_md,
        args.focal_cluster_id,
        docs_u,
        tfo,
        group_df,
        hash_df,
        cluster_df,
        yearly,
        global_focal_start,
    )

    summary = {
        "script": "10_font_tracking__run02__temporal_hypothesis_one_shot_v1.py",
        "timestamp_utc": datetime.utcnow().isoformat() + "Z",
        "focal_cluster_id": args.focal_cluster_id,
        "global_focal_start": global_focal_start.strftime("%Y-%m-%d") if pd.notna(global_focal_start) else None,
        "counts": {
            "docs_unique": int(docs_u["doc_sha256"].nunique()),
            "clusters": int(docs_u["cluster_id_num"].nunique()),
            "cases": int(docs_u["case_id"].nunique()) if "case_id" in docs_u.columns else None,
            "tracking_rows": int(len(tfo)),
            "tracking_hashes": int(tfo["object_sha256"].nunique()),
            "tracking_groups": int(tfo["object_pdf_group"].nunique()),
        },
        "outputs": {
            "group_ranges_csv": str(out_group),
            "hash_ranges_csv": str(out_hash),
            "cluster_ranges_csv": str(out_cluster),
            "yearly_csv": str(out_yearly),
            "top_signals_csv": str(out_top),
            "plot_group_ranges_png": str(args.out_dir / "10r2_plot_group_timespan_ranges.png"),
            "plot_hash_gaps_png": str(args.out_dir / "10r2_plot_hash_gap_days.png"),
            "plot_yearly_png": str(args.out_dir / "10r2_plot_yearly_docs_focal_vs_nonfocal.png"),
            "report_md": str(out_md),
        }
    }
    out_json = args.out_dir / "10r2_temporal_summary.json"
    out_json.write_text(json.dumps(summary, indent=2), encoding="utf-8")
    log(f"[ok] {out_json}")

    log("\n=== DONE: Run 2 temporal outputs generated ===")


if __name__ == "__main__":
    main()
