Files
iAOP/core/llm-gateway/backends.py
T
bot_dev1 818f1db860 feat(#58): GPU 后端实现(NVIDIA Triton/ONNX,PRD 5.6 推理后端可插拔)
落地父 EPIC #8 / Issue #58 要求的 NVIDIA GPU(5090)推理后端,严格实现
Issue #57 定义的 InferenceBackend 抽象接口(load_model/infer/health_check/unload)。

新增 core/llm-gateway/gpu_backend.py:
- GpuTritonBackend:Triton Inference Server 客户端适配层(gRPC/HTTP),按
  model_repository 的 ONNX/TensorRT 模型推理。
- 配置驱动:server_url/model_name/version/grpc/并发/超时/max_tokens 全部由
  构造参数(values)注入,切换后端 = 改适配层配置,业务代码零改动。
- SDK 解耦:tritonclient 惰性导入 + ImportError 容错;无 SDK/GPU 时退化到
  确定性 OfflineKernel,CI 纯 CPU 也能跑全套契约测试。
- fail-closed:未 load_model 即 infer 抛 RuntimeError(生产严格)。
- 健康探针:is_server_live + is_model_ready 双判定,供 #61 监控消费。
- 审计:infer 返回 InferResult(text/token 计数/latency_ms)。

注册表接入:default_registry() 登记 gpu-triton,build_backend 可配置切换。

新增 core/llm-gateway/tests/test_gpu_backend.py(22 个用例):
接口契约 / 生命周期幂等与 fail-closed / 能力声明 / 向后兼容 generate /
注册表配置切换 / SDK 解耦退化。

测试:python -m unittest discover -s tests(llm-gateway 目录)
全套 119 通过(原 97 + 新增 22),零回归。
2026-08-04 21:55:40 +08:00

305 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- 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)