150 lines
4.5 KiB
Python
150 lines
4.5 KiB
Python
"""大模型调用客户端: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": "请确认服务可用。"},
|
||
],
|
||
)
|