feat(harness): add factory.py — InferenceDeps dataclass + build_inference_deps
组装一次推理所需的全套依赖的工厂函数: - TreeIndex 加载(FileNotFoundError if missing) - TreeEnvironment 构建 - SkillRegistry 按需发现 - SearchToolDispatcher 装配 - PromptManager + prompt_builder 闭包 测试覆盖:正常路径、缺失树文件、skills 注入、frozen 不可变性。 Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -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,
|
||||||
|
)
|
||||||
@@ -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]
|
||||||
Reference in New Issue
Block a user