#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""G4.4 canary — STRICT view-mixing detection (Task 4.4).

Detects same-character-ID full-body / partial-body verbs (`stands` /
`seated` / `leaning` / `walking` / `running` / `kneeling` / `lying` —
G4.4 ground-truth) **+** body-part close-up co-occurrence (`focus on` /
`close on` / `tight on` / `detail on` + `<id>'s <body-part>` —
`_ID_BODY_PART_TRIGGERS` cross-import from G4.3) within a single
t2i_prompt = view-mixing 시점 혼합 violation.

Override O-7 — paired-string substring carry: same composite-ID body-part
focus rule cross-references `_CONTINUITY_ID_POLICY_CROSS_REF_LITERAL`
(`id_policy.body_part_focus_rule`). Cross-card paired drift is enforced
by unit + integration tests; this canary inherits the same import boundary.

Override O-17 — regex brittleness: this regex/substring detector is a
1차 detector only. False-positive examples: "C01O02 stands in profile, hand
resting on the table" (full-body + body-part are natural partial
description, not close-up). If baseline >0, escalate to LLM-validator
(gpt-5.4-mini judge) as P1 follow-up.
`feedback_no_regex_postprocessing.md` carry.

Scope: ALL pinned (si, shi) shots — view-mixing pattern detection itself
is intent-agnostic (R2-B4 carry).

Exit criteria (analyzer, OUT-OF-BAND):
  candidate.metrics.view_mixing_count == 0  (STRICT, Override O-1)

Usage (config mode):
    python scripts/canary/g4_4_view_mixing.py \\
        --config canary_config.json \\
        --cp-root data/projects/<pid>/checkpoints/scene_detail \\
        --prompt-version 19.<timestamp> \\
        --role candidate \\
        --output results/view_mixing_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_4_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()
# G4.3 cross-import — 4 trigger alternation (Trap #8 module-level constant).
_ID_BODY_PART_TRIGGERS = _CONSTS["_ID_BODY_PART_TRIGGERS"]
_CONTINUITY_ID_POLICY_CROSS_REF_LITERAL = _CONSTS[
    "_CONTINUITY_ID_POLICY_CROSS_REF_LITERAL"
]

# G4.4 ground-truth: full-body / partial-body verbs (R1+R2 audit ground-truth)
_VIEW_MIXING_FULL_BODY_VERBS: tuple[str, ...] = (
    "stands",
    "seated",
    "leaning",
    "walking",
    "running",
    "kneeling",
    "lying",
)

_TRIGGER_ALT = "|".join(re.escape(t) for t in _ID_BODY_PART_TRIGGERS)
# Override O-17 carry: regex is 1차 detector. If baseline ≠ 0, escalate to LLM-validator (gpt-5.4-mini judge) per P1 follow-up.
_BODY_PART_FOCUS_PATTERN = re.compile(
    r"\b(?:" + _TRIGGER_ALT + r")\s+C\d{2}(?:O\d{2})?'s\s+\w+",
    re.IGNORECASE,
)
_ID_REGEX = re.compile(r"\bC\d{2}(?:O\d{2})?\b")
_WINDOW_CHARS = 50


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


def detect_view_mixing(prompt: str) -> list[dict]:
    """Return list of view-mixing violations.

    Each violation = {id_token, full_body_position, body_part_position}.
    Same composite-ID + full-body verb (within ±50 chars window) + body-part
    trigger close-up of same id elsewhere (outside that window) →
    view-mixing violation.

    `feedback_no_regex_postprocessing.md` carry — char-level proximity
    (NOT sentence regex). Override O-17 — 1차 detector only.
    """
    violations: list[dict] = []
    prompt_lower = prompt.lower()

    # ID positions per id_token
    id_positions: dict[str, list[tuple[int, int]]] = {}
    for match in _ID_REGEX.finditer(prompt):
        id_token = match.group()
        id_positions.setdefault(id_token, []).append(
            (match.start(), match.end())
        )

    # Body-part focus matches with id token extraction
    bp_focus_matches: list[dict] = []
    for match in _BODY_PART_FOCUS_PATTERN.finditer(prompt):
        id_match = _ID_REGEX.search(match.group())
        if id_match:
            bp_focus_matches.append({
                "id_token": id_match.group(),
                "start": match.start(),
                "end": match.end(),
            })

    # For each id_token full-body verb appearance window, check for SAME id
    # body-part close-up elsewhere (outside the same ±50 window).
    for id_token, positions in id_positions.items():
        for id_start, id_end in positions:
            win_start = max(0, id_start - _WINDOW_CHARS)
            win_end = min(len(prompt), id_end + _WINDOW_CHARS)
            window = prompt_lower[win_start:win_end]
            has_full_body = any(
                v in window for v in _VIEW_MIXING_FULL_BODY_VERBS
            )
            if not has_full_body:
                continue
            for bp in bp_focus_matches:
                if bp["id_token"] != id_token:
                    continue
                if (
                    bp["start"] >= id_start - _WINDOW_CHARS
                    and bp["end"] <= id_end + _WINDOW_CHARS
                ):
                    # Same window — likely same camera position; not mixing.
                    continue
                violations.append({
                    "id_token": id_token,
                    "full_body_position": [id_start, id_end],
                    "body_part_position": [bp["start"], bp["end"]],
                })
    return violations


def main() -> int:
    parser = argparse.ArgumentParser(
        description=(
            "G4.4 canary: view-mixing detection (all shots, cross-import "
            "G4.3 _ID_BODY_PART_TRIGGERS — Override O-17 regex 1차 only)."
        )
    )
    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(
        "--continuity-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 = 0
    view_mixing_count = 0
    view_mixing_per_shot: dict[str, int] = {}
    violations_detail: list[dict] = []
    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
        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))

        # _validate_scene_preflight() — strict per R3 patches (Trap #1, #3).
        variations = _validate_scene_preflight(
            scene, si_int, shi_int, measurement_failures
        )
        if variations is None:
            continue

        scene_set_out.append({
            "scene_index": si_int,
            "shot_index": shi_int,
        })

        total_shots += 1
        shot_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
            prompt = var.get("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
            shot_violations = detect_view_mixing(prompt)
            shot_count += len(shot_violations)
            for v in shot_violations:
                violations_detail.append({
                    "scene_index": si_int,
                    "shot_index": shi_int,
                    "variation_index": vi,
                    "id_token": v["id_token"],
                    "full_body_position": v["full_body_position"],
                    "body_part_position": v["body_part_position"],
                })

        view_mixing_count += shot_count
        view_mixing_per_shot[f"{si_int}_{shi_int}"] = shot_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": total_shots,
            "view_mixing_count": view_mixing_count,
            "view_mixing_per_shot": view_mixing_per_shot,
            "violations_detail": violations_detail,
            "trigger_phrases_used": list(_ID_BODY_PART_TRIGGERS),
            "full_body_verbs_used": list(_VIEW_MIXING_FULL_BODY_VERBS),
        },
        "measurement_failures": measurement_failures,
        "exit_criteria_threshold": 0,
        "exit_criteria_pass": (
            not measurement_failures and view_mixing_count == 0
        ),
        "measurement_scripts": {
            "view_mixing": (
                "scripts/canary/g4_4_view_mixing.py "
                "(all shots — cross-import G4.3 _ID_BODY_PART_TRIGGERS, "
                "Override O-17 regex 1차 only — paired ref literal: "
                f"{_CONTINUITY_ID_POLICY_CROSS_REF_LITERAL})"
            )
        },
    }

    emit_output(out, args)

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

    if measurement_failures:
        return 1
    if args.role == "candidate" and view_mixing_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": 0,
            "view_mixing_count": 0,
            "view_mixing_per_shot": {},
            "violations_detail": [],
            "trigger_phrases_used": list(_ID_BODY_PART_TRIGGERS),
            "full_body_verbs_used": list(_VIEW_MIXING_FULL_BODY_VERBS),
        },
        "measurement_failures": measurement_failures,
        "exit_criteria_threshold": 0,
        "exit_criteria_pass": False,
        "measurement_scripts": {
            "view_mixing": (
                "scripts/canary/g4_4_view_mixing.py (all shots)"
            )
        },
    }


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