aps-agent/server/agent_core/providers.py

78 lines
4.6 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.

# ============================================================
# 模型 Provider(moduleId: core-providers, 可重生 ✅)
# 屏蔽 DeepSeek / KIMI 差异(两家均为 OpenAI 兼容接口,plan.md §11.2)
# 铁律:LLM 只有提议权(§3.2)——本模块只产出 JSON,绝不执行任何动作
# ============================================================
from __future__ import annotations # 前向类型引用
import json # LLM 输出解析
import os # 环境变量读取
from typing import Any # 类型标注
import httpx # 异步 HTTP 客户端
# 各 Provider 的默认端点与模型(.env 未配置时的回退值)
_DEFAULTS = {
"deepseek": {"base_url": "https://api.deepseek.com/v1", "model": "deepseek-chat"}, # DeepSeek 默认
"kimi": {"base_url": "https://api.moonshot.cn/v1", "model": "moonshot-v1-8k"}, # KIMI(Moonshot) 默认
}
class ModelProvider:
"""OpenAI 兼容的对话接口封装:chat_json() 强制 JSON 输出(P0:纯提议,无副作用)。"""
def __init__(self) -> None:
"""从环境读取配置;密钥缺失时 self.enabled=False(上游自动降级为规则解析)。"""
self.provider = os.environ.get("LLM_PROVIDER", "").strip().lower() # 提供方标识
self.api_key = os.environ.get("LLM_API_KEY", "").strip() # API 密钥(不入库)
defaults = _DEFAULTS.get(self.provider, {}) # 该提供方的默认值
self.base_url = os.environ.get("LLM_BASE_URL", "").strip() or defaults.get("base_url", "") # 端点
self.model = os.environ.get("LLM_MODEL", "").strip() or defaults.get("model", "") # 模型名
# 可用性判定:提供方合法 + 密钥存在 + 端点已知
self.enabled = bool(self.provider in _DEFAULTS and self.api_key and self.base_url)
async def chat_json(self, system: str, user: str, timeout: float = 20.0) -> dict[str, Any] | None:
"""单轮对话并强制 JSON 输出;任何异常返回 None(由上游降级,保证离线可用)。
Args:
system: 系统提示词(含意图枚举与输出 Schema 说明)
user: 用户输入
timeout: 请求超时秒数(意图识别要快,默认 20s)
Returns:
解析后的 dict;失败(网络/超时/非法 JSON)返回 None
"""
if not self.enabled: # 未配置 → 直接告知不可用
return None
try:
async with httpx.AsyncClient(timeout=timeout) as client: # 短连接客户端
resp = await client.post( # OpenAI 兼容 chat/completions
f"{self.base_url.rstrip('/')}/chat/completions",
headers={"Authorization": f"Bearer {self.api_key}"}, # Bearer 鉴权
json={
"model": self.model, # 模型名
"messages": [ # 单轮:system + user
{"role": "system", "content": system},
{"role": "user", "content": user},
],
"temperature": 0.1, # 低温:意图识别要稳定不要创意
"response_format": {"type": "json_object"}, # 强制 JSON(两家均支持)
},
)
resp.raise_for_status() # 非 2xx 抛错 → 走降级
content = resp.json()["choices"][0]["message"]["content"] # 取回复正文
return json.loads(content) # 解析 JSON(失败抛错 → 降级)
except Exception: # 网络/超时/格式错误统一兜底
return None # 上游据此降级为规则解析
# ---------------- 进程级单例(配置只读,无状态可共享) ----------------
_provider: ModelProvider | None = None # 单例槽
def get_provider() -> ModelProvider:
"""获取全局 ModelProvider 单例(懒加载,.env 已由 main 加载)。"""
global _provider # 引用模块级槽
if _provider is None: # 首次调用
_provider = ModelProvider() # 按环境构建
return _provider # 返回单例