Files
SystemSimulationApp/tests/test_reactflow_project_schema.py

228 lines
8.5 KiB
Python

from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
from xml.etree import ElementTree as ET
from fastapi import HTTPException
from app.main import (
ReactFlowParameterScientificNotation,
ReactFlowProjectPayload,
build_reactflow_system_xml,
compile_reactflow_network,
load_reactflow_project,
reactflow_project_storage_data,
save_reactflow_project,
)
def medium_project(**root_overrides: object) -> ReactFlowProjectPayload:
payload: dict[str, object] = {
"projectSchemaVersion": 1,
"name": "current-medium-project",
"nodes": [
{
"id": "air_1",
"position": {"x": 10, "y": 20},
"data": {
"label": "Air 1",
"componentType": "amesim_ideal_air_medium",
"modelType": "amesim_ideal_air_medium",
"modelVersion": "0.2.0",
"ports": [],
"parameters": {"gi": 1},
},
}
],
}
payload.update(root_overrides)
return ReactFlowProjectPayload(**payload)
class ReactFlowProjectSchemaTests(unittest.TestCase):
def test_current_project_version_is_accepted(self) -> None:
project = medium_project()
self.assertEqual(project.projectSchemaVersion, 1)
self.assertEqual(
project.model_dump(mode="json")["projectSchemaVersion"],
1,
)
self.assertEqual(project.nodes[0].data.modelVersion, "0.2.0")
def test_missing_node_model_version_can_be_loaded_for_inspection(self) -> None:
project = medium_project()
project.nodes[0].data.modelVersion = None
reloaded = ReactFlowProjectPayload.model_validate(
project.model_dump(mode="json", exclude_none=True)
)
self.assertIsNone(reloaded.nodes[0].data.modelVersion)
def test_execution_rejects_missing_or_mismatched_node_model_version(self) -> None:
missing = medium_project()
missing.nodes[0].data.modelVersion = None
for execute in (build_reactflow_system_xml, compile_reactflow_network):
with self.subTest(case="missing", execute=execute.__name__):
with self.assertRaisesRegex(
ValueError,
"COMPONENT_MODEL_VERSION_MISSING",
):
execute(missing)
mismatch = medium_project()
mismatch.nodes[0].data.modelVersion = "0.1.0"
for execute in (build_reactflow_system_xml, compile_reactflow_network):
with self.subTest(case="mismatch", execute=execute.__name__):
with self.assertRaisesRegex(
ValueError,
"COMPONENT_MODEL_VERSION_MISMATCH",
):
execute(mismatch)
def test_execution_rejects_divergent_component_and_model_types(self) -> None:
project = medium_project()
project.nodes[0].data.componentType = "amesim_helium_medium"
for execute in (build_reactflow_system_xml, compile_reactflow_network):
with self.subTest(execute=execute.__name__):
with self.assertRaisesRegex(
ValueError,
"COMPONENT_MODEL_TYPE_MISMATCH",
):
execute(project)
def test_unsupported_project_version_is_rejected(self) -> None:
for version in (0, 3, -1, "1"):
with self.subTest(version=version):
with self.assertRaises(ValueError):
medium_project(projectSchemaVersion=version)
def test_missing_project_version_is_rejected(self) -> None:
with self.assertRaises(ValueError):
ReactFlowProjectPayload(name="missing-project-version")
def test_removed_compatibility_markers_are_rejected(self) -> None:
for field in (
"mediumReferenceVersion",
"amesimParameterEncodingVersion",
"presentationLayoutVersion",
):
with self.subTest(field=field):
with self.assertRaises(ValueError):
medium_project(**{field: 1})
def test_string_port_definition_is_rejected(self) -> None:
with self.assertRaises(ValueError):
ReactFlowProjectPayload(
projectSchemaVersion=1,
name="invalid-string-port",
nodes=[
{
"id": "tank_1",
"data": {
"componentType": "tank",
"modelType": "tank",
"ports": ["port_a"],
},
}
],
)
def test_xml_export_fills_registered_medium_defaults(self) -> None:
root = ET.fromstring(build_reactflow_system_xml(medium_project()))
parameters = {
parameter.get("name"): float(str(parameter.get("value")))
for parameter in root.findall("./Components/Component/Parameter")
}
self.assertEqual(parameters["gi"], 1.0)
self.assertEqual(parameters["property_model"], 0.0)
self.assertIsNone(root.get("projectSchemaVersion"))
def test_storage_preserves_parameter_display_metadata(self) -> None:
project = medium_project()
project.nodes[0].data.parameterUnits = {
"gi": "",
"property_model": "",
}
project.nodes[0].data.parameterScientificNotation = {
"gi": ReactFlowParameterScientificNotation(text="1e0", unit="")
}
data = reactflow_project_storage_data(project)
stored_node = data["nodes"][0]["data"]
self.assertEqual(stored_node["modelVersion"], "0.2.0")
self.assertEqual(stored_node["parameterUnits"], project.nodes[0].data.parameterUnits)
self.assertEqual(
stored_node["parameterScientificNotation"],
{"gi": {"text": "1e0", "unit": ""}},
)
def test_project_storage_round_trip_preserves_current_contract(self) -> None:
project = medium_project()
project.nodes[0].data.parameterUnits = {"gi": ""}
project.nodes[0].data.parameterScientificNotation = {
"gi": ReactFlowParameterScientificNotation(text="1e0", unit="")
}
project = ReactFlowProjectPayload.model_validate(
{
**project.model_dump(mode="json"),
"edges": [
{
"id": "contact-edge",
"source": "air_1",
"target": "air_1",
"sourceHandle": "definition",
"targetHandle": "definition",
"data": {
"isContactEdge": True,
"futureDisplayMetadata": "preserved",
},
}
],
}
)
with tempfile.TemporaryDirectory() as directory:
storage = Path(directory)
with patch("app.main.PROJECT_STORAGE_DIR", storage):
result = save_reactflow_project("medium-project", project)
loaded = load_reactflow_project("medium-project")
self.assertEqual(result["id"], "medium-project")
self.assertEqual(loaded, reactflow_project_storage_data(project))
self.assertTrue(loaded["edges"][0]["data"]["isContactEdge"])
self.assertEqual(
loaded["edges"][0]["data"]["futureDisplayMetadata"],
"preserved",
)
def test_project_load_rejects_corrupt_or_unsupported_data(self) -> None:
with tempfile.TemporaryDirectory() as directory:
storage = Path(directory)
invalid_cases = {
"corrupt": "{not-json",
"future": json.dumps({"projectSchemaVersion": 3}),
}
with patch("app.main.PROJECT_STORAGE_DIR", storage):
for project_id, text in invalid_cases.items():
with self.subTest(project_id=project_id):
(storage / f"{project_id}.json").write_text(
text,
encoding="utf-8",
)
with self.assertRaises(HTTPException) as context:
load_reactflow_project(project_id)
self.assertEqual(context.exception.status_code, 422)
if __name__ == "__main__":
unittest.main()