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