Files
dwarf-fortress-annals/dfannals/llm.py
T

107 lines
3.6 KiB
Python
Raw 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.
"""调用用户自建网关(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}")