"""W21B-w5 (D2): BgSpacePartitionStep gate + path tests.

opt-in (default OFF) producer that, per fp_id, builds deterministic candidate
edges from the dossier, adjudicates the borderline ones with the pass-2 edge
judge over the projection-card visible_items, and assembles the
space_partition_plan. Default path: provider OFF → every edge ``skipped``
(provider_disabled) → each bg keeps its own plate, NO real LLM call. A
test-injected provider exercises the real-call accounting + the strong-parent
grouping. Image / DB / ImageAsset write 0.
"""
from __future__ import annotations

from unittest.mock import MagicMock, patch

from app.core.steps.bg_space_partition_step import (
    PROMPT_VERSION,
    SCHEMA_VERSION,
    BgSpacePartitionStep,
)


def _visible(num):
    return {
        "marker_number": num,
        "marker_layer": "base_structural_unit",
        "expected_label": "x",
        "visibility": "visible",
        "horizontal_band": "left",
        "depth_band": "foreground",
        "evidence": "ev",
        "source_ref": "fp",
        "confidence": 0.8,
    }


def _card_entry(bg_id, shot_id, fp_id, *, state="pass", items=(1, 2)):
    return {
        "bg_id": bg_id,
        "shot_id": shot_id,
        "fp_id": fp_id,
        "card_state": state,
        "card": {"vlm_output": {"visible_items": [_visible(n) for n in items]}},
    }


def _cp_map(**extra):
    # fp_x: XB01[1,2] XB02[1,2] (same room) XB03[3] (different room).
    cp = {
        "base_location_dossier": {"data": {"dossiers": {"fp_x": {
            "fp_id": "fp_x",
            "base_marker_inventory": [
                {"number": 1, "base_layer_decision": "base_structural_unit", "label": "a"},
                {"number": 2, "base_layer_decision": "base_structural_unit", "label": "b"},
                {"number": 3, "base_layer_decision": "base_structural_unit", "label": "c"},
                {"number": 9, "base_layer_decision": "base_persistent_furniture", "label": "sofa"},
            ],
            "per_bg_render_facts_by_bg_id": {
                "XB01": {"bg_id": "XB01", "fp_id": "fp_x", "target_unit_marker_numbers": [1, 2]},
                "XB02": {"bg_id": "XB02", "fp_id": "fp_x", "target_unit_marker_numbers": [1, 2]},
                "XB03": {"bg_id": "XB03", "fp_id": "fp_x", "target_unit_marker_numbers": [3]},
            },
        }}}},
        "floor_plan_geometry_readback": {"data": {"per_fp": {"fp_x": {"ok": 1}}}},
        "shot_projection_card": {"data": {"cards": {
            "XB01::S1": _card_entry("XB01", "S1", "fp_x"),
            "XB02::S1": _card_entry("XB02", "S1", "fp_x"),
            "XB03::S1": _card_entry("XB03", "S1", "fp_x", items=(3,)),
        }}},
    }
    cp.update(extra)
    return cp


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


def _patches(tmp_path, *, enabled=True, dossier=True, real_provider=False, mode="on",
             edge_cap=40, llm_cap=120):
    return [
        patch("app.core.config.settings.projects_dir", str(tmp_path)),
        patch("app.core.config.settings.background_mode", mode),
        patch("app.core.config.settings.base_location_dossier_enabled", dossier),
        patch("app.core.config.settings.bg_space_partition_enabled", enabled),
        patch("app.core.config.settings.bg_space_partition_real_provider_enabled", real_provider),
        patch("app.core.config.settings.bg_space_partition_edge_cap_per_fp", edge_cap),
        patch("app.core.config.settings.bg_space_partition_llm_call_cap", llm_cap),
    ]


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


def _same_space_provider(calls):
    def provider(*, prompt_bundle):
        calls["n"] += 1
        return {
            "edge_state": "same_space",
            "confidence": 0.9,
            "evidence": "shared sink + table visible in both",
            "shared_distinctive_features": ["sink", "table"],
        }
    return provider


# ──────────────────────────── constants / gates ────────────────────────────


def test_module_constants_locked():
    assert SCHEMA_VERSION == 1
    assert PROMPT_VERSION == "1"


def test_default_off_returns_not_applicable(tmp_path):
    step = _new_step(cp_map=_cp_map())
    result = _run(step, _patches(tmp_path, enabled=False))
    assert result["applicable_count"] == 0
    assert result["data"] == {}


def test_not_applicable_when_background_mode_off(tmp_path):
    step = _new_step(cp_map=_cp_map())
    result = _run(step, _patches(tmp_path, mode="off"))
    assert result["applicable_count"] == 0


def test_not_applicable_when_dossier_disabled(tmp_path):
    step = _new_step(cp_map=_cp_map())
    result = _run(step, _patches(tmp_path, dossier=False))
    assert result["applicable_count"] == 0


def test_not_applicable_when_no_dossiers(tmp_path):
    cp = _cp_map(base_location_dossier={"data": {"dossiers": {}}})
    step = _new_step(cp_map=cp)
    result = _run(step, _patches(tmp_path))
    assert result["applicable_count"] == 0


# ──────────────────── provider OFF (default, inert) path ────────────────────


def test_provider_off_skips_all_edges_each_bg_own_plate(tmp_path):
    step = _new_step(cp_map=_cp_map())
    result = _run(step, _patches(tmp_path))  # real_provider False (default)
    assert result["applicable_count"] == 1
    assert result["data"]["llm_call_count"] == 0
    assert result["data"]["image_api_call_count"] == 0
    fp = result["data"]["per_fp"]["fp_x"]
    assert fp["status"] == "ok"
    # one candidate (XB01~XB02 same_space) — but provider disabled → skipped.
    assert all(
        j["judge_status"] == "skipped" and j["skip_reason"] == "provider_disabled"
        for j in fp["edge_judgements"]
    )
    # no strong edge → every bg its own render_new plate.
    actions = fp["space_partition_plan"]["render_actions"]
    assert {bg: a["render_action"] for bg, a in actions.items()} == {
        "XB01": "render_new_plate",
        "XB02": "render_new_plate",
        "XB03": "render_new_plate",
    }


def test_geometry_missing_blocks_fp(tmp_path):
    cp = _cp_map(floor_plan_geometry_readback={"data": {"per_fp": {}}})
    step = _new_step(cp_map=cp)
    result = _run(step, _patches(tmp_path))
    fp = result["data"]["per_fp"]["fp_x"]
    assert fp["status"] == "blocked"
    assert fp["fallback_reason"] == "geometry_missing"
    assert result["data"]["llm_call_count"] == 0


# ──────────────────── injected provider (real accounting) ───────────────────


def test_injected_provider_groups_same_space(tmp_path):
    step = _new_step(cp_map=_cp_map())
    calls = {"n": 0}
    step.set_edge_judge_provider_for_testing(_same_space_provider(calls))
    result = _run(step, _patches(tmp_path, real_provider=True))
    # one candidate edge (XB01~XB02) judged.
    assert calls["n"] == 1
    assert result["data"]["llm_call_count"] == 1
    fp = result["data"]["per_fp"]["fp_x"]
    plan = fp["space_partition_plan"]
    actions = {bg: a["render_action"] for bg, a in plan["render_actions"].items()}
    # XB01 anchor (render_new), XB02 reuses it, XB03 own plate.
    assert actions["XB01"] == "render_new_plate"
    assert actions["XB02"] == "reuse_existing_plate"
    assert plan["render_actions"]["XB02"]["reuse_target_bg_id"] == "XB01"
    assert actions["XB03"] == "render_new_plate"
    assert plan["ref_tree_parents"]["XB02"] == ["XB01"]
    # cross-zone XB03 never inherits a parent.
    assert plan["ref_tree_parents"]["XB03"] == []


def test_r3_unjudgeable_card_skips_edge_no_provider_call(tmp_path):
    # XB02 card is needs_review (not pass) → edge XB01~XB02 not judgeable.
    cp = _cp_map()
    cp["shot_projection_card"]["data"]["cards"]["XB02::S1"]["card_state"] = "needs_review"
    step = _new_step(cp_map=cp)
    calls = {"n": 0}
    step.set_edge_judge_provider_for_testing(_same_space_provider(calls))
    result = _run(step, _patches(tmp_path, real_provider=True))
    assert calls["n"] == 0  # R3 precondition stops the call
    fp = result["data"]["per_fp"]["fp_x"]
    assert any(
        j["skip_reason"] == "card_not_judgeable" for j in fp["edge_judgements"]
    )
    # no strong edge → XB02 falls back to its own plate.
    assert fp["space_partition_plan"]["render_actions"]["XB02"]["render_action"] == (
        "render_new_plate"
    )


def test_edge_cap_per_fp_skips_remaining(tmp_path):
    step = _new_step(cp_map=_cp_map())
    calls = {"n": 0}
    step.set_edge_judge_provider_for_testing(_same_space_provider(calls))
    result = _run(step, _patches(tmp_path, real_provider=True, edge_cap=0))
    assert calls["n"] == 0
    fp = result["data"]["per_fp"]["fp_x"]
    assert any(
        j["skip_reason"] == "edge_cap_per_fp" for j in fp["edge_judgements"]
    )


def test_edge_cap_counts_judged_not_candidate_ordinal(tmp_path):
    # Codex Required: the per-fp cap bounds judge SPEND, not candidate position.
    # Early unjudgeable edges must NOT consume the cap — with cap=1 a later
    # judgeable edge still gets its one call. fp_x2: XB01[1,2] XB02[1,2]
    # XB04[1,2] XB03[3]. XB01 card is needs_review → edges (XB01,XB02) and
    # (XB01,XB04) skip on R3 first; (XB02,XB04) is the only judgeable edge and
    # must be judged under cap=1.
    cp = _cp_map()
    dossier = cp["base_location_dossier"]["data"]["dossiers"]["fp_x"]
    dossier["per_bg_render_facts_by_bg_id"]["XB04"] = {
        "bg_id": "XB04", "fp_id": "fp_x", "target_unit_marker_numbers": [1, 2],
    }
    cards = cp["shot_projection_card"]["data"]["cards"]
    cards["XB04::S1"] = _card_entry("XB04", "S1", "fp_x")
    cards["XB01::S1"]["card_state"] = "needs_review"  # XB01 unjudgeable

    step = _new_step(cp_map=cp)
    calls = {"n": 0}
    step.set_edge_judge_provider_for_testing(_same_space_provider(calls))
    result = _run(step, _patches(tmp_path, real_provider=True, edge_cap=1))
    # exactly one judge call — the later judgeable (XB02,XB04) edge.
    assert calls["n"] == 1
    fp = result["data"]["per_fp"]["fp_x"]
    judged_pairs = {
        frozenset((j["bg_a"], j["bg_b"]))
        for j in fp["edge_judgements"]
        if j.get("judge_status") == "judged"
    }
    assert judged_pairs == {frozenset(("XB02", "XB04"))}


def test_llm_call_cap_skips_remaining(tmp_path):
    step = _new_step(cp_map=_cp_map())
    calls = {"n": 0}
    step.set_edge_judge_provider_for_testing(_same_space_provider(calls))
    result = _run(step, _patches(tmp_path, real_provider=True, llm_cap=0))
    assert calls["n"] == 0
    fp = result["data"]["per_fp"]["fp_x"]
    assert any(j["skip_reason"] == "llm_call_cap" for j in fp["edge_judgements"])


def test_judge_error_fails_closed_per_edge(tmp_path):
    step = _new_step(cp_map=_cp_map())

    def bad_provider(*, prompt_bundle):
        raise RuntimeError("boom")

    step.set_edge_judge_provider_for_testing(bad_provider)
    result = _run(step, _patches(tmp_path, real_provider=True))
    fp = result["data"]["per_fp"]["fp_x"]
    # the failing call is still counted, the edge is skipped, no group forms.
    assert result["data"]["llm_call_count"] == 1
    assert any(
        j["skip_reason"].startswith("judge_error") for j in fp["edge_judgements"]
    )
    assert fp["space_partition_plan"]["render_actions"]["XB02"]["render_action"] == (
        "render_new_plate"
    )


def test_result_has_provenance(tmp_path):
    step = _new_step(cp_map=_cp_map())
    result = _run(step, _patches(tmp_path))
    assert result["schema_version"] == SCHEMA_VERSION
    assert isinstance(result["config_hash"], str) and result["config_hash"]
