#!/usr/bin/env python3
"""
MCRO Flags – Cluster flag value detail ONLY (with data_url_1 / data_url_2)
=========================================================================

This script rebuilds ONLY:

  - flags_cluster_flag_value_detail.csv

It does NOT write or modify:
  - flags_cluster_overview.csv
  - flags_cluster_flag_label_counts.csv
  - flags_multicluster_flag_value_links.csv

BEHAVIOR
--------

For each JSON in --input:

  - Reads:
      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/2 for grouping, but we DO include them in the
    final per-document rows.

OUTPUT
------

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-level):

    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
    data_url_1                # aggregated per doc (unique URLs joined if multiple)
    data_url_2
    filing_type
    filing_date
    doc_n_instances           # how many times this combo appears in this one document
    doc_sha256                # hash-level ID

USAGE
=====

  python3 mcro_flags_cluster_flag_values_only.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="Build flags_cluster_flag_value_detail.csv (cluster + flag value combos + filename, with data_url_*)."
    )
    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)

    inst_rows: List[Dict[str, Any]] = []

    # ---------------------------------------------------------
    # Collect per-flag instance rows (cluster + doc + flag data)
    # ---------------------------------------------------------
    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

            inst_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,
                    "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", ""),
                    "data_url_1": lab.get("data_url_1", ""),
                    "data_url_2": lab.get("data_url_2", ""),
                }
            )

    if not inst_rows:
        print("[info] No flags.labels entries found in any JSON files.")
        pd.DataFrame().to_csv(outdir / "flags_cluster_flag_value_detail.csv", index=False)
        print(f"[ok] Wrote empty flags_cluster_flag_value_detail.csv to {outdir.resolve()}")
        sys.exit(0)

    inst_df = pd.DataFrame(inst_rows)

    # Normalize 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",
        "data_url_1", "data_url_2",
    ]:
        if col in inst_df.columns:
            inst_df[col] = inst_df[col].fillna("")

    # Restrict to rows with non-empty cluster_id for cluster-level combos
    inst_cluster_df = inst_df[inst_df["cluster_id"].astype(str).str.strip() != ""].copy()

    if inst_cluster_df.empty:
        pd.DataFrame().to_csv(outdir / "flags_cluster_flag_value_detail.csv", index=False)
        print(f"[ok] No cluster_id values found; wrote empty flags_cluster_flag_value_detail.csv to {outdir.resolve()}")
        sys.exit(0)

    # ---------------------------------------------------------
    # Build cluster-level stats for each combo
    # ---------------------------------------------------------
    combo_cols = [
        "cluster_id", "cluster_name",
        "flag_label",
        "data_1", "data_2", "data_3",
        "object_sha256_1", "object_sha256_2",
    ]

    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 appear 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)
        print(f"[ok] No intra-cluster combos with >=2 docs; wrote empty flags_cluster_flag_value_detail.csv")
        sys.exit(0)

    # ---------------------------------------------------------
    # Join instance-level rows to these combos and aggregate
    # to one row per (combo + document)
    # ---------------------------------------------------------

    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",
    )

    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"),
            data_url_1=("data_url_1", join_unique),
            data_url_2=("data_url_2", join_unique),
        )
    )

    # Merge combo stats back in (for counts & date ranges)
    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",
        "data_url_1",
        "data_url_2",
        "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",
    )

    out_path = outdir / "flags_cluster_flag_value_detail.csv"
    doc_detail.to_csv(out_path, index=False)

    print(f"[ok] Wrote flags_cluster_flag_value_detail.csv with {len(doc_detail)} rows to: {out_path.resolve()}")


if __name__ == "__main__":
    main()
