use crate::math::solvers::solvertraits::{ContFunc, OptimizerSolution, SolutionStatus};
use crate::utils::errors::{QSError, Result};
pub struct Bisection<P> {
_p: std::marker::PhantomData<P>,
lower: f64,
upper: f64,
ftol: f64,
max_iter: i64,
}
impl<P> Bisection<P>
where
P: ContFunc<f64>,
{
#[must_use]
pub const fn new(lower: f64, upper: f64, max_iter: i64) -> Self {
Self {
_p: std::marker::PhantomData,
lower,
upper,
ftol: 1e-12,
max_iter,
}
}
#[must_use]
pub const fn with_ftol(mut self, ftol: f64) -> Self {
self.ftol = ftol;
self
}
#[must_use]
pub const fn with_lower(mut self, lower: f64) -> Self {
self.lower = lower;
self
}
#[must_use]
pub const fn with_upper(mut self, upper: f64) -> Self {
self.upper = upper;
self
}
pub fn solve(&self, f: &P) -> Result<OptimizerSolution<f64>> {
let mut low = self.lower;
let mut high = self.upper;
let mut f_low = f.call(&low)?;
let f_high = f.call(&high)?;
if f_low.signum() == f_high.signum() {
return Err(QSError::SolverErr(
"Bisection requires a sign change over the bracket.".into(),
));
}
for _ in 0..self.max_iter {
let mid = 0.5 * (low + high);
let f_mid = f.call(&mid)?;
if f_mid.abs() < self.ftol {
return Ok(OptimizerSolution {
x: mid,
f: f_mid,
status: SolutionStatus::Converged,
});
}
if f_mid.signum() == f_low.signum() {
low = mid;
f_low = f_mid;
} else {
high = mid;
}
}
let mid = 0.5 * (low + high);
Ok(OptimizerSolution {
x: mid,
f: f.call(&mid)?,
status: SolutionStatus::NotConverged,
})
}
}