84 lines
2.6 KiB
Python
84 lines
2.6 KiB
Python
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])
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|