use nalgebra::DVector;
#[derive(Debug, Clone)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct StepSizeController {
rtol: f64,
atol: f64,
dt_min: f64,
dt_max: f64,
safety: f64,
order: f64,
prev_error_norm: Option<f64>,
}
impl StepSizeController {
pub fn new(rtol: f64, atol: f64, dt_min: f64, dt_max: f64, order: f64) -> Self {
Self {
rtol,
atol,
dt_min,
dt_max,
safety: 0.9,
order,
prev_error_norm: None,
}
}
pub fn error_norm(&self, error: &DVector<f64>, reference: &DVector<f64>) -> f64 {
let n = error.len().max(1);
let sum_sq: f64 = error
.iter()
.zip(reference.iter())
.map(|(e, u)| {
let scale = self.atol + self.rtol * u.abs();
(e / scale).powi(2)
})
.sum();
(sum_sq / n as f64).sqrt()
}
pub fn accept(&self, error_norm: f64) -> bool {
error_norm <= 1.0
}
pub fn next_dt(&mut self, current_dt: f64, error_norm: f64) -> f64 {
let error_norm = error_norm.max(1e-12);
let beta1 = 0.7 / self.order;
let beta2 = 0.4 / self.order;
let factor = match self.prev_error_norm {
Some(prev) => {
let prev = prev.max(1e-12);
self.safety * error_norm.powf(-beta1) * prev.powf(beta2)
}
None => self.safety * error_norm.powf(-1.0 / self.order),
};
let factor = factor.clamp(0.2, 5.0);
self.prev_error_norm = Some(error_norm);
(current_dt * factor).clamp(self.dt_min, self.dt_max)
}
pub fn dt_min(&self) -> f64 {
self.dt_min
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_controller() -> StepSizeController {
StepSizeController::new(1e-6, 1e-9, 1e-8, 1.0, 4.0)
}
#[test]
fn error_norm_zero_when_error_is_zero() {
let controller = make_controller();
let error = DVector::from_vec(vec![0.0, 0.0, 0.0]);
let reference = DVector::from_vec(vec![1.0, 2.0, 3.0]);
assert_eq!(controller.error_norm(&error, &reference), 0.0);
}
#[test]
fn error_norm_scales_with_tolerance() {
let controller = StepSizeController::new(0.0, 1.0, 1e-8, 1.0, 4.0);
let error = DVector::from_vec(vec![1.0, 1.0]);
let reference = DVector::from_vec(vec![100.0, 100.0]); assert!((controller.error_norm(&error, &reference) - 1.0).abs() < 1e-12);
}
#[test]
fn accept_at_threshold() {
let controller = make_controller();
assert!(controller.accept(1.0));
assert!(controller.accept(0.5));
assert!(!controller.accept(1.0001));
}
#[test]
fn next_dt_shrinks_on_large_error() {
let mut controller = make_controller();
let dt = controller.next_dt(0.1, 100.0); assert!(dt < 0.1, "expected shrink, got dt={dt}");
}
#[test]
fn next_dt_grows_on_small_error() {
let mut controller = make_controller();
let dt = controller.next_dt(0.1, 0.01); assert!(dt > 0.1, "expected growth, got dt={dt}");
}
#[test]
fn next_dt_respects_dt_max_clamp() {
let mut controller = StepSizeController::new(1e-6, 1e-9, 1e-8, 0.2, 4.0);
let dt = controller.next_dt(0.1, 1e-9); assert!(dt <= 0.2, "expected clamp to dt_max=0.2, got dt={dt}");
}
#[test]
fn next_dt_respects_dt_min_clamp() {
let mut controller = StepSizeController::new(1e-6, 1e-9, 0.05, 1.0, 4.0);
let dt = controller.next_dt(0.1, 1e6); assert!(dt >= 0.05, "expected clamp to dt_min=0.05, got dt={dt}");
}
#[test]
fn next_dt_growth_factor_is_bounded() {
let mut controller = make_controller();
let dt = controller.next_dt(0.1, 1e-12);
assert!(
dt <= 0.1 * 5.0 + 1e-12,
"expected growth capped at 5x, got dt={dt}"
);
}
#[test]
fn dt_min_accessor_matches_constructor() {
let controller = StepSizeController::new(1e-6, 1e-9, 1e-7, 1.0, 4.0);
assert_eq!(controller.dt_min(), 1e-7);
}
#[test]
fn pi_memory_affects_second_call() {
let mut controller = make_controller();
let dt1 = controller.next_dt(0.1, 0.5);
let dt2 = controller.next_dt(dt1, 0.5);
assert!(controller.prev_error_norm.is_some());
let _ = dt2; }
#[cfg(feature = "serde")]
#[test]
fn serde_roundtrip_preserves_parameters() {
let mut controller = make_controller();
let _ = controller.next_dt(0.1, 0.5);
let json = serde_json::to_string(&controller).unwrap();
let restored: StepSizeController = serde_json::from_str(&json).unwrap();
assert_eq!(restored.dt_min(), controller.dt_min());
assert_eq!(restored.prev_error_norm, controller.prev_error_norm);
}
}