Files
2026Technology-Competition/src/genesis/inference/engine.py
T
lhl 25fc472d9b feat(inference): chat_structured 接入 jsonschema 真校验(T1 架构审查整改)
- pyproject.toml 新增 jsonschema>=4.23 依赖
- engine.py: json.loads 后 jsonschema.validate,校验失败按解析失败重试(带错误信息)
- 新增 3 用例:违规重试成功/耗尽 parse_error 含详情/合法一次通过
- 全量 163 passed / 100.00% 覆盖(936 stmts/238 br),fail_under=99 达标
2026-08-12 09:32:02 +08:00

203 lines
7.7 KiB
Python
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.
from __future__ import annotations
import json
import time
from typing import Any, Callable
import jsonschema
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
last_error_code: 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)
last_error_code = exc.error_code # 同源:取最后一次失败异常
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, error_code=last_error_code,
)
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)
if self._estimator(rendered) > self._max_context_tokens:
variables = self._apply_truncation(rendered, variables)
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
last_error_code: 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)
if schema:
# 真 schema 校验:不合 schema 时按解析失败重试(T1)
jsonschema.validate(instance=data, schema=schema)
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}"
last_error_code = "LLM_PARSE_ERROR"
# 带错误信息重试
base_rendered = base_rendered + f"\n\n上次解析失败:{exc}。请重新输出合法 JSON。"
except jsonschema.ValidationError as exc:
last_error = f"校验失败: {exc.message}"
last_error_code = "LLM_PARSE_ERROR"
# 带校验错误信息重试
base_rendered = base_rendered + f"\n\n上次校验失败:{exc.message}。请重新输出符合 JSON Schema 的 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), error_code=exc.error_code,
)
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, error_code=last_error_code,
)