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, 2, -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": 2}), } 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()