from __future__ import annotations from types import SimpleNamespace import unittest from unittest.mock import patch from app.simulation.components.amesim.boundary.sources import AmesimPnpl01 from app.simulation.components.amesim.flow.pipes import AmesimPnl0002 from app.simulation.components.amesim.mechanical.translational import ( AmesimLstp00a, ) from app.simulation.components.experimental.storage.cylinder import Cylinder from app.simulation.components.experimental.storage.tank import Tank from app.simulation.core.medium import IdealGasMedium from app.simulation.solvers.algebraic import ( AlgebraicSolveDiagnostics, PressureFlowSolver, ) from app.simulation.solvers.algebraic_blocks import ( StreamPressureBlockSolver, _BlockSolveAttempt, ) from app.simulation.systems.network import SimulationNetwork def _pnl0002_solver( *, include_unselected_island: bool = False, ) -> tuple[PressureFlowSolver, AmesimPnl0002]: medium = IdealGasMedium() left = Cylinder("left", medium, V=0.1, p0=500_000.0) right = Tank("right", medium, V=0.1, p0=100_000.0) pipe = AmesimPnl0002( "pipe", medium, p0=300_000.0, T0=300.0, ) network = SimulationNetwork("pnl0002-equation-blocks") for component in (left, right, pipe): network.add_component(component) network.connect("left", "port_b", "pipe", "port_1") network.connect("pipe", "port_2", "right", "port_a") if include_unselected_island: isolated = Cylinder("isolated", medium, V=0.2, p0=700_000.0) plug = AmesimPnpl01("isolated_plug") network.add_component(isolated) network.add_component(plug) network.connect("isolated", "port_b", "isolated_plug", "port_1") for component in network.dynamic_components(): component.refresh_thermodynamic_ports() solver = PressureFlowSolver(network) solver.solve() return solver, pipe class StreamPressureBlockSolverTests(unittest.TestCase): def test_pnl0002_uses_two_blocks_with_one_shared_component(self) -> None: solver, pipe = _pnl0002_solver() block_solver = StreamPressureBlockSolver(solver, ("pipe",)) self.assertTrue(block_solver.available, block_solver.fallback_reason) self.assertEqual(len(block_solver.blocks), 2) self.assertTrue( all("pipe" in block.scope_components for block in block_solver.blocks) ) pipe_component_evaluations = [ evaluation for block in block_solver.blocks for evaluation in block.component_evaluations if getattr(evaluation.evaluate, "__self__", None) is pipe ] self.assertEqual(len(pipe_component_evaluations), 2) self.assertTrue( {unknown.id for unknown in block_solver.blocks[0].unknowns}.isdisjoint( unknown.id for unknown in block_solver.blocks[1].unknowns ) ) result = block_solver.solve(scale_context=solver.scale_context()) self.assertFalse(result.used_global_fallback) self.assertEqual(len(result.diagnostics), 1) self.assertAlmostEqual( pipe.port_1.m_flow, pipe.port_mass_flow( pipe.port_1.p, pipe.properties().p, pipe.properties().T, port_name="port_1", ), places=12, ) self.assertAlmostEqual( pipe.port_2.m_flow, pipe.port_mass_flow( pipe.port_2.p, pipe.properties().p, pipe.properties().T, port_name="port_2", ), places=12, ) def test_selected_block_seeding_preserves_every_unselected_unknown(self) -> None: solver, _pipe = _pnl0002_solver(include_unselected_island=True) block_solver = StreamPressureBlockSolver(solver, ("pipe",)) unselected = tuple( unknown for unknown in solver.unknowns if unknown.component in {"isolated", "isolated_plug"} ) for index, unknown in enumerate(unselected, start=1): unknown.write(10_000.0 * index) expected = tuple(unknown.read() for unknown in unselected) result = block_solver.solve(scale_context=solver.scale_context()) self.assertFalse(result.used_global_fallback) self.assertEqual( tuple(unknown.read() for unknown in unselected), expected, ) def test_causal_secondary_pressure_mutation_fuses_to_verified_path(self) -> None: solver, _pipe = _pnl0002_solver() block_solver = StreamPressureBlockSolver(solver, ("pipe",)) first = block_solver.solve(scale_context=solver.scale_context()) self.assertFalse(first.used_global_fallback) expected_pressures = tuple( unknown.read() for unknown in block_solver._selected_unknowns if unknown.variable == "p" ) original_seed = block_solver._seed_selected_blocks mutated_pressure = next( unknown for unknown in block_solver._selected_unknowns if unknown.variable == "p" ) def seed_then_mutate(entry_values): seeded = original_seed(entry_values) mutated_pressure.write(mutated_pressure.read() + 10_000.0) return seeded with patch.object( block_solver, "_seed_selected_blocks", side_effect=seed_then_mutate, ): result = block_solver.solve(scale_context=solver.scale_context()) self.assertFalse(result.used_global_fallback) self.assertTrue(result.diagnostics[0].residual_verified_this_solve) actual_pressures = tuple( unknown.read() for unknown in block_solver._selected_unknowns if unknown.variable == "p" ) for expected, actual in zip(expected_pressures, actual_pressures): self.assertAlmostEqual( actual, expected, delta=1.0e-12 * max(abs(expected), 1.0), ) execution = block_solver.causal_execution_diagnostics() self.assertFalse(execution["enabled"]) self.assertEqual( execution["disabledReason"], "causalSecondaryRuntimeGateFailed", ) self.assertEqual(execution["legacyFallbackCount"], 1) def test_sparse_block_reports_actual_residual_evaluations(self) -> None: solver, pipe = _pnl0002_solver() block_solver = StreamPressureBlockSolver(solver, ("pipe",)) pipe.port_1.m_flow += 0.01 with patch.object( block_solver, "_seed_selected_blocks", return_value=None, ): result = block_solver.solve(scale_context=solver.scale_context()) sparse = result.diagnostics[0] self.assertEqual(sparse.jacobian_mode, "blockSparse") self.assertGreater(sparse.evaluations, 0) self.assertGreater(sparse.residual_evaluations, sparse.evaluations) self.assertFalse(sparse.dense_fallback_used) def test_failed_sparse_block_restores_its_original_unknowns(self) -> None: solver, pipe = _pnl0002_solver() block_solver = StreamPressureBlockSolver(solver, ("pipe",)) block = block_solver.blocks[0] pipe.port_1.m_flow += 0.01 expected = tuple(unknown.read() for unknown in block.unknowns) def failed_sparse(_fun, x0, **_kwargs): return SimpleNamespace( x=x0 + 123.0, success=False, status=-1, message="forced sparse block failure", nfev=1, ) with patch("scipy.optimize.least_squares", side_effect=failed_sparse): attempt = block_solver._solve_block( block, solver.scale_context(), ) self.assertIsNone(attempt.diagnostics) self.assertEqual(attempt.optimizer_evaluations, 1) self.assertEqual( tuple(unknown.read() for unknown in block.unknowns), expected, ) def test_block_failure_restores_full_snapshot_before_global_fallback(self) -> None: solver, pipe = _pnl0002_solver(include_unselected_island=True) block_solver = StreamPressureBlockSolver(solver, ("pipe",)) pipe.port_1.m_flow += 0.01 solver.network.components["isolated_plug"].port_1.p = 12_345.0 expected = tuple(unknown.read() for unknown in solver.unknowns) original_global_solve = solver.solve def checked_global_solve(*args, **kwargs): self.assertEqual( tuple(unknown.read() for unknown in solver.unknowns), expected, ) return original_global_solve(*args, **kwargs) def failed_block(block, _scale_context, **_kwargs): for unknown in block.unknowns: unknown.write(unknown.read() + 321.0) return _BlockSolveAttempt( diagnostics=None, optimizer_evaluations=2, residual_evaluations=7, failure_reason="blockResidualNotConverged", ) with patch.object( block_solver, "_seed_selected_blocks", return_value=None, ), patch.object( block_solver, "_solve_block", side_effect=failed_block, ), patch.object( solver, "solve", side_effect=checked_global_solve, ) as global_solve: result = block_solver.solve(scale_context=solver.scale_context()) global_solve.assert_called_once() self.assertTrue(result.used_global_fallback) self.assertEqual(len(result.diagnostics), 1) fallback_diagnostics = result.diagnostics[0] self.assertIn( fallback_diagnostics.jacobian_mode, { "seeded", "sparse", "dense", "sparseThenDense", "blockSparse", }, ) self.assertGreaterEqual( fallback_diagnostics.residual_evaluations, fallback_diagnostics.evaluations, ) self.assertGreaterEqual(fallback_diagnostics.evaluations, 2) self.assertGreaterEqual(fallback_diagnostics.residual_evaluations, 7) self.assertTrue(fallback_diagnostics.block_fallback_used) self.assertIn( "blockResidualNotConverged", fallback_diagnostics.block_fallback_reason or "", ) def test_seed_exception_restores_snapshot_before_global_fallback(self) -> None: solver, _pipe = _pnl0002_solver(include_unselected_island=True) block_solver = StreamPressureBlockSolver(solver, ("pipe",)) expected = tuple(unknown.read() for unknown in solver.unknowns) original_global_solve = solver.solve def broken_seed(_entry_values) -> None: solver.unknowns[0].write(solver.unknowns[0].read() + 123_456.0) raise ValueError("forced seed failure") def checked_global_solve(*args, **kwargs): self.assertEqual( tuple(unknown.read() for unknown in solver.unknowns), expected, ) return original_global_solve(*args, **kwargs) with patch.object( block_solver, "_seed_selected_blocks", side_effect=broken_seed, ), patch.object( solver, "solve", side_effect=checked_global_solve, ) as global_solve: result = block_solver.solve(scale_context=solver.scale_context()) global_solve.assert_called_once() self.assertTrue(result.used_global_fallback) def test_failed_global_fallback_does_not_leak_candidate_state(self) -> None: solver, _pipe = _pnl0002_solver(include_unselected_island=True) block_solver = StreamPressureBlockSolver(solver, ("pipe",)) expected = tuple(unknown.read() for unknown in solver.unknowns) def failed_block(block, _scale_context, **_kwargs): for unknown in block.unknowns: unknown.write(unknown.read() + 123.0) return _BlockSolveAttempt( diagnostics=None, optimizer_evaluations=1, residual_evaluations=4, failure_reason="blockResidualNotConverged", ) def failed_global_solve(*_args, **_kwargs): for unknown in solver.unknowns: unknown.write(unknown.read() - 456.0) raise RuntimeError("forced global fallback failure") with patch.object( block_solver, "_seed_selected_blocks", return_value=None, ), patch.object( block_solver, "_seeded_diagnostics", return_value=None, ), patch.object( block_solver, "_solve_block", side_effect=failed_block, ), patch.object( solver, "solve", side_effect=failed_global_solve, ): with self.assertRaisesRegex( RuntimeError, "forced global fallback failure", ): block_solver.solve(scale_context=solver.scale_context()) self.assertEqual( tuple(unknown.read() for unknown in solver.unknowns), expected, ) def test_local_to_global_fallback_aggregates_attempt_chain_once(self) -> None: solver, _pipe = _pnl0002_solver() block_solver = StreamPressureBlockSolver(solver, ("pipe",)) accepted_global = AlgebraicSolveDiagnostics( success=True, message="accepted global result", evaluations=3, pressure_scale=400_000.0, flow_scale=0.01, max_scaled_residual=0.25, max_raw_residual=5.0, residual_evaluations=11, jacobian_mode="dense", ) failed_local = _BlockSolveAttempt( diagnostics=None, optimizer_evaluations=2, residual_evaluations=7, failure_reason="blockResidualNotConverged", ) with patch.object( block_solver, "_seeded_diagnostics", return_value=None, ), patch.object( block_solver, "_solve_block", return_value=failed_local, ), patch.object( solver, "solve", return_value=accepted_global, ) as global_solve: result = block_solver.solve(scale_context=solver.scale_context()) global_solve.assert_called_once() self.assertTrue(result.used_global_fallback) self.assertEqual(len(result.diagnostics), 1) aggregate = result.diagnostics[0] self.assertEqual(aggregate.evaluations, 5) self.assertEqual(aggregate.residual_evaluations, 18) self.assertEqual(aggregate.jacobian_mode, "dense") self.assertEqual(aggregate.max_scaled_residual, 0.25) self.assertTrue(aggregate.block_fallback_used) self.assertEqual( aggregate.block_fallback_reason, "blockResidualNotConverged", ) def test_failed_global_fallback_restores_lstp_causal_cache_for_base_errors( self, ) -> None: class ForcedFatalError(BaseException): pass for error_type in (MemoryError, ForcedFatalError): with self.subTest(error_type=error_type.__name__): solver, _pipe = _pnl0002_solver() block_solver = StreamPressureBlockSolver(solver, ("pipe",)) block_solver.fallback_reason = "forcedUntrustedStructure" contact = AmesimLstp00a( "contact", IdealGasMedium(), gap0=0.0, kcont=1.0e6, rcont=0.0, Pdis=1.0e-6, discContactOption=1.0, ) contact.port_1.x = 1.25 contact.port_2.x = 1.0 contact.port_1.v = 0.5 contact.port_2.v = -0.25 contact.set_causal_contact(penetration=0.25, force=12.0) solver._causal_contact_components = (contact,) expected_unknowns = tuple( unknown.read() for unknown in solver.unknowns ) causal_names = tuple( name for name in vars(contact) if name.startswith("_causal_") ) expected_causal = tuple( getattr(contact, name) for name in causal_names ) expected_last = solver.last_diagnostics def failed_global_solve(*_args, **_kwargs): solver.unknowns[0].write( solver.unknowns[0].read() + 123_456.0 ) contact.clear_causal_contact() solver.last_diagnostics = None raise error_type("forced global fallback failure") with patch.object( solver, "solve", side_effect=failed_global_solve, ): with self.assertRaises(error_type): block_solver.solve(scale_context=solver.scale_context()) self.assertEqual( tuple(unknown.read() for unknown in solver.unknowns), expected_unknowns, ) self.assertEqual( tuple(getattr(contact, name) for name in causal_names), expected_causal, ) self.assertIs(solver.last_diagnostics, expected_last) if __name__ == "__main__": unittest.main()