From c0fa997186345c1bad29443edf7fe645d3b488e1 Mon Sep 17 00:00:00 2001 From: "dayuan.jiang" Date: Mon, 5 Oct 2026 18:57:02 +0900 Subject: [PATCH] fix: older defects (batch C) and the second batch's review Chats: - New Chat right after an answer saves that chat once. Saves run one at a time and read the chat on screen when their turn comes; a save scheduled for a chat that is no longer on screen is dropped. A chat whose id was still on its way to the URL no longer comes back after New Chat (the next answer went into it). - Crossing the 768 px breakpoint keeps the chat panel: a streaming answer, unsaved messages and attachments stay. The panel gets the sizes of each side, and a panel collapsed on desktop opens on mobile. - The chat's export waits for its own reply: an edit's history export still on its way no longer answers it with the older diagram, and two file saves at once no longer swap results. - A second edit in one answer is previewed on the first edit's result. - Stop also ends a running screenshot check; a chat that cannot be saved (storage full) can be left with "Continue without saving". - Small diagrams with shapes count as diagrams; the tool card no longer crashes on malformed operations. Quota and providers: - Requests that reach the server's own endpoints count toward the quota: EdgeOne (always its own endpoint now), a private base URL whatever key header is sent, keyless Ollama without a URL. With the quota on, a redirect is followed only to a public address. The output cap applies to these requests too. - Stop records the tokens of the steps that finished; the screenshot check counts its tokens without counting a request. - EdgeOne configured only by AI_PROVIDER works, also in the admin Test, which forwards the access code. Azure set up only in the admin panel works in chat. The Test sends a Bedrock session token. - The admin panel's Test of an entry without a URL uses the server's URL as the server does (no private address check for it); the admin panel no longer writes an Ollama URL. MCP server: - Write tools and start_session run one at a time, so two at once never drop each other's change; a cancelled call waiting its turn is skipped. get_diagram and export_diagram keep the session they started with. - Export to .drawio first gets the user's latest edits from the browser. - History thumbnails: one that arrives after the next AI write is dropped; a sync reply keeps the image; a version that changed only page settings is its own entry. - A diagram over the 10 MB limit is saved without its image, or the user is told to download it (the server now answers 413 instead of cutting the connection). - Labels holding text like id='1' or parent='1' are no longer read as attributes (a layer or a parent was deleted). A broken bare file is refused. - After a sync reply the tab no longer sends its autosave copy again. Desktop and files: - A newer switch of the same preset is not rolled back by an older one that failed. .env values with escaped quotes are read whole. - MCP saved files: a file that could not be read stays protected while a folder without permission hides it, and is saved again once deleted. - The desktop app reports "no chats" only when the count was read and no model settings are stored. --- app/[lang]/admin/models-section.tsx | 6 + app/[lang]/page.tsx | 13 +- app/api/admin/test-model/route.ts | 14 +- app/api/chat/route.ts | 85 ++-- app/api/validate-diagram/route.ts | 44 ++- app/api/validate-model/route.ts | 34 +- components/chat-message-display.tsx | 14 +- components/chat-panel.tsx | 136 ++++--- components/chat/ToolCallCard.tsx | 14 +- components/model-config-dialog.tsx | 34 +- contexts/diagram-context.tsx | 236 ++++++----- electron/main/app-menu.ts | 9 +- electron/main/env-loader.ts | 16 +- hooks/use-diagram-tool-handlers.ts | 20 +- hooks/use-session-manager.ts | 261 ++++++------ hooks/use-validate-diagram.ts | 15 + lib/admin/providers.ts | 7 +- lib/ai-providers.ts | 49 ++- lib/dynamo-quota-manager.ts | 5 +- lib/i18n/dictionaries/en.json | 2 + lib/i18n/dictionaries/ja.json | 2 + lib/i18n/dictionaries/zh-Hant.json | 2 + lib/i18n/dictionaries/zh.json | 2 + lib/session-storage.ts | 9 +- lib/ssrf-protection.ts | 44 ++- lib/utils.ts | 7 +- packages/mcp-server/src/exclusive.ts | 25 ++ packages/mcp-server/src/history.ts | 17 +- packages/mcp-server/src/http-server.ts | 47 ++- packages/mcp-server/src/index.ts | 129 +++--- packages/mcp-server/src/load-diagram.ts | 10 +- packages/mcp-server/src/new-diagram.ts | 10 +- packages/mcp-server/src/pages.ts | 15 +- packages/mcp-server/src/persistence.ts | 34 +- packages/mcp-server/src/preview/preview.js | 61 ++- packages/mcp-server/src/xml-attributes.ts | 27 ++ packages/mcp-server/src/xml-validation.ts | 45 +-- packages/mcp-server/tests/exclusive.test.ts | 59 +++ packages/mcp-server/tests/http-server.test.ts | 133 ++++++- .../mcp-server/tests/load-diagram.test.ts | 7 + packages/mcp-server/tests/persistence.test.ts | 37 ++ packages/mcp-server/tests/wrap-cells.test.ts | 15 + .../mcp-server/tests/xml-validation.test.ts | 27 ++ tests/e2e/chat.spec.ts | 87 +++- tests/e2e/diagram-content.spec.ts | 50 +++ tests/e2e/history-restore.spec.ts | 86 ++++ tests/e2e/provider-models.spec.ts | 29 ++ tests/unit/admin-providers.test.ts | 12 +- tests/unit/admin-test-model.test.ts | 13 + tests/unit/ai-providers-credentials.test.ts | 45 +++ tests/unit/app-menu.test.ts | 16 + .../chat-message-display-preview.test.tsx | 68 ++++ tests/unit/chat-route-abort.test.ts | 139 +++++++ tests/unit/chat-route-edgeone.test.ts | 60 +++ tests/unit/chat-route-errors.test.ts | 122 ++++++ tests/unit/chat-route-quota.test.ts | 61 +++ tests/unit/diagram-context.test.tsx | 119 ++++++ tests/unit/env-loader.test.ts | 20 + tests/unit/mcp-preview-recovery.test.ts | 372 ++++++++++++++++++ tests/unit/ssrf-protection.test.ts | 75 +++- tests/unit/tool-call-card.test.tsx | 43 ++ tests/unit/use-diagram-tool-handlers.test.tsx | 71 ++++ tests/unit/use-session-manager.test.tsx | 170 ++++++++ tests/unit/utils.test.ts | 34 +- tests/unit/validate-diagram-route.test.ts | 60 ++- tests/unit/validate-model-bedrock.test.ts | 34 ++ tests/unit/validate-model-route.test.ts | 119 +++++- 67 files changed, 3152 insertions(+), 531 deletions(-) create mode 100644 packages/mcp-server/src/exclusive.ts create mode 100644 packages/mcp-server/src/xml-attributes.ts create mode 100644 packages/mcp-server/tests/exclusive.test.ts create mode 100644 tests/unit/chat-message-display-preview.test.tsx create mode 100644 tests/unit/chat-route-abort.test.ts create mode 100644 tests/unit/diagram-context.test.tsx create mode 100644 tests/unit/mcp-preview-recovery.test.ts create mode 100644 tests/unit/tool-call-card.test.tsx create mode 100644 tests/unit/use-session-manager.test.tsx create mode 100644 tests/unit/validate-model-bedrock.test.ts diff --git a/app/[lang]/admin/models-section.tsx b/app/[lang]/admin/models-section.tsx index b7a16d57..242d3ab8 100644 --- a/app/[lang]/admin/models-section.tsx +++ b/app/[lang]/admin/models-section.tsx @@ -33,6 +33,7 @@ import { import { Switch } from "@/components/ui/switch" import { useDictionary } from "@/hooks/use-dictionary" import { formatMessage } from "@/lib/i18n/utils" +import { STORAGE_KEYS } from "@/lib/storage" import { FIXED_CRED_PROVIDERS, generateId, @@ -88,6 +89,11 @@ function ProviderDetail({ try { const data = await adminFetch("/api/admin/test-model", password, { method: "POST", + // EdgeOne's function also checks the access code + headers: { + "x-access-code": + localStorage.getItem(STORAGE_KEYS.accessCode) || "", + }, body: JSON.stringify({ provider, modelId }), }) setTestResults((prev) => ({ diff --git a/app/[lang]/page.tsx b/app/[lang]/page.tsx index 9c32a2b2..79d0a2e8 100644 --- a/app/[lang]/page.tsx +++ b/app/[lang]/page.tsx @@ -107,8 +107,8 @@ export default function Home() { resetDrawioReady() } - // Check mobile. The draw.io iframe is not remounted when crossing the - // breakpoint (only the chat panel is), so its ready state stays as is. + // Check mobile. No panel is remounted when crossing the breakpoint, so + // the draw.io ready state and the chat's turn stay as they are. useEffect(() => { const checkMobile = () => { setIsMobile(window.innerWidth < 768) @@ -119,6 +119,14 @@ export default function Home() { return () => window.removeEventListener("resize", checkMobile) }, []) + // Give the chat panel the size of this side of the breakpoint. It is + // open on both sides: the mobile panel cannot be collapsed, and one + // collapsed on desktop comes back open + useEffect(() => { + chatPanelRef.current?.resize(isMobile ? 50 : 33) + setIsChatVisible(true) + }, [isMobile]) + const toggleChatPanel = () => { const panel = chatPanelRef.current if (panel) { @@ -212,7 +220,6 @@ export default function Home() { {/* Chat Panel */} _BASE_URL: test that endpoint, not - // another one - baseUrl: resolved.baseUrl || globalBaseUrl(resolved.provider), + // 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, diff --git a/app/api/chat/route.ts b/app/api/chat/route.ts index ccc0b346..3c66756b 100644 --- a/app/api/chat/route.ts +++ b/app/api/chat/route.ts @@ -13,6 +13,7 @@ import { z } from "zod" import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" import { CACHE_POINT, + edgeOneEndpoint, getAIModel, getServerProvider, SINGLE_SYSTEM_PROVIDERS, @@ -189,15 +190,16 @@ async function handleChatRequest(req: Request): Promise { } // 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 - const isEdgeOne = (serverModelConfig.provider || provider) === "edgeone" + // 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" - // For EdgeOne provider, construct full URL from request origin - // because createOpenAI needs absolute URL, not relative path - if (isEdgeOne && !baseUrl) { - const origin = req.headers.get("origin") || new URL(req.url).origin - baseUrl = `${origin}/api/edgeai` - } + // 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 @@ -212,8 +214,12 @@ async function handleChatRequest(req: Request): Promise { const cookieHeader = req.headers.get("cookie") const clientOverrides = { - // Server model provider takes precedence over client header - provider: serverModelConfig.provider || provider, + // 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 @@ -274,18 +280,24 @@ async function handleChatRequest(req: Request): Promise { // === SERVER-SIDE QUOTA CHECK START === // Quota is opt-in (DYNAMODB_QUOTA_TABLE) and counts what runs on the - // server's keys, or on its keyless Ollama or EdgeOne. Decided by the key - // actually used: a key header the provider never reads must not skip it. - // EdgeOne never reads one; keyless Ollama at a private address is the - // server's own network. + // 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 usesClientBaseUrl = + resolvedProvider !== "bedrock" && resolvedProvider !== "edgeone" const onServerEndpoint = - (resolvedProvider === "edgeone" && !clientBaseUrl) || - (resolvedProvider === "ollama" && - !clientOverrides.apiKey && - (!clientBaseUrl || (await isPrivateUrl(clientBaseUrl)))) + process.env.NEXT_AI_DRAWIO_DESKTOP !== "1" && + (resolvedProvider === "edgeone" || + (resolvedProvider === "ollama" && + !clientBaseUrl && + !clientOverrides.apiKey) || + (usesClientBaseUrl && + !!clientBaseUrl && + (await isPrivateUrl(clientBaseUrl)))) const countsQuota = isQuotaEnabled() && (onServerCredentials || onServerEndpoint) && @@ -317,11 +329,11 @@ async function handleChatRequest(req: Request): Promise { ) // 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 it can only - // lower it + // 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, + onServerCredentials || onServerEndpoint, ) console.log(`[maxOutputTokens] ${maxOutputTokens}`) @@ -353,8 +365,13 @@ async function handleChatRequest(req: Request): Promise { ${userInputText} """` - // Convert UIMessages to ModelMessages and add system message - const modelMessages = await convertToModelMessages(messages) + // 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) { @@ -541,6 +558,8 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on const allMessages = [...systemMessages, ...enhancedMessages] + // Set by onAbort, which records the finished steps' tokens itself + let stopped = false const result = streamText({ model, // The system messages carry cache points, so they go in messages. @@ -606,7 +625,7 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on // Record token usage for server-side quota tracking (if enabled) // Use totalUsage (cumulative across all steps) instead of usage (final step only) // inputTokens already includes cache reads and writes in AI SDK 6 - if (countsQuota && totalUsage) { + if (countsQuota && totalUsage && !stopped) { const totalTokens = (totalUsage.inputTokens || 0) + (totalUsage.outputTokens || 0) @@ -618,7 +637,23 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on console.error(error) // what AI SDK does without an onError endTrace() }, - onAbort: () => 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: { diff --git a/app/api/validate-diagram/route.ts b/app/api/validate-diagram/route.ts index 2f8a3c71..3fba6628 100644 --- a/app/api/validate-diagram/route.ts +++ b/app/api/validate-diagram/route.ts @@ -6,6 +6,12 @@ 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, @@ -78,6 +84,35 @@ export async function POST(req: Request): Promise { ) } + // 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 { @@ -120,7 +155,14 @@ export async function POST(req: Request): Promise { ], maxOutputTokens: 1024, abortSignal: AbortSignal.timeout(timeout), - onFinish: ({ output }) => { + 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}`, diff --git a/app/api/validate-model/route.ts b/app/api/validate-model/route.ts index 907c671e..3a396abe 100644 --- a/app/api/validate-model/route.ts +++ b/app/api/validate-model/route.ts @@ -3,7 +3,12 @@ import { NextResponse } from "next/server" import { z } from "zod" import { checkAccessCode, rejectCrossSite } from "@/lib/access-code" import { checkAdminAuth } from "@/lib/admin/auth" -import { getAIModel, usesServerCredentials } from "@/lib/ai-providers" +import { + edgeOneEndpoint, + getAIModel, + globalBaseUrl, + usesServerCredentials, +} from "@/lib/ai-providers" import { classifyLLMError } from "@/lib/llm-errors" import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection" import type { ProviderName } from "@/lib/types/model-config" @@ -19,8 +24,11 @@ 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 @@ -47,11 +55,11 @@ export async function POST(req: Request) { const { provider, apiKey, - baseUrl, modelId, awsAccessKeyId, awsSecretAccessKey, awsRegion, + awsSessionToken, // Note: Express Mode only needs vertexApiKey vertexApiKey, } = body @@ -62,9 +70,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 && !allowPrivateUrls() && (await isPrivateUrl(baseUrl))) { + if ( + baseUrl && + !serverUrl && + !allowPrivateUrls() && + (await isPrivateUrl(baseUrl)) + ) { return NextResponse.json( { valid: false, error: "Invalid base URL" }, { status: 400 }, @@ -122,9 +147,12 @@ export async function POST(req: Request) { 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" && { diff --git a/components/chat-message-display.tsx b/components/chat-message-display.tsx index de2ecd6b..c149561b 100644 --- a/components/chat-message-display.tsx +++ b/components/chat-message-display.tsx @@ -211,7 +211,7 @@ export function ChatMessageDisplay({

) } - const { chartXML, loadDiagram: onDisplayChart } = useDiagram() + const { chartXML, chartXMLRef, loadDiagram: onDisplayChart } = useDiagram() const messagesEndRef = useRef(null) const scrollTopRef = useRef(null) const previousXML = useRef("") @@ -422,10 +422,12 @@ export function ChatMessageDisplay({ // Previous messages are already processed and won't change const messagesToProcess = messages.length > 0 ? [messages[messages.length - 1]] : [] - // The diagram without streamed previews. Undoing a failed edit's - // preview below changes it before chartXML catches up, and an edit - // streaming right after must start from the undone diagram. - let baseXml = chartXML + // The diagram without streamed previews, as loaded last: the tool + // handler's result of an earlier edit is there before the chartXML + // state catches up. Undoing a failed edit's preview below changes it + // too, and an edit streaming right after must start from the undone + // diagram. + let baseXml = chartXMLRef.current messagesToProcess.forEach((message) => { // Messages restored from a saved session were applied before it was @@ -587,7 +589,7 @@ export function ChatMessageDisplay({ }) } }) - }, [messages, handleDisplayChart, chartXML]) + }, [messages, handleDisplayChart, chartXMLRef]) return ( diff --git a/components/chat-panel.tsx b/components/chat-panel.tsx index 8e4f045f..b8cb36de 100644 --- a/components/chat-panel.tsx +++ b/components/chat-panel.tsx @@ -110,7 +110,7 @@ export default function ChatPanel({ loadDiagram: onDisplayChart, handleExport: onExport, handleExportWithoutHistory, - resolverRef, + exportResolversRef, chartXML, chartXMLRef: liveChartXMLRef, latestSvg, @@ -128,21 +128,15 @@ export default function ChatPanel({ const urlSessionId = searchParams.get("session") const onFetchChart = (saveToHistory = true) => { + // Waits for the reply to its own export, by its tag + const tag = saveToHistory ? onExport() : handleExportWithoutHistory() return Promise.race([ new Promise((resolve) => { - resolverRef.current = resolve - if (saveToHistory) { - onExport() - } else { - handleExportWithoutHistory() - } + if (tag) exportResolversRef.current[tag] = resolve }), new Promise((_, reject) => { - const currentResolver = resolverRef.current setTimeout(() => { - if (resolverRef.current === currentResolver) { - resolverRef.current = null - } + delete exportResolversRef.current[tag] reject(new Error("Chart export timed out after 10 seconds")) }, 10000) }), @@ -335,7 +329,8 @@ export default function ChatPanel({ const validationRetryCountRef = useRef(0) // VLM validation hook using AI SDK's useObject - const { validateWithFallback } = useValidateDiagram() + const { validateWithFallback, cancel: cancelValidation } = + useValidateDiagram() // Diagram tool handlers (display_diagram, edit_diagram, append_diagram) const { handleToolCall } = useDiagramToolHandlers({ @@ -352,6 +347,7 @@ export default function ChatPanel({ validateDiagram: validateWithFallback, enableVlmValidation: vlmValidationEnabled, sessionId, + isStopped: () => stoppedRef.current, onValidationStateChange: handleValidationStateChange, }) @@ -675,6 +671,7 @@ export default function ChatPanel({ isAvailable: sessionIsAvailable, currentSessionId, saveCurrentSession, + getChatGeneration, } = sessionManager // Use ref for saveCurrentSession to avoid infinite loop @@ -699,13 +696,14 @@ export default function ChatPanel({ clearTimeout(localStorageDebounceRef.current) } - // Capture current session ID at schedule time to verify at save time - const scheduledForSessionId = currentSessionId + // Capture the chat on screen at schedule time; the save is dropped + // if another chat is on screen by the time it runs + const scheduledForChat = getChatGeneration() // Capture whether there's a REAL diagram NOW (not just empty template) const hasDiagramNow = isRealDiagram(chartXMLRef.current) // Check if this session was just loaded without a diagram const isNodiagramSession = - justLoadedSessionIdRef.current === scheduledForSessionId + justLoadedSessionIdRef.current === currentSessionId // Debounce: save after 1 second of no changes localStorageDebounceRef.current = setTimeout(async () => { @@ -717,7 +715,7 @@ export default function ChatPanel({ }) await saveCurrentSessionRef.current( sessionData, - scheduledForSessionId, + scheduledForChat, ) } } catch (error) { @@ -737,6 +735,7 @@ export default function ChatPanel({ status, sessionIsAvailable, currentSessionId, + getChatGeneration, buildSessionData, ]) @@ -921,25 +920,33 @@ export default function ChatPanel({ } } + // The current chat could not be saved (storage full). The list where + // old chats can be deleted shows only in an empty chat, so let the user + // go on without saving (same toast id: it replaces the plain message) + const offerToContinueUnsaved = useCallback( + (proceed: () => void) => { + toast.error(dict.errors.sessionSaveFailedLeave, { + id: "session-save-failed", + duration: 15000, + action: { + label: dict.errors.continueWithoutSaving, + onClick: proceed, + }, + }) + }, + [dict], + ) + // Handle session switching from history dropdown const handleSelectSession = useCallback( async (sessionId: string) => { if (!sessionManager.isAvailable) return - // Save current session before switching (also a diagram drawn - // without messages); if that failed (storage full), stay on it - if (messages.length > 0 || isRealDiagram(chartXMLRef.current)) { - const sessionData = await buildSessionData({ - withThumbnail: true, - }) - if (!(await sessionManager.saveCurrentSession(sessionData))) { - return - } - } - // Switch to selected session - const sessionData = await sessionManager.switchSession(sessionId) - if (sessionData) { + const open = async () => { + const sessionData = + await sessionManager.switchSession(sessionId) + if (!sessionData) return const hasRealDiagram = isRealDiagram(sessionData.diagramXml) justLoadedSessionRef.current = true @@ -957,8 +964,29 @@ export default function ChatPanel({ syncUIWithSession(sessionData) router.replace(`?session=${sessionId}`, { scroll: false }) } + + // Save current session before switching (also a diagram drawn + // without messages); if that failed (storage full), stay on it + // unless the user goes on without saving it + if (messages.length > 0 || isRealDiagram(chartXMLRef.current)) { + const sessionData = await buildSessionData({ + withThumbnail: true, + }) + if (!(await sessionManager.saveCurrentSession(sessionData))) { + offerToContinueUnsaved(open) + return + } + } + await open() }, - [sessionManager, messages, buildSessionData, syncUIWithSession, router], + [ + sessionManager, + messages, + buildSessionData, + syncUIWithSession, + router, + offerToContinueUnsaved, + ], ) // Handle session deletion from history dropdown @@ -976,20 +1004,7 @@ export default function ChatPanel({ [sessionManager, syncUIWithSession, router, pathname], ) - const handleNewChat = useCallback(async () => { - // Save current session before creating new one (also a diagram - // drawn without messages) - if ( - sessionManager.isAvailable && - (messages.length > 0 || isRealDiagram(chartXMLRef.current)) - ) { - const sessionData = await buildSessionData({ withThumbnail: true }) - // Not saved (storage full): keep the chat on screen - if (!(await sessionManager.saveCurrentSession(sessionData))) return - // Refresh sessions list to ensure dropdown shows the saved session - await sessionManager.refreshSessions() - } - + const startNewChat = useCallback(() => { // Clear session manager state BEFORE clearing URL to prevent race condition // (otherwise the URL update effect would restore the old session URL) sessionManager.clearCurrentSession() @@ -1021,14 +1036,38 @@ export default function ChatPanel({ setMessages, setSessionId, sessionManager, - messages, router, dict.dialogs.clearSuccess, - buildSessionData, setDiagramHistory, pathname, ]) + const handleNewChat = useCallback(async () => { + // Save current session before creating new one (also a diagram + // drawn without messages) + if ( + sessionManager.isAvailable && + (messages.length > 0 || isRealDiagram(chartXMLRef.current)) + ) { + const sessionData = await buildSessionData({ withThumbnail: true }) + // Not saved (storage full): keep the chat on screen, unless the + // user goes on without saving it + if (!(await sessionManager.saveCurrentSession(sessionData))) { + offerToContinueUnsaved(startNewChat) + return + } + // Refresh sessions list to ensure dropdown shows the saved session + await sessionManager.refreshSessions() + } + startNewChat() + }, [ + sessionManager, + messages, + buildSessionData, + offerToContinueUnsaved, + startNewChat, + ]) + // Handle sending a template directly (called from TemplatePanel) const handleSendTemplate = useCallback( async (template: { prompt: string }) => { @@ -1089,6 +1128,9 @@ export default function ChatPanel({ // Handle stop button click const handleStop = useCallback(() => { stoppedRef.current = true + // A running screenshot check holds up the chat (the SDK waits for + // the tool handler): end it, so the call gets its result now + cancelValidation() const lastMessage = messages[messages.length - 1] // Calls the tool handler already took can still show as streaming: // the messages update at most every 150 ms (useChat throttle) @@ -1111,7 +1153,7 @@ export default function ChatPanel({ }) stop() - }, [messages, addToolOutput, stop]) + }, [messages, addToolOutput, stop, cancelValidation]) // Send chat message with headers const sendChatMessage = ( diff --git a/components/chat/ToolCallCard.tsx b/components/chat/ToolCallCard.tsx index 801d4cee..f4c7c258 100644 --- a/components/chat/ToolCallCard.tsx +++ b/components/chat/ToolCallCard.tsx @@ -20,9 +20,15 @@ interface ToolCallCardProps { } function OperationsDisplay({ operations }: { operations: DiagramOperation[] }) { + // Streamed or invalid input can hold anything: show only what React can + // render (an object in place of a string would crash the whole chat) + const shown = operations.filter( + (op) => typeof (op as { operation?: unknown })?.operation === "string", + ) + const text = (value: unknown) => (typeof value === "string" ? value : "") return (
- {operations.map((op, index) => ( + {shown.map((op, index) => (
- cell_id: {op.cell_id} + cell_id: {text(op.cell_id)}
- {op.new_xml && ( + {text(op.new_xml) && (
-                                {op.new_xml}
+                                {text(op.new_xml)}
                             
)} diff --git a/components/model-config-dialog.tsx b/components/model-config-dialog.tsx index a13c95b7..10b522ec 100644 --- a/components/model-config-dialog.tsx +++ b/components/model-config-dialog.tsx @@ -203,6 +203,7 @@ export function ModelConfigDialog({ p?.awsAccessKeyId, p?.awsSecretAccessKey, p?.awsRegion, + p?.awsSessionToken, p?.vertexApiKey, ]) } @@ -451,6 +452,9 @@ export function ModelConfigDialog({ awsSecretAccessKey: selectedProvider.awsSecretAccessKey, awsRegion: selectedProvider.awsRegion, + // Temporary AWS credentials, as the chat sends + awsSessionToken: + selectedProvider.awsSessionToken, // Vertex AI credentials (Express Mode) vertexApiKey: selectedProvider.vertexApiKey, }), @@ -492,19 +496,18 @@ export function ModelConfigDialog({ validationWarning: undefined, } } + // A newer test started: its own results and spinners count, + // whatever the credentials are now (they may have come back) + if (run !== validationRunRef.current) return // Credentials changed during the test: drop the result. A - // change made in this tab already reset the spinners (and a - // newer test may show its own); one from another tab did - // not, so the latest test clears its own (model ids are - // unique, whatever provider is shown). + // change in another tab left the spinner on, so clear it + // (model ids are unique, whatever provider is shown). if (credentialsOf(selectedProviderId) !== askedWith) { - if (run === validationRunRef.current) { - setValidatingModelIds((prev) => { - const next = new Set(prev) - next.delete(model.id) - return next - }) - } + setValidatingModelIds((prev) => { + const next = new Set(prev) + next.delete(model.id) + return next + }) return } // So did this model's id: the result is for the old one @@ -532,13 +535,10 @@ export function ModelConfigDialog({ }) }), ) + if (run !== validationRunRef.current) return if (credentialsOf(selectedProviderId) !== askedWith) { - // The status line belongs to the latest test, and to the - // provider shown now - if ( - run === validationRunRef.current && - selectedProviderIdRef.current === selectedProviderId - ) { + // The status line is about the provider shown now + if (selectedProviderIdRef.current === selectedProviderId) { setValidationStatus("idle") } return diff --git a/contexts/diagram-context.tsx b/contexts/diagram-context.tsx index 6f57b7b3..8ec4b05f 100644 --- a/contexts/diagram-context.tsx +++ b/contexts/diagram-context.tsx @@ -21,9 +21,14 @@ interface DiagramContextType { diagramHistory: { svg: string; xml: string }[] setDiagramHistory: (history: { svg: string; xml: string }[]) => void loadDiagram: (chart: string, skipValidation?: boolean) => string | null - handleExport: () => void - handleExportWithoutHistory: () => void - resolverRef: React.MutableRefObject<((value: string) => void) | null> + // Both return the export's tag (empty when draw.io is not there yet) + handleExport: () => string + handleExportWithoutHistory: () => string + // Pending exports by tag; a history or plain export's resolver gets the + // first page's XML + exportResolversRef: React.MutableRefObject< + Record void> + > drawioRef: React.MutableRefObject handleDiagramExport: (data: EventExport) => void handleDiagramAutoSave: (data: { xml?: string }) => void @@ -45,12 +50,10 @@ interface DiagramContextType { const DiagramContext = createContext(undefined) -// Exports for thumbnails, validation PNGs, history entries and file saves -// carry a tag in the request's `message` field. draw.io echoes the request -// back in the export event, so each result reaches its own caller; untagged -// exports (chat-panel's onFetchChart) resolve resolverRef. Thumbnail, -// validation and history tags end in a request number, so a late result -// never answers a newer request. +// Every export carries a tag in the request's `message` field. draw.io +// echoes the request back in the export event, so each result reaches its +// own caller. Tags end in a request number, so a late result never answers +// a newer request. type ExportTag = "thumbnail" | "validation" export function DiagramProvider({ children }: { children: React.ReactNode }) { @@ -63,11 +66,10 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) { const [showSaveDialog, setShowSaveDialog] = useState(false) const hasCalledOnLoadRef = useRef(false) const drawioRef = useRef(null) - const resolverRef = useRef<((value: string) => void) | null>(null) - // Pending thumbnail and validation PNG exports, keyed by their export tag - const taggedResolversRef = useRef void>>( - {}, - ) + // Pending exports, keyed by their export tag + const exportResolversRef = useRef< + Record void> + >({}) // Pending history exports: the document each one was asked for const historyXmlRef = useRef(new Map()) const exportSeqRef = useRef(0) @@ -97,32 +99,28 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) { setChartXML(xml) } - // Track if we're expecting an export for file save (stores raw export data) - const saveResolverRef = useRef<{ - resolver: ((data: string, fullDiagramXML?: string) => void) | null - format: ExportFormat | null - }>({ resolver: null, format: null }) - const handleExport = () => { - if (drawioRef.current) { - // Save this export to history, with the document shown now: - // chartXML can change before the result comes back - const tag = `history-${++exportSeqRef.current}` - historyXmlRef.current.set(tag, chartXMLRef.current) - drawioRef.current.exportDiagram({ - format: "xmlsvg", - message: tag, - }) - } + if (!drawioRef.current) return "" + // Save this export to history, with the document shown now: + // chartXML can change before the result comes back + const tag = `history-${++exportSeqRef.current}` + historyXmlRef.current.set(tag, chartXMLRef.current) + drawioRef.current.exportDiagram({ + format: "xmlsvg", + message: tag, + }) + return tag } const handleExportWithoutHistory = () => { - if (drawioRef.current) { - // Export without saving to history (for edit_diagram fetching current state) - drawioRef.current.exportDiagram({ - format: "xmlsvg", - }) - } + if (!drawioRef.current) return "" + // Export without saving to history (for edit_diagram fetching current state) + const tag = `fetch-${++exportSeqRef.current}` + drawioRef.current.exportDiagram({ + format: "xmlsvg", + message: tag, + }) + return tag } // Export with a tag in `message` (draw.io echoes it back in the export @@ -137,11 +135,11 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) { const id = `${tag}-${++exportSeqRef.current}` const finish = (value: string | null) => { clearTimeout(timer) - delete taggedResolversRef.current[id] + delete exportResolversRef.current[id] resolve(value) } const timer = setTimeout(() => finish(null), timeoutMs) - taggedResolversRef.current[id] = finish + exportResolversRef.current[id] = finish drawioRef.current?.exportDiagram({ format, message: id }) }) @@ -213,16 +211,11 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) { } const handleDiagramExport = (data: EventExport) => { - // Tagged exports (thumbnail, validation PNG, file save) go only to - // their own caller, so they never take the result meant for resolverRef + // Thumbnail, validation PNG and file save exports go only to their + // own caller const tag = data.message?.message - if (/^(thumbnail|validation)-/.test(tag ?? "")) { - taggedResolversRef.current[tag as string]?.(data.data) - return - } - if (tag === "save") { - saveResolverRef.current.resolver?.(data.data, data.xml) - saveResolverRef.current = { resolver: null, format: null } + if (/^(thumbnail|validation|save)-/.test(tag ?? "")) { + exportResolversRef.current[tag as string]?.(data.data, data.xml) return } @@ -256,9 +249,12 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) { }) } - if (resolverRef.current) { - resolverRef.current(extractedXML) - resolverRef.current = null + // The chat's own export (onFetchChart), not another one in flight + const resolve = + tag !== undefined ? exportResolversRef.current[tag] : undefined + if (resolve) { + delete exportResolversRef.current[tag as string] + resolve(extractedXML) } } @@ -297,85 +293,87 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) { const drawioFormat = format === "drawio" || format === "xmlsvg" ? "xmlsvg" : format - // Set up the resolver before triggering export - saveResolverRef.current = { - resolver: (exportData: string, fullDiagramXML?: string) => { - let fileContent: string | Blob - let mimeType: string - let extension: string + // Each save has its own tag, so two at once never swap results + const tag = `save-${++exportSeqRef.current}` + exportResolversRef.current[tag] = ( + exportData: string, + fullDiagramXML?: string, + ) => { + delete exportResolversRef.current[tag] + let fileContent: string | Blob + let mimeType: string + let extension: string - if (format === "drawio") { - // Prefer the complete document from the export event so all pages are saved. - const xml = fullDiagramXML?.trim() - ? fullDiagramXML - : extractDiagramXML(exportData) - fileContent = - normalizeToMxfile(xml, { - pageId: "page-1", - pageName: "Page-1", - }) ?? xml - mimeType = "application/xml" - extension = ".drawio" - } else if (format === "png") { - // PNG data comes as base64 data URL - fileContent = exportData - mimeType = "image/png" - extension = ".png" - } else if (format === "xmlsvg") { - // Editable SVG: pass data URL directly (like PNG) - fileContent = exportData - mimeType = "image/svg+xml" - extension = ".drawio.svg" - } else { - // SVG format (view-only) - fileContent = exportData - mimeType = "image/svg+xml" - extension = ".svg" - } + if (format === "drawio") { + // Prefer the complete document from the export event so all pages are saved. + const xml = fullDiagramXML?.trim() + ? fullDiagramXML + : extractDiagramXML(exportData) + fileContent = + normalizeToMxfile(xml, { + pageId: "page-1", + pageName: "Page-1", + }) ?? xml + mimeType = "application/xml" + extension = ".drawio" + } else if (format === "png") { + // PNG data comes as base64 data URL + fileContent = exportData + mimeType = "image/png" + extension = ".png" + } else if (format === "xmlsvg") { + // Editable SVG: pass data URL directly (like PNG) + fileContent = exportData + mimeType = "image/svg+xml" + extension = ".drawio.svg" + } else { + // SVG format (view-only) + fileContent = exportData + mimeType = "image/svg+xml" + extension = ".svg" + } - // Log save event to Langfuse (flags the trace) - logSaveToLangfuse(filename, format, sessionId) + // Log save event to Langfuse (flags the trace) + logSaveToLangfuse(filename, format, sessionId) - // Handle download - let url: string - if ( - typeof fileContent === "string" && - fileContent.startsWith("data:") - ) { - // Already a data URL (PNG) - url = fileContent - } else { - const blob = new Blob([fileContent], { type: mimeType }) - url = URL.createObjectURL(blob) - } + // Handle download + let url: string + if ( + typeof fileContent === "string" && + fileContent.startsWith("data:") + ) { + // Already a data URL (PNG) + url = fileContent + } else { + const blob = new Blob([fileContent], { type: mimeType }) + url = URL.createObjectURL(blob) + } - const a = document.createElement("a") - a.href = url - a.download = `${filename}${extension}` - document.body.appendChild(a) - a.click() - document.body.removeChild(a) + const a = document.createElement("a") + a.href = url + a.download = `${filename}${extension}` + document.body.appendChild(a) + a.click() + document.body.removeChild(a) - // Show success toast after download is initiated - if (successMessage) { - toast.success(successMessage, { - position: "bottom-left", - duration: 2500, - }) - } + // Show success toast after download is initiated + if (successMessage) { + toast.success(successMessage, { + position: "bottom-left", + duration: 2500, + }) + } - // Delay URL revocation to ensure download completes - if (!url.startsWith("data:")) { - setTimeout(() => URL.revokeObjectURL(url), 100) - } - }, - format, + // Delay URL revocation to ensure download completes + if (!url.startsWith("data:")) { + setTimeout(() => URL.revokeObjectURL(url), 100) + } } // Export diagram - callback will be handled in handleDiagramExport drawioRef.current.exportDiagram({ format: drawioFormat, - message: "save", + message: tag, }) } @@ -407,7 +405,7 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) { loadDiagram, handleExport, handleExportWithoutHistory, - resolverRef, + exportResolversRef, drawioRef, handleDiagramExport, handleDiagramAutoSave, diff --git a/electron/main/app-menu.ts b/electron/main/app-menu.ts index 6d6d073f..ae129559 100644 --- a/electron/main/app-menu.ts +++ b/electron/main/app-menu.ts @@ -32,6 +32,9 @@ export function rebuildAppMenu(): void { buildAppMenu() } +// Number of the latest preset switch +let lastSwitch = 0 + /** * Apply a preset and restart the server so it takes effect. * If the restart fails, go back to the previous preset and restart again, @@ -41,6 +44,7 @@ export function rebuildAppMenu(): void { export async function switchPreset( id: string, ): Promise> { + const switchNumber = ++lastSwitch const previousPresetId = getCurrentPresetId() const env = applyPresetToEnv(id) if (!env) { @@ -60,8 +64,9 @@ export async function switchPreset( console.error("Failed to restart server:", error) const reason = error instanceof Error ? error.message : String(error) - // Another preset was chosen meanwhile: its own restart follows - if (getCurrentPresetId() !== id) { + // A newer switch started meanwhile (also of this same preset): its + // own restart follows, and undoing would lose that choice + if (switchNumber !== lastSwitch) { throw new Error( `The server could not be restarted.\n\nError: ${reason}`, ) diff --git a/electron/main/env-loader.ts b/electron/main/env-loader.ts index ddf59085..bb3ff9da 100644 --- a/electron/main/env-loader.ts +++ b/electron/main/env-loader.ts @@ -28,6 +28,20 @@ export function loadEnvFile(): void { console.log("No .env file found, using system environment variables") } +/** + * Index of the quote that closes a value starting with a quote, or -1. A + * backslash before the quote character escapes it, as in dotenv; the + * backslash stays in the value. + */ +function findClosingQuote(value: string): number { + const quote = value[0] + for (let i = 1; i < value.length; i++) { + if (value[i] === "\\" && value[i + 1] === quote) i++ + else if (value[i] === quote) return i + } + return -1 +} + /** * Parse and load environment variables from a file */ @@ -50,7 +64,7 @@ function loadEnvFromFile(filePath: string): void { const quote = value[0] const closingQuote = - quote === '"' || quote === "'" ? value.indexOf(quote, 1) : -1 + quote === '"' || quote === "'" ? findClosingQuote(value) : -1 if ( closingQuote > 0 && /^\s*(#.*)?$/.test(value.slice(closingQuote + 1)) diff --git a/hooks/use-diagram-tool-handlers.ts b/hooks/use-diagram-tool-handlers.ts index 66dbd7f3..55313b85 100644 --- a/hooks/use-diagram-tool-handlers.ts +++ b/hooks/use-diagram-tool-handlers.ts @@ -64,6 +64,9 @@ interface UseDiagramToolHandlersParams { validateDiagram?: ValidateDiagramFn enableVlmValidation?: boolean sessionId?: string + // The user pressed Stop: a screenshot check that has not started is + // skipped (one already running is cancelled by the caller) + isStopped?: () => boolean onValidationStateChange?: ( toolCallId: string, state: ValidationState, @@ -90,6 +93,7 @@ export function useDiagramToolHandlers({ validateDiagram, enableVlmValidation = true, sessionId, + isStopped, onValidationStateChange, }: UseDiagramToolHandlersParams) { // Helper to update validation state @@ -257,7 +261,11 @@ ${finalXml} await new Promise((resolve) => setTimeout(resolve, 100)) capturedPngData = await captureValidationPng() - if (capturedPngData) { + // Stopped while the screenshot was taken: no check. The + // chat waits for this handler, so it must end now. + if (isStopped?.()) { + updateValidationState(toolCall.toolCallId, "skipped") + } else if (capturedPngData) { if (DEBUG) { console.log( "[display_diagram] Captured PNG for validation", @@ -363,6 +371,16 @@ ${finalXml} updateValidationState(toolCall.toolCallId, "skipped") } } catch (error) { + // Cancelled by Stop: the diagram stays, unchecked + if ((error as Error)?.name === "AbortError") { + updateValidationState(toolCall.toolCallId, "skipped") + addToolOutput({ + tool: "display_diagram", + toolCallId: toolCall.toolCallId, + output: "Successfully displayed the diagram.", + }) + return + } // VLM validation error - log but don't block the user console.warn( "[display_diagram] VLM validation error:", diff --git a/hooks/use-session-manager.ts b/hooks/use-session-manager.ts index a1b63780..b0cfe308 100644 --- a/hooks/use-session-manager.ts +++ b/hooks/use-session-manager.ts @@ -13,10 +13,12 @@ import { getSession, isIndexedDBAvailable, migrateFromLocalStorage, + readSessionCount, type SessionMetadata, type StoredMessage, saveSession, } from "@/lib/session-storage" +import { STORAGE_KEYS } from "@/lib/storage" export interface SessionData { messages: StoredMessage[] @@ -37,14 +39,17 @@ export interface UseSessionManagerReturn { // Actions switchSession: (id: string) => Promise deleteSession: (id: string) => Promise<{ wasCurrentSession: boolean }> - // forSessionId: optional session ID to verify save targets correct session (prevents stale debounce writes) + // chatGeneration: getChatGeneration() when the save was scheduled (by + // default, now); the save is dropped if another chat is on screen when + // its turn comes // Resolves to false when the save failed (the user was told) saveCurrentSession: ( data: SessionData, - forSessionId?: string | null, + chatGeneration?: number, ) => Promise refreshSessions: () => Promise clearCurrentSession: () => void + getChatGeneration: () => number } // Reading the session list loads every stored session in full, and window @@ -79,6 +84,20 @@ export function useSessionManager( const isInitializedRef = useRef(false) // Sequence guard for URL changes - prevents out-of-order async resolution const urlChangeSequenceRef = useRef(0) + // The chat on screen, read by saves that run after a render or a wait + const currentSessionRef = useRef(null) + // Goes up each time another chat is put on screen (creating the + // session of the chat on screen does not count) + const chatGenerationRef = useRef(0) + // Saves run one at a time, so two saves of a new chat create it once + const saveQueueRef = useRef>(Promise.resolve()) + + const changeChat = useCallback((session: ChatSession | null) => { + chatGenerationRef.current++ + currentSessionRef.current = session + setCurrentSession(session) + setCurrentSessionId(session?.id ?? null) + }, []) // Load sessions list const refreshSessions = useCallback(async () => { @@ -115,18 +134,32 @@ export function useSessionManager( const metadata = await getAllSessionMetadata() setSessions(metadata) // The desktop app may try its other port next launch, where - // an older version may have saved the chats - window.electronAPI - ?.chatsLoaded?.(metadata.length) - .catch(() => {}) + // an older version may have saved the chats: only when this + // origin surely has none (a failed read is not "none") and + // keeps no model settings or keys either + if (window.electronAPI?.chatsLoaded) { + const count = await readSessionCount() + // The app saves an empty config on its first load; the + // providers are what holds the keys + let hasSettings = true + try { + const config = JSON.parse( + localStorage.getItem(STORAGE_KEYS.modelConfigs) ?? + "{}", + ) + hasSettings = (config.providers?.length ?? 0) > 0 + } catch { + // Unreadable: treat as settings, and stay + } + if (count !== null && !hasSettings) { + window.electronAPI.chatsLoaded(count).catch(() => {}) + } + } // Only load a session if initialSessionId is provided (from URL param) if (initialSessionId) { const session = await getSession(initialSessionId) - if (session) { - setCurrentSession(session) - setCurrentSessionId(session.id) - } + if (session) changeChat(session) // If session not found, stay in blank state (URL has invalid session ID) } // If no initialSessionId, start with blank state (no auto-restore) @@ -138,7 +171,7 @@ export function useSessionManager( } init() - }, [initialSessionId]) + }, [initialSessionId, changeChat]) // Handle URL session ID changes after initialization // Note: intentionally NOT including currentSessionId in deps to avoid race conditions @@ -153,6 +186,7 @@ export function useSessionManager( async function handleSessionIdChange() { if (initialSessionId) { + const generation = chatGenerationRef.current // URL has session ID - load it const session = await getSession(initialSessionId) @@ -161,16 +195,13 @@ export function useSessionManager( if (currentSequence !== urlChangeSequenceRef.current) { return } + // Another chat was put on screen meanwhile (New Chat right + // after this one got its session id in the URL): keep it + if (generation !== chatGenerationRef.current) return - if (session) { - // Only update if the session is different from current - setCurrentSessionId((current) => { - if (current !== session.id) { - setCurrentSession(session) - return session.id - } - return current - }) + // Only update if the session is different from current + if (session && currentSessionRef.current?.id !== session.id) { + changeChat(session) } } // Removed: else clause that clears session @@ -179,7 +210,7 @@ export function useSessionManager( } handleSessionIdChange() - }, [initialSessionId, isAvailable]) + }, [initialSessionId, isAvailable, changeChat]) // Refresh sessions on window focus (multi-tab sync), at most once per interval const lastFocusRefreshRef = useRef(0) @@ -201,9 +232,11 @@ export function useSessionManager( async (id: string): Promise => { if (id === currentSessionId) return null - // Save current session first if it has messages - if (currentSession && currentSession.messages.length > 0) { - await saveSession(currentSession) + // Save current session first if it has messages (as saved + // last: the caller may have just saved it) + const current = currentSessionRef.current + if (current && current.messages.length > 0) { + await saveSession(current) } // Load the target session @@ -213,9 +246,7 @@ export function useSessionManager( return null } - // Update state - setCurrentSession(session) - setCurrentSessionId(session.id) + changeChat(session) return { messages: session.messages, @@ -225,7 +256,7 @@ export function useSessionManager( diagramHistory: session.diagramHistory, } }, - [currentSessionId, currentSession], + [currentSessionId, changeChat], ) // Delete a session @@ -235,112 +266,121 @@ export function useSessionManager( await deleteSessionFromDB(id) // If deleting current session, clear state (caller will show new empty session) - if (wasCurrentSession) { - setCurrentSession(null) - setCurrentSessionId(null) - } + if (wasCurrentSession) changeChat(null) await refreshSessions() return { wasCurrentSession } }, - [currentSessionId, refreshSessions], + [currentSessionId, refreshSessions, changeChat], ) // Save current session data (debounced externally by caller) - // forSessionId: if provided, verify save targets correct session (prevents stale debounce writes) const saveCurrentSession = useCallback( - async ( - data: SessionData, - forSessionId?: string | null, - ): Promise => { - // If forSessionId is provided, verify it matches current session - // This prevents stale debounced saves from overwriting a newly switched session - if ( - forSessionId !== undefined && - forSessionId !== currentSessionId - ) { - return true - } - // Nothing can be stored without IndexedDB - if (!isIndexedDBAvailable()) return true + (data: SessionData, chatGeneration?: number): Promise => { + // The data is of the chat on screen when the save was asked for + const generation = chatGeneration ?? chatGenerationRef.current + const run = async (): Promise => { + // That chat is no longer on screen (leaving it saved it) + if (generation !== chatGenerationRef.current) return true + // Nothing can be stored without IndexedDB + if (!isIndexedDBAvailable()) return true + // The user may put another chat on screen while this one is + // written; the stored copy is still right, the state is not + const stillOnScreen = () => + chatGenerationRef.current === generation + const currentSession = currentSessionRef.current - if (!currentSession) { - // Create a new session if none exists - const newSession: ChatSession = { - ...createEmptySession(), + if (!currentSession) { + // Create a new session if none exists + const newSession: ChatSession = { + ...createEmptySession(), + messages: data.messages, + xmlSnapshots: data.xmlSnapshots, + diagramXml: data.diagramXml, + thumbnailDataUrl: data.thumbnailDataUrl, + diagramHistory: data.diagramHistory, + title: extractTitle(data.messages), + } + // Without a stored session, keep no session id (it would end + // up in the URL and point to nothing after a reload) + if (!(await saveSession(newSession))) { + notifySaveFailed(dict.errors.sessionSaveFailed) + return false + } + await enforceSessionLimit() + if (stillOnScreen()) { + currentSessionRef.current = newSession + setCurrentSession(newSession) + setCurrentSessionId(newSession.id) + } + await refreshSessions() + return true + } + + // Update existing session + const updatedSession: ChatSession = { + ...currentSession, messages: data.messages, xmlSnapshots: data.xmlSnapshots, diagramXml: data.diagramXml, - thumbnailDataUrl: data.thumbnailDataUrl, - diagramHistory: data.diagramHistory, - title: extractTitle(data.messages), + thumbnailDataUrl: + data.thumbnailDataUrl ?? + currentSession.thumbnailDataUrl, + diagramHistory: + data.diagramHistory ?? currentSession.diagramHistory, + updatedAt: Date.now(), + // Update title if it's still default and we have messages + title: + currentSession.title === "New Chat" && + data.messages.length > 0 + ? extractTitle(data.messages) + : currentSession.title, } - // Without a stored session, keep no session id (it would end - // up in the URL and point to nothing after a reload) - if (!(await saveSession(newSession))) { + + if (!(await saveSession(updatedSession))) { notifySaveFailed(dict.errors.sessionSaveFailed) return false } - await enforceSessionLimit() - setCurrentSession(newSession) - setCurrentSessionId(newSession.id) - await refreshSessions() + if (stillOnScreen()) { + currentSessionRef.current = updatedSession + setCurrentSession(updatedSession) + } + + // Update sessions list metadata + setSessions((prev) => + prev.map((s) => + s.id === updatedSession.id + ? { + ...s, + title: updatedSession.title, + updatedAt: updatedSession.updatedAt, + messageCount: updatedSession.messages.length, + hasDiagram: + !!updatedSession.diagramXml && + updatedSession.diagramXml.trim().length > + 0, + thumbnailDataUrl: + updatedSession.thumbnailDataUrl, + } + : s, + ), + ) return true } - - // Update existing session - const updatedSession: ChatSession = { - ...currentSession, - messages: data.messages, - xmlSnapshots: data.xmlSnapshots, - diagramXml: data.diagramXml, - thumbnailDataUrl: - data.thumbnailDataUrl ?? currentSession.thumbnailDataUrl, - diagramHistory: - data.diagramHistory ?? currentSession.diagramHistory, - updatedAt: Date.now(), - // Update title if it's still default and we have messages - title: - currentSession.title === "New Chat" && - data.messages.length > 0 - ? extractTitle(data.messages) - : currentSession.title, - } - - if (!(await saveSession(updatedSession))) { - notifySaveFailed(dict.errors.sessionSaveFailed) - return false - } - setCurrentSession(updatedSession) - - // Update sessions list metadata - setSessions((prev) => - prev.map((s) => - s.id === updatedSession.id - ? { - ...s, - title: updatedSession.title, - updatedAt: updatedSession.updatedAt, - messageCount: updatedSession.messages.length, - hasDiagram: - !!updatedSession.diagramXml && - updatedSession.diagramXml.trim().length > 0, - thumbnailDataUrl: updatedSession.thumbnailDataUrl, - } - : s, - ), - ) - return true + const result = saveQueueRef.current.then(run) + saveQueueRef.current = result.catch(() => {}) + return result }, - [currentSession, currentSessionId, refreshSessions, dict], + [refreshSessions, dict], ) // Clear current session state (for starting fresh without loading another session) const clearCurrentSession = useCallback(() => { - setCurrentSession(null) - setCurrentSessionId(null) - }, []) + changeChat(null) + }, [changeChat]) + + const getChatGeneration = useCallback(() => chatGenerationRef.current, []) return { sessions, @@ -353,5 +393,6 @@ export function useSessionManager( saveCurrentSession, refreshSessions, clearCurrentSession, + getChatGeneration, } } diff --git a/hooks/use-validate-diagram.ts b/hooks/use-validate-diagram.ts index c4593c37..d3402400 100644 --- a/hooks/use-validate-diagram.ts +++ b/hooks/use-validate-diagram.ts @@ -103,9 +103,22 @@ export function useValidateDiagram(options: UseValidateDiagramOptions = {}) { [submit], ) + /** + * End a running check (the user pressed Stop): its promise rejects with + * an AbortError, so the tool handler can finish at once. + */ + const cancel = useCallback(() => { + const pending = pendingValidationRef.current + if (!pending) return + pendingValidationRef.current = null + stop() + pending.reject(new DOMException("Validation cancelled", "AbortError")) + }, [stop]) + /** * Validate with fallback - returns default valid result on error. * Use this to avoid blocking the user on validation failures. + * A cancelled check is passed on as its AbortError. */ const validateWithFallback = useCallback( async ( @@ -115,6 +128,7 @@ export function useValidateDiagram(options: UseValidateDiagramOptions = {}) { try { return await validate(imageData, sessionId) } catch (error) { + if ((error as Error)?.name === "AbortError") throw error console.warn( "[useValidateDiagram] Validation failed, using fallback:", error, @@ -130,6 +144,7 @@ export function useValidateDiagram(options: UseValidateDiagramOptions = {}) { validate, validateWithFallback, stop, + cancel, // State isValidating: isLoading, diff --git a/lib/admin/providers.ts b/lib/admin/providers.ts index ff6ca8d2..39a48cae 100644 --- a/lib/admin/providers.ts +++ b/lib/admin/providers.ts @@ -257,12 +257,7 @@ export function deriveEnvUpdates( if (p.baseUrl) updates.GOOGLE_VERTEX_BASE_URL = p.baseUrl } else if (p.provider === "ollama") { if (p.apiKey) updates.OLLAMA_API_KEY = p.apiKey - // A key without a URL is an Ollama Cloud key, as its Test sends - // it; chat sends a server key to OLLAMA_BASE_URL or local Ollama - if (p.baseUrl || p.apiKey) { - updates.OLLAMA_BASE_URL = - p.baseUrl || PROVIDER_INFO.ollama.defaultBaseUrl || null - } + if (p.baseUrl) updates.OLLAMA_BASE_URL = p.baseUrl } else { const env = credEnvNames(p.provider, index) if (env.key && p.apiKey) updates[env.key] = p.apiKey diff --git a/lib/ai-providers.ts b/lib/ai-providers.ts index 8012dc77..25529f51 100644 --- a/lib/ai-providers.ts +++ b/lib/ai-providers.ts @@ -21,6 +21,7 @@ import { adminProvidersToConfig, loadAdminProviders, } from "@/lib/admin/providers" +import { getApiEndpoint } from "@/lib/base-path" import { redirectGuardedFetch } from "@/lib/ssrf-protection" import { normalizeBaseUrl, @@ -100,6 +101,9 @@ export interface ClientOverrides { awsSessionToken?: string | null // Vertex AI config vertexApiKey?: string | null // Express Mode API key + // baseUrl is the server's own

_BASE_URL (the admin panel's Test), + // not one a user chose: no redirect guard + trustedBaseUrl?: boolean // Custom headers (e.g., for EdgeOne cookie auth) headers?: Record // Custom env var name(s) for server models @@ -569,6 +573,7 @@ function detectProvider(): ProviderName | null { function validateProviderCredentials( provider: ProviderName, customApiKeyEnv?: string | string[], + customBaseUrlEnv?: string, ): void { // Handle array of env var names - at least one must be set if (Array.isArray(customApiKeyEnv)) { @@ -604,9 +609,12 @@ function validateProviderCredentials( } } - // Azure requires either AZURE_BASE_URL or AZURE_RESOURCE_NAME in addition to API key + // Azure requires either AZURE_BASE_URL or AZURE_RESOURCE_NAME in addition + // to API key, or a server model's own URL variable (an admin panel entry) if (provider === "azure") { - const hasBaseUrl = !!process.env.AZURE_BASE_URL + const hasBaseUrl = + !!process.env.AZURE_BASE_URL || + !!(customBaseUrlEnv && process.env[customBaseUrlEnv]) const hasResourceName = !!process.env.AZURE_RESOURCE_NAME if (!hasBaseUrl && !hasResourceName) { throw new Error( @@ -879,13 +887,20 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { // Only validate server credentials if client isn't providing their own API key if (!isClientOverride) { - validateProviderCredentials(provider, overrides?.apiKeyEnv) + validateProviderCredentials( + provider, + overrides?.apiKeyEnv, + overrides?.baseUrlEnv, + ) } console.log(`[AI Provider] Initializing ${provider} with model: ${modelId}`) // Requests to a base URL the client chose must not follow redirects - const guardedFetch = overrides?.baseUrl ? redirectGuardedFetch() : undefined + const guardedFetch = + overrides?.baseUrl && !overrides.trustedBaseUrl + ? redirectGuardedFetch() + : undefined // Build provider-specific options from environment variables let providerOptions = buildProviderOptions(provider, modelId) let model: LanguageModel @@ -1091,20 +1106,30 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { return { model, providerOptions, modelId, provider } } +/** + * The deployment's EdgeOne Pages function, as an absolute URL (the SDK + * needs one), under the deployment's base path + */ +export function edgeOneEndpoint(req: Request): string { + const origin = req.headers.get("origin") || new URL(req.url).origin + return `${origin}${getApiEndpoint("/api/edgeai")}` +} + /** * The server's

_BASE_URL for a provider, which getAIModel uses for a * server model without a URL variable of its own (an admin panel entry - * without a URL). Bedrock, EdgeOne and Ollama (the panel writes - * OLLAMA_BASE_URL itself) have none. + * without a URL). None for Bedrock and EdgeOne, and none for Ollama and + * Vertex AI, whose variables the panel writes itself (before a save they + * still hold the entry's previous URL). */ export function globalBaseUrl(provider: ProviderName): string | undefined { - if (["bedrock", "edgeone", "ollama"].includes(provider)) return undefined + if (["bedrock", "edgeone", "ollama", "vertexai"].includes(provider)) { + return undefined + } const name = - provider === "vertexai" - ? "GOOGLE_VERTEX_BASE_URL" - : provider === "gateway" - ? "AI_GATEWAY_BASE_URL" - : `${provider.toUpperCase()}_BASE_URL` + provider === "gateway" + ? "AI_GATEWAY_BASE_URL" + : `${provider.toUpperCase()}_BASE_URL` return process.env[name] || undefined } diff --git a/lib/dynamo-quota-manager.ts b/lib/dynamo-quota-manager.ts index 981868b4..f423c8ed 100644 --- a/lib/dynamo-quota-manager.ts +++ b/lib/dynamo-quota-manager.ts @@ -64,10 +64,13 @@ interface QuotaCheckResult { * Check all quotas and increment request count atomically. * Uses composite key (PK=user, SK=date) for per-day tracking. * Each day automatically gets a new item - no explicit reset needed. + * A request limit of 0 means none; increment 0 checks the limits without + * counting a request (the screenshot check). */ export async function checkAndIncrementRequest( ip: string, limits: QuotaLimits, + increment = 1, ): Promise { // Skip if quota tracking not enabled if (!client || !TABLE) { @@ -99,7 +102,7 @@ export async function checkAndIncrementRequest( attribute_not_exists(tpmCount) OR tpmCount < :tpmLimit) `, ExpressionAttributeValues: { - ":one": { N: "1" }, + ":one": { N: String(increment) }, ":minute": { S: currentMinute }, ":reqLimit": { N: String(limits.requests || 999999) }, ":tokenLimit": { N: String(limits.tokens || 999999) }, diff --git a/lib/i18n/dictionaries/en.json b/lib/i18n/dictionaries/en.json index d1b7c41d..fa5ee712 100644 --- a/lib/i18n/dictionaries/en.json +++ b/lib/i18n/dictionaries/en.json @@ -187,6 +187,8 @@ "failedToRecordFeedback": "Failed to record your feedback. Please try again.", "storageUpdateFailed": "Chat cleared but browser storage could not be updated", "sessionSaveFailed": "Could not save this chat. Browser storage may be full: delete old chats from history and try again.", + "sessionSaveFailedLeave": "Could not save this chat. Browser storage may be full. You can go on without saving it, then delete old chats from the list in the new chat.", + "continueWithoutSaving": "Continue without saving", "llm": { "invalid_api_key": "The provider rejected the API key. Check it in model settings.", "forbidden": "The provider refused the request. The key may not have access to this model or region.", diff --git a/lib/i18n/dictionaries/ja.json b/lib/i18n/dictionaries/ja.json index 87861452..97c5c754 100644 --- a/lib/i18n/dictionaries/ja.json +++ b/lib/i18n/dictionaries/ja.json @@ -187,6 +187,8 @@ "failedToRecordFeedback": "フィードバックの記録に失敗しました。もう一度お試しください。", "storageUpdateFailed": "チャットはクリアされましたが、ブラウザストレージを更新できませんでした", "sessionSaveFailed": "このチャットを保存できませんでした。ブラウザのストレージがいっぱいの可能性があります。履歴から古いチャットを削除して、もう一度お試しください。", + "sessionSaveFailedLeave": "このチャットを保存できませんでした。ブラウザのストレージがいっぱいの可能性があります。保存せずに続けて、新しいチャットの一覧から古いチャットを削除できます。", + "continueWithoutSaving": "保存せずに続ける", "llm": { "invalid_api_key": "プロバイダーが API キーを拒否しました。モデル設定で確認してください。", "forbidden": "プロバイダーがリクエストを拒否しました。このキーにはこのモデルまたはリージョンの利用権限がない可能性があります。", diff --git a/lib/i18n/dictionaries/zh-Hant.json b/lib/i18n/dictionaries/zh-Hant.json index cd6d8bb9..d3521c80 100644 --- a/lib/i18n/dictionaries/zh-Hant.json +++ b/lib/i18n/dictionaries/zh-Hant.json @@ -187,6 +187,8 @@ "failedToRecordFeedback": "記錄您的回饋失敗。請重試。", "storageUpdateFailed": "聊天已清除,但無法更新瀏覽器儲存空間", "sessionSaveFailed": "無法儲存這個對話。瀏覽器儲存空間可能已滿,請在歷史紀錄裡刪除舊對話後重試。", + "sessionSaveFailedLeave": "無法儲存這個對話,瀏覽器儲存空間可能已滿。可以不儲存它、直接繼續,再在新對話的列表裡刪除舊對話。", + "continueWithoutSaving": "不儲存,繼續", "llm": { "invalid_api_key": "服務商拒絕了這個 API Key,請在模型設定中檢查。", "forbidden": "服務商拒絕了這次請求。這個 Key 可能沒有使用該模型或該地區的權限。", diff --git a/lib/i18n/dictionaries/zh.json b/lib/i18n/dictionaries/zh.json index 2c048fa4..1d8f9adb 100644 --- a/lib/i18n/dictionaries/zh.json +++ b/lib/i18n/dictionaries/zh.json @@ -187,6 +187,8 @@ "failedToRecordFeedback": "记录您的反馈失败。请重试。", "storageUpdateFailed": "聊天已清除,但无法更新浏览器存储", "sessionSaveFailed": "无法保存这个对话。浏览器存储空间可能已满,请在历史记录里删除旧对话后重试。", + "sessionSaveFailedLeave": "无法保存这个对话,浏览器存储空间可能已满。可以不保存它、直接继续,再在新对话的列表里删除旧对话。", + "continueWithoutSaving": "不保存,继续", "llm": { "invalid_api_key": "服务商拒绝了这个 API Key,请在模型设置里检查。", "forbidden": "服务商拒绝了这次请求。这个 Key 可能没有使用该模型或该地区的权限。", diff --git a/lib/session-storage.ts b/lib/session-storage.ts index 71b4acaa..33098367 100644 --- a/lib/session-storage.ts +++ b/lib/session-storage.ts @@ -199,13 +199,18 @@ export async function deleteSession(id: string): Promise { } export async function getSessionCount(): Promise { - if (!isIndexedDBAvailable()) return 0 + return (await readSessionCount()) ?? 0 +} + +/** The number of saved chats, or null when it could not be read */ +export async function readSessionCount(): Promise { + if (!isIndexedDBAvailable()) return null try { const db = await getDB() return await db.count(STORE_NAME) } catch (error) { console.error("Failed to get session count:", error) - return 0 + return null } } diff --git a/lib/ssrf-protection.ts b/lib/ssrf-protection.ts index 6b3e433b..80357cf3 100644 --- a/lib/ssrf-protection.ts +++ b/lib/ssrf-protection.ts @@ -118,24 +118,52 @@ export function allowPrivateUrls(): boolean { /** A redirect the guard below refused; its text is safe to show */ export class RedirectRefusedError extends Error { - constructor() { - super("Redirects are not allowed for custom base URLs") + constructor(message = "Redirects are not allowed for custom base URLs") { + super(message) this.name = "RedirectRefusedError" } } +const MAX_REDIRECTS = 5 + /** * A fetch for requests to a base URL the client chose. With private URLs * blocked, a public URL could still redirect the request to an internal - * host, so redirects are refused. Undefined when private URLs are allowed. + * host, so redirects are refused. With private URLs allowed but the quota + * on (DYNAMODB_QUOTA_TABLE), a request to a private address counts as the + * server's: redirects are followed only to public addresses, or a public + * URL could reach the server's own network uncounted. Undefined otherwise. */ export function redirectGuardedFetch(): typeof fetch | undefined { - if (allowPrivateUrls()) return undefined + const blockAll = !allowPrivateUrls() + if (!blockAll && !process.env.DYNAMODB_QUOTA_TABLE) return undefined return async (input, init) => { - const response = await fetch(input, { ...init, redirect: "manual" }) - if (response.status >= 300 && response.status < 400) { - throw new RedirectRefusedError() + let url = input instanceof Request ? input.url : String(input) + let next = init + for (let hop = 0; hop <= MAX_REDIRECTS; hop++) { + const response = await fetch(url, { ...next, redirect: "manual" }) + const location = response.headers.get("location") + if (response.status < 300 || response.status >= 400 || !location) { + return response + } + if (blockAll) throw new RedirectRefusedError() + url = new URL(location, url).toString() + if (await isPrivateUrl(url)) { + throw new RedirectRefusedError( + "Redirects to private addresses are not allowed", + ) + } + // As fetch itself does: 303, and 301 or 302 after a POST, go on + // as a GET without the body + const method = (next?.method ?? "GET").toUpperCase() + if ( + response.status === 303 || + ((response.status === 301 || response.status === 302) && + method === "POST") + ) { + next = { ...next, method: "GET", body: undefined } + } } - return response + throw new RedirectRefusedError("Too many redirects") } } diff --git a/lib/utils.ts b/lib/utils.ts index ad72509f..3f648a1b 100644 --- a/lib/utils.ts +++ b/lib/utils.ts @@ -1,6 +1,7 @@ import { type ClassValue, clsx } from "clsx" import * as pako from "pako" import { twMerge } from "tailwind-merge" +import { hasCells } from "@/packages/mcp-server/src/pages.ts" export function cn(...inputs: ClassValue[]) { return twMerge(clsx(inputs)) @@ -17,12 +18,14 @@ export function cn(...inputs: ClassValue[]) { export const MIN_REAL_DIAGRAM_LENGTH = 300 /** - * Check if diagram XML represents a real diagram (not just empty template). + * Check if diagram XML represents a real diagram (not just empty template): + * it has a shape (however short), or is long enough to hold pages worth + * keeping. * @param xml - The diagram XML string to check * @returns true if the XML is a real diagram with content */ export function isRealDiagram(xml: string | undefined | null): boolean { - return !!xml && xml.length > MIN_REAL_DIAGRAM_LENGTH + return !!xml && (hasCells(xml) || xml.length > MIN_REAL_DIAGRAM_LENGTH) } // ============================================================================ diff --git a/packages/mcp-server/src/exclusive.ts b/packages/mcp-server/src/exclusive.ts new file mode 100644 index 00000000..364f1c15 --- /dev/null +++ b/packages/mcp-server/src/exclusive.ts @@ -0,0 +1,25 @@ +/** + * A queue for tool handlers: each call waits until the previous one ended. + * A tool call the client cancelled while it waited (the MCP SDK aborts its + * extra.signal, the handler's last argument) is skipped. + */ +export function createExclusive() { + let tail: Promise = Promise.resolve() + return function exclusive Promise>( + handler: T, + ): T { + return ((...args: unknown[]) => { + const extra = args.at(-1) as { signal?: AbortSignal } | undefined + const run = tail.then(() => + extra?.signal?.aborted + ? { + content: [{ type: "text", text: "Cancelled." }], + isError: true, + } + : handler(...args), + ) + tail = run.catch(() => {}) + return run + }) as T + } +} diff --git a/packages/mcp-server/src/history.ts b/packages/mcp-server/src/history.ts index 9b597cf3..253a223e 100644 --- a/packages/mcp-server/src/history.ts +++ b/packages/mcp-server/src/history.ts @@ -3,7 +3,6 @@ * Stores {xml, svg} entries in a circular buffer */ -import { contentFingerprint } from "./edit-gate.ts" import { log } from "./logger.ts" const MAX_HISTORY = 20 @@ -17,14 +16,6 @@ interface HistoryEntry { let nextEntryId = 0 const historyStore = new Map() -// The same pages and cells; a document without pages has an empty -// fingerprint and is compared as text only -function sameDiagram(a: string, b: string): boolean { - if (a === b) return true - const fingerprint = contentFingerprint(a) - return fingerprint !== "" && fingerprint === contentFingerprint(b) -} - export function addHistory(sessionId: string, xml: string, svg = ""): number { let history = historyStore.get(sessionId) if (!history) { @@ -32,10 +23,10 @@ export function addHistory(sessionId: string, xml: string, svg = ""): number { historyStore.set(sessionId, history) } - // Dedupe: skip if same as last entry, also when only re-serialized - // (the browser's copy of the same diagram) + // Dedupe: skip if same as last entry (a change of page settings or + // background only is a new version) const last = history[history.length - 1] - if (last && sameDiagram(last.xml, xml)) { + if (last && last.xml === xml) { if (svg && !last.svg) last.svg = svg return history.length - 1 } @@ -79,7 +70,7 @@ export function updateLastHistorySvg( const history = historyStore.get(sessionId) if (!history || history.length === 0) return false const last = history[history.length - 1] - if (!last.svg && sameDiagram(last.xml, shownXml)) { + if (!last.svg && last.xml === shownXml) { last.svg = svg return true } diff --git a/packages/mcp-server/src/http-server.ts b/packages/mcp-server/src/http-server.ts index fc3cac01..3b40cf70 100644 --- a/packages/mcp-server/src/http-server.ts +++ b/packages/mcp-server/src/http-server.ts @@ -20,17 +20,28 @@ function readBody( // across two chunks. const chunks: Buffer[] = [] let size = 0 + let tooLarge = false req.on("data", (chunk: Buffer) => { + if (tooLarge) return size += chunk.length if (size > MAX_BODY_BYTES) { - res.writeHead(413, { "Content-Type": "application/json" }) - res.end(JSON.stringify({ error: "Payload too large" })) - req.destroy() + // Read the rest without keeping it and answer at the end: a + // connection closed mid-upload reaches the browser as a network + // error, without this answer + tooLarge = true + chunks.length = 0 return } chunks.push(chunk) }) - req.on("end", () => cb(Buffer.concat(chunks).toString("utf8"))) + req.on("end", () => { + if (tooLarge) { + res.writeHead(413, { "Content-Type": "application/json" }) + res.end(JSON.stringify({ error: "Payload too large" })) + return + } + cb(Buffer.concat(chunks).toString("utf8")) + }) } import { contentFingerprint } from "./edit-gate.ts" @@ -125,6 +136,8 @@ interface SessionState { // Version of the last write the browser did not make itself (AI edit, // restore). A browser push based on an older version is rejected. serverVersion?: number + // The XML of that write: what a thumbnail taken after loading it shows + serverXml?: string lastUpdated: Date lastPolled?: number // Last browser poll; an open tab keeps the session alive svg?: string // Cached SVG from last browser save @@ -192,11 +205,15 @@ export function setState( version: newVersion, stateId: existing?.stateId ?? randomUUID(), serverVersion: fromBrowser ? existing?.serverVersion : newVersion, + serverXml: fromBrowser ? existing?.serverXml : xml, lastUpdated: new Date(), lastPolled: existing?.lastPolled, // The image of this XML, never an older one's: a write without an - // image (AI write, sync reply) leaves none until the browser sends it - svg: svg || undefined, + // image (AI write, sync reply) leaves none until the browser sends + // it, unless it is the same XML + svg: + svg || + (existing && existing.xml === xml ? existing.svg : undefined), syncRequested: undefined, // Clear sync request when browser pushes state exportFormat: existing?.exportFormat, // Preserve pending export request exportXml: existing?.exportXml, // Preserve pending projection @@ -696,18 +713,26 @@ function handleHistorySvgApi( readBody(req, res, (body) => { try { - const { sessionId, svg } = JSON.parse(body) + const { sessionId, svg, stateId, version } = JSON.parse(body) if (!sessionId || !svg) { res.writeHead(400, { "Content-Type": "application/json" }) res.end(JSON.stringify({ error: "sessionId and svg required" })) return } - // The browser took it of the diagram it just loaded: the state + // The browser took it of the server write it loaded, named by + // the state and version. One that arrives after the next server + // write, or for a state since lost, is dropped; a browser write + // since (a sync reply) leaves that write's image valid. const state = stateStore.get(sessionId) - if (state) { - updateLastHistorySvg(sessionId, svg, state.xml) - state.svg = svg + if ( + state && + state.stateId === stateId && + state.serverVersion === version && + state.serverXml !== undefined + ) { + updateLastHistorySvg(sessionId, svg, state.serverXml) + if (state.xml === state.serverXml) state.svg = svg } res.writeHead(200, { "Content-Type": "application/json" }) res.end(JSON.stringify({ success: true })) diff --git a/packages/mcp-server/src/index.ts b/packages/mcp-server/src/index.ts index 5d131a52..35d5e26e 100644 --- a/packages/mcp-server/src/index.ts +++ b/packages/mcp-server/src/index.ts @@ -28,6 +28,7 @@ import { installDomPolyfill } from "./dom.ts" import { DRAWING_GUIDE } from "./drawing-guide.ts" import { editDiagram, targetPageXml } from "./edit-diagram.ts" import { checkEditGate, markPageSeen } from "./edit-gate.ts" +import { createExclusive } from "./exclusive.ts" import { addHistory } from "./history.ts" import { type ExportFormat, @@ -142,6 +143,18 @@ const server = new McpServer( { instructions: INSTRUCTIONS }, ) +// The tools that write the diagram, and start_session, run one at a time: +// two writes at once would both build on the same document, and the second +// would drop the first one's change. start_session in the queue keeps a +// session switch from landing in the middle of a write. +const exclusive = createExclusive() +const registerWriteTool = ((name: string, config: any, handler: any) => + server.registerTool( + name, + config, + exclusive(handler), + )) as typeof server.registerTool + // Shared Zod schema fragment for page-targeting parameters. // Every multi-page-aware tool reuses these three optional fields so the LLM // learns one consistent interface. @@ -254,7 +267,7 @@ server.registerTool( ) // Tool: start_session -server.registerTool( +registerWriteTool( "start_session", { title: "Start session", @@ -312,7 +325,7 @@ server.registerTool( ) // Tool: create_new_diagram -server.registerTool( +registerWriteTool( "create_new_diagram", { title: "Create new diagram", @@ -430,7 +443,7 @@ Rules: cells are siblings (never nested), ids are unique per page and start from ) // Tool: load_diagram -server.registerTool( +registerWriteTool( "load_diagram", { title: "Load .drawio file", @@ -549,7 +562,7 @@ server.registerTool( ) // Tool: edit_diagram -server.registerTool( +registerWriteTool( "edit_diagram", { title: "Edit diagram", @@ -791,13 +804,15 @@ server.registerTool( } } + // start_session may replace currentSession while this waits + const session = currentSession // Request browser to push fresh state and wait for it (an // expired session first gets its saved file back to sync) let staleNote = "" - restoreSavedSession(currentSession.id) - const syncRequested = requestSync(currentSession.id) + restoreSavedSession(session.id) + const syncRequested = requestSync(session.id) if (syncRequested) { - const synced = await waitForSync(currentSession.id) + const synced = await waitForSync(session.id) if (!synced) { log.warn("get_diagram: sync timeout - state may be stale") staleNote = @@ -808,13 +823,13 @@ server.registerTool( // Fetch latest state from browser, re-normalising to mxfile so a // bare pushed back by the embed/sync path doesn't // strip page structure (see edit_diagram for the same guard). - const browserState = sessionState(currentSession.id) + const browserState = sessionState(session.id) if (browserState?.xml) { - currentSession.xml = + session.xml = normalizeToMxfile(browserState.xml) ?? browserState.xml } - if (!currentSession.xml) { + if (!session.xml) { return { content: [ { @@ -834,8 +849,8 @@ server.registerTool( // The model is now looking at the current state. Record the raw // store value — the gate's fast path is plain string equality // against the store, with a structural comparison as fallback. - const liveXml = browserState?.xml || currentSession.xml - const doc = parseMxfile(currentSession.xml) + const liveXml = browserState?.xml || session.xml + const doc = parseMxfile(session.xml) const pages = doc ? listPagesFromDoc(doc) : [] const pageList = pages.length ? `Pages (${pages.length}): ${pages.map((p) => `[${p.index}] id=${p.id} name="${p.name}" cells=${p.cellCount}`).join(" | ")}` @@ -843,19 +858,19 @@ server.registerTool( // No selector → return full mxfile if (!hasPageSelector(pageSelector)) { - currentSession.lastSeenXml = liveXml + session.lastSeenXml = liveXml return { content: [ { type: "text", - text: `Current diagram XML:\n\n${currentSession.xml}\n\n${pageList}${staleNote}`, + text: `Current diagram XML:\n\n${session.xml}\n\n${pageList}${staleNote}`, }, ], } } // Selector → return a single-page projection - const projection = projectPage(currentSession.xml, pageSelector) + const projection = projectPage(session.xml, pageSelector) if (!projection.ok) { return { content: [ @@ -872,13 +887,13 @@ server.registerTool( } // One page shown counts for all only if the others are as the // model saw them last - currentSession.lastSeenXml = markPageSeen( - currentSession.lastSeenXml, + session.lastSeenXml = markPageSeen( + session.lastSeenXml, liveXml, pageSelector, ) const otherPagesNote = - currentSession.lastSeenXml === liveXml + session.lastSeenXml === liveXml ? "" : `\n\nNote: ${OTHER_PAGES_UNSEEN} Call get_diagram without a page selector before editing.` return { @@ -1148,15 +1163,48 @@ server.registerTool( } } + // start_session may replace currentSession while this waits + const session = currentSession + + // Detect format from extension if not specified + const lowerPath = path.toLowerCase() + const detectedFormat = + format || + (lowerPath.endsWith(".drawio.svg") + ? "drawio.svg" + : lowerPath.endsWith(".png") + ? "png" + : lowerPath.endsWith(".svg") + ? "svg" + : "drawio") + + // The .drawio file is written from the state, so get the + // user's latest edits into it first, as get_diagram does (the + // images are made by the browser from its canvas) + let syncNote = "" + if (detectedFormat === "drawio") { + restoreSavedSession(session.id) + if (!requestSync(session.id)) { + syncNote = + "\n\nNote: the preview was not reachable, so the file may not include the user's latest manual edits." + } else if (!(await waitForSync(session.id))) { + log.warn( + "export_diagram: sync timeout - state may be stale", + ) + syncNote = + "\n\nNote: the browser did not respond, so the file may not include the user's latest manual edits (is the preview tab open?)." + } + } + // Fetch latest state, re-normalised to mxfile so a page // selector works on a bare pushed by the browser - const browserState = sessionState(currentSession.id) + const browserState = sessionState(session.id) if (browserState?.xml) { - currentSession.xml = + session.xml = normalizeToMxfile(browserState.xml) ?? browserState.xml } - if (!currentSession.xml) { + if (!session.xml) { return { content: [ { @@ -1177,18 +1225,6 @@ server.registerTool( const fs = await import("node:fs/promises") const nodePath = await import("node:path") - // Detect format from extension if not specified - const lowerPath = path.toLowerCase() - const detectedFormat = - format || - (lowerPath.endsWith(".drawio.svg") - ? "drawio.svg" - : lowerPath.endsWith(".png") - ? "png" - : lowerPath.endsWith(".svg") - ? "svg" - : "drawio") - // .drawio path - write XML directly (no browser round-trip). if (detectedFormat === "drawio") { let filePath = path @@ -1197,12 +1233,9 @@ server.registerTool( } const absolutePath = nodePath.resolve(filePath) - let outXml = currentSession.xml + let outXml = session.xml if (hasPageSelector(pageSelector)) { - const projection = projectPage( - currentSession.xml, - pageSelector, - ) + const projection = projectPage(session.xml, pageSelector) if (!projection.ok) { return { content: [ @@ -1226,7 +1259,7 @@ server.registerTool( content: [ { type: "text", - text: `Diagram exported successfully!\n\nFile: ${absolutePath}\nSize: ${outXml.length} characters`, + text: `Diagram exported successfully!\n\nFile: ${absolutePath}\nSize: ${outXml.length} characters${syncNote}`, }, ], } @@ -1248,7 +1281,7 @@ server.registerTool( const browserFormat = detectedFormat === "drawio.svg" ? "xmlsvg" : detectedFormat - const state = sessionState(currentSession.id) + const state = sessionState(session.id) if (!state) { return { content: [ @@ -1260,8 +1293,8 @@ server.registerTool( isError: true, } } - if (previewStalled(currentSession.id)) { - return previewStalledError(currentSession.id) + if (previewStalled(session.id)) { + return previewStalledError(session.id) } // ----------------------------------------------------------------- @@ -1281,10 +1314,10 @@ server.registerTool( let projectionXml: string | undefined let pngPageId: string | undefined if (hasPageSelector(pageSelector) && detectedFormat === "png") { - pngPageId = pageIdFor(currentSession.xml, pageSelector) + pngPageId = pageIdFor(session.xml, pageSelector) } if (hasPageSelector(pageSelector) && !pngPageId) { - const projection = projectPage(currentSession.xml, pageSelector) + const projection = projectPage(session.xml, pageSelector) if (!projection.ok) { return { content: [ @@ -1303,7 +1336,7 @@ server.registerTool( } const exportData = await exportViaBrowser( - currentSession.id, + session.id, browserFormat, projectionXml, pngPageId ? { pageId: pngPageId } : undefined, @@ -1493,7 +1526,7 @@ server.registerTool( ) // Tool: add_page -server.registerTool( +registerWriteTool( "add_page", { title: "Add page", @@ -1610,7 +1643,7 @@ server.registerTool( ) // Tool: rename_page -server.registerTool( +registerWriteTool( "rename_page", { title: "Rename page", @@ -1692,7 +1725,7 @@ server.registerTool( ) // Tool: delete_page -server.registerTool( +registerWriteTool( "delete_page", { title: "Delete page", diff --git a/packages/mcp-server/src/load-diagram.ts b/packages/mcp-server/src/load-diagram.ts index c1171dc2..f1609854 100644 --- a/packages/mcp-server/src/load-diagram.ts +++ b/packages/mcp-server/src/load-diagram.ts @@ -50,14 +50,16 @@ export function decompressPageContent(compressed: string): string | null { * any compressed pages. */ export function parseDrawioFileContent(content: string): LoadResult { - const trimmed = content.trim() + let trimmed = content.trim() if (!trimmed) return { ok: false, error: "File is empty." } if (isMxGraphModel(trimmed)) { const normalized = normalizeToMxfile(trimmed) - return normalized - ? { ok: true, xml: normalized } - : { ok: false, error: "Failed to parse XML." } + if (!normalized) { + return { ok: false, error: "Failed to parse XML." } + } + // Parsed below like any , so a broken model is an error + trimmed = normalized } if (!isMxFile(trimmed)) { return { diff --git a/packages/mcp-server/src/new-diagram.ts b/packages/mcp-server/src/new-diagram.ts index f761a235..2cb7b82b 100644 --- a/packages/mcp-server/src/new-diagram.ts +++ b/packages/mcp-server/src/new-diagram.ts @@ -3,6 +3,7 @@ * and the web app's display_diagram tool. */ import { normalizeToMxfile, wrapCellsInModel } from "./pages.ts" +import { readAttributes } from "./xml-attributes.ts" import { validateAndFixXml } from "./xml-validation.ts" export type NewDiagram = @@ -23,12 +24,9 @@ export function reservedIdError(input: string): string | null { /<(mxCell|UserObject|object)\b((?:\s+[\w:.-]+\s*=\s*(?:"[^"]*"|'[^']*'))*)\s*\/?>/g, ) for (const [, tag, attrText] of tags) { - const attrs = new Map() - for (const [, name, double, single] of attrText.matchAll( - /([\w:.-]+)\s*=\s*(?:"([^"]*)"|'([^']*)')/g, - )) { - attrs.set(name, double ?? single) - } + const attrs = new Map( + readAttributes(attrText).map((a) => [a.name, a.value]), + ) const id = attrs.get("id") if (id !== "0" && id !== "1") continue // A wrapper's id is its cell's; an mxCell counts as a shape or edge diff --git a/packages/mcp-server/src/pages.ts b/packages/mcp-server/src/pages.ts index a80ed43a..1b83e55f 100644 --- a/packages/mcp-server/src/pages.ts +++ b/packages/mcp-server/src/pages.ts @@ -17,6 +17,7 @@ * - how to add/rename/delete pages without re-parsing ad-hoc. */ +import { readAttributes } from "./xml-attributes.ts" import { getXmlSyntaxError } from "./xml-syntax.ts" export interface PageInfo { @@ -125,15 +126,13 @@ export function wrapCellsInModel(xml: string): string { if (end !== -1 && /^(\s*<\/[^>]+>)*\s*$/.test(content.slice(end))) { content = content.slice(0, end) } + // The root cells come with the wrapper (a label holding id='1' is not + // an id) content = content - .replace( - /]*\bid\s*=\s*["']0["'][^>]*(?:\/>|>\s*<\/mxCell>)/g, - "", - ) - .replace( - /]*\bid\s*=\s*["']1["'][^>]*(?:\/>|>\s*<\/mxCell>)/g, - "", - ) + .replace(/]*?(?:\/>|>\s*<\/mxCell>)/g, (cell) => { + const id = readAttributes(cell).find((a) => a.name === "id")?.value + return id === "0" || id === "1" ? "" : cell + }) .trim() return `${ROOT_CELLS}${content}` } diff --git a/packages/mcp-server/src/persistence.ts b/packages/mcp-server/src/persistence.ts index 38e47e04..1d35330c 100644 --- a/packages/mcp-server/src/persistence.ts +++ b/packages/mcp-server/src/persistence.ts @@ -38,6 +38,16 @@ export function defaultDataDir(): string | null { return dir ? expandHome(dir) : join(homedir(), ".next-ai-drawio") } +/** The file surely does not exist (not merely out of reach) */ +function isGone(path: string): boolean { + try { + statSync(path) + return false + } catch (error) { + return (error as NodeJS.ErrnoException).code === "ENOENT" + } +} + export class Autosaver { private pending = new Map< string, @@ -57,22 +67,23 @@ export class Autosaver { // Saved files that could not be read back: never written over, since // the session then shows something else than what they hold. Cleared - // once the file is read, or is gone. + // once the file is read, or is surely gone (a folder without permission + // also makes a file look missing). private unreadable = new Set() /** The session's saved diagram, or null. */ load(sessionId: string): string | null { const path = this.pathFor(sessionId) if (!path) return null - if (!existsSync(path)) { - this.unreadable.delete(path) - return null - } try { const xml = readFileSync(path, "utf-8") this.unreadable.delete(path) return xml } catch (error) { + if ((error as NodeJS.ErrnoException).code === "ENOENT") { + this.unreadable.delete(path) + return null + } log.warn(`Could not read the saved diagram ${path}: ${error}`) this.unreadable.add(path) return null @@ -102,10 +113,15 @@ export class Autosaver { const path = this.pathFor(sessionId) if (!entry || !this.dir || !path) return if (this.unreadable.has(path)) { - log.warn( - `Not saving ${path}: it could not be read, so it may hold work this session does not show`, - ) - return + // Deleted meanwhile: nothing left to protect + if (isGone(path)) { + this.unreadable.delete(path) + } else { + log.warn( + `Not saving ${path}: it could not be read, so it may hold work this session does not show`, + ) + return + } } try { const isNew = !existsSync(path) diff --git a/packages/mcp-server/src/preview/preview.js b/packages/mcp-server/src/preview/preview.js index 425cb159..10b2ec67 100644 --- a/packages/mcp-server/src/preview/preview.js +++ b/packages/mcp-server/src/preview/preview.js @@ -7,11 +7,16 @@ let stateId = null; // the last one the server has let latestXml = null; let pushFailing = false; // the last push could not reach the server +// After recovery replaced the canvas, until draw.io reports the load: an +// autosave still on its way belongs to the canvas being replaced +let awaitingLoad = false; let pollSeq = 0, lastHandledPoll = 0; // polls overlap; older answers are dropped let pendingSvgExport = null; let pendingSvgBase = 0; // version the pending autosave was based on let pendingSvgStateId = null; // and the state it belonged to -let pendingAiSvg = false; +// The latest thumbnail export of a loaded server write: its number (echoed +// by draw.io), the state and version it showed, and the XML loaded +let thumbExportSeq = 0, thumbExport = null; let pendingMcpExport = null; // 'png', 'svg' or 'xmlsvg' when MCP requested export let mcpExportSeq = 0; // number of the latest MCP export let mcpExportId = null; // the server's id for it, sent back with the result @@ -26,11 +31,16 @@ window.addEventListener('message', (e) => { if (msg.event === 'init') { isReady = true; if (pendingXml) { loadDiagram(pendingXml); pendingXml = null; } + } else if (msg.event === 'load') { + awaitingLoad = false; } else if ((msg.event === 'save' || msg.event === 'autosave') && msg.xml) { // Ignore autosave while a single-page projection is on screen // for a page-targeted export — otherwise we'd push the // transient projection back as the canonical session state. if (projectionExportActive) return; + // An edit of the canvas that recovery is replacing: kept in + // History, never over the recovered diagram + if (awaitingLoad) { pushState(msg.xml, '', currentVersion, 'recover'); return; } // Also an edit undone back to what the server has latestXml = msg.xml; if (msg.xml === lastXml) return; @@ -109,17 +119,21 @@ window.addEventListener('message', (e) => { // Handle SVG export let svg = msg.data; if (!svg.startsWith('data:')) svg = 'data:image/svg+xml;base64,' + btoa(unescape(encodeURIComponent(svg))); - if (pendingSvgExport) { - const xml = pendingSvgExport; - pendingSvgExport = null; - pushState(xml, svg, pendingSvgBase, 'edit', pendingSvgStateId); - } else if (pendingAiSvg) { - pendingAiSvg = false; + if (msg.message && msg.message.thumbExport) { + // Only for the latest load, and only if the canvas still + // shows it: the export pictures the canvas as it is now + const t = thumbExport; + if (!t || msg.message.thumbExport !== t.n || latestXml !== t.xml) return; + thumbExport = null; fetch('/api/history-svg', { method: 'POST', headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ sessionId, svg }) + body: JSON.stringify({ sessionId, svg, stateId: t.stateId, version: t.version }) }).catch(() => {}); + } else if (pendingSvgExport) { + const xml = pendingSvgExport; + pendingSvgExport = null; + pushState(xml, svg, pendingSvgBase, 'edit', pendingSvgStateId); } } } catch {} @@ -131,9 +145,12 @@ function loadDiagram(xml, capturePreview = false) { latestXml = xml; iframe.contentWindow.postMessage(JSON.stringify({ action: 'load', xml, autosave: 1 }), '*'); if (capturePreview) { + // A server write: currentVersion is its version + const t = { n: ++thumbExportSeq, stateId, version: currentVersion, xml }; + thumbExport = t; setTimeout(() => { - pendingAiSvg = true; - iframe.contentWindow.postMessage(JSON.stringify({ action: 'export', format: 'svg' }), '*'); + if (thumbExport !== t) return; // a newer load takes its own + iframe.contentWindow.postMessage(JSON.stringify({ action: 'export', format: 'svg', thumbExport: t.n }), '*'); }, 500); } } @@ -177,6 +194,19 @@ async function pushState(xml, svg = '', baseVersion = currentVersion, source = ' if (sid !== stateId) return; currentVersion = d.version; lastXml = xml; + // The canvas changed while this edit was on its way, to + // something no pending autosave will send (an undo back to the + // previous version): send it now. A sync reply is draw.io's + // export of the canvas, in another format than its autosave. + if (latestXml && latestXml !== xml && pendingSvgExport !== latestXml && source === 'edit') { + pushState(latestXml); + } + } + // Over the server's size limit: the image is most of it, so try once + // without it + else if (r.status === 413) { + if (svg) pushState(xml, '', baseVersion, source, sid); + else showNotice('This diagram is too large to save to the MCP server (over 10 MB). Use Download to keep it.'); } // 409: the AI wrote a newer version, or the server lost the state // this push was based on; the next poll sorts it out @@ -216,6 +246,7 @@ function recoverState(s) { // server was down are saved now) if (projectionShown && mine) { iframe.contentWindow.postMessage(JSON.stringify({ action: 'load', xml: mine, autosave: 1 }), '*'); + expectLoad(); } if (mine && mine !== s.xml) pushState(mine, '', s.version); } else { @@ -223,10 +254,18 @@ function recoverState(s) { // missed, a saved file): show that, and keep this tab's copy in // History unless it is the same loadDiagram(s.xml, true); + expectLoad(); if (mine && mine !== s.xml) pushState(mine, '', s.version, 'recover'); } } +// Until draw.io reports the load (its messages come in order), an autosave +// is from the canvas being replaced; in case no report comes, not for long +function expectLoad() { + awaitingLoad = true; + setTimeout(() => { awaitingLoad = false; }, 5000); +} + let pendingSyncExport = false; let pendingSyncBase = 0; // version the pending sync export was taken at let pendingSyncStateId = null; // and the state it belonged to @@ -368,7 +407,7 @@ saveConfirmBtn.onclick = () => { saveConfirmBtn.textContent = 'Exporting...'; if (format === 'drawio') { - // Use lastXml directly instead of requesting export (avoids race with SVG exports). + // Use the XML directly instead of requesting export (avoids race with SVG exports). // session.xml is canonically after the multi-page refactor, // so no wrapper injection is needed. The legacy fallback below // remains only for documents that somehow slipped past diff --git a/packages/mcp-server/src/xml-attributes.ts b/packages/mcp-server/src/xml-attributes.ts new file mode 100644 index 00000000..1fa679ff --- /dev/null +++ b/packages/mcp-server/src/xml-attributes.ts @@ -0,0 +1,27 @@ +/** + * The attributes of one tag as written: name="value" or name='value' + * pairs. A quoted value is read whole, so text inside it such as + * value="Use parent='1'" is never taken for an attribute. + */ +export interface TagAttribute { + name: string + value: string + // The attribute's text in the tag, with the whitespace before it + start: number + end: number +} + +export function readAttributes(tag: string): TagAttribute[] { + const attributes: TagAttribute[] = [] + for (const m of tag.matchAll( + /\s*([A-Za-z_:][\w:.-]*)\s*=\s*(?:"([^"]*)"|'([^']*)')/g, + )) { + attributes.push({ + name: m[1], + value: m[2] ?? m[3], + start: m.index, + end: m.index + m[0].length, + }) + } + return attributes +} diff --git a/packages/mcp-server/src/xml-validation.ts b/packages/mcp-server/src/xml-validation.ts index eb88c9cd..6962a59d 100644 --- a/packages/mcp-server/src/xml-validation.ts +++ b/packages/mcp-server/src/xml-validation.ts @@ -3,6 +3,7 @@ * Copied from lib/utils.ts to avoid cross-package imports */ +import { readAttributes } from "./xml-attributes.ts" import { getXmlSyntaxError } from "./xml-syntax.ts" // ============================================================================ @@ -156,16 +157,10 @@ function replaceInOpeningTags( /** Check for duplicate structural attributes in a tag */ function checkDuplicateAttributes(xml: string): string | null { const structuralSet = new Set(STRUCTURAL_ATTRS) - const tagPattern = /<[^>]+>/g - let tagMatch - while ((tagMatch = tagPattern.exec(xml)) !== null) { - const tag = tagMatch[0] - const attrPattern = /\s([a-zA-Z_:][a-zA-Z0-9_:.-]*)\s*=/g + for (const [tag] of xml.matchAll(/<[^>]+>/g)) { const attributes = new Map() - let attrMatch - while ((attrMatch = attrPattern.exec(tag)) !== null) { - const attrName = attrMatch[1] - attributes.set(attrName, (attributes.get(attrName) || 0) + 1) + for (const { name } of readAttributes(tag)) { + attributes.set(name, (attributes.get(name) || 0) + 1) } const duplicates = Array.from(attributes.entries()) .filter(([name, count]) => count > 1 && structuralSet.has(name)) @@ -596,27 +591,23 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } { // 3. Fix duplicate attributes let dupAttrFixed = false + const structural = new Set(STRUCTURAL_ATTRS) fixed = fixed.replace(/<[^>]+>/g, (tag) => { - let newTag = tag - for (const attr of STRUCTURAL_ATTRS) { - const attrRegex = new RegExp( - `\\s${attr}\\s*=\\s*["'][^"']*["']`, - "gi", - ) - const matches = tag.match(attrRegex) - if (matches && matches.length > 1) { - let firstKept = false - newTag = newTag.replace(attrRegex, (m) => { - if (!firstKept) { - firstKept = true - return m - } - dupAttrFixed = true - return "" - }) + // Keep the first of each, drop the later ones + const seen = new Set() + let newTag = "" + let last = 0 + for (const attr of readAttributes(tag)) { + if (!structural.has(attr.name)) continue + if (!seen.has(attr.name)) { + seen.add(attr.name) + continue } + newTag += tag.slice(last, attr.start) + last = attr.end + dupAttrFixed = true } - return newTag + return newTag + tag.slice(last) }) if (dupAttrFixed) { fixes.push("Removed duplicate structural attributes") diff --git a/packages/mcp-server/tests/exclusive.test.ts b/packages/mcp-server/tests/exclusive.test.ts new file mode 100644 index 00000000..dce451ae --- /dev/null +++ b/packages/mcp-server/tests/exclusive.test.ts @@ -0,0 +1,59 @@ +/** + * Tests for the queue the write tools run in (index.ts registerWriteTool). + */ +import { describe, expect, it } from "vitest" +import { createExclusive } from "../src/exclusive.ts" + +const extra = (signal = new AbortController().signal) => ({ signal }) +const tick = () => new Promise((r) => setTimeout(r, 5)) + +describe("createExclusive", () => { + it("runs the calls one at a time, in order", async () => { + const exclusive = createExclusive() + const events: string[] = [] + // Reads the document, waits, then writes it back + const addPage = exclusive(async (args: { name: string }, _extra) => { + events.push(`read ${args.name}`) + await tick() + events.push(`write ${args.name}`) + return { content: [] } + }) + await Promise.all([ + addPage({ name: "A" }, extra()), + addPage({ name: "B" }, extra()), + ]) + expect(events).toEqual(["read A", "write A", "read B", "write B"]) + }) + + it("goes on after a call that threw", async () => { + const exclusive = createExclusive() + const failing = exclusive(async (_extra: unknown) => { + throw new Error("broken") + }) + const working = exclusive(async (_extra: unknown) => ({ content: [] })) + const first = failing(extra()) + const second = working(extra()) + await expect(first).rejects.toThrow("broken") + await expect(second).resolves.toEqual({ content: [] }) + }) + + it("skips a call cancelled while it waited", async () => { + const exclusive = createExclusive() + let ran = false + const slow = exclusive(async (_extra: unknown) => { + await tick() + return { content: [] } + }) + const deletePage = exclusive(async (_extra: unknown) => { + ran = true + return { content: [] } + }) + const cancel = new AbortController() + const first = slow(extra()) + const second = deletePage(extra(cancel.signal)) + cancel.abort() + await first + expect(await second).toMatchObject({ isError: true }) + expect(ran).toBe(false) + }) +}) diff --git a/packages/mcp-server/tests/http-server.test.ts b/packages/mcp-server/tests/http-server.test.ts index 75d539a0..60c09011 100644 --- a/packages/mcp-server/tests/http-server.test.ts +++ b/packages/mcp-server/tests/http-server.test.ts @@ -467,22 +467,82 @@ describe("history restore", () => { const page = (cellId: string) => `` + // A thumbnail the tab took after loading the server write at `version` + const thumbnail = (id: string, svg: string, version: number) => + postJson("/api/history-svg", { + sessionId: id, + svg, + stateId: getState(id)?.stateId, + version, + }) + it("gives a thumbnail only to the entry it shows", async () => { const id = "mcp-history-thumb" - setState(id, page("shown")) + const version = setState(id, page("shown")) // The last entry is another diagram (a tab's copy kept on recovery) addHistory(id, page("other")) - await postJson("/api/history-svg", { - sessionId: id, - svg: "SVG-OF-SHOWN", - }) + await thumbnail(id, "SVG-OF-SHOWN", version) expect(getHistory(id).at(-1)?.svg).toBe("") addHistory(id, page("shown")) + await thumbnail(id, "SVG-OF-SHOWN", version) + expect(getHistory(id).at(-1)?.svg).toBe("SVG-OF-SHOWN") + expect(getState(id)?.svg).toBe("SVG-OF-SHOWN") + }) + + it("drops a thumbnail that arrives after the next AI write", async () => { + const id = "mcp-history-thumb-late" + const first = setState(id, page("first")) + addHistory(id, page("first")) + setState(id, page("second")) + addHistory(id, page("second")) + await thumbnail(id, "SVG-OF-FIRST", first) + expect(getHistory(id).map((e) => e.svg)).toEqual(["", ""]) + expect(getState(id)?.svg).toBeUndefined() + }) + + it("drops a thumbnail of a state the server has since lost", async () => { + const id = "mcp-history-thumb-state" + const version = setState(id, page("shown")) + addHistory(id, page("shown")) await postJson("/api/history-svg", { sessionId: id, svg: "SVG-OF-SHOWN", + stateId: "another-state", + version, }) - expect(getHistory(id).at(-1)?.svg).toBe("SVG-OF-SHOWN") + expect(getHistory(id).at(-1)?.svg).toBe("") + }) + + it("keeps a thumbnail in time after a sync reply", async () => { + const id = "mcp-history-thumb-sync" + const version = setState(id, page("ai")) + addHistory(id, page("ai")) + // draw.io's copy of the same diagram, sent back for a sync + const synced = page("ai").replace( + "", + '', + ) + await postJson("/api/state", { + sessionId: id, + xml: synced, + baseVersion: version, + source: "sync", + stateId: getState(id)?.stateId, + }) + await thumbnail(id, "SVG-OF-AI", version) + expect(getHistory(id).at(-1)?.svg).toBe("SVG-OF-AI") + // The state's own image is of the synced XML only + expect(getState(id)?.svg).toBeUndefined() + }) + + it("keeps the image when a write repeats the same XML", async () => { + const id = "mcp-history-thumb-same" + const version = setState(id, page("same")) + await thumbnail(id, "SVG-OF-SAME", version) + setState(id, page("same"), undefined, true) + expect(getState(id)?.svg).toBe("SVG-OF-SAME") + setState(id, page("changed"), undefined, true) + expect(getState(id)?.svg).toBeUndefined() }) it("never pairs the image of an older diagram with a newer one", async () => { @@ -515,14 +575,26 @@ describe("history restore", () => { expect(getHistory(id).map((e) => e.xml)).toContain(cleared) }) - it("adds no entry for a re-serialized copy of the last one", () => { + it("adds no entry for a copy of the last one", () => { const id = "mcp-history-dedupe" addHistory(id, page("same")) + addHistory(id, page("same"), "SVG") + expect(getHistory(id)).toHaveLength(1) + // The missing image is filled in + expect(getHistory(id)[0].svg).toBe("SVG") + }) + + it("keeps a version that changed only the background", () => { + const id = "mcp-history-background" + addHistory(id, page("same")) addHistory( id, - page("same").replace("", ''), + page("same").replace( + "", + '', + ), ) - expect(getHistory(id)).toHaveLength(1) + expect(getHistory(id)).toHaveLength(2) }) it("keeps manual edits in history before restoring", async () => { @@ -545,3 +617,46 @@ describe("history restore", () => { expect(getHistory(id).map((e) => e.xml)).toContain(doc("manual")) }) }) + +describe("bodies over the size limit", () => { + it("answers 413 after reading the whole body", async () => { + const mib = Buffer.alloc(1024 * 1024, "x") + // Still sending when the limit is passed, as a browser would be. + // A browser whose upload is cut off reports a network error. + const result = await new Promise<{ status?: number; sent: boolean }>( + (resolve) => { + let sent = false + const req = http.request( + { + host: "127.0.0.1", + port, + path: "/api/state", + method: "POST", + headers: { + host: `localhost:${port}`, + "content-type": "application/json", + }, + }, + (res) => { + res.resume() + res.on("end", () => + resolve({ status: res.statusCode, sent }), + ) + }, + ) + req.on("error", () => resolve({ sent })) + const writeNext = (i: number) => { + if (i === 15) { + req.end(() => { + sent = true + }) + return + } + req.write(mib, () => setTimeout(() => writeNext(i + 1), 10)) + } + writeNext(0) + }, + ) + expect(result).toEqual({ status: 413, sent: true }) + }) +}) diff --git a/packages/mcp-server/tests/load-diagram.test.ts b/packages/mcp-server/tests/load-diagram.test.ts index b2b8cbbe..d0264e6e 100644 --- a/packages/mcp-server/tests/load-diagram.test.ts +++ b/packages/mcp-server/tests/load-diagram.test.ts @@ -66,6 +66,13 @@ describe("parseDrawioFileContent", () => { } }) + it("rejects a bare mxGraphModel that is not closed", () => { + const r = parseDrawioFileContent( + MODEL_XML.replace("", ""), + ) + expect(r.ok).toBe(false) + }) + it("decompresses a compressed mxfile into plain XML pages", () => { const r = parseDrawioFileContent(COMPRESSED_MXFILE) expect(r.ok).toBe(true) diff --git a/packages/mcp-server/tests/persistence.test.ts b/packages/mcp-server/tests/persistence.test.ts index b766b269..85e63199 100644 --- a/packages/mcp-server/tests/persistence.test.ts +++ b/packages/mcp-server/tests/persistence.test.ts @@ -130,6 +130,43 @@ describe("Autosaver", () => { expect(readFileSync(path, "utf-8")).toBe(DIAGRAM) }) + it("saves again when the unreadable file is deleted during the session", () => { + const saver = new Autosaver(tempDir(), 10) + saver.schedule("mcp-live", DIAGRAM) + saver.flush() + const path = saver.pathFor("mcp-live") as string + chmodSync(path, 0o000) + expect(saver.load("mcp-live")).toBeNull() + // The session goes on (load is not called again); the user removes + // the broken file + chmodSync(path, 0o644) + rmSync(path) + saver.schedule("mcp-live", DIAGRAM) + saver.flush() + expect(readFileSync(path, "utf-8")).toBe(DIAGRAM) + }) + + it("keeps protecting a file it cannot even look at", () => { + // A folder without permission makes the file look missing; it is not + const dir = tempDir() + const saver = new Autosaver(dir, 10) + saver.schedule("mcp-hidden", DIAGRAM) + saver.flush() + const path = saver.pathFor("mcp-hidden") as string + chmodSync(path, 0o000) + expect(saver.load("mcp-hidden")).toBeNull() + chmodSync(dir, 0o000) + try { + expect(saver.load("mcp-hidden")).toBeNull() + } finally { + chmodSync(dir, 0o755) + } + chmodSync(path, 0o644) + saver.schedule("mcp-hidden", BLANK) + saver.flush() + expect(readFileSync(path, "utf-8")).toBe(DIAGRAM) + }) + it("does nothing when saving is off", () => { const saver = new Autosaver(null) expect(saver.pathFor("mcp-x")).toBeNull() diff --git a/packages/mcp-server/tests/wrap-cells.test.ts b/packages/mcp-server/tests/wrap-cells.test.ts index 3d8378b8..51c9ac12 100644 --- a/packages/mcp-server/tests/wrap-cells.test.ts +++ b/packages/mcp-server/tests/wrap-cells.test.ts @@ -165,3 +165,18 @@ describe("prepareNewDiagram", () => { expect(out.ok).toBe(true) }) }) + +describe("labels that look like attributes", () => { + const layer = `` + + it("keep their cell when the root cells are stripped", () => { + const wrapped = wrapCellsInModel(ROOTS + layer) + expect(wrapped).toBe( + `${ROOTS}${layer}`, + ) + }) + + it("are not taken for a reserved id", () => { + expect(reservedIdError(layer)).toBeNull() + }) +}) diff --git a/packages/mcp-server/tests/xml-validation.test.ts b/packages/mcp-server/tests/xml-validation.test.ts index 67c20264..4f246f67 100644 --- a/packages/mcp-server/tests/xml-validation.test.ts +++ b/packages/mcp-server/tests/xml-validation.test.ts @@ -326,3 +326,30 @@ describe("text directly under a page", () => { expect(r.error).toMatch(/not-base64/) }) }) + +describe("attributes inside quoted values", () => { + const labelled = `` + + it("are not duplicates of the real ones", () => { + const r = validateAndFixXml(model(labelled)) + expect(r.valid).toBe(true) + expect(r.fixed ?? model(labelled)).toContain(labelled) + }) + + it("are kept when a real duplicate is removed", () => { + // The bare & makes the repair run on the whole document + const cell = `` + const r = validateAndFixXml(model(cell + BROKEN_CELL)) + expect(r.valid).toBe(true) + expect(r.fixed).toContain( + ``, + ) + }) + + it("leave two cells with an unbalanced quote their ids and parents", () => { + const broken = `` + const r = validateAndFixXml(model(broken)) + expect(r.fixed).toContain(`id="5" value="B" vertex="1" parent="1"`) + expect(r.fixed).toMatch(/id="4"[^>]*vertex="1" parent="1"/) + }) +}) diff --git a/tests/e2e/chat.spec.ts b/tests/e2e/chat.spec.ts index b2d125a3..554a8efb 100644 --- a/tests/e2e/chat.spec.ts +++ b/tests/e2e/chat.spec.ts @@ -1,4 +1,10 @@ -import { expect, getIframe, test } from "./lib/fixtures" +import { + expect, + getChatInput, + getIframe, + sendMessage, + test, +} from "./lib/fixtures" test.describe("Chat Panel", () => { test.beforeEach(async ({ page }) => { @@ -20,3 +26,82 @@ test.describe("Chat Panel", () => { expect(src).toBeTruthy() }) }) + +test.describe("Crossing the mobile breakpoint", () => { + test.beforeEach(async ({ page }) => { + // A text answer that arrives in parts over a few seconds + await page.addInitScript(() => { + const realFetch = window.fetch + window.fetch = async (input, init) => { + const url = input instanceof Request ? input.url : String(input) + if (!url.endsWith("/api/chat")) return realFetch(input, init) + const events = [ + { type: "start" }, + { type: "text-start", id: "t" }, + { type: "text-delta", id: "t", delta: "Once upon" }, + { type: "text-delta", id: "t", delta: " a time." }, + { type: "text-end", id: "t" }, + { type: "finish" }, + ] + const body = new ReadableStream({ + async start(controller) { + for (const event of events) { + controller.enqueue( + new TextEncoder().encode( + `data: ${JSON.stringify(event)}\n\n`, + ), + ) + await new Promise((r) => setTimeout(r, 1500)) + } + controller.enqueue( + new TextEncoder().encode("data: [DONE]\n\n"), + ) + controller.close() + }, + }) + return new Response(body, { + headers: { "content-type": "text/event-stream" }, + }) + } + }) + await page.setViewportSize({ width: 1280, height: 800 }) + await page.goto("/", { waitUntil: "networkidle" }) + await getIframe(page).waitFor({ state: "visible", timeout: 30000 }) + }) + + test("keeps the chat and its streaming answer", async ({ page }) => { + const chat = page.locator('[data-panel-id="chat-panel"]') + await sendMessage(page, "Tell me a story") + await expect(page.getByText("Once upon")).toBeVisible({ + timeout: 10000, + }) + + await page.setViewportSize({ width: 600, height: 900 }) + await expect(page.getByText("Tell me a story")).toBeVisible() + // Half the height on mobile + await expect + .poll(async () => (await chat.boundingBox())?.height ?? 0) + .toBeCloseTo(450, -1) + + await page.setViewportSize({ width: 1280, height: 800 }) + // A third of the width on desktop + await expect + .poll(async () => (await chat.boundingBox())?.width ?? 0) + .toBeCloseTo(1280 / 3, -1) + await expect(page.getByText("Once upon a time.")).toBeVisible({ + timeout: 10000, + }) + await expect(page.getByText("Tell me a story")).toBeVisible() + }) + + test("opens a chat collapsed on desktop", async ({ page }) => { + await page.locator("button:has(svg.lucide-panel-right-close)").click() + await expect(getChatInput(page)).toBeHidden() + + await page.setViewportSize({ width: 600, height: 900 }) + await expect(getChatInput(page)).toBeVisible() + + await page.setViewportSize({ width: 1280, height: 800 }) + await expect(getChatInput(page)).toBeVisible() + }) +}) diff --git a/tests/e2e/diagram-content.spec.ts b/tests/e2e/diagram-content.spec.ts index a0df3dec..27d130f1 100644 --- a/tests/e2e/diagram-content.spec.ts +++ b/tests/e2e/diagram-content.spec.ts @@ -409,6 +409,14 @@ const drawReply = (id: string, xml: string) => { const call = toolCallEvents(id, "display_diagram", { xml }) return `${sse([{ type: "start" }, call.start, ...call.deltas, call.done, { type: "finish" }])}data: [DONE]\n\n` } +const textReply = (text: string) => + `${sse([ + { type: "start" }, + { type: "text-start", id: "t" }, + { type: "text-delta", id: "t", delta: text }, + { type: "text-end", id: "t" }, + { type: "finish" }, + ])}data: [DONE]\n\n` // SSE comments keep a stream open without sending anything const KEEP_OPEN = Array(20).fill(":\n\n") @@ -685,3 +693,45 @@ test("stopping during the screenshot check starts no new request", async ({ await p.waitForTimeout(5000) expect(chatRequests).toBe(1) }) + +test("stopping during the screenshot check lets the next message go at once", async ({ + page: p, +}) => { + // The check was still running when the user stopped; it held up the + // chat until it ended, and its call never got a result + await p.addInitScript(() => { + localStorage.setItem("next-ai-draw-io-vlm-validation-enabled", "true") + }) + const bodies: Array<{ messages: any[] }> = [] + await p.route("**/api/chat", async (route) => { + bodies.push(route.request().postDataJSON()) + const n = bodies.length + await route.fulfill({ + status: 200, + contentType: "text/event-stream", + body: + n === 1 + ? drawReply("d1", cell("a", "Alpha", 40)) + : textReply("Second answer"), + }) + }) + let checking = false + await p.route("**/api/validate-diagram", async (route) => { + checking = true + // Much longer than this test waits for the second answer + await new Promise((r) => setTimeout(r, 30000)) + await route.fulfill({ status: 200, body: "{}" }).catch(() => {}) + }) + await p.goto("/", { waitUntil: "networkidle" }) + await getIframe(p).waitFor({ state: "visible", timeout: 30000 }) + await sendMessage(p, "Draw a box") + await expect.poll(() => checking, { timeout: 15000 }).toBe(true) + await p.getByRole("button", { name: "Stop generation" }).click() + await sendMessage(p, "Thanks") + await expect(p.getByText("Second answer")).toBeVisible({ timeout: 8000 }) + // The drawing call had its result when the next message was sent + const draw = bodies[1].messages + .flatMap((m: any) => m.parts ?? []) + .find((part: any) => part.type === "tool-display_diagram") + expect(draw?.state).toBe("output-available") +}) diff --git a/tests/e2e/history-restore.spec.ts b/tests/e2e/history-restore.spec.ts index 4162a974..37b5982c 100644 --- a/tests/e2e/history-restore.spec.ts +++ b/tests/e2e/history-restore.spec.ts @@ -86,6 +86,40 @@ test.describe("History and Session Restore", () => { ).toBeVisible() }) + test("new chat can go on without saving when storage is full", async ({ + page, + }) => { + // Old chats can only be deleted from the empty chat's list, so the + // user must be able to get there + await page.route("**/api/chat", async (route) => { + await route.fulfill({ + status: 200, + contentType: "text/event-stream", + body: createMockSSEResponse( + SINGLE_BOX_XML, + "Created your test diagram.", + ), + }) + }) + await page.goto("/", { waitUntil: "networkidle" }) + await getIframe(page).waitFor({ state: "visible", timeout: 30000 }) + await sendMessage(page, "Create a test diagram") + await waitForText(page, "Created your test diagram.") + await page.evaluate(() => { + IDBObjectStore.prototype.put = () => { + throw new DOMException("Storage is full", "QuotaExceededError") + } + }) + await page.locator('[data-testid="new-chat-button"]').click() + await page + .getByRole("button", { name: "Continue without saving" }) + .click({ timeout: 5000 }) + await expect( + page.locator('text="Created your test diagram."'), + ).toHaveCount(0, { timeout: 5000 }) + await expect(page.getByText("Paper to Diagram")).toBeVisible() + }) + // A diagram drawn by hand, without chat messages: loaded into draw.io // directly, then moved with an arrow key, which draw.io reports as an // edit like any manual change @@ -338,3 +372,55 @@ test.describe("History and Session Restore", () => { }) }) }) + +/** Number of chats stored in this origin's IndexedDB */ +const countSessions = (page: Page) => + page.evaluate( + () => + new Promise((resolve, reject) => { + const open = indexedDB.open("next-ai-drawio") + open.onerror = () => reject(open.error) + open.onsuccess = () => { + const db = open.result + if (!db.objectStoreNames.contains("sessions")) { + db.close() + return resolve(0) + } + const count = db + .transaction("sessions", "readonly") + .objectStore("sessions") + .count() + count.onsuccess = () => { + db.close() + resolve(count.result) + } + } + }), + ) + +test("new chat right after an answer saves that chat once", async ({ + page, +}) => { + test.setTimeout(180_000) + await page.route("**/api/chat", async (route) => { + await route.fulfill({ + status: 200, + contentType: "text/event-stream", + body: createMockSSEResponse(SINGLE_BOX_XML, "Drew the box."), + }) + }) + await page.goto("/", { waitUntil: "networkidle" }) + await getIframe(page).waitFor({ state: "visible", timeout: 30000 }) + const newChat = page.locator('[data-testid="new-chat-button"]') + // The auto-save runs a second after the answer; New Chat around then + // waits for its thumbnail while the auto-save starts + for (let run = 1; run <= 10; run++) { + await sendMessage(page, `Draw box ${run}`) + await waitForText(page, "Drew the box.") + await page.waitForTimeout(500 + run * 100) + await newChat.click() + await expect(page.getByText("Drew the box.")).toHaveCount(0) + await page.waitForTimeout(2500) + expect(await countSessions(page), `run ${run}`).toBe(run) + } +}) diff --git a/tests/e2e/provider-models.spec.ts b/tests/e2e/provider-models.spec.ts index a70a2df9..75cadff2 100644 --- a/tests/e2e/provider-models.spec.ts +++ b/tests/e2e/provider-models.spec.ts @@ -302,6 +302,35 @@ test("an older test does not end a newer one's spinners", async ({ page }) => { await expect(dialog.locator('[title="1.0 s"]')).toHaveCount(1) }) +test("an older test touches nothing, also when the key came back", async ({ + page, +}) => { + const releases: Array<() => void> = [] + await page.route("**/api/validate-model", async (route) => { + const n = releases.length + await new Promise((r) => releases.push(r)) + await route.fulfill({ + status: 200, + // The older test's result would say 9.0 s + json: { valid: true, responseTime: n === 0 ? 9000 : 1000 }, + }) + }) + const dialog = await openQwenSettings(page, TWO_PROVIDERS) + await dialog.getByRole("button", { name: "Test", exact: true }).click() + await expect.poll(() => releases.length).toBe(1) + // The key changes and comes back, and the user tests again + await dialog.locator("#api-key").fill("other-key") + await dialog.locator("#api-key").fill("test-key") + await dialog.getByRole("button", { name: "Test", exact: true }).click() + await expect.poll(() => releases.length).toBe(2) + releases[0]() + await page.waitForTimeout(500) + await expect(dialog.locator(".animate-spin").first()).toBeVisible() + await expect(dialog.locator('[title="9.0 s"]')).toHaveCount(0) + releases[1]() + await expect(dialog.locator('[title="1.0 s"]')).toHaveCount(1) +}) + test("no spinner stays after another tab's change while elsewhere", async ({ page, }) => { diff --git a/tests/unit/admin-providers.test.ts b/tests/unit/admin-providers.test.ts index 90fc2664..c6064234 100644 --- a/tests/unit/admin-providers.test.ts +++ b/tests/unit/admin-providers.test.ts @@ -69,15 +69,15 @@ describe("deriveEnvUpdates", () => { expect(updates.ADMIN_OPENAI_API_KEY_2).toBe("sk-second") }) - it("sends an Ollama key without a URL to Ollama Cloud, like its Test", () => { - // Chat sends a server Ollama key to OLLAMA_BASE_URL, or to local - // Ollama without one; the Test sends it to Ollama Cloud - const cloud = deriveEnvUpdates( + it("writes an Ollama URL only when the entry has one", () => { + // Without one, the operator's own OLLAMA_BASE_URL (or local Ollama) + // stays, also for the AI_PROVIDER=ollama default model + const keyOnly = deriveEnvUpdates( [provider({ provider: "ollama", apiKey: "ollama-key" })], [], ) - expect(cloud.OLLAMA_API_KEY).toBe("ollama-key") - expect(cloud.OLLAMA_BASE_URL).toBe("https://ollama.com/api") + expect(keyOnly.OLLAMA_API_KEY).toBe("ollama-key") + expect(keyOnly.OLLAMA_BASE_URL ?? null).toBeNull() const own = deriveEnvUpdates( [ provider({ diff --git a/tests/unit/admin-test-model.test.ts b/tests/unit/admin-test-model.test.ts index 42ce23d5..a9254b0f 100644 --- a/tests/unit/admin-test-model.test.ts +++ b/tests/unit/admin-test-model.test.ts @@ -48,6 +48,8 @@ describe("admin Test of an entry without a URL", () => { process.env.OPENAI_BASE_URL = "https://operator-proxy.example.com/v1" await test({ provider: "openai", apiKey: "panel-key" }) expect(sent.body.baseUrl).toBe("https://operator-proxy.example.com/v1") + // The server's own URL, tested without the rules for typed URLs + expect(sent.body.serverBaseUrl).toBe(true) process.env.AI_GATEWAY_BASE_URL = "https://gateway.example.com/v3/ai" await test({ provider: "gateway", apiKey: "k" }) @@ -65,4 +67,15 @@ describe("admin Test of an entry without a URL", () => { await test({ provider: "deepseek", apiKey: "k" }) expect(sent.body.baseUrl).toBeUndefined() }) + + it("does not use Vertex's variable, which the panel writes itself", async () => { + // Before a save it still holds the entry's previous URL + process.env.GOOGLE_VERTEX_BASE_URL = "https://old-proxy.example.com" + try { + await test({ provider: "vertexai", vertexApiKey: "new-key" }) + expect(sent.body.baseUrl).toBeUndefined() + } finally { + delete process.env.GOOGLE_VERTEX_BASE_URL + } + }) }) diff --git a/tests/unit/ai-providers-credentials.test.ts b/tests/unit/ai-providers-credentials.test.ts index df155cdb..d74ed915 100644 --- a/tests/unit/ai-providers-credentials.test.ts +++ b/tests/unit/ai-providers-credentials.test.ts @@ -28,6 +28,14 @@ vi.mock("@ai-sdk/openai", () => { } }) +vi.mock("@ai-sdk/azure", () => { + const mockModel = { modelId: "test-model" } + const mockProviderFn = vi.fn(() => mockModel) as any + mockProviderFn.chat = vi.fn(() => mockModel) + mockProviderFn.responses = vi.fn(() => mockModel) + return { createAzure: vi.fn(() => mockProviderFn) } +}) + vi.mock("@ai-sdk/amazon-bedrock", () => { const mockProviderFn = vi.fn(() => ({ modelId: "test-model" })) return { createAmazonBedrock: vi.fn(() => mockProviderFn) } @@ -446,6 +454,43 @@ describe("whose keys a request uses", () => { } }) + it("runs an Azure entry set up only in the admin panel", async () => { + // No AZURE_BASE_URL or AZURE_RESOURCE_NAME: the entry's own + // variables hold the key and the resource URL + process.env.ADMIN_AZURE_API_KEY = "panel-key" + process.env.ADMIN_AZURE_BASE_URL = "https://res.openai.azure.com/openai" + try { + const { createAzure } = await import("@ai-sdk/azure") + expect(() => + getAIModel({ + provider: "azure", + modelId: "gpt-4o", + apiKeyEnv: "ADMIN_AZURE_API_KEY", + baseUrlEnv: "ADMIN_AZURE_BASE_URL", + }), + ).not.toThrow() + expect(createAzure).toHaveBeenLastCalledWith( + expect.objectContaining({ + apiKey: "panel-key", + baseURL: "https://res.openai.azure.com/openai", + }), + ) + // Without any URL it still says what is missing + delete process.env.ADMIN_AZURE_BASE_URL + expect(() => + getAIModel({ + provider: "azure", + modelId: "gpt-4o", + apiKeyEnv: "ADMIN_AZURE_API_KEY", + baseUrlEnv: "ADMIN_AZURE_BASE_URL", + }), + ).toThrow(/AZURE_BASE_URL/) + } finally { + delete process.env.ADMIN_AZURE_API_KEY + delete process.env.ADMIN_AZURE_BASE_URL + } + }) + it("needs a base URL with a user's Azure key", () => { // The SDK would otherwise read the server's AZURE_RESOURCE_NAME process.env.AZURE_RESOURCE_NAME = "operator-resource" diff --git a/tests/unit/app-menu.test.ts b/tests/unit/app-menu.test.ts index c340fbf2..8ab3caac 100644 --- a/tests/unit/app-menu.test.ts +++ b/tests/unit/app-menu.test.ts @@ -64,4 +64,20 @@ describe("switchPreset", () => { await toC expect(state.current).toBe("C") }) + + it("keeps a newer choice of the same preset", async () => { + // A, then B, C, and B again while the first restart is pending + const first = switchPreset("B").catch(() => {}) + const second = switchPreset("C").catch(() => {}) + const third = switchPreset("B") + // The first restart fails: the current preset is B again, but it is + // the third switch's, which must not be undone + state.restarts[0].reject(new Error("timed out")) + await new Promise((r) => setTimeout(r, 0)) + for (const r of state.restarts.slice(1)) r.resolve() + await first + await second + await third + expect(state.current).toBe("B") + }) }) diff --git a/tests/unit/chat-message-display-preview.test.tsx b/tests/unit/chat-message-display-preview.test.tsx new file mode 100644 index 00000000..a837ef6f --- /dev/null +++ b/tests/unit/chat-message-display-preview.test.tsx @@ -0,0 +1,68 @@ +import { render } from "@testing-library/react" +import { describe, expect, it, vi } from "vitest" +import en from "@/lib/i18n/dictionaries/en.json" + +const page = (cells: string) => + `${cells}` +const box = (id: string) => + `` + +// The first edit's result is loaded (the ref has it); the chartXML state +// has not caught up yet +const BEFORE_FIRST_EDIT = page(box("a")) +const AFTER_FIRST_EDIT = page(box("a") + box("b")) + +vi.mock("@/contexts/diagram-context", () => ({ + useDiagram: () => ({ + chartXML: BEFORE_FIRST_EDIT, + chartXMLRef: { current: AFTER_FIRST_EDIT }, + loadDiagram: vi.fn(() => null), + }), +})) +vi.mock("@/hooks/use-dictionary", () => ({ useDictionary: () => en })) + +import { ChatMessageDisplay } from "@/components/chat-message-display" + +// jsdom has no layout +Element.prototype.scrollIntoView = () => {} + +describe("the streaming preview of a second edit", () => { + it("starts from the first edit's result", () => { + const editDiagramOriginalXmlRef = { current: new Map() } + const messages = [ + { + id: "m1", + role: "assistant", + parts: [ + { + type: "tool-edit_diagram", + toolCallId: "edit-2", + state: "input-streaming", + input: { + operations: [ + { + operation: "add", + cell_id: "c", + new_xml: box("c"), + }, + ], + }, + }, + ], + }, + ] as any + render( + {}} + setFiles={() => {}} + processedToolCallsRef={{ current: new Set() }} + editDiagramOriginalXmlRef={editDiagramOriginalXmlRef} + status="streaming" + />, + ) + expect(editDiagramOriginalXmlRef.current.get("edit-2")).toBe( + AFTER_FIRST_EDIT, + ) + }) +}) diff --git a/tests/unit/chat-route-abort.test.ts b/tests/unit/chat-route-abort.test.ts new file mode 100644 index 00000000..5090e15b --- /dev/null +++ b/tests/unit/chat-route-abort.test.ts @@ -0,0 +1,139 @@ +// @vitest-environment node +import { afterEach, describe, expect, it, vi } from "vitest" + +const quota = vi.hoisted(() => ({ recorded: [] as number[] })) +vi.mock("@/lib/dynamo-quota-manager", () => ({ + isQuotaEnabled: () => true, + checkAndIncrementRequest: async () => ({ allowed: true }), + recordTokenUsage: async (_ip: string, tokens: number) => { + quota.recorded.push(tokens) + }, +})) + +// No DNS in tests: only loopback addresses are private +vi.mock("@/lib/ssrf-protection", async (importOriginal) => ({ + ...(await importOriginal()), + isPrivateUrl: async (url: string) => + /^https?:\/\/(127\.0\.0\.1|localhost)\b/.test(url), +})) + +import { POST as chat } from "@/app/api/chat/route" + +afterEach(() => { + quota.recorded = [] + vi.unstubAllGlobals() +}) + +const sse = (chunks: object[], end = true) => + chunks.map((c) => `data: ${JSON.stringify(c)}\n\n`).join("") + + (end ? "data: [DONE]\n\n" : "") + +describe("a request stopped after a finished step", () => { + it("counts that step's tokens", async () => { + // Step 1 asks for a shape library (run on the server) and reports + // its usage; step 2 never ends, and the user stops + let call = 0 + vi.stubGlobal( + "fetch", + vi.fn(async (_url: string, init?: RequestInit) => { + call++ + if (call === 1) { + return new Response( + sse([ + { + id: "c1", + choices: [ + { + index: 0, + delta: { + role: "assistant", + tool_calls: [ + { + index: 0, + id: "call_1", + type: "function", + function: { + name: "get_shape_library", + arguments: + '{"library":"aws4"}', + }, + }, + ], + }, + finish_reason: null, + }, + ], + }, + { + id: "c1", + choices: [ + { + index: 0, + delta: {}, + finish_reason: "tool_calls", + }, + ], + usage: { + prompt_tokens: 1200, + completion_tokens: 30, + }, + }, + ]), + { headers: { "content-type": "text/event-stream" } }, + ) + } + // Never ends, until the request is aborted (as fetch does) + const body = new ReadableStream({ + start(controller) { + init?.signal?.addEventListener("abort", () => + controller.error( + new DOMException("aborted", "AbortError"), + ), + ) + }, + }) + return new Response(body, { + headers: { "content-type": "text/event-stream" }, + }) + }), + ) + const stop = new AbortController() + const res = await chat( + new Request("http://localhost/api/chat", { + method: "POST", + signal: stop.signal, + headers: { + "Content-Type": "application/json", + "x-forwarded-for": "203.0.113.7", + // The server's own network: counted + "x-ai-provider": "glm", + "x-ai-base-url": "http://127.0.0.1:9000/v1", + "x-ai-api-key": "dummy", + "x-ai-model": "glm-5", + }, + body: JSON.stringify({ + messages: [ + { + id: "u1", + role: "user", + parts: [{ type: "text", text: "Draw AWS" }], + }, + ], + xml: "", + }), + }), + ) + const reader = res.body?.getReader() + // Read until the second step has started + await vi.waitFor(() => expect(call).toBe(2), { timeout: 3000 }) + stop.abort() + // The answer stream ends; the SDK handles the stop as it is read + while ( + reader && + !(await reader.read().catch(() => ({ done: true }))).done + ) { + // drain + } + await vi.waitFor(() => expect(quota.recorded).toEqual([1230])) + }) +}) diff --git a/tests/unit/chat-route-edgeone.test.ts b/tests/unit/chat-route-edgeone.test.ts index 09b7f39e..d77f2440 100644 --- a/tests/unit/chat-route-edgeone.test.ts +++ b/tests/unit/chat-route-edgeone.test.ts @@ -77,3 +77,63 @@ describe("EdgeOne as a server model", () => { expect(calls[0]?.headers.get("cookie")).toBe("eo_token=t; eo_time=1") }) }) + +const send = (headers: Record) => + chat( + new Request("http://localhost/api/chat", { + method: "POST", + headers: { "Content-Type": "application/json", ...headers }, + body: JSON.stringify({ + messages: [ + { + id: "u1", + role: "user", + parts: [{ type: "text", text: "Draw two boxes" }], + }, + ], + xml: "", + }), + }), + ).then((r) => r.text()) + +describe("EdgeOne endpoints", () => { + it("works when the deployment names EdgeOne only in AI_PROVIDER", async () => { + process.env.AI_PROVIDER = "edgeone" + process.env.AI_MODEL = "@tx/deepseek-ai/deepseek-v3-0324" + await send({}) + expect(calls[0]?.url).toBe( + "http://localhost/api/edgeai/chat/completions", + ) + }) + + it("always calls the site's own function, whatever URL the request names", async () => { + // Another host would get the user's EdgeOne cookies + await send({ + "x-ai-provider": "edgeone", + "x-ai-model": "@tx/deepseek-ai/deepseek-v3-0324", + "x-ai-base-url": "https://elsewhere.example/api/edgeai", + cookie: "eo_token=t", + }) + expect(calls[0]?.url).toBe( + "http://localhost/api/edgeai/chat/completions", + ) + }) + + it("keeps the deployment's base path", async () => { + const savedPath = process.env.NEXT_PUBLIC_BASE_PATH + process.env.NEXT_PUBLIC_BASE_PATH = "/draw" + try { + await send({ + "x-ai-provider": "edgeone", + "x-ai-model": "@tx/deepseek-ai/deepseek-v3-0324", + }) + expect(calls[0]?.url).toBe( + "http://localhost/draw/api/edgeai/chat/completions", + ) + } finally { + if (savedPath === undefined) + delete process.env.NEXT_PUBLIC_BASE_PATH + else process.env.NEXT_PUBLIC_BASE_PATH = savedPath + } + }) +}) diff --git a/tests/unit/chat-route-errors.test.ts b/tests/unit/chat-route-errors.test.ts index e6eee6c8..e89e6b94 100644 --- a/tests/unit/chat-route-errors.test.ts +++ b/tests/unit/chat-route-errors.test.ts @@ -89,6 +89,24 @@ describe("provider error texts in the stream", () => { // The SDK retries a refused connection twice, waiting between }, 20_000) + it("shows the server's keyless Ollama error on the web too", async () => { + // No key, no money involved; round three hid this text + process.env.AI_PROVIDER = "ollama" + process.env.AI_MODEL = "llama3" + process.env.OLLAMA_BASE_URL = "http://ollama.internal:11434/api" + vi.stubGlobal( + "fetch", + vi.fn(async () => { + throw Object.assign(new TypeError("fetch failed"), { + cause: new Error("connect ECONNREFUSED 10.0.0.9:11434"), + }) + }), + ) + expect(await streamedError({})).not.toBe( + "The provider returned an error.", + ) + }, 20_000) + it("shows EdgeOne's own daily quota explanation", async () => { // The function answers 429, which the SDK retries with a wait; // the status does not decide whether the text is shown @@ -124,3 +142,107 @@ describe("provider error texts in the stream", () => { expect(message).toBe("The provider returned an error.") }) }) + +describe("the output cap", () => { + it("holds for the server's own keyless endpoints too", async () => { + process.env.MAX_OUTPUT_TOKENS = "8000" + const sent: string[] = [] + vi.stubGlobal( + "fetch", + vi.fn(async (_url: string, init?: RequestInit) => { + sent.push(String(init?.body ?? "")) + return new Response("{}", { status: 400 }) + }), + ) + try { + const res = await chat( + new Request("http://localhost/api/chat", { + method: "POST", + headers: { + "Content-Type": "application/json", + "x-ai-provider": "ollama", + "x-ai-base-url": "http://127.0.0.1:11434/api", + "x-ai-model": "llama3", + "x-max-output-tokens": "200000", + }, + body: JSON.stringify({ + messages: [ + { + id: "u1", + role: "user", + parts: [{ type: "text", text: "Draw" }], + }, + ], + xml: "", + }), + }), + ) + await res.text() + expect(sent[0]).toContain('"max_output_tokens":8000') + } finally { + delete process.env.MAX_OUTPUT_TOKENS + } + }) +}) + +describe("a tool call that never got its result", () => { + it("is left out of the prompt instead of failing every later message", async () => { + // Stop while the screenshot check ran left display_diagram without + // a result, and the chat was saved like that + process.env.AI_PROVIDER = "openai" + process.env.AI_MODEL = "gpt-5.5" + process.env.OPENAI_API_KEY = "server-key" + const sent: string[] = [] + vi.stubGlobal( + "fetch", + vi.fn(async (_url: string, init?: RequestInit) => { + sent.push(String(init?.body ?? "")) + return new Response( + JSON.stringify({ error: { message: "x" } }), + { + status: 400, + headers: { "Content-Type": "application/json" }, + }, + ) + }), + ) + const res = await chat( + new Request("http://localhost/api/chat", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + messages: [ + { + id: "u1", + role: "user", + parts: [{ type: "text", text: "Draw a box" }], + }, + { + id: "a1", + role: "assistant", + parts: [ + { + type: "tool-display_diagram", + toolCallId: "call-without-result", + state: "input-available", + input: { xml: "" }, + }, + ], + }, + { + id: "u2", + role: "user", + parts: [{ type: "text", text: "Make it red" }], + }, + ], + xml: "", + }), + }), + ) + await res.text() + // The request reached the model, without the unanswered call + expect(sent).toHaveLength(1) + expect(sent[0]).toContain("Make it red") + expect(sent[0]).not.toContain("call-without-result") + }) +}) diff --git a/tests/unit/chat-route-quota.test.ts b/tests/unit/chat-route-quota.test.ts index fd3e6ad4..74933b8f 100644 --- a/tests/unit/chat-route-quota.test.ts +++ b/tests/unit/chat-route-quota.test.ts @@ -140,6 +140,67 @@ describe("chat quota", () => { expect(quota.checks).toBe(1) }) + it("counts the server's network whatever key header comes along", async () => { + // A keyless Ollama or a local SGLang ignores a dummy key + for (const headers of [ + { + "x-ai-provider": "ollama", + "x-ai-base-url": "http://127.0.0.1:11434/api", + "x-ai-api-key": "dummy", + "x-ai-model": "llama3.2", + }, + { + "x-ai-provider": "openai", + "x-ai-base-url": "http://127.0.0.1:30000/v1", + "x-ai-api-key": "dummy", + "x-ai-model": "m", + }, + ]) { + expect((await send(headers)).status).toBe(429) + } + expect(quota.checks).toBe(2) + }) + + it("does not count a provider that never uses the base URL header", async () => { + // Bedrock on the user's own AWS keys goes to AWS, whatever the + // leftover base URL says + const res = await send({ + "x-ai-provider": "bedrock", + "x-ai-model": "amazon.nova-lite-v1:0", + "x-ai-base-url": "http://127.0.0.1:8080", + "x-aws-access-key-id": "id", + "x-aws-secret-access-key": "secret", + "x-aws-region": "us-east-1", + }) + expect(res.status).not.toBe(429) + expect(quota.checks).toBe(0) + }) + + it("counts EdgeOne even with a base URL header", async () => { + const res = await send({ + "x-ai-provider": "edgeone", + "x-ai-model": "@tx/deepseek-ai/deepseek-v3-0324", + "x-ai-base-url": "https://this-site.example/api/edgeai", + }) + expect(res.status).toBe(429) + expect(quota.checks).toBe(1) + }) + + it("never counts in the desktop app, where every endpoint is the user's", async () => { + process.env.NEXT_AI_DRAWIO_DESKTOP = "1" + try { + const res = await send({ + "x-ai-provider": "ollama", + "x-ai-base-url": "http://127.0.0.1:11434/api", + "x-ai-model": "llama3.2", + }) + expect(res.status).not.toBe(429) + expect(quota.checks).toBe(0) + } finally { + delete process.env.NEXT_AI_DRAWIO_DESKTOP + } + }) + it("does not count Ollama on the user's own server", async () => { const res = await send({ "x-ai-provider": "ollama", diff --git a/tests/unit/diagram-context.test.tsx b/tests/unit/diagram-context.test.tsx new file mode 100644 index 00000000..3b69cd27 --- /dev/null +++ b/tests/unit/diagram-context.test.tsx @@ -0,0 +1,119 @@ +import { deflateRawSync } from "node:zlib" +import { act, renderHook } from "@testing-library/react" +import type React from "react" +import { afterEach, describe, expect, it, vi } from "vitest" +import { DiagramProvider, useDiagram } from "@/contexts/diagram-context" + +vi.mock("sonner", () => ({ toast: { success: vi.fn() } })) + +// The provider with a stand-in draw.io that records each export request +function setup() { + const { result } = renderHook(() => useDiagram(), { + wrapper: ({ children }: { children: React.ReactNode }) => ( + {children} + ), + }) + const requests: { format: string; message: string }[] = [] + result.current.drawioRef.current = { + exportDiagram: (r: any) => requests.push(r), + load: vi.fn(), + } as any + // draw.io's reply to a request: it echoes the request in `message` + const reply = (request: { message: string }, data: string, xml = "") => + act(() => + result.current.handleDiagramExport({ + event: "export", + data, + xml, + format: "xmlsvg", + message: request, + } as any), + ) + return { result, requests, reply } +} + +// An editable SVG as draw.io exports it: the diagram, compressed, in its +// content attribute +const svgOf = (label: string) => { + const model = `` + const packed = deflateRawSync( + Buffer.from(encodeURIComponent(model)), + ).toString("base64") + const content = `${packed}` + .replaceAll("&", "&") + .replaceAll("<", "<") + .replaceAll(">", ">") + .replaceAll('"', """) + const svg = `` + return `data:image/svg+xml;base64,${btoa(svg)}` +} + +afterEach(() => { + vi.restoreAllMocks() +}) + +describe("exports in flight at the same time", () => { + it("give the chat's export only its own reply", () => { + const { result, requests, reply } = setup() + // An edit's history export is still on its way when the chat exports + act(() => { + result.current.handleExport() + }) + let tag = "" + const got: string[] = [] + act(() => { + tag = result.current.handleExportWithoutHistory() + result.current.exportResolversRef.current[tag] = (xml) => + got.push(xml) + }) + reply(requests[0], svgOf("older")) + expect(got).toEqual([]) + reply(requests[1], svgOf("current")) + expect(got).toHaveLength(1) + expect(got[0]).toContain('value="current"') + expect(result.current.exportResolversRef.current[tag]).toBeUndefined() + }) + + it("save each file with its own result", async () => { + const { result, requests, reply } = setup() + const saved: { name: string; href: string }[] = [] + vi.spyOn(HTMLAnchorElement.prototype, "click").mockImplementation( + function (this: HTMLAnchorElement) { + saved.push({ name: this.download, href: this.href }) + }, + ) + const blobs = new Map() + URL.createObjectURL = vi.fn((blob: Blob) => { + const url = `blob:test-${blobs.size}` + blobs.set(url, blob) + return url + }) + URL.revokeObjectURL = vi.fn() + vi.stubGlobal( + "fetch", + vi.fn(async () => new Response("{}")), + ) + + const twoPages = + '' + act(() => { + result.current.saveDiagramToFile("doc", "drawio") + result.current.saveDiagramToFile("pic", "png") + }) + // The PNG answers first + reply(requests[1], "data:image/png;base64,iVBORw0KGgo=") + reply(requests[0], svgOf("doc"), twoPages) + + expect(saved.map((s) => s.name)).toEqual(["pic.png", "doc.drawio"]) + expect(saved[0].href).toMatch(/^data:image\/png/) + const file = blobs.get(saved[1].href) + const text = await new Promise((resolve) => { + const reader = new FileReader() + reader.onload = () => resolve(String(reader.result)) + reader.readAsText(file as Blob) + }) + expect(text).toContain('name="A"') + expect(text).toContain('name="B"') + vi.unstubAllGlobals() + }) +}) diff --git a/tests/unit/env-loader.test.ts b/tests/unit/env-loader.test.ts index 51337175..2cffdeed 100644 --- a/tests/unit/env-loader.test.ts +++ b/tests/unit/env-loader.test.ts @@ -25,6 +25,9 @@ const KEYS = [ "T_HASH", "T_AFTER", "T_JOINED", + "T_ESC_HASH", + "T_ESC_INNER", + "T_ESC_COMMENT", ] afterEach(() => { for (const k of KEYS) delete process.env[k] @@ -72,4 +75,21 @@ describe("loadEnvFile", () => { expect(process.env.T_JOINED).toBe(`"a"b`) expect(process.env.T_HASH).toBe("http://host/#/x") }) + + it("does not end a quoted value at an escaped quote, like dotenv", () => { + dir.path = mkdtempSync(join(tmpdir(), "env-loader-")) + writeFileSync( + join(dir.path, ".env"), + [ + 'T_ESC_HASH="abc\\" #def"', + 'T_ESC_INNER="a # \\"b\\""', + 'T_ESC_COMMENT="x\\"y" # c', + ].join("\n"), + ) + loadEnvFile() + // Expected values from dotenv 16.6.1, which keeps the backslashes + expect(process.env.T_ESC_HASH).toBe('abc\\" #def') + expect(process.env.T_ESC_INNER).toBe('a # \\"b\\"') + expect(process.env.T_ESC_COMMENT).toBe('x\\"y') + }) }) diff --git a/tests/unit/mcp-preview-recovery.test.ts b/tests/unit/mcp-preview-recovery.test.ts new file mode 100644 index 00000000..379c0090 --- /dev/null +++ b/tests/unit/mcp-preview-recovery.test.ts @@ -0,0 +1,372 @@ +import { readFileSync } from "node:fs" +import { join } from "node:path" +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest" + +// The MCP preview page's script, run in this document with a stubbed +// server (fetch) and draw.io iframe (its postMessage), so the tab's side of +// a recreated session can be driven step by step +const dir = join(process.cwd(), "packages/mcp-server/src/preview") +const DRAWIO = "https://embed.diagrams.net" +const html = readFileSync(join(dir, "index.html"), "utf8") + .replace("{{CSS}}", "") + .replace("{{SESSION_BADGE}}", "") + .replaceAll("{{DISABLED}}", "") + .replace("{{DRAWIO_URL}}", "about:blank") + .replace("{{SESSION_JSON}}", '"mcp-test"') + .replace("{{ORIGIN_JSON}}", JSON.stringify(DRAWIO)) +const scripts = [...html.matchAll(/