Compare commits

...
Author SHA1 Message Date
dayuan.jiang a46787c1b8 fix(mcp-server): fix XSS and crashes, make XML validation strict
- Validate and escape the mcp session id; only serve localhost Host/Origin
- Malformed URLs and session ids return errors instead of crashing the process
- Strict XML syntax check with saxes (linkedom never reports parse errors)
- autoFixXml no longer corrupts valid XML; attribute newlines serialized as entities
- Sessions stay alive while polled; browser pushes carry a base version (409 on conflict)
- Page tools respect the edit gate; UTF-8 bodies decoded correctly
- Export replies matched to requests and serialized; xml sync export handled
- UserObject/object cells addressable by id; history restored by stable id; logs off stdout
2026-10-03 17:45:41 +09:00
dayuan.jiang 95f4b4b92b fix(electron): decrypt keys after ready and harden navigation and IPC
- Apply preset env after app ready, so Windows/Linux get decrypted keys
- Never re-encrypt ciphertext; restore env when switching or removing presets
- Block navigation away from the app, open external links in the browser, check IPC senders
- Keep inherited proxy settings, default NO_PROXY for localhost
- Serialize server start/restart, kill stuck processes, follow port changes
- Atomic config writes, keep corrupt files as backups, remember the server port
- Menu and settings window stay in sync; dev script gets the decrypted preset env
- Use app.isPackaged, parse inline .env comments, drop .env files from the bundle
2026-10-03 17:45:41 +09:00
dayuan.jiang 3193d20e00 fix(model-config): keep model selection valid and fix admin panel edge cases
- Fall back to the default server model when a saved one disappears
- Sync model config across tabs
- Validation uses the base path and sends the access code
- Model ids edited as drafts (no empty, duplicate or padded ids)
- Credential changes reset validation; stale validation results are dropped
- Admin: generateId over HTTP, env-locked group switches, discard and toggle fixes,
  clearing a secret field keeps the saved key, first provider not auto-default when .env sets AI_MODEL
- Model selector items use unique values
2026-10-03 17:45:41 +09:00
dayuan.jiang 87edf2e19d fix(chat-input): stop template dialogs from sending and fix attachment races
- Template dialogs no longer submit the outer chat form
- Sending is blocked while files or URLs are still extracting
- File and URL extraction no longer drop or resurrect entries
- IME composition Enter no longer sends
- Tool call cards show the error text; keyboard handling on cards fixed
- Template import available when empty, edit dialog resets, saved templates refresh
- Only png/jpeg/gif/webp images accepted, SVG sent as text; PDF objects released
- parse-url request sends the access code
2026-10-03 17:45:41 +09:00
dayuan.jiang 5c7613ea09 fix(diagram): fix autosave staleness and XML repair corrupting valid diagrams
- Autosave guard reads refs, so edits after a theme or dark mode switch are kept
- Duplicate-id check and rename run per page; repair loop no longer quadratic
- autoFixXml no longer breaks style values, rich text " or single-line cells
- extractCompleteMxCells keeps the cell after a self-closing cell
- Better truncation detection; object/UserObject wrapped cells are editable
- Exports for thumbnail, PNG and save are routed by tag instead of a shared resolver
- History stores the full document; storage errors are reported, no auto-deletion of chats
- IndexedDB connection reopens after errors; focus refresh throttled
- Keep ?session= on locale redirect, map zh-Hant to zh-tw for draw.io
2026-10-03 17:45:41 +09:00
dayuan.jiang 79b4c52741 fix(chat): keep saved diagrams and pages when restoring, editing and retrying
- Restored sessions no longer replay the last display_diagram over the saved diagram
- Failed or stopped edit_diagram restores the canvas
- Message snapshots keep the full multi-page document
- "Improve with suggestions" uses the normal send path (headers, xml, retry counters)
- Editing a message keeps its file/URL sections; cached example edits work
- New chat's first autosave no longer resets the UI
- Validation retries counted per user turn; validate-diagram sends the access code
- Cached examples only match the example files on an empty canvas
- Template sends keep attachments and wait for extraction
2026-10-03 17:45:41 +09:00
dayuan.jiang 528b6e54c8 fix(api): require access codes and limit sizes on helper routes
- Shared checkAccessCode for validate-diagram, validate-model, parse-url, verify-access-code
- parse-url: 5 MB streamed body limit; validate-diagram: 5 MB image limit
- validate-model refuses redirects when private URLs are blocked
- Admin settings state shared across module instances via globalThis
- Server model ids: unique slugs (non-ASCII names encoded), duplicates rejected
- Panel Bedrock credentials stored as ADMIN_AWS_* so the DynamoDB client keeps its own
- Locale redirect keeps basePath and query; EdgeOne function drops open CORS and checks the access code
- Providers payload reports whether .env sets a default model
2026-10-03 17:45:41 +09:00
dayuan.jiang 366480426d fix(chat): close credential leaks and harden the chat route
- Vertex: a client-supplied base URL only works with the client's own Vertex key
- Accept only data: URLs for file parts in every message, so the server never downloads them
- Output budget retry accounts for the thinking budget Bedrock/Anthropic add, and reads
  Volcengine, DashScope, SGLang and vLLM rejections; falls back to 16000 once
- x-max-output-tokens can only lower the budget on server credentials
- On server credentials only server models or AI_MODEL entries can be used
- Drop tool results together with the invalid tool calls they belong to
- Count quota tokens as input + output (cached tokens were counted twice)
- Private-URL check for custom base URLs, end Langfuse traces on error/abort/early return
- Fix repairToolCall ordering and placeholder, align edit_diagram prompt with operations
- Panel Bedrock keys are read from ADMIN_AWS_*; forward the access code to EdgeOne
- isMinimalDiagram only treats root cells as an empty canvas
2026-10-03 17:45:41 +09:00
renovate[bot] a45e5b6796 fix(deps): update minor and patch dependencies (#948)
Co-authored-by: renovate[bot] <29139614+renovate[bot]@users.noreply.github.com>
2026-10-02 06:42:01 +00:00
94 changed files with 5374 additions and 1986 deletions
+14 -2
View File
@@ -35,6 +35,7 @@ import { useDictionary } from "@/hooks/use-dictionary"
import { formatMessage } from "@/lib/i18n/utils"
import {
FIXED_CRED_PROVIDERS,
generateId,
PROVIDER_INFO,
type ProviderName,
SUGGESTED_MODELS,
@@ -225,6 +226,7 @@ function ProviderDetail({
</Button>
{suggestions.length > 0 && (
<Select
value=""
disabled={disabled}
onValueChange={(v) => addModel(v)}
>
@@ -390,12 +392,14 @@ function ProviderDetail({
export function ModelsSection({
providers,
envProviders,
envHasDefaultModel,
disabled,
password,
onChange,
}: {
providers: AdminProvider[]
envProviders: EnvProvider[]
envHasDefaultModel: boolean
disabled: boolean
password: string
onChange: (providers: AdminProvider[]) => void
@@ -409,10 +413,16 @@ export function ModelsSection({
const addProvider = (provider: ProviderName) => {
const newProvider: AdminProvider = {
id: crypto.randomUUID(),
// generateId works over plain HTTP; crypto.randomUUID needs HTTPS
id: generateId(),
provider,
models: [],
isDefault: providers.length === 0,
// Only the very first provider becomes the default, and only when
// the env config has no default that it would replace on save
isDefault:
providers.length === 0 &&
!envProviders.some((p) => p.isDefault) &&
!envHasDefaultModel,
}
onChange([...providers, newProvider])
setSelectedId(newProvider.id)
@@ -496,7 +506,9 @@ export function ModelsSection({
))}
</div>
<div className="border-t p-2">
{/* Always empty so picking the same type again still fires */}
<Select
value=""
disabled={disabled}
onValueChange={(v) => addProvider(v as ProviderName)}
>
+49 -14
View File
@@ -37,6 +37,19 @@ import { SettingField } from "./setting-field"
const NAV_GROUP_IDS = ["models", ...SETTING_GROUPS.map((g) => g.id)]
// For each toggleable group, whether any of its settings has a value (from
// the settings file or the environment)
function groupsWithValues(map: SettingsMap): Record<string, boolean> {
const result: Record<string, boolean> = {}
for (const group of SETTING_GROUPS) {
if (!group.toggleable) continue
result[group.id] = !!SETTINGS_BY_GROUP.get(group.id)?.some(
(d) => map[d.key]?.source !== "default",
)
}
return result
}
export default function AdminPage() {
const dict = useDictionary()
// Localized group title/description, keyed by group id
@@ -62,6 +75,8 @@ export default function AdminPage() {
// Models section state
const [providers, setProviders] = useState<AdminProvider[]>([])
const [envProviders, setEnvProviders] = useState<EnvProvider[]>([])
// Whether .env itself sets AI_MODEL (a default the panel would override)
const [envHasDefaultModel, setEnvHasDefaultModel] = useState(false)
const [savedProviders, setSavedProviders] = useState<string>("[]")
const providersDirty = JSON.stringify(providers) !== savedProviders
@@ -88,15 +103,13 @@ export default function AdminPage() {
const map: SettingsMap = {}
for (const s of data.settings) map[s.key] = s
setSettings(map)
// Seed each toggle once from whether the group has configured
// values; don't stomp a user's explicit toggle on later saves
// A group stays on while it still has values (e.g. from env vars
// that saving can't remove); a user's explicit "on" for a group
// with no values yet is kept across saves
setEnabledGroups((prev) => {
const next = { ...prev }
for (const group of SETTING_GROUPS) {
if (!group.toggleable || group.id in next) continue
next[group.id] = !!SETTINGS_BY_GROUP.get(group.id)?.some(
(d) => map[d.key]?.source !== "default",
)
const next = groupsWithValues(map)
for (const id of Object.keys(next)) {
next[id] = next[id] || !!prev[id]
}
return next
})
@@ -108,10 +121,12 @@ export default function AdminPage() {
(data: {
providers: AdminProvider[]
envProviders?: EnvProvider[]
envHasDefaultModel?: boolean
}) => {
setProviders(data.providers)
setSavedProviders(JSON.stringify(data.providers))
setEnvProviders(data.envProviders ?? [])
setEnvHasDefaultModel(!!data.envHasDefaultModel)
},
[],
)
@@ -181,8 +196,9 @@ export default function AdminPage() {
return () => observer.disconnect()
}, [authedPassword])
// value undefined drops the pending change (back to the saved value)
const handleChange = useCallback(
(key: string, value: string | null) => {
(key: string, value: string | null | undefined) => {
setSaveMessage(null)
setErrors((prev) => {
if (!(key in prev)) return prev
@@ -201,7 +217,7 @@ export default function AdminPage() {
value === "" &&
(!state || state.source !== "file") &&
!isSecretValue(state?.value)
if (isRevert || isNoop) {
if (value === undefined || isRevert || isNoop) {
const next = { ...prev }
delete next[key]
return next
@@ -225,9 +241,10 @@ export default function AdminPage() {
const next = { ...prev }
for (const key of keys) {
if (!enabled) {
// Stage deletion only for values currently set
if (settings[key]?.source !== "default")
next[key] = null
// Stage deletion of saved values; drop unsaved input
if (settings[key]?.source === "default")
delete next[key]
else next[key] = null
} else if (next[key] === null) {
delete next[key]
}
@@ -447,6 +464,7 @@ export default function AdminPage() {
<ModelsSection
providers={providers}
envProviders={envProviders}
envHasDefaultModel={envHasDefaultModel}
disabled={!writable || saving}
password={authedPassword}
onChange={(next) => {
@@ -462,6 +480,11 @@ export default function AdminPage() {
const defs = SETTINGS_BY_GROUP.get(group.id) ?? []
const groupOff =
group.toggleable && !enabledGroups[group.id]
// Values from env vars can't be removed here, so the
// group can't be turned off from the panel
const envLocked = defs.some(
(d) => settings[d.key]?.source === "env",
)
const fieldsDisabled = !writable || saving || !!groupOff
const gt = groupText(group.id)
const title = gt?.title ?? group.title
@@ -480,6 +503,11 @@ export default function AdminPage() {
</h2>
{group.toggleable && (
<label
title={
envLocked
? dict.admin.sourceEnvTitle
: undefined
}
className={cn(
"flex cursor-pointer items-center gap-2 rounded-full border px-3 py-1.5 text-xs font-medium transition-colors motion-reduce:transition-none",
enabledGroups[group.id]
@@ -494,7 +522,11 @@ export default function AdminPage() {
checked={
!!enabledGroups[group.id]
}
disabled={!writable || saving}
disabled={
!writable ||
saving ||
envLocked
}
aria-label={formatMessage(
dict.admin.enableGroup,
{ group: title },
@@ -579,6 +611,9 @@ export default function AdminPage() {
setPending({})
setErrors({})
setProviders(JSON.parse(savedProviders))
setEnabledGroups(
groupsWithValues(settings),
)
}}
>
{dict.admin.discard}
+10 -5
View File
@@ -73,8 +73,10 @@ export function SecretInput({
}) {
const dict = useDictionary()
const [show, setShow] = useState(false)
// The stored marker as it was at mount, to revert to on empty
const [original] = useState(value)
// The stored marker to revert to on empty. Refreshed whenever the parent
// passes server state (a marker or nothing), e.g. after a save.
const [original, setOriginal] = useState(value)
if (typeof value !== "string" && value !== original) setOriginal(value)
const hadStored = isSecretValue(original)
const text = typeof value === "string" ? value : ""
const placeholder = isSecretValue(value)
@@ -146,7 +148,8 @@ export function SettingField({
pendingValue: string | null | undefined
error?: string
disabled: boolean
onChange: (value: string | null) => void
// undefined drops the pending change (back to the saved value)
onChange: (value: string | null | undefined) => void
}) {
const dict = useDictionary()
const isDirty = pendingValue !== undefined
@@ -226,16 +229,18 @@ export function SettingField({
case "secret":
control = (
<div className="w-full max-w-md">
{/* Clearing a saved secret reverts to it; the X button deletes */}
<SecretInput
id={inputId}
keepOnEmpty={source === "file"}
value={
isDirty
? (pendingValue ?? "")
: (secretState ?? currentValue)
: (secretState ?? undefined)
}
disabled={disabled}
onChange={(v) =>
onChange(typeof v === "string" ? v : "")
onChange(typeof v === "string" ? v : undefined)
}
/>
</div>
+12 -17
View File
@@ -37,7 +37,6 @@ export default function Home() {
)
const chatPanelRef = useRef<ImperativePanelHandle>(null)
const isMobileRef = useRef(false)
// Load preferences from localStorage after mount
useEffect(() => {
@@ -48,7 +47,9 @@ export default function Home() {
const currentLocale = pathParts[0]
if (currentLocale !== savedLocale) {
pathParts[0] = savedLocale
router.replace(`/${pathParts.join("/")}`)
// Keep the query (e.g. ?session=) and hash
const { search, hash } = window.location
router.replace(`/${pathParts.join("/")}${search}${hash}`)
return // Wait for redirect
}
}
@@ -106,27 +107,17 @@ export default function Home() {
resetDrawioReady()
}
// Check mobile - reset draw.io before crossing breakpoint
const isInitialRenderRef = useRef(true)
// 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.
useEffect(() => {
const checkMobile = () => {
const newIsMobile = window.innerWidth < 768
if (
!isInitialRenderRef.current &&
newIsMobile !== isMobileRef.current
) {
setIsDrawioReady(false)
resetDrawioReady()
}
isMobileRef.current = newIsMobile
isInitialRenderRef.current = false
setIsMobile(newIsMobile)
setIsMobile(window.innerWidth < 768)
}
checkMobile()
window.addEventListener("resize", checkMobile)
return () => window.removeEventListener("resize", checkMobile)
}, [resetDrawioReady])
}, [])
const toggleChatPanel = () => {
const panel = chatPanelRef.current
@@ -193,7 +184,11 @@ export default function Home() {
noExitBtn: true,
dark:
darkMode || drawioUi === "dark",
lang: currentLang,
// draw.io names Traditional Chinese "zh-tw"
lang:
currentLang === "zh-Hant"
? "zh-tw"
: currentLang,
// Enable offline mode in Electron to disable external service calls
...(isElectron && {
offline: true,
+8 -1
View File
@@ -7,7 +7,11 @@ import {
mergeSecrets,
validateAdminProviders,
} from "@/lib/admin/providers"
import { isSettingsWritable, saveSettings } from "@/lib/admin/settings"
import {
getEnvFallback,
isSettingsWritable,
saveSettings,
} from "@/lib/admin/settings"
import { loadEnvServerModelsConfig } from "@/lib/server-model-config"
export const runtime = "nodejs"
@@ -33,6 +37,9 @@ async function payload() {
models: p.models,
isDefault: !!p.default && !adminHasDefault,
})) ?? [],
// Whether .env sets a default model. getEnvFallback skips the value
// the panel overlays onto process.env, so a panel default doesn't count.
envHasDefaultModel: !!getEnvFallback("AI_MODEL"),
}
}
+87 -96
View File
@@ -12,13 +12,17 @@ import fs from "fs/promises"
import { jsonrepair } from "jsonrepair"
import path from "path"
import { z } from "zod"
import { checkAccessCode } from "@/lib/access-code"
import {
getAIModel,
SINGLE_SYSTEM_PROVIDERS,
supportsPromptCaching,
usesServerCredentials,
} from "@/lib/ai-providers"
import { findCachedResponse } from "@/lib/cached-responses"
import {
dropInvalidToolCalls,
fixToolInputJson,
isMinimalDiagram,
replaceHistoricalToolInputs,
validateFileParts,
@@ -29,6 +33,7 @@ import {
recordTokenUsage,
} from "@/lib/dynamo-quota-manager"
import {
endTrace,
getTelemetryConfig,
setTraceInput,
setTraceOutput,
@@ -38,7 +43,11 @@ import {
resolveMaxOutputTokens,
withOutputTokenLimitFallback,
} from "@/lib/output-token-limit"
import { findServerModelById } from "@/lib/server-model-config"
import {
type FlattenedServerModel,
findServerModelById,
} from "@/lib/server-model-config"
import { allowPrivateUrls, isPrivateUrl } from "@/lib/ssrf-protection"
import { getSystemPrompt } from "@/lib/system-prompts"
import { getUserIdFromRequest } from "@/lib/user-id"
@@ -76,24 +85,14 @@ function createCachedStreamResponse(xml: string): Response {
return createUIMessageStreamResponse({ stream })
}
// Responses streamed from the model, whose trace streamText's callbacks end
const modelStreamResponses = new WeakSet<Response>()
// Inner handler function
async function handleChatRequest(req: Request): Promise<Response> {
// Check for access code
const accessCodes =
process.env.ACCESS_CODE_LIST?.split(",")
.map((code) => code.trim())
.filter(Boolean) || []
if (accessCodes.length > 0) {
const accessCodeHeader = req.headers.get("x-access-code")
if (!accessCodeHeader || !accessCodes.includes(accessCodeHeader)) {
return Response.json(
{
error: "Invalid or missing access code. Please configure it in Settings.",
},
{ status: 401 },
)
}
}
const accessDenied = checkAccessCode(req)
if (accessDenied) return accessDenied
const body = await req.json()
const { messages, xml, previousXml, sessionId } = body
@@ -192,6 +191,15 @@ async function handleChatRequest(req: Request): Promise<Response> {
baseUrl = `${origin}/api/edgeai`
}
// Same rule as validate-model: with ALLOW_PRIVATE_URLS=false a request may
// not point the server at a private or internal address
if (baseUrl && !allowPrivateUrls() && (await isPrivateUrl(baseUrl))) {
return Response.json(
{ error: "Private or internal base URLs are not allowed." },
{ status: 400 },
)
}
// Get cookie header for EdgeOne authentication (eo_token, eo_time)
const cookieHeader = req.headers.get("cookie")
@@ -201,8 +209,9 @@ async function handleChatRequest(req: Request): Promise<Response> {
baseUrlEnv?: string
provider?: string
} = {}
let serverModel: FlattenedServerModel | null = null
if (selectedModelId?.startsWith("server:")) {
const serverModel = await findServerModelById(selectedModelId)
serverModel = await findServerModelById(selectedModelId)
console.log(
`[Server Model Lookup] ID: ${selectedModelId}, Found: ${!!serverModel}, Provider: ${serverModel?.provider}`,
)
@@ -221,7 +230,8 @@ async function handleChatRequest(req: Request): Promise<Response> {
provider: serverModelConfig.provider || provider,
baseUrl,
apiKey: req.headers.get("x-ai-api-key"),
modelId: req.headers.get("x-ai-model"),
// A server model runs the model it was configured with, whatever the header says
modelId: serverModel?.modelId || req.headers.get("x-ai-model"),
// AWS Bedrock credentials
awsAccessKeyId: req.headers.get("x-aws-access-key-id"),
awsSecretAccessKey: req.headers.get("x-aws-secret-access-key"),
@@ -231,11 +241,14 @@ async function handleChatRequest(req: Request): Promise<Response> {
...serverModelConfig,
// Vertex AI credentials (Express Mode)
vertexApiKey: req.headers.get("x-vertex-api-key"),
// Pass cookies for EdgeOne Pages authentication
...(provider === "edgeone" &&
cookieHeader && {
headers: { cookie: cookieHeader },
}),
// Pass cookies for EdgeOne Pages authentication, and the access code,
// which the EdgeOne function checks too
...(provider === "edgeone" && {
headers: {
...(cookieHeader && { cookie: cookieHeader }),
"x-access-code": req.headers.get("x-access-code") || "",
},
}),
}
// Read minimal style preference from header
@@ -254,12 +267,32 @@ async function handleChatRequest(req: Request): Promise<Response> {
provider: resolvedProvider,
} = getAIModel(clientOverrides)
// On the server's own keys, only run models the server offers: a server
// model picked by id (its model name is fixed above) or one in AI_MODEL.
// With their own key, users can run any model.
const onServerCredentials = usesServerCredentials(
resolvedProvider,
clientOverrides,
)
const envModels =
process.env.AI_MODEL?.split(",").map((m) => m.trim()) || []
if (onServerCredentials && !serverModel && !envModels.includes(modelId)) {
return Response.json(
{
error: `Model "${modelId}" is not available on this server. Add your own API key in Settings to use it.`,
},
{ status: 400 },
)
}
// Retry with a smaller budget if the provider rejects the requested one
const model = withOutputTokenLimitFallback(baseModel)
// User setting wins over server env, so desktop users can raise it themselves
// The user setting can raise the budget only on their own key (desktop users
// can still raise it themselves); on the server's keys it can only lower it
const maxOutputTokens = resolveMaxOutputTokens(
req.headers.get("x-max-output-tokens"),
onServerCredentials,
)
console.log(`[maxOutputTokens] ${maxOutputTokens}`)
@@ -340,32 +373,9 @@ ${userInputText}
)
// Filter out tool-calls with invalid inputs (from failed repair or interrupted streaming)
// Bedrock API rejects messages where toolUse.input is not a valid JSON object
enhancedMessages = enhancedMessages
.map((msg: any) => {
if (msg.role !== "assistant" || !Array.isArray(msg.content)) {
return msg
}
const filteredContent = msg.content.filter((part: any) => {
if (part.type === "tool-call") {
// Check if input is a valid object (not null, undefined, or empty)
if (
!part.input ||
typeof part.input !== "object" ||
Object.keys(part.input).length === 0
) {
console.warn(
`[route.ts] Filtering out tool-call with invalid input:`,
{ toolName: part.toolName, input: part.input },
)
return false
}
}
return true
})
return { ...msg, content: filteredContent }
})
.filter((msg: any) => msg.content && msg.content.length > 0)
// and their results. Bedrock API rejects messages where toolUse.input is not a valid
// JSON object, and every provider rejects a tool result whose call is gone.
enhancedMessages = dropInvalidToolCalls(enhancedMessages)
// DEBUG: Log modelMessages structure (what's being sent to AI)
console.log("[route.ts] Model messages count:", enhancedMessages.length)
@@ -410,7 +420,7 @@ ${userInputText}
contentParts.push({
type: "image",
image: filePart.url,
mimeType: filePart.mediaType,
mediaType: filePart.mediaType,
})
}
@@ -471,7 +481,7 @@ ${previousXml}
${xml || ""}
"""
IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on the canvas right now. The user can manually add, delete, or modify shapes directly in draw.io. Always count and describe elements based on the CURRENT XML, not on what you previously generated. If both previous and current XML are shown, compare them to understand what the user changed. When using edit_diagram, COPY search patterns exactly from the CURRENT XML - attribute order matters!`
IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on the canvas right now. The user can manually add, delete, or modify shapes directly in draw.io. Always count and describe elements based on the CURRENT XML, not on what you previously generated. If both previous and current XML are shown, compare them to understand what the user changed.`
const systemMessages = isSingleSystemProvider
? [
@@ -528,23 +538,11 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
error.name === "AI_InvalidToolInputError"
) {
try {
// Pre-process to fix common LLM JSON errors that jsonrepair can't handle
let inputToRepair = toolCall.input
if (typeof inputToRepair === "string") {
// Fix `:=` instead of `: ` (LLM sometimes generates this)
inputToRepair = inputToRepair.replace(/:=/g, ": ")
// Fix `= "` instead of `: "`
inputToRepair = inputToRepair.replace(/=\s*"/g, ': "')
// Fix inconsistent quote escaping in XML attributes within JSON strings
// Pattern: attribute="value\" where opening quote is unescaped but closing is escaped
// Example: y="-20\" should be y=\"-20\"
inputToRepair = inputToRepair.replace(
/(\w+)="([^"]*?)\\"/g,
'$1=\\"$2\\"',
)
}
// Use jsonrepair to fix truncated JSON
const repairedInput = jsonrepair(inputToRepair)
// Pre-process to fix common LLM JSON errors that jsonrepair can't handle,
// then use jsonrepair to fix truncated JSON
const repairedInput = jsonrepair(
fixToolInputJson(toolCall.input),
)
console.log(
`[repairToolCall] Repaired truncated JSON for tool: ${toolCall.toolName}`,
)
@@ -554,26 +552,8 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
`[repairToolCall] Failed to repair JSON for tool: ${toolCall.toolName}`,
repairError,
)
// Return a placeholder input to avoid API errors in multi-step
// The tool will fail gracefully on client side
if (toolCall.toolName === "edit_diagram") {
return {
...toolCall,
input: {
operations: [],
_error: "JSON repair failed - no operations to apply",
},
}
}
if (toolCall.toolName === "display_diagram") {
return {
...toolCall,
input: {
xml: "",
_error: "JSON repair failed - empty diagram",
},
}
}
// Keep the original error, so the model and the client see why
// the input was rejected and the model can retry the call
return null
}
}
@@ -596,7 +576,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)
// Include all 4 token types: input, output, cache read, cache write
// inputTokens already includes cache reads and writes in AI SDK 6
if (
isQuotaEnabled() &&
!hasOwnApiKey &&
@@ -605,12 +585,16 @@ IMPORTANT: The "Current diagram XML" is the SINGLE SOURCE OF TRUTH for what's on
) {
const totalTokens =
(totalUsage.inputTokens || 0) +
(totalUsage.outputTokens || 0) +
(totalUsage.cachedInputTokens || 0) +
(totalUsage.inputTokenDetails?.cacheWriteTokens || 0)
(totalUsage.outputTokens || 0)
recordTokenUsage(userId, totalTokens)
}
},
// onFinish is skipped when the stream fails or is aborted, so end the trace here
onError: ({ error }) => {
console.error(error) // what AI SDK does without an onError
endTrace()
},
onAbort: () => endTrace(),
tools: {
// Client-side tool that will be executed on the client
display_diagram: {
@@ -782,7 +766,7 @@ Call this tool to get shape names and usage syntax for a specific library.`,
}),
})
return result.toUIMessageStreamResponse({
const response = result.toUIMessageStreamResponse({
sendReasoning: true,
messageMetadata: ({ part }) => {
if (part.type === "finish") {
@@ -796,6 +780,8 @@ Call this tool to get shape names and usage syntax for a specific library.`,
return undefined
},
})
modelStreamResponses.add(response)
return response
}
// Helper to categorize errors and return appropriate response
@@ -862,11 +848,16 @@ function handleError(error: unknown): Response {
// Wrap handler with error handling
async function safeHandler(req: Request): Promise<Response> {
let response: Response
try {
return await handleChatRequest(req)
response = await handleChatRequest(req)
} catch (error) {
return handleError(error)
response = handleError(error)
}
// Early returns, cache hits and errors never reach streamText's callbacks,
// so their Langfuse trace has to be ended here
if (!modelStreamResponses.has(response)) endTrace()
return response
}
// Wrap with Langfuse observe (if configured)
+40 -1
View File
@@ -1,9 +1,11 @@
import { extractFromHtml } from "@extractus/article-extractor"
import { NextResponse } from "next/server"
import TurndownService from "turndown"
import { checkAccessCode } from "@/lib/access-code"
import { isPrivateUrl } from "@/lib/ssrf-protection"
const MAX_CONTENT_LENGTH = 150000 // Match PDF limit
const MAX_RESPONSE_BYTES = 5 * 1024 * 1024
const EXTRACT_TIMEOUT_MS = 15000
const USER_AGENT = "Mozilla/5.0 (compatible; NextAIDrawio/1.0)"
@@ -32,7 +34,36 @@ function detectCharset(
}
}
// Read the response body, giving up once it passes MAX_RESPONSE_BYTES so a
// huge download can't exhaust server memory. Returns null when too large.
async function readLimitedBody(
response: Response,
): Promise<ArrayBuffer | null> {
if (Number(response.headers.get("content-length")) > MAX_RESPONSE_BYTES) {
return null
}
if (!response.body) return new ArrayBuffer(0)
const reader = response.body.getReader()
const chunks: Uint8Array[] = []
let total = 0
while (true) {
const { done, value } = await reader.read()
if (done) break
total += value.byteLength
if (total > MAX_RESPONSE_BYTES) {
await reader.cancel()
return null
}
chunks.push(value)
}
return new Blob(chunks as BlobPart[]).arrayBuffer()
}
export async function POST(req: Request) {
const accessError = checkAccessCode(req)
if (accessError) return accessError
try {
const { url } = await req.json()
@@ -97,7 +128,15 @@ export async function POST(req: Request) {
)
}
const buffer = await response.arrayBuffer()
const buffer = await readLimitedBody(response)
if (!buffer) {
return NextResponse.json(
{
error: `Page exceeds the ${MAX_RESPONSE_BYTES / 1024 / 1024} MB download limit`,
},
{ status: 413 },
)
}
const charset = detectCharset(contentType, buffer)
html = new TextDecoder(charset).decode(buffer)
} catch (err: any) {
+15
View File
@@ -4,6 +4,7 @@
*/
import { streamObject } from "ai"
import { checkAccessCode } from "@/lib/access-code"
import { getValidationModel } from "@/lib/ai-providers"
import { VALIDATION_SYSTEM_PROMPT } from "@/lib/validation-prompts"
import {
@@ -13,6 +14,9 @@ import {
export const maxDuration = 30
// Data URL length cap (~3.75 MB of PNG), well above a normal diagram capture
const MAX_IMAGE_DATA_LENGTH = 5 * 1024 * 1024
interface ValidateDiagramRequest {
imageData: string // Base64 PNG data URL
sessionId?: string
@@ -44,6 +48,10 @@ function createStreamingResponse(result: ValidationResult): Response {
}
export async function POST(req: Request): Promise<Response> {
// Uses the server's model credentials, so require the access code
const accessError = checkAccessCode(req)
if (accessError) return accessError
try {
// Check if VLM validation is enabled (default: true)
const enableValidation = process.env.ENABLE_VLM_VALIDATION !== "false"
@@ -72,6 +80,13 @@ export async function POST(req: Request): Promise<Response> {
)
}
if (imageData.length > MAX_IMAGE_DATA_LENGTH) {
return Response.json(
{ error: "Image data too large" },
{ status: 413 },
)
}
// Get the validation model
let model
try {
+49 -4
View File
@@ -10,6 +10,7 @@ import { createOpenRouter } from "@openrouter/ai-sdk-provider"
import { generateText } from "ai"
import { NextResponse } from "next/server"
import { createOllama } from "ollama-ai-provider-v2"
import { checkAccessCode } from "@/lib/access-code"
import {
AIHUBMIX_APP_CODE,
isAihubmixStandardBaseURL,
@@ -33,7 +34,24 @@ interface ValidateRequest {
vertexApiKey?: string // Express Mode API key
}
// With private URLs blocked, a public baseUrl could still redirect the
// request to an internal host, so redirects are refused in that case.
function redirectGuardedFetch(): typeof fetch | undefined {
if (allowPrivateUrls()) return undefined
return async (input, init) => {
const response = await fetch(input, { ...init, redirect: "manual" })
if (response.status >= 300 && response.status < 400) {
throw new Error("Redirects are not allowed for custom base URLs")
}
return response
}
}
export async function POST(req: Request) {
// Lets the server send requests to arbitrary URLs, so require the access code
const accessError = checkAccessCode(req)
if (accessError) return accessError
try {
const body: ValidateRequest = await req.json()
const {
@@ -91,6 +109,7 @@ export async function POST(req: Request) {
)
}
const guardedFetch = redirectGuardedFetch()
let model: any
switch (provider) {
@@ -98,6 +117,7 @@ export async function POST(req: Request) {
const openai = createOpenAI({
apiKey,
...(baseUrl && { baseURL: baseUrl }),
fetch: guardedFetch,
})
model = openai.chat(modelId)
break
@@ -107,6 +127,7 @@ export async function POST(req: Request) {
const anthropic = createAnthropic({
apiKey,
baseURL: baseUrl || "https://api.anthropic.com/v1",
fetch: guardedFetch,
})
model = anthropic(modelId)
break
@@ -116,6 +137,7 @@ export async function POST(req: Request) {
const google = createGoogleGenerativeAI({
apiKey,
...(baseUrl && { baseURL: baseUrl }),
fetch: guardedFetch,
})
model = google(modelId)
break
@@ -125,6 +147,7 @@ export async function POST(req: Request) {
const vertex = createVertex({
apiKey: vertexApiKey,
...(baseUrl && { baseURL: baseUrl }),
fetch: guardedFetch,
})
model = vertex(modelId)
break
@@ -134,6 +157,7 @@ export async function POST(req: Request) {
const azure = createOpenAI({
apiKey,
baseURL: baseUrl,
fetch: guardedFetch,
})
model = azure.chat(modelId)
break
@@ -153,6 +177,7 @@ export async function POST(req: Request) {
const openrouter = createOpenRouter({
apiKey,
...(baseUrl && { baseURL: baseUrl }),
fetch: guardedFetch,
})
model = openrouter(modelId)
break
@@ -174,6 +199,7 @@ export async function POST(req: Request) {
const aihubmixCompatible = createOpenAI({
apiKey,
baseURL: baseUrl,
fetch: guardedFetch,
})
model = aihubmixCompatible.chat(modelId)
}
@@ -185,6 +211,7 @@ export async function POST(req: Request) {
const ds = createDeepSeek({
apiKey,
...(baseUrl && { baseURL: baseUrl }),
fetch: guardedFetch,
})
model = ds(modelId)
} else {
@@ -197,6 +224,7 @@ export async function POST(req: Request) {
const sf = createOpenAI({
apiKey,
baseURL: baseUrl || "https://api.siliconflow.cn/v1",
fetch: guardedFetch,
})
model = sf.chat(modelId)
break
@@ -213,6 +241,7 @@ export async function POST(req: Request) {
baseUrl ||
process.env.OLLAMA_BASE_URL ||
"https://ollama.com/api",
fetch: guardedFetch,
...(ollamaApiKey && {
headers: { Authorization: `Bearer ${ollamaApiKey}` },
}),
@@ -225,6 +254,7 @@ export async function POST(req: Request) {
const gw = createGateway({
apiKey,
...(baseUrl && { baseURL: baseUrl }),
fetch: guardedFetch,
})
model = gw(modelId)
break
@@ -232,13 +262,16 @@ export async function POST(req: Request) {
case "edgeone": {
// EdgeOne uses OpenAI-compatible API via Edge Functions
// Need to pass cookies for EdgeOne Pages authentication
// Need to pass cookies for EdgeOne Pages authentication,
// and the access code, which the edge function also checks
const cookieHeader = req.headers.get("cookie") || ""
const edgeone = createOpenAI({
apiKey: "edgeone", // EdgeOne doesn't require API key
baseURL: baseUrl || "/api/edgeai",
fetch: guardedFetch,
headers: {
cookie: cookieHeader,
"x-access-code": req.headers.get("x-access-code") || "",
},
})
model = edgeone.chat(modelId)
@@ -250,6 +283,7 @@ export async function POST(req: Request) {
const sglang = createOpenAI({
apiKey: apiKey || "not-needed",
baseURL: baseUrl || "http://127.0.0.1:8000/v1",
fetch: guardedFetch,
})
model = sglang.chat(modelId)
break
@@ -267,12 +301,14 @@ export async function POST(req: Request) {
const doubao = createDeepSeek({
apiKey,
baseURL: doubaoBaseUrl,
fetch: guardedFetch,
})
model = doubao(modelId)
} else {
const doubao = createOpenAI({
apiKey,
baseURL: doubaoBaseUrl,
fetch: guardedFetch,
})
model = doubao.chat(modelId)
}
@@ -286,7 +322,7 @@ export async function POST(req: Request) {
try {
// Initiate a streaming request (required for QwQ-32B and certain Qwen3 models)
const response = await fetch(
const response = await (guardedFetch ?? fetch)(
`${baseURL}/chat/completions`,
{
method: "POST",
@@ -307,9 +343,15 @@ export async function POST(req: Request) {
)
if (!response.ok) {
const errorText = await response.text()
// Log the body but return only the status: the
// caller chooses baseUrl, so the body may come from
// any host the server can reach
console.error(
"[validate-model] ModelScope error body:",
await response.text(),
)
throw new Error(
`ModelScope API error (${response.status}): ${errorText}`,
`ModelScope API error (${response.status})`,
)
}
@@ -360,12 +402,14 @@ export async function POST(req: Request) {
const minimax = createAnthropic({
apiKey,
baseURL: minimaxBaseUrl,
fetch: guardedFetch,
})
model = minimax.chat(modelId)
} else {
const minimax = createOpenAI({
apiKey,
baseURL: minimaxBaseUrl,
fetch: guardedFetch,
})
model = minimax.chat(modelId)
}
@@ -398,6 +442,7 @@ export async function POST(req: Request) {
const openai = createOpenAI({
apiKey,
baseURL,
fetch: guardedFetch,
})
model = openai.chat(modelId)
break
+4 -24
View File
@@ -1,29 +1,9 @@
import { checkAccessCode } from "@/lib/access-code"
export async function POST(req: Request) {
const accessCodes =
process.env.ACCESS_CODE_LIST?.split(",")
.map((code) => code.trim())
.filter(Boolean) || []
// If no access codes configured, verification always passes
if (accessCodes.length === 0) {
return Response.json({
valid: true,
message: "No access code required",
})
}
const accessCodeHeader = req.headers.get("x-access-code")
if (!accessCodeHeader) {
if (checkAccessCode(req)) {
return Response.json(
{ valid: false, message: "Access code is required" },
{ status: 401 },
)
}
if (!accessCodes.includes(accessCodeHeader)) {
return Response.json(
{ valid: false, message: "Invalid access code" },
{ valid: false, message: "Invalid or missing access code" },
{ status: 401 },
)
}
+68 -36
View File
@@ -11,7 +11,9 @@ import {
} from "lucide-react"
import type React from "react"
import {
type Dispatch,
forwardRef,
type SetStateAction,
useCallback,
useEffect,
useImperativeHandle,
@@ -41,9 +43,20 @@ import { FilePreviewList } from "./file-preview-list"
const MAX_IMAGE_SIZE = 2 * 1024 * 1024 // 2MB
const MAX_FILES = 5
// Image formats every supported model provider accepts (SVG is read as text)
const SUPPORTED_IMAGE_TYPES = [
"image/png",
"image/jpeg",
"image/gif",
"image/webp",
]
function isValidFileType(file: File): boolean {
return file.type.startsWith("image/") || isPdfFile(file) || isTextFile(file)
return (
SUPPORTED_IMAGE_TYPES.includes(file.type) ||
isPdfFile(file) ||
isTextFile(file)
)
}
function formatFileSize(bytes: number): string {
@@ -164,7 +177,7 @@ interface ChatInputProps {
{ text: string; charCount: number; isExtracting: boolean }
>
urlData?: Map<string, UrlData>
onUrlChange?: (data: Map<string, UrlData>) => void
onUrlChange?: Dispatch<SetStateAction<Map<string, UrlData>>>
sessionId?: string
error?: Error | null
@@ -244,6 +257,11 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
// Allow retry when there's an error (even if status is still "streaming" or "submitted")
const isDisabled =
(status === "streaming" || status === "submitted") && !error
// Block sending until attached files and URLs have their text, otherwise
// their content would be silently dropped
const isExtractingAttachments =
files.some((file) => pdfData.get(file)?.isExtracting) ||
Array.from(urlData?.values() ?? []).some((d) => d.isExtracting)
const adjustTextareaHeight = useCallback(() => {
const textarea = textareaRef.current
@@ -281,6 +299,9 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
}
const handleKeyDown = (e: React.KeyboardEvent) => {
// Enter that confirms an IME candidate must not send the message
if (e.nativeEvent.isComposing || e.keyCode === 229) return
const shouldSend =
sendShortcut === "enter"
? e.key === "Enter" &&
@@ -292,7 +313,12 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
if (shouldSend) {
e.preventDefault()
const form = e.currentTarget.closest("form")
if (form && input.trim() && !isDisabled) {
if (
form &&
input.trim() &&
!isDisabled &&
!isExtractingAttachments
) {
form.requestSubmit()
}
}
@@ -380,13 +406,9 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
if (isDisabled) return
const droppedFiles = e.dataTransfer.files
const supportedFiles = Array.from(droppedFiles).filter((file) =>
isValidFileType(file),
)
// Let validateFiles show a toast for unsupported types
const { validFiles, errors } = validateFiles(
supportedFiles,
Array.from(e.dataTransfer.files),
files.length,
dict,
)
@@ -401,33 +423,34 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
setIsExtractingUrl(true)
// Use functional updates so a removal or send made while extracting
// is not overwritten when the request finishes
try {
const existing = urlData
? new Map(urlData)
: new Map<string, UrlData>()
existing.set(url, {
url,
title: url,
content: "",
charCount: 0,
isExtracting: true,
})
onUrlChange(existing)
onUrlChange((prev) =>
new Map(prev).set(url, {
url,
title: url,
content: "",
charCount: 0,
isExtracting: true,
}),
)
const data = await extractUrlContent(url)
const newUrlData = new Map(existing)
newUrlData.set(url, data)
onUrlChange(newUrlData)
// Skip if the URL was removed while extracting
onUrlChange((prev) =>
prev.has(url) ? new Map(prev).set(url, data) : prev,
)
setShowUrlDialog(false)
} catch (error) {
// Remove the URL from the data map on error
const newUrlData = urlData
? new Map(urlData)
: new Map<string, UrlData>()
newUrlData.delete(url)
onUrlChange(newUrlData)
onUrlChange((prev) => {
const next = new Map(prev)
next.delete(url)
return next
})
showErrorToast(
<span className="text-muted-foreground">
{error instanceof Error
@@ -463,11 +486,12 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
urlData={urlData}
onRemoveUrl={
onUrlChange
? (url) => {
const next = new Map(urlData)
next.delete(url)
onUrlChange(next)
}
? (url) =>
onUrlChange((prev) => {
const next = new Map(prev)
next.delete(url)
return next
})
: undefined
}
/>
@@ -559,7 +583,7 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
ref={fileInputRef}
className="hidden"
onChange={handleFileChange}
accept="image/*,.pdf,application/pdf,text/*,.md,.markdown,.json,.csv,.xml,.yaml,.yml,.toml"
accept="image/png,image/jpeg,image/gif,image/webp,.svg,.pdf,application/pdf,text/*,.md,.markdown,.json,.csv,.xml,.yaml,.yml,.toml"
multiple
disabled={isDisabled}
/>
@@ -588,7 +612,11 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
) : (
<Button
type="submit"
disabled={isDisabled || !input.trim()}
disabled={
isDisabled ||
isExtractingAttachments ||
!input.trim()
}
size="sm"
className="h-8 px-4 rounded-xl font-medium shadow-sm"
aria-label={dict.chat.send}
@@ -629,7 +657,11 @@ export const ChatInput = forwardRef<ChatInputRef, ChatInputProps>(
<TemplateCreateDialog
open={showSaveAsTemplate}
onOpenChange={setShowSaveAsTemplate}
onSuccess={() => setShowSaveAsTemplate(false)}
onSuccess={() => {
setShowSaveAsTemplate(false)
// Let the template list in the lobby reload
window.dispatchEvent(new Event("templatesChanged"))
}}
initialPrompt={input.trim()}
/>
</form>
+41 -5
View File
@@ -129,12 +129,14 @@ const getMessageTextContent = (message: UIMessage): string => {
.join("\n")
}
// Matches the [PDF: ...], [File: ...] and [URL: ...] sections appended to the user's text
export const APPENDED_FILE_SECTIONS_PATTERN =
/\n\n\[(PDF|File|URL):\s*[^\]]+\]\n[\s\S]*$/
// Get only the user's original text, excluding appended file content
const getUserOriginalText = (message: UIMessage): string => {
const fullText = getMessageTextContent(message)
// Strip out [PDF: ...], [File: ...], and [URL: ...] sections that were appended
const filePattern = /\n\n\[(PDF|File|URL):\s*[^\]]+\]\n[\s\S]*$/
return fullText.replace(filePattern, "").trim()
return fullText.replace(APPENDED_FILE_SECTIONS_PATTERN, "").trim()
}
interface SessionMetadata {
@@ -458,6 +460,11 @@ export function ChatMessageDisplay({
messages.length > 0 ? [messages[messages.length - 1]] : []
messagesToProcess.forEach((message) => {
// Messages restored from a saved session were applied before it was
// saved; the saved diagram is authoritative, so don't replay them
const isRestoredMessage =
loadedMessageIdsRef?.current.has(message.id) ?? false
if (message.parts) {
message.parts.forEach((part) => {
if (part.type?.startsWith("tool-")) {
@@ -475,6 +482,8 @@ export function ChatMessageDisplay({
})
}
if (isRestoredMessage) return
if (
part.type === "tool-display_diagram" &&
input?.xml
@@ -541,6 +550,32 @@ export function ChatMessageDisplay({
part.type === "tool-edit_diagram" &&
input?.operations
) {
// Failed or stopped: drop the queued preview. If the original
// XML is still stored, the tool handler never ran (user pressed
// stop), so undo the streamed preview here.
if (state === "output-error") {
if (
pendingEditRef.current?.toolCallId ===
toolCallId &&
editDebounceTimeoutRef.current
) {
clearTimeout(editDebounceTimeoutRef.current)
editDebounceTimeoutRef.current = null
pendingEditRef.current = null
}
const originalXml =
editDiagramOriginalXmlRef.current.get(
toolCallId,
)
if (originalXml) {
editDiagramOriginalXmlRef.current.delete(
toolCallId,
)
onDisplayChart(originalXml, true)
}
return
}
const completeOps = getCompleteOperations(
input.operations as DiagramOperation[],
)
@@ -610,9 +645,10 @@ export function ChatMessageDisplay({
origXml,
pending.operations,
)
handleDisplayChart(
// Load the full document so other pages stay intact
onDisplayChart(
editedXml,
false,
true,
)
lastProcessedXmlRef.current.set(
pending.toolCallId +
+103 -50
View File
@@ -32,6 +32,7 @@ import { useSessionManager } from "@/hooks/use-session-manager"
import { useValidateDiagram } from "@/hooks/use-validate-diagram"
import { getApiEndpoint } from "@/lib/base-path"
import { findCachedResponse } from "@/lib/cached-responses"
import { isMinimalDiagram } from "@/lib/chat-helpers"
import type { DrawioTheme } from "@/lib/drawio-themes"
import { formatMessage } from "@/lib/i18n/utils"
import { isPdfFile, isTextFile } from "@/lib/pdf-utils"
@@ -40,9 +41,12 @@ import { STORAGE_KEYS } from "@/lib/storage"
import type { UrlData } from "@/lib/url-utils"
import { type FileData, useFileProcessor } from "@/lib/use-file-processor"
import { useQuotaManager } from "@/lib/use-quota-manager"
import { cn, formatXML, isRealDiagram } from "@/lib/utils"
import { cn, formatXML, isRealDiagram, wrapWithMxFile } from "@/lib/utils"
import type { ValidationState } from "./chat/ValidationCard"
import { ChatMessageDisplay } from "./chat-message-display"
import {
APPENDED_FILE_SECTIONS_PATTERN,
ChatMessageDisplay,
} from "./chat-message-display"
import { DevXmlSimulator } from "./dev-xml-simulator"
// localStorage keys for persistence
@@ -107,6 +111,18 @@ function hasToolErrors(messages: ChatMessage[]): boolean {
return lastToolPart?.state === TOOL_ERROR_STATE
}
/**
* Snapshots keep the full multi-page document, but the model only sees and
* edits the first page, so give it the first page's mxGraphModel.
* Older snapshots already hold a single mxGraphModel and are returned as is.
*/
function getFirstPageXml(xml: string): string {
if (!xml.includes("<mxfile")) return xml
const doc = new DOMParser().parseFromString(xml, "text/xml")
const model = doc.querySelector("diagram")?.querySelector("mxGraphModel")
return model ? formatXML(new XMLSerializer().serializeToString(model)) : xml
}
export default function ChatPanel({
isVisible,
onToggleVisibility,
@@ -336,19 +352,8 @@ export default function ChatPanel({
localStorage.setItem(STORAGE_KEYS.maxOutputTokens, digitsOnly)
}, [])
// Ref to store the sendMessage function for use in callbacks
const sendMessageRef = useRef<typeof sendMessage | null>(null)
// Callback to improve diagram with validation suggestions
const handleImproveWithSuggestions = useCallback((feedback: string) => {
if (sendMessageRef.current) {
// Send the feedback as a new user message to trigger regeneration
sendMessageRef.current({
role: "user",
parts: [{ type: "text", text: feedback }],
})
}
}, [])
// Failed VLM validations in the current user turn (reset on user action)
const validationRetryCountRef = useRef(0)
// VLM validation hook using AI SDK's useObject
const { validateWithFallback } = useValidateDiagram()
@@ -357,6 +362,7 @@ export default function ChatPanel({
const { handleToolCall } = useDiagramToolHandlers({
partialXmlRef,
editDiagramOriginalXmlRef,
validationRetryCountRef,
chartXMLRef,
onDisplayChart,
onFetchChart,
@@ -518,11 +524,6 @@ export default function ChatPanel({
},
})
// Store sendMessage in ref for use in callbacks (like handleImproveWithSuggestions)
useEffect(() => {
sendMessageRef.current = sendMessage
}, [sendMessage])
// Ref to track latest messages for unload persistence
const messagesRef = useRef(messages)
useEffect(() => {
@@ -531,6 +532,9 @@ export default function ChatPanel({
// Track last synced session ID to detect external changes (e.g., URL back/forward)
const lastSyncedSessionIdRef = useRef<string | null>(null)
// Messages array from our latest save. A session holding this exact array was
// created by our own save, so it must not be treated as an external switch.
const lastSavedMessagesRef = useRef<unknown[] | null>(null)
// Helper: Sync UI state with session data (eliminates duplication)
// Track message IDs that are being loaded from session (to skip animations/scroll)
@@ -597,8 +601,10 @@ export default function ChatPanel({
thumbnailDataUrl = latestSvgRef.current
}
}
const messages = sanitizeMessages(messagesRef.current)
lastSavedMessagesRef.current = messages
return {
messages: sanitizeMessages(messagesRef.current),
messages,
xmlSnapshots: Array.from(xmlSnapshotsRef.current.entries()),
diagramXml: currentDiagramXml,
thumbnailDataUrl,
@@ -651,8 +657,13 @@ export default function ChatPanel({
// Skip if session ID hasn't changed (our own saves don't change the ID)
if (newSessionId === lastSyncedSessionIdRef.current) return
// Our own save created this session; the UI already shows its content
const isOwnNewSession =
newSession?.messages === lastSavedMessagesRef.current
// Update last synced ID
lastSyncedSessionIdRef.current = newSessionId
if (isOwnNewSession) return
// Sync UI with new session
if (newSession) {
@@ -793,12 +804,23 @@ export default function ChatPanel({
const onFormSubmit = async (e: React.FormEvent<HTMLFormElement>) => {
e.preventDefault()
const isProcessing = status === "streaming" || status === "submitted"
if (input.trim() && !isProcessing) {
// Check if input matches a cached example (only when no messages yet)
if (messages.length === 0) {
// Attachments still extracting have no text yet. Template sends call
// requestSubmit() and skip the disabled send button, so check here too.
const isExtracting =
files.some((f) => pdfData.get(f)?.isExtracting) ||
Array.from(urlData.values()).some((d) => d.isExtracting)
if (input.trim() && !isProcessing && !isExtracting) {
// Check if input matches a cached example (only when no messages
// yet and the canvas is empty, same rule as the server)
if (
messages.length === 0 &&
isMinimalDiagram(chartXMLRef.current || "")
) {
// Pass the file name so a user's own file never matches an example
const cached = findCachedResponse(
input.trim(),
files.length > 0,
files.length === 1 ? files[0].name : undefined,
)
if (cached) {
// Add user message and fake assistant response to messages
@@ -834,6 +856,11 @@ export default function ChatPanel({
],
},
] as any)
// Snapshot the canvas before the example so editing this message works
xmlSnapshotsRef.current.set(
0,
chartXMLRef.current || wrapWithMxFile(""),
)
setInput("")
sessionStorage.removeItem(SESSION_STORAGE_INPUT_KEY)
setFiles([])
@@ -843,9 +870,6 @@ export default function ChatPanel({
}
try {
let chartXml = await onFetchChart()
chartXml = formatXML(chartXml)
// Build user text by concatenating input with pre-extracted text
// (Backend only reads first text part, so we must combine them)
const parts: any[] = []
@@ -860,20 +884,7 @@ export default function ChatPanel({
// Add the combined text as the first part
parts.unshift({ type: "text", text: userText })
// Get previous XML from the last snapshot (before this message)
const snapshotKeys = Array.from(
xmlSnapshotsRef.current.keys(),
).sort((a, b) => b - a)
const previousXml =
snapshotKeys.length > 0
? xmlSnapshotsRef.current.get(snapshotKeys[0]) || ""
: ""
// Save XML snapshot for this message (will be at index = current messages.length)
const messageIndex = messages.length
xmlSnapshotsRef.current.set(messageIndex, chartXml)
sendChatMessage(parts, chartXml, previousXml, sessionId)
await sendWithCurrentDiagram(parts)
// Token count is tracked in onFinish with actual server usage
setInput("")
@@ -882,10 +893,37 @@ export default function ChatPanel({
setUrlData(new Map())
} catch (error) {
console.error("Error fetching chart data:", error)
toast.error(dict.errors.failedToExport)
}
}
}
// Export the current diagram, snapshot it for this message, and send
const sendWithCurrentDiagram = async (parts: any[]) => {
const chartXml = formatXML(await onFetchChart())
const previousXml = getPreviousXml(messages.length)
// Snapshot the full multi-page document (kept fresh by autosave) so
// regenerate/edit can restore every page; the model gets page 1 only
xmlSnapshotsRef.current.set(
messages.length,
chartXMLRef.current || chartXml,
)
sendChatMessage(parts, chartXml, previousXml, sessionId)
}
// Send VLM validation feedback as a new user message through the normal send path
const handleImproveWithSuggestions = async (feedback: string) => {
if (status === "streaming" || status === "submitted") return
try {
await sendWithCurrentDiagram([{ type: "text", text: feedback }])
} catch (error) {
console.error("Error fetching chart data:", error)
toast.error(dict.errors.failedToExport)
}
}
// Handle session switching from history dropdown
const handleSelectSession = useCallback(
async (sessionId: string) => {
@@ -989,10 +1027,9 @@ export default function ChatPanel({
// Handle sending a template directly (called from TemplatePanel)
const handleSendTemplate = useCallback(
async (template: { prompt: string }) => {
// Keep attachments: they are sent along with the template prompt
flushSync(() => {
setInput(template.prompt)
setFiles([])
setUrlData(new Map())
})
const formElement = document.getElementById(
@@ -1002,7 +1039,7 @@ export default function ChatPanel({
formElement.requestSubmit()
}
},
[setInput, setFiles, setUrlData],
[setInput],
)
const handleInputChange = (
@@ -1017,13 +1054,15 @@ export default function ChatPanel({
}
// Helper functions for message actions (regenerate/edit)
// Extract previous XML snapshot before a given message index
// Extract previous XML snapshot (first page, as sent to the model) before a given message index
const getPreviousXml = (beforeIndex: number): string => {
const snapshotKeys = Array.from(xmlSnapshotsRef.current.keys())
.filter((k) => k < beforeIndex)
.sort((a, b) => b - a)
return snapshotKeys.length > 0
? xmlSnapshotsRef.current.get(snapshotKeys[0]) || ""
? getFirstPageXml(
xmlSnapshotsRef.current.get(snapshotKeys[0]) || "",
)
: ""
}
@@ -1075,6 +1114,7 @@ export default function ChatPanel({
// Reset all retry/continuation state on user-initiated message
autoRetryCountRef.current = 0
continuationRetryCountRef.current = 0
validationRetryCountRef.current = 0
partialXmlRef.current = ""
const config = getSelectedAIConfig()
@@ -1223,7 +1263,12 @@ export default function ChatPanel({
})
// Now send the message after state is guaranteed to be updated
sendChatMessage(userParts, savedXml, previousXml, sessionId)
sendChatMessage(
userParts,
getFirstPageXml(savedXml),
previousXml,
sessionId,
)
}
const handleEditMessage = async (messageIndex: number, newText: string) => {
@@ -1250,10 +1295,13 @@ export default function ChatPanel({
// Clean up snapshots for messages after the user message (they will be removed)
cleanupSnapshotsAfter(messageIndex)
// Create new parts with updated text
// Create new parts with updated text. The edit box only shows the typed
// text, so keep the appended PDF/file/URL content
const newParts = message.parts?.map((part: any) => {
if (part.type === "text") {
return { ...part, text: newText }
const appended =
part.text.match(APPENDED_FILE_SECTIONS_PATTERN)?.[0] ?? ""
return { ...part, text: newText + appended }
}
return part
}) || [{ type: "text", text: newText }]
@@ -1266,7 +1314,12 @@ export default function ChatPanel({
})
// Now send the edited message after state is guaranteed to be updated
sendChatMessage(newParts, savedXml, previousXml, sessionId)
sendChatMessage(
newParts,
getFirstPageXml(savedXml),
previousXml,
sessionId,
)
}
// Collapsed view (desktop only)
+2
View File
@@ -194,6 +194,8 @@ export function ChatLobby({
className="group w-full flex items-center gap-3 p-3 rounded-xl border border-border/60 bg-card hover:bg-accent/50 hover:border-primary/30 transition-all duration-200 cursor-pointer text-left"
onClick={() => onSelectSession(session.id)}
onKeyDown={(e) => {
// Ignore keys bubbling up from the delete button
if (e.target !== e.currentTarget) return
if (
e.key === "Enter" ||
e.key === " "
+3
View File
@@ -55,6 +55,9 @@ export function TemplateCreateDialog({
const handleSubmit = async (e: React.FormEvent) => {
e.preventDefault()
// React submit events bubble through the portal; keep them away from
// the chat form this dialog may be rendered in
e.stopPropagation()
const trimmedPrompt = prompt.trim()
if (!trimmedPrompt) {
+6 -3
View File
@@ -39,16 +39,16 @@ export function TemplateEditDialog({
const [isSubmitting, setIsSubmitting] = useState(false)
const [error, setError] = useState<string | null>(null)
// Populate form when template changes
// Populate form each time the dialog opens, dropping any cancelled edits
useEffect(() => {
if (template) {
if (open && template) {
setTitle(template.title || "")
setDescription(template.description || "")
setPrompt(template.prompt || "")
setPinned(template.pinned || false)
setError(null)
}
}, [template])
}, [open, template])
const handleOpenChange = (newOpen: boolean) => {
if (!newOpen) {
@@ -59,6 +59,9 @@ export function TemplateEditDialog({
const handleSubmit = async (e: React.FormEvent) => {
e.preventDefault()
// React submit events bubble through the portal; keep them away from
// any form this dialog may be rendered in
e.stopPropagation()
if (!template) return
+42 -18
View File
@@ -110,6 +110,10 @@ export function TemplatePanel({
useEffect(() => {
loadTemplates()
// Reload when a template is saved elsewhere, e.g. from the chat input
window.addEventListener("templatesChanged", loadTemplates)
return () =>
window.removeEventListener("templatesChanged", loadTemplates)
}, [loadTemplates])
const handleCreateSuccess = () => {
@@ -302,6 +306,28 @@ export function TemplatePanel({
}
}
// Shared by the empty state and the list, so import works in both
const importInput = (
<input
ref={fileInputRef}
type="file"
accept="application/json,.json"
onChange={handleImport}
className="hidden"
/>
)
const importMessageBox = importMessage && (
<div
className={`text-xs px-3 py-2 rounded-lg ${
importMessage.type === "success"
? "bg-green-100 text-green-800 dark:bg-green-900/30 dark:text-green-400"
: "bg-red-100 text-red-800 dark:bg-red-900/30 dark:text-red-400"
}`}
>
{importMessage.text}
</div>
)
// Empty state: no templates at all
if (!loading && templates.length === 0) {
return (
@@ -332,6 +358,18 @@ export function TemplatePanel({
<Plus className="w-4 h-4" />
{dict.templates.createFirst}
</button>
<button
type="button"
onClick={() => fileInputRef.current?.click()}
className="mt-2 inline-flex items-center gap-1.5 px-3 py-1.5 rounded-md text-xs font-medium text-muted-foreground hover:text-foreground hover:bg-muted transition-colors"
>
<Upload className="w-3.5 h-3.5" />
{dict.templates.importTemplates}
</button>
{importInput}
{importMessageBox && (
<div className="mt-3">{importMessageBox}</div>
)}
<TemplateCreateDialog
open={createDialogOpen}
@@ -389,27 +427,11 @@ export function TemplatePanel({
<Upload className="w-3.5 h-3.5" />
{dict.templates.importTemplates}
</button>
<input
ref={fileInputRef}
type="file"
accept="application/json,.json"
onChange={handleImport}
className="hidden"
/>
{importInput}
</div>
{/* Import message */}
{importMessage && (
<div
className={`text-xs px-3 py-2 rounded-lg ${
importMessage.type === "success"
? "bg-green-100 text-green-800 dark:bg-green-900/30 dark:text-green-400"
: "bg-red-100 text-red-800 dark:bg-red-900/30 dark:text-red-400"
}`}
>
{importMessage.text}
</div>
)}
{importMessageBox}
<div className="space-y-2">
{loading
@@ -447,6 +469,8 @@ export function TemplatePanel({
handleTemplateClick(template)
}
onKeyDown={(e) => {
// Ignore keys bubbling up from the action buttons
if (e.target !== e.currentTarget) return
if (
e.key === "Enter" ||
e.key === " "
+28 -34
View File
@@ -66,7 +66,7 @@ export function ToolCallCard({
dict,
}: ToolCallCardProps) {
const callId = part.toolCallId
const { state, input, output } = part
const { state, input, output, errorText } = part
// Default to expanded for all states (user can manually collapse if needed)
const isExpanded = expandedTools[callId] ?? true
const toolName = part.type?.replace("tool-", "")
@@ -92,6 +92,14 @@ export function ToolCallCard({
}
}
// Incomplete XML means the output hit the length limit, unless the user
// stopped the generation themselves
const isTruncated =
state === "output-error" &&
errorText !== "Stopped by user" &&
(toolName === "display_diagram" || toolName === "append_diagram") &&
!isMxCellXmlComplete(input?.xml)
const handleCopy = () => {
let textToCopy = ""
@@ -161,22 +169,15 @@ export function ToolCallCard({
</>
)}
{state === "output-error" &&
(() => {
// Check if this is a truncation (incomplete XML) vs real error
const isTruncated =
(toolName === "display_diagram" ||
toolName === "append_diagram") &&
!isMxCellXmlComplete(input?.xml)
return isTruncated ? (
<span className="text-xs font-medium text-yellow-600 bg-yellow-50 px-2 py-0.5 rounded-full">
Truncated
</span>
) : (
<span className="text-xs font-medium text-red-600 bg-red-50 px-2 py-0.5 rounded-full">
Error
</span>
)
})()}
(isTruncated ? (
<span className="text-xs font-medium text-yellow-600 bg-yellow-50 px-2 py-0.5 rounded-full">
Truncated
</span>
) : (
<span className="text-xs font-medium text-red-600 bg-red-50 px-2 py-0.5 rounded-full">
Error
</span>
))}
{input && Object.keys(input).length > 0 && (
<button
type="button"
@@ -224,23 +225,16 @@ export function ToolCallCard({
) : null}
</div>
)}
{output &&
state === "output-error" &&
(() => {
const isTruncated =
(toolName === "display_diagram" ||
toolName === "append_diagram") &&
!isMxCellXmlComplete(input?.xml)
return (
<div
className={`px-4 py-3 border-t border-border/40 text-sm ${isTruncated ? "text-yellow-600" : "text-red-600"}`}
>
{isTruncated
? "Output truncated due to length limits. Try a simpler request or increase the maxOutputLength."
: output}
</div>
)
})()}
{/* AI SDK stores tool errors in errorText */}
{state === "output-error" && (errorText || output) && (
<div
className={`px-4 py-3 border-t border-border/40 text-sm whitespace-pre-wrap break-words ${isTruncated ? "text-yellow-600" : "text-red-600"}`}
>
{isTruncated
? "Output truncated due to length limits. Try a simpler request or increase Max Output Tokens in Settings."
: (errorText ?? output)}
</div>
)}
{/* Show get_shape_library output on success */}
{output &&
toolName === "get_shape_library" &&
+1
View File
@@ -13,4 +13,5 @@ export interface ToolPartLike {
operations?: DiagramOperation[]
} & Record<string, unknown>
output?: string
errorText?: string
}
+99 -42
View File
@@ -56,6 +56,7 @@ import { useDictionary } from "@/hooks/use-dictionary"
import type { UseModelConfigReturn } from "@/hooks/use-model-config"
import { getApiEndpoint } from "@/lib/base-path"
import { formatMessage } from "@/lib/i18n/utils"
import { STORAGE_KEYS } from "@/lib/storage"
import type { ProviderConfig, ProviderName } from "@/lib/types/model-config"
import { PROVIDER_INFO, SUGGESTED_MODELS } from "@/lib/types/model-config"
import { cn } from "@/lib/utils"
@@ -133,6 +134,14 @@ export function ModelConfigDialog({
modelId: string
message: string
} | null>(null)
// Model ID being typed; written to the config only when valid on blur
const [modelIdDraft, setModelIdDraft] = useState<{
id: string
value: string
} | null>(null)
// Bumped on every credential edit so a running test can tell that its
// results belong to the old credentials
const credentialsVersionRef = useRef(0)
const [dynamicSuggestedModels, setDynamicSuggestedModels] = useState<
Partial<Record<ProviderName, string[]>>
>({})
@@ -157,6 +166,11 @@ export function ModelConfigDialog({
(p) => p.id === selectedProviderId,
)
// Discard an unfinished model ID edit when the dialog closes
useEffect(() => {
if (!open) setModelIdDraft(null)
}, [open])
// Cleanup validation reset timeout on unmount
useEffect(() => {
return () => {
@@ -253,9 +267,9 @@ export function ModelConfigDialog({
field: keyof ProviderConfig,
value: string | boolean,
) => {
if (!selectedProviderId) return
updateProvider(selectedProviderId, { [field]: value })
// Reset validation when credentials change
if (!selectedProviderId || !selectedProvider) return
const updates: Partial<ProviderConfig> = { [field]: value }
// Reset validation of the provider and its models when credentials change
const credentialFields = [
"apiKey",
"baseUrl",
@@ -265,9 +279,17 @@ export function ModelConfigDialog({
"vertexApiKey",
]
if (credentialFields.includes(field)) {
credentialsVersionRef.current++
setValidationStatus("idle")
updateProvider(selectedProviderId, { validated: false })
setValidatingModelIndex(null)
updates.validated = false
updates.models = selectedProvider.models.map((m) => ({
...m,
validated: undefined,
validationError: undefined,
}))
}
updateProvider(selectedProviderId, updates)
}
// Handle adding a model to current provider
@@ -337,6 +359,7 @@ export function ModelConfigDialog({
let allValid = true
let errorCount = 0
const credentialsVersion = credentialsVersionRef.current
// Validate each model
for (let i = 0; i < selectedProvider.models.length; i++) {
@@ -346,26 +369,37 @@ export function ModelConfigDialog({
try {
// For EdgeOne, construct baseUrl from current origin
const baseUrl = isEdgeOne
? `${window.location.origin}/api/edgeai`
? `${window.location.origin}${getApiEndpoint("/api/edgeai")}`
: selectedProvider.baseUrl
const response = await fetch("/api/validate-model", {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({
provider: selectedProvider.provider,
apiKey: selectedProvider.apiKey,
baseUrl,
modelId: model.modelId,
// AWS Bedrock credentials
awsAccessKeyId: selectedProvider.awsAccessKeyId,
awsSecretAccessKey: selectedProvider.awsSecretAccessKey,
awsRegion: selectedProvider.awsRegion,
// Vertex AI credentials (Express Mode)
vertexApiKey: selectedProvider.vertexApiKey,
}),
})
const data = await response.json()
const response = await fetch(
getApiEndpoint("/api/validate-model"),
{
method: "POST",
headers: {
"Content-Type": "application/json",
"x-access-code":
localStorage.getItem(STORAGE_KEYS.accessCode) ||
"",
},
body: JSON.stringify({
provider: selectedProvider.provider,
apiKey: selectedProvider.apiKey,
baseUrl,
modelId: model.modelId,
// AWS Bedrock credentials
awsAccessKeyId: selectedProvider.awsAccessKeyId,
awsSecretAccessKey:
selectedProvider.awsSecretAccessKey,
awsRegion: selectedProvider.awsRegion,
// Vertex AI credentials (Express Mode)
vertexApiKey: selectedProvider.vertexApiKey,
}),
},
)
const data = await response.json().catch(() => ({}))
// Credentials changed during the test: drop the results
if (credentialsVersionRef.current !== credentialsVersion) return
if (data.valid) {
updateModel(selectedProviderId, model.id, {
@@ -377,10 +411,15 @@ export function ModelConfigDialog({
errorCount++
updateModel(selectedProviderId, model.id, {
validated: false,
validationError: data.error || "Validation failed",
validationError:
data.error ||
(response.ok
? "Validation failed"
: `Request failed (${response.status})`),
})
}
} catch {
if (credentialsVersionRef.current !== credentialsVersion) return
allValid = false
errorCount++
updateModel(selectedProviderId, model.id, {
@@ -615,7 +654,9 @@ export function ModelConfigDialog({
{/* Add Provider */}
<div className="p-3 border-t border-border-subtle">
{/* Always empty so picking the same type again still fires */}
<Select
value=""
onValueChange={(v) =>
handleAddProvider(v as ProviderName)
}
@@ -837,6 +878,7 @@ export function ModelConfigDialog({
<Plus className="h-3.5 w-3.5" />
</Button>
<Select
value=""
onValueChange={(value) => {
if (value) {
handleAddModel(
@@ -989,7 +1031,10 @@ export function ModelConfigDialog({
</div>
<Input
value={
model.modelId
modelIdDraft?.id ===
model.id
? modelIdDraft.value
: model.modelId
}
title={
model.modelId
@@ -1007,24 +1052,14 @@ export function ModelConfigDialog({
null,
)
}
if (
selectedProviderId
) {
updateModel(
selectedProviderId,
model.id,
{
modelId:
e
.target
.value,
validated:
undefined,
validationError:
undefined,
},
)
}
setModelIdDraft(
{
id: model.id,
value: e
.target
.value,
},
)
}}
onKeyDown={(
e,
@@ -1041,6 +1076,10 @@ export function ModelConfigDialog({
) => {
const newModelId =
e.target.value.trim()
// Drop the draft; an invalid ID falls back to the saved one
setModelIdDraft(
null,
)
// Helper to show error with shake
const showError =
@@ -1135,6 +1174,24 @@ export function ModelConfigDialog({
setEditError(
null,
)
if (
selectedProviderId &&
newModelId !==
model.modelId
) {
updateModel(
selectedProviderId,
model.id,
{
modelId:
newModelId,
validated:
undefined,
validationError:
undefined,
},
)
}
}}
className="flex-1 min-w-0 font-mono text-sm h-8 border-0 bg-transparent focus-visible:bg-background focus-visible:ring-1"
/>
+12 -6
View File
@@ -264,9 +264,13 @@ export function ModelSelector({
(model) => (
<ModelSelectorItem
key={model.id}
value={
model.modelId
}
// Unique value so same-named models highlight
// separately; keywords keep search by name
value={model.id}
keywords={[
model.modelId,
providerLabel,
]}
onSelect={() =>
handleSelect(
model.id,
@@ -351,9 +355,11 @@ export function ModelSelector({
(model) => (
<ModelSelectorItem
key={model.id}
value={
model.modelId
}
value={model.id}
keywords={[
model.modelId,
providerLabel,
]}
onSelect={() =>
handleSelect(
model.id,
+72 -73
View File
@@ -1,7 +1,7 @@
"use client"
import type React from "react"
import { createContext, useContext, useEffect, useRef, useState } from "react"
import { createContext, useContext, useRef, useState } from "react"
import type { DrawIoEmbedRef, EventExport } from "react-drawio"
import { toast } from "sonner"
import type { ExportFormat } from "@/components/save-dialog"
@@ -42,6 +42,12 @@ interface DiagramContextType {
const DiagramContext = createContext<DiagramContextType | undefined>(undefined)
// Exports for thumbnails, validation PNGs 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.
type ExportTag = "thumbnail" | "validation"
export function DiagramProvider({ children }: { children: React.ReactNode }) {
const [chartXML, setChartXML] = useState<string>("")
const [latestSvg, setLatestSvg] = useState<string>("")
@@ -53,8 +59,10 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
const hasCalledOnLoadRef = useRef(false)
const drawioRef = useRef<DrawIoEmbedRef | null>(null)
const resolverRef = useRef<((value: string) => void) | null>(null)
// Resolver for PNG export (used for VLM validation)
const pngResolverRef = useRef<((value: string) => void) | null>(null)
// Pending thumbnail and validation PNG exports, keyed by their export tag
const taggedResolversRef = useRef<
Partial<Record<ExportTag, (value: string) => void>>
>({})
// Track if we're expecting an export for history (user-initiated)
const expectHistoryExportRef = useRef<boolean>(false)
// Track latest chartXML for restoration after remount
@@ -76,10 +84,12 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
setIsDrawioReady(false)
}
// Keep chartXMLRef in sync with state for restoration after remount
useEffect(() => {
chartXMLRef.current = chartXML
}, [chartXML])
// Update chartXML and its ref together, so callbacks that read the ref
// (export handler, autosave) see the new value right away
const updateChartXML = (xml: string) => {
chartXMLRef.current = xml
setChartXML(xml)
}
// Track if we're expecting an export for file save (stores raw export data)
const saveResolverRef = useRef<{
@@ -106,64 +116,52 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
}
}
// Export with a tag in `message` (draw.io echoes it back in the export
// event) and wait for that result. Resolves to null on timeout, which is
// expected occasionally.
const requestTaggedExport = (
tag: ExportTag,
format: "xmlsvg" | "png",
timeoutMs: number,
) =>
new Promise<string | null>((resolve) => {
const finish = (value: string | null) => {
clearTimeout(timer)
if (taggedResolversRef.current[tag] === finish) {
delete taggedResolversRef.current[tag]
}
resolve(value)
}
const timer = setTimeout(() => finish(null), timeoutMs)
taggedResolversRef.current[tag] = finish
drawioRef.current?.exportDiagram({ format, message: tag })
})
// Get current diagram as SVG for thumbnail (used by session storage)
const getThumbnailSvg = async (): Promise<string | null> => {
if (!drawioRef.current) return null
// Don't export if diagram is empty
if (!isRealDiagram(chartXML)) return null
try {
const svgData = await Promise.race([
new Promise<string>((resolve) => {
resolverRef.current = resolve
drawioRef.current?.exportDiagram({ format: "xmlsvg" })
}),
new Promise<string>((_, reject) =>
setTimeout(() => reject(new Error("Export timeout")), 3000),
),
])
if (!isRealDiagram(chartXMLRef.current)) return null
// xmlsvg exports return an SVG data URL
const svgData = await requestTaggedExport("thumbnail", "xmlsvg", 3000)
if (svgData?.startsWith("data:image/svg")) {
// Update latestSvg so it's available for future saves
if (svgData?.includes("<svg")) {
setLatestSvg(svgData)
return svgData
}
return null
} catch {
// Timeout is expected occasionally - don't log as error
return null
setLatestSvg(svgData)
return svgData
}
return null
}
// Capture current diagram as PNG for VLM validation
const captureValidationPng = async (): Promise<string | null> => {
if (!drawioRef.current) return null
// Don't export if diagram is empty
if (!isRealDiagram(chartXML)) return null
if (!isRealDiagram(chartXMLRef.current)) return null
try {
const pngData = await Promise.race([
new Promise<string>((resolve) => {
pngResolverRef.current = resolve
drawioRef.current?.exportDiagram({ format: "png" })
}),
new Promise<string>((_, reject) =>
setTimeout(
() => reject(new Error("PNG export timeout")),
5000,
),
),
])
// PNG data should be a base64 data URL
if (pngData?.startsWith("data:image/png")) {
return pngData
}
return null
} catch {
// Timeout is expected occasionally - don't log as error
return null
}
const pngData = await requestTaggedExport("validation", "png", 5000)
// PNG data should be a base64 data URL
return pngData?.startsWith("data:image/png") ? pngData : null
}
const loadDiagram = (
@@ -193,7 +191,7 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
}
// Keep chartXML in sync even when diagrams are injected (e.g., display_diagram tool)
setChartXML(xmlToLoad)
updateChartXML(xmlToLoad)
if (drawioRef.current) {
drawioRef.current.load({
@@ -205,24 +203,17 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
}
const handleDiagramExport = (data: EventExport) => {
// Handle PNG export for VLM validation
if (pngResolverRef.current && data.data?.startsWith("data:image/png")) {
pngResolverRef.current(data.data)
pngResolverRef.current = null
// Tagged exports (thumbnail, validation PNG, file save) go only to
// their own caller, so they never take the result meant for resolverRef
const tag = data.message?.message
if (tag === "thumbnail" || tag === "validation") {
taggedResolversRef.current[tag]?.(data.data)
return
}
// Handle save to file if requested (process raw data before extraction)
if (saveResolverRef.current.resolver) {
const format = saveResolverRef.current.format
saveResolverRef.current.resolver(data.data, data.xml)
if (tag === "save") {
saveResolverRef.current.resolver?.(data.data, data.xml)
saveResolverRef.current = { resolver: null, format: null }
// For non-xmlsvg formats, skip XML extraction as it will fail
// Only drawio (which uses xmlsvg internally) has the content attribute
// xmlsvg is saved directly as SVG file, no need for extraction
if (format === "png" || format === "svg" || format === "xmlsvg") {
return
}
return
}
// Don't write chartXML here: exports don't change the diagram, and
@@ -236,12 +227,15 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
// Limit to 20 entries to prevent memory leaks during long sessions
const MAX_HISTORY_SIZE = 20
if (expectHistoryExportRef.current) {
// Store the full multi-page document (extractedXML is only the
// first page), so restoring a version keeps every page
const historyXml = chartXMLRef.current || extractedXML
setDiagramHistory((prev) => {
const newHistory = [
...prev,
{
svg: data.data,
xml: extractedXML,
xml: historyXml,
},
]
// Keep only the last MAX_HISTORY_SIZE entries (circular buffer)
@@ -256,14 +250,16 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
}
}
// react-drawio registers this callback once per iframe mount, so it must
// read refs: state captured in its closure would stay stale after a remount
const handleDiagramAutoSave = (data: { xml?: string }) => {
if (!data?.xml) return
// Don't overwrite a pending restore - if we have a real diagram in state
// but DrawIO isn't ready yet, it means we're waiting to restore
if (!isDrawioReady && isRealDiagram(chartXML)) {
// Don't overwrite a pending restore - if we have a real diagram but
// DrawIO hasn't loaded yet, it means we're waiting to restore
if (!hasCalledOnLoadRef.current && isRealDiagram(chartXMLRef.current)) {
return
}
setChartXML(data.xml)
updateChartXML(data.xml)
}
const clearDiagram = () => {
@@ -365,7 +361,10 @@ export function DiagramProvider({ children }: { children: React.ReactNode }) {
}
// Export diagram - callback will be handled in handleDiagramExport
drawioRef.current.exportDiagram({ format: drawioFormat })
drawioRef.current.exportDiagram({
format: drawioFormat,
message: "save",
})
}
// Log save event to Langfuse (just flags the trace, doesn't send content)
+43 -23
View File
@@ -67,41 +67,62 @@ const MODEL_ALIASES: Record<string, string> = {
"deepseek-v3-0324": "@tx/deepseek-ai/deepseek-v3-0324",
}
const CORS_HEADERS = {
"Access-Control-Allow-Origin": "*",
"Access-Control-Allow-Methods": "POST, OPTIONS",
"Access-Control-Allow-Headers": "Content-Type, Authorization",
}
/**
* Create standardized response with CORS headers
* Create standardized JSON response
*/
function createResponse(body: any, status = 200, extraHeaders = {}): Response {
return new Response(JSON.stringify(body), {
status,
headers: {
"Content-Type": "application/json",
...CORS_HEADERS,
...extraHeaders,
},
})
}
/**
* Handle OPTIONS request for CORS preflight
*/
function handleOptionsRequest(): Response {
return new Response(null, {
headers: {
...CORS_HEADERS,
"Access-Control-Max-Age": "86400",
},
})
// Only the app's own server (/api/chat, /api/validate-model) calls this
// function, so no CORS headers are sent: other sites' pages can't call it
// from a browser and spend the deployment's Edge AI quota.
// Same rule as lib/access-code.ts, but reading the edge function's env.
// No codes configured (or env unavailable) means no check.
function hasValidAccessCode(request: Request, env: any): boolean {
const accessCodes: string[] =
env?.ACCESS_CODE_LIST?.split(",")
.map((code: string) => code.trim())
.filter(Boolean) || []
if (accessCodes.length === 0) return true
const accessCode = request.headers.get("x-access-code")
return !!accessCode && accessCodes.includes(accessCode)
}
export async function onRequest({ request, env: _env }: any) {
if (request.method === "OPTIONS") {
return handleOptionsRequest()
export async function onRequest({ request, env }: any) {
// Requiring JSON also makes any cross-site browser request need a CORS
// preflight, which fails without CORS headers
if (
request.method !== "POST" ||
!request.headers.get("content-type")?.includes("application/json")
) {
return createResponse(
{
error: {
message: "Expected a POST request with a JSON body",
type: "invalid_request_error",
},
},
400,
)
}
if (!hasValidAccessCode(request, env)) {
return createResponse(
{
error: {
message: "Invalid or missing access code",
type: "invalid_request_error",
},
},
401,
)
}
request.headers.delete("accept-encoding")
@@ -153,7 +174,7 @@ export async function onRequest({ request, env: _env }: any) {
type: "invalid_request_error",
},
},
429,
400,
)
}
@@ -216,7 +237,6 @@ export async function onRequest({ request, env: _env }: any) {
"Cache-Control": "no-cache, no-store, no-transform",
"X-Accel-Buffering": "no",
Connection: "keep-alive",
...CORS_HEADERS,
},
})
} catch (error: any) {
+57 -26
View File
@@ -32,6 +32,55 @@ export function rebuildAppMenu(): void {
buildAppMenu()
}
/**
* Apply a preset and restart the server so it takes effect.
* If the restart fails, go back to the previous preset and restart again,
* so the running server always matches the saved current preset.
* Throws an error describing the outcome on failure.
*/
export async function switchPreset(
id: string,
): Promise<Record<string, string>> {
const previousPresetId = getCurrentPresetId()
const env = applyPresetToEnv(id)
if (!env) {
throw new Error("Preset not found")
}
rebuildAppMenu()
// In development, scripts/electron-dev.mjs restarts the Next.js dev server
if (!app.isPackaged) {
return env
}
try {
await restartNextServer()
return env
} catch (error) {
console.error("Failed to restart server:", error)
const reason = error instanceof Error ? error.message : String(error)
// Revert to previous preset on failure
if (!previousPresetId || !applyPresetToEnv(previousPresetId)) {
setCurrentPreset(null)
}
// Rebuild menu to restore previous checkmark state
rebuildAppMenu()
try {
await restartNextServer()
} catch (retryError) {
console.error("Failed to restart server again:", retryError)
throw new Error(
`The server could not be restarted.\n\nPlease restart the app.\n\nError: ${reason}`,
)
}
throw new Error(
`The server could not be restarted.\n\nThe previous configuration has been restored.\n\nError: ${reason}`,
)
}
}
/**
* Get the menu template with translations
*/
@@ -192,32 +241,14 @@ function buildConfigMenu(
type: "radio",
checked: preset.id === currentPresetId,
click: async () => {
const previousPresetId = getCurrentPresetId()
const env = applyPresetToEnv(preset.id)
if (env) {
try {
await restartNextServer()
rebuildAppMenu() // Rebuild menu to update checkmarks
} catch (error) {
console.error("Failed to restart server:", error)
// Revert to previous preset on failure
if (previousPresetId) {
applyPresetToEnv(previousPresetId)
} else {
setCurrentPreset(null)
}
// Rebuild menu to restore previous checkmark state
rebuildAppMenu()
// Show error dialog to notify user
dialog.showErrorBox(
"Configuration Error",
`Failed to apply preset "${preset.name}". The server could not be restarted.\n\nThe previous configuration has been restored.\n\nError: ${error instanceof Error ? error.message : String(error)}`,
)
}
try {
await switchPreset(preset.id)
} catch (error) {
// Show error dialog to notify user
dialog.showErrorBox(
"Configuration Error",
`Failed to apply preset "${preset.name}". ${error instanceof Error ? error.message : String(error)}`,
)
}
},
}))
+123 -69
View File
@@ -1,5 +1,11 @@
import { randomUUID } from "node:crypto"
import { existsSync, mkdirSync, readFileSync, writeFileSync } from "node:fs"
import {
existsSync,
mkdirSync,
readFileSync,
renameSync,
writeFileSync,
} from "node:fs"
import path from "node:path"
import { app, safeStorage } from "electron"
@@ -30,7 +36,9 @@ let hasWarnedAboutPlaintext = false
* Warns if encryption is not available (API key stored in plaintext)
*/
function encryptValue(value: string): string {
if (!value) {
// Already encrypted (a value that could not be decrypted): keep it as is
// instead of wrapping it in a second layer of encryption
if (!value || value.startsWith(ENCRYPTED_PREFIX)) {
return value
}
@@ -61,6 +69,7 @@ function encryptValue(value: string): string {
/**
* Decrypt a sensitive value using safeStorage
* Returns the original value if it's not encrypted or decryption fails
* (so saving writes the stored ciphertext back unchanged)
*/
function decryptValue(value: string): string {
if (!value || !value.startsWith(ENCRYPTED_PREFIX)) {
@@ -179,6 +188,15 @@ export function loadPresets(): ConfigPresetsFile {
return data
} catch (error) {
console.error("Failed to load config presets:", error)
// Move the unreadable file aside so the next save can't overwrite
// the user's presets with an empty list
const backupPath = `${configPath}.corrupt-${Date.now()}`
try {
renameSync(configPath, backupPath)
console.error(`Unreadable config presets moved to ${backupPath}`)
} catch (renameError) {
console.error("Failed to back up config presets:", renameError)
}
return {
version: 1,
currentPresetId: null,
@@ -211,7 +229,11 @@ export function savePresets(data: ConfigPresetsFile): void {
}
try {
writeFileSync(configPath, JSON.stringify(dataToSave, null, 2), "utf-8")
// Write a temp file and rename it, so a crash mid-write can't leave
// a truncated config file
const tempPath = `${configPath}.tmp`
writeFileSync(tempPath, JSON.stringify(dataToSave, null, 2), "utf-8")
renameSync(tempPath, configPath)
} catch (error) {
console.error("Failed to save config presets:", error)
throw error
@@ -307,9 +329,10 @@ export function deletePreset(id: string): boolean {
data.presets.splice(index, 1)
// Clear current preset if it was deleted
// Clear current preset (and its env vars) if it was deleted
if (data.currentPresetId === id) {
data.currentPresetId = null
setPresetEnv(null)
}
savePresets(data)
@@ -322,13 +345,15 @@ export function deletePreset(id: string): boolean {
export function setCurrentPreset(id: string | null): boolean {
const data = loadPresets()
let preset: ConfigPreset | null = null
if (id !== null) {
const preset = data.presets.find((p) => p.id === id)
preset = data.presets.find((p) => p.id === id) || null
if (!preset) {
return false
}
}
setPresetEnv(preset)
data.currentPresetId = id
savePresets(data)
return true
@@ -365,78 +390,23 @@ const PROVIDER_ENV_MAP: Record<string, { apiKey: string; baseUrl: string }> = {
}
/**
* Apply preset environment variables to the current process
* Returns the environment variables that were applied
*/
export function applyPresetToEnv(id: string): Record<string, string> | null {
const data = loadPresets()
const preset = data.presets.find((p) => p.id === id)
if (!preset) {
return null
}
const appliedEnv: Record<string, string> = {}
const provider = preset.config.AI_PROVIDER?.toLowerCase()
for (const [key, value] of Object.entries(preset.config)) {
if (value !== undefined && value !== "") {
// Map generic AI_API_KEY to provider-specific key
if (
key === "AI_API_KEY" &&
provider &&
PROVIDER_ENV_MAP[provider]
) {
const providerApiKey = PROVIDER_ENV_MAP[provider].apiKey
if (providerApiKey) {
process.env[providerApiKey] = value
appliedEnv[providerApiKey] = value
}
}
// Map generic AI_BASE_URL to provider-specific key
else if (
key === "AI_BASE_URL" &&
provider &&
PROVIDER_ENV_MAP[provider]
) {
const providerBaseUrl = PROVIDER_ENV_MAP[provider].baseUrl
if (providerBaseUrl) {
process.env[providerBaseUrl] = value
appliedEnv[providerBaseUrl] = value
}
}
// Apply other env vars directly
else {
process.env[key] = value
appliedEnv[key] = value
}
}
}
// Set as current preset
data.currentPresetId = id
savePresets(data)
return appliedEnv
}
/**
* Get environment variables from current preset
* Map a preset's config to environment variables
* Maps generic AI_API_KEY/AI_BASE_URL to provider-specific keys
*/
export function getCurrentPresetEnv(): Record<string, string> {
const preset = getCurrentPreset()
if (!preset) {
return {}
}
function presetToEnv(preset: ConfigPreset): Record<string, string> {
const env: Record<string, string> = {}
const provider = preset.config.AI_PROVIDER?.toLowerCase()
for (const [key, value] of Object.entries(preset.config)) {
if (value !== undefined && value !== "") {
// A key that could not be decrypted is useless to the server
if (value.startsWith(ENCRYPTED_PREFIX)) {
console.warn(
`Preset "${preset.name}": ${key} could not be decrypted. Please enter it again in Settings.`,
)
}
// Map generic AI_API_KEY to provider-specific key
if (
else if (
key === "AI_API_KEY" &&
provider &&
PROVIDER_ENV_MAP[provider]
@@ -466,6 +436,90 @@ export function getCurrentPresetEnv(): Record<string, string> {
return env
}
/**
* Values that env vars had before a preset first set them
* (from the system or .env files), and the keys the active preset set
*/
const originalEnv: Record<string, string | undefined> = {}
let presetEnvKeys: string[] = []
/**
* Replace the env vars of the previous preset with those of the given preset
* (null leaves no preset applied). Restoring first means switching presets
* never leaves the previous preset's base URL, model or key behind.
*/
function setPresetEnv(preset: ConfigPreset | null): Record<string, string> {
for (const key of presetEnvKeys) {
if (originalEnv[key] === undefined) {
delete process.env[key]
} else {
process.env[key] = originalEnv[key]
}
}
const env = preset ? presetToEnv(preset) : {}
for (const [key, value] of Object.entries(env)) {
if (!(key in originalEnv)) {
originalEnv[key] = process.env[key]
}
process.env[key] = value
}
presetEnvKeys = Object.keys(env)
writeDevPresetEnv(env)
return env
}
const DEV_ENV_FILE_NAME = "dev-preset-env.json"
/**
* Development only: write the active preset's env vars (decrypted and mapped)
* for scripts/electron-dev.mjs, which restarts the Next.js dev server when
* this file changes. The dev server can't decrypt the config file itself.
*/
function writeDevPresetEnv(env: Record<string, string>): void {
if (app.isPackaged) {
return
}
try {
const filePath = path.join(app.getPath("userData"), DEV_ENV_FILE_NAME)
writeFileSync(filePath, JSON.stringify(env, null, 2), {
encoding: "utf-8",
mode: 0o600,
})
} catch (error) {
console.error("Failed to write dev preset env:", error)
}
}
/**
* Apply preset environment variables to the current process
* Returns the environment variables that were applied
*/
export function applyPresetToEnv(id: string): Record<string, string> | null {
const data = loadPresets()
const preset = data.presets.find((p) => p.id === id)
if (!preset) {
return null
}
const appliedEnv = setPresetEnv(preset)
// Set as current preset
data.currentPresetId = id
savePresets(data)
return appliedEnv
}
/**
* Apply the saved current preset's environment variables (used at startup)
*/
export function applyCurrentPresetToEnv(): void {
setPresetEnv(getCurrentPreset())
}
/**
* Get user's preferred locale from config
* Returns undefined if not set
+10 -6
View File
@@ -48,12 +48,16 @@ function loadEnvFromFile(filePath: string): void {
const key = trimmed.slice(0, equalIndex).trim()
let value = trimmed.slice(equalIndex + 1).trim()
// Remove surrounding quotes
if (
(value.startsWith('"') && value.endsWith('"')) ||
(value.startsWith("'") && value.endsWith("'"))
) {
value = value.slice(1, -1)
const quote = value[0]
const closingQuote =
quote === '"' || quote === "'" ? value.indexOf(quote, 1) : -1
if (closingQuote > 0) {
// Quoted value: keep what's inside the quotes and drop
// anything after them (e.g. a comment)
value = value.slice(1, closingQuote)
} else {
// Unquoted value: drop an inline comment ("value # comment")
value = value.replace(/\s+#.*$/, "")
}
// Don't override existing environment variables
+48 -20
View File
@@ -1,12 +1,17 @@
import { app, BrowserWindow, dialog, shell } from "electron"
import { buildAppMenu } from "./app-menu"
import { getCurrentPresetEnv } from "./config-manager"
import { applyCurrentPresetToEnv } from "./config-manager"
import { loadEnvFile } from "./env-loader"
import { registerIpcHandlers } from "./ipc-handlers"
import { startNextServer, stopNextServer } from "./next-server"
import { applyProxyToEnv } from "./proxy-manager"
import { registerSettingsWindowHandlers } from "./settings-window"
import { createWindow, getMainWindow } from "./window-manager"
import {
createWindow,
getAppUrl,
getMainWindow,
isAppUrl,
} from "./window-manager"
// Single instance lock
const gotTheLock = app.requestSingleInstanceLock()
@@ -28,16 +33,14 @@ if (!gotTheLock) {
// Apply proxy settings from saved config
applyProxyToEnv()
// Apply saved preset environment variables (overrides .env)
const presetEnv = getCurrentPresetEnv()
for (const [key, value] of Object.entries(presetEnv)) {
process.env[key] = value
}
const isDev = process.env.NODE_ENV === "development"
let serverUrl: string | null = null
const isDev = !app.isPackaged
app.whenReady().then(async () => {
// Apply saved preset environment variables (overrides .env).
// Must run after ready: on Windows and Linux safeStorage can't
// decrypt the API key before that.
applyCurrentPresetToEnv()
// Register IPC handlers
registerIpcHandlers()
registerSettingsWindowHandlers()
@@ -46,6 +49,7 @@ if (!gotTheLock) {
buildAppMenu()
try {
let serverUrl: string
if (isDev) {
// Development: use the dev server URL
serverUrl =
@@ -69,8 +73,9 @@ if (!gotTheLock) {
app.on("activate", () => {
if (BrowserWindow.getAllWindows().length === 0) {
if (serverUrl) {
createWindow(serverUrl)
const appUrl = getAppUrl()
if (appUrl) {
createWindow(appUrl)
}
}
})
@@ -87,24 +92,47 @@ if (!gotTheLock) {
stopNextServer()
})
// Pages allowed inside app windows: the app server and draw.io
const isInAppUrl = (url: string): boolean => {
if (isAppUrl(url)) return true
try {
const { hostname } = new URL(url)
return ["diagrams.net", "draw.io"].some(
(domain) =>
hostname === domain || hostname.endsWith(`.${domain}`),
)
} catch {
return false
}
}
const isWebUrl = (url: string): boolean =>
url.startsWith("http://") || url.startsWith("https://")
// Open external links in default browser
app.on("web-contents-created", (_, contents) => {
contents.setWindowOpenHandler(({ url }) => {
// Allow diagrams.net iframe
if (
url.includes("diagrams.net") ||
url.includes("draw.io") ||
url.startsWith("http://localhost") ||
url.startsWith("http://127.0.0.1")
) {
if (isInAppUrl(url)) {
return { action: "allow" }
}
// Open other links in external browser
if (url.startsWith("http://") || url.startsWith("https://")) {
if (isWebUrl(url)) {
shell.openExternal(url)
return { action: "deny" }
}
return { action: "allow" }
})
// Clicking a plain link would otherwise replace the app page with
// an external site that keeps the preload API
contents.on("will-navigate", (event) => {
if (isInAppUrl(event.url)) {
return
}
event.preventDefault()
if (isWebUrl(event.url)) {
shell.openExternal(event.url)
}
})
})
}
+89 -49
View File
@@ -1,7 +1,12 @@
import { app, BrowserWindow, dialog, ipcMain } from "electron"
import { rebuildAppMenu } from "./app-menu"
import {
applyPresetToEnv,
app,
BrowserWindow,
dialog,
type IpcMainInvokeEvent,
ipcMain,
} from "electron"
import { rebuildAppMenu, switchPreset } from "./app-menu"
import {
type ConfigPreset,
createPreset,
deletePreset,
@@ -20,6 +25,7 @@ import {
type ProxyConfig,
saveProxyConfig,
} from "./proxy-manager"
import { isAppUrl } from "./window-manager"
/**
* Allowed configuration keys for presets
@@ -48,13 +54,32 @@ function sanitizePresetConfig(
return sanitized
}
/**
* Register an IPC handler that only answers the app's own pages
* (the main window on the app server, or the local settings page).
* A main window that somehow ends up on an external site still gets the
* preload API, so its calls must be rejected here.
*/
function handle<Args extends unknown[]>(
channel: string,
listener: (event: IpcMainInvokeEvent, ...args: Args) => unknown,
): void {
ipcMain.handle(channel, (event, ...args) => {
const url = event.senderFrame?.url
if (!isAppUrl(url) && !url?.startsWith("file://")) {
throw new Error(`Blocked "${channel}" from untrusted page: ${url}`)
}
return listener(event, ...(args as Args))
})
}
/**
* Register all IPC handlers
*/
export function registerIpcHandlers(): void {
// ==================== App Info ====================
ipcMain.handle("get-version", () => {
handle("get-version", () => {
return app.getVersion()
})
@@ -81,7 +106,7 @@ export function registerIpcHandlers(): void {
// ==================== File Dialogs ====================
ipcMain.handle("dialog-open-file", async (event) => {
handle("dialog-open-file", async (event) => {
const win = BrowserWindow.fromWebContents(event.sender)
if (!win) return null
@@ -108,9 +133,9 @@ export function registerIpcHandlers(): void {
}
})
ipcMain.handle("dialog-save-file", async (event, data: string) => {
handle("dialog-save-file", async (event, data: string) => {
const win = BrowserWindow.fromWebContents(event.sender)
if (!win) return false
if (!win || typeof data !== "string") return false
const result = await dialog.showSaveDialog(win, {
filters: [
@@ -135,28 +160,28 @@ export function registerIpcHandlers(): void {
// ==================== Config Presets ====================
ipcMain.handle("config-presets:get-all", () => {
handle("config-presets:get-all", () => {
return getAllPresets()
})
ipcMain.handle("config-presets:get-current", () => {
handle("config-presets:get-current", () => {
return getCurrentPreset()
})
ipcMain.handle("config-presets:get-current-id", () => {
handle("config-presets:get-current-id", () => {
return getCurrentPresetId()
})
ipcMain.handle(
handle(
"config-presets:save",
(
async (
_event,
preset: Omit<ConfigPreset, "id" | "createdAt" | "updatedAt"> & {
id?: string
},
) => {
// Validate preset name
if (typeof preset.name !== "string" || !preset.name.trim()) {
if (typeof preset?.name !== "string" || !preset.name.trim()) {
throw new Error("Invalid preset name")
}
@@ -165,42 +190,48 @@ export function registerIpcHandlers(): void {
if (preset.id) {
// Update existing preset
return updatePreset(preset.id, {
const updated = updatePreset(preset.id, {
name: preset.name.trim(),
config: sanitizedConfig,
})
// Re-apply the active preset so the edit takes effect
if (updated && updated.id === getCurrentPresetId()) {
await switchPreset(updated.id)
} else {
rebuildAppMenu()
}
return updated
}
// Create new preset
return createPreset({
const created = createPreset({
name: preset.name.trim(),
config: sanitizedConfig,
})
rebuildAppMenu()
return created
},
)
ipcMain.handle("config-presets:delete", (_event, id: string) => {
return deletePreset(id)
handle("config-presets:delete", async (_event, id: string) => {
const wasCurrent = id === getCurrentPresetId()
// Deleting the active preset also clears its env vars
const deleted = deletePreset(id)
rebuildAppMenu()
// Restart so the server stops using the deleted preset
if (deleted && wasCurrent && app.isPackaged) {
await restartNextServer()
}
return deleted
})
ipcMain.handle("config-presets:apply", async (_event, id: string) => {
const env = applyPresetToEnv(id)
if (!env) {
return { success: false, error: "Preset not found" }
}
const isDev = process.env.NODE_ENV === "development"
if (isDev) {
// In development mode, the config file change will trigger
// the file watcher in electron-dev.mjs to restart Next.js
// We just need to save the preset (already done in applyPresetToEnv)
return { success: true, env, devMode: true }
}
// Production mode: restart the Next.js server to apply new environment variables
handle("config-presets:apply", async (_event, id: string) => {
try {
await restartNextServer()
return { success: true, env }
const env = await switchPreset(id)
// In development mode, electron-dev.mjs restarts Next.js
return app.isPackaged
? { success: true, env }
: { success: true, env, devMode: true }
} catch (error) {
return {
success: false,
@@ -212,30 +243,39 @@ export function registerIpcHandlers(): void {
}
})
ipcMain.handle(
"config-presets:set-current",
(_event, id: string | null) => {
return setCurrentPreset(id)
},
)
handle("config-presets:set-current", (_event, id: string | null) => {
return setCurrentPreset(id)
})
// ==================== Proxy Settings ====================
ipcMain.handle("get-proxy", () => {
handle("get-proxy", () => {
return getProxyConfig()
})
ipcMain.handle("set-proxy", async (_event, config: ProxyConfig) => {
handle("set-proxy", async (_event, config: ProxyConfig) => {
const isOptionalString = (value: unknown) =>
value === undefined || typeof value === "string"
if (
typeof config !== "object" ||
config === null ||
!isOptionalString(config.httpProxy) ||
!isOptionalString(config.httpsProxy)
) {
return { success: false, error: "Invalid proxy settings" }
}
try {
// Save config to file
saveProxyConfig(config)
saveProxyConfig({
httpProxy: config.httpProxy,
httpsProxy: config.httpsProxy,
})
// Apply to current process environment
applyProxyToEnv()
const isDev = process.env.NODE_ENV === "development"
if (isDev) {
if (!app.isPackaged) {
// In development, env vars are already applied
// Next.js dev server may need manual restart
return { success: true, devMode: true }
@@ -257,11 +297,11 @@ export function registerIpcHandlers(): void {
// ==================== User Locale ====================
ipcMain.handle("get-user-locale", () => {
handle("get-user-locale", () => {
return getUserLocale()
})
ipcMain.handle("set-user-locale", (_event, locale: string) => {
handle("set-user-locale", (_event, locale: string) => {
// Validate locale is one of the supported values
if (!["en", "zh", "ja", "zh-Hant"].includes(locale)) {
return { success: false, error: "Invalid locale" }
+69 -42
View File
@@ -6,10 +6,22 @@ import {
getAllocatedPort,
getServerUrl,
isPortAvailable,
saveServerPort,
} from "./port-manager"
import { setAppUrl } from "./window-manager"
let serverProcess: UtilityProcess | null = null
// Start and restart run one at a time, so overlapping calls (e.g. two quick
// preset switches) can't leave two servers running
let serverQueue: Promise<unknown> = Promise.resolve()
function runExclusive<T>(task: () => Promise<T>): Promise<T> {
const result = serverQueue.then(task)
serverQueue = result.catch(() => {})
return result
}
/**
* Get the path to the standalone server resources
* In packaged app: resources/standalone
@@ -45,7 +57,11 @@ async function waitForServer(url: string, timeout = 30000): Promise<void> {
* Start the Next.js standalone server using Electron's utilityProcess
* This API is designed for running Node.js code in the background
*/
export async function startNextServer(): Promise<string> {
export function startNextServer(): Promise<string> {
return runExclusive(startServer)
}
async function startServer(): Promise<string> {
const resourcePath = getResourcePath()
const serverPath = path.join(resourcePath, "server.js")
@@ -73,6 +89,11 @@ export async function startNextServer(): Promise<string> {
NODE_USE_ENV_PROXY: "1",
}
// Keep requests to local model servers (e.g. Ollama) off the proxy
if (!process.env.NO_PROXY && !process.env.no_proxy) {
env.NO_PROXY = "localhost,127.0.0.1,[::1]"
}
// Set cache directory to a writable location (user's app data folder)
// This is necessary because the packaged app might be on a read-only volume
if (app.isPackaged) {
@@ -96,28 +117,33 @@ export async function startNextServer(): Promise<string> {
// Use Electron's utilityProcess API for running Node.js in background
// This is the recommended way to run Node.js code in Electron
serverProcess = utilityProcess.fork(serverPath, [], {
const proc = utilityProcess.fork(serverPath, [], {
cwd: resourcePath,
env,
stdio: "pipe",
})
serverProcess = proc
serverProcess.stdout?.on("data", (data) => {
proc.stdout?.on("data", (data) => {
console.log(`[Next.js] ${data.toString().trim()}`)
})
serverProcess.stderr?.on("data", (data) => {
proc.stderr?.on("data", (data) => {
console.error(`[Next.js Error] ${data.toString().trim()}`)
})
serverProcess.on("exit", (code) => {
proc.on("exit", (code) => {
console.log(`Next.js server exited with code ${code}`)
serverProcess = null
// An old server can exit after a new one started; keep the new one
if (serverProcess === proc) {
serverProcess = null
}
})
const url = getServerUrl()
await waitForServer(url)
console.log(`Next.js server started at ${url}`)
saveServerPort(port)
return url
}
@@ -126,39 +152,36 @@ export async function startNextServer(): Promise<string> {
* Stop the Next.js server process and wait for it to exit
*/
export async function stopNextServer(): Promise<void> {
if (serverProcess) {
console.log("Stopping Next.js server...")
const proc = serverProcess
if (!proc) {
return
}
console.log("Stopping Next.js server...")
serverProcess = null
// Create a promise that resolves when the process exits
const exitPromise = new Promise<void>((resolve) => {
const proc = serverProcess
if (!proc) {
resolve()
return
}
const onExit = () => {
resolve()
}
proc.once("exit", onExit)
// Timeout after 5 seconds
setTimeout(() => {
proc.removeListener("exit", onExit)
resolve()
}, 5000)
// Resolves true when the process exits, false after the timeout
const waitForExit = (ms: number) =>
new Promise<boolean>((resolve) => {
proc.once("exit", () => resolve(true))
setTimeout(() => resolve(false), ms)
})
serverProcess.kill()
serverProcess = null
proc.kill()
// Wait for process to exit
await exitPromise
// Additional wait for OS to release port
await new Promise((resolve) => setTimeout(resolve, 500))
// Next.js waits for open requests (e.g. a streaming reply) before it
// exits, so force kill it if it is still running after 5 seconds
if (!(await waitForExit(5000)) && proc.pid) {
console.warn("Next.js server did not exit in time, force killing it")
try {
process.kill(proc.pid, "SIGKILL")
} catch (error) {
console.error("Failed to force kill Next.js server:", error)
}
await waitForExit(2000)
}
// Additional wait for OS to release port
await new Promise((resolve) => setTimeout(resolve, 500))
}
/**
@@ -184,15 +207,19 @@ async function waitForServerStop(timeout = 5000): Promise<void> {
/**
* Restart the Next.js server with new environment variables
*/
export async function restartNextServer(): Promise<string> {
console.log("Restarting Next.js server...")
export function restartNextServer(): Promise<string> {
return runExclusive(async () => {
console.log("Restarting Next.js server...")
// Stop the current server and wait for it to exit
await stopNextServer()
// Stop the current server and wait for it to exit
await stopNextServer()
// Wait for the port to be released
await waitForServerStop()
// Wait for the port to be released
await waitForServerStop()
// Start the server again
return startNextServer()
// Start the server again, and follow it if it moved to another port
const url = await startServer()
setAppUrl(url)
return url
})
}
+50 -1
View File
@@ -1,4 +1,6 @@
import { readFileSync, writeFileSync } from "node:fs"
import net from "node:net"
import path from "node:path"
import { app } from "electron"
/**
@@ -23,6 +25,38 @@ const PORT_CONFIG = {
*/
let allocatedPort: number | null = null
/**
* File that remembers the production port from the last launch, so the app
* keeps the same origin (and its localStorage) instead of switching between
* the legacy and new port depending on which one is free at startup
*/
function getSavedPortPath(): string {
return path.join(app.getPath("userData"), "server-port.json")
}
function loadSavedPort(): number | null {
try {
const { port } = JSON.parse(readFileSync(getSavedPortPath(), "utf-8"))
return Number.isInteger(port) ? port : null
} catch {
return null
}
}
/**
* Remember the port the production server started on
*/
export function saveServerPort(port: number): void {
if (!app.isPackaged || port === loadSavedPort()) {
return
}
try {
writeFileSync(getSavedPortPath(), JSON.stringify({ port }), "utf-8")
} catch (error) {
console.error("Failed to save server port:", error)
}
}
/**
* Check if a specific port is available
*/
@@ -44,7 +78,8 @@ export function isPortAvailable(port: number): Promise<boolean> {
/**
* Find an available port
* - In development: uses fixed port (6002)
* - In production: uses fixed port (13370) to preserve localStorage
* - In production: uses the port from the last launch, then the legacy
* port (61337), then 13370, to preserve localStorage
* - Falls back to sequential ports if preferred port is unavailable
* - Last resort: lets the OS assign a port (port 0)
*
@@ -69,6 +104,20 @@ export async function findAvailablePort(reuseExisting = true): Promise<number> {
allocatedPort = null
}
// In production, use the port from the last launch first
if (!isDev) {
const savedPort = loadSavedPort()
if (savedPort !== null) {
if (await isPortAvailable(savedPort)) {
allocatedPort = savedPort
return savedPort
}
console.warn(
`Port ${savedPort} from the last launch is unavailable. Data saved under it will not show on the new port.`,
)
}
}
// In production, try legacy port first to preserve existing users' localStorage
if (!isDev) {
const legacyPort = PORT_CONFIG.legacyProduction
+18 -5
View File
@@ -13,18 +13,22 @@ function getConfigPath(): string {
/**
* Load proxy configuration from JSON file
* Returns null if the user never saved proxy settings (or the file is invalid)
*/
export function loadProxyConfig(): ProxyConfig {
export function loadProxyConfig(): ProxyConfig | null {
try {
const configPath = getConfigPath()
if (fs.existsSync(configPath)) {
const data = fs.readFileSync(configPath, "utf-8")
return JSON.parse(data) as ProxyConfig
const data = JSON.parse(fs.readFileSync(configPath, "utf-8"))
if (data && typeof data === "object" && !Array.isArray(data)) {
return data as ProxyConfig
}
console.error("Ignoring invalid proxy config:", data)
}
} catch (error) {
console.error("Failed to load proxy config:", error)
}
return {}
return null
}
/**
@@ -33,7 +37,11 @@ export function loadProxyConfig(): ProxyConfig {
export function saveProxyConfig(config: ProxyConfig): void {
try {
const configPath = getConfigPath()
fs.writeFileSync(configPath, JSON.stringify(config, null, 2), "utf-8")
// Write a temp file and rename it, so a crash mid-write can't leave
// a truncated file
const tempPath = `${configPath}.tmp`
fs.writeFileSync(tempPath, JSON.stringify(config, null, 2), "utf-8")
fs.renameSync(tempPath, configPath)
} catch (error) {
console.error("Failed to save proxy config:", error)
throw error
@@ -47,6 +55,11 @@ export function saveProxyConfig(config: ProxyConfig): void {
export function applyProxyToEnv(): void {
const config = loadProxyConfig()
// No saved settings: keep proxy vars inherited from the system or .env
if (!config) {
return
}
if (config.httpProxy) {
process.env.HTTP_PROXY = config.httpProxy
process.env.http_proxy = config.httpProxy
+38 -1
View File
@@ -3,6 +3,9 @@ import { app, BrowserWindow, screen } from "electron"
let mainWindow: BrowserWindow | null = null
// URL of the app server the main window loads
let appUrl: string | null = null
/**
* Get the icon path based on platform
* Note: electron-builder converts icon.png during packaging,
@@ -28,6 +31,7 @@ function getIconPath(): string | undefined {
* Create the main application window
*/
export function createWindow(serverUrl: string): BrowserWindow {
appUrl = serverUrl
const { width, height } = screen.getPrimaryDisplay().workAreaSize
mainWindow = new BrowserWindow({
@@ -56,7 +60,7 @@ export function createWindow(serverUrl: string): BrowserWindow {
})
// Open DevTools in development
if (process.env.NODE_ENV === "development") {
if (!app.isPackaged) {
mainWindow.webContents.openDevTools()
}
@@ -93,3 +97,36 @@ export function createWindow(serverUrl: string): BrowserWindow {
export function getMainWindow(): BrowserWindow | null {
return mainWindow
}
/**
* Get the app server URL the main window loads
*/
export function getAppUrl(): string | null {
return appUrl
}
/**
* Point the main window at a new app server URL
* (the restarted server can come up on a different port)
*/
export function setAppUrl(url: string): void {
if (url === appUrl) {
return
}
appUrl = url
mainWindow?.loadURL(url)
}
/**
* Check if a URL belongs to the app server (same origin)
*/
export function isAppUrl(url: string | undefined): boolean {
if (!url || !appUrl) {
return false
}
try {
return new URL(url).origin === new URL(appUrl).origin
} catch {
return false
}
}
+7 -6
View File
@@ -213,6 +213,9 @@ async function savePreset() {
}
})
// closeModal() clears editingPresetId, so remember it for the toast
const isEdit = Boolean(editingPresetId)
try {
saveBtn.disabled = true
saveBtn.innerHTML = '<span class="loading"></span>'
@@ -220,10 +223,7 @@ async function savePreset() {
await window.settingsAPI.savePreset(preset)
await loadPresets()
closeModal()
showToast(
editingPresetId ? "Preset updated" : "Preset created",
"success",
)
showToast(isEdit ? "Preset updated" : "Preset created", "success")
} catch (error) {
console.error("Failed to save preset:", error)
showToast("Failed to save preset", "error")
@@ -265,8 +265,6 @@ async function applyPreset(id) {
const result = await window.settingsAPI.applyPreset(id)
if (result.success) {
currentPresetId = id
renderPresets()
showToast("Preset applied, server restarting...", "success")
} else {
showToast(result.error || "Failed to apply preset", "error")
@@ -274,6 +272,9 @@ async function applyPreset(id) {
} catch (error) {
console.error("Failed to apply preset:", error)
showToast("Failed to apply preset", "error")
} finally {
// Reload to show the active preset and reset the Apply button
await loadPresets()
}
}
+2 -1
View File
@@ -12,7 +12,8 @@ AI_PROVIDER=bedrock
AI_MODEL=global.anthropic.claude-sonnet-4-5-20250929-v1:0
# Output limit, all providers (default: 64000). Shared by reasoning and the diagram XML,
# so a thinking model can spend it all before the tool call. Users can override it in Settings.
# so a thinking model can spend it all before the tool call. Users can lower it in Settings,
# and raise it only when they use their own API key, so this also caps cost on server keys.
# If a model's own ceiling is lower, the request is retried with that ceiling automatically.
# MAX_OUTPUT_TOKENS=64000
+41 -30
View File
@@ -1,5 +1,4 @@
import type { MutableRefObject } from "react"
import { useRef } from "react"
import type { DiagramOperation } from "@/components/chat/types"
import type {
ValidationState,
@@ -48,6 +47,8 @@ type ValidateDiagramFn = (
interface UseDiagramToolHandlersParams {
partialXmlRef: MutableRefObject<string>
editDiagramOriginalXmlRef: MutableRefObject<Map<string, string>>
// Failed VLM validations in the current user turn (reset on each user message)
validationRetryCountRef: MutableRefObject<number>
chartXMLRef: MutableRefObject<string>
onDisplayChart: (xml: string, skipValidation?: boolean) => string | null
onFetchChart: (saveToHistory?: boolean) => Promise<string>
@@ -72,6 +73,7 @@ interface UseDiagramToolHandlersParams {
export function useDiagramToolHandlers({
partialXmlRef,
editDiagramOriginalXmlRef,
validationRetryCountRef,
chartXMLRef,
onDisplayChart,
onFetchChart,
@@ -82,9 +84,6 @@ export function useDiagramToolHandlers({
sessionId,
onValidationStateChange,
}: UseDiagramToolHandlersParams) {
// Track validation retry count per tool call
const validationRetryCountRef = useRef<Map<string, number>>(new Map())
// Helper to update validation state
const updateValidationState = (
toolCallId: string,
@@ -232,17 +231,15 @@ ${finalXml}
)
}
const retryCount =
validationRetryCountRef.current.get(
toolCall.toolCallId,
) || 0
// Each retry is a new tool call, so count attempts per user turn
const attempt = validationRetryCountRef.current + 1
// Notify UI that we're validating (include the image)
updateValidationState(
toolCall.toolCallId,
"validating",
{
attempt: retryCount + 1,
attempt,
maxAttempts: MAX_VALIDATION_RETRIES,
imageData: capturedPngData,
},
@@ -254,17 +251,14 @@ ${finalXml}
)
if (!result.valid) {
if (retryCount < MAX_VALIDATION_RETRIES) {
validationRetryCountRef.current.set(
toolCall.toolCallId,
retryCount + 1,
)
if (attempt < MAX_VALIDATION_RETRIES) {
validationRetryCountRef.current = attempt
const feedback =
formatValidationFeedback(result)
if (DEBUG) {
console.log(
`[display_diagram] Validation failed (attempt ${retryCount + 1}/${MAX_VALIDATION_RETRIES}):`,
`[display_diagram] Validation failed (attempt ${attempt}/${MAX_VALIDATION_RETRIES}):`,
result.issues,
)
}
@@ -274,7 +268,7 @@ ${finalXml}
toolCall.toolCallId,
"failed",
{
attempt: retryCount + 1,
attempt,
maxAttempts: MAX_VALIDATION_RETRIES,
result,
imageData: capturedPngData,
@@ -285,19 +279,17 @@ ${finalXml}
tool: "display_diagram",
toolCallId: toolCall.toolCallId,
state: "output-error",
errorText: `[Validation attempt ${retryCount + 1}/${MAX_VALIDATION_RETRIES}]\n${feedback}`,
errorText: `[Validation attempt ${attempt}/${MAX_VALIDATION_RETRIES}]\n${feedback}`,
})
return
} else {
// Max retries reached - accept the diagram with warning
// Last attempt - accept the diagram with warning
if (DEBUG) {
console.log(
"[display_diagram] Max validation retries reached, accepting diagram",
)
}
validationRetryCountRef.current.delete(
toolCall.toolCallId,
)
validationRetryCountRef.current = 0
// Notify UI that we're accepting with issues (include the image)
updateValidationState(
@@ -314,10 +306,8 @@ ${finalXml}
return
}
} else {
// Validation passed - clean up retry count
validationRetryCountRef.current.delete(
toolCall.toolCallId,
)
// Validation passed - reset retry count
validationRetryCountRef.current = 0
if (DEBUG) {
console.log(
"[display_diagram] Validation passed!",
@@ -382,12 +372,17 @@ ${finalXml}
}
let currentXml = ""
// Use the original XML captured during streaming (shared with chat-message-display)
// This ensures we apply operations to the same base XML that streaming used
const originalXml = editDiagramOriginalXmlRef.current.get(
toolCall.toolCallId,
)
// On failure, undo the streaming preview so the canvas matches the XML
// reported back to the model
const restoreOriginal = () => {
if (originalXml) onDisplayChart(originalXml, true)
}
try {
// Use the original XML captured during streaming (shared with chat-message-display)
// This ensures we apply operations to the same base XML that streaming used
const originalXml = editDiagramOriginalXmlRef.current.get(
toolCall.toolCallId,
)
if (originalXml) {
currentXml = originalXml
} else {
@@ -416,6 +411,7 @@ ${finalXml}
)
.join("\n")
restoreOriginal()
addToolOutput({
tool: "edit_diagram",
toolCallId: toolCall.toolCallId,
@@ -441,6 +437,7 @@ Please check the cell IDs and retry.`,
"[edit_diagram] Validation error:",
validationError,
)
restoreOriginal()
addToolOutput({
tool: "edit_diagram",
toolCallId: toolCall.toolCallId,
@@ -472,6 +469,7 @@ Please fix the operations to avoid structural issues.`,
const errorMessage =
error instanceof Error ? error.message : String(error)
restoreOriginal()
addToolOutput({
tool: "edit_diagram",
toolCallId: toolCall.toolCallId,
@@ -496,6 +494,19 @@ Please check cell IDs and retry, or use display_diagram to regenerate.`,
) => {
const { xml } = toolCall.input as { xml: string }
// Nothing to continue: loading the fragment alone would replace the whole diagram
if (!partialXmlRef.current) {
addToolOutput({
tool: "append_diagram",
toolCallId: toolCall.toolCallId,
state: "output-error",
errorText: `ERROR: There is no truncated diagram to continue, so append_diagram cannot be used now.
Use display_diagram to create the complete diagram, or edit_diagram to change the current one.`,
})
return
}
// Detect if LLM incorrectly started fresh instead of continuing
// LLM should only output bare mxCells now, so wrapper tags indicate error
const trimmed = xml.trim()
+57 -29
View File
@@ -101,6 +101,15 @@ function saveConfig(config: MultiModelConfig): void {
localStorage.setItem(STORAGE_KEYS.modelConfigs, JSON.stringify(config))
}
/**
* Server model to fall back to: the one marked default, else the first one
*/
function defaultServerModelId(
serverModels: FlattenedServerModel[],
): string | undefined {
return (serverModels.find((m) => m.isDefault) ?? serverModels[0])?.id
}
export interface UseModelConfigReturn {
// State
config: MultiModelConfig
@@ -144,6 +153,16 @@ export function useModelConfig(): UseModelConfigReturn {
setIsLoaded(true)
}, [])
// Pick up config changes saved by other tabs, so this tab neither shows a
// stale model nor overwrites their changes on its next save
useEffect(() => {
const handleStorage = (e: StorageEvent) => {
if (e.key === STORAGE_KEYS.modelConfigs) setConfig(loadConfig())
}
window.addEventListener("storage", handleStorage)
return () => window.removeEventListener("storage", handleStorage)
}, [])
// Load server models on mount (if any)
useEffect(() => {
if (typeof window === "undefined") return
@@ -165,17 +184,18 @@ export function useModelConfig(): UseModelConfigReturn {
setServerModels(raw)
setServerLoaded(true)
// Auto-select default server model if no model is currently selected
// Auto-select the default server model if no model is selected,
// or if the saved server model is gone (renamed or removed)
setConfig((prev) => {
if (!prev.selectedModelId && raw.length > 0) {
const defaultModel = raw.find((m) => m.isDefault)
if (defaultModel) {
return { ...prev, selectedModelId: defaultModel.id }
}
// If no default marked, use first server model
return { ...prev, selectedModelId: raw[0].id }
}
return prev
const id = prev.selectedModelId
const isStale =
id?.startsWith("server:") &&
!raw.some((m) => m.id === id)
if (id && !isStale) return prev
const fallback = defaultServerModelId(raw)
return fallback === id
? prev
: { ...prev, selectedModelId: fallback }
})
})
.catch((error) => {
@@ -260,24 +280,31 @@ export function useModelConfig(): UseModelConfigReturn {
[],
)
const deleteProvider = useCallback((providerId: string) => {
setConfig((prev) => {
const provider = prev.providers.find((p) => p.id === providerId)
const modelIds = provider?.models.map((m) => m.id) || []
const deleteProvider = useCallback(
(providerId: string) => {
setConfig((prev) => {
const provider = prev.providers.find((p) => p.id === providerId)
const modelIds = provider?.models.map((m) => m.id) || []
// Clear selected model if it belongs to deleted provider
const newSelectedId =
prev.selectedModelId && modelIds.includes(prev.selectedModelId)
? undefined
: prev.selectedModelId
// Fall back to the default server model if the selected model
// belongs to the deleted provider
const newSelectedId =
prev.selectedModelId &&
modelIds.includes(prev.selectedModelId)
? defaultServerModelId(serverModels)
: prev.selectedModelId
return {
...prev,
providers: prev.providers.filter((p) => p.id !== providerId),
selectedModelId: newSelectedId,
}
})
}, [])
return {
...prev,
providers: prev.providers.filter(
(p) => p.id !== providerId,
),
selectedModelId: newSelectedId,
}
})
},
[serverModels],
)
const addModel = useCallback(
(providerId: string, modelId: string): ModelConfig => {
@@ -334,14 +361,15 @@ export function useModelConfig(): UseModelConfigReturn {
}
: p,
),
// Clear selected model if it was deleted
// Fall back to the default server model if the selected model
// was deleted
selectedModelId:
prev.selectedModelId === modelConfigId
? undefined
? defaultServerModelId(serverModels)
: prev.selectedModelId,
}))
},
[],
[serverModels],
)
const resetConfig = useCallback(() => {
+32 -4
View File
@@ -1,6 +1,8 @@
"use client"
import { useCallback, useEffect, useRef, useState } from "react"
import { toast } from "sonner"
import { useDictionary } from "@/hooks/use-dictionary"
import {
type ChatSession,
createEmptySession,
@@ -44,6 +46,15 @@ export interface UseSessionManagerReturn {
clearCurrentSession: () => void
}
// Reading the session list loads every stored session in full, and window
// focus also fires each time the user clicks back from the draw.io iframe
const FOCUS_REFRESH_INTERVAL_MS = 30_000
function notifySaveFailed(message: string) {
// Same id, so repeated failures update one toast instead of stacking
toast.error(message, { id: "session-save-failed", duration: 8000 })
}
interface UseSessionManagerOptions {
/** Session ID from URL param - if provided, load this session; if null, start blank */
initialSessionId?: string | null
@@ -53,6 +64,7 @@ export function useSessionManager(
options: UseSessionManagerOptions = {},
): UseSessionManagerReturn {
const { initialSessionId } = options
const dict = useDictionary()
const [sessions, setSessions] = useState<SessionMetadata[]>([])
const [currentSessionId, setCurrentSessionId] = useState<string | null>(
null,
@@ -163,9 +175,15 @@ export function useSessionManager(
handleSessionIdChange()
}, [initialSessionId, isAvailable])
// Refresh sessions on window focus (multi-tab sync)
// Refresh sessions on window focus (multi-tab sync), at most once per interval
const lastFocusRefreshRef = useRef(0)
useEffect(() => {
const handleFocus = () => {
const now = Date.now()
if (now - lastFocusRefreshRef.current < FOCUS_REFRESH_INTERVAL_MS) {
return
}
lastFocusRefreshRef.current = now
refreshSessions()
}
window.addEventListener("focus", handleFocus)
@@ -238,6 +256,8 @@ export function useSessionManager(
) {
return
}
// Nothing can be stored without IndexedDB
if (!isIndexedDBAvailable()) return
if (!currentSession) {
// Create a new session if none exists
@@ -250,7 +270,12 @@ export function useSessionManager(
diagramHistory: data.diagramHistory,
title: extractTitle(data.messages),
}
await saveSession(newSession)
// 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
}
await enforceSessionLimit()
setCurrentSession(newSession)
setCurrentSessionId(newSession.id)
@@ -277,7 +302,10 @@ export function useSessionManager(
: currentSession.title,
}
await saveSession(updatedSession)
if (!(await saveSession(updatedSession))) {
notifySaveFailed(dict.errors.sessionSaveFailed)
return
}
setCurrentSession(updatedSession)
// Update sessions list metadata
@@ -298,7 +326,7 @@ export function useSessionManager(
),
)
},
[currentSession, currentSessionId, refreshSessions],
[currentSession, currentSessionId, refreshSessions, dict],
)
// Clear current session state (for starting fresh without loading another session)
+3
View File
@@ -6,6 +6,7 @@
import { experimental_useObject as useObject } from "@ai-sdk/react"
import { useCallback, useRef } from "react"
import { getSelectedAIConfig } from "@/hooks/use-model-config"
import { getApiEndpoint } from "@/lib/base-path"
import {
type ValidationResult,
@@ -39,6 +40,8 @@ export function useValidateDiagram(options: UseValidateDiagramOptions = {}) {
const { object, submit, isLoading, error, stop } = useObject({
api: getApiEndpoint("/api/validate-diagram"),
schema: ValidationResultSchema,
// Resolved per request so a changed access code is picked up
headers: () => ({ "x-access-code": getSelectedAIConfig().accessCode }),
onFinish: ({
object,
error: finishError,
+22
View File
@@ -0,0 +1,22 @@
/**
* Check the x-access-code header against ACCESS_CODE_LIST.
* Returns a 401 response to send back when the check fails, or null when the
* request may continue (including when no access codes are configured).
*/
export function checkAccessCode(req: Request): Response | null {
const accessCodes =
process.env.ACCESS_CODE_LIST?.split(",")
.map((code) => code.trim())
.filter(Boolean) || []
if (accessCodes.length === 0) return null
const accessCodeHeader = req.headers.get("x-access-code")
if (accessCodeHeader && accessCodes.includes(accessCodeHeader)) return null
return Response.json(
{
error: "Invalid or missing access code. Please configure it in Settings.",
},
{ status: 401 },
)
}
+19 -7
View File
@@ -2,6 +2,7 @@ import { z } from "zod"
import {
ProviderNameSchema,
type ServerModelsConfig,
slugify,
} from "@/lib/server-model-config"
import {
FIXED_CRED_PROVIDERS,
@@ -182,12 +183,15 @@ export function validateAdminProviders(
return `${PROVIDER_INFO[single].label} is already configured in AI_MODELS_CONFIG / ai-models.json and shares global credentials. Manage it via the environment configuration instead.`
}
}
// Server model ids are built from the slugified name, so names must
// stay distinct after slugifying ("OpenAI" and "openai" would collide)
const names = list.map((p) => displayName(p))
if (new Set(names).size !== names.length) {
return "Provider display names must be unique."
const slugs = names.map(slugify)
if (new Set(slugs).size !== slugs.length) {
return "Provider display names must be unique (ignoring case and punctuation)."
}
const envNames = new Set(envProviders.map((p) => p.name))
const clash = names.find((n) => envNames.has(n))
const envSlugs = new Set(envProviders.map((p) => slugify(p.name)))
const clash = names.find((_, i) => envSlugs.has(slugs[i]))
if (clash) {
return `"${clash}" is already defined in AI_MODELS_CONFIG / ai-models.json. Use a different display name.`
}
@@ -240,10 +244,14 @@ export function deriveEnvUpdates(
indexByProvider.set(p.provider, index + 1)
if (p.provider === "bedrock") {
if (p.awsAccessKeyId) updates.AWS_ACCESS_KEY_ID = p.awsAccessKeyId
// ADMIN_ names keep the standard AWS_* vars untouched, so other
// AWS clients (e.g. the DynamoDB quota table) keep their own
// credentials instead of picking up the panel's Bedrock keys
if (p.awsAccessKeyId)
updates.ADMIN_AWS_ACCESS_KEY_ID = p.awsAccessKeyId
if (p.awsSecretAccessKey)
updates.AWS_SECRET_ACCESS_KEY = p.awsSecretAccessKey
if (p.awsRegion) updates.AWS_REGION = p.awsRegion
updates.ADMIN_AWS_SECRET_ACCESS_KEY = p.awsSecretAccessKey
if (p.awsRegion) updates.ADMIN_AWS_REGION = p.awsRegion
} else if (p.provider === "vertexai") {
if (p.vertexApiKey) updates.GOOGLE_VERTEX_API_KEY = p.vertexApiKey
if (p.baseUrl) updates.GOOGLE_VERTEX_BASE_URL = p.baseUrl
@@ -284,6 +292,10 @@ function derivedEnvKeys(list: StoredAdminProvider[]): string[] {
const index = indexByProvider.get(p.provider) ?? 0
indexByProvider.set(p.provider, index + 1)
if (p.provider === "bedrock") {
keys.add("ADMIN_AWS_ACCESS_KEY_ID")
keys.add("ADMIN_AWS_SECRET_ACCESS_KEY")
keys.add("ADMIN_AWS_REGION")
// Written by older versions; listed so the next save clears them
keys.add("AWS_ACCESS_KEY_ID")
keys.add("AWS_SECRET_ACCESS_KEY")
keys.add("AWS_REGION")
+34 -19
View File
@@ -10,13 +10,27 @@ interface SettingsFile {
values: Record<string, string>
}
// Original env values snapshotted before the first overlay, so removing a
// key from the settings file restores the env default. null = was unset.
const originalEnv: Record<string, string | null> = {}
// Keys currently overlaid, so we can restore ones removed from the file.
let overlaidKeys = new Set<string>()
interface SettingsState {
// Original env values snapshotted before the first overlay, so removing
// a key from the settings file restores the env default. null = was unset.
originalEnv: Record<string, string | null>
// Keys currently overlaid, so we can restore ones removed from the file.
overlaidKeys: Set<string>
cachedSettings: Record<string, string> | null
}
let cachedSettings: Record<string, string> | null = null
// Kept on globalThis because the build can load this module more than once
// (instrumentation.ts and the API routes get separate copies); per-module
// state would make a route forget what instrumentation overlaid at startup.
const globalState = globalThis as typeof globalThis & {
__adminSettingsState?: SettingsState
}
globalState.__adminSettingsState ??= {
originalEnv: {},
overlaidKeys: new Set(),
cachedSettings: null,
}
const state = globalState.__adminSettingsState
export function getSettingsPath(): string {
const custom = process.env.SETTINGS_FILE
@@ -25,7 +39,7 @@ export function getSettingsPath(): string {
}
export function loadSettings(): Record<string, string> {
if (cachedSettings) return cachedSettings
if (state.cachedSettings) return state.cachedSettings
try {
const raw = fs.readFileSync(getSettingsPath(), "utf8")
const parsed = JSON.parse(raw) as SettingsFile
@@ -43,21 +57,22 @@ export function loadSettings(): Record<string, string> {
for (const [key, value] of Object.entries(rawValues)) {
if (typeof value === "string") values[key] = value
}
cachedSettings = values
state.cachedSettings = values
} catch (err: any) {
if (err?.code !== "ENOENT") {
console.error("[admin-settings] Failed to read settings file:", err)
}
cachedSettings = {}
state.cachedSettings = {}
}
return cachedSettings
return state.cachedSettings
}
export function applyToEnv(): void {
const values = loadSettings()
const { originalEnv } = state
// Restore env for keys that were overlaid before but are now gone
for (const key of overlaidKeys) {
for (const key of state.overlaidKeys) {
if (!(key in values)) {
const original = originalEnv[key]
if (original === null) delete process.env[key]
@@ -72,12 +87,12 @@ export function applyToEnv(): void {
process.env[key] = value
}
overlaidKeys = new Set(Object.keys(values))
state.overlaidKeys = new Set(Object.keys(values))
}
// The effective env value if the file entry were removed (for fallback display)
export function getEnvFallback(key: string): string | null {
if (overlaidKeys.has(key)) return originalEnv[key] ?? null
if (state.overlaidKeys.has(key)) return state.originalEnv[key] ?? null
return process.env[key] ?? null
}
@@ -101,7 +116,7 @@ export function saveSettings(updates: Record<string, string | null>): void {
fs.writeFileSync(tmpPath, JSON.stringify(data, null, 2), { mode: 0o600 })
fs.renameSync(tmpPath, filePath)
cachedSettings = current
state.cachedSettings = current
applyToEnv()
}
@@ -122,13 +137,13 @@ export function isSettingsWritable(): boolean {
// Test-only: reset module state
export function _resetForTests(): void {
cachedSettings = null
state.cachedSettings = null
writableCache = null
for (const key of overlaidKeys) {
const original = originalEnv[key]
for (const key of state.overlaidKeys) {
const original = state.originalEnv[key]
if (original === null) delete process.env[key]
else if (original !== undefined) process.env[key] = original
}
overlaidKeys = new Set()
for (const key of Object.keys(originalEnv)) delete originalEnv[key]
state.overlaidKeys = new Set()
state.originalEnv = {}
}
+85 -16
View File
@@ -10,6 +10,10 @@ import { aihubmix, createAihubmix } from "@aihubmix/ai-sdk-provider"
import { fromNodeProviderChain } from "@aws-sdk/credential-providers"
import { createOpenRouter } from "@openrouter/ai-sdk-provider"
import { createOllama, ollama } from "ollama-ai-provider-v2"
import {
adminProvidersToConfig,
loadAdminProviders,
} from "@/lib/admin/providers"
import { PROVIDER_INFO, type ProviderName } from "@/lib/types/model-config"
export type { ProviderName }
@@ -824,8 +828,16 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
// Use client-provided credentials if available, otherwise fall back to IAM/env vars
const hasClientCredentials =
overrides?.awsAccessKeyId && overrides?.awsSecretAccessKey
// Keys from the admin panel. The ADMIN_ names keep them out of the
// default AWS credential chain, which other clients such as the
// DynamoDB quota manager use with their own credentials.
const adminAccessKeyId = process.env.ADMIN_AWS_ACCESS_KEY_ID
const adminSecretAccessKey = process.env.ADMIN_AWS_SECRET_ACCESS_KEY
const bedrockRegion =
overrides?.awsRegion || process.env.AWS_REGION || "us-west-2"
overrides?.awsRegion ||
process.env.ADMIN_AWS_REGION ||
process.env.AWS_REGION ||
"us-west-2"
const bedrockProvider = hasClientCredentials
? createAmazonBedrock({
@@ -836,10 +848,16 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
sessionToken: overrides.awsSessionToken,
}),
})
: createAmazonBedrock({
region: bedrockRegion,
credentialProvider: fromNodeProviderChain(),
})
: adminAccessKeyId && adminSecretAccessKey
? createAmazonBedrock({
region: bedrockRegion,
accessKeyId: adminAccessKeyId,
secretAccessKey: adminSecretAccessKey,
})
: createAmazonBedrock({
region: bedrockRegion,
credentialProvider: fromNodeProviderChain(),
})
model = bedrockProvider(modelId)
// Add Anthropic beta options if using Claude models via Bedrock
if (modelId.includes("anthropic.claude")) {
@@ -872,8 +890,9 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
// for compatibility (most proxies don't support /responses endpoint)
const customOpenAI = createOpenAI({ apiKey, baseURL })
model = customOpenAI.chat(modelId)
} else if (overrides?.apiKey) {
// Custom API key but official OpenAI endpoint, use Responses API
} else if (overrides?.apiKey || overrides?.apiKeyEnv) {
// Custom API key (the client's, or a server model's own env var)
// but official OpenAI endpoint, use Responses API
// to support reasoning for gpt-5, o1, o3, o4 models
const customOpenAI = createOpenAI({ apiKey })
model = customOpenAI(modelId)
@@ -928,7 +947,9 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
overrides?.baseUrl,
serverBaseUrl,
)
if (baseURL || overrides?.apiKey) {
// The default instance only reads GOOGLE_GENERATIVE_AI_API_KEY, so a
// server model's own env var (apiKeyEnv) needs a custom instance too
if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) {
const customGoogle = createGoogleGenerativeAI({
apiKey,
...(baseURL && { baseURL }),
@@ -941,8 +962,11 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
}
case "vertexai": {
// Express Mode: Use API key for authentication
const vertexApiKey =
overrides?.vertexApiKey || process.env.GOOGLE_VERTEX_API_KEY
// SECURITY: a client base URL only ever gets the client's key, so the
// server's GOOGLE_VERTEX_API_KEY is never sent to a client-chosen host
const vertexApiKey = overrides?.baseUrl
? overrides.vertexApiKey
: overrides?.vertexApiKey || process.env.GOOGLE_VERTEX_API_KEY
if (!vertexApiKey) {
throw new Error(
@@ -951,9 +975,13 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
)
}
// Support custom base URL from env or client override
const baseURL =
overrides?.baseUrl || process.env.GOOGLE_VERTEX_BASE_URL
// Support custom base URL from env or client override.
// A client key only goes to the client's URL or the official one.
const baseURL = resolveBaseURL(
overrides?.vertexApiKey,
overrides?.baseUrl,
process.env.GOOGLE_VERTEX_BASE_URL,
)
const vertexProvider = createVertex({
apiKey: vertexApiKey,
@@ -1079,7 +1107,7 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
overrides?.baseUrl,
serverBaseUrl,
)
if (baseURL || overrides?.apiKey) {
if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) {
const customDeepSeek = createDeepSeek({
apiKey,
...(baseURL && { baseURL }),
@@ -1241,7 +1269,7 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
)
// Only use custom configuration if explicitly set (local dev or custom Gateway)
// Otherwise undefined → AI SDK uses Vercel default (https://ai-gateway.vercel.sh/v1/ai) + OIDC
if (baseURL || overrides?.apiKey) {
if (baseURL || overrides?.apiKey || overrides?.apiKeyEnv) {
const customGateway = createGateway({
apiKey,
...(baseURL && { baseURL }),
@@ -1430,6 +1458,36 @@ export function getAIModel(overrides?: ClientOverrides): ModelConfig {
return { model, providerOptions, headers, modelId, provider }
}
/**
* Whether the call is paid for by the server's own credentials (env keys or
* IAM role) rather than credentials sent with the request. Mirrors which key
* each branch of getAIModel ends up using.
*/
export function usesServerCredentials(
provider: ProviderName,
overrides?: ClientOverrides,
): boolean {
switch (provider) {
case "bedrock":
return !(overrides?.awsAccessKeyId && overrides?.awsSecretAccessKey)
case "vertexai":
return !overrides?.vertexApiKey
case "edgeone":
// The platform's own endpoint, no key involved
return false
case "ollama":
// Only a server key costs money; a keyless local server or the
// client's own server does not
return (
!overrides?.baseUrl &&
!overrides?.apiKey &&
!!(overrides?.apiKeyEnv || process.env.OLLAMA_API_KEY)
)
default:
return !overrides?.apiKey
}
}
/**
* Check if a model supports prompt caching.
* Currently only Claude models on Bedrock support prompt caching.
@@ -1464,6 +1522,17 @@ export function getValidationModel(): ReturnType<typeof getAIModel>["model"] {
)
}
const { model } = getAIModel({ modelId })
// A default set in the admin panel becomes AI_PROVIDER/AI_MODEL, but its key
// lives in an ADMIN_-prefixed env var. Point at it the way the chat route
// does for server models, or the standard env var is required instead.
const panelDefault = adminProvidersToConfig(
loadAdminProviders(),
).providers.find((p) => p.default && p.provider === process.env.AI_PROVIDER)
const { model } = getAIModel({
modelId,
apiKeyEnv: panelDefault?.apiKeyEnv,
baseUrlEnv: panelDefault?.baseUrlEnv,
})
return model
}
+10
View File
@@ -1,6 +1,8 @@
export interface CachedResponse {
promptText: string
hasImage: boolean
// Name of the bundled example file the prompt is sent with
fileName?: string
xml: string
}
@@ -254,6 +256,7 @@ export const CACHED_EXAMPLE_RESPONSES: CachedResponse[] = [
{
promptText: "Replicate this in aws style",
hasImage: true,
fileName: "architecture.png",
xml: `<mxCell id="2" value="AWS" style="sketch=0;outlineConnect=0;gradientColor=none;html=1;whiteSpace=wrap;fontSize=12;fontStyle=0;container=1;pointerEvents=0;collapsible=0;recursiveResize=0;shape=mxgraph.aws4.group;grIcon=mxgraph.aws4.group_aws_cloud;strokeColor=#232F3E;fillColor=none;verticalAlign=top;align=left;spacingLeft=30;fontColor=#232F3E;dashed=0;rounded=1;arcSize=5;" vertex="1" parent="1">
<mxGeometry x="340" y="40" width="880" height="520" as="geometry"/>
</mxCell>
@@ -318,6 +321,7 @@ export const CACHED_EXAMPLE_RESPONSES: CachedResponse[] = [
{
promptText: "Replicate this flowchart.",
hasImage: true,
fileName: "example.png",
xml: `<mxCell id="2" value="Lamp doesn't work" style="rounded=1;whiteSpace=wrap;html=1;fillColor=#ffcccc;strokeColor=#000000;strokeWidth=2;fontSize=18;fontStyle=0;" vertex="1" parent="1">
<mxGeometry x="140" y="40" width="180" height="60" as="geometry"/>
</mxCell>
@@ -379,6 +383,7 @@ export const CACHED_EXAMPLE_RESPONSES: CachedResponse[] = [
{
promptText: "Summarize this paper as a diagram",
hasImage: true,
fileName: "chain-of-thought.txt",
xml: `<mxCell id="title_bg" parent="1"
style="rounded=1;whiteSpace=wrap;html=1;fillColor=#1a237e;strokeColor=none;arcSize=8;"
value="" vertex="1">
@@ -879,14 +884,19 @@ export const CACHED_EXAMPLE_RESPONSES: CachedResponse[] = [
},
]
// Examples that come with a file only match when that exact example file is
// attached, so a user's own file with the same prompt still goes to the model.
// Callers that can't tell file names (the server) only get text-only examples.
export function findCachedResponse(
promptText: string,
hasImage: boolean,
fileName?: string,
): CachedResponse | undefined {
return CACHED_EXAMPLE_RESPONSES.find(
(c) =>
c.promptText === promptText &&
c.hasImage === hasImage &&
(!c.fileName || c.fileName === fileName) &&
c.xml !== "",
)
}
+96 -43
View File
@@ -6,25 +6,37 @@ export const MAX_FILE_SIZE = 2 * 1024 * 1024 // 2MB
export const MAX_FILES = 5
// Helper function to validate file parts in messages
// Checks every message, since history is sent to the model too
export function validateFileParts(messages: any[]): {
valid: boolean
error?: string
} {
const lastMessage = messages[messages.length - 1]
const fileParts =
lastMessage?.parts?.filter((p: any) => p.type === "file") || []
for (const message of messages) {
const fileParts =
message?.parts?.filter((p: any) => p.type === "file") || []
if (fileParts.length > MAX_FILES) {
return {
valid: false,
error: `Too many files. Maximum ${MAX_FILES} allowed.`,
if (fileParts.length > MAX_FILES) {
return {
valid: false,
error: `Too many files. Maximum ${MAX_FILES} allowed.`,
}
}
}
for (const filePart of fileParts) {
// Data URLs format: data:image/png;base64,<data>
// Base64 increases size by ~33%, so we check the decoded size
if (filePart.url?.startsWith("data:")) {
for (const filePart of fileParts) {
// The client sends files inline. Any other URL would be downloaded
// by the server (AI SDK does that for models without URL support).
if (
typeof filePart.url !== "string" ||
!filePart.url.startsWith("data:")
) {
return {
valid: false,
error: "Files must be uploaded inline as data URLs.",
}
}
// Data URLs format: data:image/png;base64,<data>
// Base64 increases size by ~33%, so we check the decoded size
const base64Data = filePart.url.split(",")[1]
if (base64Data) {
const sizeInBytes = Math.ceil((base64Data.length * 3) / 4)
@@ -42,48 +54,89 @@ export function validateFileParts(messages: any[]): {
}
// Helper function to check if diagram is minimal/empty
// Empty means no mxCell besides the root cells "0" and "1". Cells drawn in
// draw.io get random ids, so checking for id="2" is not enough.
export function isMinimalDiagram(xml: string): boolean {
const stripped = xml.replace(/\s/g, "")
return !stripped.includes('id="2"')
return !/<mxCell\b[^>]*\bid="(?![01]")/.test(xml)
}
// A tool-call input providers accept: a non-empty JSON object
function isValidToolInput(input: unknown): boolean {
return !!input && typeof input === "object" && Object.keys(input).length > 0
}
// Helper function to replace historical tool call XML with placeholders
// This reduces token usage and forces LLM to rely on the current diagram XML (source of truth)
// Also fixes invalid/undefined inputs from interrupted streaming
// Tool calls with invalid inputs are left for dropInvalidToolCalls to remove
export function replaceHistoricalToolInputs(messages: any[]): any[] {
return messages.map((msg) => {
if (msg.role !== "assistant" || !Array.isArray(msg.content)) {
return msg
}
const replacedContent = msg.content
.map((part: any) => {
if (part.type === "tool-call") {
const toolName = part.toolName
// Fix invalid/undefined inputs from interrupted streaming
if (
!part.input ||
typeof part.input !== "object" ||
Object.keys(part.input).length === 0
) {
// Skip tool calls with invalid inputs entirely
return null
}
if (
toolName === "display_diagram" ||
toolName === "edit_diagram"
) {
return {
...part,
input: {
placeholder:
"[XML content replaced - see current diagram XML in system context]",
},
}
}
const replacedContent = msg.content.map((part: any) => {
if (
part.type === "tool-call" &&
isValidToolInput(part.input) &&
(part.toolName === "display_diagram" ||
part.toolName === "edit_diagram")
) {
return {
...part,
input: {
placeholder:
"[XML content replaced - see current diagram XML in system context]",
},
}
return part
})
.filter(Boolean) // Remove null entries (invalid tool calls)
}
return part
})
return { ...msg, content: replacedContent }
})
}
// Remove tool-calls with invalid inputs (from failed repair or interrupted streaming),
// together with their tool-results: providers reject a result whose call is missing.
// Messages left empty are removed too (Bedrock rejects empty content arrays).
export function dropInvalidToolCalls(messages: any[]): any[] {
const droppedIds = new Set<string>()
return messages
.map((msg) => {
if (!Array.isArray(msg.content)) return msg
const content = msg.content.filter((part: any) => {
if (
msg.role === "assistant" &&
part.type === "tool-call" &&
!isValidToolInput(part.input)
) {
console.warn(
`[chat-helpers] Dropping tool-call with invalid input:`,
{ toolName: part.toolName, input: part.input },
)
droppedIds.add(part.toolCallId)
return false
}
// Results always come after their call, so the id is known by now
return !(
part.type === "tool-result" &&
droppedIds.has(part.toolCallId)
)
})
return { ...msg, content }
})
.filter((msg) => !Array.isArray(msg.content) || msg.content.length > 0)
}
// Fix common LLM JSON mistakes in tool-call input before jsonrepair runs
export function fixToolInputJson(input: string): string {
return (
input
// Inconsistent quote escaping in XML attributes inside JSON strings:
// y="-20\" (opening quote unescaped, closing escaped) becomes y=\"-20\".
// Must run before the key fix below, which would rewrite the `="`.
.replace(/(\w+)="([^"]*?)\\"/g, '$1=\\"$2\\"')
// `:=` instead of `: `
.replace(/:=/g, ": ")
// `"key"= "` instead of `"key": "`, only for JSON keys
.replace(/"(\w+)"\s*=\s*"/g, '"$1": "')
)
}
+2 -1
View File
@@ -188,7 +188,8 @@
"failedToExport": "Error fetching chart data",
"failedToLoadExample": "Error loading example image",
"failedToRecordFeedback": "Failed to record your feedback. Please try again.",
"storageUpdateFailed": "Chat cleared but browser storage could not be updated"
"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."
},
"quota": {
"dailyLimit": "Daily Quota Reached",
+2 -1
View File
@@ -188,7 +188,8 @@
"failedToExport": "チャートデータの取得エラー",
"failedToLoadExample": "例の画像の読み込みエラー",
"failedToRecordFeedback": "フィードバックの記録に失敗しました。もう一度お試しください。",
"storageUpdateFailed": "チャットはクリアされましたが、ブラウザストレージを更新できませんでした"
"storageUpdateFailed": "チャットはクリアされましたが、ブラウザストレージを更新できませんでした",
"sessionSaveFailed": "このチャットを保存できませんでした。ブラウザのストレージがいっぱいの可能性があります。履歴から古いチャットを削除して、もう一度お試しください。"
},
"quota": {
"dailyLimit": "1日の割当量に達しました",
+2 -1
View File
@@ -188,7 +188,8 @@
"failedToExport": "取得圖表資料時出錯",
"failedToLoadExample": "載入範例圖片時出錯",
"failedToRecordFeedback": "記錄您的回饋失敗。請重試。",
"storageUpdateFailed": "聊天已清除,但無法更新瀏覽器儲存空間"
"storageUpdateFailed": "聊天已清除,但無法更新瀏覽器儲存空間",
"sessionSaveFailed": "無法儲存這個對話。瀏覽器儲存空間可能已滿,請在歷史紀錄裡刪除舊對話後重試。"
},
"quota": {
"dailyLimit": "已達每日配額",
+2 -1
View File
@@ -188,7 +188,8 @@
"failedToExport": "获取图表数据时出错",
"failedToLoadExample": "加载示例图片时出错",
"failedToRecordFeedback": "记录您的反馈失败。请重试。",
"storageUpdateFailed": "聊天已清除,但无法更新浏览器存储"
"storageUpdateFailed": "聊天已清除,但无法更新浏览器存储",
"sessionSaveFailed": "无法保存这个对话。浏览器存储空间可能已满,请在历史记录里删除旧对话后重试。"
},
"quota": {
"dailyLimit": "已达每日配额",
+8 -1
View File
@@ -51,8 +51,15 @@ export function setTraceOutput(output: string) {
if (!isLangfuseEnabled()) return
updateActiveTrace({ output })
endTrace()
}
// End the observe() wrapper span (AI SDK creates its own child spans with usage).
// It uses endOnExit: false, so every request path has to end it, or the trace
// is never exported: stream finish, stream error/abort, and early returns.
export function endTrace() {
if (!isLangfuseEnabled()) return
// End the observe() wrapper span (AI SDK creates its own child spans with usage)
const activeSpan = api.trace.getActiveSpan()
if (activeSpan) {
activeSpan.end()
+117 -38
View File
@@ -22,6 +22,12 @@ export const MAX_OUTPUT_TOKENS_LIMIT = 200000
*/
const MIN_USABLE_OUTPUT_TOKENS = 1024
/**
* Retry budget when a rejection names the budget parameter but no number we can
* read. It is the default from before 64000, which these providers ran with.
*/
const FALLBACK_OUTPUT_TOKENS = 16000
/** Status codes that can carry a complaint about the requested budget. */
const BUDGET_REJECTION_STATUSES = new Set([400, 422])
@@ -29,24 +35,8 @@ function usableLimit(value: number): number | null {
return value >= MIN_USABLE_OUTPUT_TOKENS ? value : null
}
/**
* A budget this large exceeds what some models accept. Providers reject it with a
* 400 that names the real limit, so we parse the number out and retry once
* instead of failing the turn.
*
* Formats seen in the wild:
* - Bedrock: "The maximum tokens you requested exceeds the model limit of 4096."
* - OpenRouter: "This endpoint's maximum context length is 64000 tokens. However,
* you requested about 64025 tokens (25 of text input, 64000 in the output)."
* Note this one is an input+output ceiling, so the input has to be subtracted.
* - Anthropic: "max_tokens: 200000 > 64000, which is the maximum allowed..."
* - OpenAI: "This model supports at most 16384 completion tokens"
*
* Every pattern names tokens explicitly. A generic one (an earlier draft matched
* "lower than N") would reinterpret unrelated failures, and retrying on a bogus
* number turns a readable error into an empty diagram.
*/
export function parseOutputTokenLimit(error: unknown): number | null {
/** Message and body of an error that may be about the budget, or null. */
function rejectionText(error: unknown): string | null {
const err = error as {
message?: unknown
responseBody?: unknown
@@ -66,24 +56,109 @@ export function parseOutputTokenLimit(error: unknown): number | null {
typeof err?.responseBody === "string" ? err.responseBody : "",
].join(" ")
if (!text) return null
return text.trim() ? text : null
}
/**
* A budget this large exceeds what some models accept. Providers reject it with a
* 400 that names the real limit, so we parse the number out and retry once
* instead of failing the turn.
*
* Formats seen in the wild:
* - Bedrock: "The maximum tokens you requested exceeds the model limit of 4096."
* - OpenRouter: "This endpoint's maximum context length is 64000 tokens. However,
* you requested about 64025 tokens (25 of text input, 64000 in the output)."
* Note this one is an input+output ceiling, so the input has to be subtracted.
* vLLM and SGLang send the same kind of ceiling, with the input written as
* "6000 in the messages", "has 6000 input tokens" or "6000 tokens from the input".
* - Anthropic: "max_tokens: 200000 > 64000, which is the maximum allowed..."
* - OpenAI: "This model supports at most 16384 completion tokens"
* - Volcengine Ark: "The parameter `max_tokens` specified in the request are not
* valid: integer above maximum value, expected a value <= 32768, but got 64000"
* - DashScope: "Range of max_tokens should be [1, 8192]"
*
* Every pattern names tokens explicitly. A generic one (an earlier draft matched
* "lower than N") would reinterpret unrelated failures, and retrying on a bogus
* number turns a readable error into an empty diagram.
*/
function readCeiling(text: string): number | null {
// Combined input+output ceiling: subtract the input the provider counted,
// plus a small margin because its estimate is approximate.
const context = text.match(/maximum context length is (\d+)/i)
const context = text.match(/maximum context length (?:is|of) (\d+)/i)
if (context) {
const input = text.match(/(\d+) of text input/i)
return usableLimit(
Number(context[1]) - (input ? Number(input[1]) : 0) - 1024,
)
const input =
text.match(/(\d+) of text input/i) ||
text.match(/(\d+) in the messages/i) ||
text.match(/(\d+) tokens from the input/i) ||
text.match(/(\d+) input tokens/i)
return Number(context[1]) - (input ? Number(input[1]) : 0) - 1024
}
const output =
text.match(/model limit of (\d+)/i) ||
text.match(/> (\d+), which is the maximum/i) ||
text.match(/at most (\d+) completion tokens/i)
text.match(/at most (\d+) completion tokens/i) ||
text.match(/max_\w*tokens.*?expected a value (?:<=|\\u003c=) (\d+)/i) ||
text.match(/Range of max_tokens should be \[1,\s*(\d+)\]/i)
return output ? usableLimit(Number(output[1])) : null
return output ? Number(output[1]) : null
}
/** The usable output ceiling named in a rejection, or null. */
export function parseOutputTokenLimit(error: unknown): number | null {
const text = rejectionText(error)
const ceiling = text ? readCeiling(text) : null
return ceiling === null ? null : usableLimit(ceiling)
}
/**
* Thinking budget the provider adds on top of maxOutputTokens. Bedrock and
* Anthropic send maxOutputTokens + budgetTokens as max_tokens, so a ceiling in
* their rejection covers both.
*/
function thinkingBudget(providerOptions: unknown): number {
const options = providerOptions as
| {
bedrock?: {
reasoningConfig?: { type?: string; budgetTokens?: unknown }
}
anthropic?: {
thinking?: { type?: string; budgetTokens?: unknown }
}
}
| undefined
const config =
options?.bedrock?.reasoningConfig ?? options?.anthropic?.thinking
return config?.type === "enabled" && typeof config.budgetTokens === "number"
? config.budgetTokens
: 0
}
/**
* The budget to retry with after a rejection, or null to surface the error.
*/
export function retryOutputTokens(
error: unknown,
params: { maxOutputTokens?: number; providerOptions?: unknown },
): number | null {
const requested = params.maxOutputTokens
const text = rejectionText(error)
if (!requested || !text) return null
const ceiling = readCeiling(text)
if (ceiling !== null) {
// The ceiling applies to what was actually sent, thinking included,
// so the retry has to leave room for the thinking too.
const thinking = thinkingBudget(params.providerOptions)
if (ceiling >= requested + thinking) return null
return usableLimit(ceiling - thinking)
}
// Names the budget parameter, but in a format we cannot read a number from
if (/max_\w*tokens/i.test(text) && requested > FALLBACK_OUTPUT_TOKENS) {
return FALLBACK_OUTPUT_TOKENS
}
return null
}
/**
@@ -103,17 +178,15 @@ export function withOutputTokenLimitFallback(
try {
return await doStream()
} catch (error) {
const limit = parseOutputTokenLimit(error)
const requested = params.maxOutputTokens
if (!limit || !requested || limit >= requested) throw error
const retry = retryOutputTokens(error, params)
if (!retry) throw error
console.warn(
`[maxOutputTokens] ${requested} rejected, retrying with ${limit}`,
`[maxOutputTokens] ${params.maxOutputTokens} rejected, retrying with ${retry}`,
)
return await inner.doStream({
...params,
maxOutputTokens: limit,
maxOutputTokens: retry,
})
}
},
@@ -135,11 +208,17 @@ function validBudget(value: string | null | undefined): number | null {
* desktop app too), then server env, then the default. Both sources go through
* the same validation, so a typo in either falls back instead of reaching the
* provider.
*
* On the server's credentials the user setting can only lower the server value,
* so MAX_OUTPUT_TOKENS keeps capping what the server pays for.
*/
export function resolveMaxOutputTokens(headerValue: string | null): number {
return (
validBudget(headerValue) ??
validBudget(process.env.MAX_OUTPUT_TOKENS) ??
DEFAULT_MAX_OUTPUT_TOKENS
)
export function resolveMaxOutputTokens(
headerValue: string | null,
usesServerCredentials: boolean,
): number {
const header = validBudget(headerValue)
const server =
validBudget(process.env.MAX_OUTPUT_TOKENS) ?? DEFAULT_MAX_OUTPUT_TOKENS
if (header === null) return server
return usesServerCredentials ? Math.min(header, server) : header
}
+6 -3
View File
@@ -1,4 +1,4 @@
import { extractText, getDocumentProxy } from "unpdf"
import { extractText } from "unpdf"
// Maximum characters allowed for extracted text (configurable via env)
const DEFAULT_MAX_EXTRACTED_CHARS = 150000 // 150k chars
@@ -14,6 +14,7 @@ const TEXT_EXTENSIONS = [
".json",
".csv",
".xml",
".svg",
".html",
".css",
".js",
@@ -43,8 +44,10 @@ const TEXT_EXTENSIONS = [
*/
export async function extractPdfText(file: File): Promise<string> {
const buffer = await file.arrayBuffer()
const pdf = await getDocumentProxy(new Uint8Array(buffer))
const { text } = await extractText(pdf, { mergePages: true })
// Pass raw bytes so unpdf destroys the PDF document when it is done
const { text } = await extractText(new Uint8Array(buffer), {
mergePages: true,
})
return text as string
}
+16 -2
View File
@@ -47,11 +47,14 @@ export interface FlattenedServerModel {
/**
* Convert provider name to URL-safe slug for use in model ID
* e.g., "OpenAI Production" → "openai-production"
* e.g., "OpenAI Production" → "openai-production", "主力" → "4e3b-529b"
* Non-ASCII characters become their hex code point so CJK names stay
* distinct; the id is sent in HTTP headers, which must be ASCII.
*/
function slugify(name: string): string {
export function slugify(name: string): string {
return name
.toLowerCase()
.replace(/[^\p{ASCII}]/gu, (c) => `-${c.codePointAt(0)?.toString(16)}-`)
.replace(/[^a-z0-9]+/g, "-")
.replace(/^-|-$/g, "")
}
@@ -189,6 +192,7 @@ export async function loadFlattenedServerModels(): Promise<
const defaultModelId = process.env.AI_MODEL
const flattened: FlattenedServerModel[] = []
const seenIds = new Set<string>()
for (const p of cfg.providers) {
const providerLabel =
@@ -199,6 +203,16 @@ export async function loadFlattenedServerModels(): Promise<
for (const modelId of p.models) {
const id = `server:${nameSlug}:${modelId}`
// Names that differ only in case or punctuation share a slug.
// A repeated id would always resolve to the first provider's
// credentials, so drop it instead.
if (seenIds.has(id)) {
console.warn(
`[server-model-config] Skipping duplicate model id "${id}". Provider names must differ in letters or digits.`,
)
continue
}
seenIds.add(id)
// Default model priority:
// 1. From ai-models.json: first model of provider with default: true
+31 -23
View File
@@ -1,5 +1,6 @@
import { type DBSchema, type IDBPDatabase, openDB } from "idb"
import { nanoid } from "nanoid"
import { toast } from "sonner"
import type { Template } from "./template-storage"
// Constants
@@ -61,6 +62,7 @@ let dbPromise: Promise<IDBPDatabase<ChatSessionDB>> | null = null
async function getDB(): Promise<IDBPDatabase<ChatSessionDB>> {
if (!dbPromise) {
// A failed or lost connection is not cached: the next call reopens it
dbPromise = openDB<ChatSessionDB>(DB_NAME, DB_VERSION, {
upgrade(db, oldVersion) {
if (oldVersion < 1) {
@@ -88,6 +90,28 @@ async function getDB(): Promise<IDBPDatabase<ChatSessionDB>> {
}
}
},
blocked() {
// An older tab keeps the DB open, so the upgrade has to wait
toast.warning(
"Please close other tabs of this app to finish updating chat storage.",
{ id: "idb-upgrade-blocked", duration: 10000 },
)
},
blocking(_currentVersion, _blockedVersion, event) {
// Another tab needs to upgrade the DB: close our connection so
// it is not stuck, and reopen on the next call
const db = event.target as IDBDatabase
db.close()
dbPromise = null
},
terminated() {
// The browser closed the connection (e.g. Safari after a long
// time in the background)
dbPromise = null
},
}).catch((error) => {
dbPromise = null
throw error
})
}
return dbPromise
@@ -145,6 +169,8 @@ export async function getSession(id: string): Promise<ChatSession | null> {
}
}
// Returns false on failure (e.g. storage quota exceeded). Other sessions are
// never deleted automatically; the caller tells the user instead.
export async function saveSession(session: ChatSession): Promise<boolean> {
if (!isIndexedDBAvailable()) return false
try {
@@ -152,29 +178,11 @@ export async function saveSession(session: ChatSession): Promise<boolean> {
await db.put(STORE_NAME, session)
return true
} catch (error) {
// Handle quota exceeded
if (
error instanceof DOMException &&
error.name === "QuotaExceededError"
) {
console.warn("Storage quota exceeded, deleting oldest session...")
await deleteOldestSession()
// Retry once
try {
const db = await getDB()
await db.put(STORE_NAME, session)
return true
} catch (retryError) {
console.error(
"Failed to save session after cleanup:",
retryError,
)
return false
}
} else {
console.error("Failed to save session:", error)
return false
}
console.error("Failed to save session:", error)
// Reopen the connection next time in case it was lost (Safari reports
// "Connection to Indexed Database server lost" without closing it)
dbPromise = null
return false
}
}
+1 -1
View File
@@ -41,7 +41,7 @@ parameters: {
tool name: edit_diagram
description: Edit specific parts of the EXISTING diagram. Use this when making small targeted changes like adding/removing elements, changing labels, or adjusting properties. This is more efficient than regenerating the entire diagram.
parameters: {
edits: Array<{search: string, replace: string}>
operations: Array<{operation: "update" | "add" | "delete", cell_id: string, new_xml?: string}>
}
---Tool3---
tool name: append_diagram
+6 -1
View File
@@ -1,5 +1,6 @@
import { z } from "zod"
import { getApiEndpoint } from "@/lib/base-path"
import { STORAGE_KEYS } from "@/lib/storage"
export interface UrlData {
url: string
@@ -18,7 +19,11 @@ const UrlResponseSchema = z.object({
export async function extractUrlContent(url: string): Promise<UrlData> {
const response = await fetch(getApiEndpoint("/api/parse-url"), {
method: "POST",
headers: { "Content-Type": "application/json" },
headers: {
"Content-Type": "application/json",
"x-access-code":
localStorage.getItem(STORAGE_KEYS.accessCode) || "",
},
body: JSON.stringify({ url }),
})
+55 -61
View File
@@ -27,78 +27,72 @@ export function useFileProcessor() {
const handleFileChange = async (newFiles: File[]) => {
setFiles(newFiles)
// Extract text immediately for new PDF/text files
for (const file of newFiles) {
const needsExtraction =
(isPdfFile(file) || isTextFile(file)) && !pdfData.has(file)
if (needsExtraction) {
// Mark as extracting
setPdfData((prev) => {
const next = new Map(prev)
next.set(file, {
text: "",
charCount: 0,
isExtracting: true,
})
return next
})
const pending = newFiles.filter(
(file) =>
(isPdfFile(file) || isTextFile(file)) && !pdfData.has(file),
)
// Extract text asynchronously
try {
let text: string
if (isPdfFile(file)) {
text = await extractPdfText(file)
} else {
text = await extractTextFileContent(file)
}
// Before any await: drop data for removed files and mark every new
// file as extracting, so queued files also block sending
setPdfData((prev) => {
const next = new Map<File, FileData>()
for (const file of newFiles) {
const existing = prev.get(file)
if (existing) next.set(file, existing)
}
for (const file of pending) {
next.set(file, { text: "", charCount: 0, isExtracting: true })
}
return next
})
// Check character limit
if (text.length > MAX_EXTRACTED_CHARS) {
const limitK = MAX_EXTRACTED_CHARS / 1000
toast.error(
`${file.name}: Content exceeds ${limitK}k character limit (${(text.length / 1000).toFixed(1)}k chars)`,
)
setPdfData((prev) => {
const next = new Map(prev)
next.delete(file)
return next
})
// Remove the file from the list
setFiles((prev) => prev.filter((f) => f !== file))
continue
}
// Extract one file at a time
for (const file of pending) {
try {
let text: string
if (isPdfFile(file)) {
text = await extractPdfText(file)
} else {
text = await extractTextFileContent(file)
}
setPdfData((prev) => {
const next = new Map(prev)
next.set(file, {
text,
charCount: text.length,
isExtracting: false,
})
return next
})
} catch (error) {
console.error("Failed to extract text:", error)
toast.error(`Failed to read file: ${file.name}`)
// Check character limit
if (text.length > MAX_EXTRACTED_CHARS) {
const limitK = MAX_EXTRACTED_CHARS / 1000
toast.error(
`${file.name}: Content exceeds ${limitK}k character limit (${(text.length / 1000).toFixed(1)}k chars)`,
)
setPdfData((prev) => {
const next = new Map(prev)
next.delete(file)
return next
})
// Remove the file from the list
setFiles((prev) => prev.filter((f) => f !== file))
continue
}
setPdfData((prev) => {
// The file was removed while extracting
if (!prev.has(file)) return prev
const next = new Map(prev)
next.set(file, {
text,
charCount: text.length,
isExtracting: false,
})
return next
})
} catch (error) {
console.error("Failed to extract text:", error)
toast.error(`Failed to read file: ${file.name}`)
setPdfData((prev) => {
const next = new Map(prev)
next.delete(file)
return next
})
}
}
// Clean up pdfData for removed files
setPdfData((prev) => {
const next = new Map(prev)
for (const key of prev.keys()) {
if (!newFiles.includes(key)) {
next.delete(key)
}
}
return next
})
}
return {
+262 -193
View File
@@ -76,6 +76,17 @@ export function isMxCellXmlComplete(xml: string | undefined | null): boolean {
// No valid ending found at all
if (lastValidEnd === -1) return false
// If the last mxCell has no </mxCell> after it, it must be self-closing.
// Otherwise the trailing "/>" belongs to a child such as <mxGeometry .../>
// and the output was cut off before the cell was closed.
const lastCellStart = trimmed.lastIndexOf("<mxCell")
if (
lastCellStart > lastMxCellClose &&
!/^<mxCell\b[^<]*\/>/.test(trimmed.slice(lastCellStart))
) {
return false
}
// Check what comes after the last valid ending
// For />: add 2 chars, for </mxCell>: add 9 chars
const endOffset = lastMxCellClose > lastSelfClose ? 9 : 2
@@ -95,36 +106,12 @@ export function isMxCellXmlComplete(xml: string | undefined | null): boolean {
export function extractCompleteMxCells(xml: string | undefined | null): string {
if (!xml) return ""
const completeCells: Array<{ index: number; text: string }> = []
// Match self-closing <mxCell ... /> or <mxCell ...>...</mxCell>, in document order.
// The lazy [^>]*? tries "/>" first, so a self-closing cell never swallows
// the following cells up to the next </mxCell>.
const cellPattern = /<mxCell\b[^>]*?(?:\/>|>[\s\S]*?<\/mxCell>)/g
// Match self-closing mxCell tags: <mxCell ... />
// Also match mxCell with nested mxGeometry: <mxCell ...>...<mxGeometry .../></mxCell>
const selfClosingPattern = /<mxCell\s+[^>]*\/>/g
const nestedPattern = /<mxCell\s+[^>]*>[\s\S]*?<\/mxCell>/g
// Find all self-closing mxCell elements
let match: RegExpExecArray | null
while ((match = selfClosingPattern.exec(xml)) !== null) {
completeCells.push({ index: match.index, text: match[0] })
}
// Find all mxCell elements with nested content (like mxGeometry)
while ((match = nestedPattern.exec(xml)) !== null) {
completeCells.push({ index: match.index, text: match[0] })
}
// Sort by position to maintain order
completeCells.sort((a, b) => a.index - b.index)
// Remove duplicates (a self-closing match might overlap with nested match)
const seen = new Set<number>()
const uniqueCells = completeCells.filter((cell) => {
if (seen.has(cell.index)) return false
seen.add(cell.index)
return true
})
return uniqueCells.map((c) => c.text).join("\n")
return (xml.match(cellPattern) || []).join("\n")
}
// ============================================================================
@@ -487,6 +474,31 @@ export interface ApplyOperationsResult {
errors: OperationError[]
}
/**
* draw.io wraps cells that have links, tooltips or custom data in
* <object>/<UserObject>, and the wrapper carries the id instead of the mxCell.
*/
function getCellWrapper(cell: Element): Element | null {
const parent = cell.parentElement
return parent?.tagName === "object" || parent?.tagName === "UserObject"
? parent
: null
}
/** Id of a cell, read from its wrapper when the mxCell has none */
function getCellId(cell: Element): string | null {
return (
cell.getAttribute("id") ||
getCellWrapper(cell)?.getAttribute("id") ||
null
)
}
/** Element to replace or remove for a cell (the wrapper if there is one) */
function getCellNode(cell: Element): Element {
return getCellWrapper(cell) || cell
}
/**
* Apply diagram operations (update/add/delete) using ID-based lookup.
* This replaces the text-matching approach with direct DOM manipulation.
@@ -535,12 +547,14 @@ export function applyDiagramOperations(
}
}
// Build a map of cell IDs to elements
// Build a map of cell IDs to elements (wrapper elements for wrapped cells)
const cellMap = new Map<string, Element>()
root.querySelectorAll("mxCell").forEach((cell) => {
const id = cell.getAttribute("id")
if (id) cellMap.set(id, cell)
const id = getCellId(cell)
if (id) cellMap.set(id, getCellNode(cell))
})
// Cells removed by delete operations in this batch
const deletedIds = new Set<string>()
// Process each operation
for (const op of operations) {
@@ -580,7 +594,7 @@ export function applyDiagramOperations(
}
// Validate ID matches
const newCellId = newCell.getAttribute("id")
const newCellId = getCellId(newCell)
if (newCellId !== op.cell_id) {
errors.push({
type: "update",
@@ -590,8 +604,8 @@ export function applyDiagramOperations(
continue
}
// Import and replace the node
const importedNode = doc.importNode(newCell, true)
// Import and replace the node (with its wrapper, if any)
const importedNode = doc.importNode(getCellNode(newCell), true)
existingCell.parentNode?.replaceChild(importedNode, existingCell)
// Update the map with the new element
@@ -632,7 +646,7 @@ export function applyDiagramOperations(
}
// Validate ID matches
const newCellId = newCell.getAttribute("id")
const newCellId = getCellId(newCell)
if (newCellId !== op.cell_id) {
errors.push({
type: "add",
@@ -642,8 +656,8 @@ export function applyDiagramOperations(
continue
}
// Import and append the node
const importedNode = doc.importNode(newCell, true)
// Import and append the node (with its wrapper, if any)
const importedNode = doc.importNode(getCellNode(newCell), true)
root.appendChild(importedNode)
// Add to map
@@ -661,8 +675,15 @@ export function applyDiagramOperations(
const existingCell = cellMap.get(op.cell_id)
if (!existingCell) {
// Cell not found - might have been cascade-deleted by a previous operation
// Skip silently instead of erroring (AI may redundantly list children/edges)
// Cells cascade-deleted earlier in this batch are skipped silently
// (AI may redundantly list children/edges)
if (!deletedIds.has(op.cell_id)) {
errors.push({
type: "delete",
cellId: op.cell_id,
message: `Cell with id="${op.cell_id}" not found`,
})
}
continue
}
@@ -679,7 +700,7 @@ export function applyDiagramOperations(
`mxCell[parent="${cellId}"]`,
)
children.forEach((child) => {
const childId = child.getAttribute("id")
const childId = getCellId(child)
if (childId && childId !== "0" && childId !== "1") {
collectDescendants(childId)
}
@@ -696,7 +717,7 @@ export function applyDiagramOperations(
`mxCell[source="${cellId}"], mxCell[target="${cellId}"]`,
)
referencingEdges.forEach((edge) => {
const edgeId = edge.getAttribute("id")
const edgeId = getCellId(edge)
// Protect root cells from being added via edge references
if (edgeId && edgeId !== "0" && edgeId !== "1") {
// Recurse to collect edge's children (like labels)
@@ -718,6 +739,7 @@ export function applyDiagramOperations(
if (cell) {
cell.parentNode?.removeChild(cell)
cellMap.delete(cellId)
deletedIds.add(cellId)
}
}
}
@@ -758,24 +780,89 @@ function checkDuplicateAttributes(xml: string): string | null {
return null
}
/** Check for duplicate IDs in XML */
function checkDuplicateIds(xml: string): string | null {
const idPattern = /\bid\s*=\s*["']([^"']+)["']/gi
/** Matches one <diagram> page of a document (the last one may be unclosed) */
const PAGE_PATTERN = /<diagram\b[\s\S]*?(?:<\/diagram>|$)/g
const ID_ATTR_PATTERN = /\bid\s*=\s*["']([^"']+)["']/gi
/**
* Split XML into pages. Ids only need to be unique within a page: every
* page of a multi-page document has its own root cells "0" and "1".
*/
function splitPages(xml: string): string[] {
return xml.match(PAGE_PATTERN) || [xml]
}
/** Ids that appear more than once, with their counts */
function findDuplicateIds(xml: string): Map<string, number> {
const ids = new Map<string, number>()
let idMatch
while ((idMatch = idPattern.exec(xml)) !== null) {
const id = idMatch[1]
ids.set(id, (ids.get(id) || 0) + 1)
for (const match of xml.matchAll(ID_ATTR_PATTERN)) {
ids.set(match[1], (ids.get(match[1]) || 0) + 1)
}
const duplicateIds = Array.from(ids.entries())
.filter(([, count]) => count > 1)
.map(([id, count]) => `'${id}' (${count}x)`)
if (duplicateIds.length > 0) {
return `Invalid XML: Found duplicate ID(s): ${duplicateIds.slice(0, 3).join(", ")}. All id attributes must be unique.`
return new Map(Array.from(ids).filter(([, count]) => count > 1))
}
/** Check for duplicate IDs in XML (per page) */
function checkDuplicateIds(xml: string): string | null {
for (const page of splitPages(xml)) {
const duplicateIds = Array.from(findDuplicateIds(page)).map(
([id, count]) => `'${id}' (${count}x)`,
)
if (duplicateIds.length > 0) {
return `Invalid XML: Found duplicate ID(s): ${duplicateIds.slice(0, 3).join(", ")}. All id attributes must be unique.`
}
}
return null
}
/** Rename repeated ids in one page (keeps the first occurrence) */
function renameDuplicateIds(xml: string): { xml: string; renamed: number } {
const duplicateIds = findDuplicateIds(xml)
if (duplicateIds.size === 0) return { xml, renamed: 0 }
const idCounters = new Map<string, number>()
const renamedXml = xml.replace(ID_ATTR_PATTERN, (match, id) => {
if (!duplicateIds.has(id)) return match
const count = idCounters.get(id) || 0
idCounters.set(id, count + 1)
if (count === 0) return match // Keep first occurrence
// Rename subsequent occurrences (the id sits just before the closing quote)
return `${match.slice(0, -id.length - 1)}${id}_dup${count}${match.slice(-1)}`
})
return { xml: renamedXml, renamed: duplicateIds.size }
}
/**
* Returns a function telling whether a position is inside a quoted attribute
* value. Positions must be queried in increasing order: the scan resumes where
* it stopped instead of starting over, which keeps large documents fast.
*/
function createQuoteTracker(str: string): (pos: number) => boolean {
let i = 0
let inQuote = false
let quoteChar = ""
return (pos: number) => {
for (; i < pos && i < str.length; i++) {
const c = str[i]
if (inQuote) {
if (c === quoteChar) inQuote = false
} else if (c === '"' || c === "'") {
// Only quotes that follow "=" open an attribute value
let j = i - 1
while (j >= 0 && /\s/.test(str[j])) j--
if (j >= 0 && str[j] === "=") {
inQuote = true
quoteChar = c
}
}
}
return inQuote
}
}
/** Check for tag mismatches using parsed tags */
function checkTagMismatches(xml: string): string | null {
const xmlWithoutComments = xml.replace(/<!--[\s\S]*?-->/g, "")
@@ -1088,13 +1175,19 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
// 3b. Fix malformed attribute values where &quot; is used as delimiter instead of actual quotes
// Pattern: attr=&quot;value&quot; should become attr="value" (the &quot; was meant to be the quote delimiter)
// This commonly happens with dashPattern=&quot;1 1;&quot;
const malformedQuotePattern = /(\s[a-zA-Z][a-zA-Z0-9_:-]*)=&quot;/
if (malformedQuotePattern.test(fixed)) {
// Replace =&quot; with =" and trailing &quot; before next attribute or tag end with "
fixed = fixed.replace(
/(\s[a-zA-Z][a-zA-Z0-9_:-]*)=&quot;([^&]*?)&quot;/g,
'$1="$2"',
)
// Matches inside another attribute value are kept: rich text labels like
// value="&lt;font color=&quot;#ff0000&quot;&gt;..." are valid.
const isInsideQuotesFor3b = createQuoteTracker(fixed)
let malformedQuotesFixed = false
fixed = fixed.replace(
/(\s[a-zA-Z][a-zA-Z0-9_:-]*)=&quot;([^&]*?)&quot;/g,
(match: string, attr: string, value: string, offset: number) => {
if (isInsideQuotesFor3b(offset)) return match
malformedQuotesFixed = true
return `${attr}="${value}"`
},
)
if (malformedQuotesFixed) {
fixes.push(
'Fixed malformed attribute quotes (=&quot;...&quot; to ="...")',
)
@@ -1108,9 +1201,11 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
}
// 3d. Fix missing space between attributes like vertex="1"parent="1"
const missingSpacePattern = /("[^"]*")([a-zA-Z][a-zA-Z0-9_:-]*=)/g
// Requires name=" right after the quote, so the opening quote of a value
// such as style="rounded=1;..." is not mistaken for a closing one.
const missingSpacePattern = /"([a-zA-Z_:][\w:.-]*=")/g
if (missingSpacePattern.test(fixed)) {
fixed = fixed.replace(/("[^"]*")([a-zA-Z][a-zA-Z0-9_:-]*=)/g, "$1 $2")
fixed = fixed.replace(missingSpacePattern, '" $1')
fixes.push("Added missing space between attributes")
}
@@ -1240,32 +1335,13 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
"mxPoint",
"Array",
"Object",
// Wrappers of cells with links, tooltips or custom data
"object",
"UserObject",
"mxRectangle",
])
// Helper: Check if a position is inside a quoted attribute value
// by counting unescaped quotes before that position
const isInsideQuotes = (str: string, pos: number): boolean => {
let inQuote = false
let quoteChar = ""
for (let i = 0; i < pos && i < str.length; i++) {
const c = str[i]
if (inQuote) {
if (c === quoteChar) inQuote = false
} else if (c === '"' || c === "'") {
// Check if this quote is part of an attribute (preceded by =)
// Look back for = sign
let j = i - 1
while (j >= 0 && /\s/.test(str[j])) j--
if (j >= 0 && str[j] === "=") {
inQuote = true
quoteChar = c
}
}
}
return inQuote
}
const isInsideQuotesFor8c = createQuoteTracker(fixed)
const foreignTagPattern = /<\/?([a-zA-Z][a-zA-Z0-9_]*)[^>]*>/g
let foreignMatch
const foreignTags = new Set<string>()
@@ -1280,7 +1356,7 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
// Skip if this is a valid draw.io tag
if (validDrawioTags.has(tagName)) continue
// Skip if this tag is inside a quoted attribute value
if (isInsideQuotes(fixed, foreignMatch.index)) continue
if (isInsideQuotesFor8c(foreignMatch.index)) continue
foreignTags.add(tagName)
foreignTagPositions.push({
@@ -1352,10 +1428,11 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
>()
// Match full tags to detect self-closing by checking if ends with />
const fullTagPattern = /<(\/?[a-zA-Z][a-zA-Z0-9]*)[^>]*>/g
const isInsideQuotesFor10b = createQuoteTracker(fixed)
let tagCountMatch
while ((tagCountMatch = fullTagPattern.exec(fixed)) !== null) {
// Skip tags inside quoted attribute values (e.g., value="<b>Title</b>")
if (isInsideQuotes(fixed, tagCountMatch.index)) continue
if (isInsideQuotesFor10b(tagCountMatch.index)) continue
const fullMatch = tagCountMatch[0] // e.g., "<mxCell .../>" or "</mxCell>"
const tagPart = tagCountMatch[1] // e.g., "mxCell" or "/mxCell"
@@ -1445,125 +1522,112 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
// 11. Fix nested mxCell by flattening
// Pattern A: <mxCell id="X">...<mxCell id="X">...</mxCell></mxCell> (duplicate ID)
// Pattern B: <mxCell id="X">...<mxCell id="Y">...</mxCell></mxCell> (different ID - true nesting)
const lines = fixed.split("\n")
let newLines: string[] = []
let nestedFixed = 0
let extraClosingToRemove = 0
// These passes work line by line and would break valid cells written on a
// single line, so each one runs only when cells are really nested.
if (checkNestedMxCells(fixed)) {
const lines = fixed.split("\n")
const newLines: string[] = []
let nestedFixed = 0
let extraClosingToRemove = 0
// First pass: fix duplicate ID nesting (same as before)
for (let i = 0; i < lines.length; i++) {
const line = lines[i]
const nextLine = lines[i + 1]
// First pass: fix duplicate ID nesting (same as before)
for (let i = 0; i < lines.length; i++) {
const line = lines[i]
const nextLine = lines[i + 1]
// Check if current line and next line are both mxCell opening tags with same ID
if (
nextLine &&
/<mxCell\s/.test(line) &&
/<mxCell\s/.test(nextLine) &&
!line.includes("/>") &&
!nextLine.includes("/>")
) {
const id1 = line.match(/\bid\s*=\s*["']([^"']+)["']/)?.[1]
const id2 = nextLine.match(/\bid\s*=\s*["']([^"']+)["']/)?.[1]
// Check if current line and next line are both mxCell opening tags with same ID
if (
nextLine &&
/<mxCell\s/.test(line) &&
/<mxCell\s/.test(nextLine) &&
!line.includes("/>") &&
!nextLine.includes("/>")
) {
const id1 = line.match(/\bid\s*=\s*["']([^"']+)["']/)?.[1]
const id2 = nextLine.match(/\bid\s*=\s*["']([^"']+)["']/)?.[1]
if (id1 && id1 === id2) {
nestedFixed++
extraClosingToRemove++ // Need to remove one </mxCell> later
continue // Skip this duplicate opening line
if (id1 && id1 === id2) {
nestedFixed++
extraClosingToRemove++ // Need to remove one </mxCell> later
continue // Skip this duplicate opening line
}
}
}
// Remove extra </mxCell> if we have pending removals
if (extraClosingToRemove > 0 && /^\s*<\/mxCell>\s*$/.test(line)) {
extraClosingToRemove--
continue // Skip this closing tag
}
newLines.push(line)
}
if (nestedFixed > 0) {
fixed = newLines.join("\n")
fixes.push(`Flattened ${nestedFixed} duplicate-ID nested mxCell(s)`)
}
// Second pass: fix true nesting (different IDs)
// Insert </mxCell> before nested child to close parent
const lines2 = fixed.split("\n")
newLines = []
let trueNestedFixed = 0
let cellDepth = 0
let pendingCloseRemoval = 0
for (let i = 0; i < lines2.length; i++) {
const line = lines2[i]
const trimmed = line.trim()
// Track mxCell depth
const isOpenCell = /<mxCell\s/.test(trimmed) && !trimmed.endsWith("/>")
const isCloseCell = trimmed === "</mxCell>"
if (isOpenCell) {
if (cellDepth > 0) {
// Found nested cell - insert closing tag for parent before this line
const indent = line.match(/^(\s*)/)?.[1] || ""
newLines.push(indent + "</mxCell>")
trueNestedFixed++
pendingCloseRemoval++ // Need to remove one </mxCell> later
// Remove extra </mxCell> if we have pending removals
if (extraClosingToRemove > 0 && /^\s*<\/mxCell>\s*$/.test(line)) {
extraClosingToRemove--
continue // Skip this closing tag
}
cellDepth = 1 // Reset to 1 since we just opened a new cell
newLines.push(line)
} else if (isCloseCell) {
if (pendingCloseRemoval > 0) {
pendingCloseRemoval--
// Skip this extra closing tag
}
if (nestedFixed > 0) {
fixed = newLines.join("\n")
fixes.push(`Flattened ${nestedFixed} duplicate-ID nested mxCell(s)`)
}
}
if (checkNestedMxCells(fixed)) {
// Second pass: fix true nesting (different IDs)
// Insert </mxCell> before nested child to close parent
const lines2 = fixed.split("\n")
const newLines: string[] = []
let trueNestedFixed = 0
let cellDepth = 0
let pendingCloseRemoval = 0
for (let i = 0; i < lines2.length; i++) {
const line = lines2[i]
const trimmed = line.trim()
// Track mxCell depth
const isOpenCell =
/<mxCell\s/.test(trimmed) && !trimmed.endsWith("/>")
const isCloseCell = trimmed === "</mxCell>"
if (isOpenCell) {
if (cellDepth > 0) {
// Found nested cell - insert closing tag for parent before this line
const indent = line.match(/^(\s*)/)?.[1] || ""
newLines.push(indent + "</mxCell>")
trueNestedFixed++
pendingCloseRemoval++ // Need to remove one </mxCell> later
}
cellDepth = 1 // Reset to 1 since we just opened a new cell
newLines.push(line)
} else if (isCloseCell) {
if (pendingCloseRemoval > 0) {
pendingCloseRemoval--
// Skip this extra closing tag
} else {
cellDepth = Math.max(0, cellDepth - 1)
newLines.push(line)
}
} else {
cellDepth = Math.max(0, cellDepth - 1)
newLines.push(line)
}
} else {
newLines.push(line)
}
if (trueNestedFixed > 0) {
fixed = newLines.join("\n")
fixes.push(`Fixed ${trueNestedFixed} true nested mxCell(s)`)
}
}
if (trueNestedFixed > 0) {
fixed = newLines.join("\n")
fixes.push(`Fixed ${trueNestedFixed} true nested mxCell(s)`)
// 12. Fix duplicate IDs by appending suffix, page by page (ids such as the
// root cells "0" and "1" legitimately repeat across pages)
let renamedIds = 0
const renamePage = (page: string) => {
const { xml: renamed, renamed: count } = renameDuplicateIds(page)
renamedIds += count
return renamed
}
// 12. Fix duplicate IDs by appending suffix
const seenIds = new Map<string, number>()
const duplicateIds: string[] = []
// First pass: find duplicates
const idPattern = /\bid\s*=\s*["']([^"']+)["']/gi
let idMatch
while ((idMatch = idPattern.exec(fixed)) !== null) {
const id = idMatch[1]
seenIds.set(id, (seenIds.get(id) || 0) + 1)
}
// Find which IDs are duplicated
for (const [id, count] of seenIds) {
if (count > 1) duplicateIds.push(id)
}
// Second pass: rename duplicates (keep first occurrence, rename others)
if (duplicateIds.length > 0) {
const idCounters = new Map<string, number>()
fixed = fixed.replace(/\bid\s*=\s*["']([^"']+)["']/gi, (match, id) => {
if (!duplicateIds.includes(id)) return match
const count = idCounters.get(id) || 0
idCounters.set(id, count + 1)
if (count === 0) return match // Keep first occurrence
// Rename subsequent occurrences
const newId = `${id}_dup${count}`
return match.replace(id, newId)
})
fixes.push(`Renamed ${duplicateIds.length} duplicate ID(s)`)
fixed = /<diagram\b/.test(fixed)
? fixed.replace(PAGE_PATTERN, renamePage)
: renamePage(fixed)
if (renamedIds > 0) {
fixes.push(`Renamed ${renamedIds} duplicate ID(s)`)
}
// 9. Fix empty id attributes by generating unique IDs
@@ -1673,6 +1737,11 @@ export function validateAndFixXml(xml: string): {
}
}
/**
* Decode an xmlsvg export (SVG data URL) into uncompressed diagram XML.
* Only the first page is returned; for the full multi-page document use the
* autosaved chartXML instead.
*/
export function extractDiagramXML(xml_svg_string: string): string {
try {
// 1. Parse the SVG string (using built-in DOMParser in a browser-like environment)
+414 -441
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -44,7 +44,7 @@
"@aws-sdk/client-dynamodb": "^3.957.0",
"@aws-sdk/credential-providers": "^3.943.0",
"@extractus/article-extractor": "^8.0.18",
"@formatjs/intl-localematcher": "^0.8.0",
"@formatjs/intl-localematcher": "^0.9.0",
"@langfuse/client": "^4.4.9",
"@langfuse/otel": "^4.4.4",
"@langfuse/tracing": "^4.4.9",
+38 -19
View File
@@ -12,6 +12,7 @@
"@modelcontextprotocol/sdk": "^1.0.4",
"linkedom": "^0.18.0",
"open": "^11.0.0",
"saxes": "^6.0.0",
"zod": "^4.0.0"
},
"bin": {
@@ -523,9 +524,9 @@
"license": "MIT"
},
"node_modules/@modelcontextprotocol/sdk": {
"version": "1.30.0",
"resolved": "https://registry.npmjs.org/@modelcontextprotocol/sdk/-/sdk-1.30.0.tgz",
"integrity": "sha512-xKd8OIzlqNzcqcNumGAa6g+PW2kjD5vrpcKOnfldAUPP3j7lnqMPwlTXQm8gF+UwH72z0lqaRbjr9hqGz0eITA==",
"version": "1.31.0",
"resolved": "https://registry.npmjs.org/@modelcontextprotocol/sdk/-/sdk-1.31.0.tgz",
"integrity": "sha512-UvTMgnNlnIBO/22ob2RcVGDlcvOslQs8T59+FTGdA0L27a39fdGF/EDETNtDVK4DZGpwomlsYpRdA8UXcVL/pw==",
"license": "MIT",
"dependencies": {
"@hono/node-server": "^1.19.9 || ^2.0.5",
@@ -899,13 +900,13 @@
"license": "MIT"
},
"node_modules/@types/node": {
"version": "24.13.3",
"resolved": "https://registry.npmjs.org/@types/node/-/node-24.13.3.tgz",
"integrity": "sha512-Dh8vAsV36ig5wa9OX4pXvMc9D3Veibfw2wix0CUwYODLD8nkj9UsLjASr49nPg+2eKzxhBV+v7L8pXvT4e639Q==",
"version": "24.19.1",
"resolved": "https://registry.npmjs.org/@types/node/-/node-24.19.1.tgz",
"integrity": "sha512-aS3/DG0oM05K0RIXXP+hKjinGG5IgSSVGzswZxW3O0sS3pH4/fycXundUC9XsszgKCk4gHXylTEK6hyFxVxnoQ==",
"dev": true,
"license": "MIT",
"dependencies": {
"undici-types": "~7.18.0"
"undici-types": ">=7.24.0 <7.24.7"
}
},
"node_modules/@vitest/expect": {
@@ -2569,9 +2570,9 @@
}
},
"node_modules/open": {
"version": "11.0.2",
"resolved": "https://registry.npmjs.org/open/-/open-11.0.2.tgz",
"integrity": "sha512-RWqF+pBSkqecEvCKOn8QYhaNdRMJDZRIrlS/7rTDdLHaPcfXGCZ/h8zb413NfvdeAV0MR7T1yJcA34/q+CSm1Q==",
"version": "11.0.4",
"resolved": "https://registry.npmjs.org/open/-/open-11.0.4.tgz",
"integrity": "sha512-++Zlftm0kVLPmzC06t6epuWmcRMDbI4z5P3NNX979WA/k23+NtSOynEGzsVfZwguKw2mi5umVgnBlJQMwRz4Pg==",
"license": "MIT",
"dependencies": {
"default-browser": "^5.5.1",
@@ -2834,6 +2835,18 @@
"integrity": "sha512-YZo3K82SD7Riyi0E1EQPojLz7kpepnSQI9IyPbHHg1XXXevb5dJI7tpyN2ADxGcQbHG7vcyRHk0cbwqcQriUtg==",
"license": "MIT"
},
"node_modules/saxes": {
"version": "6.0.0",
"resolved": "https://registry.npmjs.org/saxes/-/saxes-6.0.0.tgz",
"integrity": "sha512-xAg7SOnEhrm5zI3puOOKyy1OMcMlIJZYNJY7xLBwSze0UjhPLnWfj2GF2EpT0jmzaJKIWKHLsaSSajf35bcYnA==",
"license": "ISC",
"dependencies": {
"xmlchars": "^2.2.0"
},
"engines": {
"node": ">=v12.22.7"
}
},
"node_modules/send": {
"version": "1.2.1",
"resolved": "https://registry.npmjs.org/send/-/send-1.2.1.tgz",
@@ -3080,9 +3093,9 @@
"optional": true
},
"node_modules/tsx": {
"version": "4.23.13",
"resolved": "https://registry.npmjs.org/tsx/-/tsx-4.23.13.tgz",
"integrity": "sha512-BL5MGkRln6aDYhb0xbQlEAGw743BaZYWdbWtdJOBriYJboKgUUYCadFp2/FpBBZquBC/ezNBn7wMMPx7FDZUDw==",
"version": "4.23.15",
"resolved": "https://registry.npmjs.org/tsx/-/tsx-4.23.15.tgz",
"integrity": "sha512-Yiex1Ovn8z2xPpOWckIiysV1SSyRMY9BkLF++q0yKiDxCqRhosKfMg3janKkiLBwZ5c/YryloKwGZcrEmtwxKw==",
"dev": true,
"license": "MIT",
"dependencies": {
@@ -3133,9 +3146,9 @@
"license": "ISC"
},
"node_modules/undici-types": {
"version": "7.18.2",
"resolved": "https://registry.npmjs.org/undici-types/-/undici-types-7.18.2.tgz",
"integrity": "sha512-AsuCzffGHJybSaRrmr5eHr81mwJU3kjw6M+uprWvCXiNeN9SOGwQ3Jn8jb8m3Z6izVgknn1R0FTCEAP2QrLY/w==",
"version": "7.24.6",
"resolved": "https://registry.npmjs.org/undici-types/-/undici-types-7.24.6.tgz",
"integrity": "sha512-WRNW+sJgj5OBN4/0JpHFqtqzhpbnV0GuB+OozA9gCL7a993SmU+1JBZCzLNxYsbMfIeDL+lTsphD5jN5N+n0zg==",
"dev": true,
"license": "MIT"
},
@@ -3379,10 +3392,16 @@
"url": "https://github.com/sponsors/sindresorhus"
}
},
"node_modules/xmlchars": {
"version": "2.2.0",
"resolved": "https://registry.npmjs.org/xmlchars/-/xmlchars-2.2.0.tgz",
"integrity": "sha512-JZnDKK8B0RCDw84FNdDAIpZK+JuJw+s7Lz8nksI7SIuU3UXJJslUthsi+uWBUYOwPFwW7W7PRLRfUKpxjtjFCw==",
"license": "MIT"
},
"node_modules/zod": {
"version": "4.5.4",
"resolved": "https://registry.npmjs.org/zod/-/zod-4.5.4.tgz",
"integrity": "sha512-sC95tT5iHHH9gtpj6A81kh+NEaRAUFN+qlUPDUbRfOMvNf5QCBqsb3WgvnpVtK5Y+4UfA6KqufotuTvMGiTlsA==",
"version": "4.6.5",
"resolved": "https://registry.npmjs.org/zod/-/zod-4.6.5.tgz",
"integrity": "sha512-v5l/aFXZQeai4awLbOpSoHecE9UiMrnfx75tEXLjNonXVARxQ5mOeipTjROUchszUNCqnE+hqAMujRsRHsut2Q==",
"license": "MIT",
"funding": {
"url": "https://github.com/sponsors/colinhacks"
+1
View File
@@ -41,6 +41,7 @@
"@modelcontextprotocol/sdk": "^1.0.4",
"linkedom": "^0.18.0",
"open": "^11.0.0",
"saxes": "^6.0.0",
"zod": "^4.0.0"
},
"devDependencies": {
+53 -31
View File
@@ -7,6 +7,8 @@
* first page is targeted (the "active page by convention" — see pages.ts).
*/
import { getXmlSyntaxError } from "./dom.js"
import { log } from "./logger.js"
import { findPageElement, hasPageSelector, type PageSelector } from "./pages.js"
export interface DiagramOperation {
@@ -26,6 +28,18 @@ export interface ApplyOperationsResult {
errors: OperationError[]
}
// Cells with links, tooltips or custom data are stored as
// <UserObject id="..."><mxCell .../></UserObject> (or <object>): the id sits
// on the wrapper, so the wrapper is treated as the cell.
const CELL_SELECTOR = "mxCell, UserObject, object"
/** Read parent/source/target, which a wrapped cell keeps on its inner mxCell. */
function cellAttr(cell: Element, name: string): string | null {
const inner =
cell.tagName === "mxCell" ? cell : cell.querySelector("mxCell")
return inner?.getAttribute(name) ?? null
}
/**
* Apply diagram operations (update/add/delete) using ID-based lookup.
*
@@ -43,12 +57,8 @@ export function applyDiagramOperations(
): ApplyOperationsResult {
const errors: OperationError[] = []
// Parse the XML
const parser = new DOMParser()
const doc = parser.parseFromString(xmlContent, "text/xml")
// Check for parse errors
const parseError = doc.querySelector("parsererror")
// Check for syntax errors, then parse the XML
const parseError = getXmlSyntaxError(xmlContent)
if (parseError) {
return {
result: xmlContent,
@@ -56,11 +66,13 @@ export function applyDiagramOperations(
{
type: "update",
cellId: "",
message: `XML parse error: ${parseError.textContent}`,
message: `XML parse error: ${parseError}`,
},
],
}
}
const parser = new DOMParser()
const doc = parser.parseFromString(xmlContent, "text/xml")
// Locate the <root> element to operate on.
//
@@ -132,10 +144,12 @@ export function applyDiagramOperations(
// Build a map of cell IDs to elements (scoped to the resolved page).
const cellMap = new Map<string, Element>()
root.querySelectorAll("mxCell").forEach((cell) => {
root.querySelectorAll(CELL_SELECTOR).forEach((cell) => {
const id = cell.getAttribute("id")
if (id) cellMap.set(id, cell)
})
// Ids deleted so far in this batch; deleting one again is a no-op
const deletedIds = new Set<string>()
// Process each operation
for (const op of operations) {
@@ -164,7 +178,7 @@ export function applyDiagramOperations(
`<wrapper>${op.new_xml}</wrapper>`,
"text/xml",
)
const newCell = newDoc.querySelector("mxCell")
const newCell = newDoc.querySelector(CELL_SELECTOR)
if (!newCell) {
errors.push({
type: "update",
@@ -216,7 +230,7 @@ export function applyDiagramOperations(
`<wrapper>${op.new_xml}</wrapper>`,
"text/xml",
)
const newCell = newDoc.querySelector("mxCell")
const newCell = newDoc.querySelector(CELL_SELECTOR)
if (!newCell) {
errors.push({
type: "add",
@@ -256,8 +270,15 @@ export function applyDiagramOperations(
const existingCell = cellMap.get(op.cell_id)
if (!existingCell) {
// Cell not found - might have been cascade-deleted by a previous operation
// Skip silently instead of erroring (AI may redundantly list children/edges)
// Skip cells already cascade-deleted by a previous operation
// (AI may redundantly list children/edges); warn otherwise
if (!deletedIds.has(op.cell_id)) {
errors.push({
type: "delete",
cellId: op.cell_id,
message: `Cell with id="${op.cell_id}" not found`,
})
}
continue
}
@@ -270,17 +291,17 @@ export function applyDiagramOperations(
cellsToDelete.add(cellId)
// Find children (cells where parent === cellId)
// Scoped to `root` so other pages' cells with the same parent id
// (notably "1") are never touched.
const children = root!.querySelectorAll(
`mxCell[parent="${cellId}"]`,
)
children.forEach((child) => {
const childId = child.getAttribute("id")
if (childId && childId !== "0" && childId !== "1") {
// cellMap only holds this page's cells, so other pages' cells
// with the same parent id (notably "1") are never touched.
for (const [childId, child] of cellMap) {
if (
childId !== "0" &&
childId !== "1" &&
cellAttr(child, "parent") === cellId
) {
collectDescendants(childId)
}
})
}
}
// Collect the target cell and all its descendants
@@ -289,23 +310,23 @@ export function applyDiagramOperations(
// Find edges referencing any of the cells to be deleted
// Also recursively collect children of those edges (e.g., edge labels)
for (const cellId of cellsToDelete) {
const referencingEdges = root.querySelectorAll(
`mxCell[source="${cellId}"], mxCell[target="${cellId}"]`,
)
referencingEdges.forEach((edge) => {
const edgeId = edge.getAttribute("id")
for (const [edgeId, edge] of cellMap) {
// Protect root cells from being added via edge references
if (edgeId && edgeId !== "0" && edgeId !== "1") {
if (edgeId === "0" || edgeId === "1") continue
if (
cellAttr(edge, "source") === cellId ||
cellAttr(edge, "target") === cellId
) {
// Recurse to collect edge's children (like labels)
collectDescendants(edgeId)
}
})
}
}
// Log what will be deleted
// Log what will be deleted (stderr: stdout carries JSON-RPC)
if (cellsToDelete.size > 1) {
console.log(
`[applyDiagramOperations] Cascade delete "${op.cell_id}" → deleting ${cellsToDelete.size} cells: ${Array.from(cellsToDelete).join(", ")}`,
log.debug(
`Cascade delete "${op.cell_id}" → deleting ${cellsToDelete.size} cells: ${Array.from(cellsToDelete).join(", ")}`,
)
}
@@ -316,6 +337,7 @@ export function applyDiagramOperations(
cell.parentNode?.removeChild(cell)
cellMap.delete(cellId)
}
deletedIds.add(cellId)
}
}
}
+89
View File
@@ -0,0 +1,89 @@
/**
* DOM setup for Node.
*
* linkedom gives us a DOM with querySelector, but it is lenient: it never
* reports syntax errors (no <parsererror>), and its serializer writes raw
* newlines inside attribute values, which the browser reads back as spaces.
* saxes, a strict XML parser, checks well-formedness the way draw.io's
* DOMParser will, and serializeXml writes attribute values safely.
*/
import { DOMParser } from "linkedom"
import { SaxesParser } from "saxes"
/**
* Returns the first XML syntax error as "line:column: message", or null if
* the XML is well-formed. Surrounding whitespace is ignored because every
* caller trims before the XML reaches the browser.
*/
export function getXmlSyntaxError(xml: string): string | null {
let error: string | null = null
const parser = new SaxesParser()
parser.on("error", (err) => {
error ??= err.message
})
parser.write(xml.trim()).close()
return error
}
const ESCAPES: Record<string, string> = {
"&": "&amp;",
"<": "&lt;",
">": "&gt;",
'"': "&quot;",
"\t": "&#9;",
"\n": "&#xa;",
"\r": "&#xd;",
}
const escapeChars = (text: string, chars: RegExp) =>
text.replace(chars, (c) => ESCAPES[c])
/**
* Serialize a linkedom node as XML. Attribute values escape tabs and line
* breaks too, so multi-line labels (value="a&#xa;b") survive a round trip.
*/
export function serializeXml(node: Node): string {
switch (node.nodeType) {
case 9: {
// Document
const root = (node as Document).documentElement
return root ? serializeXml(root) : ""
}
case 1: {
// Element
const el = node as Element
let out = `<${el.tagName}`
for (const attr of Array.from(el.attributes)) {
out += ` ${attr.name}="${escapeChars(attr.value, /[&<>"\t\n\r]/g)}"`
}
if (el.childNodes.length === 0) return `${out}/>`
out += ">"
for (const child of Array.from(el.childNodes)) {
out += serializeXml(child)
}
return `${out}</${el.tagName}>`
}
case 3:
// Text
return escapeChars(node.textContent ?? "", /[&<>]/g)
case 4:
// CDATA
return `<![CDATA[${node.textContent ?? ""}]]>`
case 8:
// Comment
return `<!--${node.textContent ?? ""}-->`
default:
return ""
}
}
class XMLSerializerPolyfill {
serializeToString(node: Node): string {
return serializeXml(node)
}
}
/** Install the DOMParser and XMLSerializer globals the XML helpers use. */
export function installDomPolyfill(): void {
;(globalThis as any).DOMParser = DOMParser
;(globalThis as any).XMLSerializer = XMLSerializerPolyfill
}
+15 -9
View File
@@ -6,7 +6,15 @@
import { log } from "./logger.js"
const MAX_HISTORY = 20
const historyStore = new Map<string, Array<{ xml: string; svg: string }>>()
interface HistoryEntry {
id: number // Stable across shifts of the circular buffer
xml: string
svg: string
}
let nextEntryId = 0
const historyStore = new Map<string, HistoryEntry[]>()
export function addHistory(sessionId: string, xml: string, svg = ""): number {
let history = historyStore.get(sessionId)
@@ -21,7 +29,7 @@ export function addHistory(sessionId: string, xml: string, svg = ""): number {
return history.length - 1
}
history.push({ xml, svg })
history.push({ id: nextEntryId++, xml, svg })
// Circular buffer
if (history.length > MAX_HISTORY) {
@@ -32,18 +40,16 @@ export function addHistory(sessionId: string, xml: string, svg = ""): number {
return history.length - 1
}
export function getHistory(
sessionId: string,
): Array<{ xml: string; svg: string }> {
export function getHistory(sessionId: string): HistoryEntry[] {
return historyStore.get(sessionId) || []
}
/** Look up an entry by its id; the array index shifts as old entries drop. */
export function getHistoryEntry(
sessionId: string,
index: number,
): { xml: string; svg: string } | undefined {
const history = historyStore.get(sessionId)
return history?.[index]
id: number,
): HistoryEntry | undefined {
return historyStore.get(sessionId)?.find((entry) => entry.id === id)
}
export function clearHistory(sessionId: string): void {
+162 -53
View File
@@ -12,7 +12,9 @@ function readBody(
res: http.ServerResponse,
cb: (body: string) => void,
): void {
let body = ""
// Decode once at the end: a multi-byte UTF-8 character can be split
// across two chunks.
const chunks: Buffer[] = []
let size = 0
req.on("data", (chunk: Buffer) => {
size += chunk.length
@@ -22,9 +24,9 @@ function readBody(
req.destroy()
return
}
body += chunk
chunks.push(chunk)
})
req.on("end", () => cb(body))
req.on("end", () => cb(Buffer.concat(chunks).toString("utf8")))
}
import {
@@ -62,9 +64,11 @@ function normalizeUrl(url: string): string {
return url.replace(/\/$/, "")
}
function isLikelyMcpSessionId(sessionId: string): boolean {
// Keep this cheap and conservative to avoid creating state for arbitrary IDs.
return sessionId.startsWith("mcp-") && sessionId.length <= 128
// Session ids look like "mcp-<base36 time>-<base36 random>" (start_session).
// Only this charset is accepted, because ids are written into the page's
// HTML and script and into the redirect Location header.
function isValidSessionId(sessionId: string): boolean {
return /^mcp-[a-z0-9-]{1,64}$/.test(sessionId)
}
// Find the most recent active session (for auto-redirect when no sessionId provided)
@@ -80,7 +84,7 @@ function getMostRecentSessionId(): string | null {
function ensureSessionStateInitialized(sessionId: string): void {
if (!sessionId) return
if (!isLikelyMcpSessionId(sessionId)) return
if (!isValidSessionId(sessionId)) return
if (stateStore.has(sessionId)) return
setState(sessionId, DEFAULT_DIAGRAM_XML)
@@ -89,7 +93,11 @@ function ensureSessionStateInitialized(sessionId: string): void {
interface SessionState {
xml: string
version: number
// 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
lastUpdated: Date
lastPolled?: number // Last browser poll; an open tab keeps the session alive
svg?: string // Cached SVG from last browser save
syncRequested?: number // Timestamp when sync requested, cleared when browser responds
exportFormat?: "png" | "svg" // Set by MCP tool to request browser export
@@ -108,13 +116,20 @@ export function getState(sessionId: string): SessionState | undefined {
return stateStore.get(sessionId)
}
export function setState(sessionId: string, xml: string, svg?: string): number {
export function setState(
sessionId: string,
xml: string,
svg?: string,
fromBrowser = false,
): number {
const existing = stateStore.get(sessionId)
const newVersion = (existing?.version || 0) + 1
stateStore.set(sessionId, {
xml,
version: newVersion,
serverVersion: fromBrowser ? existing?.serverVersion : newVersion,
lastUpdated: new Date(),
lastPolled: existing?.lastPolled,
svg: svg || existing?.svg, // Preserve cached SVG if not provided
syncRequested: undefined, // Clear sync request when browser pushes state
exportFormat: existing?.exportFormat, // Preserve pending export request
@@ -222,7 +237,11 @@ export function stopHttpServer(): void {
function cleanupExpiredSessions(): void {
const now = Date.now()
for (const [sessionId, state] of stateStore) {
if (now - state.lastUpdated.getTime() > SESSION_TTL) {
const lastActive = Math.max(
state.lastUpdated.getTime(),
state.lastPolled ?? 0,
)
if (now - lastActive > SESSION_TTL) {
stateStore.delete(sessionId)
clearHistory(sessionId)
log.info(`Cleaned up expired session: ${sessionId}`)
@@ -245,7 +264,48 @@ function handleRequest(
req: http.IncomingMessage,
res: http.ServerResponse,
): void {
const url = new URL(req.url || "/", `http://localhost:${serverPort}`)
// A bad request must never take down the MCP process
try {
routeRequest(req, res)
} catch (err) {
log.error("HTTP request failed:", err)
if (!res.headersSent) res.writeHead(500)
res.end()
}
}
// Serve only requests addressed to localhost, sent by a localhost page or by
// a non-browser client (no Origin header). This blocks DNS rebinding and
// scripts on other websites.
function isLocalRequest(req: http.IncomingMessage): boolean {
const isLocalHost = (host: string) =>
/^(localhost|127\.0\.0\.1)(:\d+)?$/.test(host)
const origin = req.headers.origin
return (
isLocalHost(req.headers.host ?? "") &&
(origin === undefined || isLocalHost(origin.replace(/^http:\/\//, "")))
)
}
function routeRequest(
req: http.IncomingMessage,
res: http.ServerResponse,
): void {
let url: URL
try {
url = new URL(req.url || "/", `http://localhost:${serverPort}`)
} catch {
// e.g. "//" is not a valid URL path
res.writeHead(400)
res.end("Bad Request")
return
}
if (!isLocalRequest(req)) {
res.writeHead(403)
res.end("Forbidden")
return
}
const requestOrigin = req.headers.origin
if (requestOrigin === `http://localhost:${serverPort}`) {
@@ -262,12 +322,19 @@ function handleRequest(
if (url.pathname === "/" || url.pathname === "/index.html") {
const sessionId = url.searchParams.get("mcp") || ""
if (sessionId && !isValidSessionId(sessionId)) {
res.writeHead(400)
res.end("Invalid session id")
return
}
// Auto-redirect to most recent session if no sessionId provided
if (!sessionId) {
const recentSessionId = getMostRecentSessionId()
if (recentSessionId) {
res.writeHead(302, { Location: `/?mcp=${recentSessionId}` })
res.writeHead(302, {
Location: `/?mcp=${encodeURIComponent(recentSessionId)}`,
})
res.end()
return
}
@@ -305,6 +372,9 @@ function handleStateApi(
}
ensureSessionStateInitialized(sessionId)
const state = stateStore.get(sessionId)
// Polling counts as activity, so a session stays alive while its
// tab is open
if (state) state.lastPolled = Date.now()
res.writeHead(200, { "Content-Type": "application/json" })
res.end(
JSON.stringify({
@@ -320,9 +390,11 @@ function handleStateApi(
try {
const data = JSON.parse(body)
const { sessionId } = data
if (!sessionId) {
if (!sessionId || !isValidSessionId(sessionId)) {
res.writeHead(400, { "Content-Type": "application/json" })
res.end(JSON.stringify({ error: "sessionId required" }))
res.end(
JSON.stringify({ error: "valid sessionId required" }),
)
return
}
@@ -342,7 +414,25 @@ function handleStateApi(
return
}
const version = setState(sessionId, data.xml, data.svg)
// 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.
const current = stateStore.get(sessionId)
if (
typeof data.baseVersion === "number" &&
data.baseVersion < (current?.serverVersion ?? 0)
) {
res.writeHead(409, { "Content-Type": "application/json" })
res.end(
JSON.stringify({
error: "Diagram changed on the server",
version: current?.version,
}),
)
return
}
const version = setState(sessionId, data.xml, data.svg, true)
res.writeHead(200, { "Content-Type": "application/json" })
res.end(JSON.stringify({ success: true, version }))
} catch {
@@ -378,7 +468,11 @@ function handleHistoryApi(
res.writeHead(200, { "Content-Type": "application/json" })
res.end(
JSON.stringify({
entries: history.map((entry, i) => ({ index: i, svg: entry.svg })),
entries: history.map((entry, i) => ({
index: i,
id: entry.id,
svg: entry.svg,
})),
count: history.length,
}),
)
@@ -396,16 +490,14 @@ function handleRestoreApi(
readBody(req, res, (body) => {
try {
const { sessionId, index } = JSON.parse(body)
if (!sessionId || index === undefined) {
const { sessionId, id } = JSON.parse(body)
if (!sessionId || typeof id !== "number") {
res.writeHead(400, { "Content-Type": "application/json" })
res.end(
JSON.stringify({ error: "sessionId and index required" }),
)
res.end(JSON.stringify({ error: "sessionId and id required" }))
return
}
const entry = getHistoryEntry(sessionId, index)
const entry = getHistoryEntry(sessionId, id)
if (!entry) {
res.writeHead(404, { "Content-Type": "application/json" })
res.end(JSON.stringify({ error: "Entry not found" }))
@@ -415,7 +507,7 @@ function handleRestoreApi(
const newVersion = setState(sessionId, entry.xml)
addHistory(sessionId, entry.xml, entry.svg)
log.info(`Restored session ${sessionId} to index ${index}`)
log.info(`Restored session ${sessionId} to history entry ${id}`)
res.writeHead(200, { "Content-Type": "application/json" })
res.end(JSON.stringify({ success: true, newVersion }))
@@ -697,10 +789,11 @@ function getHtmlPage(sessionId: string): string {
</div>
</div>
<script>
const sessionId = "${sessionId}";
const sessionId = ${JSON.stringify(sessionId).replace(/</g, "\\u003c")};
const iframe = document.getElementById('drawio');
let currentVersion = 0, isReady = false, pendingXml = null, lastXml = null;
let pendingSvgExport = null;
let pendingSvgBase = 0; // version the pending autosave was based on
let pendingAiSvg = false;
let pendingMcpExport = null; // 'png' or 'svg' when MCP requested export
let projectionExportActive = false; // page-targeted export: showing a transient single-page projection
@@ -718,18 +811,29 @@ function getHtmlPage(sessionId: string): string {
// for a page-targeted export — otherwise we'd push the
// transient projection back as the canonical session state.
if (projectionExportActive) return;
// Request SVG export, then push state with SVG
// 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.
pendingSvgExport = msg.xml;
pendingSvgBase = currentVersion;
iframe.contentWindow.postMessage(JSON.stringify({ action: 'export', format: 'svg' }), '*');
// Fallback if export doesn't respond
setTimeout(() => { if (pendingSvgExport === msg.xml) { pushState(msg.xml, ''); pendingSvgExport = null; } }, 2000);
setTimeout(() => { if (pendingSvgExport === msg.xml) { pushState(msg.xml, '', pendingSvgBase); 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.
if (pendingSyncExport && msg.xml) {
pendingSyncExport = false;
pushState(msg.xml, '');
}
} else if (msg.event === 'export' && msg.data) {
// Handle MCP server export request (png/svg)
// Verify the response matches the requested format to avoid capturing
// unrelated exports (autosave SVG, sync XML)
if (pendingMcpExport) {
// Handle MCP server export request (png/svg). fireExport tags
// the request with mcpExport and draw.io echoes the request
// back in msg.message, which tells it apart from autosave and
// preview SVG exports.
if (msg.message && msg.message.mcpExport) {
const d = msg.data;
const isPng = pendingMcpExport === 'png' && (d.startsWith('data:image/png') || (typeof d === 'string' && d.length > 100 && !d.startsWith('<')));
const isPng = pendingMcpExport === 'png' && d.startsWith('data:image/png');
const isSvg = pendingMcpExport === 'svg' && (d.startsWith('data:image/svg') || d.startsWith('<svg'));
if (isPng || isSvg) {
pendingMcpExport = null;
@@ -741,8 +845,8 @@ function getHtmlPage(sessionId: string): string {
// Page-targeted export: restore the user's real
// multi-page document now that we have the image.
restoreFromProjection();
return;
}
return;
}
// Handle file download export (PNG/SVG only, drawio uses lastXml directly)
if (pendingDownload && (pendingDownload.format === 'png' || pendingDownload.format === 'svg')) {
@@ -761,19 +865,13 @@ function getHtmlPage(sessionId: string): string {
saveConfirmBtn.textContent = 'Save';
return;
}
// Handle sync export (XML format) - server requested fresh state
if (pendingSyncExport && !msg.data.startsWith('data:') && !msg.data.startsWith('<svg')) {
pendingSyncExport = false;
pushState(msg.data, '');
return;
}
// 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);
pushState(xml, svg, pendingSvgBase);
} else if (pendingAiSvg) {
pendingAiSvg = false;
fetch('/api/history-svg', {
@@ -814,15 +912,17 @@ function getHtmlPage(sessionId: string): string {
}
}
async function pushState(xml, svg = '') {
async function pushState(xml, svg = '', baseVersion = currentVersion) {
if (!sessionId) return;
try {
const r = await fetch('/api/state', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ sessionId, xml, svg })
body: JSON.stringify({ sessionId, xml, svg, baseVersion })
});
if (r.ok) { const d = await r.json(); currentVersion = d.version; lastXml = xml; }
// 409: the AI wrote a newer version; load it now
else if (r.status === 409) poll();
} catch (e) { console.error('Push failed:', e); }
}
@@ -830,14 +930,22 @@ function getHtmlPage(sessionId: string): string {
async function poll() {
if (!sessionId) return;
const knownVersion = currentVersion;
try {
const r = await fetch('/api/state?sessionId=' + encodeURIComponent(sessionId));
if (!r.ok) return;
const s = await r.json();
// Handle sync request - server needs fresh state
if (s.syncRequested && !pendingSyncExport) {
// Handle sync request - server needs fresh state. Reset after a
// while in case draw.io never answers, so later syncs still run.
if (s.syncRequested && !pendingSyncExport && isReady) {
pendingSyncExport = true;
iframe.contentWindow.postMessage(JSON.stringify({ action: 'export', format: 'xml' }), '*');
setTimeout(() => { pendingSyncExport = false; }, 5000);
}
// The server lost this session (e.g. it expired) and rebuilt it
// with a blank diagram: push back what the browser shows.
if (s.version < knownVersion && lastXml) {
pushState(lastXml);
}
// Load new diagram from server (before export, so we export latest).
// While a page-targeted projection is on screen, skip the reload
@@ -862,9 +970,10 @@ function getHtmlPage(sessionId: string): string {
if (s.exportFormat && !pendingMcpExport && isReady) {
pendingMcpExport = s.exportFormat;
const fireExport = () => {
// mcpExport is echoed back in msg.message (see the handler)
const exportOpts = pendingMcpExport === 'png'
? { action: 'export', format: 'png', scale: 2 }
: { action: 'export', format: 'svg' };
? { action: 'export', format: 'png', scale: 2, mcpExport: true }
: { action: 'export', format: 'svg', mcpExport: true };
iframe.contentWindow.postMessage(JSON.stringify(exportOpts), '*');
};
if (s.exportXml) {
@@ -962,7 +1071,7 @@ function getHtmlPage(sessionId: string): string {
const historyEmpty = document.getElementById('history-empty');
const restoreBtn = document.getElementById('restore-btn');
const cancelBtn = document.getElementById('cancel-btn');
let historyData = [], selectedIdx = null;
let historyData = [], selectedId = null;
historyBtn.onclick = async () => {
if (!sessionId) return;
@@ -977,7 +1086,7 @@ function getHtmlPage(sessionId: string): string {
historyModal.classList.add('open');
};
cancelBtn.onclick = () => { historyModal.classList.remove('open'); selectedIdx = null; restoreBtn.disabled = true; };
cancelBtn.onclick = () => { historyModal.classList.remove('open'); selectedId = null; restoreBtn.disabled = true; };
historyModal.onclick = (e) => { if (e.target === historyModal) cancelBtn.onclick(); };
function renderHistory() {
@@ -989,30 +1098,30 @@ function getHtmlPage(sessionId: string): string {
historyGrid.style.display = 'grid';
historyEmpty.style.display = 'none';
historyGrid.innerHTML = historyData.map((e, i) => \`
<div class="history-item" data-idx="\${e.index}">
<div class="history-item" data-id="\${e.id}">
<div class="thumb">\${e.svg ? \`<img src="\${e.svg}">\` : '#' + e.index}</div>
<div class="label">#\${e.index}</div>
</div>
\`).join('');
historyGrid.querySelectorAll('.history-item').forEach(item => {
item.onclick = () => {
const idx = parseInt(item.dataset.idx);
if (selectedIdx === idx) { selectedIdx = null; restoreBtn.disabled = true; }
else { selectedIdx = idx; restoreBtn.disabled = false; }
historyGrid.querySelectorAll('.history-item').forEach(el => el.classList.toggle('selected', parseInt(el.dataset.idx) === selectedIdx));
const id = parseInt(item.dataset.id);
if (selectedId === id) { selectedId = null; restoreBtn.disabled = true; }
else { selectedId = id; restoreBtn.disabled = false; }
historyGrid.querySelectorAll('.history-item').forEach(el => el.classList.toggle('selected', parseInt(el.dataset.id) === selectedId));
};
});
}
restoreBtn.onclick = async () => {
if (selectedIdx === null) return;
if (selectedId === null) return;
restoreBtn.disabled = true;
restoreBtn.textContent = 'Restoring...';
try {
const r = await fetch('/api/restore', {
method: 'POST',
headers: { 'Content-Type': 'application/json' },
body: JSON.stringify({ sessionId, index: selectedIdx })
body: JSON.stringify({ sessionId, id: selectedId })
});
if (r.ok) { cancelBtn.onclick(); await poll(); }
else { alert('Restore failed'); }
+55 -45
View File
@@ -18,24 +18,6 @@
* surface.
*/
// Setup DOM polyfill for Node.js (required for XML operations)
import { DOMParser } from "linkedom"
;(globalThis as any).DOMParser = DOMParser
// Create XMLSerializer polyfill using outerHTML
class XMLSerializerPolyfill {
serializeToString(node: any): string {
if (node.outerHTML !== undefined) {
return node.outerHTML
}
if (node.documentElement) {
return node.documentElement.outerHTML
}
return ""
}
}
;(globalThis as any).XMLSerializer = XMLSerializerPolyfill
import { createRequire } from "node:module"
import { McpServer } from "@modelcontextprotocol/sdk/server/mcp.js"
import { StdioServerTransport } from "@modelcontextprotocol/sdk/server/stdio.js"
@@ -45,6 +27,7 @@ import {
applyDiagramOperations,
type DiagramOperation,
} from "./diagram-operations.js"
import { installDomPolyfill } from "./dom.js"
import { checkEditGate } from "./edit-gate.js"
import { addHistory } from "./history.js"
import {
@@ -72,6 +55,9 @@ import {
} from "./pages.js"
import { validateAndFixXml } from "./xml-validation.js"
// DOMParser/XMLSerializer globals for the XML helpers (Node has neither)
installDomPolyfill()
// Server configuration
const config = {
port: parseInt(process.env.PORT || "6002", 10),
@@ -908,6 +894,47 @@ server.registerTool(
},
)
// The browser bridge has one export slot per session, so export requests
// run one at a time: a concurrent call waits for the previous one.
let exportQueue: Promise<unknown> = Promise.resolve()
/**
* Ask the browser to export (optionally via a page projection) and poll for
* the resulting image data. Resolves to undefined on timeout.
*/
function exportViaBrowser(
sessionId: string,
format: "png" | "svg",
projectionXml?: string,
): Promise<string | undefined> {
const run = exportQueue.then(async () => {
requestExport(sessionId, format, projectionXml)
// A projection export does an extra load + render round-trip in the
// browser, so give it a longer window. Re-read the live store entry
// each tick: setState() (from a concurrent autosave or tool call)
// replaces the Map entry with a new object, so a captured reference
// would go stale and never observe the browser's exportData.
const timeoutMs = projectionXml ? 15000 : 10000
const start = Date.now()
let exportData: string | undefined
while (Date.now() - start < timeoutMs) {
exportData = getState(sessionId)?.exportData
if (exportData) break
await new Promise((r) => setTimeout(r, 200))
}
const live = getState(sessionId)
if (live) {
live.exportData = undefined
live.exportFormat = undefined
live.exportXml = undefined
}
return exportData
})
exportQueue = run.catch(() => {})
return run
}
// Tool: export_diagram
server.registerTool(
"export_diagram",
@@ -1079,34 +1106,12 @@ server.registerTool(
projectionXml = projection.xml
}
// Ask the browser to export (optionally via a page projection) and
// poll for the resulting image data.
requestExport(
const exportData = await exportViaBrowser(
currentSession.id,
detectedFormat as "png" | "svg",
projectionXml,
)
// A projection export does an extra load + render round-trip in the
// browser, so give it a longer window. Re-read the live store entry
// each tick: setState() (from a concurrent autosave or tool call)
// replaces the Map entry with a new object, so a captured reference
// would go stale and never observe the browser's exportData.
const timeoutMs = projectionXml ? 15000 : 10000
const start = Date.now()
let exportData: string | undefined
while (Date.now() - start < timeoutMs) {
exportData = getState(currentSession.id)?.exportData
if (exportData) break
await new Promise((r) => setTimeout(r, 200))
}
const live = getState(currentSession.id)
if (live) {
live.exportData = undefined
live.exportFormat = undefined
live.exportXml = undefined
}
if (!exportData) {
return {
content: [
@@ -1215,15 +1220,20 @@ async function loadMxfileForMutation(): Promise<
doc,
writeBack: (newDoc: Document) => {
const newXml = serializeMxfile(newDoc)
// The store may hold user edits the model has not seen yet.
const sawLatest = checkEditGate(
sessionRef.lastSeenXml,
browserState?.xml ?? "",
).ok
// Save history before overwriting so the user can undo.
addHistory(sessionRef.id, sessionRef.xml, browserState?.svg || "")
sessionRef.xml = newXml
sessionRef.version++
setState(sessionRef.id, newXml)
// The model just wrote this exact state, so mark it as seen —
// subsequent edit_diagram calls don't need a redundant
// get_diagram round-trip.
sessionRef.lastSeenXml = newXml
// The model just wrote this exact state. If it had seen the state
// it built on, mark the result as seen so edit_diagram needs no
// extra get_diagram; otherwise edit_diagram must ask for one.
sessionRef.lastSeenXml = sawLatest ? newXml : ""
addHistory(sessionRef.id, newXml, "")
},
}
+2 -1
View File
@@ -9,6 +9,7 @@
*/
import { inflateRawSync } from "node:zlib"
import { DOMParser } from "linkedom"
import { getXmlSyntaxError } from "./dom.js"
import {
isMxFile,
isMxGraphModel,
@@ -82,7 +83,7 @@ export function parseDrawioFileContent(content: string): LoadResult {
}
const inner = new DOMParser().parseFromString(xml, "text/xml")
if (
inner.querySelector("parsererror") ||
getXmlSyntaxError(xml) ||
inner.documentElement?.tagName !== "mxGraphModel"
) {
return {
+4 -3
View File
@@ -18,6 +18,7 @@
*/
import { DOMParser } from "linkedom"
import { getXmlSyntaxError } from "./dom.js"
export interface PageInfo {
id: string
@@ -110,8 +111,8 @@ export function normalizeToMxfile(
*/
export function parseMxfile(xml: string): Document | null {
try {
if (getXmlSyntaxError(xml)) return null
const doc = new DOMParser().parseFromString(xml, "text/xml")
if (doc.querySelector("parsererror")) return null
if (doc.documentElement?.tagName !== "mxfile") return null
return doc as unknown as Document
} catch {
@@ -258,12 +259,12 @@ export function addPageToDoc(
}
const snippet = `<wrapper><diagram id="${escapeAttr(id)}" name="${escapeAttr(name)}">${inner}</diagram></wrapper>`
const tempDoc = new DOMParser().parseFromString(snippet, "text/xml")
if (tempDoc.querySelector("parsererror")) {
if (getXmlSyntaxError(snippet)) {
throw new Error(
"Failed to parse new page xml — make sure it is a valid <mxGraphModel>",
)
}
const tempDoc = new DOMParser().parseFromString(snippet, "text/xml")
const newDiagram = tempDoc.querySelector("diagram")
if (!newDiagram) {
throw new Error("Failed to construct <diagram> element for new page")
+109 -109
View File
@@ -3,6 +3,8 @@
* Copied from lib/utils.ts to avoid cross-package imports
*/
import { getXmlSyntaxError } from "./dom.js"
// ============================================================================
// Constants
// ============================================================================
@@ -10,9 +12,6 @@
/** Maximum XML size to process (1MB) - larger XMLs may cause performance issues */
const MAX_XML_SIZE = 1_000_000
/** Maximum iterations for aggressive cell dropping to prevent infinite loops */
const MAX_DROP_ITERATIONS = 10
/** Structural attributes that should not be duplicated in draw.io */
const STRUCTURAL_ATTRS = [
"edge",
@@ -91,6 +90,21 @@ function parseXmlTags(xml: string): ParsedTag[] {
return tags
}
/** Rewrite every opening tag with fn, leaving text and closing tags as is. */
function replaceInOpeningTags(
xml: string,
fn: (tag: string) => string,
): string {
let out = ""
let last = 0
for (const { tag, isClosing, startIndex, endIndex } of parseXmlTags(xml)) {
if (isClosing) continue
out += xml.slice(last, startIndex) + fn(tag)
last = endIndex + 1
}
return out + xml.slice(last)
}
// ============================================================================
// Validation Helper Functions
// ============================================================================
@@ -128,8 +142,7 @@ function checkDuplicateAttributes(xml: string): string | null {
* scope the cell-ID uniqueness check per <diagram>, and additionally check
* that the <diagram> ids themselves are unique.
*
* The legacy regex-based check is kept as a fallback for non-mxfile inputs
* and for XML that won't DOM-parse.
* The legacy regex-based check is kept as a fallback for non-mxfile inputs.
*/
function checkDuplicateIds(xml: string): string | null {
// The DOM-aware path only matters for <mxfile> wrappers; for legacy
@@ -142,51 +155,47 @@ function checkDuplicateIds(xml: string): string | null {
if (mightBeMxFile)
try {
const doc = new DOMParser().parseFromString(xml, "text/xml")
if (!doc.querySelector("parsererror")) {
const rootEl = doc.documentElement
if (rootEl && rootEl.tagName === "mxfile") {
const diagrams = doc.querySelectorAll("diagram")
const rootEl = doc.documentElement
if (rootEl && rootEl.tagName === "mxfile") {
const diagrams = doc.querySelectorAll("diagram")
// 1) <diagram> ids must be unique across the file.
const diagramIds = new Map<string, number>()
diagrams.forEach((d) => {
const id = d.getAttribute("id")
if (id)
diagramIds.set(id, (diagramIds.get(id) || 0) + 1)
})
const dupDiagrams = Array.from(diagramIds.entries())
.filter(([, c]) => c > 1)
.map(([id]) => `'${id}'`)
if (dupDiagrams.length > 0) {
return `Invalid XML: Found duplicate <diagram> id(s): ${dupDiagrams.slice(0, 3).join(", ")}. Each page must have a unique id.`
}
// 2) Within each page, mxCell ids must be unique.
for (let i = 0; i < diagrams.length; i++) {
const diagram = diagrams[i]
const pageId =
diagram.getAttribute("id") || `(index ${i})`
const cells = diagram.querySelectorAll("mxCell")
const cellIds = new Map<string, number>()
cells.forEach((c) => {
const id = c.getAttribute("id")
if (id) cellIds.set(id, (cellIds.get(id) || 0) + 1)
})
const dups = Array.from(cellIds.entries())
.filter(([, c]) => c > 1)
.map(([id, count]) => `'${id}' (${count}x)`)
if (dups.length > 0) {
return `Invalid XML: Found duplicate cell ID(s) in page "${pageId}": ${dups.slice(0, 3).join(", ")}. All mxCell ids must be unique within a page.`
}
}
return null
// 1) <diagram> ids must be unique across the file.
const diagramIds = new Map<string, number>()
diagrams.forEach((d) => {
const id = d.getAttribute("id")
if (id) diagramIds.set(id, (diagramIds.get(id) || 0) + 1)
})
const dupDiagrams = Array.from(diagramIds.entries())
.filter(([, c]) => c > 1)
.map(([id]) => `'${id}'`)
if (dupDiagrams.length > 0) {
return `Invalid XML: Found duplicate <diagram> id(s): ${dupDiagrams.slice(0, 3).join(", ")}. Each page must have a unique id.`
}
// 2) Within each page, mxCell ids must be unique.
for (let i = 0; i < diagrams.length; i++) {
const diagram = diagrams[i]
const pageId = diagram.getAttribute("id") || `(index ${i})`
const cells = diagram.querySelectorAll("mxCell")
const cellIds = new Map<string, number>()
cells.forEach((c) => {
const id = c.getAttribute("id")
if (id) cellIds.set(id, (cellIds.get(id) || 0) + 1)
})
const dups = Array.from(cellIds.entries())
.filter(([, c]) => c > 1)
.map(([id, count]) => `'${id}' (${count}x)`)
if (dups.length > 0) {
return `Invalid XML: Found duplicate cell ID(s) in page "${pageId}": ${dups.slice(0, 3).join(", ")}. All mxCell ids must be unique within a page.`
}
}
return null
}
} catch {
// fall through to regex
}
// Legacy regex-based check for bare <mxGraphModel> and parse-error cases.
// Legacy regex-based check for bare <mxGraphModel> inputs.
const idPattern = /\bid\s*=\s*["']([^"']+)["']/gi
const ids = new Map<string, number>()
let idMatch
@@ -315,14 +324,11 @@ export function validateMxCellStructure(xml: string): string | null {
)
}
// 0. First use DOM parser to catch syntax errors (most accurate)
// 0. DOM-based checks. Syntax errors are caught by the strict check at
// the end: linkedom's DOMParser never reports them.
try {
const parser = new DOMParser()
const doc = parser.parseFromString(xml, "text/xml")
const parseError = doc.querySelector("parsererror")
if (parseError) {
return `Invalid XML: The XML contains syntax errors (likely unescaped special characters like <, >, & in attribute values). Please escape special characters: use &lt; for <, &gt; for >, &amp; for &, &quot; for ". Regenerate the diagram with properly escaped values.`
}
// DOM-based checks for nested mxCell
const allCells = doc.querySelectorAll("mxCell")
@@ -404,6 +410,14 @@ export function validateMxCellStructure(xml: string): string | null {
return nestedCellError
}
// 11. Strict XML syntax check, run last so the checks above can give
// more specific messages. Catches what they miss, e.g. duplicate or
// unquoted attributes, which make draw.io refuse to load the diagram.
const syntaxError = getXmlSyntaxError(xml)
if (syntaxError) {
return `Invalid XML: syntax error at ${syntaxError} Escape special characters in attribute values (&lt; for <, &amp; for &, &quot; for "), quote every attribute value, and do not repeat an attribute.`
}
return null
}
@@ -494,13 +508,21 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
}
}
// 6. Fix malformed attribute quotes
const malformedQuotePattern = /(\s[a-zA-Z][a-zA-Z0-9_:-]*)=&quot;/
if (malformedQuotePattern.test(fixed)) {
fixed = fixed.replace(
/(\s[a-zA-Z][a-zA-Z0-9_:-]*)=&quot;([^&]*?)&quot;/g,
'$1="$2"',
)
// 6. Fix malformed attribute quotes (name=&quot;value&quot;). Quoted
// values are matched first and kept, so &quot; inside a rich-text
// label like value="&lt;font style=&quot;...&quot;&gt;" is left alone.
let quotesFixed = false
fixed = replaceInOpeningTags(fixed, (tag) =>
tag.replace(
/("[^"]*"|'[^']*')|(\s[a-zA-Z][a-zA-Z0-9_:-]*)=&quot;([^&]*?)&quot;/g,
(match, quoted, name, value) => {
if (quoted) return match
quotesFixed = true
return `${name}="${value}"`
},
),
)
if (quotesFixed) {
fixes.push("Fixed malformed attribute quotes")
}
@@ -511,10 +533,21 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
fixes.push("Fixed malformed closing tags")
}
// 8. Fix missing space between attributes
const missingSpacePattern = /("[^"]*")([a-zA-Z][a-zA-Z0-9_:-]*=)/g
if (missingSpacePattern.test(fixed)) {
fixed = fixed.replace(/("[^"]*")([a-zA-Z][a-zA-Z0-9_:-]*=)/g, "$1 $2")
// 8. Fix missing space between attributes (id="2"vertex="1"). Every
// quoted value is consumed whole, so quotes always pair up within one
// attribute.
let spaceAdded = false
fixed = replaceInOpeningTags(fixed, (tag) =>
tag.replace(
/("[^"]*"|'[^']*')([a-zA-Z_:])?/g,
(match, quoted, next) => {
if (!next) return match
spaceAdded = true
return `${quoted} ${next}`
},
),
)
if (spaceAdded) {
fixes.push("Added missing space between attributes")
}
@@ -632,6 +665,9 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
"Array",
"Object",
"mxRectangle",
// Wrappers draw.io writes for cells with links, tooltips or data
"UserObject",
"object",
])
const foreignTagPattern = /<\/?([a-zA-Z][a-zA-Z0-9_]*)[^>]*>/g
let foreignMatch
@@ -796,8 +832,10 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
fixes.push(`Flattened ${nestedFixed} duplicate-ID nested mxCell(s)`)
}
// 21. Fix true nested mxCell (different IDs)
const lines2 = fixed.split("\n")
// 21. Fix true nested mxCell (different IDs). Runs only when the nesting
// check finds real nesting, because this line-based rewrite can break
// valid cells written over several lines.
const lines2 = checkNestedMxCells(fixed) ? fixed.split("\n") : []
newLines = []
let trueNestedFixed = 0
let cellDepth = 0
@@ -807,7 +845,11 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
const line = lines2[i]
const trimmed = line.trim()
const isOpenCell = /<mxCell\s/.test(trimmed) && !trimmed.endsWith("/>")
// A line holding a whole cell (<mxCell ...>...</mxCell>) opens nothing
const isOpenCell =
/<mxCell\s/.test(trimmed) &&
!trimmed.endsWith("/>") &&
!trimmed.endsWith("</mxCell>")
const isCloseCell = trimmed === "</mxCell>"
if (isOpenCell) {
@@ -860,9 +902,11 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
if (duplicateIds.length > 0) {
const idCounters = new Map<string, number>()
// Rebuild from the captured parts so only the value changes (an id
// like "d" or "i" also occurs in the attribute name itself)
fixed = fixed.replace(
/\bid\s*=\s*["']([^"']+)["']/gi,
(match, id) => {
/(\bid\s*=\s*["'])([^"']+)(["'])/gi,
(match, before, id, after) => {
if (!duplicateIds.includes(id)) return match
const count = idCounters.get(id) || 0
@@ -870,8 +914,7 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
if (count === 0) return match
const newId = `${id}_dup${count}`
return match.replace(id, newId)
return `${before}${id}_dup${count}${after}`
},
)
fixes.push(`Renamed ${duplicateIds.length} duplicate ID(s)`)
@@ -892,49 +935,6 @@ export function autoFixXml(xml: string): { fixed: string; fixes: string[] } {
fixes.push(`Generated ${emptyIdCount} missing ID(s)`)
}
// 24. Aggressive: drop broken mxCell elements
if (typeof DOMParser !== "undefined") {
let droppedCells = 0
let maxIterations = MAX_DROP_ITERATIONS
while (maxIterations-- > 0) {
const parser = new DOMParser()
const doc = parser.parseFromString(fixed, "text/xml")
const parseError = doc.querySelector("parsererror")
if (!parseError) break
const errText = parseError.textContent || ""
const match = errText.match(/(\d+):\d+:/)
if (!match) break
const errLine = parseInt(match[1], 10) - 1
const lines = fixed.split("\n")
let cellStart = errLine
let cellEnd = errLine
while (cellStart > 0 && !lines[cellStart].includes("<mxCell")) {
cellStart--
}
while (cellEnd < lines.length - 1) {
if (
lines[cellEnd].includes("</mxCell>") ||
lines[cellEnd].trim().endsWith("/>")
) {
break
}
cellEnd++
}
lines.splice(cellStart, cellEnd - cellStart + 1)
fixed = lines.join("\n")
droppedCells++
}
if (droppedCells > 0) {
fixes.push(`Dropped ${droppedCells} unfixable mxCell element(s)`)
}
}
return { fixed, fixes }
}
@@ -0,0 +1,85 @@
/**
* Tests for edit_diagram operations on cells that draw.io wraps in
* <UserObject> or <object> (cells with links, tooltips or custom data).
* The id sits on the wrapper; the inner mxCell has none.
*/
import { beforeAll, describe, expect, it, vi } from "vitest"
import { installDomPolyfill } from "../src/dom.js"
beforeAll(() => {
installDomPolyfill()
})
import { applyDiagramOperations } from "../src/diagram-operations.js"
const DOC = `<mxfile><diagram id="p" name="Page-1"><mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><UserObject id="a" label="A" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject><mxCell id="b" value="B" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><object id="e1" label="" tooltip="t"><mxCell edge="1" source="b" target="a" parent="1"><mxGeometry relative="1" as="geometry"/></mxCell></object><mxCell id="child" value="C" vertex="1" parent="a"><mxGeometry as="geometry"/></mxCell></root></mxGraphModel></diagram></mxfile>`
describe("wrapped cells", () => {
it("deletes a UserObject cell with its edges and children", () => {
const { result, errors } = applyDiagramOperations(DOC, [
{ operation: "delete", cell_id: "a" },
])
expect(errors).toEqual([])
expect(result).not.toContain('id="a"')
expect(result).not.toContain('id="e1"')
expect(result).not.toContain('id="child"')
expect(result).toContain('id="b"')
})
it("cascades to a wrapped edge when deleting a plain cell", () => {
const { result, errors } = applyDiagramOperations(DOC, [
{ operation: "delete", cell_id: "b" },
{ operation: "delete", cell_id: "e1" },
])
// e1 was already removed by the cascade, so no warning for it
expect(errors).toEqual([])
expect(result).not.toContain('id="e1"')
expect(result).toContain('id="a"')
})
it("warns when deleting a cell that does not exist", () => {
const { errors } = applyDiagramOperations(DOC, [
{ operation: "delete", cell_id: "missing" },
])
expect(errors).toHaveLength(1)
expect(errors[0]).toMatchObject({ type: "delete", cellId: "missing" })
})
it("updates a UserObject cell", () => {
const { result, errors } = applyDiagramOperations(DOC, [
{
operation: "update",
cell_id: "a",
new_xml: `<UserObject id="a" label="A2" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject>`,
},
])
expect(errors).toEqual([])
expect(result).toContain('label="A2"')
expect(result.match(/id="a"/g)).toHaveLength(1)
})
it("refuses to add a cell whose id a UserObject already uses", () => {
const { errors } = applyDiagramOperations(DOC, [
{
operation: "add",
cell_id: "a",
new_xml: `<mxCell id="a" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell>`,
},
])
expect(errors[0]?.message).toContain("already exists")
})
})
describe("cascade delete logging", () => {
it("does not write cascade logs to stdout (the JSON-RPC channel)", () => {
const plain = `<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="x" vertex="1" parent="1"/><mxCell id="y" vertex="1" parent="1"/><mxCell id="e" edge="1" source="x" target="y" parent="1"/></root></mxGraphModel>`
const spy = vi.spyOn(console, "log").mockImplementation(() => {})
const { result } = applyDiagramOperations(plain, [
{ operation: "delete", cell_id: "x" },
])
expect(result).not.toContain('id="e"')
expect(spy).not.toHaveBeenCalled()
spy.mockRestore()
})
})
+2 -2
View File
@@ -9,11 +9,11 @@
* reads as a user edit.
*/
import { DOMParser } from "linkedom"
import { beforeAll, describe, expect, it } from "vitest"
import { installDomPolyfill } from "../src/dom.js"
beforeAll(() => {
;(globalThis as any).DOMParser = DOMParser
installDomPolyfill()
})
import { checkEditGate, contentFingerprint } from "../src/edit-gate.js"
@@ -0,0 +1,213 @@
/**
* Tests for the embedded HTTP server (browser bridge).
*
* The server runs in-process on a random high port (never 6002, which is
* also the default port of the Next.js dev server). Requests go through
* node:http so tests can set raw paths and Host/Origin headers.
*/
import http from "node:http"
import { afterAll, beforeAll, describe, expect, it } from "vitest"
import { addHistory } from "../src/history.js"
import {
getState,
setState,
shutdown,
startHttpServer,
} from "../src/http-server.js"
let port = 0
beforeAll(async () => {
port = await startHttpServer(40000 + Math.floor(Math.random() * 10000))
})
afterAll(() => {
shutdown()
})
interface Response {
status: number
headers: http.IncomingHttpHeaders
body: string
}
/** Send a request; `body` may be split into several writes. */
function request(
path: string,
opts: {
method?: string
headers?: Record<string, string>
body?: Buffer[]
} = {},
): Promise<Response> {
return new Promise((resolve, reject) => {
const req = http.request(
{
host: "127.0.0.1",
port,
path,
method: opts.method ?? "GET",
headers: { host: `localhost:${port}`, ...opts.headers },
},
(res) => {
const chunks: Buffer[] = []
res.on("data", (c: Buffer) => chunks.push(c))
res.on("end", () =>
resolve({
status: res.statusCode ?? 0,
headers: res.headers,
body: Buffer.concat(chunks).toString("utf8"),
}),
)
},
)
req.on("error", reject)
const parts = opts.body ?? []
// Pause between parts so the server reads them as separate chunks
const writeNext = (i: number) => {
if (i >= parts.length) return req.end()
req.write(parts[i])
setTimeout(() => writeNext(i + 1), 30)
}
writeNext(0)
})
}
const postJson = (path: string, data: unknown, headers = {}) =>
request(path, {
method: "POST",
headers: { "content-type": "application/json", ...headers },
body: [Buffer.from(JSON.stringify(data))],
})
describe("session id in the page URL", () => {
it("rejects a session id that could inject script", async () => {
const res = await request(`/?mcp=${encodeURIComponent('";alert(1)//')}`)
expect(res.status).toBe(400)
expect(res.body).not.toContain("alert")
})
it("writes a valid session id into the page script as a JSON string", async () => {
const res = await request("/?mcp=mcp-test-page")
expect(res.status).toBe(200)
expect(res.body).toContain('const sessionId = "mcp-test-page";')
})
})
describe("requests that used to crash the process", () => {
it("answers 400 for a path that is not a valid URL", async () => {
const res = await request("//")
expect(res.status).toBe(400)
// The server is still alive
expect((await request("/api/state?sessionId=mcp-alive")).status).toBe(
200,
)
})
it("never creates sessions with ids unsafe for the Location header", async () => {
const badId = "mcp-中"
await request(`/api/state?sessionId=${encodeURIComponent(badId)}`)
expect(getState(badId)).toBeUndefined()
const post = await postJson("/api/state", {
sessionId: badId,
xml: "<mxfile/>",
})
expect(post.status).toBe(400)
expect(getState(badId)).toBeUndefined()
const res = await request("/")
expect([200, 302]).toContain(res.status)
})
})
describe("request origin checks", () => {
it("refuses a foreign Host header (DNS rebinding)", async () => {
const res = await request("/api/state?sessionId=mcp-alive", {
headers: { host: `evil.example:${port}` },
})
expect(res.status).toBe(403)
})
it("refuses writes from another website", async () => {
const res = await postJson(
"/api/state",
{ sessionId: "mcp-csrf", xml: "<mxfile/>" },
{ origin: "https://evil.example" },
)
expect(res.status).toBe(403)
expect(getState("mcp-csrf")).toBeUndefined()
})
it("accepts writes from the page itself", async () => {
const res = await postJson(
"/api/state",
{ sessionId: "mcp-same-origin", xml: "<mxfile/>" },
{ origin: `http://localhost:${port}` },
)
expect(res.status).toBe(200)
})
})
describe("POST /api/state", () => {
it("decodes UTF-8 characters split across body chunks", async () => {
const xml = `<mxfile>${"数据".repeat(30000)}</mxfile>`
const body = Buffer.from(JSON.stringify({ sessionId: "mcp-utf8", xml }))
// Cut inside a 3-byte character
const cut = body.indexOf(Buffer.from("数")) + 1
const res = await request("/api/state", {
method: "POST",
headers: { "content-type": "application/json" },
body: [body.subarray(0, cut), body.subarray(cut)],
})
expect(res.status).toBe(200)
expect(getState("mcp-utf8")?.xml).toBe(xml)
})
it("rejects a browser push based on a version older than an AI write", async () => {
const id = "mcp-conflict"
setState(id, "<mxfile>user v1</mxfile>", undefined, true)
const aiVersion = setState(id, "<mxfile>AI edit</mxfile>")
const stale = await postJson("/api/state", {
sessionId: id,
xml: "<mxfile>user edit on old version</mxfile>",
baseVersion: aiVersion - 1,
})
expect(stale.status).toBe(409)
expect(getState(id)?.xml).toBe("<mxfile>AI edit</mxfile>")
// Pushes based on the AI version are accepted, including a second
// push sent before the first one's response updated the browser
for (const xml of ["<mxfile>a</mxfile>", "<mxfile>b</mxfile>"]) {
const ok = await postJson("/api/state", {
sessionId: id,
xml,
baseVersion: aiVersion,
})
expect(ok.status).toBe(200)
expect(getState(id)?.xml).toBe(xml)
}
})
})
describe("history restore", () => {
it("restores the entry the user picked after older entries drop", async () => {
const id = "mcp-history"
setState(id, "<mxfile/>")
for (let i = 0; i < 20; i++) addHistory(id, `<mxfile>${i}</mxfile>`)
const list = await request(`/api/history?sessionId=${id}`)
const picked = JSON.parse(list.body).entries[5]
// A new AI edit shifts the buffer before the user clicks Restore
addHistory(id, "<mxfile>new</mxfile>")
const res = await postJson("/api/restore", {
sessionId: id,
id: picked.id,
})
expect(res.status).toBe(200)
expect(getState(id)?.xml).toBe("<mxfile>5</mxfile>")
})
})
@@ -10,18 +10,11 @@
import { deflateRawSync } from "node:zlib"
import { DOMParser } from "linkedom"
import { beforeAll, describe, expect, it } from "vitest"
import { installDomPolyfill } from "../src/dom.js"
// Install the DOM polyfills exactly as index.ts does at runtime.
beforeAll(() => {
;(globalThis as any).DOMParser = DOMParser
class XMLSerializerPolyfill {
serializeToString(node: any): string {
if (node.outerHTML !== undefined) return node.outerHTML
if (node.documentElement) return node.documentElement.outerHTML
return ""
}
}
;(globalThis as any).XMLSerializer = XMLSerializerPolyfill
installDomPolyfill()
})
import {
+2 -10
View File
@@ -15,21 +15,13 @@
* (diagram-operations.ts) — i.e. the layers underneath the MCP tool surface.
*/
import { DOMParser } from "linkedom"
import { beforeAll, describe, expect, it } from "vitest"
import { installDomPolyfill } from "../src/dom.js"
// Install the DOM polyfill exactly as index.ts does at runtime — the
// helpers under test rely on it.
beforeAll(() => {
;(globalThis as any).DOMParser = DOMParser
class XMLSerializerPolyfill {
serializeToString(node: any): string {
if (node.outerHTML !== undefined) return node.outerHTML
if (node.documentElement) return node.documentElement.outerHTML
return ""
}
}
;(globalThis as any).XMLSerializer = XMLSerializerPolyfill
installDomPolyfill()
})
import { applyDiagramOperations } from "../src/diagram-operations.js"
@@ -0,0 +1,186 @@
/**
* Tests for XML syntax checking, autoFixXml and the XML serializer.
*
* linkedom (the DOM used in Node) parses leniently and never reports syntax
* errors, so validation relies on the strict check in dom.ts. autoFixXml
* runs on the whole document whenever any check fails, so its steps must
* leave valid parts of the document untouched.
*/
import { beforeAll, describe, expect, it } from "vitest"
import { getXmlSyntaxError, installDomPolyfill } from "../src/dom.js"
beforeAll(() => {
installDomPolyfill()
})
import { addPageToDoc, parseMxfile, serializeMxfile } from "../src/pages.js"
import { validateAndFixXml } from "../src/xml-validation.js"
/** Bare model with the root cells plus the given cells. */
const model = (cells: string) =>
`<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/>${cells}</root></mxGraphModel>`
// A bare & makes the first validation fail, which triggers autoFixXml on
// the whole document.
const BROKEN_CELL = `<mxCell id="9" value="R&D" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell>`
describe("getXmlSyntaxError", () => {
it("accepts well-formed XML", () => {
expect(getXmlSyntaxError(model(""))).toBeNull()
})
it.each([
["duplicate attribute", `<a style="x" style="y"/>`],
["unquoted attribute", `<a id=2/>`],
["missing space between attributes", `<a id="2"vertex="1"/>`],
["bare ampersand", `<a v="R&D"/>`],
["unclosed tag", `<a><b></a>`],
["plain text", `hello`],
])("reports %s", (_name, xml) => {
expect(getXmlSyntaxError(xml)).toMatch(/^\d+:\d+: /)
})
})
describe("validateAndFixXml", () => {
it("rejects a duplicate style attribute", () => {
const r = validateAndFixXml(
model(
`<mxCell id="2" style="a=1;" style="b=1;" vertex="1" parent="1"/>`,
),
)
expect(r.valid).toBe(false)
expect(r.error).toContain("duplicate attribute: style")
})
it("rejects an unquoted attribute value", () => {
const r = validateAndFixXml(
model(`<mxCell id=2 vertex="1" parent="1"/>`),
)
expect(r.valid).toBe(false)
})
it("keeps style values intact while fixing another cell", () => {
const cells = `<mxCell id="2" style="shape=cylinder3;whiteSpace=wrap;" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><mxCell id="3" style="edgeStyle=orthogonalEdgeStyle;" edge="1" parent="1" source="2" target="2"><mxGeometry relative="1" as="geometry"/></mxCell>`
const r = validateAndFixXml(model(cells + BROKEN_CELL))
expect(r.valid).toBe(true)
expect(r.fixed).toContain('style="shape=cylinder3;whiteSpace=wrap;"')
expect(r.fixed).toContain('style="edgeStyle=orthogonalEdgeStyle;"')
expect(r.fixed).toContain('value="R&amp;D"')
})
it("keeps &quot; inside rich-text labels", () => {
const rich = `<mxCell id="4" value="&lt;font style=&quot;color: red;&quot;&gt;Hi&lt;/font&gt;" style="html=1;" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell>`
const r = validateAndFixXml(model(rich + BROKEN_CELL))
expect(r.valid).toBe(true)
expect(r.fixed).toContain(
'value="&lt;font style=&quot;color: red;&quot;&gt;Hi&lt;/font&gt;"',
)
expect(getXmlSyntaxError(r.fixed ?? "")).toBeNull()
})
it("keeps UserObject and object wrappers", () => {
const wrapped = `<UserObject id="u" label="L" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject><object id="o" label="O"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></object>`
const r = validateAndFixXml(model(wrapped + BROKEN_CELL))
expect(r.valid).toBe(true)
expect(r.fixed).toContain('<UserObject id="u"')
expect(r.fixed).toContain('<object id="o"')
})
it("leaves one-cell-per-line XML alone while fixing another cell", () => {
const xml = [
"<mxGraphModel>",
"<root>",
'<mxCell id="0"/>',
'<mxCell id="1" parent="0"/>',
BROKEN_CELL,
'<mxCell id="3" value="B" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell>',
"</root>",
"</mxGraphModel>",
].join("\n")
const r = validateAndFixXml(xml)
expect(r.valid).toBe(true)
expect(r.fixes).toEqual(["Escaped unescaped & characters"])
})
it("leaves cells split over two lines alone while fixing another cell", () => {
const xml = [
"<mxGraphModel><root>",
'<mxCell id="0"/><mxCell id="1" parent="0"/>',
'<mxCell id="2" value="A" vertex="1" parent="1">',
' <mxGeometry as="geometry"/></mxCell>',
'<mxCell id="3" value="B" vertex="1" parent="1">',
' <mxGeometry as="geometry"/></mxCell>',
BROKEN_CELL,
"</root></mxGraphModel>",
].join("\n")
const r = validateAndFixXml(xml)
expect(r.valid).toBe(true)
expect(r.fixes).toEqual(["Escaped unescaped & characters"])
})
it("renames duplicate short ids without touching the attribute name", () => {
const r = validateAndFixXml(
model(
`<mxCell id="d" vertex="1" parent="1"/><mxCell id="d" vertex="1" parent="1"/><mxCell id="i" vertex="1" parent="1"/><mxCell id="i" vertex="1" parent="1"/>`,
),
)
expect(r.valid).toBe(true)
expect(r.fixed).toContain('id="d_dup1"')
expect(r.fixed).toContain('id="i_dup1"')
})
it("adds a missing space between attributes", () => {
const r = validateAndFixXml(
model(
`<mxCell id="2"value="a" style="x=1;" vertex="1" parent="1"/>`,
),
)
expect(r.valid).toBe(true)
expect(r.fixed).toContain('<mxCell id="2" value="a" style="x=1;"')
})
it("fixes attribute values quoted with &quot;", () => {
const r = validateAndFixXml(
model(
`<mxCell id="2" value=&quot;Hello&quot; vertex="1" parent="1"/>`,
),
)
expect(r.valid).toBe(true)
expect(r.fixed).toContain('value="Hello"')
})
})
describe("XML serializer and strict parsing in page helpers", () => {
it("keeps line breaks and tabs in attribute values", () => {
const xml = `<mxfile><diagram id="p" name="Page-1"><mxGraphModel><root><mxCell id="0"/><mxCell id="2" value="Multi-Head&#xa;Attention&#9;x" vertex="1" parent="0"/></root></mxGraphModel></diagram></mxfile>`
const out = serializeMxfile(parseMxfile(xml) as Document)
expect(out).toContain('value="Multi-Head&#xa;Attention&#9;x"')
expect(out).not.toMatch(/value="[^"]*\n/)
})
it("escapes special characters in attributes and text", () => {
const xml = `<mxfile><diagram id="p" name="R&amp;D">a &lt; b<mxGraphModel><root><mxCell id="0" value="&lt;b&gt; &amp; &quot;"/></root></mxGraphModel></diagram></mxfile>`
const out = serializeMxfile(parseMxfile(xml) as Document)
expect(out).toBe(xml)
})
it("parseMxfile returns null for malformed XML", () => {
expect(
parseMxfile(
`<mxfile><diagram id="p" name="a" name="b"></diagram></mxfile>`,
),
).toBeNull()
})
it("addPageToDoc rejects malformed page XML", () => {
const doc = parseMxfile(
`<mxfile><diagram id="p" name="Page-1">${model("")}</diagram></mxfile>`,
) as Document
expect(() =>
addPageToDoc(doc, {
xml: model(`<mxCell id=2 vertex="1" parent="1"/>`),
}),
).toThrow()
})
})
+6 -7
View File
@@ -48,13 +48,12 @@ export function proxy(request: NextRequest) {
if (pathnameIsMissingLocale) {
const locale = getLocale(request)
// Redirect to localized path
return NextResponse.redirect(
new URL(
`/${locale}${pathname.startsWith("/") ? "" : "/"}${pathname}`,
request.url,
),
)
// Redirect to localized path. Cloning nextUrl keeps the basePath
// (NEXT_PUBLIC_BASE_PATH) and query string, which
// new URL("/...", request.url) would drop.
const url = request.nextUrl.clone()
url.pathname = `/${locale}${pathname}`
return NextResponse.redirect(url)
}
}
+89 -68
View File
@@ -2,7 +2,7 @@
/**
* Development script for running Electron with Next.js
* 1. Reads preset configuration (if exists)
* 1. Reads the active preset's env vars (if any)
* 2. Starts Next.js dev server with preset env vars
* 3. Waits for it to be ready
* 4. Compiles Electron TypeScript
@@ -47,39 +47,41 @@ function getUserDataPath() {
}
/**
* Load preset configuration from config file
* File where the Electron main process (in development) writes the active
* preset's env vars, already decrypted and mapped to provider-specific keys
* (see writeDevPresetEnv in electron/main/config-manager.ts)
*/
function loadPresetConfig() {
const configPath = path.join(getUserDataPath(), "config-presets.json")
if (!existsSync(configPath)) {
console.log("📋 No preset configuration found, using .env.local")
return null
}
const PRESET_ENV_FILE = "dev-preset-env.json"
/**
* Read the active preset's env vars as JSON text (null if not available)
*/
function readPresetEnvFile() {
try {
const content = readFileSync(configPath, "utf-8")
const data = JSON.parse(content)
if (!data.currentPresetId) {
console.log("📋 No active preset, using .env.local")
return null
}
const preset = data.presets.find((p) => p.id === data.currentPresetId)
if (!preset) {
console.log("📋 Active preset not found, using .env.local")
return null
}
console.log(`📋 Using preset: "${preset.name}"`)
return preset.config
} catch (error) {
console.error("Failed to load preset config:", error.message)
const content = readFileSync(
path.join(getUserDataPath(), PRESET_ENV_FILE),
"utf-8",
)
JSON.parse(content) // Ignore a half-written file
return content
} catch {
return null
}
}
/**
* Load the active preset's env vars
*/
function loadPresetEnv(content) {
const env = content ? JSON.parse(content) : {}
if (Object.keys(env).length === 0) {
console.log("📋 No active preset, using .env.local")
return null
}
console.log(`📋 Using preset env: ${Object.keys(env).join(", ")}`)
return env
}
/**
* Wait for the Next.js server to be ready
*/
@@ -128,6 +130,18 @@ function runCommand(command, args, options = {}) {
})
}
/**
* Kill a process started with shell: true. On Windows, kill() only ends the
* cmd.exe wrapper and leaves next dev running, so kill the whole tree.
*/
function killProcess(proc) {
if (process.platform === "win32" && proc.pid) {
spawn("taskkill", ["/pid", String(proc.pid), "/T", "/F"])
} else {
proc.kill()
}
}
/**
* Start Next.js dev server with preset environment
*/
@@ -164,7 +178,8 @@ async function main() {
console.log("🚀 Starting Electron development environment...\n")
// Load preset configuration
const presetEnv = loadPresetConfig()
let presetEnvContent = readPresetEnvFile()
const presetEnv = loadPresetEnv(presetEnvContent)
// Start Next.js dev server with preset env
console.log("1. Starting Next.js development server...")
@@ -176,7 +191,7 @@ async function main() {
console.log("")
} catch (err) {
console.error("\n❌ Next.js server failed to start:", err.message)
nextProcess.kill()
killProcess(nextProcess)
process.exit(1)
}
@@ -186,7 +201,7 @@ async function main() {
await runCommand("npm", ["run", "electron:compile"])
} catch (err) {
console.error("❌ Electron compilation failed:", err.message)
nextProcess.kill()
killProcess(nextProcess)
process.exit(1)
}
@@ -203,76 +218,82 @@ async function main() {
},
})
// Watch for preset config changes
const configPath = path.join(getUserDataPath(), "config-presets.json")
// Watch for preset env changes
const userDataPath = getUserDataPath()
let configWatcher = null
let restartPending = false
function setupConfigWatcher() {
if (!existsSync(path.dirname(configPath))) {
if (!existsSync(userDataPath)) {
// Directory doesn't exist yet, check again later
setTimeout(setupConfigWatcher, 5000)
return
}
try {
// Watch the directory, since the file may not exist yet
configWatcher = watch(
configPath,
userDataPath,
{ persistent: false },
async (eventType) => {
if (eventType === "change" && !restartPending) {
restartPending = true
async (_eventType, filename) => {
if (filename !== PRESET_ENV_FILE || restartPending) return
// Only restart when the preset env vars really changed
const newContent = readPresetEnvFile()
if (newContent === null || newContent === presetEnvContent)
return
restartPending = true
presetEnvContent = newContent
console.log(
"\n🔄 Preset configuration changed, restarting Next.js server...",
)
// Kill current Next.js process
killProcess(nextProcess)
// Wait a bit for process to die
await new Promise((r) => setTimeout(r, 1000))
// Reload preset and restart
nextProcess = startNextServer(loadPresetEnv(newContent))
try {
await waitForServer(NEXT_URL)
console.log(
"\n🔄 Preset configuration changed, restarting Next.js server...",
"✅ Next.js server restarted with new configuration\n",
)
} catch (err) {
console.error(
"❌ Failed to restart Next.js:",
err.message,
)
// Kill current Next.js process
nextProcess.kill()
// Wait a bit for process to die
await new Promise((r) => setTimeout(r, 1000))
// Reload preset and restart
const newPresetEnv = loadPresetConfig()
nextProcess = startNextServer(newPresetEnv)
try {
await waitForServer(NEXT_URL)
console.log(
"✅ Next.js server restarted with new configuration\n",
)
} catch (err) {
console.error(
"❌ Failed to restart Next.js:",
err.message,
)
}
restartPending = false
}
restartPending = false
},
)
console.log("👀 Watching for preset configuration changes...")
} catch (_err) {
// File might not exist yet, that's ok
// Directory might not be ready yet, try again later
setTimeout(setupConfigWatcher, 5000)
}
}
// Start watching after a delay (config file might not exist yet)
// Start watching after a delay (user data directory might not exist yet)
setTimeout(setupConfigWatcher, 2000)
electronProcess.on("close", (code) => {
console.log(`\nElectron exited with code ${code}`)
if (configWatcher) configWatcher.close()
nextProcess.kill()
killProcess(nextProcess)
process.exit(code || 0)
})
electronProcess.on("error", (err) => {
console.error("Electron error:", err)
if (configWatcher) configWatcher.close()
nextProcess.kill()
killProcess(nextProcess)
process.exit(1)
})
@@ -280,8 +301,8 @@ async function main() {
const cleanup = () => {
console.log("\n🛑 Shutting down...")
if (configWatcher) configWatcher.close()
electronProcess.kill()
nextProcess.kill()
killProcess(electronProcess)
killProcess(nextProcess)
process.exit(0)
}
+9
View File
@@ -73,6 +73,15 @@ mkdirSync(targetDir, { recursive: true })
console.log("Copying standalone directory...")
copyDereferenced(standaloneDir, targetDir)
// Next.js copies the build machine's .env files into standalone; don't ship
// them, they can hold the builder's API keys
for (const entry of readdirSync(targetDir)) {
if (entry.startsWith(".env")) {
console.warn(`Removing ${entry} so it is not packaged into the app`)
rmSync(join(targetDir, entry))
}
}
// Copy static files
console.log("Copying static files...")
const targetStaticDir = join(targetDir, ".next", "static")
+71
View File
@@ -0,0 +1,71 @@
// @vitest-environment node
import fs from "fs"
import os from "os"
import path from "path"
import { afterEach, beforeEach, describe, expect, it } from "vitest"
import { GET, PUT } from "@/app/api/admin/providers/route"
import { _resetForTests } from "@/lib/admin/settings"
let tmpDir: string
beforeEach(() => {
tmpDir = fs.mkdtempSync(path.join(os.tmpdir(), "admin-providers-route-"))
process.env.SETTINGS_FILE = path.join(tmpDir, "settings.json")
process.env.ADMIN_PASSWORD = "pw"
process.env.AI_MODELS_CONFIG_PATH = path.join(tmpDir, "none.json")
_resetForTests()
})
afterEach(() => {
_resetForTests()
delete process.env.SETTINGS_FILE
delete process.env.ADMIN_PASSWORD
delete process.env.AI_MODELS_CONFIG_PATH
delete process.env.AI_MODEL
fs.rmSync(tmpDir, { recursive: true, force: true })
})
const headers = { "x-admin-password": "pw" }
async function saveDefaultPanelProvider() {
const res = await PUT(
new Request("http://localhost/api/admin/providers", {
method: "PUT",
headers: { ...headers, "Content-Type": "application/json" },
body: JSON.stringify({
providers: [
{
id: "p1",
provider: "openai",
apiKey: "sk-test",
models: ["gpt-panel"],
isDefault: true,
},
],
}),
}),
)
expect(res.status).toBe(200)
}
async function envHasDefaultModel(): Promise<boolean> {
const res = await GET(
new Request("http://localhost/api/admin/providers", { headers }),
)
return (await res.json()).envHasDefaultModel
}
describe("envHasDefaultModel", () => {
it("is true when .env sets AI_MODEL, even after a panel default", async () => {
process.env.AI_MODEL = "gpt-env"
expect(await envHasDefaultModel()).toBe(true)
await saveDefaultPanelProvider()
expect(await envHasDefaultModel()).toBe(true)
})
it("ignores the AI_MODEL the panel default writes", async () => {
await saveDefaultPanelProvider()
expect(process.env.AI_MODEL).toBe("gpt-panel")
expect(await envHasDefaultModel()).toBe(false)
})
})
+47 -4
View File
@@ -69,7 +69,7 @@ describe("deriveEnvUpdates", () => {
expect(updates.ADMIN_OPENAI_API_KEY_2).toBe("sk-second")
})
it("maps bedrock credentials to AWS env vars", () => {
it("maps bedrock credentials to ADMIN_AWS_* env vars", () => {
const updates = deriveEnvUpdates(
[
provider({
@@ -83,9 +83,26 @@ describe("deriveEnvUpdates", () => {
],
[],
)
expect(updates.AWS_ACCESS_KEY_ID).toBe("AKIA123")
expect(updates.AWS_SECRET_ACCESS_KEY).toBe("secret")
expect(updates.AWS_REGION).toBe("us-west-2")
expect(updates.ADMIN_AWS_ACCESS_KEY_ID).toBe("AKIA123")
expect(updates.ADMIN_AWS_SECRET_ACCESS_KEY).toBe("secret")
expect(updates.ADMIN_AWS_REGION).toBe("us-west-2")
// Standard AWS vars are left to the environment
expect(updates.AWS_ACCESS_KEY_ID).toBeUndefined()
})
it("clears AWS_* bedrock keys written by older versions", () => {
const bedrock = provider({
provider: "bedrock",
apiKey: undefined,
awsAccessKeyId: "AKIA123",
awsSecretAccessKey: "secret",
models: ["claude-x"],
})
const updates = deriveEnvUpdates([bedrock], [bedrock])
expect(updates.AWS_ACCESS_KEY_ID).toBeNull()
expect(updates.AWS_SECRET_ACCESS_KEY).toBeNull()
expect(updates.AWS_REGION).toBeNull()
expect(updates.ADMIN_AWS_ACCESS_KEY_ID).toBe("AKIA123")
})
it("clears keys owned by the previous list when providers are removed", () => {
@@ -315,6 +332,32 @@ describe("validateAdminProviders", () => {
expect(validateAdminProviders(list)).toMatch(/unique/)
})
it("rejects names that differ only in case or punctuation", () => {
const list = [
provider({ id: "p1", name: "Open AI" }),
provider({ id: "p2", name: "open-ai" }),
]
expect(validateAdminProviders(list)).toMatch(/unique/)
})
it("rejects a case-only clash with an env-configured name", () => {
expect(
validateAdminProviders([provider({ name: "openai" })], {
providers: [
{ name: "OpenAI", provider: "openai", models: ["gpt-x"] },
],
}),
).toMatch(/already defined/)
})
it("accepts distinct CJK names", () => {
const list = [
provider({ id: "p1", provider: "deepseek", name: "主力" }),
provider({ id: "p2", provider: "deepseek", name: "备用" }),
]
expect(validateAdminProviders(list)).toBeNull()
})
it("rejects multiple defaults", () => {
const list = [
provider({ id: "p1", isDefault: true }),
+87
View File
@@ -0,0 +1,87 @@
import { cleanup, fireEvent, render, screen } from "@testing-library/react"
import type { ReactNode } from "react"
import { afterEach, describe, expect, it, vi } from "vitest"
import { SecretInput, SettingField } from "@/app/[lang]/admin/setting-field"
import { DictionaryProvider } from "@/hooks/use-dictionary"
import type { SettingDef } from "@/lib/admin/settings-registry"
import type { Dictionary } from "@/lib/i18n/dictionaries"
import en from "@/lib/i18n/dictionaries/en.json"
const STORED = { isSet: true as const, hint: "…abcd" }
function withDict(node: ReactNode) {
return (
<DictionaryProvider dictionary={en as unknown as Dictionary}>
{node}
</DictionaryProvider>
)
}
function typeInto(label: string, text: string) {
fireEvent.change(screen.getByLabelText(label, { selector: "input" }), {
target: { value: text },
})
}
afterEach(cleanup)
describe("SecretInput", () => {
it("reverts to a key saved after mount instead of deleting it", () => {
const onChange = vi.fn()
const props = { id: "secret", keepOnEmpty: true, onChange }
// New provider: nothing stored at mount, then saved
const { rerender } = render(
withDict(
<>
<label htmlFor="secret">secret</label>
<SecretInput {...props} value={undefined} />
</>,
),
)
rerender(
withDict(
<>
<label htmlFor="secret">secret</label>
<SecretInput {...props} value={STORED} />
</>,
),
)
rerender(
withDict(
<>
<label htmlFor="secret">secret</label>
<SecretInput {...props} value="abc" />
</>,
),
)
typeInto("secret", "")
expect(onChange).toHaveBeenLastCalledWith(STORED)
})
})
describe("SettingField secret", () => {
const def: SettingDef = {
key: "LANGFUSE_SECRET_KEY",
group: "observability",
type: "secret",
label: "Langfuse Secret Key",
}
it("drops the pending change when a saved secret is typed over and cleared", () => {
const onChange = vi.fn()
const props = {
def,
state: { key: def.key, source: "file" as const, value: STORED },
disabled: false,
onChange,
}
const { rerender } = render(
withDict(<SettingField {...props} pendingValue={undefined} />),
)
rerender(withDict(<SettingField {...props} pendingValue="a" />))
typeInto(en.admin.settings.LANGFUSE_SECRET_KEY.label, "")
expect(onChange).toHaveBeenLastCalledWith(undefined)
})
})
+19 -1
View File
@@ -1,7 +1,7 @@
import fs from "fs"
import os from "os"
import path from "path"
import { afterEach, beforeEach, describe, expect, it } from "vitest"
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
import {
_resetForTests,
applyToEnv,
@@ -100,6 +100,24 @@ describe("applyToEnv / saveSettings", () => {
expect(process.env.TEST_ADMIN_VAR).toBeUndefined()
})
it("a second module instance can remove a key the first one overlaid", async () => {
// instrumentation.ts and API routes load separate copies in a build
process.env.TEST_ADMIN_VAR = "from-env"
fs.writeFileSync(
process.env.SETTINGS_FILE!,
JSON.stringify({ version: 1, values: { TEST_ADMIN_VAR: "abc" } }),
)
applyToEnv()
expect(process.env.TEST_ADMIN_VAR).toBe("abc")
vi.resetModules()
const second = await import("@/lib/admin/settings")
expect(second.getValueSource("TEST_ADMIN_VAR")).toBe("file")
second.saveSettings({ TEST_ADMIN_VAR: null })
expect(process.env.TEST_ADMIN_VAR).toBe("from-env")
expect(second.getEnvFallback("TEST_ADMIN_VAR")).toBe("from-env")
})
it("persists across cache reset (file round-trip)", () => {
saveSettings({ TEST_ADMIN_VAR: "persisted" })
_resetForTests()
+274
View File
@@ -0,0 +1,274 @@
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
import {
getAIModel,
getValidationModel,
usesServerCredentials,
} from "@/lib/ai-providers"
const settings = vi.hoisted(() => ({ values: {} as Record<string, string> }))
vi.mock("@/lib/admin/settings", () => ({
loadSettings: () => settings.values,
}))
vi.mock("@ai-sdk/google-vertex", () => {
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
return { createVertex: vi.fn(() => mockProviderFn) }
})
vi.mock("@ai-sdk/openai", () => {
const mockModel = { modelId: "test-model" }
const mockProviderFn = vi.fn(() => mockModel) as any
mockProviderFn.chat = vi.fn(() => mockModel)
return {
createOpenAI: vi.fn(() => mockProviderFn),
openai: vi.fn(() => mockModel),
}
})
vi.mock("@ai-sdk/amazon-bedrock", () => {
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
return { createAmazonBedrock: vi.fn(() => mockProviderFn) }
})
vi.mock("@aws-sdk/credential-providers", () => ({
fromNodeProviderChain: vi.fn(() => "node-chain"),
}))
vi.mock("@openrouter/ai-sdk-provider", () => {
const mockProviderFn = vi.fn(() => ({ modelId: "test-model" }))
return { createOpenRouter: vi.fn(() => mockProviderFn) }
})
const ENV_KEYS = [
"GOOGLE_VERTEX_API_KEY",
"GOOGLE_VERTEX_BASE_URL",
"OPENAI_API_KEY",
"OPENAI_BASE_URL",
"OPENROUTER_API_KEY",
"ADMIN_OPENAI_API_KEY",
"ADMIN_OPENROUTER_API_KEY",
"OLLAMA_API_KEY",
"ADMIN_AWS_ACCESS_KEY_ID",
"ADMIN_AWS_SECRET_ACCESS_KEY",
"ADMIN_AWS_REGION",
"AWS_REGION",
"AI_PROVIDER",
"AI_MODEL",
"VALIDATION_MODEL",
]
const savedEnv: Record<string, string | undefined> = {}
beforeEach(() => {
for (const key of ENV_KEYS) {
savedEnv[key] = process.env[key]
delete process.env[key]
}
settings.values = {}
vi.clearAllMocks()
})
afterEach(() => {
for (const key of ENV_KEYS) {
if (savedEnv[key] === undefined) delete process.env[key]
else process.env[key] = savedEnv[key]
}
})
describe("Vertex AI key security", () => {
it("never sends the server key to a client base URL", () => {
process.env.GOOGLE_VERTEX_API_KEY = "server-vertex-key"
// Any x-ai-api-key passes the outer guard; the branch must still refuse
expect(() =>
getAIModel({
provider: "vertexai",
apiKey: "x",
baseUrl: "https://attacker.example",
modelId: "gemini-2.5-flash",
}),
).toThrow("Vertex AI requires an API key")
})
it("sends the client key to the client base URL", async () => {
process.env.GOOGLE_VERTEX_API_KEY = "server-vertex-key"
const { createVertex } = await import("@ai-sdk/google-vertex")
getAIModel({
provider: "vertexai",
vertexApiKey: "client-key",
baseUrl: "https://my-proxy.example",
modelId: "gemini-2.5-flash",
})
expect(createVertex).toHaveBeenCalledWith({
apiKey: "client-key",
baseURL: "https://my-proxy.example",
})
})
it("does not send the client key to the server's base URL", async () => {
process.env.GOOGLE_VERTEX_BASE_URL = "https://server-proxy.internal"
const { createVertex } = await import("@ai-sdk/google-vertex")
getAIModel({
provider: "vertexai",
vertexApiKey: "client-key",
modelId: "gemini-2.5-flash",
})
expect(createVertex).toHaveBeenCalledWith({ apiKey: "client-key" })
})
it("still uses the server key and base URL without client overrides", async () => {
process.env.GOOGLE_VERTEX_API_KEY = "server-vertex-key"
process.env.GOOGLE_VERTEX_BASE_URL = "https://server-proxy.internal"
const { createVertex } = await import("@ai-sdk/google-vertex")
getAIModel({ provider: "vertexai", modelId: "gemini-2.5-flash" })
expect(createVertex).toHaveBeenCalledWith({
apiKey: "server-vertex-key",
baseURL: "https://server-proxy.internal",
})
})
})
describe("Bedrock admin panel credentials", () => {
it("uses the ADMIN_AWS_* keys when the client sends none", async () => {
process.env.ADMIN_AWS_ACCESS_KEY_ID = "panel-id"
process.env.ADMIN_AWS_SECRET_ACCESS_KEY = "panel-secret"
process.env.ADMIN_AWS_REGION = "eu-west-1"
process.env.AWS_REGION = "us-east-1"
const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock")
getAIModel({ provider: "bedrock", modelId: "amazon.nova-lite-v1:0" })
expect(createAmazonBedrock).toHaveBeenCalledWith({
region: "eu-west-1",
accessKeyId: "panel-id",
secretAccessKey: "panel-secret",
})
})
it("prefers the client's keys and region", async () => {
process.env.ADMIN_AWS_ACCESS_KEY_ID = "panel-id"
process.env.ADMIN_AWS_SECRET_ACCESS_KEY = "panel-secret"
process.env.ADMIN_AWS_REGION = "eu-west-1"
const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock")
getAIModel({
provider: "bedrock",
modelId: "amazon.nova-lite-v1:0",
awsAccessKeyId: "client-id",
awsSecretAccessKey: "client-secret",
awsRegion: "ap-northeast-1",
})
expect(createAmazonBedrock).toHaveBeenCalledWith({
region: "ap-northeast-1",
accessKeyId: "client-id",
secretAccessKey: "client-secret",
})
})
it("falls back to the default AWS credential chain", async () => {
process.env.AWS_REGION = "us-east-1"
const { createAmazonBedrock } = await import("@ai-sdk/amazon-bedrock")
getAIModel({ provider: "bedrock", modelId: "amazon.nova-lite-v1:0" })
expect(createAmazonBedrock).toHaveBeenCalledWith({
region: "us-east-1",
credentialProvider: "node-chain",
})
})
})
describe("usesServerCredentials", () => {
it("is true when no key comes with the request", () => {
expect(usesServerCredentials("openai", {})).toBe(true)
expect(usesServerCredentials("openai", { apiKey: "k" })).toBe(false)
})
it("looks at the credential each provider actually uses", () => {
// A stray x-ai-api-key does not replace the IAM role or Vertex key
expect(usesServerCredentials("bedrock", { apiKey: "x" })).toBe(true)
expect(
usesServerCredentials("bedrock", {
awsAccessKeyId: "id",
awsSecretAccessKey: "secret",
}),
).toBe(false)
expect(usesServerCredentials("vertexai", { apiKey: "x" })).toBe(true)
expect(usesServerCredentials("vertexai", { vertexApiKey: "k" })).toBe(
false,
)
})
it("treats keyless EdgeOne and local Ollama as free", () => {
expect(usesServerCredentials("edgeone", {})).toBe(false)
expect(usesServerCredentials("ollama", {})).toBe(false)
expect(
usesServerCredentials("ollama", {
baseUrl: "http://localhost:11434",
}),
).toBe(false)
process.env.OLLAMA_API_KEY = "server-ollama-key"
expect(usesServerCredentials("ollama", {})).toBe(true)
})
})
describe("server model apiKeyEnv", () => {
it("uses the custom env var on the official OpenAI endpoint", async () => {
process.env.ADMIN_OPENAI_API_KEY = "panel-key"
const { createOpenAI, openai } = await import("@ai-sdk/openai")
getAIModel({
provider: "openai",
modelId: "gpt-4o",
apiKeyEnv: "ADMIN_OPENAI_API_KEY",
})
// The default instance would read OPENAI_API_KEY instead
expect(openai).not.toHaveBeenCalled()
expect(createOpenAI).toHaveBeenCalledWith({ apiKey: "panel-key" })
})
})
describe("getValidationModel", () => {
it("uses the admin panel default's ADMIN_ key", async () => {
settings.values = {
ADMIN_PROVIDERS: JSON.stringify([
{
id: "p1",
provider: "openrouter",
name: "My OpenRouter",
apiKey: "panel-key",
models: ["openai/gpt-4o"],
isDefault: true,
},
]),
}
// What deriveEnvUpdates writes for that panel config
process.env.AI_PROVIDER = "openrouter"
process.env.AI_MODEL = "openai/gpt-4o"
process.env.ADMIN_OPENROUTER_API_KEY = "panel-key"
const { createOpenRouter } = await import("@openrouter/ai-sdk-provider")
expect(() => getValidationModel()).not.toThrow()
expect(createOpenRouter).toHaveBeenCalledWith({ apiKey: "panel-key" })
})
it("uses the standard env vars without a panel default", async () => {
process.env.AI_PROVIDER = "openrouter"
process.env.AI_MODEL = "openai/gpt-4o"
process.env.OPENROUTER_API_KEY = "env-key"
const { createOpenRouter } = await import("@openrouter/ai-sdk-provider")
getValidationModel()
expect(createOpenRouter).toHaveBeenCalledWith({ apiKey: "env-key" })
})
})
+156
View File
@@ -0,0 +1,156 @@
// @vitest-environment node
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
import { POST as parseUrl } from "@/app/api/parse-url/route"
import { POST as validateDiagram } from "@/app/api/validate-diagram/route"
import { POST as validateModel } from "@/app/api/validate-model/route"
import { POST as verifyAccessCode } from "@/app/api/verify-access-code/route"
import { checkAccessCode } from "@/lib/access-code"
// Treat every URL as public so no test hits DNS
vi.mock("@/lib/ssrf-protection", async (importOriginal) => ({
...(await importOriginal<typeof import("@/lib/ssrf-protection")>()),
isPrivateUrl: async () => false,
}))
function post(path: string, body: unknown, accessCode?: string): Request {
return new Request(`http://localhost${path}`, {
method: "POST",
headers: {
"Content-Type": "application/json",
...(accessCode ? { "x-access-code": accessCode } : {}),
},
body: JSON.stringify(body),
})
}
beforeEach(() => {
process.env.ACCESS_CODE_LIST = "secret, other"
})
afterEach(() => {
delete process.env.ACCESS_CODE_LIST
delete process.env.ALLOW_PRIVATE_URLS
vi.unstubAllGlobals()
})
describe("checkAccessCode", () => {
it("passes when no access codes are configured", () => {
delete process.env.ACCESS_CODE_LIST
expect(checkAccessCode(post("/x", {}))).toBeNull()
})
it("rejects a missing or wrong code and accepts a listed one", () => {
expect(checkAccessCode(post("/x", {}))?.status).toBe(401)
expect(checkAccessCode(post("/x", {}, "nope"))?.status).toBe(401)
expect(checkAccessCode(post("/x", {}, "other"))).toBeNull()
})
})
describe("routes that spend server resources require the access code", () => {
it("parse-url", async () => {
const res = await parseUrl(
post("/api/parse-url", { url: "https://example.com" }),
)
expect(res.status).toBe(401)
})
it("validate-diagram", async () => {
const res = await validateDiagram(
post("/api/validate-diagram", {
imageData: "data:image/png;base64,AAAA",
}),
)
expect(res.status).toBe(401)
})
it("validate-model", async () => {
const res = await validateModel(
post("/api/validate-model", {
provider: "openai",
apiKey: "sk",
modelId: "m",
}),
)
expect(res.status).toBe(401)
})
it("verify-access-code", async () => {
const bad = await verifyAccessCode(post("/api/verify-access-code", {}))
expect(bad.status).toBe(401)
expect((await bad.json()).valid).toBe(false)
const good = await verifyAccessCode(
post("/api/verify-access-code", {}, "secret"),
)
expect((await good.json()).valid).toBe(true)
})
})
describe("size limits", () => {
it("parse-url stops reading a body over the download limit", async () => {
const chunk = new Uint8Array(1024 * 1024)
let sent = 0
const body = new ReadableStream<Uint8Array>({
pull(controller) {
sent++
controller.enqueue(chunk)
},
})
vi.stubGlobal(
"fetch",
vi.fn(
async () =>
new Response(body, {
headers: { "content-type": "text/html" },
}),
),
)
const res = await parseUrl(
post("/api/parse-url", { url: "https://example.com" }, "secret"),
)
expect(res.status).toBe(413)
expect(sent).toBeLessThan(10)
})
it("validate-diagram rejects oversized image data", async () => {
const imageData = `data:image/png;base64,${"A".repeat(6 * 1024 * 1024)}`
const res = await validateDiagram(
post("/api/validate-diagram", { imageData }, "secret"),
)
expect(res.status).toBe(413)
})
})
describe("validate-model redirects", () => {
it("refuses redirects when private URLs are blocked", async () => {
process.env.ALLOW_PRIVATE_URLS = "false"
vi.stubGlobal(
"fetch",
vi.fn(
async () =>
new Response(null, {
status: 302,
headers: { location: "http://169.254.169.254/" },
}),
),
)
const res = await validateModel(
post(
"/api/validate-model",
{
provider: "openai",
apiKey: "sk",
modelId: "m",
baseUrl: "https://attacker.example/v1",
},
"secret",
),
)
const data = await res.json()
expect(data.valid).toBe(false)
expect(data.error).toMatch(/Redirects are not allowed/)
expect(fetch).toHaveBeenCalledTimes(1)
})
})
+25 -2
View File
@@ -14,12 +14,35 @@ describe("findCachedResponse", () => {
expect(result?.xml).toContain("Transformer Architecture")
})
it("returns cached response for exact match with image", () => {
const result = findCachedResponse("Replicate this in aws style", true)
it("returns cached response for exact match with the example file", () => {
const result = findCachedResponse(
"Replicate this in aws style",
true,
"architecture.png",
)
expect(result).toBeDefined()
expect(result?.xml).toContain("AWS")
})
it("returns undefined when the user attached their own file", () => {
expect(
findCachedResponse("Replicate this flowchart.", true, "mine.png"),
).toBeUndefined()
expect(
findCachedResponse(
"Summarize this paper as a diagram",
true,
"thesis.pdf",
),
).toBeUndefined()
})
it("returns undefined for file examples when the file name is unknown", () => {
// The server only knows whether a file is attached, not which one
const result = findCachedResponse("Replicate this in aws style", true)
expect(result).toBeUndefined()
})
it("returns undefined for non-matching prompt", () => {
const result = findCachedResponse(
"random prompt that doesn't exist",
+162 -2
View File
@@ -1,6 +1,11 @@
// @vitest-environment node
import { convertToModelMessages } from "ai"
import { jsonrepair } from "jsonrepair"
import { describe, expect, it } from "vitest"
import {
dropInvalidToolCalls,
fixToolInputJson,
isMinimalDiagram,
replaceHistoricalToolInputs,
validateFileParts,
@@ -65,6 +70,29 @@ describe("validateFileParts", () => {
expect(result.valid).toBe(false)
expect(result.error).toContain("exceeds")
})
it("rejects file URLs the server would have to download", () => {
for (const url of [
"http://10.0.0.5/secret.png",
"https://example.com/a.png",
undefined,
]) {
const messages = [{ role: "user", parts: [{ type: "file", url }] }]
expect(validateFileParts(messages).valid).toBe(false)
}
})
it("checks files in earlier messages too", () => {
const messages = [
{
role: "user",
parts: [{ type: "file", url: "http://169.254.169.254/x" }],
},
{ role: "assistant", parts: [{ type: "text", text: "ok" }] },
{ role: "user", parts: [{ type: "text", text: "hello" }] },
]
expect(validateFileParts(messages).valid).toBe(false)
})
})
describe("isMinimalDiagram", () => {
@@ -83,6 +111,18 @@ describe("isMinimalDiagram", () => {
const xml = ' <mxCell id="0"/> <mxCell id="1" parent="0"/> '
expect(isMinimalDiagram(xml)).toBe(true)
})
it("returns false for a shape drawn in draw.io with a random id", () => {
const xml =
'<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="xY3kQ9-1" value="" style="rounded=0;" vertex="1" parent="1"><mxGeometry x="10" y="10" width="120" height="60" as="geometry"/></mxCell></root></mxGraphModel>'
expect(isMinimalDiagram(xml)).toBe(false)
})
it("does not mistake ids that start with 0 or 1 for root cells", () => {
const xml =
'<mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="10"/>'
expect(isMinimalDiagram(xml)).toBe(false)
})
})
describe("replaceHistoricalToolInputs", () => {
@@ -124,7 +164,7 @@ describe("replaceHistoricalToolInputs", () => {
)
})
it("removes tool calls with invalid inputs", () => {
it("leaves tool calls with invalid inputs for dropInvalidToolCalls", () => {
const messages = [
{
role: "assistant",
@@ -143,7 +183,7 @@ describe("replaceHistoricalToolInputs", () => {
},
]
const result = replaceHistoricalToolInputs(messages)
expect(result[0].content).toHaveLength(0)
expect(result[0].content).toEqual(messages[0].content)
})
it("preserves non-assistant messages", () => {
@@ -169,3 +209,123 @@ describe("replaceHistoricalToolInputs", () => {
expect(result[0].content[0].input).toEqual({ foo: "bar" })
})
})
describe("dropInvalidToolCalls", () => {
it("drops an invalid tool-call together with its tool-result", () => {
const messages = [
{ role: "user", content: [{ type: "text", text: "draw" }] },
{
role: "assistant",
content: [
{
type: "tool-call",
toolCallId: "call-1",
toolName: "display_diagram",
input: undefined,
},
],
},
{
role: "tool",
content: [
{
type: "tool-result",
toolCallId: "call-1",
toolName: "display_diagram",
output: { type: "error-text", value: "Stopped" },
},
],
},
{ role: "user", content: [{ type: "text", text: "again" }] },
]
const result = dropInvalidToolCalls(messages)
expect(result.map((m) => m.role)).toEqual(["user", "user"])
})
it("keeps valid calls and results in the same messages", () => {
const messages = [
{
role: "assistant",
content: [
{ type: "text", text: "Here you go" },
{
type: "tool-call",
toolCallId: "bad",
toolName: "edit_diagram",
input: "{broken",
},
{
type: "tool-call",
toolCallId: "good",
toolName: "display_diagram",
input: { xml: "<mxCell/>" },
},
],
},
{
role: "tool",
content: [
{ type: "tool-result", toolCallId: "bad", output: {} },
{ type: "tool-result", toolCallId: "good", output: {} },
],
},
]
const result = dropInvalidToolCalls(messages)
expect(result[0].content.map((p: any) => p.toolCallId)).toEqual([
undefined,
"good",
])
expect(result[1].content.map((p: any) => p.toolCallId)).toEqual([
"good",
])
})
it("cleans up a tool call the user stopped before its input arrived", async () => {
// handleStop turns a still-streaming call into output-error with no input
const modelMessages = await convertToModelMessages([
{ role: "user", parts: [{ type: "text", text: "draw" }] },
{
role: "assistant",
parts: [
{
type: "tool-display_diagram",
toolCallId: "call-1",
state: "output-error",
input: undefined,
errorText: "Stopped by user",
} as any,
],
},
{ role: "user", parts: [{ type: "text", text: "again" }] },
])
const result = dropInvalidToolCalls(modelMessages)
expect(result.map((m) => m.role)).toEqual(["user", "user"])
})
it("leaves messages with string content alone", () => {
const messages = [{ role: "system", content: "You are..." }]
expect(dropInvalidToolCalls(messages)).toEqual(messages)
})
})
describe("fixToolInputJson", () => {
it("fixes an attribute whose closing quote alone is escaped", () => {
const input =
'{"xml": "<mxCell id=\\"2\\" vertex=\\"1\\"><mxGeometry x=\\"10\\" y="-20\\" as=\\"geometry\\"/></mxCell>"}'
const parsed = JSON.parse(jsonrepair(fixToolInputJson(input)))
expect(parsed.xml).toContain('y="-20"')
expect(parsed.xml).toContain('id="2"')
})
it("fixes = used instead of : after a JSON key", () => {
const input = '{"xml"= "<mxCell id=\\"2\\"/>"}'
const parsed = JSON.parse(jsonrepair(fixToolInputJson(input)))
expect(parsed.xml).toBe('<mxCell id="2"/>')
})
it("leaves well-formed input unchanged", () => {
const input =
'{"operations": [{"operation": "add", "cell_id": "a", "new_xml": "<mxCell id=\\"a\\" value=\\"x=1\\"/>"}]}'
expect(fixToolInputJson(input)).toBe(input)
})
})
+47
View File
@@ -0,0 +1,47 @@
// @vitest-environment node
import { describe, expect, it } from "vitest"
import { onRequest } from "@/edge-functions/api/edgeai/chat/completions"
function request(headers: Record<string, string>): Request {
return new Request("http://localhost/api/edgeai/chat/completions", {
method: "POST",
headers,
// Non-streaming requests return a mock reply without calling AI
body: JSON.stringify({ messages: [{ role: "user", content: "hi" }] }),
})
}
const json = { "Content-Type": "application/json" }
describe("EdgeOne chat completions function", () => {
it("sends no CORS headers", async () => {
const res = await onRequest({ request: request(json), env: {} })
expect(res.status).toBe(200)
expect(res.headers.get("access-control-allow-origin")).toBeNull()
})
it("rejects non-JSON requests", async () => {
const res = await onRequest({
request: request({ "Content-Type": "text/plain" }),
env: {},
})
expect(res.status).toBe(400)
})
it("checks the access code when ACCESS_CODE_LIST is set", async () => {
const env = { ACCESS_CODE_LIST: "secret" }
const missing = await onRequest({ request: request(json), env })
expect(missing.status).toBe(401)
const ok = await onRequest({
request: request({ ...json, "x-access-code": "secret" }),
env,
})
expect(ok.status).toBe(200)
})
it("lets requests through when env is unavailable", async () => {
const res = await onRequest({ request: request(json) })
expect(res.status).toBe(200)
})
})
+212 -13
View File
@@ -3,6 +3,7 @@ import {
DEFAULT_MAX_OUTPUT_TOKENS,
parseOutputTokenLimit,
resolveMaxOutputTokens,
retryOutputTokens,
withOutputTokenLimitFallback,
} from "@/lib/output-token-limit"
@@ -98,34 +99,204 @@ describe("parseOutputTokenLimit", () => {
}
expect(parseOutputTokenLimit(error)).toBeNull()
})
it("reads the ceiling from a Volcengine Ark rejection", () => {
const error = {
message:
"The parameter `max_tokens` specified in the request are not valid: integer above maximum value, expected a value <= 32768, but got 64000 instead.",
statusCode: 400,
}
expect(parseOutputTokenLimit(error)).toBe(32768)
// Same message JSON-escaped in the response body
expect(
parseOutputTokenLimit({
message: "Bad request",
responseBody:
'{"error":{"message":"The parameter `max_tokens` specified in the request are not valid: integer above maximum value, expected a value \\u003c= 16384, but got 64000 instead."}}',
}),
).toBe(16384)
})
it("reads the ceiling from a DashScope rejection", () => {
const error = {
message:
"<400> InternalError.Algo.InvalidParameter: Range of max_tokens should be [1, 8192]",
}
expect(parseOutputTokenLimit(error)).toBe(8192)
})
it("subtracts the input in SGLang and vLLM context rejections", () => {
// SGLang
expect(
parseOutputTokenLimit({
message:
"Requested token count exceeds the model's maximum context length of 32768 tokens. You requested a total of 70000 tokens: 6000 tokens from the input messages and 64000 tokens for the completion.",
}),
).toBe(32768 - 6000 - 1024)
// vLLM, older wording
expect(
parseOutputTokenLimit({
message:
"This model's maximum context length is 32768 tokens. However, you requested 70000 tokens (6000 in the messages, 64000 in the completion).",
}),
).toBe(32768 - 6000 - 1024)
// vLLM, newer wording
expect(
parseOutputTokenLimit({
message:
"This model's maximum context length is 32768 tokens and your request has 6000 input tokens (64000 > 32768 - 6000).",
}),
).toBe(32768 - 6000 - 1024)
})
})
describe("retryOutputTokens", () => {
const bedrockLimit = Object.assign(
new Error(
"The maximum tokens you requested exceeds the model limit of 64000.",
),
{ statusCode: 400 },
)
it("leaves room for the Bedrock thinking budget the provider adds", () => {
// 64000 + 12000 thinking was sent, so the ceiling of 64000 is below it
expect(
retryOutputTokens(bedrockLimit, {
maxOutputTokens: 64000,
providerOptions: {
bedrock: {
reasoningConfig: {
type: "enabled",
budgetTokens: 12000,
},
},
},
}),
).toBe(52000)
})
it("leaves room for the Anthropic thinking budget the provider adds", () => {
const error = Object.assign(
new Error(
"max_tokens: 76000 > 64000, which is the maximum allowed number of output tokens",
),
{ statusCode: 400 },
)
expect(
retryOutputTokens(error, {
maxOutputTokens: 64000,
providerOptions: {
anthropic: {
thinking: { type: "enabled", budgetTokens: 12000 },
},
},
}),
).toBe(52000)
})
it("does not retry when the ceiling covers what was sent", () => {
expect(
retryOutputTokens(bedrockLimit, { maxOutputTokens: 64000 }),
).toBeNull()
})
it("does not retry when the thinking budget leaves no usable room", () => {
const error = Object.assign(new Error("model limit of 16000"), {
statusCode: 400,
})
expect(
retryOutputTokens(error, {
maxOutputTokens: 64000,
providerOptions: {
bedrock: {
reasoningConfig: {
type: "enabled",
budgetTokens: 15500,
},
},
},
}),
).toBeNull()
})
it("falls back to 16000 when the budget is named but no number can be read", () => {
const error = Object.assign(
new Error("max_tokens (64000) exceeds the limit for this model"),
{ statusCode: 400 },
)
expect(retryOutputTokens(error, { maxOutputTokens: 64000 })).toBe(16000)
// Nothing to gain when the request was already that small
expect(retryOutputTokens(error, { maxOutputTokens: 16000 })).toBeNull()
})
it("does not fall back for errors that do not name the budget", () => {
const error = Object.assign(new Error("temperature must be <= 2"), {
statusCode: 400,
})
expect(retryOutputTokens(error, { maxOutputTokens: 64000 })).toBeNull()
// Not a bad request, even though it names the budget
const auth = Object.assign(new Error("max_tokens: invalid API key"), {
statusCode: 401,
})
expect(retryOutputTokens(auth, { maxOutputTokens: 64000 })).toBeNull()
})
it("does not fall back when a ceiling was found but is too small", () => {
const error = Object.assign(
new Error(
"max_tokens: 64000 > 512, which is the maximum allowed number of output tokens",
),
{ statusCode: 400 },
)
expect(retryOutputTokens(error, { maxOutputTokens: 64000 })).toBeNull()
})
})
describe("resolveMaxOutputTokens", () => {
it("uses a valid header value", () => {
expect(resolveMaxOutputTokens("32000")).toBe(32000)
expect(resolveMaxOutputTokens("32000", false)).toBe(32000)
expect(resolveMaxOutputTokens("32000", true)).toBe(32000)
})
it("falls back to the default for missing or bogus values", () => {
expect(resolveMaxOutputTokens(null)).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
expect(resolveMaxOutputTokens("")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
expect(resolveMaxOutputTokens("abc")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
expect(resolveMaxOutputTokens("0")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
expect(resolveMaxOutputTokens("-5")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
expect(resolveMaxOutputTokens("1.5")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
// Above the sanity ceiling, e.g. an extra zero
expect(resolveMaxOutputTokens("640000")).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
for (const value of [null, "", "abc", "0", "-5", "1.5", "640000"]) {
// "640000" is above the sanity ceiling, e.g. an extra zero
expect(resolveMaxOutputTokens(value, false)).toBe(
DEFAULT_MAX_OUTPUT_TOKENS,
)
}
})
it("uses the env value when no header is sent, and validates it too", () => {
const original = process.env.MAX_OUTPUT_TOKENS
try {
process.env.MAX_OUTPUT_TOKENS = "24000"
expect(resolveMaxOutputTokens(null)).toBe(24000)
// Header still wins
expect(resolveMaxOutputTokens("8000")).toBe(8000)
expect(resolveMaxOutputTokens(null, true)).toBe(24000)
// A lower header still wins
expect(resolveMaxOutputTokens("8000", true)).toBe(8000)
process.env.MAX_OUTPUT_TOKENS = "-1"
expect(resolveMaxOutputTokens(null)).toBe(DEFAULT_MAX_OUTPUT_TOKENS)
expect(resolveMaxOutputTokens(null, true)).toBe(
DEFAULT_MAX_OUTPUT_TOKENS,
)
} finally {
if (original === undefined) delete process.env.MAX_OUTPUT_TOKENS
else process.env.MAX_OUTPUT_TOKENS = original
}
})
it("lets the header raise the budget only on the user's own credentials", () => {
const original = process.env.MAX_OUTPUT_TOKENS
try {
process.env.MAX_OUTPUT_TOKENS = "16000"
expect(resolveMaxOutputTokens("200000", true)).toBe(16000)
expect(resolveMaxOutputTokens("200000", false)).toBe(200000)
// Without MAX_OUTPUT_TOKENS the default is the cap
delete process.env.MAX_OUTPUT_TOKENS
expect(resolveMaxOutputTokens("100000", true)).toBe(
DEFAULT_MAX_OUTPUT_TOKENS,
)
} finally {
if (original === undefined) delete process.env.MAX_OUTPUT_TOKENS
else process.env.MAX_OUTPUT_TOKENS = original
@@ -178,6 +349,34 @@ describe("withOutputTokenLimitFallback", () => {
expect(calls.map((c) => c.maxOutputTokens)).toEqual([64000, 4096])
})
it("subtracts the thinking budget from the retry", async () => {
const [model, calls] = fakeModel([
() =>
Promise.reject(
Object.assign(
new Error(
"The maximum tokens you requested exceeds the model limit of 64000.",
),
{ statusCode: 400 },
),
),
() => Promise.resolve(STREAM_OK),
])
const wrapped = withOutputTokenLimitFallback(model)
await wrapped.doStream({
prompt: [],
maxOutputTokens: 64000,
providerOptions: {
bedrock: {
reasoningConfig: { type: "enabled", budgetTokens: 12000 },
},
},
} as any)
expect(calls.map((c) => c.maxOutputTokens)).toEqual([64000, 52000])
})
it("does not retry an error it cannot attribute to the budget", async () => {
const [model, calls] = fakeModel([
() =>
+16
View File
@@ -0,0 +1,16 @@
import { describe, expect, it } from "vitest"
import { isTextFile } from "@/lib/pdf-utils"
describe("isTextFile", () => {
it("treats SVG files as text so their markup is sent to the model", () => {
const svg = new File(["<svg/>"], "diagram.svg", {
type: "image/svg+xml",
})
expect(isTextFile(svg)).toBe(true)
})
it("does not treat raster images as text", () => {
const png = new File(["x"], "photo.png", { type: "image/png" })
expect(isTextFile(png)).toBe(false)
})
})
+48
View File
@@ -4,6 +4,7 @@ import {
loadFlattenedServerModels,
type ServerModelsConfig,
ServerModelsConfigSchema,
slugify,
} from "@/lib/server-model-config"
const ORIGINAL_ENV = { ...process.env }
@@ -233,3 +234,50 @@ describe("loadFlattenedServerModels", () => {
expect(models[0].apiKeyEnv).toEqual(["OPENAI_KEY_1", "OPENAI_KEY_2"])
})
})
describe("slugify", () => {
it("keeps ASCII names readable", () => {
expect(slugify("OpenAI Production")).toBe("openai-production")
})
it("gives distinct ASCII slugs to distinct CJK names", () => {
const slugs = ["主力", "备用", "DeepSeek 官方", "DeepSeek 备用"].map(
slugify,
)
expect(new Set(slugs).size).toBe(4)
for (const slug of slugs) expect(slug).toMatch(/^[a-z0-9-]+$/)
})
})
describe("loadFlattenedServerModels id collisions", () => {
it("drops a model whose id repeats an earlier provider's", async () => {
const config: ServerModelsConfig = {
providers: [
{ name: "OpenAI", provider: "openai", models: ["gpt-4o"] },
{
name: "openai",
provider: "openai",
models: ["gpt-4o"],
apiKeyEnv: "OTHER_KEY",
},
{
name: "主力",
provider: "deepseek",
models: ["deepseek-chat"],
},
{
name: "备用",
provider: "deepseek",
models: ["deepseek-chat"],
},
],
}
process.env.AI_MODELS_CONFIG = JSON.stringify(config)
const models = await loadFlattenedServerModels()
const ids = models.map((m) => m.id)
expect(new Set(ids).size).toBe(ids.length)
expect(ids).toHaveLength(3)
expect(models[0].apiKeyEnv).toBeUndefined()
})
})
+25
View File
@@ -0,0 +1,25 @@
import { afterEach, describe, expect, it, vi } from "vitest"
import { STORAGE_KEYS } from "@/lib/storage"
import { extractUrlContent } from "@/lib/url-utils"
describe("extractUrlContent", () => {
afterEach(() => {
vi.unstubAllGlobals()
localStorage.clear()
})
it("sends the saved access code with the request", async () => {
localStorage.setItem(STORAGE_KEYS.accessCode, "secret")
const body = { title: "T", content: "body", charCount: 4 }
const fetchMock = vi
.fn()
.mockResolvedValue(new Response(JSON.stringify(body)))
vi.stubGlobal("fetch", fetchMock)
const data = await extractUrlContent("https://example.com")
expect(data.content).toBe("body")
const headers = fetchMock.mock.calls[0][1].headers
expect(headers["x-access-code"]).toBe("secret")
})
})
+102
View File
@@ -0,0 +1,102 @@
import { act, renderHook } from "@testing-library/react"
import { beforeEach, describe, expect, it, vi } from "vitest"
import { extractPdfText, extractTextFileContent } from "@/lib/pdf-utils"
import { useFileProcessor } from "@/lib/use-file-processor"
vi.mock("sonner", () => ({ toast: { error: vi.fn() } }))
vi.mock("@/lib/pdf-utils", async (importOriginal) => ({
...(await importOriginal<typeof import("@/lib/pdf-utils")>()),
extractPdfText: vi.fn(),
extractTextFileContent: vi.fn(),
}))
// A promise we can resolve from the test, to control extraction timing
function deferred<T>() {
let resolve!: (value: T) => void
const promise = new Promise<T>((r) => {
resolve = r
})
return { promise, resolve }
}
const pdfFile = () =>
new File(["%PDF"], "slow.pdf", { type: "application/pdf" })
const textFile = () => new File(["notes"], "notes.txt", { type: "text/plain" })
describe("useFileProcessor", () => {
beforeEach(() => {
vi.mocked(extractPdfText).mockReset()
vi.mocked(extractTextFileContent).mockReset()
})
it("marks queued files as extracting before the first one finishes", async () => {
const pdf = deferred<string>()
vi.mocked(extractPdfText).mockReturnValue(pdf.promise)
vi.mocked(extractTextFileContent).mockResolvedValue("notes")
const a = pdfFile()
const b = textFile()
const { result } = renderHook(() => useFileProcessor())
let done!: Promise<void>
act(() => {
done = result.current.handleFileChange([a, b])
})
expect(result.current.pdfData.get(a)?.isExtracting).toBe(true)
expect(result.current.pdfData.get(b)?.isExtracting).toBe(true)
await act(async () => {
pdf.resolve("pdf text")
await done
})
expect(result.current.pdfData.get(a)?.text).toBe("pdf text")
expect(result.current.pdfData.get(b)?.text).toBe("notes")
})
it("keeps text of a file added while an earlier file is extracting", async () => {
const pdf = deferred<string>()
vi.mocked(extractPdfText).mockReturnValue(pdf.promise)
vi.mocked(extractTextFileContent).mockResolvedValue("notes")
const a = pdfFile()
const b = textFile()
const { result } = renderHook(() => useFileProcessor())
let first!: Promise<void>
act(() => {
first = result.current.handleFileChange([a])
})
await act(async () => {
await result.current.handleFileChange([a, b])
})
expect(result.current.pdfData.get(b)?.text).toBe("notes")
await act(async () => {
pdf.resolve("pdf text")
await first
})
expect(result.current.pdfData.get(a)?.text).toBe("pdf text")
expect(result.current.pdfData.get(b)?.text).toBe("notes")
})
it("does not bring back a file removed while extracting", async () => {
const pdf = deferred<string>()
vi.mocked(extractPdfText).mockReturnValue(pdf.promise)
const a = pdfFile()
const { result } = renderHook(() => useFileProcessor())
let first!: Promise<void>
act(() => {
first = result.current.handleFileChange([a])
})
await act(async () => {
await result.current.handleFileChange([])
})
await act(async () => {
pdf.resolve("pdf text")
await first
})
expect(result.current.pdfData.has(a)).toBe(false)
expect(result.current.files).toEqual([])
})
})
+133
View File
@@ -0,0 +1,133 @@
import { act, cleanup, renderHook, waitFor } from "@testing-library/react"
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"
import { useModelConfig } from "@/hooks/use-model-config"
import type { FlattenedServerModel } from "@/lib/server-model-config"
import { STORAGE_KEYS } from "@/lib/storage"
import type { MultiModelConfig } from "@/lib/types/model-config"
const SERVER_MODELS: FlattenedServerModel[] = [
{
id: "server:openai-main:gpt-4o-mini",
modelId: "gpt-4o-mini",
provider: "openai",
providerLabel: "OpenAI Main",
isDefault: false,
},
{
id: "server:openai-main:gpt-4o",
modelId: "gpt-4o",
provider: "openai",
providerLabel: "OpenAI Main",
isDefault: true,
},
]
const USER_CONFIG: MultiModelConfig = {
version: 1,
providers: [
{
id: "p1",
provider: "openai",
apiKey: "sk-test",
models: [{ id: "m1", modelId: "gpt-4o" }],
},
],
}
function storeConfig(config: MultiModelConfig) {
localStorage.setItem(STORAGE_KEYS.modelConfigs, JSON.stringify(config))
}
async function renderLoaded() {
const hook = renderHook(() => useModelConfig())
await waitFor(() => expect(hook.result.current.isLoaded).toBe(true))
return hook
}
beforeEach(() => {
localStorage.clear()
vi.stubGlobal(
"fetch",
vi.fn(async () => ({
ok: true,
json: async () => ({ models: SERVER_MODELS }),
})),
)
})
afterEach(() => {
cleanup()
vi.unstubAllGlobals()
})
describe("useModelConfig server model selection", () => {
it("replaces a saved server model that no longer exists", async () => {
storeConfig({
...USER_CONFIG,
selectedModelId: "server:openai-production:gpt-4o",
})
const { result } = await renderLoaded()
expect(result.current.selectedModelId).toBe("server:openai-main:gpt-4o")
})
it("keeps a saved server model that still exists", async () => {
storeConfig({
...USER_CONFIG,
selectedModelId: "server:openai-main:gpt-4o-mini",
})
const { result } = await renderLoaded()
expect(result.current.selectedModelId).toBe(
"server:openai-main:gpt-4o-mini",
)
})
it("keeps a selected user model", async () => {
storeConfig({ ...USER_CONFIG, selectedModelId: "m1" })
const { result } = await renderLoaded()
expect(result.current.selectedModelId).toBe("m1")
})
it("falls back to the default server model when the selected model is deleted", async () => {
storeConfig({ ...USER_CONFIG, selectedModelId: "m1" })
const { result } = await renderLoaded()
act(() => result.current.deleteModel("p1", "m1"))
expect(result.current.selectedModelId).toBe("server:openai-main:gpt-4o")
})
it("falls back to the default server model when the selected provider is deleted", async () => {
storeConfig({ ...USER_CONFIG, selectedModelId: "m1" })
const { result } = await renderLoaded()
act(() => result.current.deleteProvider("p1"))
expect(result.current.selectedModelId).toBe("server:openai-main:gpt-4o")
})
})
describe("useModelConfig across tabs", () => {
it("reloads the config when another tab saves it", async () => {
storeConfig({ ...USER_CONFIG, selectedModelId: "m1" })
const { result } = await renderLoaded()
const fromOtherTab: MultiModelConfig = {
...USER_CONFIG,
providers: [
...USER_CONFIG.providers,
{
id: "p2",
provider: "anthropic",
apiKey: "sk-ant",
models: [{ id: "m2", modelId: "claude-sonnet-4-5" }],
},
],
selectedModelId: "m2",
}
act(() => {
storeConfig(fromOtherTab)
window.dispatchEvent(
new StorageEvent("storage", { key: STORAGE_KEYS.modelConfigs }),
)
})
expect(result.current.selectedModelId).toBe("m2")
expect(result.current.config.providers).toHaveLength(2)
})
})
+185 -1
View File
@@ -1,5 +1,14 @@
import { describe, expect, it } from "vitest"
import { cn, isMxCellXmlComplete, wrapWithMxFile } from "@/lib/utils"
import {
applyDiagramOperations,
autoFixXml,
cn,
extractCompleteMxCells,
isMxCellXmlComplete,
validateAndFixXml,
validateMxCellStructure,
wrapWithMxFile,
} from "@/lib/utils"
describe("isMxCellXmlComplete", () => {
it("returns false for empty/null input", () => {
@@ -33,6 +42,31 @@ describe("isMxCellXmlComplete", () => {
expect(isMxCellXmlComplete(xml)).toBe(false)
})
it("returns false when output stops after a child of an open mxCell", () => {
const xml = `<mxCell id="2" value="A" vertex="1" parent="1">
<mxGeometry x="0" y="0" width="80" height="40" as="geometry"/>
</mxCell>
<mxCell id="3" value="B" vertex="1" parent="1">
<mxGeometry x="100" y="0" width="80" height="40" as="geometry"/>`
expect(isMxCellXmlComplete(xml)).toBe(false)
})
it("returns false when output stops after </mxGeometry> of an open mxCell", () => {
const xml = `<mxCell id="e1" edge="1" parent="1" source="2" target="3">
<mxGeometry relative="1" as="geometry">
<mxPoint x="10" y="10" as="sourcePoint"/>
</mxGeometry>`
expect(isMxCellXmlComplete(xml)).toBe(false)
})
it("returns true for a self-closing last mxCell with > in its value", () => {
const xml = `<mxCell id="2" value="A" vertex="1" parent="1">
<mxGeometry as="geometry"/>
</mxCell>
<mxCell id="3" value="A -> B" vertex="1" parent="1"/></root>`
expect(isMxCellXmlComplete(xml)).toBe(true)
})
it("returns true for multiple complete mxCells", () => {
const xml = `<mxCell id="2" value="A" vertex="1" parent="1"/>
<mxCell id="3" value="B" vertex="1" parent="1"/>`
@@ -84,3 +118,153 @@ describe("cn (class name utility)", () => {
expect(cn("text-red-500", "text-blue-500")).toBe("text-blue-500")
})
})
describe("extractCompleteMxCells", () => {
it("keeps the cell right after self-closing root cells", () => {
const xml = `<mxfile><diagram id="p1"><mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="2" value="A" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><mxCell id="3" value="B" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></root></mxGraphModel></diagram></mxfile>`
const ids = [
...extractCompleteMxCells(xml).matchAll(/<mxCell id="([^"]+)"/g),
].map((m) => m[1])
expect(ids).toEqual(["0", "1", "2", "3"])
})
it("drops an incomplete trailing cell", () => {
const xml = `<mxCell id="2" vertex="1" parent="1"/><mxCell id="3" vertex="1" parent="1"><mxGeometry as="geometry"/>`
expect(extractCompleteMxCells(xml)).toBe(
'<mxCell id="2" vertex="1" parent="1"/>',
)
})
})
const page = (id: string, cells: string) =>
`<diagram name="${id}" id="${id}"><mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/>${cells}</root></mxGraphModel></diagram>`
describe("duplicate ids in multi-page documents", () => {
const shape = (id: string, value = "Box") =>
`<mxCell id="${id}" value="${value}" vertex="1" parent="1"><mxGeometry x="0" y="0" width="80" height="40" as="geometry"/></mxCell>`
it("accepts the same ids on different pages", () => {
const xml = `<mxfile>${page("p1", shape("2"))}${page("p2", shape("2"))}</mxfile>`
expect(validateMxCellStructure(xml)).toBeNull()
})
it("still reports duplicate ids within one page", () => {
const xml = `<mxfile>${page("p1", shape("2") + shape("2"))}${page("p2", "")}</mxfile>`
expect(validateMxCellStructure(xml)).toContain("duplicate ID")
})
it("does not rename the root cells of other pages when fixing", () => {
const xml = `<mxfile>${page("p1", shape("2", "R&D"))}${page("p2", shape("3"))}</mxfile>`
const result = validateAndFixXml(xml)
expect(result.valid).toBe(true)
expect(result.fixed).not.toContain("_dup")
expect(result.fixed).toContain("R&amp;D")
})
it("renames a duplicate id within a page", () => {
const xml = `<mxfile>${page("p1", shape("d") + shape("d"))}</mxfile>`
const { fixed } = autoFixXml(xml)
expect(fixed).toContain('<mxCell id="d" ')
expect(fixed).toContain('<mxCell id="d_dup1" ')
})
})
describe("autoFixXml", () => {
it("does not insert a space at the start of style values", () => {
const xml = `<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="2" value="R&D" style="rounded=1;whiteSpace=wrap;" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></root></mxGraphModel>`
const { fixed } = autoFixXml(xml)
expect(fixed).toContain('style="rounded=1;whiteSpace=wrap;"')
})
it("adds a missing space between attributes", () => {
const xml = `<mxCell id="2" vertex="1"parent="1"/>`
expect(autoFixXml(xml).fixed).toContain('vertex="1" parent="1"')
})
it("keeps &quot; inside rich text labels", () => {
const label = "&lt;font color=&quot;#ff0000&quot;&gt;Hello&lt;/font&gt;"
const xml = `<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><mxCell id="2" value="${label}" style="html=1;" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><mxCell id="3" value="Q&A" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></root></mxGraphModel>`
const result = validateAndFixXml(xml)
expect(result.valid).toBe(true)
expect(result.fixed).toContain(`value="${label}"`)
})
it("fixes an attribute delimited by &quot;", () => {
const xml = `<mxCell id="2" dashPattern=&quot;1 1;&quot; vertex="1" parent="1"/>`
expect(autoFixXml(xml).fixed).toContain('dashPattern="1 1;"')
})
it("keeps cells written on one line next to multi-line cells", () => {
const xml = `<mxGraphModel><root>
<mxCell id="0"/>
<mxCell id="1" parent="0"/>
<mxCell id="2" value="Q&A" vertex="1" parent="1">
<mxGeometry x="0" y="0" width="80" height="40" as="geometry"/>
</mxCell>
<mxCell id="e1" edge="1" parent="1" source="2" target="3"><mxGeometry relative="1" as="geometry"/></mxCell>
<mxCell id="3" value="B" vertex="1" parent="1">
<mxGeometry x="200" y="0" width="80" height="40" as="geometry"/>
</mxCell>
</root></mxGraphModel>`
const result = validateAndFixXml(xml)
expect(result.valid).toBe(true)
for (const id of ["2", "e1", "3"]) {
expect(result.fixed).toContain(`<mxCell id="${id}"`)
}
})
it("keeps object and UserObject wrappers", () => {
const xml = `<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><UserObject id="2" label="Docs" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject><object id="3" label="A&B" owner="me"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></object></root></mxGraphModel>`
const result = validateAndFixXml(xml)
expect(result.valid).toBe(true)
expect(result.fixed).toContain('<UserObject id="2"')
expect(result.fixed).toContain('<object id="3"')
})
})
describe("applyDiagramOperations with wrapped cells", () => {
const xml = `<mxfile><diagram id="p1"><mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/><UserObject id="5" label="Docs" link="https://example.com"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject><mxCell id="6" value="B" vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell><mxCell id="e1" edge="1" parent="1" source="5" target="6"><mxGeometry relative="1" as="geometry"/></mxCell></root></mxGraphModel></diagram></mxfile>`
it("deletes a wrapped cell and its edges", () => {
const { result, errors } = applyDiagramOperations(xml, [
{ operation: "delete", cell_id: "5" },
{ operation: "delete", cell_id: "e1" },
])
expect(errors).toEqual([])
expect(result).not.toContain("UserObject")
expect(result).not.toContain('id="e1"')
expect(result).toContain('id="6"')
})
it("rejects adding a cell with the id of a wrapped cell", () => {
const { errors } = applyDiagramOperations(xml, [
{
operation: "add",
cell_id: "5",
new_xml: '<mxCell id="5" vertex="1" parent="1"/>',
},
])
expect(errors[0]?.message).toContain("already exists")
})
it("updates a wrapped cell", () => {
const { result, errors } = applyDiagramOperations(xml, [
{
operation: "update",
cell_id: "5",
new_xml:
'<UserObject id="5" label="New" link="https://example.org"><mxCell vertex="1" parent="1"><mxGeometry as="geometry"/></mxCell></UserObject>',
},
])
expect(errors).toEqual([])
expect(result).toContain('label="New"')
expect(result).not.toContain('label="Docs"')
})
it("reports deleting a cell that does not exist", () => {
const { errors } = applyDiagramOperations(xml, [
{ operation: "delete", cell_id: "missing" },
])
expect(errors[0]?.message).toContain("not found")
})
})