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
This commit is contained in:
dayuan.jiang
2026-10-03 17:45:41 +09:00
parent 95f4b4b92b
commit a46787c1b8
16 changed files with 999 additions and 272 deletions
+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": {
@@ -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",
+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()
})
})