use crate::root_finding::{RootReport, RootTermination, same_sign};
use crate::scalar::{Numeric, ScalarFn};
use crate::utils::error_codes::CalcError;
pub struct Bisection<T = f64> {
xtol: T,
ftol: T,
max_iterations: usize,
}
impl<T: Numeric> Default for Bisection<T> {
fn default() -> Self {
let tol = T::EPSILON * T::from_f64(4.0);
Bisection {
xtol: tol,
ftol: tol,
max_iterations: 100,
}
}
}
impl<T: Numeric> Bisection<T> {
#[must_use]
pub fn with_xtol(mut self, xtol: T) -> Self {
self.xtol = xtol;
self
}
#[must_use]
pub fn with_ftol(mut self, ftol: T) -> Self {
self.ftol = ftol;
self
}
#[must_use]
pub fn with_max_iterations(mut self, max_iterations: usize) -> Self {
self.max_iterations = max_iterations;
self
}
pub fn solve<F: ScalarFn>(&self, f: &F, a: T, b: T) -> Result<RootReport<T>, CalcError> {
let fa = f.eval(a);
let fb = f.eval(b);
if !fa.is_finite() || !fb.is_finite() {
return Err(CalcError::NonFiniteValue);
}
if fa == T::ZERO {
return Ok(RootReport {
root: a,
residual: fa,
iterations: 0,
termination: RootTermination::ResidualTolerance,
});
}
if fb == T::ZERO {
return Ok(RootReport {
root: b,
residual: fb,
iterations: 0,
termination: RootTermination::ResidualTolerance,
});
}
if same_sign(fa, fb) {
return Err(CalcError::InvalidBracket);
}
let (mut lo, mut flo, mut hi) = if a <= b { (a, fa, b) } else { (b, fb, a) };
for iter in 1..=self.max_iterations {
let mid = lo + (hi - lo) * T::HALF;
let fmid = f.eval(mid);
if !fmid.is_finite() {
return Err(CalcError::NonFiniteValue);
}
if fmid.abs() <= self.ftol {
return Ok(RootReport {
root: mid,
residual: fmid,
iterations: iter,
termination: RootTermination::ResidualTolerance,
});
}
if (hi - lo) <= self.xtol * (T::ONE + mid.abs()) {
return Ok(RootReport {
root: mid,
residual: fmid,
iterations: iter,
termination: RootTermination::BracketWidth,
});
}
if same_sign(fmid, flo) {
lo = mid;
flo = fmid;
} else {
hi = mid;
}
}
Err(CalcError::DidNotConverge)
}
}