feat(#58): GPU 后端实现(NVIDIA Triton/ONNX,PRD 5.6 推理后端可插拔) #100
@@ -16,6 +16,9 @@
|
|||||||
高利害信度阈值 → 人工确认。
|
高利害信度阈值 → 人工确认。
|
||||||
- gateway 混合网关主编排(EPIC #6 主体):路由 → 生成 → 溯源校验 →
|
- gateway 混合网关主编排(EPIC #6 主体):路由 → 生成 → 溯源校验 →
|
||||||
DLP 出站防线,端到端闭环。
|
DLP 出站防线,端到端闭环。
|
||||||
|
- backends 推理后端抽象(Issue #57,PRD 5.6):``InferenceBackend`` 抽象接口
|
||||||
|
(``load_model / infer / health_check / unload``),5090 实现(Triton/ONNX)
|
||||||
|
与昇腾实现(ACL/CANN)均实现该接口;业务代码仅依赖接口,不感知硬件。
|
||||||
|
|
||||||
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
|
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
|
||||||
"""
|
"""
|
||||||
@@ -45,10 +48,17 @@ from .hallucination import (
|
|||||||
GuardVerdict,
|
GuardVerdict,
|
||||||
HallucinationGuard,
|
HallucinationGuard,
|
||||||
)
|
)
|
||||||
|
from .backends import (
|
||||||
|
BackendCapabilities,
|
||||||
|
BackendHealth,
|
||||||
|
InferResult,
|
||||||
|
InferenceBackend,
|
||||||
|
build_backend,
|
||||||
|
default_registry,
|
||||||
|
)
|
||||||
from .gateway import (
|
from .gateway import (
|
||||||
CloudBackend,
|
CloudBackend,
|
||||||
GatewayResult,
|
GatewayResult,
|
||||||
InferenceBackend,
|
|
||||||
LLMGateway,
|
LLMGateway,
|
||||||
LocalBackend,
|
LocalBackend,
|
||||||
)
|
)
|
||||||
@@ -62,7 +72,10 @@ __all__ = [
|
|||||||
"PromptVersion", "PromptChange", "PromptRegistry", "validate_semver",
|
"PromptVersion", "PromptChange", "PromptRegistry", "validate_semver",
|
||||||
# hallucination
|
# hallucination
|
||||||
"GuardVerdict", "HallucinationGuard",
|
"GuardVerdict", "HallucinationGuard",
|
||||||
|
# backends (Issue #57)
|
||||||
|
"BackendCapabilities", "BackendHealth", "InferResult", "InferenceBackend",
|
||||||
|
"build_backend", "default_registry",
|
||||||
# gateway
|
# gateway
|
||||||
"InferenceBackend", "LocalBackend", "CloudBackend",
|
"LocalBackend", "CloudBackend",
|
||||||
"GatewayResult", "LLMGateway",
|
"GatewayResult", "LLMGateway",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,304 @@
|
|||||||
|
# -*- coding: utf-8 -*-
|
||||||
|
"""iAOP-Core · LLM 网关 —— 推理后端抽象接口(Issue #57,PRD 5.6)。
|
||||||
|
|
||||||
|
PRD 5.6「⑥ 部署底座」明确要求:
|
||||||
|
|
||||||
|
定义统一 ``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
|
||||||
|
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Dict, Iterator, List, Optional, Sequence
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# 值对象:能力声明 / 健康状态 / 推理结果
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
@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)
|
||||||
+18
-61
@@ -12,12 +12,15 @@
|
|||||||
- `router`(SensitivityRouter):敏感度分级路由(local / cloud / block);
|
- `router`(SensitivityRouter):敏感度分级路由(local / cloud / block);
|
||||||
- `prompts`(PromptRegistry):提示词模板版本绑定(可复现);
|
- `prompts`(PromptRegistry):提示词模板版本绑定(可复现);
|
||||||
- `guard`(HallucinationGuard):引用溯源 + 信度阈值 → 人工确认;
|
- `guard`(HallucinationGuard):引用溯源 + 信度阈值 → 人工确认;
|
||||||
- `backends`(LocalBackend / CloudBackend):推理后端抽象(可注入)。
|
- `backends`(InferenceBackend / LocalBackend / CloudBackend):推理后端抽象
|
||||||
|
(可注入)。接口定义已提炼到 `backends.py`(Issue #57,对齐 PRD 5.6)。
|
||||||
|
|
||||||
设计说明:
|
设计说明:
|
||||||
- 本版提供**编排闭环 + 后端抽象接口**,本地 70B / 云端 API 的具体接入
|
- 本版提供**编排闭环 + 后端抽象接口**,本地 70B / 云端 API 的具体接入
|
||||||
由子任务 #44 / #45 实现;`LocalBackend` / `CloudBackend` 默认内置一个
|
由子任务 #44 / #45 实现;`LocalBackend` / `CloudBackend` 默认内置一个
|
||||||
最小实现(返回固定占位答案 + 回显引用),供端到端测试与演示。
|
最小实现(返回固定占位答案 + 回显引用),供端到端测试与演示。
|
||||||
|
- 推理后端契约(`loadModel / infer / health_check / unload`)见 `backends.py`,
|
||||||
|
本模块仅消费其 `generate` / `name`,业务代码不感知具体硬件。
|
||||||
|
|
||||||
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
|
测试:`python -m unittest discover -s tests -v`(在 core/llm-gateway 目录下执行)。
|
||||||
"""
|
"""
|
||||||
@@ -26,71 +29,25 @@ from __future__ import annotations
|
|||||||
import uuid
|
import uuid
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import Callable, Dict, List, Optional, Sequence
|
from typing import Dict, List, Optional, Sequence
|
||||||
|
|
||||||
from .dlp import DlpEngine
|
from .dlp import DlpEngine
|
||||||
from .router import RouteDecision, RouteTarget, SensitivityRouter
|
from .router import RouteDecision, RouteTarget, SensitivityRouter
|
||||||
from .prompts import PromptRegistry
|
from .prompts import PromptRegistry
|
||||||
from .hallucination import GuardVerdict, HallucinationGuard
|
from .hallucination import GuardVerdict, HallucinationGuard
|
||||||
|
# 推理后端抽象(Issue #57):契约定义在 backends.py,这里仅做再导出,
|
||||||
# ---------------------------------------------------------------------------
|
# 保持 ``from .gateway import InferenceBackend/LocalBackend/CloudBackend`` 的
|
||||||
# 推理后端抽象(Issue #44 / #45 将实现具体后端,业务代码只依赖本接口)
|
# 向后兼容(既有 import 路径与 ``LLMGateway`` 依赖均不变)。
|
||||||
# ---------------------------------------------------------------------------
|
from .backends import (
|
||||||
|
BackendCapabilities,
|
||||||
|
BackendHealth,
|
||||||
class InferenceBackend:
|
CloudBackend,
|
||||||
"""推理后端接口抽象(对齐 PRD 5.6 InferenceBackend 思想)。
|
InferResult,
|
||||||
|
InferenceBackend,
|
||||||
业务代码只依赖本接口,不感知具体硬件/厂商;切换后端 = 换实现。
|
LocalBackend,
|
||||||
子任务 #44(本地 70B)、#45(云端 Qwen/DeepSeek)将各自实现本接口。
|
build_backend,
|
||||||
"""
|
default_registry,
|
||||||
|
)
|
||||||
name: str = "base"
|
|
||||||
|
|
||||||
def generate(self, prompt: str, context: Sequence[str]) -> str:
|
|
||||||
"""根据 prompt 与 RAG 上下文生成回答。子类实现。"""
|
|
||||||
raise NotImplementedError
|
|
||||||
|
|
||||||
|
|
||||||
class LocalBackend(InferenceBackend):
|
|
||||||
"""本地 70B 后端占位实现:数据不出厂(敏感/核心走此通道)。
|
|
||||||
|
|
||||||
子任务 #44 将替换为真实本地模型推理封装(vLLM/TGI 等)。
|
|
||||||
"""
|
|
||||||
|
|
||||||
name = "local-70b"
|
|
||||||
|
|
||||||
def __init__(self, echo_context: bool = True) -> None:
|
|
||||||
self.echo_context = echo_context
|
|
||||||
|
|
||||||
def generate(self, prompt: str, context: Sequence[str]) -> str:
|
|
||||||
head = f"[本地70B占位] {prompt[:40]}"
|
|
||||||
refs = ""
|
|
||||||
if self.echo_context:
|
|
||||||
for i, src in enumerate(context[:3], 1):
|
|
||||||
refs += f"\n[来源: {src}]"
|
|
||||||
return head + refs
|
|
||||||
|
|
||||||
|
|
||||||
class CloudBackend(InferenceBackend):
|
|
||||||
"""云端 API 后端占位实现:仅接收 DLP 放行的脱敏/通用内容。
|
|
||||||
|
|
||||||
子任务 #45 将替换为 Qwen/DeepSeek API 接入 + 安全网关。
|
|
||||||
"""
|
|
||||||
|
|
||||||
name = "cloud-api"
|
|
||||||
|
|
||||||
def __init__(self, echo_context: bool = True) -> None:
|
|
||||||
self.echo_context = echo_context
|
|
||||||
|
|
||||||
def generate(self, prompt: str, context: Sequence[str]) -> str:
|
|
||||||
head = f"[云端API占位] {prompt[:40]}"
|
|
||||||
refs = ""
|
|
||||||
if self.echo_context:
|
|
||||||
for i, src in enumerate(context[:3], 1):
|
|
||||||
refs += f"\n[来源: {src}]"
|
|
||||||
return head + refs
|
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# 网关输出
|
# 网关输出
|
||||||
|
|||||||
@@ -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}>")
|
||||||
@@ -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()
|
||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user