fix(security): check request sources, regions and endpoints

- Bedrock: a request's AWS region must be a region name. It becomes part
  of the endpoint's host name, so a value such as
  "us-east-1.attacker.example/" sent the server's bearer token or signed
  request to another host.
- MCP preview server: only the preview page itself (Origin equal to the
  Host) or a non-browser client may call it; a page on another localhost
  port could replace the diagram with a plain text POST. History builds
  its thumbnails element by element and shows only SVG data images, so a
  stored value can no longer run script in the preview.
- chat, validate-model, validate-diagram, provider-models and parse-url
  take JSON bodies only, so another website cannot make the user's own
  server (the desktop app, a local install) run models with their keys;
  the desktop app also refuses a foreign Host (DNS rebinding).
- The model list reads at most 2 MB, also through the Gateway SDK, and
  answers only with its own error texts: the URL is the caller's and may
  be an internal address.
- An admin panel provider with its own key and no URL no longer inherits
  the global <P>_BASE_URL, which may be a proxy for another key; OpenAI
  then gets the official endpoint, as its Test. Azure keeps the server's
  resource.
This commit is contained in:
dayuan.jiang
2026-10-05 17:02:06 +09:00
parent 50c7ad3ec4
commit 4731394f32
19 changed files with 488 additions and 59 deletions
+3 -1
View File
@@ -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<Response>()
const DEBUG_LLM_PAYLOAD = process.env.DEBUG_LLM_PAYLOAD === "true"
async function handleChatRequest(req: Request): Promise<Response> {
const crossSite = rejectCrossSite(req)
if (crossSite) return crossSite
// Check for access code
const accessDenied = checkAccessCode(req)
if (accessDenied) return accessDenied
+5 -28
View File
@@ -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<ArrayBuffer | null> {
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(
{
+23 -4
View File
@@ -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<ProviderName>([
* 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.",
})
}
}
+3 -1
View File
@@ -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<Response> {
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
+3 -1
View File
@@ -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)
+1 -1
View File
@@ -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"],
+28
View File
@@ -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
+6 -1
View File
@@ -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 <P>_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 } : {}),
})
}
+19 -3
View File
@@ -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
+51 -6
View File
@@ -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<ListedModel[]> {
const fetchFn = sizeLimitedFetch(unlimitedFetch)
const base = normalizeBaseUrl(baseUrl || listFallbackUrl(provider, apiKey))
const bearer: Record<string, string> = 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.`,
)
}
+29
View File
@@ -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<ArrayBuffer | null> {
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()
}
+8 -7
View File
@@ -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}`)
)
}
+21 -6
View File
@@ -382,12 +382,27 @@ function renderHistory() {
}
historyGrid.style.display = 'grid';
historyEmpty.style.display = 'none';
historyGrid.innerHTML = historyData.map((e, i) => `
<div class="history-item" data-id="${e.id}">
<div class="thumb">${e.svg ? `<img src="${e.svg}">` : '#' + e.index}</div>
<div class="label">#${e.index}</div>
</div>
`).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);
@@ -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: "<mxfile/>" },
{ 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", "<mxfile>kept</mxfile>")
for (const path of ["/api/state", "/api/history-svg"]) {
const res = await postJson(
path,
{
sessionId: "mcp-other-port",
xml: "<mxfile>replaced</mxfile>",
svg: "x",
},
{ origin: "http://localhost:3000" },
)
expect(res.status).toBe(403)
}
expect(getState("mcp-other-port")?.xml).toBe("<mxfile>kept</mxfile>")
expect(getState("mcp-other-port")?.svg).toBeUndefined()
})
})
+23
View File
@@ -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: [] }),
@@ -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({
+13
View File
@@ -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
+95
View File
@@ -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/)
})
})
+55
View File
@@ -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.")
})
})