107 lines
3.6 KiB
Python
107 lines
3.6 KiB
Python
"""调用用户自建网关(OpenAI 兼容)生成文本。
|
||
|
||
要点:
|
||
- 密钥运行时从 Pi 的 auth.json 读取,只在内存中,永不打印或写入产物。
|
||
- 目标是推理模型,思考与正文共用 max_tokens,因此默认给足预算。
|
||
- 传输层用 http.client(方案显式限定 http/https),不依赖第三方库。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import http.client
|
||
import json
|
||
import time
|
||
from dataclasses import dataclass
|
||
from urllib.parse import urlsplit
|
||
|
||
from dfannals.config import MAX_TOKENS_CHAPTER, MODEL, Gateway, load_gateway
|
||
|
||
|
||
@dataclass
|
||
class Reply:
|
||
content: str
|
||
reasoning_tokens: int
|
||
total_tokens: int
|
||
model: str
|
||
|
||
|
||
class GatewayError(RuntimeError):
|
||
pass
|
||
|
||
|
||
def _endpoint(base_url: str) -> tuple[type[http.client.HTTPConnection], str, str]:
|
||
"""解析网关地址,只接受 http/https。"""
|
||
parts = urlsplit(base_url)
|
||
if parts.scheme not in ("http", "https"):
|
||
raise GatewayError(f"网关地址必须是 http(s):{base_url!r}")
|
||
if not parts.netloc:
|
||
raise GatewayError(f"网关地址缺少主机:{base_url!r}")
|
||
cls = http.client.HTTPSConnection if parts.scheme == "https" else http.client.HTTPConnection
|
||
path = parts.path.rstrip("/") + "/chat/completions"
|
||
return cls, parts.netloc, path
|
||
|
||
|
||
def chat(
|
||
messages: list[dict],
|
||
*,
|
||
gateway: Gateway | None = None,
|
||
model: str = MODEL,
|
||
max_tokens: int = MAX_TOKENS_CHAPTER,
|
||
temperature: float = 1.0,
|
||
timeout: int = 900,
|
||
retries: int = 3,
|
||
) -> Reply:
|
||
gw = gateway or load_gateway()
|
||
conn_cls, host, path = _endpoint(gw.base_url)
|
||
body = json.dumps(
|
||
{
|
||
"model": model,
|
||
"messages": messages,
|
||
"max_tokens": max_tokens,
|
||
"temperature": temperature,
|
||
}
|
||
).encode()
|
||
headers = {
|
||
"Content-Type": "application/json",
|
||
# 密钥仅在此处作为请求头使用,不落盘、不打印
|
||
"Authorization": "Bearer " + gw.key,
|
||
}
|
||
|
||
last: str | None = None
|
||
for attempt in range(1, retries + 1):
|
||
conn = conn_cls(host, timeout=timeout)
|
||
try:
|
||
conn.request("POST", path, body=body, headers=headers)
|
||
resp = conn.getresponse()
|
||
raw = resp.read()
|
||
if resp.status != 200:
|
||
last = f"HTTP {resp.status}: {raw[:400].decode('utf-8', 'replace')}"
|
||
if resp.status in (400, 401, 403, 404):
|
||
break
|
||
else:
|
||
data = json.loads(raw)
|
||
choice = data["choices"][0]["message"]
|
||
usage = data.get("usage") or {}
|
||
content = (choice.get("content") or "").strip()
|
||
if not content:
|
||
raise GatewayError(
|
||
"模型只产出了思考、没有正文;"
|
||
f"thinking_tokens={usage.get('completion_thinking_tokens')},请加大 max_tokens"
|
||
)
|
||
return Reply(
|
||
content=content,
|
||
reasoning_tokens=int(usage.get("completion_thinking_tokens") or 0),
|
||
total_tokens=int(usage.get("total_tokens") or 0),
|
||
model=data.get("model", model),
|
||
)
|
||
except (OSError, http.client.HTTPException, json.JSONDecodeError, KeyError) as exc:
|
||
last = f"{type(exc).__name__}: {exc}"
|
||
except GatewayError as exc:
|
||
last = str(exc)
|
||
break
|
||
finally:
|
||
conn.close()
|
||
if attempt < retries:
|
||
time.sleep(3 * attempt)
|
||
|
||
raise GatewayError(f"网关调用失败({gw.base_url},model={model}):{last}")
|