#!/usr/bin/env python3
"""
MCRO PDF Group Overview (per pdf_group / case / filing_type)
===========================================================

This script builds a *single* big CSV:

    pdf_group_overview.csv

One row per combination of:

    (pdf_group, case_id, cluster_id, cluster_name, filing_type)

It includes:

  * pdf_group = "all"          -> ALL documents in the corpus
  * all real pdf_groups C1/T1… -> based on tracking.groups[*].pdf_group for
                                  docs where tracking.tracked_flag == true

The goal is to overload this table with useful, *simple* per-group counts
and stats, so you can slice/dice and build more focused summary tables
afterwards.

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

Top-level (per doc):

  doc_sha256 or sha256        : unique doc identifier
  filename
  filing_type
  filing_date                 : "YYYY-MM-DD"
  pdf_page_count              : numeric or string
  pdf_file_size               : (optional) numeric (bytes) or string

  case : {
      "case_id"       : ...,
      "cluster_id"    : ...,
      "cluster_name"  : ...,
      ...
  }

  language : {
      "terms": [
          {
            "search_group": ...,
            "search_term":  ...,
            "quantity":     int,
          },
          ...
      ],
      ...
  }

  objects : [
      {
        "object_type": ".ttf" / ".jpg" / ...,
        "object_sha256": "...",
        "object_font_name": "...",   # fonts only
        ...
      },
      ...
  ]

  signatures : {
      "pdfsig": {
          # Possible shapes:
          # 1) {"summary": {"num_signatures": 2}, "signatures": [...]}
          # 2) {"signatures": [ ... ]}  # then len(signatures)
          # 3) Some other dict -> we fallback to 1 signature if present.
      },
      ...
  }

  tracking : {
      "tracked_flag": true / "true" / "TRUE",
      "groups": [
          {
            "pdf_group": "C1",
            "font_group": "...",
            ...
          },
          ...
      ]
  }

OUTPUT
======

pdf_group_overview.csv with columns:

  pdf_group
  case_id
  cluster_id
  cluster_name
  filing_type

  n_docs
  filing_date          # earliest filing_date in this combo
  pdf_pages            # max pdf_page_count in this combo
  pdf_file_size        # avg pdf_file_size in this combo (if present)
  filing_type_count    # same as n_docs

  n_docs_2_or_more_pdf_sigs

  n_term_occurrences
  n_unique_terms
  n_unique_search_groups

  n_objects
  n_unique_object_sha256
  n_unique_font_objects_sha256
  n_unique_font_object_types
  n_unique_object_font_names

USAGE
=====

  python3 mcro_pdf_group_overview_big.py \
      --input  /path/to/json_dir \
      --outdir /path/to/pdf_group_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):
    """Ensure x is a list (None -> [], scalar -> [scalar])."""
    if x is None:
        return []
    if isinstance(x, list):
        return x
    return [x]


def safe_get(d: Dict[str, Any], key: str, default=None):
    if isinstance(d, dict):
        return d.get(key, default)
    return default


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


def compute_n_signatures(pdfsig: Any) -> int:
    """
    Try to infer the number of signatures from a signatures.pdfsig dict.

    Heuristics:
      1) If "signatures" is a list -> len(list)
      2) Else if "summary.num_signatures" or "summary.n_signatures" exists -> that int
      3) Else if pdfsig is a non-empty dict -> assume 1
      4) Else -> 0
    """
    if not isinstance(pdfsig, dict) or not pdfsig:
        return 0

    # 1) "signatures": [ ... ]
    sig_list = pdfsig.get("signatures")
    if isinstance(sig_list, list):
        return len(sig_list)

    # 2) summary.num_signatures / summary.n_signatures
    summary = pdfsig.get("summary") or {}
    for key in ("num_signatures", "n_signatures"):
        if key in summary:
            try:
                return int(summary[key])
            except Exception:
                pass

    # 3) Fallback: any non-empty dict -> 1
    return 1


def main():
    ap = argparse.ArgumentParser(description="MCRO PDF group overview (per pdf_group / case / filing_type).")
    ap.add_argument(
        "--input",
        required=True,
        help="Directory containing JSON files.",
    )
    ap.add_argument(
        "--outdir",
        required=True,
        help="Output directory for pdf_group_overview.csv.",
    )
    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 under --input.",
    )
    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)
    overview_path = outdir / "pdf_group_overview.csv"

    docs_rows: List[Dict[str, Any]] = []
    terms_rows: List[Dict[str, Any]] = []
    objects_rows: List[Dict[str, Any]] = []
    sig_rows: List[Dict[str, Any]] = []

    n_parsed = 0
    n_failed = 0

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

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

        doc_sha256 = data.get("doc_sha256") or data.get("sha256") or ""
        filename = data.get("filename", "")
        filing_type = (data.get("filing_type") or "").strip()
        filing_date = (data.get("filing_date") or "").strip()
        pdf_page_count = data.get("pdf_page_count", None)
        pdf_file_size = data.get("pdf_file_size", None)  # optional

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

        # tracking / pdf_groups
        tracking = data.get("tracking") or {}
        tracked_flag = tracking.get("tracked_flag")
        is_tracked = (
            tracked_flag is True
            or (isinstance(tracked_flag, str) and tracked_flag.strip().lower() == "true")
        )
        groups = as_list(tracking.get("groups")) if is_tracked else []

        # language / terms
        language = data.get("language") or {}
        terms = as_list(language.get("terms"))

        # objects
        objects = as_list(data.get("objects"))

        # signatures
        signatures = data.get("signatures") or {}
        pdfsig = safe_get(signatures, "pdfsig", {})
        n_sigs = compute_n_signatures(pdfsig)

        # -------------------------
        # Helper: add rows for a given pdf_group label
        # -------------------------
        def add_for_pdf_group(pg_label: str):
            # docs row
            docs_rows.append(
                {
                    "pdf_group": pg_label,
                    "doc_sha256": doc_sha256,
                    "filename": filename,
                    "filing_type": filing_type,
                    "filing_date": filing_date,
                    "pdf_page_count": pdf_page_count,
                    "pdf_file_size": pdf_file_size,
                    "case_id": case_id,
                    "cluster_id": cluster_id,
                    "cluster_name": cluster_name,
                }
            )

            # terms rows
            for term in terms:
                if not isinstance(term, dict):
                    continue
                search_group = term.get("search_group")
                search_term = term.get("search_term")
                quantity = term.get("quantity")
                terms_rows.append(
                    {
                        "pdf_group": pg_label,
                        "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,
                        "search_group": search_group,
                        "search_term": search_term,
                        "quantity": quantity,
                    }
                )

            # objects rows
            for obj in objects:
                if not isinstance(obj, dict):
                    continue
                row = {
                    "pdf_group": pg_label,
                    "doc_sha256": doc_sha256,
                    "filename": filename,
                    "case_id": case_id,
                    "cluster_id": cluster_id,
                    "cluster_name": cluster_name,
                    "filing_type": filing_type,
                }
                for k, v in obj.items():
                    row[k] = v
                objects_rows.append(row)

            # signatures rows
            # (we only care about n_signatures per doc here)
            if n_sigs > 0:
                sig_rows.append(
                    {
                        "pdf_group": pg_label,
                        "doc_sha256": doc_sha256,
                        "filename": filename,
                        "case_id": case_id,
                        "cluster_id": cluster_id,
                        "cluster_name": cluster_name,
                        "filing_type": filing_type,
                        "n_signatures": n_sigs,
                    }
                )

        # 1) ALWAYS add "all" pdf_group row for every doc
        add_for_pdf_group("all")

        # 2) For tracked docs with pdf_groups, add one row per group
        if is_tracked and groups:
            for g in groups:
                if not isinstance(g, dict):
                    continue
                raw_pg = g.get("pdf_group")
                if raw_pg is None:
                    continue
                pg_label = str(raw_pg).strip()
                if not pg_label:
                    continue
                add_for_pdf_group(pg_label)

        n_parsed += 1

    # ------------------------------------------------------------------
    # Build DataFrames
    # ------------------------------------------------------------------
    if not docs_rows:
        overview_path.write_text("", encoding="utf-8")
        print(f"[info] No docs found; wrote empty {overview_path}", file=sys.stderr)
        sys.exit(0)

    docs_df = pd.DataFrame(docs_rows)
    terms_df = pd.DataFrame(terms_rows) if terms_rows else pd.DataFrame()
    objects_df = pd.DataFrame(objects_rows) if objects_rows else pd.DataFrame()
    sig_df = pd.DataFrame(sig_rows) if sig_rows else pd.DataFrame()

    # Normalize obvious fields
    for col in ["pdf_group", "case_id", "cluster_id", "cluster_name", "filing_type"]:
        if col in docs_df.columns:
            docs_df[col] = docs_df[col].fillna("").astype(str)

    # Grouping keys for overview rows
    group_keys = ["pdf_group", "case_id", "cluster_id", "cluster_name", "filing_type"]

    # Ensure doc-level uniqueness before grouping
    docs_overview = docs_df.drop_duplicates(
        subset=["pdf_group", "doc_sha256", "case_id", "cluster_id", "cluster_name", "filing_type"]
    ).copy()

    # ------------------------------------------------------------------
    # Core doc-level grouping
    # ------------------------------------------------------------------
    def agg_filing_date(s: pd.Series):
        s = s.dropna()
        return s.min() if len(s) else pd.NA

    def agg_pdf_pages(s: pd.Series):
        nums = pd.to_numeric(s, errors="coerce")
        nums = nums.dropna()
        if len(nums) == 0:
            return pd.NA
        # Max pages is a good canonical value; clones will share it.
        return int(nums.max())

    def agg_pdf_file_size(s: pd.Series):
        nums = pd.to_numeric(s, errors="coerce")
        nums = nums.dropna()
        if len(nums) == 0:
            return pd.NA
        # Average size in bytes
        return float(nums.mean())

    docs_group = (
        docs_overview.groupby(group_keys, as_index=False)
        .agg(
            n_docs=("doc_sha256", "nunique"),
            filing_date=("filing_date", agg_filing_date),
            pdf_pages=("pdf_page_count", agg_pdf_pages),
            pdf_file_size=("pdf_file_size", agg_pdf_file_size),
        )
    )

    # We'll also set filing_type_count = n_docs later.

    # ------------------------------------------------------------------
    # Term stats per group
    # ------------------------------------------------------------------
    if not terms_df.empty:
        terms_df["quantity"] = pd.to_numeric(terms_df["quantity"], errors="coerce").fillna(0).astype(int)

        # Align grouping keys
        for col in ["pdf_group", "case_id", "cluster_id", "cluster_name", "filing_type"]:
            if col in terms_df.columns:
                terms_df[col] = terms_df[col].fillna("").astype(str)

        term_group = (
            terms_df.groupby(group_keys, as_index=False)
            .agg(
                n_term_occurrences=("quantity", "sum"),
                n_unique_terms=("search_term", lambda s: int(pd.Series(s).dropna().nunique())),
                n_unique_search_groups=("search_group", lambda s: int(pd.Series(s).dropna().nunique())),
            )
        )
    else:
        term_group = pd.DataFrame(columns=group_keys + ["n_term_occurrences", "n_unique_terms", "n_unique_search_groups"])

    # ------------------------------------------------------------------
    # Object stats per group
    # ------------------------------------------------------------------
    if not objects_df.empty:
        for col in ["pdf_group", "case_id", "cluster_id", "cluster_name", "filing_type"]:
            if col in objects_df.columns:
                objects_df[col] = objects_df[col].fillna("").astype(str)

        # All objects
        obj_group_all = (
            objects_df.groupby(group_keys, as_index=False)
            .agg(
                n_objects=("doc_sha256", "size"),
                n_unique_object_sha256=("object_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
            )
        )

        # Font-like objects
        font_types = {".cff", ".cid", ".otf", ".pfa", ".ttf"}
        if "object_type" in objects_df.columns:
            font_df = objects_df[
                objects_df["object_type"].astype(str).str.lower().isin(font_types)
            ].copy()
        else:
            font_df = objects_df.iloc[0:0].copy()  # empty

        if not font_df.empty:
            font_group = (
                font_df.groupby(group_keys, as_index=False)
                .agg(
                    n_unique_font_objects_sha256=(
                        "object_sha256",
                        lambda s: int(pd.Series(s).dropna().nunique()),
                    ),
                    n_unique_font_object_types=(
                        "object_type",
                        lambda s: int(pd.Series(s).dropna().nunique()),
                    ),
                    n_unique_object_font_names=(
                        "object_font_name",
                        lambda s: int(pd.Series(s).dropna().nunique()),
                    ),
                )
            )
        else:
            font_group = pd.DataFrame(
                columns=group_keys
                + [
                    "n_unique_font_objects_sha256",
                    "n_unique_font_object_types",
                    "n_unique_object_font_names",
                ]
            )
    else:
        obj_group_all = pd.DataFrame(
            columns=group_keys + ["n_objects", "n_unique_object_sha256"]
        )
        font_group = pd.DataFrame(
            columns=group_keys
            + [
                "n_unique_font_objects_sha256",
                "n_unique_font_object_types",
                "n_unique_object_font_names",
            ]
        )

    # ------------------------------------------------------------------
    # Signature stats per group (docs with 2+ signatures)
    # ------------------------------------------------------------------
    if not sig_df.empty:
        for col in ["pdf_group", "case_id", "cluster_id", "cluster_name", "filing_type"]:
            if col in sig_df.columns:
                sig_df[col] = sig_df[col].fillna("").astype(str)

        sig_df["n_signatures"] = pd.to_numeric(sig_df["n_signatures"], errors="coerce").fillna(0).astype(int)
        sig_2plus = sig_df[sig_df["n_signatures"] >= 2]

        if not sig_2plus.empty:
            sig_group = (
                sig_2plus.groupby(group_keys, as_index=False)
                .agg(
                    n_docs_2_or_more_pdf_sigs=("doc_sha256", "nunique"),
                )
            )
        else:
            sig_group = pd.DataFrame(columns=group_keys + ["n_docs_2_or_more_pdf_sigs"])
    else:
        sig_group = pd.DataFrame(columns=group_keys + ["n_docs_2_or_more_pdf_sigs"])

    # ------------------------------------------------------------------
    # Merge everything into one big overview table
    # ------------------------------------------------------------------
    overview = docs_group.merge(term_group, on=group_keys, how="left")
    overview = overview.merge(obj_group_all, on=group_keys, how="left")
    overview = overview.merge(font_group, on=group_keys, how="left")
    overview = overview.merge(sig_group, on=group_keys, how="left")

    # Fill numeric NaNs with 0
    for col in [
        "n_term_occurrences",
        "n_unique_terms",
        "n_unique_search_groups",
        "n_objects",
        "n_unique_object_sha256",
        "n_unique_font_objects_sha256",
        "n_unique_font_object_types",
        "n_unique_object_font_names",
        "n_docs_2_or_more_pdf_sigs",
    ]:
        if col in overview.columns:
            overview[col] = overview[col].fillna(0).astype(int)

    # filing_type_count = n_docs (per row)
    overview["filing_type_count"] = overview["n_docs"].astype(int)

    # Column order
    col_order = [
        "pdf_group",
        "case_id",
        "cluster_id",
        "cluster_name",
        "filing_type",
        "n_docs",
        "filing_type_count",
        "filing_date",
        "pdf_pages",
        "pdf_file_size",
        "n_docs_2_or_more_pdf_sigs",
        "n_term_occurrences",
        "n_unique_terms",
        "n_unique_search_groups",
        "n_objects",
        "n_unique_object_sha256",
        "n_unique_font_objects_sha256",
        "n_unique_font_object_types",
        "n_unique_object_font_names",
    ]

    for c in col_order:
        if c not in overview.columns:
            overview[c] = pd.NA

    overview = overview[col_order]

    # Sort for sanity
    overview = overview.sort_values(
        by=["pdf_group", "cluster_id", "case_id", "filing_type"],
        kind="mergesort",
    )

    overview.to_csv(overview_path, index=False)

    print(f"[ok] Parsed JSON files: {n_parsed} (failed: {n_failed})", file=sys.stderr)
    print(f"[ok] Wrote pdf_group_overview.csv to: {overview_path.resolve()}", file=sys.stderr)


if __name__ == "__main__":
    main()

