Files
iomgaa a7ca6d15ed feat(harness): InferenceDepsRouter per-video 路由器
按 video_id 懒加载 InferenceDeps 并缓存,路由 dispatch/prompt_builder:
- create_dispatch: 按 session_id 路由到对应视频的工具调度
- create_prompt_builder: 自动注册 qid→vid 映射并路由 prompt 构建
- 三元组 (video_id, skills_dir, prompts_dir) 缓存键

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
2026-07-09 12:20:32 -04:00

183 lines
5.9 KiB
Python

"""InferenceDepsRouter 单元测试。"""
from __future__ import annotations
from pathlib import Path
from unittest.mock import AsyncMock, MagicMock
import pytest
from app.harness.deps_router import InferenceDepsRouter
from core.types import GeneratedQuestion
# ---------------------------------------------------------------------------
# 测试辅助
# ---------------------------------------------------------------------------
def _make_router(
store_dir: Path | None = None,
) -> InferenceDepsRouter:
"""构造一个 InferenceDepsRouter 实例,所有外部依赖用 MagicMock。"""
return InferenceDepsRouter(
store_dir=store_dir or Path("/fake/store"),
embed_provider=MagicMock(),
llm=MagicMock(),
vlm=MagicMock(),
ocr=None,
default_prompts_dir=Path("/fake/prompts"),
default_skills_dir=Path("/fake/skills"),
skill_mode="none",
verify_vision=False,
anchor=False,
assemble_mode="concat",
)
def _make_qa(
question_id: str = "q1",
video_id: str = "vid1",
) -> GeneratedQuestion:
"""构造 GeneratedQuestion 测试实例。"""
return GeneratedQuestion(
question_id=question_id,
video_id=video_id,
task_type="Action Reasoning",
question="test?",
options=("A. a", "B. b", "C. c", "D. d"),
answer="A",
source_nodes=(),
difficulty="medium",
)
def _make_fake_deps() -> MagicMock:
"""构造 InferenceDeps 替身。"""
deps = MagicMock()
deps.prompt_builder.return_value = ("system_prompt", "user_prompt")
deps.tool_dispatch_fn = AsyncMock(return_value="tool_result")
return deps
# =========================================================================
# prompt_builder 测试
# =========================================================================
class TestPromptBuilder:
"""create_prompt_builder 返回的闭包测试。"""
def test_registers_qid_to_vid_mapping(self) -> None:
"""prompt_builder 调用时自动注册 question_id → video_id 映射。"""
router = _make_router()
fake_deps = _make_fake_deps()
router._build_deps = MagicMock(return_value=fake_deps)
builder = router.create_prompt_builder()
qa = _make_qa(question_id="q42", video_id="vid99")
builder(qa)
assert router._qid_to_vid["q42"] == "vid99"
def test_returns_prompt_from_deps(self) -> None:
"""prompt_builder 返回 deps.prompt_builder 的结果。"""
router = _make_router()
fake_deps = _make_fake_deps()
router._build_deps = MagicMock(return_value=fake_deps)
builder = router.create_prompt_builder()
qa = _make_qa()
result = builder(qa)
assert result == ("system_prompt", "user_prompt")
fake_deps.prompt_builder.assert_called_once_with(qa)
# =========================================================================
# dispatch 测试
# =========================================================================
class TestDispatch:
"""create_dispatch 返回的闭包测试。"""
@pytest.mark.asyncio
async def test_routes_by_session_id(self) -> None:
"""dispatch 按 session_id 查找 video_id 并路由到对应 deps。"""
router = _make_router()
fake_deps = _make_fake_deps()
router._build_deps = MagicMock(return_value=fake_deps)
# 先通过 prompt_builder 注册映射
builder = router.create_prompt_builder()
qa = _make_qa(question_id="q1", video_id="vid1")
builder(qa)
# 然后 dispatch
dispatch = router.create_dispatch()
result = await dispatch("search", {"query": "test"}, context={"session_id": "q1"})
assert result == "tool_result"
fake_deps.tool_dispatch_fn.assert_called_once_with(
"search", {"query": "test"}, context={"session_id": "q1"}
)
@pytest.mark.asyncio
async def test_unknown_session_id_raises_key_error(self) -> None:
"""未注册的 session_id 抛出 KeyError。"""
router = _make_router()
dispatch = router.create_dispatch()
with pytest.raises(KeyError, match="未注册的 session_id"):
await dispatch("search", {}, context={"session_id": "unknown"})
@pytest.mark.asyncio
async def test_missing_session_id_raises_key_error(self) -> None:
"""context 中缺少 session_id 抛出 KeyError。"""
router = _make_router()
dispatch = router.create_dispatch()
with pytest.raises(KeyError, match="未注册的 session_id"):
await dispatch("search", {}, context={})
# =========================================================================
# 缓存测试
# =========================================================================
class TestCaching:
"""InferenceDeps 缓存行为测试。"""
def test_same_video_reuses_cached_deps(self) -> None:
"""同一 video_id 复用缓存的 InferenceDeps。"""
router = _make_router()
fake_deps = _make_fake_deps()
router._build_deps = MagicMock(return_value=fake_deps)
builder = router.create_prompt_builder()
qa1 = _make_qa(question_id="q1", video_id="vid1")
qa2 = _make_qa(question_id="q2", video_id="vid1")
builder(qa1)
builder(qa2)
# _build_deps 只调用一次(第二次命中缓存)
assert router._build_deps.call_count == 1
def test_clear_cache_empties_both(self) -> None:
"""clear_cache 清空 deps 缓存和 qid→vid 映射表。"""
router = _make_router()
fake_deps = _make_fake_deps()
router._build_deps = MagicMock(return_value=fake_deps)
builder = router.create_prompt_builder()
builder(_make_qa())
assert len(router._deps_cache) == 1
assert len(router._qid_to_vid) == 1
router.clear_cache()
assert len(router._deps_cache) == 0
assert len(router._qid_to_vid) == 0