From 8bd5ea777c3c46fdf1f745057800388d81e77b85 Mon Sep 17 00:00:00 2001 From: Mumit Khan Date: Sun, 26 Jul 2026 20:56:19 -0700 Subject: [PATCH] Add custom OpenAI-compatible providers --- README.md | 2 + coworker/providers/__init__.py | 4 + coworker/providers/registry.py | 135 ++++++++++++++++-- coworker/providers/router.py | 10 +- coworker/server/app.py | 7 + coworker/server/manager.py | 127 +++++++++++++++- surfaces/gui/src/api.ts | 17 ++- .../gui/src/components/ManageTabs.test.tsx | 92 ++++++++++++ surfaces/gui/src/components/ManageTabs.tsx | 92 +++++++++++- surfaces/gui/src/providers/ProviderSetup.tsx | 41 ++++-- tests/test_provider_router.py | 105 ++++++++++++++ tests/test_server.py | 26 ++++ 12 files changed, 622 insertions(+), 36 deletions(-) create mode 100644 surfaces/gui/src/components/ManageTabs.test.tsx diff --git a/README.md b/README.md index d96cf139..fdc2c2f6 100644 --- a/README.md +++ b/README.md @@ -56,6 +56,8 @@ Model access is yours: pick a provider, paste your key, switch anytime. Supporte A curated model list marks what we've verified for tool-calling work. Adding any model string works at your own risk. +Need another Chat Completions gateway? In **Settings → Models**, choose **Custom provider** and enter its endpoint, Bearer API key, and a short route prefix. Add the gateway's model IDs manually; for example, an OpenRouter route named `openrouter` uses `openrouter:anthropic/claude-…`. This supports standard OpenAI-compatible APIs without OpenWorker maintaining their full catalog. + ## Privacy OpenWorker is local-first. Everything lives on your machine: the agent loop, your conversations, connector tokens, and model keys - all in the app's local secret store. The only cloud piece is a small service that brokers OAuth handshakes for connectors. You can always use the App without signing-in - use the connectors via manually-created credentials/API-keys. diff --git a/coworker/providers/__init__.py b/coworker/providers/__init__.py index 6b34c141..2e10ed42 100644 --- a/coworker/providers/__init__.py +++ b/coworker/providers/__init__.py @@ -15,6 +15,8 @@ build_provider_client, detect_provider, get_descriptor, + is_custom_provider, + is_removed_custom_provider, provider_descriptors, provider_names, verify_provider_key, @@ -38,6 +40,8 @@ "provider_descriptors", "provider_names", "get_descriptor", + "is_custom_provider", + "is_removed_custom_provider", "build_provider_client", "detect_provider", "verify_provider_key", diff --git a/coworker/providers/registry.py b/coworker/providers/registry.py index b5c628c6..2acecee8 100644 --- a/coworker/providers/registry.py +++ b/coworker/providers/registry.py @@ -24,6 +24,9 @@ from .openai_provider import OpenAIProvider DEFAULT_OLLAMA_URL = "http://localhost:11434" +# The index deliberately lives in the SecretStore too: it names profiles which contain +# credentials, and keeping the two together makes a future secure-store backend a drop-in. +CUSTOM_PROVIDERS_PROFILE = "providers:custom" @dataclass(frozen=True) @@ -193,6 +196,50 @@ def _compat( ) +def _custom_provider_descriptor(name: str, title: str) -> ProviderDescriptor: + """Descriptor synthesized for a user-defined OpenAI-compatible endpoint. + + Unlike the built-ins there is no catalog or capability claim here. The endpoint and + bearer key are required; model ids are supplied separately by the user. + """ + return ProviderDescriptor( + name=name, + title=title, + needs_key=True, + fields=[ + ProviderField("api_key", f"{title} API key", secret=True), + ProviderField( + "base_url", + "Endpoint", + required=True, + placeholder="https://…/v1", + help="An HTTPS OpenAI-compatible API base URL ending in /v1. HTTP is allowed only for local gateways.", + ), + ], + build=_openai_compat(title, ""), + blurb="Uses the standard OpenAI Chat Completions API. Add model IDs manually below.", + ) + + +def _disabled_custom_provider_descriptor(name: str, title: str) -> ProviderDescriptor: + """Keep a removed route recognizable so old sessions fail closed, never reroute.""" + + def build(profile: dict[str, Any], secrets: Any) -> ProviderClient: + raise RuntimeError( + f"{title} was removed from Settings ▸ Models. Configure that provider again " + "before continuing this session." + ) + + return ProviderDescriptor( + name=name, + title=title, + needs_key=True, + fields=[], + build=build, + blurb="Removed custom provider.", + ) + + DESCRIPTORS: list[ProviderDescriptor] = [ ProviderDescriptor( name="openai", @@ -357,23 +404,78 @@ def _compat( _BY_NAME = {d.name: d for d in DESCRIPTORS} -def provider_descriptors() -> list[ProviderDescriptor]: - return list(DESCRIPTORS) - - -def provider_names() -> list[str]: - return [d.name for d in DESCRIPTORS] - - -def get_descriptor(name: str) -> Optional[ProviderDescriptor]: - return _BY_NAME.get(name) +def _custom_provider_titles(secrets: Any = None) -> dict[str, str]: + """Return the validated custom-provider index without exposing credentials.""" + if secrets is None: + return {} + raw = secrets.get(CUSTOM_PROVIDERS_PROFILE) or {} + providers = raw.get("providers") if isinstance(raw, dict) else None + if not isinstance(providers, dict): + return {} + return { + name: title + for name, title in providers.items() + if isinstance(name, str) + and isinstance(title, str) + and name not in _BY_NAME + and title.strip() + } + + +def _removed_custom_provider_titles(secrets: Any = None) -> dict[str, str]: + """Routes kept as tombstones after deletion to protect existing sessions.""" + if secrets is None: + return {} + raw = secrets.get(CUSTOM_PROVIDERS_PROFILE) or {} + removed = raw.get("removed") if isinstance(raw, dict) else None + if not isinstance(removed, dict): + return {} + return { + name: title + for name, title in removed.items() + if isinstance(name, str) + and isinstance(title, str) + and name not in _BY_NAME + and title.strip() + } + + +def is_custom_provider(name: str, secrets: Any = None) -> bool: + return name in _custom_provider_titles(secrets) + + +def is_removed_custom_provider(name: str, secrets: Any = None) -> bool: + return name in _removed_custom_provider_titles(secrets) + + +def provider_descriptors(secrets: Any = None) -> list[ProviderDescriptor]: + custom = [ + _custom_provider_descriptor(name, title) + for name, title in _custom_provider_titles(secrets).items() + ] + return [*DESCRIPTORS, *custom] + + +def provider_names(secrets: Any = None) -> list[str]: + return [d.name for d in provider_descriptors(secrets)] + + +def get_descriptor(name: str, secrets: Any = None) -> Optional[ProviderDescriptor]: + builtin = _BY_NAME.get(name) + if builtin is not None: + return builtin + title = _custom_provider_titles(secrets).get(name) + if title: + return _custom_provider_descriptor(name, title) + title = _removed_custom_provider_titles(secrets).get(name) + return _disabled_custom_provider_descriptor(name, title) if title else None def build_provider_client( name: str, profile: dict[str, Any], secrets: Any ) -> ProviderClient: """Build a `ProviderClient` for `name` from its stored profile. Unknown → OpenAI default.""" - descriptor = _BY_NAME.get(name) or _BY_NAME["openai"] + descriptor = get_descriptor(name, secrets) or _BY_NAME["openai"] return descriptor.build(profile or {}, secrets) @@ -399,6 +501,7 @@ def verify_provider_key( api_key: Optional[str] = None, base_url: Optional[str] = None, timeout: float = 10.0, + descriptor: Optional[ProviderDescriptor] = None, ) -> dict[str, Any]: """Validate a provider's credentials with one cheap, read-only call (list models) — the same pattern connectors use to validate tokens. Transient: callers pass the key directly so a user @@ -406,7 +509,7 @@ def verify_provider_key( """ import httpx - d = _BY_NAME.get(name) or _BY_NAME["openai"] + d = descriptor or _BY_NAME.get(name) or _BY_NAME["openai"] key = (api_key or "").strip() try: if name == "anthropic": @@ -446,6 +549,14 @@ def verify_provider_key( if resp.status_code < 300: return {"ok": True} + if resp.status_code in (404, 405) and d.name not in _BY_NAME: + # Chat Completions compatibility does not require the optional model-list + # endpoint. The gateway was reached, so do not call a usable custom route + # invalid merely because it cannot provide a catalog. + return { + "ok": True, + "warning": "Endpoint reached, but it does not expose /models. Add a model ID and try a chat to finish checking it.", + } if resp.status_code in (401, 403): if name == "ollama": return {"ok": False, "error": "Server rejected the request."} diff --git a/coworker/providers/router.py b/coworker/providers/router.py index c2c56c47..4a6d7f19 100644 --- a/coworker/providers/router.py +++ b/coworker/providers/router.py @@ -51,7 +51,7 @@ def _provider_name(self, model: str) -> str: """ if ":" in model: prefix = model.split(":", 1)[0] - if get_descriptor(prefix) is not None: + if get_descriptor(prefix, getattr(self, "_secrets", None)) is not None: return prefix return self._default @@ -68,14 +68,14 @@ def _client_for(self, model: str) -> ProviderClient: return client @staticmethod - def _bare(model: str) -> str: + def _bare(model: str, secrets: Any = None) -> str: """Strip a KNOWN provider prefix; the underlying SDK wants the bare model name. A model whose first segment isn't a provider (e.g. `qwen2.5-coder:32b` — a version tag, not a prefix) is returned unchanged, so the colon isn't mistaken for a provider separator. """ if ":" in model: prefix, rest = model.split(":", 1) - if get_descriptor(prefix) is not None: + if get_descriptor(prefix, secrets) is not None: return rest return model @@ -98,7 +98,7 @@ def complete( ): self._note_use(model) return self._client_for(model).complete( - model=self._bare(model), messages=messages, tools=tools, **settings + model=self._bare(model, self._secrets), messages=messages, tools=tools, **settings ) def stream( @@ -111,7 +111,7 @@ def stream( ): self._note_use(model) return self._client_for(model).stream( - model=self._bare(model), messages=messages, tools=tools, **settings + model=self._bare(model, self._secrets), messages=messages, tools=tools, **settings ) def capabilities(self, model: str): diff --git a/coworker/server/app.py b/coworker/server/app.py index 3b054945..83d1e8f0 100644 --- a/coworker/server/app.py +++ b/coworker/server/app.py @@ -1308,6 +1308,13 @@ def providers_set(body: dict) -> dict[str, Any]: return {"ok": False, "error": "name required"} return manager.set_provider(name, (body or {}).get("fields")) + @app.post("/v1/providers/custom") + def providers_custom_create(body: dict) -> dict[str, Any]: + body = body or {} + return manager.create_custom_provider( + body.get("name", ""), body.get("title", ""), body.get("fields") + ) + @app.delete("/v1/providers/{name}") def providers_remove(name: str) -> dict[str, Any]: return manager.remove_provider(name) diff --git a/coworker/server/manager.py b/coworker/server/manager.py index b2d9be0c..e230a56e 100644 --- a/coworker/server/manager.py +++ b/coworker/server/manager.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio +import ipaddress import json import logging import os @@ -17,6 +18,7 @@ import time from pathlib import Path from typing import Any, Optional +from urllib.parse import urlsplit from ..agent import build_engine from ..agents import get_agent @@ -76,9 +78,12 @@ ProviderClient, ProviderRouter, get_descriptor, + is_custom_provider, + is_removed_custom_provider, provider_descriptors, verify_provider_key, ) +from ..providers.registry import CUSTOM_PROVIDERS_PROFILE from ..secrets import SecretStore, state_dir from ..sessions import SessionRecord from ..skills import SkillLoader @@ -88,6 +93,28 @@ logger = logging.getLogger("coworker.manager") +def _custom_endpoint_error(url: str) -> Optional[str]: + """Allow HTTPS everywhere; plaintext HTTP only for an explicitly local gateway.""" + try: + parsed = urlsplit((url or "").strip()) + except ValueError: + return "Enter a valid HTTPS endpoint." + if parsed.scheme == "https" and parsed.hostname: + return None + if parsed.scheme != "http" or not parsed.hostname: + return "Endpoint must be an HTTPS URL (for example, https://gateway.example/v1)." + host = parsed.hostname.lower() + if host == "localhost": + return None + try: + address = ipaddress.ip_address(host) + except ValueError: + return "HTTP endpoints are allowed only for localhost or a private IP address." + if address.is_loopback or address.is_private: + return None + return "HTTP endpoints are allowed only for localhost or a private IP address." + + def _grants_of(engine) -> dict[str, Any]: """The engine's session-scoped "Always allow" approvals, in persistable shape.""" tools = sorted(getattr(engine.permissions, "session_allow_tools", None) or ()) @@ -1394,7 +1421,7 @@ def get_providers(self) -> list[dict[str, Any]]: import os out: list[dict[str, Any]] = [] - for d in provider_descriptors(): + for d in provider_descriptors(self.secrets): profile = self.secrets.get(f"provider:{d.name}") or {} if d.needs_key: configured = bool(profile.get("api_key")) or bool( @@ -1420,10 +1447,56 @@ def get_providers(self) -> list[dict[str, Any]]: "last_used_at": (self._prefs.get("provider_last_used") or {}).get( d.name ), + "is_custom": is_custom_provider(d.name, self.secrets), } ) return out + def create_custom_provider( + self, name: str, title: str, fields: Optional[dict[str, Any]] + ) -> dict[str, Any]: + """Create a named OpenAI-compatible provider, then save its first credentials. + + The short slug is the stable model-route prefix (``slug:model-id``), so it is + deliberately immutable and restricted to a transport-safe subset. + """ + name = (name or "").strip().lower() + title = (title or "").strip() + if not re.fullmatch(r"[a-z][a-z0-9_-]{0,31}", name): + return { + "ok": False, + "error": "Route prefix must start with a letter and use only lowercase letters, numbers, - or _.", + } + if not title or len(title) > 80: + return {"ok": False, "error": "Enter a provider name (up to 80 characters)."} + if get_descriptor(name) is not None or is_custom_provider(name, self.secrets): + return {"ok": False, "error": "That route prefix is already in use."} + fields = fields or {} + base_url = str(fields.get("base_url") or "").strip() + endpoint_error = _custom_endpoint_error(base_url) + if endpoint_error: + return {"ok": False, "error": endpoint_error} + if not str(fields.get("api_key") or "").strip(): + return {"ok": False, "error": "Enter an API key."} + + index = dict(self.secrets.get(CUSTOM_PROVIDERS_PROFILE) or {}) + providers = dict(index.get("providers") or {}) + removed = dict(index.get("removed") or {}) + previous_removed = dict(removed) + providers[name] = title + removed.pop(name, None) # re-creating a route clears its fail-closed tombstone + self.secrets.put( + CUSTOM_PROVIDERS_PROFILE, {"providers": providers, "removed": removed} + ) + saved = self.set_provider(name, fields) + if not saved.get("ok"): + providers.pop(name, None) + self.secrets.put( + CUSTOM_PROVIDERS_PROFILE, + {"providers": providers, "removed": previous_removed}, + ) + return saved + def pick_native_folder(self) -> dict[str, Any]: """Open the OS folder picker FROM THE SIDECAR — the browser GUI can't obtain absolute paths from web file dialogs, but the sidecar is local and can (the desktop shell uses @@ -1510,9 +1583,11 @@ def set_provider( ) -> dict[str, Any]: """Store a provider's config in its `provider:` SecretStore profile and rebuild its cached client. Merges provided fields into any existing profile.""" - d = get_descriptor(name) + d = get_descriptor(name, self.secrets) if d is None: return {"ok": False, "error": f"unknown provider: {name}"} + if is_removed_custom_provider(name, self.secrets): + return {"ok": False, "error": "provider was removed; create it again first"} fields = fields or {} profile = dict(self.secrets.get(f"provider:{name}") or {}) for f in d.fields: @@ -1528,6 +1603,10 @@ def set_provider( missing = [f.label for f in d.fields if f.required and not profile.get(f.key)] if missing: return {"ok": False, "error": "missing: " + ", ".join(missing)} + if is_custom_provider(name, self.secrets): + endpoint_error = _custom_endpoint_error(str(profile.get("base_url") or "")) + if endpoint_error: + return {"ok": False, "error": endpoint_error} # A (re)pasted key stamps its save date — Settings shows "key added " so stale # keys are visible. Endpoint-only saves keep the original stamp. if isinstance(fields.get("api_key"), str) and fields["api_key"].strip(): @@ -1555,10 +1634,35 @@ def remove_provider(self, name: str) -> dict[str, Any]: """Forget a provider's stored config (Settings ▸ Models "Remove key"). The whole `provider:` profile goes — key, endpoint, key_set_at — so the provider reads as never configured. Curated models stay; they just gray out until a new key.""" - d = get_descriptor(name) + d = get_descriptor(name, self.secrets) if d is None: return {"ok": False, "error": f"unknown provider: {name}"} self.secrets.delete(f"provider:{name}") + if is_custom_provider(name, self.secrets): + index = dict(self.secrets.get(CUSTOM_PROVIDERS_PROFILE) or {}) + providers = dict(index.get("providers") or {}) + removed = dict(index.get("removed") or {}) + title = providers.pop(name, d.title) + removed[name] = title + self.secrets.put( + CUSTOM_PROVIDERS_PROFILE, + {"providers": providers, "removed": removed}, + ) + # A deleted custom route must not remain selectable and accidentally fall + # through to the default OpenAI provider as an unknown ``prefix:model``. + prefix = name + ":" + self._prefs["models"] = [ + m for m in self._prefs.get("models", []) if not str(m).startswith(prefix) + ] + self._prefs["hidden_models"] = [ + m + for m in self._prefs.get("hidden_models", []) + if not str(m).startswith(prefix) + ] + if str(self.model).startswith(prefix): + self.model = "gpt-5.6-sol" + self._prefs["default_model"] = self.model + self._save_prefs() self._refresh_provider(name) return {"ok": True, "provider": name} @@ -1570,7 +1674,7 @@ def verify_provider( the key blank (e.g. testing an already-configured provider).""" import os - d = get_descriptor(name) + d = get_descriptor(name, self.secrets) if d is None: return {"ok": False, "error": f"unknown provider: {name}"} fields = fields or {} @@ -1581,18 +1685,27 @@ def verify_provider( base_url = (fields.get("base_url") or profile.get("base_url") or "").strip() if d.needs_key and not api_key: return {"ok": False, "error": "Enter an API key to test."} - return verify_provider_key(name, api_key=api_key, base_url=base_url) + # A verification call sends the key to the supplied endpoint too. Apply the + # same transport policy as Save so an unsaved edit cannot accidentally + # send a credential over plaintext to a remote host. + if is_custom_provider(name, self.secrets): + endpoint_error = _custom_endpoint_error(base_url) + if endpoint_error: + return {"ok": False, "error": endpoint_error} + return verify_provider_key( + name, api_key=api_key, base_url=base_url, descriptor=d + ) def _model_provider(self, model: str) -> str: """The provider a model string routes to (known `prefix:` or the OpenAI default).""" if ":" in (model or ""): prefix = model.split(":", 1)[0] - if get_descriptor(prefix) is not None: + if get_descriptor(prefix, self.secrets) is not None: return prefix return "openai" def _provider_configured(self, name: str) -> bool: - d = get_descriptor(name) + d = get_descriptor(name, self.secrets) if d is None: return False if not d.needs_key: diff --git a/surfaces/gui/src/api.ts b/surfaces/gui/src/api.ts index c700cdd0..78843960 100644 --- a/surfaces/gui/src/api.ts +++ b/surfaces/gui/src/api.ts @@ -1279,6 +1279,7 @@ export interface ProviderInfo { blurb?: string; // one-line note under the title ("Uses X's OpenAI-compatible API…") key_set_at?: string | null; // ISO date the key was last (re)saved — absent for env-only config last_used_at?: number | null; // epoch secs the provider last served a completion + is_custom?: boolean; // user-created OpenAI-compatible provider } export async function getProviders(): Promise { @@ -1298,6 +1299,20 @@ export async function setProvider( return res.json(); } +/** Create a named, user-configured OpenAI-compatible provider. */ +export async function createCustomProvider( + name: string, + title: string, + fields: Record, +): Promise<{ ok: boolean; error?: string; provider?: string }> { + const res = await fetch(`${httpBase()}/v1/providers/custom`, { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ name, title, fields }), + }); + return res.json(); +} + /** Forget a provider's stored config (Settings ▸ Models "Remove key…"). */ export async function removeProvider(name: string): Promise<{ ok: boolean; error?: string }> { const res = await fetch(`${httpBase()}/v1/providers/${encodeURIComponent(name)}`, { @@ -1310,7 +1325,7 @@ export async function removeProvider(name: string): Promise<{ ok: boolean; error export async function verifyProvider( name: string, fields: Record, -): Promise<{ ok: boolean; error?: string }> { +): Promise<{ ok: boolean; error?: string; warning?: string }> { const res = await fetch(`${httpBase()}/v1/providers/verify`, { method: "POST", headers: { "Content-Type": "application/json" }, diff --git a/surfaces/gui/src/components/ManageTabs.test.tsx b/surfaces/gui/src/components/ManageTabs.test.tsx new file mode 100644 index 00000000..39f4072d --- /dev/null +++ b/surfaces/gui/src/components/ManageTabs.test.tsx @@ -0,0 +1,92 @@ +import { afterEach, expect, it, vi } from "vitest"; +import { cleanup, fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { ModelsTab } from "./ManageTabs"; + +afterEach(() => { + cleanup(); + vi.unstubAllGlobals(); +}); + +it("creates a custom provider and opens its manual-model configuration", async () => { + let created = false; + const calls: Array<{ url: string; method: string; body?: unknown }> = []; + vi.stubGlobal( + "fetch", + vi.fn(async (url: string, init?: RequestInit) => { + const method = (init?.method || "GET").toUpperCase(); + calls.push({ + url, + method, + body: init?.body ? JSON.parse(String(init.body)) : undefined, + }); + if (url.includes("/v1/providers/custom") && method === "POST") { + created = true; + return { json: async () => ({ ok: true, provider: "openrouter" }) } as Response; + } + if (url.includes("/v1/providers")) { + const providers: any[] = [ + { + name: "openai", + title: "OpenAI", + needs_key: true, + fields: [{ key: "api_key", label: "OpenAI API key", secret: true, required: true }], + configured: false, + values: {}, + suggested_models: [], + recommended_model: "gpt-5.6-sol", + }, + ]; + if (created) { + providers.push({ + name: "openrouter", + title: "OpenRouter", + needs_key: true, + fields: [ + { key: "api_key", label: "OpenRouter API key", secret: true, required: true }, + { key: "base_url", label: "Endpoint", secret: false, required: true, placeholder: "https://…/v1" }, + ], + configured: true, + values: { base_url: "https://openrouter.ai/api/v1" }, + suggested_models: [], + recommended_model: null, + is_custom: true, + }); + } + return { json: async () => providers } as Response; + } + if (url.includes("/v1/settings")) { + return { + json: async () => ({ + model: "gpt-5.6-sol", + models: ["gpt-5.6-sol"], + model_labels: {}, + source: null, + }), + } as Response; + } + return { json: async () => ({ ok: true }) } as Response; + }), + ); + + render(); + fireEvent.click(await screen.findByTestId("set-provider-custom-add")); + fireEvent.change(screen.getByLabelText("Provider name"), { target: { value: "OpenRouter" } }); + fireEvent.change(screen.getByLabelText("Route prefix"), { target: { value: "openrouter" } }); + fireEvent.change(screen.getByLabelText("Endpoint"), { target: { value: "https://openrouter.ai/api/v1" } }); + fireEvent.change(screen.getByLabelText("API key"), { target: { value: "or-test" } }); + fireEvent.click(screen.getByRole("button", { name: "Add provider" })); + + await waitFor(() => { + expect(calls.find((c) => c.url.includes("/v1/providers/custom")))?.toMatchObject({ + method: "POST", + body: { + name: "openrouter", + title: "OpenRouter", + fields: { base_url: "https://openrouter.ai/api/v1", api_key: "or-test" }, + }, + }); + }); + expect((await screen.findByTestId("set-field-base_url") as HTMLInputElement).value).toBe( + "https://openrouter.ai/api/v1", + ); +}); diff --git a/surfaces/gui/src/components/ManageTabs.tsx b/surfaces/gui/src/components/ManageTabs.tsx index f9ab2992..b6498102 100644 --- a/surfaces/gui/src/components/ManageTabs.tsx +++ b/surfaces/gui/src/components/ManageTabs.tsx @@ -6,6 +6,7 @@ import { connectManaged, connectMcpBacked, connectMcp, + createCustomProvider, deleteMcpServer, disallowUser, getMcpServers, @@ -77,12 +78,21 @@ const EXAMPLE = `{ // per-provider ModelChecklist / read-only model preview (form view). export function ModelsTab() { const [settings, setSettings] = useState(null); + const [addingCustom, setAddingCustom] = useState(false); + const [openCustom, setOpenCustom] = useState(null); const refreshSettings = () => getSettings().then(setSettings).catch(() => setSettings(null)); const ps = useProviderSetup({ onSaved: refreshSettings }); useEffect(() => { refreshSettings(); }, []); + useEffect(() => { + if (openCustom && ps.providers.some((p) => p.name === openCustom)) { + setOpenCustom(null); + ps.openProvider(openCustom); + } + }, [openCustom, ps.providers]); + if (!settings) return
Loading…
; const info = ps.info; @@ -91,7 +101,25 @@ export function ModelsTab() { if (ps.sel === null) { return (
- + {addingCustom ? ( + setAddingCustom(false)} + onCreated={async (name) => { + setAddingCustom(false); + setOpenCustom(name); + await ps.refreshProviders(); + refreshSettings(); + }} + /> + ) : ( + setAddingCustom(true)} + /> + )}
); @@ -108,10 +136,11 @@ export function ModelsTab() { className="text-[12.5px] text-danger/80 hover:text-danger hover:underline underline-offset-2" data-testid="set-remove-key" onClick={() => { - if (window.confirm(`Remove the ${info?.title} key from this computer?`)) ps.removeKey(); + const what = info?.is_custom ? "provider and its key" : "key"; + if (window.confirm(`Remove the ${info?.title} ${what} from this computer?`)) ps.removeKey(); }} > - Remove key… + {info?.is_custom ? "Remove provider…" : "Remove key…"} ) : null } @@ -171,6 +200,63 @@ export function ModelsTab() { ); } +/** First-run form for a generic OpenAI-compatible provider. It intentionally does + * not fetch a model catalog: users add the exact model ids their gateway exposes. */ +function CustomProviderForm({ + onCancel, + onCreated, +}: { + onCancel: () => void; + onCreated: (name: string) => Promise; +}) { + const [title, setTitle] = useState(""); + const [name, setName] = useState(""); + const [baseUrl, setBaseUrl] = useState(""); + const [apiKey, setApiKey] = useState(""); + const [error, setError] = useState(null); + const [saving, setSaving] = useState(false); + const input = "w-full px-3 py-2 rounded-lg border border-line bg-panel text-[13.5px] outline-none focus:border-accent"; + + const save = async () => { + setSaving(true); + setError(null); + const result: { ok: boolean; error?: string; provider?: string } = await createCustomProvider( + name, + title, + { base_url: baseUrl, api_key: apiKey }, + ).catch(() => ({ ok: false, error: "unreachable" })); + setSaving(false); + if (!result.ok || !result.provider) { + setError(result.error || "Couldn’t add provider."); + return; + } + await onCreated(result.provider); + }; + + return ( +
+
Add OpenAI-compatible provider
+

+ This provider uses Chat Completions with Bearer authentication. Use HTTPS unless the gateway is local; you’ll add its model IDs manually next. +

+ + setTitle(e.target.value)} /> + + setName(e.target.value.toLowerCase())} /> +

Models will be selected as openrouter:model-id.

+ + setBaseUrl(e.target.value)} /> + + setApiKey(e.target.value)} /> + {error &&
{error}
} +
+ + +
+
+ ); +} + // The gallery view's "In the composer's picker" card: every curated model across providers, // with its provider tag. Unticking removes it from the picker; adding happens from a // provider's card (the ModelChecklist there has the suggested list + free-type add). diff --git a/surfaces/gui/src/providers/ProviderSetup.tsx b/surfaces/gui/src/providers/ProviderSetup.tsx index eef8c8b7..6986498d 100644 --- a/surfaces/gui/src/providers/ProviderSetup.tsx +++ b/surfaces/gui/src/providers/ProviderSetup.tsx @@ -133,7 +133,9 @@ export function useProviderSetup(opts?: { onSaved?: () => void }): ProviderSetup setFields(next); setDirty(!!draft && Object.values(draft).some(Boolean)); setVerify({ state: "idle" }); - setShowEndpoint(false); + // A custom provider cannot work without its endpoint, so keep that field + // visible instead of hiding it behind the built-ins' expert disclosure. + setShowEndpoint(!!p?.fields.find((f) => f.key === "base_url")?.required); }; const backToGallery = () => { @@ -151,14 +153,17 @@ export function useProviderSetup(opts?: { onSaved?: () => void }): ProviderSetup const runTestAndSave = async (): Promise => { if (!sel) return false; setVerify({ state: "testing" }); - const res = await verifyProvider(sel, fields).catch(() => ({ ok: false, error: "unreachable" })); + const res: { ok: boolean; error?: string; warning?: string } = await verifyProvider( + sel, + fields, + ).catch(() => ({ ok: false, error: "unreachable" })); if (!res.ok) { setVerify({ state: "error", msg: res.error || "couldn't verify" }); return false; } if (dirty || !info?.configured) await setProvider(sel, fields).catch(() => {}); if (!info?.needs_key) setKeylessOk((s) => new Set(s).add(sel)); - setVerify({ state: "ok" }); + setVerify({ state: "ok", msg: res.warning }); setDirty(false); setDrafts((d) => ({ ...d, [sel]: {} })); await refreshProviders(); @@ -167,11 +172,15 @@ export function useProviderSetup(opts?: { onSaved?: () => void }): ProviderSetup // the timeout would fire its stale closure (dirty/fields from before the save) and // re-stash the just-saved key as a draft — the state-restore bug (owner catch // 2026-07-19). This return path clears the draft unconditionally. - backTimer.current = window.setTimeout(() => { - setDrafts((d) => ({ ...d, [sel]: {} })); - setSel(null); - setVerify({ state: "idle" }); - }, 900); + // A catalog-less gateway is still usable, but its warning contains the next action; + // leave the form open instead of making it disappear after the usual success animation. + if (!res.warning) { + backTimer.current = window.setTimeout(() => { + setDrafts((d) => ({ ...d, [sel]: {} })); + setSel(null); + setVerify({ state: "idle" }); + }, 900); + } return true; }; @@ -270,11 +279,13 @@ export function ProviderCards({ tp, gridClass = "grid grid-cols-2 gap-2.5", lastUsed = false, + onAddCustom, }: { ps: ProviderSetupState; tp: string; // testid prefix ("ob" onboarding, "set" settings) gridClass?: string; lastUsed?: boolean; + onAddCustom?: () => void; }) { const card = "flex items-center gap-2.5 rounded-xl border border-line bg-panel px-3 py-2.5 text-left hover:border-lineStrong transition-colors"; @@ -295,6 +306,19 @@ export function ProviderCards({ ))} + {onAddCustom && ( + + )} ); } @@ -465,6 +489,7 @@ export function ProviderForm({ {/* Error line: fixed height so failures never reflow the form. */}
{ps.verify.state === "error" && {ps.verify.msg}} + {ps.verify.state === "ok" && ps.verify.msg && {ps.verify.msg}}
{footer} diff --git a/tests/test_provider_router.py b/tests/test_provider_router.py index 81e53d30..c8e4077b 100644 --- a/tests/test_provider_router.py +++ b/tests/test_provider_router.py @@ -5,6 +5,8 @@ from types import SimpleNamespace +import pytest + from coworker.providers import ( AssistantTurn, ModelCapabilities, @@ -310,6 +312,109 @@ def test_manager_provider_config(tmp_path, monkeypatch): assert mgr.set_provider("nope", {})["ok"] is False # unknown provider rejected +def test_custom_openai_compatible_provider_routes_and_persists(tmp_path, monkeypatch): + """A custom route owns its key/endpoint and strips only its own prefix.""" + monkeypatch.setenv("COWORKER_STATE_DIR", str(tmp_path / "state")) + from coworker.server.manager import SessionManager + + mgr = SessionManager(data_dir=tmp_path) + created = mgr.create_custom_provider( + "openrouter", + "OpenRouter", + {"api_key": "or-test", "base_url": "https://openrouter.ai/api/v1"}, + ) + assert created["ok"] is True and created["provider"] == "openrouter" + + providers = {p["name"]: p for p in mgr.get_providers()} + assert providers["openrouter"]["is_custom"] is True + assert providers["openrouter"]["values"] == { + "base_url": "https://openrouter.ai/api/v1" + } + assert "or-test" not in str(providers) + + router = mgr.provider + assert router._provider_name("openrouter:anthropic/claude-example") == "openrouter" + assert ( + router._bare("openrouter:anthropic/claude-example", mgr.secrets) + == "anthropic/claude-example" + ) + client = router._client_for("openrouter:anthropic/claude-example") + assert client._api_key == "or-test" and client._base_url == "https://openrouter.ai/api/v1" + + mgr.add_model("openrouter:anthropic/claude-example") + assert "openrouter:anthropic/claude-example" in mgr.get_settings()["models"] + assert mgr.remove_provider("openrouter")["ok"] is True + assert "openrouter" not in {p["name"] for p in mgr.get_providers()} + assert "openrouter:anthropic/claude-example" not in mgr.get_settings()["models"] + # Existing sessions retain the full routed model id. After removal that must + # fail locally, never fall through to whichever provider is now the default. + assert router._provider_name("openrouter:anthropic/claude-example") == "openrouter" + with pytest.raises(RuntimeError, match="was removed"): + router.complete(model="openrouter:anthropic/claude-example", messages=[]) + + # Re-creating the same route intentionally clears its tombstone. + assert mgr.create_custom_provider( + "openrouter", + "OpenRouter", + {"api_key": "or-recreated", "base_url": "https://openrouter.ai/api/v1"}, + )["ok"] + assert router._client_for("openrouter:anthropic/claude-example")._api_key == "or-recreated" + + +def test_custom_provider_requires_safe_route_and_complete_config(tmp_path, monkeypatch): + monkeypatch.setenv("COWORKER_STATE_DIR", str(tmp_path / "state")) + from coworker.server.manager import SessionManager + + mgr = SessionManager(data_dir=tmp_path) + assert not mgr.create_custom_provider("Bad Name", "Gateway", {})["ok"] + assert not mgr.create_custom_provider( + "gateway", "Gateway", {"api_key": "x", "base_url": "gateway.example/v1"} + )["ok"] + assert not mgr.create_custom_provider( + "remote-http", "Remote HTTP", {"api_key": "x", "base_url": "http://example.com/v1"} + )["ok"] + assert mgr.create_custom_provider( + "local-http", "Local HTTP", {"api_key": "x", "base_url": "http://127.0.0.1:8000/v1"} + )["ok"] + + +def test_custom_provider_verify_allows_missing_models_endpoint(tmp_path, monkeypatch): + monkeypatch.setenv("COWORKER_STATE_DIR", str(tmp_path / "state")) + from coworker.server.manager import SessionManager + + mgr = SessionManager(data_dir=tmp_path) + assert mgr.create_custom_provider( + "gateway", "Gateway", {"api_key": "test", "base_url": "https://gateway.example/v1"} + )["ok"] + + monkeypatch.setattr( + "httpx.get", lambda *args, **kwargs: SimpleNamespace(status_code=404) + ) + checked = mgr.verify_provider("gateway", {}) + assert checked["ok"] is True and "does not expose /models" in checked["warning"] + + +def test_custom_provider_verify_applies_endpoint_transport_policy(tmp_path, monkeypatch): + monkeypatch.setenv("COWORKER_STATE_DIR", str(tmp_path / "state")) + from coworker.server.manager import SessionManager + + mgr = SessionManager(data_dir=tmp_path) + assert mgr.create_custom_provider( + "gateway", "Gateway", {"api_key": "test", "base_url": "https://gateway.example/v1"} + )["ok"] + + # Verify accepts unsaved form fields, so it must independently reject an + # unsafe replacement URL before it makes any request with the key. + monkeypatch.setattr( + "httpx.get", lambda *args, **kwargs: pytest.fail("unsafe endpoint was requested") + ) + checked = mgr.verify_provider( + "gateway", {"base_url": "http://example.com/v1", "api_key": "test"} + ) + assert checked["ok"] is False + assert "localhost or a private IP" in checked["error"] + + def test_manager_curated_models(tmp_path, monkeypatch): """No seed list: the picker is the curated matrix filtered to key-holding providers, plus user-added custom ids. A fresh install shows only the (not-yet-usable) default. diff --git a/tests/test_server.py b/tests/test_server.py index af870c51..1a8e23e9 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -989,6 +989,32 @@ def test_provider_set_and_remove_roundtrip(tmp_path): assert not client.delete("/v1/providers/nope").json()["ok"] +def test_custom_provider_rest_create_and_delete(tmp_path, monkeypatch): + """The public Settings API creates a separately-routed custom provider.""" + monkeypatch.setenv("COWORKER_STATE_DIR", str(tmp_path / "state")) + manager = SessionManager(workspace=tmp_path, provider=ScriptedProvider([])) + client = TestClient(create_app(manager)) + + created = client.post( + "/v1/providers/custom", + json={ + "name": "openrouter", + "title": "OpenRouter", + "fields": { + "api_key": "or-test", + "base_url": "https://openrouter.ai/api/v1", + }, + }, + ).json() + assert created == {"ok": True, "provider": "openrouter", "recommended_model": None} + provider = {p["name"]: p for p in client.get("/v1/providers").json()}["openrouter"] + assert provider["is_custom"] is True and provider["values"]["base_url"].endswith("/v1") + assert "or-test" not in str(provider) + + assert client.delete("/v1/providers/openrouter").json()["ok"] is True + assert "openrouter" not in {p["name"] for p in client.get("/v1/providers").json()} + + def test_always_allow_grants_survive_restart(tmp_path): """"Always allow" is session-scoped, and the session outlives the process — a restart (fresh manager over the same store) must not re-ask for an approved command