"""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