use crate::errors::CurveError;
use super::Interpolator;
#[derive(Debug, Clone)]
pub struct PiecewiseConstantForward {
times: Vec<f64>,
discounts: Vec<f64>,
forwards: Vec<f64>,
}
impl PiecewiseConstantForward {
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 discounts = Vec::with_capacity(n);
for (i, &(t, d)) in knots.iter().enumerate() {
if !t.is_finite() {
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 });
}
}
times.push(t);
discounts.push(d);
}
let mut forwards = Vec::with_capacity(n - 1);
for i in 0..n - 1 {
let dt = times[i + 1] - times[i];
let f = (discounts[i].ln() - discounts[i + 1].ln()) / dt;
forwards.push(f);
}
Ok(Self {
times,
discounts,
forwards,
})
}
#[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]
pub fn forward_rate(&self, t: f64) -> f64 {
let n = self.times.len();
if t <= self.times[0] {
return self.forwards[0];
}
if t > self.times[n - 1] {
return self.forwards[n - 2];
}
let i = self.locate_right_continuous(t);
self.forwards[i]
}
#[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 locate_right_continuous(&self, t: f64) -> usize {
let i = self.locate(t);
i
}
}
impl Interpolator for PiecewiseConstantForward {
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] {
let f0 = self.forwards[0];
return self.discounts[0] * (-f0 * (t - self.times[0])).exp();
}
if t >= self.times[n - 1] {
let f_last = self.forwards[n - 2];
return self.discounts[n - 1] * (-f_last * (t - self.times[n - 1])).exp();
}
let i = self.locate(t);
let t_lo = self.times[i];
let f_i = self.forwards[i];
self.discounts[i] * (-f_i * (t - t_lo)).exp()
}
fn deriv(&self, t: f64) -> Option<f64> {
let f = self.forward_rate(t);
let d = self.eval(t);
Some(-f * d)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_empty() {
let err = PiecewiseConstantForward::new(&[]).unwrap_err();
assert!(matches!(err, CurveError::TooFewNodes { found: 0 }));
}
#[test]
fn rejects_single_knot() {
let err = PiecewiseConstantForward::new(&[(0.0, 1.0)]).unwrap_err();
assert!(matches!(err, CurveError::TooFewNodes { found: 1 }));
}
#[test]
fn rejects_non_monotone_times() {
let err =
PiecewiseConstantForward::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 =
PiecewiseConstantForward::new(&[(0.0, 1.0), (1.0, 0.95), (1.0, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::DuplicateNode { .. }));
}
#[test]
fn rejects_negative_discount() {
let err = PiecewiseConstantForward::new(&[(0.0, 1.0), (1.0, -0.5)]).unwrap_err();
assert!(matches!(
err,
CurveError::NonPositiveDiscount { at_index: 1, .. }
));
}
#[test]
fn rejects_zero_discount() {
let err = PiecewiseConstantForward::new(&[(0.0, 1.0), (1.0, 0.0)]).unwrap_err();
assert!(matches!(
err,
CurveError::NonPositiveDiscount { at_index: 1, .. }
));
}
#[test]
fn rejects_nan_discount() {
let err = PiecewiseConstantForward::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 = PiecewiseConstantForward::new(&[(0.0, 1.0), (f64::NAN, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::InvalidTime { .. }));
}
#[test]
fn rejects_inf_time() {
let err = PiecewiseConstantForward::new(&[(0.0, 1.0), (f64::INFINITY, 0.9)]).unwrap_err();
assert!(matches!(err, CurveError::InvalidTime { .. }));
}
fn flat_forward_knots() -> [(f64, f64); 3] {
[
(0.0, 1.0),
(1.0, (-0.04_f64).exp()),
(2.0, (-0.10_f64).exp()),
]
}
#[test]
fn knot_reproduction_exact() {
let knots = flat_forward_knots();
let interp = PiecewiseConstantForward::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 forward_rate_segment_values() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
assert!((interp.forward_rate(0.5) - 0.04).abs() < 1e-15);
assert!((interp.forward_rate(1.5) - 0.06).abs() < 1e-15);
}
#[test]
fn forward_rate_flat_extrapolation() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
assert!((interp.forward_rate(-1.0) - 0.04).abs() < 1e-15);
assert!((interp.forward_rate(f64::NEG_INFINITY) - 0.04).abs() < 1e-15);
assert!((interp.forward_rate(3.0) - 0.06).abs() < 1e-15);
assert!((interp.forward_rate(f64::INFINITY) - 0.06).abs() < 1e-15);
}
#[test]
fn eval_midsegment_matches_exponential_formula() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
let v = interp.eval(1.5);
let expected = (-0.07_f64).exp();
assert!(
(v - expected).abs() < 1e-15,
"v = {v}, expected = {expected}"
);
}
#[test]
fn eval_first_segment_matches_exponential_formula() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
let v = interp.eval(0.25);
let expected = (-0.01_f64).exp();
assert!((v - expected).abs() < 1e-15);
}
#[test]
fn eval_extrapolation_extends_boundary_segment() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
let v_left = interp.eval(-0.5);
let expected_left = (0.02_f64).exp();
assert!((v_left - expected_left).abs() < 1e-15);
let v_right = interp.eval(2.5);
let expected_right = (-0.13_f64).exp();
assert!((v_right - expected_right).abs() < 1e-15);
}
#[test]
fn monotone_decreasing_input_produces_monotone_output() {
let knots = [
(0.0, 1.0),
(0.25, 0.9875),
(0.5, 0.9752),
(1.0, 0.9512),
(2.0, 0.9048),
(5.0, 0.7788),
];
let interp = PiecewiseConstantForward::new(&knots).unwrap();
let mut prev = interp.eval(0.0);
let mut t = 0.01_f64;
while t <= 5.0 {
let v = interp.eval(t);
assert!(
v <= prev + 1e-15,
"non-monotone at t = {t}: prev={prev}, v={v}"
);
prev = v;
t += 0.01;
}
}
#[test]
fn deriv_equals_minus_f_times_d() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
let t = 0.5_f64;
let d = interp.eval(t);
let f = interp.forward_rate(t);
let dy = interp.deriv(t).unwrap();
assert!((dy - (-f * d)).abs() < 1e-15);
let t = 1.5_f64;
let d = interp.eval(t);
let f = interp.forward_rate(t);
let dy = interp.deriv(t).unwrap();
assert!((dy - (-f * d)).abs() < 1e-15);
}
#[test]
fn deriv_finite_difference_interior() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
let t = 0.5_f64;
let dy = interp.deriv(t).unwrap();
let h = 1e-6_f64;
let fd = (interp.eval(t + h) - interp.eval(t - h)) / (2.0 * h);
assert!((dy - fd).abs() < 1e-8, "analytic={dy}, fd={fd}");
}
#[test]
fn deriv_in_extrapolation_region_is_flat_forward() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
let t = -0.5_f64;
let d = interp.eval(t);
let f0 = 0.04_f64;
let expected = -f0 * d;
let dy = interp.deriv(t).unwrap();
assert!((dy - expected).abs() < 1e-15);
let t = 2.5_f64;
let d = interp.eval(t);
let f_last = 0.06_f64;
let expected = -f_last * d;
let dy = interp.deriv(t).unwrap();
assert!((dy - expected).abs() < 1e-15);
}
#[test]
fn eval_matches_log_linear_on_same_knots_interior() {
use super::super::LogLinear;
let knots = [(0.0, 1.0), (0.5, 0.97), (1.0, 0.95), (2.0, 0.90)];
let pcf = PiecewiseConstantForward::new(&knots).unwrap();
let ll = LogLinear::new(&knots).unwrap();
let mut t = 0.0_f64;
while t <= 2.0 {
let a = pcf.eval(t);
let b = ll.eval(t);
assert!((a - b).abs() < 1e-14, "t={t}: pcf={a}, ll={b}");
t += 0.01;
}
}
#[test]
fn build_trait_method_equivalent_to_new() {
let knots = flat_forward_knots();
let a = PiecewiseConstantForward::new(&knots).unwrap();
let b = <PiecewiseConstantForward as Interpolator>::build(&knots).unwrap();
assert!((a.eval(1.5) - b.eval(1.5)).abs() < 1e-15);
}
#[test]
fn len_and_is_empty() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
assert_eq!(interp.len(), 3);
assert!(!interp.is_empty());
}
#[test]
fn clone_yields_equivalent_interpolant() {
let interp = PiecewiseConstantForward::new(&flat_forward_knots()).unwrap();
let copy = interp.clone();
assert!((interp.eval(1.5) - copy.eval(1.5)).abs() < 1e-15);
assert!((interp.forward_rate(0.5) - copy.forward_rate(0.5)).abs() < 1e-15);
}
}