#!/usr/bin/env python3
import csv, json, re, sys
from pathlib import Path
from collections import defaultdict, Counter
from typing import Dict, Tuple, List, Any, Set

SUFFIX_RE = re.compile(
    r"(?:,\s*)?(?:jr\.?|sr\.?|ii|iii|iv|v|vi|vii|viii|ix|x)\b\.?",
    re.IGNORECASE
)

WS_RE = re.compile(r"\s+")

def normalize_whitespace(s: str) -> str:
    return WS_RE.sub(" ", (s or "").strip())

def tokenize_for_match(raw: str) -> List[str]:
    s = normalize_whitespace(raw)
    s_no_suffix = SUFFIX_RE.sub("", s)
    s_no_suffix = normalize_whitespace(s_no_suffix)
    if not s_no_suffix:
        return []
    return s_no_suffix.lower().split()

def variant_form(raw: str) -> str:
    return normalize_whitespace(raw or "")

def read_csv(path: Path):
    with path.open("r", encoding="utf-8-sig", newline="") as f:
        r = csv.DictReader(f)
        headers = r.fieldnames or []
        rows = [dict(x) for x in r]
    return headers, rows

def write_csv(path: Path, headers: List[str], rows: List[Dict[str, Any]]):
    with path.open("w", encoding="utf-8", newline="") as f:
        w = csv.DictWriter(f, fieldnames=headers, extrasaction="ignore")
        w.writeheader()
        for row in rows:
            w.writerow(row)

class DSU:
    def __init__(self, n: int):
        self.p = list(range(n))
        self.r = [0]*n
    def find(self, x: int) -> int:
        while self.p[x] != x:
            self.p[x] = self.p[self.p[x]]
            x = self.p[x]
        return x
    def union(self, a: int, b: int):
        ra, rb = self.find(a), self.find(b)
        if ra == rb: return
        if self.r[ra] < self.r[rb]:
            self.p[ra] = rb
        elif self.r[rb] < self.r[ra]:
            self.p[rb] = ra
        else:
            self.p[rb] = ra
            self.r[ra] += 1

def main(argv: List[str]):
    import argparse
    ap = argparse.ArgumentParser(description="Cluster defendants by first-name + shared-token; count variants incl. suffix/case.")
    ap.add_argument("input_csv")
    ap.add_argument("output_csv")
    ap.add_argument("--json", dest="json_out")
    ap.add_argument("--defendant-col", default="Defendant")
    ap.add_argument("--case-col", default="Case Number")
    ap.add_argument("--keep-blank", action="store_true")
    args = ap.parse_args(argv[1:])

    in_path = Path(args.input_csv); out_path = Path(args.output_csv)
    headers, rows = read_csv(in_path)
    if args.defendant_col not in headers:
        sys.exit(f'Missing defendant column: "{args.defendant_col}"')
    if args.case_col not in headers:
        sys.exit(f'Missing case column: "{args.case_col}"')

    row_infos: List[Dict[str, Any]] = []
    for i, r in enumerate(rows):
        raw_def = (r.get(args.defendant_col, "") or "").strip()
        case_no = (r.get(args.case_col, "") or "").strip()
        tokens = tokenize_for_match(raw_def)
        if not tokens:
            if not args.keep_blank:
                continue
            first_tok = ""
            other_set: Set[str] = set()
        else:
            first_tok = tokens[0]
            other_set = set(tokens[1:])
        row_infos.append({
            "idx": i,
            "case": case_no,
            "first": first_tok,
            "others": other_set,
            "variant": variant_form(raw_def),
        })

    by_first: Dict[str, List[Dict[str, Any]]] = defaultdict(list)
    for info in row_infos:
        by_first[info["first"]].append(info)

    components: List[List[Dict[str, Any]]] = []
    for first, bucket in by_first.items():
        n = len(bucket)
        if n == 0:
            continue
        if n == 1:
            components.append([bucket[0]])
            continue
        dsu = DSU(n)
        index_by_other: Dict[str, List[int]] = defaultdict(list)
        for idx_local, info in enumerate(bucket):
            for tok in info["others"]:
                index_by_other[tok].append(idx_local)
        for tok, idx_list in index_by_other.items():
            if len(idx_list) > 1:
                base = idx_list[0]
                for j in idx_list[1:]:
                    dsu.union(base, j)
        comp_map: Dict[int, List[Dict[str, Any]]] = defaultdict(list)
        for idx_local, info in enumerate(bucket):
            root = dsu.find(idx_local)
            comp_map[root].append(info)
        components.extend(comp_map.values())

    clusters: List[Dict[str, Any]] = []
    for comp in components:
        if not comp:
            continue
        variants = Counter()
        cases: Set[str] = set()
        indices: List[int] = []
        first_tok = comp[0]["first"]
        all_others: Set[str] = set()
        for info in comp:
            if info["variant"]:
                variants[info["variant"]] += 1
            if info["case"]:
                cases.add(info["case"])
            indices.append(info["idx"])
            all_others.update(info["others"])
        clusters.append({
            "first": first_tok,
            "others_set": all_others,
            "variants": variants,
            "case_numbers": cases,
            "row_indices": indices,
        })

    def cluster_key_str(c: Dict[str, Any]) -> str:
        first = c["first"]
        others = " ".join(sorted(c["others_set"])) if c["others_set"] else ""
        return f"first:{first} | tokens:{others}".strip()

    ranking = sorted(
        clusters,
        key=lambda c: (-len(c["case_numbers"]), cluster_key_str(c))
    )
    cluster_id_by_ref = {id(c): rank for rank, c in enumerate(ranking, start=1)}

    out_headers = list(headers)
    for col in ["ClusterID", "ClusterKey", "ClusterSize", "VariantCount", "CaseCount"]:
        if col not in out_headers:
            out_headers.append(col)

    for c in ranking:
        cid = cluster_id_by_ref[id(c)]
        cluster_size = len(c["row_indices"])
        case_count = len(c["case_numbers"])
        variant_cnt = len(c["variants"])
        key_str = cluster_key_str(c)
        for idx in c["row_indices"]:
            rows[idx]["ClusterID"]    = str(cid)
            rows[idx]["ClusterKey"]   = key_str
            rows[idx]["ClusterSize"]  = str(cluster_size)
            rows[idx]["VariantCount"] = str(variant_cnt)
            rows[idx]["CaseCount"]    = str(case_count)

    write_csv(out_path, out_headers, rows)

    if args.json_out:
        clusters_list = []
        for c in ranking:
            cid = cluster_id_by_ref[id(c)]
            clusters_list.append({
                "cluster_id": cid,
                "cluster_key": cluster_key_str(c),
                "case_count": len(c["case_numbers"]),
                "row_count": len(c["row_indices"]),
                "variant_count": len(c["variants"]),
                "variants": [{"name": n, "count": cnt} for n, cnt in sorted(c["variants"].items(), key=lambda x: (-x[1], x[0]))],
                "case_numbers": sorted(c["case_numbers"])
            })
        payload = {
            "source_csv": in_path.name,
            "defendant_col": args.defendant_col,
            "case_col": args.case_col,
            "cluster_count": len(clusters_list),
            "clusters": clusters_list
        }
        Path(args.json_out).write_text(json.dumps(payload, indent=2, ensure_ascii=False), encoding="utf-8")
        print(f"Wrote clusters JSON: {args.json_out}")

    print(f"Wrote augmented CSV: {out_path}  (clusters: {len(ranking)})")

if __name__ == "__main__":
    main(sys.argv)
