Skip to main content

yield_curves/
linear.rs

1//! Piecewise-linear interpolation with flat extrapolation.
2
3use crate::{validate::validate_and_sort, YieldCurveError, YieldCurveInterpolator};
4
5/// Piecewise-linear interpolation between anchor points. Extrapolates flat
6/// outside the observed range.
7///
8/// O(log n) evaluation via binary search.
9#[derive(Debug, Clone)]
10pub struct LinearCurve {
11    points: Vec<(f64, f64)>,
12}
13
14impl LinearCurve {
15    /// Fits from `(t_years, rate)` anchor points. Requires at least 2 points.
16    pub fn fit(points: &[(f64, f64)]) -> Result<Self, YieldCurveError> {
17        let points = validate_and_sort(points, "linear", 2)?;
18        Ok(Self { points })
19    }
20}
21
22impl YieldCurveInterpolator for LinearCurve {
23    fn rate_at(&self, t_years: f64) -> f64 {
24        let pts = &self.points;
25        if t_years <= pts[0].0 {
26            return pts[0].1;
27        }
28        if t_years >= pts.last().unwrap().0 {
29            return pts.last().unwrap().1;
30        }
31        let idx = pts.partition_point(|p| p.0 < t_years);
32        let (x0, y0) = pts[idx - 1];
33        let (x1, y1) = pts[idx];
34        let t = (t_years - x0) / (x1 - x0);
35        y0 + t * (y1 - y0)
36    }
37
38    fn method_name(&self) -> &'static str {
39        "linear"
40    }
41
42    fn observed_range(&self) -> (f64, f64) {
43        (self.points[0].0, self.points.last().unwrap().0)
44    }
45}
46
47#[cfg(test)]
48mod tests {
49    use super::*;
50
51    fn approx_eq(a: f64, b: f64, eps: f64) -> bool {
52        (a - b).abs() < eps
53    }
54
55    #[test]
56    fn passes_through_anchors() {
57        let curve = LinearCurve::fit(&[(0.25, 13.0), (0.5, 13.5), (1.0, 14.0)]).unwrap();
58        assert!(approx_eq(curve.rate_at(0.25), 13.0, 1e-10));
59        assert!(approx_eq(curve.rate_at(0.5), 13.5, 1e-10));
60        assert!(approx_eq(curve.rate_at(1.0), 14.0, 1e-10));
61    }
62
63    #[test]
64    fn midpoint_is_average() {
65        let curve = LinearCurve::fit(&[(0.0, 10.0), (1.0, 20.0)]).unwrap();
66        assert!(approx_eq(curve.rate_at(0.5), 15.0, 1e-10));
67    }
68
69    #[test]
70    fn flat_extrapolation() {
71        let curve = LinearCurve::fit(&[(0.25, 13.0), (1.0, 14.0)]).unwrap();
72        assert_eq!(curve.rate_at(0.1), 13.0);
73        assert_eq!(curve.rate_at(5.0), 14.0);
74    }
75
76    #[test]
77    fn rejects_single_point() {
78        let err = LinearCurve::fit(&[(0.25, 13.0)]).unwrap_err();
79        assert!(matches!(err, YieldCurveError::InsufficientData { .. }));
80    }
81
82    #[test]
83    fn rejects_nan() {
84        let err = LinearCurve::fit(&[(0.25, f64::NAN), (0.5, 13.5)]).unwrap_err();
85        assert!(matches!(err, YieldCurveError::InvalidPoint(_)));
86    }
87
88    #[test]
89    fn rejects_duplicate_vertex() {
90        let err = LinearCurve::fit(&[(0.25, 13.0), (0.25, 13.5)]).unwrap_err();
91        assert!(matches!(err, YieldCurveError::InvalidPoint(_)));
92    }
93
94    #[test]
95    fn method_name_stable() {
96        let curve = LinearCurve::fit(&[(1.0, 1.0), (2.0, 2.0)]).unwrap();
97        assert_eq!(curve.method_name(), "linear");
98    }
99
100    #[test]
101    fn observed_range_reflects_input() {
102        let curve = LinearCurve::fit(&[(0.25, 13.0), (0.5, 13.5), (3.0, 14.0)]).unwrap();
103        assert_eq!(curve.observed_range(), (0.25, 3.0));
104    }
105}