Files
llm-speed-tester/llm_providers.py
T

372 lines
15 KiB
Python
Raw Permalink Normal View History

# -*- coding: utf-8 -*-
"""大模型提供商适配器:OpenAI 兼容 / Anthropic / Google Gemini
统一接口: call_stream(cfg, prompt, gen, log, should_stop) -> metrics(dict)
度量指标:
prompt_tokens 上文/提示词 token 数(来自 provider usage
output_tokens 输出 token 数
cached_tokens 命中缓存的 token 数(provider 返回时才有)
ttft_ms 首字延迟(time to first token
decode_ms 解码阶段耗时(首字 -> 结束)
total_ms 总耗时(连接开始 -> 结束)
prefill_speed 预填充速度 = prompt_tokens / ttfttok/s
decode_speed 解码速度 = output_tokens / decode_timetok/s
"""
import json
import time
import requests
import config
class ProviderError(Exception):
"""API 调用失败"""
class StopRequested(Exception):
"""用户请求停止"""
DEFAULT_URLS = {
"openai": "https://api.openai.com/v1",
"anthropic": "https://api.anthropic.com",
"google": "https://generativelanguage.googleapis.com",
}
PROVIDER_LABELS = {
"openai": "OpenAI 兼容",
"anthropic": "Anthropic",
"google": "Google Gemini",
}
def _parse_sse_line(line):
line = (line or "").strip()
if not line.startswith("data:"):
return None
data = line[5:].strip()
if not data or data == "[DONE]":
return None
try:
return json.loads(data)
except Exception:
return None
def _metrics(start, first_token_at, end, prompt_tokens, output_tokens,
cached_tokens, output_chars, prompt_chars):
if first_token_at is None:
first_token_at = end # 未收到正文但请求完成(如纯 usage 响应)
ttft_ms = (first_token_at - start) * 1000
decode_ms = (end - first_token_at) * 1000
total_ms = (end - start) * 1000
prefill = (prompt_tokens / (ttft_ms / 1000)) if prompt_tokens and ttft_ms > 0 else None
decode = (output_tokens / (decode_ms / 1000)) if output_tokens and decode_ms > 0 else None
# 防“幻影 token”:usage 报了输出 token,但整个流没收到任何文本(如推理全用满 max_tokens
# 或提供商只回 usage 不正文),此时解码时间无意义,不报出数百万的荒谬速度
if output_tokens and decode is not None and not output_chars:
decode = None
return {
"prompt_tokens": int(prompt_tokens or 0),
"output_tokens": int(output_tokens or 0),
"cached_tokens": int(cached_tokens or 0),
"prompt_chars": int(prompt_chars or 0),
"output_chars": int(output_chars or 0),
"ttft_ms": round(ttft_ms, 1),
"decode_ms": round(decode_ms, 1),
"total_ms": round(total_ms, 1),
"prefill_speed": round(prefill, 1) if prefill else None,
"decode_speed": round(decode, 1) if decode else None,
}
# ───────────────────────── OpenAI 兼容 ─────────────────────────
def stream_openai(cfg, prompt, gen, log, should_stop=None):
base = (cfg.get("base_url") or DEFAULT_URLS["openai"]).rstrip("/")
url = base + "/chat/completions"
headers = {
"Authorization": "Bearer " + (cfg.get("api_key") or ""),
"Content-Type": "application/json",
}
payload = {
"model": cfg["model"],
"messages": [{"role": "user", "content": prompt}],
"stream": True,
"max_tokens": int(gen.get("max_tokens", 256)),
"temperature": float(cfg.get("temperature", 0.7)),
}
use_usage = True # stream_options include_usage(部分网关不支持则自动去掉重试)
def build_payload():
p = dict(payload)
if use_usage:
p["stream_options"] = {"include_usage": True}
return p
start = time.time()
first_token_at = None
output_chars = 0
prompt_tokens = output_tokens = cached_tokens = 0
event_count = 0
resp = None
try:
while True:
if should_stop and should_stop():
raise StopRequested()
resp = requests.post(url, json=build_payload(), headers=headers, stream=True,
timeout=(config.CONNECT_TIMEOUT, config.STREAM_READ_TIMEOUT))
if resp.status_code == 200:
break
err = resp.text[:400]
code = resp.status_code
resp.close()
resp = None
if use_usage and _bad_stream_options(err):
use_usage = False
if log:
log("WARN", "提供商不支持 stream_options=include_usage,已去掉参数重试")
continue
raise ProviderError("HTTP %s: %s" % (code, err))
for obj in _iter_json(resp):
if should_stop and should_stop():
raise StopRequested()
event_count += 1
if obj.get("choices"):
delta = obj["choices"][0].get("delta") or {}
# 兼容推理型模型思维链:DeepSeek 用 reasoning_content,部分 vLLM(Qwen3.8-27B-FP8) 用 reasoningAnthropic/Gemini 分别在各自适配器处理
piece = (delta.get("content") or delta.get("reasoning_content")
or delta.get("reasoning") or "")
if piece:
if first_token_at is None:
first_token_at = time.time()
output_chars += len(piece)
usage = obj.get("usage")
if usage:
prompt_tokens = usage.get("prompt_tokens") or 0
output_tokens = usage.get("completion_tokens") or 0
details = usage.get("prompt_tokens_details") or {}
if isinstance(details, dict):
cached_tokens = details.get("cached_tokens") or 0
except StopRequested:
raise
except Exception as e:
raise ProviderError("流式请求异常: %s" % e)
finally:
if resp is not None:
resp.close()
if event_count == 0:
raise ProviderError("未收到任何输出内容(HTTP 200 但响应流为空)")
end = time.time()
return _metrics(start, first_token_at, end, prompt_tokens, output_tokens,
cached_tokens, output_chars, len(prompt))
def _bad_stream_options(err: str):
err = (err or "").lower()
return ("stream_options" in err or "unknown parameter" in err or "unknown field" in err
or "additional properties" in err)
def _iter_json(resp):
"""解析 SSE data: 行,逐个返回 JSON 对象"""
for raw in resp.iter_lines(decode_unicode=True):
obj = _parse_sse_line(raw)
if obj is not None:
yield obj
# ───────────────────────── Anthropic ─────────────────────────
def stream_anthropic(cfg, prompt, gen, log, should_stop=None):
base = (cfg.get("base_url") or DEFAULT_URLS["anthropic"]).rstrip("/")
url = base + "/v1/messages"
headers = {
"x-api-key": cfg.get("api_key") or "",
"anthropic-version": "2023-06-01",
"Content-Type": "application/json",
}
payload = {
"model": cfg["model"],
"messages": [{"role": "user", "content": prompt}],
"max_tokens": int(gen.get("max_tokens", 256)),
"temperature": float(cfg.get("temperature", 0.7)),
"stream": True,
}
start = time.time()
first_token_at = None
output_chars = 0
prompt_tokens = output_tokens = 0
event_count = 0
resp = None
try:
if should_stop and should_stop():
raise StopRequested()
resp = requests.post(url, json=payload, headers=headers, stream=True,
timeout=(config.CONNECT_TIMEOUT, config.STREAM_READ_TIMEOUT))
if resp.status_code != 200:
err = resp.text[:400]
resp.close()
resp = None
raise ProviderError("HTTP %s: %s" % (resp.status_code, err))
for obj in _iter_json(resp):
if should_stop and should_stop():
raise StopRequested()
event_count += 1
etype = obj.get("type")
if etype == "message_start":
usage = (obj.get("message") or {}).get("usage") or {}
prompt_tokens = usage.get("input_tokens") or 0
elif etype == "content_block_delta":
delta = obj.get("delta") or {}
# 兼容 extended thinkingthinking 文本也算输出
text = delta.get("text") or delta.get("thinking") or ""
if text:
if first_token_at is None:
first_token_at = time.time()
output_chars += len(text)
elif etype == "message_delta":
usage = obj.get("usage") or {}
output_tokens = usage.get("output_tokens") or output_tokens
except StopRequested:
raise
except Exception as e:
raise ProviderError("流式请求异常: %s" % e)
finally:
if resp is not None:
resp.close()
if event_count == 0:
raise ProviderError("未收到任何输出内容(HTTP 200 但响应流为空)")
end = time.time()
return _metrics(start, first_token_at, end, prompt_tokens, output_tokens,
0, output_chars, len(prompt))
# ───────────────────────── Google Gemini ─────────────────────────
def stream_google(cfg, prompt, gen, log, should_stop=None):
base = (cfg.get("base_url") or DEFAULT_URLS["google"]).rstrip("/")
model = cfg["model"]
url = "%s/v1beta/models/%s:streamGenerateContent" % (base, model)
params = {"alt": "sse", "key": cfg.get("api_key") or ""}
headers = {"Content-Type": "application/json"}
payload = {
"contents": [{"parts": [{"text": prompt}]}],
"generationConfig": {
"temperature": float(cfg.get("temperature", 0.7)),
"maxOutputTokens": int(gen.get("max_tokens", 256)),
"candidateCount": 1,
},
}
start = time.time()
first_token_at = None
output_chars = 0
prompt_tokens = output_tokens = cached_tokens = 0
event_count = 0
resp = None
try:
if should_stop and should_stop():
raise StopRequested()
resp = requests.post(url, params=params, json=payload, headers=headers, stream=True,
timeout=(config.CONNECT_TIMEOUT, config.STREAM_READ_TIMEOUT))
if resp.status_code != 200:
err = resp.text[:400]
resp.close()
resp = None
raise ProviderError("HTTP %s: %s" % (resp.status_code, err))
for obj in _iter_json(resp):
if should_stop and should_stop():
raise StopRequested()
event_count += 1
cands = obj.get("candidates") or []
if cands:
parts = (cands[0].get("content") or {}).get("parts") or []
for part in parts:
# 兼容 thinking 模型:thought 文本也算输出
text = part.get("text") or part.get("thought") or ""
if text:
if first_token_at is None:
first_token_at = time.time()
output_chars += len(text)
um = obj.get("usageMetadata") or {}
if um:
prompt_tokens = um.get("promptTokenCount") or 0
output_tokens = um.get("candidatesTokenCount") or 0
cached_tokens = um.get("cachedContentTokenCount") or 0
except StopRequested:
raise
except Exception as e:
raise ProviderError("流式请求异常: %s" % e)
finally:
if resp is not None:
resp.close()
if event_count == 0:
raise ProviderError("未收到任何输出内容(HTTP 200 但响应流为空)")
end = time.time()
return _metrics(start, first_token_at, end, prompt_tokens, output_tokens,
cached_tokens, output_chars, len(prompt))
# ───────────────────────── 模型列表获取 ─────────────────────────
def list_models(cfg):
"""实时获取该接口(Base URL + API Key)下的所有模型 ID 列表。
支持 OpenAI 兼容 / Anthropic / Google Gemini 三种提供商。"""
provider = cfg.get("provider", "openai")
api_key = cfg.get("api_key") or ""
if not api_key:
raise ProviderError("请填写 API Key")
base = (cfg.get("base_url") or DEFAULT_URLS.get(provider, "")).rstrip("/")
if provider == "openai":
# GET {base}/models Authorization: Bearer <key> → data[].id
headers = {"Authorization": "Bearer " + api_key}
resp = requests.get(base + "/models", headers=headers,
timeout=config.CONNECT_TIMEOUT)
if resp.status_code != 200:
raise ProviderError("HTTP %s: %s" % (resp.status_code, resp.text[:400]))
data = resp.json()
return [m.get("id") for m in (data.get("data") or []) if m.get("id")]
if provider == "anthropic":
# GET {base}/v1/models x-api-key + anthropic-version → data[].id
headers = {"x-api-key": api_key, "anthropic-version": "2023-06-01"}
resp = requests.get(base + "/v1/models", headers=headers,
timeout=config.CONNECT_TIMEOUT)
if resp.status_code != 200:
raise ProviderError("HTTP %s: %s" % (resp.status_code, resp.text[:400]))
data = resp.json()
return [m.get("id") for m in (data.get("data") or []) if m.get("id")]
if provider == "google":
# GET {base}/v1beta/models?key=<key> → models[].name(形如 models/gemini-1.5-pro
resp = requests.get(base + "/v1beta/models", params={"key": api_key},
timeout=config.CONNECT_TIMEOUT)
if resp.status_code != 200:
raise ProviderError("HTTP %s: %s" % (resp.status_code, resp.text[:400]))
data = resp.json()
out = []
for m in (data.get("models") or []):
name = m.get("name") or ""
if name.startswith("models/"):
name = name[len("models/"):]
if name and name not in out:
out.append(name)
return out
raise ProviderError("不支持的提供商类型: %s" % provider)
# ───────────────────────── 统一入口 ─────────────────────────
def call_stream(cfg, prompt, gen, log=None, should_stop=None):
provider = cfg.get("provider", "openai")
fn = {"openai": stream_openai, "anthropic": stream_anthropic, "google": stream_google}.get(provider)
if fn is None:
raise ProviderError("不支持的提供商类型: %s" % provider)
return fn(cfg, prompt, gen, log, should_stop)