from __future__ import annotations import math from dataclasses import dataclass from typing import Callable, Literal, Sequence CancellationCheck = Callable[[], bool] AcceptedStepCallback = Callable[[float], None] IntegrationStatus = Literal["completed", "cancelled", "failed"] DenseState = Callable[[float], list[float]] @dataclass(frozen=True) class StateTransition: """A state reset located inside an accepted integration step.""" time: float state: list[float] StateTransitionHandler = Callable[ [float, list[float], float, list[float], DenseState], StateTransition | None, ] _MAX_STATE_TRANSITIONS_AT_SAME_TIME = 64 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 first_step: float | None = None @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: time = float(time) if times and time <= times[-1]: return times.append(time) for index, value in enumerate(state): states[index].append(float(value)) def _append_or_replace_solution_sample( times: list[float], states: list[list[float]], time: float, state: list[float], ) -> None: """Store a reset state even when its event time was already sampled.""" time = float(time) if times and time == times[-1]: times[-1] = time for index, value in enumerate(state): states[index][-1] = float(value) return _append_solution_sample(times, states, time, state) def _normalize_state_transition( transition: StateTransition, before_time: float, after_time: float, state_size: int, ) -> StateTransition: """Validate and normalize a transition returned for an accepted step.""" if not isinstance(transition, StateTransition): raise TypeError( "State transition handlers must return StateTransition or None." ) transition_time = float(transition.time) if not math.isfinite(transition_time): raise ValueError("State transition times must be finite numbers.") tolerance = 16.0 * max( math.ulp(before_time), math.ulp(after_time), math.ulp(transition_time), ) if ( transition_time < before_time - tolerance or transition_time > after_time + tolerance ): raise ValueError( "State transition time must lie inside the accepted integration step." ) transition_time = min(max(transition_time, before_time), after_time) transition_state = [float(value) for value in transition.state] if len(transition_state) != state_size: raise ValueError( "State transition reset state must have the same size as the ODE state." ) if not all(math.isfinite(value) for value in transition_state): raise ValueError("State transition reset states must contain finite numbers.") return StateTransition(time=transition_time, state=transition_state) def _is_repeated_state_transition( transition: StateTransition, last_transition: StateTransition | None, ) -> bool: """Suppress only the exact reset that was just applied. A second reset at the same instant is meaningful when it produces a different state (for example, two constraints becoming active together). """ return ( last_transition is not None and transition.time == last_transition.time and transition.state == last_transition.state ) def _next_same_time_transition_count( transition: StateTransition, last_transition: StateTransition | None, previous_count: int, ) -> int: count = ( previous_count + 1 if last_transition is not None and transition.time == last_transition.time else 1 ) if count > _MAX_STATE_TRANSITIONS_AT_SAME_TIME: raise RuntimeError( "State transition handler exceeded " f"{_MAX_STATE_TRANSITIONS_AT_SAME_TIME} chained resets at the same time." ) return count def _align_transition_with_exact_endpoint( transition: StateTransition, requested_time: float, exact_endpoint: float | None, ) -> StateTransition: """Keep an event reported at a breakpoint on that exact public timestamp.""" if exact_endpoint is not None and requested_time == exact_endpoint: return StateTransition( time=float(exact_endpoint), state=list(transition.state), ) return transition def _normalize_breakpoints( config: SolveIVPConfig, breakpoints: Sequence[float] | None, ) -> list[float]: """Return sorted, unique breakpoints strictly inside the integration span.""" if breakpoints is None or len(breakpoints) == 0: return [] if config.t_stop < config.t_start: raise ValueError("Segmented integration requires t_stop to follow t_start.") normalized: list[float] = [] for raw_breakpoint in breakpoints: breakpoint = float(raw_breakpoint) if not math.isfinite(breakpoint): raise ValueError("Integration breakpoints must be finite numbers.") if config.t_start < breakpoint < config.t_stop: normalized.append(breakpoint) normalized.sort() return [ breakpoint for index, breakpoint in enumerate(normalized) if index == 0 or breakpoint != normalized[index - 1] ] 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, state_transition_handler: StateTransitionHandler | 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 last_transition: StateTransition | None = None same_time_transition_count = 0 last_reported_step: float | None = None def report_step(time: float) -> None: nonlocal last_reported_step if accepted_step_callback is None: return if last_reported_step is not None and time <= last_reported_step: return accepted_step_callback(float(time)) last_reported_step = float(time) try: for target_time in t_eval[1:]: while current_time < target_time: 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)) next_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) ] next_time = current_time + dt transition: StateTransition | None = None if state_transition_handler is not None: step_start = current_time step_state = list(state) def dense_state(time: float) -> list[float]: fraction = (float(time) - step_start) / (next_time - step_start) return [ before + fraction * (after - before) for before, after in zip(step_state, next_state) ] candidate = state_transition_handler( step_start, list(step_state), next_time, list(next_state), dense_state, ) if candidate is not None: candidate = _normalize_state_transition( candidate, step_start, next_time, len(state), ) if not _is_repeated_state_transition( candidate, last_transition, ): transition = candidate if transition is not None: same_time_transition_count = _next_same_time_transition_count( transition, last_transition, same_time_transition_count, ) current_time = transition.time state = list(transition.state) last_transition = transition _append_or_replace_solution_sample( times, states, current_time, state, ) else: current_time = next_time state = next_state report_step(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 _runge_kutta_4_segmented( rhs: Callable[[float, list[float]], list[float]], initial_state: list[float], config: SolveIVPConfig, t_eval: list[float] | None, breakpoints: Sequence[float], cancel_check: CancellationCheck | None = None, accepted_step_callback: AcceptedStepCallback | None = None, state_transition_handler: StateTransitionHandler | None = None, ) -> ODESolution: """RK4 fallback that never evaluates a pre-breakpoint step at the breakpoint.""" if t_eval is None: point_count = max( 2, int((config.t_stop - config.t_start) / max(config.max_step, 1e-6)) + 1, ) sample_step = (config.t_stop - config.t_start) / (point_count - 1) sample_times = [ config.t_start + index * sample_step for index in range(point_count) ] else: sample_times = [float(time) for time in t_eval] state = [float(value) for value in initial_state] states = [[value] for value in state] times = [float(config.t_start)] current_time = float(config.t_start) sample_index = 0 while ( sample_index < len(sample_times) and sample_times[sample_index] <= config.t_start ): sample_index += 1 status: IntegrationStatus = "completed" message = "Integrated with built-in RK4 fallback because SciPy is unavailable." error: Exception | None = None last_transition: StateTransition | None = None same_time_transition_count = 0 last_reported_step: float | None = None def report_step(time: float) -> None: nonlocal last_reported_step if accepted_step_callback is None: return if last_reported_step is not None and time <= last_reported_step: return accepted_step_callback(float(time)) last_reported_step = float(time) def advance_to( target_time: float, reported_terminal_time: float | None = None ) -> None: nonlocal current_time, last_transition, same_time_transition_count, state while current_time < target_time: 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)) next_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) ] next_time = current_time + dt transition: StateTransition | None = None if state_transition_handler is not None: step_start = current_time step_state = list(state) def dense_state(time: float) -> list[float]: fraction = (float(time) - step_start) / (next_time - step_start) return [ before + fraction * (after - before) for before, after in zip(step_state, next_state) ] candidate = state_transition_handler( step_start, list(step_state), next_time, list(next_state), dense_state, ) if candidate is not None: requested_time = float(candidate.time) candidate = _normalize_state_transition( candidate, step_start, next_time, len(state), ) candidate = _align_transition_with_exact_endpoint( candidate, requested_time, reported_terminal_time, ) if not _is_repeated_state_transition( candidate, last_transition, ): transition = candidate if transition is not None: same_time_transition_count = _next_same_time_transition_count( transition, last_transition, same_time_transition_count, ) current_time = transition.time state = list(transition.state) last_transition = transition _append_or_replace_solution_sample( times, states, current_time, state, ) else: current_time = next_time state = next_state report_time = current_time if ( reported_terminal_time is not None and current_time >= target_time ): report_time = reported_terminal_time report_step(report_time) try: segment_ends = [*breakpoints, float(config.t_stop)] for segment_index, segment_end in enumerate(segment_ends): is_breakpoint = segment_index < len(breakpoints) integration_end = ( math.nextafter(segment_end, -math.inf) if is_breakpoint else segment_end ) while ( sample_index < len(sample_times) and sample_times[sample_index] <= integration_end ): sample_time = float(sample_times[sample_index]) advance_to(sample_time) _append_solution_sample(times, states, sample_time, state) sample_index += 1 advance_to( integration_end, segment_end if is_breakpoint else None, ) if is_breakpoint: current_time = float(segment_end) report_step(current_time) while ( sample_index < len(sample_times) and sample_times[sample_index] <= segment_end ): sample_time = float(sample_times[sample_index]) _append_solution_sample(times, states, sample_time, state) sample_index += 1 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, breakpoints: Sequence[float] = (), state_transition_handler: StateTransitionHandler | None = None, ) -> ODESolution: """Initial stepwise integration path for breakpoints and state resets. Known V1 limitation: an adaptive solver can evaluate a trial state outside the algebraic or thermodynamic model domain. Such an RHS exception still aborts the run here; recoverable trial failures are not yet restored to the last accepted state and retried with a smaller step. This is not specific to BDF, although implicit Newton/Jacobian probes make it especially visible. """ 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 = [float(time) for time in (t_eval or [])] sample_index = 0 while ( sample_index < len(sample_times) and sample_times[sample_index] <= config.t_start ): sample_index += 1 def cancellable_rhs(time, state): if cancel_check(): raise _IntegrationCancelled return rhs(float(time), [float(value) for value in state]) status: IntegrationStatus = "completed" message = "The solver successfully reached the end of the integration interval." error: Exception | None = None last_transition: StateTransition | None = None same_time_transition_count = 0 integration_progressed = False last_reported_step: float | None = None def cancellation_message() -> str: return ( "Simulation was stopped before reaching the requested end time." if integration_progressed else "Simulation was stopped before integration started." ) def report_step(time: float) -> None: nonlocal last_reported_step if accepted_step_callback is None: return if last_reported_step is not None and time <= last_reported_step: return accepted_step_callback(float(time)) last_reported_step = float(time) segment_ends = [*breakpoints, float(config.t_stop)] for segment_index, segment_end in enumerate(segment_ends): if cancel_check(): status = "cancelled" message = cancellation_message() break is_breakpoint = segment_index < len(breakpoints) integration_end = ( math.nextafter(segment_end, -math.inf) if is_breakpoint else segment_end ) has_integration_interval = integration_end > last_accepted_time while has_integration_interval and last_accepted_time < integration_end: if cancel_check(): status = "cancelled" message = cancellation_message() break solver_options = { "rtol": config.rtol, "atol": config.atol, "max_step": config.max_step, } if config.first_step is not None: solver_options["first_step"] = min( config.first_step, integration_end - last_accepted_time, ) try: solver = solver_type( cancellable_rhs, last_accepted_time, np.asarray(last_accepted_state, dtype=float), integration_end, **solver_options, ) except _IntegrationCancelled: status = "cancelled" message = cancellation_message() break except Exception as exc: status = "failed" message = str(exc) error = exc break restart_at_transition = False while solver.status == "running": if cancel_check(): status = "cancelled" message = ( "Simulation was stopped before reaching the requested end time." ) break step_start_time = last_accepted_time step_start_state = list(last_accepted_state) 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 integration_progressed = True if solver.status == "failed": status = "failed" message = str(step_message or "Integration step failed.") break step_end_time = float(solver.t) step_end_state = [float(value) for value in solver.y] dense_output = ( solver.dense_output() if sample_times or state_transition_handler is not None else None ) transition: StateTransition | None = None if state_transition_handler is not None: assert dense_output is not None def dense_state(time: float) -> list[float]: return [float(value) for value in dense_output(float(time))] try: candidate = state_transition_handler( step_start_time, list(step_start_state), step_end_time, list(step_end_state), dense_state, ) if candidate is not None: requested_time = float(candidate.time) candidate = _normalize_state_transition( candidate, step_start_time, step_end_time, len(last_accepted_state), ) candidate = _align_transition_with_exact_endpoint( candidate, requested_time, float(segment_end) if is_breakpoint else None, ) if not _is_repeated_state_transition( candidate, last_transition, ): transition = candidate except Exception as exc: status = "failed" message = str(exc) error = exc break if transition is not None: try: same_time_transition_count = ( _next_same_time_transition_count( transition, last_transition, same_time_transition_count, ) ) except Exception as exc: status = "failed" message = str(exc) error = exc break while ( sample_index < len(sample_times) and sample_times[sample_index] < transition.time ): sample_time = float(sample_times[sample_index]) assert dense_output is not None sample_state = [ float(value) for value in dense_output(sample_time) ] _append_solution_sample( times, states, sample_time, sample_state, ) sample_index += 1 last_accepted_time = transition.time last_accepted_state = list(transition.state) last_transition = transition _append_or_replace_solution_sample( times, states, last_accepted_time, last_accepted_state, ) while ( sample_index < len(sample_times) and sample_times[sample_index] <= last_accepted_time ): sample_index += 1 report_step(last_accepted_time) restart_at_transition = last_accepted_time < integration_end break last_accepted_time = step_end_time last_accepted_state = step_end_state reported_time = ( float(segment_end) if is_breakpoint and solver.status == "finished" else last_accepted_time ) if sample_times: assert dense_output is not None while ( sample_index < len(sample_times) and sample_times[sample_index] <= last_accepted_time ): 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, reported_time, last_accepted_state, ) report_step(reported_time) if status != "completed" or not restart_at_transition: break if status != "completed": break if is_breakpoint: # The old equation is integrated only to the representable point just # left of the event. The continuous state is then lifted to the exact # event time, where the freshly constructed next solver sees the new # equation immediately. last_accepted_time = float(segment_end) if sample_times: while ( sample_index < len(sample_times) and sample_times[sample_index] <= segment_end ): sample_time = float(sample_times[sample_index]) _append_solution_sample( times, states, sample_time, last_accepted_state, ) sample_index += 1 elif not has_integration_interval: _append_solution_sample( times, states, last_accepted_time, last_accepted_state, ) report_step(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, breakpoints: Sequence[float] | None = None, state_transition_handler: StateTransitionHandler | None = None, ): """Integrate an ODE, optionally restarting at equation discontinuities. Breakpoints are interpreted as right-continuous equation changes: the old equation is integrated to the floating-point left limit, then a fresh solver starts at the exact breakpoint with the unchanged continuous state. A state transition handler inspects every accepted step using its dense interpolant. When it returns a transition, samples before the event retain the pre-event trajectory, the reset state is stored at the event, and a fresh solver continues from that state. """ if ( state_transition_handler is not None and config.t_stop < config.t_start ): raise ValueError( "State transition handling does not support reverse integration." ) if config.t_stop == config.t_start: 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.", ) normalized_breakpoints = _normalize_breakpoints(config, breakpoints) try: from scipy.integrate import solve_ivp except ImportError: if normalized_breakpoints: return _runge_kutta_4_segmented( rhs, initial_state, config, t_eval, normalized_breakpoints, cancel_check, accepted_step_callback, state_transition_handler, ) return _runge_kutta_4( rhs, initial_state, config, t_eval, cancel_check, accepted_step_callback, state_transition_handler, ) if ( cancel_check is not None or normalized_breakpoints or state_transition_handler is not None ): return _integrate_scipy_stepwise( rhs, initial_state, config, t_eval, cancel_check or (lambda: False), accepted_step_callback, normalized_breakpoints, state_transition_handler, ) solve_options = { "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, } if config.first_step is not None: solve_options["first_step"] = config.first_step return solve_ivp(**solve_options)