#!/usr/bin/env python3
import argparse, os, json, collections
from pathlib import Path
import pandas as pd

def iter_json(dir_path: Path):
    for p in sorted(dir_path.glob("*.json")):
        try:
            with p.open("r", encoding="utf-8") as f:
                yield p.name, json.load(f)
        except Exception:
            # Skip unreadable/bad JSON files
            continue

def extract_flag_labels(doc):
    """
    Return (labels_list, labels_set) from doc["flags"].
    Accepts either [{'label': '...'}, ...] or ['...','...'] defensively.
    """
    flags = doc.get("flags")
    if not flags or not isinstance(flags, list):
        return [], set()

    labels = []
    for item in flags:
        if isinstance(item, dict):
            lab = item.get("label")
        else:
            lab = str(item)
        if lab and isinstance(lab, str):
            lab = lab.strip()
            if lab:
                labels.append(lab)
    return labels, set(labels)

def main():
    ap = argparse.ArgumentParser(description="Analyze flags across document JSON files.")
    ap.add_argument("--input", required=True, help="Directory containing document JSON files")
    ap.add_argument("--outdir", required=True, help="Directory to write CSV reports")
    args = ap.parse_args()

    input_dir = Path(args.input)
    outdir = Path(args.outdir)
    outdir.mkdir(parents=True, exist_ok=True)

    # Counters
    doc_counter = collections.Counter()   # flag -> number of docs containing this flag (unique per doc)
    occ_counter = collections.Counter()   # flag -> total occurrences across all docs (counts duplicates within a doc)
    all_docs = []                         # track per-doc summary rows
    # For co-occurrence (per document unique sets)
    # We'll fill a set of all flags observed first, then populate a matrix
    all_flag_labels = set()
    per_doc_flagsets = {}  # filename -> set(labels)

    docs_scanned = 0
    for fname, doc in iter_json(input_dir):
        docs_scanned += 1
        labels_list, labels_set = extract_flag_labels(doc)
        # per-document presence
        for lab in labels_set:
            doc_counter[lab] += 1
            all_flag_labels.add(lab)
        # total occurrences
        for lab in labels_list:
            occ_counter[lab] += 1

        per_doc_flagsets[fname] = labels_set
        all_docs.append({
            "json_filename": fname,
            "flags_count_unique": len(labels_set),
            "flags_list": "; ".join(sorted(labels_set)) if labels_set else ""
        })

    # --- Summary dataframe ---
    flags_sorted = sorted(all_flag_labels)
    rows = []
    total_docs = docs_scanned if docs_scanned else 1
    for lab in flags_sorted:
        docs_with = doc_counter.get(lab, 0)
        occs = occ_counter.get(lab, 0)
        share = round((docs_with / total_docs) * 100.0, 2)
        avg_occ_per_doc_with_flag = round((occs / docs_with), 3) if docs_with else 0.0
        rows.append({
            "flag_label": lab,
            "docs_with_flag": docs_with,
            "total_docs": total_docs,
            "share_of_all_docs_pct": share,
            "total_flag_occurrences": occs,
            "avg_occurrences_per_doc_with_flag": avg_occ_per_doc_with_flag
        })
    df_summary = pd.DataFrame(rows).sort_values(
        ["docs_with_flag", "total_flag_occurrences", "flag_label"],
        ascending=[False, False, True]
    )

    # --- Per-document dataframe ---
    df_docs = pd.DataFrame(all_docs).sort_values(
        ["flags_count_unique", "json_filename"],
        ascending=[False, True]
    )

    # --- Co-occurrence matrix ---
    # Matrix[i,j] = number of docs that contain both flag_i and flag_j
    labels_list = sorted(all_flag_labels)
    co_data = {lab: collections.Counter() for lab in labels_list}
    for flagset in per_doc_flagsets.values():
        if not flagset:
            continue
        f = sorted(flagset)
        for i, a in enumerate(f):
            co_data[a][a] += 1  # diagonal: docs where 'a' appears
            for b in f[i+1:]:
                co_data[a][b] += 1
                co_data[b][a] += 1
    # Build DataFrame
    co_rows = []
    for a in labels_list:
        row = {"flag_label": a}
        for b in labels_list:
            row[b] = co_data[a].get(b, 0)
        co_rows.append(row)
    df_co = pd.DataFrame(co_rows)
    df_co = df_co[["flag_label"] + labels_list]  # order columns

    # --- Write CSVs ---
    df_summary.to_csv(outdir / "flags_summary.csv", index=False)
    df_docs.to_csv(outdir / "flags_docs.csv", index=False)
    df_co.to_csv(outdir / "flags_cooccurrence.csv", index=False)

    print("[OK] Wrote:")
    for name in ["flags_summary.csv", "flags_docs.csv", "flags_cooccurrence.csv"]:
        print("  -", outdir / name)
    print(f"[Stats] docs_scanned={docs_scanned} unique_flags={len(all_flag_labels)}")

if __name__ == "__main__":
    main()
