#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""G4.3 canary — STRICT body-part focus pattern counter (Task 4.1).

Detects `<trigger>` + composite ID + `'s` + body part. The 4 trigger phrases
(`focus on`, `close on`, `tight on`, `detail on`) are imported from
`render_prompt_card._ID_BODY_PART_TRIGGERS` (Override R1R2-B2 — drift ban).

Scope: ALL shots of pinned (si, shi) tuples — `focus on` pattern itself is
the intentional close-up trigger, so this is not a close-framing-only metric.

Exit criteria (analyzer, OUT-OF-BAND):
  candidate.metrics.body_part_focus_pattern_count == 0  (STRICT, R1R2-B2 / R3-I1)
  baseline value is informational only.

Usage (config mode):
    python scripts/canary/g4_3_body_part_focus.py \\
        --config canary_config.json \\
        --cp-root data/projects/<pid>/checkpoints/scene_detail \\
        --prompt-version 17.<timestamp> \\
        --role candidate \\
        --output results/body_part_focus_candidate.json
"""
from __future__ import annotations

import argparse
import json
import re
import sys
from datetime import datetime
from pathlib import Path

# Make scripts/canary importable AND backend/ on sys.path so the production
# render_prompt_card module resolves.
sys.path.insert(0, str(Path(__file__).resolve().parent))
from _g4_3_common import (  # noqa: E402
    _load_scenes_or_fail,
    _validate_scene_preflight,
    build_pinning_scene_set,
    is_close_framing,
    load_cp_manifest,
    load_pinning_from_args,
    render_prompt_card_imports,
    validate_pinning,
)

_CONSTS = render_prompt_card_imports()
_ID_BODY_PART_TRIGGERS = _CONSTS["_ID_BODY_PART_TRIGGERS"]
_ID_REPRODUCTION_SURFACES = _CONSTS["_ID_REPRODUCTION_SURFACES"]

# Override R1R2-B2 — 4-trigger alternation (`re.escape` for safety) followed
# by composite ID + `'s` + body-part word. Captures "focus on C03's eye",
# "close on C12O04's hand", etc. Case-insensitive.
_TRIGGER_ALT = "|".join(re.escape(t) for t in _ID_BODY_PART_TRIGGERS)
BODY_PART_FOCUS_PATTERN = re.compile(
    r"\b(?:" + _TRIGGER_ALT + r")\s+C\d{2}(?:O\d{2})?'s\s+\w+",
    re.IGNORECASE,
)

# Used only for scene_set diagnostics (`has_composite_id` / `has_reproduction_surface`).
_COMPOSITE_ID_RE = re.compile(r"\bC\d{2}(?:O\d{2})?\b")


def iso8601_now() -> str:
    return datetime.utcnow().isoformat() + "Z"


def main() -> int:
    parser = argparse.ArgumentParser(
        description="G4.3 canary: body-part focus pattern counter (4 triggers, all shots)."
    )
    parser.add_argument("--config", type=str, default=None)
    parser.add_argument("--pid", type=str, default=None)
    parser.add_argument("--cp-root", type=str, default=None,
                        help="Directory containing scene_detail manifest.json")
    parser.add_argument("--prompt-version", type=str, required=True)
    parser.add_argument("--role", type=str, default="candidate",
                        choices=["baseline", "candidate"])
    parser.add_argument("--scene-index-list", type=str, default=None)
    parser.add_argument("--shot-index-list-per-scene", type=str, default=None)
    parser.add_argument("--model-routing", type=str, default=None)
    parser.add_argument("--prompt-source-mode", type=str, default=None,
                        choices=[None, "file", "db"])
    parser.add_argument("--card-commit-hash", type=str, default=None)
    parser.add_argument("--id-policy-card-snapshot-hash", type=str, default=None)
    parser.add_argument("--output", type=str, required=True)
    args = parser.parse_args()

    pinning = load_pinning_from_args(args)
    validate_pinning(pinning)

    cp_root_str: str | None = None
    if args.config:
        with Path(args.config).open("r", encoding="utf-8") as fp:
            cfg = json.load(fp)
        cp_root_str = cfg.get("cp_root")
    if not cp_root_str:
        cp_root_str = args.cp_root
    if not cp_root_str:
        print("[g4_3_body_part_focus] ERROR: --cp-root (or config.cp_root) required",
              file=sys.stderr)
        return 1
    cp_root = Path(cp_root_str)
    if not cp_root.exists():
        print(f"[g4_3_body_part_focus] ERROR: --cp-root not found: {cp_root}",
              file=sys.stderr)
        return 1

    cp = load_cp_manifest(cp_root)
    measurement_failures: list[str] = []

    scenes = _load_scenes_or_fail(cp, measurement_failures)
    if scenes is None:
        out = {
            "timestamp": iso8601_now(),
            "prompt_version": args.prompt_version,
            "role": args.role,
            "pinning": pinning,
            "scene_set": [],
            "metrics": {
                "total_shots": 0,
                "body_part_focus_pattern_count": 0,
                "body_part_focus_per_shot": {},
                "trigger_phrases_used": list(_ID_BODY_PART_TRIGGERS),
            },
            "measurement_failures": measurement_failures,
            "measurement_scripts": {
                "body_part_focus": "scripts/canary/g4_3_body_part_focus.py (all shots — 4 trigger alternation R1R2-B2)"
            },
        }
        _emit(out, args)
        return 1

    pinning_set = build_pinning_scene_set(pinning)
    seen_in_cp: set[tuple[int, int]] = set()

    total_shots = 0
    body_part_focus_count = 0
    body_part_focus_per_shot: dict[str, int] = {}
    scene_set_out: list[dict] = []

    for scene in scenes:
        if not isinstance(scene, dict):
            measurement_failures.append("scene entry is not dict")
            continue
        si = scene.get("scene_index")
        shi = scene.get("_shot_index")
        if si is None or shi is None:
            continue
        if (int(si), int(shi)) not in pinning_set:
            continue
        seen_in_cp.add((int(si), int(shi)))

        variations = _validate_scene_preflight(scene, si, shi, measurement_failures)
        if variations is None:
            continue

        cf_flag, cf_err = is_close_framing(scene)
        if cf_err is not None:
            measurement_failures.append(f"s{si}_sh{shi}: {cf_err}")

        # Diagnostic flags (advisory — not gating).
        has_composite_id = False
        has_reproduction_surface = False
        for var in variations:
            if not isinstance(var, dict):
                continue
            p = var.get("t2i_prompt")
            if isinstance(p, str):
                if _COMPOSITE_ID_RE.search(p):
                    has_composite_id = True
                pl = p.lower()
                for surf in _ID_REPRODUCTION_SURFACES:
                    if surf.lower() in pl:
                        has_reproduction_surface = True
                        break

        scene_set_out.append({
            "scene_index": si,
            "shot_index": shi,
            "is_close_framing": cf_flag,
            "has_composite_id": has_composite_id,
            "has_reproduction_surface": has_reproduction_surface,
        })

        total_shots += 1
        shot_count = 0
        for vi, var in enumerate(variations):
            if not isinstance(var, dict):
                measurement_failures.append(
                    f"s{si}_sh{shi}_v{vi}: variation not a dict"
                )
                continue
            prompt = var.get("t2i_prompt")
            if not isinstance(prompt, str) or not prompt.strip():
                measurement_failures.append(
                    f"s{si}_sh{shi}_v{vi}: t2i_prompt missing, empty, or not str"
                )
                continue
            matches = BODY_PART_FOCUS_PATTERN.findall(prompt)
            shot_count += len(matches)
        body_part_focus_count += shot_count
        body_part_focus_per_shot[f"{si}_{shi}"] = shot_count

    # Pinned-tuple post-loop validation (R2-B6 carry — false STRICT pass 차단).
    for (si, shi) in pinning_set:
        if (si, shi) not in seen_in_cp:
            measurement_failures.append(
                f"pinning (s{si}_sh{shi}): not present in CP data.scenes"
            )

    out = {
        "timestamp": iso8601_now(),
        "prompt_version": args.prompt_version,
        "role": args.role,
        "pinning": pinning,
        "scene_set": scene_set_out,
        "metrics": {
            "total_shots": total_shots,
            "body_part_focus_pattern_count": body_part_focus_count,
            "body_part_focus_per_shot": body_part_focus_per_shot,
            "trigger_phrases_used": list(_ID_BODY_PART_TRIGGERS),
        },
        "measurement_failures": measurement_failures,
        "measurement_scripts": {
            "body_part_focus": "scripts/canary/g4_3_body_part_focus.py (all shots — 4 trigger alternation R1R2-B2)"
        },
    }

    _emit(out, args)

    print(
        f"[g4_3_body_part_focus] role={args.role} prompt_version={args.prompt_version}\n"
        f"  total_shots={total_shots}  body_part_focus_pattern_count={body_part_focus_count}\n"
        f"  measurement_failures: {len(measurement_failures)}\n"
        f"Exit criteria (STRICT): candidate.metrics.body_part_focus_pattern_count == 0",
        file=sys.stderr,
    )

    if measurement_failures:
        return 1
    if args.role == "candidate" and body_part_focus_count > 0:
        return 1
    return 0


def _emit(out: dict, args: argparse.Namespace) -> None:
    if args.output == "-":
        json.dump(out, sys.stdout, ensure_ascii=False, indent=2, sort_keys=True)
        sys.stdout.write("\n")
    else:
        out_path = Path(args.output)
        out_path.parent.mkdir(parents=True, exist_ok=True)
        with out_path.open("w", encoding="utf-8") as fp:
            json.dump(out, fp, ensure_ascii=False, indent=2, sort_keys=True)


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