v2.5.0 模型实时获取:接口配置区「📋 查看模型」按钮,Base URL+API Key 下实时拉取模型列表(GET /models),搜索/点击填入,仍支持手动输入
This commit is contained in:
@@ -307,6 +307,56 @@ def stream_google(cfg, prompt, gen, log, should_stop=None):
|
||||
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):
|
||||
|
||||
Reference in New Issue
Block a user