from __future__ import annotations import base64 import json import os import zlib from pathlib import Path from typing import Iterable from urllib.parse import unquote from xml.etree import ElementTree as ET from fastapi import FastAPI, HTTPException, Request, Response from fastapi.responses import HTMLResponse from pydantic import BaseModel app = FastAPI(title="System Simulation App") INDEX_TEMPLATE = Path(__file__).parent / "static" / "index.html" class DiagramPayload(BaseModel): xml: str @app.get("/") def index(request: Request) -> HTMLResponse: html = INDEX_TEMPLATE.read_text(encoding="utf-8") drawio_base_url = os.getenv("DRAWIO_BASE_URL") or default_drawio_base_url(request) html = html.replace("__DRAWIO_BASE_URL__", json.dumps(drawio_base_url)) return HTMLResponse(html) @app.post("/api/system-xml") def export_system_xml(payload: DiagramPayload) -> Response: try: model = read_graph_model(payload.xml) system_xml = build_system_xml(model) except ValueError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc return Response(content=system_xml, media_type="application/xml") def default_drawio_base_url(request: Request) -> str: hostname = request.url.hostname or "127.0.0.1" if ":" in hostname and not hostname.startswith("["): hostname = f"[{hostname}]" return f"{request.url.scheme}://{hostname}:8080/" def read_graph_model(xml_text: str) -> ET.Element: xml_text = xml_text.strip() if not xml_text: raise ValueError("Diagram XML is empty.") try: root = ET.fromstring(xml_text) except ET.ParseError as exc: raise ValueError("Diagram XML is invalid.") from exc if tag_name(root) == "mxGraphModel": return root if tag_name(root) != "mxfile": raise ValueError("Expected mxGraphModel or mxfile XML.") diagram = next((child for child in root if tag_name(child) == "diagram"), None) if diagram is None: raise ValueError("mxfile does not contain a diagram.") embedded_model = next( (child for child in diagram if tag_name(child) == "mxGraphModel"), None, ) if embedded_model is not None: return embedded_model if not diagram.text: raise ValueError("diagram node is empty.") return decode_diagram_payload(diagram.text.strip()) def decode_diagram_payload(payload: str) -> ET.Element: if payload.startswith("<"): return parse_graph_model(payload) if payload.startswith("%3C"): return parse_graph_model(unquote(payload)) try: padded = payload + ("=" * (-len(payload) % 4)) compressed = base64.b64decode(padded) try: decoded = zlib.decompress(compressed, -15) except zlib.error: decoded = zlib.decompress(compressed) return parse_graph_model(unquote(decoded.decode("utf-8"))) except Exception as exc: # noqa: BLE001 - keep API error small and stable. raise ValueError("Unable to decode draw.io diagram payload.") from exc def parse_graph_model(xml_text: str) -> ET.Element: try: model = ET.fromstring(xml_text) except ET.ParseError as exc: raise ValueError("Decoded diagram XML is invalid.") from exc if tag_name(model) != "mxGraphModel": raise ValueError("Decoded diagram is not an mxGraphModel.") return model def build_system_xml(model: ET.Element) -> bytes: root = next((child for child in model if tag_name(child) == "root"), None) if root is None: raise ValueError("mxGraphModel does not contain a root node.") system = ET.Element("System") components_node = ET.SubElement(system, "Components") connections_node = ET.SubElement(system, "Connections") for component in read_components(root): component_node = ET.SubElement( components_node, "Component", { "id": component["id"], "type": component["type"], "label": component["label"], }, ) for name, value in component["properties"].items(): ET.SubElement( component_node, "Property", {"name": name, "value": value}, ) for edge in read_connections(root): ET.SubElement(connections_node, "Connection", edge) ET.indent(system, space=" ") return ET.tostring(system, encoding="utf-8", xml_declaration=True) def read_components(root: ET.Element) -> Iterable[dict]: for element in root: if tag_name(element) != "object": continue cell = next((child for child in element if tag_name(child) == "mxCell"), None) if cell is None or cell.attrib.get("vertex") != "1": continue component_id = element.attrib.get("id") or cell.attrib.get("id") if not component_id: continue properties = { key: value for key, value in element.attrib.items() if key not in {"id", "label", "componentType"} } yield { "id": component_id, "type": element.attrib.get("componentType", "component"), "label": element.attrib.get("label", ""), "properties": properties, } def read_connections(root: ET.Element) -> Iterable[dict[str, str]]: for cell in root.iter(): if tag_name(cell) != "mxCell" or cell.attrib.get("edge") != "1": continue source = cell.attrib.get("source") target = cell.attrib.get("target") if not source or not target: continue edge = {"source": source, "target": target} if "id" in cell.attrib: edge["id"] = cell.attrib["id"] yield edge def tag_name(element: ET.Element) -> str: return element.tag.rsplit("}", 1)[-1]