40 KiB
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.tomlfail_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:
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:
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 增加:
"httpx>=0.28",
"jinja2>=3.1",
- Step 4: 实现
src/genesis/inference/types.py:
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:
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:
"""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: 提交
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:
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:
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: 提交
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) -> Noneget(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:
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:
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: 提交
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:
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:
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: 提交
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(...) -> StructuredResultchat:渲染 → 裁剪判断 → 主模型调用 → 失败再取 fallback → statuschat_structured:带 schema 提示构造 prompt → 逐次解析(最多 retry_count+1 次)→ parse_error 带 raw_text
-
Step 1: 写失败测试
tests/inference_helpers.py:
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:
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 追加:
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§4llm_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: 提交
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。