# -*- coding: utf-8 -*- """推理后端统一抽象接口(PRD 5.6「⑥ 部署底座」)。 设计约束(对齐产品设计文档 5.6): - 统一接口 `load_model / infer / health / unload`,与底层硬件无关; - NVIDIA 5090 实现(Triton/ONNX)与华为昇腾实现(ACL/CANN)均实现本接口; - 业务代码只 import 本模块,不 import 任何具体后端; - 切换后端 = 修改适配层配置(见 config/backends.template.yaml),业务代码零改动。 """ from abc import ABC, abstractmethod from dataclasses import dataclass, field import json import time import urllib.request @dataclass class InferRequest: """统一的推理请求载荷(业务侧构造,与硬件无关)。""" prompt: str max_tokens: int = 1024 temperature: float = 0.1 extra: dict = field(default_factory=dict) def to_payload(self) -> dict: """转换为后端服务 OpenAI 兼容的请求体。""" return { "model": self.extra.get("model") if self.extra.get("model") else None, "prompt": self.prompt, "max_tokens": self.max_tokens, "temperature": self.temperature, } @dataclass class InferResult: """统一的推理结果(业务侧消费,不感知硬件)。""" text: str backend: str latency_ms: float = 0.0 meta: dict = field(default_factory=dict) class InferenceBackend(ABC): """推理后端统一接口(对应 PRD 5.6:loadModel / infer / health / unload)。""" #: 后端标识:gpu(NVIDIA 5090)| npu(华为昇腾) backend_name = "base" def __init__(self, endpoint: str = "", model: str = "iaop-default", timeout_seconds: float = 10.0): self.endpoint = (endpoint or "").rstrip("/") self.model = model self.timeout = float(timeout_seconds) # ---- 统一接口(业务代码仅依赖以下四个方法) ---- @abstractmethod def load_model(self, model_name: str | None = None) -> dict: """加载/热载模型(对应 PRD 的 loadModel)。返回加载状态。""" @abstractmethod def infer(self, request: InferRequest) -> InferResult: """执行一次推理(对应 PRD 的 infer)。返回统一结果对象。""" @abstractmethod def health(self) -> dict: """健康巡检(对应 PRD 的 health)。返回后端状态与版本信息。""" @abstractmethod def unload(self) -> dict: """卸载模型/释放显存(对应 PRD 的 unload)。返回卸载结果。""" # ---- 内部工具:HTTP JSON 调用(各后端共享) ---- def _post_json(self, path: str, payload: dict) -> dict: """向后端服务发起 JSON POST 请求,返回解析后的响应体。""" url = self.endpoint + path data = json.dumps(payload).encode("utf-8") req = urllib.request.Request( url, data=data, headers={"Content-Type": "application/json"} ) started = time.monotonic() with urllib.request.urlopen(req, timeout=self.timeout) as resp: raw = resp.read().decode("utf-8") meta = {"http_ms": round((time.monotonic() - started) * 1000, 2)} body = json.loads(raw) if raw else {} if isinstance(body, dict): body.setdefault("_http_ms", meta["http_ms"]) return body def _healthz(self) -> dict: """通用 /health 探针封装。""" if not self.endpoint: return { "status": "ok", "backend": self.backend_name, "model": self.model, "reason": "dry-run(未配置 endpoint,适配层就绪)", } return self._post_json("/health", {"model": self.model})