use crate::errors::CurveError;
use super::Interpolator;
#[derive(Debug, Clone)]
pub struct LinearInZero {
times: Vec<f64>,
zero_rates: Vec<f64>,
}
impl LinearInZero {
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 zero_rates = Vec::with_capacity(n);
for (i, &(t, d)) in knots.iter().enumerate() {
if !t.is_finite() || t < 0.0 {
return Err(CurveError::InvalidTime { t });
}
if !d.is_finite() || d <= 0.0 {
return Err(CurveError::NonPositiveDiscount {
at_index: i,
value: d,
});
}
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 });
}
}
#[allow(clippy::float_cmp)]
let is_anchor = t == 0.0;
let z_i = if is_anchor {
#[allow(clippy::float_cmp)]
let d_is_unit = d == 1.0;
if !d_is_unit {
return Err(CurveError::AnchorNotUnit);
}
f64::NAN
} else {
-d.ln() / t
};
times.push(t);
zero_rates.push(z_i);
}
if zero_rates[0].is_nan() {
zero_rates[0] = zero_rates[1];
}
Ok(Self { times, zero_rates })
}
#[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
}
#[inline]
fn zero_rate_at(&self, t: f64) -> f64 {
let n = self.times.len();
if t <= self.times[0] {
return self.zero_rates[0];
}
if t >= self.times[n - 1] {
return self.zero_rates[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.zero_rates[i] + w * self.zero_rates[i + 1]
}
}
impl Interpolator for LinearInZero {
fn build(knots: &[(f64, f64)]) -> Result<Self, CurveError> {
Self::new(knots)
}
fn eval(&self, t: f64) -> f64 {
let z = self.zero_rate_at(t);
(-z * t).exp()
}
fn deriv(&self, t: f64) -> Option<f64> {
let n = self.times.len();
let discount = self.eval(t);
let zero = self.zero_rate_at(t);
if t <= self.times[0] || t >= self.times[n - 1] {
return Some(-zero * discount);
}
let idx = self.locate(t);
let t_lo = self.times[idx];
let t_hi = self.times[idx + 1];
let slope_z = (self.zero_rates[idx + 1] - self.zero_rates[idx]) / (t_hi - t_lo);
Some(discount * (-(slope_z * t + zero)))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_empty() {
let err = LinearInZero::new(&[]).unwrap_err();
assert!(matches!(err, CurveError::TooFewNodes { found: 0 }));
}
#[test]
fn rejects_single_knot() {
let err = LinearInZero::new(&[(0.0, 1.0)]).unwrap_err();
assert!(matches!(err, CurveError::TooFewNodes { found: 1 }));
}
#[test]
fn rejects_non_monotone_times() {
let err = LinearInZero::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 = LinearInZero::new(&[(0.0, 1.0), (1.0, 0.95), (1.0, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::DuplicateNode { .. }));
}
#[test]
fn rejects_negative_value() {
let err = LinearInZero::new(&[(0.0, 1.0), (1.0, -0.5)]).unwrap_err();
assert!(matches!(
err,
CurveError::NonPositiveDiscount { at_index: 1, .. }
));
}
#[test]
fn rejects_zero_value() {
let err = LinearInZero::new(&[(0.0, 1.0), (1.0, 0.0)]).unwrap_err();
assert!(matches!(
err,
CurveError::NonPositiveDiscount { at_index: 1, .. }
));
}
#[test]
fn rejects_nan_value() {
let err = LinearInZero::new(&[(0.0, 1.0), (1.0, f64::NAN)]).unwrap_err();
assert!(matches!(
err,
CurveError::NonPositiveDiscount { at_index: 1, .. }
));
}
#[test]
fn rejects_nan_time() {
let err = LinearInZero::new(&[(0.0, 1.0), (f64::NAN, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::InvalidTime { .. }));
}
#[test]
fn rejects_inf_time() {
let err = LinearInZero::new(&[(0.0, 1.0), (f64::INFINITY, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::InvalidTime { .. }));
}
#[test]
fn rejects_negative_time() {
let err = LinearInZero::new(&[(-0.5, 1.0), (1.0, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::InvalidTime { .. }));
}
#[test]
fn rejects_anchor_not_unit_discount() {
let err = LinearInZero::new(&[(0.0, 0.99), (1.0, 0.95)]).unwrap_err();
assert!(matches!(err, CurveError::AnchorNotUnit));
}
#[test]
fn knot_reproduction_exact() {
let knots = [
(0.0_f64, 1.0_f64),
(0.5, (-0.05_f64 * 0.5).exp()),
(1.0, (-0.04_f64 * 1.0).exp()),
(2.0, (-0.045_f64 * 2.0).exp()),
(5.0, (-0.05_f64 * 5.0).exp()),
];
let interp = LinearInZero::new(&knots).unwrap();
for &(t, d) in &knots {
let v = interp.eval(t);
assert!((v - d).abs() < 1e-15, "knot ({t}, {d}) -> {v}");
}
}
#[test]
fn midpoint_constant_zero_rate() {
let interp = LinearInZero::new(&[
(0.0_f64, 1.0_f64),
(1.0, (-0.04_f64 * 1.0).exp()),
(2.0, (-0.04_f64 * 2.0).exp()),
])
.unwrap();
let v = interp.eval(1.5);
let expected = (-0.04_f64 * 1.5).exp();
assert!((v - expected).abs() < 1e-15, "got {v}, expected {expected}");
}
#[test]
fn midpoint_varying_zero_rate() {
let interp = LinearInZero::new(&[
(1.0_f64, (-0.05_f64 * 1.0).exp()),
(2.0, (-0.04_f64 * 2.0).exp()),
])
.unwrap();
let v = interp.eval(1.5);
let expected = (-0.045_f64 * 1.5).exp();
assert!(
(v - expected).abs() < 1e-15,
"midpoint z = 0.045 check: got {v}, expected {expected}",
);
assert!(
(v - 0.9348_f64).abs() < 1e-3,
"exp(-0.0675) magnitude: got {v}",
);
}
#[test]
fn anchor_t_zero_extension() {
let interp = LinearInZero::new(&[(0.0_f64, 1.0_f64), (1.0, (-0.05_f64).exp())]).unwrap();
assert!((interp.eval(0.0) - 1.0).abs() < 1e-15);
let v = interp.eval(0.25);
let expected = (-0.05_f64 * 0.25).exp();
assert!((v - expected).abs() < 1e-15);
}
#[test]
fn flat_extrapolation_in_z_right() {
let interp = LinearInZero::new(&[
(1.0_f64, (-0.04_f64 * 1.0).exp()),
(2.0, (-0.05_f64 * 2.0).exp()),
])
.unwrap();
let v = interp.eval(3.0);
let expected = (-0.05_f64 * 3.0).exp();
assert!((v - expected).abs() < 1e-15, "got {v}, expected {expected}");
let v_large = interp.eval(20.0);
let expected_large = (-0.05_f64 * 20.0).exp();
assert!((v_large - expected_large).abs() < 1e-15);
}
#[test]
fn flat_extrapolation_in_z_left() {
let interp = LinearInZero::new(&[
(1.0_f64, (-0.04_f64 * 1.0).exp()),
(2.0, (-0.05_f64 * 2.0).exp()),
])
.unwrap();
let v = interp.eval(0.5);
let expected = (-0.04_f64 * 0.5).exp();
assert!((v - expected).abs() < 1e-15, "got {v}, expected {expected}");
}
#[test]
fn monotone_decreasing_input_produces_positive_output() {
let knots = [
(0.0_f64, 1.0_f64),
(0.25, (-0.05_f64 * 0.25).exp()),
(0.5, (-0.05_f64 * 0.5).exp()),
(1.0, (-0.05_f64 * 1.0).exp()),
(2.0, (-0.05_f64 * 2.0).exp()),
(5.0, (-0.05_f64 * 5.0).exp()),
];
let interp = LinearInZero::new(&knots).unwrap();
let mut t = 0.0_f64;
while t <= 5.0 {
let v = interp.eval(t);
assert!(v > 0.0 && v <= 1.0 + 1e-15, "out of (0, 1] at t = {t}: {v}");
t += 0.01;
}
}
#[test]
fn deriv_finite_difference_interior() {
let knots = [
(1.0_f64, (-0.05_f64 * 1.0).exp()),
(2.0, (-0.04_f64 * 2.0).exp()),
(4.0, (-0.045_f64 * 4.0).exp()),
];
let interp = LinearInZero::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-7, "analytic={dy_dt}, fd={fd}");
}
#[test]
fn deriv_extrapolation_uses_flat_zero_rate() {
let interp = LinearInZero::new(&[
(1.0_f64, (-0.04_f64 * 1.0).exp()),
(2.0, (-0.05_f64 * 2.0).exp()),
])
.unwrap();
let t = 3.0_f64;
let d = interp.deriv(t).unwrap();
let expected = -0.05_f64 * interp.eval(t);
assert!((d - expected).abs() < 1e-12);
}
#[test]
fn build_trait_method_equivalent_to_new() {
let knots = [(0.0_f64, 1.0_f64), (1.0, (-0.04_f64).exp())];
let a = LinearInZero::new(&knots).unwrap();
let b = <LinearInZero as Interpolator>::build(&knots).unwrap();
assert!((a.eval(0.5) - b.eval(0.5)).abs() < 1e-15);
}
#[test]
fn len_and_is_empty() {
let interp = LinearInZero::new(&[
(0.0_f64, 1.0_f64),
(1.0, (-0.04_f64).exp()),
(2.0, (-0.045_f64 * 2.0).exp()),
])
.unwrap();
assert_eq!(interp.len(), 3);
assert!(!interp.is_empty());
}
#[test]
fn clone_yields_equivalent_interpolant() {
let interp = LinearInZero::new(&[(0.0_f64, 1.0_f64), (1.0, (-0.04_f64).exp())]).unwrap();
let copy = interp.clone();
assert!((interp.eval(0.5) - copy.eval(0.5)).abs() < 1e-15);
}
}