"""experiment_background_render_payload_slice — W13 dry-run.

Consumes a prior W12 floor_plan_generation_slice run (its three Stage outputs
plus its run_meta) and assembles, per background id, the payload that
production ``background_render`` would receive immediately before calling
``gpt-image-2``. This stage is deterministic, makes no image API call, makes
no DB write, and does not modify production code.

Flow:
  1. Load W12 artifacts (director_set_layout_brief / floor_plan_prompt_candidate
     / per_bg_render_reference_instruction / run_meta).
  2. Resolve the W11 run dir via W12 args + derived_from (robust resolution).
  3. Walk derived_from chain (W11 -> W7 -> ... -> W3) to obtain source_bundle
     project_id / episode_id; then load the production floor_plan_render
     manifest for png_path / png_exists per fp_id. No PNG bytes are read.
  4. For each bg, build a payload preview: model = gpt-image-2,
     api_method_preview (images.edit if layout_only + fp_id, else
     images.generate), floor_plan_ref_path / floor_plan_ref_exists, decorated
     use/ignore/unlisted label lists, and the final prompt text composed of
     the W12 final_prompt_assembly_preview plus a literal floor-plan
     reference instruction block.
  5. Run six structural invariants. Render HTML with the per-bg table as the
     first content section and raw JSON inside collapsed <details>.

CLI:
  --derive-render-payload-from <W12_run_dir>   (required)
  --output-root <path>                          (defaults to scripts_output)
  --diag-print-imports
"""
from __future__ import annotations

import argparse
import json
import sys
from datetime import datetime
from pathlib import Path
from typing import Any, Dict, List, Optional

_REPO_ROOT = Path(__file__).resolve().parents[2]
_SCRIPTS_DIR = _REPO_ROOT / "backend" / "scripts"
if str(_SCRIPTS_DIR) not in sys.path:
    sys.path.insert(0, str(_SCRIPTS_DIR))

from experiment_background_pipeline_slice import (  # type: ignore
    KST,
    PLAN_VERSION,
    W6_IMAGE_BACKEND,
    _check_image_imports_present,
    _check_production_diff_empty,
    _load_production_floor_plan_context,
    _maybe_print_imports,
    _resolve_source_bundle_via_derived_from_chain,
)

W13_STAGE = "w13_background_render_payload_slice"
W13_IMAGE_BACKEND = W6_IMAGE_BACKEND  # carry-only; this stage makes no image call

_DEFAULT_OUTPUT_ROOT = (
    _REPO_ROOT / "scripts_output" / "background_render_payload_slice_experiment"
)
_W11_PARENT_DIR = (
    _REPO_ROOT / "scripts_output" / "background_pipeline_slice_experiment"
)


def _run_id() -> str:
    import secrets

    return datetime.now(KST).strftime("%Y%m%d_%H%M") + "_" + secrets.token_hex(3)


def _parse_args(argv):
    p = argparse.ArgumentParser(
        description="W13 background_render_payload_slice — payload preview only, no image API call"
    )
    p.add_argument(
        "--derive-render-payload-from",
        required=True,
        help="Path to a prior W12 success run dir (containing director_set_layout_brief.json, "
             "floor_plan_prompt_candidate.json, per_bg_render_reference_instruction.json, run_meta.json).",
    )
    p.add_argument("--output-root", default=str(_DEFAULT_OUTPUT_ROOT))
    p.add_argument("--diag-print-imports", action="store_true")
    return p.parse_args(argv)


# ─────────────────────────────────────────────────────────────────────────────
# W12 artifact loader
# ─────────────────────────────────────────────────────────────────────────────

def _load_w12_artifacts(prev_run_dir: Path) -> dict:
    required = {
        "brief": "director_set_layout_brief.json",
        "candidate": "floor_plan_prompt_candidate.json",
        "per_bg": "per_bg_render_reference_instruction.json",
        "run_meta": "run_meta.json",
    }
    out: Dict[str, Any] = {}
    missing: List[str] = []
    for key, fname in required.items():
        p = prev_run_dir / fname
        if not p.exists():
            missing.append(fname)
            continue
        out[key] = json.loads(p.read_text())
    out["_missing"] = missing
    out["_prev_run_id"] = prev_run_dir.name
    return out


def _resolve_w11_run_dir(
    *, w12_args_value: Optional[str],
    derived_from: Optional[str],
    w12_run_dir: Path,
) -> Optional[Path]:
    """Robust resolution of the W11 run dir referenced by the W12 run.

    Tries (in order): absolute path / repo-root relative / W12-run-dir relative
    / run-id lookup under the standard background_pipeline_slice_experiment
    parent directory. Returns ``None`` if no candidate has a run_meta.json.
    """
    candidates: List[Path] = []
    if w12_args_value:
        raw = Path(w12_args_value)
        if raw.is_absolute():
            candidates.append(raw)
        else:
            candidates.append(_REPO_ROOT / w12_args_value)
            candidates.append(w12_run_dir / w12_args_value)
            try:
                candidates.append((w12_run_dir / w12_args_value).resolve())
            except OSError:
                pass
    if derived_from:
        candidates.append(_W11_PARENT_DIR / derived_from)
    seen: set = set()
    for c in candidates:
        key = str(c)
        if key in seen:
            continue
        seen.add(key)
        try:
            if c.exists() and (c / "run_meta.json").exists():
                return c
        except OSError:
            continue
    return None


# ─────────────────────────────────────────────────────────────────────────────
# Payload assembly
# ─────────────────────────────────────────────────────────────────────────────

def _fp_id_to_number_label_map(candidates: dict) -> Dict[str, Dict[int, str]]:
    out: Dict[str, Dict[int, str]] = {}
    for fp_id, fp in (candidates or {}).items():
        m: Dict[int, str] = {}
        for e in (fp.get("candidate_numbered_elements") or []):
            num = e.get("number")
            if isinstance(num, int):
                m[num] = e.get("label") or ""
        out[fp_id] = m
    return out


def _decorate_nums_with_labels(nums: List[int], label_map: Dict[int, str]) -> List[str]:
    """Pure structural label join. Returns list of '#N label' strings."""
    decorated: List[str] = []
    for n in nums or []:
        if not isinstance(n, int):
            continue
        label = label_map.get(n) or ""
        if label:
            decorated.append(f"#{n} {label}")
        else:
            decorated.append(f"#{n}")
    return decorated


def _build_reference_block(use_labels: List[str], ignore_labels: List[str],
                           unlisted_labels: List[str]) -> str:
    """Deterministic literal template. No scenario hardcode; all labels come
    dynamically from the W12 candidate data."""
    lines = ["Floor-plan reference (layout-only):"]
    if use_labels:
        lines.append(
            "- Use these numbered anchors as layout cues only: "
            + ", ".join(use_labels) + "."
        )
    else:
        lines.append("- No numbered anchors are used for this background.")
    if ignore_labels:
        lines.append(
            "- Do not draw or anchor on these numbered elements: "
            + ", ".join(ignore_labels) + "."
        )
    if unlisted_labels:
        lines.append(
            "- The following numbered elements are unlisted candidate elements "
            "— they are not explicit anchors. Do NOT add them unless the base "
            "prompt above already requests them: " + ", ".join(unlisted_labels) + "."
        )
    return "\n".join(lines)


def _build_per_bg_payloads(
    *, per_bg: Dict[str, dict],
    candidates: Dict[str, dict],
    fp_context: dict,
) -> Dict[str, dict]:
    fp_label_map = _fp_id_to_number_label_map(candidates)
    fp_candidate_numbers: Dict[str, set] = {
        fp_id: set(label_map.keys()) for fp_id, label_map in fp_label_map.items()
    }
    prod_fps = (fp_context or {}).get("floor_plans") or {}

    payloads: Dict[str, dict] = {}
    for bg_id, instr in (per_bg or {}).items():
        fp_id = instr.get("fp_id") or ""
        role = instr.get("floor_plan_ref_role") or "not_used"
        use_nums: List[int] = [n for n in (instr.get("use_numbered_elements") or []) if isinstance(n, int)]
        ignore_nums: List[int] = [n for n in (instr.get("ignore_numbered_elements") or []) if isinstance(n, int)]

        label_map = fp_label_map.get(fp_id, {})
        candidate_set = fp_candidate_numbers.get(fp_id, set())
        unlisted_nums = sorted(candidate_set - set(use_nums) - set(ignore_nums))

        use_labels = _decorate_nums_with_labels(sorted(use_nums), label_map)
        ignore_labels = _decorate_nums_with_labels(sorted(ignore_nums), label_map)
        unlisted_labels = _decorate_nums_with_labels(unlisted_nums, label_map)

        prod_fp = prod_fps.get(fp_id) if fp_id else None
        floor_plan_ref_path = ((prod_fp or {}).get("png_path") if prod_fp else "") or ""
        floor_plan_ref_exists = bool(prod_fp and prod_fp.get("png_exists"))

        if role == "not_used" or not fp_id:
            api_method_preview = "images.generate"
            ref_block = ""
        else:
            api_method_preview = "images.edit"
            ref_block = _build_reference_block(use_labels, ignore_labels, unlisted_labels)

        base = instr.get("final_prompt_assembly_preview") or ""
        if ref_block:
            final_prompt = base.rstrip() + "\n\n" + ref_block
        else:
            final_prompt = base

        warnings: List[str] = []
        if role == "layout_only" and not use_nums:
            warnings.append("layout_only_but_use_empty")
        if role == "layout_only" and not fp_id:
            warnings.append("layout_only_but_fp_id_empty")
        if api_method_preview == "images.edit":
            if not floor_plan_ref_path:
                warnings.append("edit_preview_missing_floor_plan_ref_path")
            elif not floor_plan_ref_exists:
                warnings.append("edit_preview_floor_plan_ref_not_on_disk")

        payloads[bg_id] = {
            "bg_id": bg_id,
            "fp_id": fp_id,
            "image_model": W13_IMAGE_BACKEND,
            "api_method_preview": api_method_preview,
            "floor_plan_ref_role": role,
            "floor_plan_ref_path": floor_plan_ref_path,
            "floor_plan_ref_exists": floor_plan_ref_exists,
            "prior_bg_ref_role": instr.get("prior_bg_ref_role") or "none",
            "applies_to_shots": list(instr.get("applies_to_shots") or []),
            "camera_axis_used": instr.get("camera_axis_used") or "",
            "camera_axis_source": instr.get("camera_axis_source") or "",
            "visible_zone_scope": list(instr.get("visible_zone_scope") or []),
            "use_numbered_elements_with_labels": use_labels,
            "ignore_numbered_elements_with_labels": ignore_labels,
            "unlisted_candidate_elements_with_labels": unlisted_labels,
            "render_prompt_appendix": instr.get("render_prompt_appendix") or "",
            "final_prompt_text_for_image_preview": final_prompt,
            "warnings_advisories": warnings,
        }
    return payloads


# ─────────────────────────────────────────────────────────────────────────────
# Compatibility report
# ─────────────────────────────────────────────────────────────────────────────

def _build_w13_compatibility_report(
    *, per_bg: Dict[str, dict],
    payloads: Dict[str, dict],
    candidates: Dict[str, dict],
    fp_context: dict,
    production_diff_empty: bool,
    db_write_count: int,
    image_import_seen: bool,
    missing_inputs: List[str],
    prev_run_id: str,
    stage_status: str,
) -> dict:
    inv: Dict[str, Dict[str, Any]] = {}

    inv["w12_inputs_present"] = {
        "pass": not missing_inputs,
        "detail": {
            "prev_run_id": prev_run_id,
            "missing_inputs": list(missing_inputs),
            "stage_status": stage_status,
        },
    }

    per_bg_ids = set((per_bg or {}).keys())
    payload_ids = set((payloads or {}).keys())
    inv["all_bg_payloads_covered"] = {
        "pass": bool(per_bg_ids) and per_bg_ids == payload_ids,
        "detail": {
            "per_bg_count": len(per_bg_ids),
            "payload_count": len(payload_ids),
            "missing_in_payloads": sorted(per_bg_ids - payload_ids),
            "extra_in_payloads": sorted(payload_ids - per_bg_ids),
        },
    }

    fp_candidate_numbers: Dict[str, set] = {}
    for fp_id, fp in (candidates or {}).items():
        s: set = set()
        for e in (fp.get("candidate_numbered_elements") or []):
            n = e.get("number")
            if isinstance(n, int):
                s.add(n)
        fp_candidate_numbers[fp_id] = s

    label_unresolved: List[str] = []
    for bg_id, instr in (per_bg or {}).items():
        fp_id = instr.get("fp_id") or ""
        if not fp_id:
            continue
        s = fp_candidate_numbers.get(fp_id, set())
        for n in (instr.get("use_numbered_elements") or []):
            if isinstance(n, int) and n not in s:
                label_unresolved.append(f"{bg_id}:use:#{n}")
        for n in (instr.get("ignore_numbered_elements") or []):
            if isinstance(n, int) and n not in s:
                label_unresolved.append(f"{bg_id}:ignore:#{n}")

    edit_missing_refs: List[str] = []
    for bg_id, payload in (payloads or {}).items():
        if payload.get("api_method_preview") != "images.edit":
            continue
        if not payload.get("floor_plan_ref_path"):
            edit_missing_refs.append(f"{bg_id}:no_path")
        elif not payload.get("floor_plan_ref_exists"):
            edit_missing_refs.append(f"{bg_id}:not_on_disk")
    inv["labels_and_floor_plan_refs_resolved"] = {
        "pass": not label_unresolved and not edit_missing_refs,
        "detail": {
            "label_unresolved": label_unresolved,
            "edit_missing_refs": edit_missing_refs,
        },
    }

    wrong_model = [
        b for b, p in (payloads or {}).items()
        if p.get("image_model") != W13_IMAGE_BACKEND
    ]
    inv["model_is_gpt_image_2"] = {
        "pass": not wrong_model,
        "detail": {
            "wrong_model_bg_ids": wrong_model,
            "expected": W13_IMAGE_BACKEND,
        },
    }

    inv["no_image_api_call_and_zero_count"] = {
        "pass": not image_import_seen,
        "detail": {
            "image_import_seen": image_import_seen,
            "image_generation_count": 0,
        },
    }

    inv["production_diff_zero_and_db_write_zero"] = {
        "pass": production_diff_empty and db_write_count == 0,
        "detail": {
            "production_diff_empty": production_diff_empty,
            "db_write_count": db_write_count,
        },
    }

    all_pass = all(v["pass"] for v in inv.values())
    return {"invariants": inv, "all_pass": all_pass}


# ─────────────────────────────────────────────────────────────────────────────
# HTML render
# ─────────────────────────────────────────────────────────────────────────────

def _render_w13_html(run_meta: dict, payloads: Dict[str, dict], report: dict,
                     run_dir: Path) -> None:
    def esc(x):
        return (str(x).replace("&", "&amp;").replace("<", "&lt;").replace(">", "&gt;"))

    inv = (report or {}).get("invariants", {}) or {}
    inv_rows = "".join(
        f"<tr><td>{esc(k)}</td>"
        f"<td class=\"{'pass' if v['pass'] else 'fail'}\">{'PASS' if v['pass'] else 'FAIL'}</td>"
        f"<td><pre>{esc(json.dumps(v.get('detail'), ensure_ascii=False))[:400]}</pre></td></tr>"
        for k, v in inv.items()
    )

    per_bg_rows = ""
    for bg_id, p in payloads.items():
        use_labels = ", ".join(esc(s) for s in p.get("use_numbered_elements_with_labels") or [])
        ignore_labels = ", ".join(esc(s) for s in p.get("ignore_numbered_elements_with_labels") or [])
        unlisted_labels = ", ".join(esc(s) for s in p.get("unlisted_candidate_elements_with_labels") or [])
        warnings = p.get("warnings_advisories") or []
        warn_cell = (
            "<br><small class=\"warn\">" + esc(", ".join(warnings)) + "</small>"
            if warnings else ""
        )
        per_bg_rows += (
            f"<tr><td>{esc(bg_id)}{warn_cell}</td>"
            f"<td>{esc(p.get('fp_id'))}</td>"
            f"<td>{esc(p.get('api_method_preview'))}</td>"
            f"<td>{'Y' if p.get('floor_plan_ref_exists') else 'N'}</td>"
            f"<td>{use_labels}</td>"
            f"<td>{ignore_labels}</td>"
            f"<td>{unlisted_labels}</td>"
            f"<td>{esc(p.get('prior_bg_ref_role'))}</td>"
            f"<td><pre>{esc(p.get('final_prompt_text_for_image_preview') or '')[:2400]}</pre></td></tr>"
        )

    html = f"""<!doctype html><html><head><meta charset=\"utf-8\">
<title>W13 background_render_payload_slice {esc(run_meta.get('run_id'))}</title>
<style>body{{font-family:sans-serif;margin:1.5em}}
table{{border-collapse:collapse;margin:0.5em 0}} td,th{{border:1px solid #ccc;padding:4px 8px;vertical-align:top}}
.pass{{color:#080}} .fail{{color:#b00}} .warn{{color:#a60}}
pre{{white-space:pre-wrap;font-size:0.85em;max-width:72ch}}
section{{margin:1.5em 0}}</style></head>
<body>
<h1>W13 — background_render_payload_slice {esc(run_meta.get('run_id'))}</h1>
<p>stage: <b>{esc(run_meta.get('stage'))}</b>
| run_status: <b>{esc(run_meta.get('run_status'))}</b>
| exit_code: {esc(run_meta.get('exit_code'))}
| derived_from(W12): {esc(run_meta.get('derived_from'))}
| image_generation_count: <b>{esc(run_meta.get('image_generation_count'))}</b>
| image_generation_backend: <b>{esc(run_meta.get('image_generation_backend'))}</b></p>

<section><h2>1. Per-bg render payload preview</h2>
<table><tr>
<th>bg_id</th><th>fp_id</th><th>api_method_preview</th><th>fp ref exists</th>
<th>use labels</th><th>ignore labels</th><th>unlisted labels</th>
<th>prior role</th><th>final prompt preview</th>
</tr>{per_bg_rows}</table></section>

<section><h2>2. Invariants</h2>
<table><tr><th>invariant</th><th>status</th><th>detail</th></tr>{inv_rows}</table></section>

<details><summary>raw run_meta.json</summary>
<pre>{esc(json.dumps(run_meta, ensure_ascii=False, indent=2))}</pre></details>
<details><summary>raw payloads.json</summary>
<pre>{esc(json.dumps(payloads, ensure_ascii=False, indent=2))[:200000]}</pre></details>
</body></html>"""
    (run_dir / "index.html").write_text(html)


# ─────────────────────────────────────────────────────────────────────────────
# main
# ─────────────────────────────────────────────────────────────────────────────

def main(argv=None) -> int:
    args = _parse_args(argv)
    run_id = _run_id()
    out_root = Path(args.output_root)
    run_dir = out_root / run_id
    run_dir.mkdir(parents=True, exist_ok=True)

    prev_run_dir = Path(args.derive_render_payload_from)
    if not prev_run_dir.is_absolute():
        prev_run_dir = Path.cwd() / prev_run_dir

    artifacts = _load_w12_artifacts(prev_run_dir)
    missing = list(artifacts.get("_missing", []))
    outputs: List[str] = []
    failed: List[str] = []
    run_status = "succeeded"
    exit_code = 0
    stage_status = "generated"

    run_meta: Dict[str, Any] = {
        "run_id": run_id,
        "stage": W13_STAGE,
        "plan_version": PLAN_VERSION,
        "generated_at": datetime.now(KST).isoformat(),
        "image_generation_count": 0,
        "image_generation_backend": W13_IMAGE_BACKEND,
        "args": vars(args),
        "derived_from": prev_run_dir.name,
        "outputs": outputs,
        "run_status": run_status,
        "exit_code": exit_code,
        "failed_invariants": failed,
    }

    if missing:
        failed.append("w12_inputs_missing")
        run_meta["run_status"] = "validation_failed"
        run_meta["exit_code"] = 1
        run_meta["failed_invariants"] = failed
        (run_dir / "run_meta.json").write_text(
            json.dumps(run_meta, ensure_ascii=False, indent=2)
        )
        return 1

    w12_brief = artifacts["brief"]
    w12_candidates = (artifacts["candidate"].get("candidate_floor_plans") or {})
    w12_per_bg = (artifacts["per_bg"].get("per_bg_render_reference_instructions") or {})
    w12_run_meta = artifacts["run_meta"]

    # Resolve W11 run dir → walk the derived_from chain to source_bundle.json.
    w12_args = (w12_run_meta or {}).get("args") or {}
    w11_run_dir = _resolve_w11_run_dir(
        w12_args_value=w12_args.get("derive_floor_plan_candidate_from"),
        derived_from=w12_run_meta.get("derived_from"),
        w12_run_dir=prev_run_dir,
    )
    fp_context: Dict[str, Any] = {"floor_plans": {}}
    chain_status = "unresolved"
    chain_info: Dict[str, Any] = {}
    if w11_run_dir is None:
        chain_status = "w11_run_dir_unresolved"
    else:
        chain_info["w11_run_id"] = w11_run_dir.name
        try:
            resolved = _resolve_source_bundle_via_derived_from_chain(w11_run_dir)
            source_bundle = resolved["source_bundle"]
            chain_info["chain"] = resolved["chain"]
            chain_info["root_run_id"] = resolved["root_run_id"]
            project_id = source_bundle.get("project_id") or ""
            episode_id = source_bundle.get("episode_id") or ""
            chain_info["project_id"] = project_id
            chain_info["episode_id"] = episode_id
            fp_context = _load_production_floor_plan_context(project_id, episode_id)
            chain_status = "resolved"
        except FileNotFoundError as exc:
            chain_status = "source_bundle_unresolved"
            chain_info["error"] = str(exc)[:300]

    run_meta["chain_status"] = chain_status
    run_meta["chain_info"] = chain_info

    payloads = _build_per_bg_payloads(
        per_bg=w12_per_bg, candidates=w12_candidates, fp_context=fp_context,
    )
    (run_dir / "background_render_payload_preview.json").write_text(
        json.dumps({"payloads": payloads}, ensure_ascii=False, indent=2)
    )
    outputs.append("background_render_payload_preview.json")

    report = _build_w13_compatibility_report(
        per_bg=w12_per_bg, payloads=payloads,
        candidates=w12_candidates, fp_context=fp_context,
        production_diff_empty=_check_production_diff_empty(),
        db_write_count=0,
        image_import_seen=_check_image_imports_present(),
        missing_inputs=missing,
        prev_run_id=prev_run_dir.name,
        stage_status=stage_status,
    )
    (run_dir / "w13_compatibility_report.json").write_text(
        json.dumps(report, ensure_ascii=False, indent=2)
    )
    outputs.append("w13_compatibility_report.json")

    for name, v in report["invariants"].items():
        if not v["pass"] and name not in failed:
            failed.append(name)
    if failed:
        run_status = "validation_failed"
        exit_code = 1

    run_meta["run_status"] = run_status
    run_meta["exit_code"] = exit_code
    run_meta["failed_invariants"] = failed
    run_meta["outputs"] = outputs
    _render_w13_html(run_meta, payloads, report, run_dir)
    outputs.append("index.html")
    run_meta["outputs"] = outputs
    (run_dir / "run_meta.json").write_text(
        json.dumps(run_meta, ensure_ascii=False, indent=2)
    )
    _maybe_print_imports(args)
    return exit_code


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