#!/usr/bin/env python3
"""Extract a normalized intermediate representation (IR) from a draw.io file.
The deterministic half of the draw.io import flow: this script never makes a
design decision. It decodes whatever draw.io wrote (raw XML, deflate+base64
payloads, PNG/SVG files with an embedded ``mxfile``), flattens the mxGraphModel
into absolute-positioned nodes and edges, and reports structural signals — hubs,
containers, depth, cycles, leaf clusters — that the skill uses to pick a diagram
type and a level of detail.
Usage:
python3 drawio_extract.py [--page N|NAME] [--json]
[--max-rows N] [--out PATH]
Default output is a compact Markdown digest meant to be read into context.
``--json`` emits the full IR instead (every node, every edge, every style).
Exit codes: 0 ok, 2 unreadable / unsupported input.
"""
from __future__ import annotations
import argparse
import base64
import html
import json
import re
import struct
import sys
import zlib
from dataclasses import dataclass, field, asdict
from pathlib import Path
from typing import Any
from urllib.parse import unquote
from xml.etree import ElementTree as ET
# --------------------------------------------------------------------------
# container / payload decoding
# --------------------------------------------------------------------------
PNG_MAGIC = b"\x89PNG\r\n\x1a\n"
MAX_INPUT_BYTES = 32 * 1024 * 1024
MAX_XML_BYTES = 64 * 1024 * 1024
class PayloadTooLarge(ValueError):
"""Raised when compressed metadata expands beyond the supported limit."""
def _fail(msg: str) -> "NoReturn": # type: ignore[valid-type]
print(f"drawio_extract: {msg}", file=sys.stderr)
raise SystemExit(2)
def _reject_unsafe_xml(xml: str, source: str) -> None:
"""Reject declarations that can make XML parsing expand external data."""
upper = xml.upper()
if " bytes:
"""Decompress without allowing a small payload to expand without bound."""
decompressor = zlib.decompressobj(wbits)
output = bytearray()
chunk = data
while chunk:
remaining = limit + 1 - len(output)
if remaining <= 0:
raise PayloadTooLarge(f"decoded payload exceeds {limit} bytes")
output.extend(decompressor.decompress(chunk, remaining))
if len(output) > limit:
raise PayloadTooLarge(f"decoded payload exceeds {limit} bytes")
chunk = decompressor.unconsumed_tail
if not decompressor.eof:
raise zlib.error("incomplete compressed payload")
remaining = limit + 1 - len(output)
if remaining <= 0:
raise PayloadTooLarge(f"decoded payload exceeds {limit} bytes")
output.extend(decompressor.flush(remaining))
if len(output) > limit:
raise PayloadTooLarge(f"decoded payload exceeds {limit} bytes")
return bytes(output)
def _inflate(payload: str) -> str | None:
"""Undo draw.io's base64 + raw-deflate + URL-encoding pipeline."""
try:
raw = base64.b64decode(payload, validate=False)
except Exception:
return None
for wbits in (-15, 15, 47):
try:
text = _decompress_limited(raw, wbits).decode("utf-8", "replace")
except PayloadTooLarge:
_fail(
f"decoded diagram exceeds the {MAX_XML_BYTES // (1024 * 1024)} MiB limit"
)
except Exception:
continue
# draw.io URL-encodes before deflating; unquote is a no-op if it didn't.
return unquote(text)
return None
def _png_embedded_xml(data: bytes) -> str | None:
"""Pull the ``mxfile`` tEXt/zTXt chunk out of a draw.io-exported PNG."""
pos = len(PNG_MAGIC)
while pos + 8 <= len(data):
(length,) = struct.unpack(">I", data[pos : pos + 4])
ctype = data[pos + 4 : pos + 8]
body_end = pos + 8 + length
chunk_end = body_end + 4
if chunk_end > len(data):
_fail("PNG has a truncated metadata chunk")
body = data[pos + 8 : body_end]
pos = chunk_end
if ctype not in (b"tEXt", b"zTXt", b"iTXt"):
if ctype == b"IEND":
break
continue
key, _, rest = body.partition(b"\x00")
if key.lower() != b"mxfile":
continue
try:
if ctype == b"tEXt":
value = rest
elif ctype == b"zTXt":
value = _decompress_limited(rest[1:], 15)
else: # iTXt: compression flag, method, lang, translated key, text
flag = rest[0:1]
tail = rest[2:].split(b"\x00", 2)[-1]
value = _decompress_limited(tail, 15) if flag == b"\x01" else tail
except PayloadTooLarge:
_fail(
f"embedded PNG diagram exceeds the "
f"{MAX_XML_BYTES // (1024 * 1024)} MiB limit"
)
except (IndexError, ValueError, zlib.error):
_fail("PNG has invalid compressed draw.io metadata")
return unquote(value.decode("utf-8", "replace"))
return None
def _svg_embedded_xml(text: str) -> str | None:
for match in re.finditer(
r"\bcontent\s*=\s*([\"'])(.*?)\1", text, flags=re.IGNORECASE | re.DOTALL
):
candidate = html.unescape(match.group(2))
if " str:
"""Return the ```` (or bare ````) XML for any input."""
size = path.stat().st_size
if size > MAX_INPUT_BYTES:
_fail(
f"{path.name}: input is {size} bytes; maximum is "
f"{MAX_INPUT_BYTES // (1024 * 1024)} MiB"
)
data = path.read_bytes()
if data.startswith(PNG_MAGIC):
xml = _png_embedded_xml(data)
if not xml:
_fail(f"{path.name}: PNG has no embedded draw.io diagram")
return xml
text = data.decode("utf-8", "replace").lstrip("").strip()
if "|
|", re.IGNORECASE)
TAG_RE = re.compile(r"<[^>]+>")
def parse_style(style: str | None) -> dict[str, str]:
out: dict[str, str] = {}
if not style:
return out
for part in style.split(";"):
part = part.strip()
if not part:
continue
key, sep, value = part.partition("=")
out[key.strip()] = value.strip() if sep else "1"
return out
def clean_label(value: str | None) -> str:
"""draw.io labels are often HTML fragments; flatten to plain text lines."""
if not value:
return ""
text = BR_RE.sub("\n", value)
text = TAG_RE.sub("", text)
text = html.unescape(text)
text = text.replace("\xa0", " ")
lines = [re.sub(r"[ \t]+", " ", ln).strip() for ln in text.split("\n")]
return "\n".join(ln for ln in lines if ln).strip()
SHAPE_FAMILIES = (
("mxgraph.aws", "aws"),
("mxgraph.azure", "azure"),
("mxgraph.gcp", "gcp"),
("mxgraph.kubernetes", "kubernetes"),
("mxgraph.cisco", "network"),
("mxgraph.veeam", "infra"),
("mxgraph.flowchart", "flowchart"),
("mxgraph.bpmn", "bpmn"),
("mxgraph.er", "er"),
("mxgraph.sysml", "uml"),
("mxgraph.archimate", "archimate"),
)
# style key -> canonical shape name, checked in order
SHAPE_KEYS = (
("swimlane", "swimlane"),
("ellipse", "ellipse"),
("rhombus", "rhombus"),
("triangle", "triangle"),
("cylinder", "cylinder"),
("cylinder3", "cylinder"),
("hexagon", "hexagon"),
("cloud", "cloud"),
("actor", "actor"),
("umlActor", "actor"),
("note", "note"),
("card", "card"),
("step", "step"),
("process", "process"),
("parallelogram", "parallelogram"),
("document", "document"),
("datastore", "cylinder"),
("umlLifeline", "lifeline"),
("umlFrame", "frame"),
("table", "table"),
("tableRow", "table-row"),
("partialRectangle", "table-row"),
("image", "image"),
("text", "text"),
("group", "group"),
)
def classify_shape(style: dict[str, str]) -> str:
raw = style.get("shape", "")
if raw:
for key, name in SHAPE_KEYS:
if raw == key or raw.startswith(key):
return name
for prefix, family in SHAPE_FAMILIES:
if raw.startswith(prefix):
return f"icon:{family}"
return f"shape:{raw}"
for key, name in SHAPE_KEYS:
if key in style:
return name
if style.get("ellipse") == "1":
return "ellipse"
return "rect"
def shape_family(shape: str) -> str:
if shape.startswith("icon:"):
return shape.split(":", 1)[1]
if shape.startswith("shape:"):
return "custom"
return shape
# --------------------------------------------------------------------------
# IR model
# --------------------------------------------------------------------------
@dataclass
class Node:
id: str
label: str = ""
shape: str = "rect"
parent: str | None = None
depth: int = 0
x: float = 0.0
y: float = 0.0
w: float = 0.0
h: float = 0.0
fill: str = ""
stroke: str = ""
font_color: str = ""
dashed: bool = False
rounded: bool = False
container: bool = False
children: list[str] = field(default_factory=list)
link: str = ""
attrs: dict[str, str] = field(default_factory=dict)
in_degree: int = 0
out_degree: int = 0
@dataclass
class Edge:
id: str
source: str | None
target: str | None
label: str = ""
dashed: bool = False
bidirectional: bool = False
undirected: bool = False
style_name: str = ""
waypoints: int = 0
stroke: str = ""
@dataclass
class Page:
id: str
name: str
index: int
nodes: list[Node] = field(default_factory=list)
edges: list[Edge] = field(default_factory=list)
@property
def node_map(self) -> dict[str, Node]:
return {n.id: n for n in self.nodes}
def _num(geom: ET.Element | None, key: str) -> float:
if geom is None:
return 0.0
try:
return float(geom.get(key, "0") or 0)
except ValueError:
return 0.0
def parse_page(diagram: ET.Element, index: int) -> Page:
name = diagram.get("name") or f"Page-{index + 1}"
page = Page(id=diagram.get("id") or f"page-{index}", name=name, index=index)
model = diagram.find(".//mxGraphModel")
if model is None:
text = (diagram.text or "").strip()
inflated = _inflate(text) if text else None
if not inflated:
return page
_reject_unsafe_xml(inflated, f"page {index}")
model = ET.fromstring(inflated)
if model.tag != "mxGraphModel":
found = model.find(".//mxGraphModel")
if found is None:
return page
model = found
root = model.find("root")
if root is None:
return page
# Pass 1: collect raw cells, unwrapping