"""experiment_floor_plan_grid_layout_model_compare_slice — W16c.

Sibling comparison wave for W16. The user rejected the W16 GPT-5.5 grid
layout (room placement infeasible). This wave asks Gemini 3.1 Pro Preview
the same prompt with the same W15e inputs and renders the result with the
same deterministic SVG so the two models can be compared head-to-head.

The original W16 GPT-5.5 fail-closed contract stays intact: this script
imports W16 helpers but never modifies the W16 module's default paths or
defaults.

CLI:
  --derive-grid-layout-from <W15e_run_dir>      (required)
  --target-fp-ids fp_l05_01                     (default, restricted)
  --generate                                     (default off — dry-run)
  --model gemini/gemini-3.1-pro-preview          (default; no fallback)
  --output-root <path>
  --diag-print-imports

No image API call, no DB write, no production manifest mutation, no
commit, and no auto retry — a single LLM attempt only.
"""
from __future__ import annotations

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

_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,
    _check_image_imports_present,
    _check_production_diff_empty,
    _load_backend_env,
    _maybe_print_imports,
)
from experiment_floor_plan_grid_layout_slice import (  # type: ignore
    GRID_COLS,
    GRID_ROWS,
    W16_ALLOWED_TARGET_FP_IDS,
    W16_SYSTEM_PROMPT,
    _build_w16_compatibility_report,
    _build_w16_llm_input,
    _load_w15e_artifacts,
    _placeholder_dry_run_layout,
    _render_w16_html,
    _render_w16_svg,
    _resolve_targets,
)

W16C_STAGE = "w16c_floor_plan_grid_layout_model_compare_slice"
W16C_EXPECTED_MODEL_ID = "gemini/gemini-3.1-pro-preview"
W16C_MODEL_INVARIANT_KEY = "model_is_gemini_3_1_pro_when_generated"

_DEFAULT_OUTPUT_ROOT = (
    _REPO_ROOT / "scripts_output" / "floor_plan_grid_layout_model_compare_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=(
            "W16c — Gemini 3.1 Pro Preview comparison render of the W16 grid "
            "layout prompt. No image API call. No production mutation."
        )
    )
    p.add_argument(
        "--derive-grid-layout-from", required=True,
        help="Path to a prior W15e success run dir.",
    )
    p.add_argument(
        "--target-fp-ids", default="fp_l05_01",
        help="Comma-separated fp_id subset. Wave 1 only allows fp_l05_01.",
    )
    p.add_argument(
        "--generate", action="store_true",
        help="Actual LLM call (Gemini 3.1 Pro Preview). Default off — dry-run.",
    )
    p.add_argument("--model", default=W16C_EXPECTED_MODEL_ID)
    p.add_argument("--output-root", default=str(_DEFAULT_OUTPUT_ROOT))
    p.add_argument("--diag-print-imports", action="store_true")
    return p.parse_args(argv)


# ─────────────────────────────────────────────────────────────────────────────
# Gemini caller — fail-closed, no fallback, no retry by default
# ─────────────────────────────────────────────────────────────────────────────

def _generate_w16c_via_gemini(llm_input: dict, *, model: str = W16C_EXPECTED_MODEL_ID,
                              retry_once: bool = False) -> dict:
    """Call Gemini via litellm. Fail-closed:
    - missing GEMINI_API_KEY / GOOGLE_API_KEY → RuntimeError
    - model id that is NOT exactly `W16C_EXPECTED_MODEL_ID`
      (`gemini/gemini-3.1-pro-preview`) → RuntimeError BEFORE litellm call,
      mirroring the W16 GPT-5.5 exact-id guard so a typo'd Gemini id
      cannot incur a paid call.
    - first litellm exception → RuntimeError (no auto retry by default)
    No fallback to other providers."""
    if not (os.environ.get("GEMINI_API_KEY") or os.environ.get("GOOGLE_API_KEY")):
        raise RuntimeError("missing GEMINI_API_KEY / GOOGLE_API_KEY env var")
    requested = (model or "").strip()
    if requested != W16C_EXPECTED_MODEL_ID:
        raise RuntimeError(
            f"w16c routing refuses non-exact gemini model: got '{requested}', "
            f"required exact id '{W16C_EXPECTED_MODEL_ID}' (no fallback allowed)"
        )

    import litellm  # lazy import — dry-run path must not import it

    user_prompt = json.dumps(llm_input, ensure_ascii=False)
    last_exc: Optional[Exception] = None
    attempts = 2 if retry_once else 1
    for _ in range(attempts):
        try:
            resp = litellm.completion(
                model=model,
                messages=[
                    {"role": "system", "content": W16_SYSTEM_PROMPT},
                    {"role": "user", "content": user_prompt},
                ],
                response_format={"type": "json_object"},
            )
            raw_text = resp.choices[0].message.content
            decoder = json.JSONDecoder()
            stripped = (raw_text or "").lstrip()
            parsed, _ = decoder.raw_decode(stripped)
            return parsed
        except Exception as exc:  # noqa: BLE001
            last_exc = exc
            continue
    raise RuntimeError(f"w16c_gemini_failed: {last_exc!s}"[:400])


# ─────────────────────────────────────────────────────────────────────────────
# Compatibility wrapper — reuses W16 invariants, swaps the model invariant key
# ─────────────────────────────────────────────────────────────────────────────

def build_w16c_compatibility_report(
    *, grid_layout: Dict[str, Any],
    topology_brief: Dict[str, Any],
    candidate: Dict[str, Any],
    target_fp_ids: Set[str],
    production_diff_empty: bool,
    db_write_count: int,
    image_import_seen: bool,
    image_api_call_count: int,
    svg_emitted: bool,
    model_used: Optional[str],
    stage_status: str,
    missing_inputs: List[str],
    prev_run_id: str,
    svg_paths_by_fp: Optional[Dict[str, Any]] = None,
) -> dict:
    return _build_w16_compatibility_report(
        grid_layout=grid_layout,
        topology_brief=topology_brief,
        candidate=candidate,
        target_fp_ids=target_fp_ids,
        production_diff_empty=production_diff_empty,
        db_write_count=db_write_count,
        image_import_seen=image_import_seen,
        image_api_call_count=image_api_call_count,
        svg_emitted=svg_emitted,
        model_used=model_used,
        stage_status=stage_status,
        missing_inputs=missing_inputs,
        prev_run_id=prev_run_id,
        svg_paths_by_fp=svg_paths_by_fp,
        expected_model_id=W16C_EXPECTED_MODEL_ID,
        model_invariant_key=W16C_MODEL_INVARIANT_KEY,
    )


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

def main(argv=None) -> int:
    args = _parse_args(argv)
    if args.generate:
        _load_backend_env()
    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_grid_layout_from)
    if not prev_run_dir.is_absolute():
        prev_run_dir = Path.cwd() / prev_run_dir

    artifacts = _load_w15e_artifacts(prev_run_dir)
    missing = list(artifacts.get("_missing", []))

    target_fp_ids, invalid_targets = _resolve_targets(args.target_fp_ids)

    failed_invariants: List[str] = []
    run_status = "succeeded"
    exit_code = 0
    model_used: Optional[str] = None
    stage_status = "dry_run"

    run_meta: Dict[str, Any] = {
        "run_id": run_id,
        "stage": W16C_STAGE,
        "plan_version": PLAN_VERSION,
        "generated_at": datetime.now(KST).isoformat(),
        "args": vars(args),
        "derived_from": prev_run_dir.name,
        "target_fp_ids": sorted(target_fp_ids),
        "invalid_targets": invalid_targets,
        "model_used": None,
        "expected_model": W16C_EXPECTED_MODEL_ID,
        "image_api_call_count": 0,
        "image_generation_count": 0,
        "outputs": [],
        "run_status": run_status,
        "exit_code": exit_code,
        "failed_invariants": failed_invariants,
        "stage_status": stage_status,
    }

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

    topology_brief = artifacts["topology"]
    candidate = artifacts["candidate"]

    grid_layout_by_fp: Dict[str, Any] = {}
    if args.generate:
        try:
            for fp_id in sorted(target_fp_ids):
                llm_input = _build_w16_llm_input(
                    topology_brief=topology_brief,
                    candidate=candidate,
                    target_fp_id=fp_id,
                )
                parsed = _generate_w16c_via_gemini(
                    llm_input, model=args.model, retry_once=False,
                )
                by_fp = (parsed or {}).get("grid_layout_by_fp") or {}
                if fp_id not in by_fp and "grid" in (parsed or {}):
                    by_fp = {fp_id: parsed}
                grid_layout_by_fp.update(
                    {k: v for k, v in by_fp.items() if k == fp_id}
                )
            model_used = args.model
            stage_status = "generated"
        except Exception as exc:  # noqa: BLE001
            failed_invariants.append("llm_call_failed")
            run_meta["llm_error"] = str(exc)[:400]
            run_status = "validation_failed"
            exit_code = 1
            stage_status = "llm_failed"
    else:
        skeleton = _placeholder_dry_run_layout(target_fp_ids=target_fp_ids)
        grid_layout_by_fp = skeleton.get("grid_layout_by_fp") or {}
        model_used = None
        stage_status = "dry_run"

    grid_layout = {"grid_layout_by_fp": grid_layout_by_fp}
    (run_dir / "grid_layout_plan.json").write_text(
        json.dumps(grid_layout, ensure_ascii=False, indent=2)
    )
    run_meta["outputs"].append("grid_layout_plan.json")

    svg_paths_by_fp: Dict[str, Path] = {}
    for fp_id, fp_layout in grid_layout_by_fp.items():
        if not isinstance(fp_layout, dict):
            continue
        units = fp_layout.get("units") or []
        if not units:
            continue
        try:
            svg_paths_by_fp[fp_id] = _render_w16_svg(
                fp_id=fp_id, fp_layout=fp_layout, run_dir=run_dir,
            )
        except Exception as exc:  # noqa: BLE001
            run_meta.setdefault("svg_errors", []).append({
                "fp_id": fp_id, "error": str(exc)[:240],
            })
    svg_emitted = bool(svg_paths_by_fp)

    report = build_w16c_compatibility_report(
        grid_layout=grid_layout,
        topology_brief=topology_brief,
        candidate=candidate,
        target_fp_ids=target_fp_ids,
        production_diff_empty=_check_production_diff_empty(),
        db_write_count=0,
        image_import_seen=_check_image_imports_present(),
        image_api_call_count=0,
        svg_emitted=svg_emitted,
        model_used=model_used,
        stage_status=stage_status,
        missing_inputs=missing,
        prev_run_id=prev_run_dir.name,
        svg_paths_by_fp=svg_paths_by_fp,
    )
    (run_dir / "grid_layout_compatibility_report.json").write_text(
        json.dumps(report, ensure_ascii=False, indent=2)
    )
    run_meta["outputs"].append("grid_layout_compatibility_report.json")

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

    run_meta["model_used"] = model_used
    run_meta["stage_status"] = stage_status
    run_meta["run_status"] = run_status
    run_meta["exit_code"] = exit_code
    run_meta["failed_invariants"] = failed_invariants

    _render_w16_html(
        run_meta=run_meta, grid_layout=grid_layout,
        report=report, svg_paths_by_fp=svg_paths_by_fp, run_dir=run_dir,
    )
    run_meta["outputs"].append("index.html")
    for fp_id, p in svg_paths_by_fp.items():
        try:
            run_meta["outputs"].append(str(p.relative_to(run_dir)))
        except ValueError:
            run_meta["outputs"].append(p.name)

    (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__":
    sys.exit(main())
