aps-agent/server/agent_core/providers.py

128 lines
5.4 KiB
Python
Raw Permalink 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 兼容端点差异
# 铁律:LLM 只有提议权(§3.2)——本模块只产出文本/JSON,绝不执行任何动作
# ============================================================
from __future__ import annotations
import json
import os
from typing import Any
import httpx
# 已知提供方默认端点与模型(也可被 LLM_BASE_URL / LLM_MODEL 覆盖)
_DEFAULTS = {
"deepseek": {"base_url": "https://api.deepseek.com/v1", "model": "deepseek-chat"},
"kimi": {"base_url": "https://api.moonshot.cn/v1", "model": "moonshot-v1-8k"},
"moonshot": {"base_url": "https://api.moonshot.cn/v1", "model": "moonshot-v1-8k"},
"doubao": {"base_url": "https://ark.cn-beijing.volces.com/api/v3", "model": ""},
"openai": {"base_url": "https://api.openai.com/v1", "model": "gpt-4o-mini"},
"compatible": {"base_url": "", "model": ""}, # 完全自定义:必须配 BASE_URL + MODEL
}
class ModelProvider:
"""OpenAI 兼容对话接口:chat_json / chat_text;失败返回 None 由上游降级。"""
def __init__(self) -> None:
self.provider = os.environ.get("LLM_PROVIDER", "").strip().lower()
self.api_key = os.environ.get("LLM_API_KEY", "").strip()
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", "")
# 任意 OpenAI 兼容:有密钥 + 端点 + 模型即可(不限于白名单名)
known = self.provider in _DEFAULTS or bool(self.base_url)
self.enabled = bool(known and self.api_key and self.base_url and self.model)
async def chat_json(self, system: str, user: str, timeout: float = 20.0) -> dict[str, Any] | None:
if not self.enabled:
return None
try:
async with httpx.AsyncClient(timeout=timeout) as client:
resp = await client.post(
f"{self.base_url.rstrip('/')}/chat/completions",
headers={"Authorization": f"Bearer {self.api_key}"},
json={
"model": self.model,
"messages": [
{"role": "system", "content": system},
{"role": "user", "content": user},
],
"temperature": 0.1,
"response_format": {"type": "json_object"},
},
)
resp.raise_for_status()
content = resp.json()["choices"][0]["message"]["content"]
return json.loads(content)
except Exception:
# 部分国产端点不支持 response_format → 再试无强制 JSON
try:
async with httpx.AsyncClient(timeout=timeout) as client:
resp = await client.post(
f"{self.base_url.rstrip('/')}/chat/completions",
headers={"Authorization": f"Bearer {self.api_key}"},
json={
"model": self.model,
"messages": [
{"role": "system", "content": system + "\n只输出合法 JSON 对象。"},
{"role": "user", "content": user},
],
"temperature": 0.1,
},
)
resp.raise_for_status()
content = resp.json()["choices"][0]["message"]["content"]
# 容错:从 markdown 代码块里抠 JSON
text = content.strip()
if "```" in text:
text = re_strip_fence(text)
return json.loads(text)
except Exception:
return None
async def chat_text(self, system: str, user: str, timeout: float = 45.0) -> str | None:
if not self.enabled:
return None
try:
async with httpx.AsyncClient(timeout=timeout) as client:
resp = await client.post(
f"{self.base_url.rstrip('/')}/chat/completions",
headers={"Authorization": f"Bearer {self.api_key}"},
json={
"model": self.model,
"messages": [
{"role": "system", "content": system},
{"role": "user", "content": user},
],
"temperature": 0.2,
},
)
resp.raise_for_status()
return resp.json()["choices"][0]["message"]["content"]
except Exception:
return None
def re_strip_fence(text: str) -> str:
import re
m = re.search(r"```(?:json)?\s*([\s\S]*?)```", text)
return (m.group(1) if m else text).strip()
_provider: ModelProvider | None = None
def get_provider() -> ModelProvider:
global _provider
if _provider is None:
_provider = ModelProvider()
return _provider
def reset_provider() -> None:
"""测试用:清空单例以便重读环境变量。"""
global _provider
_provider = None