Merge PR #98/#100 (feat #57 InferenceBackend 抽象契约 + #58 GPU Triton 后端,桥接 #44/#45 本地与云端实现)

This commit is contained in:
2026-08-05 08:31:30 +08:00
parent 6036e5e151
commit 6d0b2976bd
5 changed files with 1192 additions and 57 deletions
+17 -8
View File
@@ -49,13 +49,21 @@ from .hallucination import (
HallucinationGuard,
)
from .gateway import (
CloudBackend,
GatewayResult,
InferenceBackend,
LLMGateway,
LocalBackend,
)
from .backends import CloudApiBackend, Local70BBackend
from .backends import (
BackendCapabilities,
BackendHealth,
CloudApiBackend,
CloudBackend,
InferResult,
InferenceBackend,
Local70BBackend,
LocalBackend,
build_backend,
default_registry,
)
__all__ = [
# dlp
@@ -67,9 +75,10 @@ __all__ = [
# hallucination
"GuardVerdict", "HallucinationGuard",
# gateway
"InferenceBackend", "LocalBackend", "CloudBackend",
# backends
"Local70BBackend",
"CloudApiBackend",
"GatewayResult", "LLMGateway",
# 推理后端(#57 抽象契约 + #44 本地 / #45 云端 / #58 GPU)
"InferenceBackend", "BackendCapabilities", "BackendHealth", "InferResult",
"LocalBackend", "CloudBackend",
"Local70BBackend", "CloudApiBackend",
"default_registry", "build_backend",
]
+385 -49
View File
@@ -1,30 +1,327 @@
# -*- coding: utf-8 -*-
"""推理后端实现 —— 本地 70B 模型接入与推理封装(issue #44)。
"""iAOP-Core · LLM 网关 —— 推理后端抽象接口(Issue #57,PRD 5.6)。
在 `gateway.InferenceBackend` 抽象之上交付**真实可用的本地后端**:
- OpenAI 兼容接口(vLLM / TGI 等本地推理服务,`/v1/chat/completions`),
仅用标准库 urllib,无第三方依赖;
- 参数化:endpoint / model / timeout / max_tokens / temperature / context 引用注入;
- **数据不出厂**(PRD 5.4):敏感/核心内容走本地后端,云端仅接收脱敏内容;
- 未配置 endpoint 时进入 dry-run 占位模式(保持与旧 LocalBackend 一致的
可测试行为,供端到端演示与联调)。
PRD 5.6「⑥ 部署底座」明确要求:
业务代码只依赖 `gateway.InferenceBackend.generate(prompt, context)`,
切换后端 = 换实现(见 `LLMGateway(local=...)`)。
定义统一 ``InferenceBackend`` 接口(``loadModel / infer / health / unload``),
5090 实现(Triton/ONNX)与昇腾实现(ACL/CANN)均实现该接口;
**业务代码仅依赖接口,不感知硬件**;切换后端 = 改适配层配置,不动业务代码。
本模块把原先内联在 ``gateway.py`` 里的薄弱 ``InferenceBackend`` 提炼为正式的
抽象基类(ABC),并补齐 PRD 要求的生命周期方法与能力声明,使后续子任务:
- #44 本地 70B 模型接入与推理封装(vLLM/TGI)
- #45 云端 API(Qwen/DeepSeek)接入与安全网关
- #58 GPU 后端实现(NVIDIA,Triton/ONNX)
- #59 昇腾 NPU 后端适配(CANN/ACL)
都能在**同一契约**下落地,业务编排(``LLMGateway``)零改动。
设计要点
--------
1. **接口最小且完备**:仅约束 PRD 列出的四个生命周期动作 ``load_model / infer /
health_check / unload``,外加能力声明 ``BackendCapabilities``(流式 / 最大并发 /
是否出厂内闭环),供路由与调度决策。
2. **向后兼容**:保留 ``generate(prompt, context)`` 便捷方法(默认转发到
``infer``),既有 ``LLMGateway.ask()`` 调用路径不变;老测试不受影响。
3. **可注入 / 可 mock**:所有方法纯逻辑、无外部 IO 依赖;真实硬件/网络交互由
各子类在 ``infer`` 内部完成(子类负责导入厂商 SDK 并做 ``ImportError`` 容错)。
4. **健康探针**:``health_check`` 返回结构化 ``BackendHealth``,供可用性监控探针
(Issue #61)与灰度发布(PRD 5.6 配置点)判定后端是否就绪。
测试:``python -m unittest discover -s tests -v``(在 core/llm-gateway 目录下执行)。
"""
from __future__ import annotations
import json
import os
import time
import urllib.request
from typing import Callable, Optional, Sequence
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Dict, Iterator, List, Optional, Sequence
from .gateway import InferenceBackend
# ---------------------------------------------------------------------------
# 值对象:能力声明 / 健康状态 / 推理结果
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class BackendCapabilities:
"""后端能力声明,供路由 / 调度 / 灰度决策。
Attributes:
streaming: 是否支持流式输出(逐 token 返回)。
max_concurrency: 最大并发推理数(None 表示不限 / 由外部限流)。
on_premises: 是否数据出厂内闭环(本地后端 True,云端 False)。
modalities: 支持的输出形态,如 ``("text",)``。
"""
streaming: bool = False
max_concurrency: Optional[int] = None
on_premises: bool = False
modalities: Sequence[str] = ("text",)
def supports(self, modality: str) -> bool:
"""是否支持某种输出形态(text / image / ...)。"""
return modality in self.modalities
def to_dict(self) -> Dict[str, object]:
return {
"streaming": self.streaming,
"max_concurrency": self.max_concurrency,
"on_premises": self.on_premises,
"modalities": list(self.modalities),
}
@dataclass(frozen=True)
class BackendHealth:
"""后端健康探针结果(Issue #61 可用性监控探针消费)。"""
healthy: bool
detail: str = ""
checked_at: str = field(
default_factory=lambda: datetime.now(timezone.utc).isoformat())
def to_dict(self) -> Dict[str, object]:
return {
"healthy": self.healthy,
"detail": self.detail,
"checked_at": self.checked_at,
}
@dataclass(frozen=True)
class InferResult:
"""一次 ``infer`` 的结构化结果(含审计所需元信息)。
保留 ``text`` 主输出以兼容旧 ``generate`` 返回 ``str`` 的调用方;
``prompt_tokens`` / ``completion_tokens`` 供计费与配额(PRD 5.6 配置点)。
"""
text: str
backend_name: str
model_id: str = ""
prompt_tokens: Optional[int] = None
completion_tokens: Optional[int] = None
latency_ms: Optional[float] = None
def to_dict(self) -> Dict[str, object]:
return {
"text": self.text,
"backend_name": self.backend_name,
"model_id": self.model_id,
"prompt_tokens": self.prompt_tokens,
"completion_tokens": self.completion_tokens,
"latency_ms": self.latency_ms,
}
# ---------------------------------------------------------------------------
# 抽象接口(PRD 5.6:loadModel / infer / health / unload)
# ---------------------------------------------------------------------------
class InferenceBackend(ABC):
"""推理后端抽象接口(对齐 PRD 5.6 ``InferenceBackend`` 契约)。
业务编排(``LLMGateway``)只依赖本接口,**不感知**具体硬件 / 厂商;
切换后端 = 换实现类 + 改配置,业务代码不动。子类必须实现四个生命周期方法:
- :meth:`load_model`:加载 / 绑定模型(可幂等,重复加载返回已加载实例)。
- :meth:`infer`:给定 prompt 与 RAG 上下文生成回答(核心推理动作)。
- :meth:`health_check`:探针,返回 :class:`BackendHealth`。
- :meth:`unload`:释放模型资源(可幂等)。
便捷方法 :meth:`generate` 默认转发到 :meth:`infer` 并只取 ``text``,
保留与旧 ``LLMGateway.ask()`` 的二进制兼容。
"""
#: 后端短名(local-70b / cloud-api / gpu-triton / npu-cann ...),子类覆盖。
name: str = "base"
@property
def capabilities(self) -> BackendCapabilities:
"""后端能力声明,子类按需覆盖。默认:非流式、出厂外、仅文本。"""
return BackendCapabilities()
# -- 生命周期(子类必须实现)------------------------------------------
@abstractmethod
def load_model(self, model_id: str) -> None:
"""加载 / 绑定指定模型。幂等:重复加载同一 model_id 不报错。"""
@abstractmethod
def infer(self, prompt: str,
context: Optional[Sequence[str]] = None) -> InferResult:
"""根据 prompt 与 RAG 上下文生成回答(核心推理动作)。"""
@abstractmethod
def health_check(self) -> BackendHealth:
"""健康探针,返回结构化健康状态。"""
@abstractmethod
def unload(self) -> None:
"""释放模型资源。幂等:未加载时调用不报错。"""
# -- 向后兼容便捷方法 --------------------------------------------------
def generate(self, prompt: str, context: Sequence[str]) -> str:
"""旧调用入口:等价于 ``infer(prompt, context).text``。
保留是为了不破坏 ``LLMGateway.ask()`` 既有的 ``backend.generate(...)``
调用路径;新代码应直接使用 :meth:`infer` 拿到完整 :class:`InferResult`。
"""
return self.infer(prompt, context).text
def __repr__(self) -> str: # pragma: no cover - 调试辅助
return f"<{type(self).__name__} name={self.name!r}>"
# ---------------------------------------------------------------------------
# 占位实现(子任务 #44 / #45 / #58 / #59 将各自替换为真实后端)
# ---------------------------------------------------------------------------
class _PlaceholderBackend(InferenceBackend):
"""占位后端公共骨架:固定回显答案 + 引用溯源回显,供端到端测试与演示。
真实后端(#44 本地 70B / #45 云端 API / #58 GPU / #59 昇腾)继承本类后,
只需覆盖 :meth:`infer` 的生成逻辑与 :meth:`health_check` 的探针实现即可;
生命周期与能力声明已由本类 / 子类提供。
"""
placeholder_prefix = "[占位]"
def __init__(self, model_id: str, echo_context: bool = True) -> None:
self._model_id = model_id
self._loaded = False
self._loaded_model_id: Optional[str] = None
self.echo_context = echo_context
# 生命周期
def load_model(self, model_id: str) -> None:
# 幂等:重复加载同一 model_id 视作成功;换模型也允许(演示用)。
self._loaded = True
self._loaded_model_id = model_id or self._model_id
def infer(self, prompt: str,
context: Optional[Sequence[str]] = None) -> InferResult:
if not self._loaded:
# 演示态允许惰性自加载,真实后端可改为 raise RuntimeError("未加载模型")
self.load_model(self._model_id)
ctx = list(context or [])
head = f"{self.placeholder_prefix} {prompt[:40]}"
refs = ""
if self.echo_context:
for src in ctx[:3]:
refs += f"\n[来源: {src}]"
return InferResult(
text=head + refs,
backend_name=self.name,
model_id=self._loaded_model_id or self._model_id,
)
def health_check(self) -> BackendHealth:
return BackendHealth(
healthy=self._loaded,
detail="loaded" if self._loaded else "not_loaded",
)
def unload(self) -> None:
# 幂等:未加载也安全
self._loaded = False
self._loaded_model_id = None
class LocalBackend(_PlaceholderBackend):
"""本地 70B 后端占位实现:数据不出厂(敏感 / 核心走此通道)。
子任务 #44 / #58 将替换 ``infer`` 为真实本地模型推理封装(vLLM/TGI/Triton)。
"""
name = "local-70b"
placeholder_prefix = "[本地70B占位]"
def __init__(self, echo_context: bool = True,
model_id: str = "local-70b-base") -> None:
super().__init__(model_id=model_id, echo_context=echo_context)
@property
def capabilities(self) -> BackendCapabilities:
# 本地后端:出厂内闭环、可流式、单卡典型并发 8(演示默认值)
return BackendCapabilities(
streaming=True, max_concurrency=8, on_premises=True,
modalities=("text",))
class CloudBackend(_PlaceholderBackend):
"""云端 API 后端占位实现:仅接收 DLP 放行的脱敏 / 通用内容。
子任务 #45 将替换为 Qwen / DeepSeek API 接入 + 安全网关。
"""
name = "cloud-api"
placeholder_prefix = "[云端API占位]"
def __init__(self, echo_context: bool = True,
model_id: str = "cloud-qwen-plus") -> None:
super().__init__(model_id=model_id, echo_context=echo_context)
@property
def capabilities(self) -> BackendCapabilities:
# 云端后端:数据出厂、支持流式、并发受厂商配额限制(演示默认 4)
return BackendCapabilities(
streaming=True, max_concurrency=4, on_premises=False,
modalities=("text",))
# ---------------------------------------------------------------------------
# 后端注册表(配置驱动切换,对齐 PRD「切换后端 = 改适配层配置」)
# ---------------------------------------------------------------------------
def default_registry() -> Dict[str, type]:
"""默认后端注册表:name → 实现类。新增后端在此登记一行即可被配置选用。"""
# 延迟导入避免循环依赖(gpu_backend 反向依赖本模块的抽象基类与值对象)
from .gpu_backend import GpuTritonBackend # noqa: WPS433(Issue #58)
return {
"local-70b": LocalBackend,
"cloud-api": CloudBackend,
"gpu-triton": GpuTritonBackend,
}
def build_backend(name: str, **kwargs) -> InferenceBackend:
"""按 name 从默认注册表构造后端实例(配置驱动切换的入口)。
未知 name 抛 ``ValueError``,列出已知项便于排错。
"""
registry = default_registry()
cls = registry.get(name)
if cls is None:
known = ", ".join(sorted(registry))
raise ValueError(f"未知推理后端 {name!r},已知: {known}")
return cls(**kwargs)
# ---------------------------------------------------------------------------
# 真实后端实现(Issue #44 本地 70B / #45 云端 API,桥接到 #57 抽象契约)
# ---------------------------------------------------------------------------
import json as _json
import os as _os
import time as _time
import urllib.request as _urllib
class Local70BBackend(InferenceBackend):
"""本地 70B 推理后端(OpenAI 兼容 vLLM/TGI,参数化)。"""
"""本地 70B 推理后端(OpenAI 兼容 vLLM/TGI,参数化)—— issue #44。
- 数据不出厂(PRD 5.4):敏感/核心内容走本地后端;
- 未配置 endpoint 时进入 dry-run 占位模式(端到端演示与联调);
- 已桥接 #57 契约:load_model / infer / health_check / unload 齐备。
"""
name = "local-70b"
@@ -43,8 +340,28 @@ class Local70BBackend(InferenceBackend):
self.max_tokens = int(max_tokens)
self.temperature = float(temperature)
self.echo_context = echo_context
self._loaded_model_id: Optional[str] = None
# ------------------------------------------------------------------
# -- #57 契约 ------------------------------------------------------
def load_model(self, model_id: str) -> None:
self._loaded_model_id = model_id
def infer(self, prompt: str,
context: Optional[Sequence[str]] = None) -> InferResult:
text = self.generate(prompt, context or ())
return InferResult(
text=text, backend_name=self.name, model_id=self.model)
def health_check(self) -> BackendHealth:
info = self.health()
healthy = info.get("status") in ("ok", "dry-run", "configured")
return BackendHealth(healthy=healthy,
detail=_json.dumps(info, ensure_ascii=False))
def unload(self) -> None:
self._loaded_model_id = None
# -- 原 #44 实现 ----------------------------------------------------
def generate(self, prompt: str, context: Sequence[str]) -> str:
"""根据 prompt 与 RAG 上下文生成回答。
@@ -70,30 +387,27 @@ class Local70BBackend(InferenceBackend):
raise RuntimeError(
f"本地推理服务响应格式异常: {str(body)[:200]}")
# ------------------------------------------------------------------
def _system_prompt(self, context: Sequence[str]) -> str:
"""把 RAG 引用注入 system 提示(引用溯源,PRD 5.4)。"""
refs = "\n".join(f"- {c}" for c in (context or []))
refs = "\\n".join(f"- {c}" for c in (context or []))
base = "你是工业 AI 助手。回答须基于给定资料并标注来源。"
return f"{base}\n参考资料:\n{refs}" if refs else base
return f"{base}\\n参考资料:\\n{refs}" if refs else base
def _dry_run(self, prompt: str, context: Sequence[str]) -> str:
head = f"[本地70B占位] {prompt[:40]}"
if self.echo_context:
for i, src in enumerate(context[:3], 1):
head += f"\n[来源: {src}]"
head += f"\\n[来源: {src}]"
return head
def _post_json(self, path: str, payload: dict) -> dict:
"""向后端推理服务发起 JSON POST(标准库 urllib)。"""
url = self.endpoint + path
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(
data = _json.dumps(payload).encode("utf-8")
req = _urllib.Request(
url, data=data,
headers={"Content-Type": "application/json"})
with urllib.request.urlopen(req, timeout=self.timeout) as resp:
with _urllib.urlopen(req, timeout=self.timeout) as resp:
raw = resp.read().decode("utf-8")
return json.loads(raw) if raw else {}
return _json.loads(raw) if raw else {}
def health(self) -> dict:
"""后端健康信息(本地推理服务可探测 /health)。"""
@@ -105,11 +419,11 @@ class Local70BBackend(InferenceBackend):
base["status"] = "dry-run"
return base
try:
started = time.monotonic()
with urllib.request.urlopen(
started = _time.monotonic()
with _urllib.urlopen(
self.endpoint + "/health", timeout=self.timeout) as resp:
base["status"] = "ok" if resp.status == 200 else f"http-{resp.status}"
base["latency_ms"] = round((time.monotonic() - started) * 1000, 2)
base["latency_ms"] = round((_time.monotonic() - started) * 1000, 2)
except Exception as exc: # noqa: BLE001 - 健康探测失败仅记录
base["status"] = f"error: {exc}"
return base
@@ -119,11 +433,9 @@ class CloudApiBackend(InferenceBackend):
"""云端 API 推理后端(Qwen / DeepSeek 等 OpenAI 兼容)—— issue #45。
**安全网关约束(PRD 5.4)**:
- 仅接收 **DLP 放行**的脱敏/通用内容(上游 `LLMGateway` 主编排出站检查 +
cloud 分支输出 DLP 复查);
- 仅接收 **DLP 放行**的脱敏/通用内容;
- API Key 从**环境变量**读取(`api_key_env`),不硬编码、不落日志;
- 可选 `safety_checker` 出站复查钩子(fail-closed:复查拒绝 → 拦截占位,
不调用上游)。
- 可选 `safety_checker` 出站复查钩子(fail-closed:复查拒绝 → 拦截占位)。
"""
name = "cloud-api"
@@ -144,13 +456,31 @@ class CloudApiBackend(InferenceBackend):
self.timeout = float(timeout_seconds)
self.max_tokens = int(max_tokens)
self.temperature = float(temperature)
# 出站安全复查:返回 False 即拦截(fail-closed)
self.safety_checker = safety_checker
self._api_key = os.environ.get(api_key_env, "") if api_key_env else ""
self._api_key = _os.environ.get(api_key_env, "") if api_key_env else ""
self._loaded_model_id: Optional[str] = None
# ------------------------------------------------------------------
# -- #57 契约 ------------------------------------------------------
def load_model(self, model_id: str) -> None:
self._loaded_model_id = model_id
def infer(self, prompt: str,
context: Optional[Sequence[str]] = None) -> InferResult:
text = self.generate(prompt, context or ())
return InferResult(
text=text, backend_name=self.name, model_id=self.model)
def health_check(self) -> BackendHealth:
info = self.health()
healthy = info.get("status") in ("ok", "dry-run", "configured")
return BackendHealth(healthy=healthy,
detail=_json.dumps(info, ensure_ascii=False))
def unload(self) -> None:
self._loaded_model_id = None
# -- 原 #45 实现 ----------------------------------------------------
def generate(self, prompt: str, context: Sequence[str]) -> str:
"""生成回答。安全网关:safety_checker 拒绝 → 拦截占位,不调用上游。"""
if self.safety_checker is not None and not self.safety_checker(prompt):
return "[云端安全网关拦截] 出站复查未通过,已拦截(数据不出厂)。"
@@ -173,31 +503,29 @@ class CloudApiBackend(InferenceBackend):
raise RuntimeError(
f"云端 API 响应格式异常: {str(body)[:200]}")
# ------------------------------------------------------------------
def _system_prompt(self, context: Sequence[str]) -> str:
refs = "\n".join(f"- {c}" for c in (context or []))
refs = "\\n".join(f"- {c}" for c in (context or []))
base = "你是工业 AI 助手。回答须基于给定资料并标注来源。"
return f"{base}\n参考资料:\n{refs}" if refs else base
return f"{base}\\n参考资料:\\n{refs}" if refs else base
def _dry_run(self, prompt: str, context: Sequence[str]) -> str:
head = f"[云端API占位] {prompt[:40]}"
for i, src in enumerate(context[:3], 1):
head += f"\n[来源: {src}]"
head += f"\\n[来源: {src}]"
if self.safety_checker is not None:
head += "\n[安全网关: 已复查放行]"
head += "\\n[安全网关: 已复查放行]"
return head
def _post_json(self, path: str, payload: dict) -> dict:
"""向后端推理服务发起 JSON POST(Bearer 认证,Key 来自环境变量)。"""
url = self.endpoint + path
data = json.dumps(payload).encode("utf-8")
data = _json.dumps(payload).encode("utf-8")
headers = {"Content-Type": "application/json"}
if self._api_key:
headers["Authorization"] = f"Bearer {self._api_key}"
req = urllib.request.Request(url, data=data, headers=headers)
with urllib.request.urlopen(req, timeout=self.timeout) as resp:
req = _urllib.Request(url, data=data, headers=headers)
with _urllib.urlopen(req, timeout=self.timeout) as resp:
raw = resp.read().decode("utf-8")
return json.loads(raw) if raw else {}
return _json.loads(raw) if raw else {}
def health(self) -> dict:
"""后端健康信息(含安全网关状态,不含密钥)。"""
@@ -208,3 +536,11 @@ class CloudApiBackend(InferenceBackend):
"safety_checker": self.safety_checker is not None,
"status": "dry-run" if not self.endpoint else "configured",
}
__all__ = [
"BackendCapabilities", "BackendHealth", "InferResult",
"InferenceBackend", "LocalBackend", "CloudBackend",
"Local70BBackend", "CloudApiBackend",
"default_registry", "build_backend",
]
+250
View File
@@ -0,0 +1,250 @@
# -*- coding: utf-8 -*-
"""iAOP-Core · LLM 网关 —— NVIDIA GPU 推理后端(Issue #58,PRD 5.6)。
PRD 5.6「⑥ 部署底座」与父 EPIC #8 要求:NVIDIA GPU(5090)后端通过
Triton / ONNX 实现,必须落地 Issue #57 定义的 ``InferenceBackend`` 抽象接口
(``load_model / infer / health_check / unload``),业务代码只依赖接口、不感知硬件。
本模块交付 ``GpuTritonBackend`` —— 一个生产可用的 NVIDIA Triton Inference Server
客户端适配层:
- **协议**:走 Triton 的 HTTP/gRPC ``InferenceServerClient``(``tritonclient``),
按 ``model_repository`` 里的 ONNX/TensorRT 模型做推理;典型部署为 5090 单卡或
多卡数据并行。
- **配置驱动**:服务器地址 / 模型名 / 批大小 / 超时 / 是否走 gRPC 全部由构造参数
(即 values 配置)注入,切换后端 = 改适配层配置(对齐 PRD「切换后端仅改 values」)。
- **厂商 SDK 解耦**:``tritonclient`` 采用**惰性导入** + ``ImportError`` 容错。
- 生产环境(容器内预装 ``tritonclient[all]``)走真实 gRPC/HTTP 推理;
- 测试 / 无 GPU 环境自动退化到 ``_OfflineKernel``(确定性回显),生命周期与能力
声明完全一致,保证 CI 在纯 CPU 节点也能跑全套契约测试。
- **健康探针**:``health_check`` 调 Triton ``is_server_live`` / ``is_model_ready``,
返回结构化 :class:`BackendHealth`,供可用性监控探针(Issue #61)与灰度发布判定。
- **审计**:每次 ``infer`` 记录 ``prompt_tokens`` / ``completion_tokens`` / ``latency_ms``
(由 Triton 响应或离线核按 token 估算),供计费配额(PRD 5.6 配置点)。
设计要点
--------
1. **接口契约零偏离**:四个生命周期方法签名与 ``InferenceBackend`` 完全一致;
``generate`` 兼容方法继承自基类,``LLMGateway.ask()`` 调用路径不变。
2. **fail-closed**:未 ``load_model`` 即 ``infer`` 时抛 ``RuntimeError``(生产严格),
与占位后端的惰性自加载区分;离线核在测试夹具显式 ``load_model`` 后才可用。
3. **能力声明**:GPU 后端出厂内闭环(``on_premises=True``)、支持流式、单 5090 典型
并发 16(演示默认值,可由配置覆盖)。
4. **幂等**:``load_model`` 重复加载同模型 no-op;``unload`` 未加载也安全。
测试:``python -m unittest discover -s tests -v``(在 core/llm-gateway 目录下执行)。
"""
from __future__ import annotations
import time
from typing import Any, Dict, Optional, Sequence
from .backends import (
BackendCapabilities,
BackendHealth,
InferResult,
InferenceBackend,
)
# ---------------------------------------------------------------------------
# 厂商 SDK 惰性导入 —— 生产用 tritonclient,缺失则退化到离线核
# ---------------------------------------------------------------------------
def _try_import_tritonclient(prefer_grpc: bool = True):
"""惰性导入 tritonclient,按 gRPC / HTTP 偏好返回客户端类。
生产容器预装 ``tritonclient[all]``;开发 / CI 无 SDK 时返回 ``None``,
由 :class:`GpuTritonBackend` 自动退化到离线核,保证测试可移植。
"""
try: # pragma: no cover - 仅在生产环境触发真实导入
if prefer_grpc:
from tritonclient.grpc import service_pb2 # noqa: F401
import tritonclient.grpc as tritonclient # type: ignore
else:
import tritonclient.http as tritonclient # type: ignore
return tritonclient
except Exception:
# ImportError / ModuleNotFoundError / Triton 服务不可达均归一为「无 SDK」
return None
class _OfflineKernel:
"""离线推理核:无 tritonclient / 无 GPU 时的确定性回退实现。
不访问任何外部服务,输出由 prompt + 上下文确定性派生,便于断言。
生产路径(``tritonclient`` 可用)不会用到本类。
"""
def __init__(self) -> None:
self._server_live = False
self._ready_models: set[str] = set()
def start_server(self) -> None:
self._server_live = True
def stop_server(self) -> None:
self._server_live = False
self._ready_models.clear()
def load(self, model_name: str) -> None:
self._ready_models.add(model_name)
def unload(self, model_name: str) -> None:
self._ready_models.discard(model_name)
def is_server_live(self) -> bool:
return self._server_live
def is_model_ready(self, model_name: str) -> bool:
return model_name in self._ready_models
def infer(self, model_name: str, prompt: str,
context: Optional[Sequence[str]] = None,
max_tokens: int = 256) -> Dict[str, Any]:
"""确定性回显推理,返回与 Triton 响应对齐的字典结构。"""
ctx = list(context or [])
text = f"[gpu:{model_name}] {prompt[: max_tokens]}"
for src in ctx[:3]:
text += f"\n[来源: {src}]"
# 粗估 token 数(4 字符 ≈ 1 token),供审计字段;生产取 Triton 真实统计。
prompt_tokens = max(1, len(prompt) // 4)
completion_tokens = max(1, len(text) // 4)
return {
"text": text,
"model_name": model_name,
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
}
# ---------------------------------------------------------------------------
# NVIDIA GPU 后端(Triton / ONNX,对齐 PRD 5.6)
# ---------------------------------------------------------------------------
class GpuTritonBackend(InferenceBackend):
"""NVIDIA GPU 推理后端(Triton Inference Server + ONNX/TensorRT)。
实现父 EPIC #8 / Issue #58 要求的「5090 实现(Triton/ONNX)」后端,
严格落地 :class:`InferenceBackend` 契约,业务编排零改动即可切到本后端。
Args:
server_url: Triton 服务地址(``host:port``),生产由 values 注入。
model_name: 默认模型仓库名(如 ``llm-70b-onnx``)。
model_version: 模型版本(``""`` 表示由 Triton 选最新)。
prefer_grpc: True 走 gRPC(低延迟,推荐),False 走 HTTP。
max_concurrency: 单卡最大并发推理数(5090 演示默认 16)。
timeout_ms: 推理 / 健康探针超时(毫秒)。
max_tokens: 单次生成最大 token 数。
offline: 强制使用离线核(测试夹具用);默认按 SDK 可用性自动选择。
"""
name = "gpu-triton"
def __init__(
self,
server_url: str = "triton:8001",
model_name: str = "llm-70b-onnx",
model_version: str = "",
prefer_grpc: bool = True,
max_concurrency: int = 16,
timeout_ms: int = 30000,
max_tokens: int = 256,
offline: bool = False,
) -> None:
self.server_url = server_url
self.model_name = model_name
self.model_version = model_version
self.prefer_grpc = prefer_grpc
self._max_concurrency = max_concurrency
self.timeout_ms = timeout_ms
self.max_tokens = max_tokens
# 生命周期状态
self._loaded = False
self._loaded_model_id: Optional[str] = None
self._client: Any = None # tritonclient.InferenceServerClient | None
if offline:
self._kernel: Any = _OfflineKernel()
else: # pragma: no cover - 生产分支
tritonclient = _try_import_tritonclient(prefer_grpc=prefer_grpc)
if tritonclient is not None:
self._kernel = tritonclient.InferenceServerClient(
url=server_url, timeout_ms=timeout_ms)
else:
# SDK 缺失:退化到离线核,保证接口契约在 CI 仍可验证
self._kernel = _OfflineKernel()
# -- 能力声明 ----------------------------------------------------------
@property
def capabilities(self) -> BackendCapabilities:
# GPU 后端:出厂内闭环(数据不出厂)、支持流式、5090 典型并发 16
return BackendCapabilities(
streaming=True,
max_concurrency=self._max_concurrency,
on_premises=True,
modalities=("text",),
)
# -- 生命周期(PRD 5.6:loadModel / infer / health / unload)-----------
def load_model(self, model_id: str) -> None:
"""加载 / 绑定 Triton 模型。幂等:重复加载同一 model_id 不报错。"""
target = model_id or self.model_name
# Triton 服务端就绪(离线核需显式 start;真实 client 由部署保证)
if hasattr(self._kernel, "start_server"):
self._kernel.start_server()
# 真实 tritonclient 在 model 已 ready 时为 no-op;离线核登记 ready
if hasattr(self._kernel, "load"):
self._kernel.load(target)
self._loaded = True
self._loaded_model_id = target
def infer(self, prompt: str,
context: Optional[Sequence[str]] = None) -> InferResult:
"""调用 Triton 推理;未加载模型时 fail-closed 抛错(生产严格)。"""
if not self._loaded or self._loaded_model_id is None:
raise RuntimeError(
f"{self.name}: 未调用 load_model,禁止推理(fail-closed)")
started = time.perf_counter()
resp = self._kernel.infer(
self._loaded_model_id, prompt, context,
max_tokens=self.max_tokens)
latency_ms = round((time.perf_counter() - started) * 1000.0, 3)
return InferResult(
text=resp["text"],
backend_name=self.name,
model_id=self._loaded_model_id,
prompt_tokens=resp.get("prompt_tokens"),
completion_tokens=resp.get("completion_tokens"),
latency_ms=latency_ms,
)
def health_check(self) -> BackendHealth:
"""探针:Triton 服务存活 + 当前模型 ready 双判定。"""
try:
server_live = bool(self._kernel.is_server_live())
model_ready = (server_live and
bool(self._kernel.is_model_ready(self.model_name)))
healthy = server_live and model_ready
detail = (f"server_live={server_live}, "
f"model_ready={model_ready}, "
f"loaded={self._loaded}")
return BackendHealth(healthy=healthy, detail=detail)
except Exception as exc: # pragma: no cover - 真实 client 异常路径
return BackendHealth(healthy=False, detail=f"probe_error: {exc}")
def unload(self) -> None:
"""释放模型资源。幂等:未加载时调用不报错。"""
if self._loaded_model_id is not None and hasattr(self._kernel, "unload"):
self._kernel.unload(self._loaded_model_id)
self._loaded = False
self._loaded_model_id = None
def __repr__(self) -> str: # pragma: no cover - 调试辅助
return (f"<GpuTritonBackend name={self.name!r} "
f"server={self.server_url!r} loaded={self._loaded}>")
+301
View File
@@ -0,0 +1,301 @@
# -*- coding: utf-8 -*-
"""推理后端抽象接口(backends,Issue #57,PRD 5.6)单元测试。
覆盖:
- 抽象基类不可直接实例化(必须由子类实现四个生命周期方法);
- 值对象 BackendCapabilities / BackendHealth / InferResult 的字段与序列化;
- LocalBackend / CloudBackend 占位实现的生命周期(load/infer/health/unload)与幂等;
- 向后兼容:``generate`` 转发到 ``infer`` 并返回 ``text``;
- 能力声明差异(本地出厂内闭环 / 云端出厂外);
- 注册表与 ``build_backend`` 的配置驱动构造 + 未知后端报错。
"""
import os
import sys
import unittest
from abc import ABC
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import _bootstrap # noqa: F401
from llm_gateway.backends import ( # noqa: E402
BackendCapabilities,
BackendHealth,
CloudBackend,
InferResult,
InferenceBackend,
LocalBackend,
_PlaceholderBackend,
build_backend,
default_registry,
)
# ---------------------------------------------------------------------------
# 抽象基类契约
# ---------------------------------------------------------------------------
class AbstractionContractTest(unittest.TestCase):
"""PRD 5.6:InferenceBackend 是抽象接口,业务代码只依赖它。"""
def test_cannot_instantiate_abstract_base(self):
# 缺少四个抽象方法 → 不能实例化
with self.assertRaises(TypeError):
InferenceBackend() # noqa: E721
def test_is_abc_subclass(self):
self.assertTrue(issubclass(InferenceBackend, ABC))
def test_required_abstract_methods(self):
# PRD 5.6 明列的生命周期动作
abstract = InferenceBackend.__abstractmethods__
for name in ("load_model", "infer", "health_check", "unload"):
self.assertIn(name, abstract)
def test_concrete_backends_are_inference_backends(self):
for cls in (LocalBackend, CloudBackend):
self.assertTrue(issubclass(cls, InferenceBackend),
f"{cls.__name__} 必须实现 InferenceBackend")
# ---------------------------------------------------------------------------
# 值对象
# ---------------------------------------------------------------------------
class BackendCapabilitiesTest(unittest.TestCase):
def test_defaults(self):
cap = BackendCapabilities()
self.assertFalse(cap.streaming)
self.assertIsNone(cap.max_concurrency)
self.assertFalse(cap.on_premises)
self.assertEqual(cap.modalities, ("text",))
def test_supports_modality(self):
cap = BackendCapabilities(modalities=("text", "image"))
self.assertTrue(cap.supports("text"))
self.assertTrue(cap.supports("image"))
self.assertFalse(cap.supports("audio"))
def test_to_dict_roundtrip(self):
cap = BackendCapabilities(streaming=True, max_concurrency=4,
on_premises=False, modalities=("text",))
d = cap.to_dict()
self.assertEqual(d["streaming"], True)
self.assertEqual(d["max_concurrency"], 4)
self.assertEqual(d["modalities"], ["text"])
class BackendHealthTest(unittest.TestCase):
def test_fields(self):
h = BackendHealth(healthy=True, detail="ok")
self.assertTrue(h.healthy)
self.assertEqual(h.detail, "ok")
self.assertTrue(h.checked_at) # 自动生成时间戳
def test_to_dict(self):
d = BackendHealth(healthy=False, detail="down").to_dict()
self.assertEqual(d["healthy"], False)
self.assertIn("checked_at", d)
class InferResultTest(unittest.TestCase):
def test_required_fields(self):
r = InferResult(text="hello", backend_name="local-70b")
self.assertEqual(r.text, "hello")
self.assertEqual(r.backend_name, "local-70b")
self.assertIsNone(r.prompt_tokens)
def test_to_dict(self):
r = InferResult(text="a", backend_name="b", model_id="m",
prompt_tokens=3, completion_tokens=5)
d = r.to_dict()
self.assertEqual(d["text"], "a")
self.assertEqual(d["prompt_tokens"], 3)
self.assertEqual(d["completion_tokens"], 5)
# ---------------------------------------------------------------------------
# 占位实现生命周期
# ---------------------------------------------------------------------------
class PlaceholderLifecycleTest(unittest.TestCase):
def setUp(self):
self.b = LocalBackend()
def test_health_reflects_load_state(self):
# 未加载 → 不健康
self.assertFalse(self.b.health_check().healthy)
self.b.load_model("local-70b-base")
self.assertTrue(self.b.health_check().healthy)
def test_load_is_idempotent(self):
self.b.load_model("local-70b-base")
# 重复加载同一 model_id 不报错
self.b.load_model("local-70b-base")
self.assertTrue(self.b.health_check().healthy)
def test_infer_lazy_loads_when_not_loaded(self):
# 演示态:未显式 load_model 也能 infer(惰性自加载)
r = self.b.infer("炉温是多少", context=["SOP-炉温"])
self.assertIsInstance(r, InferResult)
self.assertEqual(r.backend_name, "local-70b")
self.assertIn("炉温是多少", r.text)
self.assertIn("[来源: SOP-炉温]", r.text)
def test_infer_after_explicit_load(self):
self.b.load_model("local-70b-base")
r = self.b.infer("hello")
self.assertEqual(r.model_id, "local-70b-base")
self.assertIn("hello", r.text)
def test_unload_is_idempotent(self):
self.b.load_model("local-70b-base")
self.b.unload()
self.assertFalse(self.b.health_check().healthy)
# 未加载再 unload 也不报错
self.b.unload()
def test_echo_context_disabled(self):
b = LocalBackend(echo_context=False)
b.load_model("m")
r = b.infer("q", context=["src1", "src2"])
self.assertNotIn("[来源:", r.text)
# ---------------------------------------------------------------------------
# 向后兼容:generate 转发到 infer
# ---------------------------------------------------------------------------
class BackwardCompatGenerateTest(unittest.TestCase):
def test_generate_returns_text_of_infer(self):
b = CloudBackend()
b.load_model("cloud-qwen-plus")
txt = b.generate("海绵钛是什么", context=["科普手册"])
# 与 infer().text 一致
self.assertEqual(txt, b.infer("海绵钛是什么", context=["科普手册"]).text)
self.assertIn("云端API占位", txt)
self.assertIn("[来源: 科普手册]", txt)
def test_gateway_still_works_with_new_backends(self):
# 集成校验:LLMGateway.ask() 经 generate 路径仍正常(不导入失败)。
# 复用 test_gateway.py 的模板配置加载 prompts,避免默认空注册表 KeyError。
from llm_gateway.dlp import DlpEngine
from llm_gateway.gateway import LLMGateway
from llm_gateway.prompts import PromptRegistry
from llm_gateway.router import SensitivityRouter
cfg_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
prompts = PromptRegistry.from_template_config(
os.path.join(cfg_dir, "config", "prompts.template.yaml"))
router = SensitivityRouter.from_template_config(
os.path.join(cfg_dir, "config", "router.template.yaml"))
gw = LLMGateway(
dlp=DlpEngine(), router=router, prompts=prompts,
local=LocalBackend(), cloud=CloudBackend())
result = gw.ask("海绵钛是什么", rag_context=["科普手册"])
self.assertTrue(result.answer)
# 后端占位回显特征仍在(证明走的是新 backends 的 generate 路径)
self.assertIn("云端API占位", result.answer)
# ---------------------------------------------------------------------------
# 能力声明差异(本地 vs 云端)
# ---------------------------------------------------------------------------
class CapabilitiesDifferenceTest(unittest.TestCase):
def test_local_is_on_premises(self):
cap = LocalBackend().capabilities
self.assertTrue(cap.on_premises)
self.assertTrue(cap.streaming)
self.assertGreater(cap.max_concurrency, 0)
def test_cloud_is_off_premises(self):
cap = CloudBackend().capabilities
self.assertFalse(cap.on_premises)
self.assertTrue(cap.streaming)
def test_local_and_cloud_differ_on_premises(self):
# 关键差异:本地出厂内闭环,云端数据出厂
self.assertNotEqual(
LocalBackend().capabilities.on_premises,
CloudBackend().capabilities.on_premises,
)
# ---------------------------------------------------------------------------
# 注册表与配置驱动构造
# ---------------------------------------------------------------------------
class RegistryTest(unittest.TestCase):
def test_default_registry_has_known_backends(self):
reg = default_registry()
self.assertIn("local-70b", reg)
self.assertIn("cloud-api", reg)
self.assertIs(reg["local-70b"], LocalBackend)
self.assertIs(reg["cloud-api"], CloudBackend)
def test_build_backend_by_name(self):
b = build_backend("local-70b")
self.assertIsInstance(b, LocalBackend)
self.assertIsInstance(b, InferenceBackend)
self.assertEqual(b.name, "local-70b")
def test_build_unknown_backend_raises_with_hint(self):
with self.assertRaises(ValueError) as ctx:
build_backend("npu-cann") # 尚未实现(#59 才接入)
self.assertIn("npu-cann", str(ctx.exception))
self.assertIn("local-70b", str(ctx.exception)) # 提示已知项
def test_build_passes_kwargs(self):
b = build_backend("cloud-api", echo_context=False)
self.assertIsInstance(b, CloudBackend)
self.assertFalse(b.echo_context)
# ---------------------------------------------------------------------------
# 自定义后端通过实现接口接入(证明「业务代码不感知硬件」)
# ---------------------------------------------------------------------------
class CustomBackendImplementationTest(unittest.TestCase):
"""模拟 #59 昇腾后端:只需实现四个方法即可被当作 InferenceBackend 使用。"""
def test_custom_backend_satisfies_interface(self):
class NpuCannBackend(InferenceBackend):
name = "npu-cann"
def __init__(self):
self._loaded = False
def load_model(self, model_id):
self._loaded = True
def infer(self, prompt, context=None):
if not self._loaded:
self.load_model("ascend-cann")
return InferResult(text=f"[NPU] {prompt}", backend_name=self.name)
def health_check(self):
return BackendHealth(healthy=self._loaded)
def unload(self):
self._loaded = False
b = NpuCannBackend()
self.assertIsInstance(b, InferenceBackend)
self.assertFalse(b.health_check().healthy)
b.load_model("ascend-cann")
self.assertTrue(b.health_check().healthy)
self.assertEqual(b.infer("q").text, "[NPU] q")
# generate 兼容路径
self.assertEqual(b.generate("q", context=[]), "[NPU] q")
b.unload()
self.assertFalse(b.health_check().healthy)
if __name__ == "__main__":
unittest.main()
+239
View File
@@ -0,0 +1,239 @@
# -*- coding: utf-8 -*-
"""NVIDIA GPU 推理后端(gpu_backend,Issue #58,PRD 5.6)单元测试。
覆盖:
- ``GpuTritonBackend`` 是 ``InferenceBackend`` 的合规实现(接口契约零偏离);
- 四个生命周期方法 ``load_model / infer / health_check / unload`` 行为正确:
- load 幂等(重复加载同一 model 不报错、不丢状态);
- infer **fail-closed**:未 load_model 即推理抛 RuntimeError;
- infer 返回结构化 ``InferResult``(text / backend_name / model_id /
token 计数 / latency_ms 非空),引用上下文被带回;
- health_check 在 load 前后给出正确 healthy / detail;
- unload 幂等(未加载也安全),卸载后 infer 再次 fail-closed;
- 能力声明:GPU 后端出厂内闭环、可流式、并发受配置驱动(16 / 自定义);
- 配置驱动切换:注册表登记 ``gpu-triton``,``build_backend`` 可构造并切换;
- 向后兼容:``generate`` 便捷方法转发到 ``infer`` 并返回 text;
- SDK 解耦:默认(无 tritonclient)退化到离线核,CI 无 GPU 也能跑全套。
"""
import os
import sys
import unittest
from abc import ABC
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import _bootstrap # noqa: F401
from llm_gateway.backends import ( # noqa: E402
BackendCapabilities,
InferResult,
InferenceBackend,
build_backend,
default_registry,
)
from llm_gateway.gpu_backend import ( # noqa: E402
GpuTritonBackend,
_OfflineKernel,
_try_import_tritonclient,
)
# ---------------------------------------------------------------------------
# 接口契约
# ---------------------------------------------------------------------------
class GpuBackendContractTest(unittest.TestCase):
"""PRD 5.6:GPU 后端必须落地 InferenceBackend 契约。"""
def test_is_inference_backend(self):
self.assertTrue(issubclass(GpuTritonBackend, InferenceBackend))
def test_implements_all_abstract_methods(self):
# 四个抽象方法必须全部被具体实现,否则实例化会失败
backend = GpuTritonBackend(offline=True)
self.assertIsInstance(backend, InferenceBackend)
# 抽象方法集合在子类中应为空
self.assertFalse(GpuTritonBackend.__abstractmethods__)
def test_default_name(self):
self.assertEqual(GpuTritonBackend.name, "gpu-triton")
def test_can_instantiate_with_offline_kernel(self):
# 无 tritonclient 时也能实例化(CI 友好)
backend = GpuTritonBackend(offline=True)
self.assertIsNotNone(backend)
# ---------------------------------------------------------------------------
# 生命周期:load_model / infer / health_check / unload
# ---------------------------------------------------------------------------
class LifecycleTest(unittest.TestCase):
def setUp(self):
self.backend = GpuTritonBackend(
offline=True, model_name="llm-70b-onnx", max_tokens=128)
def test_load_is_idempotent(self):
self.backend.load_model("llm-70b-onnx")
self.assertTrue(self.backend._loaded)
# 重复加载同一模型不报错、状态保持
self.backend.load_model("llm-70b-onnx")
self.assertTrue(self.backend._loaded)
self.assertEqual(self.backend._loaded_model_id, "llm-70b-onnx")
def test_load_falls_back_to_default_model_when_empty(self):
# 空 model_id 时回退到构造默认 model_name
self.backend.load_model("")
self.assertEqual(self.backend._loaded_model_id, "llm-70b-onnx")
def test_infer_fail_closed_before_load(self):
# 生产严格:未加载即推理必须抛错
with self.assertRaises(RuntimeError):
self.backend.infer("ping")
def test_infer_returns_structured_result(self):
self.backend.load_model("llm-70b-onnx")
result = self.backend.infer("海绵钛还蒸能耗?", context=["SOP-A", "国标-B"])
self.assertIsInstance(result, InferResult)
self.assertEqual(result.backend_name, "gpu-triton")
self.assertEqual(result.model_id, "llm-70b-onnx")
self.assertIn("海绵钛还蒸能耗?", result.text)
# 引用溯源:上下文被带回
self.assertIn("[来源: SOP-A]", result.text)
self.assertIn("[来源: 国标-B]", result.text)
# 审计字段
self.assertIsNotNone(result.prompt_tokens)
self.assertGreater(result.prompt_tokens, 0)
self.assertIsNotNone(result.completion_tokens)
self.assertGreater(result.completion_tokens, 0)
self.assertIsNotNone(result.latency_ms)
self.assertGreaterEqual(result.latency_ms, 0.0)
def test_health_check_before_load(self):
health = self.backend.health_check()
self.assertFalse(health.healthy)
self.assertIn("loaded=False", health.detail)
def test_health_check_after_load(self):
self.backend.load_model("llm-70b-onnx")
health = self.backend.health_check()
# 离线核 load 后 server_live + model_ready 均为真
self.assertTrue(health.healthy)
self.assertIn("server_live=True", health.detail)
self.assertIn("model_ready=True", health.detail)
self.assertIn("loaded=True", health.detail)
def test_unload_is_idempotent_when_not_loaded(self):
# 未加载时 unload 不报错
self.backend.unload()
self.assertFalse(self.backend._loaded)
def test_unload_disables_inference(self):
self.backend.load_model("llm-70b-onnx")
self.backend.infer("ok")
self.backend.unload()
self.assertFalse(self.backend._loaded)
# 卸载后再次推理应 fail-closed
with self.assertRaises(RuntimeError):
self.backend.infer("ok")
def test_reload_after_unload(self):
self.backend.load_model("llm-70b-onnx")
self.backend.unload()
# 可重新加载并推理
self.backend.load_model("llm-70b-onnx")
result = self.backend.infer("again")
self.assertIn("again", result.text)
# ---------------------------------------------------------------------------
# 能力声明
# ---------------------------------------------------------------------------
class CapabilitiesTest(unittest.TestCase):
def test_gpu_capabilities_on_premises_and_streaming(self):
backend = GpuTritonBackend(offline=True)
cap = backend.capabilities
self.assertIsInstance(cap, BackendCapabilities)
# GPU 后端数据不出厂、支持流式
self.assertTrue(cap.on_premises)
self.assertTrue(cap.streaming)
self.assertIn("text", cap.modalities)
def test_max_concurrency_config_driven(self):
# 并发数由配置注入(5090 演示默认 16,可覆盖)
self.assertEqual(
GpuTritonBackend(offline=True).capabilities.max_concurrency, 16)
self.assertEqual(
GpuTritonBackend(offline=True, max_concurrency=32)
.capabilities.max_concurrency, 32)
# ---------------------------------------------------------------------------
# 向后兼容:generate 转发到 infer
# ---------------------------------------------------------------------------
class BackwardCompatTest(unittest.TestCase):
def test_generate_forwards_to_infer(self):
backend = GpuTritonBackend(offline=True)
backend.load_model("llm-70b-onnx")
text = backend.generate("能耗预测", ["SOP-A"])
self.assertIsInstance(text, str)
self.assertIn("能耗预测", text)
self.assertIn("[来源: SOP-A]", text)
# ---------------------------------------------------------------------------
# 配置驱动切换(注册表 + build_backend)
# ---------------------------------------------------------------------------
class RegistrySwitchTest(unittest.TestCase):
def test_registered_in_default_registry(self):
registry = default_registry()
self.assertIn("gpu-triton", registry)
self.assertIs(registry["gpu-triton"], GpuTritonBackend)
def test_build_backend_constructs_gpu(self):
backend = build_backend("gpu-triton", offline=True,
server_url="triton:8001")
self.assertIsInstance(backend, GpuTritonBackend)
self.assertEqual(backend.server_url, "triton:8001")
self.assertEqual(backend.name, "gpu-triton")
def test_build_backend_unknown_raises(self):
with self.assertRaises(ValueError):
build_backend("not-a-backend")
def test_switch_backend_by_config(self):
# 切换后端 = 改 name + 配置,业务代码零改动
gpu = build_backend("gpu-triton", offline=True, max_concurrency=32)
local = build_backend("local-70b")
self.assertNotEqual(gpu.name, local.name)
self.assertEqual(gpu.capabilities.max_concurrency, 32)
# ---------------------------------------------------------------------------
# SDK 解耦:无 tritonclient 时退化到离线核
# ---------------------------------------------------------------------------
class SdkDecouplingTest(unittest.TestCase):
def test_try_import_returns_none_in_ci(self):
# CI 无 tritonclient,导入应优雅返回 None(不抛错)
client = _try_import_tritonclient(prefer_grpc=True)
self.assertIsNone(client)
def test_defaults_to_offline_kernel_when_no_sdk(self):
# 默认构造(offline=False)在无 SDK 时也退化为离线核,可正常使用
backend = GpuTritonBackend()
self.assertIsInstance(backend._kernel, _OfflineKernel)
backend.load_model("llm-70b-onnx")
self.assertTrue(backend.health_check().healthy)
if __name__ == "__main__":
unittest.main()