finance_solution/util/
root_find.rs1use crate::util::error::{FinanceError, FinanceResult};
13
14pub 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 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 (2.0 * xm * s, 1.0 - s)
81 } else {
82 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 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}