diff --git a/README.md b/README.md index 5130c8e..a1a0069 100644 --- a/README.md +++ b/README.md @@ -4,7 +4,7 @@ > 把本机已经登录的消费级 AI 客户端,接成 OpenAI 兼容接口,给 Codex、OpenCode、Cherry Studio、NextChat 等用。默认打开 Work Buddy / CodeBuddy、QClaw、千问办公(QwenWork)、TraeWork 四个通道;管理页下拉选其中一个。一次请求只走一个通道。 -当前版本 **2.1.1**。这个项目只适合本机自用,不要公开部署,也不要把登录凭据、API Key、数据库文件发给别人。 +当前版本 **2.1.2**。这个项目只适合本机自用,不要公开部署,也不要把登录凭据、API Key、数据库文件发给别人。 ## 这是什么? @@ -105,7 +105,7 @@ python server.py 4. 在客户端里填: - Base URL:`http://127.0.0.1:8787/v1` - API Key:刚复制的 Key - - 模型:WorkBuddy 用 `auto` 即可;QClaw 用 `auto`;千问办公用 `auto` 或 `qwork-advanced`;TraeWork 用 `auto` 或 `qwen-3.7-plus` + - 模型:WorkBuddy 用 `auto` 即可;QClaw 用 `auto`;千问办公用 `auto` 或 `qwork-advanced`;TraeWork 用 `auto` 或 `qwen-3.7-plus`。上游加了新模型时,到「模型配置」点「一键读取供应模型」;各通道目录分开保存,选错通道仍会 400/403。 管理页打不开或要远程访问时: diff --git a/README_EN.md b/README_EN.md index 0e0b69b..89282fc 100644 --- a/README_EN.md +++ b/README_EN.md @@ -4,7 +4,7 @@ > Local consumer AI clients → one OpenAI-compatible API for Codex, OpenCode, Cherry Studio, NextChat, and similar agents. Work Buddy / CodeBuddy, QClaw, QwenWork, and TraeWork are on by default; pick one in the UI dropdown. Each request stays on one channel. -Release **2.1.1**. Local use only. Do not expose this on the public internet, and do not share credentials, API keys, or the database. +Release **2.1.2**. Local use only. Do not expose this on the public internet, and do not share credentials, API keys, or the database. ## What is this? @@ -78,7 +78,7 @@ The database migrates on startup. Existing keys stay on `workbuddy`. Startup no | API Key | Created in the UI, bound to one channel | | Model | WorkBuddy `auto`; QClaw `auto`; QwenWork `qwork-advanced` | -Unprefixed `auto` follows the key’s channel. Use a separate key per channel. +Unprefixed `auto` follows the key’s channel. Use a separate key per channel. On the Models page, “一键读取供应模型” refreshes each channel’s supplier list separately; a TraeWork-only id such as Doubao is never merged into WorkBuddy. ### Reasoning effort diff --git a/catalog.py b/catalog.py new file mode 100644 index 0000000..2f63201 --- /dev/null +++ b/catalog.py @@ -0,0 +1,252 @@ +"""Per-channel supplier model catalogs. + +Fetch+parse of each source's list is separate from persist and from chat I/O. +WorkBuddy and QwenWork have no supplier-list HTTP; they stay on the existing +static/admin catalog and are reported as fallback. +""" + +from __future__ import annotations + +import sqlite3 +import time +from typing import Any, Awaitable, Callable + +import database as db + +CATALOG_SETTING = "channel_catalogs" +REFRESH_SETTING = "channel_catalog_refresh" + +Fetcher = Callable[[dict], Awaitable[list[dict]]] + + +def _load_map(key: str) -> dict: + try: + value = db.get_setting(key, {}) or {} + except sqlite3.OperationalError: + return {} + return value if isinstance(value, dict) else {} + + +def stored_catalog(channel: str) -> list[dict] | None: + items = _load_map(CATALOG_SETTING).get(channel) + if isinstance(items, list) and items: + return [item for item in items if isinstance(item, dict) and item.get("id")] + return None + + +def save_catalog(channel: str, models: list[dict]) -> None: + catalogs = _load_map(CATALOG_SETTING) + catalogs[channel] = models + db.set_setting(CATALOG_SETTING, catalogs) + + +def normalize_models(rows: Any) -> list[dict]: + if not isinstance(rows, list): + return [] + models: list[dict] = [] + seen: set[str] = set() + for row in rows: + if isinstance(row, str): + mid = row.strip() + name = mid + description = "" + elif isinstance(row, dict): + mid = str( + row.get("id") + or row.get("model_id") + or row.get("model_name") + or row.get("name") + or "" + ).strip() + name = str(row.get("name") or row.get("display_id") or row.get("display_name") or mid) + description = str(row.get("description") or "") + else: + continue + if not mid or mid in seen: + continue + seen.add(mid) + item = {"id": mid, "name": name or mid} + if description: + item["description"] = description + models.append(item) + return models + + +def models_for(channel: str, fallback: list[dict]) -> list[dict]: + stored = stored_catalog(channel) + if stored: + return stored + return list(fallback) + + +def workbuddy_fallback_models() -> list[dict]: + import proxy + + try: + models = db.get_setting("models", proxy.DEFAULT_MODELS) + except sqlite3.OperationalError: + return list(proxy.DEFAULT_MODELS) + if isinstance(models, list) and models: + return models + return list(proxy.DEFAULT_MODELS) + + +def _fallback_models(channel: str, provider) -> list[dict]: + if channel == "workbuddy": + return workbuddy_fallback_models() + if provider is not None: + stored = stored_catalog(channel) + if stored: + return stored + try: + return list(provider.list_models()) + except Exception: + pass + return [] + + +def _status_row( + channel: str, + *, + mode: str, + models: list[dict], + message: str = "", + display_name: str = "", +) -> dict: + return { + "channel": channel, + "display_name": display_name or channel, + "mode": mode, + "message": message, + "count": len(models), + "models": models, + "updated_at": int(time.time()), + } + + +async def _pick_account(provider) -> dict | None: + if provider is None: + return None + picker = getattr(provider, "pick_account_with_fallback", None) + if picker is None: + return None + return await picker() + + +async def _fetch_qclaw(account: dict) -> list[dict]: + from providers.qclaw.jprx import fetch_supplier_models + + return await fetch_supplier_models(account) + + +async def _fetch_traework(account: dict) -> list[dict]: + from providers.traework.models import fetch_supplier_models + + return await fetch_supplier_models(account) + + +LIVE_FETCHERS: dict[str, Fetcher] = { + "qclaw": _fetch_qclaw, + "traework": _fetch_traework, +} + + +async def refresh_one(channel: str) -> dict: + import providers + + provider = providers.get_provider(channel) + display_name = getattr(provider, "display_name", channel) if provider else channel + fallback = _fallback_models(channel, provider) + fetcher = LIVE_FETCHERS.get(channel) + if fetcher is None: + return _status_row( + channel, + mode="fallback", + models=fallback, + message="no supplier-list API", + display_name=display_name, + ) + account = await _pick_account(provider) + if not account: + return _status_row( + channel, + mode="fallback", + models=fallback, + message="no usable account", + display_name=display_name, + ) + try: + fetched = normalize_models(await fetcher(account)) + except Exception as exc: + return _status_row( + channel, + mode="fallback", + models=fallback, + message=str(exc)[:240], + display_name=display_name, + ) + if not fetched: + return _status_row( + channel, + mode="fallback", + models=fallback, + message="empty supplier list", + display_name=display_name, + ) + save_catalog(channel, fetched) + return _status_row( + channel, + mode="live", + models=fetched, + message="", + display_name=display_name, + ) + + +async def refresh_supplier_catalogs() -> dict: + import providers + + sources = [] + for channel in providers.enabled_provider_ids(): + sources.append(await refresh_one(channel)) + db.set_setting( + REFRESH_SETTING, + { + item["channel"]: { + "mode": item["mode"], + "message": item["message"], + "count": item["count"], + "updated_at": item["updated_at"], + } + for item in sources + }, + ) + return {"sources": sources} + + +def catalog_snapshot() -> dict: + import providers + + refresh = _load_map(REFRESH_SETTING) + sources = [] + for channel in providers.enabled_provider_ids(): + provider = providers.get_provider(channel) + if channel == "workbuddy": + models = workbuddy_fallback_models() + elif provider is not None: + models = list(provider.list_models()) + else: + models = [] + meta = refresh.get(channel) if isinstance(refresh.get(channel), dict) else {} + sources.append( + { + "channel": channel, + "display_name": getattr(provider, "display_name", channel) if provider else channel, + "mode": meta.get("mode") or ("fallback" if channel not in LIVE_FETCHERS else "static"), + "message": meta.get("message") or "", + "count": len(models), + "models": models, + "updated_at": meta.get("updated_at"), + } + ) + return {"sources": sources} diff --git a/docs/releases/v2.1.2.md b/docs/releases/v2.1.2.md new file mode 100644 index 0000000..09ea32e --- /dev/null +++ b/docs/releases/v2.1.2.md @@ -0,0 +1,25 @@ +# Buddy2api v2.1.2 + +发布日期:2026-08-28 + +管理页「模型配置」增加一键读取各通道供应模型,目录按通道保存,不再靠手改表格追新模型。 + +## 一键读取供应模型 + +- 一次操作覆盖当前启用的全部通道,并标明每个来源是在线读取还是回退本地/后台列表。 +- QClaw 走官方 cmd `4320`;TraeWork 走 `/api/remote/v1/models`,按官方 `data.list[].models[]`(`name` 为模型 id)解析,豆包等独有 id 只挂在 TraeWork。 +- WorkBuddy、QwenWork 没有供应商列表接口,回退现有后台或静态目录,不编造远程目录。 +- 目录按通道写入,不合并。`GET /v1/models` 仍是 WorkBuddy 裸 id 加 `workbuddy/`,其它通道只有 `channel/`。 +- 模型不属于当前 API Key 绑定的通道时,继续 400/403,不会换通道重试。 +- WorkBuddy 模型表和别名仍可手改。 + +## 升级说明 + +- 无数据库迁移;刷新结果写在 settings 里。 +- Docker 用户需要重新构建镜像并重启服务。 +- 思考强度策略未改:非法档 400;DeepSeek V4 未指定时默认 `high`;其它通道不补默认档。 + +## 验证 + +- 完整测试集:`241 passed`。 +- 已用本机账号实测一键读取:QClaw 在线 11 个模型;TraeWork 官方列表含豆包,补解析后可入库。 diff --git a/providers/qclaw/__init__.py b/providers/qclaw/__init__.py index 8357b6e..b9998e6 100644 --- a/providers/qclaw/__init__.py +++ b/providers/qclaw/__init__.py @@ -22,16 +22,19 @@ class QClawProvider: checkin_supported = False def list_models(self) -> list[dict]: - return [{"id": item} for item in STATIC_MODELS] + import catalog + + return catalog.models_for(self.id, [{"id": item} for item in STATIC_MODELS]) def alias_map(self) -> dict[str, str]: return dict(ALIASES) def accepts_model(self, inner: str) -> bool: value = (inner or "").strip() - if value in STATIC_MODELS or value in ALIASES: + if value in ALIASES or value.startswith("pool-"): return True - return value.startswith("pool-") + ids = {str(item.get("id")) for item in self.list_models() if isinstance(item, dict)} + return value in ids def translate_model(self, model: str) -> str: return chat.translate_model(model) diff --git a/providers/qclaw/jprx.py b/providers/qclaw/jprx.py index c523d9a..c90fd1b 100644 --- a/providers/qclaw/jprx.py +++ b/providers/qclaw/jprx.py @@ -129,24 +129,41 @@ async def time_sync(account: dict) -> str: return str(server_time or "") -async def list_remote_models(account: dict) -> list[dict]: - data, token = await post_cmd(CMD_MODEL_LIST, account) - apply_new_token(account, token) - rows = data.get("model_status_list") or data.get("models") or [] +def parse_model_list(data: dict) -> list[dict]: + rows = [] + if isinstance(data, dict): + rows = data.get("model_status_list") or data.get("models") or [] + elif isinstance(data, list): + rows = data models = [] + seen: set[str] = set() for row in rows: - if not isinstance(row, dict): + if isinstance(row, str): + mid = row.strip() + name = mid + description = "" + elif isinstance(row, dict): + mid = str(row.get("id") or row.get("model_id") or "").strip() + name = row.get("name") or row.get("display_id") or mid + description = row.get("description") or "" + else: continue - mid = str(row.get("id") or "").strip() - if not mid: + if not mid or mid in seen: continue - models.append( - { - "id": mid, - "name": row.get("name") or row.get("display_id") or mid, - "description": row.get("description") or "", - } - ) + seen.add(mid) + models.append({"id": mid, "name": name, "description": description}) + return models + + +async def fetch_supplier_models(account: dict) -> list[dict]: + """Live cmd 4320 list. Empty means no remote ids; caller decides fallback.""" + data, token = await post_cmd(CMD_MODEL_LIST, account) + apply_new_token(account, token) + return parse_model_list(data) + + +async def list_remote_models(account: dict) -> list[dict]: + models = await fetch_supplier_models(account) return models or [{"id": item, "name": item} for item in STATIC_MODELS] diff --git a/providers/qwenwork/__init__.py b/providers/qwenwork/__init__.py index a9e9b0e..da2fcb8 100644 --- a/providers/qwenwork/__init__.py +++ b/providers/qwenwork/__init__.py @@ -27,14 +27,19 @@ class QwenWorkProvider: checkin_supported = False def list_models(self) -> list[dict]: - return [{"id": item} for item in STATIC_MODELS] + import catalog + + return catalog.models_for(self.id, [{"id": item} for item in STATIC_MODELS]) def alias_map(self) -> dict[str, str]: return dict(ALIASES) def accepts_model(self, inner: str) -> bool: value = (inner or "").strip() - return value in STATIC_MODELS or value in ALIASES + if value in ALIASES: + return True + ids = {str(item.get("id")) for item in self.list_models() if isinstance(item, dict)} + return value in ids def translate_model(self, model: str) -> str: return chat.translate_model(model) diff --git a/providers/traework/__init__.py b/providers/traework/__init__.py index 01c4d57..0814dcb 100644 --- a/providers/traework/__init__.py +++ b/providers/traework/__init__.py @@ -18,7 +18,9 @@ class TraeWorkProvider: checkin_supported = True def list_models(self) -> list[dict]: - return [{"id": item} for item in STATIC_MODELS] + import catalog + + return catalog.models_for(self.id, [{"id": item} for item in STATIC_MODELS]) def alias_map(self) -> dict[str, str]: return dict(ALIASES) diff --git a/providers/traework/chat.py b/providers/traework/chat.py index 494368b..a28f296 100644 --- a/providers/traework/chat.py +++ b/providers/traework/chat.py @@ -30,8 +30,14 @@ def translate_model(model: str) -> str: def accepts_model(inner: str) -> bool: + import catalog + value = (inner or "").strip() - return value in STATIC_MODELS or value in ALIASES + if value in ALIASES: + return True + models = catalog.models_for(CHANNEL_ID, [{"id": item} for item in STATIC_MODELS]) + ids = {str(item.get("id")) for item in models if isinstance(item, dict)} + return value in ids def _last_user_text(payload: dict) -> str: diff --git a/providers/traework/models.py b/providers/traework/models.py new file mode 100644 index 0000000..b2f945c --- /dev/null +++ b/providers/traework/models.py @@ -0,0 +1,93 @@ +"""TraeWork supplier-list fetch. Isolated from chat I/O.""" + +from __future__ import annotations + +import httpx + +from providers.traework.constants import AGENT_API, MODELS_PATH +from providers.traework.token import TraeWorkAuthError, auth_headers + + +def parse_supplier_models(payload) -> list[dict]: + rows = _rows_from_payload(payload) + models: list[dict] = [] + seen: set[str] = set() + for row in rows: + if isinstance(row, str): + mid = row.strip() + name = mid + elif isinstance(row, dict): + mid = str( + row.get("id") + or row.get("model_name") + or row.get("name") + or row.get("model") + or "" + ).strip() + name = str(row.get("display_name") or row.get("name") or row.get("model_name") or mid) + else: + continue + if not mid or mid in seen: + continue + seen.add(mid) + models.append({"id": mid, "name": name or mid}) + return models + + +def _flatten_model_groups(rows: list) -> list: + """TraeWork returns function buckets: data.list[].models[], not a flat id list.""" + flattened: list = [] + for row in rows: + if not isinstance(row, dict): + flattened.append(row) + continue + nested = row.get("models") + grouped = isinstance(nested, list) and not ( + row.get("id") or row.get("model_name") or row.get("model") + ) + if grouped: + flattened.extend(nested) + continue + flattened.append(row) + return flattened + + +def _rows_from_payload(payload) -> list: + if isinstance(payload, list): + return _flatten_model_groups(payload) + if not isinstance(payload, dict): + return [] + data = payload.get("data") if isinstance(payload.get("data"), (dict, list)) else payload + if isinstance(data, list): + return _flatten_model_groups(data) + if not isinstance(data, dict): + return [] + for key in ("models", "model_list", "items", "list", "model_infos"): + rows = data.get(key) + if isinstance(rows, list): + return _flatten_model_groups(rows) + nested = data.get("data") + if isinstance(nested, list): + return _flatten_model_groups(nested) + if isinstance(nested, dict): + for key in ("models", "model_list", "items", "list"): + rows = nested.get(key) + if isinstance(rows, list): + return _flatten_model_groups(rows) + return [] + + +async def fetch_supplier_models(account: dict) -> list[dict]: + headers = auth_headers(account) + url = f"{AGENT_API}{MODELS_PATH}" + async with httpx.AsyncClient(timeout=30.0) as client: + response = await client.get(url, headers=headers) + if response.status_code >= 400: + raise TraeWorkAuthError(f"models HTTP {response.status_code}") + try: + payload = response.json() + except ValueError as exc: + raise TraeWorkAuthError("models response is not JSON") from exc + if isinstance(payload, dict) and payload.get("code") not in (None, 0): + raise TraeWorkAuthError(str(payload.get("message") or payload.get("code"))) + return parse_supplier_models(payload) diff --git a/server.py b/server.py index f43963c..86af03d 100644 --- a/server.py +++ b/server.py @@ -28,6 +28,7 @@ import database as db import auth_manager +import catalog import proxy import responses import providers @@ -285,14 +286,8 @@ async def health(): } -@app.get("/v1/models") -async def list_models( - authorization: str | None = Header(default=None), - x_api_key: str | None = Header(default=None, alias="X-Api-Key"), -): - await run_in_threadpool( - lambda: _check_client_auth(authorization, x_api_key, consume_quota=False) - ) +def collect_v1_models() -> list[dict]: + """Aggregate per-channel catalogs for GET /v1/models. WorkBuddy is bare + namespaced.""" data = [] workbuddy = providers.get_provider("workbuddy") wb_models = workbuddy.list_models() if workbuddy else db.get_setting("models", proxy.DEFAULT_MODELS) @@ -327,7 +322,18 @@ async def list_models( "owned_by": "buddy2api", "channel": channel, }) - return {"object": "list", "data": data} + return data + + +@app.get("/v1/models") +async def list_models( + authorization: str | None = Header(default=None), + x_api_key: str | None = Header(default=None, alias="X-Api-Key"), +): + await run_in_threadpool( + lambda: _check_client_auth(authorization, x_api_key, consume_quota=False) + ) + return {"object": "list", "data": collect_v1_models()} @app.post("/v1/chat/completions") @@ -990,6 +996,18 @@ async def admin_get_models(authorization: str | None = Header(default=None)): return db.get_setting("models", proxy.DEFAULT_MODELS) +@app.get("/admin/models/catalogs") +async def admin_get_model_catalogs(authorization: str | None = Header(default=None)): + _check_admin(authorization) + return catalog.catalog_snapshot() + + +@app.post("/admin/models/refresh") +async def admin_refresh_models(authorization: str | None = Header(default=None)): + _check_admin(authorization) + return await catalog.refresh_supplier_catalogs() + + @app.put("/admin/models") async def admin_update_models( request: Request, diff --git a/tests/test_models_refresh.py b/tests/test_models_refresh.py new file mode 100644 index 0000000..9f83ffb --- /dev/null +++ b/tests/test_models_refresh.py @@ -0,0 +1,269 @@ +import asyncio +import json +import os +from pathlib import Path + +import pytest + +import credential_crypto +import database as db +import providers +import proxy +import router +import server +from providers.protocol import KeyChannelMismatch, UnknownModel +from providers.qclaw.constants import STATIC_MODELS as QCLAW_STATIC +from providers.qwenwork.constants import STATIC_MODELS as QWEN_STATIC +from providers.traework.constants import STATIC_MODELS as TRAE_STATIC + +QCLAW_NEW_ID = "qclaw-live-only-model" +TRAE_NEW_DOUBAO = "Doubao-Seed-2.2-Pro" +TRAE_DOUBAO_TURBO = "Doubao-Seed-2.1-Turbo" +TRAE_DOUBAO_CODE = "Doubao-Seed-2.0-Code" + +QCLAW_HTTP_PAYLOAD = { + "ret": 0, + "data": { + "resp": { + "common": {"code": 0}, + "data": { + "model_status_list": [ + {"id": "default", "name": "Default"}, + {"id": "pool-glm-5.2", "name": "GLM 5.2"}, + {"id": QCLAW_NEW_ID, "name": "QClaw Live Only"}, + ] + }, + } + }, +} + +TRAEWORK_HTTP_PAYLOAD = { + "code": 0, + "message": "success", + "data": { + "list": [ + { + "function": "solo_coder", + "models": [ + {"name": TRAE_DOUBAO_CODE, "display_name": TRAE_DOUBAO_CODE}, + {"name": TRAE_DOUBAO_TURBO, "display_name": TRAE_DOUBAO_TURBO}, + {"name": TRAE_NEW_DOUBAO, "display_name": TRAE_NEW_DOUBAO}, + {"name": "qwen-3.6-plus", "display_name": "qwen-3.6-plus"}, + ], + } + ] + }, +} + + +class _FakeResponse: + def __init__(self, payload, status=200): + self.status_code = status + self._payload = payload + self.headers = {} + self.content = b"{}" + + def json(self): + return self._payload + + +def _ids(models): + return {str(item.get("id") if isinstance(item, dict) else item) for item in models} + + +@pytest.fixture() +def isolated_db(tmp_path, monkeypatch): + path = tmp_path / "gateway.db" + monkeypatch.setattr(db, "DB_PATH", path) + monkeypatch.setenv("CB_GATEWAY_MASTER_KEY", "pytest-master-key") + credential_crypto.reset_cache() + db.init_db() + yield path + credential_crypto.reset_cache() + + +@pytest.fixture() +def all_channels(monkeypatch): + monkeypatch.setenv("CB_GATEWAY_PROVIDERS", "workbuddy,qclaw,qwenwork,traework") + yield + monkeypatch.delenv("CB_GATEWAY_PROVIDERS", raising=False) + + +def _seed_live_accounts(): + db.add_account( + { + "name": "qc", + "uid": "qc-1", + "provider": "qclaw", + "status": "active", + "access_token": "sk-qclaw", + "refresh_token": "jwt-qclaw", + "extra": {"guid": "guid-1"}, + } + ) + db.add_account( + { + "name": "tw", + "uid": "tw-1", + "provider": "traework", + "status": "active", + "access_token": "tok-trae", + "expires_at": 9_999_999_999_999, + "extra": {"device_id": "dev-1"}, + } + ) + + +def _install_supplier_http(monkeypatch): + requested = [] + + class FakeAsyncClient: + def __init__(self, *args, **kwargs): + pass + + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + async def post(self, url, **kwargs): + requested.append(("POST", str(url))) + if "/data/4320/forward" in str(url): + return _FakeResponse(QCLAW_HTTP_PAYLOAD) + raise AssertionError(f"unexpected POST {url}") + + async def get(self, url, **kwargs): + requested.append(("GET", str(url))) + if "/api/remote/v1/models" in str(url): + return _FakeResponse(TRAEWORK_HTTP_PAYLOAD) + raise AssertionError(f"unexpected GET {url}") + + monkeypatch.setattr("providers.qclaw.jprx.httpx.AsyncClient", FakeAsyncClient) + monkeypatch.setattr("providers.traework.models.httpx.AsyncClient", FakeAsyncClient) + return requested + + +def _by_channel(result): + return {item["channel"]: item for item in result["sources"]} + + +def test_admin_models_page_has_one_click_control(): + html = (Path(__file__).resolve().parents[1] / "web" / "index.html").read_text(encoding="utf-8") + assert "一键读取供应模型" in html + assert "/admin/models/refresh" in html + assert "syncSources" in html + + +def test_supplier_catalog_refresh_keeps_channels_distinct(isolated_db, all_channels, monkeypatch): + assert QCLAW_NEW_ID not in QCLAW_STATIC + assert TRAE_NEW_DOUBAO not in TRAE_STATIC + assert TRAE_DOUBAO_TURBO in TRAE_STATIC + + qclaw = providers.get_provider("qclaw") + traework = providers.get_provider("traework") + workbuddy = providers.get_provider("workbuddy") + qwenwork = providers.get_provider("qwenwork") + + assert QCLAW_NEW_ID not in _ids(qclaw.list_models()) + assert TRAE_NEW_DOUBAO not in _ids(traework.list_models()) + assert not qclaw.accepts_model(QCLAW_NEW_ID) + assert not traework.accepts_model(TRAE_NEW_DOUBAO) + + _seed_live_accounts() + requested = _install_supplier_http(monkeypatch) + + chat_hits = [] + + async def boom_chat(payload, api_key_info): + chat_hits.append((payload.get("model"), api_key_info)) + raise AssertionError("chat should not run") + + for channel in providers.enabled_provider_ids(): + monkeypatch.setattr(providers.get_provider(channel), "chat_completions", boom_chat) + + monkeypatch.setattr(server, "ALLOW_NO_ADMIN_AUTH", True) + result = asyncio.run(server.admin_refresh_models()) + sources = _by_channel(result) + + assert sources["qclaw"]["mode"] == "live" + assert sources["traework"]["mode"] == "live" + assert sources["workbuddy"]["mode"] == "fallback" + assert sources["workbuddy"]["message"] == "no supplier-list API" + assert sources["qwenwork"]["mode"] == "fallback" + assert sources["qwenwork"]["message"] == "no supplier-list API" + + assert QCLAW_NEW_ID in _ids(sources["qclaw"]["models"]) + assert TRAE_NEW_DOUBAO in _ids(sources["traework"]["models"]) + assert TRAE_DOUBAO_TURBO in _ids(sources["traework"]["models"]) + + wb_ids = _ids(sources["workbuddy"]["models"]) + assert QCLAW_NEW_ID not in wb_ids + assert TRAE_NEW_DOUBAO not in wb_ids + assert TRAE_DOUBAO_TURBO not in wb_ids + assert wb_ids == _ids(proxy.DEFAULT_MODELS) + + qwen_ids = _ids(sources["qwenwork"]["models"]) + assert qwen_ids == set(QWEN_STATIC) + assert TRAE_NEW_DOUBAO not in qwen_ids + assert QCLAW_NEW_ID not in qwen_ids + + assert any("/data/4320/forward" in url for method, url in requested if method == "POST") + assert any("/api/remote/v1/models" in url for method, url in requested if method == "GET") + assert not any("copilot.tencent.com" in url for _, url in requested) + assert not any("qwenwork.cn" in url for _, url in requested) + + assert QCLAW_NEW_ID in _ids(qclaw.list_models()) + assert TRAE_NEW_DOUBAO in _ids(traework.list_models()) + assert qclaw.accepts_model(QCLAW_NEW_ID) + assert traework.accepts_model(TRAE_NEW_DOUBAO) + assert traework.accepts_model(TRAE_DOUBAO_TURBO) + assert not workbuddy.accepts_model(TRAE_NEW_DOUBAO) + assert not workbuddy.accepts_model(TRAE_DOUBAO_TURBO) + assert not qclaw.accepts_model(TRAE_NEW_DOUBAO) + assert not qwenwork.accepts_model(TRAE_NEW_DOUBAO) + assert not qwenwork.accepts_model(QCLAW_NEW_ID) + + bound = router.bind({"model": "traework/" + TRAE_NEW_DOUBAO}, {"default_channel": "traework"}) + assert bound.channel == "traework" + assert bound.inner == TRAE_NEW_DOUBAO + bound = router.bind({"model": "qclaw/" + QCLAW_NEW_ID}, {"default_channel": "qclaw"}) + assert bound.channel == "qclaw" + assert bound.inner == QCLAW_NEW_ID + + def attempt_chat(payload, key): + bound = router.bind(payload, key) + return asyncio.run(router.chat_after_bind(bound, payload, key)) + + with pytest.raises(UnknownModel): + attempt_chat({"model": TRAE_NEW_DOUBAO, "messages": [{"role": "user", "content": "hi"}]}, {"default_channel": "workbuddy"}) + with pytest.raises(UnknownModel): + attempt_chat({"model": TRAE_DOUBAO_TURBO, "messages": [{"role": "user", "content": "hi"}]}, {"default_channel": "workbuddy"}) + with pytest.raises(KeyChannelMismatch): + attempt_chat({"model": "traework/" + TRAE_NEW_DOUBAO, "messages": [{"role": "user", "content": "hi"}]}, {"default_channel": "workbuddy"}) + with pytest.raises(UnknownModel): + attempt_chat({"model": QCLAW_NEW_ID, "messages": [{"role": "user", "content": "hi"}]}, {"default_channel": "workbuddy"}) + + assert chat_hits == [] + + payload = {"object": "list", "data": server.collect_v1_models()} + data = payload["data"] + by_id = {item["id"]: item for item in data} + + assert by_id["qclaw/" + QCLAW_NEW_ID]["channel"] == "qclaw" + assert by_id["traework/" + TRAE_NEW_DOUBAO]["channel"] == "traework" + assert by_id["traework/" + TRAE_DOUBAO_TURBO]["channel"] == "traework" + assert TRAE_NEW_DOUBAO not in by_id + assert TRAE_DOUBAO_TURBO not in by_id + assert QCLAW_NEW_ID not in by_id + assert "glm-5.2" in by_id + assert by_id["glm-5.2"]["channel"] == "workbuddy" + assert by_id["workbuddy/glm-5.2"]["channel"] == "workbuddy" + + evidence = os.environ.get("BUDDY2API_EVIDENCE_DIR") + if evidence: + Path(evidence).mkdir(parents=True, exist_ok=True) + Path(evidence, "v1-models.json").write_text( + json.dumps(payload, ensure_ascii=False, indent=2), + encoding="utf-8", + ) diff --git a/tests/test_traework.py b/tests/test_traework.py index 95d4930..61cace4 100644 --- a/tests/test_traework.py +++ b/tests/test_traework.py @@ -73,6 +73,42 @@ def test_bind_traework_when_enabled(traework_enabled): router.bind({"model": "glm-5.2"}, {"default_channel": "traework"}) +def test_parse_supplier_models_official_grouped_list(): + from providers.traework.models import parse_supplier_models + + parsed = parse_supplier_models( + { + "code": 0, + "message": "success", + "data": { + "list": [ + { + "function": "solo_coder", + "models": [ + { + "name": "Doubao-Seed-2.0-Code", + "display_name": "Doubao-Seed-2.0-Code", + "is_default": False, + }, + { + "name": "Doubao-Seed-Code", + "display_name": "Doubao-Seed-Code", + }, + { + "name": "qwen-3.6-plus", + "display_name": "qwen-3.6-plus", + }, + ], + } + ] + }, + } + ) + ids = [item["id"] for item in parsed] + assert ids == ["Doubao-Seed-2.0-Code", "Doubao-Seed-Code", "qwen-3.6-plus"] + assert "function" not in ids + + def test_translate_auto(): assert translate_model("auto") == "qwen-3.7-plus" diff --git a/version.py b/version.py index 5b0431e..b777579 100644 --- a/version.py +++ b/version.py @@ -1 +1 @@ -VERSION = "2.1.1" +VERSION = "2.1.2" diff --git a/web/index.html b/web/index.html index f274700..5c56bd8 100644 --- a/web/index.html +++ b/web/index.html @@ -388,7 +388,7 @@ template:`
-
B2
Buddy 2 API
Local model gateway · v2.1.1
+
B2
Buddy 2 API
Local model gateway · v2.1.2
@@ -813,11 +813,12 @@

自定义路径

` }).component('mdls',{props:['token','toast'],setup(p){ - const m=ref([]),al=ref({}),ld=ref(true); + const m=ref([]),al=ref({}),ld=ref(true),sources=ref([]),syncing=ref(false); const BUILTIN=['gpt-4o','gpt-4o-mini','gpt-4-turbo','gpt-4','gpt-3.5-turbo','claude-3.5-sonnet','claude-3-haiku','deepseek-chat','deepseek-coder','moonshot-v1-128k','moonshot-v1-32k']; const newKey=ref(''),newVal=ref(''); - async function load(){ld.value=true;try{m.value=await api.get('/admin/models',p.token);al.value=await api.get('/admin/aliases',p.token)}catch(e){p.toast('失败','err')}ld.value=false} + async function load(){ld.value=true;try{m.value=await api.get('/admin/models',p.token);al.value=await api.get('/admin/aliases',p.token);try{const snap=await api.get('/admin/models/catalogs',p.token);sources.value=snap.sources||[]}catch(e){sources.value=[]}}catch(e){p.toast('失败','err')}ld.value=false} async function save(){try{await api.put('/admin/models',m.value,p.token);await api.put('/admin/aliases',al.value,p.token);p.toast('已保存')}catch(e){p.toast('失败','err')}} + async function syncSources(){if(syncing.value)return;syncing.value=true;try{const r=await api.post('/admin/models/refresh',{},p.token);sources.value=r.sources||[];p.toast('已读取各通道供应模型');try{m.value=await api.get('/admin/models',p.token)}catch(e){}}catch(e){p.toast(apiErr(e,'读取供应模型失败'),'err')}syncing.value=false} function add(){m.value.push({id:'',name:''})}function rm(i){m.value.splice(i,1)} function isBuiltin(k){return BUILTIN.includes(k)} function addAlias(){ @@ -828,14 +829,28 @@

自定义路径

al.value={...al.value,[k]:v};newKey.value='';newVal.value='';p.toast('已添加') } function rmAlias(k){const o={...al.value};delete o[k];al.value=o} - onMounted(load);return{m,al,ld,load,save,add,rm,newKey,newVal,addAlias,rmAlias,isBuiltin,I} + function modeText(s){return s.mode==='live'?'在线读取':(s.mode==='fallback'?'回退本地':'未读取')} + function modeClass(s){return s.mode==='live'?'ok':(s.mode==='fallback'?'warn':'')} + const otherSources=computed(()=>sources.value.filter(s=>s.channel!=='workbuddy')); + onMounted(load);return{m,al,ld,sources,otherSources,syncing,load,save,syncSources,add,rm,newKey,newVal,addAlias,rmAlias,isBuiltin,modeText,modeClass,I} },template:`
-

模型配置

/v1/models 列表 & 别名映射

-
+

模型配置

/v1/models 按通道列出;一键读取各来源供应模型,独有模型(如 TraeWork 豆包)只挂在该通道

+
+
各通道供应模型在线读取或回退本地/后台列表,互不合并
+ + +
通道方式数量说明
{{s.display_name||s.channel}} {{s.channel}}{{modeText(s)}}{{s.count||0}}{{s.message||'—'}}
+
+
{{s.display_name||s.channel}}仅 {{s.channel}}/id ,客户端选错通道会直接拒绝
+ + + +
ID名称
{{s.channel}}/{{x.id}}{{x.name||x.id}}
暂无
+