mirror of
https://github.com/DayuanJiang/next-ai-draw-io.git
synced 2026-10-10 19:49:52 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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(
|
||||
{
|
||||
|
||||
@@ -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.",
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"],
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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.`,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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}`)
|
||||
)
|
||||
}
|
||||
|
||||
|
||||
@@ -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()
|
||||
})
|
||||
})
|
||||
|
||||
|
||||
@@ -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({
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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/)
|
||||
})
|
||||
})
|
||||
@@ -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.")
|
||||
})
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user