Files
Video-Tree-TRM5/tests/unit/test_distractor_selector.py
T

102 lines
4.0 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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_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