import math import sys import types import unittest from unittest.mock import patch from app.simulation.solvers.solver import SolveIVPConfig, 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_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_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))) if __name__ == "__main__": unittest.main()