228 lines
8.5 KiB
Python
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()
|