"""tests/unit/test_correctness_unit_view.py — correctness 三对象口径单元测试。 验证 unit_correctness_view(逐题 per_q → unit 折叠:AR pair 双向 AND、非 AR single) 及 core.evolution 的 pair_block / compute_accuracy 消费 unit 口径时,混格池中 AR pair 折叠为单元、W/L 与准确率分母不被 P/Q 单题计分污染(核心算法保真 #5)。 """ from __future__ import annotations from app.harness.question_units import build_units, unit_correctness_view from core.evolution.validate import classify_quadrants, compute_accuracy, pair_block from core.types import GeneratedQuestion def _single(qid: str, task_type: str = "temporal") -> GeneratedQuestion: """构造一条非 AR single 题(unit_id 回填为 question_id)。""" return GeneratedQuestion( question_id=qid, video_id=f"v_{qid}", task_type=task_type, question="Q?", options=("A", "B", "C", "D"), answer="A", source_nodes=(), difficulty="easy", ) def _pair(pair_id: str, task_type: str = "temporal") -> tuple[GeneratedQuestion, GeneratedQuestion]: """构造一个 AR 孪生对(original + mirror,共享 pair_id → unit_id=pair_id)。""" common = { "video_id": f"v_{pair_id}", "task_type": task_type, "question": "Q?", "options": ("A", "B", "C", "D"), "answer": "A", "source_nodes": (), "difficulty": "easy", "pair_id": pair_id, "flip_axis": "before_after", } orig = GeneratedQuestion(question_id=f"{pair_id}_o", question_role="pair_original", **common) mirror = GeneratedQuestion(question_id=f"{pair_id}_m", question_role="pair_mirror", **common) return orig, mirror class TestUnitCorrectnessView: """unit_correctness_view:逐题对错折叠成 unit_id → bool。""" def test_mixed_pool_and_semantics(self) -> None: """混格:single 直取、AR pair 双向 AND。""" s0 = _single("s0") s1 = _single("s1") p1o, p1m = _pair("p1") p2o, p2m = _pair("p2") units = build_units([s0, s1, p1o, p1m, p2o, p2m]) per_q = { "s0": True, "s1": False, "p1_o": True, "p1_m": True, "p2_o": True, "p2_m": False, } view = unit_correctness_view(units, per_q) # single 的 unit_id 等于 question_id;pair 的 unit_id 等于 pair_id assert view == {"s0": True, "s1": False, "p1": True, "p2": False} def test_missing_per_q_raises(self) -> None: """任一题缺 per_q → KeyError(禁静默兜底)。""" p1o, p1m = _pair("p1") units = build_units([p1o, p1m]) import pytest with pytest.raises(KeyError): unit_correctness_view(units, {"p1_o": True}) class TestPairBlockUnitFold: """pair_block 消费 unit 口径:混格 W/L 不被 P/Q 单题计分污染。""" def test_pair_partial_improvement_not_counted(self) -> None: """AR pair 基线(F,F)→候选(T,F):单元仍错,不计 W(保真 #5)。""" p1o, p1m = _pair("p1") s0 = _single("s0") units = build_units([p1o, p1m, s0]) unit_ids = [u.unit_id for u in units] baseline_per_q = {"p1_o": False, "p1_m": False, "s0": False} candidate_per_q = {"p1_o": True, "p1_m": False, "s0": True} b_units = unit_correctness_view(units, baseline_per_q) c_units = unit_correctness_view(units, candidate_per_q) result = pair_block(b_units, c_units, unit_ids) # 只有 s0 单元发生 F→T 翻转;pair 单元双向 AND 后仍错,不计 W assert result.w == 1 assert result.l == 0 assert result.observed["p1"] == (False, False) assert result.observed["s0"] == (False, True) def test_pair_full_flip_counts_once(self) -> None: """AR pair 两题齐翻(F,F)→(T,T):单元计 1 次 W(非 2)。""" p1o, p1m = _pair("p1") units = build_units([p1o, p1m]) b_units = unit_correctness_view(units, {"p1_o": False, "p1_m": False}) c_units = unit_correctness_view(units, {"p1_o": True, "p1_m": True}) result = pair_block(b_units, c_units, [u.unit_id for u in units]) assert result.w == 1 assert result.l == 0 class TestComputeAccuracyUnitDenominator: """compute_accuracy 分母按 unit 数(非逐题)。""" def test_denominator_is_unit_count(self) -> None: """1 pair(错) + 1 single(对) → 1/2;逐题会误算 2/3。""" p1o, p1m = _pair("p1") s0 = _single("s0") units = build_units([p1o, p1m, s0]) unit_ids = [u.unit_id for u in units] view = unit_correctness_view(units, {"p1_o": True, "p1_m": False, "s0": True}) assert compute_accuracy(view, unit_ids) == 0.5 class TestClassifyQuadrantsUnitKeys: """classify_quadrants 按 unit_id 分桶(pair 单元只出现一次)。""" def test_pair_unit_single_bucket(self) -> None: """pair 单元 F→F 落 persistent_fails,只记 unit_id 一次。""" observed = {"p1": (False, False), "s0": (False, True)} qc = classify_quadrants(observed) assert qc.improvements == ["s0"] assert qc.persistent_fails == ["p1"]