From 3166c5ab2098cac24c42e8aae05c199e49a0812a Mon Sep 17 00:00:00 2001 From: Mr-Shaw-Yihan <1179647539@qq.com> Date: Thu, 24 Sep 2026 07:47:26 +0800 Subject: [PATCH 1/2] feat(llm): protocol-first BYOK onboarding, live model discovery, and adaptive token budget Describe gateways by the wire protocol they speak (OpenAI-compatible chat completions vs Anthropic-compatible messages) instead of a vendor brand, add three discovery/validation endpoints (list protocols, ping the selected model before saving, discover models merged with litellm capability metadata), and rewrite the API-key modal into a linear three-step form (endpoint -> model -> verify). Fixes concrete integration problems: - stream_llm/call_llm had no api_base channel, so an OpenAI-compatible gateway (e.g. a Qwen Token Plan key, sk-sp-) was routed to platform.openai.com; the form's base URL + protocol + model now flow through x-api-base / x-api-protocol / x-model headers into analyze._llm_kwargs and on to litellm, with the base URL normalised per protocol (/v1 suffix vs bare root). - Anthropic-compatible gateways often expose no /v1/models (e.g. Alibaba Bailian's /apps/anthropic), so connection checks ping the selected model with a 1-token request and the form falls back to manual model entry. - max_tokens was hard-coded to 4096 in both call helpers, truncating models with larger output windows; it now derives from the model's real output window via suggest_max_tokens (clamped, floored at 4096). - capability lookups missed when a gateway reuses a model name litellm only knows under another prefix; added a cross-prefix catalog fallback. Everything the form renders from a backend payload (protocol names, wire descriptions, grey-out reasons, connection diagnosis) ships as `label` plus a `_zh` twin, following the convention activity/contributing already use, so the language toggle reaches it. The rewritten modal also keeps the three promises the previous one made: the key stays in local storage, the OpenRouter signup link, and the free tier of 20 requests/day. llm.py keeps its eager `import litellm` + drop_params invariant (locked by test_llm.py); the registry imports litellm lazily. API keys are masked in every response and redacted from error strings and log lines alike. Fully backward compatible: without the new headers, behavior is unchanged. Tests: 809 passed (test_providers, test_providers_router rewritten for the protocol design; test_llm extended). ruff clean; new backend files are 0-flag under the project's own long_functions/too_many_params/docstrings/typing_coverage scanners. Frontend tsc + eslint + build green. --- backend/app.py | 5 +- backend/models.py | 10 + backend/routers/analyze.py | 30 +- backend/routers/providers.py | 291 ++++++++++++++ backend/services/llm.py | 65 +++- backend/services/providers.py | 219 +++++++++++ frontend/src/components/ApiKeyModal.tsx | 484 ++++++++++++++++++++---- frontend/src/lib/api.ts | 32 +- frontend/src/lib/protocols.ts | 90 +++++ tests/test_llm.py | 122 ++++++ tests/test_providers.py | 154 ++++++++ tests/test_providers_router.py | 225 +++++++++++ 12 files changed, 1625 insertions(+), 102 deletions(-) create mode 100644 backend/routers/providers.py create mode 100644 backend/services/providers.py create mode 100644 frontend/src/lib/protocols.ts create mode 100644 tests/test_providers.py create mode 100644 tests/test_providers_router.py 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("1 · 连接信息", "1 · Endpoint")}
+
+ + onBaseChange(e.target.value)} + placeholder={selected?.base_hint || "https://api.example.com/v1"} + /> +

+ {t( + "OpenAI 系需以 /v1 结尾,Anthropic 系填根地址;系统会自动归一。", + "OpenAI-style ends with /v1, Anthropic-style uses the root; auto-normalised." + )} +

+
+
+ + + {selected && ( +

+ {t(selected.wire_zh, selected.wire)} +

)} - className="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" - /> - - +
+ +
+ onKeyChange(e.target.value)} + placeholder={t("输入 API Key", "Enter API Key")} + autoComplete="off" + /> + +
+

+ {t( + "Key 只保存在你浏览器的本地存储,仅作为请求头发给你选择的端点;服务器不会保存它。", + "The key is stored only in your browser's local storage and sent as a request header to the endpoint you choose; the server does not store it." + )} +

+
+ {t( + "还没有 Key?去 OpenRouter 几分钟注册一个(sk-or- 开头)→", + "No key yet? Get one from OpenRouter in a couple of minutes (it starts with sk-or-) →" + )} + +
+ + {/* Step 2 — model */} +
{t("2 · 模型", "2 · Model")}
+ + + {disc.kind === "ok" && ( +
+ {t("可自动发现", "Auto-discovery available")} · {disc.count} {t("个模型", "models")} +
+ )} + {disc.kind === "none" && ( +
+ {disc.msg} +
+ )} + {disc.kind === "error" && ( +
+ {disc.msg} +
+ )} + + {disc.kind === "idle" && ( +

+ {t( + "先探测该端点能否自动列出模型;不能则需手动填写模型 ID。", + "Probe whether this endpoint can list its models; otherwise enter a model ID." + )} +

+ )} + + {disc.kind === "ok" && !manualMode && ( +
+ + +
+ )} + + {(manualMode || disc.kind === "none") && ( +
+ + onManual(e.target.value)} + placeholder={t("如 qwen3.8-flash", "e.g. qwen3.8-flash")} + /> + {disc.kind === "ok" && ( + + )} +
+ )} -

+ {/* Step 3 — verify */} +

{t("3 · 验证", "3 · Verify")}
+ + {!bareModel && ( +

+ {t("请先在第 2 步选择或填写模型。", "Choose or enter a model in step 2 first.")} +

+ )} + + {check && ( +
+ {check.ok ? ( + + {t("连接成功", "Connected")} + {check.masked_key && ` · ${check.masked_key}`} + + ) : ( + {t(check.diagnosis_zh || "", check.diagnosis || "")} + )} +
+ )} + + {/* Footer */} + + {!check?.ok && ( +

+ {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"] From 93a797eb72c9a2c34cbff644a5d48381d1968da8 Mon Sep 17 00:00:00 2001 From: Mr-Shaw-Yihan <1179647539@qq.com> Date: Thu, 24 Sep 2026 07:47:26 +0800 Subject: [PATCH 2/2] docs: document the three-step BYOK form and OPENAI_API_BASE README (EN + CN) now spell out the endpoint -> model -> verify flow: Base URL normalisation per API format, model discovery with a manual fallback for gateways that expose no list, and a connection test that pings only the selected model. Both still note that an OpenRouter key needs nothing but the key. .env.example documents OPENAI_API_BASE, the server-side counterpart of the form's Base URL field. --- .env.example | 6 ++++++ README.md | 7 ++++++- README_CN.md | 7 ++++++- 3 files changed, 18 insertions(+), 2 deletions(-) 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`)。 ## 路线图