nex_math/backend/services/llm_client.py

150 lines
4.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""大模型调用客户端:OpenAI / DeepSeek / 阿里百炼 qwen / OpenAI 兼容。"""
from __future__ import annotations
import httpx
PROVIDER_PRESETS = {
"openai": {
"label": "OpenAI",
"base_url": "https://api.openai.com/v1",
"model": "gpt-4o-mini",
},
"deepseek": {
"label": "DeepSeek",
"base_url": "https://api.deepseek.com/v1",
"model": "deepseek-chat",
},
"qwen": {
"label": "阿里千问(百炼)",
"base_url": "https://dashscope.aliyuncs.com/compatible-mode/v1",
"model": "qwen-plus",
},
"openai_compatible": {
"label": "OpenAI 兼容(自定义)",
"base_url": "",
"model": "",
},
}
def mask_key(api_key: str) -> str:
if not api_key:
return ""
if len(api_key) <= 8:
return "****"
return f"****{api_key[-4:]}"
def _content_text(value) -> str:
"""兼容 content 为字符串或 OpenAI 多段数组([{"type":"text","text":...}])。"""
if value is None:
return ""
if isinstance(value, str):
return value
if isinstance(value, list):
chunks: list[str] = []
for part in value:
if isinstance(part, dict):
text = part.get("text")
if isinstance(text, str):
chunks.append(text)
else:
chunks.append(str(part))
return "".join(chunks)
if isinstance(value, dict):
text = value.get("text")
return text if isinstance(text, str) else ""
return str(value)
def _message_text(message: dict) -> str:
"""OpenAI 兼容网关常见差异:正文可能在 content 或 reasoning/reasoning_content。"""
text = _content_text(message.get("content"))
if text.strip():
return text.strip()
for key in ("reasoning_content", "reasoning"):
fallback = _content_text(message.get(key))
if fallback.strip():
return fallback.strip()
refusal = message.get("refusal")
if isinstance(refusal, str) and refusal.strip():
raise ValueError(f"模型拒绝回答:{refusal[:200]}")
return ""
def chat_completion(
*,
base_url: str,
api_key: str,
model: str,
messages: list[dict[str, str]],
temperature: float = 0.3,
max_tokens: int = 1600,
timeout: float = 120.0,
) -> str:
if not base_url or not api_key or not model:
raise ValueError("请先完整填写模型配置(地址、API Key、模型名)。")
endpoint = f"{base_url.rstrip('/')}/chat/completions"
payload = {
"model": model,
"messages": messages,
"temperature": temperature,
"max_tokens": max_tokens,
}
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
}
try:
with httpx.Client(timeout=timeout) as client:
response = client.post(endpoint, json=payload, headers=headers)
except httpx.HTTPError as exc:
raise ValueError(f"无法连接模型服务:{exc.__class__.__name__}") from exc
if response.status_code != 200:
detail = response.text[:300]
raise ValueError(f"模型服务返回 {response.status_code}:{detail}")
try:
data = response.json()
message = data["choices"][0].get("message")
if not isinstance(message, dict):
raise KeyError("message")
content = _message_text(message)
if not content:
finish_reason = data["choices"][0].get("finish_reason")
raise ValueError(
"模型没有返回可见文本"
f"(finish_reason={finish_reason})。"
"若模型启用了推理模式,请提高 max_tokens 或缩短提示。"
)
return content
except ValueError:
raise
except (KeyError, IndexError, TypeError) as exc:
raise ValueError("模型返回格式不符合 Chat Completions 规范") from exc
def quick_test(
*,
base_url: str,
api_key: str,
model: str,
temperature: float = 0.0,
) -> str:
return chat_completion(
base_url=base_url,
api_key=api_key,
model=model,
temperature=temperature,
max_tokens=256,
timeout=45.0,
messages=[
{
"role": "system",
"content": "你是连通性测试助手,只回复两个字:正常",
},
{"role": "user", "content": "请确认服务可用。"},
],
)