from __future__ import annotations from app.simulation.core.errors import RecoverableTrialStateError 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 | Sequence[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 _project_nearby_pre_transition_sample( times: list[float], states: list[list[float]], transition: StateTransition, config: SolveIVPConfig, ) -> None: """Resolve a sample/event ordering that is below solver time precision. An adaptive dense interpolant can place a discontinuous impact a few nanoseconds after its analytically coincident output sample. Keep the located event and restart time unchanged, but report that ambiguous sample on the reset side of the discontinuity. """ if not times or not math.isfinite(config.max_step): return time_gap = float(transition.time) - times[-1] tolerance = max( 64.0 * math.ulp(max(abs(float(transition.time)), 1.0)), min( abs(float(config.max_step) * float(config.rtol)), 1.0e-8, ), ) if not 0.0 < time_gap <= tolerance: return for index, value in enumerate(transition.state): states[index][-1] = float(value) 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, jac_sparsity=None, ) -> ODESolution: """Initial stepwise integration path for breakpoints and state resets. Recoverable physical-domain failures from rejected integrator trial states restore the last accepted state and rebuild the same solver with a smaller maximum/first step. Structural, algebraic, and ordinary model errors still fail immediately. """ 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 segment_max_step = float(config.max_step) recoverable_retry_count = 0 last_recoverable_error: RecoverableTrialStateError | None = None 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": segment_max_step, } if jac_sparsity is not None and config.method in {"BDF", "Radau"}: solver_options["jac_sparsity"] = jac_sparsity requested_first_step = ( 0.1 * segment_max_step if last_recoverable_error is not None else config.first_step ) if requested_first_step is not None: solver_options["first_step"] = min( requested_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 RecoverableTrialStateError as exc: recoverable_retry_count += 1 last_recoverable_error = exc next_step = 0.5 * segment_max_step minimum_step = 64.0 * math.ulp(max(abs(last_accepted_time), 1.0)) if recoverable_retry_count > 16 or next_step <= minimum_step: status = "failed" message = str(exc) error = exc break segment_max_step = next_step continue except Exception as exc: status = "failed" message = str(exc) error = exc break restart_at_transition = False restart_after_recoverable = 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 RecoverableTrialStateError as exc: recoverable_retry_count += 1 last_recoverable_error = exc attempted_step = segment_max_step next_step = 0.5 * attempted_step minimum_step = 64.0 * math.ulp( max(abs(last_accepted_time), 1.0) ) if recoverable_retry_count > 16 or next_step <= minimum_step: status = "failed" message = str(exc) error = exc break segment_max_step = next_step restart_after_recoverable = True break except Exception as exc: status = "failed" message = str(exc) error = exc break integration_progressed = True if solver.status == "failed": if last_recoverable_error is not None: recoverable_retry_count += 1 next_step = 0.5 * segment_max_step minimum_step = 64.0 * math.ulp( max(abs(last_accepted_time), 1.0) ) if ( recoverable_retry_count <= 16 and next_step > minimum_step ): segment_max_step = next_step restart_after_recoverable = True break 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 _project_nearby_pre_transition_sample( times, states, transition, config, ) 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 recoverable_retry_count = 0 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": break if restart_after_recoverable: continue if 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, jac_sparsity=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, jac_sparsity, ) 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 if jac_sparsity is not None and config.method in {"BDF", "Radau"}: solve_options["jac_sparsity"] = jac_sparsity return solve_ivp(**solve_options)