110 lines
4.3 KiB
Python
110 lines
4.3 KiB
Python
"""distractor_selector 区间选择纯逻辑与编排(mock VLM)测试。"""
|
||
|
||
import pytest
|
||
|
||
from app.question_gen.distractor_selector import (
|
||
SelectorConfig,
|
||
_select_in_interval,
|
||
build_grounded_options,
|
||
)
|
||
from core.types import LLMResponse
|
||
|
||
|
||
def test_select_three_in_interval_by_highest_score():
|
||
# correct=0.90, 区间 = [0.90-0.35, 0.90-0.05] = [0.55, 0.85]
|
||
cands = ["a", "b", "c", "d", "e"]
|
||
scores = [0.84, 0.70, 0.60, 0.50, 0.88] # e=0.88 太接近(>0.85)剔除, d=0.50 太低剔除
|
||
chosen = _select_in_interval(0.90, cands, scores, delta_low=0.05, delta_high=0.35)
|
||
assert chosen == ["a", "b", "c"] # 落区间的按分数降序取 3(最难)
|
||
|
||
|
||
def test_select_returns_none_when_fewer_than_three():
|
||
cands = ["a", "b"]
|
||
scores = [0.80, 0.70]
|
||
assert _select_in_interval(0.90, cands, scores, 0.05, 0.35) is None
|
||
|
||
|
||
def test_select_excludes_out_of_band():
|
||
cands = ["hi", "lo", "ok1", "ok2", "ok3"]
|
||
scores = [0.89, 0.10, 0.80, 0.75, 0.70] # hi>上界, lo<下界
|
||
chosen = _select_in_interval(0.90, cands, scores, 0.05, 0.35)
|
||
assert chosen == ["ok1", "ok2", "ok3"]
|
||
|
||
|
||
class _FakeVLM:
|
||
"""按队列返回预设响应的 mock VLM。"""
|
||
|
||
def __init__(self, responses: list[str]):
|
||
self._responses = list(responses)
|
||
self.calls = 0
|
||
|
||
async def chat_with_images(self, messages, images, *, session_id=None, parent_call_id=None):
|
||
self.calls += 1
|
||
content = self._responses.pop(0)
|
||
return LLMResponse(
|
||
content=content, thinking="", model="fake", provider="fake",
|
||
prompt_tokens=0, completion_tokens=0, latency_ms=0,
|
||
ttft_ms=None, max_inter_token_ms=None, cache_hit=False, call_id="c",
|
||
)
|
||
|
||
|
||
class _Material:
|
||
subtitle_sentences = ["厨师先炒后蒸"]
|
||
frame_paths = ["/f1.jpg", "/f2.jpg"]
|
||
cross_l2_texts: list = []
|
||
source_nodes = ("n1",)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_generate_pool_degrades_on_malformed_response():
|
||
from app.question_gen.distractor_selector import _generate_pool
|
||
vlm = _FakeVLM(['["not", "a", "dict"]']) # 顶层是 list 不是 dict
|
||
out = await _generate_pool(vlm, "?", "蒸", _Material(), 4, session_id="s")
|
||
assert out == []
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_score_options_raises_valueerror_on_non_numeric():
|
||
from app.question_gen.distractor_selector import _score_options
|
||
vlm = _FakeVLM(['{"scores": [0.9, null, 0.5]}']) # null 不是数值
|
||
with pytest.raises(ValueError):
|
||
await _score_options(vlm, "?", ["蒸", "炒", "煮"], _Material(), session_id="s")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_build_grounded_options_happy_path():
|
||
pool = '{"distractors": ["炒", "煮", "炸", "烤"]}'
|
||
scores = '{"scores": [0.90, 0.80, 0.70, 0.60, 0.20]}' # 正解0.90; 炒0.80 煮0.70 炸0.60 落区间, 烤0.20 剔除
|
||
vlm = _FakeVLM([pool, scores])
|
||
cfg = SelectorConfig(candidate_pool_size=4, delta_low=0.05, delta_high=0.35)
|
||
out = await build_grounded_options(
|
||
vlm=vlm, question="厨师最终用哪种方式?", correct_text="蒸",
|
||
material=_Material(), config=cfg, session_id="s",
|
||
)
|
||
assert out.hard_fail is False
|
||
assert out.answer == "A"
|
||
assert out.options[0] == "A. 蒸"
|
||
assert {o[3:] for o in out.options[1:]} == {"炒", "煮", "炸"}
|
||
assert out.observation["hard_fail"] is False
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_build_grounded_options_hard_fail_keeps_observation():
|
||
# 所有候选都在负空间(分数极低),退火后仍不足 3 个 → hard_fail。
|
||
# VLM 只被调 2 次(首轮 pool+score)+ 1 次退火 pool + 1 次退火 score = 4 次;
|
||
# δ_high 放宽轮次是纯重选,不调 VLM。退火 pool 打分含正解,共 4 个分数。
|
||
pool = '{"distractors": ["x", "y", "z"]}'
|
||
scores = '{"scores": [0.90, 0.05, 0.04, 0.03]}'
|
||
pool2 = '{"distractors": ["p", "q", "r"]}'
|
||
scores2 = '{"scores": [0.90, 0.05, 0.04, 0.03]}'
|
||
vlm = _FakeVLM([pool, scores, pool2, scores2])
|
||
cfg = SelectorConfig(candidate_pool_size=3, delta_low=0.05, delta_high=0.35)
|
||
out = await build_grounded_options(
|
||
vlm=vlm, question="?", correct_text="蒸",
|
||
material=_Material(), config=cfg, session_id="s",
|
||
)
|
||
assert out.hard_fail is True
|
||
assert out.options is None
|
||
assert out.observation["hard_fail"] is True
|
||
assert out.observation["pool_size"] == 6 # 首轮 3 + 退火追加 3
|