- 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 达标
203 lines
7.7 KiB
Python
203 lines
7.7 KiB
Python
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,
|
||
)
|