第 1 卷第 1 章:1–9 年
This commit is contained in:
+106
@@ -0,0 +1,106 @@
|
||||
"""调用用户自建网关(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}")
|
||||
Reference in New Issue
Block a user