Skip to main content

finance_solution/util/
root_find.rs

1//! Scalar root-finding helpers used by implied vol and other solvers.
2//!
3//! | Method | Use in this crate |
4//! |--------|-------------------|
5//! | Newton–Raphson | Primary IV step when vega is reliable |
6//! | **Brent** | Robust bracketed fallback after Newton |
7//!
8//! Brent combines bisection reliability with secant / inverse-quadratic speed on
9//! smooth problems. We previously used plain bisection as the fallback; Brent is
10//! the standard upgrade for the same bracketed contract.
11
12use crate::util::error::{FinanceError, FinanceResult};
13
14/// Find a root of continuous `f` on bracket `[lo, hi]` (Brent 1973).
15///
16/// Requires opposite signs at the endpoints (or a zero on an endpoint).
17pub fn brent_root(
18    lo: f64,
19    hi: f64,
20    mut f: impl FnMut(f64) -> f64,
21    tol: f64,
22    max_iter: usize,
23) -> FinanceResult<f64> {
24    let mut a = lo;
25    let mut b = hi;
26    let mut fa = f(a);
27    let mut fb = f(b);
28    if !fa.is_finite() || !fb.is_finite() {
29        return Err(FinanceError::Unsolvable {
30            message: "Brent: non-finite function value at bracket",
31        });
32    }
33    if fa == 0.0 {
34        return Ok(a);
35    }
36    if fb == 0.0 {
37        return Ok(b);
38    }
39    if fa * fb > 0.0 {
40        return Err(FinanceError::Unsolvable {
41            message: "Brent: bracket does not change sign",
42        });
43    }
44
45    let mut c = a;
46    let mut fc = fa;
47    let mut d = b - a;
48    let mut e = d;
49
50    for _ in 0..max_iter {
51        if fb == 0.0 {
52            return Ok(b);
53        }
54        if fa * fb > 0.0 {
55            a = c;
56            fa = fc;
57            d = b - a;
58            e = d;
59        }
60        if fa.abs() < fb.abs() {
61            // swap a <-> b
62            c = b;
63            b = a;
64            a = c;
65            fc = fb;
66            fb = fa;
67            fa = fc;
68        }
69
70        let tol1 = 2.0 * f64::EPSILON * b.abs() + 0.5 * tol;
71        let xm = 0.5 * (a - b);
72        if xm.abs() <= tol1 {
73            return Ok(b);
74        }
75
76        if e.abs() >= tol1 && fa.abs() > fb.abs() {
77            let s = fb / fa;
78            let (mut p, mut q) = if (a - c).abs() <= f64::EPSILON {
79                // linear interpolation (secant)
80                (2.0 * xm * s, 1.0 - s)
81            } else {
82                // inverse quadratic
83                let q0 = fa / fc;
84                let r = fb / fc;
85                let p = s * (2.0 * xm * q0 * (q0 - r) - (b - a) * (r - 1.0));
86                let q = (q0 - 1.0) * (r - 1.0) * (s - 1.0);
87                (p, q)
88            };
89            if p > 0.0 {
90                q = -q;
91            } else {
92                p = -p;
93            }
94            let min1 = 3.0 * xm * q.abs() - (tol1 * q).abs();
95            let min2 = (e * q).abs();
96            if 2.0 * p < min1.min(min2) {
97                e = d;
98                d = p / q;
99            } else {
100                d = xm;
101                e = d;
102            }
103        } else {
104            d = xm;
105            e = d;
106        }
107
108        c = b;
109        fc = fb;
110        if d.abs() > tol1 {
111            b += d;
112        } else {
113            b += xm.signum() * tol1;
114            if xm == 0.0 {
115                b += if a > b { tol1 } else { -tol1 };
116            }
117        }
118        fb = f(b);
119        if !fb.is_finite() {
120            return Err(FinanceError::Unsolvable {
121                message: "Brent: non-finite function value",
122            });
123        }
124    }
125
126    Err(FinanceError::Unsolvable {
127        message: "Brent: max iterations exceeded",
128    })
129}
130
131#[cfg(test)]
132mod tests {
133    use super::*;
134
135    #[test]
136    fn sqrt_two() {
137        let r = brent_root(1.0, 2.0, |x| x * x - 2.0, 1e-14, 100).unwrap();
138        assert!((r - 2.0_f64.sqrt()).abs() < 1e-10);
139    }
140
141    #[test]
142    fn no_sign_change_err() {
143        assert!(brent_root(1.0, 2.0, |x| x + 1.0, 1e-8, 50).is_err());
144    }
145
146    #[test]
147    fn endpoint_root() {
148        let r = brent_root(0.0, 2.0, |x| x - 2.0, 1e-12, 50).unwrap();
149        assert!((r - 2.0).abs() < 1e-12);
150    }
151
152    #[test]
153    fn cubic_root() {
154        // x^3 - x - 2 = 0, root ≈ 1.521
155        let r = brent_root(1.0, 2.0, |x| x * x * x - x - 2.0, 1e-12, 100).unwrap();
156        assert!((r * r * r - r - 2.0).abs() < 1e-10);
157    }
158
159    #[test]
160    fn sine_root() {
161        let r = brent_root(3.0, 3.2, |x| x.sin(), 1e-14, 100).unwrap();
162        assert!((r - std::f64::consts::PI).abs() < 1e-10);
163    }
164}