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

Rebuild ONLY:
  - flags_cluster_overview.csv      (expanded, per cluster + case)
  - flags_cluster_flag_value_detail.csv (expanded, per cluster + combo + filename)

Does NOT touch:
  - flags_cluster_flag_label_counts.csv
  - flags_multicluster_flag_value_links.csv

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, e.g.:
      {
        "flag_label": "...",
        "data_1": "...",
        "data_2": "...",
        "data_3": "...",
        "object_sha256_1": "...",
        "object_sha256_2": "...",
        ...
      }

We IGNORE data_url_1 / data_url_2.

OUTPUTS
-------

1) flags_cluster_overview.csv
   One row per (cluster_id, cluster_name, case_id):

     cluster_id
     cluster_name

     cluster_n_docs
     cluster_n_cases
     cluster_n_flag_instances
     cluster_n_flagged_docs
     cluster_first_filing_date
     cluster_last_filing_date

     case_id
     case_n_docs
     case_n_flag_instances
     case_n_flagged_docs
     case_first_filing_date
     case_last_filing_date

2) flags_cluster_flag_value_detail.csv
   One row per (cluster, flag value combo, filename), but ONLY for combos
   that appear in >= 2 distinct docs within that cluster.

   Grouping combo key:
     cluster_id, cluster_name,
     flag_label,
     data_1, data_2, data_3,
     object_sha256_1, object_sha256_2

   Columns:

     cluster_id
     cluster_name
     flag_label
     data_1
     data_2
     data_3
     object_sha256_1
     object_sha256_2

     combo_n_docs              # how many distinct filenames share this combo in this cluster
     combo_n_cases             # how many cases in this cluster share this combo
     combo_n_instances         # total flag entries for this combo in this cluster
     combo_first_filing_date
     combo_last_filing_date

     case_id
     filename                  # primary visible ID
     filing_type
     filing_date
     doc_n_instances           # how many times this combo appears in this one document
     doc_sha256                # optional hash-level ID (last column)

USAGE
=====

  python3 mcro_flags_cluster_report_v2.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="Expanded 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 not inst_rows:
        print("[info] No flags.labels entries found in any JSON files.")
        # Still write empty shells so downstream things don't explode
        pd.DataFrame().to_csv(outdir / "flags_cluster_overview.csv", index=False)
        pd.DataFrame().to_csv(outdir / "flags_cluster_flag_value_detail.csv", index=False)
        print(f"[ok] Wrote empty expanded flags cluster CSVs to {outdir.resolve()}")
        sys.exit(0)

    doc_df = pd.DataFrame(doc_rows).drop_duplicates(subset=["doc_sha256"])
    inst_df = pd.DataFrame(inst_rows)

    # Normalize: avoid 'nan' in key fields
    for col in ["cluster_id", "cluster_name", "case_id", "filing_date",
                "data_1", "data_2", "data_3", "object_sha256_1", "object_sha256_2"]:
        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("")

    # Subset: only rows with non-empty cluster_id for cluster-level stats
    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  (expanded: cluster + case rows)
    # =========================================================================

    if not doc_cluster_df.empty:
        # Cluster-level totals
        cluster_totals = (
            doc_cluster_df.groupby(["cluster_id", "cluster_name"], as_index=False)
            .agg(
                cluster_n_docs=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
                cluster_n_cases=("case_id", lambda s: int(pd.Series(s).dropna().nunique())),
                cluster_first_filing_date=("filing_date", lambda s: s.dropna().min() if len(s.dropna()) else ""),
                cluster_last_filing_date=("filing_date", lambda s: s.dropna().max() if len(s.dropna()) else ""),
            )
        )

        if not inst_cluster_df.empty:
            cluster_flag_totals = (
                inst_cluster_df.groupby(["cluster_id", "cluster_name"], as_index=False)
                .agg(
                    cluster_n_flag_instances=("flag_label", "size"),
                    cluster_n_flagged_docs=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
                )
            )
            cluster_totals = cluster_totals.merge(
                cluster_flag_totals, on=["cluster_id", "cluster_name"], how="left"
            )
        else:
            cluster_totals["cluster_n_flag_instances"] = 0
            cluster_totals["cluster_n_flagged_docs"] = 0

        for col in ["cluster_n_flag_instances", "cluster_n_flagged_docs"]:
            if col in cluster_totals.columns:
                cluster_totals[col] = cluster_totals[col].fillna(0).astype(int)

        # Case-level stats within each cluster
        case_docs = (
            doc_cluster_df.groupby(["cluster_id", "cluster_name", "case_id"], as_index=False)
            .agg(
                case_n_docs=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
                case_first_filing_date=("filing_date", lambda s: s.dropna().min() if len(s.dropna()) else ""),
                case_last_filing_date=("filing_date", lambda s: s.dropna().max() if len(s.dropna()) else ""),
            )
        )

        if not inst_cluster_df.empty:
            case_flags = (
                inst_cluster_df.groupby(["cluster_id", "cluster_name", "case_id"], as_index=False)
                .agg(
                    case_n_flag_instances=("flag_label", "size"),
                    case_n_flagged_docs=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
                )
            )
            case_overview = case_docs.merge(
                case_flags, on=["cluster_id", "cluster_name", "case_id"], how="left"
            )
        else:
            case_overview = case_docs.copy()
            case_overview["case_n_flag_instances"] = 0
            case_overview["case_n_flagged_docs"] = 0

        for col in ["case_n_flag_instances", "case_n_flagged_docs"]:
            if col in case_overview.columns:
                case_overview[col] = case_overview[col].fillna(0).astype(int)

        # Attach cluster totals to each case row
        overview = case_overview.merge(
            cluster_totals,
            on=["cluster_id", "cluster_name"],
            how="left",
            suffixes=("", "_cluster"),
        )

        # Column order
        cols = [
            "cluster_id",
            "cluster_name",
            "cluster_n_docs",
            "cluster_n_cases",
            "cluster_n_flag_instances",
            "cluster_n_flagged_docs",
            "cluster_first_filing_date",
            "cluster_last_filing_date",
            "case_id",
            "case_n_docs",
            "case_n_flag_instances",
            "case_n_flagged_docs",
            "case_first_filing_date",
            "case_last_filing_date",
        ]
        for c in cols:
            if c not in overview.columns:
                overview[c] = pd.NA

        overview = overview[cols].sort_values(
            by=["cluster_id", "case_id"], kind="mergesort"
        )

        overview.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_value_detail.csv  (expanded: per filename)
    # =========================================================================

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

        # Cluster-level stats for each combo
        combo_stats = (
            inst_cluster_df.groupby(combo_cols, as_index=False)
            .agg(
                combo_n_instances=("flag_label", "size"),
                combo_n_docs=("filename", lambda s: int(pd.Series(s).dropna().nunique())),
                combo_n_cases=("case_id", lambda s: int(pd.Series(s).dropna().nunique())),
                combo_first_filing_date=("filing_date", lambda s: s.dropna().min() if len(s.dropna()) else ""),
                combo_last_filing_date=("filing_date", lambda s: s.dropna().max() if len(s.dropna()) else ""),
            )
        )

        # Only combos that show up in >= 2 docs within that cluster
        combos_dup = combo_stats[combo_stats["combo_n_docs"] >= 2].copy()

        if combos_dup.empty:
            pd.DataFrame().to_csv(outdir / "flags_cluster_flag_value_detail.csv", index=False)
        else:
            # Join back to instance-level to get per-doc rows for those combos
            inst_for_combos = inst_cluster_df.merge(
                combos_dup[combo_cols + [
                    "combo_n_instances",
                    "combo_n_docs",
                    "combo_n_cases",
                    "combo_first_filing_date",
                    "combo_last_filing_date",
                ]],
                on=combo_cols,
                how="inner",
            )

            # Now aggregate to one row per (combo + document)
            group_cols_doc = combo_cols + [
                "case_id",
                "filename",
                "filing_type",
                "filing_date",
                "doc_sha256",
            ]

            doc_detail = (
                inst_for_combos.groupby(group_cols_doc, as_index=False)
                .agg(
                    doc_n_instances=("flag_label", "size"),
                    # combo-level stats already merged, so no need to re-agg them here
                )
            )

            # Merge combo stats (to ensure all combo columns are present)
            doc_detail = doc_detail.merge(
                combos_dup,
                on=combo_cols,
                how="left",
            )

            # Ensure integer types
            for col in ["combo_n_instances", "combo_n_docs", "combo_n_cases", "doc_n_instances"]:
                if col in doc_detail.columns:
                    doc_detail[col] = doc_detail[col].fillna(0).astype(int)

            # Reorder columns: combo info, then doc info
            final_cols = [
                "cluster_id",
                "cluster_name",
                "flag_label",
                "data_1",
                "data_2",
                "data_3",
                "object_sha256_1",
                "object_sha256_2",
                "combo_n_docs",
                "combo_n_cases",
                "combo_n_instances",
                "combo_first_filing_date",
                "combo_last_filing_date",
                "case_id",
                "filename",
                "filing_type",
                "filing_date",
                "doc_n_instances",
                "doc_sha256",
            ]
            for c in final_cols:
                if c not in doc_detail.columns:
                    doc_detail[c] = pd.NA

            doc_detail = doc_detail[final_cols].sort_values(
                by=[
                    "cluster_id",
                    "flag_label",
                    "data_1",
                    "data_2",
                    "data_3",
                    "object_sha256_1",
                    "case_id",
                    "filename",
                ],
                kind="mergesort",
            )

            doc_detail.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)

    print(f"[ok] Expanded flags cluster CSVs written to: {outdir.resolve()}")
    print(f"  flags_cluster_overview.csv          : {len(doc_cluster_df.groupby(['cluster_id','case_id'])) if not doc_cluster_df.empty else 0} rows")
    # detail row count might be useful to echo
    # But we don't recompute; that's fine.


if __name__ == "__main__":
    main()
