from __future__ import annotations from app.simulation.core.errors import RecoverableTrialStateError from app.simulation.performance import profile_phase 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]] JacobianCallable = Callable[[float, object], object] @dataclass(frozen=True) class SolverActivitySnapshot: """Low-cost, additive view of work inside an integration task. ``accepted_time`` deliberately changes only after an accepted solver step. Trial evaluations may continue to advance ``activity_sequence`` and ``current_trial_time`` while that public progress value stays fixed. """ activity_sequence: int activity_kind: str current_trial_time: float | None rhs_call_count: int accepted_step_sequence: int accepted_time: float | None solver_step_sequence: int jacobian_evaluation_count: int thermofluid_closure_count: int def as_dict(self) -> dict[str, object]: return { "activitySequence": self.activity_sequence, "activityKind": self.activity_kind, "currentTrialTime": self.current_trial_time, "rhsCallCount": self.rhs_call_count, "acceptedStepSequence": self.accepted_step_sequence, "acceptedTime": self.accepted_time, "solverStepSequence": self.solver_step_sequence, "jacobianEvaluationCount": self.jacobian_evaluation_count, "thermofluidClosureCount": self.thermofluid_closure_count, } class SolverActivityTracker: """Single-writer activity telemetry for a solver worker. The solver thread is the only writer and the stream thread only snapshots scalar attributes. The sequence is published last, so a reader never treats partially published fields as a newer completed activity update. Passing no tracker to :func:`integrate_ode` is the zero-cost opt-out path. """ __slots__ = ( "_accepted_step_sequence", "_accepted_time", "_activity_kind", "_activity_sequence", "_current_trial_time", "_jacobian_evaluation_count", "_rhs_call_count", "_solver_step_sequence", "_thermofluid_closure_count", ) def __init__(self) -> None: self._activity_sequence = 0 self._activity_kind = "idle" self._current_trial_time: float | None = None self._rhs_call_count = 0 self._accepted_step_sequence = 0 self._accepted_time: float | None = None self._solver_step_sequence = 0 self._jacobian_evaluation_count = 0 self._thermofluid_closure_count = 0 def _publish(self, kind: str, time: float | None = None) -> None: self._activity_kind = kind if time is not None: self._current_trial_time = float(time) self._activity_sequence += 1 def start_integration(self, time: float) -> None: self._accepted_time = float(time) self._publish("solver_initialization", time) def record_phase(self, kind: str, time: float | None = None) -> None: self._publish(kind, time) def record_solver_step(self, time: float) -> None: self._solver_step_sequence += 1 self._publish("solver_step", time) def record_rhs(self, time: float) -> None: self._rhs_call_count += 1 self._publish("rhs", time) def record_jacobian(self, time: float) -> None: self._jacobian_evaluation_count += 1 self._publish("jacobian", time) def record_thermofluid_closure(self, time: float) -> None: self._thermofluid_closure_count += 1 self._publish("thermofluid_closure", time) def record_accepted_step(self, time: float) -> None: accepted_time = float(time) if ( self._accepted_time is not None and accepted_time <= self._accepted_time ): return self._accepted_step_sequence += 1 self._accepted_time = accepted_time self._publish("accepted_step", accepted_time) def snapshot(self) -> SolverActivitySnapshot: # ``activity_sequence`` is read last because writers publish it last. activity_kind = self._activity_kind current_trial_time = self._current_trial_time rhs_call_count = self._rhs_call_count accepted_step_sequence = self._accepted_step_sequence accepted_time = self._accepted_time solver_step_sequence = self._solver_step_sequence jacobian_evaluation_count = self._jacobian_evaluation_count thermofluid_closure_count = self._thermofluid_closure_count activity_sequence = self._activity_sequence return SolverActivitySnapshot( activity_sequence=activity_sequence, activity_kind=activity_kind, current_trial_time=current_trial_time, rhs_call_count=rhs_call_count, accepted_step_sequence=accepted_step_sequence, accepted_time=accepted_time, solver_step_sequence=solver_step_sequence, jacobian_evaluation_count=jacobian_evaluation_count, thermofluid_closure_count=thermofluid_closure_count, ) @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 _MAX_RECOVERABLE_RETRIES = 16 _RECOVERABLE_RETRY_FACTOR = 0.5 class IntegrationCancelled(Exception): """Internal control-flow signal shared by RHS and Jacobian evaluation.""" 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 RecoverableRetryDiagnostics: """One recoverable trial failure and the step cap chosen for its retry.""" phase: Literal["constructor", "step", "solver-status"] attempted_step: float reason: str next_max_step: float | None = None next_first_step: float | None = None def as_dict(self) -> dict[str, object]: result: dict[str, object] = { "phase": self.phase, "attemptedStep": self.attempted_step, "reason": self.reason, } if self.next_max_step is not None: result["nextMaxStep"] = self.next_max_step if self.next_first_step is not None: result["nextFirstStep"] = self.next_first_step return result @dataclass(frozen=True) class SolverSegmentDiagnostics: """Work performed by implicit solver instances inside one event segment.""" start_time: float requested_stop_time: float simulated_until: float nfev: int = 0 njev: int = 0 nlu: int = 0 accepted_step_count: int = 0 solver_start_count: int = 0 state_transition_count: int = 0 state_transition_times: tuple[float, ...] = () recoverable_retry_count: int = 0 jacobian_evaluation_count: int = 0 jacobian_full_build_count: int = 0 jacobian_secant_reuse_count: int = 0 jacobian_audit_failure_count: int = 0 finite_difference_rhs_evaluation_count: int = 0 jacobian_base_rhs_evaluation_count: int = 0 jacobian_jv_audit_rhs_evaluation_count: int = 0 exact_column_build_count: int = 0 exact_column_fallback_count: int = 0 jacobian_assembly_seconds: float = 0.0 recoverable_retries: tuple[RecoverableRetryDiagnostics, ...] = () def as_dict(self) -> dict[str, object]: result: dict[str, object] = { "startTime": self.start_time, "requestedStopTime": self.requested_stop_time, "simulatedUntil": self.simulated_until, "nfev": self.nfev, "njev": self.njev, "nlu": self.nlu, "acceptedStepCount": self.accepted_step_count, "solverStartCount": self.solver_start_count, "stateTransitionCount": self.state_transition_count, "recoverableRetryCount": self.recoverable_retry_count, } if self.state_transition_times: result["stateTransitionTimes"] = list( self.state_transition_times ) if self.recoverable_retries: result["recoverableRetries"] = [ retry.as_dict() for retry in self.recoverable_retries ] if ( self.jacobian_evaluation_count or self.finite_difference_rhs_evaluation_count or self.jacobian_assembly_seconds ): result.update( { "jacobianEvaluationCount": self.jacobian_evaluation_count, "jacobianFullBuildCount": self.jacobian_full_build_count, "jacobianSecantReuseCount": self.jacobian_secant_reuse_count, "jacobianAuditFailureCount": self.jacobian_audit_failure_count, "finiteDifferenceRhsEvaluationCount": ( self.finite_difference_rhs_evaluation_count ), "jacobianBaseRhsEvaluationCount": ( self.jacobian_base_rhs_evaluation_count ), "jacobianJvAuditRhsEvaluationCount": ( self.jacobian_jv_audit_rhs_evaluation_count ), "exactColumnBuildCount": self.exact_column_build_count, "exactColumnFallbackCount": ( self.exact_column_fallback_count ), "jacobianAssemblySeconds": self.jacobian_assembly_seconds, } ) return result _JACOBIAN_DIAGNOSTIC_KEYS = ( "jacobianEvaluationCount", "fullBuildCount", "secantReuseCount", "auditFailureCount", "finiteDifferenceRhsEvaluationCount", "baseRhsEvaluationCount", "jvAuditEvaluationCount", "exactColumnBuildCount", "exactColumnFallbackCount", "assemblySeconds", ) def _jacobian_diagnostic_snapshot( jac: JacobianCallable | None, ) -> dict[str, float]: diagnostics = getattr(jac, "diagnostics", None) if diagnostics is None: return {key: 0.0 for key in _JACOBIAN_DIAGNOSTIC_KEYS} values = diagnostics() return { key: float(values.get(key, 0.0)) for key in _JACOBIAN_DIAGNOSTIC_KEYS } def _positive_finite_step(value: object) -> float | None: if value is None: return None try: candidate = abs(float(value)) except (TypeError, ValueError, OverflowError): return None return candidate if candidate > 0.0 and math.isfinite(candidate) else None def _smallest_positive_finite_step(*values: object) -> float: """Return a conservative step bound from configuration candidates.""" candidates = [ candidate for value in values if (candidate := _positive_finite_step(value)) is not None ] if not candidates: raise ValueError("No positive finite integration step is available.") return min(candidates) def _solver_attempted_step( solver: object, *, segment_max_step: float, remaining_interval: float, ) -> float: """Snapshot the real trial scale before calling ``solver.step()``. SciPy exposes the proposed step as ``h_abs``. ``step_size`` is the prior accepted step, so it is only a fallback for solvers without a valid ``h_abs``; it must not reduce an otherwise valid failed-trial estimate. """ configured_cap = _smallest_positive_finite_step( segment_max_step, remaining_interval, ) for attribute in ("h_abs", "step_size"): try: candidate = _positive_finite_step( getattr(solver, attribute, None) ) except Exception: # A third-party OdeSolver may implement these as fragile # properties. The configured cap remains a safe fallback. continue if candidate is not None: return min(candidate, configured_cap) return configured_cap def _recoverable_retry_steps( attempted_step: float, *, last_accepted_time: float, ) -> tuple[float, float] | None: """Return strictly smaller max/first steps, or None at machine precision.""" next_step = _RECOVERABLE_RETRY_FACTOR * attempted_step minimum_step = 64.0 * math.ulp(max(abs(last_accepted_time), 1.0)) if ( not math.isfinite(next_step) or next_step <= minimum_step or next_step >= attempted_step ): return None # This first step is intentionally one-shot. Keeping it equal to the new # cap makes both controls strictly smaller than the failed trial scale. return next_step, next_step @dataclass(frozen=True) class ODESolution: t: list[float] y: list[list[float]] success: bool message: str status: IntegrationStatus = "completed" error: Exception | None = None solver_segments: tuple[SolverSegmentDiagnostics, ...] = () 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, jac: JacobianCallable | None = None, activity_tracker: SolverActivityTracker | None = 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}") implicit_jac = jac if config.method in {"BDF", "Radau"} else None solver_jac = implicit_jac if implicit_jac is not None and activity_tracker is not None: original_jacobian = implicit_jac def activity_jacobian(time, state): activity_tracker.record_jacobian(float(time)) try: return original_jacobian(time, state) finally: activity_tracker.record_phase("solver_step", float(time)) solver_jac = activity_jacobian 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 normalized_state = [float(value) for value in state] derivative = rhs(float(time), normalized_state) observer = getattr(implicit_jac, "observe", None) if observer is not None: observer(float(time), normalized_state, derivative) return derivative 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 solver_segments: list[SolverSegmentDiagnostics] = [] 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_start_time = last_accepted_time segment_max_step = float(config.max_step) recoverable_retry_count = 0 last_recoverable_error: RecoverableTrialStateError | None = None retry_first_step: float | None = None segment_nfev = 0 segment_njev = 0 segment_nlu = 0 segment_accepted_steps = 0 segment_solver_starts = 0 segment_state_transitions = 0 segment_state_transition_times: list[float] = [] segment_recoverable_retries = 0 segment_recoverable_retry_diagnostics: list[ RecoverableRetryDiagnostics ] = [] jacobian_work_start = _jacobian_diagnostic_snapshot(implicit_jac) 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 config.method in {"BDF", "Radau"}: if solver_jac is not None: solver_options["jac"] = solver_jac elif jac_sparsity is not None: solver_options["jac_sparsity"] = jac_sparsity requested_first_step = ( retry_first_step if retry_first_step 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: constructor_attempted_step = _smallest_positive_finite_step( segment_max_step, integration_end - last_accepted_time, solver_options.get("first_step"), ) start_segment = getattr(implicit_jac, "start_segment", None) if start_segment is not None: start_segment() if activity_tracker is not None: activity_tracker.record_phase( "solver_initialization", last_accepted_time, ) 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 segment_recoverable_retries += 1 last_recoverable_error = exc retry_steps = ( _recoverable_retry_steps( constructor_attempted_step, last_accepted_time=last_accepted_time, ) if recoverable_retry_count <= _MAX_RECOVERABLE_RETRIES else None ) segment_recoverable_retry_diagnostics.append( RecoverableRetryDiagnostics( phase="constructor", attempted_step=constructor_attempted_step, reason=str(exc), next_max_step=( retry_steps[0] if retry_steps is not None else None ), next_first_step=( retry_steps[1] if retry_steps is not None else None ), ) ) if retry_steps is None: status = "failed" message = str(exc) error = exc break segment_max_step, retry_first_step = retry_steps continue except Exception as exc: status = "failed" message = str(exc) error = exc break segment_solver_starts += 1 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: attempted_step = _solver_attempted_step( solver, segment_max_step=segment_max_step, remaining_interval=( integration_end - last_accepted_time ), ) if activity_tracker is not None: activity_tracker.record_solver_step( last_accepted_time ) 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 segment_recoverable_retries += 1 last_recoverable_error = exc retry_steps = ( _recoverable_retry_steps( attempted_step, last_accepted_time=last_accepted_time, ) if recoverable_retry_count <= _MAX_RECOVERABLE_RETRIES else None ) segment_recoverable_retry_diagnostics.append( RecoverableRetryDiagnostics( phase="step", attempted_step=attempted_step, reason=str(exc), next_max_step=( retry_steps[0] if retry_steps is not None else None ), next_first_step=( retry_steps[1] if retry_steps is not None else None ), ) ) if retry_steps is None: status = "failed" message = str(exc) error = exc break segment_max_step, retry_first_step = retry_steps 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 segment_recoverable_retries += 1 retry_steps = ( _recoverable_retry_steps( attempted_step, last_accepted_time=last_accepted_time, ) if recoverable_retry_count <= _MAX_RECOVERABLE_RETRIES else None ) failure_reason = str( step_message or last_recoverable_error ) segment_recoverable_retry_diagnostics.append( RecoverableRetryDiagnostics( phase="solver-status", attempted_step=attempted_step, reason=failure_reason, next_max_step=( retry_steps[0] if retry_steps is not None else None ), next_first_step=( retry_steps[1] if retry_steps is not None else None ), ) ) if retry_steps is not None: segment_max_step, retry_first_step = retry_steps restart_after_recoverable = True break status = "failed" message = str(step_message or "Integration step failed.") break # A returned running/finished status means this step was # accepted. Any prior recoverable failure is now historical: # it must not influence an event restart or an ordinary later # solver failure. The reduced cap is local to the failed # trial: after one accepted retry step, let this solver grow # adaptively again and ensure a later event restart receives # the configured maximum. The retry-specific first step is # likewise strictly one-shot. if retry_first_step is not None: segment_max_step = float(config.max_step) try: solver.max_step = segment_max_step except (AttributeError, TypeError, ValueError): # Third-party OdeSolver-compatible test doubles may not # expose a writable cap. SciPy's supported solvers do. pass last_recoverable_error = None retry_first_step = None recoverable_retry_count = 0 segment_accepted_steps += 1 step_end_time = float(solver.t) step_end_state = [float(value) for value in solver.y] crosses_sample = ( sample_index < len(sample_times) and sample_times[sample_index] <= step_end_time ) dense_output = ( solver.dense_output() if crosses_sample 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: segment_state_transitions += 1 segment_state_transition_times.append( float(transition.time) ) 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 reported_time = ( float(segment_end) if is_breakpoint and solver.status == "finished" else last_accepted_time ) if sample_times: while ( sample_index < len(sample_times) and sample_times[sample_index] <= last_accepted_time ): assert dense_output is not None 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) segment_nfev += int(getattr(solver, "nfev", 0)) segment_njev += int(getattr(solver, "njev", 0)) segment_nlu += int(getattr(solver, "nlu", 0)) if status != "completed": break if restart_after_recoverable: continue if not restart_at_transition: break jacobian_work_end = _jacobian_diagnostic_snapshot(implicit_jac) jacobian_work = { key: jacobian_work_end[key] - jacobian_work_start[key] for key in _JACOBIAN_DIAGNOSTIC_KEYS } solver_segments.append( SolverSegmentDiagnostics( start_time=float(segment_start_time), requested_stop_time=float(segment_end), simulated_until=float( segment_end if status == "completed" else last_accepted_time ), nfev=segment_nfev, njev=segment_njev, nlu=segment_nlu, accepted_step_count=segment_accepted_steps, solver_start_count=segment_solver_starts, state_transition_count=segment_state_transitions, state_transition_times=tuple( segment_state_transition_times ), recoverable_retry_count=segment_recoverable_retries, recoverable_retries=tuple( segment_recoverable_retry_diagnostics ), jacobian_evaluation_count=int( jacobian_work["jacobianEvaluationCount"] ), jacobian_full_build_count=int( jacobian_work["fullBuildCount"] ), jacobian_secant_reuse_count=int( jacobian_work["secantReuseCount"] ), jacobian_audit_failure_count=int( jacobian_work["auditFailureCount"] ), finite_difference_rhs_evaluation_count=int( jacobian_work["finiteDifferenceRhsEvaluationCount"] ), jacobian_base_rhs_evaluation_count=int( jacobian_work["baseRhsEvaluationCount"] ), jacobian_jv_audit_rhs_evaluation_count=int( jacobian_work["jvAuditEvaluationCount"] ), exact_column_build_count=int( jacobian_work["exactColumnBuildCount"] ), exact_column_fallback_count=int( jacobian_work["exactColumnFallbackCount"] ), jacobian_assembly_seconds=jacobian_work["assemblySeconds"], ) ) 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, solver_segments=tuple(solver_segments), ) @profile_phase("simulation.integration") 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, jac: JacobianCallable | None = None, recoverable_trial_retries: bool = False, activity_tracker: SolverActivityTracker | 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. ``recoverable_trial_retries`` opts an eventless/cancellation-free caller into the stepwise path so a ``RecoverableTrialStateError`` can rebuild the solver from its last accepted state. It defaults to false to preserve the direct ``solve_ivp`` path for ordinary callers. ``activity_tracker`` is optional and additive. When omitted, the numerical call path and callback behavior are unchanged. """ integration_rhs = rhs integration_accepted_step_callback = accepted_step_callback if activity_tracker is not None: activity_tracker.start_integration(config.t_start) original_rhs = rhs def activity_rhs(time, state): numeric_time = float(time) activity_tracker.record_rhs(numeric_time) try: return original_rhs(time, state) finally: activity_tracker.record_phase("solver_step", numeric_time) integration_rhs = activity_rhs def activity_accepted_step(time: float) -> None: activity_tracker.record_accepted_step(float(time)) if accepted_step_callback is not None: accepted_step_callback(float(time)) integration_accepted_step_callback = activity_accepted_step 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( integration_rhs, initial_state, config, t_eval, normalized_breakpoints, cancel_check, integration_accepted_step_callback, state_transition_handler, ) return _runge_kutta_4( integration_rhs, initial_state, config, t_eval, cancel_check, integration_accepted_step_callback, state_transition_handler, ) if ( cancel_check is not None or normalized_breakpoints or state_transition_handler is not None or recoverable_trial_retries ): return _integrate_scipy_stepwise( integration_rhs, initial_state, config, t_eval, cancel_check or (lambda: False), integration_accepted_step_callback, normalized_breakpoints, state_transition_handler, jac_sparsity, jac, activity_tracker, ) implicit_jac = jac if config.method in {"BDF", "Radau"} else None solve_rhs = integration_rhs if implicit_jac is not None: observer = getattr(implicit_jac, "observe", None) if observer is not None: def observed_rhs(time, state): derivative = integration_rhs(time, state) observer(float(time), state, derivative) return derivative solve_rhs = observed_rhs start_segment = getattr(implicit_jac, "start_segment", None) if start_segment is not None: start_segment() solve_jac = implicit_jac if implicit_jac is not None and activity_tracker is not None: original_jacobian = implicit_jac def activity_jacobian(time, state): activity_tracker.record_jacobian(float(time)) try: return original_jacobian(time, state) finally: activity_tracker.record_phase("solver_step", float(time)) solve_jac = activity_jacobian solve_options = { "fun": solve_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 solve_jac is not None: solve_options["jac"] = solve_jac elif jac_sparsity is not None and config.method in {"BDF", "Radau"}: solve_options["jac_sparsity"] = jac_sparsity direct_solution = solve_ivp(**solve_options) if activity_tracker is not None and len(direct_solution.t): activity_tracker.record_accepted_step(float(direct_solution.t[-1])) return direct_solution