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 ( IntegrationCancelled, SolveIVPConfig, StateTransition, integrate_ode, ) def _counting_fixed_step_solver( step_size: float, dense_output_times: list[float], ): import numpy as np class FixedStepSolver: 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 = 0 self.njev = 0 self.nlu = 0 def step(self): self.t = min(self.t + step_size, self.t_bound) self.y = np.asarray([self.t], dtype=float) if self.t >= self.t_bound: self.status = "finished" return None def dense_output(self): dense_output_times.append(self.t) return lambda time: np.asarray([float(time)], dtype=float) return FixedStepSolver class IntegrateOdeTests(unittest.TestCase): def test_generic_solver_keeps_canonical_default_tolerance(self) -> None: self.assertEqual(SolveIVPConfig().atol, 1.0e-8) def test_stepwise_jacobian_cancellation_returns_cancelled_solution(self) -> None: class CancellingJacobian: def start_segment(self) -> None: pass def observe(self, *_args) -> None: pass def __call__(self, _time, _state): raise IntegrationCancelled result = integrate_ode( rhs=lambda _time, state: [-float(state[0])], initial_state=[1.0], config=SolveIVPConfig(t_stop=1.0, method="BDF"), cancel_check=lambda: False, jac=CancellingJacobian(), ) self.assertFalse(result.success) self.assertEqual(result.status, "cancelled") self.assertEqual(result.t, [0.0]) def test_stepwise_explicit_solver_ignores_callable_jacobian(self) -> None: class UnexpectedJacobian: call_count = 0 def start_segment(self) -> None: self.call_count += 1 def observe(self, *_args) -> None: self.call_count += 1 def __call__(self, _time, _state): self.call_count += 1 raise AssertionError("Explicit solver evaluated Jacobian") jacobian = UnexpectedJacobian() result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig(t_stop=0.01, method="RK45"), state_transition_handler=lambda *_args: None, jac=jacobian, ) self.assertTrue(result.success, result.message) self.assertEqual(jacobian.call_count, 0) 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_direct_implicit_solver_receives_callable_jacobian(self) -> None: calls: list[dict[str, object]] = [] observed: list[tuple[float, list[float], list[float]]] = [] class RecordingJacobian: def __init__(self) -> None: self.segment_count = 0 def start_segment(self) -> None: self.segment_count += 1 def observe(self, time, state, derivative) -> None: observed.append( ( float(time), [float(value) for value in state], [float(value) for value in derivative], ) ) def __call__(self, _time, _state): return [[1.0]] def fake_solve_ivp(**kwargs): calls.append(kwargs) kwargs["fun"](0.0, [2.0]) 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 jacobian = RecordingJacobian() with patch.dict( sys.modules, {"scipy": scipy_module, "scipy.integrate": integrate_module}, ): integrate_ode( rhs=lambda _time, state: [2.0 * state[0]], initial_state=[1.0], config=SolveIVPConfig(t_start=0.0, t_stop=1.0, method="BDF"), jac_sparsity=[[True]], jac=jacobian, ) self.assertIs(calls[0]["jac"], jacobian) self.assertNotIn("jac_sparsity", calls[0]) self.assertEqual(jacobian.segment_count, 1) self.assertEqual(observed, [(0.0, [2.0], [4.0])]) 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_stepwise_solvers_build_dense_output_only_when_crossing_a_sample( self, ) -> None: import scipy.integrate for method in ("BDF", "Radau", "RK45"): with self.subTest(method=method): dense_output_times: list[float] = [] fixed_step_solver = _counting_fixed_step_solver( 0.2, dense_output_times, ) with patch.object(scipy.integrate, method, fixed_step_solver): result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method=method, max_step=1.0, ), t_eval=[0.0, 0.75, 1.0], cancel_check=lambda: False, ) self.assertTrue(result.success, result.message) self.assertEqual(result.t, [0.0, 0.75, 1.0]) self.assertEqual(result.y, [[0.0, 0.75, 1.0]]) self.assertEqual(dense_output_times, [0.8, 1.0]) self.assertEqual( result.solver_segments[0].accepted_step_count, 5, ) def test_stepwise_solvers_keep_dense_output_for_state_event_detection( self, ) -> None: import scipy.integrate for method in ("BDF", "Radau", "RK45"): with self.subTest(method=method): dense_output_times: list[float] = [] inspected_steps = 0 fixed_step_solver = _counting_fixed_step_solver( 0.25, dense_output_times, ) def inspect_state_event( previous_time, _previous_state, current_time, _current_state, dense_state, ): nonlocal inspected_steps inspected_steps += 1 midpoint = 0.5 * (previous_time + current_time) self.assertAlmostEqual(dense_state(midpoint)[0], midpoint) return None with patch.object(scipy.integrate, method, fixed_step_solver): result = integrate_ode( rhs=lambda _time, _state: [1.0], initial_state=[0.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method=method, max_step=1.0, ), t_eval=[0.0, 1.0], state_transition_handler=inspect_state_event, ) self.assertTrue(result.success, result.message) self.assertEqual(len(dense_output_times), 4) self.assertEqual(inspected_steps, 4) def test_stepwise_dense_output_preserves_adjacent_float_samples(self) -> None: import scipy.integrate dense_output_times: list[float] = [] adjacent_time = math.nextafter(0.5, math.inf) half_interval_solver = _counting_fixed_step_solver( 0.5, dense_output_times, ) with patch.object(scipy.integrate, "RK45", half_interval_solver): 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=1.0, ), t_eval=[0.0, 0.5, adjacent_time, 1.0], cancel_check=lambda: False, ) self.assertTrue(result.success, result.message) self.assertEqual(result.t, [0.0, 0.5, adjacent_time, 1.0]) self.assertEqual(result.y[0], result.t) self.assertEqual(dense_output_times, [0.5, 1.0]) def test_dense_output_pruning_preserves_cancelled_partial_samples(self) -> None: import scipy.integrate cancellation_requested = False dense_output_times: list[float] = [] fixed_step_solver = _counting_fixed_step_solver( 0.2, dense_output_times, ) def request_cancel(time: float) -> None: nonlocal cancellation_requested cancellation_requested = time >= 0.4 with patch.object(scipy.integrate, "BDF", fixed_step_solver): 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=1.0, ), t_eval=[0.0, 0.3, 0.8, 1.0], cancel_check=lambda: cancellation_requested, accepted_step_callback=request_cancel, ) self.assertFalse(result.success) self.assertEqual(result.status, "cancelled") self.assertEqual(result.t, [0.0, 0.3, 0.4]) self.assertEqual(result.y, [[0.0, 0.3, 0.4]]) self.assertEqual(dense_output_times, [0.4]) 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_stepwise_implicit_solver_restarts_callable_jacobian(self) -> None: import numpy as np import scipy.integrate for method in ("BDF", "Radau"): with self.subTest(method=method): received: list[dict[str, object]] = [] class RecordingJacobian: def __init__(self) -> None: self.segment_count = 0 self.observation_count = 0 def start_segment(self) -> None: self.segment_count += 1 def observe(self, _time, _state, _derivative) -> None: self.observation_count += 1 def __call__(self, _time, _state): return [[0.0]] class RecordingSolver: def __init__(self, fun, t0, y0, t_bound, **kwargs): received.append(kwargs) self.fun = fun self.t = float(t0) self.y = np.asarray(y0, dtype=float) self.t_bound = float(t_bound) self.status = "running" self.nfev = 0 self.njev = 0 self.nlu = 0 def step(self): self.fun(self.t, self.y) self.t = self.t_bound self.status = "finished" return None def dense_output(self): state = self.y.copy() return lambda _time: state.copy() jacobian = RecordingJacobian() with patch.object(scipy.integrate, method, RecordingSolver): result = integrate_ode( rhs=lambda _time, _state: [0.0], initial_state=[1.0], config=SolveIVPConfig( t_start=0.0, t_stop=1.0, method=method, max_step=1.0, ), t_eval=[0.0, 0.4, 1.0], breakpoints=[0.4], jac_sparsity=[[True]], jac=jacobian, ) self.assertTrue(result.success, result.message) self.assertEqual(jacobian.segment_count, 2) self.assertEqual(jacobian.observation_count, 2) self.assertEqual(len(received), 2) self.assertTrue( all(options.get("jac") is jacobian for options in received) ) self.assertTrue( all("jac_sparsity" not in options for options in received) ) 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()