from __future__ import annotations import base64 import json import os import urllib.error import urllib.request import zlib from pathlib import Path from typing import Iterable from urllib.parse import unquote, urljoin 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" DRAWIO_UPSTREAM_URL = ( os.getenv("DRAWIO_UPSTREAM_URL") or os.getenv("DRAWIO_BASE_URL") or "http://127.0.0.1:8081/" ) COMPONENT_LIBRARY = [ { "title": "\u8d2e\u7bb1", "w": 80, "h": 80, "xml": ( '' '' "" ), }, { "title": "\u9600\u95e8", "w": 72, "h": 56, "xml": ( '' '' ), }, { "title": "\u7ba1\u6bb5", "w": 96, "h": 36, "xml": ( '' '' "" ), }, ] class DiagramPayload(BaseModel): xml: str @app.get("/") def index(request: Request) -> HTMLResponse: html = INDEX_TEMPLATE.read_text(encoding="utf-8") drawio_base_url = default_drawio_base_url(request) html = html.replace("__DRAWIO_BASE_URL__", json.dumps(drawio_base_url)) return HTMLResponse(html, headers={"Cache-Control": "no-store"}) @app.get("/api/component-library.drawiolib") def component_library() -> Response: library = ET.Element("mxlibrary", {"title": "System Components"}) library.text = json.dumps(COMPONENT_LIBRARY, ensure_ascii=False) return Response( content=ET.tostring(library, encoding="utf-8"), media_type="application/xml", headers={ "Access-Control-Allow-Origin": "*", "Cache-Control": "no-store", }, ) @app.api_route("/drawio", methods=["GET", "POST", "HEAD", "OPTIONS"]) @app.api_route("/drawio/{path:path}", methods=["GET", "POST", "HEAD", "OPTIONS"]) async def drawio_proxy(request: Request, path: str = "") -> Response: upstream_url = build_drawio_upstream_url(path, request.url.query) body = await request.body() headers = { key: value for key, value in request.headers.items() if key.lower() not in { "host", "connection", "content-length", "accept-encoding", } } upstream_request = urllib.request.Request( upstream_url, data=body if body else None, headers=headers, method=request.method, ) try: with urllib.request.urlopen(upstream_request, timeout=30) as upstream_response: content = upstream_response.read() return Response( content=content, status_code=upstream_response.status, headers=filter_proxy_headers(upstream_response.headers), ) except urllib.error.HTTPError as exc: return Response( content=exc.read(), status_code=exc.code, headers=filter_proxy_headers(exc.headers), ) except urllib.error.URLError as exc: raise HTTPException( status_code=502, detail=f"Unable to reach draw.io upstream: {exc.reason}", ) from exc @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}]" port = f":{request.url.port}" if request.url.port else "" return f"{request.url.scheme}://{hostname}{port}/drawio/" def build_drawio_upstream_url(path: str, query: str) -> str: target = urljoin(DRAWIO_UPSTREAM_URL, path) if query: target = f"{target}?{query}" return target def filter_proxy_headers(headers) -> dict[str, str]: excluded = { "connection", "content-encoding", "content-length", "keep-alive", "proxy-authenticate", "proxy-authorization", "te", "trailer", "transfer-encoding", "upgrade", } return { key: value for key, value in headers.items() if key.lower() not in excluded } 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]