= {}
+
+ for (const [key, value] of Object.entries(body.values)) {
+ const def = SETTINGS_BY_KEY.get(key)
+ if (!def) {
+ errors[key] = "Unknown setting"
+ continue
+ }
+ if (value === null || value === "") {
+ updates[key] = null
+ continue
+ }
+ if (typeof value !== "string") {
+ errors[key] = "Value must be a string"
+ continue
+ }
+ const error = validateValue(def, value)
+ if (error) {
+ errors[key] = error
+ continue
+ }
+ updates[key] = value
+ }
+
+ if (Object.keys(errors).length > 0) {
+ return Response.json({ errors }, { status: 400 })
+ }
+
+ saveSettings(updates)
+
+ return Response.json({
+ writable: true,
+ settings: serializeSettings(),
+ })
+}
diff --git a/app/api/admin/test-model/route.ts b/app/api/admin/test-model/route.ts
new file mode 100644
index 00000000..51999ca8
--- /dev/null
+++ b/app/api/admin/test-model/route.ts
@@ -0,0 +1,84 @@
+import { POST as validateModel } from "@/app/api/validate-model/route"
+import { checkAdminAuth } from "@/lib/admin/auth"
+import {
+ AdminProviderSchema,
+ loadAdminProviders,
+ mergeSecrets,
+} from "@/lib/admin/providers"
+import { globalBaseUrl } from "@/lib/ai-providers"
+
+export const runtime = "nodejs"
+export const dynamic = "force-dynamic"
+
+// Test a model with the client's CURRENT provider state (which may be
+// unsaved). Secret fields arrive either as plaintext (newly typed) or as
+// masked {isSet} markers, which are resolved against settings.json — so
+// testing works both before and after saving.
+export async function POST(req: Request) {
+ const authError = checkAdminAuth(req)
+ if (authError) return authError
+
+ let body: { provider?: unknown; modelId?: string }
+ try {
+ body = await req.json()
+ } catch {
+ return Response.json({ error: "Invalid JSON body" }, { status: 400 })
+ }
+
+ const parsed = AdminProviderSchema.safeParse(body.provider)
+ if (!parsed.success || !body.modelId) {
+ return Response.json(
+ { valid: false, error: "Invalid provider or model" },
+ { status: 400 },
+ )
+ }
+
+ // SECURITY: a stored secret is only resolved from an {isSet} marker if
+ // the endpoint it would be sent to (provider + baseUrl) still matches
+ // the stored entry. Otherwise a tampered baseUrl could exfiltrate the
+ // stored key to an arbitrary host. Mismatches must re-supply plaintext.
+ const stored = loadAdminProviders().find((p) => p.id === parsed.data.id)
+ const sameEndpoint =
+ stored &&
+ stored.provider === parsed.data.provider &&
+ (stored.baseUrl ?? "") === (parsed.data.baseUrl ?? "") &&
+ (stored.awsRegion ?? "") === (parsed.data.awsRegion ?? "")
+ const [resolved] = mergeSecrets(
+ [parsed.data],
+ sameEndpoint && stored ? [stored] : [],
+ )
+
+ const serverUrl = globalBaseUrl(resolved.provider)
+ return validateModel(
+ new Request(new URL("/api/validate-model", req.url), {
+ method: "POST",
+ headers: {
+ "Content-Type": "application/json",
+ // Checked again there, in place of an access code
+ "x-admin-password": req.headers.get("x-admin-password") || "",
+ // The EdgeOne function checks the access code and Pages
+ // cookies, and its URL is built from the page's origin
+ "x-access-code": req.headers.get("x-access-code") || "",
+ cookie: req.headers.get("cookie") || "",
+ ...(req.headers.get("origin") && {
+ origin: req.headers.get("origin") as string,
+ }),
+ },
+ body: JSON.stringify({
+ provider: resolved.provider,
+ apiKey: resolved.apiKey,
+ // Without a URL of its own, chat sends the entry's key to
+ // the server's _BASE_URL: test that endpoint, not
+ // another one. It is the server's own, which chat uses
+ // without the checks for a URL a user typed.
+ baseUrl: resolved.baseUrl || serverUrl,
+ ...(!resolved.baseUrl && serverUrl && { serverBaseUrl: true }),
+ modelId: body.modelId,
+ awsAccessKeyId: resolved.awsAccessKeyId,
+ awsSecretAccessKey: resolved.awsSecretAccessKey,
+ awsRegion: resolved.awsRegion,
+ vertexApiKey: resolved.vertexApiKey,
+ }),
+ }),
+ )
+}
diff --git a/app/api/chat/route.ts b/app/api/chat/route.ts
index 872b8702..c5ea47a2 100644
--- a/app/api/chat/route.ts
+++ b/app/api/chat/route.ts
@@ -4,40 +4,66 @@ import {
createUIMessageStream,
createUIMessageStreamResponse,
InvalidToolInputError,
- LoadAPIKeyError,
stepCountIs,
streamText,
} from "ai"
-import fs from "fs/promises"
import { jsonrepair } from "jsonrepair"
import path from "path"
import { z } from "zod"
+import { checkAccessCode, rejectCrossSite } from "@/lib/access-code"
import {
+ CACHE_POINT,
+ edgeOneEndpoint,
getAIModel,
- supportsImageInput,
+ getServerProvider,
+ SINGLE_SYSTEM_PROVIDERS,
supportsPromptCaching,
+ usesServerCredentials,
+ usesServerEndpoint,
} from "@/lib/ai-providers"
import { findCachedResponse } from "@/lib/cached-responses"
import {
- isMinimalDiagram,
+ dropInvalidToolCalls,
+ fixToolInputJson,
replaceHistoricalToolInputs,
validateFileParts,
} from "@/lib/chat-helpers"
+import { withDeprecatedParamsFallback } from "@/lib/deprecated-params"
import {
checkAndIncrementRequest,
isQuotaEnabled,
recordTokenUsage,
} from "@/lib/dynamo-quota-manager"
import {
+ endTrace,
getTelemetryConfig,
setTraceInput,
setTraceOutput,
wrapWithObserve,
} from "@/lib/langfuse"
+import { classifyLLMError, streamErrorText } from "@/lib/llm-errors"
+import {
+ resolveMaxOutputTokens,
+ withOutputTokenLimitFallback,
+} from "@/lib/output-token-limit"
+import {
+ type FlattenedServerModel,
+ findServerModelById,
+} from "@/lib/server-model-config"
+import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection"
import { getSystemPrompt } from "@/lib/system-prompts"
+import { normalizeBaseUrl } from "@/lib/types/model-config"
import { getUserIdFromRequest } from "@/lib/user-id"
+import { hasCells } from "@/packages/mcp-server/src/pages.ts"
+import {
+ getShapeLibrary,
+ SHAPE_LIBRARY_LIST,
+} from "@/packages/mcp-server/src/shape-library.ts"
+import { SWIMLANE_EXAMPLE } from "@/packages/mcp-server/src/xml-examples.ts"
-export const maxDuration = 120
+// No explicit cap: a reasoning model can spend minutes planning before it emits
+// the tool call, so take whatever the host allows. Vercel's own default is 300s,
+// which is also where Node's response-body timeout on the upstream stream lands.
// Helper function to create cached stream response
function createCachedStreamResponse(xml: string): Response {
@@ -69,26 +95,25 @@ function createCachedStreamResponse(xml: string): Response {
return createUIMessageStreamResponse({ stream })
}
-// Inner handler function
-async function handleChatRequest(req: Request): Promise {
- // Check for access code
- const accessCodes =
- process.env.ACCESS_CODE_LIST?.split(",")
- .map((code) => code.trim())
- .filter(Boolean) || []
- if (accessCodes.length > 0) {
- const accessCodeHeader = req.headers.get("x-access-code")
- if (!accessCodeHeader || !accessCodes.includes(accessCodeHeader)) {
- return Response.json(
- {
- error: "Invalid or missing access code. Please configure it in Settings.",
- },
- { status: 401 },
- )
- }
- }
+// Responses streamed from the model, whose trace streamText's callbacks end
+const modelStreamResponses = new WeakSet()
- const { messages, xml, previousXml, sessionId } = await req.json()
+// Inner handler function
+const DEBUG_LLM_PAYLOAD = process.env.DEBUG_LLM_PAYLOAD === "true"
+
+async function handleChatRequest(req: Request): Promise {
+ const crossSite = rejectCrossSite(req)
+ if (crossSite) return crossSite
+ // Check for access code
+ const accessDenied = checkAccessCode(req)
+ if (accessDenied) return accessDenied
+
+ const body = await req.json()
+ const { messages, xml, previousXml, sessionId } = body
+ const customSystemMessage =
+ typeof body.customSystemMessage === "string"
+ ? body.customSystemMessage.slice(0, 5000)
+ : ""
// Get user ID for Langfuse tracking and quota
const userId = getUserIdFromRequest(req)
@@ -114,14 +139,165 @@ async function handleChatRequest(req: Request): Promise {
userId: userId,
})
- // === SERVER-SIDE QUOTA CHECK START ===
- // Quota is opt-in: only enabled when DYNAMODB_QUOTA_TABLE env var is set
- const hasOwnApiKey = !!(
- req.headers.get("x-ai-provider") && req.headers.get("x-ai-api-key")
+ // === FILE VALIDATION START ===
+ const fileValidation = validateFileParts(messages)
+ if (!fileValidation.valid) {
+ return Response.json({ error: fileValidation.error }, { status: 400 })
+ }
+ // === FILE VALIDATION END ===
+
+ // === CACHE CHECK START ===
+ const isFirstMessage = messages.length === 1
+ const isEmptyDiagram = !xml || !hasCells(xml)
+
+ if (isFirstMessage && isEmptyDiagram) {
+ const lastMessage = messages[0]
+ const textPart = lastMessage.parts?.find((p: any) => p.type === "text")
+ const filePart = lastMessage.parts?.find((p: any) => p.type === "file")
+
+ const cached = findCachedResponse(textPart?.text || "", !!filePart)
+
+ if (cached) {
+ return createCachedStreamResponse(cached.xml)
+ }
+ }
+ // === CACHE CHECK END ===
+
+ // Read client AI provider overrides from headers
+ const provider = req.headers.get("x-ai-provider")
+ let baseUrl = req.headers.get("x-ai-base-url")
+ const selectedModelId = req.headers.get("x-selected-model-id")
+
+ // Check if this is a server model with custom env var names
+ let serverModelConfig: {
+ apiKeyEnv?: string | string[]
+ baseUrlEnv?: string
+ provider?: string
+ } = {}
+ let serverModel: FlattenedServerModel | null = null
+ if (selectedModelId?.startsWith("server:")) {
+ serverModel = await findServerModelById(selectedModelId)
+ console.log(
+ `[Server Model Lookup] ID: ${selectedModelId}, Found: ${!!serverModel}, Provider: ${serverModel?.provider}`,
+ )
+ if (serverModel) {
+ serverModelConfig = {
+ apiKeyEnv: serverModel.apiKeyEnv,
+ baseUrlEnv: serverModel.baseUrlEnv,
+ // Use actual provider from config (client header may have incorrect value due to ID format change)
+ provider: serverModel.provider,
+ }
+ }
+ }
+
+ // A server model's provider comes from its config: for one set up in
+ // the admin panel the header holds the provider name's slug. Without
+ // either, the server's own AI_PROVIDER.
+ const isEdgeOne =
+ (serverModelConfig.provider || provider || getServerProvider()) ===
+ "edgeone"
+
+ // EdgeOne is this deployment's own function, whatever URL the request
+ // names: another host would get the user's EdgeOne cookies, and the
+ // quota counts it. Absolute, as the SDK needs.
+ if (isEdgeOne) baseUrl = edgeOneEndpoint(req)
+
+ // Same rule as validate-model: with ALLOW_PRIVATE_URLS=false a request may
+ // not point the server at a private or internal address
+ if (baseUrl && !allowPrivateUrls() && (await isPrivateUrl(baseUrl))) {
+ return Response.json(
+ { error: "Private or internal base URLs are not allowed." },
+ { status: 400 },
+ )
+ }
+
+ // Get cookie header for EdgeOne authentication (eo_token, eo_time)
+ const cookieHeader = req.headers.get("cookie")
+
+ const clientOverrides = {
+ // Server model provider takes precedence over client header; EdgeOne
+ // named only in AI_PROVIDER is named here, for its own base URL
+ provider:
+ serverModelConfig.provider ||
+ provider ||
+ (isEdgeOne ? "edgeone" : null),
+ baseUrl,
+ apiKey: req.headers.get("x-ai-api-key"),
+ // A server model runs the model it was configured with, whatever the header says
+ modelId: serverModel?.modelId || req.headers.get("x-ai-model"),
+ // AWS Bedrock credentials
+ awsAccessKeyId: req.headers.get("x-aws-access-key-id"),
+ awsSecretAccessKey: req.headers.get("x-aws-secret-access-key"),
+ awsRegion: req.headers.get("x-aws-region"),
+ awsSessionToken: req.headers.get("x-aws-session-token"),
+ // Server model custom env var names
+ ...serverModelConfig,
+ // Vertex AI credentials (Express Mode)
+ vertexApiKey: req.headers.get("x-vertex-api-key"),
+ // Pass cookies for EdgeOne Pages authentication, and the access code,
+ // which the EdgeOne function checks too
+ ...(isEdgeOne && {
+ headers: {
+ ...(cookieHeader && { cookie: cookieHeader }),
+ "x-access-code": req.headers.get("x-access-code") || "",
+ },
+ }),
+ }
+
+ // Read minimal style preference from header
+ const minimalStyle = req.headers.get("x-minimal-style") === "true"
+
+ console.log(
+ `[Client Overrides] provider: ${clientOverrides.provider}, modelId: ${clientOverrides.modelId}`,
)
- // Skip quota check if: quota disabled, user has own API key, or is anonymous
- if (isQuotaEnabled() && !hasOwnApiKey && userId !== "anonymous") {
+ // Get AI model with optional client overrides
+ const {
+ model: baseModel,
+ providerOptions,
+ modelId,
+ provider: resolvedProvider,
+ } = getAIModel(clientOverrides)
+
+ // On the server's own keys, only run models the server offers: a server
+ // model picked by id (its model name is fixed above) or one in AI_MODEL
+ // on AI_PROVIDER. With their own key, users can run any model.
+ const onServerCredentials = usesServerCredentials(
+ resolvedProvider,
+ clientOverrides,
+ )
+ const envModels =
+ process.env.AI_MODEL?.split(",").map((m) => m.trim()) || []
+ const offeredInEnv =
+ envModels.includes(modelId) && resolvedProvider === getServerProvider()
+ if (onServerCredentials && !serverModel && !offeredInEnv) {
+ return Response.json(
+ {
+ error: `Model "${modelId}" is not available on this server. Add your own API key in Settings to use it.`,
+ },
+ { status: 400 },
+ )
+ }
+
+ // === SERVER-SIDE QUOTA CHECK START ===
+ // Quota is opt-in (DYNAMODB_QUOTA_TABLE) and counts what runs on the
+ // server's keys, or on the server's own endpoints: EdgeOne, its keyless
+ // Ollama, and anything at a private address (the server's network,
+ // which ignores a dummy key header). Bedrock and EdgeOne never use the
+ // base URL header. In the desktop app every endpoint is the user's.
+ const clientBaseUrl = normalizeBaseUrl(
+ req.headers.get("x-ai-base-url") ?? "",
+ )
+ const onServerEndpoint = await usesServerEndpoint(
+ resolvedProvider,
+ clientBaseUrl,
+ clientOverrides.apiKey,
+ )
+ const countsQuota =
+ isQuotaEnabled() &&
+ (onServerCredentials || onServerEndpoint) &&
+ userId !== "anonymous"
+ if (countsQuota) {
const quotaCheck = await checkAndIncrementRequest(userId, {
requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10,
tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000,
@@ -141,67 +317,20 @@ async function handleChatRequest(req: Request): Promise {
}
// === SERVER-SIDE QUOTA CHECK END ===
- // === FILE VALIDATION START ===
- const fileValidation = validateFileParts(messages)
- if (!fileValidation.valid) {
- return Response.json({ error: fileValidation.error }, { status: 400 })
- }
- // === FILE VALIDATION END ===
+ // Retry once if the provider rejects the requested budget, or (newer
+ // Claude models) the sampling or thinking settings
+ const model = withOutputTokenLimitFallback(
+ withDeprecatedParamsFallback(baseModel),
+ )
- // === CACHE CHECK START ===
- const isFirstMessage = messages.length === 1
- const isEmptyDiagram = !xml || xml.trim() === "" || isMinimalDiagram(xml)
-
- if (isFirstMessage && isEmptyDiagram) {
- const lastMessage = messages[0]
- const textPart = lastMessage.parts?.find((p: any) => p.type === "text")
- const filePart = lastMessage.parts?.find((p: any) => p.type === "file")
-
- const cached = findCachedResponse(textPart?.text || "", !!filePart)
-
- if (cached) {
- return createCachedStreamResponse(cached.xml)
- }
- }
- // === CACHE CHECK END ===
-
- // Read client AI provider overrides from headers
- const provider = req.headers.get("x-ai-provider")
- let baseUrl = req.headers.get("x-ai-base-url")
-
- // For EdgeOne provider, construct full URL from request origin
- // because createOpenAI needs absolute URL, not relative path
- if (provider === "edgeone" && !baseUrl) {
- const origin = req.headers.get("origin") || new URL(req.url).origin
- baseUrl = `${origin}/api/edgeai`
- }
-
- // Get cookie header for EdgeOne authentication (eo_token, eo_time)
- const cookieHeader = req.headers.get("cookie")
-
- const clientOverrides = {
- provider,
- baseUrl,
- apiKey: req.headers.get("x-ai-api-key"),
- modelId: req.headers.get("x-ai-model"),
- // AWS Bedrock credentials
- awsAccessKeyId: req.headers.get("x-aws-access-key-id"),
- awsSecretAccessKey: req.headers.get("x-aws-secret-access-key"),
- awsRegion: req.headers.get("x-aws-region"),
- awsSessionToken: req.headers.get("x-aws-session-token"),
- // Pass cookies for EdgeOne Pages authentication
- ...(provider === "edgeone" &&
- cookieHeader && {
- headers: { cookie: cookieHeader },
- }),
- }
-
- // Read minimal style preference from header
- const minimalStyle = req.headers.get("x-minimal-style") === "true"
-
- // Get AI model with optional client overrides
- const { model, providerOptions, headers, modelId } =
- getAIModel(clientOverrides)
+ // The user setting can raise the budget only on their own key (in the
+ // desktop app every key is the user's); on the server's keys or own
+ // endpoints it can only lower it
+ const maxOutputTokens = resolveMaxOutputTokens(
+ req.headers.get("x-max-output-tokens"),
+ onServerCredentials || onServerEndpoint,
+ )
+ console.log(`[maxOutputTokens] ${maxOutputTokens}`)
// Check if model supports prompt caching
const shouldCache = supportsPromptCaching(modelId)
@@ -211,22 +340,19 @@ async function handleChatRequest(req: Request): Promise {
// Get the appropriate system prompt based on model (extended for Opus/Haiku 4.5)
const systemMessage = getSystemPrompt(modelId, minimalStyle)
+ const finalSystemMessage = customSystemMessage
+ ? `${systemMessage}\n\n## Custom Instructions\n${customSystemMessage}`
+ : systemMessage
// Extract file parts (images) from the last user message
const fileParts =
lastUserMessage?.parts?.filter((part: any) => part.type === "file") ||
[]
- // Check if user is sending images to a model that doesn't support them
- // AI SDK silently drops unsupported parts, so we need to catch this early
- if (fileParts.length > 0 && !supportsImageInput(modelId)) {
- return Response.json(
- {
- error: `The model "${modelId}" does not support image input. Please use a vision-capable model (e.g., GPT-4o, Claude, Gemini) or remove the image.`,
- },
- { status: 400 },
- )
- }
+ // Note: we used to pre-emptively reject images for models we guessed were
+ // text-only (by name matching). That heuristic misfired on newer models
+ // (see issue #874), so we now let the request through and surface the real
+ // provider error if the model genuinely can't accept images.
// User input only - XML is now in a separate cached system message
const formattedUserInput = `User input:
@@ -234,39 +360,46 @@ async function handleChatRequest(req: Request): Promise {
${userInputText}
"""`
- // Convert UIMessages to ModelMessages and add system message
- const modelMessages = await convertToModelMessages(messages)
-
- // DEBUG: Log incoming messages structure
- console.log("[route.ts] Incoming messages count:", messages.length)
- messages.forEach((msg: any, idx: number) => {
- console.log(
- `[route.ts] Message ${idx} role:`,
- msg.role,
- "parts count:",
- msg.parts?.length,
- )
- if (msg.parts) {
- msg.parts.forEach((part: any, partIdx: number) => {
- if (
- part.type === "tool-invocation" ||
- part.type === "tool-result"
- ) {
- console.log(`[route.ts] Part ${partIdx}:`, {
- type: part.type,
- toolName: part.toolName,
- hasInput: !!part.input,
- inputType: typeof part.input,
- inputKeys:
- part.input && typeof part.input === "object"
- ? Object.keys(part.input)
- : null,
- })
- }
- })
- }
+ // Convert UIMessages to ModelMessages and add system message. A tool
+ // call that never got its result (the user stopped while it ran) is
+ // left out: the SDK would refuse this and every later request of the
+ // chat (MissingToolResultsError)
+ const modelMessages = await convertToModelMessages(messages, {
+ ignoreIncompleteToolCalls: true,
})
+ // DEBUG_LLM_PAYLOAD=true logs the incoming message structure
+ if (DEBUG_LLM_PAYLOAD) {
+ console.log("[route.ts] Incoming messages count:", messages.length)
+ messages.forEach((msg: any, idx: number) => {
+ console.log(
+ `[route.ts] Message ${idx} role:`,
+ msg.role,
+ "parts count:",
+ msg.parts?.length,
+ )
+ if (msg.parts) {
+ msg.parts.forEach((part: any, partIdx: number) => {
+ if (
+ part.type === "tool-invocation" ||
+ part.type === "tool-result"
+ ) {
+ console.log(`[route.ts] Part ${partIdx}:`, {
+ type: part.type,
+ toolName: part.toolName,
+ hasInput: !!part.input,
+ inputType: typeof part.input,
+ inputKeys:
+ part.input && typeof part.input === "object"
+ ? Object.keys(part.input)
+ : null,
+ })
+ }
+ })
+ }
+ })
+ }
+
// Replace historical tool call XML with placeholders to reduce tokens
// Disabled by default - some models (e.g. minimax) copy placeholders instead of generating XML
const enableHistoryReplace =
@@ -283,61 +416,43 @@ ${userInputText}
)
// Filter out tool-calls with invalid inputs (from failed repair or interrupted streaming)
- // Bedrock API rejects messages where toolUse.input is not a valid JSON object
- enhancedMessages = enhancedMessages
- .map((msg: any) => {
- if (msg.role !== "assistant" || !Array.isArray(msg.content)) {
- return msg
- }
- const filteredContent = msg.content.filter((part: any) => {
- if (part.type === "tool-call") {
- // Check if input is a valid object (not null, undefined, or empty)
- if (
- !part.input ||
- typeof part.input !== "object" ||
- Object.keys(part.input).length === 0
- ) {
- console.warn(
- `[route.ts] Filtering out tool-call with invalid input:`,
- { toolName: part.toolName, input: part.input },
- )
- return false
- }
- }
- return true
- })
- return { ...msg, content: filteredContent }
- })
- .filter((msg: any) => msg.content && msg.content.length > 0)
+ // and their results. Bedrock API rejects messages where toolUse.input is not a valid
+ // JSON object, and every provider rejects a tool result whose call is gone.
+ enhancedMessages = dropInvalidToolCalls(enhancedMessages)
- // DEBUG: Log modelMessages structure (what's being sent to AI)
- console.log("[route.ts] Model messages count:", enhancedMessages.length)
- enhancedMessages.forEach((msg: any, idx: number) => {
- console.log(
- `[route.ts] ModelMsg ${idx} role:`,
- msg.role,
- "content count:",
- msg.content?.length,
- )
- if (msg.content) {
- msg.content.forEach((part: any, partIdx: number) => {
- if (part.type === "tool-call" || part.type === "tool-result") {
- console.log(`[route.ts] Content ${partIdx}:`, {
- type: part.type,
- toolName: part.toolName,
- hasInput: !!part.input,
- inputType: typeof part.input,
- inputValue:
- part.input === undefined
- ? "undefined"
- : part.input === null
- ? "null"
- : "object",
- })
- }
- })
- }
- })
+ // DEBUG_LLM_PAYLOAD=true logs what is sent to the model
+ if (DEBUG_LLM_PAYLOAD) {
+ console.log("[route.ts] Model messages count:", enhancedMessages.length)
+ enhancedMessages.forEach((msg: any, idx: number) => {
+ console.log(
+ `[route.ts] ModelMsg ${idx} role:`,
+ msg.role,
+ "content count:",
+ msg.content?.length,
+ )
+ if (msg.content) {
+ msg.content.forEach((part: any, partIdx: number) => {
+ if (
+ part.type === "tool-call" ||
+ part.type === "tool-result"
+ ) {
+ console.log(`[route.ts] Content ${partIdx}:`, {
+ type: part.type,
+ toolName: part.toolName,
+ hasInput: !!part.input,
+ inputType: typeof part.input,
+ inputValue:
+ part.input === undefined
+ ? "undefined"
+ : part.input === null
+ ? "null"
+ : "object",
+ })
+ }
+ })
+ }
+ })
+ }
// Update the last message with user input only (XML moved to separate cached system message)
if (enhancedMessages.length >= 1) {
@@ -353,7 +468,7 @@ ${userInputText}
contentParts.push({
type: "image",
image: filePart.url,
- mimeType: filePart.mediaType,
+ mediaType: filePart.mediaType,
})
}
@@ -373,9 +488,7 @@ ${userInputText}
if (enhancedMessages[i].role === "assistant") {
enhancedMessages[i] = {
...enhancedMessages[i],
- providerOptions: {
- bedrock: { cachePoint: { type: "default" } },
- },
+ providerOptions: CACHE_POINT,
}
break // Only cache the last assistant message
}
@@ -383,40 +496,75 @@ ${userInputText}
}
// System messages with multiple cache breakpoints for optimal caching:
- // - Breakpoint 1: Static instructions (~1500 tokens) - rarely changes
+ // - Breakpoint 1: System instructions + custom instructions - changes when user updates custom system message
// - Breakpoint 2: Current XML context - changes per diagram, but constant within a conversation turn
- // This allows: if only user message changes, both system caches are reused
- // if XML changes, instruction cache is still reused
- const systemMessages = [
- // Cache breakpoint 1: Instructions (rarely change)
- {
- role: "system" as const,
- content: systemMessage,
- ...(shouldCache && {
- providerOptions: {
- bedrock: { cachePoint: { type: "default" } },
- },
- }),
- },
- // Cache breakpoint 2: Previous and Current diagram XML context
- {
- role: "system" as const,
- content: `${previousXml ? `Previous diagram XML (before user's last message):\n"""xml\n${previousXml}\n"""\n\n` : ""}Current diagram XML (AUTHORITATIVE - the source of truth):\n"""xml\n${xml || ""}\n"""\n\nIMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on the canvas right now. The user can manually add, delete, or modify shapes directly in draw.io. Always count and describe elements based on the CURRENT XML, not on what you previously generated. If both previous and current XML are shown, compare them to understand what the user changed. When using edit_diagram, COPY search patterns exactly from the CURRENT XML - attribute order matters!`,
- ...(shouldCache && {
- providerOptions: {
- bedrock: { cachePoint: { type: "default" } },
- },
- }),
- },
- ]
+ // Some providers (e.g. MiniMax) don't support multiple system messages
+ // Merge them into a single system message for compatibility
+ // Also merge for OpenAI-compatible providers with custom base URLs (e.g. vLLM, LMStudio)
+ // because open-source model chat templates (Qwen, Llama, etc.) typically reject multiple system messages
+ const isCustomOpenAIEndpoint =
+ resolvedProvider === "openai" &&
+ !!(
+ baseUrl ||
+ process.env.OPENAI_BASE_URL ||
+ (serverModelConfig.baseUrlEnv &&
+ process.env[serverModelConfig.baseUrlEnv])
+ )
+ const isSingleSystemProvider =
+ SINGLE_SYSTEM_PROVIDERS.has(resolvedProvider) || isCustomOpenAIEndpoint
+
+ const xmlContext = `${
+ previousXml
+ ? `Previous diagram XML (before user's last message):
+"""xml
+${previousXml}
+"""
+
+`
+ : ""
+ }Current diagram XML (AUTHORITATIVE - the source of truth):
+"""xml
+${xml || ""}
+"""
+
+IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on the canvas right now. The user can manually add, delete, or modify shapes directly in draw.io. Always count and describe elements based on the CURRENT XML, not on what you previously generated. If both previous and current XML are shown, compare them to understand what the user changed.`
+
+ const systemMessages = isSingleSystemProvider
+ ? [
+ {
+ role: "system" as const,
+ content: `${finalSystemMessage}\n\n${xmlContext}`,
+ },
+ ]
+ : [
+ // Cache breakpoint 1: Instructions (+ optional custom instructions)
+ {
+ role: "system" as const,
+ content: finalSystemMessage,
+ ...(shouldCache && { providerOptions: CACHE_POINT }),
+ },
+ // Cache breakpoint 2: Previous and Current diagram XML context
+ {
+ role: "system" as const,
+ content: xmlContext,
+ ...(shouldCache && { providerOptions: CACHE_POINT }),
+ },
+ ]
const allMessages = [...systemMessages, ...enhancedMessages]
+ // Set by onAbort, which records the finished steps' tokens itself
+ let stopped = false
const result = streamText({
model,
- ...(process.env.MAX_OUTPUT_TOKENS && {
- maxOutputTokens: parseInt(process.env.MAX_OUTPUT_TOKENS, 10),
- }),
+ // The system messages carry cache points, so they go in messages.
+ // A client's own system messages have string content and were
+ // dropped by the empty-content filter above.
+ allowSystemInMessages: true,
+ abortSignal: req.signal,
+ // Must be sent: unset means the provider's own default, and Bedrock's is
+ // 4096, enough for a small diagram, so larger ones were cut off mid-attribute.
+ maxOutputTokens,
stopWhen: stepCountIs(5),
// Repair truncated tool calls when maxOutputTokens is reached mid-JSON
experimental_repairToolCall: async ({ toolCall, error }) => {
@@ -434,16 +582,11 @@ ${userInputText}
error.name === "AI_InvalidToolInputError"
) {
try {
- // Pre-process to fix common LLM JSON errors that jsonrepair can't handle
- let inputToRepair = toolCall.input
- if (typeof inputToRepair === "string") {
- // Fix `:=` instead of `: ` (LLM sometimes generates this)
- inputToRepair = inputToRepair.replace(/:=/g, ": ")
- // Fix `= "` instead of `: "`
- inputToRepair = inputToRepair.replace(/=\s*"/g, ': "')
- }
- // Use jsonrepair to fix truncated JSON
- const repairedInput = jsonrepair(inputToRepair)
+ // Pre-process to fix common LLM JSON errors that jsonrepair can't handle,
+ // then use jsonrepair to fix truncated JSON
+ const repairedInput = jsonrepair(
+ fixToolInputJson(toolCall.input),
+ )
console.log(
`[repairToolCall] Repaired truncated JSON for tool: ${toolCall.toolName}`,
)
@@ -453,26 +596,8 @@ ${userInputText}
`[repairToolCall] Failed to repair JSON for tool: ${toolCall.toolName}`,
repairError,
)
- // Return a placeholder input to avoid API errors in multi-step
- // The tool will fail gracefully on client side
- if (toolCall.toolName === "edit_diagram") {
- return {
- ...toolCall,
- input: {
- operations: [],
- _error: "JSON repair failed - no operations to apply",
- },
- }
- }
- if (toolCall.toolName === "display_diagram") {
- return {
- ...toolCall,
- input: {
- xml: "",
- _error: "JSON repair failed - empty diagram",
- },
- }
- }
+ // Keep the original error, so the model and the client see why
+ // the input was rejected and the model can retry the call
return null
}
}
@@ -481,7 +606,6 @@ ${userInputText}
},
messages: allMessages,
...(providerOptions && { providerOptions }), // This now includes all reasoning configs
- ...(headers && { headers }),
// Langfuse telemetry config (returns undefined if not configured)
...(getTelemetryConfig({ sessionId: validSessionId, userId }) && {
experimental_telemetry: getTelemetryConfig({
@@ -495,21 +619,36 @@ ${userInputText}
// Record token usage for server-side quota tracking (if enabled)
// Use totalUsage (cumulative across all steps) instead of usage (final step only)
- // Include all 4 token types: input, output, cache read, cache write
- if (
- isQuotaEnabled() &&
- !hasOwnApiKey &&
- userId !== "anonymous" &&
- totalUsage
- ) {
+ // inputTokens already includes cache reads and writes in AI SDK 6
+ if (countsQuota && totalUsage && !stopped) {
const totalTokens =
(totalUsage.inputTokens || 0) +
- (totalUsage.outputTokens || 0) +
- (totalUsage.cachedInputTokens || 0) +
- (totalUsage.inputTokenDetails?.cacheWriteTokens || 0)
+ (totalUsage.outputTokens || 0)
recordTokenUsage(userId, totalTokens)
}
},
+ // onFinish is skipped when the stream fails or is aborted, so end the trace here
+ onError: ({ error }) => {
+ console.error(error) // what AI SDK does without an onError
+ endTrace()
+ },
+ onAbort: ({ steps }) => {
+ stopped = true
+ endTrace()
+ // Stopped (or disconnected) after some steps finished: their
+ // tokens were used, or stopping every request after a costly
+ // first step would get around the token limits
+ if (countsQuota) {
+ const tokens = steps.reduce(
+ (sum, step) =>
+ sum +
+ (step.usage.inputTokens || 0) +
+ (step.usage.outputTokens || 0),
+ 0,
+ )
+ if (tokens > 0) recordTokenUsage(userId, tokens)
+ }
+ },
tools: {
// Client-side tool that will be executed on the client
display_diagram: {
@@ -524,21 +663,7 @@ VALIDATION RULES (XML will be rejected if violated):
6. Escape special chars in values: < > & "
Example (generate ONLY this - no wrapper tags):
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
+${SWIMLANE_EXAMPLE}
Notes:
- For AWS diagrams, use **AWS 2025 icons**.
@@ -616,14 +741,7 @@ Example: If previous output ended with ' streamErrorText(error, onServerCredentials),
messageMetadata: ({ part }) => {
if (part.type === "finish") {
const usage = (part as any).totalUsage
@@ -695,63 +784,28 @@ Call this tool to get shape names and usage syntax for a specific library.`,
return undefined
},
})
+ modelStreamResponses.add(response)
+ return response
}
-// Helper to categorize errors and return appropriate response
+// Errors before the stream starts, as JSON the chat panel reads
function handleError(error: unknown): Response {
console.error("Error in chat route:", error)
const isDev = process.env.NODE_ENV === "development"
-
- // Check for specific AI SDK error types
- if (APICallError.isInstance(error)) {
- return Response.json(
- {
- error: error.message,
- ...(isDev && {
- details: error.responseBody,
- stack: error.stack,
- }),
- },
- { status: error.statusCode || 500 },
- )
- }
-
- if (LoadAPIKeyError.isInstance(error)) {
- return Response.json(
- {
- error: "Authentication failed. Please check your API key.",
- ...(isDev && {
- stack: error.stack,
- }),
- },
- { status: 401 },
- )
- }
-
- // Fallback for other errors with safety filter
- const message =
- error instanceof Error ? error.message : "An unexpected error occurred"
- const status = (error as any)?.statusCode || (error as any)?.status || 500
-
- // Prevent leaking API keys, tokens, or other sensitive data
- const lowerMessage = message.toLowerCase()
- const safeMessage =
- lowerMessage.includes("key") ||
- lowerMessage.includes("token") ||
- lowerMessage.includes("sig") ||
- lowerMessage.includes("signature") ||
- lowerMessage.includes("secret") ||
- lowerMessage.includes("password") ||
- lowerMessage.includes("credential")
- ? "Authentication failed. Please check your credentials."
- : message
+ const classified = classifyLLMError(error)
+ const status =
+ (error as { statusCode?: number })?.statusCode ||
+ (error as { status?: number })?.status ||
+ (classified.code === "invalid_api_key" ? 401 : 500)
return Response.json(
{
- error: safeMessage,
+ ...classified,
...(isDev && {
- details: message,
+ details: APICallError.isInstance(error)
+ ? error.responseBody
+ : undefined,
stack: error instanceof Error ? error.stack : undefined,
}),
},
@@ -761,11 +815,16 @@ function handleError(error: unknown): Response {
// Wrap handler with error handling
async function safeHandler(req: Request): Promise {
+ let response: Response
try {
- return await handleChatRequest(req)
+ response = await handleChatRequest(req)
} catch (error) {
- return handleError(error)
+ response = handleError(error)
}
+ // Early returns, cache hits and errors never reach streamText's callbacks,
+ // so their Langfuse trace has to be ended here
+ if (!modelStreamResponses.has(response)) endTrace()
+ return response
}
// Wrap with Langfuse observe (if configured)
diff --git a/app/api/log-save/route.ts b/app/api/log-save/route.ts
index fc73fb2b..eb30e0fe 100644
--- a/app/api/log-save/route.ts
+++ b/app/api/log-save/route.ts
@@ -4,7 +4,7 @@ import { getLangfuseClient } from "@/lib/langfuse"
const saveSchema = z.object({
filename: z.string().min(1).max(255),
- format: z.enum(["drawio", "png", "svg"]),
+ format: z.enum(["drawio", "png", "svg", "xmlsvg"]),
sessionId: z.string().min(1).max(200).optional(),
})
diff --git a/app/api/parse-url/route.ts b/app/api/parse-url/route.ts
index f5278e65..33a15c4b 100644
--- a/app/api/parse-url/route.ts
+++ b/app/api/parse-url/route.ts
@@ -1,62 +1,46 @@
-import { extract } from "@extractus/article-extractor"
+import { extractFromHtml } from "@extractus/article-extractor"
import { NextResponse } from "next/server"
import TurndownService from "turndown"
+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
+const MAX_RESPONSE_BYTES = 5 * 1024 * 1024
const EXTRACT_TIMEOUT_MS = 15000
+const USER_AGENT = "Mozilla/5.0 (compatible; NextAIDrawio/1.0)"
-// SSRF protection - block private/internal addresses
-function isPrivateUrl(urlString: string): boolean {
+// Detect the page's charset so non-UTF-8 pages (Shift_JIS/GBK/EUC/Big5, common
+// on CJK sites) are decoded correctly. Response.text() always assumes UTF-8 and
+// would produce mojibake; the article-extractor library does the same detection
+// when it fetches the page itself, which we no longer rely on.
+function detectCharset(
+ contentType: string | null,
+ buffer: ArrayBuffer,
+): string {
+ // 1. HTTP Content-Type header charset (most authoritative).
+ const headerCharset = contentType?.match(/charset=([^;]+)/i)?.[1]?.trim()
+ // 2. / in the first bytes of the document.
+ const head = new TextDecoder("utf-8").decode(buffer.slice(0, 4096))
+ const metaCharset =
+ head.match(/]+charset=["']?\s*([\w-]+)/i)?.[1] ||
+ head.match(/]+content=["'][^"']*charset=([\w-]+)/i)?.[1]
+ const charset = (headerCharset || metaCharset || "utf-8").toLowerCase()
+ // TextDecoder throws on unknown encoding labels; fall back to UTF-8.
try {
- const url = new URL(urlString)
- const hostname = url.hostname.toLowerCase()
-
- // Block localhost
- if (
- hostname === "localhost" ||
- hostname === "127.0.0.1" ||
- hostname === "::1"
- ) {
- return true
- }
-
- // Block AWS/cloud metadata endpoints
- if (
- hostname === "169.254.169.254" ||
- hostname === "metadata.google.internal"
- ) {
- return true
- }
-
- // Check for private IPv4 ranges
- const ipv4Match = hostname.match(
- /^(\d{1,3})\.(\d{1,3})\.(\d{1,3})\.(\d{1,3})$/,
- )
- if (ipv4Match) {
- const [, a, b] = ipv4Match.map(Number)
- if (a === 10) return true // 10.0.0.0/8
- if (a === 172 && b >= 16 && b <= 31) return true // 172.16.0.0/12
- if (a === 192 && b === 168) return true // 192.168.0.0/16
- if (a === 169 && b === 254) return true // 169.254.0.0/16 (link-local)
- if (a === 127) return true // 127.0.0.0/8 (loopback)
- }
-
- // Block common internal hostnames
- if (
- hostname.endsWith(".local") ||
- hostname.endsWith(".internal") ||
- hostname.endsWith(".localhost")
- ) {
- return true
- }
-
- return false
+ new TextDecoder(charset)
+ return charset
} catch {
- return true // Invalid URL - block it
+ return "utf-8"
}
}
export async function POST(req: Request) {
+ const crossSite = rejectCrossSite(req)
+ if (crossSite) return crossSite
+ const accessError = checkAccessCode(req)
+ if (accessError) return accessError
+
try {
const { url } = await req.json()
@@ -77,28 +61,61 @@ export async function POST(req: Request) {
)
}
- // SSRF protection
- if (isPrivateUrl(url)) {
+ // SSRF protection: parse-url has no use case for fetching internal
+ // hosts, so private URLs are always rejected. ALLOW_PRIVATE_URLS only
+ // governs LLM provider baseUrl overrides (validate-model, chat).
+ if (await isPrivateUrl(url)) {
return NextResponse.json(
{ error: "Cannot access private/internal URLs" },
{ status: 400 },
)
}
-
- // Extract article content with timeout to avoid tying up server resources
+ // Fetch the page ourselves so we control redirect handling. The
+ // article-extractor library follows redirects internally and ignores a
+ // `redirect` option, which would let a public URL 302 to an internal
+ // host and bypass the SSRF check above. `redirect: "error"` rejects any
+ // redirect outright.
const controller = new AbortController()
const timeoutId = setTimeout(() => {
controller.abort()
}, EXTRACT_TIMEOUT_MS)
- let article
+ let html: string
try {
- article = await extract(url, undefined, {
- headers: {
- "User-Agent": "Mozilla/5.0 (compatible; NextAIDrawio/1.0)",
- },
+ const response = await fetch(url, {
+ headers: { "User-Agent": USER_AGENT },
+ redirect: "error",
signal: controller.signal,
})
+
+ const contentType = response.headers.get("content-type")
+ if (contentType?.includes("application/pdf")) {
+ return NextResponse.json(
+ {
+ error: "PDF URLs are not supported. Please download and upload the PDF file directly",
+ },
+ { status: 422 },
+ )
+ }
+
+ if (!response.ok) {
+ return NextResponse.json(
+ { error: "Could not fetch URL content" },
+ { status: 400 },
+ )
+ }
+
+ const buffer = await readLimitedBody(response, MAX_RESPONSE_BYTES)
+ 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) {
if (err?.name === "AbortError") {
return NextResponse.json(
@@ -106,9 +123,26 @@ export async function POST(req: Request) {
{ status: 504 },
)
}
- throw err
+ // Redirects are rejected with a TypeError ("failed to fetch" /
+ // "unexpected redirect") when redirect: "error" is set.
+ return NextResponse.json(
+ { error: "Could not fetch URL content" },
+ { status: 400 },
+ )
} finally {
clearTimeout(timeoutId)
+ // Ends a download left unread (too large, PDF, error status);
+ // a body already read is not affected
+ controller.abort()
+ }
+
+ // extractFromHtml throws (not returns null) on empty/non-HTML bodies,
+ // so map any parse error to the same 400 as the no-content case.
+ let article: Awaited>
+ try {
+ article = await extractFromHtml(html, url)
+ } catch {
+ article = null
}
if (!article || !article.content) {
diff --git a/app/api/provider-models/route.ts b/app/api/provider-models/route.ts
new file mode 100644
index 00000000..6b362e73
--- /dev/null
+++ b/app/api/provider-models/route.ts
@@ -0,0 +1,79 @@
+import { NextResponse } from "next/server"
+import { checkAccessCode, rejectCrossSite } from "@/lib/access-code"
+import { classifyLLMError } from "@/lib/llm-errors"
+import {
+ canListModels,
+ listProviderModels,
+ ModelListError,
+} from "@/lib/provider-models"
+import {
+ allowPrivateUrls,
+ isPrivateUrl,
+ RedirectRefusedError,
+ redirectGuardedFetch,
+} from "@/lib/ssrf-protection"
+import type { ProviderName } from "@/lib/types/model-config"
+
+export const runtime = "nodejs"
+
+// Public lists need no key
+const NO_KEY_NEEDED = new Set([
+ "ollama",
+ "openrouter",
+ "aihubmix",
+])
+
+/**
+ * The models a provider offers, for the "Fetch models" button in model
+ * settings. Answers { models: null } for providers that cannot list them,
+ * 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
+
+ const { provider, apiKey, baseUrl } = (await req.json()) as {
+ provider: ProviderName
+ apiKey?: string
+ baseUrl?: string
+ }
+ if (!canListModels(provider)) {
+ return NextResponse.json({ models: null })
+ }
+ // SECURITY: Block SSRF attacks via custom baseUrl
+ if (baseUrl && !allowPrivateUrls() && (await isPrivateUrl(baseUrl))) {
+ return NextResponse.json({ error: "Invalid base URL" }, { status: 400 })
+ }
+ if (!apiKey && !NO_KEY_NEEDED.has(provider)) {
+ return NextResponse.json(
+ { error: "API key is required" },
+ { status: 400 },
+ )
+ }
+
+ try {
+ const models = await listProviderModels(
+ provider,
+ { apiKey, baseUrl },
+ (baseUrl && redirectGuardedFetch()) || fetch,
+ )
+ return NextResponse.json({ models })
+ } catch (error) {
+ console.warn("[provider-models] Listing failed:", error)
+ // 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 isOwn = (e: unknown): e is Error =>
+ e instanceof ModelListError || e instanceof RedirectRefusedError
+ const cause = (error as { cause?: unknown })?.cause
+ const own = isOwn(error) ? error : isOwn(cause) ? cause : null
+ const { code } = classifyLLMError(own ?? error)
+ return NextResponse.json({
+ code,
+ error: own?.message ?? "The model list request failed.",
+ })
+ }
+}
diff --git a/app/api/server-models/route.ts b/app/api/server-models/route.ts
new file mode 100644
index 00000000..49ea12bb
--- /dev/null
+++ b/app/api/server-models/route.ts
@@ -0,0 +1,14 @@
+import { NextResponse } from "next/server"
+import { loadFlattenedServerModels } from "@/lib/server-model-config"
+
+// Use dynamic rendering to read AI_MODEL/AI_PROVIDER env vars at runtime
+// This ensures Docker users can set these values when starting containers
+export const dynamic = "force-dynamic"
+
+export async function GET() {
+ const models = await loadFlattenedServerModels()
+ return NextResponse.json({
+ models,
+ hasConfig: models.length > 0,
+ })
+}
diff --git a/app/api/validate-diagram/route.ts b/app/api/validate-diagram/route.ts
new file mode 100644
index 00000000..3fba6628
--- /dev/null
+++ b/app/api/validate-diagram/route.ts
@@ -0,0 +1,184 @@
+/**
+ * API endpoint for VLM-based diagram validation.
+ * Accepts a PNG image and streams validation results using useObject-compatible format.
+ */
+
+import { Output, streamText } from "ai"
+import { checkAccessCode, rejectCrossSite } from "@/lib/access-code"
+import { getValidationModel } from "@/lib/ai-providers"
+import {
+ checkAndIncrementRequest,
+ isQuotaEnabled,
+ recordTokenUsage,
+} from "@/lib/dynamo-quota-manager"
+import { getUserIdFromRequest } from "@/lib/user-id"
+import { VALIDATION_SYSTEM_PROMPT } from "@/lib/validation-prompts"
+import {
+ type ValidationResult,
+ ValidationResultSchema,
+} from "@/lib/validation-schema"
+
+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
+}
+
+// Default valid result for disabled/error cases
+const DEFAULT_VALID_RESULT: ValidationResult = {
+ valid: true,
+ issues: [],
+ suggestions: [],
+}
+
+/** A fixed result in the text format useObject reads */
+function createStreamingResponse(result: ValidationResult): Response {
+ return new Response(JSON.stringify(result), {
+ headers: { "Content-Type": "text/plain; charset=utf-8" },
+ })
+}
+
+export async function POST(req: Request): Promise {
+ const crossSite = rejectCrossSite(req)
+ if (crossSite) return crossSite
+ // Uses the server's model credentials, so require the access code
+ const accessError = checkAccessCode(req)
+ if (accessError) return accessError
+
+ try {
+ // Check if VLM validation is enabled (default: true)
+ const enableValidation = process.env.ENABLE_VLM_VALIDATION !== "false"
+ if (!enableValidation) {
+ return createStreamingResponse(DEFAULT_VALID_RESULT)
+ }
+
+ const body: ValidateDiagramRequest = await req.json()
+ const { imageData, sessionId } = body
+
+ if (!imageData) {
+ return Response.json(
+ { error: "Missing imageData" },
+ { status: 400 },
+ )
+ }
+
+ // Validate image data format
+ if (
+ !imageData.startsWith("data:image/png;base64,") &&
+ !imageData.startsWith("data:image/")
+ ) {
+ return Response.json(
+ { error: "Invalid image data format" },
+ { status: 400 },
+ )
+ }
+
+ if (imageData.length > MAX_IMAGE_DATA_LENGTH) {
+ return Response.json(
+ { error: "Image data too large" },
+ { status: 413 },
+ )
+ }
+
+ // It runs the server's vision model: with the quota on, the daily
+ // and per-minute token limits apply, and its tokens are counted. Not
+ // the request limit, which is for chats: the day's last chat still
+ // gets its check, and a check does not count as a chat.
+ const userId = getUserIdFromRequest(req)
+ const countsQuota = isQuotaEnabled() && userId !== "anonymous"
+ if (countsQuota) {
+ const quotaCheck = await checkAndIncrementRequest(
+ userId,
+ {
+ requests: 0,
+ tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000,
+ tpm: Number(process.env.TPM_LIMIT) || 20000,
+ },
+ 0,
+ )
+ if (!quotaCheck.allowed) {
+ return Response.json(
+ {
+ error: quotaCheck.error,
+ type: quotaCheck.type,
+ used: quotaCheck.used,
+ limit: quotaCheck.limit,
+ },
+ { status: 429 },
+ )
+ }
+ }
+
+ // Get the validation model
+ let model
+ try {
+ model = getValidationModel()
+ } catch (error) {
+ console.warn(
+ "[validate-diagram] Validation model not available:",
+ error,
+ )
+ // Return valid if no vision model is configured
+ return createStreamingResponse(DEFAULT_VALID_RESULT)
+ }
+
+ // Parse timeout with validation (minimum 1000ms, default 10000ms)
+ const timeout =
+ Math.max(
+ 1000,
+ parseInt(process.env.VALIDATION_TIMEOUT || "10000", 10),
+ ) || 10000
+
+ // Stream the VLM response for useObject consumption
+ const result = streamText({
+ model,
+ output: Output.object({ schema: ValidationResultSchema }),
+ system: VALIDATION_SYSTEM_PROMPT,
+ messages: [
+ {
+ role: "user",
+ content: [
+ {
+ type: "image",
+ image: imageData,
+ },
+ {
+ type: "text",
+ text: "Please analyze this diagram for visual quality issues.",
+ },
+ ],
+ },
+ ],
+ maxOutputTokens: 1024,
+ abortSignal: AbortSignal.timeout(timeout),
+ onFinish: ({ output, totalUsage }) => {
+ if (countsQuota && totalUsage) {
+ recordTokenUsage(
+ userId,
+ (totalUsage.inputTokens || 0) +
+ (totalUsage.outputTokens || 0),
+ )
+ }
+ if (sessionId && output) {
+ console.log(
+ `[validate-diagram] Session ${sessionId}: valid=${output.valid}, issues=${output.issues?.length ?? 0}`,
+ )
+ }
+ },
+ })
+
+ return result.toTextStreamResponse()
+ } catch (error) {
+ // Log with session context if available
+ const errorMessage =
+ error instanceof Error ? error.message : String(error)
+ console.error("[validate-diagram] Error:", errorMessage)
+
+ // On error, return valid to not block the user
+ return createStreamingResponse(DEFAULT_VALID_RESULT)
+ }
+}
diff --git a/app/api/validate-model/route.ts b/app/api/validate-model/route.ts
index b8b258e6..1109a976 100644
--- a/app/api/validate-model/route.ts
+++ b/app/api/validate-model/route.ts
@@ -1,78 +1,28 @@
-import { createAmazonBedrock } from "@ai-sdk/amazon-bedrock"
-import { createAnthropic } from "@ai-sdk/anthropic"
-import { createDeepSeek, deepseek } from "@ai-sdk/deepseek"
-import { createGateway } from "@ai-sdk/gateway"
-import { createGoogleGenerativeAI } from "@ai-sdk/google"
-import { createOpenAI } from "@ai-sdk/openai"
-import { createOpenRouter } from "@openrouter/ai-sdk-provider"
-import { generateText } from "ai"
+import { streamText, tool } from "ai"
import { NextResponse } from "next/server"
-import { createOllama } from "ollama-ai-provider-v2"
+import { z } from "zod"
+import { checkAccessCode, rejectCrossSite } from "@/lib/access-code"
+import { checkAdminAuth } from "@/lib/admin/auth"
+import {
+ edgeOneEndpoint,
+ getAIModel,
+ globalBaseUrl,
+ usesServerCredentials,
+ usesServerEndpoint,
+} from "@/lib/ai-providers"
+import {
+ checkAndIncrementRequest,
+ isQuotaEnabled,
+} from "@/lib/dynamo-quota-manager"
+import { classifyLLMError } from "@/lib/llm-errors"
+import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection"
+import { normalizeBaseUrl, type ProviderName } from "@/lib/types/model-config"
+import { getUserIdFromRequest } from "@/lib/user-id"
export const runtime = "nodejs"
-/**
- * SECURITY: Check if URL points to private/internal network (SSRF protection)
- * Blocks: localhost, private IPs, link-local, AWS metadata service
- */
-function isPrivateUrl(urlString: string): boolean {
- try {
- const url = new URL(urlString)
- const hostname = url.hostname.toLowerCase()
-
- // Block localhost
- if (
- hostname === "localhost" ||
- hostname === "127.0.0.1" ||
- hostname === "::1"
- ) {
- return true
- }
-
- // Block AWS/cloud metadata endpoints
- if (
- hostname === "169.254.169.254" ||
- hostname === "metadata.google.internal"
- ) {
- return true
- }
-
- // Check for private IPv4 ranges
- const ipv4Match = hostname.match(
- /^(\d{1,3})\.(\d{1,3})\.(\d{1,3})\.(\d{1,3})$/,
- )
- if (ipv4Match) {
- const [, a, b] = ipv4Match.map(Number)
- // 10.0.0.0/8
- if (a === 10) return true
- // 172.16.0.0/12
- if (a === 172 && b >= 16 && b <= 31) return true
- // 192.168.0.0/16
- if (a === 192 && b === 168) return true
- // 169.254.0.0/16 (link-local)
- if (a === 169 && b === 254) return true
- // 127.0.0.0/8 (loopback)
- if (a === 127) return true
- }
-
- // Block common internal hostnames
- if (
- hostname.endsWith(".local") ||
- hostname.endsWith(".internal") ||
- hostname.endsWith(".localhost")
- ) {
- return true
- }
-
- return false
- } catch {
- // Invalid URL - block it
- return true
- }
-}
-
interface ValidateRequest {
- provider: string
+ provider: ProviderName
apiKey: string
baseUrl?: string
modelId: string
@@ -80,19 +30,44 @@ interface ValidateRequest {
awsAccessKeyId?: string
awsSecretAccessKey?: string
awsRegion?: string
+ awsSessionToken?: string
+ // Vertex AI specific
+ vertexApiKey?: string // Express Mode API key
+ // Set by the admin panel's Test: baseUrl is the server's _BASE_URL
+ serverBaseUrl?: boolean
}
+const TEST_TIMEOUT_MS = 15_000
+
+// Drawing works through tool calls, so the test asks for one
+const PING_TOOL = tool({
+ description: "Report that the connection works.",
+ inputSchema: z.object({}),
+})
+
+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)
+ if (accessError && checkAdminAuth(req)) return accessError
+
try {
const body: ValidateRequest = await req.json()
const {
provider,
apiKey,
- baseUrl,
modelId,
awsAccessKeyId,
awsSecretAccessKey,
awsRegion,
+ awsSessionToken,
+ // Note: Express Mode only needs vertexApiKey
+ vertexApiKey,
} = body
if (!provider || !modelId) {
@@ -101,9 +76,26 @@ export async function POST(req: Request) {
{ status: 400 },
)
}
+ // EdgeOne is this site's own function, as in the chat; the admin
+ // panel's Test sends no URL, and a relative one cannot be fetched
+ const baseUrl =
+ provider === "edgeone" ? edgeOneEndpoint(req) : body.baseUrl
+ // The admin panel's Test of an entry without a URL sends the
+ // server's own
_BASE_URL, which chat uses as it is: not a URL a
+ // user chose, so no private-address or redirect rules
+ const serverUrl =
+ body.serverBaseUrl === true &&
+ !!baseUrl &&
+ baseUrl === globalBaseUrl(provider) &&
+ !checkAdminAuth(req)
// SECURITY: Block SSRF attacks via custom baseUrl
- if (baseUrl && isPrivateUrl(baseUrl)) {
+ if (
+ baseUrl &&
+ !serverUrl &&
+ !allowPrivateUrls() &&
+ (await isPrivateUrl(baseUrl))
+ ) {
return NextResponse.json(
{ valid: false, error: "Invalid base URL" },
{ status: 400 },
@@ -121,278 +113,133 @@ export async function POST(req: Request) {
{ status: 400 },
)
}
+ } else if (provider === "vertexai") {
+ if (!vertexApiKey) {
+ return NextResponse.json(
+ {
+ valid: false,
+ error: "Vertex AI API key is required for Express Mode",
+ },
+ { status: 400 },
+ )
+ }
} else if (provider !== "ollama" && provider !== "edgeone" && !apiKey) {
return NextResponse.json(
{ valid: false, error: "API key is required" },
{ status: 400 },
)
}
-
- let model: any
-
- switch (provider) {
- case "openai": {
- const openai = createOpenAI({
- apiKey,
- ...(baseUrl && { baseURL: baseUrl }),
- })
- model = openai.chat(modelId)
- break
- }
-
- case "anthropic": {
- const anthropic = createAnthropic({
- apiKey,
- baseURL: baseUrl || "https://api.anthropic.com/v1",
- })
- model = anthropic(modelId)
- break
- }
-
- case "google": {
- const google = createGoogleGenerativeAI({
- apiKey,
- ...(baseUrl && { baseURL: baseUrl }),
- })
- model = google(modelId)
- break
- }
-
- case "azure": {
- const azure = createOpenAI({
- apiKey,
- baseURL: baseUrl,
- })
- model = azure.chat(modelId)
- break
- }
-
- case "bedrock": {
- const bedrock = createAmazonBedrock({
- accessKeyId: awsAccessKeyId,
- secretAccessKey: awsSecretAccessKey,
- region: awsRegion,
- })
- model = bedrock(modelId)
- break
- }
-
- case "openrouter": {
- const openrouter = createOpenRouter({
- apiKey,
- ...(baseUrl && { baseURL: baseUrl }),
- })
- model = openrouter(modelId)
- break
- }
-
- case "deepseek": {
- if (baseUrl || apiKey) {
- const ds = createDeepSeek({
- apiKey,
- ...(baseUrl && { baseURL: baseUrl }),
- })
- model = ds(modelId)
- } else {
- model = deepseek(modelId)
- }
- break
- }
-
- case "siliconflow": {
- const sf = createOpenAI({
- apiKey,
- baseURL: baseUrl || "https://api.siliconflow.cn/v1",
- })
- model = sf.chat(modelId)
- break
- }
-
- case "ollama": {
- const ollama = createOllama({
- baseURL: baseUrl || "http://localhost:11434",
- })
- model = ollama(modelId)
- break
- }
-
- case "gateway": {
- const gw = createGateway({
- apiKey,
- ...(baseUrl && { baseURL: baseUrl }),
- })
- model = gw(modelId)
- break
- }
-
- case "edgeone": {
- // EdgeOne uses OpenAI-compatible API via Edge Functions
- // Need to pass cookies for EdgeOne Pages authentication
- const cookieHeader = req.headers.get("cookie") || ""
- const edgeone = createOpenAI({
- apiKey: "edgeone", // EdgeOne doesn't require API key
- baseURL: baseUrl || "/api/edgeai",
- headers: {
- cookie: cookieHeader,
- },
- })
- model = edgeone.chat(modelId)
- break
- }
-
- case "sglang": {
- // SGLang is OpenAI-compatible
- const sglang = createOpenAI({
- apiKey: apiKey || "not-needed",
- baseURL: baseUrl || "http://127.0.0.1:8000/v1",
- })
- model = sglang.chat(modelId)
- break
- }
-
- case "doubao": {
- // ByteDance Doubao: use DeepSeek for DeepSeek/Kimi models, OpenAI for others
- const doubaoBaseUrl =
- baseUrl || "https://ark.cn-beijing.volces.com/api/v3"
- const lowerModelId = modelId.toLowerCase()
- if (
- lowerModelId.includes("deepseek") ||
- lowerModelId.includes("kimi")
- ) {
- const doubao = createDeepSeek({
- apiKey,
- baseURL: doubaoBaseUrl,
- })
- model = doubao(modelId)
- } else {
- const doubao = createOpenAI({
- apiKey,
- baseURL: doubaoBaseUrl,
- })
- model = doubao.chat(modelId)
- }
- break
- }
-
- case "modelscope": {
- const baseURL =
- baseUrl || "https://api-inference.modelscope.cn/v1"
- const startTime = Date.now()
-
- try {
- // Initiate a streaming request (required for QwQ-32B and certain Qwen3 models)
- const response = await fetch(
- `${baseURL}/chat/completions`,
- {
- method: "POST",
- headers: {
- "Content-Type": "application/json",
- Authorization: `Bearer ${apiKey}`,
- },
- body: JSON.stringify({
- model: modelId,
- messages: [
- { role: "user", content: "Say 'OK'" },
- ],
- max_tokens: 20,
- stream: true,
- enable_thinking: false,
- }),
- },
- )
-
- if (!response.ok) {
- const errorText = await response.text()
- throw new Error(
- `ModelScope API error (${response.status}): ${errorText}`,
- )
- }
-
- const contentType =
- response.headers.get("content-type") || ""
- const isValidStreamingResponse =
- response.status === 200 &&
- (contentType.includes("text/event-stream") ||
- contentType.includes("application/json"))
-
- if (!isValidStreamingResponse) {
- throw new Error(
- `Unexpected response format: ${contentType}`,
- )
- }
-
- const responseTime = Date.now() - startTime
-
- if (response.body) {
- response.body.cancel().catch(() => {
- /* Ignore cancellation errors */
- })
- }
-
- return NextResponse.json({
- valid: true,
- responseTime,
- note: "ModelScope model validated (using streaming API)",
- })
- } catch (error) {
- console.error(
- "[validate-model] ModelScope validation failed:",
- error,
- )
- throw error
- }
- }
-
- default:
- return NextResponse.json(
- { valid: false, error: `Unknown provider: ${provider}` },
- { status: 400 },
- )
+ // The Test button checks the user's own provider. On the server's
+ // keys (Ollama Cloud without a key or URL) anyone could run any model.
+ if (
+ usesServerCredentials(provider, {
+ apiKey,
+ baseUrl,
+ awsAccessKeyId,
+ awsSecretAccessKey,
+ vertexApiKey,
+ })
+ ) {
+ return NextResponse.json(
+ { valid: false, error: "API key is required" },
+ { status: 400 },
+ )
}
- // Make a minimal test request
- const startTime = Date.now()
- await generateText({
- model,
- prompt: "Say 'OK'",
- maxOutputTokens: 20,
+ // On the deployment's own endpoints a Test runs a model as a chat
+ // does, so with the quota on it counts as a chat request (an
+ // admin's Test of the server's URL does not)
+ const userId = getUserIdFromRequest(req)
+ if (
+ isQuotaEnabled() &&
+ !serverUrl &&
+ userId !== "anonymous" &&
+ (await usesServerEndpoint(
+ provider,
+ normalizeBaseUrl(body.baseUrl ?? ""),
+ apiKey,
+ ))
+ ) {
+ const quotaCheck = await checkAndIncrementRequest(userId, {
+ requests: Number(process.env.DAILY_REQUEST_LIMIT) || 10,
+ tokens: Number(process.env.DAILY_TOKEN_LIMIT) || 200000,
+ tpm: Number(process.env.TPM_LIMIT) || 20000,
+ })
+ if (!quotaCheck.allowed) {
+ return NextResponse.json(
+ { valid: false, error: quotaCheck.error },
+ { status: 429 },
+ )
+ }
+ }
+
+ // The same model the chat would use. A client base URL makes it
+ // refuse redirects to internal hosts.
+ const { model } = getAIModel({
+ provider,
+ modelId,
+ apiKey,
+ baseUrl,
+ trustedBaseUrl: serverUrl,
+ awsAccessKeyId,
+ awsSecretAccessKey,
+ awsRegion,
+ // Temporary AWS credentials need it, as in the chat
+ awsSessionToken,
+ vertexApiKey,
+ // EdgeOne checks the Pages cookies and the access code
+ ...(provider === "edgeone" && {
+ headers: {
+ cookie: req.headers.get("cookie") || "",
+ "x-access-code": req.headers.get("x-access-code") || "",
+ },
+ }),
})
+
+ // Streaming, like the chat (some models only stream). Stop at the
+ // first tool call; a reasoning model that runs out of tokens first
+ // proves the connection but not tool support.
+ const startTime = Date.now()
+ const result = streamText({
+ model,
+ prompt: "Call the ping tool.",
+ tools: { ping: PING_TOOL },
+ maxOutputTokens: 1024,
+ maxRetries: 0,
+ abortSignal: AbortSignal.timeout(TEST_TIMEOUT_MS),
+ })
+ let calledTool = false
+ let finishReason: string | undefined
+ for await (const part of result.fullStream) {
+ if (part.type === "error") throw part.error
+ // The timeout ends the stream with an abort part, not an error
+ if (part.type === "abort") {
+ const timeout = new Error(
+ `The model did not answer within ${TEST_TIMEOUT_MS / 1000} s.`,
+ )
+ timeout.name = "TimeoutError"
+ throw timeout
+ }
+ if (part.type === "tool-call") {
+ calledTool = true
+ break
+ }
+ if (part.type === "finish") finishReason = part.finishReason
+ }
const responseTime = Date.now() - startTime
return NextResponse.json({
valid: true,
responseTime,
+ ...(!calledTool &&
+ finishReason !== "length" && { warning: NO_TOOL_CALL_WARNING }),
})
} catch (error) {
console.error("[validate-model] Error:", error)
- let errorMessage = "Validation failed"
- if (error instanceof Error) {
- // Extract meaningful error message
- if (
- error.message.includes("401") ||
- error.message.includes("Unauthorized")
- ) {
- errorMessage = "Invalid API key"
- } else if (
- error.message.includes("404") ||
- error.message.includes("not found")
- ) {
- errorMessage = "Model not found"
- } else if (
- error.message.includes("429") ||
- error.message.includes("rate limit")
- ) {
- errorMessage = "Rate limited - try again later"
- } else if (error.message.includes("ECONNREFUSED")) {
- errorMessage = "Cannot connect to server"
- } else {
- errorMessage = error.message.slice(0, 100)
- }
- }
-
+ const { code, message } = classifyLLMError(error)
return NextResponse.json(
- { valid: false, error: errorMessage },
+ { valid: false, code, error: message },
{ status: 200 }, // Return 200 so client can read error message
)
}
diff --git a/app/api/verify-access-code/route.ts b/app/api/verify-access-code/route.ts
index d69f59d1..55cbc94b 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/biome.json b/biome.json
index 32874167..bf56b8b0 100644
--- a/biome.json
+++ b/biome.json
@@ -1,12 +1,18 @@
{
- "$schema": "https://biomejs.dev/schemas/2.3.10/schema.json",
+ "$schema": "https://biomejs.dev/schemas/2.4.14/schema.json",
"vcs": {
"enabled": true,
"clientKind": "git",
"useIgnoreFile": true
},
"files": {
- "ignoreUnknown": false
+ "ignoreUnknown": false,
+ "includes": [
+ "**",
+ "!public",
+ "!packages/mcp-server/src/preview",
+ "!lib/model-catalog.json"
+ ]
},
"formatter": {
"enabled": true,
diff --git a/components/ai-elements/model-selector.tsx b/components/ai-elements/model-selector.tsx
index 1b71cb70..7164f44b 100644
--- a/components/ai-elements/model-selector.tsx
+++ b/components/ai-elements/model-selector.tsx
@@ -1,5 +1,6 @@
import { Cloud } from "lucide-react"
-import type { ComponentProps, ReactNode } from "react"
+import type { ComponentProps, ElementRef, ReactNode } from "react"
+import { useEffect, useRef, useState } from "react"
import {
Command,
CommandDialog,
@@ -69,20 +70,62 @@ export type ModelSelectorListProps = ComponentProps
export const ModelSelectorList = ({
className,
...props
-}: ModelSelectorListProps) => (
-
-
- {/* Bottom shadow indicator for scrollable content */}
-
-
-)
+}: ModelSelectorListProps) => {
+ const listRef = useRef>(null)
+ const [showShadow, setShowShadow] = useState(false)
+
+ useEffect(() => {
+ const listElement = listRef.current
+ if (!listElement) return
+
+ const checkScroll = () => {
+ const { scrollTop, scrollHeight, clientHeight } = listElement
+ // Show shadow if there is more content below
+ // Using a small threshold to handle fractional pixel rendering
+ setShowShadow(
+ scrollHeight > Math.ceil(scrollTop + clientHeight) + 1,
+ )
+ }
+
+ // Initial check
+ checkScroll()
+
+ // Event listeners
+ listElement.addEventListener("scroll", checkScroll)
+ window.addEventListener("resize", checkScroll)
+
+ // Observe content changes (e.g. async loading of items)
+ const observer = new MutationObserver(checkScroll)
+ observer.observe(listElement, { childList: true, subtree: true })
+
+ return () => {
+ listElement.removeEventListener("scroll", checkScroll)
+ window.removeEventListener("resize", checkScroll)
+ observer.disconnect()
+ }
+ }, [])
+
+ return (
+
+
+ {/* Bottom shadow indicator for scrollable content */}
+
+
+ )
+}
export type ModelSelectorEmptyProps = ComponentProps
@@ -169,3 +212,27 @@ export const ModelSelectorName = ({
}: ModelSelectorNameProps) => (
)
+
+export type ModelSelectorSectionHeaderProps = {
+ icon: ReactNode
+ label: string
+ className?: string
+}
+
+export const ModelSelectorSectionHeader = ({
+ icon,
+ label,
+ className,
+}: ModelSelectorSectionHeaderProps) => (
+
+
+ {icon}
+
+ {label}
+
+)
diff --git a/components/chat-example-panel.tsx b/components/chat-example-panel.tsx
index a74f42cc..4e721292 100644
--- a/components/chat-example-panel.tsx
+++ b/components/chat-example-panel.tsx
@@ -141,9 +141,6 @@ export default function ExamplePanel({
{dict.examples.mcpServer}
-
- {dict.examples.preview}
-
{dict.examples.mcpDescription}
diff --git a/components/chat-input.tsx b/components/chat-input.tsx
index 6848dc58..0f375925 100644
--- a/components/chat-input.tsx
+++ b/components/chat-input.tsx
@@ -1,17 +1,28 @@
"use client"
import {
+ BookmarkPlus,
Download,
History,
Image as ImageIcon,
Link,
- Loader2,
Send,
+ Square,
} from "lucide-react"
import type React from "react"
-import { useCallback, useEffect, useRef, useState } from "react"
+import {
+ type Dispatch,
+ forwardRef,
+ type SetStateAction,
+ useCallback,
+ useEffect,
+ useImperativeHandle,
+ useRef,
+ useState,
+} from "react"
import { toast } from "sonner"
import { ButtonWithTooltip } from "@/components/button-with-tooltip"
+import { TemplateCreateDialog } from "@/components/chat/TemplateCreateDialog"
import { ErrorToast } from "@/components/error-toast"
import { HistoryDialog } from "@/components/history-dialog"
import { ModelSelector } from "@/components/model-selector"
@@ -27,13 +38,25 @@ import { isPdfFile, isTextFile } from "@/lib/pdf-utils"
import { STORAGE_KEYS } from "@/lib/storage"
import type { FlattenedModel } from "@/lib/types/model-config"
import { extractUrlContent, type UrlData } from "@/lib/url-utils"
+import { isRealDiagram } from "@/lib/utils"
import { FilePreviewList } from "./file-preview-list"
const MAX_IMAGE_SIZE = 2 * 1024 * 1024 // 2MB
const MAX_FILES = 5
+// Image formats every supported model provider accepts (SVG is read as text)
+const SUPPORTED_IMAGE_TYPES = [
+ "image/png",
+ "image/jpeg",
+ "image/gif",
+ "image/webp",
+]
function isValidFileType(file: File): boolean {
- return file.type.startsWith("image/") || isPdfFile(file) || isTextFile(file)
+ return (
+ SUPPORTED_IMAGE_TYPES.includes(file.type) ||
+ isPdfFile(file) ||
+ isTextFile(file)
+ )
}
function formatFileSize(bytes: number): string {
@@ -137,11 +160,16 @@ function showValidationErrors(errors: string[], dict: any) {
}
}
+export interface ChatInputRef {
+ focus: () => void
+}
+
interface ChatInputProps {
input: string
status: "submitted" | "streaming" | "ready" | "error"
onSubmit: (e: React.FormEvent) => void
onChange: (e: React.ChangeEvent) => void
+ onStop?: () => void
files?: File[]
onFileChange?: (files: File[]) => void
pdfData?: Map<
@@ -149,7 +177,7 @@ interface ChatInputProps {
{ text: string; charCount: number; isExtracting: boolean }
>
urlData?: Map
- onUrlChange?: (data: Map) => void
+ onUrlChange?: Dispatch>>
sessionId?: string
error?: Error | null
@@ -157,122 +185,230 @@ interface ChatInputProps {
models?: FlattenedModel[]
selectedModelId?: string
onModelSelect?: (modelId: string | undefined) => void
- showUnvalidatedModels?: boolean
onConfigureModels?: () => void
+ showUnvalidatedModels?: boolean
+ // Focus control props
+ shouldFocus?: boolean
+ onFocused?: () => void
}
-export function ChatInput({
- input,
- status,
- onSubmit,
- onChange,
- files = [],
- onFileChange = () => {},
- pdfData = new Map(),
- urlData,
- onUrlChange,
- sessionId,
- error = null,
- models = [],
- selectedModelId,
- onModelSelect = () => {},
- showUnvalidatedModels = false,
- onConfigureModels = () => {},
-}: ChatInputProps) {
- const dict = useDictionary()
- const {
- diagramHistory,
- saveDiagramToFile,
- showSaveDialog,
- setShowSaveDialog,
- } = useDiagram()
+export const ChatInput = forwardRef(
+ function ChatInput(
+ {
+ input,
+ status,
+ onSubmit,
+ onChange,
+ onStop,
+ files = [],
+ onFileChange = () => {},
+ pdfData = new Map(),
+ urlData,
+ onUrlChange,
+ sessionId,
+ error = null,
+ models = [],
+ selectedModelId,
+ onModelSelect = () => {},
+ onConfigureModels,
+ showUnvalidatedModels = false,
+ shouldFocus = false,
+ onFocused,
+ },
+ ref,
+ ) {
+ const dict = useDictionary()
+ const {
+ chartXML,
+ diagramHistory,
+ saveDiagramToFile,
+ showSaveDialog,
+ setShowSaveDialog,
+ } = useDiagram()
- const textareaRef = useRef(null)
- const fileInputRef = useRef(null)
- const [isDragging, setIsDragging] = useState(false)
- const [showHistory, setShowHistory] = useState(false)
- const [showUrlDialog, setShowUrlDialog] = useState(false)
- const [isExtractingUrl, setIsExtractingUrl] = useState(false)
- const [sendShortcut, setSendShortcut] = useState("ctrl-enter")
- // Allow retry when there's an error (even if status is still "streaming" or "submitted")
- const isDisabled =
- (status === "streaming" || status === "submitted") && !error
+ const textareaRef = useRef(null)
+ const fileInputRef = useRef(null)
+ const [isDragging, setIsDragging] = useState(false)
- const adjustTextareaHeight = useCallback(() => {
- const textarea = textareaRef.current
- if (textarea) {
- textarea.style.height = "auto"
- textarea.style.height = `${Math.min(textarea.scrollHeight, 200)}px`
- }
- }, [])
- // Handle programmatic input changes (e.g., setInput("") after form submission)
- useEffect(() => {
- adjustTextareaHeight()
- }, [input, adjustTextareaHeight])
+ // Expose focus method via ref
+ useImperativeHandle(ref, () => ({
+ focus: () => {
+ textareaRef.current?.focus()
+ },
+ }))
- // Load send shortcut preference from localStorage and listen for changes
- useEffect(() => {
- const stored = localStorage.getItem(STORAGE_KEYS.sendShortcut)
- if (stored) setSendShortcut(stored)
+ // Focus the textarea when shouldFocus becomes true
+ // Use setTimeout to ensure focus happens after drawio iframe settles
+ useEffect(() => {
+ if (shouldFocus) {
+ const timer = setTimeout(() => {
+ textareaRef.current?.focus()
+ onFocused?.()
+ }, 150)
+ return () => clearTimeout(timer)
+ }
+ }, [shouldFocus, onFocused])
- const handleChange = (e: CustomEvent) =>
- setSendShortcut(e.detail)
- window.addEventListener(
- "sendShortcutChange",
- handleChange as EventListener,
- )
- return () =>
- window.removeEventListener(
+ const [showHistory, setShowHistory] = useState(false)
+ const [showUrlDialog, setShowUrlDialog] = useState(false)
+ const [showSaveAsTemplate, setShowSaveAsTemplate] = useState(false)
+ const [isExtractingUrl, setIsExtractingUrl] = useState(false)
+ const [sendShortcut, setSendShortcut] = useState("ctrl-enter")
+ // Allow retry when there's an error (even if status is still "streaming" or "submitted")
+ const isDisabled =
+ (status === "streaming" || status === "submitted") && !error
+ // Block sending until attached files and URLs have their text, otherwise
+ // their content would be silently dropped
+ const isExtractingAttachments =
+ files.some((file) => pdfData.get(file)?.isExtracting) ||
+ Array.from(urlData?.values() ?? []).some((d) => d.isExtracting)
+
+ const adjustTextareaHeight = useCallback(() => {
+ const textarea = textareaRef.current
+ if (textarea) {
+ textarea.style.height = "auto"
+ textarea.style.height = `${Math.min(textarea.scrollHeight, 200)}px`
+ }
+ }, [])
+ // Handle programmatic input changes (e.g., setInput("") after form submission)
+ useEffect(() => {
+ adjustTextareaHeight()
+ }, [input, adjustTextareaHeight])
+
+ // Load send shortcut preference from localStorage and listen for changes
+ useEffect(() => {
+ const stored = localStorage.getItem(STORAGE_KEYS.sendShortcut)
+ if (stored) setSendShortcut(stored)
+
+ const handleChange = (e: CustomEvent) =>
+ setSendShortcut(e.detail)
+ window.addEventListener(
"sendShortcutChange",
handleChange as EventListener,
)
- }, [])
+ return () =>
+ window.removeEventListener(
+ "sendShortcutChange",
+ handleChange as EventListener,
+ )
+ }, [])
- const handleChange = (e: React.ChangeEvent) => {
- onChange(e)
- adjustTextareaHeight()
- }
+ const handleChange = (e: React.ChangeEvent) => {
+ onChange(e)
+ adjustTextareaHeight()
+ }
- const handleKeyDown = (e: React.KeyboardEvent) => {
- const shouldSend =
- sendShortcut === "enter"
- ? e.key === "Enter" && !e.shiftKey && !e.ctrlKey && !e.metaKey
- : (e.metaKey || e.ctrlKey) && e.key === "Enter"
+ const handleKeyDown = (e: React.KeyboardEvent) => {
+ // Enter that confirms an IME candidate must not send the message
+ if (e.nativeEvent.isComposing || e.keyCode === 229) return
- if (shouldSend) {
- e.preventDefault()
- const form = e.currentTarget.closest("form")
- if (form && input.trim() && !isDisabled) {
- form.requestSubmit()
+ const shouldSend =
+ sendShortcut === "enter"
+ ? e.key === "Enter" &&
+ !e.shiftKey &&
+ !e.ctrlKey &&
+ !e.metaKey
+ : (e.metaKey || e.ctrlKey) && e.key === "Enter"
+
+ if (shouldSend) {
+ e.preventDefault()
+ const form = e.currentTarget.closest("form")
+ if (
+ form &&
+ input.trim() &&
+ !isDisabled &&
+ !isExtractingAttachments
+ ) {
+ form.requestSubmit()
+ }
}
}
- }
- const handlePaste = async (e: React.ClipboardEvent) => {
- if (isDisabled) return
+ const handlePaste = async (e: React.ClipboardEvent) => {
+ if (isDisabled) return
- const items = e.clipboardData.items
- const imageItems = Array.from(items).filter((item) =>
- item.type.startsWith("image/"),
- )
+ const items = e.clipboardData.items
+ const imageItems = Array.from(items).filter((item) =>
+ item.type.startsWith("image/"),
+ )
- if (imageItems.length > 0) {
- const imageFiles = (
- await Promise.all(
- imageItems.map(async (item, index) => {
- const file = item.getAsFile()
- if (!file) return null
- return new File(
- [file],
- `pasted-image-${Date.now()}-${index}.${file.type.split("/")[1]}`,
- { type: file.type },
- )
- }),
+ if (imageItems.length > 0) {
+ const imageFiles = (
+ await Promise.all(
+ imageItems.map(async (item, index) => {
+ const file = item.getAsFile()
+ if (!file) return null
+ return new File(
+ [file],
+ `pasted-image-${Date.now()}-${index}.${file.type.split("/")[1]}`,
+ { type: file.type },
+ )
+ }),
+ )
+ ).filter((f): f is File => f !== null)
+
+ const { validFiles, errors } = validateFiles(
+ imageFiles,
+ files.length,
+ dict,
)
- ).filter((f): f is File => f !== null)
+ showValidationErrors(errors, dict)
+ if (validFiles.length > 0) {
+ onFileChange([...files, ...validFiles])
+ }
+ }
+ }
+ const handleFileChange = (e: React.ChangeEvent) => {
+ const newFiles = Array.from(e.target.files || [])
const { validFiles, errors } = validateFiles(
- imageFiles,
+ newFiles,
+ files.length,
+ dict,
+ )
+ showValidationErrors(errors, dict)
+ if (validFiles.length > 0) {
+ onFileChange([...files, ...validFiles])
+ }
+
+ if (fileInputRef.current) {
+ fileInputRef.current.value = ""
+ }
+ }
+
+ const handleRemoveFile = (fileToRemove: File) => {
+ onFileChange(files.filter((file) => file !== fileToRemove))
+ if (fileInputRef.current) {
+ fileInputRef.current.value = ""
+ }
+ }
+
+ const triggerFileInput = () => {
+ fileInputRef.current?.click()
+ }
+
+ const handleDragOver = (e: React.DragEvent) => {
+ e.preventDefault()
+ e.stopPropagation()
+ setIsDragging(true)
+ }
+
+ const handleDragLeave = (e: React.DragEvent) => {
+ e.preventDefault()
+ e.stopPropagation()
+ setIsDragging(false)
+ }
+
+ const handleDrop = (e: React.DragEvent) => {
+ e.preventDefault()
+ e.stopPropagation()
+ setIsDragging(false)
+
+ if (isDisabled) return
+
+ // Let validateFiles show a toast for unsupported types
+ const { validFiles, errors } = validateFiles(
+ Array.from(e.dataTransfer.files),
files.length,
dict,
)
@@ -281,278 +417,253 @@ export function ChatInput({
onFileChange([...files, ...validFiles])
}
}
- }
- const handleFileChange = (e: React.ChangeEvent) => {
- const newFiles = Array.from(e.target.files || [])
- const { validFiles, errors } = validateFiles(
- newFiles,
- files.length,
- dict,
- )
- showValidationErrors(errors, dict)
- if (validFiles.length > 0) {
- onFileChange([...files, ...validFiles])
+ const handleUrlExtract = async (url: string) => {
+ if (!onUrlChange) return
+
+ setIsExtractingUrl(true)
+
+ // Use functional updates so a removal or send made while extracting
+ // is not overwritten when the request finishes
+ try {
+ onUrlChange((prev) =>
+ new Map(prev).set(url, {
+ url,
+ title: url,
+ content: "",
+ charCount: 0,
+ isExtracting: true,
+ }),
+ )
+
+ const data = await extractUrlContent(url)
+
+ // Skip if the URL was removed while extracting
+ onUrlChange((prev) =>
+ prev.has(url) ? new Map(prev).set(url, data) : prev,
+ )
+
+ setShowUrlDialog(false)
+ } catch (error) {
+ // Remove the URL from the data map on error
+ onUrlChange((prev) => {
+ const next = new Map(prev)
+ next.delete(url)
+ return next
+ })
+ showErrorToast(
+
+ {error instanceof Error
+ ? error.message
+ : "Failed to extract URL content"}
+ ,
+ )
+ } finally {
+ setIsExtractingUrl(false)
+ }
}
- if (fileInputRef.current) {
- fileInputRef.current.value = ""
- }
- }
-
- const handleRemoveFile = (fileToRemove: File) => {
- onFileChange(files.filter((file) => file !== fileToRemove))
- if (fileInputRef.current) {
- fileInputRef.current.value = ""
- }
- }
-
- const triggerFileInput = () => {
- fileInputRef.current?.click()
- }
-
- const handleDragOver = (e: React.DragEvent) => {
- e.preventDefault()
- e.stopPropagation()
- setIsDragging(true)
- }
-
- const handleDragLeave = (e: React.DragEvent) => {
- e.preventDefault()
- e.stopPropagation()
- setIsDragging(false)
- }
-
- const handleDrop = (e: React.DragEvent) => {
- e.preventDefault()
- e.stopPropagation()
- setIsDragging(false)
-
- if (isDisabled) return
-
- const droppedFiles = e.dataTransfer.files
- const supportedFiles = Array.from(droppedFiles).filter((file) =>
- isValidFileType(file),
- )
-
- const { validFiles, errors } = validateFiles(
- supportedFiles,
- files.length,
- dict,
- )
- showValidationErrors(errors, dict)
- if (validFiles.length > 0) {
- onFileChange([...files, ...validFiles])
- }
- }
-
- const handleUrlExtract = async (url: string) => {
- if (!onUrlChange) return
-
- setIsExtractingUrl(true)
-
- try {
- const existing = urlData
- ? new Map(urlData)
- : new Map()
- existing.set(url, {
- url,
- title: url,
- content: "",
- charCount: 0,
- isExtracting: true,
- })
- onUrlChange(existing)
-
- const data = await extractUrlContent(url)
-
- const newUrlData = new Map(existing)
- newUrlData.set(url, data)
- onUrlChange(newUrlData)
-
- setShowUrlDialog(false)
- } catch (error) {
- // Remove the URL from the data map on error
- const newUrlData = urlData
- ? new Map(urlData)
- : new Map()
- newUrlData.delete(url)
- onUrlChange(newUrlData)
- showErrorToast(
-
- {error instanceof Error
- ? error.message
- : "Failed to extract URL content"}
- ,
- )
- } finally {
- setIsExtractingUrl(false)
- }
- }
-
- return (
-