Skip to main content

fin_primitives/options/
surface.rs

1//! # Module: options::surface
2//!
3//! ## Responsibility
4//! Volatility surface: a grid of implied vols over (strike, expiry) space,
5//! with bilinear interpolation, ATM vol, term structure, and smile extraction.
6
7/// A single point on the volatility surface.
8#[derive(Debug, Clone, Copy, serde::Serialize, serde::Deserialize)]
9pub struct VolPoint {
10    /// Strike price.
11    pub strike: f64,
12    /// Time to expiry in years.
13    pub expiry: f64,
14    /// Implied volatility at this (strike, expiry).
15    pub implied_vol: f64,
16}
17
18/// Volatility smile at a fixed expiry: a set of (strike, implied_vol) pairs.
19#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
20pub struct VolSmile {
21    /// Time to expiry (years) for this smile.
22    pub expiry: f64,
23    /// (strike, implied_vol) pairs sorted by strike ascending.
24    pub points: Vec<(f64, f64)>,
25}
26
27/// A volatility surface built from a collection of `VolPoint`s.
28///
29/// Internally stores a sorted grid of unique strikes and expiries, with
30/// bilinear interpolation for queries inside the grid.
31#[derive(Debug, Clone)]
32pub struct VolSurface {
33    /// Unique strikes, sorted ascending.
34    strikes: Vec<f64>,
35    /// Unique expiries, sorted ascending.
36    expiries: Vec<f64>,
37    /// Grid[i_expiry][i_strike] = implied_vol.
38    grid: Vec<Vec<f64>>,
39}
40
41impl VolSurface {
42    /// Build a `VolSurface` from a collection of `VolPoint`s.
43    ///
44    /// Duplicate (strike, expiry) pairs are averaged. Gaps in the grid are
45    /// filled with the nearest known value (nearest-neighbour fallback).
46    pub fn from_points(points: Vec<VolPoint>) -> Self {
47        // Collect unique strikes and expiries
48        let mut strike_set: Vec<f64> = points.iter().map(|p| p.strike).collect();
49        strike_set.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
50        strike_set.dedup_by(|a, b| (*a - *b).abs() < 1e-12);
51
52        let mut expiry_set: Vec<f64> = points.iter().map(|p| p.expiry).collect();
53        expiry_set.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
54        expiry_set.dedup_by(|a, b| (*a - *b).abs() < 1e-12);
55
56        let n_exp = expiry_set.len();
57        let n_str = strike_set.len();
58
59        // Accumulators for averaging duplicate entries
60        let mut sum_grid = vec![vec![0.0_f64; n_str]; n_exp];
61        let mut cnt_grid = vec![vec![0_u32; n_str]; n_exp];
62
63        for p in &points {
64            let i_exp = expiry_set
65                .iter()
66                .position(|&e| (e - p.expiry).abs() < 1e-12)
67                .unwrap_or(0);
68            let i_str = strike_set
69                .iter()
70                .position(|&s| (s - p.strike).abs() < 1e-12)
71                .unwrap_or(0);
72            sum_grid[i_exp][i_str] += p.implied_vol;
73            cnt_grid[i_exp][i_str] += 1;
74        }
75
76        // Average, filling zeros with NaN for gap detection
77        let mut grid = vec![vec![f64::NAN; n_str]; n_exp];
78        for i_exp in 0..n_exp {
79            for i_str in 0..n_str {
80                let c = cnt_grid[i_exp][i_str];
81                if c > 0 {
82                    grid[i_exp][i_str] = sum_grid[i_exp][i_str] / f64::from(c);
83                }
84            }
85        }
86
87        // Fill NaN gaps with nearest known value (simple sweep)
88        Self::fill_gaps(&mut grid, n_exp, n_str);
89
90        Self { strikes: strike_set, expiries: expiry_set, grid }
91    }
92
93    /// Bilinear interpolation for (strike, expiry).
94    ///
95    /// Returns `None` if the query is outside the grid boundaries.
96    pub fn interpolate(&self, strike: f64, expiry: f64) -> Option<f64> {
97        if self.strikes.is_empty() || self.expiries.is_empty() {
98            return None;
99        }
100        // Check bounds
101        if strike < *self.strikes.first()? || strike > *self.strikes.last()? {
102            return None;
103        }
104        if expiry < *self.expiries.first()? || expiry > *self.expiries.last()? {
105            return None;
106        }
107
108        let (i0, i1, t_s) = bracket(&self.strikes, strike);
109        let (j0, j1, t_e) = bracket(&self.expiries, expiry);
110
111        // Bilinear interpolation
112        let v00 = self.grid[j0][i0];
113        let v10 = self.grid[j0][i1];
114        let v01 = self.grid[j1][i0];
115        let v11 = self.grid[j1][i1];
116
117        if v00.is_nan() || v10.is_nan() || v01.is_nan() || v11.is_nan() {
118            return None;
119        }
120
121        let v = (1.0 - t_e) * ((1.0 - t_s) * v00 + t_s * v10)
122            + t_e * ((1.0 - t_s) * v01 + t_s * v11);
123        Some(v)
124    }
125
126    /// Implied vol at-the-money (spot = strike) for a given expiry.
127    ///
128    /// Uses interpolation across the strike axis at the nearest available
129    /// expiry (or interpolates between expiries). Returns `None` outside grid.
130    ///
131    /// For a pure ATM query the surface must contain strikes that bracket
132    /// the spot level; this implementation takes ATM as the midpoint of the
133    /// strike range at the given expiry as a proxy when no spot is provided.
134    /// Use `interpolate` with `strike = spot` for a proper ATM vol lookup.
135    pub fn atm_vol(&self, expiry: f64) -> Option<f64> {
136        if self.strikes.is_empty() || self.expiries.is_empty() {
137            return None;
138        }
139        // Use the middle strike as ATM proxy
140        let mid_idx = self.strikes.len() / 2;
141        let atm_strike = self.strikes[mid_idx];
142        self.interpolate(atm_strike, expiry)
143    }
144
145    /// Returns (expiry, atm_vol) pairs for all grid expiries, sorted by expiry.
146    pub fn term_structure(&self) -> Vec<(f64, f64)> {
147        self.expiries
148            .iter()
149            .enumerate()
150            .filter_map(|(j, &exp)| {
151                let mid = self.strikes.len() / 2;
152                let vol = self.grid[j][mid];
153                if vol.is_nan() { None } else { Some((exp, vol)) }
154            })
155            .collect()
156    }
157
158    /// Returns a `VolSmile` at the given expiry (interpolated between grid expiries).
159    ///
160    /// Returns `None` if the expiry is outside the grid.
161    pub fn smile(&self, expiry: f64) -> Option<VolSmile> {
162        if expiry < *self.expiries.first()? || expiry > *self.expiries.last()? {
163            return None;
164        }
165        let pts: Vec<(f64, f64)> = self
166            .strikes
167            .iter()
168            .filter_map(|&k| self.interpolate(k, expiry).map(|v| (k, v)))
169            .collect();
170        if pts.is_empty() {
171            return None;
172        }
173        Some(VolSmile { expiry, points: pts })
174    }
175
176    // Fill NaN cells with the nearest non-NaN value via a simple forward/backward pass.
177    fn fill_gaps(grid: &mut Vec<Vec<f64>>, n_exp: usize, n_str: usize) {
178        // Forward pass over strikes for each expiry row
179        for row in grid.iter_mut().take(n_exp) {
180            let mut last = f64::NAN;
181            for j in 0..n_str {
182                if !row[j].is_nan() {
183                    last = row[j];
184                } else if !last.is_nan() {
185                    row[j] = last;
186                }
187            }
188            // Backward pass
189            let mut last = f64::NAN;
190            for j in (0..n_str).rev() {
191                if !row[j].is_nan() {
192                    last = row[j];
193                } else if !last.is_nan() {
194                    row[j] = last;
195                }
196            }
197        }
198        // Forward pass over expiries for each strike column
199        for i in 0..n_str {
200            let mut last = f64::NAN;
201            for j in 0..n_exp {
202                if !grid[j][i].is_nan() {
203                    last = grid[j][i];
204                } else if !last.is_nan() {
205                    grid[j][i] = last;
206                }
207            }
208            let mut last = f64::NAN;
209            for j in (0..n_exp).rev() {
210                if !grid[j][i].is_nan() {
211                    last = grid[j][i];
212                } else if !last.is_nan() {
213                    grid[j][i] = last;
214                }
215            }
216        }
217    }
218}
219
220/// Returns (lower_idx, upper_idx, fraction) for bilinear interpolation.
221fn bracket(sorted: &[f64], x: f64) -> (usize, usize, f64) {
222    let n = sorted.len();
223    if n == 1 {
224        return (0, 0, 0.0);
225    }
226    // Binary search for insertion point
227    let pos = sorted.partition_point(|&v| v <= x);
228    if pos == 0 {
229        return (0, 0, 0.0);
230    }
231    if pos >= n {
232        return (n - 1, n - 1, 0.0);
233    }
234    let lo = pos - 1;
235    let hi = pos;
236    let span = sorted[hi] - sorted[lo];
237    let t = if span.abs() < 1e-15 { 0.0 } else { (x - sorted[lo]) / span };
238    (lo, hi, t)
239}
240
241// ─── tests ────────────────────────────────────────────────────────────────────
242
243#[cfg(test)]
244mod tests {
245    use super::*;
246
247    fn flat_surface(vol: f64) -> VolSurface {
248        let strikes = [80.0, 90.0, 100.0, 110.0, 120.0];
249        let expiries = [0.25, 0.5, 1.0, 2.0];
250        let points: Vec<VolPoint> = strikes
251            .iter()
252            .flat_map(|&k| {
253                expiries.iter().map(move |&e| VolPoint {
254                    strike: k,
255                    expiry: e,
256                    implied_vol: vol,
257                })
258            })
259            .collect();
260        VolSurface::from_points(points)
261    }
262
263    fn skewed_surface() -> VolSurface {
264        // vol = 0.20 + 0.05 * (100 - K)/100 + 0.10 * T
265        let strikes = [80.0, 90.0, 100.0, 110.0, 120.0];
266        let expiries = [0.25, 0.5, 1.0, 2.0];
267        let points: Vec<VolPoint> = strikes
268            .iter()
269            .flat_map(|&k| {
270                expiries.iter().map(move |&e| VolPoint {
271                    strike: k,
272                    expiry: e,
273                    implied_vol: 0.20 + 0.05 * (100.0 - k) / 100.0 + 0.10 * e,
274                })
275            })
276            .collect();
277        VolSurface::from_points(points)
278    }
279
280    #[test]
281    fn flat_surface_interpolate_on_grid() {
282        let surf = flat_surface(0.20);
283        let v = surf.interpolate(100.0, 1.0).unwrap();
284        assert!((v - 0.20).abs() < 1e-10, "flat surface on-grid: {v}");
285    }
286
287    #[test]
288    fn flat_surface_interpolate_between_grid() {
289        let surf = flat_surface(0.20);
290        // Midpoint between strikes and expiries should still return 0.20
291        let v = surf.interpolate(95.0, 0.75).unwrap();
292        assert!((v - 0.20).abs() < 1e-10, "flat surface off-grid: {v}");
293    }
294
295    #[test]
296    fn interpolate_outside_returns_none_high_strike() {
297        let surf = flat_surface(0.20);
298        assert!(surf.interpolate(200.0, 1.0).is_none());
299    }
300
301    #[test]
302    fn interpolate_outside_returns_none_low_strike() {
303        let surf = flat_surface(0.20);
304        assert!(surf.interpolate(10.0, 1.0).is_none());
305    }
306
307    #[test]
308    fn interpolate_outside_returns_none_high_expiry() {
309        let surf = flat_surface(0.20);
310        assert!(surf.interpolate(100.0, 5.0).is_none());
311    }
312
313    #[test]
314    fn interpolate_outside_returns_none_low_expiry() {
315        let surf = flat_surface(0.20);
316        assert!(surf.interpolate(100.0, 0.01).is_none());
317    }
318
319    #[test]
320    fn skewed_surface_on_grid_point() {
321        let surf = skewed_surface();
322        // strike=100, expiry=1.0 → vol = 0.20 + 0 + 0.10 = 0.30
323        let v = surf.interpolate(100.0, 1.0).unwrap();
324        assert!((v - 0.30).abs() < 1e-10, "skewed on-grid: {v}");
325    }
326
327    #[test]
328    fn skewed_surface_bilinear_accuracy() {
329        let surf = skewed_surface();
330        // Midpoint of (90, 0.5) and (100, 1.0): should be ~average
331        let v00 = 0.20 + 0.05 * (100.0 - 90.0) / 100.0 + 0.10 * 0.5; // 0.255
332        let v10 = 0.20 + 0.05 * (100.0 - 100.0) / 100.0 + 0.10 * 0.5; // 0.250
333        let v01 = 0.20 + 0.05 * (100.0 - 90.0) / 100.0 + 0.10 * 1.0; // 0.305
334        let v11 = 0.20 + 0.05 * (100.0 - 100.0) / 100.0 + 0.10 * 1.0; // 0.300
335        let expected = 0.25 * (v00 + v10 + v01 + v11); // bilinear at t=0.5, s=0.5
336        let v = surf.interpolate(95.0, 0.75).unwrap();
337        assert!((v - expected).abs() < 0.005, "bilinear: {v:.4} vs {expected:.4}");
338    }
339
340    #[test]
341    fn atm_vol_on_grid_expiry() {
342        let surf = flat_surface(0.20);
343        let v = surf.atm_vol(1.0).unwrap();
344        assert!((v - 0.20).abs() < 1e-10);
345    }
346
347    #[test]
348    fn atm_vol_off_grid_expiry() {
349        let surf = flat_surface(0.25);
350        let v = surf.atm_vol(0.75).unwrap();
351        assert!((v - 0.25).abs() < 1e-10);
352    }
353
354    #[test]
355    fn term_structure_sorted() {
356        let surf = flat_surface(0.20);
357        let ts = surf.term_structure();
358        assert!(!ts.is_empty());
359        for w in ts.windows(2) {
360            assert!(w[0].0 < w[1].0, "term structure not sorted");
361        }
362    }
363
364    #[test]
365    fn term_structure_flat() {
366        let surf = flat_surface(0.20);
367        for (_, vol) in surf.term_structure() {
368            assert!((vol - 0.20).abs() < 1e-10);
369        }
370    }
371
372    #[test]
373    fn smile_on_grid_expiry() {
374        let surf = flat_surface(0.20);
375        let smile = surf.smile(1.0).unwrap();
376        assert_eq!(smile.expiry, 1.0);
377        assert!(!smile.points.is_empty());
378        for (_, v) in &smile.points {
379            assert!((v - 0.20).abs() < 1e-10);
380        }
381    }
382
383    #[test]
384    fn smile_outside_returns_none() {
385        let surf = flat_surface(0.20);
386        assert!(surf.smile(10.0).is_none());
387    }
388
389    #[test]
390    fn smile_strikes_sorted() {
391        let surf = skewed_surface();
392        let smile = surf.smile(0.5).unwrap();
393        for w in smile.points.windows(2) {
394            assert!(w[0].0 <= w[1].0, "smile strikes not sorted");
395        }
396    }
397
398    #[test]
399    fn from_points_single_point() {
400        let pts = vec![VolPoint { strike: 100.0, expiry: 1.0, implied_vol: 0.20 }];
401        let surf = VolSurface::from_points(pts);
402        // Should be able to query the exact point
403        let v = surf.interpolate(100.0, 1.0).unwrap();
404        assert!((v - 0.20).abs() < 1e-10);
405    }
406
407    #[test]
408    fn from_points_duplicate_averaged() {
409        let pts = vec![
410            VolPoint { strike: 100.0, expiry: 1.0, implied_vol: 0.20 },
411            VolPoint { strike: 100.0, expiry: 1.0, implied_vol: 0.30 },
412        ];
413        let surf = VolSurface::from_points(pts);
414        let v = surf.interpolate(100.0, 1.0).unwrap();
415        assert!((v - 0.25).abs() < 1e-10, "duplicates should average: {v}");
416    }
417
418    #[test]
419    fn smile_vol_decreases_with_strike_for_skewed() {
420        // In skewed_surface, lower strike → higher vol
421        let surf = skewed_surface();
422        let smile = surf.smile(1.0).unwrap();
423        for w in smile.points.windows(2) {
424            assert!(w[0].1 >= w[1].1, "vol should decrease with strike: {} < {}", w[0].1, w[1].1);
425        }
426    }
427}