Files
nba-fan-hub/llm.py
T

89 lines
2.8 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.
# -*- coding: utf-8 -*-
"""DeepSeek 大模型调用层(OpenAI 兼容 /chat/completions,支持 function calling"""
import json
import logging
import time
import requests
from config import LLM_BASE_URL, LLM_API_KEY, LLM_MODEL, LLM_TIMEOUT, LLM_MAX_TOKENS, LLM_TEMPERATURE
log = logging.getLogger("llm")
URL = f"{LLM_BASE_URL}/chat/completions"
HEADERS = {"Authorization": f"Bearer {LLM_API_KEY}", "Content-Type": "application/json"}
def _post(payload, timeout=LLM_TIMEOUT):
resp = requests.post(URL, headers=HEADERS, json=payload, timeout=timeout)
resp.raise_for_status()
return resp.json()
def chat(messages, tools=None, temperature=LLM_TEMPERATURE, max_tokens=LLM_MAX_TOKENS):
"""基础对话。返回完整 OpenAI 响应 dict。
tools 为 None 时不带工具;带工具时模型可能返回 tool_calls。"""
payload = {
"model": LLM_MODEL,
"messages": messages,
"temperature": temperature,
"max_tokens": max_tokens,
"stream": False,
}
if tools:
payload["tools"] = tools
payload["tool_choice"] = "auto"
for attempt in range(3):
try:
return _post(payload)
except requests.exceptions.HTTPError as e:
body = e.response.text[:400] if e.response is not None else ""
code = e.response.status_code if e.response is not None else 0
if code in (400, 429, 500, 502, 503) and attempt < 2:
log.warning("LLM HTTP %s(第%s次重试): %s", code, attempt + 1, body)
time.sleep(2 * (attempt + 1))
continue
log.error("LLM HTTP %s: %s", code, body)
raise
raise RuntimeError("LLM 请求失败")
def parse_content(resp):
"""提取 assistant 文本内容"""
try:
return resp["choices"][0]["message"].get("content") or ""
except Exception:
return ""
def extract_tool_calls(resp):
"""提取 [{name, arguments(dict), id}]"""
calls = []
try:
msg = resp["choices"][0]["message"]
for tc in msg.get("tool_calls") or []:
try:
args = json.loads(tc["function"].get("arguments") or "{}")
except json.JSONDecodeError:
args = {}
calls.append({"name": tc["function"]["name"], "arguments": args, "id": tc["id"]})
except Exception:
pass
return calls
def extract_reasoning(resp):
"""提取思考模式下的 reasoning_content(回传时必须带上)"""
try:
return resp["choices"][0]["message"].get("reasoning_content") or ""
except Exception:
return ""
def count_tokens(messages):
"""粗略估算 token 数(中英混合,1字≈1token"""
total = 0
for m in messages:
total += len(m.get("content") or "") // 2 + 8
return total