#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""G4.3 canary — demographic descriptor presence ratio (Task 4.4).

Measures the ratio of FIRST id appearances (per
`(scene_index, shot_index, variation_index, id_token)` tuple) that have any
ethnicity component (10 entries) OR age band (6 entries) within ±50 chars
window. Override R1R2-B3 — char-level proximity, NO sentence regex.

Scope: pinned (si, shi) tuples — this canary measures the denominator
(all first-appearance id tuples in pinned shots) and the numerator (those
with descriptor present). Composite-ID-only shot scope per spec §5.1 is
realised structurally — denominator is naturally 0 for shots with no
composite IDs (R1R2-B3 / R3-I1).

Override R4-M3 carry — ratio floor:
  if total_first_id_appearances < 10, append measurement_failure
  ("baseline scene set too small for stable ratio per Override R2-M6")
  and exit 1. Do NOT silently return ratio = 0.0.

Exit criteria (analyzer, OUT-OF-BAND):
  candidate.metrics.demographic_descriptor_present_ratio
    >= baseline.metrics.demographic_descriptor_present_ratio  (no degradation)

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

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

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_ETHNICITY_COMPONENTS = _CONSTS["_ID_ETHNICITY_COMPONENTS"]
_ID_AGE_BANDS = _CONSTS["_ID_AGE_BANDS"]
_ID_REPRODUCTION_SURFACES = _CONSTS["_ID_REPRODUCTION_SURFACES"]

# bare C## OR composite C##O## — both forms count as first-appearance ids.
_ID_REGEX = re.compile(r"\bC\d{2}(?:O\d{2})?\b")
WINDOW_CHARS = 50

# R4-M3: minimum sample for stable ratio. < 10 → measurement_failure (no
# silent 0.0). Mirrors R2-M6 minimum scene-set size.
RATIO_MIN_SAMPLE = 10


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


def first_appearance_descriptor_present(prompt: str) -> tuple[dict[str, int], set[str]]:
    """Return (first_pos_by_id_token, descriptor_present_id_token_set).

    For each id token, only the FIRST occurrence in the prompt is recorded.
    descriptor_present is the subset of first-occurrence id tokens whose
    [start-50, end+50] window contains any ethnicity OR age-band substring
    (case-insensitive, char-level).
    """
    seen_ids: dict[str, int] = {}
    descriptor_present: set[str] = set()
    pl = prompt.lower()

    for match in _ID_REGEX.finditer(prompt):
        id_token = match.group()
        if id_token in seen_ids:
            continue  # not first appearance — skip per (si, shi, vi, id) tuple key
        seen_ids[id_token] = match.start()

        win_start = max(0, match.start() - WINDOW_CHARS)
        win_end = min(len(prompt), match.end() + WINDOW_CHARS)
        window = pl[win_start:win_end]

        matched = False
        for ethnicity in _ID_ETHNICITY_COMPONENTS:
            if ethnicity.lower() in window:
                descriptor_present.add(id_token)
                matched = True
                break
        if matched:
            continue
        for age in _ID_AGE_BANDS:
            if age.lower() in window:
                descriptor_present.add(id_token)
                break

    return seen_ids, descriptor_present


def main() -> int:
    parser = argparse.ArgumentParser(
        description="G4.3 canary: demographic descriptor presence ratio (composite-ID first-appearance, ±50 chars)."
    )
    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("--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_demographic_descriptor_present] ERROR: --cp-root required",
              file=sys.stderr)
        return 1
    cp_root = Path(cp_root_str)
    if not cp_root.exists():
        print(f"[g4_3_demographic_descriptor_present] 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,
                "shots_with_composite_id": 0,
                "demographic_descriptor_total_first_id_appearances": 0,
                "demographic_descriptor_present_first_appearances": 0,
                "demographic_descriptor_present_ratio": 0.0,
            },
            "measurement_failures": measurement_failures,
            "measurement_scripts": {
                "demographic_descriptor_present": (
                    "scripts/canary/g4_3_demographic_descriptor_present.py "
                    "(composite-ID shots only, 4-tuple denominator — R1R2-B3 / R4-M3)"
                )
            },
        }
        _emit(out, args)
        return 1

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

    total_shots = 0
    shots_with_composite_id = 0
    total_first_id_appearances = 0
    descriptor_present_first_appearances = 0
    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)))

        # R3R4-B2: silent absorb 차단. NO `scene.get("t2i_variations", [])` —
        # _validate_scene_preflight() returns None on missing/non-list.
        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.
        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 _ID_REGEX.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
        if has_composite_id:
            shots_with_composite_id += 1

        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
            seen, descriptor_present = first_appearance_descriptor_present(prompt)
            # 4-tuple denominator: (si, shi, vi, id_token) — within a single
            # variation each unique id contributes once. Across variations the
            # same id is counted again because vi differs (R1R2-B3).
            total_first_id_appearances += len(seen)
            descriptor_present_first_appearances += len(descriptor_present)

    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"
            )

    # R4-M3 floor — < 10 first-appearance samples means ratio is statistically
    # unstable. Append measurement_failure and exit 1 (NO silent 0.0).
    if total_first_id_appearances < RATIO_MIN_SAMPLE:
        measurement_failures.append(
            f"baseline scene set too small for stable ratio per Override R2-M6: "
            f"total_first_id_appearances={total_first_id_appearances} < "
            f"RATIO_MIN_SAMPLE={RATIO_MIN_SAMPLE}"
        )

    if total_first_id_appearances > 0:
        ratio = (
            descriptor_present_first_appearances / total_first_id_appearances
        )
    else:
        ratio = 0.0

    out = {
        "timestamp": iso8601_now(),
        "prompt_version": args.prompt_version,
        "role": args.role,
        "pinning": pinning,
        "scene_set": scene_set_out,
        "metrics": {
            "total_shots": total_shots,
            "shots_with_composite_id": shots_with_composite_id,
            "demographic_descriptor_total_first_id_appearances": total_first_id_appearances,
            "demographic_descriptor_present_first_appearances": descriptor_present_first_appearances,
            "demographic_descriptor_present_ratio": ratio,
        },
        "measurement_failures": measurement_failures,
        "measurement_scripts": {
            "demographic_descriptor_present": (
                "scripts/canary/g4_3_demographic_descriptor_present.py "
                "(composite-ID shots only, 4-tuple denominator — R1R2-B3 / R4-M3)"
            )
        },
    }

    _emit(out, args)

    print(
        f"[g4_3_demographic_descriptor_present] role={args.role} prompt_version={args.prompt_version}\n"
        f"  total_shots={total_shots}  shots_with_composite_id={shots_with_composite_id}\n"
        f"  total_first_id_appearances={total_first_id_appearances}\n"
        f"  descriptor_present_first_appearances={descriptor_present_first_appearances}\n"
        f"  ratio={ratio:.4f}\n"
        f"  measurement_failures: {len(measurement_failures)}\n"
        f"Exit criteria (analyzer): candidate.ratio >= baseline.ratio (degradation block)",
        file=sys.stderr,
    )

    if measurement_failures:
        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())
