From 528b6e54c8eccd60d7dd61d8ac785c7cdf3945e5 Mon Sep 17 00:00:00 2001 From: "dayuan.jiang" Date: Sat, 3 Oct 2026 17:45:21 +0900 Subject: [PATCH] fix(api): require access codes and limit sizes on helper routes - Shared checkAccessCode for validate-diagram, validate-model, parse-url, verify-access-code - parse-url: 5 MB streamed body limit; validate-diagram: 5 MB image limit - validate-model refuses redirects when private URLs are blocked - Admin settings state shared across module instances via globalThis - Server model ids: unique slugs (non-ASCII names encoded), duplicates rejected - Panel Bedrock credentials stored as ADMIN_AWS_* so the DynamoDB client keeps its own - Locale redirect keeps basePath and query; EdgeOne function drops open CORS and checks the access code - Providers payload reports whether .env sets a default model --- app/api/admin/providers/route.ts | 9 +- app/api/parse-url/route.ts | 41 ++++- app/api/validate-diagram/route.ts | 15 ++ app/api/validate-model/route.ts | 53 +++++- app/api/verify-access-code/route.ts | 28 +--- edge-functions/api/edgeai/chat/completions.ts | 66 +++++--- lib/admin/providers.ts | 26 ++- lib/admin/settings.ts | 53 +++--- lib/server-model-config.ts | 18 +- proxy.ts | 13 +- tests/unit/admin-providers-route.test.ts | 71 ++++++++ tests/unit/admin-providers.test.ts | 51 +++++- tests/unit/admin-settings.test.ts | 20 ++- tests/unit/api-access-code.test.ts | 156 ++++++++++++++++++ tests/unit/edgeone-function.test.ts | 47 ++++++ tests/unit/server-model-config.test.ts | 48 ++++++ 16 files changed, 622 insertions(+), 93 deletions(-) create mode 100644 tests/unit/admin-providers-route.test.ts create mode 100644 tests/unit/api-access-code.test.ts create mode 100644 tests/unit/edgeone-function.test.ts diff --git a/app/api/admin/providers/route.ts b/app/api/admin/providers/route.ts index 216f08c..13e2fe1 100644 --- a/app/api/admin/providers/route.ts +++ b/app/api/admin/providers/route.ts @@ -7,7 +7,11 @@ import { mergeSecrets, validateAdminProviders, } from "@/lib/admin/providers" -import { isSettingsWritable, saveSettings } from "@/lib/admin/settings" +import { + getEnvFallback, + isSettingsWritable, + saveSettings, +} from "@/lib/admin/settings" import { loadEnvServerModelsConfig } from "@/lib/server-model-config" export const runtime = "nodejs" @@ -33,6 +37,9 @@ async function payload() { models: p.models, isDefault: !!p.default && !adminHasDefault, })) ?? [], + // Whether .env sets a default model. getEnvFallback skips the value + // the panel overlays onto process.env, so a panel default doesn't count. + envHasDefaultModel: !!getEnvFallback("AI_MODEL"), } } diff --git a/app/api/parse-url/route.ts b/app/api/parse-url/route.ts index 794b711..fb03032 100644 --- a/app/api/parse-url/route.ts +++ b/app/api/parse-url/route.ts @@ -1,9 +1,11 @@ import { extractFromHtml } from "@extractus/article-extractor" import { NextResponse } from "next/server" import TurndownService from "turndown" +import { checkAccessCode } from "@/lib/access-code" import { isPrivateUrl } from "@/lib/ssrf-protection" const MAX_CONTENT_LENGTH = 150000 // Match PDF limit +const MAX_RESPONSE_BYTES = 5 * 1024 * 1024 const EXTRACT_TIMEOUT_MS = 15000 const USER_AGENT = "Mozilla/5.0 (compatible; NextAIDrawio/1.0)" @@ -32,7 +34,36 @@ 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 accessError = checkAccessCode(req) + if (accessError) return accessError + try { const { url } = await req.json() @@ -97,7 +128,15 @@ export async function POST(req: Request) { ) } - const buffer = await response.arrayBuffer() + const buffer = await readLimitedBody(response) + if (!buffer) { + return NextResponse.json( + { + error: `Page exceeds the ${MAX_RESPONSE_BYTES / 1024 / 1024} MB download limit`, + }, + { status: 413 }, + ) + } const charset = detectCharset(contentType, buffer) html = new TextDecoder(charset).decode(buffer) } catch (err: any) { diff --git a/app/api/validate-diagram/route.ts b/app/api/validate-diagram/route.ts index 6146401..41fd12d 100644 --- a/app/api/validate-diagram/route.ts +++ b/app/api/validate-diagram/route.ts @@ -4,6 +4,7 @@ */ import { streamObject } from "ai" +import { checkAccessCode } from "@/lib/access-code" import { getValidationModel } from "@/lib/ai-providers" import { VALIDATION_SYSTEM_PROMPT } from "@/lib/validation-prompts" import { @@ -13,6 +14,9 @@ import { export const maxDuration = 30 +// Data URL length cap (~3.75 MB of PNG), well above a normal diagram capture +const MAX_IMAGE_DATA_LENGTH = 5 * 1024 * 1024 + interface ValidateDiagramRequest { imageData: string // Base64 PNG data URL sessionId?: string @@ -44,6 +48,10 @@ function createStreamingResponse(result: ValidationResult): Response { } export async function POST(req: Request): Promise { + // Uses the server's model credentials, so require the access code + const accessError = checkAccessCode(req) + if (accessError) return accessError + try { // Check if VLM validation is enabled (default: true) const enableValidation = process.env.ENABLE_VLM_VALIDATION !== "false" @@ -72,6 +80,13 @@ export async function POST(req: Request): Promise { ) } + if (imageData.length > MAX_IMAGE_DATA_LENGTH) { + return Response.json( + { error: "Image data too large" }, + { status: 413 }, + ) + } + // Get the validation model let model try { diff --git a/app/api/validate-model/route.ts b/app/api/validate-model/route.ts index 5d08fc8..dfb3e0a 100644 --- a/app/api/validate-model/route.ts +++ b/app/api/validate-model/route.ts @@ -10,6 +10,7 @@ import { createOpenRouter } from "@openrouter/ai-sdk-provider" import { generateText } from "ai" import { NextResponse } from "next/server" import { createOllama } from "ollama-ai-provider-v2" +import { checkAccessCode } from "@/lib/access-code" import { AIHUBMIX_APP_CODE, isAihubmixStandardBaseURL, @@ -33,7 +34,24 @@ interface ValidateRequest { vertexApiKey?: string // Express Mode API key } +// With private URLs blocked, a public baseUrl could still redirect the +// request to an internal host, so redirects are refused in that case. +function redirectGuardedFetch(): typeof fetch | undefined { + if (allowPrivateUrls()) return undefined + return async (input, init) => { + const response = await fetch(input, { ...init, redirect: "manual" }) + if (response.status >= 300 && response.status < 400) { + throw new Error("Redirects are not allowed for custom base URLs") + } + return response + } +} + export async function POST(req: Request) { + // Lets the server send requests to arbitrary URLs, so require the access code + const accessError = checkAccessCode(req) + if (accessError) return accessError + try { const body: ValidateRequest = await req.json() const { @@ -91,6 +109,7 @@ export async function POST(req: Request) { ) } + const guardedFetch = redirectGuardedFetch() let model: any switch (provider) { @@ -98,6 +117,7 @@ export async function POST(req: Request) { const openai = createOpenAI({ apiKey, ...(baseUrl && { baseURL: baseUrl }), + fetch: guardedFetch, }) model = openai.chat(modelId) break @@ -107,6 +127,7 @@ export async function POST(req: Request) { const anthropic = createAnthropic({ apiKey, baseURL: baseUrl || "https://api.anthropic.com/v1", + fetch: guardedFetch, }) model = anthropic(modelId) break @@ -116,6 +137,7 @@ export async function POST(req: Request) { const google = createGoogleGenerativeAI({ apiKey, ...(baseUrl && { baseURL: baseUrl }), + fetch: guardedFetch, }) model = google(modelId) break @@ -125,6 +147,7 @@ export async function POST(req: Request) { const vertex = createVertex({ apiKey: vertexApiKey, ...(baseUrl && { baseURL: baseUrl }), + fetch: guardedFetch, }) model = vertex(modelId) break @@ -134,6 +157,7 @@ export async function POST(req: Request) { const azure = createOpenAI({ apiKey, baseURL: baseUrl, + fetch: guardedFetch, }) model = azure.chat(modelId) break @@ -153,6 +177,7 @@ export async function POST(req: Request) { const openrouter = createOpenRouter({ apiKey, ...(baseUrl && { baseURL: baseUrl }), + fetch: guardedFetch, }) model = openrouter(modelId) break @@ -174,6 +199,7 @@ export async function POST(req: Request) { const aihubmixCompatible = createOpenAI({ apiKey, baseURL: baseUrl, + fetch: guardedFetch, }) model = aihubmixCompatible.chat(modelId) } @@ -185,6 +211,7 @@ export async function POST(req: Request) { const ds = createDeepSeek({ apiKey, ...(baseUrl && { baseURL: baseUrl }), + fetch: guardedFetch, }) model = ds(modelId) } else { @@ -197,6 +224,7 @@ export async function POST(req: Request) { const sf = createOpenAI({ apiKey, baseURL: baseUrl || "https://api.siliconflow.cn/v1", + fetch: guardedFetch, }) model = sf.chat(modelId) break @@ -213,6 +241,7 @@ export async function POST(req: Request) { baseUrl || process.env.OLLAMA_BASE_URL || "https://ollama.com/api", + fetch: guardedFetch, ...(ollamaApiKey && { headers: { Authorization: `Bearer ${ollamaApiKey}` }, }), @@ -225,6 +254,7 @@ export async function POST(req: Request) { const gw = createGateway({ apiKey, ...(baseUrl && { baseURL: baseUrl }), + fetch: guardedFetch, }) model = gw(modelId) break @@ -232,13 +262,16 @@ export async function POST(req: Request) { case "edgeone": { // EdgeOne uses OpenAI-compatible API via Edge Functions - // Need to pass cookies for EdgeOne Pages authentication + // Need to pass cookies for EdgeOne Pages authentication, + // and the access code, which the edge function also checks const cookieHeader = req.headers.get("cookie") || "" const edgeone = createOpenAI({ apiKey: "edgeone", // EdgeOne doesn't require API key baseURL: baseUrl || "/api/edgeai", + fetch: guardedFetch, headers: { cookie: cookieHeader, + "x-access-code": req.headers.get("x-access-code") || "", }, }) model = edgeone.chat(modelId) @@ -250,6 +283,7 @@ export async function POST(req: Request) { const sglang = createOpenAI({ apiKey: apiKey || "not-needed", baseURL: baseUrl || "http://127.0.0.1:8000/v1", + fetch: guardedFetch, }) model = sglang.chat(modelId) break @@ -267,12 +301,14 @@ export async function POST(req: Request) { const doubao = createDeepSeek({ apiKey, baseURL: doubaoBaseUrl, + fetch: guardedFetch, }) model = doubao(modelId) } else { const doubao = createOpenAI({ apiKey, baseURL: doubaoBaseUrl, + fetch: guardedFetch, }) model = doubao.chat(modelId) } @@ -286,7 +322,7 @@ export async function POST(req: Request) { try { // Initiate a streaming request (required for QwQ-32B and certain Qwen3 models) - const response = await fetch( + const response = await (guardedFetch ?? fetch)( `${baseURL}/chat/completions`, { method: "POST", @@ -307,9 +343,15 @@ export async function POST(req: Request) { ) if (!response.ok) { - const errorText = await response.text() + // Log the body but return only the status: the + // caller chooses baseUrl, so the body may come from + // any host the server can reach + console.error( + "[validate-model] ModelScope error body:", + await response.text(), + ) throw new Error( - `ModelScope API error (${response.status}): ${errorText}`, + `ModelScope API error (${response.status})`, ) } @@ -360,12 +402,14 @@ export async function POST(req: Request) { const minimax = createAnthropic({ apiKey, baseURL: minimaxBaseUrl, + fetch: guardedFetch, }) model = minimax.chat(modelId) } else { const minimax = createOpenAI({ apiKey, baseURL: minimaxBaseUrl, + fetch: guardedFetch, }) model = minimax.chat(modelId) } @@ -398,6 +442,7 @@ export async function POST(req: Request) { const openai = createOpenAI({ apiKey, baseURL, + fetch: guardedFetch, }) model = openai.chat(modelId) break diff --git a/app/api/verify-access-code/route.ts b/app/api/verify-access-code/route.ts index d69f59d..55cbc94 100644 --- a/app/api/verify-access-code/route.ts +++ b/app/api/verify-access-code/route.ts @@ -1,29 +1,9 @@ +import { checkAccessCode } from "@/lib/access-code" + export async function POST(req: Request) { - const accessCodes = - process.env.ACCESS_CODE_LIST?.split(",") - .map((code) => code.trim()) - .filter(Boolean) || [] - - // If no access codes configured, verification always passes - if (accessCodes.length === 0) { - return Response.json({ - valid: true, - message: "No access code required", - }) - } - - const accessCodeHeader = req.headers.get("x-access-code") - - if (!accessCodeHeader) { + if (checkAccessCode(req)) { return Response.json( - { valid: false, message: "Access code is required" }, - { status: 401 }, - ) - } - - if (!accessCodes.includes(accessCodeHeader)) { - return Response.json( - { valid: false, message: "Invalid access code" }, + { valid: false, message: "Invalid or missing access code" }, { status: 401 }, ) } diff --git a/edge-functions/api/edgeai/chat/completions.ts b/edge-functions/api/edgeai/chat/completions.ts index eafd4de..e2c558f 100644 --- a/edge-functions/api/edgeai/chat/completions.ts +++ b/edge-functions/api/edgeai/chat/completions.ts @@ -67,41 +67,62 @@ const MODEL_ALIASES: Record = { "deepseek-v3-0324": "@tx/deepseek-ai/deepseek-v3-0324", } -const CORS_HEADERS = { - "Access-Control-Allow-Origin": "*", - "Access-Control-Allow-Methods": "POST, OPTIONS", - "Access-Control-Allow-Headers": "Content-Type, Authorization", -} - /** - * Create standardized response with CORS headers + * Create standardized JSON response */ function createResponse(body: any, status = 200, extraHeaders = {}): Response { return new Response(JSON.stringify(body), { status, headers: { "Content-Type": "application/json", - ...CORS_HEADERS, ...extraHeaders, }, }) } -/** - * Handle OPTIONS request for CORS preflight - */ -function handleOptionsRequest(): Response { - return new Response(null, { - headers: { - ...CORS_HEADERS, - "Access-Control-Max-Age": "86400", - }, - }) +// Only the app's own server (/api/chat, /api/validate-model) calls this +// function, so no CORS headers are sent: other sites' pages can't call it +// from a browser and spend the deployment's Edge AI quota. +// Same rule as lib/access-code.ts, but reading the edge function's env. +// No codes configured (or env unavailable) means no check. +function hasValidAccessCode(request: Request, env: any): boolean { + const accessCodes: string[] = + env?.ACCESS_CODE_LIST?.split(",") + .map((code: string) => code.trim()) + .filter(Boolean) || [] + if (accessCodes.length === 0) return true + const accessCode = request.headers.get("x-access-code") + return !!accessCode && accessCodes.includes(accessCode) } -export async function onRequest({ request, env: _env }: any) { - if (request.method === "OPTIONS") { - return handleOptionsRequest() +export async function onRequest({ request, env }: any) { + // Requiring JSON also makes any cross-site browser request need a CORS + // preflight, which fails without CORS headers + if ( + request.method !== "POST" || + !request.headers.get("content-type")?.includes("application/json") + ) { + return createResponse( + { + error: { + message: "Expected a POST request with a JSON body", + type: "invalid_request_error", + }, + }, + 400, + ) + } + + if (!hasValidAccessCode(request, env)) { + return createResponse( + { + error: { + message: "Invalid or missing access code", + type: "invalid_request_error", + }, + }, + 401, + ) } request.headers.delete("accept-encoding") @@ -153,7 +174,7 @@ export async function onRequest({ request, env: _env }: any) { type: "invalid_request_error", }, }, - 429, + 400, ) } @@ -216,7 +237,6 @@ export async function onRequest({ request, env: _env }: any) { "Cache-Control": "no-cache, no-store, no-transform", "X-Accel-Buffering": "no", Connection: "keep-alive", - ...CORS_HEADERS, }, }) } catch (error: any) { diff --git a/lib/admin/providers.ts b/lib/admin/providers.ts index 67d5034..39a48ca 100644 --- a/lib/admin/providers.ts +++ b/lib/admin/providers.ts @@ -2,6 +2,7 @@ import { z } from "zod" import { ProviderNameSchema, type ServerModelsConfig, + slugify, } from "@/lib/server-model-config" import { FIXED_CRED_PROVIDERS, @@ -182,12 +183,15 @@ export function validateAdminProviders( return `${PROVIDER_INFO[single].label} is already configured in AI_MODELS_CONFIG / ai-models.json and shares global credentials. Manage it via the environment configuration instead.` } } + // Server model ids are built from the slugified name, so names must + // stay distinct after slugifying ("OpenAI" and "openai" would collide) const names = list.map((p) => displayName(p)) - if (new Set(names).size !== names.length) { - return "Provider display names must be unique." + const slugs = names.map(slugify) + if (new Set(slugs).size !== slugs.length) { + return "Provider display names must be unique (ignoring case and punctuation)." } - const envNames = new Set(envProviders.map((p) => p.name)) - const clash = names.find((n) => envNames.has(n)) + const envSlugs = new Set(envProviders.map((p) => slugify(p.name))) + const clash = names.find((_, i) => envSlugs.has(slugs[i])) if (clash) { return `"${clash}" is already defined in AI_MODELS_CONFIG / ai-models.json. Use a different display name.` } @@ -240,10 +244,14 @@ export function deriveEnvUpdates( indexByProvider.set(p.provider, index + 1) if (p.provider === "bedrock") { - if (p.awsAccessKeyId) updates.AWS_ACCESS_KEY_ID = p.awsAccessKeyId + // ADMIN_ names keep the standard AWS_* vars untouched, so other + // AWS clients (e.g. the DynamoDB quota table) keep their own + // credentials instead of picking up the panel's Bedrock keys + if (p.awsAccessKeyId) + updates.ADMIN_AWS_ACCESS_KEY_ID = p.awsAccessKeyId if (p.awsSecretAccessKey) - updates.AWS_SECRET_ACCESS_KEY = p.awsSecretAccessKey - if (p.awsRegion) updates.AWS_REGION = p.awsRegion + updates.ADMIN_AWS_SECRET_ACCESS_KEY = p.awsSecretAccessKey + if (p.awsRegion) updates.ADMIN_AWS_REGION = p.awsRegion } else if (p.provider === "vertexai") { if (p.vertexApiKey) updates.GOOGLE_VERTEX_API_KEY = p.vertexApiKey if (p.baseUrl) updates.GOOGLE_VERTEX_BASE_URL = p.baseUrl @@ -284,6 +292,10 @@ function derivedEnvKeys(list: StoredAdminProvider[]): string[] { const index = indexByProvider.get(p.provider) ?? 0 indexByProvider.set(p.provider, index + 1) if (p.provider === "bedrock") { + keys.add("ADMIN_AWS_ACCESS_KEY_ID") + keys.add("ADMIN_AWS_SECRET_ACCESS_KEY") + keys.add("ADMIN_AWS_REGION") + // Written by older versions; listed so the next save clears them keys.add("AWS_ACCESS_KEY_ID") keys.add("AWS_SECRET_ACCESS_KEY") keys.add("AWS_REGION") diff --git a/lib/admin/settings.ts b/lib/admin/settings.ts index 664c370..f2ecce4 100644 --- a/lib/admin/settings.ts +++ b/lib/admin/settings.ts @@ -10,13 +10,27 @@ interface SettingsFile { values: Record } -// Original env values snapshotted before the first overlay, so removing a -// key from the settings file restores the env default. null = was unset. -const originalEnv: Record = {} -// Keys currently overlaid, so we can restore ones removed from the file. -let overlaidKeys = new Set() +interface SettingsState { + // Original env values snapshotted before the first overlay, so removing + // a key from the settings file restores the env default. null = was unset. + originalEnv: Record + // Keys currently overlaid, so we can restore ones removed from the file. + overlaidKeys: Set + cachedSettings: Record | null +} -let cachedSettings: Record | null = null +// Kept on globalThis because the build can load this module more than once +// (instrumentation.ts and the API routes get separate copies); per-module +// state would make a route forget what instrumentation overlaid at startup. +const globalState = globalThis as typeof globalThis & { + __adminSettingsState?: SettingsState +} +globalState.__adminSettingsState ??= { + originalEnv: {}, + overlaidKeys: new Set(), + cachedSettings: null, +} +const state = globalState.__adminSettingsState export function getSettingsPath(): string { const custom = process.env.SETTINGS_FILE @@ -25,7 +39,7 @@ export function getSettingsPath(): string { } export function loadSettings(): Record { - if (cachedSettings) return cachedSettings + if (state.cachedSettings) return state.cachedSettings try { const raw = fs.readFileSync(getSettingsPath(), "utf8") const parsed = JSON.parse(raw) as SettingsFile @@ -43,21 +57,22 @@ export function loadSettings(): Record { for (const [key, value] of Object.entries(rawValues)) { if (typeof value === "string") values[key] = value } - cachedSettings = values + state.cachedSettings = values } catch (err: any) { if (err?.code !== "ENOENT") { console.error("[admin-settings] Failed to read settings file:", err) } - cachedSettings = {} + state.cachedSettings = {} } - return cachedSettings + return state.cachedSettings } export function applyToEnv(): void { const values = loadSettings() + const { originalEnv } = state // Restore env for keys that were overlaid before but are now gone - for (const key of overlaidKeys) { + for (const key of state.overlaidKeys) { if (!(key in values)) { const original = originalEnv[key] if (original === null) delete process.env[key] @@ -72,12 +87,12 @@ export function applyToEnv(): void { process.env[key] = value } - overlaidKeys = new Set(Object.keys(values)) + state.overlaidKeys = new Set(Object.keys(values)) } // The effective env value if the file entry were removed (for fallback display) export function getEnvFallback(key: string): string | null { - if (overlaidKeys.has(key)) return originalEnv[key] ?? null + if (state.overlaidKeys.has(key)) return state.originalEnv[key] ?? null return process.env[key] ?? null } @@ -101,7 +116,7 @@ export function saveSettings(updates: Record): void { fs.writeFileSync(tmpPath, JSON.stringify(data, null, 2), { mode: 0o600 }) fs.renameSync(tmpPath, filePath) - cachedSettings = current + state.cachedSettings = current applyToEnv() } @@ -122,13 +137,13 @@ export function isSettingsWritable(): boolean { // Test-only: reset module state export function _resetForTests(): void { - cachedSettings = null + state.cachedSettings = null writableCache = null - for (const key of overlaidKeys) { - const original = originalEnv[key] + for (const key of state.overlaidKeys) { + const original = state.originalEnv[key] if (original === null) delete process.env[key] else if (original !== undefined) process.env[key] = original } - overlaidKeys = new Set() - for (const key of Object.keys(originalEnv)) delete originalEnv[key] + state.overlaidKeys = new Set() + state.originalEnv = {} } diff --git a/lib/server-model-config.ts b/lib/server-model-config.ts index f00d148..6b5220f 100644 --- a/lib/server-model-config.ts +++ b/lib/server-model-config.ts @@ -47,11 +47,14 @@ export interface FlattenedServerModel { /** * Convert provider name to URL-safe slug for use in model ID - * e.g., "OpenAI Production" → "openai-production" + * e.g., "OpenAI Production" → "openai-production", "主力" → "4e3b-529b" + * Non-ASCII characters become their hex code point so CJK names stay + * distinct; the id is sent in HTTP headers, which must be ASCII. */ -function slugify(name: string): string { +export function slugify(name: string): string { return name .toLowerCase() + .replace(/[^\p{ASCII}]/gu, (c) => `-${c.codePointAt(0)?.toString(16)}-`) .replace(/[^a-z0-9]+/g, "-") .replace(/^-|-$/g, "") } @@ -189,6 +192,7 @@ export async function loadFlattenedServerModels(): Promise< const defaultModelId = process.env.AI_MODEL const flattened: FlattenedServerModel[] = [] + const seenIds = new Set() for (const p of cfg.providers) { const providerLabel = @@ -199,6 +203,16 @@ export async function loadFlattenedServerModels(): Promise< for (const modelId of p.models) { const id = `server:${nameSlug}:${modelId}` + // Names that differ only in case or punctuation share a slug. + // A repeated id would always resolve to the first provider's + // credentials, so drop it instead. + if (seenIds.has(id)) { + console.warn( + `[server-model-config] Skipping duplicate model id "${id}". Provider names must differ in letters or digits.`, + ) + continue + } + seenIds.add(id) // Default model priority: // 1. From ai-models.json: first model of provider with default: true diff --git a/proxy.ts b/proxy.ts index 7b616fe..fdf43a4 100644 --- a/proxy.ts +++ b/proxy.ts @@ -48,13 +48,12 @@ export function proxy(request: NextRequest) { if (pathnameIsMissingLocale) { const locale = getLocale(request) - // Redirect to localized path - return NextResponse.redirect( - new URL( - `/${locale}${pathname.startsWith("/") ? "" : "/"}${pathname}`, - request.url, - ), - ) + // Redirect to localized path. Cloning nextUrl keeps the basePath + // (NEXT_PUBLIC_BASE_PATH) and query string, which + // new URL("/...", request.url) would drop. + const url = request.nextUrl.clone() + url.pathname = `/${locale}${pathname}` + return NextResponse.redirect(url) } } diff --git a/tests/unit/admin-providers-route.test.ts b/tests/unit/admin-providers-route.test.ts new file mode 100644 index 0000000..c0ad3d1 --- /dev/null +++ b/tests/unit/admin-providers-route.test.ts @@ -0,0 +1,71 @@ +// @vitest-environment node +import fs from "fs" +import os from "os" +import path from "path" +import { afterEach, beforeEach, describe, expect, it } from "vitest" +import { GET, PUT } from "@/app/api/admin/providers/route" +import { _resetForTests } from "@/lib/admin/settings" + +let tmpDir: string + +beforeEach(() => { + tmpDir = fs.mkdtempSync(path.join(os.tmpdir(), "admin-providers-route-")) + process.env.SETTINGS_FILE = path.join(tmpDir, "settings.json") + process.env.ADMIN_PASSWORD = "pw" + process.env.AI_MODELS_CONFIG_PATH = path.join(tmpDir, "none.json") + _resetForTests() +}) + +afterEach(() => { + _resetForTests() + delete process.env.SETTINGS_FILE + delete process.env.ADMIN_PASSWORD + delete process.env.AI_MODELS_CONFIG_PATH + delete process.env.AI_MODEL + fs.rmSync(tmpDir, { recursive: true, force: true }) +}) + +const headers = { "x-admin-password": "pw" } + +async function saveDefaultPanelProvider() { + const res = await PUT( + new Request("http://localhost/api/admin/providers", { + method: "PUT", + headers: { ...headers, "Content-Type": "application/json" }, + body: JSON.stringify({ + providers: [ + { + id: "p1", + provider: "openai", + apiKey: "sk-test", + models: ["gpt-panel"], + isDefault: true, + }, + ], + }), + }), + ) + expect(res.status).toBe(200) +} + +async function envHasDefaultModel(): Promise { + const res = await GET( + new Request("http://localhost/api/admin/providers", { headers }), + ) + return (await res.json()).envHasDefaultModel +} + +describe("envHasDefaultModel", () => { + it("is true when .env sets AI_MODEL, even after a panel default", async () => { + process.env.AI_MODEL = "gpt-env" + expect(await envHasDefaultModel()).toBe(true) + await saveDefaultPanelProvider() + expect(await envHasDefaultModel()).toBe(true) + }) + + it("ignores the AI_MODEL the panel default writes", async () => { + await saveDefaultPanelProvider() + expect(process.env.AI_MODEL).toBe("gpt-panel") + expect(await envHasDefaultModel()).toBe(false) + }) +}) diff --git a/tests/unit/admin-providers.test.ts b/tests/unit/admin-providers.test.ts index 9a552d4..0602d96 100644 --- a/tests/unit/admin-providers.test.ts +++ b/tests/unit/admin-providers.test.ts @@ -69,7 +69,7 @@ describe("deriveEnvUpdates", () => { expect(updates.ADMIN_OPENAI_API_KEY_2).toBe("sk-second") }) - it("maps bedrock credentials to AWS env vars", () => { + it("maps bedrock credentials to ADMIN_AWS_* env vars", () => { const updates = deriveEnvUpdates( [ provider({ @@ -83,9 +83,26 @@ describe("deriveEnvUpdates", () => { ], [], ) - expect(updates.AWS_ACCESS_KEY_ID).toBe("AKIA123") - expect(updates.AWS_SECRET_ACCESS_KEY).toBe("secret") - expect(updates.AWS_REGION).toBe("us-west-2") + expect(updates.ADMIN_AWS_ACCESS_KEY_ID).toBe("AKIA123") + expect(updates.ADMIN_AWS_SECRET_ACCESS_KEY).toBe("secret") + expect(updates.ADMIN_AWS_REGION).toBe("us-west-2") + // Standard AWS vars are left to the environment + expect(updates.AWS_ACCESS_KEY_ID).toBeUndefined() + }) + + it("clears AWS_* bedrock keys written by older versions", () => { + const bedrock = provider({ + provider: "bedrock", + apiKey: undefined, + awsAccessKeyId: "AKIA123", + awsSecretAccessKey: "secret", + models: ["claude-x"], + }) + const updates = deriveEnvUpdates([bedrock], [bedrock]) + expect(updates.AWS_ACCESS_KEY_ID).toBeNull() + expect(updates.AWS_SECRET_ACCESS_KEY).toBeNull() + expect(updates.AWS_REGION).toBeNull() + expect(updates.ADMIN_AWS_ACCESS_KEY_ID).toBe("AKIA123") }) it("clears keys owned by the previous list when providers are removed", () => { @@ -315,6 +332,32 @@ describe("validateAdminProviders", () => { expect(validateAdminProviders(list)).toMatch(/unique/) }) + it("rejects names that differ only in case or punctuation", () => { + const list = [ + provider({ id: "p1", name: "Open AI" }), + provider({ id: "p2", name: "open-ai" }), + ] + expect(validateAdminProviders(list)).toMatch(/unique/) + }) + + it("rejects a case-only clash with an env-configured name", () => { + expect( + validateAdminProviders([provider({ name: "openai" })], { + providers: [ + { name: "OpenAI", provider: "openai", models: ["gpt-x"] }, + ], + }), + ).toMatch(/already defined/) + }) + + it("accepts distinct CJK names", () => { + const list = [ + provider({ id: "p1", provider: "deepseek", name: "主力" }), + provider({ id: "p2", provider: "deepseek", name: "备用" }), + ] + expect(validateAdminProviders(list)).toBeNull() + }) + it("rejects multiple defaults", () => { const list = [ provider({ id: "p1", isDefault: true }), diff --git a/tests/unit/admin-settings.test.ts b/tests/unit/admin-settings.test.ts index 8610751..c9fe084 100644 --- a/tests/unit/admin-settings.test.ts +++ b/tests/unit/admin-settings.test.ts @@ -1,7 +1,7 @@ import fs from "fs" import os from "os" import path from "path" -import { afterEach, beforeEach, describe, expect, it } from "vitest" +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest" import { _resetForTests, applyToEnv, @@ -100,6 +100,24 @@ describe("applyToEnv / saveSettings", () => { expect(process.env.TEST_ADMIN_VAR).toBeUndefined() }) + it("a second module instance can remove a key the first one overlaid", async () => { + // instrumentation.ts and API routes load separate copies in a build + process.env.TEST_ADMIN_VAR = "from-env" + fs.writeFileSync( + process.env.SETTINGS_FILE!, + JSON.stringify({ version: 1, values: { TEST_ADMIN_VAR: "abc" } }), + ) + applyToEnv() + expect(process.env.TEST_ADMIN_VAR).toBe("abc") + + vi.resetModules() + const second = await import("@/lib/admin/settings") + expect(second.getValueSource("TEST_ADMIN_VAR")).toBe("file") + second.saveSettings({ TEST_ADMIN_VAR: null }) + expect(process.env.TEST_ADMIN_VAR).toBe("from-env") + expect(second.getEnvFallback("TEST_ADMIN_VAR")).toBe("from-env") + }) + it("persists across cache reset (file round-trip)", () => { saveSettings({ TEST_ADMIN_VAR: "persisted" }) _resetForTests() diff --git a/tests/unit/api-access-code.test.ts b/tests/unit/api-access-code.test.ts new file mode 100644 index 0000000..10432a7 --- /dev/null +++ b/tests/unit/api-access-code.test.ts @@ -0,0 +1,156 @@ +// @vitest-environment node +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest" +import { POST as parseUrl } from "@/app/api/parse-url/route" +import { POST as validateDiagram } from "@/app/api/validate-diagram/route" +import { POST as validateModel } from "@/app/api/validate-model/route" +import { POST as verifyAccessCode } from "@/app/api/verify-access-code/route" +import { checkAccessCode } from "@/lib/access-code" + +// Treat every URL as public so no test hits DNS +vi.mock("@/lib/ssrf-protection", async (importOriginal) => ({ + ...(await importOriginal()), + isPrivateUrl: async () => false, +})) + +function post(path: string, body: unknown, accessCode?: string): Request { + return new Request(`http://localhost${path}`, { + method: "POST", + headers: { + "Content-Type": "application/json", + ...(accessCode ? { "x-access-code": accessCode } : {}), + }, + body: JSON.stringify(body), + }) +} + +beforeEach(() => { + process.env.ACCESS_CODE_LIST = "secret, other" +}) + +afterEach(() => { + delete process.env.ACCESS_CODE_LIST + delete process.env.ALLOW_PRIVATE_URLS + vi.unstubAllGlobals() +}) + +describe("checkAccessCode", () => { + it("passes when no access codes are configured", () => { + delete process.env.ACCESS_CODE_LIST + expect(checkAccessCode(post("/x", {}))).toBeNull() + }) + + it("rejects a missing or wrong code and accepts a listed one", () => { + expect(checkAccessCode(post("/x", {}))?.status).toBe(401) + expect(checkAccessCode(post("/x", {}, "nope"))?.status).toBe(401) + expect(checkAccessCode(post("/x", {}, "other"))).toBeNull() + }) +}) + +describe("routes that spend server resources require the access code", () => { + it("parse-url", async () => { + const res = await parseUrl( + post("/api/parse-url", { url: "https://example.com" }), + ) + expect(res.status).toBe(401) + }) + + it("validate-diagram", async () => { + const res = await validateDiagram( + post("/api/validate-diagram", { + imageData: "data:image/png;base64,AAAA", + }), + ) + expect(res.status).toBe(401) + }) + + it("validate-model", async () => { + const res = await validateModel( + post("/api/validate-model", { + provider: "openai", + apiKey: "sk", + modelId: "m", + }), + ) + expect(res.status).toBe(401) + }) + + it("verify-access-code", async () => { + const bad = await verifyAccessCode(post("/api/verify-access-code", {})) + expect(bad.status).toBe(401) + expect((await bad.json()).valid).toBe(false) + + const good = await verifyAccessCode( + post("/api/verify-access-code", {}, "secret"), + ) + expect((await good.json()).valid).toBe(true) + }) +}) + +describe("size limits", () => { + it("parse-url stops reading a body over the download limit", async () => { + const chunk = new Uint8Array(1024 * 1024) + let sent = 0 + const body = new ReadableStream({ + pull(controller) { + sent++ + controller.enqueue(chunk) + }, + }) + vi.stubGlobal( + "fetch", + vi.fn( + async () => + new Response(body, { + headers: { "content-type": "text/html" }, + }), + ), + ) + + const res = await parseUrl( + post("/api/parse-url", { url: "https://example.com" }, "secret"), + ) + expect(res.status).toBe(413) + expect(sent).toBeLessThan(10) + }) + + it("validate-diagram rejects oversized image data", async () => { + const imageData = `data:image/png;base64,${"A".repeat(6 * 1024 * 1024)}` + const res = await validateDiagram( + post("/api/validate-diagram", { imageData }, "secret"), + ) + expect(res.status).toBe(413) + }) +}) + +describe("validate-model redirects", () => { + it("refuses redirects when private URLs are blocked", async () => { + process.env.ALLOW_PRIVATE_URLS = "false" + vi.stubGlobal( + "fetch", + vi.fn( + async () => + new Response(null, { + status: 302, + headers: { location: "http://169.254.169.254/" }, + }), + ), + ) + + const res = await validateModel( + post( + "/api/validate-model", + { + provider: "openai", + apiKey: "sk", + modelId: "m", + baseUrl: "https://attacker.example/v1", + }, + "secret", + ), + ) + const data = await res.json() + expect(data.valid).toBe(false) + expect(data.error).toMatch(/Redirects are not allowed/) + expect(fetch).toHaveBeenCalledTimes(1) + }) +}) diff --git a/tests/unit/edgeone-function.test.ts b/tests/unit/edgeone-function.test.ts new file mode 100644 index 0000000..e0d9649 --- /dev/null +++ b/tests/unit/edgeone-function.test.ts @@ -0,0 +1,47 @@ +// @vitest-environment node +import { describe, expect, it } from "vitest" +import { onRequest } from "@/edge-functions/api/edgeai/chat/completions" + +function request(headers: Record): Request { + return new Request("http://localhost/api/edgeai/chat/completions", { + method: "POST", + headers, + // Non-streaming requests return a mock reply without calling AI + body: JSON.stringify({ messages: [{ role: "user", content: "hi" }] }), + }) +} + +const json = { "Content-Type": "application/json" } + +describe("EdgeOne chat completions function", () => { + it("sends no CORS headers", async () => { + const res = await onRequest({ request: request(json), env: {} }) + expect(res.status).toBe(200) + expect(res.headers.get("access-control-allow-origin")).toBeNull() + }) + + it("rejects non-JSON requests", async () => { + const res = await onRequest({ + request: request({ "Content-Type": "text/plain" }), + env: {}, + }) + expect(res.status).toBe(400) + }) + + it("checks the access code when ACCESS_CODE_LIST is set", async () => { + const env = { ACCESS_CODE_LIST: "secret" } + const missing = await onRequest({ request: request(json), env }) + expect(missing.status).toBe(401) + + const ok = await onRequest({ + request: request({ ...json, "x-access-code": "secret" }), + env, + }) + expect(ok.status).toBe(200) + }) + + it("lets requests through when env is unavailable", async () => { + const res = await onRequest({ request: request(json) }) + expect(res.status).toBe(200) + }) +}) diff --git a/tests/unit/server-model-config.test.ts b/tests/unit/server-model-config.test.ts index 80d0c87..d94c87b 100644 --- a/tests/unit/server-model-config.test.ts +++ b/tests/unit/server-model-config.test.ts @@ -4,6 +4,7 @@ import { loadFlattenedServerModels, type ServerModelsConfig, ServerModelsConfigSchema, + slugify, } from "@/lib/server-model-config" const ORIGINAL_ENV = { ...process.env } @@ -233,3 +234,50 @@ describe("loadFlattenedServerModels", () => { expect(models[0].apiKeyEnv).toEqual(["OPENAI_KEY_1", "OPENAI_KEY_2"]) }) }) + +describe("slugify", () => { + it("keeps ASCII names readable", () => { + expect(slugify("OpenAI Production")).toBe("openai-production") + }) + + it("gives distinct ASCII slugs to distinct CJK names", () => { + const slugs = ["主力", "备用", "DeepSeek 官方", "DeepSeek 备用"].map( + slugify, + ) + expect(new Set(slugs).size).toBe(4) + for (const slug of slugs) expect(slug).toMatch(/^[a-z0-9-]+$/) + }) +}) + +describe("loadFlattenedServerModels id collisions", () => { + it("drops a model whose id repeats an earlier provider's", async () => { + const config: ServerModelsConfig = { + providers: [ + { name: "OpenAI", provider: "openai", models: ["gpt-4o"] }, + { + name: "openai", + provider: "openai", + models: ["gpt-4o"], + apiKeyEnv: "OTHER_KEY", + }, + { + name: "主力", + provider: "deepseek", + models: ["deepseek-chat"], + }, + { + name: "备用", + provider: "deepseek", + models: ["deepseek-chat"], + }, + ], + } + process.env.AI_MODELS_CONFIG = JSON.stringify(config) + + const models = await loadFlattenedServerModels() + const ids = models.map((m) => m.id) + expect(new Set(ids).size).toBe(ids.length) + expect(ids).toHaveLength(3) + expect(models[0].apiKeyEnv).toBeUndefined() + }) +})