#![allow(dead_code)]
pub(crate) fn rk4_integrate(
f: &dyn Fn(f64, &[f64]) -> Vec<f64>,
y0: &[f64],
t_start: f64,
t_end: f64,
dt: f64,
) -> Vec<(f64, Vec<f64>)> {
assert!(dt > 0.0, "step size dt must be positive");
assert!(t_end > t_start, "t_end must be greater than t_start");
assert!(!y0.is_empty(), "initial state y0 must be non-empty");
let n = y0.len();
let num_steps = ((t_end - t_start) / dt).ceil() as usize;
let mut result = Vec::with_capacity(num_steps + 1);
let mut t = t_start;
let mut y = y0.to_vec();
result.push((t, y.clone()));
for _ in 0..num_steps {
let h = dt.min(t_end - t);
if h <= 0.0 {
break;
}
let k1 = f(t, &y);
let y_tmp: Vec<f64> = (0..n).map(|i| y[i] + k1[i] * h * 0.5).collect();
let k2 = f(t + h * 0.5, &y_tmp);
let y_tmp: Vec<f64> = (0..n).map(|i| y[i] + k2[i] * h * 0.5).collect();
let k3 = f(t + h * 0.5, &y_tmp);
let y_tmp: Vec<f64> = (0..n).map(|i| y[i] + k3[i] * h).collect();
let k4 = f(t + h, &y_tmp);
for i in 0..n {
y[i] += h / 6.0 * (k1[i] + 2.0 * k2[i] + 2.0 * k3[i] + k4[i]);
}
t += h;
result.push((t, y.clone()));
}
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn exponential_decay() {
let result = rk4_integrate(&|_t, y| vec![-y[0]], &[1.0], 0.0, 1.0, 0.001);
let last = result.last().unwrap();
let expected = (-1.0_f64).exp(); assert!(
(last.1[0] - expected).abs() < 1e-8,
"expected y(1) ≈ {expected}, got {}",
last.1[0]
);
}
#[test]
fn linear_growth() {
let result = rk4_integrate(&|_t, _y| vec![1.0], &[0.0], 0.0, 5.0, 0.1);
let last = result.last().unwrap();
assert!(
(last.0 - 5.0).abs() < 1e-10,
"expected t = 5.0, got {}",
last.0
);
assert!(
(last.1[0] - 5.0).abs() < 1e-8,
"expected y(5) = 5.0, got {}",
last.1[0]
);
}
#[test]
fn harmonic_oscillator() {
let result = rk4_integrate(
&|_t, y| vec![y[1], -y[0]],
&[1.0, 0.0],
0.0,
2.0 * std::f64::consts::PI,
0.001,
);
let last = result.last().unwrap();
assert!(
(last.1[0] - 1.0).abs() < 1e-6,
"expected x(2π) ≈ 1.0, got {}",
last.1[0]
);
assert!(
last.1[1].abs() < 1e-5,
"expected x'(2π) ≈ 0.0, got {}",
last.1[1]
);
}
#[test]
fn includes_initial_state() {
let result = rk4_integrate(&|_t, y| vec![-y[0]], &[1.0], 0.0, 0.1, 0.1);
assert_eq!(result[0].0, 0.0);
assert_eq!(result[0].1, vec![1.0]);
assert!(result.len() >= 2);
}
#[test]
#[should_panic(expected = "step size dt must be positive")]
fn zero_dt_panics() {
rk4_integrate(&|_t, y| vec![-y[0]], &[1.0], 0.0, 1.0, 0.0);
}
#[test]
#[should_panic(expected = "t_end must be greater than t_start")]
fn bad_range_panics() {
rk4_integrate(&|_t, y| vec![-y[0]], &[1.0], 1.0, 0.0, 0.1);
}
}