diff --git a/apps/desktop/src/demo/api.ts b/apps/desktop/src/demo/api.ts index cc21e8f..7323793 100644 --- a/apps/desktop/src/demo/api.ts +++ b/apps/desktop/src/demo/api.ts @@ -1,6 +1,7 @@ import type { CallDetail, CursorHarnessStatus, + ExternalApiSettings, LlmCall, Model, Overview, @@ -86,6 +87,7 @@ let harnessStatus: CursorHarnessStatus = { let detailed = true; let portSettings = { proxy_port: 0, service_port: 0 }; +let externalApiSettings: ExternalApiSettings = { enabled: false, api_key: "" }; let proxySettings: ProxySettings = { mode: "default", address: "", @@ -148,6 +150,11 @@ export function installDemoApi() { portSettings = body as typeof portSettings; return json(portSettings); } + if (path === "/settings/external-api" && method === "GET") return json(externalApiSettings); + if (path === "/settings/external-api") { + externalApiSettings = body as ExternalApiSettings; + return json(externalApiSettings); + } if (path === "/settings/storage/statistics" && method === "GET") return json(storage); if (path === "/settings/storage/statistics") { const scope = (body as { scope?: string } | null)?.scope ?? "details"; diff --git a/apps/desktop/src/features/settings/ExternalApiSettingsCard.module.scss b/apps/desktop/src/features/settings/ExternalApiSettingsCard.module.scss new file mode 100644 index 0000000..5999919 --- /dev/null +++ b/apps/desktop/src/features/settings/ExternalApiSettingsCard.module.scss @@ -0,0 +1,38 @@ +@use "../../styles/typography" as type; + +.content { + display: grid; + gap: 16px; + padding: 16px; +} + +.row, .address { + display: flex; + align-items: center; + justify-content: space-between; + gap: 16px; + min-width: 0; +} + +.description { + display: grid; + gap: 5px; + min-width: 0; + + small { + color: var(--vscode-descriptionForeground); + font-size: type.$font-size-xs; + } +} + +.address { + flex-wrap: wrap; + + code { + flex: 1 1 240px; + min-width: 0; + overflow-wrap: anywhere; + color: var(--vscode-descriptionForeground); + font-size: type.$font-size-xs; + } +} diff --git a/apps/desktop/src/features/settings/ExternalApiSettingsCard.tsx b/apps/desktop/src/features/settings/ExternalApiSettingsCard.tsx new file mode 100644 index 0000000..d96d0e9 --- /dev/null +++ b/apps/desktop/src/features/settings/ExternalApiSettingsCard.tsx @@ -0,0 +1,62 @@ +import { useEffect, useState } from "react"; +import { api, type ExternalApiSettings } from "../../shared/api"; +import { Button } from "../../shared/ui/Button"; +import { FormField, SecretTextInput } from "../../shared/ui/FormControls"; +import { Switch } from "../../shared/ui/Switch"; +import { TitledCard } from "../../shared/ui/TitledCard"; +import { useMessage } from "../../shared/ui/message"; +import styles from "./ExternalApiSettingsCard.module.scss"; + +export function ExternalApiSettingsCard({ servicePort }: { servicePort: number }) { + const message = useMessage(); + const [saved, setSaved] = useState(null); + const [draft, setDraft] = useState({ enabled: false, api_key: "" }); + const [saving, setSaving] = useState(false); + + useEffect(() => { + void api.externalApiSettings().then((settings) => { + setSaved(settings); + setDraft(settings); + }).catch((cause: unknown) => message(cause instanceof Error ? cause.message : String(cause))); + }, [message]); + + const save = async () => { + try { + setSaving(true); + const settings = await api.setExternalApiSettings(draft); + setSaved(settings); + setDraft(settings); + message(t("外部 API 设置已保存")); + } catch (cause) { + message(cause instanceof Error ? cause.message : String(cause)); + } finally { + setSaving(false); + } + }; + + const address = `http://127.0.0.1:${servicePort}/byok/v1`; + const changed = saved && (saved.enabled !== draft.enabled || saved.api_key !== draft.api_key); + + return void save()}> + {saving ? t("保存中…") : t("保存")} + }> +
+
+
+ {t("开启外部 API")} + {t("允许本机应用使用密钥调用已配置的模型,包括插件模型。")} +
+ setDraft((current) => ({ ...current, enabled }))} /> +
+ + setDraft((current) => ({ ...current, api_key: event.target.value }))} /> + +
+ {t("基础地址")} + {address} +
+
+
; +} diff --git a/apps/desktop/src/features/settings/SettingsPage.tsx b/apps/desktop/src/features/settings/SettingsPage.tsx index 212c2fd..20b0fd9 100644 --- a/apps/desktop/src/features/settings/SettingsPage.tsx +++ b/apps/desktop/src/features/settings/SettingsPage.tsx @@ -4,6 +4,7 @@ import { PageContent } from "../../shell/layout/PageContent"; import { LegacyModelImport } from "../models/LegacyModelImport"; import { AppLifecycleSettingsCard } from "./AppLifecycleSettingsCard"; import { CommitSettingsCard } from "./CommitSettingsCard"; +import { ExternalApiSettingsCard } from "./ExternalApiSettingsCard"; import { PricingSettingsCard } from "./PricingSettingsCard"; import { ProxySettingsCard } from "./ProxySettingsCard"; import { TabSettingsCard } from "./TabSettingsCard"; @@ -224,6 +225,7 @@ export function SettingsPage() { + void saveProxy()} /> void saveTab()} /> diff --git a/apps/desktop/src/i18n/generated/catalog.json b/apps/desktop/src/i18n/generated/catalog.json index be12b12..e0edae8 100644 --- a/apps/desktop/src/i18n/generated/catalog.json +++ b/apps/desktop/src/i18n/generated/catalog.json @@ -131,7 +131,7 @@ }, { "file": "features/settings/SettingsPage.tsx", - "line": 182, + "line": 183, "column": 82 }, { @@ -372,7 +372,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 99, + "line": 100, "column": 38 } ] @@ -401,7 +401,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 277, + "line": 279, "column": 26 } ] @@ -537,6 +537,23 @@ } ] }, + "13a494f82d7f26df": { + "source": "开启外部 API", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 46, + "column": 20 + }, + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 49, + "column": 24 + } + ] + }, "13a9ac7a68c5fd96": { "source": "CA 仅保存在本机,用于安全解析 Cursor 的 HTTPS 请求。", "kind": "text", @@ -764,7 +781,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 281, + "line": 283, "column": 131 } ] @@ -790,17 +807,17 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 69, + "line": 70, "column": 42 }, { "file": "features/settings/SettingsPage.tsx", - "line": 187, + "line": 188, "column": 22 }, { "file": "features/settings/SettingsPage.tsx", - "line": 214, + "line": 215, "column": 58 } ] @@ -899,7 +916,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 254, + "line": 256, "column": 43 } ] @@ -923,7 +940,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 124, + "line": 125, "column": 15 } ] @@ -947,7 +964,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 304, + "line": 306, "column": 24 } ] @@ -1135,12 +1152,12 @@ }, { "file": "features/settings/SettingsPage.tsx", - "line": 178, + "line": 179, "column": 81 }, { "file": "features/settings/SettingsPage.tsx", - "line": 296, + "line": 298, "column": 22 }, { @@ -1246,7 +1263,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 148, + "line": 149, "column": 15 } ] @@ -1396,7 +1413,7 @@ }, { "file": "features/settings/SettingsPage.tsx", - "line": 239, + "line": 241, "column": 27 } ] @@ -1469,7 +1486,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 306, + "line": 308, "column": 42 } ] @@ -1534,7 +1551,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 288, + "line": 290, "column": 14 } ] @@ -1563,7 +1580,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 239, + "line": 241, "column": 39 } ] @@ -1630,12 +1647,12 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 246, + "line": 248, "column": 22 }, { "file": "features/settings/SettingsPage.tsx", - "line": 252, + "line": 254, "column": 26 } ] @@ -1717,7 +1734,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 155, + "line": 156, "column": 66 } ] @@ -1729,7 +1746,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 272, + "line": 274, "column": 52 } ] @@ -2179,7 +2196,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 176, + "line": 177, "column": 26 } ] @@ -2302,7 +2319,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 232, + "line": 234, "column": 78 } ] @@ -2348,7 +2365,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 236, + "line": 238, "column": 21 } ] @@ -2446,6 +2463,18 @@ } ] }, + "5f600b307b4eb0fb": { + "source": "API 密钥", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 52, + "column": 25 + } + ] + }, "5f8d556a9c47da3c": { "source": "已关闭开机启动", "kind": "text", @@ -2649,7 +2678,7 @@ "refs": [ { "file": "shared/api.ts", - "line": 527, + "line": 532, "column": 43 } ] @@ -2776,7 +2805,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 99, + "line": 100, "column": 55 } ] @@ -2940,7 +2969,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 166, + "line": 167, "column": 16 } ] @@ -2964,7 +2993,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 264, + "line": 266, "column": 26 } ] @@ -2988,7 +3017,7 @@ "refs": [ { "file": "shared/api.ts", - "line": 522, + "line": 527, "column": 43 } ] @@ -3012,7 +3041,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 297, + "line": 299, "column": 23 } ] @@ -3029,6 +3058,18 @@ } ] }, + "7cc1cf151d967bda": { + "source": "允许本机应用使用密钥调用已配置的模型,包括插件模型。", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 47, + "column": 19 + } + ] + }, "7cea2f3c46565d29": { "source": "OpenAI 额外参数", "kind": "text", @@ -3067,7 +3108,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 142, + "line": 143, "column": 83 } ] @@ -3084,6 +3125,18 @@ } ] }, + "7ed9f96f513813a4": { + "source": "外部 API 设置已保存", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 29, + "column": 15 + } + ] + }, "7f3c8312816fe26a": { "source": "刷新中…", "kind": "text", @@ -3142,7 +3195,7 @@ "refs": [ { "file": "shared/api.ts", - "line": 456, + "line": 461, "column": 21 } ] @@ -3247,6 +3300,18 @@ } ] }, + "857d282a0213d785": { + "source": "外部 API", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 40, + "column": 29 + } + ] + }, "864597982c308d72": { "source": "已开启静默启动", "kind": "text", @@ -3387,7 +3452,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 235, + "line": 237, "column": 22 } ] @@ -3514,7 +3579,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 158, + "line": 159, "column": 7 } ] @@ -3670,6 +3735,18 @@ } ] }, + "99ac861433209a53": { + "source": "开启后,所有外部请求都必须提供此密钥。", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 52, + "column": 44 + } + ] + }, "9a84733cc9ab1706": { "source": "资源详情", "kind": "text", @@ -3838,7 +3915,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 243, + "line": 245, "column": 26 } ] @@ -3874,7 +3951,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 62, + "line": 63, "column": 34 } ] @@ -3971,6 +4048,11 @@ "line": 213, "column": 22 }, + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 41, + "column": 27 + }, { "file": "features/settings/PricingSettingsCard.tsx", "line": 92, @@ -3983,7 +4065,7 @@ }, { "file": "features/settings/SettingsPage.tsx", - "line": 179, + "line": 180, "column": 133 }, { @@ -4377,6 +4459,18 @@ } ] }, + "b1968db970026087": { + "source": "基础地址", + "kind": "text", + "placeholders": [], + "refs": [ + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 57, + "column": 18 + } + ] + }, "b254ff315d861346": { "source": "请重试初始化", "kind": "text", @@ -4495,7 +4589,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 155, + "line": 156, "column": 45 } ] @@ -4521,7 +4615,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 280, + "line": 282, "column": 22 } ] @@ -4679,7 +4773,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 281, + "line": 283, "column": 31 } ] @@ -4691,7 +4785,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 272, + "line": 274, "column": 40 } ] @@ -4977,12 +5071,12 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 164, + "line": 165, "column": 22 }, { "file": "features/settings/SettingsPage.tsx", - "line": 170, + "line": 171, "column": 20 } ] @@ -4994,7 +5088,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 201, + "line": 202, "column": 21 } ] @@ -5030,7 +5124,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 157, + "line": 158, "column": 7 } ] @@ -5398,7 +5492,7 @@ }, { "file": "features/settings/SettingsPage.tsx", - "line": 316, + "line": 318, "column": 30 }, { @@ -5662,7 +5756,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 75, + "line": 76, "column": 17 } ] @@ -5756,7 +5850,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 161, + "line": 162, "column": 26 } ] @@ -5871,7 +5965,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 220, + "line": 221, "column": 16 } ] @@ -6032,7 +6126,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 247, + "line": 249, "column": 21 } ] @@ -6118,7 +6212,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 188, + "line": 189, "column": 21 } ] @@ -6316,7 +6410,7 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 307, + "line": 309, "column": 38 } ] @@ -6331,6 +6425,11 @@ "line": 169, "column": 24 }, + { + "file": "features/settings/ExternalApiSettingsCard.tsx", + "line": 41, + "column": 15 + }, { "file": "features/settings/PricingSettingsCard.tsx", "line": 92, @@ -6343,7 +6442,7 @@ }, { "file": "features/settings/SettingsPage.tsx", - "line": 179, + "line": 180, "column": 121 }, { @@ -6372,17 +6471,17 @@ "refs": [ { "file": "features/settings/SettingsPage.tsx", - "line": 70, + "line": 71, "column": 46 }, { "file": "features/settings/SettingsPage.tsx", - "line": 200, + "line": 201, "column": 22 }, { "file": "features/settings/SettingsPage.tsx", - "line": 215, + "line": 216, "column": 58 } ] diff --git a/apps/desktop/src/i18n/locales/en-US.json b/apps/desktop/src/i18n/locales/en-US.json index 5f00647..d09d8d9 100644 --- a/apps/desktop/src/i18n/locales/en-US.json +++ b/apps/desktop/src/i18n/locales/en-US.json @@ -34,6 +34,7 @@ "124be3f86f197802": "Token usage", "12ae77e6202d063e": "Custom Headers", "133340e53175128a": "Test all", + "13a494f82d7f26df": "Enable external API", "13a9ac7a68c5fd96": "The CA is stored only on this device and is used to securely inspect Cursor HTTPS requests.", "13b61c5f697b6700": "Cache hit rate", "146da2e2a991493e": "Fetching…", @@ -163,6 +164,7 @@ "5c55a67935af8f45": "All", "5c62e36c152dfc7c": "Plugin runtime initialized", "5d59857bf039cac9": "Cursor Assistant v{version}", + "5f600b307b4eb0fb": "API key", "5f8d556a9c47da3c": "Launch at login disabled", "5f9acfb945229062": "Are you sure you no longer want to see this ad?", "5fd2ec5a6e9b654c": "Total: {cost}", @@ -205,10 +207,12 @@ "7a3cec4ca715de80": "Call statistics", "7ba2d6728fe2531b": "Confirm clear", "7c10d97162c96dbd": "Validating the plugin runtime", + "7cc1cf151d967bda": "Allow local apps to call configured models, including plugin models, with an API key.", "7cea2f3c46565d29": "OpenAI extra parameters", "7d9f043f8f7ab45c": "Version {version} is available in Settings", "7e0891860c9e6374": "TAB service address is required", "7e7df68f2a82e09e": "Importing the same configuration again will not create duplicate models. Existing models are skipped automatically.", + "7ed9f96f513813a4": "External API settings saved", "7f3c8312816fe26a": "Refreshing…", "7f68ebad19ba6bcd": "Check for updates", "802b0faf0ceb513e": "{label}: {percent}% left", @@ -221,6 +225,7 @@ "842b9f11cdd96bda": "Launch at login", "843ac7e15a5047a7": "Confirm legacy model configuration import", "844b8cc8dff7c1d8": "Default", + "857d282a0213d785": "External API", "864597982c308d72": "Silent start enabled", "86de7c4ee8fa7689": "Sync models", "8716e1344b0daddb": "Cursor official", @@ -252,6 +257,7 @@ "9845c165151daee3": "Used", "9850ed41a5bfbb0c": "{count} selected", "997ec8201c2adeda": "Open terminal to install CA", + "99ac861433209a53": "All external requests must provide this key when enabled.", "9a84733cc9ab1706": "Resource details", "9b1b7ed518ee401d": "This will open the tutorial in your system browser. Continue?", "9b9bc9cd7c76406f": "Open authorization page", @@ -298,6 +304,7 @@ "affb73206cfa035d": "Expires: {time}", "b06325c5660f0c29": "Direct", "b16c3b2ecedd6fe1": "Cursor integration is active. Add a model configuration to use a BYOK model.", + "b1968db970026087": "Base URL", "b254ff315d861346": "Try initializing again", "b2617bf9ae663752": "Group settings", "b4411558b932266f": "Provider type", diff --git a/apps/desktop/src/i18n/locales/pt-BR.json b/apps/desktop/src/i18n/locales/pt-BR.json index 1c9a98a..1ed853a 100644 --- a/apps/desktop/src/i18n/locales/pt-BR.json +++ b/apps/desktop/src/i18n/locales/pt-BR.json @@ -34,6 +34,7 @@ "124be3f86f197802": "Uso de tokens", "12ae77e6202d063e": "Cabeçalhos personalizados", "133340e53175128a": "Testar tudo", + "13a494f82d7f26df": "Ativar API externa", "13a9ac7a68c5fd96": "A CA é armazenada apenas neste dispositivo e é usada para inspecionar com segurança as requisições HTTPS do Cursor.", "13b61c5f697b6700": "Taxa de acerto de cache", "146da2e2a991493e": "Obtendo…", @@ -163,6 +164,7 @@ "5c55a67935af8f45": "Todos", "5c62e36c152dfc7c": "Runtime de plugins inicializado com sucesso", "5d59857bf039cac9": "Assistente Cursor v{version}", + "5f600b307b4eb0fb": "Chave da API", "5f8d556a9c47da3c": "Inicialização com o sistema desativada", "5f9acfb945229062": "Tem certeza de que não deseja mais ver este anúncio?", "5fd2ec5a6e9b654c": "Total: {cost}", @@ -205,10 +207,12 @@ "7a3cec4ca715de80": "Estatísticas de chamadas", "7ba2d6728fe2531b": "Confirmar limpeza", "7c10d97162c96dbd": "Validando o runtime de plugins", + "7cc1cf151d967bda": "Permite que aplicativos locais acessem os modelos configurados, incluindo modelos de plugins, com uma chave de API.", "7cea2f3c46565d29": "Parâmetros adicionais da OpenAI", "7d9f043f8f7ab45c": "Versão {version} disponível nas Configurações", "7e0891860c9e6374": "O endereço do serviço TAB é obrigatório", "7e7df68f2a82e09e": "Importar a mesma configuração novamente não criará modelos duplicados. Modelos existentes são ignorados automaticamente.", + "7ed9f96f513813a4": "Configurações da API externa salvas", "7f3c8312816fe26a": "Atualizando…", "7f68ebad19ba6bcd": "Verificar atualizações", "802b0faf0ceb513e": "{label}: {percent}% restante", @@ -221,6 +225,7 @@ "842b9f11cdd96bda": "Iniciar com o sistema", "843ac7e15a5047a7": "Confirmar importação da configuração legada de modelos", "844b8cc8dff7c1d8": "Padrão", + "857d282a0213d785": "API externa", "864597982c308d72": "Inicialização silenciosa ativada", "86de7c4ee8fa7689": "Sincronizar modelos", "8716e1344b0daddb": "Cursor Oficial", @@ -252,6 +257,7 @@ "9845c165151daee3": "Utilizado", "9850ed41a5bfbb0c": "{count} selecionados", "997ec8201c2adeda": "Abrir terminal para instalar CA", + "99ac861433209a53": "Todas as solicitações externas devem fornecer esta chave quando a API estiver ativada.", "9a84733cc9ab1706": "Detalhes do recurso", "9b1b7ed518ee401d": "O tutorial de uso será aberto no navegador padrão. Deseja continuar?", "9b9bc9cd7c76406f": "Abrir página de autorização", @@ -298,6 +304,7 @@ "affb73206cfa035d": "Expira em: {time}", "b06325c5660f0c29": "Conexão direta", "b16c3b2ecedd6fe1": "A integração com o Cursor está ativa. Adicione uma configuração de modelo para usar modelos BYOK.", + "b1968db970026087": "URL base", "b254ff315d861346": "Tente inicializar novamente", "b2617bf9ae663752": "Configurações do grupo", "b4411558b932266f": "Tipo de provedor", diff --git a/apps/desktop/src/i18n/locales/zh-CN.json b/apps/desktop/src/i18n/locales/zh-CN.json index 36c9f67..2375022 100644 --- a/apps/desktop/src/i18n/locales/zh-CN.json +++ b/apps/desktop/src/i18n/locales/zh-CN.json @@ -34,6 +34,7 @@ "124be3f86f197802": "Token 消耗", "12ae77e6202d063e": "自定义 Headers", "133340e53175128a": "一键测试", + "13a494f82d7f26df": "开启外部 API", "13a9ac7a68c5fd96": "CA 仅保存在本机,用于安全解析 Cursor 的 HTTPS 请求。", "13b61c5f697b6700": "缓存命中率", "146da2e2a991493e": "获取中…", @@ -163,6 +164,7 @@ "5c55a67935af8f45": "全部", "5c62e36c152dfc7c": "插件运行时初始化完成", "5d59857bf039cac9": "Cursor 助手 v{version}", + "5f600b307b4eb0fb": "API 密钥", "5f8d556a9c47da3c": "已关闭开机启动", "5f9acfb945229062": "你确认不想再看到此广告吗?", "5fd2ec5a6e9b654c": "合计:{cost}", @@ -205,10 +207,12 @@ "7a3cec4ca715de80": "调用统计", "7ba2d6728fe2531b": "确认清理", "7c10d97162c96dbd": "正在验证插件运行时", + "7cc1cf151d967bda": "允许本机应用使用密钥调用已配置的模型,包括插件模型。", "7cea2f3c46565d29": "OpenAI 额外参数", "7d9f043f8f7ab45c": "发现新版本 {version},可在设置中安装", "7e0891860c9e6374": "TAB 服务地址不能为空", "7e7df68f2a82e09e": "重复导入相同配置不会创建重复模型;已经存在的模型会自动跳过。", + "7ed9f96f513813a4": "外部 API 设置已保存", "7f3c8312816fe26a": "刷新中…", "7f68ebad19ba6bcd": "检查更新", "802b0faf0ceb513e": "{label} 剩余 {percent}%", @@ -221,6 +225,7 @@ "842b9f11cdd96bda": "开机启动", "843ac7e15a5047a7": "确认导入旧版模型配置", "844b8cc8dff7c1d8": "默认", + "857d282a0213d785": "外部 API", "864597982c308d72": "已开启静默启动", "86de7c4ee8fa7689": "同步模型", "8716e1344b0daddb": "Cursor 官方", @@ -252,6 +257,7 @@ "9845c165151daee3": "已使用", "9850ed41a5bfbb0c": "已选 {count} 项", "997ec8201c2adeda": "打开终端安装 CA", + "99ac861433209a53": "开启后,所有外部请求都必须提供此密钥。", "9a84733cc9ab1706": "资源详情", "9b1b7ed518ee401d": "将在系统浏览器中打开使用教程,是否继续?", "9b9bc9cd7c76406f": "打开授权网页", @@ -298,6 +304,7 @@ "affb73206cfa035d": "到期时间:{time}", "b06325c5660f0c29": "直连", "b16c3b2ecedd6fe1": "Cursor 接管已生效;添加模型配置后即可使用 BYOK 模型。", + "b1968db970026087": "基础地址", "b254ff315d861346": "请重试初始化", "b2617bf9ae663752": "分组设置", "b4411558b932266f": "上游类型", diff --git a/apps/desktop/src/shared/api.ts b/apps/desktop/src/shared/api.ts index 0c70886..bb1f36b 100644 --- a/apps/desktop/src/shared/api.ts +++ b/apps/desktop/src/shared/api.ts @@ -113,6 +113,11 @@ export interface PortSettings { service_port: number; } +export interface ExternalApiSettings { + enabled: boolean; + api_key: string; +} + export interface StatisticsStorage { call_count: number; trace_count: number; @@ -541,6 +546,8 @@ export const api = { setObservability: (detailed: boolean) => request<{ detailed: boolean }>("/settings/observability", { method: "PUT", body: JSON.stringify({ detailed }) }), ports: () => request("/settings/ports"), setPorts: (settings: PortSettings) => request("/settings/ports", { method: "PUT", body: JSON.stringify(settings) }), + externalApiSettings: () => request("/settings/external-api"), + setExternalApiSettings: (settings: ExternalApiSettings) => request("/settings/external-api", { method: "PUT", body: JSON.stringify(settings) }), statisticsStorage: () => request("/settings/storage/statistics"), clearStatisticsStorage: (scope: StatisticsStorageScope) => request("/settings/storage/statistics", { method: "DELETE", body: JSON.stringify({ scope }) }), proxySettings: () => request("/settings/proxy"), diff --git a/server/src/api/byok/direct.rs b/server/src/api/byok/direct.rs new file mode 100644 index 0000000..b0e3d43 --- /dev/null +++ b/server/src/api/byok/direct.rs @@ -0,0 +1,379 @@ +//! Forwards matching HTTP protocols without projecting provider messages. +use std::time::Duration; + +use async_stream::try_stream; +use axum::{ + body::Body, + http::{header, HeaderMap, StatusCode}, + response::Response, +}; +use bytes::Bytes; +use eventsource_stream::Eventsource; +use futures_util::StreamExt; +use serde_json::Value; +use tokio_stream::wrappers::ReceiverStream; + +use crate::{ + model::{ModelConfig, NewLlmCall, Usage}, + network::NetworkClients, + provider::{custom_headers, CallRecorder, FinishReason}, + store::Store, + Error, Result, +}; + +use super::protocol::Protocol; + +#[derive(Clone)] +pub struct NativeForwarder { + store: Store, + clients: NetworkClients, + request_timeout: Duration, + stream_idle_timeout: Duration, +} + +impl NativeForwarder { + pub fn new( + store: Store, + clients: NetworkClients, + request_timeout: Duration, + stream_idle_timeout: Duration, + ) -> Self { + Self { + store, + clients, + request_timeout, + stream_idle_timeout, + } + } + + pub(super) async fn forward( + &self, + protocol: Protocol, + incoming_headers: &HeaderMap, + mut body: Value, + model: &ModelConfig, + ) -> Result { + let call_id = format!("external-api:{}", uuid::Uuid::new_v4()); + let request_url = model.request_url()?; + let new_call = NewLlmCall { + call_id: call_id.clone(), + run_id: call_id.clone(), + conversation_id: call_id.clone(), + provider_call_index: 0, + model_hash: model.model_hash.clone(), + provider_type: model.provider_type(), + provider_url: request_url.clone(), + request_type: model.provider_type(), + request_url: request_url.clone(), + model_id: model.model_id.clone(), + display_name: model.display_name.clone(), + reasoning_effort: model.reasoning_effort.clone(), + fast: false, + message_count: body + .get("messages") + .or_else(|| body.get("input")) + .and_then(Value::as_array) + .map_or(1, Vec::len), + tool_count: body + .get("tools") + .and_then(Value::as_array) + .map_or(0, Vec::len), + detailed: false, + }; + body["model"] = Value::String(model.model_id.clone()); + let mut upstream_headers = HeaderMap::new(); + for (name, value) in incoming_headers { + if !matches!( + name.as_str(), + "authorization" + | "x-api-key" + | "host" + | "content-length" + | "content-type" + | "connection" + | "transfer-encoding" + | "accept-encoding" + | "cookie" + | "proxy-authorization" + | "proxy-authenticate" + | "te" + | "trailer" + | "upgrade" + ) { + upstream_headers.append(name, value.clone()); + } + } + if model.custom_headers_enabled { + upstream_headers.extend(custom_headers(&model.custom_headers, &call_id)?); + } + match protocol { + Protocol::Messages => { + upstream_headers.insert( + "x-api-key", + model.api_key.parse().map_err(|error| { + Error::Config(format!("invalid API key header: {error}")) + })?, + ); + if !upstream_headers.contains_key("anthropic-version") { + upstream_headers.insert("anthropic-version", "2023-06-01".parse().unwrap()); + } + } + Protocol::Chat | Protocol::Responses => { + upstream_headers.insert( + header::AUTHORIZATION, + format!("Bearer {}", model.api_key) + .parse() + .map_err(|error| { + Error::Config(format!("invalid API key header: {error}")) + })?, + ); + } + } + let logged_headers = upstream_headers + .iter() + .filter(|(name, _)| !crate::model::is_sensitive_header(name.as_str())) + .filter_map(|(name, value)| { + value + .to_str() + .ok() + .map(|value| (name.as_str().to_owned(), Value::String(value.to_owned()))) + }) + .collect::>(); + let client = self.clients.provider_client(self.request_timeout).await?; + let recorder = CallRecorder::start(self.store.clone(), new_call).await?; + recorder + .request(Value::Object(logged_headers), &body) + .await?; + let response = match client + .post(&request_url) + .headers(upstream_headers) + .json(&body) + .send() + .await + { + Ok(response) => response, + Err(error) => { + let error = Error::Http(error); + recorder.failed(&error).await?; + return Err(error); + } + }; + let status = response.status(); + recorder.response_headers(status.as_u16()).await?; + let headers = response.headers().clone(); + let stream = body.get("stream").and_then(Value::as_bool).unwrap_or(false); + if stream && status.is_success() { + Ok(self.stream_response(protocol, response, recorder, status, &headers)) + } else { + let bytes = match response.bytes().await { + Ok(bytes) => bytes, + Err(error) => { + let error = Error::Http(error); + recorder.failed(&error).await?; + return Err(error); + } + }; + recorder.response_chunk(&bytes).await?; + if status.is_success() { + if let Ok(value) = serde_json::from_slice::(&bytes) { + let summary = NativeSummary::from_value(protocol, &value); + if let Some(usage) = summary.usage { + recorder.usage(usage).await?; + } + if let Some(message) = summary.failure { + recorder.failed(&Error::Provider(message)).await?; + } else { + recorder + .completed(summary.finish.unwrap_or(FinishReason::Stop)) + .await?; + } + } else { + recorder.completed(FinishReason::Stop).await?; + } + } else { + recorder + .failed(&Error::Provider(format!("upstream returned {status}"))) + .await?; + } + Ok(build_response(status, &headers, Body::from(bytes))) + } + } + + fn stream_response( + &self, + protocol: Protocol, + response: reqwest::Response, + recorder: CallRecorder, + status: StatusCode, + headers: &HeaderMap, + ) -> Response { + let idle_timeout = self.stream_idle_timeout; + let output = try_stream! { + let _cancel = recorder.cancel_on_drop(); + let (sender, receiver) = tokio::sync::mpsc::channel::(8); + let observer = tokio::spawn(observe_events(protocol, receiver)); + let mut upstream = response.bytes_stream(); + loop { + let next = tokio::time::timeout(idle_timeout, upstream.next()).await; + let chunk = match next { + Ok(Some(Ok(chunk))) => chunk, + Ok(None) => break, + Ok(Some(Err(error))) => { + let error = Error::Http(error); + recorder.failed(&error).await?; + Err::(error)? + } + Err(_) => { + let error = Error::Provider("native upstream stream idle timeout".into()); + recorder.failed(&error).await?; + Err::(error)? + } + }; + recorder.response_chunk(&chunk).await?; + let _ = sender.send(chunk.clone()).await; + yield chunk; + } + drop(sender); + if let Ok(summary) = observer.await { + if let Some(usage) = summary.usage { recorder.usage(usage).await?; } + if let Some(message) = summary.failure { + recorder.failed(&Error::Provider(message)).await?; + } else { + recorder.completed(summary.finish.unwrap_or(FinishReason::Stop)).await?; + } + } else { + recorder.completed(FinishReason::Stop).await?; + } + }; + let output = output.map(|result: Result| result); + build_response(status, headers, Body::from_stream(output)) + } +} + +fn build_response(status: StatusCode, headers: &HeaderMap, body: Body) -> Response { + let mut response = Response::new(body); + *response.status_mut() = status; + for (name, value) in headers { + if !matches!( + name.as_str(), + "content-length" | "content-encoding" | "transfer-encoding" | "connection" + ) { + response.headers_mut().append(name, value.clone()); + } + } + response +} + +#[derive(Default)] +struct NativeSummary { + usage: Option, + finish: Option, + failure: Option, +} + +impl NativeSummary { + fn from_value(protocol: Protocol, value: &Value) -> Self { + let mut summary = Self::default(); + summary.observe(protocol, value); + summary + } + + fn observe(&mut self, protocol: Protocol, value: &Value) { + if matches!( + value.get("type").and_then(Value::as_str), + Some("error" | "response.failed") + ) || value.get("error").is_some_and(|error| !error.is_null()) + { + self.failure = Some( + value + .pointer("/error/message") + .or_else(|| value.pointer("/response/error/message")) + .or_else(|| value.get("message")) + .and_then(Value::as_str) + .unwrap_or("upstream error event") + .to_owned(), + ); + } + let usage = match protocol { + Protocol::Chat => value.get("usage"), + Protocol::Responses => value + .pointer("/response/usage") + .or_else(|| value.get("usage")), + Protocol::Messages => value + .pointer("/message/usage") + .or_else(|| value.get("usage")), + } + .filter(|usage| !usage.is_null()) + .map(|usage| crate::provider::native_usage(protocol.provider_type(), usage)); + if let Some(usage) = usage { + let current = self.usage.get_or_insert_default(); + for (target, incoming) in [ + (&mut current.input_tokens, usage.input_tokens), + ( + &mut current.context_input_tokens, + usage.context_input_tokens, + ), + (&mut current.output_tokens, usage.output_tokens), + (&mut current.total_tokens, usage.total_tokens), + (&mut current.cache_read_tokens, usage.cache_read_tokens), + (&mut current.cache_write_tokens, usage.cache_write_tokens), + (&mut current.reasoning_tokens, usage.reasoning_tokens), + ] { + if incoming.is_some() { + *target = incoming; + } + } + } + let reason = match protocol { + Protocol::Chat => value + .pointer("/choices/0/finish_reason") + .and_then(Value::as_str), + Protocol::Responses => value + .pointer("/response/incomplete_details/reason") + .and_then(Value::as_str), + Protocol::Messages => value + .pointer("/delta/stop_reason") + .or_else(|| value.get("stop_reason")) + .and_then(Value::as_str), + }; + self.finish = match reason { + Some("tool_calls" | "tool_use") => Some(FinishReason::ToolUse), + Some("length" | "max_tokens" | "max_output_tokens") => Some(FinishReason::Length), + Some(_) => Some(FinishReason::Stop), + None => self.finish, + }; + if matches!(protocol, Protocol::Responses) + && (value.pointer("/item/type").and_then(Value::as_str) == Some("function_call") + || value + .pointer("/response/output") + .and_then(Value::as_array) + .is_some_and(|items| { + items.iter().any(|item| { + item.get("type").and_then(Value::as_str) == Some("function_call") + }) + })) + { + self.finish = Some(FinishReason::ToolUse); + } + } +} + +async fn observe_events( + protocol: Protocol, + receiver: tokio::sync::mpsc::Receiver, +) -> NativeSummary { + let source = ReceiverStream::new(receiver) + .map(Ok::<_, Error>) + .eventsource(); + futures_util::pin_mut!(source); + let mut summary = NativeSummary::default(); + while let Some(event) = source.next().await { + let Ok(event) = event else { + break; + }; + if let Ok(value) = serde_json::from_str::(&event.data) { + summary.observe(protocol, &value); + } + } + summary +} diff --git a/server/src/api/byok/mod.rs b/server/src/api/byok/mod.rs new file mode 100644 index 0000000..6a376fd --- /dev/null +++ b/server/src/api/byok/mod.rs @@ -0,0 +1,201 @@ +//! External OpenAI and Anthropic compatible API. +mod direct; +mod models; +mod output; +mod protocol; + +use std::sync::Arc; + +use axum::{ + extract::State, + http::{header, HeaderMap, StatusCode}, + response::{IntoResponse, Response}, + routing::{get, post}, + Json, Router, +}; +use serde_json::{json, Value}; +use tokio_util::sync::CancellationToken; + +use crate::{ + model::ModelInvocation, plugin::PluginRegistry, provider::Provider, store::Store, Error, +}; + +use self::protocol::Protocol; +pub use direct::NativeForwarder; + +#[derive(Clone)] +struct ApiState { + store: Store, + plugins: PluginRegistry, + provider: Arc, + native: Option, +} + +pub fn router( + store: Store, + plugins: PluginRegistry, + provider: Arc, + native: Option, +) -> Router { + Router::new() + .route("/byok/v1/models", get(list_models)) + .route("/byok/v1/chat/completions", post(chat)) + .route("/byok/v1/responses", post(responses)) + .route("/byok/v1/messages", post(messages)) + .with_state(ApiState { + store, + plugins, + provider, + native, + }) +} + +fn api_error(status: StatusCode, message: impl std::fmt::Display) -> Response { + (status, Json(json!({"error":{"message":message.to_string(),"type":if status.is_client_error() { "invalid_request_error" } else { "server_error" }}}))).into_response() +} + +async fn authorized( + state: &ApiState, + headers: &HeaderMap, +) -> std::result::Result<(), Box> { + let settings = state + .store + .external_api_settings() + .await + .map_err(|error| Box::new(api_error(StatusCode::INTERNAL_SERVER_ERROR, error)))?; + if !settings.enabled { + return Err(Box::new(api_error( + StatusCode::FORBIDDEN, + "external API is disabled", + ))); + } + let supplied = headers + .get(header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.strip_prefix("Bearer ")) + .or_else(|| { + headers + .get("x-api-key") + .and_then(|value| value.to_str().ok()) + }) + .unwrap_or_default(); + if supplied.is_empty() || !constant_time_eq(supplied.as_bytes(), settings.api_key.as_bytes()) { + return Err(Box::new(api_error( + StatusCode::UNAUTHORIZED, + "invalid API key", + ))); + } + Ok(()) +} + +fn constant_time_eq(left: &[u8], right: &[u8]) -> bool { + if left.len() != right.len() { + return false; + } + let difference = left + .iter() + .zip(right) + .fold(0_u8, |difference, (left, right)| { + difference | (left ^ right) + }); + difference == 0 +} + +async fn list_models(State(state): State, headers: HeaderMap) -> Response { + if let Err(error) = authorized(&state, &headers).await { + return *error; + } + match models::list(&state.store, &state.plugins).await { + Ok(models) => Json(models::response(&models)).into_response(), + Err(error) => api_error(StatusCode::CONFLICT, error), + } +} + +async fn chat( + State(state): State, + headers: HeaderMap, + Json(body): Json, +) -> Response { + generate(state, headers, body, Protocol::Chat).await +} + +async fn responses( + State(state): State, + headers: HeaderMap, + Json(body): Json, +) -> Response { + generate(state, headers, body, Protocol::Responses).await +} + +async fn messages( + State(state): State, + headers: HeaderMap, + Json(body): Json, +) -> Response { + generate(state, headers, body, Protocol::Messages).await +} + +async fn generate( + state: ApiState, + headers: HeaderMap, + body: Value, + protocol: Protocol, +) -> Response { + if let Err(error) = authorized(&state, &headers).await { + return *error; + } + let Some(public_model_id) = body + .get("model") + .and_then(Value::as_str) + .filter(|id| !id.is_empty()) + else { + return api_error(StatusCode::BAD_REQUEST, "model is required"); + }; + let models = match models::list(&state.store, &state.plugins).await { + Ok(models) => models, + Err(error) => return api_error(StatusCode::CONFLICT, error), + }; + let Some(model) = models + .iter() + .find(|model| model.public_id == public_model_id) + else { + return api_error(StatusCode::NOT_FOUND, "model not found"); + }; + if let Some(native) = &state.native { + match state.store.model(&model.internal_id).await { + Ok(Some(config)) if protocol.provider_type() == config.provider_type() => { + return match native.forward(protocol, &headers, body, &config).await { + Ok(response) => response, + Err(error) => api_error(StatusCode::BAD_GATEWAY, error), + }; + } + Ok(_) => {} + Err(error) => return api_error(StatusCode::INTERNAL_SERVER_ERROR, error), + } + } + let parsed = match protocol::parse(protocol, &body) { + Ok(parsed) => parsed, + Err(error) => return api_error(StatusCode::BAD_REQUEST, error), + }; + let id = format!("external-api:{}", uuid::Uuid::new_v4()); + let mut request = parsed.request; + request.model.model_id = model.internal_id.clone(); + let invocation = ModelInvocation { + call_id: id.clone(), + run_id: id.clone(), + conversation_id: id.clone(), + provider_call_index: 0, + request, + }; + let cancellation = CancellationToken::new(); + let stream = state.provider.stream(invocation, cancellation.clone()); + if parsed.stream { + output::streamed(protocol, stream, cancellation, id, parsed.public_model_id) + } else { + match output::complete(protocol, stream, &id, &parsed.public_model_id).await { + Ok(response) => response.into_response(), + Err(Error::Protocol(message)) => api_error(StatusCode::BAD_REQUEST, message), + Err(error) => api_error(StatusCode::BAD_GATEWAY, error), + } + } +} diff --git a/server/src/api/byok/models.rs b/server/src/api/byok/models.rs new file mode 100644 index 0000000..5eb2e07 --- /dev/null +++ b/server/src/api/byok/models.rs @@ -0,0 +1,66 @@ +use std::collections::HashSet; + +use serde_json::{json, Value}; + +use crate::{plugin::PluginRegistry, store::Store, Result}; + +#[derive(Clone)] +pub(super) struct ListedModel { + pub public_id: String, + pub internal_id: String, + pub display_name: String, +} + +pub(super) async fn list(store: &Store, plugins: &PluginRegistry) -> Result> { + let mut models = Vec::new(); + let mut ids = HashSet::new(); + let mut configured = store.models().await?; + configured.sort_by(|left, right| { + left.sort_order + .cmp(&right.sort_order) + .then_with(|| left.display_name.cmp(&right.display_name)) + .then_with(|| left.model_hash.cmp(&right.model_hash)) + }); + for model in configured { + let group = model + .group_name + .as_deref() + .filter(|value| !value.trim().is_empty()) + .map(str::to_owned) + .unwrap_or_else(|| { + url::Url::parse(&model.base_url) + .ok() + .and_then(|url| url.host_str().map(str::to_owned)) + .unwrap_or_else(|| model.base_url.clone()) + }); + let public_id = format!("{group}/{}", model.model_id); + if ids.insert(public_id.clone()) { + models.push(ListedModel { + public_id, + internal_id: model.model_hash, + display_name: model.display_name, + }); + } + } + for model in plugins.configured_models().await { + let public_id = format!( + "plugin:{}:{}/{}", + model.plugin_id, model.provider_id, model.model_id + ); + if ids.insert(public_id.clone()) { + models.push(ListedModel { + public_id, + internal_id: model.id, + display_name: model.display_name, + }); + } + } + Ok(models) +} + +pub(super) fn response(models: &[ListedModel]) -> Value { + json!({"object":"list","data":models.iter().map(|model| json!({ + "id":model.public_id,"object":"model","created":0,"owned_by":"cursor-byok", + "name":model.display_name + })).collect::>()}) +} diff --git a/server/src/api/byok/output.rs b/server/src/api/byok/output.rs new file mode 100644 index 0000000..13c0cbf --- /dev/null +++ b/server/src/api/byok/output.rs @@ -0,0 +1,444 @@ +use std::{ + collections::{BTreeMap, HashSet}, + convert::Infallible, +}; + +use axum::{ + response::{ + sse::{Event, KeepAlive, Sse}, + IntoResponse, Response, + }, + Json, +}; +use futures_util::StreamExt; +use serde_json::{json, Value}; +use tokio_util::sync::CancellationToken; + +use crate::{ + model::Usage, + provider::{FinishReason, ModelEvent, ProviderStream}, + Error, Result, +}; + +use super::protocol::Protocol; + +#[derive(Default)] +struct ToolOutput { + id: String, + name: String, + arguments: String, +} + +#[derive(Default)] +struct Output { + text: String, + thinking: String, + tools: BTreeMap, + usage: Option, + finish: Option, +} + +impl Output { + fn update(&mut self, event: &ModelEvent) { + match event { + ModelEvent::TextDelta(delta) => self.text.push_str(delta), + ModelEvent::ThinkingDelta(delta) => self.thinking.push_str(delta), + ModelEvent::ToolCallStart { + index, + call_id, + name, + } => { + self.tools.insert( + *index, + ToolOutput { + id: call_id.clone(), + name: name.clone(), + arguments: String::new(), + }, + ); + } + ModelEvent::ToolCallArgumentsDelta { index, delta } => { + if let Some(tool) = self.tools.get_mut(index) { + tool.arguments.push_str(delta); + } + } + ModelEvent::Usage(usage) => self.usage = Some(*usage), + ModelEvent::Done(reason) => self.finish = Some(*reason), + _ => {} + } + } + + fn json(&self, protocol: Protocol, id: &str, model: &str) -> Value { + let usage = self.usage_json(protocol); + let finish = match self.finish.unwrap_or(FinishReason::Stop) { + FinishReason::Stop => "stop", + FinishReason::Length => "length", + FinishReason::ToolUse => "tool_calls", + }; + match protocol { + Protocol::Chat => json!({"id":id,"object":"chat.completion","created":0,"model":model, + "choices":[{"index":0,"message":{"role":"assistant","content":self.text, + "tool_calls":self.tools.values().map(|tool| json!({"id":tool.id,"type":"function", + "function":{"name":tool.name,"arguments":tool.arguments}})).collect::>()},"finish_reason":finish}], + "usage":usage}), + Protocol::Responses => { + let mut items = Vec::new(); + if !self.text.is_empty() { + items.push(json!({"id":format!("msg_{id}"),"type":"message","status":"completed", + "role":"assistant","content":[{"type":"output_text","text":self.text,"annotations":[]}]})); + } + items.extend(self.tools.values().map(|tool| json!({"type":"function_call","id":format!("fc_{}",tool.id), + "call_id":tool.id,"name":tool.name,"arguments":tool.arguments,"status":"completed"}))); + json!({"id":id,"object":"response","created_at":0,"status":"completed","model":model, + "output":items,"output_text":self.text, + "usage":usage}) + } + Protocol::Messages => { + let mut content = Vec::new(); + if !self.text.is_empty() { + content.push(json!({"type":"text","text":self.text})); + } + content.extend(self.tools.values().map(|tool| { + json!({"type":"tool_use","id":tool.id,"name":tool.name, + "input":serde_json::from_str::(&tool.arguments).unwrap_or(json!({}))}) + })); + json!({"id":id,"type":"message","role":"assistant","model":model,"content":content, + "stop_reason":if self.tools.is_empty() { match self.finish { Some(FinishReason::Length) => "max_tokens", _ => "end_turn" } } else { "tool_use" }, + "stop_sequence":null,"usage":usage}) + } + } + } + + fn usage_json(&self, protocol: Protocol) -> Value { + let usage = self.usage.unwrap_or_default(); + let input = usage + .context_input_tokens + .or(usage.input_tokens) + .unwrap_or(0); + let output = usage.output_tokens.unwrap_or(0); + match protocol { + Protocol::Chat => { + let mut value = json!({"prompt_tokens":input,"completion_tokens":output, + "total_tokens":usage.total_tokens.unwrap_or(input.saturating_add(output))}); + if let Some(cached) = usage.cache_read_tokens { + value["prompt_tokens_details"] = json!({"cached_tokens":cached}); + } + if let Some(reasoning) = usage.reasoning_tokens { + value["completion_tokens_details"] = json!({"reasoning_tokens":reasoning}); + } + value + } + Protocol::Responses => { + let mut value = json!({"input_tokens":input,"output_tokens":output, + "total_tokens":usage.total_tokens.unwrap_or(input.saturating_add(output))}); + if let Some(cached) = usage.cache_read_tokens { + value["input_tokens_details"] = json!({"cached_tokens":cached}); + } + if let Some(reasoning) = usage.reasoning_tokens { + value["output_tokens_details"] = json!({"reasoning_tokens":reasoning}); + } + value + } + Protocol::Messages => { + let uncached = + usage + .context_input_tokens + .map_or(usage.input_tokens.unwrap_or(0), |total| { + total + .saturating_sub(usage.cache_read_tokens.unwrap_or(0)) + .saturating_sub(usage.cache_write_tokens.unwrap_or(0)) + }); + let mut value = json!({"input_tokens":uncached,"output_tokens":output}); + if let Some(cached) = usage.cache_read_tokens { + value["cache_read_input_tokens"] = json!(cached); + } + if let Some(written) = usage.cache_write_tokens { + value["cache_creation_input_tokens"] = json!(written); + } + value + } + } + } +} + +pub(super) async fn complete( + protocol: Protocol, + mut stream: ProviderStream, + id: &str, + model: &str, +) -> Result> { + let mut output = Output::default(); + while let Some(event) = stream.next().await { + output.update(&event?); + } + if output.finish.is_none() { + return Err(Error::Provider( + "model stream ended without completion".into(), + )); + } + Ok(Json(output.json(protocol, id, model))) +} + +struct CancelOnDrop(CancellationToken); +impl Drop for CancelOnDrop { + fn drop(&mut self) { + self.0.cancel(); + } +} + +pub(super) fn streamed( + protocol: Protocol, + mut provider: ProviderStream, + cancellation: CancellationToken, + id: String, + model: String, +) -> Response { + let events = async_stream::stream! { + let _cancel = CancelOnDrop(cancellation); + let mut output = Output::default(); + if matches!(protocol, Protocol::Chat) { + yield Ok::(Event::default().data(chat_chunk(&id, &model, json!({"role":"assistant"}), Value::Null).to_string())); + } else if matches!(protocol, Protocol::Responses) { + yield Ok(Event::default().event("response.created").data(json!({"type":"response.created","response":{"id":id,"status":"in_progress","model":model}}).to_string())); + yield Ok(Event::default().event("response.in_progress").data(json!({"type":"response.in_progress","response":{"id":id,"status":"in_progress","model":model}}).to_string())); + } else { + yield Ok(Event::default().event("message_start").data(json!({"type":"message_start","message":{"id":id,"type":"message","role":"assistant","model":model,"content":[],"usage":{"input_tokens":0,"output_tokens":0}}}).to_string())); + } + let mut text_started = false; + let mut text_closed = false; + let mut closed_tools = HashSet::new(); + while let Some(result) = provider.next().await { + match result { + Ok(event) => { + if matches!(event, ModelEvent::TextDelta(_)) && !text_started && !matches!(protocol, Protocol::Chat) { + for (name, value) in stream_events(protocol, &ModelEvent::TextStart, &id, &model, &output, &mut text_started) { + yield Ok(Event::default().event(name).data(value.to_string())); + } + } + if matches!(event, ModelEvent::Done(_)) { + if text_started && !text_closed { + for (name, value) in stream_events(protocol, &ModelEvent::TextEnd, &id, &model, &output, &mut text_started) { + yield Ok(Event::default().event(name).data(value.to_string())); + } + } + for index in output.tools.keys().filter(|index| !closed_tools.contains(*index)).copied().collect::>() { + for (name, value) in stream_events(protocol, &ModelEvent::ToolCallEnd { index }, &id, &model, &output, &mut text_started) { + yield Ok(Event::default().event(name).data(value.to_string())); + } + } + } + output.update(&event); + for (name, value) in stream_events(protocol, &event, &id, &model, &output, &mut text_started) { + yield Ok(Event::default().event(name).data(value.to_string())); + } + if matches!(event, ModelEvent::TextEnd | ModelEvent::Done(_)) { + text_closed = true; + } + if let ModelEvent::ToolCallEnd { index } = event { + closed_tools.insert(index); + } + } + Err(error) => { + yield Ok(Event::default().event("error").data(json!({"error":{"message":error.to_string(),"type":"upstream_error"}}).to_string())); + return; + } + } + } + if output.finish.is_some() && matches!(protocol, Protocol::Chat) { + yield Ok(Event::default().data("[DONE]")); + } + }; + Sse::new(events) + .keep_alive(KeepAlive::default()) + .into_response() +} + +fn chat_chunk(id: &str, model: &str, delta: Value, finish: Value) -> Value { + json!({"id":id,"object":"chat.completion.chunk","created":0,"model":model, + "choices":[{"index":0,"delta":delta,"finish_reason":finish}]}) +} + +fn stream_events( + protocol: Protocol, + event: &ModelEvent, + id: &str, + model: &str, + output: &Output, + text_started: &mut bool, +) -> Vec<(&'static str, Value)> { + match protocol { + Protocol::Chat => match event { + ModelEvent::TextDelta(delta) => vec![( + "message", + chat_chunk(id, model, json!({"content":delta}), Value::Null), + )], + ModelEvent::ToolCallStart { + index, + call_id, + name, + } => vec![( + "message", + chat_chunk( + id, + model, + json!({"tool_calls":[{"index":index,"id":call_id,"type":"function","function":{"name":name,"arguments":""}}]}), + Value::Null, + ), + )], + ModelEvent::ToolCallArgumentsDelta { index, delta } => vec![( + "message", + chat_chunk( + id, + model, + json!({"tool_calls":[{"index":index,"function":{"arguments":delta}}]}), + Value::Null, + ), + )], + ModelEvent::Done(reason) => { + let finish = match reason { + FinishReason::Stop => "stop", + FinishReason::Length => "length", + FinishReason::ToolUse => "tool_calls", + }; + let mut events = vec![("message", chat_chunk(id, model, json!({}), json!(finish)))]; + if output.usage.is_some() { + events.push(( + "message", + json!({"id":id,"object":"chat.completion.chunk","created":0,"model":model, + "choices":[],"usage":output.usage_json(protocol)}), + )); + } + events + } + _ => Vec::new(), + }, + Protocol::Responses => match event { + ModelEvent::TextStart if !*text_started => { + *text_started = true; + vec![ + ( + "response.output_item.added", + json!({"type":"response.output_item.added","output_index":0,"item":{"type":"message","id":format!("msg_{id}"),"status":"in_progress","role":"assistant","content":[]}}), + ), + ( + "response.content_part.added", + json!({"type":"response.content_part.added","output_index":0,"content_index":0,"part":{"type":"output_text","text":"","annotations":[]}}), + ), + ] + } + ModelEvent::TextDelta(delta) => vec![( + "response.output_text.delta", + json!({"type":"response.output_text.delta","delta":delta,"output_index":0,"content_index":0}), + )], + ModelEvent::TextEnd if *text_started => vec![ + ( + "response.output_text.done", + json!({"type":"response.output_text.done","text":output.text,"output_index":0,"content_index":0}), + ), + ( + "response.content_part.done", + json!({"type":"response.content_part.done","output_index":0,"content_index":0,"part":{"type":"output_text","text":output.text,"annotations":[]}}), + ), + ( + "response.output_item.done", + json!({"type":"response.output_item.done","output_index":0,"item":{"type":"message","id":format!("msg_{id}"),"status":"completed","role":"assistant","content":[{"type":"output_text","text":output.text,"annotations":[]}]}}), + ), + ], + ModelEvent::ToolCallStart { + index, + call_id, + name, + } => vec![( + "response.output_item.added", + json!({"type":"response.output_item.added","output_index":index + usize::from(*text_started),"item":{"type":"function_call","id":format!("fc_{call_id}"),"call_id":call_id,"name":name,"arguments":"","status":"in_progress"}}), + )], + ModelEvent::ToolCallArgumentsDelta { index, delta } => vec![( + "response.function_call_arguments.delta", + json!({"type":"response.function_call_arguments.delta","output_index":index + usize::from(*text_started),"delta":delta}), + )], + ModelEvent::ToolCallEnd { index } => { + let tool = output.tools.get(index); + let arguments = tool.map(|tool| tool.arguments.as_str()).unwrap_or_default(); + let call_id = tool.map(|tool| tool.id.as_str()).unwrap_or_default(); + vec![ + ( + "response.function_call_arguments.done", + json!({"type":"response.function_call_arguments.done","output_index":index + usize::from(*text_started),"arguments":arguments}), + ), + ( + "response.output_item.done", + json!({"type":"response.output_item.done","output_index":index + usize::from(*text_started),"item":{"type":"function_call","id":format!("fc_{call_id}"),"call_id":call_id,"name":tool.map(|tool| tool.name.as_str()).unwrap_or("tool"),"arguments":arguments,"status":"completed"}}), + ), + ] + } + ModelEvent::Done(_) => vec![( + "response.completed", + json!({"type":"response.completed","response":output.json(protocol,id,model)}), + )], + _ => Vec::new(), + }, + Protocol::Messages => match event { + ModelEvent::TextStart => { + if *text_started { + Vec::new() + } else { + *text_started = true; + vec![( + "content_block_start", + json!({"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}), + )] + } + } + ModelEvent::TextDelta(delta) => { + let mut events = Vec::new(); + if !*text_started { + *text_started = true; + events.push(("content_block_start",json!({"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}))); + } + events.push(("content_block_delta",json!({"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":delta}}))); + events + } + ModelEvent::TextEnd => { + if *text_started { + vec![( + "content_block_stop", + json!({"type":"content_block_stop","index":0}), + )] + } else { + Vec::new() + } + } + ModelEvent::ToolCallStart { + index, + call_id, + name, + } => vec![( + "content_block_start", + json!({"type":"content_block_start","index":index + usize::from(*text_started),"content_block":{"type":"tool_use","id":call_id,"name":name,"input":{}}}), + )], + ModelEvent::ToolCallArgumentsDelta { index, delta } => vec![( + "content_block_delta", + json!({"type":"content_block_delta","index":index + usize::from(*text_started),"delta":{"type":"input_json_delta","partial_json":delta}}), + )], + ModelEvent::ToolCallEnd { index } => vec![( + "content_block_stop", + json!({"type":"content_block_stop","index":index + usize::from(*text_started)}), + )], + ModelEvent::Done(reason) => { + let stop = match reason { + FinishReason::Stop => "end_turn", + FinishReason::Length => "max_tokens", + FinishReason::ToolUse => "tool_use", + }; + vec![ + ( + "message_delta", + json!({"type":"message_delta","delta":{"stop_reason":stop,"stop_sequence":null},"usage":output.usage_json(protocol)}), + ), + ("message_stop", json!({"type":"message_stop"})), + ] + } + _ => Vec::new(), + }, + } +} diff --git a/server/src/api/byok/protocol.rs b/server/src/api/byok/protocol.rs new file mode 100644 index 0000000..b6c8616 --- /dev/null +++ b/server/src/api/byok/protocol.rs @@ -0,0 +1,471 @@ +use base64::Engine; +use serde_json::Value; + +use crate::{ + model::{ + ContentPart, ModelRequest, ModelSpec, ProjectedContent, ProjectedMessage, PromptSpec, + ProviderType, Role, ToolCallContent, ToolDefinition, ToolResultContent, + }, + Error, Result, +}; + +#[derive(Clone, Copy, Debug)] +pub(super) enum Protocol { + Chat, + Responses, + Messages, +} + +impl Protocol { + pub(super) fn provider_type(self) -> ProviderType { + match self { + Self::Chat => ProviderType::OpenAiChat, + Self::Responses => ProviderType::OpenAiResponses, + Self::Messages => ProviderType::Anthropic, + } + } +} + +pub(super) struct ParsedRequest { + pub public_model_id: String, + pub request: ModelRequest, + pub stream: bool, +} + +pub(super) fn parse(protocol: Protocol, body: &Value) -> Result { + let public_model_id = required_string(body, "model")?.to_owned(); + let stream = body.get("stream").and_then(Value::as_bool).unwrap_or(false); + let mut prompt = PromptSpec { + instructions: String::new(), + tools: Vec::new(), + }; + let mut history = Vec::new(); + let mut model = ModelSpec::new(&public_model_id); + match protocol { + Protocol::Chat => { + parse_chat_messages(body.get("messages"), &mut prompt, &mut history)?; + parse_tools(body.get("tools"), protocol, &mut prompt)?; + model.max_output_tokens = body + .get("max_completion_tokens") + .or_else(|| body.get("max_tokens")) + .and_then(Value::as_u64); + } + Protocol::Responses => { + prompt.instructions = body + .get("instructions") + .and_then(Value::as_str) + .unwrap_or_default() + .to_owned(); + match body.get("input") { + Some(Value::String(text)) => history.push(text_message(0, Role::User, text)), + Some(input) => parse_responses_input(input, &mut history)?, + None => return Err(Error::Protocol("input is required".into())), + } + parse_tools(body.get("tools"), protocol, &mut prompt)?; + model.max_output_tokens = body.get("max_output_tokens").and_then(Value::as_u64); + } + Protocol::Messages => { + prompt.instructions = text_content(body.get("system")).unwrap_or_default(); + parse_anthropic_messages(body.get("messages"), &mut history)?; + parse_tools(body.get("tools"), protocol, &mut prompt)?; + model.max_output_tokens = body.get("max_tokens").and_then(Value::as_u64); + } + } + if history.is_empty() { + return Err(Error::Protocol( + "at least one input message is required".into(), + )); + } + Ok(ParsedRequest { + public_model_id, + request: ModelRequest { + prompt, + model, + history, + }, + stream, + }) +} + +fn required_string<'a>(value: &'a Value, field: &str) -> Result<&'a str> { + value + .get(field) + .and_then(Value::as_str) + .filter(|text| !text.is_empty()) + .ok_or_else(|| Error::Protocol(format!("{field} is required"))) +} + +fn text_message(index: usize, role: Role, text: &str) -> ProjectedMessage { + ProjectedMessage { + message_id: format!("external:{index}"), + role, + content: ProjectedContent::Parts(vec![ContentPart::Text { text: text.into() }]), + } +} + +fn text_content(value: Option<&Value>) -> Option { + match value? { + Value::String(text) => Some(text.clone()), + Value::Array(parts) => Some( + parts + .iter() + .filter_map(|part| { + part.as_str() + .or_else(|| part.get("text").and_then(Value::as_str)) + }) + .collect::>() + .join("\n"), + ), + _ => None, + } +} + +fn parse_parts(value: Option<&Value>) -> Result> { + let value = value.ok_or_else(|| Error::Protocol("message content is required".into()))?; + if let Some(text) = value.as_str() { + return Ok(vec![ContentPart::Text { text: text.into() }]); + } + let array = value + .as_array() + .ok_or_else(|| Error::Protocol("message content must be text or an array".into()))?; + array + .iter() + .map(|part| { + let kind = part.get("type").and_then(Value::as_str).unwrap_or("text"); + match kind { + "text" | "input_text" | "output_text" => Ok(ContentPart::Text { + text: required_string(part, "text")?.into(), + }), + "image_url" | "input_image" | "image" => { + let url = part + .pointer("/image_url/url") + .or_else(|| part.get("image_url")) + .or_else(|| part.pointer("/source/data")) + .and_then(Value::as_str) + .ok_or_else(|| Error::Protocol("image data URL is required".into()))?; + let data_url = if url.starts_with("data:") { + url.to_owned() + } else if let Some(mime) = + part.pointer("/source/media_type").and_then(Value::as_str) + { + format!("data:{mime};base64,{url}") + } else { + return Err(Error::Protocol( + "only base64 image data URLs are supported".into(), + )); + }; + let (prefix, encoded) = data_url + .split_once(',') + .ok_or_else(|| Error::Protocol("invalid image data URL".into()))?; + let mime_type = prefix + .strip_prefix("data:") + .and_then(|prefix| prefix.strip_suffix(";base64")) + .ok_or_else(|| Error::Protocol("invalid image data URL".into()))?; + let data = base64::engine::general_purpose::STANDARD + .decode(encoded) + .map_err(|_| Error::Protocol("invalid base64 image".into()))?; + Ok(ContentPart::Image { + mime_type: mime_type.into(), + data, + }) + } + _ => Err(Error::Protocol(format!("unsupported content type: {kind}"))), + } + }) + .collect() +} + +fn parse_chat_messages( + value: Option<&Value>, + prompt: &mut PromptSpec, + history: &mut Vec, +) -> Result<()> { + let messages = value + .and_then(Value::as_array) + .ok_or_else(|| Error::Protocol("messages must be an array".into()))?; + for (index, message) in messages.iter().enumerate() { + let role = required_string(message, "role")?; + match role { + "system" | "developer" => { + if !history.is_empty() { + return Err(Error::Protocol( + "system messages must precede conversation messages".into(), + )); + } + let text = text_content(message.get("content")) + .ok_or_else(|| Error::Protocol("system content must be text".into()))?; + if !prompt.instructions.is_empty() { + prompt.instructions.push('\n'); + } + prompt.instructions.push_str(&text); + } + "user" => history.push(ProjectedMessage { + message_id: format!("external:{index}"), + role: Role::User, + content: ProjectedContent::Parts(parse_parts(message.get("content"))?), + }), + "assistant" => { + let text = text_content(message.get("content")).unwrap_or_default(); + let calls = message + .get("tool_calls") + .and_then(Value::as_array) + .map(|calls| { + calls + .iter() + .enumerate() + .map(|(index, call)| { + let function = call.get("function").ok_or_else(|| { + Error::Protocol("tool function is required".into()) + })?; + let arguments = required_string(function, "arguments")?; + Ok(ToolCallContent { + index, + call_id: required_string(call, "id")?.into(), + name: required_string(function, "name")?.into(), + arguments: serde_json::from_str(arguments)?, + }) + }) + .collect::>>() + }) + .transpose()? + .unwrap_or_default(); + history.push(ProjectedMessage { + message_id: format!("external:{index}"), + role: Role::Assistant, + content: ProjectedContent::Assistant { + text, + thinking: String::new(), + replay_state: None, + calls, + }, + }); + } + "tool" => history.push(ProjectedMessage { + message_id: format!("external:{index}"), + role: Role::Tool, + content: ProjectedContent::ToolResult(ToolResultContent { + call_id: required_string(message, "tool_call_id")?.into(), + name: message + .get("name") + .and_then(Value::as_str) + .map(str::to_owned) + .or_else(|| { + tool_name( + history, + message + .get("tool_call_id") + .and_then(Value::as_str) + .unwrap_or_default(), + ) + }) + .unwrap_or_else(|| "tool".into()), + content: text_content(message.get("content")).unwrap_or_default(), + is_error: false, + image: None, + provider_parts: Vec::new(), + }), + }), + _ => return Err(Error::Protocol(format!("unsupported role: {role}"))), + } + } + Ok(()) +} + +fn parse_responses_input(value: &Value, history: &mut Vec) -> Result<()> { + let items = value + .as_array() + .ok_or_else(|| Error::Protocol("input must be text or an array".into()))?; + for (index, item) in items.iter().enumerate() { + match item.get("type").and_then(Value::as_str) { + Some("function_call_output") => history.push(ProjectedMessage { + message_id: format!("external:{index}"), + role: Role::Tool, + content: ProjectedContent::ToolResult(ToolResultContent { + call_id: required_string(item, "call_id")?.into(), + name: tool_name(history, required_string(item, "call_id")?) + .unwrap_or_else(|| "tool".into()), + content: text_content(item.get("output")).unwrap_or_default(), + is_error: false, + image: None, + provider_parts: Vec::new(), + }), + }), + Some("function_call") => history.push(ProjectedMessage { + message_id: format!("external:{index}"), + role: Role::Assistant, + content: ProjectedContent::Assistant { + text: String::new(), + thinking: String::new(), + replay_state: None, + calls: vec![ToolCallContent { + index: 0, + call_id: required_string(item, "call_id")?.into(), + name: required_string(item, "name")?.into(), + arguments: serde_json::from_str(required_string(item, "arguments")?)?, + }], + }, + }), + _ => { + let role = item.get("role").and_then(Value::as_str).unwrap_or("user"); + let role = match role { + "user" => Role::User, + "assistant" => Role::Assistant, + _ => return Err(Error::Protocol(format!("unsupported role: {role}"))), + }; + history.push(ProjectedMessage { + message_id: format!("external:{index}"), + role, + content: ProjectedContent::Parts(parse_parts(item.get("content"))?), + }); + } + } + } + Ok(()) +} + +fn parse_anthropic_messages( + value: Option<&Value>, + history: &mut Vec, +) -> Result<()> { + let messages = value + .and_then(Value::as_array) + .ok_or_else(|| Error::Protocol("messages must be an array".into()))?; + for (index, message) in messages.iter().enumerate() { + let role = required_string(message, "role")?; + let content = message + .get("content") + .ok_or_else(|| Error::Protocol("content is required".into()))?; + if let Some(blocks) = content.as_array() { + let mut text_parts = Vec::new(); + let mut calls = Vec::new(); + for (block_index, block) in blocks.iter().enumerate() { + match block.get("type").and_then(Value::as_str) { + Some("tool_use") => calls.push(ToolCallContent { + index: calls.len(), + call_id: required_string(block, "id")?.into(), + name: required_string(block, "name")?.into(), + arguments: block + .get("input") + .cloned() + .unwrap_or(Value::Object(Default::default())), + }), + Some("tool_result") => { + if role != "user" { + return Err(Error::Protocol( + "tool_result must be in a user message".into(), + )); + } + if !text_parts.is_empty() { + history.push(ProjectedMessage { + message_id: format!("external:{index}:{block_index}:text"), + role: Role::User, + content: ProjectedContent::Parts(std::mem::take(&mut text_parts)), + }); + } + let call_id = required_string(block, "tool_use_id")?; + history.push(ProjectedMessage { + message_id: format!("external:{index}:{block_index}:tool"), + role: Role::Tool, + content: ProjectedContent::ToolResult(ToolResultContent { + call_id: call_id.into(), + name: tool_name(history, call_id).unwrap_or_else(|| "tool".into()), + content: text_content(block.get("content")).unwrap_or_default(), + is_error: block + .get("is_error") + .and_then(Value::as_bool) + .unwrap_or(false), + image: None, + provider_parts: Vec::new(), + }), + }); + } + _ => text_parts.extend(parse_parts(Some(&Value::Array(vec![block.clone()])))?), + } + } + if role == "assistant" { + let text = text_parts + .iter() + .filter_map(|part| { + if let ContentPart::Text { text } = part { + Some(text.as_str()) + } else { + None + } + }) + .collect::>() + .join(""); + history.push(ProjectedMessage { + message_id: format!("external:{index}"), + role: Role::Assistant, + content: ProjectedContent::Assistant { + text, + thinking: String::new(), + replay_state: None, + calls, + }, + }); + } else if !text_parts.is_empty() { + history.push(ProjectedMessage { + message_id: format!("external:{index}"), + role: Role::User, + content: ProjectedContent::Parts(text_parts), + }); + } + } else { + let role = match role { + "user" => Role::User, + "assistant" => Role::Assistant, + _ => return Err(Error::Protocol(format!("unsupported role: {role}"))), + }; + history.push(ProjectedMessage { + message_id: format!("external:{index}"), + role, + content: ProjectedContent::Parts(parse_parts(Some(content))?), + }); + } + } + Ok(()) +} + +fn parse_tools(value: Option<&Value>, protocol: Protocol, prompt: &mut PromptSpec) -> Result<()> { + let Some(value) = value else { + return Ok(()); + }; + let tools = value + .as_array() + .ok_or_else(|| Error::Protocol("tools must be an array".into()))?; + for tool in tools { + let function = if matches!(protocol, Protocol::Chat) { + tool.get("function").unwrap_or(tool) + } else { + tool + }; + prompt.tools.push(ToolDefinition { + name: required_string(function, "name")?.into(), + description: function + .get("description") + .and_then(Value::as_str) + .unwrap_or_default() + .into(), + parameters: function + .get("parameters") + .or_else(|| function.get("input_schema")) + .cloned() + .unwrap_or(serde_json::json!({"type":"object","properties":{}})), + }); + } + Ok(()) +} + +fn tool_name(history: &[ProjectedMessage], call_id: &str) -> Option { + history + .iter() + .rev() + .find_map(|message| match &message.content { + ProjectedContent::Assistant { calls, .. } => calls + .iter() + .find(|call| call.call_id == call_id) + .map(|call| call.name.clone()), + _ => None, + }) +} diff --git a/server/src/api/mod.rs b/server/src/api/mod.rs index d557379..19e04ac 100644 --- a/server/src/api/mod.rs +++ b/server/src/api/mod.rs @@ -1,5 +1,6 @@ //! Exposes the HTTP and Connect API layer. +pub mod byok; pub mod cursor; mod router; diff --git a/server/src/app.rs b/server/src/app.rs index 6f97c86..55d1007 100644 --- a/server/src/app.rs +++ b/server/src/app.rs @@ -52,6 +52,17 @@ impl App { config.provider_request_timeout, config.provider_stream_idle_timeout, )); + let byok = api::byok::router( + store.clone(), + plugins.clone(), + provider.clone(), + Some(api::byok::NativeForwarder::new( + store.clone(), + clients.clone(), + config.provider_request_timeout, + config.provider_stream_idle_timeout, + )), + ); let registry = TransportRegistry::with_plugins( store.clone(), provider.clone(), @@ -69,7 +80,7 @@ impl App { config.app_version.clone(), )?; let harness = control.cursor_harness().clone(); - let mut router = api::router(registry.clone(), clients)?; + let mut router = api::router(registry.clone(), clients)?.merge(byok); router = match &config.console { Some(ConsoleSource::Directory(directory)) => { router.merge(control::web_router(control.clone(), directory)) diff --git a/server/src/control/mod.rs b/server/src/control/mod.rs index 28df780..9cad7ba 100644 --- a/server/src/control/mod.rs +++ b/server/src/control/mod.rs @@ -196,6 +196,10 @@ pub fn api_router(service: ControlService) -> Router { "/__byok-api__/api/settings/ports", get(settings::get_ports).put(settings::update_ports), ) + .route( + "/__byok-api__/api/settings/external-api", + get(settings::get_external_api).put(settings::update_external_api), + ) .route( "/__byok-api__/api/settings/storage/statistics", get(settings::get_storage).delete(settings::clear_storage), diff --git a/server/src/control/service.rs b/server/src/control/service.rs index 405b4f2..da8acfd 100644 --- a/server/src/control/service.rs +++ b/server/src/control/service.rs @@ -28,8 +28,8 @@ use crate::{ plugin::{PluginDescriptor, PluginRegistry, PluginRuntime, PluginRuntimeStatus}, provider::{is_valid_response_event, ModelEvent, Provider}, store::{ - CommitSettings, DesktopSettings, PortSettings, ProxySettings, ProxySettingsInput, - StatisticsStorage, Store, TabSettings, TokenPricingSettings, + CommitSettings, DesktopSettings, ExternalApiSettings, PortSettings, ProxySettings, + ProxySettingsInput, StatisticsStorage, Store, TabSettings, TokenPricingSettings, }, Error, Result, }; @@ -605,10 +605,13 @@ impl ControlService { .llm_calls(limit) .await? .into_iter() - .map(|call| CallSummary { - call, - call_kind: "provider_llm", - route: "local_byok", + .map(|call| { + let route = call_route(&call.run_id); + CallSummary { + call, + call_kind: "provider_llm", + route, + } }) .collect::>(); calls.extend( @@ -626,13 +629,14 @@ impl ControlService { pub async fn call(&self, call_id: &str) -> Result { if let Some(call) = self.store.llm_call(call_id).await? { let cursor_trace = self.cursor_trace_detail(&call.run_id).await?; + let route = call_route(&call.run_id); return Ok(CallDetail { request: self.store.llm_call_request(call_id).await?, response_chunks: self.store.llm_call_chunks(call_id).await?, call: CallSummary { call, call_kind: "provider_llm", - route: "local_byok", + route, }, cursor_trace, }); @@ -691,6 +695,17 @@ impl ControlService { self.store.port_settings().await } + pub async fn external_api_settings(&self) -> Result { + self.store.external_api_settings().await + } + + pub async fn set_external_api_settings( + &self, + settings: ExternalApiSettings, + ) -> Result { + self.store.set_external_api_settings(settings).await + } + pub async fn set_ports(&self, settings: PortSettings) -> Result { self.store.set_port_settings(settings).await?; Ok(settings) @@ -761,6 +776,14 @@ impl ControlService { } } +fn call_route(run_id: &str) -> &'static str { + if run_id.starts_with("external-api:") { + "external_api" + } else { + "local_byok" + } +} + fn official_call(trace: CursorRunTraceSummary) -> CallSummary { let model_id = trace.model_id.clone().unwrap_or_else(|| "Cursor".into()); let ttfb = trace diff --git a/server/src/control/settings.rs b/server/src/control/settings.rs index 003c1c9..2f64124 100644 --- a/server/src/control/settings.rs +++ b/server/src/control/settings.rs @@ -8,8 +8,8 @@ use axum::{ use serde::{Deserialize, Serialize}; use crate::store::{ - CommitPromptLocale, CommitSettings, DesktopSettings, PortSettings, ProxySettings, - ProxySettingsInput, StatisticsStorage, StatisticsStorageScope, TabSettings, + CommitPromptLocale, CommitSettings, DesktopSettings, ExternalApiSettings, PortSettings, + ProxySettings, ProxySettingsInput, StatisticsStorage, StatisticsStorageScope, TabSettings, TokenPricingSettings, }; @@ -37,6 +37,19 @@ pub async fn update_ports( Ok(Json(service.set_ports(settings).await?)) } +pub async fn get_external_api( + State(service): State, +) -> Result> { + Ok(Json(service.external_api_settings().await?)) +} + +pub async fn update_external_api( + State(service): State, + Json(settings): Json, +) -> Result> { + Ok(Json(service.set_external_api_settings(settings).await?)) +} + pub async fn get_storage(State(service): State) -> Result> { Ok(Json(service.statistics_storage().await?)) } diff --git a/server/src/provider/anthropic.rs b/server/src/provider/anthropic.rs index f9e2c0a..e2a47c8 100644 --- a/server/src/provider/anthropic.rs +++ b/server/src/provider/anthropic.rs @@ -480,7 +480,7 @@ fn required_u64(value: &Value, name: &str) -> Result { .ok_or_else(|| Error::Provider(format!("Anthropic event is missing {name}"))) } -fn anthropic_usage(value: &Value) -> Usage { +pub(crate) fn anthropic_usage(value: &Value) -> Usage { let input_tokens = value.get("input_tokens").and_then(Value::as_u64); let cache_read_tokens = value.get("cache_read_input_tokens").and_then(Value::as_u64); let cache_write_tokens = value diff --git a/server/src/provider/mod.rs b/server/src/provider/mod.rs index 80cffb1..83d1941 100644 --- a/server/src/provider/mod.rs +++ b/server/src/provider/mod.rs @@ -21,6 +21,21 @@ pub use event::*; pub use openai_chat::OpenAiChatProvider; pub use openai_responses::OpenAiResponsesProvider; pub use recorder::CallRecorder; +pub(crate) use router::custom_headers; + +pub(crate) fn native_usage( + kind: crate::model::ProviderType, + value: &serde_json::Value, +) -> crate::model::Usage { + match kind { + crate::model::ProviderType::OpenAiChat => openai_chat::openai_usage(value), + crate::model::ProviderType::OpenAiResponses => openai_responses::responses_usage(value), + crate::model::ProviderType::Anthropic => anthropic::anthropic_usage(value), + crate::model::ProviderType::Plugin => { + unreachable!("plugin calls have no native HTTP protocol") + } + } +} pub use router::{build as build_provider, ProviderRouter}; pub type ProviderStream = Pin> + Send>>; diff --git a/server/src/provider/openai_responses.rs b/server/src/provider/openai_responses.rs index 275f45a..bac1b90 100644 --- a/server/src/provider/openai_responses.rs +++ b/server/src/provider/openai_responses.rs @@ -522,7 +522,7 @@ fn required_u64(value: &Value, name: &str) -> Result { .ok_or_else(|| Error::Provider(format!("OpenAI Responses event is missing {name}"))) } -fn responses_usage(value: &Value) -> Usage { +pub(crate) fn responses_usage(value: &Value) -> Usage { let input_tokens = value.get("input_tokens").and_then(Value::as_u64); Usage { input_tokens, diff --git a/server/src/provider/recorder.rs b/server/src/provider/recorder.rs index 652fc4e..e60f176 100644 --- a/server/src/provider/recorder.rs +++ b/server/src/provider/recorder.rs @@ -41,7 +41,7 @@ pub struct CallRecorder { inner: Arc, } -pub(super) struct CancelOnDrop { +pub(crate) struct CancelOnDrop { recorder: CallRecorder, } @@ -133,7 +133,7 @@ impl CallRecorder { self.inner.finished.load(Ordering::Acquire) } - pub(super) fn cancel_on_drop(&self) -> CancelOnDrop { + pub(crate) fn cancel_on_drop(&self) -> CancelOnDrop { CancelOnDrop { recorder: self.clone(), } diff --git a/server/src/provider/router.rs b/server/src/provider/router.rs index 17bbd7f..6d8397c 100644 --- a/server/src/provider/router.rs +++ b/server/src/provider/router.rs @@ -308,7 +308,7 @@ fn root_error_message(error: &(dyn std::error::Error + 'static)) -> String { current.to_string() } -fn custom_headers( +pub(crate) fn custom_headers( value: &serde_json::Value, conversation_id: &str, ) -> Result { diff --git a/server/src/store/settings.rs b/server/src/store/settings.rs index 1a2673c..d71b13a 100644 --- a/server/src/store/settings.rs +++ b/server/src/store/settings.rs @@ -13,6 +13,7 @@ const DESKTOP_SETTINGS_KEY: &str = "desktop_lifecycle"; const COMMIT_SETTINGS_KEY: &str = "commit_settings"; const CURSOR_TAKEOVER_ENABLED_KEY: &str = "cursor_takeover_enabled"; const PRICING_SETTINGS_KEY: &str = "token_pricing"; +const EXTERNAL_API_SETTINGS_KEY: &str = "external_api"; /// Embedded default system prompts for commit message generation. pub const DEFAULT_COMMIT_PROMPT_ZH_CN: &str = include_str!("../../prompt/cursor/commit/zh-CN.md"); @@ -26,6 +27,12 @@ pub struct PortSettings { pub service_port: u16, } +#[derive(Clone, Debug, Default, Deserialize, PartialEq, Eq, Serialize)] +pub struct ExternalApiSettings { + pub enabled: bool, + pub api_key: String, +} + #[derive(Clone, Copy, Debug, Deserialize, PartialEq, Serialize)] pub struct TokenPricingSettings { pub input_per_million: f64, @@ -212,6 +219,39 @@ fn read_proxy_settings(value: &str) -> ProxySettingsSecret { } impl Store { + pub async fn external_api_settings(&self) -> Result { + let value = sqlx::query_scalar::<_, String>( + "SELECT value_json FROM service_settings WHERE setting_key = ?", + ) + .bind(EXTERNAL_API_SETTINGS_KEY) + .fetch_optional(&self.pool) + .await?; + value + .map(|value| serde_json::from_str(&value).map_err(Into::into)) + .unwrap_or_else(|| Ok(ExternalApiSettings::default())) + } + + pub async fn set_external_api_settings( + &self, + mut settings: ExternalApiSettings, + ) -> Result { + settings.api_key = settings.api_key.trim().to_owned(); + if settings.enabled && settings.api_key.is_empty() { + return Err(crate::Error::Config( + "external API key is required when enabled".into(), + )); + } + let value_json = serde_json::to_string(&settings)?; + let _write = self.writes.lock().await; + sqlx::query("INSERT INTO service_settings(setting_key, value_json, updated_at_ms) VALUES (?, ?, ?) ON CONFLICT(setting_key) DO UPDATE SET value_json = excluded.value_json, updated_at_ms = excluded.updated_at_ms") + .bind(EXTERNAL_API_SETTINGS_KEY) + .bind(value_json) + .bind(now_ms()) + .execute(&self.pool) + .await?; + Ok(settings) + } + pub(crate) async fn cursor_takeover_enabled(&self) -> Result { let value = sqlx::query_scalar::<_, String>( "SELECT value_json FROM service_settings WHERE setting_key = ?", @@ -495,9 +535,9 @@ impl Store { #[cfg(test)] mod tests { use super::{ - read_proxy_settings, CommitPromptLocale, CommitSettings, ProxyMode, ProxySettingsInput, - ProxySettingsSecret, Store, TokenPricingSettings, DEFAULT_COMMIT_PROMPT_EN_US, - DEFAULT_COMMIT_PROMPT_ZH_CN, PROXY_SETTINGS_KEY, + read_proxy_settings, CommitPromptLocale, CommitSettings, ExternalApiSettings, ProxyMode, + ProxySettingsInput, ProxySettingsSecret, Store, TokenPricingSettings, + DEFAULT_COMMIT_PROMPT_EN_US, DEFAULT_COMMIT_PROMPT_ZH_CN, PROXY_SETTINGS_KEY, }; /// The `outbound_proxy` row exactly as builds before the `system` -> `default` @@ -632,4 +672,43 @@ mod tests { assert_eq!(store.pricing_settings().await.unwrap(), custom); } + + #[tokio::test] + async fn external_api_requires_a_key_and_persists_its_switch() { + let directory = tempfile::tempdir().unwrap(); + let url = format!("sqlite://{}", directory.path().join("test.db").display()); + let store = Store::connect(&url).await.unwrap(); + + assert_eq!( + store.external_api_settings().await.unwrap(), + ExternalApiSettings::default() + ); + assert!(store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: String::new(), + }) + .await + .is_err()); + let settings = ExternalApiSettings { + enabled: true, + api_key: "local-test-key".into(), + }; + assert_eq!( + store + .set_external_api_settings(settings.clone()) + .await + .unwrap(), + settings + ); + assert_eq!( + Store::connect(&url) + .await + .unwrap() + .external_api_settings() + .await + .unwrap(), + settings + ); + } } diff --git a/server/tests/external_api.rs b/server/tests/external_api.rs new file mode 100644 index 0000000..97b5a98 --- /dev/null +++ b/server/tests/external_api.rs @@ -0,0 +1,903 @@ +#[path = "support/fake_provider.rs"] +mod fake_provider; +#[path = "support/fixtures.rs"] +mod fixtures; + +use std::{collections::HashSet, sync::Arc, time::Duration}; + +use axum::{ + body::{to_bytes, Body}, + http::{header, HeaderMap, Request, StatusCode}, + response::IntoResponse, + routing::post, + Json, Router, +}; +use cursor_server::{ + api::byok, + control, + model::{ + ModelConfigInput, ModelType, ProjectedContent, Usage, OPENAI_CHAT_ENDPOINT, + OPENAI_RESPONSES_ENDPOINT, + }, + network::NetworkClients, + plugin::{PluginRegistry, PluginRuntime}, + provider::{FinishReason, ModelEvent, ProviderRouter}, + store::ExternalApiSettings, +}; +use serde_json::{json, Value}; +use tower::ServiceExt; + +fn model_input() -> ModelConfigInput { + ModelConfigInput { + sort_order: 0, + display_name: "Example".into(), + group_name: Some("work".into()), + model_type: ModelType::OpenAi, + base_url: "https://example.com/v1".into(), + use_full_url: false, + api_key: "upstream".into(), + tooltip_data: "test".into(), + model_id: "qwen/model".into(), + reasoning_effort: None, + openai_endpoint: OPENAI_CHAT_ENDPOINT.into(), + openai_extra_params_enabled: false, + openai_extra_params: json!({}), + custom_headers_enabled: false, + custom_headers: json!({}), + anthropic_extra_params_enabled: false, + anthropic_extra_params: json!({}), + context_window_tokens: None, + max_completion_tokens: None, + anthropic_max_tokens: None, + anthropic_thinking_effort: None, + thinking_budget_tokens: None, + } +} + +async fn setup() -> ( + axum::Router, + fake_provider::FakeProvider, + cursor_server::store::Store, + tempfile::TempDir, +) { + let (directory, store) = fixtures::temp_store().await; + store.create_model(&model_input()).await.unwrap(); + let runtime = PluginRuntime::managed().unwrap(); + let plugins = PluginRegistry::managed(store.clone(), runtime.clone(), "0.1.0".into()).unwrap(); + let provider = fake_provider::FakeProvider::default(); + let shared_provider = Arc::new(provider.clone()); + let control = control::ControlService::new( + store.clone(), + shared_provider.clone(), + runtime, + plugins.clone(), + NetworkClients::new(store.clone()), + "0.1.0".into(), + ) + .unwrap(); + let router = byok::router(store.clone(), plugins, shared_provider, None) + .merge(control::api_router(control)); + (router, provider, store, directory) +} + +#[tokio::test] +async fn management_settings_enable_the_external_route_without_restart() { + let (router, _provider, _store, _directory) = setup().await; + let (status, body) = send( + router.clone(), + "GET", + "/__byok-api__/api/settings/external-api", + None, + json!({}), + ) + .await; + assert_eq!(status, StatusCode::OK); + assert_eq!( + serde_json::from_str::(&body).unwrap()["enabled"], + false + ); + let (status, _) = send( + router.clone(), + "PUT", + "/__byok-api__/api/settings/external-api", + None, + json!({"enabled":true,"api_key":"changed-key"}), + ) + .await; + assert_eq!(status, StatusCode::OK); + assert_eq!( + send( + router, + "GET", + "/byok/v1/models", + Some("changed-key"), + json!({}) + ) + .await + .0, + StatusCode::OK + ); +} + +async fn send( + router: axum::Router, + method: &str, + path: &str, + key: Option<&str>, + body: Value, +) -> (StatusCode, String) { + let mut request = Request::builder() + .method(method) + .uri(path) + .header(header::CONTENT_TYPE, "application/json"); + if let Some(key) = key { + request = request.header(header::AUTHORIZATION, format!("Bearer {key}")); + } + let response = router + .oneshot(request.body(Body::from(body.to_string())).unwrap()) + .await + .unwrap(); + let status = response.status(); + let bytes = to_bytes(response.into_body(), 1024 * 1024).await.unwrap(); + (status, String::from_utf8(bytes.to_vec()).unwrap()) +} + +#[tokio::test] +async fn disabled_and_unauthorized_requests_cannot_list_models() { + let (router, _provider, store, _directory) = setup().await; + assert_eq!( + send( + router.clone(), + "GET", + "/byok/v1/models", + Some("secret"), + json!({}) + ) + .await + .0, + StatusCode::FORBIDDEN + ); + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + assert_eq!( + send(router.clone(), "GET", "/byok/v1/models", None, json!({})) + .await + .0, + StatusCode::UNAUTHORIZED + ); + assert_eq!( + send( + router.clone(), + "GET", + "/byok/v1/models", + Some("wrong"), + json!({}) + ) + .await + .0, + StatusCode::UNAUTHORIZED + ); + let (status, body) = send(router, "GET", "/byok/v1/models", Some("secret"), json!({})).await; + assert_eq!(status, StatusCode::OK); + assert_eq!( + serde_json::from_str::(&body).unwrap()["data"][0]["id"], + "work/qwen/model" + ); +} + +#[tokio::test] +async fn duplicate_public_model_ids_use_the_first_configured_model() { + let (router, provider, store, _directory) = setup().await; + let first = store.models().await.unwrap().remove(0); + let mut duplicate = model_input(); + duplicate.sort_order = 1; + duplicate.display_name = "Backup".into(); + duplicate.base_url = "https://backup.example.com/v1".into(); + store.create_model(&duplicate).await.unwrap(); + let mut distinct = model_input(); + distinct.sort_order = 2; + distinct.model_id = "qwen/other".into(); + store.create_model(&distinct).await.unwrap(); + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + + let (status, body) = send( + router.clone(), + "GET", + "/byok/v1/models", + Some("secret"), + json!({}), + ) + .await; + assert_eq!(status, StatusCode::OK, "{body}"); + let models = serde_json::from_str::(&body).unwrap(); + let ids = models["data"] + .as_array() + .unwrap() + .iter() + .map(|model| model["id"].as_str().unwrap()) + .collect::>(); + assert_eq!(ids, ["work/qwen/model", "work/qwen/other"]); + + provider.push(vec![ + ModelEvent::TextDelta("ok".into()), + ModelEvent::Done(FinishReason::Stop), + ]); + let (status, body) = send( + router, + "POST", + "/byok/v1/chat/completions", + Some("secret"), + json!({"model":"work/qwen/model","messages":[{"role":"user","content":"hello"}]}), + ) + .await; + assert_eq!(status, StatusCode::OK, "{body}"); + assert_eq!(provider.requests()[0].model.model_id, first.model_hash); +} + +#[tokio::test] +async fn all_three_protocols_use_the_public_model_id() { + let (router, provider, store, _directory) = setup().await; + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + for (path, body) in [ + ( + "/byok/v1/chat/completions", + json!({"model":"work/qwen/model","messages":[{"role":"user","content":"hello"}]}), + ), + ( + "/byok/v1/responses", + json!({"model":"work/qwen/model","input":"hello"}), + ), + ( + "/byok/v1/messages", + json!({"model":"work/qwen/model","max_tokens":100,"messages":[{"role":"user","content":"hello"}]}), + ), + ] { + provider.push(vec![ + ModelEvent::TextStart, + ModelEvent::TextDelta("world".into()), + ModelEvent::TextEnd, + ModelEvent::Done(FinishReason::Stop), + ]); + let (status, response) = send(router.clone(), "POST", path, Some("secret"), body).await; + assert_eq!(status, StatusCode::OK, "{path}: {response}"); + assert!(response.contains("world"), "{path}: {response}"); + } + assert_eq!(provider.requests().len(), 3); + assert!(provider + .requests() + .iter() + .all(|request| request.model.model_id != "work/qwen/model")); +} + +#[tokio::test] +async fn streaming_chat_returns_incremental_sse_and_tool_calls() { + let (router, provider, store, _directory) = setup().await; + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + provider.push(vec![ + ModelEvent::TextStart, + ModelEvent::TextDelta("hello".into()), + ModelEvent::TextEnd, + ModelEvent::ToolCallStart { + index: 0, + call_id: "call_1".into(), + name: "lookup".into(), + }, + ModelEvent::ToolCallArgumentsDelta { + index: 0, + delta: "{\"q\":1}".into(), + }, + ModelEvent::ToolCallEnd { index: 0 }, + ModelEvent::Done(FinishReason::ToolUse), + ]); + let (status, body) = send(router, "POST", "/byok/v1/chat/completions", Some("secret"), + json!({"model":"work/qwen/model","stream":true,"messages":[{"role":"user","content":"hello"}]})).await; + assert_eq!(status, StatusCode::OK); + assert!(body.contains("chat.completion.chunk")); + assert!(body.contains("lookup")); + assert!(body.contains("tool_calls")); + assert!(body.contains("[DONE]")); +} + +#[tokio::test] +async fn chat_tool_result_keeps_the_assistant_function_name() { + let (router, provider, store, _directory) = setup().await; + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + provider.push(vec![ + ModelEvent::TextDelta("done".into()), + ModelEvent::Done(FinishReason::Stop), + ]); + let body = json!({"model":"work/qwen/model","messages":[ + {"role":"user","content":"find it"}, + {"role":"assistant","tool_calls":[{"id":"call_1","type":"function","function":{"name":"lookup","arguments":"{\"q\":1}"}}]}, + {"role":"tool","tool_call_id":"call_1","content":"found"} + ]}); + assert_eq!( + send( + router, + "POST", + "/byok/v1/chat/completions", + Some("secret"), + body + ) + .await + .0, + StatusCode::OK + ); + let requests = provider.requests(); + let ProjectedContent::ToolResult(result) = &requests[0].history[2].content else { + panic!("expected tool result"); + }; + assert_eq!(result.name, "lookup"); +} + +#[tokio::test] +async fn responses_and_messages_stream_with_protocol_end_events() { + let (router, provider, store, _directory) = setup().await; + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + for (path, request, terminal) in [ + ( + "/byok/v1/responses", + json!({"model":"work/qwen/model","input":"hello","stream":true}), + "response.completed", + ), + ( + "/byok/v1/messages", + json!({"model":"work/qwen/model","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"stream":true}), + "message_stop", + ), + ] { + provider.push(vec![ + ModelEvent::TextDelta("world".into()), + ModelEvent::Done(FinishReason::Stop), + ]); + let (status, body) = send(router.clone(), "POST", path, Some("secret"), request).await; + assert_eq!(status, StatusCode::OK); + assert!(body.contains(terminal), "{path}: {body}"); + } +} + +#[tokio::test] +async fn responses_stream_emits_complete_text_and_tool_item_lifecycles() { + let (router, provider, store, _directory) = setup().await; + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + provider.push(vec![ + ModelEvent::TextStart, + ModelEvent::TextDelta("hello".into()), + ModelEvent::TextEnd, + ModelEvent::ToolCallStart { + index: 0, + call_id: "call_1".into(), + name: "lookup".into(), + }, + ModelEvent::ToolCallArgumentsDelta { + index: 0, + delta: "{\"q\":1}".into(), + }, + ModelEvent::ToolCallEnd { index: 0 }, + ModelEvent::Done(FinishReason::ToolUse), + ]); + let (status, body) = send( + router, + "POST", + "/byok/v1/responses", + Some("secret"), + json!({"model":"work/qwen/model","input":"hello","stream":true}), + ) + .await; + assert_eq!(status, StatusCode::OK); + let events = body + .lines() + .filter_map(|line| line.strip_prefix("event: ")) + .collect::>(); + assert_eq!( + events, + [ + "response.created", + "response.in_progress", + "response.output_item.added", + "response.content_part.added", + "response.output_text.delta", + "response.output_text.done", + "response.content_part.done", + "response.output_item.done", + "response.output_item.added", + "response.function_call_arguments.delta", + "response.function_call_arguments.done", + "response.output_item.done", + "response.completed", + ] + ); + assert!(body.contains("\"output_index\":1"), "{body}"); +} + +#[tokio::test] +async fn messages_stream_closes_each_content_block_before_stopping() { + let (router, provider, store, _directory) = setup().await; + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + provider.push(vec![ + ModelEvent::TextStart, + ModelEvent::TextDelta("hello".into()), + ModelEvent::TextEnd, + ModelEvent::ToolCallStart { + index: 0, + call_id: "call_1".into(), + name: "lookup".into(), + }, + ModelEvent::ToolCallArgumentsDelta { + index: 0, + delta: "{\"q\":1}".into(), + }, + ModelEvent::ToolCallEnd { index: 0 }, + ModelEvent::Done(FinishReason::ToolUse), + ]); + let (status, body) = send(router, "POST", "/byok/v1/messages", Some("secret"), + json!({"model":"work/qwen/model","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"stream":true})).await; + assert_eq!(status, StatusCode::OK); + let events = body + .lines() + .filter_map(|line| line.strip_prefix("event: ")) + .collect::>(); + assert_eq!( + events, + [ + "message_start", + "content_block_start", + "content_block_delta", + "content_block_stop", + "content_block_start", + "content_block_delta", + "content_block_stop", + "message_delta", + "message_stop", + ] + ); + assert!(body.contains("\"index\":1"), "{body}"); +} + +#[tokio::test] +async fn all_protocols_preserve_cached_usage_in_streaming_and_complete_responses() { + let (router, provider, store, _directory) = setup().await; + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + let usage = Usage { + input_tokens: Some(1_000), + context_input_tokens: Some(1_000), + output_tokens: Some(50), + total_tokens: Some(1_050), + cache_read_tokens: Some(800), + cache_write_tokens: Some(20), + reasoning_tokens: Some(10), + }; + for (path, request, usage_pointer, cached_pointer, expected_input) in [ + ( + "/byok/v1/chat/completions", + json!({"model":"work/qwen/model","messages":[{"role":"user","content":"hello"}]}), + "/usage", + "/prompt_tokens_details/cached_tokens", + 1_000, + ), + ( + "/byok/v1/responses", + json!({"model":"work/qwen/model","input":"hello"}), + "/response/usage", + "/input_tokens_details/cached_tokens", + 1_000, + ), + ( + "/byok/v1/messages", + json!({"model":"work/qwen/model","max_tokens":100,"messages":[{"role":"user","content":"hello"}]}), + "/usage", + "/cache_read_input_tokens", + 180, + ), + ] { + for stream in [false, true] { + provider.push(vec![ + ModelEvent::TextDelta("world".into()), + ModelEvent::Usage(usage), + ModelEvent::Done(FinishReason::Stop), + ]); + let mut request = request.clone(); + request["stream"] = json!(stream); + let (status, body) = send(router.clone(), "POST", path, Some("secret"), request).await; + assert_eq!(status, StatusCode::OK, "{path}: {body}"); + let response = if stream { + body.lines() + .filter_map(|line| line.strip_prefix("data: ")) + .filter_map(|line| serde_json::from_str::(line).ok()) + .find(|event| match path { + "/byok/v1/chat/completions" => event.get("usage").is_some(), + "/byok/v1/responses" => event["type"] == "response.completed", + _ => event["type"] == "message_delta", + }) + .unwrap_or_else(|| panic!("missing usage event in {path}: {body}")) + } else { + serde_json::from_str::(&body).unwrap() + }; + let usage = response + .pointer(if stream { usage_pointer } else { "/usage" }) + .unwrap(); + assert_eq!( + usage.pointer(cached_pointer), + Some(&json!(800)), + "{path} stream={stream}: {body}" + ); + let input_field = if path == "/byok/v1/chat/completions" { + "prompt_tokens" + } else { + "input_tokens" + }; + assert_eq!( + usage[input_field], expected_input, + "{path} stream={stream}: {body}" + ); + if path == "/byok/v1/messages" { + assert_eq!(usage["cache_creation_input_tokens"], 20, "{body}"); + } + if stream && path == "/byok/v1/chat/completions" { + let finished = body.find("\"finish_reason\":\"stop\"").unwrap(); + let usage_position = body.find("\"cached_tokens\":800").unwrap(); + let done = body.find("[DONE]").unwrap(); + assert!(finished < usage_position && usage_position < done, "{body}"); + } + } + } +} + +#[tokio::test] +async fn all_entry_and_upstream_protocol_pairs_preserve_cache_usage() { + let (_directory, store) = fixtures::temp_store().await; + store.create_model(&model_input()).await.unwrap(); + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + let runtime = PluginRuntime::managed().unwrap(); + let plugins = PluginRegistry::managed(store.clone(), runtime, "0.1.0".into()).unwrap(); + let fake = fake_provider::FakeProvider::default(); + let inner = byok::router(store.clone(), plugins.clone(), Arc::new(fake.clone()), None); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let port = listener.local_addr().unwrap().port(); + let server = tokio::spawn(async move { axum::serve(listener, inner).await.unwrap() }); + + for (order, group, model_type, endpoint) in [ + (1, "bridge-chat", ModelType::OpenAi, OPENAI_CHAT_ENDPOINT), + ( + 2, + "bridge-responses", + ModelType::OpenAi, + OPENAI_RESPONSES_ENDPOINT, + ), + (3, "bridge-messages", ModelType::Anthropic, ""), + ] { + let mut model = model_input(); + model.sort_order = order; + model.display_name = format!("Bridge {group}"); + model.group_name = Some(group.into()); + model.model_type = model_type; + model.openai_endpoint = endpoint.into(); + model.base_url = format!("http://127.0.0.1:{port}/byok/v1"); + model.model_id = "work/qwen/model".into(); + model.api_key = "secret".into(); + store.create_model(&model).await.unwrap(); + } + let provider = ProviderRouter::new( + store.clone(), + plugins.clone(), + NetworkClients::new(store.clone()), + Duration::from_secs(10), + Duration::from_secs(10), + ); + let outer = byok::router( + store.clone(), + plugins, + Arc::new(provider), + Some(byok::NativeForwarder::new( + store.clone(), + NetworkClients::new(store.clone()), + Duration::from_secs(10), + Duration::from_secs(10), + )), + ); + for (path, request) in [ + ( + "/byok/v1/chat/completions", + json!({"messages":[{"role":"user","content":"hello"}], + "tools":[{"type":"function","function":{"name":"lookup","description":"Look up a value","parameters":{"type":"object","properties":{"q":{"type":"string"}}}}}]}), + ), + ( + "/byok/v1/responses", + json!({"input":"hello", + "tools":[{"type":"function","name":"lookup","description":"Look up a value","parameters":{"type":"object","properties":{"q":{"type":"string"}}}}]}), + ), + ( + "/byok/v1/messages", + json!({"max_tokens":100,"messages":[{"role":"user","content":"hello"}], + "tools":[{"name":"lookup","description":"Look up a value","input_schema":{"type":"object","properties":{"q":{"type":"string"}}}}]}), + ), + ] { + for group in ["bridge-chat", "bridge-responses", "bridge-messages"] { + for stream in [false, true] { + let prior_calls: HashSet<_> = store + .llm_calls(100) + .await + .unwrap() + .into_iter() + .map(|call| call.call_id) + .collect(); + fake.push(vec![ + ModelEvent::TextStart, + ModelEvent::TextDelta("world".into()), + ModelEvent::TextEnd, + ModelEvent::ToolCallStart { + index: 0, + call_id: "call_1".into(), + name: "lookup".into(), + }, + ModelEvent::ToolCallArgumentsDelta { + index: 0, + delta: "{\"q\":\"x\"}".into(), + }, + ModelEvent::ToolCallEnd { index: 0 }, + ModelEvent::Usage(Usage { + input_tokens: Some(1_000), + context_input_tokens: Some(1_000), + output_tokens: Some(50), + total_tokens: Some(1_050), + cache_read_tokens: Some(800), + ..Usage::default() + }), + ModelEvent::Done(FinishReason::ToolUse), + ]); + let mut request = request.clone(); + request["model"] = json!(format!("{group}/work/qwen/model")); + request["stream"] = json!(stream); + let (status, body) = + send(outer.clone(), "POST", path, Some("secret"), request).await; + assert_eq!( + status, + StatusCode::OK, + "{path} -> {group} stream={stream}: {body}" + ); + assert!( + body.contains("world"), + "{path} -> {group} stream={stream}: {body}" + ); + assert!( + body.contains("lookup") && body.contains("call_1"), + "{path} -> {group} stream={stream}: {body}" + ); + let requests = fake.requests(); + let forwarded = requests.last().unwrap(); + assert_eq!( + forwarded.prompt.tools.len(), + 1, + "{path} -> {group} stream={stream}" + ); + assert_eq!( + forwarded.prompt.tools[0].name, "lookup", + "{path} -> {group} stream={stream}" + ); + let calls = store.llm_calls(100).await.unwrap(); + let call = calls + .iter() + .find(|call| { + call.display_name == format!("Bridge {group}") + && !prior_calls.contains(&call.call_id) + }) + .unwrap(); + assert_eq!( + call.cache_read_tokens, + Some(800), + "{path} -> {group} stream={stream}" + ); + } + } + } + assert_eq!(fake.requests().len(), 18); + server.abort(); +} + +#[tokio::test] +async fn matching_http_protocols_forward_native_requests_and_responses() { + let (directory, store) = fixtures::temp_store().await; + store + .set_external_api_settings(ExternalApiSettings { + enabled: true, + api_key: "secret".into(), + }) + .await + .unwrap(); + let runtime = PluginRuntime::managed().unwrap(); + let plugins = PluginRegistry::managed(store.clone(), runtime, "0.1.0".into()).unwrap(); + let (sender, mut receiver) = tokio::sync::mpsc::unbounded_channel::<(HeaderMap, Value)>(); + let upstream = Router::new().route( + "/{*path}", + post(move |headers: HeaderMap, Json(body): Json| { + let sender = sender.clone(); + async move { + sender.send((headers, body.clone())).unwrap(); + if body["native_error"] == true { + return ( + StatusCode::TOO_MANY_REQUESTS, + Json(json!({"error":"native-rate-limit"})), + ) + .into_response(); + } + if body["stream"] == true { + ( + [(header::CONTENT_TYPE, "text/event-stream")], + "data: {\"native_marker\":\"untouched-stream\"}\n\n", + ) + .into_response() + } else { + Json(json!({"native_marker":"untouched-complete","echo":body})).into_response() + } + } + }), + ); + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let port = listener.local_addr().unwrap().port(); + let server = tokio::spawn(async move { axum::serve(listener, upstream).await.unwrap() }); + for (order, group, model_type, endpoint, path, request) in [ + ( + 1, + "chat", + ModelType::OpenAi, + OPENAI_CHAT_ENDPOINT, + "/byok/v1/chat/completions", + json!({"messages":[{"role":"user","content":"hi"}],"native_extension":{"keep":1}}), + ), + ( + 2, + "responses", + ModelType::OpenAi, + OPENAI_RESPONSES_ENDPOINT, + "/byok/v1/responses", + json!({"input":[{"type":"native_unsupported","value":1}],"native_extension":{"keep":2}}), + ), + ( + 3, + "messages", + ModelType::Anthropic, + "", + "/byok/v1/messages", + json!({"max_tokens":100,"messages":[{"role":"user","content":[{"type":"document","source":{"type":"url","url":"https://example.com"}}]}],"native_extension":{"keep":3}}), + ), + ] { + let mut model = model_input(); + model.sort_order = order; + model.group_name = Some(group.into()); + model.model_type = model_type; + model.openai_endpoint = endpoint.into(); + model.base_url = format!("http://127.0.0.1:{port}/byok/v1"); + model.model_id = "native-model".into(); + store.create_model(&model).await.unwrap(); + for stream in [false, true] { + let mut request = request.clone(); + request["model"] = json!(format!("{group}/native-model")); + request["stream"] = json!(stream); + let (status, body) = send( + byok::router( + store.clone(), + plugins.clone(), + Arc::new(fake_provider::FakeProvider::default()), + Some(byok::NativeForwarder::new( + store.clone(), + NetworkClients::new(store.clone()), + Duration::from_secs(10), + Duration::from_secs(10), + )), + ), + "POST", + path, + Some("secret"), + request.clone(), + ) + .await; + assert_eq!(status, StatusCode::OK, "{path} stream={stream}: {body}"); + if stream { + assert_eq!(body, "data: {\"native_marker\":\"untouched-stream\"}\n\n"); + } else { + assert_eq!( + serde_json::from_str::(&body).unwrap()["native_marker"], + "untouched-complete" + ); + } + let (upstream_headers, forwarded) = receiver.recv().await.unwrap(); + request["model"] = json!("native-model"); + assert_eq!(forwarded, request); + if model_type == ModelType::Anthropic { + assert_eq!(upstream_headers["x-api-key"], "upstream"); + assert_eq!(upstream_headers["anthropic-version"], "2023-06-01"); + assert!(!upstream_headers.contains_key(header::AUTHORIZATION)); + } else { + assert_eq!(upstream_headers[header::AUTHORIZATION], "Bearer upstream"); + assert!(!upstream_headers.contains_key("x-api-key")); + } + } + } + let mut error_request = json!({"model":"chat/native-model","messages":[],"native_error":true}); + let (status, body) = send( + byok::router( + store.clone(), + plugins, + Arc::new(fake_provider::FakeProvider::default()), + Some(byok::NativeForwarder::new( + store.clone(), + NetworkClients::new(store.clone()), + Duration::from_secs(10), + Duration::from_secs(10), + )), + ), + "POST", + "/byok/v1/chat/completions", + Some("secret"), + error_request.clone(), + ) + .await; + assert_eq!(status, StatusCode::TOO_MANY_REQUESTS); + assert_eq!( + serde_json::from_str::(&body).unwrap()["error"], + "native-rate-limit" + ); + error_request["model"] = json!("native-model"); + assert_eq!(receiver.recv().await.unwrap().1, error_request); + server.abort(); + drop(directory); +}