#!/usr/bin/env python3
"""
MCRO PDF Group Term/Overview Report (with ALL-group)
====================================================

This script is a focused, trimmed-down version of mcro_pdf_group_report.py
that ONLY produces:

  1) pdf_group_term_stats.csv
  2) pdf_group_overview.csv

Key differences vs the original full report:

  * We introduce a synthetic pdf_group called "all" which includes EVERY
    JSON document in the directory, regardless of tracking.tracked_flag.
    - "all" rows exist in BOTH outputs.
    - This gives you a clean control/comparison group over the entire corpus.

  * Existing pdf_group rows (C1, TTF-C, etc.) are preserved, driven by:
      tracking.tracked_flag == true  AND  tracking.groups[*].pdf_group

  * No other CSVs are written (no documents/objects/attorneys/signatures dumps).

ASSUMED JSON STRUCTURE (same as original)
-----------------------------------------

Top-level:
  doc_sha256 or sha256
  filename
  filing_type
  filing_date
  pdf_page_count

  case : {
      "case_id":       ...,
      "cluster_id":    ...,
      "cluster_name":  ...,
      ... (other keys ignored here)
  }

  metadata   (ignored for this script)
  language : {
      "terms": [
          {
            "search_group": ...,
            "search_term":  ...,
            "quantity":     int,
          }, ...
      ],
      ... (instances ignored)
  }

  signatures : {
      "pdfsig": { ... }
      ... (other keys ignored)
  }

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

LOGIC SUMMARY
=============

For EVERY JSON doc:
  - We add an "all" row in the term pipeline and overview pipeline.

For docs with tracking.tracked_flag == true and at least one group:
  - We add one row per (pdf_group, doc) to the overview pipeline.
  - We add one row per (pdf_group, doc, search_group, search_term) to
    the term pipeline.
  - We add one row per (pdf_group, object) and (pdf_group, pdfsig) to
    support object/pdfsig stats.

Thus, the output tables contain:

  * Rows for each real pdf_group (C1, TTF-C, etc.), as before.
  * Additional rows for pdf_group "all" covering the entire corpus.

OUTPUTS
=======

1) pdf_group_term_stats.csv
   One row per (pdf_group, search_group, search_term) with:

     pdf_group
     search_group
     search_term
     total_quantity
     doc_count
     case_count
     cluster_count
     first_seen_date
     last_seen_date

   "all" acts as a control group across 3,601 docs.

2) pdf_group_overview.csv
   One row per pdf_group (including "all") with:

     pdf_group
     n_docs
     n_cases
     n_clusters
     case_ids
     cluster_ids
     cluster_names
     min_filing_date
     max_filing_date
     min_pdf_page_count
     max_pdf_page_count
     avg_pdf_page_count
     n_unique_filing_types
     filing_types

     n_term_occurrences
     n_unique_terms
     n_unique_search_groups

     n_objects
     n_unique_object_sha256
     n_unique_object_font_names
     n_docs_with_pdfsig

     docs_per_case_avg
     docs_per_cluster_avg

USAGE
=====

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

The script is robust to missing keys and will leave those cells blank or 0.
"""

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 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 safe_get(d: Dict[str, Any], key: str, default=None):
    if isinstance(d, dict):
        return d.get(key, default)
    return default


def main():
    ap = argparse.ArgumentParser(description="MCRO PDF group term/overview report with 'all' group.")
    ap.add_argument(
        "--input",
        required=True,
        help="Directory containing JSON files.",
    )
    ap.add_argument(
        "--outdir",
        required=True,
        help="Output directory for the two 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 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)

    # Data collectors for stats
    docs_rows: List[Dict[str, Any]] = []     # For overview (per pdf_group, doc)
    objects_rows: List[Dict[str, Any]] = []  # For object stats per pdf_group
    sig_rows: List[Dict[str, Any]] = []      # For pdfsig stats per pdf_group
    terms_rows: List[Dict[str, Any]] = []    # For term stats per pdf_group

    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", "")
        filing_date = data.get("filing_date", "")
        pdf_page_count = data.get("pdf_page_count", None)

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

        # Tracking & pdf_groups (for non-"all" 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", {})

        # --------------------------------------------------------
        # 1) "all" group rows (EVERY document, independent of tracking)
        # --------------------------------------------------------
        # Overview doc row for "all"
        docs_rows.append(
            {
                "pdf_group": "all",
                "doc_sha256": doc_sha256,
                "filename": filename,
                "filing_type": filing_type,
                "filing_date": filing_date,
                "pdf_page_count": pdf_page_count,
                "case_case_id": case_id,
                "case_cluster_id": cluster_id,
                "case_cluster_name": cluster_name,
            }
        )

        # Term rows for "all"
        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": "all",
                    "doc_sha256": doc_sha256,
                    "filename": filename,
                    "case_id": case_id,
                    "cluster_id": cluster_id,
                    "cluster_name": cluster_name,
                    "filing_date": filing_date,
                    "search_group": search_group,
                    "search_term": search_term,
                    "quantity": quantity,
                }
            )

        # Object rows for "all"
        for obj in objects:
            if not isinstance(obj, dict):
                continue
            o_row = {
                "pdf_group": "all",
                "doc_sha256": doc_sha256,
                "filename": filename,
                "case_id": case_id,
            }
            for k, v in obj.items():
                o_row[k] = v
            objects_rows.append(o_row)

        # pdfsig rows for "all"
        if isinstance(pdfsig, dict):
            s_row = {
                "pdf_group": "all",
                "doc_sha256": doc_sha256,
                "filename": filename,
                "case_id": case_id,
            }
            for k, v in pdfsig.items():
                s_row[f"pdfsig_{k}"] = v
            sig_rows.append(s_row)

        # --------------------------------------------------------
        # 2) Per-real-pdf_group rows (ONLY for tracked docs with groups)
        # --------------------------------------------------------
        if is_tracked and groups:
            for group in groups:
                if not isinstance(group, dict):
                    continue
                raw_pg = group.get("pdf_group")
                if raw_pg is None or str(raw_pg).strip() == "":
                    continue

                pdf_group = str(raw_pg).strip()

                # Overview doc row for this pdf_group
                docs_rows.append(
                    {
                        "pdf_group": pdf_group,
                        "doc_sha256": doc_sha256,
                        "filename": filename,
                        "filing_type": filing_type,
                        "filing_date": filing_date,
                        "pdf_page_count": pdf_page_count,
                        "case_case_id": case_id,
                        "case_cluster_id": cluster_id,
                        "case_cluster_name": cluster_name,
                    }
                )

                # Term rows for this pdf_group
                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": pdf_group,
                            "doc_sha256": doc_sha256,
                            "filename": filename,
                            "case_id": case_id,
                            "cluster_id": cluster_id,
                            "cluster_name": cluster_name,
                            "filing_date": filing_date,
                            "search_group": search_group,
                            "search_term": search_term,
                            "quantity": quantity,
                        }
                    )

                # Object rows for this pdf_group
                for obj in objects:
                    if not isinstance(obj, dict):
                        continue
                    o_row = {
                        "pdf_group": pdf_group,
                        "doc_sha256": doc_sha256,
                        "filename": filename,
                        "case_id": case_id,
                    }
                    for k, v in obj.items():
                        o_row[k] = v
                    objects_rows.append(o_row)

                # pdfsig rows for this pdf_group
                if isinstance(pdfsig, dict):
                    s_row = {
                        "pdf_group": pdf_group,
                        "doc_sha256": doc_sha256,
                        "filename": filename,
                        "case_id": case_id,
                    }
                    for k, v in pdfsig.items():
                        s_row[f"pdfsig_{k}"] = v
                    sig_rows.append(s_row)

        n_parsed += 1

    # ------------------------------------------------------------------
    # Build DataFrames
    # ------------------------------------------------------------------
    docs_df = pd.DataFrame(docs_rows)
    objects_df = pd.DataFrame(objects_rows)
    terms_df = pd.DataFrame(terms_rows)
    sig_df = pd.DataFrame(sig_rows)

    # ------------------------------------------------------------------
    # 1) pdf_group_term_stats.csv
    # ------------------------------------------------------------------
    term_stats_path = outdir / "pdf_group_term_stats.csv"

    if not terms_df.empty:
        # Normalize quantity to int
        terms_df["quantity"] = pd.to_numeric(terms_df["quantity"], errors="coerce").fillna(0).astype(int)

        term_stats = (
            terms_df.groupby(["pdf_group", "search_group", "search_term"], as_index=False)
            .agg(
                total_quantity=("quantity", "sum"),
                doc_count=("doc_sha256", "nunique"),
                case_count=("case_id", "nunique"),
                cluster_count=("cluster_id", "nunique"),
                first_seen_date=(
                    "filing_date",
                    lambda s: s.dropna().min() if len(s.dropna()) else pd.NA,
                ),
                last_seen_date=(
                    "filing_date",
                    lambda s: s.dropna().max() if len(s.dropna()) else pd.NA,
                ),
            )
        )

        term_stats.to_csv(term_stats_path, index=False)
    else:
        # No terms at all; write empty file
        term_stats_path.write_text("", encoding="utf-8")

    # ------------------------------------------------------------------
    # 2) pdf_group_overview.csv
    # ------------------------------------------------------------------
    overview_path = outdir / "pdf_group_overview.csv"

    if not docs_df.empty:
        # Unique (pdf_group, doc) combos
        docs_overview = docs_df.drop_duplicates(subset=["pdf_group", "doc_sha256"]).copy()

        # ---- Term summary per pdf_group ----
        if not terms_df.empty:
            term_group = (
                terms_df.groupby("pdf_group", 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=["pdf_group", "n_term_occurrences", "n_unique_terms", "n_unique_search_groups"]
            )

        # ---- Object summary per pdf_group ----
        if not objects_df.empty:
            if "object_sha256" in objects_df.columns:
                obj_group = objects_df.groupby("pdf_group", as_index=False).agg(
                    n_objects=("object_sha256", "count"),
                    n_unique_object_sha256=(
                        "object_sha256",
                        lambda s: int(pd.Series(s).dropna().nunique()),
                    ),
                    n_unique_object_font_names=(
                        "object_font_name",
                        lambda s: int(pd.Series(s).dropna().nunique()),
                    )
                    if "object_font_name" in objects_df.columns
                    else ("doc_sha256", lambda s: 0),
                )
            else:
                # Fallback if object_sha256 missing (shouldn't happen)
                obj_group = objects_df.groupby("pdf_group", as_index=False).agg(
                    n_objects=("doc_sha256", "count"),
                    n_unique_object_sha256=("doc_sha256", lambda s: int(pd.Series(s).dropna().nunique())),
                    n_unique_object_font_names=("doc_sha256", lambda s: 0),
                )
        else:
            obj_group = pd.DataFrame(
                columns=["pdf_group", "n_objects", "n_unique_object_sha256", "n_unique_object_font_names"]
            )

        # ---- pdfsig summary per pdf_group ----
        if not sig_df.empty:
            sig_group = (
                sig_df.groupby("pdf_group", as_index=False)
                .agg(n_docs_with_pdfsig=("doc_sha256", "nunique"))
            )
        else:
            sig_group = pd.DataFrame(columns=["pdf_group", "n_docs_with_pdfsig"])

        # ---- Core overview (per pdf_group) ----
        overview = (
            docs_overview.groupby("pdf_group", as_index=False)
            .agg(
                n_docs=("doc_sha256", "nunique"),
                n_cases=(
                    "case_case_id",
                    lambda s: int(pd.Series(s).dropna().nunique()),
                ),
                n_clusters=(
                    "case_cluster_id",
                    lambda s: int(pd.Series(s).dropna().nunique()),
                ),
                case_ids=("case_case_id", join_unique),
                cluster_ids=("case_cluster_id", join_unique),
                cluster_names=("case_cluster_name", join_unique),
                min_filing_date=(
                    "filing_date",
                    lambda s: s.dropna().min() if len(s.dropna()) else pd.NA,
                ),
                max_filing_date=(
                    "filing_date",
                    lambda s: s.dropna().max() if len(s.dropna()) else pd.NA,
                ),
                min_pdf_page_count=(
                    "pdf_page_count",
                    lambda s: pd.to_numeric(s, errors="coerce").min(),
                ),
                max_pdf_page_count=(
                    "pdf_page_count",
                    lambda s: pd.to_numeric(s, errors="coerce").max(),
                ),
                avg_pdf_page_count=(
                    "pdf_page_count",
                    lambda s: float(pd.to_numeric(s, errors="coerce").mean() or 0.0),
                ),
                n_unique_filing_types=(
                    "filing_type",
                    lambda s: int(pd.Series(s).dropna().nunique()),
                ),
                filing_types=("filing_type", join_unique),
            )
        )

        # Merge term/object/sig summaries
        overview = overview.merge(term_group, on="pdf_group", how="left")
        overview = overview.merge(obj_group, on="pdf_group", how="left")
        overview = overview.merge(sig_group, on="pdf_group", how="left")

        # Fill NaNs with sensible defaults for counts
        for col in [
            "n_term_occurrences",
            "n_unique_terms",
            "n_unique_search_groups",
            "n_objects",
            "n_unique_object_sha256",
            "n_unique_object_font_names",
            "n_docs_with_pdfsig",
        ]:
            if col in overview.columns:
                overview[col] = overview[col].fillna(0).astype(int)

        # Extra helpful metrics for probability / density:
        # docs_per_case_avg and docs_per_cluster_avg
        def safe_ratio(num, den):
            try:
                num = float(num)
                den = float(den)
                if den <= 0:
                    return 0.0
                return num / den
            except Exception:
                return 0.0

        overview["docs_per_case_avg"] = overview.apply(
            lambda r: safe_ratio(r["n_docs"], r["n_cases"]), axis=1
        )
        overview["docs_per_cluster_avg"] = overview.apply(
            lambda r: safe_ratio(r["n_docs"], r["n_clusters"]), axis=1
        )

        # Sort by pdf_group for stable output
        overview = overview.sort_values(by=["pdf_group"], kind="mergesort")

        overview.to_csv(overview_path, index=False)
    else:
        overview_path.write_text("", encoding="utf-8")

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


if __name__ == "__main__":
    main()
