#!/usr/bin/env python3
"""Benchmark screenplay entity extraction across providers and modes."""

from __future__ import annotations

import argparse
import json
import re
from dataclasses import dataclass
from difflib import SequenceMatcher
from pathlib import Path
from typing import Dict, List, Sequence, Tuple

import extract_entities as extractor


ENTITY_SECTIONS: Sequence[Tuple[str, str]] = (
    ("characters", "character"),
    ("locations", "location"),
    ("props", "prop"),
)

PRIORITY_RANK = {"critical": 3, "high": 2, "medium": 1, "low": 0}
REFERENCE_RANK = {"required": 2, "helpful": 1, "not_needed": 0}
COMPACT_RE = re.compile(r"[\W_]+", re.UNICODE)
TOKEN_RE = re.compile(r"[\s/,:;()\-]+")


@dataclass
class RunConfig:
    key: str
    provider: str
    mode: str
    model: str
    fallback_model: str


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Benchmark screenplay entity extraction across a series.")
    parser.add_argument(
        "--input-glob",
        default="screenplay/srd part * blue revision.pdf",
        help="Glob for screenplay PDF files in episode order.",
    )
    parser.add_argument(
        "--outdir",
        default="screenplay/benchmark_results",
        help="Directory for benchmark outputs.",
    )
    parser.add_argument(
        "--configs",
        default="openai/fulltext,openai/chunked,gemini/fulltext,gemini/chunked",
        help="Comma-separated provider/mode configs.",
    )
    parser.add_argument(
        "--prompt-version",
        default=None,
        help="Prompt version to use. Defaults to manifest current_version.",
    )
    parser.add_argument(
        "--source-language",
        default="auto",
        choices=["auto", "ko", "ja", "en"],
        help="Force the screenplay language instead of auto detection.",
    )
    parser.add_argument(
        "--chunk-chars",
        type=int,
        default=14000,
        help="Approximate max characters per chunk in chunked mode.",
    )
    parser.add_argument(
        "--temperature",
        type=float,
        default=0.2,
        help="Sampling temperature for extraction.",
    )
    parser.add_argument(
        "--limit",
        type=int,
        default=None,
        help="Only process the first N episodes.",
    )
    parser.add_argument(
        "--skip-existing",
        action="store_true",
        help="Reuse existing episode outputs if they already exist on disk.",
    )
    return parser.parse_args()


def parse_configs(raw: str) -> List[RunConfig]:
    configs: List[RunConfig] = []
    for piece in raw.split(","):
        piece = piece.strip()
        if not piece:
            continue
        provider, mode = piece.split("/", 1)
        model, fallback = extractor.resolve_models(provider, None, None)
        configs.append(
            RunConfig(
                key=f"{provider}_{mode}",
                provider=provider,
                mode=mode,
                model=model,
                fallback_model=fallback,
            )
        )
    return configs


def compact_identifier(text: str) -> str:
    return COMPACT_RE.sub("", text.lower())


def token_set(text: str) -> set[str]:
    return {token for token in TOKEN_RE.split(text.lower()) if token}


def entity_identifiers(item: Dict[str, object]) -> set[str]:
    identifiers = {extractor.normalize_name(str(item["name"])), compact_identifier(str(item["name"]))}
    for alias in item.get("aliases", []):
        if isinstance(alias, str) and alias.strip():
            identifiers.add(extractor.normalize_name(alias))
            identifiers.add(compact_identifier(alias))
    return {identifier for identifier in identifiers if identifier}


def choose_canonical_name(current_name: str, candidate_name: str) -> str:
    current_compact = compact_identifier(current_name)
    candidate_compact = compact_identifier(candidate_name)
    if current_compact == candidate_compact:
        return candidate_name if len(candidate_name) > len(current_name) else current_name
    if current_compact in candidate_compact and len(candidate_compact) > len(current_compact):
        return candidate_name
    if candidate_compact in current_compact and len(current_compact) >= len(candidate_compact):
        return current_name
    return candidate_name if len(candidate_name) > len(current_name) else current_name


def similarity_score(entity: Dict[str, object], canon: Dict[str, object]) -> float:
    entity_ids = entity_identifiers(entity)
    canon_ids = set(canon["identifiers"])
    if entity_ids & canon_ids:
        return 1.0

    best = 0.0
    for left in entity_ids:
        for right in canon_ids:
            if not left or not right:
                continue
            ratio = SequenceMatcher(None, left, right).ratio()
            if len(left) > 3 and len(right) > 3 and (left in right or right in left):
                ratio = max(ratio, 0.94)
            best = max(best, ratio)

    entity_tokens = token_set(str(entity["name"]))
    canon_tokens = set(canon["name_tokens"])
    if entity_tokens and canon_tokens:
        union = entity_tokens | canon_tokens
        if union:
            best = max(best, len(entity_tokens & canon_tokens) / len(union))

    entity_traits = {trait.lower() for trait in entity.get("visual_anchor_traits", [])}
    canon_traits = {trait.lower() for trait in canon.get("visual_anchor_traits", [])}
    if entity_traits and canon_traits:
        overlap = len(entity_traits & canon_traits)
        if overlap >= 2:
            best = max(best, 0.88)

    return best


def find_canon_match(entity: Dict[str, object], canons: List[Dict[str, object]]) -> Tuple[int | None, float]:
    best_index = None
    best_score = 0.0
    for index, canon in enumerate(canons):
        score = similarity_score(entity, canon)
        if score > best_score:
            best_score = score
            best_index = index

    if best_score >= 0.93:
        return best_index, best_score
    return None, best_score


def merge_lists_preserve(items: Sequence[str], limit: int = 20) -> List[str]:
    merged = list(dict.fromkeys(item for item in items if item))
    return merged[:limit]


def create_canon(entity_type: str, entity: Dict[str, object], canon_id: str, episode_index: int, episode_label: str) -> Dict[str, object]:
    return {
        "canon_id": canon_id,
        "entity_type": entity_type,
        "name": entity["name"],
        "aliases": merge_lists_preserve(entity.get("aliases", []), limit=12),
        "identifiers": sorted(entity_identifiers(entity)),
        "name_tokens": sorted(token_set(str(entity["name"]))),
        "visual_anchor_traits": merge_lists_preserve(entity.get("visual_anchor_traits", []), limit=16),
        "variant_axes": merge_lists_preserve(entity.get("variant_axes", []), limit=12),
        "continuity_priority": entity.get("continuity_priority", "low"),
        "reference_image_priority": entity.get("reference_image_priority", "not_needed"),
        "episodes": [episode_index],
        "episode_labels": [episode_label],
        "appearance_count": 1,
    }


def merge_into_canon(canon: Dict[str, object], entity: Dict[str, object], episode_index: int, episode_label: str) -> None:
    canon["name"] = choose_canonical_name(canon["name"], entity["name"])
    canon["aliases"] = merge_lists_preserve([*canon["aliases"], *entity.get("aliases", [])], limit=12)
    canon["identifiers"] = sorted(entity_identifiers({"name": canon["name"], "aliases": canon["aliases"]}))
    canon["name_tokens"] = sorted(token_set(canon["name"]))
    canon["visual_anchor_traits"] = merge_lists_preserve(
        [*canon["visual_anchor_traits"], *entity.get("visual_anchor_traits", [])],
        limit=16,
    )
    canon["variant_axes"] = merge_lists_preserve([*canon["variant_axes"], *entity.get("variant_axes", [])], limit=12)
    if PRIORITY_RANK[entity.get("continuity_priority", "low")] > PRIORITY_RANK[canon["continuity_priority"]]:
        canon["continuity_priority"] = entity["continuity_priority"]
    if REFERENCE_RANK[entity.get("reference_image_priority", "not_needed")] > REFERENCE_RANK[canon["reference_image_priority"]]:
        canon["reference_image_priority"] = entity["reference_image_priority"]
    if episode_index not in canon["episodes"]:
        canon["episodes"].append(episode_index)
    if episode_label not in canon["episode_labels"]:
        canon["episode_labels"].append(episode_label)
    canon["appearance_count"] += 1


def relation_signature(relation: Dict[str, object]) -> str:
    participant_signature = sorted(
        f"{item.get('series_canon_id') or item['entity_name']}|{item['entity_type']}|{item['role']}"
        for item in relation["participants"]
    )
    return "||".join(
        [
            relation["relation_family"],
            relation["relation_type"],
            relation["directionality"],
            relation["temporal_scope"],
            *participant_signature,
        ]
    )


def merge_relation_memory(
    relation_memory_state: List[Dict[str, object]],
    relation: Dict[str, object],
    episode_index: int,
    episode_label: str,
) -> None:
    signature = relation_signature(relation)
    for current in relation_memory_state:
        if current["signature"] != signature:
            continue
        if PRIORITY_RANK[relation["continuity_priority"]] > PRIORITY_RANK[current["continuity_priority"]]:
            current["continuity_priority"] = relation["continuity_priority"]
        if episode_index not in current["episodes"]:
            current["episodes"].append(episode_index)
        if episode_label not in current["episode_labels"]:
            current["episode_labels"].append(episode_label)
        current["evidence"] = merge_lists_preserve([*current["evidence"], *relation["evidence"]], limit=12)
        return

    relation_memory_state.append(
        {
            "signature": signature,
            "relation_family": relation["relation_family"],
            "relation_type": relation["relation_type"],
            "directionality": relation["directionality"],
            "temporal_scope": relation["temporal_scope"],
            "continuity_priority": relation["continuity_priority"],
            "continuity_reason": relation["continuity_reason"],
            "participants": relation["participants"],
            "episodes": [episode_index],
            "episode_labels": [episode_label],
            "evidence": relation["evidence"][:12],
        }
    )


def build_series_memory(
    canon_state: Dict[str, List[Dict[str, object]]],
    relation_memory_state: Sequence[Dict[str, object]],
) -> Dict[str, object]:
    memory: Dict[str, object] = {}
    for section, entity_type in ENTITY_SECTIONS:
        entries = sorted(
            canon_state[entity_type],
            key=lambda item: (-len(item["episodes"]), item["name"]),
        )
        memory[section] = [
            {
                "canon_id": canon["canon_id"],
                "name": canon["name"],
                "aliases": canon["aliases"][:6],
                "episodes": canon["episodes"],
                "visual_anchor_traits": canon["visual_anchor_traits"][:8],
                "variant_axes": canon["variant_axes"][:8],
                "continuity_priority": canon["continuity_priority"],
            }
            for canon in entries[:120]
        ]
    relation_entries = sorted(
        relation_memory_state,
        key=lambda item: (-len(item["episodes"]), item["relation_type"]),
    )
    memory["relation_facts"] = [
        {
            "relation_family": item["relation_family"],
            "relation_type": item["relation_type"],
            "directionality": item["directionality"],
            "temporal_scope": item["temporal_scope"],
            "continuity_priority": item["continuity_priority"],
            "continuity_reason": item["continuity_reason"],
            "participants": [
                {
                    "entity_name": participant.get("series_canon_name") or participant["entity_name"],
                    "entity_type": participant["entity_type"],
                    "role": participant["role"],
                }
                for participant in item["participants"]
            ],
            "episodes": item["episodes"],
        }
        for item in relation_entries[:160]
    ]
    return memory


def entity_text_fields_valid(item: Dict[str, object], language_code: str) -> bool:
    fields = [str(item.get("description", "")), str(item.get("continuity_reason", ""))]
    text = " ".join(fields)
    if language_code == "ko":
        return any("\uac00" <= char <= "\ud7a3" for char in text)
    if language_code == "ja":
        return any(
            ("\u3040" <= char <= "\u30ff") or ("\u4e00" <= char <= "\u9fff")
            for char in text
        )
    return any(("a" <= char.lower() <= "z") for char in text)


def find_suspicious_pairs(canons: List[Dict[str, object]]) -> List[Dict[str, object]]:
    suspicious: List[Dict[str, object]] = []
    for left_index, left in enumerate(canons):
        for right in canons[left_index + 1 :]:
            score = similarity_score(
                {"name": left["name"], "aliases": left["aliases"], "visual_anchor_traits": left["visual_anchor_traits"]},
                {"name": right["name"], "aliases": right["aliases"], "identifiers": right["identifiers"], "name_tokens": right["name_tokens"], "visual_anchor_traits": right["visual_anchor_traits"]},
            )
            if score >= 0.90:
                suspicious.append(
                    {
                        "left": left["canon_id"],
                        "left_name": left["name"],
                        "right": right["canon_id"],
                        "right_name": right["name"],
                        "score": round(score, 3),
                    }
                )
    suspicious.sort(key=lambda item: item["score"], reverse=True)
    return suspicious[:12]


def score_config(summary: Dict[str, object]) -> float:
    later_total = summary["later_episode_entity_total"]
    reuse_ratio = summary["later_episode_relinked_total"] / later_total if later_total else 0.0
    return round(
        reuse_ratio * 100
        - summary["language_violations"] * 8
        - len(summary["suspicious_pairs"]) * 5
        - summary["episode_failures"] * 100,
        3,
    )


def write_json(path: Path, payload: Dict[str, object]) -> None:
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text(json.dumps(payload, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")


def load_json(path: Path) -> Dict[str, object]:
    return json.loads(path.read_text(encoding="utf-8"))


def run_config(args: argparse.Namespace, config: RunConfig, pdf_paths: Sequence[Path]) -> Dict[str, object]:
    config_dir = Path(args.outdir).expanduser().resolve() / config.key
    episodes_dir = config_dir / "episodes"
    canon_state: Dict[str, List[Dict[str, object]]] = {entity_type: [] for _, entity_type in ENTITY_SECTIONS}
    relation_memory_state: List[Dict[str, object]] = []
    counters = {entity_type: 0 for _, entity_type in ENTITY_SECTIONS}
    episode_summaries: List[Dict[str, object]] = []
    language_violations = 0

    for episode_index, pdf_path in enumerate(pdf_paths, start=1):
        episode_label = pdf_path.stem
        episode_output = episodes_dir / f"episode_{episode_index:02d}.json"
        memory_output = episodes_dir / f"episode_{episode_index:02d}_memory.json"
        memory_payload = build_series_memory(canon_state, relation_memory_state)
        write_json(memory_output, memory_payload)

        if args.skip_existing and episode_output.exists():
            result = load_json(episode_output)
        else:
            result = extractor.extract_entities_for_pdf(
                input_path=pdf_path,
                output_path=episode_output,
                sqlite_output_path=None,
                episode_key=episode_label,
                provider=config.provider,
                model=config.model,
                fallback_model=config.fallback_model,
                prompt_version=args.prompt_version,
                source_language_override=args.source_language,
                mode=config.mode,
                series_memory=memory_payload,
                chunk_chars=args.chunk_chars,
                temperature=args.temperature,
            )

        linked_episode = {
            "episode_index": episode_index,
            "episode_label": episode_label,
            "source_file": str(pdf_path),
            "summary": result["summary"],
            "extraction_metadata": result["extraction_metadata"],
            "characters": [],
            "locations": [],
            "props": [],
            "relation_facts": [],
        }
        relinked_total = 0
        total_entities = 0
        relation_total = 0

        language_code = result["extraction_metadata"]["source_language"]["code"]
        linked_entity_lookup: Dict[Tuple[str, str], Dict[str, object]] = {}
        for section, entity_type in ENTITY_SECTIONS:
            canons = canon_state[entity_type]
            for entity in result[section]:
                total_entities += 1
                if not entity_text_fields_valid(entity, language_code):
                    language_violations += 1

                match_index, similarity = find_canon_match(entity, canons)
                matched_existing = match_index is not None
                if matched_existing:
                    canon = canons[match_index]
                    merge_into_canon(canon, entity, episode_index, episode_label)
                    relinked_total += 1
                else:
                    counters[entity_type] += 1
                    canon_id = f"{entity_type}:{counters[entity_type]:03d}"
                    canon = create_canon(entity_type, entity, canon_id, episode_index, episode_label)
                    canons.append(canon)

                linked_entity = {
                    **entity,
                    "series_canon_id": canon["canon_id"],
                    "matched_existing_canon": matched_existing,
                    "match_similarity": round(similarity, 3),
                }
                linked_episode[section].append(linked_entity)
                for identifier in entity_identifiers(entity):
                    linked_entity_lookup[(entity_type, identifier)] = linked_entity

        for relation in result.get("relation_facts", []):
            participants = []
            for participant in relation["participants"]:
                lookup_key = (
                    participant["entity_type"],
                    extractor.normalize_name(str(participant["entity_name"])),
                )
                linked_entity = linked_entity_lookup.get(lookup_key)
                if not linked_entity:
                    continue
                participants.append(
                    {
                        "entity_type": participant["entity_type"],
                        "entity_name": participant["entity_name"],
                        "role": participant["role"],
                        "series_canon_id": linked_entity["series_canon_id"],
                        "series_canon_name": linked_entity["name"],
                    }
                )

            unique_canons = {item["series_canon_id"] for item in participants}
            if len(unique_canons) < 2:
                continue

            linked_relation = {
                **relation,
                "participants": participants,
            }
            linked_episode["relation_facts"].append(linked_relation)
            merge_relation_memory(relation_memory_state, linked_relation, episode_index, episode_label)
            relation_total += 1

        linked_output = episodes_dir / f"episode_{episode_index:02d}_linked.json"
        write_json(linked_output, linked_episode)
        episode_summaries.append(
            {
                "episode_index": episode_index,
                "episode_label": episode_label,
                "entity_total": total_entities,
                "relinked_total": relinked_total,
                "counts": {
                    section: len(result[section]) for section, _ in ENTITY_SECTIONS
                },
                "relation_fact_total": relation_total,
                "model_calls": len(result["extraction_metadata"]["chunk_models_used"]),
                "elapsed_seconds": result["extraction_metadata"]["elapsed_seconds"],
            }
        )

    series_canon = {
        section: sorted(canon_state[entity_type], key=lambda item: item["canon_id"])
        for section, entity_type in ENTITY_SECTIONS
    }
    write_json(config_dir / "series_canon.json", series_canon)

    suspicious_pairs = []
    for section, entity_type in ENTITY_SECTIONS:
        for item in find_suspicious_pairs(canon_state[entity_type]):
            suspicious_pairs.append({"section": section, **item})

    later_episode_entity_total = sum(item["entity_total"] for item in episode_summaries[1:])
    later_episode_relinked_total = sum(item["relinked_total"] for item in episode_summaries[1:])
    summary = {
        "config_key": config.key,
        "provider": config.provider,
        "mode": config.mode,
        "model": config.model,
        "fallback_model": config.fallback_model,
        "episodes_processed": len(episode_summaries),
        "episode_failures": 0,
        "language_violations": language_violations,
        "later_episode_entity_total": later_episode_entity_total,
        "later_episode_relinked_total": later_episode_relinked_total,
        "canon_counts": {
            section: len(series_canon[section]) for section, _ in ENTITY_SECTIONS
        },
        "relation_memory_count": len(relation_memory_state),
        "episode_summaries": episode_summaries,
        "suspicious_pairs": suspicious_pairs,
    }
    summary["score"] = score_config(summary)
    write_json(config_dir / "summary.json", summary)
    write_json(config_dir / "series_relation_memory.json", {"relation_facts": relation_memory_state})
    return summary


def build_report(summaries: Sequence[Dict[str, object]], pdf_paths: Sequence[Path], outdir: Path) -> None:
    sorted_summaries = sorted(summaries, key=lambda item: item["score"], reverse=True)
    lines = [
        "# Screenplay Entity Benchmark",
        "",
        "## Episodes",
        "",
    ]
    for path in pdf_paths:
        lines.append(f"- `{path.name}`")

    lines.extend(
        [
            "",
            "## Config Summary",
            "",
            "| Config | Provider | Mode | Model | Score | Reuse Ratio | Relation Memory | Language Violations | Suspicious Pairs | Canon Counts |",
            "| --- | --- | --- | --- | ---: | ---: | ---: | ---: | ---: | --- |",
        ]
    )

    for summary in sorted_summaries:
        later_total = summary["later_episode_entity_total"]
        reuse_ratio = summary["later_episode_relinked_total"] / later_total if later_total else 0.0
        canon_counts = summary["canon_counts"]
        lines.append(
            "| {config_key} | {provider} | {mode} | {model} | {score:.3f} | {reuse_ratio:.3f} | {relation_memory_count} | {language_violations} | {pair_count} | C:{c} / L:{l} / P:{p} |".format(
                config_key=summary["config_key"],
                provider=summary["provider"],
                mode=summary["mode"],
                model=summary["model"],
                score=summary["score"],
                reuse_ratio=reuse_ratio,
                relation_memory_count=summary["relation_memory_count"],
                language_violations=summary["language_violations"],
                pair_count=len(summary["suspicious_pairs"]),
                c=canon_counts["characters"],
                l=canon_counts["locations"],
                p=canon_counts["props"],
            )
        )

    lines.extend(["", "## Notes", ""])
    for summary in sorted_summaries:
        lines.append(f"### {summary['config_key']}")
        later_total = summary["later_episode_entity_total"]
        reuse_ratio = summary["later_episode_relinked_total"] / later_total if later_total else 0.0
        lines.append(f"- Score: `{summary['score']:.3f}`")
        lines.append(f"- Later-episode relink ratio: `{reuse_ratio:.3f}`")
        lines.append(f"- Relation memory count: `{summary['relation_memory_count']}`")
        lines.append(f"- Language violations: `{summary['language_violations']}`")
        if summary["suspicious_pairs"]:
            top = summary["suspicious_pairs"][:3]
            formatted = ", ".join(
                f"{item['section']}:{item['left_name']} <-> {item['right_name']} ({item['score']})"
                for item in top
            )
            lines.append(f"- Top suspicious unlinked pairs: {formatted}")
        else:
            lines.append("- Top suspicious unlinked pairs: none")
        lines.append("")

    report_path = outdir / "comparison_report.md"
    report_path.write_text("\n".join(lines).rstrip() + "\n", encoding="utf-8")


def main() -> int:
    args = parse_args()
    input_root = Path("/Users/jedi/Documents/TheRoad-Scene")
    pdf_paths = sorted(input_root.glob(args.input_glob))
    if args.limit:
        pdf_paths = pdf_paths[: args.limit]
    if not pdf_paths:
        raise SystemExit(f"No PDF files matched: {args.input_glob}")

    configs = parse_configs(args.configs)
    outdir = Path(args.outdir).expanduser().resolve()
    outdir.mkdir(parents=True, exist_ok=True)

    summaries = []
    for config in configs:
        print(f"=== Running {config.key} ===", flush=True)
        summaries.append(run_config(args, config, pdf_paths))

    build_report(summaries, pdf_paths, outdir)
    write_json(outdir / "comparison_summary.json", {"configs": summaries})
    print(f"Saved benchmark outputs to: {outdir}", flush=True)
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
