use crate::math::solver1d::{DerivativeSolver, Solver1D, func1d};
use crate::types::Real;
use std::cell::Cell;
pub fn f1(x: Real) -> Real {
x * x - 1.0
}
pub fn f2(x: Real) -> Real {
1.0 - x * x
}
pub fn f3(x: Real) -> Real {
(x - 1.0).atan()
}
pub fn d1(x: Real) -> Real {
2.0 * x
}
pub fn d2(x: Real) -> Real {
-2.0 * x
}
pub fn d3(x: Real) -> Real {
let u = x - 1.0;
1.0 / (1.0 + u * u)
}
pub fn dd1(_x: Real) -> Real {
2.0
}
pub fn dd2(_x: Real) -> Real {
-2.0
}
pub fn dd3(x: Real) -> Real {
let u = x - 1.0;
let d = 1.0 + u * u;
-2.0 * u / (d * d)
}
const ACCURACIES: [Real; 3] = [1.0e-4, 1.0e-6, 1.0e-8];
pub fn check_finds_known_roots<S: Solver1D>(make: impl Fn() -> S) {
for f in [f1 as fn(Real) -> Real, f2] {
for guess in [0.5, 1.5] {
for acc in ACCURACIES {
let root = make().solve(f, acc, guess, 0.1).unwrap();
assert!(
(root - 1.0).abs() <= acc,
"auto: guess={guess} acc={acc} root={root}"
);
let root = make().solve_bracketed(f, acc, guess, 0.0, 2.0).unwrap();
assert!(
(root - 1.0).abs() <= acc,
"bracketed: guess={guess} acc={acc} root={root}"
);
}
}
}
for acc in ACCURACIES {
let root = make().solve(f3, acc, 1.00001, 0.1).unwrap();
assert!((root - 1.0).abs() <= acc, "f3: acc={acc} root={root}");
}
}
pub fn check_last_call_with_root<S: Solver1D>(make: impl Fn() -> S) {
let mins = [3.0, 2.25, 1.5, 1.0];
let maxs = [7.0, 5.75, 4.5, 3.0];
let steps = [0.2, 0.2, 0.1, 0.1];
let offsets = [25.0, 11.0, 5.0, 1.0];
let guesses = [4.5, 4.5, 2.5, 2.5];
let accuracy = 1.0e-6;
for bracketed in [false, true] {
let argument = Cell::new(0.0);
for i in 0..4 {
let previous = argument.get();
let probe = |x: Real| {
argument.set(x);
previous + offsets[i] - x * x
};
let result = if bracketed {
make()
.solve_bracketed(probe, accuracy, guesses[i], mins[i], maxs[i])
.unwrap()
} else {
make().solve(probe, accuracy, guesses[i], steps[i]).unwrap()
};
assert!(
(result - argument.get()).abs() <= 2.0 * Real::EPSILON,
"bracketed={bracketed} i={i}: result={result} last_arg={}",
argument.get()
);
}
}
}
pub fn check_rejects_invalid_inputs<S: Solver1D>(make: impl Fn() -> S) {
assert!(make().solve(f1, 0.0, 0.5, 0.1).is_err());
assert!(make().solve_bracketed(f1, 1e-8, 2.5, 2.0, 3.0).is_err());
assert!(make().solve_bracketed(f1, 1e-8, 5.0, 0.0, 2.0).is_err());
}
pub fn check_honours_bounds<S: Solver1D>(make: impl Fn() -> S) {
let mut solver = make();
solver.set_upper_bound(2.0);
assert!(solver.solve_bracketed(f1, 1e-8, 1.5, 0.0, 3.0).is_err());
let mut solver = make();
solver.set_lower_bound(0.0);
solver.set_upper_bound(5.0);
let root = solver.solve(f1, 1e-10, 0.5, 0.1).unwrap();
assert!((root - 1.0).abs() <= 1e-9, "root={root}");
}
pub fn check_derivative_solver_finds_roots<S: DerivativeSolver>(make: impl Fn() -> S) {
let cases = [(f1 as fn(Real) -> Real, d1 as fn(Real) -> Real), (f2, d2)];
for (f, d) in cases {
for guess in [0.5, 1.5] {
for acc in ACCURACIES {
let root = make().solve(func1d(f, d), acc, guess, 0.1).unwrap();
assert!(
(root - 1.0).abs() <= acc,
"auto: guess={guess} acc={acc} root={root}"
);
let root = make()
.solve_bracketed(func1d(f, d), acc, guess, 0.0, 2.0)
.unwrap();
assert!(
(root - 1.0).abs() <= acc,
"bracketed: guess={guess} acc={acc} root={root}"
);
}
}
}
}
pub fn check_derivative_last_call<S: DerivativeSolver>(make: impl Fn() -> S) {
let mins = [3.0, 2.25, 1.5, 1.0];
let maxs = [7.0, 5.75, 4.5, 3.0];
let steps = [0.2, 0.2, 0.1, 0.1];
let offsets = [25.0, 11.0, 5.0, 1.0];
let guesses = [4.5, 4.5, 2.5, 2.5];
let accuracy = 1.0e-6;
for bracketed in [false, true] {
let argument = Cell::new(0.0);
for i in 0..4 {
let previous = argument.get();
let value = |x: Real| {
argument.set(x);
previous + offsets[i] - x * x
};
let derivative = |x: Real| -2.0 * x;
let g = func1d(value, derivative);
let result = if bracketed {
make()
.solve_bracketed(g, accuracy, guesses[i], mins[i], maxs[i])
.unwrap()
} else {
make().solve(g, accuracy, guesses[i], steps[i]).unwrap()
};
assert!(
(result - argument.get()).abs() <= 2.0 * Real::EPSILON,
"bracketed={bracketed} i={i}: result={result} last_arg={}",
argument.get()
);
}
}
}
pub fn check_derivative_rejects<S: DerivativeSolver>(make: impl Fn() -> S) {
assert!(make().solve(func1d(f1, d1), 0.0, 0.5, 0.1).is_err());
assert!(
make()
.solve_bracketed(func1d(f1, d1), 1e-8, 2.5, 2.0, 3.0)
.is_err()
);
assert!(
make()
.solve_bracketed(func1d(f1, d1), 1e-8, 5.0, 0.0, 2.0)
.is_err()
);
}
pub fn check_safe_derivative_solver<S: DerivativeSolver>(make: impl Fn() -> S) {
for acc in ACCURACIES {
let root = make().solve(func1d(f3, d3), acc, 1.00001, 0.1).unwrap();
assert!(
(root - 1.0).abs() <= acc,
"f3 stress: acc={acc} root={root}"
);
}
}