#!/usr/bin/env python3
"""
MCRO Flags – Cluster-focused report
===================================

Scan a directory of MCRO JSON files and produce cluster-centric reports
showing how flags (and their exact value combos) distribute across
clusters, cases, and documents.

ASSUMED JSON STRUCTURE
----------------------

Top-level keys (all optional, script is robust to missing ones):

  doc_sha256        (or sha256)
  filename
  filing_type
  filing_date

  case.case_id
  case.cluster_id
  case.cluster_name

  flags.labels      -> list of dicts, each like:
      {
        "flag_label": "...",
        "data_1": "...",
        "data_2": "...",
        "data_3": "...",
        "object_sha256_1": "...",
        "object_sha256_2": "...",
        "data_url_1": "...",
        "data_url_2": "...",
        ...
      }

We IGNORE data_url_1 / data_url_2.

OUTPUTS (all CSVs in --outdir)
------------------------------

1) flags_cluster_overview.csv
   One row per cluster_id/cluster_name:
     cluster_id, cluster_name,
     n_docs, n_cases,
     case_ids,
     case_doc_counts,      (e.g. "27-CR-23-1886:12; 27-CR-22-9999:4")
     n_flag_instances,
     n_flagged_docs

2) flags_cluster_flag_label_counts.csv
   One row per (cluster_id, cluster_name, flag_label):
     cluster_id, cluster_name, flag_label,
     n_instances, n_docs, n_cases,
     first_filing_date, last_filing_date

3) flags_cluster_flag_value_detail.csv
   One row per (cluster_id, cluster_name, flag_label,
                data_1, data_2, data_3,
                object_sha256_1, object_sha256_2)
   BUT ONLY where n_docs >= 2 (intra-cluster duplicates).

   Columns:
     cluster_id, cluster_name,
     flag_label,
     data_1, data_2, data_3,
     object_sha256_1, object_sha256_2,
     n_instances, n_docs, n_cases,
     case_ids,
     doc_sha256s,
     first_filing_date, last_filing_date

4) flags_multicluster_flag_value_links.csv
   One row per value combo across ALL clusters (no cluster in grouping),
   ONLY where n_clusters >= 2 (i.e., linking multiple clusters).

   Grouped by:
     flag_label,
     data_1, data_2, data_3,
     object_sha256_1, object_sha256_2

   Columns:
     flag_label,
     data_1, data_2, data_3,
     object_sha256_1, object_sha256_2,
     n_instances, n_docs, n_cases, n_clusters,
     cluster_ids, cluster_names,
     case_ids, doc_sha256s,
     first_filing_date, last_filing_date

USAGE
=====

  python3 mcro_flags_cluster_report.py \
      --input /path/to/json_dir \
      --outdir /path/to/flags_cluster_report \
      [--pattern "*.json"] \
      [--recurse]

"""

import argparse
import json
from pathlib import Path
from typing import Any, Dict, List
import sys

import pandas as pd


def as_list(x):
    if x is None:
        return []
    if isinstance(x, list):
        return x
    return [x]


def join_unique(vals):
    """Join unique, non-null values with '; '."""
    vals = [v for v in vals if pd.notna(v)]
    if not vals:
        return ""
    return "; ".join(sorted(set(map(str, vals))))


def main():
    ap = argparse.ArgumentParser(description="Cluster-focused flags report for MCRO JSON files.")
    ap.add_argument("--input", required=True, help="Directory containing JSON files")
    ap.add_argument("--outdir", required=True, help="Output directory for CSVs")
    ap.add_argument("--pattern", default="*.json", help='Glob pattern for JSON files (default: "*.json")')
    ap.add_argument("--recurse", action="store_true", help="Recurse into subdirectories")
    args = ap.parse_args()

    in_dir = Path(args.input)
    if args.recurse:
        files = sorted(in_dir.rglob(args.pattern))
    else:
        files = sorted(in_dir.glob(args.pattern))

    if not files:
        print(f"[warn] No files matched pattern {args.pattern} under {in_dir}", file=sys.stderr)
        sys.exit(0)

    outdir = Path(args.outdir)
    outdir.mkdir(parents=True, exist_ok=True)

    # Collect per-doc and per-flag rows
    doc_rows: List[Dict[str, Any]] = []
    inst_rows: List[Dict[str, Any]] = []

    for fp in files:
        try:
            data = json.loads(fp.read_text(encoding="utf-8"))
        except Exception as e:
            print(f"[warn] Failed to parse JSON: {fp} ({e})", file=sys.stderr)
            continue

        if not isinstance(data, dict):
            print(f"[warn] Skipping {fp}: JSON root is not an object", file=sys.stderr)
            continue

        doc_sha256 = data.get("doc_sha256") or data.get("sha256") or ""
        filename = data.get("filename", "")

        filing_type = data.get("filing_type", "")
        filing_date = data.get("filing_date", "")

        case = data.get("case") or {}
        case_id = case.get("case_id", "")
        cluster_id = case.get("cluster_id", "")
        cluster_name = case.get("cluster_name", "")

        # record doc info (we'll dedupe later)
        doc_rows.append(
            {
                "doc_sha256": doc_sha256,
                "filename": filename,
                "case_id": case_id,
                "cluster_id": cluster_id,
                "cluster_name": cluster_name,
                "filing_type": filing_type,
                "filing_date": filing_date,
            }
        )

        flags = data.get("flags") or {}
        labels = as_list(flags.get("labels"))

        if not labels:
            continue

        for lab in labels:
            if not isinstance(lab, dict):
                continue

            row = {
                "doc_sha256": doc_sha256,
                "filename": filename,
                "case_id": case_id,
                "cluster_id": cluster_id,
                "cluster_name": cluster_name,
                "filing_type": filing_type,
                "filing_date": filing_date,
                "flag_label": lab.get("flag_label", ""),
                "data_1": lab.get("data_1", ""),
                "data_2": lab.get("data_2", ""),
                "data_3": lab.get("data_3", ""),
                "object_sha256_1": lab.get("object_sha256_1", ""),
                "object_sha256_2": lab.get("object_sha256_2", ""),
            }
            inst_rows.append(row)

    # If there are no flags at all, bail with empty shells
    if not inst_rows:
        print("[info] No flags.labels entries found in any JSON files.")
        # Write empty shells so downstream logic doesn't explode
        pd.DataFrame().to_csv(outdir / "flags_cluster_overview.csv", index=False)
        pd.DataFrame().to_csv(outdir / "flags_cluster_flag_label_counts.csv", index=False)
        pd.DataFrame().to_csv(outdir / "flags_cluster_flag_value_detail.csv", index=False)
        pd.DataFrame().to_csv(outdir / "flags_multicluster_flag_value_links.csv", index=False)
        print(f"[ok] Wrote empty flags cluster report shells to {outdir.resolve()}")
        sys.exit(0)

    # Build DataFrames
    doc_df = pd.DataFrame(doc_rows)
    inst_df = pd.DataFrame(inst_rows)

    # Normalize empties: treat missing cluster_id as empty string
    for col in ["cluster_id", "cluster_name", "case_id", "filing_date"]:
        if col in doc_df.columns:
            doc_df[col] = doc_df[col].fillna("")
        if col in inst_df.columns:
            inst_df[col] = inst_df[col].fillna("")

    # For cluster-level stats, drop rows with empty cluster_id
    doc_cluster_df = doc_df[doc_df["cluster_id"].astype(str).str.strip() != ""].copy()
    inst_cluster_df = inst_df[inst_df["cluster_id"].astype(str).str.strip() != ""].copy()

    # -----------------------
    # 1) flags_cluster_overview.csv
    # -----------------------

    if not doc_cluster_df.empty:
        # Unique docs
        docs_unique = doc_cluster_df.drop_duplicates(subset=["doc_sha256"])

        # Basic per-cluster counts
        cluster_over = (
            docs_unique.groupby(["cluster_id", "cluster_name"], as_index=False)
            .agg(
                n_docs=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
                n_cases=("case_id", lambda s: int(pd.Series(s).dropna().nunique())),
                case_ids=("case_id", join_unique),
            )
        )

        # Docs per case within each cluster
        case_counts = (
            docs_unique.groupby(["cluster_id", "cluster_name", "case_id"], as_index=False)
            .agg(doc_count=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())))
        )
        # Build "case_id:doc_count" pairs
        case_counts["pair"] = case_counts.apply(
            lambda r: f"{r['case_id']}:{r['doc_count']}", axis=1
        )
        case_doc_counts = (
            case_counts.groupby(["cluster_id", "cluster_name"], as_index=False)
            .agg(case_doc_counts=("pair", lambda s: "; ".join(s)))
        )

        cluster_over = cluster_over.merge(
            case_doc_counts, on=["cluster_id", "cluster_name"], how="left"
        )

        # Add flags counts per cluster
        if not inst_cluster_df.empty:
            flag_counts = (
                inst_cluster_df.groupby(["cluster_id", "cluster_name"], as_index=False)
                .agg(
                    n_flag_instances=("flag_label", "size"),
                    n_flagged_docs=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
                )
            )
            cluster_over = cluster_over.merge(
                flag_counts, on=["cluster_id", "cluster_name"], how="left"
            )
        else:
            cluster_over["n_flag_instances"] = 0
            cluster_over["n_flagged_docs"] = 0

        # Fill NaNs
        for col in ["n_flag_instances", "n_flagged_docs"]:
            if col in cluster_over.columns:
                cluster_over[col] = cluster_over[col].fillna(0).astype(int)

        cluster_over.to_csv(outdir / "flags_cluster_overview.csv", index=False)
    else:
        pd.DataFrame().to_csv(outdir / "flags_cluster_overview.csv", index=False)

    # -----------------------
    # 2) flags_cluster_flag_label_counts.csv
    # -----------------------

    if not inst_cluster_df.empty:
        cluster_label = (
            inst_cluster_df.groupby(["cluster_id", "cluster_name", "flag_label"], as_index=False)
            .agg(
                n_instances=("flag_label", "size"),
                n_docs=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
                n_cases=("case_id", lambda s: int(pd.Series(s).dropna().nunique())),
                first_filing_date=("filing_date", lambda s: s.dropna().min() if len(s.dropna()) else ""),
                last_filing_date=("filing_date", lambda s: s.dropna().max() if len(s.dropna()) else ""),
            )
        )
        cluster_label.to_csv(outdir / "flags_cluster_flag_label_counts.csv", index=False)
    else:
        pd.DataFrame().to_csv(outdir / "flags_cluster_flag_label_counts.csv", index=False)

    # -----------------------
    # 3) flags_cluster_flag_value_detail.csv (intra-cluster duplicates)
    # -----------------------

    if not inst_cluster_df.empty:
        combo_cols_cluster = [
            "cluster_id", "cluster_name",
            "flag_label",
            "data_1", "data_2", "data_3",
            "object_sha256_1", "object_sha256_2",
        ]

        combo_cluster = (
            inst_cluster_df.groupby(combo_cols_cluster, as_index=False)
            .agg(
                n_instances=("flag_label", "size"),
                n_docs=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
                n_cases=("case_id", lambda s: int(pd.Series(s).dropna().nunique())),
                case_ids=("case_id", join_unique),
                doc_sha256s=("doc_sha256", join_unique),
                first_filing_date=("filing_date", lambda s: s.dropna().min() if len(s.dropna()) else ""),
                last_filing_date=("filing_date", lambda s: s.dropna().max() if len(s.dropna()) else ""),
            )
        )

        # Keep only combos that appear in >= 2 docs within that cluster
        combo_cluster_dup = combo_cluster[combo_cluster["n_docs"] >= 2].copy()

        combo_cluster_dup.to_csv(outdir / "flags_cluster_flag_value_detail.csv", index=False)
    else:
        pd.DataFrame().to_csv(outdir / "flags_cluster_flag_value_detail.csv", index=False)

    # -----------------------
    # 4) flags_multicluster_flag_value_links.csv (cross-cluster linkers)
    # -----------------------

    # Here we use ALL inst_df (including empty cluster_id rows), because we care about
    # how value combos link clusters together across the entire dataset.
    combo_cols_global = [
        "flag_label",
        "data_1", "data_2", "data_3",
        "object_sha256_1", "object_sha256_2",
    ]

    combo_global = (
        inst_df.groupby(combo_cols_global, as_index=False)
        .agg(
            n_instances=("flag_label", "size"),
            n_docs=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
            n_cases=("case_id", lambda s: int(pd.Series(s).dropna().nunique())),
            n_clusters=("cluster_id", lambda s: int(pd.Series(s).dropna().nunique())),
            cluster_ids=("cluster_id", join_unique),
            cluster_names=("cluster_name", join_unique),
            case_ids=("case_id", join_unique),
            doc_sha256s=("doc_sha256", join_unique),
            first_filing_date=("filing_date", lambda s: s.dropna().min() if len(s.dropna()) else ""),
            last_filing_date=("filing_date", lambda s: s.dropna().max() if len(s.dropna()) else ""),
        )
    )

    # Keep only combos that link 2+ clusters
    combo_global_multi = combo_global[combo_global["n_clusters"] >= 2].copy()

    combo_global_multi.to_csv(outdir / "flags_multicluster_flag_value_links.csv", index=False)

    print(f"[ok] Flags cluster report written to: {outdir.resolve()}")
    print(f"  flags_cluster_overview.csv               : {len(doc_cluster_df.drop_duplicates(subset=['cluster_id','cluster_name']))} clusters" if not doc_cluster_df.empty else "  flags_cluster_overview.csv               : 0 rows")
    print(f"  flags_cluster_flag_label_counts.csv      : {len(inst_cluster_df.groupby(['cluster_id','cluster_name','flag_label'])) if not inst_cluster_df.empty else 0} rows")
    print(f"  flags_cluster_flag_value_detail.csv      : {len(combo_cluster_dup) if not inst_cluster_df.empty else 0} rows (intra-cluster duplicates)")
    print(f"  flags_multicluster_flag_value_links.csv  : {len(combo_global_multi)} rows (cross-cluster linkers)")


if __name__ == "__main__":
    main()
