From 0855b35ff233a5c74a2f034d03ed0c316e1f1dc9 Mon Sep 17 00:00:00 2001 From: "dayuan.jiang" Date: Sun, 4 Oct 2026 23:03:06 +0900 Subject: [PATCH] fix(server): count quota by the key actually used, and more review fixes Found by the PR review, each with a test that failed first: - Quota: any key header skipped it, even one the provider never reads (x-aws-access-key-id with OpenAI), so a request ran on the server's key without being counted. The check now runs after the model is resolved and uses usesServerCredentials. On main already. - usesServerCredentials read the raw base URL; "/" cleans up to none, so an Ollama request ran on the server's key past the server-model check. - SGLang's default 127.0.0.1:8000 only fills the settings form. Chat and the model list used it as a real address, so the server called its own machine even with private URLs blocked. Now a base URL is required. - With a user's OpenAI key and no base URL, the SDK read the server's OPENAI_BASE_URL. The official endpoint is now passed. On main already. - The Test button refused nothing on the server's keys (Ollama Cloud), and a 15 s timeout reported "connected, no tool call". - The model list for Ollama without a base URL came from ollama.com while chat went to the server's Ollama. - Bedrock's "Too many tokens, please wait" counted as context too long. - On the server's keys the provider's error text stays in the server log; it can name the server's AWS account, role or internal hosts. - Desktop app: the preset keys are the user's own (NEXT_AI_DRAWIO_DESKTOP), so Max Output Tokens can be raised and keyless models in settings work again. A launch that found the remembered port taken no longer replaces it, which hid the user's chats and settings for good. --- app/api/chat/route.ts | 70 +++++++--------- app/api/validate-model/route.ts | 29 ++++++- electron/main/next-server.ts | 2 + electron/main/port-manager.ts | 6 +- lib/ai-providers.ts | 27 +++++- lib/llm-errors.ts | 14 +++- lib/provider-models.ts | 22 ++++- scripts/electron-dev.mjs | 3 +- tests/unit/ai-providers-credentials.test.ts | 46 +++++++++++ tests/unit/chat-route-quota.test.ts | 91 +++++++++++++++++++++ tests/unit/llm-errors.test.ts | 21 +++++ tests/unit/port-manager.test.ts | 38 +++++++++ tests/unit/provider-models.test.ts | 23 ++++++ tests/unit/validate-model-route.test.ts | 44 ++++++++++ 14 files changed, 382 insertions(+), 54 deletions(-) create mode 100644 tests/unit/chat-route-quota.test.ts create mode 100644 tests/unit/port-manager.test.ts diff --git a/app/api/chat/route.ts b/app/api/chat/route.ts index cececdf5..f957aae6 100644 --- a/app/api/chat/route.ts +++ b/app/api/chat/route.ts @@ -133,36 +133,6 @@ async function handleChatRequest(req: Request): Promise { userId: userId, }) - // === SERVER-SIDE QUOTA CHECK START === - // Quota is opt-in: only enabled when DYNAMODB_QUOTA_TABLE env var is set - const hasOwnApiKey = !!( - req.headers.get("x-ai-provider") && - (req.headers.get("x-ai-api-key") || - req.headers.get("x-aws-access-key-id") || - req.headers.get("x-vertex-api-key")) - ) - - // Skip quota check if: quota disabled, user has own API key, or is anonymous - if (isQuotaEnabled() && !hasOwnApiKey && userId !== "anonymous") { - const quotaCheck = await checkAndIncrementRequest(userId, { - requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10, - tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000, - tpm: Number(process.env.TPM_LIMIT) || 20000, - }) - if (!quotaCheck.allowed) { - return Response.json( - { - error: quotaCheck.error, - type: quotaCheck.type, - used: quotaCheck.used, - limit: quotaCheck.limit, - }, - { status: 429 }, - ) - } - } - // === SERVER-SIDE QUOTA CHECK END === - // === FILE VALIDATION START === const fileValidation = validateFileParts(messages) if (!fileValidation.valid) { @@ -292,14 +262,41 @@ 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. + const countsQuota = + isQuotaEnabled() && onServerCredentials && userId !== "anonymous" + if (countsQuota) { + const quotaCheck = await checkAndIncrementRequest(userId, { + requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10, + tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000, + tpm: Number(process.env.TPM_LIMIT) || 20000, + }) + if (!quotaCheck.allowed) { + return Response.json( + { + error: quotaCheck.error, + type: quotaCheck.type, + used: quotaCheck.used, + limit: quotaCheck.limit, + }, + { status: 429 }, + ) + } + } + // === SERVER-SIDE QUOTA CHECK END === + // Retry once if the provider rejects the requested budget, or (newer // Claude models) the sampling or thinking settings const model = withOutputTokenLimitFallback( withDeprecatedParamsFallback(baseModel), ) - // The user setting can raise the budget only on their own key (desktop users - // can still raise it themselves); on the server's keys it can only lower it + // The user setting can raise the budget only on their own key (in the + // desktop app every key is the user's); on the server's keys it can only + // lower it const maxOutputTokens = resolveMaxOutputTokens( req.headers.get("x-max-output-tokens"), onServerCredentials, @@ -587,12 +584,7 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on // Record token usage for server-side quota tracking (if enabled) // Use totalUsage (cumulative across all steps) instead of usage (final step only) // inputTokens already includes cache reads and writes in AI SDK 6 - if ( - isQuotaEnabled() && - !hasOwnApiKey && - userId !== "anonymous" && - totalUsage - ) { + if (countsQuota && totalUsage) { const totalTokens = (totalUsage.inputTokens || 0) + (totalUsage.outputTokens || 0) @@ -724,7 +716,7 @@ Call this tool to get shape names and usage syntax for a specific library.`, const response = result.toUIMessageStreamResponse({ sendReasoning: true, - onError: streamErrorText, + onError: (error) => streamErrorText(error, onServerCredentials), messageMetadata: ({ part }) => { if (part.type === "finish") { const usage = (part as any).totalUsage diff --git a/app/api/validate-model/route.ts b/app/api/validate-model/route.ts index 1f103781..c8567c1e 100644 --- a/app/api/validate-model/route.ts +++ b/app/api/validate-model/route.ts @@ -2,14 +2,15 @@ import { streamText, tool } from "ai" import { NextResponse } from "next/server" import { z } from "zod" import { checkAccessCode } from "@/lib/access-code" -import { getAIModel } from "@/lib/ai-providers" +import { getAIModel, usesServerCredentials } from "@/lib/ai-providers" import { classifyLLMError } from "@/lib/llm-errors" import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection" +import type { ProviderName } from "@/lib/types/model-config" export const runtime = "nodejs" interface ValidateRequest { - provider: string + provider: ProviderName apiKey: string baseUrl?: string modelId: string @@ -93,6 +94,22 @@ export async function POST(req: Request) { { status: 400 }, ) } + // The Test button checks the user's own provider. On the server's + // keys (Ollama Cloud without a key or URL) anyone could run any model. + if ( + usesServerCredentials(provider, { + apiKey, + baseUrl, + awsAccessKeyId, + awsSecretAccessKey, + vertexApiKey, + }) + ) { + return NextResponse.json( + { valid: false, error: "API key is required" }, + { status: 400 }, + ) + } // The same model the chat would use. A client base URL makes it // refuse redirects to internal hosts. @@ -130,6 +147,14 @@ export async function POST(req: Request) { let finishReason: string | undefined for await (const part of result.fullStream) { if (part.type === "error") throw part.error + // The timeout ends the stream with an abort part, not an error + if (part.type === "abort") { + const timeout = new Error( + `The model did not answer within ${TEST_TIMEOUT_MS / 1000} s.`, + ) + timeout.name = "TimeoutError" + throw timeout + } if (part.type === "tool-call") { calledTool = true break diff --git a/electron/main/next-server.ts b/electron/main/next-server.ts index a9847947..739fedfb 100644 --- a/electron/main/next-server.ts +++ b/electron/main/next-server.ts @@ -87,6 +87,8 @@ async function startServer(): Promise { HOSTNAME: "127.0.0.1", // Enable Node.js built-in proxy support for fetch (Node.js 24+) NODE_USE_ENV_PROXY: "1", + // The preset keys are the user's own, not a server's + NEXT_AI_DRAWIO_DESKTOP: "1", } // Keep requests to local model servers (e.g. Ollama) off the proxy diff --git a/electron/main/port-manager.ts b/electron/main/port-manager.ts index f0c5c2ac..f9e8cf7c 100644 --- a/electron/main/port-manager.ts +++ b/electron/main/port-manager.ts @@ -44,10 +44,12 @@ function loadSavedPort(): number | null { } /** - * Remember the port the production server started on + * 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. */ export function saveServerPort(port: number): void { - if (!app.isPackaged || port === loadSavedPort()) { + if (!app.isPackaged || loadSavedPort() !== null) { return } try { diff --git a/lib/ai-providers.ts b/lib/ai-providers.ts index ec39f36a..b1c3f2ea 100644 --- a/lib/ai-providers.ts +++ b/lib/ai-providers.ts @@ -682,7 +682,8 @@ function createModel( // 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 + return e.baseURL && + e.baseURL !== PROVIDER_INFO.openai.defaultBaseUrl ? openaiProvider.chat(modelId) : openaiProvider(modelId) } @@ -994,14 +995,28 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { provider === "gateway" ? "AI_GATEWAY_BASE_URL" : `${provider.toUpperCase()}_BASE_URL` + // A local default (SGLang's 127.0.0.1) only fills the settings + // form; the server must not call its own machine for it. With a + // user's key the OpenAI SDK would read the server's + // OPENAI_BASE_URL, so name the official endpoint. + const defaultUrl = PROVIDER_INFO[provider].defaultBaseUrl + const publicDefault = defaultUrl?.startsWith("https://") + ? defaultUrl + : undefined const baseURL = resolveBaseURL( overrides?.apiKey, overrides?.baseUrl, resolveBaseUrlEnv(overrides, baseUrlEnv), - SDK_KNOWS_ENDPOINT.has(provider) + SDK_KNOWS_ENDPOINT.has(provider) && + !(provider === "openai" && overrides?.apiKey) ? undefined - : PROVIDER_INFO[provider].defaultBaseUrl, + : publicDefault, ) + if (!baseURL && !SDK_KNOWS_ENDPOINT.has(provider)) { + throw new Error( + `${PROVIDER_INFO[provider].label} needs a base URL. Add it in the model settings.`, + ) + } model = createModel(provider, modelId, { apiKey, baseURL, @@ -1032,6 +1047,10 @@ export function usesServerCredentials( provider: ProviderName, overrides?: ClientOverrides, ): boolean { + // The desktop app's local server holds the user's own preset keys + if (process.env.NEXT_AI_DRAWIO_DESKTOP === "1") return false + // Cleaned like getAIModel does: "/" means no base URL + const baseUrl = normalizeBaseUrl(overrides?.baseUrl ?? "") switch (provider) { case "bedrock": return !(overrides?.awsAccessKeyId && overrides?.awsSecretAccessKey) @@ -1044,7 +1063,7 @@ export function usesServerCredentials( // Only a server key costs money; a keyless local server or the // client's own server does not return ( - !overrides?.baseUrl && + !baseUrl && !overrides?.apiKey && !!(overrides?.apiKeyEnv || process.env.OLLAMA_API_KEY) ) diff --git a/lib/llm-errors.ts b/lib/llm-errors.ts index 87fb15d0..077cc2f5 100644 --- a/lib/llm-errors.ts +++ b/lib/llm-errors.ts @@ -36,7 +36,8 @@ export interface LLMError { // error can come as 403 or 429, a context or image error as a plain 400 const SPECIFIC_TEXTS: Array<[RegExp, LLMErrorCode]> = [ [ - /context length|context window|maximum context|prompt is too long|input is too long|too many (?:input )?tokens/i, + // Not "too many tokens": that is Bedrock's throttling message + /context length|context window|maximum context|prompt is too long|input is too long|too many input tokens/i, "context_too_long", ], [ @@ -110,12 +111,19 @@ function problemDetail(body: string): string | undefined { /** * The error text for the chat stream: what went wrong with the provider as * JSON for the hint, or the text the model must read to fix a tool call. + * On the server's keys the provider's own text stays in the server log: + * it can name the server's account, role or internal hosts. */ -export function streamErrorText(error: unknown): string { +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 if (isToolCallError(error)) return (error as Error).message - return JSON.stringify(classifyLLMError(error)) + const classified = classifyLLMError(error) + if (hideDetails) { + console.error("[chat] Provider error:", error) + classified.message = "The provider returned an error." + } + return JSON.stringify(classified) } /** diff --git a/lib/provider-models.ts b/lib/provider-models.ts index 4f53df4c..b64fc805 100644 --- a/lib/provider-models.ts +++ b/lib/provider-models.ts @@ -75,6 +75,19 @@ async function getJson( return response.json() } +/** + * 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. + */ +function listFallbackUrl(provider: ProviderName): string { + if (provider === "ollama") { + return process.env.OLLAMA_BASE_URL || "http://127.0.0.1:11434/api" + } + const url = PROVIDER_INFO[provider].defaultBaseUrl + return url?.startsWith("https://") ? url : "" +} + /** * The provider's chat models, with tool support from the provider's own * data or else models.dev. Only the client's key is used, so the server's @@ -85,9 +98,7 @@ export async function listProviderModels( { apiKey, baseUrl }: { apiKey?: string; baseUrl?: string }, fetchFn: typeof fetch = fetch, ): Promise { - const base = normalizeBaseUrl( - baseUrl || PROVIDER_INFO[provider].defaultBaseUrl || "", - ) + const base = normalizeBaseUrl(baseUrl || listFallbackUrl(provider)) const bearer: Record = apiKey ? { Authorization: `Bearer ${apiKey}` } : {} @@ -170,6 +181,11 @@ export async function listProviderModels( break } default: { + if (!base) { + throw new Error( + `${PROVIDER_INFO[provider].label} needs a base URL to list its models.`, + ) + } const data = await getJson(`${base}/models`, bearer, fetchFn) models = (data.data ?? []) .map((m: { id: string }) => ({ id: m.id })) diff --git a/scripts/electron-dev.mjs b/scripts/electron-dev.mjs index 98b2463a..0ae5cddc 100644 --- a/scripts/electron-dev.mjs +++ b/scripts/electron-dev.mjs @@ -146,7 +146,8 @@ function killProcess(proc) { * Start Next.js dev server with preset environment */ function startNextServer(presetEnv) { - const env = { ...process.env } + // The preset keys are the user's own, not a server's + const env = { ...process.env, NEXT_AI_DRAWIO_DESKTOP: "1" } // Apply preset environment variables if (presetEnv) { diff --git a/tests/unit/ai-providers-credentials.test.ts b/tests/unit/ai-providers-credentials.test.ts index 50703a83..7b13eae8 100644 --- a/tests/unit/ai-providers-credentials.test.ts +++ b/tests/unit/ai-providers-credentials.test.ts @@ -1,3 +1,4 @@ +import { createOpenAI } from "@ai-sdk/openai" import { afterEach, beforeEach, describe, expect, it, vi } from "vitest" import { getAIModel, @@ -56,6 +57,9 @@ const ENV_KEYS = [ "AI_PROVIDER", "AI_MODEL", "VALIDATION_MODEL", + "NEXT_AI_DRAWIO_DESKTOP", + "SGLANG_API_KEY", + "SGLANG_BASE_URL", ] const savedEnv: Record = {} @@ -272,3 +276,45 @@ describe("getValidationModel", () => { expect(createOpenRouter).toHaveBeenCalledWith({ apiKey: "env-key" }) }) }) + +describe("whose keys a request uses", () => { + it("cleans the base URL like the request does", () => { + // "/" and a pasted path clean up to no base URL: the server's Ollama + process.env.OLLAMA_API_KEY = "server-ollama-key" + expect(usesServerCredentials("ollama", { baseUrl: "/" })).toBe(true) + expect( + usesServerCredentials("ollama", { baseUrl: "/chat/completions" }), + ).toBe(true) + }) + + it("counts the desktop app's keys as the user's own", () => { + // Electron passes the user's preset keys as server env vars + process.env.NEXT_AI_DRAWIO_DESKTOP = "1" + expect(usesServerCredentials("openai", {})).toBe(false) + }) + + it("sends a user's OpenAI key to the official endpoint", () => { + // The SDK would otherwise read the server's OPENAI_BASE_URL + process.env.OPENAI_BASE_URL = "https://operator-proxy.example.com/v1" + getAIModel({ + provider: "openai", + apiKey: "user-key", + modelId: "gpt-5.5", + }) + expect(createOpenAI).toHaveBeenLastCalledWith( + expect.objectContaining({ + apiKey: "user-key", + baseURL: "https://api.openai.com/v1", + }), + ) + // Still the Responses API, like without a base URL + const provider = vi.mocked(createOpenAI).mock.results.at(-1)?.value + expect(provider.chat).not.toHaveBeenCalled() + }) + + it("needs a base URL for SGLang instead of using 127.0.0.1", () => { + expect(() => + getAIModel({ provider: "sglang", apiKey: "k", modelId: "m" }), + ).toThrow(/base URL/) + }) +}) diff --git a/tests/unit/chat-route-quota.test.ts b/tests/unit/chat-route-quota.test.ts new file mode 100644 index 00000000..8bb65c28 --- /dev/null +++ b/tests/unit/chat-route-quota.test.ts @@ -0,0 +1,91 @@ +// @vitest-environment node +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest" + +// Quota on, and every check answers that the daily limit is used up +const quota = vi.hoisted(() => ({ checks: 0 })) +vi.mock("@/lib/dynamo-quota-manager", () => ({ + isQuotaEnabled: () => true, + checkAndIncrementRequest: async () => { + quota.checks++ + return { + allowed: false, + error: "Daily limit reached", + type: "request", + used: 10, + limit: 10, + } + }, + recordTokenUsage: async () => {}, +})) + +import { POST as chat } from "@/app/api/chat/route" + +const ENV = ["AI_PROVIDER", "AI_MODEL", "OPENAI_API_KEY"] +const saved: Record = {} + +beforeEach(() => { + for (const k of ENV) saved[k] = process.env[k] + process.env.AI_PROVIDER = "openai" + process.env.AI_MODEL = "gpt-5.5" + process.env.OPENAI_API_KEY = "server-key" + quota.checks = 0 + // No request may reach a provider + vi.stubGlobal( + "fetch", + vi.fn(async () => { + throw new Error("no network in tests") + }), + ) +}) + +afterEach(() => { + for (const k of ENV) { + if (saved[k] === undefined) delete process.env[k] + else process.env[k] = saved[k] + } + vi.unstubAllGlobals() +}) + +const send = (headers: Record) => + chat( + new Request("http://localhost/api/chat", { + method: "POST", + headers: { + "Content-Type": "application/json", + "x-forwarded-for": "203.0.113.7", + ...headers, + }, + body: JSON.stringify({ + messages: [ + { + id: "u1", + role: "user", + parts: [{ type: "text", text: "Draw two boxes" }], + }, + ], + xml: "", + }), + }), + ) + +describe("chat quota", () => { + it("counts a request whose key header the provider never reads", async () => { + // OpenAI ignores the AWS key, so this runs on the server's key + const res = await send({ + "x-ai-provider": "openai", + "x-aws-access-key-id": "x", + }) + expect(res.status).toBe(429) + expect(quota.checks).toBe(1) + }) + + it("does not count a request on the user's own key", async () => { + const res = await send({ + "x-ai-provider": "openai", + "x-ai-api-key": "user-key", + "x-ai-model": "gpt-5.5", + }) + expect(res.status).not.toBe(429) + expect(quota.checks).toBe(0) + }) +}) diff --git a/tests/unit/llm-errors.test.ts b/tests/unit/llm-errors.test.ts index 07be0d0b..113b22d7 100644 --- a/tests/unit/llm-errors.test.ts +++ b/tests/unit/llm-errors.test.ts @@ -110,6 +110,14 @@ describe("classifyLLMError", () => { expect(classifyLLMError(error).code).toBe("model_not_found") }) + it("reads Bedrock's token throttling as a rate limit", () => { + const error = apiError( + 429, + "Too many tokens, please wait before trying again.", + ) + expect(classifyLLMError(error).code).toBe("rate_limited") + }) + it("names a network error the SDK wrapped", () => { const error = new APICallError({ message: @@ -201,6 +209,19 @@ describe("streamErrorText", () => { } }) + it("hides the provider's text on the server's keys", () => { + const error = apiError( + 403, + "User: arn:aws:sts::123456789012:assumed-role/app/s is not authorized to perform: bedrock:InvokeModel", + ) + const hidden = JSON.parse(streamErrorText(error, true)) + expect(hidden.code).toBe("forbidden") + expect(hidden.message).not.toMatch(/arn:aws|123456789012/) + expect(JSON.parse(streamErrorText(error)).message).toMatch( + /not authorized/, + ) + }) + it("classifies a provider error", () => { expect(JSON.parse(streamErrorText(apiError(401, "bad key")))).toEqual({ type: "provider", diff --git a/tests/unit/port-manager.test.ts b/tests/unit/port-manager.test.ts new file mode 100644 index 00000000..42b2f0bb --- /dev/null +++ b/tests/unit/port-manager.test.ts @@ -0,0 +1,38 @@ +// @vitest-environment node +import { mkdtempSync, readFileSync, 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: { isPackaged: true, getPath: () => userData.dir }, +})) + +import { saveServerPort } from "@/electron/main/port-manager" + +const savedPort = () => + JSON.parse(readFileSync(join(userData.dir, "server-port.json"), "utf-8")) + .port + +beforeEach(() => { + userData.dir = mkdtempSync(join(tmpdir(), "port-manager-")) +}) + +describe("saveServerPort", () => { + it("remembers the port of the first launch", () => { + saveServerPort(13370) + expect(savedPort()).toBe(13370) + }) + + 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 }), + ) + saveServerPort(13371) + expect(savedPort()).toBe(61337) + }) +}) diff --git a/tests/unit/provider-models.test.ts b/tests/unit/provider-models.test.ts index 3cae9ed5..a5d4214b 100644 --- a/tests/unit/provider-models.test.ts +++ b/tests/unit/provider-models.test.ts @@ -103,6 +103,29 @@ describe("listProviderModels", () => { ]) }) + it("lists Ollama from where chat goes without a base URL", async () => { + const { fn, calls } = answer({ models: [{ name: "llama3.2" }] }) + process.env.OLLAMA_BASE_URL = "http://ollama.internal:11434" + try { + await listProviderModels("ollama", {}, fn) + } finally { + delete process.env.OLLAMA_BASE_URL + } + await listProviderModels("ollama", {}, fn) + expect(calls.map((c) => c.url)).toEqual([ + "http://ollama.internal:11434/api/tags", + "http://127.0.0.1:11434/api/tags", + ]) + }) + + it("does not use SGLang's local address as a default", async () => { + const { fn, calls } = answer({ data: [] }) + await expect( + listProviderModels("sglang", { apiKey: "k" }, fn), + ).rejects.toThrow(/base URL/) + expect(calls).toHaveLength(0) + }) + it("turns a failed request into an error with its status", async () => { const { fn } = answer({ error: "bad key" }, 401) await expect( diff --git a/tests/unit/validate-model-route.test.ts b/tests/unit/validate-model-route.test.ts index 8f8d271f..0d7f03d3 100644 --- a/tests/unit/validate-model-route.test.ts +++ b/tests/unit/validate-model-route.test.ts @@ -74,6 +74,50 @@ describe("POST /api/validate-model", () => { expect(typeof data.responseTime).toBe("number") }) + it("reports a model that did not answer in time", async () => { + // The 15 s timeout has fired: the SDK ends the stream with an + // abort part instead of throwing + const timedOut = AbortSignal.abort( + new DOMException("The operation timed out.", "TimeoutError"), + ) + const timeout = vi + .spyOn(AbortSignal, "timeout") + .mockReturnValue(timedOut) + vi.stubGlobal( + "fetch", + vi.fn(async () => { + throw timedOut.reason + }), + ) + try { + const data = await testGlm() + expect(data.valid).toBe(false) + expect(data.code).toBe("timeout") + } finally { + timeout.mockRestore() + } + }) + + it("does not run on the server's keys", async () => { + process.env.OLLAMA_API_KEY = "server-ollama-key" + try { + const res = await validateModel( + new Request("http://localhost/api/validate-model", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + provider: "ollama", + modelId: "any-cloud-model", + }), + }), + ) + expect(res.status).toBe(400) + expect((await res.json()).error).toMatch(/API key/) + } finally { + delete process.env.OLLAMA_API_KEY + } + }) + it("warns when the model answers without a tool call", async () => { streamReply({ role: "assistant", content: "OK" }) const data = await testGlm()