1use crate::{validate::validate_and_sort, YieldCurveError, YieldCurveInterpolator};
4
5#[derive(Debug, Clone)]
10pub struct LinearCurve {
11 points: Vec<(f64, f64)>,
12}
13
14impl LinearCurve {
15 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}