use crate::errors::CurveError;
use super::Interpolator;
#[derive(Debug, Clone)]
pub struct HermiteBessel {
times: Vec<f64>,
values: Vec<f64>,
slopes: Vec<f64>,
}
impl HermiteBessel {
pub fn new(knots: &[(f64, f64)]) -> Result<Self, CurveError> {
if knots.len() < 2 {
return Err(CurveError::TooFewNodes { found: knots.len() });
}
let n = knots.len();
let mut times = Vec::with_capacity(n);
let mut values = Vec::with_capacity(n);
for (i, &(t, y)) in knots.iter().enumerate() {
if !t.is_finite() {
return Err(CurveError::InvalidTime { t });
}
if !y.is_finite() {
return Err(CurveError::NonPositiveDiscount {
at_index: i,
value: y,
});
}
if i > 0 {
let prev = times[i - 1];
#[allow(clippy::float_cmp)]
let is_duplicate = t == prev;
if is_duplicate {
return Err(CurveError::DuplicateNode { t });
}
if t < prev {
return Err(CurveError::NodesNotIncreasing { at_index: i });
}
}
times.push(t);
values.push(y);
}
let slopes = bessel_slopes(×, &values);
Ok(Self {
times,
values,
slopes,
})
}
#[must_use]
#[inline]
pub fn len(&self) -> usize {
self.times.len()
}
#[must_use]
#[inline]
pub fn is_empty(&self) -> bool {
self.times.is_empty()
}
#[must_use]
#[inline]
pub fn slope_at_knot(&self, i: usize) -> Option<f64> {
self.slopes.get(i).copied()
}
#[inline]
fn locate(&self, t: f64) -> usize {
let n = self.times.len();
if t <= self.times[0] {
return 0;
}
if t >= self.times[n - 1] {
return n - 2;
}
let mut lo = 0_usize;
let mut hi = n - 1;
while hi - lo > 1 {
let mid = lo + (hi - lo) / 2;
if self.times[mid] <= t {
lo = mid;
} else {
hi = mid;
}
}
lo
}
}
fn bessel_slopes(times: &[f64], values: &[f64]) -> Vec<f64> {
let n = times.len();
let mut slopes = Vec::with_capacity(n);
if n == 2 {
let s0 = (values[1] - values[0]) / (times[1] - times[0]);
slopes.push(s0);
slopes.push(s0);
return slopes;
}
let mut h = Vec::with_capacity(n - 1);
let mut s = Vec::with_capacity(n - 1);
for i in 0..n - 1 {
let hi = times[i + 1] - times[i];
h.push(hi);
s.push((values[i + 1] - values[i]) / hi);
}
let m0 = ((2.0 * h[0] + h[1]) * s[0] - h[0] * s[1]) / (h[0] + h[1]);
slopes.push(m0);
for i in 1..n - 1 {
let m = (h[i] * s[i - 1] + h[i - 1] * s[i]) / (h[i - 1] + h[i]);
slopes.push(m);
}
let hn2 = h[n - 2];
let hn3 = h[n - 3];
let sn2 = s[n - 2];
let sn3 = s[n - 3];
let m_last = ((2.0 * hn2 + hn3) * sn2 - hn2 * sn3) / (hn3 + hn2);
slopes.push(m_last);
slopes
}
impl Interpolator for HermiteBessel {
fn build(knots: &[(f64, f64)]) -> Result<Self, CurveError> {
Self::new(knots)
}
fn eval(&self, t: f64) -> f64 {
let n = self.times.len();
if t <= self.times[0] {
return self.values[0];
}
if t >= self.times[n - 1] {
return self.values[n - 1];
}
let i = self.locate(t);
let t_lo = self.times[i];
let t_hi = self.times[i + 1];
let dt = t_hi - t_lo;
let u = (t - t_lo) / dt;
let u2 = u * u;
let u3 = u2 * u;
let h00 = 2.0 * u3 - 3.0 * u2 + 1.0;
let h10 = u3 - 2.0 * u2 + u;
let h01 = -2.0 * u3 + 3.0 * u2;
let h11 = u3 - u2;
h00 * self.values[i]
+ dt * h10 * self.slopes[i]
+ h01 * self.values[i + 1]
+ dt * h11 * self.slopes[i + 1]
}
fn deriv(&self, t: f64) -> Option<f64> {
let n = self.times.len();
if t < self.times[0] || t > self.times[n - 1] {
return Some(0.0);
}
let i = self.locate(t);
let t_lo = self.times[i];
let t_hi = self.times[i + 1];
let dt = t_hi - t_lo;
let u = (t - t_lo) / dt;
let u2 = u * u;
let term_endpoints = (6.0 * u2 - 6.0 * u) * (self.values[i] - self.values[i + 1]) / dt;
let term_left = (3.0 * u2 - 4.0 * u + 1.0) * self.slopes[i];
let term_right = (3.0 * u2 - 2.0 * u) * self.slopes[i + 1];
Some(term_endpoints + term_left + term_right)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_empty() {
let err = HermiteBessel::new(&[]).unwrap_err();
assert!(matches!(err, CurveError::TooFewNodes { found: 0 }));
}
#[test]
fn rejects_single_knot() {
let err = HermiteBessel::new(&[(0.0, 1.0)]).unwrap_err();
assert!(matches!(err, CurveError::TooFewNodes { found: 1 }));
}
#[test]
fn rejects_non_monotone_times() {
let err = HermiteBessel::new(&[(0.0, 1.0), (2.0, 0.9), (1.0, 0.95)]).unwrap_err();
assert!(matches!(
err,
CurveError::NodesNotIncreasing { at_index: 2 }
));
}
#[test]
fn rejects_duplicate_times() {
let err = HermiteBessel::new(&[(0.0, 1.0), (1.0, 0.95), (1.0, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::DuplicateNode { .. }));
}
#[test]
fn rejects_nan_time() {
let err = HermiteBessel::new(&[(0.0, 1.0), (f64::NAN, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::InvalidTime { .. }));
}
#[test]
fn rejects_inf_time() {
let err = HermiteBessel::new(&[(0.0, 1.0), (f64::INFINITY, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::InvalidTime { .. }));
}
#[test]
fn rejects_nan_value() {
let err = HermiteBessel::new(&[(0.0, 1.0), (1.0, f64::NAN)]).unwrap_err();
assert!(matches!(
err,
CurveError::NonPositiveDiscount { at_index: 1, .. }
));
}
#[test]
fn rejects_inf_value() {
let err = HermiteBessel::new(&[(0.0, 1.0), (1.0, f64::INFINITY)]).unwrap_err();
assert!(matches!(
err,
CurveError::NonPositiveDiscount { at_index: 1, .. }
));
}
#[test]
fn accepts_negative_value() {
let interp = HermiteBessel::new(&[(0.0, 1.0), (1.0, -0.5), (2.0, 0.25)]).unwrap();
assert!((interp.eval(0.0) - 1.0).abs() < 1e-14);
assert!((interp.eval(1.0) - (-0.5)).abs() < 1e-14);
}
#[test]
fn knot_reproduction_exact() {
let knots = [
(0.0, 1.0),
(0.5, 0.97),
(1.0, 0.95),
(2.5, 0.90),
(5.0, 0.80),
];
let interp = HermiteBessel::new(&knots).unwrap();
for &(t, y) in &knots {
let v = interp.eval(t);
assert!((v - y).abs() < 1e-14, "knot ({t}, {y}) -> {v}");
}
}
#[test]
fn linear_function_reproduced_exactly() {
let knots: Vec<(f64, f64)> = [0.0, 0.7, 1.3, 2.5, 4.1, 5.0]
.iter()
.map(|&t| (t, 2.0 + 3.0 * t))
.collect();
let interp = HermiteBessel::new(&knots).unwrap();
let mut t = 0.0_f64;
while t <= 5.0 {
let v = interp.eval(t);
let expected = 2.0 + 3.0 * t;
assert!(
(v - expected).abs() < 1e-13,
"t={t}, v={v}, expected={expected}"
);
t += 0.05;
}
}
#[test]
fn parabola_reproduced_at_interior_midpoints() {
let knots: Vec<(f64, f64)> = [0.0, 1.0, 2.0, 3.0, 4.0]
.iter()
.map(|&t| (t, t * t))
.collect();
let interp = HermiteBessel::new(&knots).unwrap();
for &mid in &[1.5_f64, 2.5_f64] {
let v = interp.eval(mid);
let expected = mid * mid;
assert!(
(v - expected).abs() < 1e-13,
"t={mid}, v={v}, expected={expected}"
);
}
}
#[test]
fn cubic_polynomial_well_approximated() {
let knots: Vec<(f64, f64)> = [0.0, 0.5, 1.0, 1.5, 2.0]
.iter()
.map(|&t| (t, t * t * t))
.collect();
let interp = HermiteBessel::new(&knots).unwrap();
let mut max_err = 0.0_f64;
let mut t = 0.5_f64;
while t <= 1.5 {
let v = interp.eval(t);
let expected = t * t * t;
let err = (v - expected).abs();
if err > max_err {
max_err = err;
}
t += 0.001;
}
assert!(max_err < 1.5e-2, "peak |y - t^3| on interior = {max_err}");
}
#[test]
fn c1_continuity_at_interior_knots() {
let knots = [(0.0, 0.0), (0.5, 0.25), (1.4, 1.0), (2.6, 0.5), (4.1, -0.3)];
let interp = HermiteBessel::new(&knots).unwrap();
for i in 1..knots.len() - 1 {
let (t, _) = knots[i];
let m_i = interp.slope_at_knot(i).expect("interior knot");
let right = interp.deriv(t).expect("deriv defined on interior");
let (t_lo, y_lo) = knots[i - 1];
let (t_hi, y_hi) = knots[i];
let h = t_hi - t_lo;
let u = 1.0_f64;
let u2 = u * u;
let m_prev = interp.slope_at_knot(i - 1).expect("left knot");
let left = (6.0 * u2 - 6.0 * u) * (y_lo - y_hi) / h
+ (3.0 * u2 - 4.0 * u + 1.0) * m_prev
+ (3.0 * u2 - 2.0 * u) * m_i;
assert!(
(left - m_i).abs() < 1e-14,
"left u=1 derivative at t={t}: got {left}, m_i={m_i}",
);
assert!(
(right - m_i).abs() < 1e-14,
"right u=0 derivative at t={t}: got {right}, m_i={m_i}",
);
assert!(
(left - right).abs() < 1e-14,
"C^1 violated at t={t}: left={left}, right={right}",
);
}
}
#[test]
fn deriv_consistent_with_finite_difference() {
let knots = [
(0.0, 1.0),
(1.0, 0.95),
(2.0, 0.85),
(4.0, 0.70),
(7.0, 0.55),
];
let interp = HermiteBessel::new(&knots).unwrap();
for i in 0..knots.len() - 1 {
let mid = 0.5 * (knots[i].0 + knots[i + 1].0);
let analytic = interp.deriv(mid).expect("deriv defined on interior");
let h = 1e-6_f64;
let fd = (interp.eval(mid + h) - interp.eval(mid - h)) / (2.0 * h);
assert!(
(analytic - fd).abs() < 1e-7,
"segment {i}, mid={mid}: analytic={analytic}, fd={fd}",
);
}
}
#[test]
fn flat_extrapolation_in_value() {
let interp = HermiteBessel::new(&[(0.0, 1.0), (1.0, 0.95), (2.0, 0.85)]).unwrap();
assert!((interp.eval(-1.0) - 1.0).abs() < 1e-14);
assert!((interp.eval(-100.0) - 1.0).abs() < 1e-14);
assert!((interp.eval(2.0) - 0.85).abs() < 1e-14);
assert!((interp.eval(100.0) - 0.85).abs() < 1e-14);
assert!((interp.eval(f64::INFINITY) - 0.85).abs() < 1e-14);
assert!((interp.eval(f64::NEG_INFINITY) - 1.0).abs() < 1e-14);
}
#[test]
fn deriv_zero_in_extrapolation_region() {
let interp = HermiteBessel::new(&[(0.0, 1.0), (1.0, 0.95), (2.0, 0.85)]).unwrap();
let d_left = interp.deriv(-1.0).unwrap();
let d_right = interp.deriv(3.0).unwrap();
assert!((d_left - 0.0).abs() < 1e-15);
assert!((d_right - 0.0).abs() < 1e-15);
}
#[test]
fn two_knot_case_is_linear() {
let interp = HermiteBessel::new(&[(0.0, 1.0), (2.0, 5.0)]).unwrap();
assert!((interp.slope_at_knot(0).unwrap() - 2.0).abs() < 1e-15);
assert!((interp.slope_at_knot(1).unwrap() - 2.0).abs() < 1e-15);
for &t in &[0.5_f64, 1.0, 1.5] {
let v = interp.eval(t);
let expected = 1.0 + 2.0 * t;
assert!((v - expected).abs() < 1e-14, "t={t}, v={v}");
}
}
#[test]
fn three_knot_case_reproduces_parabola_exactly() {
let knots = [(-1.0, 1.0), (0.5, 0.25), (2.0, 4.0)];
let interp = HermiteBessel::new(&knots).unwrap();
for (i, &(t, _)) in knots.iter().enumerate() {
let m = interp.slope_at_knot(i).unwrap();
assert!(
(m - 2.0 * t).abs() < 1e-13,
"knot {i}: slope={m}, expected={}",
2.0 * t
);
}
for &t in &[-0.5_f64, 0.0, 0.8, 1.3, 1.9] {
let v = interp.eval(t);
assert!((v - t * t).abs() < 1e-13, "t={t}, v={v}");
}
}
#[test]
fn slope_at_knot_out_of_range_is_none() {
let interp = HermiteBessel::new(&[(0.0, 1.0), (1.0, 0.95)]).unwrap();
assert!(interp.slope_at_knot(0).is_some());
assert!(interp.slope_at_knot(1).is_some());
assert!(interp.slope_at_knot(2).is_none());
assert!(interp.slope_at_knot(99).is_none());
}
#[test]
fn build_trait_method_equivalent_to_new() {
let knots = [(0.0, 1.0), (1.0, 0.95), (2.0, 0.85)];
let a = HermiteBessel::new(&knots).unwrap();
let b = <HermiteBessel as Interpolator>::build(&knots).unwrap();
assert!((a.eval(0.5) - b.eval(0.5)).abs() < 1e-15);
assert!((a.eval(1.7) - b.eval(1.7)).abs() < 1e-15);
}
#[test]
fn len_and_is_empty() {
let interp = HermiteBessel::new(&[(0.0, 1.0), (1.0, 0.95), (2.0, 0.9)]).unwrap();
assert_eq!(interp.len(), 3);
assert!(!interp.is_empty());
}
#[test]
fn clone_yields_equivalent_interpolant() {
let interp = HermiteBessel::new(&[(0.0, 1.0), (1.0, 0.95), (2.0, 0.85)]).unwrap();
let copy = interp.clone();
for &t in &[0.0_f64, 0.3, 0.5, 1.0, 1.4, 2.0] {
assert!((interp.eval(t) - copy.eval(t)).abs() < 1e-15);
}
}
}