From a7ca6d15ed907942f00a081b1991f39a2b2976c8 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Thu, 9 Jul 2026 12:20:32 -0400 Subject: [PATCH] =?UTF-8?q?feat(harness):=20InferenceDepsRouter=20per-vide?= =?UTF-8?q?o=20=E8=B7=AF=E7=94=B1=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 按 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) --- app/harness/deps_router.py | 183 +++++++++++++++++++++++++++++++++ tests/unit/test_deps_router.py | 182 ++++++++++++++++++++++++++++++++ 2 files changed, 365 insertions(+) create mode 100644 app/harness/deps_router.py create mode 100644 tests/unit/test_deps_router.py diff --git a/app/harness/deps_router.py b/app/harness/deps_router.py new file mode 100644 index 0000000..c9c2c37 --- /dev/null +++ b/app/harness/deps_router.py @@ -0,0 +1,183 @@ +"""按 video_id 懒加载 InferenceDeps 并路由 dispatch/prompt_builder。 + +每个视频的 TreeIndex、TreeEnvironment、SkillRegistry 等重量级对象 +只在首次访问时构建并缓存,后续同视频的请求直接复用。 +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Any + +from loguru import logger + +from app.harness.factory import InferenceDeps, build_inference_deps + +if TYPE_CHECKING: + from pathlib import Path + + from app.ports import EmbeddingProvider, OCRProvider + from core.protocols import LLMProvider, VLMProvider + from core.types import GeneratedQuestion + + +class InferenceDepsRouter: + """按 video_id 懒加载 InferenceDeps 并路由工具调度和 prompt 构建。 + + 职责: + 1. 维护 question_id → video_id 的映射表(由 prompt_builder 自动注册)。 + 2. 按 (video_id, skills_dir, prompts_dir) 三元组缓存 InferenceDeps。 + 3. 提供 create_dispatch / create_prompt_builder 工厂方法, + 返回的闭包符合 ToolDispatchFn / PromptBuilderFn Protocol。 + + 参数: + store_dir: store 根目录。 + embed_provider: 嵌入端口实例。 + llm: LLM 端口实例。 + vlm: VLM 端口实例。 + ocr: OCR 端口实例(None 不启用)。 + default_prompts_dir: 默认 prompt 文件目录。 + default_skills_dir: 默认 skill 文件目录(None 则不加载 skill)。 + skill_mode: skill 模式("auto"/"manual"/"none")。 + verify_vision: observe_frame 是否执行验证轮。 + anchor: view_node 是否启用行号锚模式。 + assemble_mode: 锚模式装配形态。 + """ + + def __init__( + self, + *, + store_dir: Path, + embed_provider: EmbeddingProvider, + llm: LLMProvider, + vlm: VLMProvider, + ocr: OCRProvider | None, + default_prompts_dir: Path, + default_skills_dir: Path | None, + skill_mode: str, + verify_vision: bool, + anchor: bool, + assemble_mode: str, + ) -> None: + self._store_dir = store_dir + self._embed = embed_provider + self._llm = llm + self._vlm = vlm + self._ocr = ocr + self._default_prompts_dir = default_prompts_dir + self._default_skills_dir = default_skills_dir + self._skill_mode = skill_mode + self._verify_vision = verify_vision + self._anchor = anchor + self._assemble_mode = assemble_mode + self._deps_cache: dict[tuple[str, str, str], InferenceDeps] = {} + self._qid_to_vid: dict[str, str] = {} + + def create_dispatch(self, *, skills_dir: Path | None = None) -> Any: + """创建工具调度闭包,按 session_id 路由到对应视频的 InferenceDeps。 + + 参数: + skills_dir: skill 文件目录覆盖(None 使用默认值)。 + + 返回: + 符合 ToolDispatchFn 签名的 async 闭包。 + """ + effective_skills = skills_dir or self._default_skills_dir + + async def _dispatch( + tool_name: str, args: dict[str, Any], *, context: dict[str, Any] + ) -> str: + """按 session_id 查找视频 → 获取缓存 deps → 委托执行。""" + session_id = context.get("session_id") + if not session_id or session_id not in self._qid_to_vid: + raise KeyError( + f"未注册的 session_id={session_id!r},已注册 {len(self._qid_to_vid)} 条映射" + ) + video_id = self._qid_to_vid[session_id] + deps = self._ensure_deps(video_id, effective_skills, self._default_prompts_dir) + return await deps.tool_dispatch_fn(tool_name, args, context=context) + + return _dispatch + + def create_prompt_builder( + self, + *, + skills_dir: Path | None = None, + prompts_dir: Path | None = None, + ) -> Any: + """创建 prompt 构建闭包,自动注册 qid→vid 映射并路由到对应视频的 deps。 + + 参数: + skills_dir: skill 文件目录覆盖(None 使用默认值)。 + prompts_dir: prompt 文件目录覆盖(None 使用默认值)。 + + 返回: + 符合 PromptBuilderFn 签名的闭包。 + """ + effective_skills = skills_dir or self._default_skills_dir + effective_prompts = prompts_dir or self._default_prompts_dir + + def _builder(qa: GeneratedQuestion) -> tuple[str, str]: + """注册 qid→vid 映射 → 获取缓存 deps → 委托构建 prompt。""" + self._qid_to_vid[qa.question_id] = qa.video_id + deps = self._ensure_deps(qa.video_id, effective_skills, effective_prompts) + return deps.prompt_builder(qa) + + return _builder + + def _ensure_deps( + self, + video_id: str, + skills_dir: Path | None, + prompts_dir: Path, + ) -> InferenceDeps: + """按 (video_id, skills_dir, prompts_dir) 三元组缓存 InferenceDeps。 + + 参数: + video_id: 视频标识。 + skills_dir: skill 文件目录。 + prompts_dir: prompt 文件目录。 + + 返回: + 缓存命中或新建的 InferenceDeps 实例。 + """ + key = (video_id, str(skills_dir), str(prompts_dir)) + if key not in self._deps_cache: + self._deps_cache[key] = self._build_deps(video_id, skills_dir, prompts_dir) + logger.debug("InferenceDeps 已缓存: video_id={}", video_id) + return self._deps_cache[key] + + def _build_deps( + self, + video_id: str, + skills_dir: Path | None, + prompts_dir: Path, + ) -> InferenceDeps: + """调用 build_inference_deps 构建 InferenceDeps 实例。 + + 参数: + video_id: 视频标识。 + skills_dir: skill 文件目录。 + prompts_dir: prompt 文件目录。 + + 返回: + 新建的 InferenceDeps 实例。 + """ + return build_inference_deps( + store_dir=self._store_dir, + video_id=video_id, + prompts_dir=prompts_dir, + skills_dir=skills_dir, + skill_mode=self._skill_mode, + embed_provider=self._embed, + llm=self._llm, + vlm=self._vlm, + ocr=self._ocr, + verify_vision=self._verify_vision, + anchor=self._anchor, + assemble_mode=self._assemble_mode, + ) + + def clear_cache(self) -> None: + """清空 deps 缓存和 qid→vid 映射表。""" + self._deps_cache.clear() + self._qid_to_vid.clear() diff --git a/tests/unit/test_deps_router.py b/tests/unit/test_deps_router.py new file mode 100644 index 0000000..2614f53 --- /dev/null +++ b/tests/unit/test_deps_router.py @@ -0,0 +1,182 @@ +"""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