Files
SystemSimulationApp/app/main.py
T

331 lines
10 KiB
Python

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": (
'<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/>'
'<object id="tank" label="\u8d2e\u7bb1" componentType="tank" '
'medium="N2" volume="1.0"><mxCell style="shape=cylinder;'
"whiteSpace=wrap;html=1;boundedLbl=1;backgroundOutline=1;"
'size=15;fillColor=#dae8fc;strokeColor=#6c8ebf;" vertex="1" '
'parent="1"><mxGeometry width="80" height="80" as="geometry"/>'
"</mxCell></object></root></mxGraphModel>"
),
},
{
"title": "\u9600\u95e8",
"w": 72,
"h": 56,
"xml": (
'<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/>'
'<object id="valve" label="\u9600\u95e8" componentType="valve" '
'valveType="manual" nominalDiameter="10"><mxCell style="rhombus;'
'whiteSpace=wrap;html=1;fillColor=#ffe6cc;strokeColor=#d79b00;" '
'vertex="1" parent="1"><mxGeometry width="72" height="56" '
'as="geometry"/></mxCell></object></root></mxGraphModel>'
),
},
{
"title": "\u7ba1\u6bb5",
"w": 96,
"h": 36,
"xml": (
'<mxGraphModel><root><mxCell id="0"/><mxCell id="1" parent="0"/>'
'<object id="pipe" label="\u7ba1\u6bb5" componentType="pipe" '
'length="1.0" diameter="0.01"><mxCell style="rounded=1;'
"whiteSpace=wrap;html=1;arcSize=12;fillColor=#d5e8d4;"
'strokeColor=#82b366;" vertex="1" parent="1"><mxGeometry '
'width="96" height="36" as="geometry"/></mxCell></object>'
"</root></mxGraphModel>"
),
},
]
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]