diff --git a/packages/mcp-server/package-lock.json b/packages/mcp-server/package-lock.json index fc7a3ada..5d84a11b 100644 --- a/packages/mcp-server/package-lock.json +++ b/packages/mcp-server/package-lock.json @@ -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": { @@ -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", @@ -3379,6 +3392,12 @@ "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.6.5", "resolved": "https://registry.npmjs.org/zod/-/zod-4.6.5.tgz", diff --git a/packages/mcp-server/package.json b/packages/mcp-server/package.json index 3143eb2e..702089d7 100644 --- a/packages/mcp-server/package.json +++ b/packages/mcp-server/package.json @@ -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": { diff --git a/packages/mcp-server/src/diagram-operations.ts b/packages/mcp-server/src/diagram-operations.ts index 4711de7b..6e40b9f6 100644 --- a/packages/mcp-server/src/diagram-operations.ts +++ b/packages/mcp-server/src/diagram-operations.ts @@ -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 +// (or ): 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 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() - 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() // Process each operation for (const op of operations) { @@ -164,7 +178,7 @@ export function applyDiagramOperations( `${op.new_xml}`, "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( `${op.new_xml}`, "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) } } } diff --git a/packages/mcp-server/src/dom.ts b/packages/mcp-server/src/dom.ts new file mode 100644 index 00000000..fd265c55 --- /dev/null +++ b/packages/mcp-server/src/dom.ts @@ -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 ), 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 = { + "&": "&", + "<": "<", + ">": ">", + '"': """, + "\t": " ", + "\n": " ", + "\r": " ", +} +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 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}` + } + case 3: + // Text + return escapeChars(node.textContent ?? "", /[&<>]/g) + case 4: + // CDATA + return `` + case 8: + // Comment + return `` + 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 +} diff --git a/packages/mcp-server/src/history.ts b/packages/mcp-server/src/history.ts index 492eda1b..22b89544 100644 --- a/packages/mcp-server/src/history.ts +++ b/packages/mcp-server/src/history.ts @@ -6,7 +6,15 @@ import { log } from "./logger.js" const MAX_HISTORY = 20 -const historyStore = new Map>() + +interface HistoryEntry { + id: number // Stable across shifts of the circular buffer + xml: string + svg: string +} + +let nextEntryId = 0 +const historyStore = new Map() 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 { diff --git a/packages/mcp-server/src/http-server.ts b/packages/mcp-server/src/http-server.ts index f23a2455..08b50bca 100644 --- a/packages/mcp-server/src/http-server.ts +++ b/packages/mcp-server/src/http-server.ts @@ -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--" (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 {