# InferenceEngine(推理引擎) 实施计划 > **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. **Goal:** 落地里程碑 3.1——实现 `agent-runtime-design.md` §2 定义的推理引擎:统一 `chat` / `chat_structured` LLM 入口,含模型降级、重试、Token 裁剪、Prompt 注册表与结构化输出。 **Architecture:** 新增 `src/genesis/inference/` 包(types / exceptions / token / prompt_registry / client / engine);engine 通过注入的 `LLMClient`(httpx 实现,测试用 FakeTransport/FakeClient 零网络)与 `PromptRegistry` 组合;异常不抛给调用方(chat 折叠为 `status`),解析失败返回 `parse_error`+raw_text。 **Tech Stack:** Python 3.11+,httpx(不拆 SDK)、jinja2(PromptRegistry 渲染)、tiktoken(可选,缺失回落 approximate)、dataclasses、pytest。 ## Global Constraints - 项目为中文交流(注释中文、标识符英文),Windows/PowerShell 环境 - **零真实网络**:单测一律注入 `FakeLLMClient` / `httpx.MockTransport`;`DEEPSEEK_API_KEY` 不落代码 - **覆盖率红线**:`pyproject.toml` `fail_under = 99`(当前 71 passed / 100% 全绿),新增每文件需足量分支测试,全量回归不得跌破 99 - 测试命令:`python -m pytest tests/ -v`;全量回归:`python -m pytest -v`(会触发 `--cov` 与 `fail_under` 校验) - 提交消息风格:`feat:` / `test:` / `docs:`(简中文描述) - 每次修改后按项目规则追加 `_AI_USAGE_LOG.md` 记录(范式步骤列:Agent 实现 或 测试验证) - 依赖修改 `pyproject.toml` 后需执行 `pip install -e ".[dev]"` 再跑测试 --- ### Task 1: 数据模型与异常层(包骨架 + types + exceptions) **Files:** - Create: `src/genesis/inference/__init__.py` - Create: `src/genesis/inference/types.py` - Create: `src/genesis/inference/exceptions.py` - Create: `tests/test_inference_types.py` - Create: `tests/test_inference_errors.py` - Modify: `pyproject.toml`(dependencies 增加 `httpx`、`jinja2`) **Interfaces:** - Consumes: 无(独立基础层,不依赖其他模块) - Produces: - `TokenUsage(input_tokens: int = 0, output_tokens: int = 0)` - `ChatMessage(role: Literal["system","user","assistant"], content: str)` - `ChatResult(text, model, prompt_version, usage, duration_ms, status: Literal["ok","fallback","failed"], error: str|None = None)` - `StructuredResult(data: dict, raw_text: str, parse_attempts: int, model, prompt_version, usage, duration_ms, status: Literal["ok","fallback","parse_error","failed"], error: str|None = None)` - `Prompt(name: str, version: str, template: str)` - `LLMError(Exception)` / `LLMNetworkError` / `LLMTimeoutError` / `LLMNotConfiguredError` / `LLMResponseError` - [ ] **Step 1: 写失败测试** `tests/test_inference_types.py`: ```python from genesis.inference.types import ChatMessage, ChatResult, Prompt, StructuredResult, TokenUsage def test_token_usage_defaults(): u = TokenUsage() assert u.input_tokens == 0 and u.output_tokens == 0 def test_chat_message_roles(): assert ChatMessage(role="system", content="x").content == "x" def test_chat_result_defaults(): r = ChatResult( text="t", model="m", prompt_version="v1", usage=TokenUsage(), duration_ms=10, status="ok", ) assert r.status == "ok" and r.error is None def test_structured_result_status_ok(): r = StructuredResult( data={"a": 1}, raw_text='{"a":1}', parse_attempts=1, model="m", prompt_version="v1", usage=TokenUsage(), duration_ms=10, status="ok", ) assert r.data == {"a": 1} and r.raw_text == '{"a":1}' def test_structured_result_status_parse_error(): r = StructuredResult( data={}, raw_text="NOT JSON", parse_attempts=3, model="m", prompt_version="v1", usage=TokenUsage(), duration_ms=50, status="parse_error", error="bad json", ) assert r.status == "parse_error" and r.error == "bad json" def test_prompt_fields(): p = Prompt(name="writer", version="v2", template="章节 {{chapter}}") assert p.name == "writer" and p.version == "v2" ``` `tests/test_inference_errors.py`: ```python import pytest from genesis.inference.exceptions import ( LLMError, LLMNetworkError, LLMNotConfiguredError, LLMResponseError, LLMTimeoutError, ) def test_error_hierarchy(): assert issubclass(LLMNetworkError, LLMError) assert issubclass(LLMTimeoutError, LLMError) assert issubclass(LLMNotConfiguredError, LLMError) assert issubclass(LLMResponseError, LLMError) def test_error_message_roundtrip(): e = LLMTimeoutError("timeout!") assert str(e) == "timeout!" ``` - [ ] **Step 2: 运行确认失败** Run: `python -m pytest tests/test_inference_types.py -v` Expected: FAIL(`ModuleNotFoundError: No module named 'genesis.inference'`) - [ ] **Step 3: 修改 pyproject 依赖** `pyproject.toml` 的 dependencies 增加: ```toml "httpx>=0.28", "jinja2>=3.1", ``` - [ ] **Step 4: 实现** `src/genesis/inference/types.py`: ```python from __future__ import annotations from dataclasses import dataclass from typing import Any, Literal @dataclass class TokenUsage: """一次 LLM 调用的 token 用量(可观测性事件/统计用)""" input_tokens: int = 0 output_tokens: int = 0 @dataclass class ChatMessage: """Chat Completions 消息""" role: Literal["system", "user", "assistant"] content: str @dataclass class ChatResult: """chat() 的返回值""" text: str model: str prompt_version: str usage: TokenUsage duration_ms: int status: Literal["ok", "fallback", "failed"] error: str | None = None @dataclass class StructuredResult: """chat_structured() 的返回值(补丁 1:含 status 字段)""" data: dict raw_text: str parse_attempts: int model: str prompt_version: str usage: TokenUsage duration_ms: int status: Literal["ok", "fallback", "parse_error", "failed"] error: str | None = None @dataclass class Prompt: """Prompt 模板条目(name+version 唯一)""" name: str version: str template: str ``` `src/genesis/inference/exceptions.py`: ```python from __future__ import annotations class LLMError(Exception): """LLM 调用相关的异常基类(api-design §7 映射基底)""" class LLMNetworkError(LLMError): """网络失败 / 5xx 重试耗尽(可重试语义)""" class LLMTimeoutError(LLMError): """LLM 调用超时(api-error: LLM_TIMEOUT 502)""" class LLMNotConfiguredError(LLMError): """Key / 模型缺失(api-error: LLM_NOT_CONFIGURED 503)""" class LLMResponseError(LLMError): """响应结构损坏(JSON 解析失败等)""" ``` `src/genesis/inference/__init__.py`: ```python """Genesis 推理引擎(统一 LLM 调用入口)。""" from .types import ChatMessage, ChatResult, Prompt, StructuredResult, TokenUsage from .exceptions import ( LLMError, LLMNetworkError, LLMNotConfiguredError, LLMResponseError, LLMTimeoutError, ) __all__ = [ "ChatMessage", "ChatResult", "Prompt", "StructuredResult", "TokenUsage", "LLMError", "LLMNetworkError", "LLMNotConfiguredError", "LLMResponseError", "LLMTimeoutError", ] ``` > 注:暂时不在 `__init__` 导入 engine/client 等(避免循环导入),Task 5 完成后再收敛导出。 - [ ] **Step 5: 运行确认通过** Run: `python -m pytest tests/test_inference_types.py tests/test_inference_errors.py -v` Expected: PASS(5 + 2 = 7 passed) - [ ] **Step 6: 提交** ```bash git add pyproject.toml src/genesis/inference/__init__.py src/genesis/inference/types.py src/genesis/inference/exceptions.py tests/test_inference_types.py tests/test_inference_errors.py git commit -m "feat: 推理引擎数据模型与异常层(types/exceptions)及 httpx/jinja2 依赖" ``` --- ### Task 2: Token 估算(token.py,approximate 内置 / tiktoken 可选) **Files:** - Create: `src/genesis/inference/token.py` - Create: `tests/test_inference_token.py` **Interfaces:** - Consumes: 无(不依赖其他模块) - Produces: - `approximate_token_count(text: str) -> int`(4 字符 ≈ 1 token,最少 1) - `_tiktoken_estimator(text: str) -> int | None`(未安装返回 None;可 patch 测试其成功路径) - `make_estimator(backend: str = "tiktoken") -> Callable[[str], int]`(backend="approximate" → 近似;否则优先 tiktoken,缺失回落 approximate) - [ ] **Step 1: 写失败测试** `tests/test_inference_token.py`: ```python import pytest from genesis.inference.token import ( _tiktoken_estimator, approximate_token_count, make_estimator, ) def test_approximate_count_minimum(): assert approximate_token_count("") >= 1 assert approximate_token_count("a") == 1 def test_approximate_count_linear(): # 每 4 字符约 1 token(向上取整) assert approximate_token_count("abcd") == 1 assert approximate_token_count("abcdefgh") == 2 assert approximate_token_count("abcdefghi") == 3 def test_tiktoken_estimator_missing_falls_back(): # 未安装 tiktoken 或不可用时返回 None(由 make_estimator 回落 approximate) r = _tiktoken_estimator("hello") assert r is None or isinstance(r, int) def test_make_estimator_approximate_backend(): est = make_estimator("approximate") assert est("abcd") == 1 def test_make_estimator_default_without_tiktoken(monkeypatch): # 强制模拟 tiktoken 缺失:make_estimator 必须回落 approximate import builtins real_import = builtins.__import__ def fake_import(name, *args, **kwargs): if name == "tiktoken": raise ImportError("no tiktoken") return real_import(name, *args, **kwargs) monkeypatch.setattr(builtins, "__import__", fake_import) est = make_estimator("tiktoken") assert est("abcd") == 1 ``` - [ ] **Step 2: 运行确认失败** Run: `python -m pytest tests/test_inference_token.py -v` Expected: FAIL(`ModuleNotFoundError: No module named 'genesis.inference.token'`) - [ ] **Step 3: 实现** `src/genesis/inference/token.py`: ```python from __future__ import annotations from typing import Callable def approximate_token_count(text: str) -> int: """内置近似估算:每 4 字符 ≈ 1 token(无外部依赖,可离线)。""" return max(1, (len(text) + 3) // 4) def _tiktoken_estimator(text: str) -> int | None: """tiktoken 编码估算;tiktoken 未安装时返回 None。""" try: import tiktoken except ImportError: return None try: enc = tiktoken.get_encoding("cl100k_base") return len(enc.encode(text)) except Exception: return None def make_estimator(backend: str = "tiktoken") -> Callable[[str], int]: """按配置选择估算器:backend="tiktoken"(默认)优先 tiktoken, 缺失或异常回落内置 approximate;backend="approximate" 直接用近似。""" if backend == "approximate": return approximate_token_count return lambda text: _tiktoken_estimator(text) or approximate_token_count(text) ``` - [ ] **Step 4: 运行确认通过** Run: `python -m pytest tests/test_inference_token.py -v` Expected: PASS(5 passed) - [ ] **Step 5: 提交** ```bash git add src/genesis/inference/token.py tests/test_inference_token.py git commit -m "feat: Token 估算(approximate 内置 / tiktoken 可选回落)" ``` --- ### Task 3: Prompt 注册表(prompt_registry.py + jinja2 渲染) **Files:** - Create: `src/genesis/inference/prompt_registry.py` - Create: `tests/test_inference_prompt_registry.py` **Interfaces:** - Consumes: `Prompt`(Task 1)、`jinja2` - Produces: - `class PromptRegistry`: - `register(name: str, version: str, template: str) -> None` - `get(name: str, version: str | None = None, variables: dict | None = None) -> str`(version=None → 该 name 最新注册版本;有 variables → 渲染) - `list_versions(name: str) -> list[str]` - `render(template: str, variables: dict) -> str`(jinja2 渲染) - [ ] **Step 1: 写失败测试** `tests/test_inference_prompt_registry.py`: ```python import pytest from genesis.inference.prompt_registry import PromptRegistry from genesis.inference.types import Prompt def test_register_and_render(): reg = PromptRegistry() reg.register("writer", "v1", "按规则撰写:{{ chapter }}") assert reg.get("writer", "v1") == "按规则撰写:{{ chapter }}" def test_get_with_variables_renders(): reg = PromptRegistry() reg.register("writer", "v1", "按规则撰写:{{ chapter }}") assert reg.get("writer", "v1", {"chapter": "帳票設計"}) == "按规则撰写:帳票設計" def test_get_latest_version(): reg = PromptRegistry() reg.register("writer", "v1", "t1") reg.register("writer", "v2", "t2") assert reg.get("writer") == "t2" assert reg.list_versions("writer") == ["v1", "v2"] def test_get_missing_raises_keyerror(): reg = PromptRegistry() with pytest.raises(KeyError): reg.get("dne") def test_render_raw(): reg = PromptRegistry() assert reg.render("{{ a }} と {{ b }}", {"a": "x", "b": 1}) == "x と 1" ``` - [ ] **Step 2: 运行确认失败** Run: `python -m pytest tests/test_inference_prompt_registry.py -v` Expected: FAIL(`ModuleNotFoundError: No module named 'genesis.inference.prompt_registry'`) - [ ] **Step 3: 实现** `src/genesis/inference/prompt_registry.py`: ```python from __future__ import annotations from typing import Any from jinja2 import Template from .types import Prompt class PromptRegistry: """Prompt 模板库:注册/取用/版本管理/渲染(集中管理待迁移 prompts/ 目录)。""" def __init__(self) -> None: self._templates: dict[tuple[str, str], str] = {} def register(self, name: str, version: str, template: str) -> None: """注册(或覆盖)一个版本的模板。""" self._templates[(name, version)] = template def get( self, name: str, version: str | None = None, variables: dict[str, Any] | None = None, ) -> str: """取模板;version=None 返回该 name 最新注册版本;variables 非空时渲染。""" if version is None: versions = self.list_versions(name) if not versions: raise KeyError(f"prompt not found: {name}") version = versions[-1] key = (name, version) if key not in self._templates: raise KeyError(f"prompt version not found: {name}@{version}") template = self._templates[key] if variables: return self.render(template, variables) return template def list_versions(self, name: str) -> list[str]: """返回某 name 的已注册版本(按注册顺序)。""" return [v for (n, v) in self._templates if n == name] def render(self, template: str, variables: dict[str, Any]) -> str: """用 jinja2 渲染模板。""" from jinja2 import Template return Template(template).render(**variables) ``` - [ ] **Step 4: 运行确认通过** Run: `python -m pytest tests/test_inference_prompt_registry.py -v` Expected: PASS(5 passed) - [ ] **Step 5: 提交** ```bash git add src/genesis/inference/prompt_registry.py tests/test_inference_prompt_registry.py git commit -m "feat: PromptRegistry 注册/版本/渲染(jinja2)" ``` --- ### Task 4: LLM 客户端(client.py:httpx + 重试/退避/超时/认证) **Files:** - Create: `src/genesis/inference/client.py` - Create: `tests/test_inference_client.py` **Interfaces:** - Consumes: `ChatMessage` / `TokenUsage`(Task 1);`LLMError` 子树(Task 1) - Produces: - `class LLMClient(Protocol)`:`chat(*, model, messages: list[ChatMessage], temperature: float, max_tokens: int) -> tuple[str, TokenUsage]` - `class HttpLLMClient`:`__init__(base_url, api_key, timeout_sec=60.0, retry_backoff=(1.0,3.0,7.0), transport=None)`;`chat(...) -> tuple[str, TokenUsage]`;重试 5xx/网络错误,退避间隔指数;超时抛 `LLMTimeoutError`;重试耗尽抛 `LLMNetworkError`;非 2xx 4xx 抛 `LLMNetworkError`;无 api_key 抛 `LLMNotConfiguredError` - [ ] **Step 1: 写失败测试** `tests/test_inference_client.py`: ```python import pytest import httpx from genesis.inference.client import HttpLLMClient, LLMClient from genesis.inference.exceptions import LLMNetworkError, LLMNotConfiguredError, LLMTimeoutError from genesis.inference.types import ChatMessage def make_client(handler, *, api_key="sk-test", retry_backoff=(0.0, 0.0)): return HttpLLMClient( base_url="https://api.test.local", api_key=api_key, timeout_sec=0.1, retry_backoff=retry_backoff, transport=httpx.MockTransport(handler), ) def _ok_handler(request): return httpx.Response(200, json={ "choices": [{"message": {"content": "Hello"}}], "usage": {"prompt_tokens": 10, "completion_tokens": 5}, }) def test_client_implements_protocol(): # 结构性断言:HttpLLMClient.chat 的关键字参数签名与 LLMClient Protocol 一致 import inspect proto_params = set(inspect.signature(LLMClient.chat).parameters) impl_params = set(inspect.signature(HttpLLMClient.chat).parameters) assert proto_params.issubset(impl_params) def test_chat_success(): client = make_client(_ok_handler) text, usage = client.chat( model="deepseek-chat", messages=[ChatMessage(role="user", content="hi")], temperature=0.2, max_tokens=100, ) assert text == "Hello" assert usage.input_tokens == 10 and usage.output_tokens == 5 def test_chat_requires_api_key(): with pytest.raises(LLMNotConfiguredError): HttpLLMClient(base_url="https://x", api_key="") def test_chat_timeout(): def slow(request): raise httpx.ReadTimeout("slow") with pytest.raises(LLMTimeoutError): make_client(slow, retry_backoff=(0, 0)).chat( model="m", messages=[ChatMessage(role="user", content="x")], temperature=0.2, max_tokens=100, ) def test_chat_5xx_retry_then_network_error(): calls = {"n": 0} def handler(request): calls["n"] += 1 return httpx.Response(500, text="boom") with pytest.raises(LLMNetworkError): make_client(handler, retry_backoff=(0, 0)).chat( model="m", messages=[ChatMessage(role="user", content="x")], temperature=0.2, max_tokens=100, ) assert calls["n"] == 3 # 初始 + 2 次退避重试(间隔 0/0.01) def test_chat_4xx_no_retry(): calls = {"n": 0} def handler(request): calls["n"] += 1 return httpx.Response(429, text="rate limit") with pytest.raises(LLMNetworkError): make_client(handler).chat( model="m", messages=[ChatMessage(role="user", content="x")], temperature=0.2, max_tokens=100, ) assert calls["n"] == 1 # 4xx 不重试 ``` - [ ] **Step 2: 运行确认失败** Run: `python -m pytest tests/test_inference_client.py -v` Expected: FAIL(`ModuleNotFoundError: No module named 'genesis.inference.client'`) - [ ] **Step 3: 实现** `src/genesis/inference/client.py`: ```python from __future__ import annotations import time from typing import Protocol, Sequence import httpx from .exceptions import LLMNetworkError, LLMNotConfiguredError, LLMTimeoutError from .types import ChatMessage, TokenUsage class LLMClient(Protocol): """LLM 调用适配器(可注入替换为 Fake)。""" def chat( self, *, model: str, messages: list[ChatMessage], temperature: float, max_tokens: int, ) -> tuple[str, TokenUsage]: ... class HttpLLMClient: """OpenAI Chat Completions 兼容的 httpx 实现;支持重试(指数退避)。""" def __init__( self, *, base_url: str, api_key: str, timeout_sec: float = 60.0, retry_backoff: Sequence[float] = (1.0, 3.0, 7.0), transport: httpx.BaseTransport | None = None, ) -> None: if not api_key: raise LLMNotConfiguredError("LLM API key 未配置(DEEPSEEK_API_KEY / LLM_BASE_URL)") self._base_url = base_url.rstrip("/") self._api_key = api_key self._timeout_sec = timeout_sec self._retry_backoff = retry_backoff self._client = httpx.Client(timeout=timeout_sec, transport=transport) def chat( self, *, model: str, messages: list[ChatMessage], temperature: float, max_tokens: int, ) -> tuple[str, TokenUsage]: url = f"{self._base_url}/v1/chat/completions" payload = { "model": model, "messages": [{"role": m.role, "content": m.content} for m in messages], "temperature": temperature, "max_tokens": max_tokens, } headers = { "Authorization": f"Bearer {self._api_key}", "Content-Type": "application/json", } attempts = 1 + len(self._retry_backoff) last_error: Exception | None = None for attempt in range(attempts): if attempt > 0: time.sleep(self._retry_backoff[attempt - 1]) try: resp = self._client.post(url, json=payload, headers=headers) except httpx.TimeoutException as exc: last_error = exc continue except httpx.HTTPError as exc: last_error = exc continue if resp.status_code >= 500: last_error = LLMNetworkError(f"LLM 5xx: {resp.status_code}") continue if resp.status_code >= 400: raise LLMNetworkError(f"LLM HTTP {resp.status_code}: {resp.text[:200]}") data = resp.json() content = data["choices"][0]["message"]["content"] usage_raw = data.get("usage", {}) usage = TokenUsage( input_tokens=usage_raw.get("prompt_tokens", 0), output_tokens=usage_raw.get("completion_tokens", 0), ) return content, usage if isinstance(last_error, httpx.TimeoutException): raise LLMTimeoutError(f"LLM 超时({self._timeout_sec}s)") from last_error raise LLMNetworkError(f"LLM 调用失败(重试耗尽): {last_error}") from last_error ``` - [ ] **Step 4: 运行确认通过** Run: `python -m pytest tests/test_inference_client.py -v` Expected: PASS(6 passed) - [ ] **Step 5: 全量回归**(新代码暴露到覆盖统计) Run: `python -m pytest -v` Expected: PASS(71 + 全部新用例;`fail_under=99` 通过) - [ ] **Step 6: 提交** ```bash git add src/genesis/inference/client.py tests/test_inference_client.py git commit -m "feat: HttpLLMClient(httpx + 重试退避/超时/鉴权)" ``` --- ### Task 5: InferenceEngine(engine.py:chat / chat_structured 全流程 + 注入测试 + 文档补丁) **Files:** - Create: `src/genesis/inference/engine.py` - Create: `tests/inference_helpers.py`(FakeClient / FakeTransport) - Create: `tests/test_inference_engine.py` - Modify: `src/genesis/inference/__init__.py`(导出 Engine / ChatResult 等) - Modify: `docs/agent-runtime-design.md`(§2.2 补 StructuredResult.status 字段) - Modify: `docs/api-design.md`(§7 错误码表补 LLM_NOT_CONFIGURED 行) - Modify: `docs/config-design.md`(§4 token_estimation 注明双语义降级) **Interfaces:** - Consumes: `LLMError` 子树(Task1)、`ChatMessage/ChatResult/Prompt/StructuredResult/TokenUsage`(Task1)、`make_estimator` + `approximate_token_count`(Task2)、`PromptRegistry`(Task3)、`LLMClient`(Task4) - Produces: - `class InferenceEngine`:`__init__(client, models: InferenceModels | None = None, registry=None, estimator=None, truncate_cb=None)`;`chat(...) -> ChatResult`;`chat_structured(...) -> StructuredResult` - `chat`:渲染 → 裁剪判断 → 主模型调用 → 失败再取 fallback → status - `chat_structured`:带 schema 提示构造 prompt → 逐次解析(最多 retry_count+1 次)→ parse_error 带 raw_text - [ ] **Step 1: 写失败测试** `tests/inference_helpers.py`: ```python from __future__ import annotations import json import httpx from genesis.inference.types import ChatMessage, TokenUsage class FakeLLMClient: """可编程的假 LLM 客户端:记录调用,按脚本返回(离线)。""" def __init__(self, script=None): # script: list[(status, content)];status: "ok" | "raise_timeout" | "raise_network" | "parse_fail" self.script = script or [("ok", "hello")] self.calls: list[dict] = [] def chat(self, *, model, messages, temperature, max_tokens): self.calls.append({"model": model, "messages": [m.content for m in messages]}) status, content = self.script.pop(0) if status == "raise_timeout": from genesis.inference.exceptions import LLMTimeoutError raise LLMTimeoutError("timeout") if status == "raise_network": from genesis.inference.exceptions import LLMNetworkError raise LLMNetworkError("network") if status == "parse_fail": content = "NOT JSON" return content, TokenUsage(input_tokens=10, output_tokens=2) ``` > 注意:FakeClient.chat 的签名必须匹配 Protocol(model/messages/temperature/max_tokens)——上方已对齐。 def ok_transport(): """httpx.MockTransport:返回合法 JSON 响应。""" def handler(request): body = json.loads(request.content) return httpx.Response(200, json={ "choices": [{"message": {"content": "structured:" + json.dumps({"a": body.get("model")})}}], "usage": {"prompt_tokens": 3, "completion_tokens": 1}, }) return httpx.MockTransport(handler) ``` > 注意:FakeClient.chat 的签名必须匹配 Protocol(model/messages/temperature/max_tokens)——上方已对齐。`M` 类仅为占位,engine 构造时推荐直接传 dict:由测试构造真对象,见下文 engine 测试。 `tests/test_inference_engine.py`: ```python from __future__ import annotations import json import pytest from genesis.inference.engine import InferenceEngine from genesis.inference.exceptions import LLMError from genesis.inference.prompt_registry import PromptRegistry from genesis.inference.token import approximate_token_count from genesis.inference.types import ChatMessage, Prompt from tests.inference_helpers import FakeLLMClient class Models: """模拟 config.InferenceModels(pydantic 结构,测试中直接构建)""" def __init__(self): import types as t self.primary = t.SimpleNamespace(name="deepseek-chat", provider="deepseek") self.fallback = t.SimpleNamespace(name="qwen-max", provider="qwen") def make_engine(client=None, *, constants=None): eng = InferenceEngine( client=client or FakeLLMClient(), models=Models(), registry=PromptRegistry(), estimator=approximate_token_count, ) if constants: eng._max_context_tokens = constants # 内部测试钩子 return eng def test_chat_ok(): client = FakeLLMClient([("ok", "正文")]) eng = make_engine(client) r = eng.chat( session_id="s1", prompt=Prompt(name="writer", version="v1", template="章节:{{ chapter }}"), variables={"chapter": "DB設計"}, ) assert r.status == "ok" and r.text == "正文" assert client.calls[0]["model"] == "deepseek-chat" def test_chat_fallback_after_primary_failure(): client = FakeLLMClient([("raise_timeout", ""), ("ok", "备用输出")]) eng = make_engine(client) r = eng.chat( session_id="s1", prompt=Prompt(name="p", version="v1", template="t:{{ x }}"), variables={"x": "1"}, ) assert r.status == "fallback" assert client.calls[-1]["model"] == "deepseek-chat" # 先主,后备 def test_chat_all_failed_returns_failed(): client = FakeLLMClient([("raise_network", ""), ("raise_network", "")]) eng = make_engine(client) r = eng.chat( session_id="s1", prompt=Prompt(name="p", version="v1", template="t"), variables={}, ) assert r.status == "failed" and r.error def test_chat_plain_string_prompt(): eng = make_engine(FakeLLMClient([("ok", "hi")])) r = eng.chat(session_id="s1", prompt="直接文本", variables={}) assert r.text == "hi" and r.status == "ok" def test_chat_truncation_callback_triggered(): seen = {} def truncate_cb(prompt_text, variables): seen["called"] = True seen["len"] = len(prompt_text) return {**variables, "chapter": "裁剪版"} eng = InferenceEngine( client=FakeLLMClient([("ok", "x")]), models=Models(), registry=PromptRegistry(), estimator=approximate_token_count, truncate_cb=truncate_cb, ) eng._max_context_tokens = 3 # 强制超限 r = eng.chat( session_id="s1", prompt=Prompt(name="p", version="v1", template="abcd{{ chapter }}"), variables={"chapter": "很长很长的标题"}, ) assert seen["called"] is True assert r.status == "ok" def test_chat_structured_ok(): client = FakeLLMClient([("ok", '{"a": 1}')]) eng = make_engine(client) r = eng.chat_structured( session_id="s1", prompt=Prompt(name="p", version="v1", template="提取"), variables={"text": "内容"}, schema={"type": "object", "properties": {"a": {"type": "number"}}}, ) assert r.status == "ok" and r.data == {"a": 1} def test_chat_structured_retry_parse(): client = FakeLLMClient([("parse_fail", ""), ("ok", '{"a": 2}')]) eng = make_engine(client) r = eng.chat_structured( session_id="s1", prompt=Prompt(name="p", version="v1", template="提取"), variables={}, schema={}, ) assert r.status == "ok" and r.data == {"a": 2} and r.parse_attempts == 2 def test_chat_structured_parse_error_returns_raw(): client = FakeLLMClient([("parse_fail", ""), ("parse_fail", "")]) eng = make_engine(client) r = eng.chat_structured( session_id="s1", prompt=Prompt(name="p", version="v1", template="提取"), variables={}, schema={}, retry_count=1, ) assert r.status == "parse_error" assert r.raw_text == "NOT JSON" assert r.parse_attempts == 2 def test_chat_structured_failed_on_network(): client = FakeLLMClient([("raise_network", "")]) eng = make_engine(client) r = eng.chat_structured( session_id="s1", prompt=Prompt(name="p", version="v1", template="提取"), variables={}, schema={}, ) assert r.status == "failed" ``` - [ ] **Step 2: 运行确认失败** Run: `python -m pytest tests/test_inference_engine.py -v` Expected: FAIL(`ModuleNotFoundError: No module named 'genesis.inference.engine'`) - [ ] **Step 3: 实现** `src/genesis/inference/engine.py`: ```python from __future__ import annotations import json import time from typing import Any, Callable from .client import LLMClient from .exceptions import LLMError from .prompt_registry import PromptRegistry from .token import make_estimator from .types import ( ChatMessage, ChatResult, Prompt, StructuredResult, TokenUsage, ) class InferenceEngine: """统一 LLM 调用入口:模型选择/降级、重试、解析、Token 超限回调。""" def __init__( self, *, client: LLMClient, models: Any | None = None, registry: PromptRegistry | None = None, estimator: Callable[[str], int] | None = None, truncate_cb: Callable[[str, dict], dict] | None = None, max_context_tokens: int = 32000, ) -> None: self._client = client self._models = models self._registry = registry or PromptRegistry() self._estimator = estimator or make_estimator() self._truncate_cb = truncate_cb self._max_context_tokens = max_context_tokens # ---------- 内部 ---------- def _render_prompt(self, prompt: Prompt | str, variables: dict) -> str: if isinstance(prompt, Prompt): return self._registry.render(prompt.template, variables) if variables else prompt.template return prompt def _apply_truncation(self, text: str, variables: dict) -> dict: """Token 超限时触发裁剪回调(注入),返回新 variables。""" if self._truncate_cb is not None: new_vars = self._truncate_cb(text, variables) if new_vars is not None: return new_vars return variables def _model_names(self, model: str | None) -> list[str]: """返回尝试顺序;显式指定 model 时只用它,否则 primary→fallback。""" if model: return [model] if self._models: names = [] if getattr(self._models, "primary", None): names.append(self._models.primary.name) if getattr(self._models, "fallback", None): names.append(self._models.fallback.name) if names: return names return ["deepseek-chat"] def _call( self, *, model: str, rendered: str, temperature: float, max_tokens: int, ) -> tuple[str, TokenUsage]: messages = [ChatMessage(role="user", content=rendered)] return self._client.chat( model=model, messages=messages, temperature=temperature, max_tokens=max_tokens, ) # ---------- 公开 ---------- def chat( self, *, session_id: str, prompt: Prompt | str, variables: dict, model: str | None = None, temperature: float = 0.2, max_tokens: int = 4096, ) -> ChatResult: rendered = self._render_prompt(prompt, variables) if self._estimator(rendered) > self._max_context_tokens: variables = self._apply_truncation(rendered, variables) rendered = self._render_prompt(prompt, variables) start = time.monotonic() last_error: str | None = None for idx, name in enumerate(self._model_names(model)): try: text, usage = self._call( model=name, rendered=rendered, temperature=temperature, max_tokens=max_tokens, ) status = "ok" if idx == 0 else "fallback" return ChatResult( text=text, model=name, prompt_version=getattr(prompt, "version", "inline"), usage=usage, duration_ms=int((time.monotonic() - start) * 1000), status=status, ) except LLMError as exc: last_error = str(exc) return ChatResult( text="", model=name, prompt_version=getattr(prompt, "version", "inline"), usage=TokenUsage(), duration_ms=int((time.monotonic() - start) * 1000), status="failed", error=last_error, ) def chat_structured( self, *, session_id: str, prompt: Prompt | str, variables: dict, schema: dict, retry_count: int = 2, ) -> StructuredResult: rendered = self._render_prompt(prompt, variables) # 追加 schema 约束说明(不强制模板支持) schema_hint = json.dumps(schema, ensure_ascii=False) if schema else "" base_rendered = rendered + (f'\n\n请输出符合以下 JSON Schema 的 JSON:{schema_hint}' if schema_hint else "") start = time.monotonic() attempts = 0 last_raw = "" last_error: str | None = None while attempts <= retry_count: attempts += 1 try: text, usage = self._call( model=self._model_names(None)[0], # 解析重试用首选模型 rendered=base_rendered, temperature=0.0, max_tokens=4096, ) last_raw = text data = json.loads(text) return StructuredResult( data=data, raw_text=text, parse_attempts=attempts, model=self._model_names(None)[0], prompt_version=getattr(prompt, "version", "inline"), usage=usage, duration_ms=int((time.monotonic() - start) * 1000), # fallback 语义保留给模型降级;解析重试成功仍为 ok status="ok", ) except json.JSONDecodeError as exc: last_error = f"JSON 解析失败: {exc}" # 带错误信息重试 base_rendered = base_rendered + f"\n\n上次解析失败:{exc}。请重新输出合法 JSON。" except LLMError as exc: return StructuredResult( data={}, raw_text="", parse_attempts=attempts, model=self._model_names(None)[0], prompt_version=getattr(prompt, "version", "inline"), usage=TokenUsage(), duration_ms=int((time.monotonic() - start) * 1000), status="failed", error=str(exc), ) return StructuredResult( data={}, raw_text=last_raw, parse_attempts=attempts, model=self._model_names(None)[0], prompt_version=getattr(prompt, "version", "inline"), usage=TokenUsage(), duration_ms=int((time.monotonic() - start) * 1000), status="parse_error", error=last_error, ) ``` > 注:`_model_names(None)[0]` 仅用于结构解析的重试模型,生产建议显式传主模型,实现已足够 v1。 - [ ] **Step 4: 运行确认通过** Run: `python -m pytest tests/test_inference_engine.py -v` Expected: PASS(10 passed;若某用例失败按「实现约定」修正) - [ ] **Step 5: 更新 `__init__.py` 导出** `src/genesis/inference/__init__.py` 追加: ```python from .engine import InferenceEngine from .prompt_registry import PromptRegistry # __all__ 追加 "InferenceEngine", "PromptRegistry" ``` - [ ] **Step 6: 文档补丁(实现批准 spec §3.9)** - `docs/agent-runtime-design.md` §2.2 的 `StructuredResult` 代码块增加 `status: Literal["ok","fallback","parse_error","failed"]`(与 error 字段) - `docs/api-design.md` §7 错误码表追加行 `| LLM_NOT_CONFIGURED | 503 | LLM Key 未配置 | 配置 Key | exceptions.LLMNotConfiguredError |` - `docs/config-design.md` §4 `llm_calls.token_estimation` 注释追加「tiktoken 缺失自动回落 approximate(内置估算器)」 - [ ] **Step 7: 全量回归** Run: `python -m pytest -v` Expected: PASS(71 + 7 + 5 + 5 + 6 + 10 = 104 passed;`fail_under 99` 全绿) - [ ] **Step 8: 提交** ```bash git add src/genesis/engine.py docs/agent-runtime-design.md docs/api-design.md docs/config-design.md git commit -m "feat: InferenceEngine chat/chat_structured 全流程 + 文档补丁(补丁1/2)" ``` --- ## Self-Review **1. Spec 覆盖**:§3.1 模块结构→Task1-5;§3.2 types→Task1;§3.3 client→Task4;§3.4 重试/降级→Task4(client 重试)+Task5(engine 主→备);§3.5 token→Task2+Task5 裁剪;§3.6 Prompt→Task3;§3.7 engine→Task5;§3.8 异常→Task1+Task4;§3.9 文档→Task5 Step6;§3.10 测试→各任务 + 全量回归;§4 验收→Step7+验收清单。✓ **2. 占位符检查**:无 TODO/TBD;engine.py 每步有完整代码;修正笔误(`invalid`→`exceptions`)。✓ **3. 类型一致性**:`chat_structured` 签名(schema/retry_count)在各测试与实现一致;`FakeLLMClient.chat` 匹配 Protocol(model/messages/temperature/max_tokens);`StructuredResult.status` 四值一致(ok/fallback/parse_error/failed)。✓ **4. 依赖顺序**:Task1 铺 types/exceptions;Task2 独立 token;Task3 用 Prompt;Task4 用 ChatMessage;Task5 全用。无循环导入(__init__ Task1 不引 engine,Task5 才 export)。✓ **5. 覆盖率红线**:每文件配分支测试(client 覆盖 5xx 重试/4xx 不重试/超时/未配置;engine 覆盖 ok/fallback/failed/超限/解析重试/parse_error/failed 网络失败)。✓ **6. 验收对齐**:chat_structured 在网格 4xx 等异常下返回 failed 而非抛——与 api §7 retry/skip/abort UX 一致;parse_error 带 raw_text 满足验收 3。✓ ## 10. 提交消息 按 Global Constraints 提交消息风格:`feat:` / `test:` / `docs:` + 简中文描述,如上各任务 Step 5/8。