From d5f31cb2532aed44d9d9c3779cb6d5976542c15b Mon Sep 17 00:00:00 2001 From: "dayuan.jiang" Date: Mon, 5 Oct 2026 10:52:37 +0900 Subject: [PATCH] fix(server): use the keys the user sent, and more review fixes Found by the second PR review: - With AWS_BEARER_TOKEN_BEDROCK set on the server, a request with the user's AWS keys ran on the server's token: the Bedrock SDK prefers it. Checked with Bedrock: invalid user keys used to get an answer. - An OpenAI key with the official URL filled in (the settings form does that) went to the Responses API. Back to main's rule: a configured base URL uses Chat Completions. - A user's Ollama key went to the server's OLLAMA_BASE_URL, for chat and for the model list. Like every other provider, it goes to the user's base URL or Ollama Cloud. - The server's keyless Ollama and EdgeOne were not counted in the quota. - AI_MODEL models on the server's keys ran on any provider with a server key, not only on AI_PROVIDER. - A user's Azure key without a base URL used the server's resource name. - The admin panel's Test button failed whenever access codes were set. - DeepSeek's errors in the stream (plain text) were shown as they were, without a hint and also on the server's keys. Bedrock's throttling in the stream was not recognised as a rate limit. - The EdgeOne function accepted text/plain; x=application/json, which other sites can send without a CORS preflight. - Desktop app: a launch that found the old port taken for a moment (the previous version still quitting after an update) remembered the new port for good. The new port is kept only when Windows reserves the old one. A failed read of the presets file moved it aside as corrupt, and a save could then replace the presets. Switching presets on the same port now reloads the page. The dev launcher no longer misses a preset change made before or during a restart. --- app/api/admin/test-model/route.ts | 6 +- app/api/chat/route.ts | 22 +++-- app/api/validate-model/route.ts | 6 +- edge-functions/api/edgeai/chat/completions.ts | 12 +-- electron/main/config-manager.ts | 27 ++++++- electron/main/port-manager.ts | 35 ++++++-- electron/main/window-manager.ts | 6 +- lib/ai-providers.ts | 48 ++++++++--- lib/llm-errors.ts | 15 +++- lib/provider-models.ts | 11 +-- scripts/electron-dev.mjs | 73 ++++++++--------- tests/unit/ai-providers-credentials.test.ts | 45 +++++++++++ tests/unit/chat-route-quota.test.ts | 50 +++++++++++- tests/unit/config-manager.test.ts | 67 +++++++++++++++ tests/unit/edgeone-function.test.ts | 15 ++++ tests/unit/llm-errors.test.ts | 21 +++++ tests/unit/port-manager.test.ts | 81 ++++++++++++++++--- tests/unit/provider-models.test.ts | 12 +++ tests/unit/validate-model-route.test.ts | 38 +++++++++ 19 files changed, 499 insertions(+), 91 deletions(-) create mode 100644 tests/unit/config-manager.test.ts diff --git a/app/api/admin/test-model/route.ts b/app/api/admin/test-model/route.ts index 7ab08ee5..c43a98ab 100644 --- a/app/api/admin/test-model/route.ts +++ b/app/api/admin/test-model/route.ts @@ -50,7 +50,11 @@ export async function POST(req: Request) { return validateModel( new Request(new URL("/api/validate-model", req.url), { method: "POST", - headers: { "Content-Type": "application/json" }, + headers: { + "Content-Type": "application/json", + // Checked again there, in place of an access code + "x-admin-password": req.headers.get("x-admin-password") || "", + }, body: JSON.stringify({ provider: resolved.provider, apiKey: resolved.apiKey, diff --git a/app/api/chat/route.ts b/app/api/chat/route.ts index f957aae6..4b9388b6 100644 --- a/app/api/chat/route.ts +++ b/app/api/chat/route.ts @@ -14,6 +14,7 @@ import { checkAccessCode } from "@/lib/access-code" import { CACHE_POINT, getAIModel, + getServerProvider, SINGLE_SYSTEM_PROVIDERS, supportsPromptCaching, usesServerCredentials, @@ -49,6 +50,7 @@ import { } from "@/lib/server-model-config" import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection" import { getSystemPrompt } from "@/lib/system-prompts" +import { normalizeBaseUrl } from "@/lib/types/model-config" import { getUserIdFromRequest } from "@/lib/user-id" import { hasCells } from "@/packages/mcp-server/src/pages.ts" import { @@ -245,15 +247,17 @@ async function handleChatRequest(req: Request): Promise { } = getAIModel(clientOverrides) // On the server's own keys, only run models the server offers: a server - // model picked by id (its model name is fixed above) or one in AI_MODEL. - // With their own key, users can run any model. + // model picked by id (its model name is fixed above) or one in AI_MODEL + // on AI_PROVIDER. With their own key, users can run any model. const onServerCredentials = usesServerCredentials( resolvedProvider, clientOverrides, ) const envModels = process.env.AI_MODEL?.split(",").map((m) => m.trim()) || [] - if (onServerCredentials && !serverModel && !envModels.includes(modelId)) { + const offeredInEnv = + envModels.includes(modelId) && resolvedProvider === getServerProvider() + if (onServerCredentials && !serverModel && !offeredInEnv) { return Response.json( { error: `Model "${modelId}" is not available on this server. Add your own API key in Settings to use it.`, @@ -264,10 +268,16 @@ async function handleChatRequest(req: Request): Promise { // === SERVER-SIDE QUOTA CHECK START === // Quota is opt-in (DYNAMODB_QUOTA_TABLE) and counts what runs on the - // server's keys. Decided by the key actually used: a key header the - // provider never reads must not skip it. + // server's keys, or on its keyless Ollama or EdgeOne. Decided by the key + // actually used: a key header the provider never reads must not skip it. + const onServerEndpoint = + (resolvedProvider === "ollama" || resolvedProvider === "edgeone") && + !clientOverrides.apiKey && + !normalizeBaseUrl(req.headers.get("x-ai-base-url") ?? "") const countsQuota = - isQuotaEnabled() && onServerCredentials && userId !== "anonymous" + isQuotaEnabled() && + (onServerCredentials || onServerEndpoint) && + userId !== "anonymous" if (countsQuota) { const quotaCheck = await checkAndIncrementRequest(userId, { requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10, diff --git a/app/api/validate-model/route.ts b/app/api/validate-model/route.ts index c8567c1e..b83aeab9 100644 --- a/app/api/validate-model/route.ts +++ b/app/api/validate-model/route.ts @@ -2,6 +2,7 @@ import { streamText, tool } from "ai" import { NextResponse } from "next/server" import { z } from "zod" import { checkAccessCode } from "@/lib/access-code" +import { checkAdminAuth } from "@/lib/admin/auth" import { getAIModel, usesServerCredentials } from "@/lib/ai-providers" import { classifyLLMError } from "@/lib/llm-errors" import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection" @@ -34,9 +35,10 @@ const NO_TOOL_CALL_WARNING = "Connected, but the model answered without calling a tool. It may not support tool calls, which drawing needs." export async function POST(req: Request) { - // Lets the server send requests to arbitrary URLs, so require the access code + // Lets the server send requests to arbitrary URLs, so require the access + // code, or the admin password (the admin panel's Test button) const accessError = checkAccessCode(req) - if (accessError) return accessError + if (accessError && checkAdminAuth(req)) return accessError try { const body: ValidateRequest = await req.json() diff --git a/edge-functions/api/edgeai/chat/completions.ts b/edge-functions/api/edgeai/chat/completions.ts index e2c558fe..19112ab9 100644 --- a/edge-functions/api/edgeai/chat/completions.ts +++ b/edge-functions/api/edgeai/chat/completions.ts @@ -97,11 +97,13 @@ function hasValidAccessCode(request: Request, env: any): boolean { export async function onRequest({ request, env }: any) { // Requiring JSON also makes any cross-site browser request need a CORS - // preflight, which fails without CORS headers - if ( - request.method !== "POST" || - !request.headers.get("content-type")?.includes("application/json") - ) { + // preflight, which fails without CORS headers. Only the type before any + // parameters counts: "text/plain; x=application/json" needs none. + const mediaType = (request.headers.get("content-type") ?? "") + .split(";")[0] + .trim() + .toLowerCase() + if (request.method !== "POST" || mediaType !== "application/json") { return createResponse( { error: { diff --git a/electron/main/config-manager.ts b/electron/main/config-manager.ts index 38bd07a1..9f780d13 100644 --- a/electron/main/config-manager.ts +++ b/electron/main/config-manager.ts @@ -159,6 +159,10 @@ function getConfigFilePath(): string { return path.join(userDataPath, CONFIG_FILE_NAME) } +// The presets file exists but the last read failed: a save now would +// replace the user's presets with the empty list that read returned +let presetsUnreadable = false + /** * Load presets from the config file * Decrypts sensitive fields automatically @@ -175,8 +179,24 @@ export function loadPresets(): ConfigPresetsFile { } } + let content: string + try { + content = readFileSync(configPath, "utf-8") + presetsUnreadable = false + } catch (error) { + // Often only for now (on Windows an antivirus scanner can hold the + // file): keep the file, and refuse saves based on this empty list + console.error("Failed to read config presets:", error) + presetsUnreadable = true + return { + version: 1, + currentPresetId: null, + presets: [], + userLocale: undefined, + } + } + try { - const content = readFileSync(configPath, "utf-8") const data = JSON.parse(content) as ConfigPresetsFile // Decrypt sensitive fields in each preset @@ -211,6 +231,11 @@ export function loadPresets(): ConfigPresetsFile { * Encrypts sensitive fields automatically */ export function savePresets(data: ConfigPresetsFile): void { + if (presetsUnreadable) { + throw new Error( + "The presets file could not be read, so it was not overwritten. Please try again.", + ) + } const configPath = getConfigFilePath() const userDataPath = app.getPath("userData") diff --git a/electron/main/port-manager.ts b/electron/main/port-manager.ts index f9e8cf7c..9daa5635 100644 --- a/electron/main/port-manager.ts +++ b/electron/main/port-manager.ts @@ -43,15 +43,29 @@ function loadSavedPort(): number | null { } } +/** + * Why the legacy port was not used at this launch: "EACCES" when the system + * reserves it (Windows excludes port ranges for Hyper-V, which can change + * on each boot), "EADDRINUSE" when another process holds it for now + */ +let legacyPortError: string | null = null + /** * Remember the port of the first production launch. A later launch that * found it taken keeps it remembered: the user's data lives under that - * origin, and the next launch goes back to it once it is free. + * origin, and the next launch goes back to it once it is free. Only the two + * fixed ports count, and 13370 only when the system reserves the legacy + * port: while another process holds it (such as the previous version still + * quitting after an update), the next launch tries it again. */ export function saveServerPort(port: number): void { if (!app.isPackaged || loadSavedPort() !== null) { return } + const fixed = + port === PORT_CONFIG.legacyProduction || + (port === PORT_CONFIG.production && legacyPortError === "EACCES") + if (!fixed) return try { writeFileSync(getSavedPortPath(), JSON.stringify({ port }), "utf-8") } catch (error) { @@ -60,23 +74,31 @@ export function saveServerPort(port: number): void { } /** - * Check if a specific port is available + * Try to listen on a port. Resolves to null when it is free, else to the + * error code. */ -export function isPortAvailable(port: number): Promise { +function portError(port: number): Promise { return new Promise((resolve) => { const server = net.createServer() server.once("error", (err: NodeJS.ErrnoException) => { console.warn(`Port ${port} unavailable: ${err.code}`) - resolve(false) + resolve(err.code ?? "unknown") }) server.once("listening", () => { server.close() - resolve(true) + resolve(null) }) server.listen(port, "127.0.0.1") }) } +/** + * Check if a specific port is available + */ +export async function isPortAvailable(port: number): Promise { + return (await portError(port)) === null +} + /** * Find an available port * - In development: uses fixed port (6002) @@ -123,7 +145,8 @@ export async function findAvailablePort(reuseExisting = true): Promise { // In production, try legacy port first to preserve existing users' localStorage if (!isDev) { const legacyPort = PORT_CONFIG.legacyProduction - if (await isPortAvailable(legacyPort)) { + legacyPortError = await portError(legacyPort) + if (legacyPortError === null) { allocatedPort = legacyPort return legacyPort } diff --git a/electron/main/window-manager.ts b/electron/main/window-manager.ts index 4660fb5d..d70650d7 100644 --- a/electron/main/window-manager.ts +++ b/electron/main/window-manager.ts @@ -106,11 +106,13 @@ export function getAppUrl(): string | null { } /** - * Point the main window at a new app server URL - * (the restarted server can come up on a different port) + * Point the main window at the restarted app server (it can come up on a + * different port). On the same port the page reloads, so it fetches the new + * preset's server models instead of sending the old preset's choice. */ export function setAppUrl(url: string): void { if (url === appUrl) { + mainWindow?.webContents.reload() return } appUrl = url diff --git a/lib/ai-providers.ts b/lib/ai-providers.ts index b1c3f2ea..044020a5 100644 --- a/lib/ai-providers.ts +++ b/lib/ai-providers.ts @@ -639,6 +639,8 @@ interface Endpoint { fetch?: typeof fetch authToken?: string // Anthropic Bearer auth resourceName?: string // Azure + // baseURL comes from the settings or env, not the provider's default + configuredBaseURL?: boolean } /** @@ -679,11 +681,10 @@ function createModel( switch (provider) { case "openai": { const openaiProvider = createOpenAI(opts) - // A custom base URL is usually a proxy that only has Chat - // Completions; the official endpoint uses the Responses API, - // which returns reasoning for the o-series and gpt-5 or later - return e.baseURL && - e.baseURL !== PROVIDER_INFO.openai.defaultBaseUrl + // A configured base URL is usually a proxy that only has Chat + // Completions; without one the Responses API is used, which + // returns reasoning for the o-series and gpt-5 or later + return e.configuredBaseURL ? openaiProvider.chat(modelId) : openaiProvider(modelId) } @@ -899,6 +900,9 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { ...(overrides?.awsSessionToken && { sessionToken: overrides.awsSessionToken, }), + // Without an apiKey the SDK reads the server's + // AWS_BEARER_TOKEN_BEDROCK, which wins over the keys + apiKey: "", }) : adminAccessKeyId && adminSecretAccessKey ? createAmazonBedrock({ @@ -955,7 +959,13 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { } case "ollama": { - const baseURL = overrides?.baseUrl || process.env.OLLAMA_BASE_URL + // Like other providers, a user's key never goes to the server's + // base URL; without a base URL it is an Ollama Cloud key + const baseURL = + overrides?.baseUrl || + (overrides?.apiKey + ? PROVIDER_INFO.ollama.defaultBaseUrl + : process.env.OLLAMA_BASE_URL) // SECURITY: When client provides a custom base URL, only use // client-provided API key. Never fall back to server OLLAMA_API_KEY // to prevent leaking server credentials to user-controlled endpoints. @@ -1003,16 +1013,24 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { const publicDefault = defaultUrl?.startsWith("https://") ? defaultUrl : undefined - const baseURL = resolveBaseURL( + const configuredBaseURL = resolveBaseURL( overrides?.apiKey, overrides?.baseUrl, resolveBaseUrlEnv(overrides, baseUrlEnv), - SDK_KNOWS_ENDPOINT.has(provider) && - !(provider === "openai" && overrides?.apiKey) - ? undefined - : publicDefault, ) - if (!baseURL && !SDK_KNOWS_ENDPOINT.has(provider)) { + const baseURL = + configuredBaseURL || + (SDK_KNOWS_ENDPOINT.has(provider) && + !(provider === "openai" && overrides?.apiKey) + ? undefined + : publicDefault) + // With a user's Azure key the SDK would read the server's + // AZURE_RESOURCE_NAME + if ( + !baseURL && + (!SDK_KNOWS_ENDPOINT.has(provider) || + (provider === "azure" && overrides?.apiKey)) + ) { throw new Error( `${PROVIDER_INFO[provider].label} needs a base URL. Add it in the model settings.`, ) @@ -1020,6 +1038,7 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { model = createModel(provider, modelId, { apiKey, baseURL, + configuredBaseURL: !!configuredBaseURL, fetch: guardedFetch, // Bearer auth for Anthropic when there is no API key authToken: @@ -1038,6 +1057,11 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { return { model, providerOptions, modelId, provider } } +/** The provider of the server's own config: AI_PROVIDER, or the one with a key */ +export function getServerProvider(): ProviderName | null { + return (process.env.AI_PROVIDER as ProviderName) || detectProvider() +} + /** * Whether the call is paid for by the server's own credentials (env keys or * IAM role) rather than credentials sent with the request. Mirrors which key diff --git a/lib/llm-errors.ts b/lib/llm-errors.ts index 077cc2f5..c587276d 100644 --- a/lib/llm-errors.ts +++ b/lib/llm-errors.ts @@ -80,7 +80,8 @@ const GENERAL_TEXTS: Array<[RegExp, LLMErrorCode]> = [ /invalid[_ ]api[_ ]key|incorrect api key|unauthorized/i, "invalid_api_key", ], - [/rate limit|too many requests/i, "rate_limited"], + // "too many tokens": Bedrock's throttling + [/rate limit|too many requests|too many tokens/i, "rate_limited"], [ /Cannot connect to API|ECONNREFUSED|ENOTFOUND|ECONNRESET|ETIMEDOUT|fetch failed/i, "cannot_connect", @@ -115,8 +116,16 @@ function problemDetail(body: string): string | undefined { * it can name the server's account, role or internal hosts. */ export function streamErrorText(error: unknown, hideDetails = false): string { - // The SDK passes an invalid tool call's error as a plain string - if (typeof error === "string") return error + // The SDK passes an invalid tool call's error as a plain string. Other + // strings come from providers (DeepSeek's SDK sends stream errors so). + if ( + typeof error === "string" && + /^(Invalid input for tool|Model tried to call unavailable tool)/.test( + error, + ) + ) { + return error + } if (isToolCallError(error)) return (error as Error).message const classified = classifyLLMError(error) if (hideDetails) { diff --git a/lib/provider-models.ts b/lib/provider-models.ts index b64fc805..f291d35f 100644 --- a/lib/provider-models.ts +++ b/lib/provider-models.ts @@ -77,11 +77,12 @@ async function getJson( /** * Where to list from without the user's base URL: where chat goes then. For - * Ollama that is the server's Ollama, else the SDK's local default; a local - * default in PROVIDER_INFO (SGLang's) only fills the settings form. + * Ollama without a key that is the server's Ollama, else the SDK's local + * default; a local default in PROVIDER_INFO (SGLang's) only fills the + * settings form. */ -function listFallbackUrl(provider: ProviderName): string { - if (provider === "ollama") { +function listFallbackUrl(provider: ProviderName, apiKey?: string): string { + if (provider === "ollama" && !apiKey) { return process.env.OLLAMA_BASE_URL || "http://127.0.0.1:11434/api" } const url = PROVIDER_INFO[provider].defaultBaseUrl @@ -98,7 +99,7 @@ export async function listProviderModels( { apiKey, baseUrl }: { apiKey?: string; baseUrl?: string }, fetchFn: typeof fetch = fetch, ): Promise { - const base = normalizeBaseUrl(baseUrl || listFallbackUrl(provider)) + const base = normalizeBaseUrl(baseUrl || listFallbackUrl(provider, apiKey)) const bearer: Record = apiKey ? { Authorization: `Bearer ${apiKey}` } : {} diff --git a/scripts/electron-dev.mjs b/scripts/electron-dev.mjs index 0ae5cddc..99ee6c84 100644 --- a/scripts/electron-dev.mjs +++ b/scripts/electron-dev.mjs @@ -224,6 +224,39 @@ async function main() { let configWatcher = null let restartPending = false + // Restart Next.js when the preset env vars really changed + async function applyPresetChange() { + if (restartPending) return + const newContent = readPresetEnvFile() + if (newContent === null || newContent === presetEnvContent) return + + restartPending = true + presetEnvContent = newContent + console.log( + "\nšŸ”„ Preset configuration changed, restarting Next.js server...", + ) + + // Kill current Next.js process + killProcess(nextProcess) + + // Wait a bit for process to die + await new Promise((r) => setTimeout(r, 1000)) + + // Reload preset and restart + nextProcess = startNextServer(loadPresetEnv(newContent)) + + try { + await waitForServer(NEXT_URL) + console.log("āœ… Next.js server restarted with new configuration\n") + } catch (err) { + console.error("āŒ Failed to restart Next.js:", err.message) + } + + restartPending = false + // A change written during the restart was skipped above + applyPresetChange() + } + function setupConfigWatcher() { if (!existsSync(userDataPath)) { // Directory doesn't exist yet, check again later @@ -236,45 +269,13 @@ async function main() { configWatcher = watch( userDataPath, { persistent: false }, - async (_eventType, filename) => { - if (filename !== PRESET_ENV_FILE || restartPending) return - - // Only restart when the preset env vars really changed - const newContent = readPresetEnvFile() - if (newContent === null || newContent === presetEnvContent) - return - - restartPending = true - presetEnvContent = newContent - console.log( - "\nšŸ”„ Preset configuration changed, restarting Next.js server...", - ) - - // Kill current Next.js process - killProcess(nextProcess) - - // Wait a bit for process to die - await new Promise((r) => setTimeout(r, 1000)) - - // Reload preset and restart - nextProcess = startNextServer(loadPresetEnv(newContent)) - - try { - await waitForServer(NEXT_URL) - console.log( - "āœ… Next.js server restarted with new configuration\n", - ) - } catch (err) { - console.error( - "āŒ Failed to restart Next.js:", - err.message, - ) - } - - restartPending = false + (_eventType, filename) => { + if (filename === PRESET_ENV_FILE) applyPresetChange() }, ) console.log("šŸ‘€ Watching for preset configuration changes...") + // Electron may have written its preset before the watch started + applyPresetChange() } catch (_err) { // Directory might not be ready yet, try again later setTimeout(setupConfigWatcher, 5000) diff --git a/tests/unit/ai-providers-credentials.test.ts b/tests/unit/ai-providers-credentials.test.ts index 7b13eae8..a9ac7f4e 100644 --- a/tests/unit/ai-providers-credentials.test.ts +++ b/tests/unit/ai-providers-credentials.test.ts @@ -36,6 +36,11 @@ vi.mock("@aws-sdk/credential-providers", () => ({ fromNodeProviderChain: vi.fn(() => "node-chain"), })) +vi.mock("ollama-ai-provider-v2", () => { + const mockProviderFn = vi.fn(() => ({ modelId: "test-model" })) + return { createOllama: vi.fn(() => mockProviderFn) } +}) + vi.mock("@openrouter/ai-sdk-provider", () => { const mockProviderFn = vi.fn(() => ({ modelId: "test-model" })) return { createOpenRouter: vi.fn(() => mockProviderFn) } @@ -60,6 +65,8 @@ const ENV_KEYS = [ "NEXT_AI_DRAWIO_DESKTOP", "SGLANG_API_KEY", "SGLANG_BASE_URL", + "AZURE_RESOURCE_NAME", + "OLLAMA_BASE_URL", ] const savedEnv: Record = {} @@ -173,6 +180,8 @@ describe("Bedrock admin panel credentials", () => { region: "ap-northeast-1", accessKeyId: "client-id", secretAccessKey: "client-secret", + // The SDK would otherwise use the server's AWS_BEARER_TOKEN_BEDROCK + apiKey: "", }) }) @@ -312,6 +321,42 @@ describe("whose keys a request uses", () => { expect(provider.chat).not.toHaveBeenCalled() }) + it("uses Chat Completions for any configured base URL", () => { + // The settings form fills in the official URL for a new provider + getAIModel({ + provider: "openai", + apiKey: "user-key", + baseUrl: "https://api.openai.com/v1", + modelId: "gpt-5.5", + }) + const provider = vi.mocked(createOpenAI).mock.results.at(-1)?.value + expect(provider.chat).toHaveBeenCalledWith("gpt-5.5") + }) + + it("sends a user's Ollama key to Ollama Cloud, not the server's Ollama", async () => { + process.env.OLLAMA_BASE_URL = "http://ollama.internal:11434/api" + const { createOllama } = await import("ollama-ai-provider-v2") + getAIModel({ provider: "ollama", apiKey: "user-key", modelId: "m" }) + expect(createOllama).toHaveBeenLastCalledWith( + expect.objectContaining({ baseURL: "https://ollama.com/api" }), + ) + // Without a key: the server's Ollama + getAIModel({ provider: "ollama", modelId: "m" }) + expect(createOllama).toHaveBeenLastCalledWith( + expect.objectContaining({ + baseURL: "http://ollama.internal:11434/api", + }), + ) + }) + + it("needs a base URL with a user's Azure key", () => { + // The SDK would otherwise read the server's AZURE_RESOURCE_NAME + process.env.AZURE_RESOURCE_NAME = "operator-resource" + expect(() => + getAIModel({ provider: "azure", apiKey: "k", modelId: "gpt-4o" }), + ).toThrow(/base URL/) + }) + it("needs a base URL for SGLang instead of using 127.0.0.1", () => { expect(() => getAIModel({ provider: "sglang", apiKey: "k", modelId: "m" }), diff --git a/tests/unit/chat-route-quota.test.ts b/tests/unit/chat-route-quota.test.ts index 8bb65c28..29a00c69 100644 --- a/tests/unit/chat-route-quota.test.ts +++ b/tests/unit/chat-route-quota.test.ts @@ -20,7 +20,15 @@ vi.mock("@/lib/dynamo-quota-manager", () => ({ import { POST as chat } from "@/app/api/chat/route" -const ENV = ["AI_PROVIDER", "AI_MODEL", "OPENAI_API_KEY"] +const ENV = [ + "AI_PROVIDER", + "AI_MODEL", + "OPENAI_API_KEY", + "OLLAMA_BASE_URL", + "OLLAMA_API_KEY", + "AI_GATEWAY_API_KEY", + "ALLOW_PRIVATE_URLS", +] const saved: Record = {} beforeEach(() => { @@ -88,4 +96,44 @@ describe("chat quota", () => { expect(res.status).not.toBe(429) expect(quota.checks).toBe(0) }) + + it("counts the server's keyless Ollama and EdgeOne", async () => { + process.env.AI_PROVIDER = "ollama" + process.env.AI_MODEL = "llama3.2" + process.env.OLLAMA_BASE_URL = "http://ollama.internal:11434/api" + expect((await send({})).status).toBe(429) + expect( + ( + await send({ + "x-ai-provider": "edgeone", + "x-ai-model": "@tx/deepseek-ai/deepseek-v3-0324", + }) + ).status, + ).toBe(429) + expect(quota.checks).toBe(2) + }) + + it("does not count Ollama on the user's own server", async () => { + const res = await send({ + "x-ai-provider": "ollama", + "x-ai-base-url": "https://ollama.example.com/api", + "x-ai-model": "llama3.2", + }) + expect(res.status).not.toBe(429) + expect(quota.checks).toBe(0) + }) +}) + +describe("server model allowlist", () => { + it("runs AI_MODEL only on the server's AI_PROVIDER", async () => { + // Another provider's server key must not run it + process.env.AI_GATEWAY_API_KEY = "server-gateway-key" + const res = await send({ + "x-ai-provider": "gateway", + "x-ai-model": "gpt-5.5", + }) + expect(res.status).toBe(400) + expect(await res.text()).toMatch(/not available on this server/) + expect(quota.checks).toBe(0) + }) }) diff --git a/tests/unit/config-manager.test.ts b/tests/unit/config-manager.test.ts new file mode 100644 index 00000000..9a125b87 --- /dev/null +++ b/tests/unit/config-manager.test.ts @@ -0,0 +1,67 @@ +// @vitest-environment node +import { existsSync, mkdtempSync, readdirSync, writeFileSync } from "node:fs" +import { tmpdir } from "node:os" +import { join } from "node:path" +import { beforeEach, describe, expect, it, vi } from "vitest" + +const userData = vi.hoisted(() => ({ dir: "" })) +vi.mock("electron", () => ({ + app: { getPath: () => userData.dir }, + safeStorage: { isEncryptionAvailable: () => false }, +})) + +// Make the next read of the presets file fail, like a file an antivirus +// scanner holds on Windows +const readFails = vi.hoisted(() => ({ next: false })) +vi.mock("node:fs", async (importOriginal) => { + const fs = await importOriginal() + return { + ...fs, + readFileSync: ((...args: Parameters) => { + if (readFails.next && String(args[0]).endsWith(".json")) { + readFails.next = false + throw Object.assign(new Error("EBUSY: resource busy"), { + code: "EBUSY", + }) + } + return fs.readFileSync(...args) + }) as typeof fs.readFileSync, + } +}) + +import { createPreset, loadPresets } from "@/electron/main/config-manager" + +const presetsFile = () => join(userData.dir, "config-presets.json") + +beforeEach(() => { + userData.dir = mkdtempSync(join(tmpdir(), "config-manager-")) + readFails.next = false +}) + +describe("config presets file", () => { + it("keeps a file it could not read for now", () => { + createPreset({ name: "Mine", config: { AI_PROVIDER: "openai" } }) + + readFails.next = true + expect(loadPresets().presets).toEqual([]) + expect(existsSync(presetsFile())).toBe(true) + + // A save based on that empty read must not replace the presets + readFails.next = true + expect(() => + createPreset({ name: "New", config: { AI_PROVIDER: "openai" } }), + ).toThrow() + expect(loadPresets().presets.map((p) => p.name)).toEqual(["Mine"]) + }) + + it("moves a file that is not JSON aside", () => { + writeFileSync(presetsFile(), "{not json") + expect(loadPresets().presets).toEqual([]) + expect(existsSync(presetsFile())).toBe(false) + expect( + readdirSync(userData.dir).some((f) => + f.startsWith("config-presets.json.corrupt-"), + ), + ).toBe(true) + }) +}) diff --git a/tests/unit/edgeone-function.test.ts b/tests/unit/edgeone-function.test.ts index e0d96491..972a2b01 100644 --- a/tests/unit/edgeone-function.test.ts +++ b/tests/unit/edgeone-function.test.ts @@ -26,6 +26,21 @@ describe("EdgeOne chat completions function", () => { env: {}, }) expect(res.status).toBe(400) + // A plain-text type that only mentions JSON needs no CORS preflight + const disguised = await onRequest({ + request: request({ + "Content-Type": "text/plain; x=application/json", + }), + env: {}, + }) + expect(disguised.status).toBe(400) + const withCharset = await onRequest({ + request: request({ + "Content-Type": "application/json; charset=utf-8", + }), + env: {}, + }) + expect(withCharset.status).toBe(200) }) it("checks the access code when ACCESS_CODE_LIST is set", async () => { diff --git a/tests/unit/llm-errors.test.ts b/tests/unit/llm-errors.test.ts index 113b22d7..9105eb54 100644 --- a/tests/unit/llm-errors.test.ts +++ b/tests/unit/llm-errors.test.ts @@ -229,6 +229,27 @@ describe("streamErrorText", () => { message: "bad key", }) }) + + it("classifies a provider error sent as plain text", () => { + // DeepSeek's SDK sends errors in the stream as a string + const text = "Insufficient Balance for account 42" + expect(JSON.parse(streamErrorText(text))).toEqual({ + type: "provider", + code: "insufficient_quota", + message: text, + }) + expect(JSON.parse(streamErrorText(text, true)).message).not.toMatch( + /account 42/, + ) + }) + + it("classifies Bedrock's throttling sent in the stream", () => { + // Bedrock's ThrottlingException as a plain object, not an API error + const throttled = { + message: "Too many tokens, please wait before trying again.", + } + expect(JSON.parse(streamErrorText(throttled)).code).toBe("rate_limited") + }) }) describe("isToolCallError", () => { diff --git a/tests/unit/port-manager.test.ts b/tests/unit/port-manager.test.ts index 42b2f0bb..9016e2f1 100644 --- a/tests/unit/port-manager.test.ts +++ b/tests/unit/port-manager.test.ts @@ -1,5 +1,5 @@ // @vitest-environment node -import { mkdtempSync, readFileSync, writeFileSync } from "node:fs" +import { existsSync, mkdtempSync, readFileSync, writeFileSync } from "node:fs" import { tmpdir } from "node:os" import { join } from "node:path" import { beforeEach, describe, expect, it, vi } from "vitest" @@ -9,29 +9,88 @@ vi.mock("electron", () => ({ app: { isPackaged: true, getPath: () => userData.dir }, })) -import { saveServerPort } from "@/electron/main/port-manager" +// Ports that fail to listen, with the error code +const busy = vi.hoisted(() => ({ ports: {} as Record })) +vi.mock("node:net", () => ({ + default: { + createServer: () => { + const handlers: Record void> = {} + const server = { + once: (event: string, cb: (arg?: unknown) => void) => { + handlers[event] = cb + return server + }, + listen: (port: number) => { + const code = busy.ports[port] + if (code) handlers.error?.({ code }) + else handlers.listening?.() + }, + close: () => {}, + } + return server + }, + }, +})) -const savedPort = () => - JSON.parse(readFileSync(join(userData.dir, "server-port.json"), "utf-8")) - .port +import { + findAvailablePort, + resetAllocatedPort, + saveServerPort, +} from "@/electron/main/port-manager" + +const portFile = () => join(userData.dir, "server-port.json") +const savedPort = () => JSON.parse(readFileSync(portFile(), "utf-8")).port + +/** One launch: pick a port, start on it, remember it if it should be */ +async function launch() { + const port = await findAvailablePort(false) + saveServerPort(port) + return port +} beforeEach(() => { userData.dir = mkdtempSync(join(tmpdir(), "port-manager-")) + busy.ports = {} + resetAllocatedPort() }) describe("saveServerPort", () => { - it("remembers the port of the first launch", () => { - saveServerPort(13370) + it("remembers the legacy port of the first launch", async () => { + expect(await launch()).toBe(61337) + expect(savedPort()).toBe(61337) + }) + + it("remembers 13370 when the system reserves the legacy port", async () => { + // Windows excludes port ranges for Hyper-V, which can change on + // each boot + busy.ports[61337] = "EACCES" + expect(await launch()).toBe(13370) expect(savedPort()).toBe(13370) }) + it("does not remember a port used while the legacy port was in use", async () => { + // For example the previous version still quitting after an update: + // the user's data is under 61337, so the next launch tries it again + busy.ports[61337] = "EADDRINUSE" + expect(await launch()).toBe(13370) + expect(existsSync(portFile())).toBe(false) + + busy.ports = {} + expect(await launch()).toBe(61337) + expect(savedPort()).toBe(61337) + }) + + it("does not remember a fallback port", async () => { + busy.ports[61337] = "EACCES" + busy.ports[13370] = "EADDRINUSE" + expect(await launch()).toBe(13371) + expect(existsSync(portFile())).toBe(false) + }) + it("keeps the remembered port when a launch had to use another", () => { // The app's data lives under the remembered port's origin; going // back to it once it is free brings the chats and settings back - writeFileSync( - join(userData.dir, "server-port.json"), - JSON.stringify({ port: 61337 }), - ) + writeFileSync(portFile(), JSON.stringify({ port: 61337 })) saveServerPort(13371) expect(savedPort()).toBe(61337) }) diff --git a/tests/unit/provider-models.test.ts b/tests/unit/provider-models.test.ts index a5d4214b..19ad7248 100644 --- a/tests/unit/provider-models.test.ts +++ b/tests/unit/provider-models.test.ts @@ -118,6 +118,18 @@ describe("listProviderModels", () => { ]) }) + it("lists Ollama Cloud with a user's key, like chat", async () => { + // The user's key must not go to the server's Ollama + const { fn, calls } = answer({ models: [{ name: "gpt-oss:120b" }] }) + process.env.OLLAMA_BASE_URL = "http://ollama.internal:11434" + try { + await listProviderModels("ollama", { apiKey: "user-key" }, fn) + } finally { + delete process.env.OLLAMA_BASE_URL + } + expect(calls[0].url).toBe("https://ollama.com/api/tags") + }) + it("does not use SGLang's local address as a default", async () => { const { fn, calls } = answer({ data: [] }) await expect( diff --git a/tests/unit/validate-model-route.test.ts b/tests/unit/validate-model-route.test.ts index 0d7f03d3..b7990613 100644 --- a/tests/unit/validate-model-route.test.ts +++ b/tests/unit/validate-model-route.test.ts @@ -1,9 +1,13 @@ // @vitest-environment node import { streamText } from "ai" import { afterEach, describe, expect, it, vi } from "vitest" +import { POST as testModel } from "@/app/api/admin/test-model/route" import { POST as validateModel } from "@/app/api/validate-model/route" import { getAIModel } from "@/lib/ai-providers" +// No saved admin providers +vi.mock("@/lib/admin/settings", () => ({ loadSettings: () => ({}) })) + // Treat every URL as public so no test hits DNS vi.mock("@/lib/ssrf-protection", async (importOriginal) => ({ ...(await importOriginal()), @@ -158,3 +162,37 @@ describe("chat requests to a client base URL", () => { expect(String(error)).toMatch(/Redirects are not allowed/) }) }) + +describe("the admin panel's Test button", () => { + it("works when access codes are set", async () => { + // The admin password stands in for the visitor access code + process.env.ACCESS_CODE_LIST = "visitor-code" + process.env.ADMIN_PASSWORD = "admin-pw" + try { + streamReply({ role: "assistant", content: "OK" }) + const res = await testModel( + new Request("http://localhost/api/admin/test-model", { + method: "POST", + headers: { + "Content-Type": "application/json", + "x-admin-password": "admin-pw", + }, + body: JSON.stringify({ + provider: { + id: "p1", + provider: "glm", + apiKey: "key", + models: ["glm-5"], + }, + modelId: "glm-5", + }), + }), + ) + expect(res.status).toBe(200) + expect((await res.json()).valid).toBe(true) + } finally { + delete process.env.ACCESS_CODE_LIST + delete process.env.ADMIN_PASSWORD + } + }) +})