Files
2026Technology-Competition/docs/superpowers/plans/2026-08-09-inference-engine.md
T

1189 lines
40 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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)、jinja2PromptRegistry 渲染)、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/<file> -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: PASS5 + 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.pyapproximate 内置 / 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
缺失或异常回落内置 approximatebackend="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: PASS5 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: PASS5 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.pyhttpx + 重试/退避/超时/认证)
**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: PASS6 passed
- [ ] **Step 5: 全量回归**(新代码暴露到覆盖统计)
Run: `python -m pytest -v`
Expected: PASS71 + 全部新用例;`fail_under=99` 通过)
- [ ] **Step 6: 提交**
```bash
git add src/genesis/inference/client.py tests/test_inference_client.py
git commit -m "feat: HttpLLMClienthttpx + 重试退避/超时/鉴权)"
```
---
### Task 5: InferenceEngineengine.pychat / 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 的签名必须匹配 Protocolmodel/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 的签名必须匹配 Protocolmodel/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.InferenceModelspydantic 结构,测试中直接构建)"""
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: PASS10 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: PASS71 + 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/TBDengine.py 每步有完整代码;修正笔误(`invalid``exceptions`)。✓
**3. 类型一致性**`chat_structured` 签名(schema/retry_count)在各测试与实现一致;`FakeLLMClient.chat` 匹配 Protocolmodel/messages/temperature/max_tokens);`StructuredResult.status` 四值一致(ok/fallback/parse_error/failed)。✓
**4. 依赖顺序**Task1 铺 types/exceptionsTask2 独立 tokenTask3 用 PromptTask4 用 ChatMessageTask5 全用。无循环导入(__init__ Task1 不引 engineTask5 才 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。