feat: 完成 issue #45 ④ 云端 API(Qwen/DeepSeek)接入与安全网关

This commit is contained in:
2026-08-05 02:37:41 +08:00
parent 6bbcd810e7
commit c5a7a0ee52
4 changed files with 249 additions and 2 deletions
+97 -1
View File
@@ -15,9 +15,10 @@
from __future__ import annotations
import json
import os
import time
import urllib.request
from typing import Optional, Sequence
from typing import Callable, Optional, Sequence
from .gateway import InferenceBackend
@@ -112,3 +113,98 @@ class Local70BBackend(InferenceBackend):
except Exception as exc: # noqa: BLE001 - 健康探测失败仅记录
base["status"] = f"error: {exc}"
return base
class CloudApiBackend(InferenceBackend):
"""云端 API 推理后端(Qwen / DeepSeek 等 OpenAI 兼容)—— issue #45。
**安全网关约束(PRD 5.4)**:
- 仅接收 **DLP 放行**的脱敏/通用内容(上游 `LLMGateway` 主编排出站检查 +
cloud 分支输出 DLP 复查);
- API Key 从**环境变量**读取(`api_key_env`),不硬编码、不落日志;
- 可选 `safety_checker` 出站复查钩子(fail-closed:复查拒绝 → 拦截占位,
不调用上游)。
"""
name = "cloud-api"
def __init__(
self,
endpoint: str = "",
api_key_env: str = "",
model: str = "deepseek-chat",
timeout_seconds: float = 60.0,
max_tokens: int = 1024,
temperature: float = 0.1,
safety_checker: Optional[Callable[[str], bool]] = None,
) -> None:
self.endpoint = (endpoint or "").rstrip("/")
self.api_key_env = api_key_env
self.model = model
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 ""
# ------------------------------------------------------------------
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 "[云端安全网关拦截] 出站复查未通过,已拦截(数据不出厂)。"
if not self.endpoint:
return self._dry_run(prompt, context)
payload = {
"model": self.model,
"messages": [
{"role": "system", "content": self._system_prompt(context)},
{"role": "user", "content": prompt},
],
"max_tokens": self.max_tokens,
"temperature": self.temperature,
}
body = self._post_json("/v1/chat/completions", payload)
try:
return body["choices"][0]["message"]["content"]
except (KeyError, IndexError, TypeError):
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 []))
base = "你是工业 AI 助手。回答须基于给定资料并标注来源。"
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}]"
if self.safety_checker is not None:
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")
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:
raw = resp.read().decode("utf-8")
return json.loads(raw) if raw else {}
def health(self) -> dict:
"""后端健康信息(含安全网关状态,不含密钥)。"""
return {
"backend": self.name, "model": self.model,
"endpoint": self.endpoint or "(dry-run)",
"api_key_configured": bool(self._api_key),
"safety_checker": self.safety_checker is not None,
"status": "dry-run" if not self.endpoint else "configured",
}