From 6e46d184b86273fca41ffe67f283799a78863996 Mon Sep 17 00:00:00 2001 From: iomgaa Date: Thu, 9 Jul 2026 05:36:07 -0400 Subject: [PATCH] =?UTF-8?q?feat(harness):=20add=20factory.py=20=E2=80=94?= =?UTF-8?q?=20InferenceDeps=20dataclass=20+=20build=5Finference=5Fdeps?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 组装一次推理所需的全套依赖的工厂函数: - TreeIndex 加载(FileNotFoundError if missing) - TreeEnvironment 构建 - SkillRegistry 按需发现 - SearchToolDispatcher 装配 - PromptManager + prompt_builder 闭包 测试覆盖:正常路径、缺失树文件、skills 注入、frozen 不可变性。 Co-Authored-By: Claude Opus 4.6 (1M context) --- app/harness/factory.py | 155 +++++++++++++++++++++++++++ tests/unit/test_factory.py | 208 +++++++++++++++++++++++++++++++++++++ 2 files changed, 363 insertions(+) create mode 100644 app/harness/factory.py create mode 100644 tests/unit/test_factory.py diff --git a/app/harness/factory.py b/app/harness/factory.py new file mode 100644 index 0000000..cd856f4 --- /dev/null +++ b/app/harness/factory.py @@ -0,0 +1,155 @@ +"""推理依赖工厂 — 组装一次推理所需的全套依赖。 + +将 TreeIndex 加载、TreeEnvironment 构建、SkillRegistry 发现、 +SearchToolDispatcher 装配、PromptManager 初始化等步骤封装为 +单一工厂函数 ``build_inference_deps``,返回不可变的 ``InferenceDeps``。 + +调用方(runner / inference)只需传入配置参数,无需了解内部装配逻辑。 +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import TYPE_CHECKING, Any + +from loguru import logger + +from app.search.prompt import PromptManager +from app.search.skills import discover_skills +from app.search.tools import SearchToolDispatcher +from app.tree.environment import TreeEnvironment +from app.tree.index import TreeIndex + +if TYPE_CHECKING: + from collections.abc import Callable + from pathlib import Path + + from app.ports import EmbeddingProvider, OCRProvider + from core.protocols import LLMProvider, VLMProvider + from core.types import GeneratedQuestion + + +@dataclass(frozen=True) +class InferenceDeps: + """跑一次推理所需的全套依赖(不含 HarnessLog,其生命周期由调用方管理)。 + + 属性: + llm: LLM 端口实例。 + tool_dispatch_fn: SearchToolDispatcher.dispatch 的绑定方法。 + prompt_builder: (GeneratedQuestion) -> (system_prompt, user_prompt)。 + """ + + llm: LLMProvider + tool_dispatch_fn: Callable[..., Any] + prompt_builder: Callable[[GeneratedQuestion], tuple[str, str]] + + +def build_inference_deps( + *, + store_dir: Path, + video_id: str, + prompts_dir: Path, + skills_dir: Path | None, + skill_mode: str, + embed_provider: EmbeddingProvider, + llm: LLMProvider, + vlm: VLMProvider, + ocr: OCRProvider | None, + verify_vision: bool, + anchor: bool, + assemble_mode: str, +) -> InferenceDeps: + """组装一次推理所需的全套依赖。 + + 参数: + store_dir: store 根目录(包含 videos/{video_id}/tree.json)。 + video_id: 视频标识。 + prompts_dir: prompt 文件目录。 + skills_dir: skill 文件目录(None 则不加载 skill)。 + skill_mode: skill 模式("auto"/"manual"/"none")。 + embed_provider: 嵌入端口实例。 + llm: LLM 端口实例。 + vlm: VLM 端口实例。 + ocr: OCR 端口实例(None 不启用)。 + verify_vision: observe_frame 是否执行验证轮。 + anchor: view_node 是否启用行号锚模式。 + assemble_mode: 锚模式装配形态。 + + 返回: + InferenceDeps 实例。 + + 异常: + FileNotFoundError: tree.json 不存在。 + """ + # Phase 1: 加载 TreeIndex + tree_path = store_dir / "videos" / video_id / "tree.json" + if not tree_path.exists(): + raise FileNotFoundError(f"树索引文件不存在: {tree_path}") + tree_index = TreeIndex.load_json(str(tree_path)) + logger.info("已加载 TreeIndex: video_id={}, L1 节点数={}", video_id, len(tree_index.roots)) + + # Phase 2: 构建 TreeEnvironment + frames_dir = store_dir / "videos" / video_id / "frames" + env = TreeEnvironment(index=tree_index, frames_dir=frames_dir) + + # Phase 3: 构建 SkillRegistry + skills = None + always_skills_text = "" + task_skill_map: dict[str, str] = {} + catalog_text = "" + if skills_dir is not None: + always_skills_text, task_skill_map, catalog_text, skills = discover_skills(skills_dir) + logger.info( + "已发现 skills: always={} 字符, task_map={} 项", + len(always_skills_text), + len(task_skill_map), + ) + + # Phase 4: 构建 SearchToolDispatcher + dispatcher = SearchToolDispatcher( + env, + tool_llm=llm, + vlm=vlm, + ocr=ocr, + prompts_dir=prompts_dir, + skills=skills, + embed_fn=embed_provider.embed, + verify_vision=verify_vision, + anchor=anchor, + assemble_mode=assemble_mode, + ) + + # Phase 5: 构建 PromptManager + _prompt_builder 闭包 + pm = PromptManager(prompts_dir) + l1_ids = [root.id for root in tree_index.roots] + + def _prompt_builder(qa: GeneratedQuestion) -> tuple[str, str]: + """为单条题目生成 (system_prompt, user_prompt)。 + + 参数: + qa: 生成的题目实例。 + + 返回: + (system_prompt, user_prompt) 二元组。 + """ + system = pm.build_inference_prompt( + skill_mode, + qa.task_type, + always_skills_text, + task_skill_map, + catalog_text, + ) + user = pm.format_user_prompt( + qa.question, + list(qa.options), + l1_ids, + qa.task_type, + ) + return system, user + + logger.info("InferenceDeps 组装完成: video_id={}, skill_mode={}", video_id, skill_mode) + return InferenceDeps( + llm=llm, + tool_dispatch_fn=dispatcher.dispatch, + prompt_builder=_prompt_builder, + ) diff --git a/tests/unit/test_factory.py b/tests/unit/test_factory.py new file mode 100644 index 0000000..af18b11 --- /dev/null +++ b/tests/unit/test_factory.py @@ -0,0 +1,208 @@ +"""app/harness/factory.py 的单元测试。 + +验证 build_inference_deps 的返回类型、字段连接、以及错误路径。 +""" + +from __future__ import annotations + +import json +from typing import TYPE_CHECKING +from unittest.mock import AsyncMock, MagicMock + +if TYPE_CHECKING: + from pathlib import Path + +import numpy as np +import pytest + +from app.harness.factory import InferenceDeps, build_inference_deps +from core.types import GeneratedQuestion + + +class TestBuildInferenceDeps: + """build_inference_deps 工厂函数测试。""" + + def test_returns_inference_deps(self, tmp_path: Path) -> None: + """用 fake adapters 验证返回类型和字段非 None。""" + # 准备一棵最小树 + vid_dir = tmp_path / "videos" / "test_vid" + vid_dir.mkdir(parents=True) + (vid_dir / "frames").mkdir() + minimal_tree = { + "metadata": {"source_path": "test", "modality": "video"}, + "roots": [ + { + "id": "L1_000", + "card": { + "scene_summary": "s", + "main_setting": "s", + "key_entities": [], + "main_actions": [], + "topic_keywords": [], + "visible_text": [], + "temporal_flow": "s", + }, + "time_range": [0, 10], + "children": [], + } + ], + } + (vid_dir / "tree.json").write_text(json.dumps(minimal_tree)) + + # prompts + prompts_dir = tmp_path / "prompts" + prompts_dir.mkdir() + (prompts_dir / "system.md").write_text("You are a search agent.") + + fake_llm = AsyncMock() + fake_vlm = AsyncMock() + fake_embed = MagicMock() + fake_embed.dim = 4 + fake_embed.embed = lambda t: np.zeros((1, 4), dtype=np.float32) + + deps = build_inference_deps( + store_dir=tmp_path, + video_id="test_vid", + prompts_dir=prompts_dir, + skills_dir=None, + skill_mode="none", + embed_provider=fake_embed, + llm=fake_llm, + vlm=fake_vlm, + ocr=None, + verify_vision=False, + anchor=False, + assemble_mode="ids", + ) + assert isinstance(deps, InferenceDeps) + assert deps.llm is fake_llm + assert callable(deps.tool_dispatch_fn) + assert callable(deps.prompt_builder) + + # 验证 prompt_builder 实际可用(连接正确) + fake_q = GeneratedQuestion( + question_id="q1", + video_id="test_vid", + task_type="Object Recognition", + question="What?", + options=("A. X", "B. Y", "C. Z", "D. W"), + answer="A", + source_nodes=(), + difficulty="medium", + ) + system, user = deps.prompt_builder(fake_q) + assert isinstance(system, str) and len(system) > 0 + assert isinstance(user, str) and "What?" in user + + def test_missing_tree_raises(self, tmp_path: Path) -> None: + """tree.json 不存在时应抛出 FileNotFoundError。""" + prompts_dir = tmp_path / "prompts" + prompts_dir.mkdir() + (prompts_dir / "system.md").write_text("x") + vid_dir = tmp_path / "videos" / "nonexist" + vid_dir.mkdir(parents=True) + + with pytest.raises(FileNotFoundError): + build_inference_deps( + store_dir=tmp_path, + video_id="nonexist", + prompts_dir=prompts_dir, + skills_dir=None, + skill_mode="none", + embed_provider=MagicMock(), + llm=AsyncMock(), + vlm=AsyncMock(), + ocr=None, + verify_vision=False, + anchor=False, + assemble_mode="ids", + ) + + def test_with_skills_dir(self, tmp_path: Path) -> None: + """提供 skills_dir 时 skill 信息应正确注入到 prompt_builder 输出。""" + # 准备树 + vid_dir = tmp_path / "videos" / "vid1" + vid_dir.mkdir(parents=True) + (vid_dir / "frames").mkdir() + minimal_tree = { + "metadata": {"source_path": "test", "modality": "video"}, + "roots": [ + { + "id": "L1_000", + "card": { + "scene_summary": "test scene", + "main_setting": "indoor", + "key_entities": [], + "main_actions": [], + "topic_keywords": [], + "visible_text": [], + "temporal_flow": "linear", + }, + "time_range": [0, 5], + "children": [], + } + ], + } + (vid_dir / "tree.json").write_text(json.dumps(minimal_tree)) + + # prompts + prompts_dir = tmp_path / "prompts" + prompts_dir.mkdir() + (prompts_dir / "system.md").write_text("Base system prompt.") + + # skills + skills_dir = tmp_path / "skills" + skills_dir.mkdir() + (skills_dir / "always_nav.md").write_text( + "---\nname: always_nav\nalways: true\n---\nAlways navigate broadly." + ) + (skills_dir / "action_skill.md").write_text( + "---\nname: action_skill\ntask_type: Action Reasoning\n---\nFocus on actions." + ) + + fake_llm = AsyncMock() + fake_vlm = AsyncMock() + fake_embed = MagicMock() + fake_embed.dim = 4 + fake_embed.embed = lambda t: np.zeros((1, 4), dtype=np.float32) + + deps = build_inference_deps( + store_dir=tmp_path, + video_id="vid1", + prompts_dir=prompts_dir, + skills_dir=skills_dir, + skill_mode="auto", + embed_provider=fake_embed, + llm=fake_llm, + vlm=fake_vlm, + ocr=None, + verify_vision=False, + anchor=False, + assemble_mode="ids", + ) + + fake_q = GeneratedQuestion( + question_id="q2", + video_id="vid1", + task_type="Action Reasoning", + question="What happened?", + options=("A. X", "B. Y", "C. Z", "D. W"), + answer="B", + source_nodes=(), + difficulty="easy", + ) + system, user = deps.prompt_builder(fake_q) + # always skill 文本和 task_type skill 文本应出现在 system prompt 中 + assert "Always navigate broadly" in system + assert "Focus on actions" in system + assert "What happened?" in user + + def test_frozen_dataclass(self) -> None: + """InferenceDeps 是 frozen dataclass,不可修改属性。""" + deps = InferenceDeps( + llm=AsyncMock(), + tool_dispatch_fn=lambda: None, + prompt_builder=lambda q: ("", ""), + ) + with pytest.raises(AttributeError): + deps.llm = AsyncMock() # type: ignore[misc]