import math import sys import types import unittest from unittest.mock import patch from app.simulation.core.errors import RecoverableTrialStateError from app.simulation.solvers.solver import ( SolveIVPConfig, StateTransition, integrate_ode, ) class IntegrateOdeTests(unittest.TestCase): def test_generic_solver_keeps_canonical_default_tolerance(self) -> None: self.assertEqual(SolveIVPConfig().atol, 1.0e-8) def test_scipy_solver_receives_step_size_controls(self) -> None: calls: list[dict[str, object]] = [] def fake_solve_ivp(**kwargs): calls.append(kwargs) return object() scipy_module = types.ModuleType("scipy") integrate_module = types.ModuleType("scipy.integrate") integrate_module.solve_ivp = fake_solve_ivp scipy_module.integrate = integrate_module with patch.dict( sys.modules, {"scipy": scipy_module, "scipy.integrate": integrate_module}, ): integrate_ode( rhs=lambda _time, state: state, initial_state=[1.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, max_step=1.0e-4, first_step=1.0e-8, ), t_eval=[0.0, 1.0], ) self.assertEqual(calls[0]["max_step"], 1.0e-4) self.assertEqual(calls[0]["first_step"], 1.0e-8) def test_scipy_solver_omits_unset_first_step(self) -> None: calls: list[dict[str, object]] = [] def fake_solve_ivp(**kwargs): calls.append(kwargs) return object() scipy_module = types.ModuleType("scipy") integrate_module = types.ModuleType("scipy.integrate") integrate_module.solve_ivp = fake_solve_ivp scipy_module.integrate = integrate_module with patch.dict( sys.modules, {"scipy": scipy_module, "scipy.integrate": integrate_module}, ): integrate_ode( rhs=lambda _time, state: state, initial_state=[1.0], config=SolveIVPConfig(t_start=0.0, t_stop=1.0), ) self.assertEqual(calls[0]["max_step"], 1.0e-3) self.assertNotIn("first_step", calls[0]) def test_stepwise_solver_rebuilds_after_recoverable_trial_failure(self) -> None: import numpy as np import scipy.integrate attempted_max_steps: list[float] = [] class RetryBdf: def __init__(self, fun, t0, y0, t_bound, **kwargs): self.fun = fun self.t = float(t0) self.y = np.asarray(y0, dtype=float) self.t_bound = float(t_bound) self.h_abs = float(kwargs["max_step"]) self.status = "running" attempted_max_steps.append(self.h_abs) def step(self): if self.h_abs > 0.25: raise RecoverableTrialStateError("trial state outside domain") self.t = self.t_bound self.status = "finished" return None def dense_output(self): state = self.y.copy() return lambda _time: state.copy() with patch.object(scipy.integrate, "BDF", RetryBdf): result = integrate_ode( rhs=lambda _time, _state: [0.0], initial_state=[1.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="BDF", max_step=1.0, ), t_eval=[0.0, 1.0], cancel_check=lambda: False, ) self.assertTrue(result.success, result.message) self.assertEqual(attempted_max_steps, [1.0, 0.5, 0.25]) self.assertEqual(result.t, [0.0, 1.0]) self.assertEqual(result.y, [[1.0, 1.0]]) def test_stepwise_solver_does_not_retry_ordinary_model_errors(self) -> None: import numpy as np import scipy.integrate attempts = 0 class FailingBdf: def __init__(self, fun, t0, y0, t_bound, **kwargs): nonlocal attempts attempts += 1 self.t = float(t0) self.y = np.asarray(y0, dtype=float) self.status = "running" def step(self): raise ValueError("structural model error") with patch.object(scipy.integrate, "BDF", FailingBdf): result = integrate_ode( rhs=lambda _time, _state: [0.0], initial_state=[1.0], config=SolveIVPConfig(t_start=0.0, t_stop=1.0, method="BDF"), cancel_check=lambda: False, ) self.assertFalse(result.success) self.assertEqual(result.message, "structural model error") self.assertEqual(attempts, 1) def test_scipy_stepwise_solver_can_cancel_before_start(self) -> None: result = integrate_ode( rhs=lambda _time, state: state, initial_state=[1.0], config=SolveIVPConfig(t_start=0.0, t_stop=1.0), cancel_check=lambda: True, ) self.assertFalse(result.success) self.assertEqual(result.status, "cancelled") self.assertEqual(result.t, [0.0]) def test_segmented_bdf_uses_left_limit_and_restarts_at_event(self) -> None: import scipy.integrate event_time = 0.5 actual_bdf = scipy.integrate.BDF starts: list[float] = [] bounds: list[float] = [] call_times: list[list[float]] = [] class RecordingBDF(actual_bdf): def __init__(self, fun, t0, y0, t_bound, **kwargs): starts.append(float(t0)) bounds.append(float(t_bound)) segment_calls: list[float] = [] call_times.append(segment_calls) def recording_fun(time, state): segment_calls.append(float(time)) return fun(time, state) super().__init__(recording_fun, t0, y0, t_bound, **kwargs) with patch.object(scipy.integrate, "BDF", RecordingBDF): result = integrate_ode( rhs=lambda time, _state: [1.0 if time < event_time else 2.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="BDF", max_step=0.1, first_step=0.8, ), t_eval=[0.0, event_time, event_time, 1.0], breakpoints=[event_time], ) self.assertTrue(result.success, result.message) self.assertEqual(starts, [0.0, event_time]) self.assertEqual(bounds[0], math.nextafter(event_time, -math.inf)) self.assertEqual(bounds[1], 1.0) self.assertTrue(call_times[0]) self.assertTrue(all(time < event_time for time in call_times[0])) self.assertTrue(any(time >= event_time for time in call_times[1])) self.assertEqual(result.t, [0.0, event_time, 1.0]) self.assertAlmostEqual(result.y[0][-1], 1.5, places=5) def test_segmented_implicit_solvers_merge_samples_and_report_progress(self) -> None: event_time = 0.4 for method in ("BDF", "Radau"): with self.subTest(method=method): callback_times: list[float] = [] result = integrate_ode( rhs=lambda time, _state: [1.0 if time < event_time else 3.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method=method, max_step=0.05, first_step=0.9, ), t_eval=[0.0, event_time, event_time, 0.7, 1.0], accepted_step_callback=callback_times.append, breakpoints=[event_time, event_time], ) self.assertTrue(result.success, result.message) self.assertEqual(result.t, [0.0, event_time, 0.7, 1.0]) self.assertAlmostEqual(result.y[0][-1], 2.2, places=5) self.assertEqual(callback_times.count(event_time), 1) self.assertTrue( all( earlier < later for earlier, later in zip( callback_times, callback_times[1:], ) ) ) def test_segmented_solver_reports_implicit_work_by_event_segment(self) -> None: import numpy as np import scipy.integrate class CountingBDF: def __init__(self, _fun, t0, y0, t_bound, **_kwargs): self.t = float(t0) self.y = np.asarray(y0, dtype=float) self.t_bound = float(t_bound) self.status = "running" self.nfev = 2 self.njev = 1 self.nlu = 0 def step(self): self.t = self.t_bound self.nfev += 3 self.nlu += 2 self.status = "finished" return None def dense_output(self): state = self.y.copy() return lambda _time: state.copy() with patch.object(scipy.integrate, "BDF", CountingBDF): result = integrate_ode( rhs=lambda _time, _state: [0.0], initial_state=[1.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="BDF", max_step=1.0, ), t_eval=[0.0, 0.4, 1.0], breakpoints=[0.4], ) self.assertTrue(result.success, result.message) self.assertEqual(len(result.solver_segments), 2) self.assertEqual( [segment.as_dict() for segment in result.solver_segments], [ { "startTime": 0.0, "requestedStopTime": 0.4, "simulatedUntil": 0.4, "nfev": 5, "njev": 1, "nlu": 2, "acceptedStepCount": 1, "solverStartCount": 1, "stateTransitionCount": 0, "recoverableRetryCount": 0, }, { "startTime": 0.4, "requestedStopTime": 1.0, "simulatedUntil": 1.0, "nfev": 5, "njev": 1, "nlu": 2, "acceptedStepCount": 1, "solverStartCount": 1, "stateTransitionCount": 0, "recoverableRetryCount": 0, }, ], ) def test_segmented_solver_can_cancel_after_crossing_a_breakpoint(self) -> None: callback_times: list[float] = [] cancellation_requested = False def record_progress(time: float) -> None: nonlocal cancellation_requested callback_times.append(time) cancellation_requested = time >= 0.55 result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="BDF", max_step=0.05, ), t_eval=[0.0, 0.2, 0.4, 0.6, 0.8, 1.0], cancel_check=lambda: cancellation_requested, accepted_step_callback=record_progress, breakpoints=[0.4, 0.8], ) self.assertFalse(result.success) self.assertEqual(result.status, "cancelled") self.assertIn(0.4, callback_times) self.assertGreater(callback_times[-1], 0.4) self.assertTrue( all( earlier < later for earlier, later in zip(callback_times, callback_times[1:]) ) ) self.assertEqual(result.t, sorted(set(result.t))) def test_state_transition_resets_at_root_and_discards_step_overshoot(self) -> None: event_time = 0.35 event_enabled = True def transition_handler( previous_time, previous_state, current_time, current_state, dense_state, ): nonlocal event_enabled if ( not event_enabled or previous_state[0] >= event_time or current_state[0] < event_time ): return None lower = previous_time upper = current_time for _iteration in range(60): middle = 0.5 * (lower + upper) if dense_state(middle)[0] >= event_time: upper = middle else: lower = middle event_enabled = False return StateTransition(time=upper, state=[0.0]) result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="BDF", max_step=0.5, ), t_eval=[0.0, event_time, 0.4, 1.0], state_transition_handler=transition_handler, ) self.assertTrue(result.success, result.message) self.assertEqual(result.t, [0.0, event_time, 0.4, 1.0]) self.assertAlmostEqual(result.y[0][1], 0.0, places=12) self.assertAlmostEqual(result.y[0][2], 0.05, places=8) self.assertAlmostEqual(result.y[0][-1], 0.65, places=8) def test_state_transitions_chain_at_same_time_until_state_repeats(self) -> None: event_time = 0.25 stage = 0 returned_reset_states: list[float] = [] def transition_handler( previous_time, _previous_state, current_time, _current_state, _dense_state, ): nonlocal stage if stage == 0 and previous_time <= event_time <= current_time: stage = 1 returned_reset_states.append(10.0) return StateTransition(time=event_time, state=[10.0]) if stage == 1 and previous_time == event_time: stage = 2 returned_reset_states.append(20.0) return StateTransition(time=event_time, state=[20.0]) if stage == 2 and previous_time == event_time: returned_reset_states.append(20.0) return StateTransition(time=event_time, state=[20.0]) return None result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="BDF", max_step=0.4, ), t_eval=[0.0, event_time, 1.0], state_transition_handler=transition_handler, ) self.assertTrue(result.success, result.message) self.assertEqual(returned_reset_states, [10.0, 20.0, 20.0]) self.assertEqual(result.t, [0.0, event_time, 1.0]) self.assertEqual(result.y[0][1], 20.0) self.assertAlmostEqual(result.y[0][-1], 20.75, places=8) def test_state_transition_chain_has_a_finite_guard(self) -> None: event_time = 0.25 reset_count = 0 def transition_handler( previous_time, _previous_state, current_time, _current_state, _dense_state, ): nonlocal reset_count if previous_time <= event_time <= current_time: reset_count += 1 return StateTransition( time=event_time, state=[float(reset_count)], ) return None result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="BDF", max_step=0.4, ), state_transition_handler=transition_handler, ) self.assertFalse(result.success) self.assertEqual(result.status, "failed") self.assertIn("64 chained resets", result.message) def test_stepwise_solver_preserves_adjacent_float_samples(self) -> None: adjacent_time = math.nextafter(0.5, math.inf) result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="RK45", max_step=0.4, ), t_eval=[0.0, 0.5, adjacent_time, 1.0], state_transition_handler=lambda *_args: None, ) self.assertTrue(result.success, result.message) self.assertEqual(result.t, [0.0, 0.5, adjacent_time, 1.0]) def test_state_transition_at_breakpoint_uses_exact_breakpoint_sample(self) -> None: event_time = 0.5 integration_left_limit = math.nextafter(event_time, -math.inf) event_enabled = True def transition_handler( previous_time, _previous_state, current_time, _current_state, _dense_state, ): nonlocal event_enabled if ( event_enabled and previous_time <= integration_left_limit <= current_time ): event_enabled = False return StateTransition(time=event_time, state=[7.0]) return None result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="BDF", max_step=0.2, ), t_eval=[0.0, event_time, 1.0], breakpoints=[event_time], state_transition_handler=transition_handler, ) self.assertTrue(result.success, result.message) self.assertEqual(result.t, [0.0, event_time, 1.0]) self.assertEqual(result.y[0][1], 7.0) self.assertAlmostEqual(result.y[0][-1], 7.5, places=8) def test_cancellation_after_state_transition_reports_partial_progress(self) -> None: event_time = 0.25 cancellation_requested = False event_enabled = True def transition_handler( previous_time, _previous_state, current_time, _current_state, _dense_state, ): nonlocal cancellation_requested, event_enabled if event_enabled and previous_time <= event_time <= current_time: event_enabled = False cancellation_requested = True return StateTransition(time=event_time, state=[0.0]) return None result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method="BDF", max_step=0.4, ), cancel_check=lambda: cancellation_requested, state_transition_handler=transition_handler, ) self.assertFalse(result.success) self.assertEqual(result.status, "cancelled") self.assertEqual( result.message, "Simulation was stopped before reaching the requested end time.", ) self.assertEqual(result.t[-1], event_time) def test_state_transition_handler_rejects_reverse_integration(self) -> None: with self.assertRaisesRegex( ValueError, "does not support reverse integration", ): integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig(t_start=1.0, t_stop=0.0), state_transition_handler=lambda *_args: None, ) if __name__ == "__main__": unittest.main()