use crate::errors::CurveError;
use super::Interpolator;
#[derive(Debug, Clone)]
pub struct Linear {
times: Vec<f64>,
values: Vec<f64>,
}
impl Linear {
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);
}
Ok(Self { times, values })
}
#[must_use]
#[inline]
pub fn len(&self) -> usize {
self.times.len()
}
#[must_use]
#[inline]
pub fn is_empty(&self) -> bool {
self.times.is_empty()
}
#[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
}
}
impl Interpolator for Linear {
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 w = (t - t_lo) / (t_hi - t_lo);
(1.0 - w) * self.values[i] + w * self.values[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];
Some((self.values[i + 1] - self.values[i]) / (t_hi - t_lo))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_empty() {
let err = Linear::new(&[]).unwrap_err();
assert!(matches!(err, CurveError::TooFewNodes { found: 0 }));
}
#[test]
fn rejects_single_knot() {
let err = Linear::new(&[(0.0, 1.0)]).unwrap_err();
assert!(matches!(err, CurveError::TooFewNodes { found: 1 }));
}
#[test]
fn rejects_non_monotone_times() {
let err = Linear::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 = Linear::new(&[(0.0, 1.0), (1.0, 0.95), (1.0, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::DuplicateNode { .. }));
}
#[test]
fn accepts_negative_value() {
let interp = Linear::new(&[(0.0, 1.0), (1.0, -0.5)]).unwrap();
assert!((interp.eval(0.0) - 1.0).abs() < 1e-15);
assert!((interp.eval(1.0) - (-0.5)).abs() < 1e-15);
}
#[test]
fn accepts_zero_value() {
let interp = Linear::new(&[(0.0, 1.0), (1.0, 0.0)]).unwrap();
assert!((interp.eval(1.0) - 0.0).abs() < 1e-15);
}
#[test]
fn rejects_nan_value() {
let err = Linear::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 = Linear::new(&[(0.0, 1.0), (1.0, f64::INFINITY)]).unwrap_err();
assert!(matches!(
err,
CurveError::NonPositiveDiscount { at_index: 1, .. }
));
}
#[test]
fn rejects_nan_time() {
let err = Linear::new(&[(0.0, 1.0), (f64::NAN, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::InvalidTime { .. }));
}
#[test]
fn rejects_inf_time() {
let err = Linear::new(&[(0.0, 1.0), (f64::INFINITY, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::InvalidTime { .. }));
}
#[test]
fn knot_reproduction_exact() {
let knots = [(0.0, 1.0), (0.5, 0.97), (1.0, 0.95), (2.0, 0.90)];
let interp = Linear::new(&knots).unwrap();
for &(t, y) in &knots {
let v = interp.eval(t);
assert!((v - y).abs() < 1e-15, "knot ({t}, {y}) -> {v}");
}
}
#[test]
fn midpoint_identity() {
let interp = Linear::new(&[(0.0, 1.0), (1.0, 2.0)]).unwrap();
let mid = interp.eval(0.5);
#[allow(clippy::float_cmp)]
let exact = mid == 1.5;
assert!(exact, "expected exact 1.5, got {mid}");
}
#[test]
fn midpoint_in_segment_is_average() {
let interp = Linear::new(&[(0.0, 1.0), (1.0, 0.95), (3.0, 0.85)]).unwrap();
let mid = f64::midpoint(1.0, 3.0);
let v = interp.eval(mid);
let expected = f64::midpoint(0.95, 0.85);
assert!((v - expected).abs() < 1e-15);
}
#[test]
fn evaluation_matches_closed_form() {
let interp = Linear::new(&[(1.0, 10.0), (3.0, 4.0)]).unwrap();
for &t in &[1.0_f64, 1.25, 1.5, 2.0, 2.75, 3.0] {
let w = (t - 1.0) / (3.0 - 1.0);
let expected = (1.0 - w) * 10.0 + w * 4.0;
let v = interp.eval(t);
assert!(
(v - expected).abs() < 1e-15,
"t={t}: got {v}, want {expected}"
);
}
}
#[test]
fn flat_extrapolation_left() {
let interp = Linear::new(&[(0.5, 0.97), (1.0, 0.95)]).unwrap();
assert!((interp.eval(0.0) - 0.97).abs() < 1e-15);
assert!((interp.eval(-100.0) - 0.97).abs() < 1e-15);
assert!((interp.eval(f64::NEG_INFINITY) - 0.97).abs() < 1e-15);
}
#[test]
fn flat_extrapolation_right() {
let interp = Linear::new(&[(0.0, 1.0), (1.0, 0.95)]).unwrap();
assert!((interp.eval(2.0) - 0.95).abs() < 1e-15);
assert!((interp.eval(100.0) - 0.95).abs() < 1e-15);
assert!((interp.eval(f64::INFINITY) - 0.95).abs() < 1e-15);
}
#[test]
fn deriv_matches_segment_slope() {
let interp = Linear::new(&[(1.0, 10.0), (3.0, 4.0)]).unwrap();
let s = interp.deriv(2.0).unwrap();
assert!((s - (-3.0)).abs() < 1e-15);
}
#[test]
fn deriv_finite_difference_interior() {
let knots = [(0.0, 1.0), (1.0, 1.2), (2.0, 0.8), (5.0, 0.5)];
let interp = Linear::new(&knots).unwrap();
let t = 1.5_f64;
let dy_dt = interp.deriv(t).unwrap();
let h = 1e-6_f64;
let fd = (interp.eval(t + h) - interp.eval(t - h)) / (2.0 * h);
assert!((dy_dt - fd).abs() < 1e-9, "analytic={dy_dt}, fd={fd}");
let expected = -0.4_f64;
assert!((dy_dt - expected).abs() < 1e-15);
}
#[test]
fn deriv_zero_in_extrapolation_region() {
let interp = Linear::new(&[(0.0, 1.0), (1.0, 0.95)]).unwrap();
let d_left = interp.deriv(-1.0).unwrap();
assert!((d_left - 0.0).abs() < 1e-15);
let d_right = interp.deriv(2.0).unwrap();
assert!((d_right - 0.0).abs() < 1e-15);
}
#[test]
fn deriv_at_knot_returns_right_slope() {
let knots = [(0.0, 1.0), (1.0, 0.95), (2.0, 0.80)];
let interp = Linear::new(&knots).unwrap();
let d_at_1 = interp.deriv(1.0).unwrap();
let expected_right = 0.80_f64 - 0.95_f64;
assert!(
(d_at_1 - expected_right).abs() < 1e-15,
"deriv at knot = {d_at_1}, expected right-slope = {expected_right}",
);
}
#[test]
fn build_trait_method_equivalent_to_new() {
let knots = [(0.0, 1.0), (1.0, 2.0)];
let a = Linear::new(&knots).unwrap();
let b = <Linear as Interpolator>::build(&knots).unwrap();
assert!((a.eval(0.5) - b.eval(0.5)).abs() < 1e-15);
assert_eq!(a.len(), b.len());
}
#[test]
fn len_and_is_empty() {
let interp = Linear::new(&[(0.0, 1.0), (1.0, 2.0), (2.0, 1.5)]).unwrap();
assert_eq!(interp.len(), 3);
assert!(!interp.is_empty());
}
#[test]
fn clone_yields_equivalent_interpolant() {
let interp = Linear::new(&[(0.0, 1.0), (1.0, 2.0)]).unwrap();
let copy = interp.clone();
assert!((interp.eval(0.5) - copy.eval(0.5)).abs() < 1e-15);
}
}