diff --git a/app/api/admin/test-model/route.ts b/app/api/admin/test-model/route.ts index c43a98ab..6d9aff98 100644 --- a/app/api/admin/test-model/route.ts +++ b/app/api/admin/test-model/route.ts @@ -5,6 +5,7 @@ import { loadAdminProviders, mergeSecrets, } from "@/lib/admin/providers" +import { globalBaseUrl } from "@/lib/ai-providers" export const runtime = "nodejs" export const dynamic = "force-dynamic" @@ -58,7 +59,10 @@ export async function POST(req: Request) { body: JSON.stringify({ provider: resolved.provider, apiKey: resolved.apiKey, - baseUrl: resolved.baseUrl, + // Without a URL of its own, chat sends the entry's key to + // the server's

_BASE_URL: test that endpoint, not + // another one + baseUrl: resolved.baseUrl || globalBaseUrl(resolved.provider), modelId: body.modelId, awsAccessKeyId: resolved.awsAccessKeyId, awsSecretAccessKey: resolved.awsSecretAccessKey, diff --git a/app/api/chat/route.ts b/app/api/chat/route.ts index a50f708b..ccc0b346 100644 --- a/app/api/chat/route.ts +++ b/app/api/chat/route.ts @@ -738,9 +738,10 @@ Call this tool to get shape names and usage syntax for a specific library.`, const response = result.toUIMessageStreamResponse({ sendReasoning: true, - // The provider's text can name the server's account or hosts - onError: (error) => - streamErrorText(error, onServerCredentials || onServerEndpoint), + // On the server's keys the provider's text can name its account. + // Keyless endpoints keep theirs: the desktop app's Ollama is the + // user's own, and EdgeOne's text is our function's explanation. + onError: (error) => streamErrorText(error, onServerCredentials), messageMetadata: ({ part }) => { if (part.type === "finish") { const usage = (part as any).totalUsage diff --git a/app/api/provider-models/route.ts b/app/api/provider-models/route.ts index 0b6cbb8a..6b362e73 100644 --- a/app/api/provider-models/route.ts +++ b/app/api/provider-models/route.ts @@ -9,6 +9,7 @@ import { import { allowPrivateUrls, isPrivateUrl, + RedirectRefusedError, redirectGuardedFetch, } from "@/lib/ssrf-protection" import type { ProviderName } from "@/lib/types/model-config" @@ -65,13 +66,10 @@ export async function POST(req: Request) { // Only our own explanations go back: the URL may be an internal // address, whose answer or host names must not reach the caller. // The Gateway SDK wraps them, keeping ours as the cause. + const isOwn = (e: unknown): e is Error => + e instanceof ModelListError || e instanceof RedirectRefusedError const cause = (error as { cause?: unknown })?.cause - const own = - error instanceof ModelListError - ? error - : cause instanceof ModelListError - ? cause - : null + const own = isOwn(error) ? error : isOwn(cause) ? cause : null const { code } = classifyLLMError(own ?? error) return NextResponse.json({ code, diff --git a/components/model-config-dialog.tsx b/components/model-config-dialog.tsx index 95cebc6a..a13c95b7 100644 --- a/components/model-config-dialog.tsx +++ b/components/model-config-dialog.tsx @@ -189,6 +189,9 @@ export function ModelConfigDialog({ selectedProviderIdRef.current = selectedProviderId const configRef = useRef(config) configRef.current = config + // Number of the latest Test click: only that test may reset the busy + // state when its credentials changed meanwhile + const validationRunRef = useRef(0) // A model list or test result belongs to the credentials it was asked // with; they can change meanwhile, here or in another tab const credentialsOf = (providerId: string) => { @@ -414,6 +417,7 @@ export function ModelConfigDialog({ let errorCount = 0 let idChanged = false const askedWith = credentialsOf(selectedProviderId) + const run = ++validationRunRef.current // For EdgeOne, construct baseUrl from current origin const baseUrl = isEdgeOne @@ -488,8 +492,21 @@ export function ModelConfigDialog({ validationWarning: undefined, } } - // Credentials changed during the test: drop the result - if (credentialsOf(selectedProviderId) !== askedWith) 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). + if (credentialsOf(selectedProviderId) !== askedWith) { + if (run === validationRunRef.current) { + 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 const current = configRef.current.providers .find((p) => p.id === selectedProviderId) @@ -515,7 +532,17 @@ export function ModelConfigDialog({ }) }), ) - if (credentialsOf(selectedProviderId) !== askedWith) 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 + ) { + setValidationStatus("idle") + } + return + } // A model whose id changed was not tested if (allValid && !idChanged) { diff --git a/electron/electron.d.ts b/electron/electron.d.ts index c10de5e3..4c54bdcc 100644 --- a/electron/electron.d.ts +++ b/electron/electron.d.ts @@ -75,6 +75,10 @@ declare global { * preset); returns a function that stops the calls */ onServerRestarted?: (callback: () => void) => () => void + /** A chat was saved: open this port next launch */ + chatSaved?: () => Promise + /** The page loaded with this many chats */ + chatsLoaded?: (count: number) => Promise } /** Settings window Electron API */ diff --git a/electron/main/env-loader.ts b/electron/main/env-loader.ts index 0d63eb69..ddf59085 100644 --- a/electron/main/env-loader.ts +++ b/electron/main/env-loader.ts @@ -51,17 +51,25 @@ function loadEnvFromFile(filePath: string): void { const quote = value[0] const closingQuote = quote === '"' || quote === "'" ? value.indexOf(quote, 1) : -1 - if (value.length > 1 && value.endsWith(quote) && closingQuote > 0) { - // Quoted from start to end: the quotes inside belong to the - // value (JSON with an apostrophe), as dotenv reads it - value = value.slice(1, -1) - } else if (closingQuote > 0) { - // Quoted value: keep what's inside the quotes and drop - // anything after them (e.g. a comment) + if ( + closingQuote > 0 && + /^\s*(#.*)?$/.test(value.slice(closingQuote + 1)) + ) { + // Quoted value, then nothing or a comment: keep what is + // inside the quotes, as dotenv reads it value = value.slice(1, closingQuote) } else { - // Unquoted value: drop an inline comment ("value # comment") + // Unquoted value: drop an inline comment ("value # comment"). + // A value quoted from start to end with quotes inside (JSON + // with an apostrophe) loses only the outer two, as in dotenv. value = value.replace(/\s+#.*$/, "") + if ( + closingQuote > 0 && + value.length > 1 && + value.endsWith(quote) + ) { + value = value.slice(1, -1) + } } // Don't override existing environment variables diff --git a/electron/main/ipc-handlers.ts b/electron/main/ipc-handlers.ts index 1a7458c1..8ccddd60 100644 --- a/electron/main/ipc-handlers.ts +++ b/electron/main/ipc-handlers.ts @@ -13,6 +13,7 @@ import { updatePreset, } from "./config-manager" import { restartNextServer } from "./next-server" +import { noteNoChats, rememberChatPort } from "./port-manager" import { applyProxyToEnv, getProxyConfig, @@ -77,6 +78,15 @@ export function registerIpcHandlers(): void { return app.getVersion() }) + // ==================== Where the chats are ==================== + + // The page saved a chat, or loaded without any: decides which port + // (and so which origin's chats) the next launch opens + handle("chat-saved", () => rememberChatPort()) + handle("chats-loaded", (_event, count: unknown) => { + if (count === 0) noteNoChats() + }) + // ==================== Window Controls ==================== ipcMain.on("window-minimize", (event) => { diff --git a/electron/main/port-manager.ts b/electron/main/port-manager.ts index 623c65bb..509bc9ab 100644 --- a/electron/main/port-manager.ts +++ b/electron/main/port-manager.ts @@ -1,4 +1,4 @@ -import { existsSync } from "node:fs" +import { existsSync, readFileSync, writeFileSync } from "node:fs" import net from "node:net" import path from "node:path" import { app } from "electron" @@ -39,6 +39,55 @@ function hasStoredData(port: number): boolean { ) } +// The two fixed production ports, the only ones whose origin (and so its +// chats and settings) is the same at every launch +const HOME_PORTS = [PORT_CONFIG.legacyProduction, PORT_CONFIG.production] + +const chatPortFile = () => path.join(app.getPath("userData"), "chat-port.json") + +/** The fixed port where a chat was last saved, if known */ +function readChatPort(): number | null { + try { + const { port } = JSON.parse(readFileSync(chatPortFile(), "utf-8")) + return HOME_PORTS.includes(port) ? port : null + } catch { + return null + } +} + +function writeChatPort(port: number): void { + try { + writeFileSync(chatPortFile(), JSON.stringify({ port })) + } catch (error) { + console.warn("Could not save the chat port:", error) + } +} + +/** + * The page saved a chat: open on this port next time. Chats of the two + * ports cannot be shown together (each origin has its own storage), so the + * app opens where the user last worked. A launch that had to use the other + * port and saved nothing does not move it. + */ +export function rememberChatPort(): void { + const port = allocatedPort + if (!app.isPackaged || port === null || !HOME_PORTS.includes(port)) return + if (readChatPort() !== port) writeChatPort(port) +} + +/** + * The page loaded without any chats. Before any chat was saved under this + * version (no file yet), the user's chats may be on the other fixed port, + * where an older version opened: try it first next time. + */ +export function noteNoChats(): void { + const port = allocatedPort + if (!app.isPackaged || port === null || !HOME_PORTS.includes(port)) return + if (existsSync(chatPortFile())) return + const other = HOME_PORTS.find((p) => p !== port) + if (other !== undefined && hasStoredData(other)) writeChatPort(other) +} + /** * Check if a specific port is available */ @@ -86,15 +135,19 @@ export async function findAvailablePort(reuseExisting = true): Promise { allocatedPort = null } - // In production, try the legacy port first to preserve existing users' - // data, unless only the new port has data: their app started on 13370 - // while Windows reserved 61337, and 61337 being free now would hide it + // In production, first the port where a chat was last saved. Without + // one, the legacy port first to preserve existing users' data, unless + // only the new port has data: their app started on 13370 while Windows + // reserved 61337, and 61337 being free now would hide it + const chatPort = isDev ? null : readChatPort() const candidates = isDev ? [preferredPort] - : hasStoredData(PORT_CONFIG.production) && - !hasStoredData(PORT_CONFIG.legacyProduction) - ? [PORT_CONFIG.production, PORT_CONFIG.legacyProduction] - : [PORT_CONFIG.legacyProduction, PORT_CONFIG.production] + : chatPort !== null + ? [chatPort, ...HOME_PORTS.filter((p) => p !== chatPort)] + : hasStoredData(PORT_CONFIG.production) && + !hasStoredData(PORT_CONFIG.legacyProduction) + ? [PORT_CONFIG.production, PORT_CONFIG.legacyProduction] + : [PORT_CONFIG.legacyProduction, PORT_CONFIG.production] for (const port of candidates) { if (await isPortAvailable(port)) { allocatedPort = port diff --git a/electron/preload/index.ts b/electron/preload/index.ts index f406fd32..f24249f7 100644 --- a/electron/preload/index.ts +++ b/electron/preload/index.ts @@ -28,6 +28,11 @@ contextBridge.exposeInMainWorld("electronAPI", { setUserLocale: (locale: string) => ipcRenderer.invoke("set-user-locale", locale), + // A chat was saved, or the page loaded with this many chats: the next + // launch opens the port where the chats are + chatSaved: () => ipcRenderer.invoke("chat-saved"), + chatsLoaded: (count: number) => ipcRenderer.invoke("chats-loaded", count), + // The server restarted on the same port (another preset) onServerRestarted: (callback: () => void) => { const listener = () => callback() diff --git a/env.example b/env.example index d800f120..bf22cff9 100644 --- a/env.example +++ b/env.example @@ -70,7 +70,7 @@ AI_MODEL=global.anthropic.claude-sonnet-4-5-20250929-v1:0 # AZURE_REASONING_SUMMARY=detailed # Ollama Configuration (Local or Cloud) -# OLLAMA_BASE_URL=https://ollama.com/api # Optional, defaults to Ollama Cloud +# OLLAMA_BASE_URL=https://ollama.com/api # Optional: Ollama Cloud; defaults to local Ollama (http://127.0.0.1:11434) # OLLAMA_API_KEY=your-ollama-cloud-api-key # Optional: For Ollama Cloud or authenticated remote instances # OLLAMA_ENABLE_THINKING=true # Optional: Enable thinking for models that support it (e.g., qwen3) diff --git a/hooks/use-diagram-tool-handlers.ts b/hooks/use-diagram-tool-handlers.ts index f1ae4860..66dbd7f3 100644 --- a/hooks/use-diagram-tool-handlers.ts +++ b/hooks/use-diagram-tool-handlers.ts @@ -122,34 +122,33 @@ export function useDiagramToolHandlers({ } processedToolCallsRef.current.add(toolCall.toolCallId) - // Only these two put their result on the canvas. Other tools - // (get_shape_library, which the server runs, still arrives here) - // leave the stored originals for the preview code to undo. - const drawsDiagram = - toolCall.toolName === "display_diagram" || - toolCall.toolName === "edit_diagram" - // Stored originals belong to previews not handled yet: this call's, - // and those of earlier calls with invalid input, which never get - // here. The first is the diagram before all of them. This call's - // result replaces those previews, so the preview code must neither - // draw them again nor undo them later. - const [originalXml] = editDiagramOriginalXmlRef.current.values() - if (drawsDiagram) { - for (const id of editDiagramOriginalXmlRef.current.keys()) { - processedToolCallsRef.current.add(id) - } - editDiagramOriginalXmlRef.current.clear() - } - + // Only display_diagram, edit_diagram and a completing append_diagram + // put their result on the canvas. Other tools (get_shape_library, + // which the server runs, still arrives here) leave the stored + // originals for the preview code to undo. if (toolCall.toolName === "display_diagram") { - await handleDisplayDiagram(toolCall, addToolOutput, originalXml) + await handleDisplayDiagram(toolCall, addToolOutput, takeOriginals()) } else if (toolCall.toolName === "edit_diagram") { - await handleEditDiagram(toolCall, addToolOutput, originalXml) + await handleEditDiagram(toolCall, addToolOutput, takeOriginals()) } else if (toolCall.toolName === "append_diagram") { handleAppendDiagram(toolCall, addToolOutput) } } + // Stored originals belong to previews not handled yet: this call's, and + // those of earlier calls with invalid input, which never get to the + // handler. The first is the diagram before all of them. A call that + // draws its result replaces those previews, so the preview code must + // neither draw them again nor undo them later. Returns that first one. + const takeOriginals = (): string | undefined => { + const [originalXml] = editDiagramOriginalXmlRef.current.values() + for (const id of editDiagramOriginalXmlRef.current.keys()) { + processedToolCallsRef.current.add(id) + } + editDiagramOriginalXmlRef.current.clear() + return originalXml + } + // originalXml: the diagram before the streamed previews, if any were drawn const handleDisplayDiagram = async ( toolCall: ToolCall, @@ -540,11 +539,17 @@ Start your continuation with the NEXT character after where it stopped.`, partialXmlRef.current = "" // Reset const prepared = prepareNewDiagram(finalXml, NEW_PAGE) + // It draws now: it takes the stored originals, as display_diagram + const originalXml = prepared.ok ? takeOriginals() : undefined const validationError = prepared.ok ? onDisplayChart(prepared.xml, true) : prepared.error if (validationError) { + // Loading failed: back to the diagram before the previews + if (prepared.ok && originalXml) { + onDisplayChart(originalXml, true) + } addToolOutput({ tool: "append_diagram", toolCallId: toolCall.toolCallId, diff --git a/hooks/use-model-config.ts b/hooks/use-model-config.ts index eda7fe9c..e05f680f 100644 --- a/hooks/use-model-config.ts +++ b/hooks/use-model-config.ts @@ -64,6 +64,28 @@ function migrateOldConfig(): MultiModelConfig | null { return config } +const isKnownProvider = (p: { provider: string }) => + Object.hasOwn(PROVIDER_INFO, p.provider) + +/** + * The stored config without providers this version does not know (saved + * by another version, or edited by hand): they would break every list of + * models. They stay in storage (saveConfig keeps them). Throws on bad JSON. + */ +function parseStoredConfig(stored: string): MultiModelConfig { + const config = JSON.parse(stored) as MultiModelConfig + const known = config.providers.filter(isKnownProvider) + if (known.length < config.providers.length) { + console.warn( + "Skipped saved providers this version does not know:", + config.providers + .filter((p) => !isKnownProvider(p)) + .map((p) => p.provider), + ) + } + return { ...config, providers: known } +} + /** * Load config from localStorage */ @@ -74,21 +96,7 @@ function loadConfig(): MultiModelConfig { const stored = localStorage.getItem(STORAGE_KEYS.modelConfigs) if (stored) { try { - const config = JSON.parse(stored) as MultiModelConfig - // A provider this version does not know (saved by another - // version, or edited by hand) would break every list of models - const known = config.providers.filter((p) => - Object.hasOwn(PROVIDER_INFO, p.provider), - ) - if (known.length < config.providers.length) { - console.warn( - "Skipped saved providers this version does not know:", - config.providers - .filter((p) => !known.includes(p)) - .map((p) => p.provider), - ) - } - return { ...config, providers: known } + return parseStoredConfig(stored) } catch { console.error("Failed to parse model config") } @@ -113,7 +121,26 @@ function loadConfig(): MultiModelConfig { */ function saveConfig(config: MultiModelConfig): void { if (typeof window === "undefined") return - localStorage.setItem(STORAGE_KEYS.modelConfigs, JSON.stringify(config)) + // Providers this version does not know are not in config: keep them, + // with their keys, for the version that saved them + let unknown: MultiModelConfig["providers"] = [] + try { + const stored = localStorage.getItem(STORAGE_KEYS.modelConfigs) + if (stored) { + unknown = (JSON.parse(stored) as MultiModelConfig).providers.filter( + (p) => !isKnownProvider(p), + ) + } + } catch { + // Unreadable: nothing to keep + } + localStorage.setItem( + STORAGE_KEYS.modelConfigs, + JSON.stringify({ + ...config, + providers: [...config.providers, ...unknown], + }), + ) } /** @@ -488,7 +515,8 @@ export function getSelectedAIConfig(): { let config: MultiModelConfig try { - config = JSON.parse(stored) + // Unknown providers would break the model lookup below + config = parseStoredConfig(stored) } catch { return { ...empty, accessCode } } diff --git a/hooks/use-session-manager.ts b/hooks/use-session-manager.ts index 8bd62535..a1b63780 100644 --- a/hooks/use-session-manager.ts +++ b/hooks/use-session-manager.ts @@ -114,6 +114,11 @@ export function useSessionManager( // Load sessions list 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(() => {}) // Only load a session if initialSessionId is provided (from URL param) if (initialSessionId) { diff --git a/lib/admin/providers.ts b/lib/admin/providers.ts index bb8445a1..ff6ca8d2 100644 --- a/lib/admin/providers.ts +++ b/lib/admin/providers.ts @@ -213,17 +213,12 @@ export function adminProvidersToConfig( indexByProvider.set(p.provider, index + 1) if (p.models.length === 0) continue const env = credEnvNames(p.provider, index) - // An entry with its own key also names its own URL variable, unset - // when the URL is empty: the global

_BASE_URL may be a proxy for - // another key, and the Test used the official endpoint. An Azure - // key belongs to one resource, so it keeps the server's. - const ownUrl = !!p.baseUrl || (!!p.apiKey && p.provider !== "azure") config.providers.push({ name: displayName(p), provider: p.provider, models: p.models, ...(env.key && p.apiKey ? { apiKeyEnv: env.key } : {}), - ...(env.url && ownUrl ? { baseUrlEnv: env.url } : {}), + ...(env.url && p.baseUrl ? { baseUrlEnv: env.url } : {}), ...(p.isDefault ? { default: true } : {}), }) } @@ -262,7 +257,12 @@ 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 - if (p.baseUrl) updates.OLLAMA_BASE_URL = p.baseUrl + // 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 + } } 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 72ab6492..8012dc77 100644 --- a/lib/ai-providers.ts +++ b/lib/ai-providers.ts @@ -997,17 +997,15 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { ? overrides?.apiKey || undefined : resolveApiKey(overrides, "OLLAMA_API_KEY") // Like other providers, a user's key never goes to the server's - // base URL. A key without a base URL is an Ollama Cloud key - // (local Ollama has no keys); without either, the SDK's local - // default. + // base URL: without a URL of their own it goes to Ollama Cloud. + // The server's key goes to OLLAMA_BASE_URL (or a server model's + // own variable), else to the SDK's local default: the desktop + // app's "Ollama (Local)" preset puts its key field there too. const baseURL = overrides?.baseUrl || (overrides?.apiKey ? PROVIDER_INFO.ollama.defaultBaseUrl - : process.env.OLLAMA_BASE_URL || - (apiKey - ? PROVIDER_INFO.ollama.defaultBaseUrl - : undefined)) + : resolveBaseUrlEnv(overrides, "OLLAMA_BASE_URL")) model = createOllama({ ...(baseURL && { baseURL }), ...(apiKey && { @@ -1043,9 +1041,8 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { : `${provider.toUpperCase()}_BASE_URL` // A local default (SGLang's 127.0.0.1) only fills the settings // form; the server must not call its own machine for it. With a - // user's key, or an admin entry's own (empty) URL variable, the - // OpenAI SDK would read the server's OPENAI_BASE_URL, so name - // the official endpoint. + // user's key the OpenAI SDK would read the server's + // OPENAI_BASE_URL, so name the official endpoint. const defaultUrl = PROVIDER_INFO[provider].defaultBaseUrl const publicDefault = defaultUrl?.startsWith("https://") ? defaultUrl @@ -1058,10 +1055,7 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { const baseURL = configuredBaseURL || (SDK_KNOWS_ENDPOINT.has(provider) && - !( - provider === "openai" && - (overrides?.apiKey || overrides?.baseUrlEnv) - ) + !(provider === "openai" && overrides?.apiKey) ? undefined : publicDefault) // With a user's Azure key the SDK would read the server's @@ -1097,6 +1091,23 @@ export function getAIModel(clientOverrides?: ClientOverrides): ModelConfig { return { model, providerOptions, modelId, provider } } +/** + * 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. + */ +export function globalBaseUrl(provider: ProviderName): string | undefined { + if (["bedrock", "edgeone", "ollama"].includes(provider)) return undefined + const name = + provider === "vertexai" + ? "GOOGLE_VERTEX_BASE_URL" + : provider === "gateway" + ? "AI_GATEWAY_BASE_URL" + : `${provider.toUpperCase()}_BASE_URL` + return process.env[name] || undefined +} + /** The provider of the server's own config: AI_PROVIDER, or the one with a key */ export function getServerProvider(): ProviderName | null { return (process.env.AI_PROVIDER as ProviderName) || detectProvider() diff --git a/lib/provider-models.ts b/lib/provider-models.ts index 183718a2..a809da75 100644 --- a/lib/provider-models.ts +++ b/lib/provider-models.ts @@ -78,16 +78,25 @@ const MAX_LIST_BYTES = 2 * 1024 * 1024 /** A fetch that reads at most MAX_LIST_BYTES of each response */ function sizeLimitedFetch(fetchFn: typeof fetch): typeof fetch { return async (input, init) => { - const response = await fetchFn(input, init) + // Ends a download that is too large (the Gateway SDK passes no + // signal of its own) + const download = new AbortController() + const signal = init?.signal + ? AbortSignal.any([init.signal, download.signal]) + : download.signal + const response = await fetchFn(input, { ...init, signal }) const body = await readLimitedBody(response, MAX_LIST_BYTES) if (body === null) { + download.abort() throw new ModelListError("The model list is too large.") } // The body is already decoded and has its own length now const headers = new Headers(response.headers) headers.delete("content-encoding") headers.delete("content-length") - return new Response(body, { + // Some statuses must have no body at all + const noBody = [101, 204, 205, 304].includes(response.status) + return new Response(noBody ? null : body, { status: response.status, statusText: response.statusText, headers, diff --git a/lib/read-limited-body.ts b/lib/read-limited-body.ts index cb53e0e5..884462f9 100644 --- a/lib/read-limited-body.ts +++ b/lib/read-limited-body.ts @@ -1,7 +1,8 @@ /** * Read a response body, giving up once it passes maxBytes, so a huge * download from a URL the client chose can't exhaust server memory. - * Returns null when it is too large. + * Returns null when it is too large; the caller then aborts the request, + * which ends the download. */ export async function readLimitedBody( response: Response, @@ -20,7 +21,9 @@ export async function readLimitedBody( if (done) break total += value.byteLength if (total > maxBytes) { - await reader.cancel() + // Not awaited: a copy of the body that Next.js keeps (its fetch + // dedupe) can hold the cancel back until it is read + reader.cancel().catch(() => {}) return null } chunks.push(value) diff --git a/lib/session-storage.ts b/lib/session-storage.ts index b693cbde..71b4acaa 100644 --- a/lib/session-storage.ts +++ b/lib/session-storage.ts @@ -176,6 +176,8 @@ export async function saveSession(session: ChatSession): Promise { try { const db = await getDB() await db.put(STORE_NAME, session) + // The desktop app opens this port (this origin's chats) next launch + window.electronAPI?.chatSaved?.().catch(() => {}) return true } catch (error) { console.error("Failed to save session:", error) diff --git a/lib/ssrf-protection.ts b/lib/ssrf-protection.ts index 88e4952e..6b3e433b 100644 --- a/lib/ssrf-protection.ts +++ b/lib/ssrf-protection.ts @@ -116,6 +116,14 @@ export function allowPrivateUrls(): boolean { return process.env.ALLOW_PRIVATE_URLS !== "false" } +/** 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") + this.name = "RedirectRefusedError" + } +} + /** * 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 @@ -126,7 +134,7 @@ export function redirectGuardedFetch(): typeof fetch | undefined { return async (input, init) => { const response = await fetch(input, { ...init, redirect: "manual" }) if (response.status >= 300 && response.status < 400) { - throw new Error("Redirects are not allowed for custom base URLs") + throw new RedirectRefusedError() } return response } diff --git a/packages/mcp-server/src/http-server.ts b/packages/mcp-server/src/http-server.ts index 6682761f..fc3cac01 100644 --- a/packages/mcp-server/src/http-server.ts +++ b/packages/mcp-server/src/http-server.ts @@ -3,6 +3,7 @@ * Serves draw.io embed with state sync and history UI */ +import { randomUUID } from "node:crypto" import { readFileSync } from "node:fs" import http from "node:http" import { dirname, join } from "node:path" @@ -105,11 +106,22 @@ function ensureSessionStateInitialized(sessionId: string): void { // A blank diagram keeps the draw.io spinner (spin=1) from waiting // forever when no load(xml) is ever sent setState(sessionId, BLANK_MXFILE, undefined, false, false) + // Nothing is known about this session: a tab that still shows it keeps + // its diagram + const state = stateStore.get(sessionId) + if (state) state.blank = true } interface SessionState { xml: string version: number + // Made when the state is created (first use, or again after it expired + // or the MCP process restarted) and kept by every write. A tab tells by + // it that the server lost what it knew, and every push names the state + // it was based on, so one based on a lost state is refused. + stateId: string + // Created blank because nothing was saved; cleared by the first write + blank?: boolean // 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 @@ -178,6 +190,7 @@ export function setState( stateStore.set(sessionId, { xml, version: newVersion, + stateId: existing?.stateId ?? randomUUID(), serverVersion: fromBrowser ? existing?.serverVersion : newVersion, lastUpdated: new Date(), lastPolled: existing?.lastPolled, @@ -444,6 +457,8 @@ function handleStateApi( JSON.stringify({ xml: state?.xml || null, version: state?.version || 0, + stateId: state?.stateId ?? null, + blank: !!state?.blank, syncRequested: !!state?.syncRequested, exportFormat: state?.exportFormat || null, exportXml: state?.exportXml || null, @@ -486,12 +501,58 @@ function handleStateApi( return } + // A push can come before the tab's first poll after a + // restart: recover the saved file first, so it is compared + // with that and never overwrites it unseen + ensureSessionStateInitialized(sessionId) + const current = stateStore.get(sessionId) + + // A tab of this version names the state its push is based + // on. Another state (the server lost the one it knew, or + // the tab has not polled yet): refused, and the tab's next + // poll decides whose diagram wins. + if (current && "stateId" in data) { + if (data.stateId !== current.stateId) { + res.writeHead(409, { + "Content-Type": "application/json", + }) + res.end( + JSON.stringify({ + error: "Session was recreated", + stateChanged: true, + version: current.version, + }), + ) + return + } + // What a recovering tab showed: kept in history only + if (data.source === "recover") { + const saved = + typeof data.xml === "string" && + !!data.xml && + data.xml !== current.xml + if (saved) { + addHistory(sessionId, data.xml, data.svg || "") + } + res.writeHead(409, { + "Content-Type": "application/json", + }) + res.end( + JSON.stringify({ + error: "Diagram changed on the server", + version: current.version, + savedToHistory: saved, + }), + ) + return + } + } + // The browser edited a version older than the latest AI write // (it has not loaded that write yet). Keep the AI write; the // browser loads it on its next poll. A sync reply is also // stale after a newer write of the browser's own (a user // edit saved while the export ran). - const current = stateStore.get(sessionId) if ( typeof data.baseVersion === "number" && (data.baseVersion < (current?.serverVersion ?? 0) || diff --git a/packages/mcp-server/src/persistence.ts b/packages/mcp-server/src/persistence.ts index 123cd9b7..38e47e04 100644 --- a/packages/mcp-server/src/persistence.ts +++ b/packages/mcp-server/src/persistence.ts @@ -56,15 +56,22 @@ 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 + // the session then shows something else than what they hold. Cleared + // once the file is read, or is gone. private unreadable = new Set() /** The session's saved diagram, or null. */ load(sessionId: string): string | null { const path = this.pathFor(sessionId) - if (!path || !existsSync(path)) return null + if (!path) return null + if (!existsSync(path)) { + this.unreadable.delete(path) + return null + } try { - return readFileSync(path, "utf-8") + const xml = readFileSync(path, "utf-8") + this.unreadable.delete(path) + return xml } catch (error) { log.warn(`Could not read the saved diagram ${path}: ${error}`) this.unreadable.add(path) @@ -93,7 +100,13 @@ export class Autosaver { const entry = this.pending.get(sessionId) this.pending.delete(sessionId) const path = this.pathFor(sessionId) - if (!entry || !this.dir || !path || this.unreadable.has(path)) return + 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 + } try { const isNew = !existsSync(path) // A blank page the browser shows before any drawing: nothing to keep diff --git a/packages/mcp-server/src/preview/preview.js b/packages/mcp-server/src/preview/preview.js index ec85a304..425cb159 100644 --- a/packages/mcp-server/src/preview/preview.js +++ b/packages/mcp-server/src/preview/preview.js @@ -1,7 +1,16 @@ const iframe = document.getElementById('drawio'); let currentVersion = 0, isReady = false, pendingXml = null, lastXml = null; +// The server state this tab is in step with (see stateId in http-server.ts); +// null until the first poll +let stateId = null; +// The newest diagram on the canvas, saved to the server or not: lastXml is +// the last one the server has +let latestXml = null; +let pushFailing = false; // the last push could not reach the server +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; let pendingMcpExport = null; // 'png', 'svg' or 'xmlsvg' when MCP requested export let mcpExportSeq = 0; // number of the latest MCP export @@ -17,19 +26,24 @@ window.addEventListener('message', (e) => { if (msg.event === 'init') { isReady = true; if (pendingXml) { loadDiagram(pendingXml); pendingXml = null; } - } else if ((msg.event === 'save' || msg.event === 'autosave') && msg.xml && msg.xml !== lastXml) { + } 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; + // Also an edit undone back to what the server has + latestXml = msg.xml; + if (msg.xml === lastXml) return; // Request SVG export, then push state with SVG. Remember the - // version this edit is based on, so the server can reject it - // if the AI wrote a newer version that is not loaded yet. + // version and state this edit is based on, so the server can + // reject it if the AI wrote a newer version that is not loaded + // yet, or if it lost that state. pendingSvgExport = msg.xml; pendingSvgBase = currentVersion; + pendingSvgStateId = stateId; iframe.contentWindow.postMessage(JSON.stringify({ action: 'export', format: 'svg' }), '*'); // Fallback if export doesn't respond - setTimeout(() => { if (pendingSvgExport === msg.xml) { pushState(msg.xml, '', pendingSvgBase); pendingSvgExport = null; } }, 2000); + setTimeout(() => { if (pendingSvgExport === msg.xml) { pushState(msg.xml, '', pendingSvgBase, 'edit', pendingSvgStateId); pendingSvgExport = null; } }, 2000); } else if (msg.event === 'export' && msg.format === 'xml') { // Sync export requested by the server (get_diagram). // draw.io returns the XML in msg.xml, with no msg.data. A late @@ -39,7 +53,7 @@ window.addEventListener('message', (e) => { // Push with the version the export was taken at: a // newer AI write may have loaded meanwhile, and this // older XML must not overwrite it. - pushState(msg.xml, '', pendingSyncBase, 'sync'); + pushState(msg.xml, '', pendingSyncBase, 'sync', pendingSyncStateId); } } else if (msg.event === 'export' && msg.data) { // Handle MCP server export request (png/svg). fireExport tags @@ -98,7 +112,7 @@ window.addEventListener('message', (e) => { if (pendingSvgExport) { const xml = pendingSvgExport; pendingSvgExport = null; - pushState(xml, svg, pendingSvgBase); + pushState(xml, svg, pendingSvgBase, 'edit', pendingSvgStateId); } else if (pendingAiSvg) { pendingAiSvg = false; fetch('/api/history-svg', { @@ -114,6 +128,7 @@ window.addEventListener('message', (e) => { function loadDiagram(xml, capturePreview = false) { if (!isReady) { pendingXml = xml; return; } lastXml = xml; + latestXml = xml; iframe.contentWindow.postMessage(JSON.stringify({ action: 'load', xml, autosave: 1 }), '*'); if (capturePreview) { setTimeout(() => { @@ -144,24 +159,27 @@ function showNotice(text) { noticeTimer = setTimeout(() => el.classList.remove('open'), 8000); } -// Same rule as hasCells in pages.ts: a cell besides the root cells, or a -// compressed page -function hasCells(xml) { - return /<(mxCell\b[^>]*\bid\s*=\s*["'](?![01]["'])|UserObject\b|object\b)|]*>\s*[^\s<]/.test(xml || ''); -} - // source is 'sync' for replies to a server sync request, 'recover' for the -// tab's copy after the server recovered the session, else 'edit' -async function pushState(xml, svg = '', baseVersion = currentVersion, source = 'edit') { +// tab's copy after the server recovered the session, else 'edit'. sid is the +// server state the push is based on. +async function pushState(xml, svg = '', baseVersion = currentVersion, source = 'edit', sid = stateId) { if (!sessionId) return; try { const r = await fetch('/api/state', { method: 'POST', headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ sessionId, xml, svg, baseVersion, source }) + body: JSON.stringify({ sessionId, xml, svg, baseVersion, source, stateId: sid }) }); - if (r.ok) { const d = await r.json(); currentVersion = d.version; lastXml = xml; } - // 409: the AI wrote a newer version; load it now + pushFailing = false; + if (r.ok) { + const d = await r.json(); + // An answer about a state this tab has left since + if (sid !== stateId) return; + currentVersion = d.version; + lastXml = xml; + } + // 409: the AI wrote a newer version, or the server lost the state + // this push was based on; the next poll sorts it out else if (r.status === 409) { const d = await r.json().catch(() => ({})); if (d.savedToHistory) { @@ -171,33 +189,63 @@ async function pushState(xml, svg = '', baseVersion = currentVersion, source = ' } poll(); } - } catch (e) { console.error('Push failed:', e); } + } catch (e) { + console.error('Push failed:', e); + if (!pushFailing) { + pushFailing = true; + showNotice("Can't reach the MCP server. Your changes are only in this tab for now; use Download to keep a copy."); + } + } +} + +// The server made a new state for this session: it expired, or the MCP +// process restarted. Decide whose diagram wins. +function recoverState(s) { + stateId = s.stateId; + // The old state's pending work is gone with it + const projectionShown = projectionExportActive; + projectionExportActive = false; + forceReload = false; + pendingMcpExport = null; + pendingSyncExport = false; + const mine = latestXml; + currentVersion = s.version; + if (s.blank || s.xml === lastXml) { + // The server knows nothing, or exactly what this tab last saved: + // the canvas can only be newer, so it wins (edits made while the + // server was down are saved now) + if (projectionShown && mine) { + iframe.contentWindow.postMessage(JSON.stringify({ action: 'load', xml: mine, autosave: 1 }), '*'); + } + if (mine && mine !== s.xml) pushState(mine, '', s.version); + } else { + // The server has a diagram this tab never showed (an AI write it + // missed, a saved file): show that, and keep this tab's copy in + // History unless it is the same + loadDiagram(s.xml, true); + if (mine && mine !== s.xml) pushState(mine, '', s.version, 'recover'); + } } let pendingSyncExport = false; let pendingSyncBase = 0; // version the pending sync export was taken at +let pendingSyncStateId = null; // and the state it belonged to let syncExportSeq = 0; // number of the latest sync export async function poll() { if (!sessionId) return; - const knownVersion = currentVersion; + const seq = ++pollSeq; try { const r = await fetch('/api/state?sessionId=' + encodeURIComponent(sessionId)); if (!r.ok) return; const s = await r.json(); - // The server lost this session (it expired, or the MCP process - // restarted) and rebuilt it. Blank: push back what the browser - // shows. From the auto-save file, which can hold an AI write this - // tab never loaded: show that, and keep this tab's copy in History - // (a push based on version 0 is refused and saved there). - if (s.version < knownVersion && lastXml) { - if (!hasCells(s.xml)) { - pushState(lastXml); - } else { - currentVersion = 0; - if (s.xml !== lastXml) pushState(lastXml, '', 0, 'recover'); - } - } + // An older answer than one already handled (the interval, the 409 + // handler and the projection restore each poll): it could name a + // state that is gone + if (seq < lastHandledPoll) return; + lastHandledPoll = seq; + if (stateId === null) stateId = s.stateId; + else if (s.stateId && s.stateId !== stateId) recoverState(s); // Load new diagram from server (before export, so we export latest). // While a page-targeted projection is on screen, only the restore // (forceReload) replaces it, so a new version doesn't fight the @@ -217,6 +265,7 @@ async function poll() { if (s.syncRequested && !pendingSyncExport && isReady && !projectionExportActive) { pendingSyncExport = true; pendingSyncBase = currentVersion; + pendingSyncStateId = stateId; // draw.io echoes the request in msg.message, so the reply can // be matched to this request const seq = ++syncExportSeq; @@ -324,7 +373,8 @@ saveConfirmBtn.onclick = () => { // so no wrapper injection is needed. The legacy fallback below // remains only for documents that somehow slipped past // normalisation (e.g. an older session loaded from external state). - let xmlData = lastXml || ''; + // The canvas as it is, also edits not saved to the server yet + let xmlData = latestXml || lastXml || ''; if (xmlData && !xmlData.includes(''; } diff --git a/packages/mcp-server/tests/http-server.test.ts b/packages/mcp-server/tests/http-server.test.ts index cc060382..75d539a0 100644 --- a/packages/mcp-server/tests/http-server.test.ts +++ b/packages/mcp-server/tests/http-server.test.ts @@ -337,6 +337,87 @@ describe("export requests", () => { }) }) +describe("a session state recreated after it was lost", () => { + const SAVED = `` + const getJson = async (id: string) => + JSON.parse((await request(`/api/state?sessionId=${id}`)).body) + + it("names each state, and says when it was made blank", async () => { + const first = await getJson("mcp-sid-blank") + expect(first.stateId).toMatch(/^[0-9a-f-]{36}$/) + expect(first.blank).toBe(true) + setState("mcp-sid-blank", "AI write") + const after = await getJson("mcp-sid-blank") + // Same state, no longer blank + expect(after.stateId).toBe(first.stateId) + expect(after.blank).toBe(false) + }) + + it("refuses a push made for another state, also before any poll", async () => { + // The MCP process restarted; the tab's push comes before its poll + onSessionRecreate((id) => (id === "mcp-sid-restart" ? SAVED : null)) + try { + for (const stateId of ["from-before", null]) { + const res = await postJson("/api/state", { + sessionId: "mcp-sid-restart", + xml: "tab's old copy", + baseVersion: 7, + stateId, + }) + expect(res.status).toBe(409) + expect(JSON.parse(res.body).stateChanged).toBe(true) + // The saved file was recovered first and is kept + expect(getState("mcp-sid-restart")?.xml).toBe(SAVED) + } + } finally { + onSessionRecreate(() => null) + } + }) + + it("accepts a push for the current state", async () => { + const { stateId, version } = await getJson("mcp-sid-ok") + const res = await postJson("/api/state", { + sessionId: "mcp-sid-ok", + xml: "user edit", + baseVersion: version, + stateId, + }) + expect(res.status).toBe(200) + expect(getState("mcp-sid-ok")?.xml).toBe("user edit") + }) + + it("keeps a recovering tab's copy in history, never on the canvas", async () => { + setState("mcp-sid-recover", SAVED) + const { stateId, version } = await getJson("mcp-sid-recover") + const before = getHistory("mcp-sid-recover").length + const res = await postJson("/api/state", { + sessionId: "mcp-sid-recover", + xml: "what the tab showed", + baseVersion: version, + stateId, + source: "recover", + }) + expect(res.status).toBe(409) + expect(JSON.parse(res.body).savedToHistory).toBe(true) + expect(getState("mcp-sid-recover")?.xml).toBe(SAVED) + expect(getHistory("mcp-sid-recover")).toHaveLength(before + 1) + expect(getHistory("mcp-sid-recover").at(-1)?.xml).toBe( + "what the tab showed", + ) + }) + + it("keeps the old rules for a tab from an older version", async () => { + // Its pushes have no stateId field + const version = setState("mcp-sid-legacy", "AI") + const res = await postJson("/api/state", { + sessionId: "mcp-sid-legacy", + xml: "edit", + baseVersion: version, + }) + expect(res.status).toBe(200) + }) +}) + describe("preview page", () => { it("shows the saved diagram of a session whose state expired", async () => { const saved = `` diff --git a/packages/mcp-server/tests/persistence.test.ts b/packages/mcp-server/tests/persistence.test.ts index 30b6047b..b766b269 100644 --- a/packages/mcp-server/tests/persistence.test.ts +++ b/packages/mcp-server/tests/persistence.test.ts @@ -8,6 +8,7 @@ import { mkdtempSync, readdirSync, readFileSync, + rmSync, utimesSync, writeFileSync, } from "node:fs" @@ -102,6 +103,33 @@ describe("Autosaver", () => { expect(readFileSync(path, "utf-8")).toBe(DIAGRAM) }) + it("saves again once the file was read, or is gone", () => { + const saver = new Autosaver(tempDir(), 10) + saver.schedule("mcp-fixed", DIAGRAM) + saver.flush() + const path = saver.pathFor("mcp-fixed") as string + chmodSync(path, 0o000) + expect(saver.load("mcp-fixed")).toBeNull() + // Permissions fixed; the session is recreated and reads the file + chmodSync(path, 0o644) + expect(saver.load("mcp-fixed")).toBe(DIAGRAM) + const edited = DIAGRAM.replace('id="a"', 'id="b"') + saver.schedule("mcp-fixed", edited) + saver.flush() + expect(readFileSync(path, "utf-8")).toBe(edited) + + // A file that could not be read and was then deleted protects + // nothing any more + chmodSync(path, 0o000) + expect(saver.load("mcp-fixed")).toBeNull() + chmodSync(path, 0o644) + rmSync(path) + expect(saver.load("mcp-fixed")).toBeNull() + saver.schedule("mcp-fixed", DIAGRAM) + 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/tests/e2e/provider-models.spec.ts b/tests/e2e/provider-models.spec.ts index faab32b7..a70a2df9 100644 --- a/tests/e2e/provider-models.spec.ts +++ b/tests/e2e/provider-models.spec.ts @@ -271,4 +271,60 @@ test("a test result for an old API key is dropped", async ({ page }) => { release() await page.waitForTimeout(500) await expect(dialog.locator('[title="1.0 s"]')).toHaveCount(0) + // The test is over: the button works again and nothing spins + await expect( + dialog.getByRole("button", { name: "Test", exact: true }), + ).toBeEnabled() + await expect(dialog.locator(".animate-spin")).toHaveCount(0) +}) + +test("an older test does not end a newer one's spinners", async ({ page }) => { + // Each validate request waits for its own release + const releases: Array<() => void> = [] + await page.route("**/api/validate-model", async (route) => { + await new Promise((r) => releases.push(r)) + await route.fulfill({ + status: 200, + json: { valid: true, responseTime: 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 user corrects the key and tests again + await dialog.locator("#api-key").fill("new-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() + 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, +}) => { + const release = await holdRoute(page, "**/api/validate-model", { + valid: true, + responseTime: 1000, + }) + const dialog = await openQwenSettings(page, TWO_PROVIDERS) + await dialog.getByRole("button", { name: "Test", exact: true }).click() + // The user looks at the other provider while another tab changes the key + await dialog.getByText("GLM (Zhipu)").first().click() + await page.evaluate(() => { + const key = "next-ai-draw-io-model-configs" + const config = JSON.parse(localStorage.getItem(key) ?? "{}") + config.providers[0].apiKey = "key-from-another-tab" + const value = JSON.stringify(config) + localStorage.setItem(key, value) + window.dispatchEvent( + new StorageEvent("storage", { key, newValue: value }), + ) + }) + release() + await page.waitForTimeout(500) + await dialog.getByText("Qwen (Alibaba)").first().click() + await expect(dialog.locator(".animate-spin")).toHaveCount(0) }) diff --git a/tests/unit/admin-providers.test.ts b/tests/unit/admin-providers.test.ts index 4c8a7c02..90fc2664 100644 --- a/tests/unit/admin-providers.test.ts +++ b/tests/unit/admin-providers.test.ts @@ -69,6 +69,34 @@ 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( + [provider({ provider: "ollama", apiKey: "ollama-key" })], + [], + ) + expect(cloud.OLLAMA_API_KEY).toBe("ollama-key") + expect(cloud.OLLAMA_BASE_URL).toBe("https://ollama.com/api") + const own = deriveEnvUpdates( + [ + provider({ + provider: "ollama", + apiKey: "k", + baseUrl: "https://ollama.internal/api", + }), + ], + [], + ) + expect(own.OLLAMA_BASE_URL).toBe("https://ollama.internal/api") + // No key: local Ollama, nothing to write + const local = deriveEnvUpdates( + [provider({ provider: "ollama", apiKey: undefined })], + [], + ) + expect(local.OLLAMA_BASE_URL ?? null).toBeNull() + }) + it("maps bedrock credentials to ADMIN_AWS_* env vars", () => { const updates = deriveEnvUpdates( [ @@ -156,29 +184,6 @@ describe("adminProvidersToConfig", () => { expect(config.providers[1].apiKeyEnv).toBe("ADMIN_OPENAI_API_KEY_2") }) - it("names its own URL variable when it has its own key, even empty", () => { - // Otherwise chat reads the global OPENAI_BASE_URL, which may be a - // proxy for another key, while the Test used the official endpoint - const own = adminProvidersToConfig([provider()]).providers[0] - expect(own.baseUrlEnv).toBe("ADMIN_OPENAI_BASE_URL") - // Without a key or URL of its own: the global key and URL, a pair - const shared = adminProvidersToConfig([provider({ apiKey: undefined })]) - .providers[0] - expect(shared.baseUrlEnv).toBeUndefined() - // An Azure key belongs to one resource: AZURE_BASE_URL stays - const azure = adminProvidersToConfig([provider({ provider: "azure" })]) - .providers[0] - expect(azure.baseUrlEnv).toBeUndefined() - expect( - adminProvidersToConfig([ - provider({ - provider: "azure", - baseUrl: "https://r.openai.azure.com/openai", - }), - ]).providers[0].baseUrlEnv, - ).toBe("ADMIN_AZURE_BASE_URL") - }) - it("skips providers without models and carries the default flag", () => { const config = adminProvidersToConfig([ provider({ id: "p1", models: [] }), diff --git a/tests/unit/admin-test-model.test.ts b/tests/unit/admin-test-model.test.ts new file mode 100644 index 00000000..42ce23d5 --- /dev/null +++ b/tests/unit/admin-test-model.test.ts @@ -0,0 +1,68 @@ +// @vitest-environment node +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest" + +// The request the admin Test hands to validate-model +const sent = vi.hoisted(() => ({ body: null as any, headers: null as any })) +vi.mock("@/app/api/validate-model/route", () => ({ + POST: async (req: Request) => { + sent.body = await req.json() + sent.headers = Object.fromEntries(req.headers) + return Response.json({ valid: true }) + }, +})) +vi.mock("@/lib/admin/auth", () => ({ checkAdminAuth: () => null })) +vi.mock("@/lib/admin/settings", () => ({ loadSettings: () => ({}) })) + +import { POST as testModel } from "@/app/api/admin/test-model/route" + +const ENV = ["OPENAI_BASE_URL", "SGLANG_BASE_URL", "AI_GATEWAY_BASE_URL"] +const saved: Record = {} +beforeEach(() => { + for (const k of ENV) { + saved[k] = process.env[k] + delete process.env[k] + } +}) +afterEach(() => { + for (const k of ENV) { + if (saved[k] === undefined) delete process.env[k] + else process.env[k] = saved[k] + } +}) + +const test = (provider: Record) => + testModel( + new Request("http://localhost/api/admin/test-model", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify({ + provider: { id: "p1", models: ["m"], ...provider }, + modelId: "m", + }), + }), + ) + +describe("admin Test of an entry without a URL", () => { + it("tests the server's

_BASE_URL, where chat sends the entry's key", async () => { + // A server model without baseUrlEnv reads the global variable + 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") + + process.env.AI_GATEWAY_BASE_URL = "https://gateway.example.com/v3/ai" + await test({ provider: "gateway", apiKey: "k" }) + expect(sent.body.baseUrl).toBe("https://gateway.example.com/v3/ai") + }) + + it("keeps the entry's own URL, and none when the server has none", async () => { + process.env.SGLANG_BASE_URL = "http://gpu-box:8000/v1" + await test({ + provider: "sglang", + apiKey: "k", + baseUrl: "http://other:8000/v1", + }) + expect(sent.body.baseUrl).toBe("http://other:8000/v1") + await test({ provider: "deepseek", apiKey: "k" }) + expect(sent.body.baseUrl).toBeUndefined() + }) +}) diff --git a/tests/unit/ai-providers-credentials.test.ts b/tests/unit/ai-providers-credentials.test.ts index 344088c0..df155cdb 100644 --- a/tests/unit/ai-providers-credentials.test.ts +++ b/tests/unit/ai-providers-credentials.test.ts @@ -378,25 +378,6 @@ describe("whose keys a request uses", () => { expect(provider.chat).not.toHaveBeenCalled() }) - it("sends an admin OpenAI key without a URL to the official endpoint", () => { - // Its URL variable is named but empty; the SDK would otherwise read - // the server's OPENAI_BASE_URL, a proxy for another key - process.env.OPENAI_BASE_URL = "https://operator-proxy.example.com/v1" - process.env.ADMIN_OPENAI_API_KEY = "panel-key" - getAIModel({ - provider: "openai", - modelId: "gpt-5.5", - apiKeyEnv: "ADMIN_OPENAI_API_KEY", - baseUrlEnv: "ADMIN_OPENAI_BASE_URL", - }) - expect(createOpenAI).toHaveBeenLastCalledWith( - expect.objectContaining({ - apiKey: "panel-key", - baseURL: "https://api.openai.com/v1", - }), - ) - }) - it("uses Chat Completions for any configured base URL", () => { // The settings form fills in the official URL for a new provider getAIModel({ @@ -425,21 +406,44 @@ describe("whose keys a request uses", () => { ) }) - it("sends the server's Ollama key without a base URL to Ollama Cloud", async () => { - // An Ollama key is an Ollama Cloud key: local Ollama has none. - // The admin panel saves it as OLLAMA_API_KEY, without a base URL. + it("sends the server's Ollama key where OLLAMA_BASE_URL says, or to local Ollama", async () => { + // The desktop app's "Ollama (Local)" preset puts its API Key field + // into OLLAMA_API_KEY; with no base URL that is the local Ollama process.env.OLLAMA_API_KEY = "server-key" const { createOllama } = await import("ollama-ai-provider-v2") getAIModel({ provider: "ollama", modelId: "m" }) - expect(createOllama).toHaveBeenLastCalledWith( - expect.objectContaining({ baseURL: "https://ollama.com/api" }), - ) - // Without a key: the SDK's local default - delete process.env.OLLAMA_API_KEY - getAIModel({ provider: "ollama", modelId: "m" }) expect(vi.mocked(createOllama).mock.lastCall?.[0]).not.toHaveProperty( "baseURL", ) + process.env.OLLAMA_BASE_URL = "https://ollama.com/api" + getAIModel({ provider: "ollama", modelId: "m" }) + expect(createOllama).toHaveBeenLastCalledWith( + expect.objectContaining({ baseURL: "https://ollama.com/api" }), + ) + }) + + it("uses a server model's own Ollama URL variable", async () => { + process.env.OLLAMA_BASE_URL = "http://other.internal:11434/api" + process.env.MY_OLLAMA_URL = "https://ollama.proxy.example/api" + process.env.MY_OLLAMA_KEY = "proxy-key" + try { + const { createOllama } = await import("ollama-ai-provider-v2") + getAIModel({ + provider: "ollama", + modelId: "m", + apiKeyEnv: "MY_OLLAMA_KEY", + baseUrlEnv: "MY_OLLAMA_URL", + }) + expect(createOllama).toHaveBeenLastCalledWith( + expect.objectContaining({ + baseURL: "https://ollama.proxy.example/api", + headers: { Authorization: "Bearer proxy-key" }, + }), + ) + } finally { + delete process.env.MY_OLLAMA_URL + delete process.env.MY_OLLAMA_KEY + } }) it("needs a base URL with a user's Azure key", () => { diff --git a/tests/unit/ai-providers.test.ts b/tests/unit/ai-providers.test.ts index f9eb7c65..f01f2195 100644 --- a/tests/unit/ai-providers.test.ts +++ b/tests/unit/ai-providers.test.ts @@ -462,11 +462,11 @@ describe("Ollama API key security", () => { expect(createOllamaMock).toHaveBeenCalledTimes(1) const callArgs = createOllamaMock.mock.calls[0][0] - // As env.example says: without OLLAMA_BASE_URL, Ollama Cloud (the - // SDK's default is the local server, which has no keys) + // The SDK's local default: the desktop app's "Ollama (Local)" + // preset puts its API Key field into OLLAMA_API_KEY + expect(callArgs).not.toHaveProperty("baseURL") expect(callArgs).toEqual( expect.objectContaining({ - baseURL: "https://ollama.com/api", headers: { Authorization: "Bearer server-key" }, }), ) diff --git a/tests/unit/chat-route-errors.test.ts b/tests/unit/chat-route-errors.test.ts new file mode 100644 index 00000000..e6eee6c8 --- /dev/null +++ b/tests/unit/chat-route-errors.test.ts @@ -0,0 +1,126 @@ +// @vitest-environment node +import { afterEach, beforeEach, describe, expect, it, vi } from "vitest" + +// 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" + +const ENV = [ + "AI_PROVIDER", + "AI_MODEL", + "OPENAI_API_KEY", + "OLLAMA_BASE_URL", + "OLLAMA_API_KEY", + "NEXT_AI_DRAWIO_DESKTOP", +] +const saved: Record = {} + +beforeEach(() => { + for (const k of ENV) saved[k] = process.env[k] + for (const k of ENV) delete process.env[k] +}) + +afterEach(() => { + for (const k of ENV) { + if (saved[k] === undefined) delete process.env[k] + else process.env[k] = saved[k] + } + vi.unstubAllGlobals() +}) + +/** Every provider request answers with this status and text */ +const providerAnswers = (status: number, body: string) => + vi.stubGlobal( + "fetch", + vi.fn( + async () => + new Response(body, { + status, + headers: { "Content-Type": "application/json" }, + }), + ), + ) + +/** The error text the chat panel gets from the stream */ +async function streamedError(headers: Record) { + const res = await 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: "", + }), + }), + ) + const text = await res.text() + const line = text.split("\n").find((l) => l.includes('"type":"error"')) + return line ? JSON.parse(JSON.parse(line.slice(6)).errorText).message : "" +} + +describe("provider error texts in the stream", () => { + it("shows the user's own local Ollama error in the desktop app", async () => { + process.env.NEXT_AI_DRAWIO_DESKTOP = "1" + process.env.AI_PROVIDER = "ollama" + process.env.AI_MODEL = "llama3" + // Ollama is not running + vi.stubGlobal( + "fetch", + vi.fn(async () => { + throw Object.assign(new TypeError("fetch failed"), { + cause: new Error("connect ECONNREFUSED 127.0.0.1:11434"), + }) + }), + ) + expect(await streamedError({})).toMatch( + /127\.0\.0\.1:11434|fetch failed/, + ) + // The SDK retries a refused connection twice, waiting between + }, 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 + providerAnswers( + 400, + JSON.stringify({ + error: { + message: + "The daily public quota has been exhausted. After deployment, you can enjoy a personal daily exclusive quota.", + }, + }), + ) + expect( + await streamedError({ + "x-ai-provider": "edgeone", + "x-ai-model": "@tx/deepseek-ai/deepseek-v3-0324", + }), + ).toMatch(/daily public quota/) + }) + + it("hides the provider's text on the server's own key", async () => { + process.env.AI_PROVIDER = "openai" + process.env.AI_MODEL = "gpt-5.5" + process.env.OPENAI_API_KEY = "server-key" + providerAnswers( + 403, + JSON.stringify({ + error: { message: "Organization org-operator is suspended" }, + }), + ) + const message = await streamedError({}) + expect(message).not.toMatch(/org-operator/) + expect(message).toBe("The provider returned an error.") + }) +}) diff --git a/tests/unit/env-loader.test.ts b/tests/unit/env-loader.test.ts index 06d6a033..51337175 100644 --- a/tests/unit/env-loader.test.ts +++ b/tests/unit/env-loader.test.ts @@ -15,7 +15,17 @@ vi.mock("electron", () => ({ import { loadEnvFile } from "@/electron/main/env-loader" -const KEYS = ["T_JSON", "T_COMMENT", "T_PLAIN", "T_DOUBLE"] +const KEYS = [ + "T_JSON", + "T_COMMENT", + "T_PLAIN", + "T_DOUBLE", + "T_QUOTED_COMMENT", + "T_KEY_COMMENT", + "T_HASH", + "T_AFTER", + "T_JOINED", +] afterEach(() => { for (const k of KEYS) delete process.env[k] }) @@ -39,4 +49,27 @@ describe("loadEnvFile", () => { expect(process.env.T_PLAIN).toBe("plain") expect(process.env.T_DOUBLE).toBe(`say "hi"`) }) + + it("drops a comment that ends with a quote, like dotenv", () => { + // Expected values checked against dotenv 16.6.1's parse + dir.path = mkdtempSync(join(tmpdir(), "env-loader-")) + writeFileSync( + join(dir.path, ".env"), + [ + `T_QUOTED_COMMENT="gpt-5" # pick "fast"`, + `T_KEY_COMMENT="sk-abc" # from "Team A"`, + `T_AFTER='a' b`, + `T_JOINED="a"b`, + // Unquoted: a # without a space before it stays in the value + // (dotenv would cut it; this loader never did) + "T_HASH=http://host/#/x", + ].join("\n"), + ) + loadEnvFile() + expect(process.env.T_QUOTED_COMMENT).toBe("gpt-5") + expect(process.env.T_KEY_COMMENT).toBe("sk-abc") + expect(process.env.T_AFTER).toBe(`'a' b`) + expect(process.env.T_JOINED).toBe(`"a"b`) + expect(process.env.T_HASH).toBe("http://host/#/x") + }) }) diff --git a/tests/unit/mcp-preview-history.test.ts b/tests/unit/mcp-preview-history.test.ts new file mode 100644 index 00000000..5b899d99 --- /dev/null +++ b/tests/unit/mcp-preview-history.test.ts @@ -0,0 +1,46 @@ +import { readFileSync } from "node:fs" +import { join } from "node:path" +import { describe, expect, it } from "vitest" + +// The MCP preview page, as the server fills it in, with its script run in +// this document (no session id, so it does not poll) +const dir = join(process.cwd(), "packages/mcp-server/src/preview") +const html = readFileSync(join(dir, "index.html"), "utf8") + .replace("{{CSS}}", "") + .replace("{{SESSION_BADGE}}", "") + .replaceAll("{{DISABLED}}", "") + .replace("{{DRAWIO_URL}}", "about:blank") + .replace("{{SESSION_JSON}}", '""') + .replace("{{ORIGIN_JSON}}", '"https://embed.diagrams.net"') +const scripts = [...html.matchAll(/