diff --git a/.env.example b/.env.example
index 86f9a43..63bfb21 100644
--- a/.env.example
+++ b/.env.example
@@ -10,5 +10,11 @@ OPENAI_API_KEY=sk-xxx
# CODEABC_MODEL=deepseek/deepseek-chat
# CODEABC_MODEL=openrouter/anthropic/claude-haiku-4.5
+# Using an OpenAI-compatible gateway (Qwen Token Plan, MiniMax, a local Ollama)?
+# Point the base URL at it: an openai/-prefixed model then goes to your gateway
+# instead of platform.openai.com. The browser form's Base URL field does the
+# same thing per request and wins over this value.
+# OPENAI_API_BASE=https://gateway.example.com/v1
+
# Frontend origin for CORS (default: http://localhost:5173)
# FRONTEND_ORIGIN=https://your-domain.com
diff --git a/README.md b/README.md
index f21f6f1..6fd016f 100644
--- a/README.md
+++ b/README.md
@@ -177,7 +177,12 @@ npm run tauri:dev
CodeABC supports two modes:
- **Free mode** (default): Limited to 20 requests per day
-- **BYOK mode**: Click the gear icon in the top-right corner to enter your own API key for unlimited use. The key is stored only in your browser's localStorage.
+- **BYOK mode**: Click the gear icon in the top-right corner for unlimited use. The key is stored only in your browser's localStorage. The form walks you through three steps:
+ 1. **Endpoint** — paste a Base URL, pick the API format (`OpenAI-compatible` or `Anthropic-compatible`), then paste your key. The Base URL is normalised for you: a missing `/v1` is added, a trailing `/v1` on an Anthropic-style base is stripped.
+ 2. **Model** — *Fetch model list* asks the gateway what it serves. Pick from the dropdown when it answers (non-chat models are greyed out with the reason); type the model id when it does not, which is the norm for Anthropic-compatible gateways.
+ 3. **Verify** — *Test connection* pings only the model you selected with a 1-token request, so it works whether or not the gateway exposes a model list. Keys come back masked.
+
+ An OpenRouter key still needs nothing but the key: paste it and CodeABC picks a fast, inexpensive model for you. Server-side configuration uses `OPENAI_API_BASE` for the endpoint (see `.env.example`).
## Project Structure
diff --git a/README_CN.md b/README_CN.md
index 666d95b..2ef4aa7 100644
--- a/README_CN.md
+++ b/README_CN.md
@@ -188,7 +188,12 @@ npm run tauri:dev
码上懂支持两种模式:
- **免费模式**(默认):每天 20 次调用
-- **自带 Key 模式**:点击右上角齿轮图标,填入你自己的 API Key,无限使用。Key 只存在浏览器本地,不会上传。
+- **自带 Key 模式**:点击右上角齿轮图标,无限使用。Key 只存在浏览器本地,不会上传。表单分三步:
+ 1. **连接信息**——填 Base URL、选 API 格式(`OpenAI 兼容` 或 `Anthropic 兼容`),再填 Key。Base URL 会自动归一:缺 `/v1` 补上,Anthropic 系尾部的 `/v1` 剥掉。
+ 2. **模型**——点「获取模型列表」问网关到底提供哪些模型:能答就下拉选(非对话模型会置灰并注明原因),答不了就手填模型 ID(Anthropic 兼容网关通常没这个接口)。
+ 3. **验证**——「测试连接」只用 1 个 token 的真实请求 ping 你选中的那个模型,所以网关有没有模型列表都能测。返回的 Key 一律脱敏。
+
+ OpenRouter 的 key 依旧只填 key 就行:粘进去,码上懂自动挑一个又快又便宜的模型。服务端配置用 `OPENAI_API_BASE` 指定端点(见 `.env.example`)。
## 路线图
diff --git a/backend/app.py b/backend/app.py
index 935e32d..e981421 100644
--- a/backend/app.py
+++ b/backend/app.py
@@ -10,7 +10,7 @@
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse
-from backend.routers import analyze, project
+from backend.routers import analyze, project, providers
from backend.services.cache import init_db
@@ -53,6 +53,9 @@ async def lifespan(app: FastAPI):
# needs to match before project's greedy /file/{path:path} route
app.include_router(analyze.router, prefix="/api")
app.include_router(project.router, prefix="/api")
+# providers routes (/providers, /providers/check, /models) are discovery-only:
+# no LLM generation, no overlap with project's greedy /file/{path:path}.
+app.include_router(providers.router, prefix="/api")
@app.get("/api/health")
diff --git a/backend/models.py b/backend/models.py
index b77d8de..1642a58 100644
--- a/backend/models.py
+++ b/backend/models.py
@@ -546,6 +546,16 @@ class EditRequest(BaseModel):
language: str = ""
+class ProtocolCheckRequest(BaseModel):
+ """Connection-check probe body; protocol, base and model expected."""
+
+ protocol: str | None = None
+ api_base: str | None = None
+ api_key: str | None = None
+ # The single model the check validates with a 1-token request.
+ model: str | None = None
+
+
class GlossaryTerm(BaseModel):
term: str
definition: str
diff --git a/backend/routers/analyze.py b/backend/routers/analyze.py
index 64761e0..d926374 100644
--- a/backend/routers/analyze.py
+++ b/backend/routers/analyze.py
@@ -90,6 +90,23 @@ async def _enforce_rate_limit(request: Request):
await cache.increment_rate_limit(ip)
+def _llm_kwargs(request: Request) -> dict[str, str | None]:
+ """BYOK overrides read from request headers, passed straight to the LLM.
+
+ All four are optional and backward-compatible: with no headers set every
+ value is ``None``, so llm falls back to its env/default resolution and
+ existing callers are unaffected. ``x-api-base``/``x-model``/``x-api-protocol``
+ let the settings form aim a bring-your-own-key gateway (e.g. a Qwen
+ OpenAI-compatible endpoint) at a chosen base URL, model, and wire protocol.
+ """
+ return {
+ "api_key": request.headers.get("x-api-key"),
+ "api_base": request.headers.get("x-api-base"),
+ "model": request.headers.get("x-model"),
+ "protocol": request.headers.get("x-api-protocol"),
+ }
+
+
@router.get("/project/{project_id}/overview")
async def get_overview(project_id: str, request: Request):
"""Generate or return cached project overview. Streams SSE."""
@@ -115,12 +132,12 @@ async def cached_stream():
await _enforce_rate_limit(request)
prompt = build_overview_prompt(proj["files"])
- api_key = request.headers.get("x-api-key")
+ llm_kwargs = _llm_kwargs(request)
async def generate():
full_response = ""
errored = False
- async for chunk in stream_llm(prompt, api_key=api_key):
+ async for chunk in stream_llm(prompt, **llm_kwargs):
full_response += chunk
# Don't stream a raw "[LLM Error: ...]" into the reader's view; once
# we see the sentinel, hold it back and report it as an error below.
@@ -175,9 +192,8 @@ async def get_annotations(project_id: str, file_path: str, request: Request):
await _enforce_rate_limit(request)
prompt = build_annotation_prompt(content, lang)
- api_key = request.headers.get("x-api-key")
- result = await call_llm(prompt, api_key=api_key)
+ result = await call_llm(prompt, **_llm_kwargs(request))
parsed = _extract_json(result)
annotations = _coerce_annotation_list(parsed)
@@ -217,8 +233,7 @@ async def ask_question(project_id: str, req: QARequest, request: Request):
file_path=req.file_path,
language=req.language,
)
- api_key = request.headers.get("x-api-key")
- answer = (await call_llm(prompt, api_key=api_key)).strip()
+ answer = (await call_llm(prompt, **_llm_kwargs(request))).strip()
if is_error_text(answer):
raise HTTPException(502, answer)
@@ -259,8 +274,7 @@ async def edit_code(project_id: str, req: EditRequest, request: Request):
file_path=req.file_path,
language=req.language,
)
- api_key = request.headers.get("x-api-key")
- raw = await call_llm(prompt, api_key=api_key)
+ raw = await call_llm(prompt, **_llm_kwargs(request))
if is_error_text(raw):
raise HTTPException(502, raw)
edited = extract_code_block(raw)
diff --git a/backend/routers/providers.py b/backend/routers/providers.py
new file mode 100644
index 0000000..63439e7
--- /dev/null
+++ b/backend/routers/providers.py
@@ -0,0 +1,291 @@
+"""Protocol registry, connection check, and live model discovery endpoints.
+
+These back the settings form: the frontend lists wire protocols, probes a
+gateway before its credentials are saved, and discovers which models an
+endpoint actually serves (merged with litellm capability metadata).
+
+The probe paths touch litellm only through the capability helpers, which import
+it lazily; httpx (already a litellm dependency) is imported inside the handlers,
+so this module carries no heavy import of its own. These routes are discovery
+and connection utilities that never run an LLM generation, so they are
+deliberately not rate-limited (unlike the analyze endpoints' free-tier cap).
+"""
+
+from __future__ import annotations
+
+from fastapi import APIRouter, Header
+
+from backend.models import ProtocolCheckRequest
+from backend.services.providers import (
+ Protocol,
+ get_protocol,
+ list_protocols,
+ mask_key,
+ model_metadata,
+ normalize_base_url,
+ redact_secret,
+)
+
+router = APIRouter(tags=["providers"])
+
+# litellm "mode" values that mean a model cannot answer a chat turn.
+_NON_CHAT_MODES = {
+ "image_generation",
+ "audio",
+ "tts",
+ "embedding",
+ "embeddings",
+ "rerank",
+ "moderation",
+}
+# Substrings that hint a model is non-chat when litellm's catalog is silent.
+# Deliberately conservative: only strong, unambiguous signals.
+_NON_CHAT_HINTS = (
+ "image",
+ "tts",
+ "realtime",
+ "embedding",
+ "rerank",
+ "whisper",
+ "dall-e",
+ "stable-diffusion",
+ "-audio",
+ "speech",
+)
+
+
+def _protocol_payload(protocol: Protocol) -> dict:
+ """Serialise a protocol for the settings form (no secrets)."""
+ return {
+ "id": protocol.id,
+ "label": protocol.label,
+ "wire": protocol.wire,
+ "label_zh": protocol.label_zh,
+ "wire_zh": protocol.wire_zh,
+ "litellm_prefix": protocol.litellm_prefix,
+ "models_path": protocol.models_path,
+ "base_hint": protocol.base_hint,
+ }
+
+
+def _auth_headers(protocol: Protocol, api_key: str) -> dict:
+ """Build the auth header a protocol expects, or {} when no key given."""
+ if not api_key:
+ return {}
+ if protocol.auth_header == "x-api-key":
+ return {"x-api-key": api_key}
+ return {"Authorization": f"Bearer {api_key}"}
+
+
+def _classify_chat(model_id: str, mode: str | None) -> tuple[bool, str, str]:
+ """Decide whether a model is chat-capable, and say why we think so.
+
+ Returns ``(chat, reason, reason_zh)``: the form greys a non-chat model out
+ and shows the reason inline, so it needs both languages.
+ """
+ if mode is not None:
+ return (
+ mode not in _NON_CHAT_MODES,
+ f"litellm reports mode '{mode}'",
+ f"litellm 标为 '{mode}' 类型",
+ )
+ lowered = model_id.lower()
+ for hint in _NON_CHAT_HINTS:
+ if hint in lowered:
+ return (
+ False,
+ f"name suggests a non-chat model ('{hint}')",
+ f"名称不像对话模型(含 '{hint}')",
+ )
+ return (
+ True,
+ "assumed chat-capable (not in litellm catalog)",
+ "litellm 目录里没有,按可对话处理",
+ )
+
+
+def _failure_diagnosis(status_code: int, url: str) -> tuple[str, str]:
+ """Translate a non-200 probe into an actionable, key-free message.
+
+ Returns the English and Chinese forms of the same diagnosis.
+ """
+ if status_code in (401, 403):
+ return (
+ f"Rejected the API key ({status_code}). Check the key, or that it "
+ "is valid for this endpoint.",
+ f"该端点拒了这个 API Key({status_code})。请确认 Key 没填错,"
+ "且对这个端点有效。",
+ )
+ if status_code == 404:
+ return (
+ f"No endpoint at {url} (404). The Base URL, API format or model "
+ "id may be wrong.",
+ f"{url} 上没有这个接口(404)。Base URL、API 格式或模型 ID 可能填错了。",
+ )
+ return (
+ f"Unexpected status {status_code} from {url}.",
+ f"{url} 返回了意外的状态码 {status_code}。",
+ )
+
+
+def _model_entry(entry: object, prefix: str) -> dict | None:
+ """Build one form row from a models-list entry, or None if it has no id."""
+ model_id = entry.get("id") if isinstance(entry, dict) else str(entry)
+ if not model_id:
+ return None
+ litellm_model = f"{prefix}{model_id}"
+ meta = model_metadata(litellm_model)
+ chat, chat_reason, chat_reason_zh = _classify_chat(model_id, meta.get("mode"))
+ return {
+ "id": model_id,
+ "litellm_model": litellm_model,
+ "chat": chat,
+ "chat_reason": chat_reason,
+ "chat_reason_zh": chat_reason_zh,
+ "mode": meta.get("mode"),
+ "max_input_tokens": meta.get("max_input_tokens"),
+ "max_output_tokens": meta.get("max_output_tokens"),
+ "input_cost_per_token": meta.get("input_cost_per_token"),
+ "output_cost_per_token": meta.get("output_cost_per_token"),
+ "metadata_source": meta.get("metadata_source"),
+ }
+
+
+async def _fetch_models(
+ base: str, models_path: str, headers: dict
+) -> tuple[list, int | None, str | None]:
+ """Fetch and unwrap a gateway's models list.
+
+ Returns ``(entries, status, error)``: on success ``error`` is None; on any
+ failure ``entries`` is empty and ``error`` is a key-free message. The HTTP
+ status lets callers tell "no models endpoint" (404) from "bad key" (401).
+ """
+ import httpx
+
+ status: int | None = None
+ try:
+ async with httpx.AsyncClient(timeout=15.0) as client:
+ resp = await client.get(f"{base}{models_path}", headers=headers)
+ status = resp.status_code
+ resp.raise_for_status()
+ payload = resp.json()
+ except Exception as exc: # network / HTTP / JSON errors all become a message
+ return [], status, str(exc)
+ data = payload.get("data", payload) if isinstance(payload, dict) else payload
+ if not isinstance(data, list):
+ return [], status, "Unexpected models response shape."
+ return data, status, None
+
+
+async def _ping_chat(
+ base: str, protocol: Protocol, headers: dict, model: str
+) -> tuple[int | None, str | None, str | None]:
+ """Send a 1-token request to validate auth when no models list exists.
+
+ Some Anthropic-compatible gateways (e.g. Alibaba Bailian's ``/apps/anthropic``)
+ expose ``/v1/messages`` but no ``/v1/models``, so a models probe 404s even
+ with a valid key. A minimal generation against a caller-supplied model is
+ the only honest way to confirm the credential there.
+
+ Returns ``(status, diagnosis, diagnosis_zh)``; both are None on success.
+ """
+ import httpx
+
+ if protocol.id == "anthropic_messages":
+ url = f"{base}/v1/messages"
+ else:
+ url = f"{base}/chat/completions"
+ body = {
+ "model": model,
+ "max_tokens": 1,
+ "messages": [{"role": "user", "content": "ping"}],
+ }
+ try:
+ async with httpx.AsyncClient(timeout=30.0) as client:
+ resp = await client.post(url, json=body, headers=headers)
+ except httpx.HTTPError as exc:
+ return None, f"Could not reach {url}: {exc}", f"连不上 {url}:{exc}"
+ if resp.status_code == 200:
+ return 200, None, None
+ diagnosis, diagnosis_zh = _failure_diagnosis(resp.status_code, url)
+ return resp.status_code, diagnosis, diagnosis_zh
+
+
+@router.get("/protocols")
+async def protocols_list() -> dict:
+ """List every supported wire protocol for the settings form."""
+ return {"protocols": [_protocol_payload(p) for p in list_protocols()]}
+
+
+@router.post("/protocols/check")
+async def check_protocol(req: ProtocolCheckRequest) -> dict:
+ """Validate base URL + key by pinging the selected model only.
+
+ A 1-token request against the caller's model is the honest test: it works
+ on gateways with or without a models list. The key is never echoed back:
+ only a masked form and a plain diagnosis.
+ """
+ protocol = get_protocol(req.protocol or "")
+ verdict: dict = {
+ "ok": False,
+ "base_url": "",
+ "status_code": None,
+ "masked_key": mask_key(req.api_key or ""),
+ "diagnosis": None,
+ "diagnosis_zh": None,
+ }
+ if protocol is None:
+ verdict["diagnosis"] = "Unknown API format."
+ verdict["diagnosis_zh"] = "未知的 API 格式。"
+ return verdict
+ if not (req.api_base or "").strip():
+ verdict["diagnosis"] = "A Base URL is required."
+ verdict["diagnosis_zh"] = "必须填写 Base URL。"
+ return verdict
+ if not (req.model or "").strip():
+ verdict["diagnosis"] = "Select or enter a model first; only that model is tested."
+ verdict["diagnosis_zh"] = "请先选择或填写模型;测试只会打到这一个模型。"
+ return verdict
+
+ base = normalize_base_url(req.api_base or "", protocol.id)
+ verdict["base_url"] = base
+ headers = _auth_headers(protocol, req.api_key or "")
+ status, diagnosis, diagnosis_zh = await _ping_chat(
+ base, protocol, headers, (req.model or "").strip()
+ )
+ verdict["status_code"] = status
+ if diagnosis is None:
+ verdict["ok"] = True
+ return verdict
+ key = req.api_key or ""
+ verdict["diagnosis"] = redact_secret(diagnosis, key)
+ verdict["diagnosis_zh"] = redact_secret(diagnosis_zh, key)
+ return verdict
+
+
+@router.get("/models")
+async def discover_models(
+ protocol: str | None = None,
+ api_base: str | None = None,
+ x_api_key: str | None = Header(None),
+) -> dict:
+ """List the models an endpoint serves, merged with capability metadata.
+
+ Combines the gateway's live models list with litellm's catalog (context
+ window, output budget, cost, mode). Non-chat models are flagged so the
+ form can disable them with a reason; unknown ones are reported honestly.
+ """
+ proto = get_protocol(protocol or "")
+ if proto is None or not (api_base or "").strip():
+ return {"models": [], "count": 0, "discoverable": False,
+ "error": "A protocol and Base URL are required."}
+
+ base = normalize_base_url(api_base or "", proto.id)
+ headers = _auth_headers(proto, x_api_key or "")
+ data, status, error = await _fetch_models(base, proto.models_path, headers)
+ if error:
+ return {"models": [], "count": 0, "discoverable": False, "status_code": status,
+ "error": redact_secret(error, x_api_key or "")}
+
+ models = [m for m in (_model_entry(e, proto.litellm_prefix) for e in data) if m is not None]
+ return {"models": models, "count": len(models), "discoverable": True}
diff --git a/backend/services/llm.py b/backend/services/llm.py
index 4e299ee..aa50faa 100644
--- a/backend/services/llm.py
+++ b/backend/services/llm.py
@@ -8,6 +8,13 @@
import litellm
+from backend.services.providers import (
+ normalize_base_url,
+ redact_secret,
+ resolve_litellm_model,
+ suggest_max_tokens,
+)
+
# Some models reject a custom temperature: gpt-5 / gpt-5-mini / gpt-5-codex and
# the o1 / o3 reasoning families only accept the default (temperature=1) and
# error on `temperature=0.3`. Since gpt-5-mini is our default OpenAI model, let
@@ -30,16 +37,22 @@
_DEFAULT_OPENROUTER_MODEL = "openrouter/deepseek/deepseek-v4-flash"
-def _resolve_model(api_key: str | None = None, model: str | None = None) -> str:
- """Pick the litellm model string, making a pasted OpenRouter key just work.
+def _resolve_model(
+ api_key: str | None = None,
+ model: str | None = None,
+ protocol: str | None = None,
+) -> str:
+ """Pick the litellm model string for a request.
- An explicit model (the ``model`` argument or the ``CODEABC_MODEL`` env var)
- always wins. Otherwise the key's shape decides the provider: an OpenRouter
- key (``sk-or-...``) routes through OpenRouter, so a non-technical user only
- has to paste their key — there is no provider or ``openrouter/`` prefix to
- remember. Anything else falls back to the OpenAI default.
+ With an explicit ``protocol`` (the settings form's BYOK path) the model is
+ a bare gateway id qualified by that protocol's litellm prefix. Without one
+ the legacy key-shape routing applies: an explicit model (argument or
+ ``CODEABC_MODEL``) wins, else an OpenRouter key (``sk-or-...``) routes
+ through OpenRouter, else the OpenAI default.
"""
override = model or os.getenv("CODEABC_MODEL")
+ if protocol:
+ return resolve_litellm_model(protocol, override or _DEFAULT_MODEL)
is_openrouter_key = bool(api_key) and api_key.startswith("sk-or-")
if override:
# a bare model name alongside an OpenRouter key -> route it accordingly
@@ -49,6 +62,14 @@ def _resolve_model(api_key: str | None = None, model: str | None = None) -> str:
return _DEFAULT_OPENROUTER_MODEL if is_openrouter_key else _DEFAULT_MODEL
+def _resolve_base(api_base: str | None, protocol: str | None) -> str | None:
+ """Pick the gateway base URL, normalised to the protocol's convention."""
+ base = api_base or os.getenv("OPENAI_API_BASE")
+ if protocol and base:
+ return normalize_base_url(base, protocol)
+ return base
+
+
# On failure the call helpers below surface this sentinel instead of raising, so
# the streaming generator can finish cleanly. Routers detect it (is_error_text)
# and turn it into a proper error the UI can show in plain language, rather than
@@ -66,19 +87,24 @@ async def stream_llm(
*,
api_key: str | None = None,
model: str | None = None,
+ api_base: str | None = None,
+ protocol: str | None = None,
) -> AsyncGenerator[str, None]:
"""Stream LLM response chunks."""
- model = _resolve_model(api_key, model)
+ model = _resolve_model(api_key, model, protocol)
kwargs: dict = {
"model": model,
"messages": [{"role": "user", "content": prompt}],
"stream": True,
"temperature": 0.3,
- "max_tokens": 4096,
+ "max_tokens": suggest_max_tokens(model),
}
if api_key:
kwargs["api_key"] = api_key
+ base = _resolve_base(api_base, protocol)
+ if base:
+ kwargs["api_base"] = base
try:
response = await litellm.acompletion(**kwargs)
@@ -87,8 +113,11 @@ async def stream_llm(
if delta.content:
yield delta.content
except Exception as e:
- logger.error(f"LLM call failed: {e}")
- yield f"{LLM_ERROR_PREFIX} {e}]"
+ # the key can appear inside a provider's error text, so scrub it from
+ # both the log line and the string the reader sees
+ detail = redact_secret(str(e), api_key or "")
+ logger.error(f"LLM call failed: {detail}")
+ yield f"{LLM_ERROR_PREFIX} {detail}]"
async def call_llm(
@@ -96,22 +125,28 @@ async def call_llm(
*,
api_key: str | None = None,
model: str | None = None,
+ api_base: str | None = None,
+ protocol: str | None = None,
) -> str:
"""Non-streaming LLM call. Returns the full response text."""
- model = _resolve_model(api_key, model)
+ model = _resolve_model(api_key, model, protocol)
kwargs: dict = {
"model": model,
"messages": [{"role": "user", "content": prompt}],
"temperature": 0.3,
- "max_tokens": 4096,
+ "max_tokens": suggest_max_tokens(model),
}
if api_key:
kwargs["api_key"] = api_key
+ base = _resolve_base(api_base, protocol)
+ if base:
+ kwargs["api_base"] = base
try:
response = await litellm.acompletion(**kwargs)
return response.choices[0].message.content or ""
except Exception as e:
- logger.error(f"LLM call failed: {e}")
- return f"{LLM_ERROR_PREFIX} {e}]"
+ detail = redact_secret(str(e), api_key or "")
+ logger.error(f"LLM call failed: {detail}")
+ return f"{LLM_ERROR_PREFIX} {detail}]"
diff --git a/backend/services/providers.py b/backend/services/providers.py
new file mode 100644
index 0000000..4d863e1
--- /dev/null
+++ b/backend/services/providers.py
@@ -0,0 +1,219 @@
+"""Wire-protocol registry and model-capability helpers for BYOK gateways.
+
+A gateway is described by the *protocol* it speaks, not by a vendor brand:
+the same base URL + API key can serve many models, and the protocol decides
+how requests are shaped and where the model list lives. Keeping the registry
+declarative means adding a protocol never touches a call site.
+
+litellm is imported lazily inside the capability helpers (not at module top)
+so this registry stays import-light and its catalog lookups are trivial to
+fake in tests. ``backend.services.llm`` still owns the eager litellm import
+and the ``drop_params`` side effect; this module deliberately does not.
+"""
+
+from __future__ import annotations
+
+from dataclasses import dataclass
+
+# Conservative output budget when litellm's catalog knows nothing about a
+# model: big enough for real answers, small enough to stay cheap.
+_FALLBACK_MAX_OUTPUT = 8192
+# Never ask for less than the historical hard-coded budget, and never more
+# than this even for models with huge output windows (keeps cost bounded).
+_MIN_MAX_TOKENS = 4096
+_CAP_MAX_TOKENS = 32768
+
+
+@dataclass(frozen=True)
+class Protocol:
+ """A wire protocol a gateway can speak.
+
+ ``litellm_prefix`` is prepended to a bare model id to form the litellm
+ model string. ``models_path`` is appended to the normalized base URL to
+ list models. ``base_has_v1`` records whether the base URL is expected to
+ end in ``/v1`` (OpenAI-style) or to be a bare root (Anthropic-style, where
+ the client appends ``/v1/messages`` itself). ``wire`` is a short, plain
+ description of the request shape shown under the format picker. Both are
+ English with ``*_zh`` twins, following the rest of the backend: the UI
+ picks the pair through ``t(zh, en)``.
+ """
+
+ id: str
+ label: str
+ wire: str
+ label_zh: str
+ wire_zh: str
+ litellm_prefix: str
+ models_path: str
+ base_hint: str
+ auth_header: str
+ base_has_v1: bool
+
+
+PROTOCOLS: tuple[Protocol, ...] = (
+ Protocol(
+ id="chat_completions",
+ label="OpenAI-compatible",
+ wire="POST {base}/chat/completions · model list GET {base}/models",
+ label_zh="OpenAI 兼容",
+ wire_zh="POST {base}/chat/completions · 模型列表 GET {base}/models",
+ litellm_prefix="openai/",
+ models_path="/models",
+ base_hint="https://api.example.com/v1",
+ auth_header="authorization",
+ base_has_v1=True,
+ ),
+ Protocol(
+ id="anthropic_messages",
+ label="Anthropic-compatible",
+ wire="POST {base}/v1/messages · most gateways expose no model list",
+ label_zh="Anthropic 兼容",
+ wire_zh="POST {base}/v1/messages · 多数网关无模型列表接口",
+ litellm_prefix="anthropic/",
+ models_path="/v1/models",
+ base_hint="https://api.anthropic.com",
+ auth_header="x-api-key",
+ base_has_v1=False,
+ ),
+)
+
+
+def list_protocols() -> list[Protocol]:
+ """Return every supported protocol, in display order."""
+ return list(PROTOCOLS)
+
+
+def get_protocol(protocol_id: str) -> Protocol | None:
+ """Look a protocol up by id, or None when unknown."""
+ for protocol in PROTOCOLS:
+ if protocol.id == protocol_id:
+ return protocol
+ return None
+
+
+def normalize_base_url(base: str, protocol_id: str) -> str:
+ """Normalise a user-typed base URL to the protocol's convention.
+
+ Kills the classic ``/v1/v1`` (or missing ``/v1``) 404: OpenAI-style
+ protocols expect the base to end in ``/v1``, while Anthropic-style expect
+ a bare root because the client appends ``/v1/messages`` itself.
+ """
+ cleaned = (base or "").strip().rstrip("/")
+ protocol = get_protocol(protocol_id)
+ if protocol is None:
+ return cleaned
+ if protocol.base_has_v1:
+ return cleaned if cleaned.endswith("/v1") else f"{cleaned}/v1"
+ if cleaned.endswith("/v1"):
+ return cleaned[: -len("/v1")].rstrip("/")
+ return cleaned
+
+
+def resolve_litellm_model(protocol_id: str, model_id: str) -> str:
+ """Compose the litellm model string for a protocol + bare model id.
+
+ Raises ValueError for an unknown protocol or an empty model id, since a
+ gateway model must always be chosen explicitly.
+ """
+ protocol = get_protocol(protocol_id)
+ if protocol is None:
+ raise ValueError(f"unknown protocol: {protocol_id}")
+ if not model_id:
+ raise ValueError(f"protocol {protocol_id} needs an explicit model id")
+ if model_id.startswith(protocol.litellm_prefix):
+ return model_id
+ return f"{protocol.litellm_prefix}{model_id}"
+
+
+def mask_key(key: str) -> str:
+ """Return a display-safe rendering of an API key (first 6 / last 4)."""
+ if not key:
+ return ""
+ if len(key) <= 10:
+ return "***"
+ return f"{key[:6]}***{key[-4:]}"
+
+
+def redact_secret(text: str, secret: str) -> str:
+ """Strip a secret out of free-form text (error messages, logs).
+
+ A credential must never survive into a reader-visible string, so replace
+ it wholesale with ``***`` rather than trying to mask it in place.
+ """
+ if not secret:
+ return text
+ return text.replace(secret, "***")
+
+
+def _load_litellm():
+ """Import litellm on first use only (see module docstring)."""
+ import litellm
+
+ return litellm
+
+
+def _catalog_info(model: str) -> tuple[dict | None, str]:
+ """Fetch litellm catalog metadata for ``model``.
+
+ Returns ``(info, source)`` where source is ``litellm_catalog`` for a
+ direct hit, ``cross_prefix`` when the bare model name was found under a
+ different provider prefix (OpenAI-compatible gateways reuse model names
+ litellm only knows elsewhere, e.g. ``openai/qwen3.8-flash`` vs
+ ``openrouter/qwen/qwen3.8-flash``), or ``(None, "unknown")``.
+
+ The fallback walks litellm's catalog (~4k keys) once per model, which costs
+ well under a millisecond each; a 260-model discovery list lands around
+ 0.1s, so no index is kept for it.
+ """
+ litellm = _load_litellm()
+ try:
+ return litellm.get_model_info(model), "litellm_catalog"
+ except Exception:
+ pass
+
+ bare = model.split("/", 1)[1] if "/" in model else model
+ for key in litellm.model_cost:
+ if key == bare or key.endswith("/" + bare):
+ try:
+ return litellm.get_model_info(key), "cross_prefix"
+ except Exception:
+ continue
+ return None, "unknown"
+
+
+def model_metadata(model: str) -> dict:
+ """Summarise capability metadata for ``model`` for API responses.
+
+ Unknown numbers are reported as None rather than guessed, so callers can
+ surface "unknown" honestly instead of fabricating precision.
+ """
+ info, source = _catalog_info(model)
+ if not info:
+ return {
+ "max_input_tokens": None,
+ "max_output_tokens": None,
+ "input_cost_per_token": None,
+ "output_cost_per_token": None,
+ "mode": None,
+ "metadata_source": "unknown",
+ }
+ return {
+ "max_input_tokens": info.get("max_input_tokens"),
+ "max_output_tokens": info.get("max_output_tokens"),
+ "input_cost_per_token": info.get("input_cost_per_token"),
+ "output_cost_per_token": info.get("output_cost_per_token"),
+ "mode": info.get("mode"),
+ "metadata_source": source,
+ }
+
+
+def suggest_max_tokens(model: str) -> int:
+ """Derive a sane ``max_tokens`` budget from the model's output window.
+
+ Replaces the historical hard-coded 4096, which silently truncated
+ reasoning models whose thinking alone exceeds that budget. Falls back to
+ a conservative default when the catalog knows nothing about the model.
+ """
+ info, _ = _catalog_info(model)
+ cap = (info or {}).get("max_output_tokens") or _FALLBACK_MAX_OUTPUT
+ return max(_MIN_MAX_TOKENS, min(int(cap), _CAP_MAX_TOKENS))
diff --git a/frontend/src/components/ApiKeyModal.tsx b/frontend/src/components/ApiKeyModal.tsx
index bc7cb41..930b69a 100644
--- a/frontend/src/components/ApiKeyModal.tsx
+++ b/frontend/src/components/ApiKeyModal.tsx
@@ -1,31 +1,199 @@
-import { useState } from "react";
+import { useEffect, useState } from "react";
import { useI18n } from "../lib/i18n";
+import {
+ listProtocols,
+ checkProtocol,
+ discoverModels,
+ bareModelId,
+ type Protocol,
+ type CheckResult,
+ type DiscoveredModel,
+} from "../lib/protocols";
interface Props {
open: boolean;
onClose: () => void;
}
+type DiscState =
+ | { kind: "idle" }
+ | { kind: "loading" }
+ | { kind: "ok"; count: number }
+ | { kind: "none"; msg: string }
+ | { kind: "error"; msg: string };
+
+function saveLocal(key: string, value: string): void {
+ try {
+ if (value) localStorage.setItem(key, value);
+ else localStorage.removeItem(key);
+ } catch {
+ // localStorage unavailable (private mode etc.) — keep in-memory only
+ }
+}
+
export default function ApiKeyModal({ open, onClose }: Props) {
if (!open) return null;
-
- return ;
+ return ;
}
-function ApiKeyDialog({ onClose }: { onClose: () => void }) {
+function ProviderForm({ onClose }: { onClose: () => void }) {
const { t } = useI18n();
- const [key, setKey] = useState(() => localStorage.getItem("codeabc_api_key") || "");
- const [saved, setSaved] = useState(false);
-
- const handleSave = () => {
- if (key.trim()) {
- localStorage.setItem("codeabc_api_key", key.trim());
- } else {
- localStorage.removeItem("codeabc_api_key");
+ const [protocols, setProtocols] = useState([]);
+ const [protocolId, setProtocolId] = useState(
+ () => localStorage.getItem("codeabc_protocol") || ""
+ );
+ const [baseUrl, setBaseUrl] = useState(
+ () => localStorage.getItem("codeabc_api_base") || ""
+ );
+ const [apiKey, setApiKey] = useState(
+ () => localStorage.getItem("codeabc_api_key") || ""
+ );
+ const [showKey, setShowKey] = useState(false);
+ const [model, setModel] = useState(
+ () => localStorage.getItem("codeabc_model") || ""
+ );
+ const [manualModel, setManualModel] = useState(() =>
+ bareModelId(localStorage.getItem("codeabc_model") || "")
+ );
+ const [manualMode, setManualMode] = useState(false);
+ const [disc, setDisc] = useState({ kind: "idle" });
+ const [models, setModels] = useState([]);
+ const [checking, setChecking] = useState(false);
+ const [check, setCheck] = useState(null);
+
+ useEffect(() => {
+ listProtocols()
+ .then((ps) => {
+ setProtocols(ps);
+ setProtocolId((cur) => cur || ps[0]?.id || "");
+ })
+ .catch(() => setProtocols([]));
+ }, []);
+
+ const selected = protocols.find((p) => p.id === protocolId) ?? null;
+
+ function qualify(p: Protocol | null, id: string): string {
+ if (!p || !id) return id;
+ const prefix = p.litellm_prefix || "";
+ if (!prefix || id.startsWith(prefix)) return id;
+ return `${prefix}${id}`;
+ }
+
+ /** Bare id sent to the backend so it can ping when no models list exists. */
+ const bareModel = manualMode ? manualModel.trim() : bareModelId(model);
+
+ async function runDiscover() {
+ if (!baseUrl.trim() || !protocolId) {
+ setDisc({
+ kind: "error",
+ msg: t("先填写 Base URL 并选择 API 格式。", "Fill in the Base URL and pick an API format first."),
+ });
+ return;
+ }
+ setDisc({ kind: "loading" });
+ setModels([]);
+ try {
+ const r = await discoverModels(protocolId, baseUrl.trim(), apiKey.trim() || undefined);
+ if (r.discoverable && (r.models?.length ?? 0) > 0) {
+ setModels(r.models ?? []);
+ setDisc({ kind: "ok", count: r.count });
+ setManualMode(false);
+ } else if (r.discoverable) {
+ setManualMode(true);
+ setDisc({
+ kind: "none",
+ msg: t("端点可达但未返回模型,请手动填写模型 ID。", "Endpoint reachable but returned no models; enter a model ID."),
+ });
+ } else if (r.status_code === 401 || r.status_code === 403) {
+ setDisc({
+ kind: "error",
+ msg: t(
+ "API Key 被该端点拒绝(401/403),请检查 Key 后重新获取。",
+ "The endpoint rejected the API key (401/403); check the key and fetch again."
+ ),
+ });
+ } else {
+ setManualMode(true);
+ setDisc({
+ kind: "none",
+ msg: t(
+ "该端点不提供模型列表接口(Anthropic 兼容端点常见),请手动填写模型 ID。",
+ "This endpoint exposes no models list (common for Anthropic-compatible gateways); enter a model ID."
+ ),
+ });
+ }
+ } catch {
+ setDisc({ kind: "error", msg: t("连不上 CodeABC 后端。", "Could not reach the CodeABC server.") });
}
- setSaved(true);
- setTimeout(() => onClose(), 800);
- };
+ setCheck(null);
+ }
+
+ async function runCheck() {
+ setChecking(true);
+ setCheck(null);
+ try {
+ const res = await checkProtocol({
+ protocol: protocolId || undefined,
+ api_base: baseUrl.trim() || undefined,
+ api_key: apiKey.trim() || undefined,
+ model: bareModel || undefined,
+ });
+ setCheck(res);
+ } catch {
+ setCheck({
+ ok: false,
+ base_url: baseUrl,
+ status_code: null,
+ masked_key: "",
+ diagnosis: "Could not reach the CodeABC server.",
+ diagnosis_zh: "连不上 CodeABC 后端。",
+ });
+ } finally {
+ setChecking(false);
+ }
+ }
+
+ function onProtocolChange(id: string) {
+ setProtocolId(id);
+ setCheck(null);
+ setDisc({ kind: "idle" });
+ setModels([]);
+ saveLocal("codeabc_protocol", id);
+ }
+
+ function onBaseChange(v: string) {
+ setBaseUrl(v);
+ setCheck(null);
+ setDisc({ kind: "idle" });
+ setModels([]);
+ saveLocal("codeabc_api_base", v.trim());
+ }
+
+ function onKeyChange(v: string) {
+ setApiKey(v);
+ setCheck(null);
+ saveLocal("codeabc_api_key", v.trim());
+ }
+
+ function pickModel(m: DiscoveredModel) {
+ setModel(m.litellm_model);
+ setManualModel(m.id);
+ setCheck(null);
+ saveLocal("codeabc_model", m.litellm_model);
+ }
+
+ function onManual(v: string) {
+ setManualModel(v);
+ setModel(qualify(selected, v.trim()));
+ setCheck(null);
+ saveLocal("codeabc_model", qualify(selected, v.trim()));
+ }
+
+ const inputCls =
+ "w-full px-4 py-3 border border-gray-300 rounded-xl text-sm " +
+ "focus:outline-none focus:ring-2 focus:ring-blue-500 focus:border-transparent";
+ const labelCls = "block text-sm font-medium text-gray-700 mb-1.5";
+ const stepCls = "text-xs font-semibold text-gray-400 uppercase tracking-wide mb-2";
return (
void }) {
onClick={onClose}
>
e.stopPropagation()}
>
-
- {t("API Key 设置", "API Key settings")}
-
-
- {t(
- "填入你自己的 API Key 可以无限使用。留空则使用免费额度(每天 20 次)。",
- "Add your own API key for unlimited use. Leave it blank to use the free tier (20/day).",
- )}
-
+ {/* Header */}
+
+
+ {t("模型服务商设置", "Model provider settings")}
+
+
+
- setKey(e.target.value)}
- placeholder={t(
- "sk-or-... (OpenRouter)或其他提供商的 Key",
- "sk-or-... (OpenRouter) or a key from another provider",
+ {/* Step 1 — endpoint */}
+
+ {t("测试连接成功后即可完成。", "Test the connection successfully to finish.")}
+
+ )}
+
{t(
- "粘贴 OpenRouter 的 Key,会自动帮你选一个又快又便宜的模型,不用懂模型。Key 只存在你浏览器本地,绝不上传。OpenAI / Claude / DeepSeek / Kimi 等 litellm 兼容的 Key 也都支持。",
- "Paste an OpenRouter key and a fast, inexpensive model is picked for you — no model knowledge needed. The key stays in your browser and is never uploaded. OpenAI / Claude / DeepSeek / Kimi and any litellm-compatible key work too.",
+ "不填 Key 也能用免费额度(每天 20 次);直接关闭本窗口即可。",
+ "Without a key the free tier (20 requests/day) still works — just close this dialog."
)}
-
-
-
-
-
);
diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts
index bafc88a..9b3fc21 100644
--- a/frontend/src/lib/api.ts
+++ b/frontend/src/lib/api.ts
@@ -16,13 +16,13 @@ function isDesktopShell(): boolean {
);
}
-const BASE =
+export const BASE =
import.meta.env.VITE_API_BASE ??
(isDesktopShell() ? "http://127.0.0.1:8000/api" : "/api");
/** fetch wrapper that turns a network-level failure in the desktop shell into a
* clear "the backend isn't running" message instead of a raw "Failed to fetch". */
-async function apiFetch(input: string, init?: RequestInit): Promise {
+export async function apiFetch(input: string, init?: RequestInit): Promise {
try {
return await window.fetch(input, init);
} catch (e) {
@@ -38,18 +38,26 @@ async function apiFetch(input: string, init?: RequestInit): Promise {
}
}
-function getHeaders(): HeadersInit {
- const headers: Record = {
- "Content-Type": "application/json",
- };
- // BYOK: attach user's API key if configured
+/** BYOK headers: the user's API key plus any provider base URL / model picked
+ * in the settings wizard. All optional — an absent header keeps the backend's
+ * existing default (key-shape model routing, OPENAI_API_BASE env fallback). */
+function byokHeaders(): Record {
+ const headers: Record = {};
const apiKey = localStorage.getItem("codeabc_api_key");
- if (apiKey) {
- headers["x-api-key"] = apiKey;
- }
+ if (apiKey) headers["x-api-key"] = apiKey;
+ const apiBase = localStorage.getItem("codeabc_api_base");
+ if (apiBase) headers["x-api-base"] = apiBase;
+ const model = localStorage.getItem("codeabc_model");
+ if (model) headers["x-model"] = model;
+ const protocol = localStorage.getItem("codeabc_protocol");
+ if (protocol) headers["x-api-protocol"] = protocol;
return headers;
}
+function getHeaders(): HeadersInit {
+ return { "Content-Type": "application/json", ...byokHeaders() };
+}
+
export interface FileInfo {
path: string;
size: number;
@@ -584,9 +592,7 @@ export async function streamOverview(
onResult: (overview: ProjectOverview) => void,
onError: (err: string) => void
) {
- const apiKey = localStorage.getItem("codeabc_api_key");
- const headers: Record = {};
- if (apiKey) headers["x-api-key"] = apiKey;
+ const headers: Record = byokHeaders();
const res = await apiFetch(`${BASE}/project/${projectId}/overview`, { headers });
if (!res.ok || !res.body) {
diff --git a/frontend/src/lib/protocols.ts b/frontend/src/lib/protocols.ts
new file mode 100644
index 0000000..a319e25
--- /dev/null
+++ b/frontend/src/lib/protocols.ts
@@ -0,0 +1,90 @@
+import { apiFetch, BASE } from "./api";
+
+export interface Protocol {
+ id: string;
+ label: string;
+ /** plain description of the request shape, shown under the format picker */
+ wire: string;
+ /** Chinese twins of label/wire — the backend registry is the source of truth */
+ label_zh: string;
+ wire_zh: string;
+ litellm_prefix: string;
+ models_path: string;
+ base_hint: string;
+}
+
+export interface CheckResult {
+ ok: boolean;
+ base_url: string;
+ status_code: number | null;
+ masked_key: string;
+ diagnosis: string | null;
+ diagnosis_zh: string | null;
+}
+
+export interface DiscoveredModel {
+ id: string;
+ litellm_model: string;
+ chat: boolean;
+ chat_reason: string;
+ chat_reason_zh: string;
+ mode: string | null;
+ max_input_tokens: number | null;
+ max_output_tokens: number | null;
+ input_cost_per_token: number | null;
+ output_cost_per_token: number | null;
+ metadata_source: string;
+}
+
+export interface ModelsResponse {
+ models: DiscoveredModel[];
+ count: number;
+ /** false when the gateway exposes no models list (e.g. 404) */
+ discoverable?: boolean;
+ /** HTTP status of the models probe when it failed */
+ status_code?: number | null;
+ error?: string;
+}
+
+/** Fetch the supported wire protocols (no secrets). */
+export async function listProtocols(): Promise {
+ const res = await apiFetch(`${BASE}/protocols`);
+ const data = await res.json();
+ return (data.protocols ?? []) as Protocol[];
+}
+
+/** Strip a litellm provider prefix: the bare id a gateway is pinged with. */
+export function bareModelId(id: string): string {
+ const slash = id.indexOf("/");
+ return slash === -1 ? id : id.slice(slash + 1);
+}
+
+/** Probe a gateway's models endpoint before its credentials are saved. */
+export async function checkProtocol(req: {
+ protocol?: string;
+ api_base?: string;
+ api_key?: string;
+ model?: string;
+}): Promise {
+ const res = await apiFetch(`${BASE}/protocols/check`, {
+ method: "POST",
+ headers: { "Content-Type": "application/json" },
+ body: JSON.stringify(req),
+ });
+ return res.json();
+}
+
+/** Discover the models an endpoint serves, merged with capability metadata. */
+export async function discoverModels(
+ protocol: string,
+ apiBase?: string,
+ apiKey?: string,
+): Promise {
+ const params = new URLSearchParams();
+ if (protocol) params.set("protocol", protocol);
+ if (apiBase) params.set("api_base", apiBase);
+ const headers: Record = {};
+ if (apiKey) headers["x-api-key"] = apiKey;
+ const res = await apiFetch(`${BASE}/models?${params.toString()}`, { headers });
+ return res.json();
+}
diff --git a/tests/test_llm.py b/tests/test_llm.py
index a90c5c7..f4b3a90 100644
--- a/tests/test_llm.py
+++ b/tests/test_llm.py
@@ -2,6 +2,8 @@
from __future__ import annotations
+import asyncio
+
import pytest
from backend.services import llm
@@ -66,3 +68,123 @@ def test_is_error_text():
assert not llm.is_error_text("")
# the call helpers emit exactly this sentinel shape
assert llm.is_error_text(f"{llm.LLM_ERROR_PREFIX} something]")
+
+
+class _FakeMessage:
+ def __init__(self, content):
+ self.content = content
+
+
+class _FakeChoice:
+ def __init__(self, content):
+ self.message = _FakeMessage(content)
+
+
+class _FakeResponse:
+ def __init__(self, content="ok"):
+ self.choices = [_FakeChoice(content)]
+
+
+@pytest.fixture
+def capture_kwargs(monkeypatch):
+ """Replace litellm.acompletion with a spy that records its kwargs."""
+ captured: dict = {}
+
+ async def fake_acompletion(**kwargs):
+ captured.update(kwargs)
+ return _FakeResponse()
+
+ monkeypatch.setattr(llm.litellm, "acompletion", fake_acompletion)
+ return captured
+
+
+def test_call_llm_passes_api_base_and_model(capture_kwargs):
+ # the wizard's chosen gateway + model must reach litellm verbatim
+ result = asyncio.run(
+ llm.call_llm(
+ "hi",
+ api_key="sk-sp-x",
+ api_base="https://gw/v1",
+ model="openai/qwen3.8-flash",
+ )
+ )
+ assert result == "ok"
+ assert capture_kwargs["api_base"] == "https://gw/v1"
+ assert capture_kwargs["model"] == "openai/qwen3.8-flash"
+
+
+def test_call_llm_uses_adaptive_max_tokens(capture_kwargs, monkeypatch):
+ # regression: max_tokens was hard-coded to 4096, truncating long answers;
+ # it now derives from the model's real output window via suggest_max_tokens.
+ monkeypatch.setattr(llm, "suggest_max_tokens", lambda model: 12345)
+ asyncio.run(llm.call_llm("hi", model="openai/x"))
+ assert capture_kwargs["max_tokens"] == 12345
+
+
+def test_call_llm_env_api_base_fallback(capture_kwargs, monkeypatch):
+ monkeypatch.setenv("OPENAI_API_BASE", "https://env-gw/v1")
+ asyncio.run(llm.call_llm("hi", model="openai/x"))
+ assert capture_kwargs["api_base"] == "https://env-gw/v1"
+
+
+def test_call_llm_no_api_base_when_unset(capture_kwargs, monkeypatch):
+ monkeypatch.delenv("OPENAI_API_BASE", raising=False)
+ asyncio.run(llm.call_llm("hi", model="openai/x"))
+ assert "api_base" not in capture_kwargs
+
+
+def test_call_llm_redacts_key_from_error(monkeypatch):
+ key = "sk-sp-SECRET123"
+
+ async def boom(**kwargs):
+ raise RuntimeError(f"auth failed for {key}")
+
+ monkeypatch.setattr(llm.litellm, "acompletion", boom)
+ result = asyncio.run(llm.call_llm("hi", api_key=key))
+ assert llm.is_error_text(result)
+ assert key not in result # the credential never reaches a reader-visible string
+ assert "***" in result
+
+
+def test_llm_kwargs_reads_byok_headers():
+ from backend.routers.analyze import _llm_kwargs
+
+ class _Req:
+ def __init__(self, headers):
+ self.headers = headers
+
+ headers = {
+ "x-api-key": "k",
+ "x-api-base": "https://gw/v1",
+ "x-model": "openai/m",
+ "x-api-protocol": "chat_completions",
+ }
+ assert _llm_kwargs(_Req(headers)) == {
+ "api_key": "k",
+ "api_base": "https://gw/v1",
+ "model": "openai/m",
+ "protocol": "chat_completions",
+ }
+ # absent headers -> all None, so llm falls back to its env/default path
+ assert _llm_kwargs(_Req({})) == {
+ "api_key": None,
+ "api_base": None,
+ "model": None,
+ "protocol": None,
+ }
+
+
+def test_resolve_model_with_protocol():
+ # the settings form's BYOK path qualifies a bare gateway model by protocol
+ assert llm._resolve_model(model="qwen3.8-flash", protocol="chat_completions") == (
+ "openai/qwen3.8-flash"
+ )
+ assert llm._resolve_model(model="claude-x", protocol="anthropic_messages") == (
+ "anthropic/claude-x"
+ )
+
+
+def test_resolve_base_normalises_per_protocol():
+ # OpenAI-style bases end in /v1; Anthropic-style are bare roots
+ assert llm._resolve_base("https://gw", "chat_completions") == "https://gw/v1"
+ assert llm._resolve_base("https://gw/v1", "anthropic_messages") == "https://gw"
diff --git a/tests/test_providers.py b/tests/test_providers.py
new file mode 100644
index 0000000..ebe4bce
--- /dev/null
+++ b/tests/test_providers.py
@@ -0,0 +1,154 @@
+"""Tests for the wire-protocol registry and model-capability helpers."""
+
+from __future__ import annotations
+
+import pytest
+
+from backend.services import providers
+
+
+class _FakeLitellm:
+ """Stand-in for litellm: a model->info catalog, raising for unknown ids."""
+
+ def __init__(self, catalog):
+ self.model_cost = catalog
+
+ def get_model_info(self, model):
+ if model in self.model_cost:
+ return self.model_cost[model]
+ raise KeyError(model)
+
+
+@pytest.fixture
+def fake_catalog(monkeypatch):
+ def _install(catalog):
+ monkeypatch.setattr(providers, "_load_litellm", lambda: _FakeLitellm(catalog))
+
+ return _install
+
+
+def test_protocols_are_well_formed():
+ protos = providers.list_protocols()
+ assert [p.id for p in protos] == ["chat_completions", "anthropic_messages"]
+ for p in protos:
+ assert p.label and p.wire and p.litellm_prefix and p.models_path and p.base_hint
+ # the settings form is bilingual, so every label carries both languages
+ assert p.label_zh and p.wire_zh
+ assert p.label_zh != p.label
+ chat = providers.get_protocol("chat_completions")
+ assert chat.base_has_v1 is True
+ assert chat.auth_header == "authorization"
+ assert chat.models_path == "/models"
+ anth = providers.get_protocol("anthropic_messages")
+ assert anth.base_has_v1 is False
+ assert anth.auth_header == "x-api-key"
+ assert anth.models_path == "/v1/models"
+
+
+def test_get_protocol_unknown():
+ assert providers.get_protocol("nope") is None
+
+
+@pytest.mark.parametrize(
+ ("raw", "expected"),
+ [
+ ("https://gw", "https://gw/v1"),
+ ("https://gw/", "https://gw/v1"),
+ ("https://gw/v1", "https://gw/v1"),
+ ("https://gw/v1/", "https://gw/v1"),
+ ("https://gw/compatible-mode/v1", "https://gw/compatible-mode/v1"),
+ ],
+)
+def test_normalize_base_chat_completions(raw, expected):
+ assert providers.normalize_base_url(raw, "chat_completions") == expected
+
+
+@pytest.mark.parametrize(
+ ("raw", "expected"),
+ [
+ ("https://api.anthropic.com", "https://api.anthropic.com"),
+ ("https://api.anthropic.com/v1", "https://api.anthropic.com"),
+ ("https://api.anthropic.com/v1/", "https://api.anthropic.com"),
+ ],
+)
+def test_normalize_base_anthropic(raw, expected):
+ assert providers.normalize_base_url(raw, "anthropic_messages") == expected
+
+
+def test_normalize_base_unknown_protocol_passthrough():
+ assert providers.normalize_base_url("https://gw/v1/", "nope") == "https://gw/v1"
+
+
+def test_resolve_litellm_model():
+ resolve = providers.resolve_litellm_model
+ assert resolve("chat_completions", "qwen3.8-flash") == "openai/qwen3.8-flash"
+ assert resolve("chat_completions", "openai/qwen3.8-flash") == "openai/qwen3.8-flash"
+ assert resolve("anthropic_messages", "claude-x") == "anthropic/claude-x"
+
+
+def test_resolve_litellm_model_errors():
+ with pytest.raises(ValueError):
+ providers.resolve_litellm_model("nope", "m")
+ with pytest.raises(ValueError):
+ providers.resolve_litellm_model("chat_completions", "")
+
+
+def test_mask_key():
+ assert providers.mask_key("sk-sp-SECRET123") == "sk-sp-***T123"
+ assert providers.mask_key("short") == "***"
+ assert providers.mask_key("") == ""
+
+
+def test_redact_secret():
+ assert providers.redact_secret("auth failed sk-ABC", "sk-ABC") == "auth failed ***"
+ assert providers.redact_secret("no secret", "") == "no secret"
+
+
+def test_suggest_max_tokens_direct_hit(fake_catalog):
+ fake_catalog({"openai/qwen3.8-flash": {"max_output_tokens": 131072}})
+ # huge output window is capped so cost stays bounded
+ assert providers.suggest_max_tokens("openai/qwen3.8-flash") == 32768
+
+
+def test_suggest_max_tokens_cross_prefix(fake_catalog):
+ # gateway reuses a model name litellm only knows under another prefix
+ fake_catalog({"openrouter/qwen/qwen3.8-flash": {"max_output_tokens": 32768}})
+ assert providers.suggest_max_tokens("openai/qwen3.8-flash") == 32768
+
+
+def test_suggest_max_tokens_unknown(fake_catalog):
+ fake_catalog({})
+ assert providers.suggest_max_tokens("openai/mystery") == 8192
+
+
+def test_suggest_max_tokens_floor(fake_catalog):
+ fake_catalog({"openai/tiny": {"max_output_tokens": 10}})
+ assert providers.suggest_max_tokens("openai/tiny") == 4096
+
+
+def test_model_metadata_cross_prefix(fake_catalog):
+ fake_catalog({"openrouter/qwen/qwen3.8-flash": {"max_input_tokens": 1_000_000}})
+ meta = providers.model_metadata("openai/qwen3.8-flash")
+ assert meta["max_input_tokens"] == 1_000_000
+ assert meta["metadata_source"] == "cross_prefix"
+
+
+def test_model_metadata_unknown(fake_catalog):
+ fake_catalog({})
+ meta = providers.model_metadata("openai/mystery")
+ assert meta["max_input_tokens"] is None
+ assert meta["metadata_source"] == "unknown"
+
+
+def test_model_metadata_matches_a_nested_bare_name(fake_catalog):
+ # the suffix index must still find a name litellm knows two segments deep
+ fake_catalog({"openrouter/qwen/qwen3.8-flash": {"max_output_tokens": 32768}})
+ meta = providers.model_metadata("openai/qwen/qwen3.8-flash")
+ assert meta["metadata_source"] == "cross_prefix"
+ assert meta["max_output_tokens"] == 32768
+
+
+def test_model_metadata_never_matches_a_partial_name(fake_catalog):
+ # "flash" is a segment of the catalog key but not a whole bare name
+ fake_catalog({"openrouter/qwen/qwen3.8-flash": {"max_output_tokens": 32768}})
+ assert providers.model_metadata("openai/flash")["metadata_source"] == "unknown"
diff --git a/tests/test_providers_router.py b/tests/test_providers_router.py
new file mode 100644
index 0000000..20285b1
--- /dev/null
+++ b/tests/test_providers_router.py
@@ -0,0 +1,225 @@
+"""Tests for the protocol check + live model discovery endpoints."""
+
+from __future__ import annotations
+
+import httpx
+import pytest
+from fastapi.testclient import TestClient
+
+from backend.app import app
+from backend.routers import providers as router_module
+
+_KEY = "sk-sp-SECRET123"
+
+
+def _client() -> TestClient:
+ # Deliberately not used as a context manager: that would fire the app's
+ # lifespan (init_db) and write a real cache DB; these routes never touch it.
+ return TestClient(app)
+
+
+class _FakeResp:
+ def __init__(self, status_code=200, payload=None):
+ self.status_code = status_code
+ self._payload = payload
+
+ def json(self):
+ return self._payload
+
+ def raise_for_status(self):
+ if self.status_code >= 400:
+ raise httpx.HTTPStatusError("boom", request=None, response=self)
+
+
+class _FakeAsyncClient:
+ last_url = ""
+ last_headers: dict = {}
+ post_resp = None
+
+ def __init__(self, resp):
+ self._resp = resp
+
+ async def __aenter__(self):
+ return self
+
+ async def __aexit__(self, *exc):
+ return False
+
+ async def get(self, url, headers=None):
+ _FakeAsyncClient.last_url = url
+ _FakeAsyncClient.last_headers = headers or {}
+ return self._resp
+
+ async def post(self, url, json=None, headers=None):
+ _FakeAsyncClient.last_url = url
+ _FakeAsyncClient.last_headers = headers or {}
+ return _FakeAsyncClient.post_resp or _FakeResp(200, {"id": "ping"})
+
+
+@pytest.fixture
+def patch_httpx(monkeypatch):
+ def _install(resp):
+ monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: _FakeAsyncClient(resp))
+ _FakeAsyncClient.last_url = ""
+ _FakeAsyncClient.last_headers = {}
+ _FakeAsyncClient.post_resp = None
+ return _install
+
+
+def test_protocols_endpoint_lists_protocols():
+ r = _client().get("/api/protocols")
+ assert r.status_code == 200
+ protos = r.json()["protocols"]
+ assert [p["id"] for p in protos] == ["chat_completions", "anthropic_messages"]
+ # registry data only — no secrets leak into the listing
+ assert all("api_key" not in p for p in protos)
+ assert protos[0]["base_hint"]
+ # both languages ship in the listing so the form can render either
+ assert all(p["label_zh"] and p["wire_zh"] for p in protos)
+
+
+def test_check_requires_base_url():
+ r = _client().post(
+ "/api/protocols/check", json={"protocol": "chat_completions", "api_key": _KEY}
+ )
+ assert r.status_code == 200
+ body = r.json()
+ assert body["ok"] is False
+ assert "Base URL" in (body["diagnosis"] or "")
+
+
+def test_check_unknown_protocol():
+ r = _client().post(
+ "/api/protocols/check", json={"protocol": "nope", "api_base": "https://gw"}
+ )
+ assert r.json()["ok"] is False
+
+
+def test_check_pings_selected_model_chat(patch_httpx):
+ patch_httpx(_FakeResp(200, {"data": []}))
+ r = _client().post(
+ "/api/protocols/check",
+ json={
+ "protocol": "chat_completions",
+ "api_base": "https://gw.example.com",
+ "api_key": _KEY,
+ "model": "qwen3.8-flash",
+ },
+ )
+ body = r.json()
+ assert body["ok"] is True
+ # base normalised to the OpenAI-style /v1 convention
+ assert body["base_url"] == "https://gw.example.com/v1"
+ # only the selected model is tested, via a 1-token chat completion
+ assert _FakeAsyncClient.last_url == "https://gw.example.com/v1/chat/completions"
+ # bearer auth for the OpenAI-style protocol
+ assert _FakeAsyncClient.last_headers.get("Authorization") == f"Bearer {_KEY}"
+ # the raw key never appears in the response body
+ assert body["masked_key"] == "sk-sp-***T123"
+ assert _KEY not in r.text
+
+
+def test_check_anthropic_pings_messages(patch_httpx):
+ patch_httpx(_FakeResp(200, {"data": []}))
+ r = _client().post(
+ "/api/protocols/check",
+ json={
+ "protocol": "anthropic_messages",
+ "api_base": "https://api.anthropic.com/v1",
+ "api_key": "sk-ant-SECRET1",
+ "model": "claude-x",
+ },
+ )
+ body = r.json()
+ assert body["ok"] is True
+ # Anthropic-style base is a bare root (client appends /v1/messages)
+ assert body["base_url"] == "https://api.anthropic.com"
+ assert _FakeAsyncClient.last_url == "https://api.anthropic.com/v1/messages"
+ assert _FakeAsyncClient.last_headers.get("x-api-key") == "sk-ant-SECRET1"
+
+
+def test_check_unauthorized_diagnosis(patch_httpx):
+ patch_httpx(_FakeResp(200, {"data": []}))
+ _FakeAsyncClient.post_resp = _FakeResp(401, {})
+ r = _client().post(
+ "/api/protocols/check",
+ json={
+ "protocol": "chat_completions",
+ "api_base": "https://gw/v1",
+ "api_key": _KEY,
+ "model": "m",
+ },
+ )
+ body = r.json()
+ assert body["ok"] is False
+ assert body["status_code"] == 401
+ assert "key" in (body["diagnosis"] or "").lower()
+ # the same diagnosis is available in Chinese for the zh UI
+ assert body["diagnosis_zh"] and body["diagnosis_zh"] != body["diagnosis"]
+
+
+def test_check_requires_model():
+ # the check tests exactly one model, so it refuses to run without one
+ r = _client().post(
+ "/api/protocols/check",
+ json={"protocol": "chat_completions", "api_base": "https://gw/v1", "api_key": _KEY},
+ )
+ body = r.json()
+ assert body["ok"] is False
+ assert "model" in (body["diagnosis"] or "").lower()
+ assert body["diagnosis_zh"]
+
+
+def test_models_endpoint_discovers_with_prefix(patch_httpx, monkeypatch):
+ monkeypatch.setattr(
+ router_module,
+ "model_metadata",
+ lambda m: {"mode": None, "max_input_tokens": 100, "max_output_tokens": 50,
+ "input_cost_per_token": None, "output_cost_per_token": None,
+ "metadata_source": "test"},
+ )
+ patch_httpx(_FakeResp(200, {"data": [{"id": "qwen3.8-flash"}, {"id": "wan2.7-image"}]}))
+ r = _client().get(
+ "/api/models",
+ params={"protocol": "chat_completions", "api_base": "https://gw/v1"},
+ headers={"x-api-key": _KEY},
+ )
+ body = r.json()
+ assert body["count"] == 2
+ chat = body["models"][0]
+ assert chat["id"] == "qwen3.8-flash"
+ assert chat["litellm_model"] == "openai/qwen3.8-flash"
+ assert chat["chat"] is True
+ img = body["models"][1]
+ assert img["chat"] is False # name hints a non-chat model
+ # the picker shows the reason inline, so it needs both languages too
+ assert chat["chat_reason_zh"] and img["chat_reason_zh"]
+
+
+def test_models_endpoint_requires_protocol_and_base():
+ r = _client().get("/api/models")
+ body = r.json()
+ assert body["models"] == []
+ assert body["error"]
+
+
+def test_models_endpoint_redacts_key_from_transport_error(monkeypatch):
+ class _Boom:
+ async def __aenter__(self):
+ return self
+
+ async def __aexit__(self, *exc):
+ return False
+
+ async def get(self, url, headers=None):
+ raise httpx.ConnectError(f"proxy denied for {_KEY}")
+
+ monkeypatch.setattr(httpx, "AsyncClient", lambda **kw: _Boom())
+ r = _client().get(
+ "/api/models",
+ params={"protocol": "chat_completions", "api_base": "https://gw/v1"},
+ headers={"x-api-key": _KEY},
+ )
+ body = r.json()
+ assert body["models"] == []
+ assert _KEY not in body["error"]