#!/usr/bin/env python3
"""Run fulltext-only scene still extraction across a screenplay series."""

from __future__ import annotations

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

import extract_entities as entity_common
import extract_scene_stills as scene_extractor


LATIN_TOKEN_RE = re.compile(r"[A-Za-z]{3,}")
NON_LATIN_RE = re.compile(r"[가-힣ぁ-ゟ゠-ヿ一-鿿]")


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


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Run scene still extraction across episodes.")
    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/scene_still_results",
        help="Directory for provider outputs and comparison reports.",
    )
    parser.add_argument(
        "--providers",
        default="openai,gemini",
        help="Comma-separated provider list.",
    )
    parser.add_argument(
        "--prompt-version",
        default=None,
        help="Prompt version to use. Defaults to the scene prompt 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(
        "--temperature",
        type=float,
        default=0.15,
        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 provider outputs and memory snapshots when present.",
    )
    return parser.parse_args()


def parse_providers(raw: str) -> List[RunConfig]:
    configs: List[RunConfig] = []
    for provider in [piece.strip() for piece in raw.split(",") if piece.strip()]:
        model, fallback = entity_common.resolve_models(provider, None, None)
        configs.append(
            RunConfig(
                key=f"{provider}_fulltext",
                provider=provider,
                model=model,
                fallback_model=fallback,
            )
        )
    return configs


def episode_sort_key(path: Path) -> tuple[int, str]:
    match = re.search(r"part\s+(\d+)", path.stem, re.IGNORECASE)
    if match:
        return int(match.group(1)), path.name
    return 9999, path.name


def find_episode_files(pattern: str, limit: int | None) -> List[Path]:
    files = sorted(Path().glob(pattern), key=episode_sort_key)
    if limit is not None:
        files = files[:limit]
    return [path.resolve() for path in files]


def english_token_hits(text: str) -> int:
    tokens = [token for token in LATIN_TOKEN_RE.findall(text) if token.lower() not in {"vip", "mm"}]
    return len(tokens)


def count_language_violations(payload: Dict[str, object]) -> int:
    source_code = payload["source_language"]["code"]
    if source_code == "en":
        return 0

    violations = 0
    for still in payload["scene_stills"]:
        texts = [
            still["beat_title"],
            still["still_frame_prompt"],
            *still["camera"].values(),
            *still["lighting"].values(),
        ]
        for text in texts:
            text = str(text)
            if english_token_hits(text) >= 2 and not NON_LATIN_RE.search(text):
                violations += 1
    return violations


def count_unresolved_visible_entities(payload: Dict[str, object]) -> int:
    unresolved = 0
    for still in payload["scene_stills"]:
        for visible in still["visible_entities"]:
            if not str(visible.get("entity_id", "")).strip():
                unresolved += 1
    return unresolved


def count_blank_blocks(payload: Dict[str, object], field: str) -> int:
    blank_count = 0
    for still in payload["scene_stills"]:
        if any(not str(value).strip() for value in still[field].values()):
            blank_count += 1
    return blank_count


def summarize_episode(payload: Dict[str, object], episode_label: str) -> Dict[str, object]:
    stills = payload["scene_stills"]
    heading_count = payload["extraction_metrics"]["detected_heading_count"]
    covered_headings = payload["extraction_metrics"]["covered_heading_catalog_count"]
    avg_visible = round(
        sum(len(item["visible_entities"]) for item in stills) / len(stills),
        2,
    ) if stills else 0.0
    return {
        "episode": episode_label,
        "source_file": payload["source_file"],
        "scene_still_count": len(stills),
        "detected_heading_count": heading_count,
        "covered_heading_catalog_count": covered_headings,
        "heading_coverage_ratio": round(covered_headings / heading_count, 4) if heading_count else 0.0,
        "characters": len(payload["entity_index"]["characters"]),
        "locations": len(payload["entity_index"]["locations"]),
        "props": len(payload["entity_index"]["props"]),
        "avg_visible_entities_per_still": avg_visible,
        "language_violation_count": count_language_violations(payload),
        "unresolved_visible_entity_count": count_unresolved_visible_entities(payload),
        "camera_blank_count": count_blank_blocks(payload, "camera"),
        "lighting_blank_count": count_blank_blocks(payload, "lighting"),
        "matched_existing_entity_count": payload["extraction_metrics"]["id_assignment"]["matched_existing"],
        "created_new_entity_count": payload["extraction_metrics"]["id_assignment"]["created_new"],
        "backfilled_entity_count": payload["extraction_metrics"]["id_assignment"]["backfilled_from_scene_refs"],
        "elapsed_seconds": payload["extraction_metadata"]["elapsed_seconds"],
        "model_used": payload["extraction_metadata"]["model_used"],
        "prompt_version": payload["extraction_metadata"]["prompt_version"],
    }


def aggregate_provider_summary(config: RunConfig, episode_summaries: Sequence[Dict[str, object]]) -> Dict[str, object]:
    total_headings = sum(item["detected_heading_count"] for item in episode_summaries)
    total_covered = sum(item["covered_heading_catalog_count"] for item in episode_summaries)
    total_stills = sum(item["scene_still_count"] for item in episode_summaries)
    total_elapsed = round(sum(item["elapsed_seconds"] for item in episode_summaries), 2)
    return {
        "config": config.key,
        "provider": config.provider,
        "model": config.model,
        "episodes": len(episode_summaries),
        "total_scene_stills": total_stills,
        "total_detected_headings": total_headings,
        "total_covered_headings": total_covered,
        "overall_heading_coverage_ratio": round(total_covered / total_headings, 4) if total_headings else 0.0,
        "avg_stills_per_episode": round(total_stills / len(episode_summaries), 2) if episode_summaries else 0.0,
        "avg_stills_per_heading": round(total_stills / total_headings, 3) if total_headings else 0.0,
        "total_language_violations": sum(item["language_violation_count"] for item in episode_summaries),
        "total_unresolved_visible_entities": sum(item["unresolved_visible_entity_count"] for item in episode_summaries),
        "total_camera_blank": sum(item["camera_blank_count"] for item in episode_summaries),
        "total_lighting_blank": sum(item["lighting_blank_count"] for item in episode_summaries),
        "total_matched_existing_entities": sum(item["matched_existing_entity_count"] for item in episode_summaries),
        "total_created_new_entities": sum(item["created_new_entity_count"] for item in episode_summaries),
        "total_backfilled_entities": sum(item["backfilled_entity_count"] for item in episode_summaries),
        "total_elapsed_seconds": total_elapsed,
        "episode_summaries": list(episode_summaries),
    }


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 build_comparison_summary(provider_summaries: Sequence[Dict[str, object]]) -> Dict[str, object]:
    ordered = sorted(
        provider_summaries,
        key=lambda item: (
            item["overall_heading_coverage_ratio"],
            item["total_scene_stills"],
            -item["total_language_violations"],
            -item["total_unresolved_visible_entities"],
        ),
        reverse=True,
    )
    return {
        "providers": provider_summaries,
        "best_overall": ordered[0]["config"] if ordered else None,
        "best_scene_count": max(provider_summaries, key=lambda item: item["total_scene_stills"])["config"] if provider_summaries else None,
        "best_fewest_issues": min(
            provider_summaries,
            key=lambda item: (
                item["total_language_violations"] + item["total_unresolved_visible_entities"] + item["total_camera_blank"] + item["total_lighting_blank"],
                -item["overall_heading_coverage_ratio"],
            ),
        )["config"] if provider_summaries else None,
    }


def build_report(provider_summaries: Sequence[Dict[str, object]]) -> str:
    lines = [
        "# Scene Still Extraction Comparison",
        "",
        "| Config | Model | Episodes | Total stills | Heading coverage | Lang flags | Unresolved IDs | Camera blanks | Lighting blanks | Matched prior entities | Backfilled entities | Elapsed (s) |",
        "| --- | --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: |",
    ]

    for item in provider_summaries:
        lines.append(
            f"| {item['config']} | {item['model']} | {item['episodes']} | {item['total_scene_stills']} | "
            f"{item['overall_heading_coverage_ratio']:.2%} | {item['total_language_violations']} | "
            f"{item['total_unresolved_visible_entities']} | {item['total_camera_blank']} | {item['total_lighting_blank']} | "
            f"{item['total_matched_existing_entities']} | {item['total_backfilled_entities']} | {item['total_elapsed_seconds']} |"
        )

    lines.extend(["", "## Episode Breakdown", ""])
    for item in provider_summaries:
        lines.append(f"### {item['config']}")
        lines.append("")
        lines.append("| Episode | Stills | Headings | Covered | Coverage | Matched prior | Backfilled | Elapsed (s) |")
        lines.append("| --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: |")
        for episode in item["episode_summaries"]:
            lines.append(
                f"| {episode['episode']} | {episode['scene_still_count']} | {episode['detected_heading_count']} | "
                f"{episode['covered_heading_catalog_count']} | {episode['heading_coverage_ratio']:.2%} | "
                f"{episode['matched_existing_entity_count']} | {episode['backfilled_entity_count']} | {episode['elapsed_seconds']} |"
            )
        lines.append("")
    return "\n".join(lines) + "\n"


def run_config(
    *,
    config: RunConfig,
    files: Sequence[Path],
    outdir: Path,
    prompt_version: str | None,
    source_language: str,
    temperature: float,
    skip_existing: bool,
) -> Dict[str, object]:
    provider_dir = outdir / config.key
    episodes_dir = provider_dir / "episodes"
    series_memory: Dict[str, object] = {}
    episode_summaries: List[Dict[str, object]] = []

    print(f"=== {config.key} ===", flush=True)
    for index, pdf_path in enumerate(files, start=1):
        episode_label = f"episode_{index:02d}"
        output_path = episodes_dir / f"{episode_label}.json"
        sqlite_path = episodes_dir / f"{episode_label}.sqlite"
        memory_path = episodes_dir / f"{episode_label}_memory.json"

        if skip_existing and output_path.exists() and memory_path.exists():
            payload = json.loads(output_path.read_text(encoding="utf-8"))
            series_memory = json.loads(memory_path.read_text(encoding="utf-8"))
            print(f"[reuse] {config.key} {episode_label}", flush=True)
        else:
            payload, series_memory = scene_extractor.extract_scene_stills_for_pdf(
                input_path=pdf_path,
                output_path=output_path,
                sqlite_output_path=sqlite_path,
                episode_key=episode_label,
                provider=config.provider,
                model=config.model,
                fallback_model=config.fallback_model,
                prompt_version=prompt_version,
                source_language_override=source_language,
                series_memory=series_memory,
                temperature=temperature,
            )
            write_json(memory_path, series_memory)

        episode_summaries.append(summarize_episode(payload, episode_label))

    summary = aggregate_provider_summary(config, episode_summaries)
    write_json(provider_dir / "summary.json", summary)
    write_json(provider_dir / "series_memory.json", series_memory)
    return summary


def main() -> int:
    args = parse_args()
    files = find_episode_files(args.input_glob, args.limit)
    if not files:
        raise SystemExit(f"No files matched: {args.input_glob}")

    outdir = Path(args.outdir).expanduser().resolve()
    provider_summaries: List[Dict[str, object]] = []
    for config in parse_providers(args.providers):
        provider_summaries.append(
            run_config(
                config=config,
                files=files,
                outdir=outdir,
                prompt_version=args.prompt_version,
                source_language=args.source_language,
                temperature=args.temperature,
                skip_existing=args.skip_existing,
            )
        )

    comparison_summary = build_comparison_summary(provider_summaries)
    write_json(outdir / "comparison_summary.json", comparison_summary)
    (outdir / "comparison_report.md").write_text(build_report(provider_summaries), encoding="utf-8")
    print(f"Saved comparison report to: {outdir / 'comparison_report.md'}", flush=True)
    return 0


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