from __future__ import annotations import os import unittest from unittest.mock import patch from app.main import compile_reactflow_network from app.simulation.solvers.algebraic import ( CAUSAL_FAST_PATH_ENVIRONMENT_VARIABLE, ) from app.simulation.systems.generic import GenericFluidSystem from tests.test_amesim_mechanical_xml import zero_force_mass_project from tests.test_amesim_pnvo001_signal_xml import ( high_pressure_helium_step_project, ) from tests.test_generic_system_xml_simulation import chain_project def _system(project) -> GenericFluidSystem: return GenericFluidSystem(compile_reactflow_network(project)) class PressureFlowCausalExecutionTests(unittest.TestCase): def test_strict_causal_rhs_matches_environment_disabled_legacy_bitwise( self, ) -> None: optimized = _system(high_pressure_helium_step_project()) with patch.dict( os.environ, {CAUSAL_FAST_PATH_ENVIRONMENT_VARIABLE: "0"}, ): legacy = _system(high_pressure_helium_step_project()) optimized_state = optimized.initial_state_vector() legacy_state = legacy.initial_state_vector() for time in (0.0, 0.041, 0.8): optimized_derivative = optimized.rhs(time, optimized_state) legacy_derivative = legacy.rhs(time, legacy_state) self.assertEqual(optimized_derivative, legacy_derivative) self.assertEqual( tuple( unknown.read() for unknown in optimized.pressure_flow_solver.unknowns ), tuple( unknown.read() for unknown in legacy.pressure_flow_solver.unknowns ), ) global_diagnostics = ( optimized.pressure_flow_solver.causal_execution_diagnostics() ) self.assertTrue(global_diagnostics["eligible"]) self.assertGreater(global_diagnostics["fastSolveCount"], 0) self.assertGreaterEqual( global_diagnostics["fullResidualAuditCount"], 1, ) secondary = ( optimized._thermofluid_closure_plan.secondary_block_solvers[0] ) secondary_diagnostics = secondary.causal_execution_diagnostics() self.assertTrue(secondary_diagnostics["eligible"]) self.assertGreater(secondary_diagnostics["fastSolveCount"], 0) legacy_diagnostics = ( legacy.pressure_flow_solver.causal_execution_diagnostics() ) self.assertFalse(legacy_diagnostics["enabled"]) self.assertEqual( legacy_diagnostics["disabledReason"], "disabledByEnvironment", ) self.assertEqual(legacy_diagnostics["fastSolveCount"], 0) def test_multiple_effort_anchors_conservatively_keep_legacy_path(self) -> None: system = _system(chain_project()) diagnostics = ( system.pressure_flow_solver.causal_execution_diagnostics() ) self.assertFalse(diagnostics["eligible"]) self.assertFalse(diagnostics["enabled"]) self.assertEqual( diagnostics["fallbackReason"], "effortGroupDoesNotHaveOneAnchor", ) def test_runtime_flow_coverage_failure_fuses_to_verified_legacy_path( self, ) -> None: system = _system(zero_force_mass_project()) state = system.initial_state_vector() system.rhs(0.0, state) solver = system.pressure_flow_solver original = solver._solve_explicit_flow_unknowns def hide_coverage(*args, **kwargs): original(*args, **kwargs) return set() with patch.object( solver, "_solve_explicit_flow_unknowns", side_effect=hide_coverage, ): diagnostics = solver.solve(effort_variables=("p",)) self.assertTrue(diagnostics.success) self.assertTrue(diagnostics.residual_verified_this_solve) execution = solver.causal_execution_diagnostics() self.assertFalse(execution["enabled"]) self.assertEqual( execution["disabledReason"], "causalRuntimeGateFailed", ) self.assertEqual(execution["legacyFallbackCount"], 1) def test_periodic_audit_failure_disables_fast_path_before_fallback(self) -> None: system = _system(zero_force_mass_project()) state = system.initial_state_vector() system.rhs(0.0, state) solver = system.pressure_flow_solver solver._causal_audit_interval = 0 original_values = solver._pressure_flow_equation_values call_count = 0 def one_bad_audit_value(): nonlocal call_count call_count += 1 values = original_values() if call_count != 1: return values return (values[0] + 1.0, *values[1:]) with patch.object( solver, "_pressure_flow_equation_values", side_effect=one_bad_audit_value, ): diagnostics = solver.solve(effort_variables=("p",)) self.assertTrue(diagnostics.success) execution = solver.causal_execution_diagnostics() self.assertFalse(execution["enabled"]) self.assertEqual( execution["disabledReason"], "causalResidualAuditFailed", ) self.assertEqual(execution["auditFailureCount"], 1) self.assertEqual(execution["legacyFallbackCount"], 1) def test_requested_audit_interrupts_periodic_fast_sequence(self) -> None: system = _system(zero_force_mass_project()) state = system.initial_state_vector() solver = system.pressure_flow_solver system.rhs(0.0, state) system.rhs(0.0, state) before = solver.causal_execution_diagnostics() self.assertEqual(before["fullResidualAuditCount"], 1) self.assertEqual(before["fastSolveCount"], 1) solver.request_causal_audit() system.rhs(0.0, state) after = solver.causal_execution_diagnostics() self.assertEqual(after["fullResidualAuditCount"], 2) self.assertEqual(after["fastSolveCount"], 1) def test_fast_solve_skips_the_full_residual_evaluator(self) -> None: system = _system(zero_force_mass_project()) state = system.initial_state_vector() solver = system.pressure_flow_solver system.rhs(0.0, state) with patch.object( solver, "_pressure_flow_equation_values", wraps=solver._pressure_flow_equation_values, ) as evaluate_all: system.rhs(0.0, state) evaluate_all.assert_not_called() self.assertTrue(solver.last_diagnostics.causal_fast_path_used) self.assertFalse( solver.last_diagnostics.residual_verified_this_solve ) def test_nonfinite_explicit_assignment_fuses_and_verifies_same_solve( self, ) -> None: system = _system(zero_force_mass_project()) state = system.initial_state_vector() system.rhs(0.0, state) solver = system.pressure_flow_solver original = solver._evaluate_explicit_flow_stage call_count = 0 def one_nonfinite_assignment(stage): nonlocal call_count call_count += 1 values = original(stage) if call_count != 1: return values return (float("nan"), *values[1:]) with patch.object( solver, "_evaluate_explicit_flow_stage", side_effect=one_nonfinite_assignment, ): diagnostics = solver.solve(effort_variables=("p",)) self.assertTrue(diagnostics.success) self.assertTrue(diagnostics.residual_verified_this_solve) execution = solver.causal_execution_diagnostics() self.assertFalse(execution["enabled"]) self.assertEqual( execution["disabledReason"], "causalRuntimeGateFailed", ) self.assertEqual(execution["legacyFallbackCount"], 1) def test_nonpositive_pressure_fuses_and_verifies_same_solve(self) -> None: system = _system(high_pressure_helium_step_project()) state = system.initial_state_vector() system.rhs(0.0, state) solver = system.pressure_flow_solver original = solver._solve_explicit_flow_unknowns pressure = next( unknown for unknown in solver.unknowns if unknown.variable == "p" ) def make_pressure_invalid(*args, **kwargs): seeded = original(*args, **kwargs) pressure.write(-1.0) return seeded with patch.object( solver, "_solve_explicit_flow_unknowns", side_effect=make_pressure_invalid, ): diagnostics = solver.solve(effort_variables=("p",)) self.assertTrue(diagnostics.success) self.assertTrue(diagnostics.residual_verified_this_solve) self.assertGreater(pressure.read(), 0.0) execution = solver.causal_execution_diagnostics() self.assertFalse(execution["enabled"]) self.assertEqual( execution["disabledReason"], "causalRuntimeGateFailed", ) self.assertEqual(execution["legacyFallbackCount"], 1) def test_default_periodic_audit_runs_after_sixty_four_fast_solves(self) -> None: system = _system(zero_force_mass_project()) state = system.initial_state_vector() solver = system.pressure_flow_solver system.rhs(0.0, state) for _iteration in range(64): system.rhs(0.0, state) before_boundary = solver.causal_execution_diagnostics() self.assertEqual(before_boundary["auditInterval"], 64) self.assertEqual(before_boundary["fastSolveCount"], 64) self.assertEqual(before_boundary["fullResidualAuditCount"], 1) system.rhs(0.0, state) after_boundary = solver.causal_execution_diagnostics() self.assertEqual(after_boundary["fastSolveCount"], 64) self.assertEqual(after_boundary["fullResidualAuditCount"], 2) if __name__ == "__main__": unittest.main()