1use serde::{Deserialize, Serialize};
28
29use crate::error::CoreError;
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
33#[serde(rename_all = "snake_case")]
34#[non_exhaustive]
35pub enum Interpolation {
36 #[default]
38 Linear,
39 NaturalCubic,
42}
43
44#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
46#[serde(rename_all = "snake_case")]
47#[non_exhaustive]
48pub enum Extrapolation {
49 #[default]
51 Clamp,
52 Linear,
55 Error,
57}
58
59#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
61#[serde(rename_all = "snake_case")]
62pub enum Side {
63 Below,
65 Above,
67}
68
69#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
71pub struct Lookup {
72 pub value: f64,
74 pub extrapolated: Option<Side>,
77}
78
79#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
84#[serde(try_from = "TableData", into = "TableData")]
85pub struct Table1D {
86 xs: Vec<f64>,
87 ys: Vec<f64>,
88 interpolation: Interpolation,
89 extrapolation: Extrapolation,
90 slopes: Vec<f64>,
92}
93
94#[derive(Serialize, Deserialize)]
96#[serde(deny_unknown_fields)]
97struct TableData {
98 x: Vec<f64>,
99 y: Vec<f64>,
100 #[serde(default)]
101 interpolation: Interpolation,
102 #[serde(default)]
103 extrapolation: Extrapolation,
104}
105
106impl TryFrom<TableData> for Table1D {
107 type Error = CoreError;
108
109 fn try_from(data: TableData) -> Result<Self, CoreError> {
110 Table1D::new(data.x, data.y, data.interpolation, data.extrapolation)
111 }
112}
113
114impl From<Table1D> for TableData {
115 fn from(table: Table1D) -> Self {
116 TableData {
117 x: table.xs,
118 y: table.ys,
119 interpolation: table.interpolation,
120 extrapolation: table.extrapolation,
121 }
122 }
123}
124
125impl Table1D {
126 pub const MIN_KNOTS: usize = 2;
128
129 pub fn new(
140 xs: Vec<f64>,
141 ys: Vec<f64>,
142 interpolation: Interpolation,
143 extrapolation: Extrapolation,
144 ) -> Result<Self, CoreError> {
145 if xs.len() != ys.len() {
146 return Err(CoreError::TableLengthMismatch {
147 xs: xs.len(),
148 ys: ys.len(),
149 });
150 }
151 if xs.len() < Self::MIN_KNOTS {
152 return Err(CoreError::TableTooShort {
153 min: Self::MIN_KNOTS,
154 got: xs.len(),
155 });
156 }
157 if let Some(index) = xs
158 .iter()
159 .zip(&ys)
160 .position(|(x, y)| !x.is_finite() || !y.is_finite())
161 {
162 return Err(CoreError::TableNotFinite { index });
163 }
164 if let Some(index) = xs.windows(2).position(|w| w[1] <= w[0]) {
165 return Err(CoreError::TableNotIncreasing { index: index + 1 });
166 }
167 if let Some(interval) = xs.windows(2).position(|w| !(w[1] - w[0]).is_finite()) {
168 return Err(CoreError::TableOverflow { interval });
169 }
170 let secants: Vec<f64> = xs
171 .windows(2)
172 .zip(ys.windows(2))
173 .map(|(x, y)| (y[1] - y[0]) / (x[1] - x[0]))
174 .collect();
175 if let Some(interval) = secants.iter().position(|s| !s.is_finite()) {
176 return Err(CoreError::TableOverflow { interval });
177 }
178 let slopes = match interpolation {
179 Interpolation::Linear => Vec::new(),
180 Interpolation::NaturalCubic => natural_spline_slopes(&xs, &secants),
181 };
182 if let Some(knot) = slopes.iter().position(|s| !s.is_finite()) {
183 return Err(CoreError::TableOverflow {
184 interval: knot.min(xs.len() - 2),
185 });
186 }
187 Ok(Self {
188 xs,
189 ys,
190 interpolation,
191 extrapolation,
192 slopes,
193 })
194 }
195
196 pub fn xs(&self) -> &[f64] {
198 &self.xs
199 }
200
201 pub fn ys(&self) -> &[f64] {
203 &self.ys
204 }
205
206 pub fn interpolation(&self) -> Interpolation {
208 self.interpolation
209 }
210
211 pub fn extrapolation(&self) -> Extrapolation {
213 self.extrapolation
214 }
215
216 pub fn domain(&self) -> (f64, f64) {
218 (self.first_x(), self.last_x())
219 }
220
221 pub fn lookup(&self, x: f64) -> Result<Lookup, CoreError> {
231 if x.is_nan() {
232 return Err(CoreError::NanLookup);
233 }
234 let (first, last) = self.domain();
235 let side = if x < first {
236 Some(Side::Below)
237 } else if x > last {
238 Some(Side::Above)
239 } else {
240 None
241 };
242 let value = match side {
243 None => self.interpolate(x),
244 Some(side) => {
245 let (x_end, y_end, slope) = match side {
246 Side::Below => (first, self.first_y(), self.end_slope(Side::Below)),
247 Side::Above => (last, self.last_y(), self.end_slope(Side::Above)),
248 };
249 match self.extrapolation {
250 Extrapolation::Clamp => y_end,
251 Extrapolation::Linear => {
252 let value = y_end + slope * (x - x_end);
253 if !value.is_finite() {
254 return Err(CoreError::ExtrapolationOverflow { x });
255 }
256 value
257 }
258 Extrapolation::Error => {
259 return Err(CoreError::OutOfRange {
260 x,
261 min: first,
262 max: last,
263 });
264 }
265 }
266 }
267 };
268 Ok(Lookup {
269 value,
270 extrapolated: side,
271 })
272 }
273
274 pub fn eval(&self, x: f64) -> Result<f64, CoreError> {
280 self.lookup(x).map(|lookup| lookup.value)
281 }
282
283 fn first_x(&self) -> f64 {
284 self.xs.first().copied().unwrap_or(f64::NAN)
285 }
286
287 fn last_x(&self) -> f64 {
288 self.xs.last().copied().unwrap_or(f64::NAN)
289 }
290
291 fn first_y(&self) -> f64 {
292 self.ys.first().copied().unwrap_or(f64::NAN)
293 }
294
295 fn last_y(&self) -> f64 {
296 self.ys.last().copied().unwrap_or(f64::NAN)
297 }
298
299 fn end_slope(&self, side: Side) -> f64 {
301 let n = self.xs.len();
302 match (self.interpolation, side) {
303 (Interpolation::Linear, Side::Below) => secant(&self.xs, &self.ys, 0),
304 (Interpolation::Linear, Side::Above) => secant(&self.xs, &self.ys, n - 2),
305 (Interpolation::NaturalCubic, Side::Below) => self.slopes[0],
306 (Interpolation::NaturalCubic, Side::Above) => self.slopes[n - 1],
307 }
308 }
309
310 fn interpolate(&self, x: f64) -> f64 {
312 let n = self.xs.len();
315 let i = self.xs.partition_point(|&xi| xi <= x).clamp(1, n - 1) - 1;
316 let (x0, x1) = (self.xs[i], self.xs[i + 1]);
317 let (y0, y1) = (self.ys[i], self.ys[i + 1]);
318 let h = x1 - x0;
319 let t = (x - x0) / h;
320 match self.interpolation {
321 Interpolation::Linear => (1.0 - t) * y0 + t * y1,
323 Interpolation::NaturalCubic => {
324 let u = 1.0 - t;
327 let h00 = (1.0 + 2.0 * t) * u * u;
328 let h10 = t * u * u;
329 let h01 = t * t * (3.0 - 2.0 * t);
330 let h11 = -t * t * u;
331 h00 * y0 + h10 * h * self.slopes[i] + h01 * y1 + h11 * h * self.slopes[i + 1]
332 }
333 }
334 }
335}
336
337fn secant(xs: &[f64], ys: &[f64], i: usize) -> f64 {
339 (ys[i + 1] - ys[i]) / (xs[i + 1] - xs[i])
340}
341
342fn natural_spline_slopes(xs: &[f64], secants: &[f64]) -> Vec<f64> {
345 let n = xs.len();
346 let mut lower = vec![0.0; n];
348 let mut diag = vec![0.0; n];
349 let mut upper = vec![0.0; n];
350 let mut rhs = vec![0.0; n];
351 let h_first = xs[1] - xs[0];
354 diag[0] = 2.0 * h_first;
355 upper[0] = h_first;
356 rhs[0] = 3.0 * h_first * secants[0];
357 for i in 1..n - 1 {
358 let h_prev = xs[i] - xs[i - 1];
359 let h_next = xs[i + 1] - xs[i];
360 lower[i] = h_next;
361 diag[i] = 2.0 * (h_prev + h_next);
362 upper[i] = h_prev;
363 rhs[i] = 3.0 * (h_next * secants[i - 1] + h_prev * secants[i]);
364 }
365 let h_last = xs[n - 1] - xs[n - 2];
366 lower[n - 1] = h_last;
367 diag[n - 1] = 2.0 * h_last;
368 rhs[n - 1] = 3.0 * h_last * secants[n - 2];
369
370 for i in 1..n {
372 let factor = lower[i] / diag[i - 1];
373 diag[i] -= factor * upper[i - 1];
374 rhs[i] -= factor * rhs[i - 1];
375 }
376 let mut slopes = vec![0.0; n];
378 slopes[n - 1] = rhs[n - 1] / diag[n - 1];
379 for i in (0..n - 1).rev() {
380 slopes[i] = (rhs[i] - upper[i] * slopes[i + 1]) / diag[i];
381 }
382 slopes
383}
384
385#[cfg(test)]
386mod tests {
387 use proptest::prelude::*;
388
389 use super::*;
390
391 fn table(xs: &[f64], ys: &[f64], i: Interpolation, e: Extrapolation) -> Table1D {
392 Table1D::new(xs.to_vec(), ys.to_vec(), i, e).unwrap()
393 }
394
395 fn second_derivative(table: &Table1D, i: usize, right: bool) -> f64 {
398 let h = table.xs[i + 1] - table.xs[i];
399 let d = secant(&table.xs, &table.ys, i);
400 let (s0, s1) = (table.slopes[i], table.slopes[i + 1]);
401 if right {
402 (-6.0 * d + 2.0 * s0 + 4.0 * s1) / h
403 } else {
404 (6.0 * d - 4.0 * s0 - 2.0 * s1) / h
405 }
406 }
407
408 #[test]
409 fn rejects_malformed_knots() {
410 let lin = Interpolation::Linear;
411 let clamp = Extrapolation::Clamp;
412 assert_eq!(
413 Table1D::new(vec![0.0], vec![1.0], lin, clamp),
414 Err(CoreError::TableTooShort { min: 2, got: 1 })
415 );
416 assert_eq!(
417 Table1D::new(vec![0.0, 1.0], vec![1.0], lin, clamp),
418 Err(CoreError::TableLengthMismatch { xs: 2, ys: 1 })
419 );
420 assert_eq!(
421 Table1D::new(vec![0.0, 1.0, 1.0], vec![1.0, 2.0, 3.0], lin, clamp),
422 Err(CoreError::TableNotIncreasing { index: 2 })
423 );
424 assert_eq!(
425 Table1D::new(vec![0.0, f64::NAN], vec![1.0, 2.0], lin, clamp),
426 Err(CoreError::TableNotFinite { index: 1 })
427 );
428 assert_eq!(
429 Table1D::new(vec![0.0, 1.0], vec![f64::INFINITY, 2.0], lin, clamp),
430 Err(CoreError::TableNotFinite { index: 0 })
431 );
432 assert_eq!(
434 Table1D::new(vec![0.0, 1e-310], vec![-1e300, 1e300], lin, clamp),
435 Err(CoreError::TableOverflow { interval: 0 })
436 );
437 assert_eq!(
439 Table1D::new(vec![-1e308, 1e308], vec![0.0, 1.0], lin, clamp),
440 Err(CoreError::TableOverflow { interval: 0 })
441 );
442 }
443
444 #[test]
445 fn subnormal_intervals_keep_the_spline_finite() {
446 let xs = vec![0.0, 1e-310, 2e-310, 4e-310];
447 let ys = vec![0.0, 1e-300, 0.0, 1e-300];
448 let t = Table1D::new(xs, ys, Interpolation::NaturalCubic, Extrapolation::Clamp).unwrap();
449 assert!(t.slopes.iter().all(|s| s.is_finite()), "{:?}", t.slopes);
450 assert_eq!(t.eval(1e-310).unwrap(), 1e-300);
451 assert!(t.eval(1.5e-310).unwrap().is_finite());
452 }
453
454 #[test]
455 fn linear_extrapolation_to_infinity_is_an_error() {
456 let flat = table(
457 &[0.0, 1.0],
458 &[2.0, 2.0],
459 Interpolation::Linear,
460 Extrapolation::Linear,
461 );
462 assert_eq!(
463 flat.lookup(f64::INFINITY),
464 Err(CoreError::ExtrapolationOverflow { x: f64::INFINITY })
465 );
466 let steep = table(
467 &[0.0, 1.0],
468 &[0.0, 1e300],
469 Interpolation::Linear,
470 Extrapolation::Linear,
471 );
472 assert!(matches!(
473 steep.lookup(1e10),
474 Err(CoreError::ExtrapolationOverflow { .. })
475 ));
476 let clamp = table(
478 &[0.0, 1.0],
479 &[2.0, 3.0],
480 Interpolation::NaturalCubic,
481 Extrapolation::Clamp,
482 );
483 assert_eq!(clamp.eval(f64::NEG_INFINITY).unwrap(), 2.0);
484 }
485
486 #[test]
487 fn linear_interpolates_and_flags_extrapolation() {
488 let t = table(
489 &[0.0, 1.0, 3.0],
490 &[0.0, 10.0, 0.0],
491 Interpolation::Linear,
492 Extrapolation::Linear,
493 );
494 assert_eq!(t.eval(0.5).unwrap(), 5.0);
495 assert_eq!(t.eval(2.0).unwrap(), 5.0);
496 assert_eq!(
497 t.lookup(3.0).unwrap(),
498 Lookup {
499 value: 0.0,
500 extrapolated: None
501 }
502 );
503 assert_eq!(
504 t.lookup(-1.0).unwrap(),
505 Lookup {
506 value: -10.0,
507 extrapolated: Some(Side::Below)
508 }
509 );
510 assert_eq!(
511 t.lookup(4.0).unwrap(),
512 Lookup {
513 value: -5.0,
514 extrapolated: Some(Side::Above)
515 }
516 );
517 assert_eq!(t.lookup(f64::NAN), Err(CoreError::NanLookup));
518
519 let clamped = table(
520 &[0.0, 1.0],
521 &[2.0, 4.0],
522 Interpolation::Linear,
523 Extrapolation::Clamp,
524 );
525 assert_eq!(clamped.eval(-5.0).unwrap(), 2.0);
526 assert_eq!(clamped.eval(f64::INFINITY).unwrap(), 4.0);
527
528 let strict = table(
529 &[0.0, 1.0],
530 &[2.0, 4.0],
531 Interpolation::Linear,
532 Extrapolation::Error,
533 );
534 assert_eq!(
535 strict.lookup(1.5),
536 Err(CoreError::OutOfRange {
537 x: 1.5,
538 min: 0.0,
539 max: 1.0
540 })
541 );
542 assert_eq!(strict.eval(1.0).unwrap(), 4.0);
543 }
544
545 #[test]
549 fn natural_spline_matches_the_three_knot_closed_form() {
550 let t = table(
551 &[0.0, 1.0, 2.0],
552 &[0.0, 1.0, 0.0],
553 Interpolation::NaturalCubic,
554 Extrapolation::Linear,
555 );
556 assert_eq!(t.slopes, vec![1.5, 0.0, -1.5]);
557 for k in 0..=10 {
558 let x = f64::from(k) / 10.0;
559 let exact = 1.5 * x - 0.5 * x.powi(3);
560 assert!((t.eval(x).unwrap() - exact).abs() < 1e-15, "x = {x}");
561 assert!(
563 (t.eval(2.0 - x).unwrap() - exact).abs() < 1e-15,
564 "x = {}",
565 2.0 - x
566 );
567 }
568 assert_eq!(t.eval(-2.0).unwrap(), -3.0);
570 assert_eq!(t.eval(3.0).unwrap(), -1.5);
571 }
572
573 fn knots() -> impl Strategy<Value = (Vec<f64>, Vec<f64>)> {
575 (2usize..24).prop_flat_map(|n| {
576 (
577 -1e3..1e3f64,
578 prop::collection::vec(1e-3..10.0f64, n - 1),
579 prop::collection::vec(-1e3..1e3f64, n),
580 )
581 .prop_map(|(start, gaps, ys)| {
582 let mut xs = vec![start];
583 for gap in gaps {
584 let last = xs[xs.len() - 1];
585 xs.push(last + gap);
586 }
587 (xs, ys)
588 })
589 })
590 }
591
592 proptest! {
593 #[test]
594 fn both_kinds_reproduce_the_knots((xs, ys) in knots()) {
595 for kind in [Interpolation::Linear, Interpolation::NaturalCubic] {
596 let t = table(&xs, &ys, kind, Extrapolation::Error);
597 for (x, y) in xs.iter().zip(&ys) {
598 prop_assert_eq!(t.lookup(*x).unwrap(), Lookup { value: *y, extrapolated: None });
599 }
600 }
601 }
602
603 #[test]
604 fn linear_stays_within_each_interval((xs, ys) in knots(), frac in 0.0..1.0f64) {
605 let t = table(&xs, &ys, Interpolation::Linear, Extrapolation::Error);
606 for i in 0..xs.len() - 1 {
607 let x = xs[i] + frac * (xs[i + 1] - xs[i]);
608 let y = t.eval(x).unwrap();
609 let (lo, hi) = (ys[i].min(ys[i + 1]), ys[i].max(ys[i + 1]));
610 prop_assert!(y >= lo - 1e-9 && y <= hi + 1e-9, "{y} outside [{lo}, {hi}]");
611 }
612 }
613
614 #[test]
615 fn natural_spline_is_c2_with_free_ends((xs, ys) in knots()) {
616 let t = table(&xs, &ys, Interpolation::NaturalCubic, Extrapolation::Clamp);
617 let n = xs.len();
618 let scale = ys.iter().fold(1.0f64, |m, y| m.max(y.abs()))
619 / xs.windows(2).fold(f64::INFINITY, |m, w| m.min(w[1] - w[0])).powi(2);
620 prop_assert!(second_derivative(&t, 0, false).abs() <= 1e-9 * scale);
621 prop_assert!(second_derivative(&t, n - 2, true).abs() <= 1e-9 * scale);
622 for i in 1..n - 1 {
623 let left = second_derivative(&t, i - 1, true);
624 let right = second_derivative(&t, i, false);
625 prop_assert!((left - right).abs() <= 1e-9 * scale, "knot {i}: {left} vs {right}");
626 }
627 }
628
629 #[test]
630 fn natural_spline_reproduces_straight_lines(
631 (xs, _) in knots(), a in -100.0..100.0f64, b in -100.0..100.0f64, frac in 0.0..1.0f64,
632 ) {
633 let ys: Vec<f64> = xs.iter().map(|x| a + b * x).collect();
634 let t = table(&xs, &ys, Interpolation::NaturalCubic, Extrapolation::Linear);
635 let (first, last) = t.domain();
636 for x in [first + frac * (last - first), first - 5.0, last + 5.0] {
637 let exact = a + b * x;
638 prop_assert!((t.eval(x).unwrap() - exact).abs() <= 1e-9 * (1.0 + exact.abs()));
639 }
640 }
641
642 #[test]
643 fn extrapolation_follows_the_policy((xs, ys) in knots(), beyond in 1e-6..100.0f64) {
644 let n = xs.len();
645 for kind in [Interpolation::Linear, Interpolation::NaturalCubic] {
646 let below = xs[0] - beyond;
647 let above = xs[n - 1] + beyond;
648
649 let clamp = table(&xs, &ys, kind, Extrapolation::Clamp);
650 prop_assert_eq!(clamp.lookup(below).unwrap(), Lookup { value: ys[0], extrapolated: Some(Side::Below) });
651 prop_assert_eq!(clamp.lookup(above).unwrap(), Lookup { value: ys[n - 1], extrapolated: Some(Side::Above) });
652
653 let strict = table(&xs, &ys, kind, Extrapolation::Error);
654 prop_assert!(matches!(strict.lookup(below), Err(CoreError::OutOfRange { .. })), "below");
655 prop_assert!(matches!(strict.lookup(above), Err(CoreError::OutOfRange { .. })), "above");
656
657 let line = table(&xs, &ys, kind, Extrapolation::Linear);
659 let slope = line.end_slope(Side::Above);
660 let y = line.lookup(above).unwrap();
661 prop_assert_eq!(y.extrapolated, Some(Side::Above));
662 prop_assert!((y.value - (ys[n - 1] + slope * beyond)).abs() <= 1e-9 * (1.0 + y.value.abs()));
663 }
664 }
665
666 #[test]
667 fn serde_round_trip_rebuilds_the_table((xs, ys) in knots()) {
668 let t = table(&xs, &ys, Interpolation::NaturalCubic, Extrapolation::Linear);
669 let json = serde_json::to_string(&t).unwrap();
670 let back: Table1D = serde_json::from_str(&json).unwrap();
671 prop_assert_eq!(back, t);
672 }
673 }
674
675 #[test]
676 fn deserializing_checks_the_knots() {
677 let bad = r#"{"x": [0.0, 0.0], "y": [1.0, 2.0]}"#;
678 assert!(serde_json::from_str::<Table1D>(bad).is_err());
679 let good = r#"{"x": [0.0, 1.0], "y": [1.0, 2.0]}"#;
680 let t: Table1D = serde_json::from_str(good).unwrap();
681 assert_eq!(t.interpolation(), Interpolation::Linear);
682 assert_eq!(t.extrapolation(), Extrapolation::Clamp);
683 }
684}