use super::MathError;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct BrentConfig {
pub xtol: f64,
pub ftol: f64,
pub max_iter: u32,
}
impl Default for BrentConfig {
fn default() -> Self {
Self {
xtol: 1e-12,
ftol: 1e-14,
max_iter: 100,
}
}
}
#[allow(clippy::many_single_char_names)]
pub fn brent_root<F>(mut f: F, a: f64, b: f64, config: BrentConfig) -> Result<f64, MathError>
where
F: FnMut(f64) -> f64,
{
let mut a = a;
let mut b = b;
let mut fa = f(a);
let mut fb = f(b);
if fa == 0.0 {
return Ok(a);
}
if fb == 0.0 {
return Ok(b);
}
if fa * fb > 0.0 {
return Err(MathError::BracketNotStraddling);
}
if fa.abs() < fb.abs() {
core::mem::swap(&mut a, &mut b);
core::mem::swap(&mut fa, &mut fb);
}
let mut c = a;
let mut fc = fa;
let mut d = a;
let mut mflag = true;
for _ in 0..config.max_iter {
if (b - a).abs() <= config.xtol || fb.abs() <= config.ftol || fb == 0.0 {
return Ok(b);
}
let mut s = if (fa - fc).abs() > f64::EPSILON && (fb - fc).abs() > f64::EPSILON {
a * fb * fc / ((fa - fb) * (fa - fc))
+ b * fa * fc / ((fb - fa) * (fb - fc))
+ c * fa * fb / ((fc - fa) * (fc - fb))
} else {
b - fb * (b - a) / (fb - fa)
};
let lo = (3.0 * a + b) / 4.0;
let bound_lo = lo.min(b);
let bound_hi = lo.max(b);
let use_bisection = !(bound_lo..=bound_hi).contains(&s)
|| (mflag && (s - b).abs() >= (b - c).abs() / 2.0)
|| (!mflag && (s - b).abs() >= (c - d).abs() / 2.0)
|| (mflag && (b - c).abs() < config.xtol)
|| (!mflag && (c - d).abs() < config.xtol);
if use_bisection {
s = f64::midpoint(a, b);
mflag = true;
} else {
mflag = false;
}
let fs = f(s);
d = c;
c = b;
fc = fb;
if fa * fs < 0.0 {
b = s;
fb = fs;
} else {
a = s;
fa = fs;
}
if fa.abs() < fb.abs() {
core::mem::swap(&mut a, &mut b);
core::mem::swap(&mut fa, &mut fb);
}
}
if (b - a).abs() <= config.xtol || fb.abs() <= config.ftol {
return Ok(b);
}
Err(MathError::NoConvergence)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn brent_config_default_values() {
let cfg = BrentConfig::default();
assert!((cfg.xtol - 1e-12).abs() < 1e-18);
assert!((cfg.ftol - 1e-14).abs() < 1e-20);
assert_eq!(cfg.max_iter, 100);
}
#[test]
fn brent_config_copy_eq() {
let cfg = BrentConfig::default();
let copy = cfg;
assert_eq!(cfg, copy);
}
#[test]
fn brent_finds_sqrt_two() {
let root = brent_root(|x| x * x - 2.0, 0.0, 2.0, BrentConfig::default()).unwrap();
assert!((root - 2.0_f64.sqrt()).abs() < 1e-12);
}
#[test]
fn brent_finds_cubic_root() {
let root = brent_root(|x| x * x * x - x - 2.0, 1.0, 2.0, BrentConfig::default()).unwrap();
assert!((root - 1.521_379_706_804_567_6).abs() < 1e-10);
}
#[test]
fn brent_finds_dottie_number() {
let root = brent_root(|x| x.cos() - x, 0.0, 1.0, BrentConfig::default()).unwrap();
assert!((root - 0.739_085_133_215_160_6).abs() < 1e-10);
}
#[test]
fn brent_endpoint_root_accepted() {
let root = brent_root(|x| x - 3.0, 3.0, 5.0, BrentConfig::default()).unwrap();
assert!((root - 3.0).abs() < 1e-15);
}
#[test]
fn brent_endpoint_root_b() {
let root = brent_root(|x| x - 5.0, 3.0, 5.0, BrentConfig::default()).unwrap();
assert!((root - 5.0).abs() < 1e-15);
}
#[test]
fn brent_rejects_no_bracket() {
let err = brent_root(|x| x * x + 1.0, -1.0, 1.0, BrentConfig::default()).unwrap_err();
assert!(matches!(err, MathError::BracketNotStraddling));
}
#[test]
fn brent_finds_transcendental_root() {
let root = brent_root(|x| (-x).exp() - x, 0.0, 1.0, BrentConfig::default()).unwrap();
assert!((root - 0.567_143_290_409_783_8).abs() < 1e-10);
}
#[test]
fn brent_respects_max_iter() {
let cfg = BrentConfig {
xtol: 1e-30,
ftol: 1e-30,
max_iter: 1,
};
let res = brent_root(|x| x * x * x - 0.123, 0.0, 1.0, cfg);
assert!(matches!(res, Err(MathError::NoConvergence)));
}
#[test]
fn brent_uses_secant_when_three_points_coincide() {
let root = brent_root(|x| x - 0.5, 0.0, 1.0, BrentConfig::default()).unwrap();
assert!((root - 0.5).abs() < 1e-12);
}
#[test]
fn brent_swapped_endpoints() {
let root = brent_root(|x| x - 2.0, 5.0, 0.0, BrentConfig::default()).unwrap();
assert!((root - 2.0).abs() < 1e-12);
}
#[test]
fn brent_tight_ftol() {
let cfg = BrentConfig {
xtol: 1e-15,
ftol: 1e-15,
max_iter: 200,
};
let f = |x: f64| (x - 3.7).powi(3);
let root = brent_root(f, 0.0, 10.0, cfg).unwrap();
assert!((root - 3.7).abs() < 1e-5);
}
#[test]
fn brent_strictly_monotonic_function() {
let root = brent_root(|x| x.exp() - 5.0, 0.0, 5.0, BrentConfig::default()).unwrap();
assert!((root - 5.0_f64.ln()).abs() < 1e-12);
}
}