#!/usr/bin/env python3
"""
MCRO Flags Report
=================

Scan a directory of MCRO JSON files and produce a focused report
on the "flags" (flags.labels) across cases/clusters/docs.

ASSUMED JSON STRUCTURE (per doc)
--------------------------------
Top-level keys used (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 for all analysis.

OUTPUTS
=======

All CSVs go into --outdir:

1) flags_instances.csv
   One row per flag instance (per label entry per doc):

      doc_sha256
      filename
      case_id
      cluster_id
      cluster_name
      filing_type
      filing_date

      flag_label
      data_1
      data_2
      data_3
      object_sha256_1
      object_sha256_2

2) flags_label_overview.csv
   One row per flag_label, with high-level stats:

      flag_label
      n_instances        (number of label entries)
      n_docs             (distinct doc_sha256)
      n_cases            (distinct case_id)
      n_clusters         (distinct cluster_id)
      first_filing_date
      last_filing_date
      n_unique_data_1
      n_unique_data_2
      n_unique_data_3
      n_unique_object_sha256_1
      n_unique_object_sha256_2

3) flags_label_value_overview.csv
   One row per unique value combo:

      flag_label
      data_1
      data_2
      data_3
      object_sha256_1
      object_sha256_2

      n_instances
      n_docs
      n_cases
      n_clusters
      case_ids           (semicolon-joined)
      cluster_ids        (semicolon-joined)
      doc_sha256s        (semicolon-joined)
      first_filing_date
      last_filing_date

USAGE
=====

  python3 mcro_flags_report.py \
      --input /path/to/json_dir \
      --outdir /path/to/flags_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="Produce MCRO flags-focused CSV reports from 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 raw flag instances
    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", "")

        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

            flag_label = lab.get("flag_label", "")

            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": 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 CSV shells so downstream scripts don't explode
        empty_df = pd.DataFrame(
            columns=[
                "doc_sha256", "filename", "case_id", "cluster_id", "cluster_name",
                "filing_type", "filing_date", "flag_label",
                "data_1", "data_2", "data_3",
                "object_sha256_1", "object_sha256_2",
            ]
        )
        empty_df.to_csv(outdir / "flags_instances.csv", index=False)
        pd.DataFrame().to_csv(outdir / "flags_label_overview.csv", index=False)
        pd.DataFrame().to_csv(outdir / "flags_label_value_overview.csv", index=False)
        print(f"[ok] Wrote empty flags report shells to {outdir.resolve()}")
        sys.exit(0)

    inst_df = pd.DataFrame(inst_rows)

    # 1) Raw instances
    inst_df.to_csv(outdir / "flags_instances.csv", index=False)

    # 2) Overview by flag_label
    label_group = (
        inst_df
        .groupby("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())),
            n_clusters=("cluster_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 ""),
            n_unique_data_1=("data_1", lambda s: int(pd.Series(s).dropna().nunique())),
            n_unique_data_2=("data_2", lambda s: int(pd.Series(s).dropna().nunique())),
            n_unique_data_3=("data_3", lambda s: int(pd.Series(s).dropna().nunique())),
            n_unique_object_sha256_1=("object_sha256_1", lambda s: int(pd.Series(s).dropna().nunique())),
            n_unique_object_sha256_2=("object_sha256_2", lambda s: int(pd.Series(s).dropna().nunique())),
        )
    )
    label_group.to_csv(outdir / "flags_label_overview.csv", index=False)

    # 3) Overview by full value combo
    combo_cols = [
        "flag_label",
        "data_1", "data_2", "data_3",
        "object_sha256_1", "object_sha256_2",
    ]

    combo_group = (
        inst_df
        .groupby(combo_cols, 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())),
            case_ids=("case_id", join_unique),
            cluster_ids=("cluster_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 ""),
        )
    )
    combo_group.to_csv(outdir / "flags_label_value_overview.csv", index=False)

    print(f"[ok] Flags report written to: {outdir.resolve()}")
    print(f"     flags_instances.csv              rows: {len(inst_df)}")
    print(f"     flags_label_overview.csv         rows: {len(label_group)}")
    print(f"     flags_label_value_overview.csv   rows: {len(combo_group)}")


if __name__ == "__main__":
    main()
