from __future__ import annotations from dataclasses import dataclass from typing import Callable, Literal CancellationCheck = Callable[[], bool] AcceptedStepCallback = Callable[[float], None] IntegrationStatus = Literal["completed", "cancelled", "failed"] class _IntegrationCancelled(Exception): pass @dataclass(frozen=True) class SolveIVPConfig: t_start: float = 0.0 t_stop: float = 20.0 method: str = "BDF" rtol: float = 1e-6 atol: float = 1e-8 max_step: float = 1e-3 @dataclass(frozen=True) class ODESolution: t: list[float] y: list[list[float]] success: bool message: str status: IntegrationStatus = "completed" error: Exception | None = None def _vector_add(a: list[float], b: list[float], scale: float = 1.0) -> list[float]: return [x + scale * y for x, y in zip(a, b)] def _append_solution_sample( times: list[float], states: list[list[float]], time: float, state: list[float], ) -> None: if times and time <= times[-1] + 1e-12: return times.append(float(time)) for index, value in enumerate(state): states[index].append(float(value)) def _runge_kutta_4( rhs: Callable[[float, list[float]], list[float]], initial_state: list[float], config: SolveIVPConfig, t_eval: list[float] | None, cancel_check: CancellationCheck | None = None, accepted_step_callback: AcceptedStepCallback | None = None, ) -> ODESolution: if t_eval is None: point_count = max( 2, int((config.t_stop - config.t_start) / max(config.max_step, 1e-6)) + 1, ) step = (config.t_stop - config.t_start) / (point_count - 1) t_eval = [config.t_start + index * step for index in range(point_count)] state = list(initial_state) states = [[value] for value in state] times = [float(t_eval[0])] current_time = float(t_eval[0]) status: IntegrationStatus = "completed" message = "Integrated with built-in RK4 fallback because SciPy is unavailable." error: Exception | None = None try: for target_time in t_eval[1:]: while current_time < target_time - 1e-15: if cancel_check is not None and cancel_check(): raise _IntegrationCancelled dt = min(config.max_step, target_time - current_time) k1 = rhs(current_time, state) k2 = rhs(current_time + 0.5 * dt, _vector_add(state, k1, 0.5 * dt)) k3 = rhs(current_time + 0.5 * dt, _vector_add(state, k2, 0.5 * dt)) k4 = rhs(current_time + dt, _vector_add(state, k3, dt)) state = [ value + (dt / 6.0) * (a + 2.0 * b + 2.0 * c + d) for value, a, b, c, d in zip(state, k1, k2, k3, k4) ] current_time += dt if accepted_step_callback is not None: accepted_step_callback(current_time) _append_solution_sample(times, states, target_time, state) except _IntegrationCancelled: status = "cancelled" message = "Simulation was stopped before reaching the requested end time." _append_solution_sample(times, states, current_time, state) except Exception as exc: status = "failed" message = str(exc) error = exc _append_solution_sample(times, states, current_time, state) return ODESolution( t=times, y=states, success=status == "completed", message=message, status=status, error=error, ) def _integrate_scipy_stepwise( rhs: Callable[[float, list[float]], list[float]], initial_state: list[float], config: SolveIVPConfig, t_eval: list[float] | None, cancel_check: CancellationCheck, accepted_step_callback: AcceptedStepCallback | None, ) -> ODESolution: import numpy as np from scipy.integrate import BDF, DOP853, LSODA, RK23, RK45, Radau solver_types = { "BDF": BDF, "DOP853": DOP853, "LSODA": LSODA, "RK23": RK23, "RK45": RK45, "Radau": Radau, } solver_type = solver_types.get(config.method) if solver_type is None: raise ValueError(f"Unsupported integration method: {config.method}") times = [float(config.t_start)] states = [[float(value)] for value in initial_state] last_accepted_time = float(config.t_start) last_accepted_state = [float(value) for value in initial_state] sample_times = list(t_eval or []) sample_index = 0 while ( sample_index < len(sample_times) and sample_times[sample_index] <= config.t_start + 1e-12 ): sample_index += 1 def cancellable_rhs(time, state): if cancel_check(): raise _IntegrationCancelled return rhs(float(time), [float(value) for value in state]) if cancel_check(): return ODESolution( t=times, y=states, success=False, message="Simulation was stopped before integration started.", status="cancelled", ) try: solver = solver_type( cancellable_rhs, config.t_start, np.asarray(initial_state, dtype=float), config.t_stop, rtol=config.rtol, atol=config.atol, max_step=config.max_step, ) except _IntegrationCancelled: return ODESolution( t=times, y=states, success=False, message="Simulation was stopped before integration started.", status="cancelled", ) except Exception as exc: return ODESolution( t=times, y=states, success=False, message=str(exc), status="failed", error=exc, ) status: IntegrationStatus = "completed" message = "The solver successfully reached the end of the integration interval." error: Exception | None = None while solver.status == "running": if cancel_check(): status = "cancelled" message = "Simulation was stopped before reaching the requested end time." break try: step_message = solver.step() except _IntegrationCancelled: status = "cancelled" message = "Simulation was stopped before reaching the requested end time." break except Exception as exc: status = "failed" message = str(exc) error = exc break if solver.status == "failed": status = "failed" message = str(step_message or "Integration step failed.") break last_accepted_time = float(solver.t) last_accepted_state = [float(value) for value in solver.y] if sample_times: dense_output = solver.dense_output() while ( sample_index < len(sample_times) and sample_times[sample_index] <= last_accepted_time + 1e-12 ): sample_time = float(sample_times[sample_index]) sample_state = [float(value) for value in dense_output(sample_time)] _append_solution_sample(times, states, sample_time, sample_state) sample_index += 1 else: _append_solution_sample( times, states, last_accepted_time, last_accepted_state, ) if accepted_step_callback is not None: accepted_step_callback(last_accepted_time) if status != "completed": _append_solution_sample( times, states, last_accepted_time, last_accepted_state, ) return ODESolution( t=times, y=states, success=status == "completed", message=message, status=status, error=error, ) def integrate_ode( rhs: Callable[[float, list[float]], list[float]], initial_state: list[float], config: SolveIVPConfig, t_eval: list[float] | None = None, cancel_check: CancellationCheck | None = None, accepted_step_callback: AcceptedStepCallback | None = None, ): """Thin wrapper around scipy.integrate.solve_ivp with a pure-Python fallback.""" if abs(config.t_stop - config.t_start) <= 1e-15: return ODESolution( t=[float(config.t_start)], y=[[value] for value in initial_state], success=True, message="Skipped integration because t_start equals t_stop.", ) try: from scipy.integrate import solve_ivp except ImportError: return _runge_kutta_4( rhs, initial_state, config, t_eval, cancel_check, accepted_step_callback, ) if cancel_check is not None: return _integrate_scipy_stepwise( rhs, initial_state, config, t_eval, cancel_check, accepted_step_callback, ) return solve_ivp( fun=rhs, t_span=(config.t_start, config.t_stop), y0=initial_state, method=config.method, rtol=config.rtol, atol=config.atol, max_step=config.max_step, t_eval=t_eval, )