use crate::symbolic::symbolic_engine::Expr;
use std::fmt;
use std::collections::HashMap;
#[derive(Debug, Clone)]
pub enum RootFindingMethod {
Bisection,
Secant,
NewtonRaphson,
Brent,
}
#[derive(Debug, Clone)]
pub enum RootFindingError {
MaxIterationsReached,
InvalidInterval,
FunctionNotContinuous,
DerivativeZero,
ToleranceNotMet,
InvalidInput(String),
}
impl fmt::Display for RootFindingError {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
match self {
RootFindingError::MaxIterationsReached => write!(f, "Maximum iterations reached"),
RootFindingError::InvalidInterval => write!(f, "Invalid interval for bisection method"),
RootFindingError::FunctionNotContinuous => {
write!(f, "Function is not continuous in the given interval")
}
RootFindingError::DerivativeZero => write!(f, "Derivative is zero"),
RootFindingError::ToleranceNotMet => write!(f, "Required tolerance not met"),
RootFindingError::InvalidInput(msg) => write!(f, "Invalid input: {}", msg),
}
}
}
impl std::error::Error for RootFindingError {}
pub trait NonlinearFunction {
fn evaluate(&self, x: f64) -> f64;
fn derivative(&self, _x: f64) -> Option<f64> {
None
}
fn name(&self) -> &str {
"unnamed_function"
}
}
pub struct ClosureFunction<F>
where
F: Fn(f64) -> f64,
{
func: F,
name: String,
}
impl<F> ClosureFunction<F>
where
F: Fn(f64) -> f64,
{
pub fn new(func: F, name: String) -> Self {
Self { func, name }
}
}
impl<F> NonlinearFunction for ClosureFunction<F>
where
F: Fn(f64) -> f64,
{
fn evaluate(&self, x: f64) -> f64 {
(self.func)(x)
}
fn name(&self) -> &str {
&self.name
}
}
pub struct FunctionWithDerivative<F, D>
where
F: Fn(f64) -> f64,
D: Fn(f64) -> f64,
{
func: F,
derivative_func: D,
name: String,
}
impl<F, D> FunctionWithDerivative<F, D>
where
F: Fn(f64) -> f64,
D: Fn(f64) -> f64,
{
pub fn new(func: F, derivative_func: D, name: String) -> Self {
Self {
func,
derivative_func,
name,
}
}
}
impl<F, D> NonlinearFunction for FunctionWithDerivative<F, D>
where
F: Fn(f64) -> f64,
D: Fn(f64) -> f64,
{
fn evaluate(&self, x: f64) -> f64 {
(self.func)(x)
}
fn derivative(&self, x: f64) -> Option<f64> {
Some((self.derivative_func)(x))
}
fn name(&self) -> &str {
&self.name
}
}
pub struct SymbolicFunction {
original_expr: Expr, expr: Expr,
derivative_expr: Option<Expr>,
variable: String,
name: String,
func: Box<dyn Fn(f64) -> f64>,
derivative_func: Option<Box<dyn Fn(f64) -> f64>>,
parameters: HashMap<String, f64>,
}
impl SymbolicFunction {
pub fn from_string(
expr_str: &str,
variable: &str,
name: Option<String>,
) -> Result<Self, RootFindingError> {
let expr = Expr::parse_expression(expr_str);
let func_name = name.unwrap_or_else(|| format!("symbolic_function({})", expr_str));
Self::from_expr(expr, variable, Some(func_name))
}
pub fn from_expr(
expr: Expr,
variable: &str,
name: Option<String>,
) -> Result<Self, RootFindingError> {
let func_name = name.unwrap_or_else(|| "symbolic_function".to_string());
let derivative_expr = expr.diff(variable);
let func = expr.lambdify1D();
let derivative_func = derivative_expr.lambdify1D();
Ok(Self {
original_expr: expr.clone(), expr: expr.clone(),
derivative_expr: Some(derivative_expr),
variable: variable.to_string(),
name: func_name,
func,
derivative_func: Some(derivative_func),
parameters: HashMap::new(),
})
}
pub fn from_string_with_params(
expr_str: &str,
variable: &str,
parameters: HashMap<String, f64>,
name: Option<String>,
) -> Result<Self, RootFindingError> {
let expr = Expr::parse_expression(expr_str);
let func_name = name.unwrap_or_else(|| format!("symbolic_function({})", expr_str));
Self::from_expr_with_params(expr, variable, parameters, Some(func_name))
}
pub fn from_expr_with_params(
expr: Expr,
variable: &str,
parameters: HashMap<String, f64>,
name: Option<String>,
) -> Result<Self, RootFindingError> {
let func_name = name.unwrap_or_else(|| "symbolic_function_with_params".to_string());
let expr_with_params = expr.set_variable_from_map(¶meters);
let derivative_expr = expr_with_params.diff(variable);
let func = expr_with_params.lambdify1D();
let derivative_func = derivative_expr.lambdify1D();
Ok(Self {
original_expr: expr.clone(), expr: expr_with_params,
derivative_expr: Some(derivative_expr),
variable: variable.to_string(),
name: func_name,
func,
derivative_func: Some(derivative_func),
parameters,
})
}
pub fn set_parameters(
&mut self,
parameters: HashMap<String, f64>,
) -> Result<(), RootFindingError> {
self.parameters = parameters;
let expr_with_params = self.original_expr.set_variable_from_map(&self.parameters);
let derivative_expr = expr_with_params.diff(&self.variable);
self.expr = expr_with_params.clone();
self.derivative_expr = Some(derivative_expr.clone());
self.func = expr_with_params.lambdify1D();
self.derivative_func = Some(derivative_expr.lambdify1D());
Ok(())
}
pub fn expression_string(&self) -> String {
self.expr.sym_to_str(&self.variable)
}
pub fn derivative_string(&self) -> Option<String> {
self.derivative_expr
.as_ref()
.map(|expr| expr.sym_to_str(&self.variable))
}
}
impl NonlinearFunction for SymbolicFunction {
fn evaluate(&self, x: f64) -> f64 {
(self.func)(x)
}
fn derivative(&self, x: f64) -> Option<f64> {
self.derivative_func.as_ref().map(|f| f(x))
}
fn name(&self) -> &str {
&self.name
}
}
#[derive(Debug, Clone)]
pub struct RootFindingResult {
pub root: f64,
pub function_value: f64,
pub iterations: usize,
pub converged: bool,
pub method: String,
}
#[derive(Debug, Clone)]
pub struct RootFindingConfig {
pub tolerance: f64,
pub max_iterations: usize,
pub verbose: bool,
}
impl Default for RootFindingConfig {
fn default() -> Self {
Self {
tolerance: 1e-10,
max_iterations: 100,
verbose: false,
}
}
}
pub struct ScalarRootFinder {
config: RootFindingConfig,
}
impl ScalarRootFinder {
pub fn new() -> Self {
Self {
config: RootFindingConfig::default(),
}
}
pub fn with_config(config: RootFindingConfig) -> Self {
Self { config }
}
pub fn set_tolerance(&mut self, tolerance: f64) {
self.config.tolerance = tolerance;
}
pub fn set_max_iterations(&mut self, max_iterations: usize) {
self.config.max_iterations = max_iterations;
}
pub fn set_verbose(&mut self, verbose: bool) {
self.config.verbose = verbose;
}
pub fn solve_symbolic_str(
&self,
expr_str: &str,
variable: &str,
method: RootFindingMethod,
initial_guess: f64,
search_range: Option<(f64, f64)>,
parameters: Option<HashMap<String, f64>>,
) -> Result<RootFindingResult, RootFindingError> {
let symbolic_func = if let Some(params) = parameters {
SymbolicFunction::from_string_with_params(expr_str, variable, params, None)?
} else {
SymbolicFunction::from_string(expr_str, variable, None)?
};
self.solve_with_method(&symbolic_func, method, initial_guess, search_range)
}
pub fn solve_symbolic_expr(
&self,
expr: Expr,
variable: &str,
method: RootFindingMethod,
initial_guess: f64,
search_range: Option<(f64, f64)>,
parameters: Option<HashMap<String, f64>>,
) -> Result<RootFindingResult, RootFindingError> {
let symbolic_func = if let Some(params) = parameters {
SymbolicFunction::from_expr_with_params(expr, variable, params, None)?
} else {
SymbolicFunction::from_expr(expr, variable, None)?
};
self.solve_with_method(&symbolic_func, method, initial_guess, search_range)
}
pub fn solve_with_method<F>(
&self,
function: &F,
method: RootFindingMethod,
initial_guess: f64,
search_range: Option<(f64, f64)>,
) -> Result<RootFindingResult, RootFindingError>
where
F: NonlinearFunction,
{
match method {
RootFindingMethod::Bisection => {
if let Some((a, b)) = search_range {
self.bisection(function, a, b)
} else {
Err(RootFindingError::InvalidInput(
"Bisection method requires search range".to_string(),
))
}
}
RootFindingMethod::Secant => {
let x1 = initial_guess + 0.01 * initial_guess.abs().max(1.0);
self.secant(function, initial_guess, x1)
}
RootFindingMethod::NewtonRaphson => self.newton_raphson(function, initial_guess),
RootFindingMethod::Brent => self.brent(function, search_range),
}
}
pub fn bisection<F>(
&self,
function: &F,
mut a: f64,
mut b: f64,
) -> Result<RootFindingResult, RootFindingError>
where
F: NonlinearFunction,
{
if a > b {
std::mem::swap(&mut a, &mut b);
}
let fa = function.evaluate(a);
let fb = function.evaluate(b);
if fa * fb > 0.0 {
return Err(RootFindingError::InvalidInterval);
}
if fa.abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: a,
function_value: fa,
iterations: 0,
converged: true,
method: "bisection".to_string(),
});
}
if fb.abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: b,
function_value: fb,
iterations: 0,
converged: true,
method: "bisection".to_string(),
});
}
let mut iterations = 0;
let mut c: f64;
let mut fc: f64;
if self.config.verbose {
println!("Bisection method for function: {}", function.name());
println!("Initial interval: [{}, {}]", a, b);
println!("Tolerance: {}", self.config.tolerance);
}
while iterations < self.config.max_iterations {
c = (a + b) / 2.0;
fc = function.evaluate(c);
if self.config.verbose {
println!(
"Iteration {}: x = {:.10}, f(x) = {:.2e}, interval = [{:.6}, {:.6}]",
iterations + 1,
c,
fc,
a,
b
);
}
if fc.abs() < self.config.tolerance || (b - a) / 2.0 < self.config.tolerance {
return Ok(RootFindingResult {
root: c,
function_value: fc,
iterations: iterations + 1,
converged: true,
method: "bisection".to_string(),
});
}
if function.evaluate(a) * fc < 0.0 {
b = c;
} else {
a = c;
}
iterations += 1;
}
Err(RootFindingError::MaxIterationsReached)
}
pub fn secant<F>(
&self,
function: &F,
mut x0: f64,
mut x1: f64,
) -> Result<RootFindingResult, RootFindingError>
where
F: NonlinearFunction,
{
let mut f0 = function.evaluate(x0);
let mut f1 = function.evaluate(x1);
if self.config.verbose {
println!("Secant method for function: {}", function.name());
println!("Initial guesses: x0 = {}, x1 = {}", x0, x1);
println!("Tolerance: {}", self.config.tolerance);
}
if f0.abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: x0,
function_value: f0,
iterations: 0,
converged: true,
method: "secant".to_string(),
});
}
if f1.abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: x1,
function_value: f1,
iterations: 0,
converged: true,
method: "secant".to_string(),
});
}
let mut iterations = 0;
while iterations < self.config.max_iterations {
if (f1 - f0).abs() < 1e-15 {
return Err(RootFindingError::DerivativeZero);
}
let x2 = x1 - f1 * (x1 - x0) / (f1 - f0);
let f2 = function.evaluate(x2);
if self.config.verbose {
println!(
"Iteration {}: x = {:.10}, f(x) = {:.2e}",
iterations + 1,
x2,
f2
);
}
if f2.abs() < self.config.tolerance || (x2 - x1).abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: x2,
function_value: f2,
iterations: iterations + 1,
converged: true,
method: "secant".to_string(),
});
}
x0 = x1;
f0 = f1;
x1 = x2;
f1 = f2;
iterations += 1;
}
Err(RootFindingError::MaxIterationsReached)
}
pub fn newton_raphson<F>(
&self,
function: &F,
mut x: f64,
) -> Result<RootFindingResult, RootFindingError>
where
F: NonlinearFunction,
{
if self.config.verbose {
println!("Newton-Raphson method for function: {}", function.name());
println!("Initial guess: x0 = {}", x);
println!("Tolerance: {}", self.config.tolerance);
}
let mut iterations = 0;
while iterations < self.config.max_iterations {
let fx = function.evaluate(x);
if self.config.verbose {
println!(
"Iteration {}: x = {:.10}, f(x) = {:.2e}",
iterations + 1,
x,
fx
);
}
if fx.abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: x,
function_value: fx,
iterations: iterations + 1,
converged: true,
method: "newton_raphson".to_string(),
});
}
let fpx = match function.derivative(x) {
Some(deriv) => deriv,
None => {
let h = 1e-8;
(function.evaluate(x + h) - function.evaluate(x - h)) / (2.0 * h)
}
};
if fpx.abs() < 1e-15 {
return Err(RootFindingError::DerivativeZero);
}
let x_new = x - fx / fpx;
if (x_new - x).abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: x_new,
function_value: function.evaluate(x_new),
iterations: iterations + 1,
converged: true,
method: "newton_raphson".to_string(),
});
}
x = x_new;
iterations += 1;
}
Err(RootFindingError::MaxIterationsReached)
}
pub fn brent<F>(
&self,
function: &F,
search_range: Option<(f64, f64)>,
) -> Result<RootFindingResult, RootFindingError>
where
F: NonlinearFunction,
{
let (mut a, mut b) = match search_range {
Some((x, y)) => (x, y),
None => {
return Err(RootFindingError::InvalidInput(
"Brent method requires search range".to_string(),
));
}
};
if a > b {
std::mem::swap(&mut a, &mut b);
}
let mut fa = function.evaluate(a);
let mut fb = function.evaluate(b);
if fa * fb > 0.0 {
return Err(RootFindingError::InvalidInterval);
}
if fa.abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: a,
function_value: fa,
iterations: 0,
converged: true,
method: "brent".to_string(),
});
}
if fb.abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: b,
function_value: fb,
iterations: 0,
converged: true,
method: "brent".to_string(),
});
}
if fa.abs() < fb.abs() {
std::mem::swap(&mut a, &mut b);
std::mem::swap(&mut fa, &mut fb);
}
let mut c = a;
let mut fc = fa;
let mut d = 0.0;
let _e = 0.0;
let mut mflag = true;
let mut iterations = 0;
if self.config.verbose {
println!("Brent method for function: {}", function.name());
println!("Initial interval: [{}, {}]", a, b);
println!("Tolerance: {}", self.config.tolerance);
}
while iterations < self.config.max_iterations {
if fb.abs() < self.config.tolerance || (b - a).abs() < self.config.tolerance {
return Ok(RootFindingResult {
root: b,
function_value: fb,
iterations: iterations + 1,
converged: true,
method: "brent".to_string(),
});
}
let mut s: f64;
if fa != fc && fb != fc {
s = a * fb * fc / ((fa - fb) * (fa - fc))
+ b * fa * fc / ((fb - fa) * (fb - fc))
+ c * fa * fb / ((fc - fa) * (fc - fb));
} else {
s = b - fb * (b - a) / (fb - fa);
}
let condition1 = s < (3.0 * a + b) / 4.0 || s > b;
let condition2 = mflag && (s - b).abs() >= (b - c).abs() / 2.0;
let condition3 = !mflag && (s - b).abs() >= (c - d).abs() / 2.0;
let condition4 = mflag && (b - c).abs() < self.config.tolerance;
let condition5 = !mflag && (c - d).abs() < self.config.tolerance;
if condition1 || condition2 || condition3 || condition4 || condition5 {
s = (a + b) / 2.0;
mflag = true;
} else {
mflag = false;
}
let fs = function.evaluate(s);
if self.config.verbose {
println!(
"Iteration {}: s = {:.10}, f(s) = {:.2e}, interval = [{:.6}, {:.6}]",
iterations + 1,
s,
fs,
a,
b
);
}
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() {
std::mem::swap(&mut a, &mut b);
std::mem::swap(&mut fa, &mut fb);
}
iterations += 1;
}
Err(RootFindingError::MaxIterationsReached)
}
pub fn solve<F>(
&self,
function: &F,
initial_guess: f64,
search_range: Option<(f64, f64)>,
) -> Result<RootFindingResult, RootFindingError>
where
F: NonlinearFunction,
{
if function.derivative(initial_guess).is_some() {
if let Ok(result) = self.newton_raphson(function, initial_guess) {
return Ok(result);
}
}
let x1 = initial_guess + 0.01 * initial_guess.abs().max(1.0);
if let Ok(result) = self.secant(function, initial_guess, x1) {
return Ok(result);
}
if let Some((a, b)) = search_range {
return self.bisection(function, a, b);
}
Err(RootFindingError::ToleranceNotMet)
}
}
impl Default for ScalarRootFinder {
fn default() -> Self {
Self::new()
}
}
pub fn bisection<F>(function: F, a: f64, b: f64, tolerance: f64) -> Result<f64, RootFindingError>
where
F: Fn(f64) -> f64,
{
let func = ClosureFunction::new(function, "bisection_function".to_string());
let mut solver = ScalarRootFinder::new();
solver.set_tolerance(tolerance);
let result = solver.bisection(&func, a, b)?;
Ok(result.root)
}
pub fn secant<F>(function: F, x0: f64, x1: f64, tolerance: f64) -> Result<f64, RootFindingError>
where
F: Fn(f64) -> f64,
{
let func = ClosureFunction::new(function, "secant_function".to_string());
let mut solver = ScalarRootFinder::new();
solver.set_tolerance(tolerance);
let result = solver.secant(&func, x0, x1)?;
Ok(result.root)
}
pub fn approx_equal(a: f64, b: f64, tolerance: f64) -> bool {
(a - b).abs() < tolerance
}
#[cfg(test)]
mod tests {
use super::*;
use std::f64::consts::PI;
#[test]
fn test_closure_function() {
let func = ClosureFunction::new(|x| x * x - 4.0, "x^2 - 4".to_string());
assert_eq!(func.evaluate(2.0), 0.0);
assert_eq!(func.evaluate(-2.0), 0.0);
assert_eq!(func.evaluate(0.0), -4.0);
assert_eq!(func.name(), "x^2 - 4");
}
#[test]
fn test_function_with_derivative() {
let func = FunctionWithDerivative::new(
|x| x * x - 4.0,
|x| 2.0 * x,
"x^2 - 4 with derivative".to_string(),
);
assert_eq!(func.evaluate(2.0), 0.0);
assert_eq!(func.derivative(2.0), Some(4.0));
assert_eq!(func.derivative(3.0), Some(6.0));
}
#[test]
fn test_bisection_simple_quadratic() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x - 4.0, "x^2 - 4".to_string());
let result = solver.bisection(&func, 0.0, 3.0).unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert!(result.converged);
assert_eq!(result.method, "bisection");
let result = solver.bisection(&func, -3.0, 0.0).unwrap();
assert!(approx_equal(result.root, -2.0, 1e-10));
assert!(result.converged);
}
#[test]
fn test_bisection_cubic() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x * x - x - 1.0, "x^3 - x - 1".to_string());
let result = solver.bisection(&func, 1.0, 2.0).unwrap();
let expected_root = 1.324717957244746;
assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
}
#[test]
fn test_bisection_trigonometric() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x.sin(), "sin(x)".to_string());
let result = solver.bisection(&func, 3.0, 4.0).unwrap();
assert!(approx_equal(result.root, PI, 1e-10));
assert!(result.converged);
}
#[test]
fn test_bisection_invalid_interval() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x + 1.0, "x^2 + 1".to_string());
let result = solver.bisection(&func, -1.0, 1.0);
assert!(matches!(result, Err(RootFindingError::InvalidInterval)));
}
#[test]
fn test_bisection_root_at_endpoint() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x - 2.0, "x - 2".to_string());
let result = solver.bisection(&func, 1.0, 2.0).unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert_eq!(result.iterations, 0);
}
#[test]
fn test_secant_simple_quadratic() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x - 4.0, "x^2 - 4".to_string());
let result = solver.secant(&func, 1.0, 3.0).unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert!(result.converged);
assert_eq!(result.method, "secant");
let result = solver.secant(&func, -1.0, -3.0).unwrap();
assert!(approx_equal(result.root, -2.0, 1e-10));
assert!(result.converged);
}
#[test]
fn test_secant_cubic() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x * x - x - 1.0, "x^3 - x - 1".to_string());
let result = solver.secant(&func, 1.0, 2.0).unwrap();
let expected_root = 1.324717957244746;
assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
let result = solver.bisection(&func, 1.0, 2.0).unwrap();
let expected_root = 1.324717957244746;
println!("result.root Bisection: {}", result.root);
assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
let result = solver.newton_raphson(&func, 1.0).unwrap();
let expected_root = 1.324717957244746;
println!("result.root Newton: {}", result.root);
assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
}
#[test]
fn test_secant_exponential() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x.exp() - 2.0, "e^x - 2".to_string());
let result = solver.secant(&func, 0.0, 1.0).unwrap();
let expected_root = 2.0_f64.ln();
assert!(approx_equal(result.root, expected_root, 1e-10));
assert!(result.converged);
}
#[test]
fn test_secant_root_at_initial_guess() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x - 2.0, "x - 2".to_string());
let result = solver.secant(&func, 2.0, 3.0).unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert_eq!(result.iterations, 0);
}
#[test]
fn test_secant_derivative_zero_error() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|_x| 1.0, "constant function".to_string());
let result = solver.secant(&func, 1.0, 1.0001);
assert!(matches!(result, Err(RootFindingError::DerivativeZero)));
}
#[test]
fn test_newton_raphson_with_derivative() {
let solver = ScalarRootFinder::new();
let func = FunctionWithDerivative::new(
|x| x * x - 4.0,
|x| 2.0 * x,
"x^2 - 4 with derivative".to_string(),
);
let result = solver.newton_raphson(&func, 1.0).unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert!(result.converged);
assert_eq!(result.method, "newton_raphson");
}
#[test]
fn test_newton_raphson_without_derivative() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x - 4.0, "x^2 - 4".to_string());
let result = solver.newton_raphson(&func, 1.0).unwrap();
assert!(approx_equal(result.root, 2.0, 1e-9));
assert!(result.converged);
}
#[test]
fn test_newton_raphson_complex_function() {
let solver = ScalarRootFinder::new();
let func = FunctionWithDerivative::new(
|x| x * x * x - 2.0 * x - 5.0,
|x| 3.0 * x * x - 2.0,
"x^3 - 2x - 5".to_string(),
);
let result = solver.newton_raphson(&func, 2.0).unwrap();
let expected_root = 2.094551481542327; assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
}
#[test]
fn test_solver_configuration() {
let config = RootFindingConfig {
tolerance: 1e-6,
max_iterations: 50,
verbose: false,
};
let solver = ScalarRootFinder::with_config(config);
let func = ClosureFunction::new(|x| x * x - 4.0, "x^2 - 4".to_string());
let result = solver.bisection(&func, 0.0, 3.0).unwrap();
assert!(approx_equal(result.root, 2.0, 1e-6));
assert!(result.converged);
}
#[test]
fn test_solver_max_iterations() {
let mut solver = ScalarRootFinder::new();
solver.set_max_iterations(5); solver.set_tolerance(1e-15);
let func = ClosureFunction::new(|x| x * x - 4.0, "x^2 - 4".to_string());
let result = solver.bisection(&func, 0.0, 3.0);
assert!(matches!(
result,
Err(RootFindingError::MaxIterationsReached)
));
}
#[test]
fn test_hybrid_solve_method() {
let solver = ScalarRootFinder::new();
let func_with_deriv = FunctionWithDerivative::new(
|x| x * x - 4.0,
|x| 2.0 * x,
"x^2 - 4 with derivative".to_string(),
);
let result = solver
.solve(&func_with_deriv, 1.0, Some((0.0, 3.0)))
.unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert_eq!(result.method, "newton_raphson");
let func_without_deriv = ClosureFunction::new(|x| x * x - 4.0, "x^2 - 4".to_string());
let result = solver
.solve(&func_without_deriv, 1.0, Some((0.0, 3.0)))
.unwrap();
assert!(approx_equal(result.root, 2.0, 1e-9));
assert_eq!(result.method, "secant");
}
#[test]
fn test_convenience_functions() {
let root = bisection(|x| x * x - 4.0, 0.0, 3.0, 1e-10).unwrap();
assert!(approx_equal(root, 2.0, 1e-10));
let root = secant(|x| x * x - 4.0, 1.0, 3.0, 1e-10).unwrap();
assert!(approx_equal(root, 2.0, 1e-10));
}
#[test]
fn test_real_world_examples() {
let solver = ScalarRootFinder::new();
let intersection_func =
ClosureFunction::new(|x| x * x - 2.0 * x - 3.0, "x^2 - 2x - 3".to_string());
let positive_root = solver.bisection(&intersection_func, 0.0, 5.0).unwrap();
assert!(approx_equal(positive_root.root, 3.0, 1e-10));
let negative_root = solver.bisection(&intersection_func, -5.0, 0.0).unwrap();
assert!(approx_equal(negative_root.root, -1.0, 1e-10));
}
}
#[cfg(test)]
mod brent_solver_tests {
use super::*;
use approx::assert_relative_eq;
use std::f64::consts::PI;
#[test]
fn test_brent_simple_quadratic() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x - 4.0, "x^2 - 4".to_string());
let result = solver.brent(&func, Some((0.0, 3.0))).unwrap();
assert_relative_eq!(result.root, 2.0, epsilon = 1e-10);
assert!(result.converged);
assert_eq!(result.method, "brent");
let result = solver.brent(&func, Some((-3.0, 0.0))).unwrap();
assert_relative_eq!(result.root, -2.0, epsilon = 1e-10);
assert!(result.converged);
}
#[test]
fn test_brent_cubic() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x * x - x - 1.0, "x^3 - x - 1".to_string());
let result = solver.brent(&func, Some((1.0, 2.0))).unwrap();
let expected_root = 1.324717957244746;
assert_relative_eq!(result.root, expected_root, epsilon = 1e-9);
assert!(result.converged);
}
#[test]
fn test_brent_trigonometric() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x.sin(), "sin(x)".to_string());
let result = solver.brent(&func, Some((3.0, 4.0))).unwrap();
assert_relative_eq!(result.root, PI, epsilon = 1e-10);
assert!(result.converged);
}
#[test]
fn test_brent_exponential() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x.exp() - 2.0, "e^x - 2".to_string());
let result = solver.brent(&func, Some((0.0, 1.0))).unwrap();
let expected_root = 2.0_f64.ln();
assert_relative_eq!(result.root, expected_root, epsilon = 1e-10);
assert!(result.converged);
}
#[test]
fn test_brent_no_search_range() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x - 4.0, "x^2 - 4".to_string());
let result = solver.brent(&func, None);
assert!(result.is_err());
match result.unwrap_err() {
RootFindingError::InvalidInput(msg) => {
assert!(msg.contains("Brent method requires search range"));
}
_ => panic!("Expected InvalidInput error"),
}
}
#[test]
fn test_brent_invalid_interval() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x + 1.0, "x^2 + 1".to_string());
let result = solver.brent(&func, Some((-1.0, 1.0)));
assert!(matches!(result, Err(RootFindingError::InvalidInterval)));
}
#[test]
fn test_brent_vs_other_methods() {
let solver = ScalarRootFinder::new();
let func = ClosureFunction::new(|x| x * x * x - 2.0 * x - 5.0, "x^3 - 2x - 5".to_string());
let brent_result = solver.brent(&func, Some((2.0, 3.0))).unwrap();
let bisection_result = solver.bisection(&func, 2.0, 3.0).unwrap();
let secant_result = solver.secant(&func, 2.0, 2.5).unwrap();
let expected_root = 2.094551481542327;
assert_relative_eq!(brent_result.root, expected_root, epsilon = 1e-9);
assert_relative_eq!(bisection_result.root, expected_root, epsilon = 1e-9);
assert_relative_eq!(secant_result.root, expected_root, epsilon = 1e-9);
assert!(brent_result.iterations <= bisection_result.iterations);
println!("Brent iterations: {}", brent_result.iterations);
println!("Bisection iterations: {}", bisection_result.iterations);
println!("Secant iterations: {}", secant_result.iterations);
}
}
#[cfg(test)]
mod symbolic_tests {
use super::*;
use std::collections::HashMap;
use std::f64::consts::E;
fn approx_equal(a: f64, b: f64, tolerance: f64) -> bool {
(a - b).abs() < tolerance
}
#[test]
fn test_symbolic_function_creation_from_string() {
let func =
SymbolicFunction::from_string("x^2 - 4", "x", Some("quadratic".to_string())).unwrap();
assert_eq!(func.name(), "quadratic");
assert!(approx_equal(func.evaluate(2.0), 0.0, 1e-10));
assert!(approx_equal(func.evaluate(-2.0), 0.0, 1e-10));
assert!(approx_equal(func.evaluate(0.0), -4.0, 1e-10));
assert!(func.derivative(2.0).is_some());
assert!(approx_equal(func.derivative(2.0).unwrap(), 4.0, 1e-10));
assert!(approx_equal(func.derivative(3.0).unwrap(), 6.0, 1e-10));
}
#[test]
fn test_symbolic_function_with_parameters() {
let mut params = HashMap::new();
params.insert("a".to_string(), 2.0);
params.insert("b".to_string(), -8.0);
let func = SymbolicFunction::from_string_with_params(
"a*x^2 + b",
"x",
params,
Some("parametric_quadratic".to_string()),
)
.unwrap();
assert!(approx_equal(func.evaluate(2.0), 0.0, 1e-10));
assert!(approx_equal(func.evaluate(-2.0), 0.0, 1e-10));
assert!(approx_equal(func.evaluate(0.0), -8.0, 1e-10));
assert!(approx_equal(func.derivative(2.0).unwrap(), 8.0, 1e-10));
assert!(approx_equal(func.derivative(1.0).unwrap(), 4.0, 1e-10));
}
#[test]
fn test_symbolic_function_parameter_update() {
let mut params = HashMap::new();
params.insert("a".to_string(), 1.0);
let mut func =
SymbolicFunction::from_string_with_params("a*x^2 - 4", "x", params, None).unwrap();
assert!(approx_equal(func.evaluate(2.0), 0.0, 1e-10));
let mut new_params = HashMap::new();
new_params.insert("a".to_string(), 4.0);
func.set_parameters(new_params).unwrap();
println!("func.evaluate(1.0) ,{}", func.evaluate(1.0));
assert!(approx_equal(func.evaluate(1.0), 0.0, 1e-10));
assert!(approx_equal(func.evaluate(2.0), 12.0, 1e-10));
}
#[test]
fn test_solve_symbolic_quadratic_newton_raphson() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"x^2 - 4",
"x",
RootFindingMethod::NewtonRaphson,
1.0, None,
None,
)
.unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert!(result.converged);
assert_eq!(result.method, "newton_raphson");
assert!(result.iterations > 0);
}
#[test]
fn test_solve_symbolic_quadratic_bisection() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"x^2 - 4",
"x",
RootFindingMethod::Bisection,
1.0, Some((0.0, 3.0)), None,
)
.unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert!(result.converged);
assert_eq!(result.method, "bisection");
}
#[test]
fn test_solve_symbolic_quadratic_secant() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"x^2 - 4",
"x",
RootFindingMethod::Secant,
1.0, None,
None,
)
.unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert!(result.converged);
assert_eq!(result.method, "secant");
}
#[test]
fn test_solve_symbolic_cubic() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"x^3 - (x + 1)",
"x",
RootFindingMethod::NewtonRaphson,
1.5, None,
None,
)
.unwrap();
println!("result.root: {}", result.root);
let expected_root = 1.324717957244746;
assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
}
#[test]
fn test_solve_symbolic_cubic2() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"x^3 - (x + 1)",
"x",
RootFindingMethod::Secant,
1.5, None,
None,
)
.unwrap();
println!("result.root: {}", result.root);
let expected_root = 1.324717957244746;
assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
}
#[test]
fn test_solve_symbolic_cubic3() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"x^3 - (x + 1)",
"x",
RootFindingMethod::Bisection,
1.5, Some((-10.0, 3.0)), None,
)
.unwrap();
println!("result.root: {}", result.root);
let expected_root = 1.324717957244746;
assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
}
#[test]
fn SymbolicFunction_from_string_with_params() {
let func = SymbolicFunction::from_string("x^3-(x+1)", "x", None).unwrap();
let F = func.func;
let exp = func.expr;
let f = F(0.0);
println!("f: {}", f);
println!("exp: {}", exp);
}
#[test]
fn test_solve_symbolic_exponential() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"exp(x) - 2",
"x",
RootFindingMethod::NewtonRaphson,
1.0, None,
None,
)
.unwrap();
let expected_root = 2.0_f64.ln();
assert!(approx_equal(result.root, expected_root, 1e-10));
assert!(result.converged);
}
#[test]
fn test_solve_symbolic_with_parameters() {
let solver = ScalarRootFinder::new();
let mut params = HashMap::new();
params.insert("a".to_string(), 3.0);
params.insert("b".to_string(), -12.0);
let result = solver
.solve_symbolic_str(
"a*x^2 + b",
"x",
RootFindingMethod::NewtonRaphson,
1.0, None,
Some(params),
)
.unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert!(result.converged);
}
#[test]
fn test_solve_symbolic_complex_polynomial() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"x^4 - 10*x^2 + 9",
"x",
RootFindingMethod::NewtonRaphson,
0.5, None,
None,
)
.unwrap();
assert!(approx_equal(result.root, 1.0, 1e-9));
assert!(result.converged);
let result2 = solver
.solve_symbolic_str(
"x^4 - 10*x^2 + 9",
"x",
RootFindingMethod::NewtonRaphson,
2.5, None,
None,
)
.unwrap();
assert!(approx_equal(result2.root, 3.0, 1e-9));
assert!(result2.converged);
}
#[test]
fn test_solve_symbolic_rational_function() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"(x^2 - 4)/(x + 1)",
"x",
RootFindingMethod::NewtonRaphson,
1.0, None,
None,
)
.unwrap();
assert!(approx_equal(result.root, 2.0, 1e-10));
assert!(result.converged);
}
#[test]
fn test_solve_symbolic_logarithmic() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"ln(x) - 1",
"x",
RootFindingMethod::Secant,
2.0, None,
None,
)
.unwrap();
assert!(approx_equal(result.root, E, 1e-10));
assert!(result.converged);
}
#[test]
fn test_solve_symbolic_mixed_functions() {
let solver = ScalarRootFinder::new();
let result = solver
.solve_symbolic_str(
"x*exp(x) - 1",
"x",
RootFindingMethod::NewtonRaphson,
0.5, None,
None,
)
.unwrap();
let expected_root = 0.567143290409784;
assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
}
#[test]
fn test_solve_symbolic_expr_object() {
let solver = ScalarRootFinder::new();
let expr = Expr::parse_expression("x^3 - (2*x + 5)");
let result = solver
.solve_symbolic_expr(
expr,
"x",
RootFindingMethod::NewtonRaphson,
2.0, None,
None,
)
.unwrap();
println!("result.root: {}", result.root);
let expected_root = 2.094551481542327;
assert!(approx_equal(result.root, expected_root, 1e-9));
assert!(result.converged);
}
#[test]
fn test_symbolic_function_expressions_strings() {
let func = SymbolicFunction::from_string("x^2 + 3*x - 4", "x", None).unwrap();
let expr_str = func.expression_string();
let deriv_str = func.derivative_string().unwrap();
assert!(expr_str.contains("x"));
assert!(deriv_str.contains("x"));
println!("Expression: {}", expr_str);
println!("Derivative: {}", deriv_str);
}
#[test]
fn test_solve_symbolic_with_different_tolerances() {
let mut solver = ScalarRootFinder::new();
solver.set_tolerance(1e-15);
let result = solver
.solve_symbolic_str(
"x^2 - 2",
"x",
RootFindingMethod::NewtonRaphson,
1.0,
None,
None,
)
.unwrap();
let expected_root = 2.0_f64.sqrt();
assert!(approx_equal(result.root, expected_root, 1e-14));
assert!(result.converged);
}
}