diff --git a/app/api/chat/route.ts b/app/api/chat/route.ts index 6bb28878..a50f708b 100644 --- a/app/api/chat/route.ts +++ b/app/api/chat/route.ts @@ -10,7 +10,7 @@ import { import { jsonrepair } from "jsonrepair" import path from "path" import { z } from "zod" -import { checkAccessCode } from "@/lib/access-code" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" import { CACHE_POINT, getAIModel, @@ -100,6 +100,8 @@ const modelStreamResponses = new WeakSet() const DEBUG_LLM_PAYLOAD = process.env.DEBUG_LLM_PAYLOAD === "true" async function handleChatRequest(req: Request): Promise { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite // Check for access code const accessDenied = checkAccessCode(req) if (accessDenied) return accessDenied diff --git a/app/api/parse-url/route.ts b/app/api/parse-url/route.ts index 54755020..33a15c4b 100644 --- a/app/api/parse-url/route.ts +++ b/app/api/parse-url/route.ts @@ -1,7 +1,8 @@ import { extractFromHtml } from "@extractus/article-extractor" import { NextResponse } from "next/server" import TurndownService from "turndown" -import { checkAccessCode } from "@/lib/access-code" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" +import { readLimitedBody } from "@/lib/read-limited-body" import { isPrivateUrl } from "@/lib/ssrf-protection" const MAX_CONTENT_LENGTH = 150000 // Match PDF limit @@ -34,33 +35,9 @@ function detectCharset( } } -// Read the response body, giving up once it passes MAX_RESPONSE_BYTES so a -// huge download can't exhaust server memory. Returns null when too large. -async function readLimitedBody( - response: Response, -): Promise { - if (Number(response.headers.get("content-length")) > MAX_RESPONSE_BYTES) { - return null - } - if (!response.body) return new ArrayBuffer(0) - - const reader = response.body.getReader() - const chunks: Uint8Array[] = [] - let total = 0 - while (true) { - const { done, value } = await reader.read() - if (done) break - total += value.byteLength - if (total > MAX_RESPONSE_BYTES) { - await reader.cancel() - return null - } - chunks.push(value) - } - return new Blob(chunks as BlobPart[]).arrayBuffer() -} - export async function POST(req: Request) { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite const accessError = checkAccessCode(req) if (accessError) return accessError @@ -128,7 +105,7 @@ export async function POST(req: Request) { ) } - const buffer = await readLimitedBody(response) + const buffer = await readLimitedBody(response, MAX_RESPONSE_BYTES) if (!buffer) { return NextResponse.json( { diff --git a/app/api/provider-models/route.ts b/app/api/provider-models/route.ts index de85421b..0b6cbb8a 100644 --- a/app/api/provider-models/route.ts +++ b/app/api/provider-models/route.ts @@ -1,7 +1,11 @@ import { NextResponse } from "next/server" -import { checkAccessCode } from "@/lib/access-code" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" import { classifyLLMError } from "@/lib/llm-errors" -import { canListModels, listProviderModels } from "@/lib/provider-models" +import { + canListModels, + listProviderModels, + ModelListError, +} from "@/lib/provider-models" import { allowPrivateUrls, isPrivateUrl, @@ -24,6 +28,8 @@ const NO_KEY_NEEDED = new Set([ * so the dialog keeps its suggested models. */ export async function POST(req: Request) { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite // Sends requests to a URL the client chose, so require the access code const accessError = checkAccessCode(req) if (accessError) return accessError @@ -56,7 +62,20 @@ export async function POST(req: Request) { return NextResponse.json({ models }) } catch (error) { console.warn("[provider-models] Listing failed:", error) - const { code, message } = classifyLLMError(error) - return NextResponse.json({ code, error: message }) + // Only our own explanations go back: the URL may be an internal + // address, whose answer or host names must not reach the caller. + // The Gateway SDK wraps them, keeping ours as the cause. + const cause = (error as { cause?: unknown })?.cause + const own = + error instanceof ModelListError + ? error + : cause instanceof ModelListError + ? cause + : null + const { code } = classifyLLMError(own ?? error) + return NextResponse.json({ + code, + error: own?.message ?? "The model list request failed.", + }) } } diff --git a/app/api/validate-diagram/route.ts b/app/api/validate-diagram/route.ts index a66ef326..2f8a3c71 100644 --- a/app/api/validate-diagram/route.ts +++ b/app/api/validate-diagram/route.ts @@ -4,7 +4,7 @@ */ import { Output, streamText } from "ai" -import { checkAccessCode } from "@/lib/access-code" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" import { getValidationModel } from "@/lib/ai-providers" import { VALIDATION_SYSTEM_PROMPT } from "@/lib/validation-prompts" import { @@ -37,6 +37,8 @@ function createStreamingResponse(result: ValidationResult): Response { } export async function POST(req: Request): Promise { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite // Uses the server's model credentials, so require the access code const accessError = checkAccessCode(req) if (accessError) return accessError diff --git a/app/api/validate-model/route.ts b/app/api/validate-model/route.ts index b83aeab9..907c671e 100644 --- a/app/api/validate-model/route.ts +++ b/app/api/validate-model/route.ts @@ -1,7 +1,7 @@ import { streamText, tool } from "ai" import { NextResponse } from "next/server" import { z } from "zod" -import { checkAccessCode } from "@/lib/access-code" +import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" import { checkAdminAuth } from "@/lib/admin/auth" import { getAIModel, usesServerCredentials } from "@/lib/ai-providers" import { classifyLLMError } from "@/lib/llm-errors" @@ -35,6 +35,8 @@ 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) { + const crossSite = rejectCrossSite(req) + if (crossSite) return crossSite // 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) diff --git a/components/provider-credentials-fields.tsx b/components/provider-credentials-fields.tsx index ea3e0e60..4b97716e 100644 --- a/components/provider-credentials-fields.tsx +++ b/components/provider-credentials-fields.tsx @@ -31,7 +31,7 @@ export type SecretField = | "vertexApiKey" // AWS regions offered for Bedrock (shared by both screens) -const AWS_REGIONS: Array<[string, string]> = [ +export const AWS_REGIONS: Array<[string, string]> = [ ["us-east-1", "N. Virginia"], ["us-east-2", "Ohio"], ["us-west-2", "Oregon"], diff --git a/lib/access-code.ts b/lib/access-code.ts index f31f7ac3..ec6cb52c 100644 --- a/lib/access-code.ts +++ b/lib/access-code.ts @@ -1,3 +1,31 @@ +/** + * Refuse a POST that a page on another website could have sent. A browser + * sends a cross-site POST without asking first (CORS preflight) only with a + * text or form body, so the routes take JSON only. In the desktop app also + * refuse a foreign Host: a site that points its own domain name at + * 127.0.0.1 (DNS rebinding) is same-origin with the local server, but its + * requests carry that domain. A request the server builds itself has no + * Host. Returns the response to send, or null when the request may go on. + */ +export function rejectCrossSite(req: Request): Response | null { + const contentType = req.headers.get("content-type") ?? "" + if (!/^\s*application\/json\b/i.test(contentType)) { + return Response.json( + { error: "Content-Type must be application/json" }, + { status: 415 }, + ) + } + const host = req.headers.get("host") + if ( + process.env.NEXT_AI_DRAWIO_DESKTOP === "1" && + host && + !/^(127\.0\.0\.1|localhost)(:\d+)?$/i.test(host) + ) { + return Response.json({ error: "Forbidden" }, { status: 403 }) + } + return null +} + /** * Check the x-access-code header against ACCESS_CODE_LIST. * Returns a 401 response to send back when the check fails, or null when the diff --git a/lib/admin/providers.ts b/lib/admin/providers.ts index 39a48cae..bb8445a1 100644 --- a/lib/admin/providers.ts +++ b/lib/admin/providers.ts @@ -213,12 +213,17 @@ export function adminProvidersToConfig( indexByProvider.set(p.provider, index + 1) if (p.models.length === 0) continue const env = credEnvNames(p.provider, index) + // An entry with its own key also names its own URL variable, unset + // when the URL is empty: the global

_BASE_URL may be a proxy for + // another key, and the Test used the official endpoint. An Azure + // key belongs to one resource, so it keeps the server's. + const ownUrl = !!p.baseUrl || (!!p.apiKey && p.provider !== "azure") config.providers.push({ name: displayName(p), provider: p.provider, models: p.models, ...(env.key && p.apiKey ? { apiKeyEnv: env.key } : {}), - ...(env.url && p.baseUrl ? { baseUrlEnv: env.url } : {}), + ...(env.url && ownUrl ? { baseUrlEnv: env.url } : {}), ...(p.isDefault ? { default: true } : {}), }) } diff --git a/lib/ai-providers.ts b/lib/ai-providers.ts index 71da8966..72ab6492 100644 --- a/lib/ai-providers.ts +++ b/lib/ai-providers.ts @@ -900,6 +900,18 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { // DynamoDB quota manager use with their own credentials. const adminAccessKeyId = process.env.ADMIN_AWS_ACCESS_KEY_ID const adminSecretAccessKey = process.env.ADMIN_AWS_SECRET_ACCESS_KEY + // The region becomes part of the endpoint's host name, so a + // request's region must be a region name, or it could send the + // server's credentials to another host + if ( + overrides?.awsRegion && + !/^[a-z]{2,4}(-[a-z]+)+-\d{1,2}$/.test(overrides.awsRegion) + ) { + throw Object.assign( + new Error(`Invalid AWS region "${overrides.awsRegion}"`), + { statusCode: 400 }, + ) + } const bedrockRegion = overrides?.awsRegion || process.env.ADMIN_AWS_REGION || @@ -1031,8 +1043,9 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { : `${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. + // user's key, or an admin entry's own (empty) URL variable, 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 @@ -1045,7 +1058,10 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { const baseURL = configuredBaseURL || (SDK_KNOWS_ENDPOINT.has(provider) && - !(provider === "openai" && overrides?.apiKey) + !( + provider === "openai" && + (overrides?.apiKey || overrides?.baseUrlEnv) + ) ? undefined : publicDefault) // With a user's Azure key the SDK would read the server's diff --git a/lib/provider-models.ts b/lib/provider-models.ts index f291d35f..183718a2 100644 --- a/lib/provider-models.ts +++ b/lib/provider-models.ts @@ -1,5 +1,6 @@ import { createGateway } from "ai" import { getModelInfo } from "@/lib/model-catalog" +import { readLimitedBody } from "@/lib/read-limited-body" import { normalizeBaseUrl, PROVIDER_INFO, @@ -56,6 +57,44 @@ export function extractAihubmixModelIds(payload: unknown): string[] { return [...ids] } +/** + * An error this module wrote itself. Only these texts reach the caller: + * the base URL is the caller's and may be an internal address, so anything + * else (a parse error quoting the body, a network error naming a host) + * stays in the server log. + */ +export class ModelListError extends Error { + constructor( + message: string, + readonly statusCode?: number, + ) { + super(message) + this.name = "ModelListError" + } +} + +const MAX_LIST_BYTES = 2 * 1024 * 1024 + +/** A fetch that reads at most MAX_LIST_BYTES of each response */ +function sizeLimitedFetch(fetchFn: typeof fetch): typeof fetch { + return async (input, init) => { + const response = await fetchFn(input, init) + const body = await readLimitedBody(response, MAX_LIST_BYTES) + if (body === null) { + throw new ModelListError("The model list is too large.") + } + // The body is already decoded and has its own length now + const headers = new Headers(response.headers) + headers.delete("content-encoding") + headers.delete("content-length") + return new Response(body, { + status: response.status, + statusText: response.statusText, + headers, + }) + } +} + /** GET a JSON list; a failed request carries its status for the error hint */ async function getJson( url: string, @@ -67,12 +106,17 @@ async function getJson( signal: AbortSignal.timeout(15_000), }) if (!response.ok) { - throw Object.assign( - new Error(`The model list request failed (${response.status})`), - { statusCode: response.status }, + throw new ModelListError( + `The model list request failed (${response.status})`, + response.status, ) } - return response.json() + const text = await response.text() + try { + return JSON.parse(text) + } catch { + throw new ModelListError("The model list was not valid JSON.") + } } /** @@ -97,8 +141,9 @@ function listFallbackUrl(provider: ProviderName, apiKey?: string): string { export async function listProviderModels( provider: ProviderName, { apiKey, baseUrl }: { apiKey?: string; baseUrl?: string }, - fetchFn: typeof fetch = fetch, + unlimitedFetch: typeof fetch = fetch, ): Promise { + const fetchFn = sizeLimitedFetch(unlimitedFetch) const base = normalizeBaseUrl(baseUrl || listFallbackUrl(provider, apiKey)) const bearer: Record = apiKey ? { Authorization: `Bearer ${apiKey}` } @@ -183,7 +228,7 @@ export async function listProviderModels( } default: { if (!base) { - throw new Error( + throw new ModelListError( `${PROVIDER_INFO[provider].label} needs a base URL to list its models.`, ) } diff --git a/lib/read-limited-body.ts b/lib/read-limited-body.ts new file mode 100644 index 00000000..cb53e0e5 --- /dev/null +++ b/lib/read-limited-body.ts @@ -0,0 +1,29 @@ +/** + * Read a response body, giving up once it passes maxBytes, so a huge + * download from a URL the client chose can't exhaust server memory. + * Returns null when it is too large. + */ +export async function readLimitedBody( + response: Response, + maxBytes: number, +): Promise { + if (Number(response.headers.get("content-length")) > maxBytes) { + return null + } + if (!response.body) return new ArrayBuffer(0) + + const reader = response.body.getReader() + const chunks: Uint8Array[] = [] + let total = 0 + while (true) { + const { done, value } = await reader.read() + if (done) break + total += value.byteLength + if (total > maxBytes) { + await reader.cancel() + return null + } + chunks.push(value) + } + return new Blob(chunks as BlobPart[]).arrayBuffer() +} diff --git a/packages/mcp-server/src/http-server.ts b/packages/mcp-server/src/http-server.ts index 86f07339..6682761f 100644 --- a/packages/mcp-server/src/http-server.ts +++ b/packages/mcp-server/src/http-server.ts @@ -337,16 +337,17 @@ function handleRequest( } } -// Serve only requests addressed to localhost, sent by a localhost page or by -// a non-browser client (no Origin header). This blocks DNS rebinding and -// scripts on other websites. +// Serve only requests addressed to localhost, sent by the preview page +// itself (Origin is the address it was opened at, the Host) or by a +// non-browser client (no Origin header). This blocks DNS rebinding, other +// websites, and pages on other localhost ports, whose plain text POSTs need +// no CORS preflight. function isLocalRequest(req: http.IncomingMessage): boolean { - const isLocalHost = (host: string) => - /^(localhost|127\.0\.0\.1)(:\d+)?$/.test(host) + const host = req.headers.host ?? "" const origin = req.headers.origin return ( - isLocalHost(req.headers.host ?? "") && - (origin === undefined || isLocalHost(origin.replace(/^http:\/\//, ""))) + /^(localhost|127\.0\.0\.1)(:\d+)?$/.test(host) && + (origin === undefined || origin === `http://${host}`) ) } diff --git a/packages/mcp-server/src/preview/preview.js b/packages/mcp-server/src/preview/preview.js index 8e6f3835..ec85a304 100644 --- a/packages/mcp-server/src/preview/preview.js +++ b/packages/mcp-server/src/preview/preview.js @@ -382,12 +382,27 @@ function renderHistory() { } historyGrid.style.display = 'grid'; historyEmpty.style.display = 'none'; - historyGrid.innerHTML = historyData.map((e, i) => ` -

-
${e.svg ? `` : '#' + e.index}
-
#${e.index}
-
- `).join(''); + // Built element by element: a stored image is never read as HTML, and + // only an SVG data URL is shown as one + historyGrid.replaceChildren(...historyData.map((e) => { + const item = document.createElement('div'); + item.className = 'history-item'; + item.dataset.id = String(e.id); + const thumb = document.createElement('div'); + thumb.className = 'thumb'; + if (typeof e.svg === 'string' && e.svg.startsWith('data:image/svg+xml;base64,')) { + const img = document.createElement('img'); + img.src = e.svg; + thumb.appendChild(img); + } else { + thumb.textContent = '#' + e.index; + } + const label = document.createElement('div'); + label.className = 'label'; + label.textContent = '#' + e.index; + item.append(thumb, label); + return item; + })); historyGrid.querySelectorAll('.history-item').forEach(item => { item.onclick = () => { const id = parseInt(item.dataset.id); diff --git a/packages/mcp-server/tests/http-server.test.ts b/packages/mcp-server/tests/http-server.test.ts index 92ba8b40..cc060382 100644 --- a/packages/mcp-server/tests/http-server.test.ts +++ b/packages/mcp-server/tests/http-server.test.ts @@ -153,6 +153,36 @@ describe("request origin checks", () => { { origin: `http://localhost:${port}` }, ) expect(res.status).toBe(200) + // Opened as 127.0.0.1, or through a forwarded port: Origin and + // Host name the same host + for (const host of [`127.0.0.1:${port}`, "localhost:7000"]) { + const page = await postJson( + "/api/state", + { sessionId: "mcp-same-origin", xml: "" }, + { origin: `http://${host}`, host }, + ) + expect(page.status).toBe(200) + } + }) + + it("refuses writes from a page on another localhost port", async () => { + // A plain text POST needs no CORS preflight, so the server must + // refuse it itself + setState("mcp-other-port", "kept") + for (const path of ["/api/state", "/api/history-svg"]) { + const res = await postJson( + path, + { + sessionId: "mcp-other-port", + xml: "replaced", + svg: "x", + }, + { origin: "http://localhost:3000" }, + ) + expect(res.status).toBe(403) + } + expect(getState("mcp-other-port")?.xml).toBe("kept") + expect(getState("mcp-other-port")?.svg).toBeUndefined() }) }) diff --git a/tests/unit/admin-providers.test.ts b/tests/unit/admin-providers.test.ts index 0602d96f..4c8a7c02 100644 --- a/tests/unit/admin-providers.test.ts +++ b/tests/unit/admin-providers.test.ts @@ -156,6 +156,29 @@ describe("adminProvidersToConfig", () => { expect(config.providers[1].apiKeyEnv).toBe("ADMIN_OPENAI_API_KEY_2") }) + it("names its own URL variable when it has its own key, even empty", () => { + // Otherwise chat reads the global OPENAI_BASE_URL, which may be a + // proxy for another key, while the Test used the official endpoint + const own = adminProvidersToConfig([provider()]).providers[0] + expect(own.baseUrlEnv).toBe("ADMIN_OPENAI_BASE_URL") + // Without a key or URL of its own: the global key and URL, a pair + const shared = adminProvidersToConfig([provider({ apiKey: undefined })]) + .providers[0] + expect(shared.baseUrlEnv).toBeUndefined() + // An Azure key belongs to one resource: AZURE_BASE_URL stays + const azure = adminProvidersToConfig([provider({ provider: "azure" })]) + .providers[0] + expect(azure.baseUrlEnv).toBeUndefined() + expect( + adminProvidersToConfig([ + provider({ + provider: "azure", + baseUrl: "https://r.openai.azure.com/openai", + }), + ]).providers[0].baseUrlEnv, + ).toBe("ADMIN_AZURE_BASE_URL") + }) + it("skips providers without models and carries the default flag", () => { const config = adminProvidersToConfig([ provider({ id: "p1", models: [] }), diff --git a/tests/unit/ai-providers-credentials.test.ts b/tests/unit/ai-providers-credentials.test.ts index 5da4d189..344088c0 100644 --- a/tests/unit/ai-providers-credentials.test.ts +++ b/tests/unit/ai-providers-credentials.test.ts @@ -1,5 +1,6 @@ import { createOpenAI } from "@ai-sdk/openai" import { afterEach, beforeEach, describe, expect, it, vi } from "vitest" +import { AWS_REGIONS } from "@/components/provider-credentials-fields" import { getAIModel, getValidationModel, @@ -189,6 +190,58 @@ describe("Bedrock admin panel credentials", () => { }) }) + it("refuses a region that is not a region name", async () => { + // It becomes part of the endpoint's host name, with the server's + // credentials too + process.env.ADMIN_AWS_ACCESS_KEY_ID = "panel-id" + process.env.ADMIN_AWS_SECRET_ACCESS_KEY = "panel-secret" + const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock") + for (const awsRegion of [ + "us-east-1.attacker.example/", + "x/#", + "US-EAST-1", + "us-east-1 ", + ]) { + expect(() => + getAIModel({ + provider: "bedrock", + modelId: "amazon.nova-lite-v1:0", + awsRegion, + }), + ).toThrow(/Invalid AWS region/) + expect(() => + getAIModel({ + provider: "bedrock", + modelId: "amazon.nova-lite-v1:0", + awsAccessKeyId: "client-id", + awsSecretAccessKey: "client-secret", + awsRegion, + }), + ).toThrow(/Invalid AWS region/) + } + expect(createAmazonBedrock).not.toHaveBeenCalled() + }) + + it("accepts every region the settings offer, and other partitions", () => { + for (const awsRegion of [ + ...AWS_REGIONS.map(([region]) => region), + "us-gov-west-1", + "cn-northwest-1", + "us-iso-east-1", + "eusc-de-east-1", + ]) { + expect(() => + getAIModel({ + provider: "bedrock", + modelId: "amazon.nova-lite-v1:0", + awsAccessKeyId: "client-id", + awsSecretAccessKey: "client-secret", + awsRegion, + }), + ).not.toThrow() + } + }) + it("falls back to the default AWS credential chain", async () => { process.env.AWS_REGION = "us-east-1" const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock") @@ -325,6 +378,25 @@ describe("whose keys a request uses", () => { expect(provider.chat).not.toHaveBeenCalled() }) + it("sends an admin OpenAI key without a URL to the official endpoint", () => { + // Its URL variable is named but empty; the SDK would otherwise read + // the server's OPENAI_BASE_URL, a proxy for another key + process.env.OPENAI_BASE_URL = "https://operator-proxy.example.com/v1" + process.env.ADMIN_OPENAI_API_KEY = "panel-key" + getAIModel({ + provider: "openai", + modelId: "gpt-5.5", + apiKeyEnv: "ADMIN_OPENAI_API_KEY", + baseUrlEnv: "ADMIN_OPENAI_BASE_URL", + }) + expect(createOpenAI).toHaveBeenLastCalledWith( + expect.objectContaining({ + apiKey: "panel-key", + baseURL: "https://api.openai.com/v1", + }), + ) + }) + it("uses Chat Completions for any configured base URL", () => { // The settings form fills in the official URL for a new provider getAIModel({ diff --git a/tests/unit/chat-route-quota.test.ts b/tests/unit/chat-route-quota.test.ts index a2b01cfd..fd3e6ad4 100644 --- a/tests/unit/chat-route-quota.test.ts +++ b/tests/unit/chat-route-quota.test.ts @@ -151,6 +151,19 @@ describe("chat quota", () => { }) }) +describe("request checks", () => { + it("refuses an AWS region that is not a region name", async () => { + process.env.AI_PROVIDER = "bedrock" + process.env.AI_MODEL = "amazon.nova-lite-v1:0" + const res = await send({ + "x-aws-region": "us-east-1.attacker.example/", + }) + expect(res.status).toBe(400) + expect(await res.text()).toMatch(/Invalid AWS region/) + expect(fetch).not.toHaveBeenCalled() + }) +}) + 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 diff --git a/tests/unit/cross-site-requests.test.ts b/tests/unit/cross-site-requests.test.ts new file mode 100644 index 00000000..8053ae3b --- /dev/null +++ b/tests/unit/cross-site-requests.test.ts @@ -0,0 +1,95 @@ +// @vitest-environment node +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest" +import { POST as chat } from "@/app/api/chat/route" +import { POST as parseUrl } from "@/app/api/parse-url/route" +import { POST as providerModels } from "@/app/api/provider-models/route" +import { POST as validateDiagram } from "@/app/api/validate-diagram/route" +import { POST as validateModel } from "@/app/api/validate-model/route" + +// The routes that run models or fetch URLs, which a page on another site +// could otherwise make the user's own server do +const ROUTES = { + chat, + "parse-url": parseUrl, + "provider-models": providerModels, + "validate-diagram": validateDiagram, + "validate-model": validateModel, +} + +const body = JSON.stringify({ + messages: [ + { id: "u1", role: "user", parts: [{ type: "text", text: "hi" }] }, + ], + url: "https://example.com", + provider: "openai", + apiKey: "k", + modelId: "gpt-5.5", + imageData: "data:image/png;base64,AAAA", +}) + +const saved = process.env.NEXT_AI_DRAWIO_DESKTOP + +beforeEach(() => { + vi.stubGlobal( + "fetch", + vi.fn(async () => { + throw new Error("no network in tests") + }), + ) +}) + +afterEach(() => { + if (saved === undefined) delete process.env.NEXT_AI_DRAWIO_DESKTOP + else process.env.NEXT_AI_DRAWIO_DESKTOP = saved + vi.unstubAllGlobals() +}) + +describe("requests another website could send", () => { + for (const [name, post] of Object.entries(ROUTES)) { + it(`${name}: refuses a text body, which needs no CORS preflight`, async () => { + // fetch(..., { mode: "no-cors", body: JSON.stringify(...) }) + // from another site arrives as text/plain + const res = await post( + new Request(`http://127.0.0.1:61337/api/${name}`, { + method: "POST", + body, + }), + ) + expect(res.status).toBe(415) + expect(fetch).not.toHaveBeenCalled() + }) + + it(`${name}: desktop app refuses a foreign Host (DNS rebinding)`, async () => { + process.env.NEXT_AI_DRAWIO_DESKTOP = "1" + const res = await post( + new Request(`http://127.0.0.1:61337/api/${name}`, { + method: "POST", + headers: { + "Content-Type": "application/json", + host: "rebind.attacker.example:61337", + }, + body, + }), + ) + expect(res.status).toBe(403) + expect(fetch).not.toHaveBeenCalled() + }) + } + + it("lets the desktop window's own requests through", async () => { + process.env.NEXT_AI_DRAWIO_DESKTOP = "1" + const res = await validateModel( + new Request("http://127.0.0.1:61337/api/validate-model", { + method: "POST", + headers: { + "Content-Type": "application/json; charset=utf-8", + host: "127.0.0.1:61337", + }, + body: JSON.stringify({ provider: "openai" }), + }), + ) + // Past the check: the route's own validation answers + expect(res.status).toBe(400) + expect(await res.text()).toMatch(/required/) + }) +}) diff --git a/tests/unit/provider-models.test.ts b/tests/unit/provider-models.test.ts index 19ad7248..a1bd9e02 100644 --- a/tests/unit/provider-models.test.ts +++ b/tests/unit/provider-models.test.ts @@ -190,4 +190,59 @@ describe("POST /api/provider-models", () => { const res = await post({ provider: "deepseek", apiKey: "k" }) expect(await res.json()).toMatchObject({ code: "invalid_api_key" }) }) + + // The base URL is the caller's, and private addresses are allowed by + // default (local Ollama), so the answer must not reveal what an + // internal address sent back + const text = (body: string) => + vi.fn( + async () => new Response(body, { status: 200 }), + ) as unknown as typeof fetch + + it("does not repeat a body that is not JSON", async () => { + vi.stubGlobal("fetch", text("ROLE-NAME-OF-THE-SERVER")) + const res = await post({ + provider: "ollama", + baseUrl: "http://169.254.169.254/latest/meta-data/x?", + }) + const data = await res.json() + expect(data.error).toBe("The model list was not valid JSON.") + expect(JSON.stringify(data)).not.toContain("ROLE") + }) + + it("stops reading a list over 2 MB, also through the Gateway SDK", async () => { + const huge = JSON.stringify({ data: [{ id: "x".repeat(3_000_000) }] }) + for (const body of [ + { provider: "ollama", baseUrl: "https://big.example.com" }, + { + provider: "gateway", + apiKey: "k", + baseUrl: "https://big.example.com/v3/ai", + }, + ]) { + vi.stubGlobal("fetch", text(huge)) + const data = await (await post(body)).json() + expect(data.error).toBe("The model list is too large.") + expect(data.models).toBeUndefined() + } + }) + + it("keeps its own explanations and hides other error texts", async () => { + // Our own: no base URL for SGLang + const own = await ( + await post({ provider: "sglang", apiKey: "k" }) + ).json() + expect(own.error).toMatch(/needs a base URL/) + // Not ours: an exception text from the network layer + vi.stubGlobal( + "fetch", + vi.fn(async () => { + throw new Error("connect ECONNREFUSED 10.1.2.3:8080") + }), + ) + const other = await ( + await post({ provider: "ollama", baseUrl: "http://10.1.2.3:8080" }) + ).json() + expect(other.error).toBe("The model list request failed.") + }) })