#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""G4.5a canary — primary framing same-character-ID detection (Rule J, Task 4.3).

Detect Rule J violations (RO-8 algorithm — STRICT 0): in close-framing shots,
when a body-part close-up trigger (`focus on` / `close on` / `tight on` /
`detail on` — `_ID_BODY_PART_TRIGGERS` cross-import from G4.3) co-occurs
with a full-body description keyword (`stands` / `seated` / `leaning` /
`전신`) and BOTH refer to the SAME character ID within ±20 token distance,
that's a primary_framing rule violation.

RO-8 algorithm (per spec §5.1 + Plan Task 4.3):
- per-shot scan, build dict `{character_id: [token_pos]}` for all C##/C##O##
  occurrences.
- For each body_part_close_up_trigger token position, find nearest
  C##/C##O## within ±20 tokens — `cid_trigger`.
- For each full_body_description_keyword position, find nearest character
  ID — `cid_body`.
- If `cid_trigger == cid_body` AND BOTH nearest-character-distances ≤ 20:
  increment violation count.

RO-5 binding: close-framing shot scope filter uses
`_SPATIAL_FRAMING_CLOSE_KEYWORDS` 9-entry production-aligned (NOT the
LLM-facing 11-entry list). `손이`/`눈이`/`얼굴이` are NOT in canary scope.

RO-11 binding: `_ID_BODY_PART_TRIGGERS` G4.3 carry — direct external
import via `_g4_5a_common.render_prompt_card_imports()`. NO alias.

PR4-B7 / RO-8: STRICT 0 — exit 1 on any violation, NOT soft warning.

Trap #1 silent-fallback ban / Trap #3 CP shape / Trap #4 pinned-tuple
post-loop validation / Trap #8 module-level constants.

Override O-17 — regex/substring detector is 1차 only. baseline 비0 시 P1
follow-up: LLM-validator (gpt-5.4-mini judge) 전환.
`feedback_no_regex_postprocessing.md` carry.

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

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

sys.path.insert(0, str(Path(__file__).resolve().parent))
from _g4_5a_common import (  # noqa: E402
    _load_scenes_or_fail,
    _validate_scene_preflight,
    build_pinning_scene_set,
    emit_output,
    load_cp_manifest,
    load_pinning_from_args,
    render_prompt_card_imports,
    resolve_cp_root,
    validate_pinning,
)

_CONSTS = render_prompt_card_imports()
# RO-5 binding: 9-entry production-aligned close-framing keywords (NOT
# LLM-facing 11-entry list).
_SPATIAL_FRAMING_CLOSE_KEYWORDS = _CONSTS["_SPATIAL_FRAMING_CLOSE_KEYWORDS"]
# RO-11 binding: G4.3 `_ID_BODY_PART_TRIGGERS` direct external import — NO alias.
# Test V7 verifies this exact import string. Spec §5.4 Trap #8: also re-import
# `compute_render_strategy_snapshot_hash` so module-level dependency is
# auditable via grep (helper is invoked transitively by preflight helper).
from app.core.steps.render_prompt_card import (  # noqa: E402,F401
    _ID_BODY_PART_TRIGGERS,
    compute_render_strategy_snapshot_hash,
)

_ID_REGEX = re.compile(r"\bC\d{2}(?:O\d{2})?\b")

# RO-8 binding: ±20 tokens nearest-character distance.
_WINDOW_TOKENS = 20

# G4.5a ground-truth: full-body description keywords (R1+R2 audit alignment
# carry; Plan Task 4.3 spec — 4 keyword anchor `stands / seated / leaning /
# 전신`). G4.4 view-mixing 7 verb list 와 다름 — G4.5a Rule J 는 4 anchor.
_FULL_BODY_KEYWORDS: tuple[str, ...] = (
    "stands",
    "seated",
    "leaning",
    "전신",
)


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


def _tokenize(text: str) -> list[str]:
    """Whitespace tokenizer for RO-8 algorithm.

    NOTE: simple whitespace split is intentional — punctuation and quotes
    are treated as part of the adjacent token. The ±20 token window is
    forgiving enough that punctuation-attachment does not change RO-8
    detection results materially. canary smoke fixture covers this.
    """
    return text.split()


def _build_cid_positions(tokens: list[str]) -> dict[str, list[int]]:
    """Build dict {character_id: [token_positions]} per RO-8 algorithm."""
    cid_positions: dict[str, list[int]] = {}
    for pos, tok in enumerate(tokens):
        m = _ID_REGEX.search(tok)
        if m is not None:
            cid = m.group()
            # Explicit init pattern — Trap #1 PR-fix-iter-1-1 strict (NO
            # `setdefault(cid, [])` fallback). `setdefault` is technically a
            # default-injecting form; we use explicit `in` check for clarity.
            if cid in cid_positions:
                cid_positions[cid].append(pos)
            else:
                cid_positions[cid] = [pos]
    return cid_positions


def _nearest_cid(
    target_pos: int,
    cid_positions: dict[str, list[int]],
) -> tuple[str | None, int]:
    """Return (cid, distance) for nearest cid within ±20 tokens.

    If no cid is within window, returns (None, sentinel). Caller checks
    distance ≤ _WINDOW_TOKENS.
    """
    best_cid: str | None = None
    best_dist = _WINDOW_TOKENS + 1
    for cid, positions in cid_positions.items():
        for cp in positions:
            d = abs(cp - target_pos)
            if d <= _WINDOW_TOKENS and d < best_dist:
                best_cid = cid
                best_dist = d
    return best_cid, best_dist


def _detect_same_character_id_violations(prompt: str) -> int:
    """RO-8 same-character-ID detection algorithm (PR4-B7 STRICT 0).

    Build per-shot {cid: [pos]} dict. For each body-part close-up trigger
    token position, find nearest cid (`cid_trigger`). For each full-body
    keyword position, find nearest cid (`cid_body`). If `cid_trigger ==
    cid_body` AND both distances ≤ 20: violation.

    Returns count of violations in this prompt.
    """
    tokens = _tokenize(prompt)
    cid_positions = _build_cid_positions(tokens)

    if not cid_positions:
        return 0

    violations = 0
    seen_pairs: set[tuple[int, int]] = set()
    # Find body-part close-up trigger positions (multi-word triggers — match
    # consecutive tokens lower-cased).
    trigger_positions: list[int] = []
    n = len(tokens)
    for pos in range(n):
        for trigger in _ID_BODY_PART_TRIGGERS:
            tlen = len(trigger.split())
            if pos + tlen > n:
                continue
            window_text = " ".join(tokens[pos:pos + tlen]).lower()
            if window_text == trigger.lower():
                trigger_positions.append(pos)
                break  # don't double-count same position with multiple triggers

    # Find full-body keyword positions (single-token substring match).
    full_body_positions: list[int] = []
    for pos, tok in enumerate(tokens):
        tok_lower = tok.lower()
        if any(kw.lower() in tok_lower for kw in _FULL_BODY_KEYWORDS):
            full_body_positions.append(pos)

    if not trigger_positions or not full_body_positions:
        return 0

    for tpos in trigger_positions:
        cid_trigger, dist_trigger = _nearest_cid(tpos, cid_positions)
        if cid_trigger is None:
            continue
        if dist_trigger > _WINDOW_TOKENS:
            continue
        for fpos in full_body_positions:
            cid_body, dist_body = _nearest_cid(fpos, cid_positions)
            if cid_body is None:
                continue
            if dist_body > _WINDOW_TOKENS:
                continue
            if cid_trigger == cid_body:
                pair_key = (tpos, fpos)
                if pair_key in seen_pairs:
                    continue
                seen_pairs.add(pair_key)
                violations += 1
    return violations


def main() -> int:
    parser = argparse.ArgumentParser(
        description=(
            "G4.5a canary: primary framing same-character-ID detection "
            "(close-framing shots only, RO-8 algorithm ±20 tokens, "
            "Rule J, STRICT 0)."
        )
    )
    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)
    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(
        "--render-strategy-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 = resolve_cp_root(args)
    cp = load_cp_manifest(cp_root)

    measurement_failures: list[str] = []
    scenes = _load_scenes_or_fail(cp, measurement_failures)
    if scenes is None:
        out = _empty_output(args, pinning, measurement_failures)
        emit_output(out, args)
        return 1

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

    total_shots_with_close_framing = 0
    primary_framing_violation_count = 0
    primary_framing_violation_per_shot: dict[str, int] = {}
    scene_set_out: list[dict] = []

    # RO-5 binding: 9-entry production-aligned scope filter.
    close_kw_lower = tuple(kw.lower() for kw in _SPATIAL_FRAMING_CLOSE_KEYWORDS)

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

        variations = _validate_scene_preflight(
            scene, si_int, shi_int, measurement_failures
        )
        if variations is None:
            continue

        # Determine close-framing shot scope filter — at least one variation
        # contains close-framing keyword (9-entry scope).
        shot_has_close_framing = False
        shot_violation_count = 0
        for vi, var in enumerate(variations):
            if not isinstance(var, dict):
                measurement_failures.append(
                    f"s{si_int}_sh{shi_int}_v{vi}: variation not a dict"
                )
                continue
            if "t2i_prompt" not in var:
                measurement_failures.append(
                    f"s{si_int}_sh{shi_int}_v{vi}: t2i_prompt key missing"
                )
                continue
            prompt = var["t2i_prompt"]
            if not isinstance(prompt, str) or not prompt.strip():
                measurement_failures.append(
                    f"s{si_int}_sh{shi_int}_v{vi}: t2i_prompt missing, "
                    f"empty, or not str"
                )
                continue
            prompt_lower = prompt.lower()
            has_close = any(kw in prompt_lower for kw in close_kw_lower)
            if not has_close:
                continue  # close-framing shot scope filter — out of scope
            shot_has_close_framing = True
            # RO-8 algorithm
            shot_violation_count += _detect_same_character_id_violations(
                prompt
            )

        scene_set_out.append({
            "scene_index": si_int,
            "shot_index": shi_int,
            "has_close_framing": shot_has_close_framing,
        })
        if shot_has_close_framing:
            total_shots_with_close_framing += 1
        if shot_violation_count > 0:
            primary_framing_violation_count += shot_violation_count
            shot_key = f"{si_int}_{shi_int}"
            # Trap #1 / PR-fix-iter-1-1 — explicit init.
            if shot_key in primary_framing_violation_per_shot:
                primary_framing_violation_per_shot[shot_key] += (
                    shot_violation_count
                )
            else:
                primary_framing_violation_per_shot[shot_key] = (
                    shot_violation_count
                )

    # Trap #4 — pinned-tuple post-loop validation.
    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_with_close_framing": total_shots_with_close_framing,
            "primary_framing_violation_count": (
                primary_framing_violation_count
            ),
            "primary_framing_violation_per_shot": (
                primary_framing_violation_per_shot
            ),
            "trigger_phrases_used": list(_ID_BODY_PART_TRIGGERS),
            "full_body_keywords_used": list(_FULL_BODY_KEYWORDS),
            "close_framing_keywords_used": list(_SPATIAL_FRAMING_CLOSE_KEYWORDS),
        },
        "measurement_failures": measurement_failures,
        "exit_criteria_threshold": 0,
        "exit_criteria_pass": (
            not measurement_failures
            and primary_framing_violation_count == 0
        ),
        "measurement_scripts": {
            "primary_framing": (
                "scripts/canary/g4_5a_primary_framing.py "
                "(close-framing shots only, RO-5 9-entry, RO-8 ±20 tokens, "
                "RO-11 G4.3 _ID_BODY_PART_TRIGGERS carry, Rule J, "
                "Override O-17 regex 1차 only)"
            )
        },
    }

    emit_output(out, args)

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

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


def _empty_output(
    args: argparse.Namespace,
    pinning: dict,
    measurement_failures: list,
) -> dict:
    return {
        "timestamp": iso8601_now(),
        "prompt_version": args.prompt_version,
        "role": args.role,
        "pinning": pinning,
        "scene_set": [],
        "metrics": {
            "total_shots_with_close_framing": 0,
            "primary_framing_violation_count": 0,
            "primary_framing_violation_per_shot": {},
            "trigger_phrases_used": list(_ID_BODY_PART_TRIGGERS),
            "full_body_keywords_used": list(_FULL_BODY_KEYWORDS),
            "close_framing_keywords_used": list(_SPATIAL_FRAMING_CLOSE_KEYWORDS),
        },
        "measurement_failures": measurement_failures,
        "exit_criteria_threshold": 0,
        "exit_criteria_pass": False,
        "measurement_scripts": {
            "primary_framing": (
                "scripts/canary/g4_5a_primary_framing.py "
                "(close-framing shots only, Rule J)"
            )
        },
    }


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