from __future__ import annotations from types import SimpleNamespace import unittest from unittest.mock import patch from scipy.optimize import least_squares as scipy_least_squares 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 ( AmesimF000, AmesimForc, AmesimLstp00a, AmesimMecmas21, ) 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 PressureFlowSolver from app.simulation.systems.network import SimulationNetwork class _CustomPnl0002(AmesimPnl0002): """External subclass: structural declarations alone are not a purity promise.""" def _branched_solver( *, custom_pipe: bool = False, include_contact: 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_type = _CustomPnl0002 if custom_pipe else AmesimPnl0002 pipe = pipe_type("pipe", medium, p0=300_000.0, T0=300.0) isolated = Cylinder("isolated", medium, V=0.2, p0=700_000.0) plug = AmesimPnpl01("isolated_plug") network = SimulationNetwork("nonlinear-equation-blocks") for component in (left, right, pipe, isolated, plug): network.add_component(component) network.connect("left", "port_b", "pipe", "port_1") network.connect("pipe", "port_2", "right", "port_a") network.connect("isolated", "port_b", "isolated_plug", "port_1") if include_contact: force = AmesimForc("contact_force") force.res.signal = 40.0 contact = AmesimLstp00a( "contact", medium, gap0=0.0, kcont=1.0e11, rcont=0.0, Pdis=1.0e-7, discContactOption=1.0, ) mass = AmesimMecmas21( "mass", medium, mass=2.0, useFriction=1.0, x0=1.0e9, ) zero = AmesimF000("zero") for component in (force, contact, mass, zero): network.add_component(component) network.connect("contact_force", "port_2", "contact", "port_1") network.connect("contact", "port_2", "mass", "port_1") network.connect("mass", "port_2", "zero", "port_1") for component in network.dynamic_components(): component.refresh_thermodynamic_ports() solver = PressureFlowSolver(network) solver.solve() return solver, pipe def _damage_flows_after_seed( solver: PressureFlowSolver, pipe: AmesimPnl0002, ports: tuple[str, ...], *, snapshots: list[tuple[float, ...]] | None = None, ) -> None: original = solver._solve_explicit_flow_unknowns def damaged(*args, **kwargs): result = original(*args, **kwargs) for index, port_name in enumerate(ports, start=1): port = pipe.get_port(port_name) port.m_flow += index * 0.01 if snapshots is not None: snapshots.append(tuple(unknown.read() for unknown in solver.unknowns)) return result solver._solve_explicit_flow_unknowns = damaged class PressureFlowEquationBlockTests(unittest.TestCase): def test_pnl0002_owner_is_safely_split_across_two_blocks(self) -> None: solver, _pipe = _branched_solver() pipe_rows = { index for index, equation in enumerate(solver.equation_templates) if equation.owner == "component" and equation.owner_id == "pipe" } owner_blocks = [ block for block in solver.equation_blocks if pipe_rows.intersection(block.equation_indices) ] self.assertTrue(solver.equation_blocks_are_trusted) self.assertEqual(len(pipe_rows), 2) self.assertEqual(len(owner_blocks), 2) self.assertTrue( set(owner_blocks[0].unknown_indices).isdisjoint( owner_blocks[1].unknown_indices ) ) def test_only_bad_block_uses_scoped_sparse_residual(self) -> None: optimized, optimized_pipe = _branched_solver() dense, dense_pipe = _branched_solver() _damage_flows_after_seed(optimized, optimized_pipe, ("port_1",)) _damage_flows_after_seed(dense, dense_pipe, ("port_1",)) dense._equation_blocks_are_trusted = False dense._jacobian_sparsity_is_trusted = False block_diagnostics = optimized.solve() dense_diagnostics = dense.solve() self.assertEqual(block_diagnostics.jacobian_mode, "blockSparse") self.assertEqual(block_diagnostics.nonlinear_block_count, 1) self.assertEqual(block_diagnostics.nonlinear_block_unknown_count, 4) self.assertFalse(block_diagnostics.block_fallback_used) self.assertLess( block_diagnostics.residual_evaluations, dense_diagnostics.residual_evaluations, ) for block_unknown, dense_unknown in zip( optimized.unknowns, dense.unknowns, ): self.assertAlmostEqual( block_unknown.read(), dense_unknown.read(), delta=1.0e-6 * max(abs(dense_unknown.read()), 1.0), ) def test_multiple_bad_blocks_are_solved_in_one_union_call(self) -> None: solver, pipe = _branched_solver() _damage_flows_after_seed(solver, pipe, ("port_1", "port_2")) with patch( "scipy.optimize.least_squares", wraps=scipy_least_squares, ) as least_squares: diagnostics = solver.solve() self.assertEqual(least_squares.call_count, 1) self.assertEqual(diagnostics.jacobian_mode, "blockSparse") self.assertEqual(diagnostics.nonlinear_block_count, 2) self.assertEqual(diagnostics.nonlinear_block_unknown_count, 8) subset_pattern = least_squares.call_args.kwargs["jac_sparsity"] self.assertEqual(subset_pattern.shape, (8, 8)) self.assertLess(subset_pattern.nnz, solver.jacobian_sparsity.nnz) def test_failed_local_candidate_restores_global_seed_before_fallback( self, ) -> None: solver, pipe = _branched_solver() seeded_snapshots: list[tuple[float, ...]] = [] _damage_flows_after_seed( solver, pipe, ("port_1",), snapshots=seeded_snapshots, ) call_count = 0 def fail_block_then_solve_global(fun, x0, **kwargs): nonlocal call_count call_count += 1 if call_count == 1: self.assertEqual(x0.shape, (4,)) return SimpleNamespace( x=x0 + 123.0, success=False, status=-1, message="forced block failure", nfev=1, ) self.assertEqual(x0.shape, (len(solver.unknowns),)) self.assertEqual( tuple(unknown.read() for unknown in solver.unknowns), seeded_snapshots[-1], ) return scipy_least_squares(fun, x0, **kwargs) with patch( "scipy.optimize.least_squares", side_effect=fail_block_then_solve_global, ): diagnostics = solver.solve() self.assertGreaterEqual(call_count, 2) self.assertTrue(diagnostics.success) self.assertTrue(diagnostics.block_fallback_used) self.assertEqual(diagnostics.nonlinear_block_count, 1) self.assertEqual( diagnostics.block_fallback_reason, "blockResidualNotConverged", ) def test_block_memory_error_restores_seed_and_remains_fatal(self) -> None: solver, pipe = _branched_solver() seeded_snapshots: list[tuple[float, ...]] = [] _damage_flows_after_seed( solver, pipe, ("port_1",), snapshots=seeded_snapshots, ) with patch( "scipy.optimize.least_squares", side_effect=MemoryError("forced allocation failure"), ): with self.assertRaises(MemoryError): solver.solve() self.assertEqual( tuple(unknown.read() for unknown in solver.unknowns), seeded_snapshots[-1], ) def test_custom_component_keeps_the_legacy_global_path(self) -> None: solver, pipe = _branched_solver(custom_pipe=True) _damage_flows_after_seed(solver, pipe, ("port_1",)) diagnostics = solver.solve() self.assertFalse(solver.equation_blocks_are_trusted) self.assertEqual( solver.equation_blocks_fallback_reason, "untrustedCustomComponent", ) self.assertEqual(diagnostics.jacobian_mode, "dense") self.assertTrue(diagnostics.block_fallback_used) self.assertEqual( diagnostics.block_fallback_reason, "untrustedCustomComponent", ) self.assertFalse(solver.causal_fast_path_eligible) self.assertEqual( solver.causal_execution_diagnostics()["fallbackReason"], "untrustedCustomComponent", ) def test_active_contact_conservatively_keeps_global_dense_fallback(self) -> None: solver, pipe = _branched_solver(include_contact=True) _damage_flows_after_seed(solver, pipe, ("port_1",)) diagnostics = solver.solve() self.assertEqual(diagnostics.jacobian_mode, "dense") self.assertTrue(diagnostics.block_fallback_used) self.assertEqual( diagnostics.block_fallback_reason, "activeCausalContact", ) contact = solver.network.components["contact"] self.assertIsNotNone(contact._causal_penetration) self.assertAlmostEqual(contact.contact_force, 40.0, places=7) self.assertFalse(solver.causal_fast_path_eligible) self.assertEqual( solver.causal_execution_diagnostics()["fallbackReason"], "activeSetCausalizationRequired", ) if __name__ == "__main__": unittest.main()