"""W21B-w5 STEP5-B (gpt-image-2 sparse): FloorPlanLightSidecarStep gate + path tests.

opt-in (default OFF). The winning design = gpt-5.5 frequency-aware sparse
selection → gpt-image-2 PURE text-to-image best-of-N render. Tests cover
gating → byte-stable not_applicable, the selection+render success path (rendered
+ provenance + best-of-N draws), and the per-fp DETAILED fallback paths
(selection-miss / selection-invalid / render-failed / unsafe-id / no-elements)
that must NEVER fail the step (failed_count held at 0). The real LLM + image client
are bypassed via the select/render overrides; no network/spend.
"""
from __future__ import annotations

from unittest.mock import MagicMock, patch

from app.core.steps.floor_plan_light_sidecar_step import (
    PROMPT_VERSION,
    SCHEMA_VERSION,
    VALIDATION_MODE,
    FloorPlanLightSidecarStep,
)


def _elements():
    # 2 rooms + 1 opening + 3 fixtures/furniture + 1 overlay.
    return [
        {"number": 1, "label": "거실", "base_layer_decision": "base_structural_unit"},
        {"number": 2, "label": "침실", "base_layer_decision": "base_structural_unit"},
        {"number": 7, "label": "현관문", "base_layer_decision": "base_opening"},
        {"number": 16, "label": "TV", "base_layer_decision": "base_persistent_fixture"},
        {"number": 23, "label": "커튼", "base_layer_decision": "base_persistent_fixture"},
        {"number": 30, "label": "침대", "base_layer_decision": "base_persistent_furniture"},
        {"number": 88, "label": "핏자국", "base_layer_decision": "state_overlay_plot_cue"},
    ]


def _cameras():
    return [
        {"bg_id": "B01", "use_numbered_elements": [1, 16, 23], "ignore_numbered_elements": []},
        {"bg_id": "B02", "use_numbered_elements": [16, 23], "ignore_numbered_elements": []},
        {"bg_id": "B03", "use_numbered_elements": [2, 30], "ignore_numbered_elements": []},
    ]


def _good_selection():
    return {
        "geometry_skeleton": [
            {"number": 1, "label": "거실", "kind": "room"},
            {"number": 2, "label": "침실", "kind": "room"},
            {"number": 7, "label": "현관문", "kind": "door"},
        ],
        "essential_numbered_markers": [
            {"number": 16, "label": "TV", "reason": "framed often", "camera_use_count": 2},
            {"number": 23, "label": "커튼", "reason": "framed often", "camera_use_count": 2},
            {"number": 30, "label": "침대", "reason": "bedroom anchor", "camera_use_count": 1},
        ],
        "high_frequency_candidates": [],
        "metadata_only": [{"number": 88, "label": "핏자국"}],
        "room_schematic_prompt": "thick black enclosed walls, sparse numbered circles " * 3,
    }


def _prompt_cp(targets: dict) -> dict:
    """floor_plan_prompt checkpoint shape."""
    return {"floor_plan_prompt": {"data": {"floor_plans": targets}}}


def _fp_targets(*fids, elements=None, cameras=None):
    out = {}
    for fid in fids:
        out[fid] = {
            "status": "ok",
            "numbered_elements": elements if elements is not None else _elements(),
            "camera_recommendations": cameras if cameras is not None else _cameras(),
        }
    return out


def _new_step(*, cp_map: dict) -> FloorPlanLightSidecarStep:
    step = FloorPlanLightSidecarStep.__new__(FloorPlanLightSidecarStep)
    step.project_id = "p"
    step.episode_id = "e"
    step.project_config = {}
    step.build_opik_metadata = MagicMock(return_value={})
    step._select_override = None
    step._render_override = None
    step._client = None
    step._load_prev_checkpoint = MagicMock(side_effect=lambda sid: cp_map.get(sid))
    return step


def _select_ok(calls):
    def select(*, bundle):
        calls["sel"] += 1
        calls["bundles"].append(bundle)
        return _good_selection()
    return select


def _select_none(calls):
    def select(*, bundle):
        calls["sel"] += 1
        return None
    return select


def _select_invalid(calls):
    def select(*, bundle):
        calls["sel"] += 1
        bad = _good_selection()
        # drop the bedroom (2) from both buckets → a room is not drawn (forbidden).
        bad["geometry_skeleton"] = [
            s for s in bad["geometry_skeleton"] if s["number"] != 2]
        return bad
    return select


def _render_ok(calls):
    def render(*, prompt, out_path, fp_id):
        calls["img"] += 1
        calls["prompts"].append(prompt)
        out_path.parent.mkdir(parents=True, exist_ok=True)
        out_path.write_bytes(b"\x89PNG nb2")
        return {"status": "ok", "png_path": str(out_path), "attempts": 1, "error": ""}
    return render


def _render_fail(calls):
    def render(*, prompt, out_path, fp_id):
        calls["img"] += 1
        return {"status": "failed", "png_path": "", "attempts": 1, "error": "boom"}
    return render


def _patches(tmp_path, *, enabled=True, mode="on", fp_version="6", best_of_n=2):
    return [
        patch("app.core.config.settings.projects_dir", str(tmp_path / "proj")),
        patch("app.core.config.settings.background_mode", mode),
        patch("app.core.config.settings.floor_plan_prompt_version", fp_version),
        patch("app.core.config.settings.floor_plan_light_sidecar_enabled", enabled),
        patch("app.core.config.settings.floor_plan_light_sidecar_best_of_n", best_of_n),
    ]


def _run(step, patches):
    for p in patches:
        p.start()
    try:
        return step._execute()
    finally:
        for p in reversed(patches):
            p.stop()


# ──────────────────────────── constants / gates ────────────────────────────
def test_constants():
    assert SCHEMA_VERSION == 2
    assert PROMPT_VERSION == "2"
    assert VALIDATION_MODE == "selection_plus_render_visual_gate"


def test_gate_background_mode_off(tmp_path):
    step = _new_step(cp_map=_prompt_cp(_fp_targets("fp_a")))
    out = _run(step, _patches(tmp_path, mode="off"))
    assert out["applicable_count"] == 0
    assert out["data"] == {}
    assert out["schema_version"] == 2


def test_gate_flag_off(tmp_path):
    step = _new_step(cp_map=_prompt_cp(_fp_targets("fp_a")))
    out = _run(step, _patches(tmp_path, enabled=False))
    assert out["applicable_count"] == 0
    assert out["data"] == {}


def test_gate_prompt_version_unsupported(tmp_path):
    step = _new_step(cp_map=_prompt_cp(_fp_targets("fp_a")))
    out = _run(step, _patches(tmp_path, fp_version="5"))
    assert out["applicable_count"] == 0
    assert out["data"] == {}


def test_gate_v7_supported(tmp_path):
    step = _new_step(cp_map=_prompt_cp(_fp_targets("fp_a")))
    step._select_override = _select_ok({"sel": 0, "bundles": []})
    step._render_override = _render_ok({"img": 0, "prompts": []})
    out = _run(step, _patches(tmp_path, fp_version="7"))
    assert out["applicable_count"] == 1
    assert "fp_a" in out["data"]["per_fp"]


def test_no_prompt_cp(tmp_path):
    step = _new_step(cp_map={})
    out = _run(step, _patches(tmp_path))
    assert out["applicable_count"] == 0
    assert out["data"] == {}


# ──────────────────────────── selection + render success ────────────────────────────
def test_rendered_with_provenance_and_best_of_n(tmp_path):
    scalls = {"sel": 0, "bundles": []}
    rcalls = {"img": 0, "prompts": []}
    step = _new_step(cp_map=_prompt_cp(_fp_targets("fp_a", "fp_b")))
    step._select_override = _select_ok(scalls)
    step._render_override = _render_ok(rcalls)
    out = _run(step, _patches(tmp_path, best_of_n=2))
    per_fp = out["data"]["per_fp"]
    assert set(per_fp) == {"fp_a", "fp_b"}
    for fid in ("fp_a", "fp_b"):
        e = per_fp[fid]
        assert e["status"] == "rendered"
        assert e["light_png_relative_path"]
        assert f"{fid}_1.png" in e["light_png_relative_path"]
        assert len(e["best_of_n_paths"]) == 2  # both draws persisted
        assert e["essential_numbers"] == [16, 23, 30]
        assert e["skeleton_numbers"] == [1, 2, 7]
        assert e["validation_mode"] == "selection_plus_render_visual_gate"
        assert e["prompt_version"] == "2"
    # one selection per fp, best_of_n image draws per fp.
    assert scalls["sel"] == 2
    assert rcalls["img"] == 4
    assert out["data"]["rendered_count"] == 2
    assert out["data"]["fallback_count"] == 0
    assert out["data"]["selection_attempt_count"] == 2
    assert out["data"]["image_api_call_count"] == 4
    assert out["data"]["best_of_n"] == 2
    assert out["completed_count"] == 1
    assert out["failed_count"] == 0
    # the selection bundle carried the camera frequency signal.
    assert '"camera_use_count": 2' in scalls["bundles"][0]["user"]


def test_only_ok_status_fps_processed(tmp_path):
    targets = _fp_targets("fp_a")
    targets["fp_bad"] = {"status": "failed", "numbered_elements": _elements()}
    step = _new_step(cp_map=_prompt_cp(targets))
    step._select_override = _select_ok({"sel": 0, "bundles": []})
    step._render_override = _render_ok({"img": 0, "prompts": []})
    out = _run(step, _patches(tmp_path))
    assert set(out["data"]["per_fp"]) == {"fp_a"}


# ──────────────────────────── fallback paths (never a step failure) ────────────────────────────
def test_selection_failed_falls_back(tmp_path):
    step = _new_step(cp_map=_prompt_cp(_fp_targets("fp_a")))
    step._select_override = _select_none({"sel": 0})
    rcalls = {"img": 0, "prompts": []}
    step._render_override = _render_ok(rcalls)
    out = _run(step, _patches(tmp_path))
    e = out["data"]["per_fp"]["fp_a"]
    assert e["status"] == "fallback"
    assert e["fallback_reason"] == "selection_failed"
    assert e["light_png_relative_path"] is None
    assert rcalls["img"] == 0  # never reached render
    assert out["failed_count"] == 0
    assert out["data"]["fallback_count"] == 1


def test_selection_invalid_falls_back(tmp_path):
    step = _new_step(cp_map=_prompt_cp(_fp_targets("fp_a")))
    step._select_override = _select_invalid({"sel": 0})
    rcalls = {"img": 0, "prompts": []}
    step._render_override = _render_ok(rcalls)
    out = _run(step, _patches(tmp_path))
    e = out["data"]["per_fp"]["fp_a"]
    assert e["status"] == "fallback"
    assert e["fallback_reason"] == "selection_invalid"
    assert any("not drawn" in d for d in e["diagnostics"])
    assert rcalls["img"] == 0  # invalid selection never renders
    assert out["failed_count"] == 0


def test_render_failure_falls_back(tmp_path):
    step = _new_step(cp_map=_prompt_cp(_fp_targets("fp_a")))
    step._select_override = _select_ok({"sel": 0, "bundles": []})
    rcalls = {"img": 0, "prompts": []}
    step._render_override = _render_fail(rcalls)
    out = _run(step, _patches(tmp_path, best_of_n=2))
    e = out["data"]["per_fp"]["fp_a"]
    assert e["status"] == "fallback"
    assert e["fallback_reason"] == "render_failed"
    assert e["light_png_relative_path"] is None
    # render miss is a per-fp detailed fallback, NOT a step failure.
    assert out["failed_count"] == 0
    assert out["completed_count"] == 1
    assert out["data"]["render_failed_count"] == 1
    assert rcalls["img"] == 2  # both draws attempted before fallback


def test_no_numbered_elements_falls_back(tmp_path):
    targets = _fp_targets("fp_a", elements=[])
    step = _new_step(cp_map=_prompt_cp(targets))
    scalls = {"sel": 0}
    step._select_override = _select_none(scalls)
    step._render_override = _render_ok({"img": 0, "prompts": []})
    out = _run(step, _patches(tmp_path))
    e = out["data"]["per_fp"]["fp_a"]
    assert e["status"] == "fallback"
    assert e["fallback_reason"] == "no_numbered_elements"
    assert scalls["sel"] == 0  # never reached selection
    assert out["failed_count"] == 0


def test_unsafe_fp_id_falls_back(tmp_path):
    targets = _fp_targets("fp_a")
    targets["../evil"] = {
        "status": "ok", "numbered_elements": _elements(),
        "camera_recommendations": _cameras(),
    }
    step = _new_step(cp_map=_prompt_cp(targets))
    step._select_override = _select_ok({"sel": 0, "bundles": []})
    step._render_override = _render_ok({"img": 0, "prompts": []})
    out = _run(step, _patches(tmp_path))
    e = out["data"]["per_fp"]["../evil"]
    assert e["status"] == "fallback"
    assert e["fallback_reason"] == "unsafe_fp_id"
    assert out["failed_count"] == 0


# ──────────────────────────── config_hash ────────────────────────────
def _config_hash_under(step, patches):
    for p in patches:
        p.start()
    try:
        return step._config_hash()
    finally:
        for p in reversed(patches):
            p.stop()


def test_config_hash_changes_with_flag(tmp_path):
    step = _new_step(cp_map={})
    h_off = _config_hash_under(step, _patches(tmp_path, enabled=False))
    h_on = _config_hash_under(step, _patches(tmp_path, enabled=True))
    assert h_off != h_on
