use std::sync::Arc;
use super::{
atom::Atom,
evaluate::{FunctionMap, PreparedEvaluator},
state::Symbol,
};
use crate::wrap_symbol;
impl Atom {
#[inline(always)]
pub fn lambdify_borrowed_thread_safe(
&self,
vars: &[&str],
) -> Box<dyn Fn(&[f64]) -> f64 + Send + Sync> {
let var_symbols: Vec<Symbol> = vars
.iter()
.map(|name| Symbol::new(wrap_symbol!(*name)))
.collect();
let prepared = Arc::new(
PreparedEvaluator::new(self, &var_symbols, &FunctionMap::new())
.expect("lambdify: failed to prepare evaluator"),
);
let n_vars = vars.len();
Box::new(move |vals: &[f64]| {
assert_eq!(
vals.len(),
n_vars,
"lambdify: expected {} argument(s), got {}",
n_vars,
vals.len()
);
prepared
.evaluate(vals)
.expect("lambdify: evaluation failed")
})
}
}
#[cfg(test)]
mod tests {
use crate::parse;
fn check(expr: &str, vars: &[&str], vals: &[f64], expected: f64) {
let atom = parse!(expr).unwrap();
let f = atom.lambdify_borrowed_thread_safe(vars);
let got = f(vals);
assert!(
(got - expected).abs() < 1e-10,
"{expr}({vals:?}) = {got}, expected {expected}"
);
}
#[test]
fn constant_expression() {
check("2+3", &[], &[], 5.0);
}
#[test]
fn single_variable_identity() {
check("x", &["x"], &[7.0], 7.0);
}
#[test]
fn linear_polynomial() {
check("3*x+2", &["x"], &[4.0], 14.0);
}
#[test]
fn two_variable_sum() {
check("x+y", &["x", "y"], &[3.0, 5.0], 8.0);
}
#[test]
fn two_variable_product() {
check("x*y", &["x", "y"], &[3.0, 5.0], 15.0);
}
#[test]
fn subtraction_via_negation() {
check("x-y", &["x", "y"], &[10.0, 3.0], 7.0);
}
#[test]
fn division() {
check("x/y", &["x", "y"], &[9.0, 3.0], 3.0);
}
#[test]
fn integer_power() {
check("x^3", &["x"], &[2.0], 8.0);
}
#[test]
fn fractional_power_as_sqrt() {
check("x^(1/2)", &["x"], &[9.0], 3.0);
}
#[test]
fn negative_power() {
check("x^-1", &["x"], &[4.0], 0.25);
}
#[test]
fn exp_function() {
let atom = parse!("exp(x)").unwrap();
let f = atom.lambdify_borrowed_thread_safe(&["x"]);
assert!((f(&[1.0]) - std::f64::consts::E).abs() < 1e-10);
}
#[test]
fn log_function() {
let atom = parse!("log(x)").unwrap();
let f = atom.lambdify_borrowed_thread_safe(&["x"]);
assert!((f(&[std::f64::consts::E]) - 1.0).abs() < 1e-10);
}
#[test]
fn sin_cos_identity() {
let atom = parse!("sin(x)^2+cos(x)^2").unwrap();
let f = atom.lambdify_borrowed_thread_safe(&["x"]);
for x in [0.0, 0.5, 1.0, 2.0, std::f64::consts::PI] {
assert!((f(&[x]) - 1.0).abs() < 1e-10, "failed at x={x}");
}
}
#[test]
fn trig_aliases() {
check("tg(x)", &["x"], &[0.0], 0.0);
check("ctg(x)", &["x"], &[std::f64::consts::FRAC_PI_4], 1.0);
}
#[test]
fn inverse_trig() {
check("arcsin(x)", &["x"], &[1.0], std::f64::consts::FRAC_PI_2);
check("arccos(x)", &["x"], &[0.0], std::f64::consts::FRAC_PI_2);
check("arctg(x)", &["x"], &[1.0], std::f64::consts::FRAC_PI_4);
}
#[test]
fn sqrt_function() {
check("sqrt(x)", &["x"], &[16.0], 4.0);
}
#[test]
fn builtin_pi_constant() {
let atom = parse!("sin(pi)").unwrap();
let f = atom.lambdify_borrowed_thread_safe(&[]);
assert!(f(&[]).abs() < 1e-10);
}
#[test]
fn builtin_e_constant() {
let atom = parse!("log(e)").unwrap();
let f = atom.lambdify_borrowed_thread_safe(&[]);
assert!((f(&[]) - 1.0).abs() < 1e-10);
}
#[test]
fn nested_expression() {
let atom = parse!("sin(x^2+1)").unwrap();
let f = atom.lambdify_borrowed_thread_safe(&["x"]);
assert!((f(&[0.0]) - 1.0_f64.sin()).abs() < 1e-10);
}
#[test]
fn multivariate_polynomial() {
check("x^2+2*x*y+y^2", &["x", "y"], &[3.0, 4.0], 49.0);
}
#[test]
fn three_variables() {
check(
"a*x^2+b*x+c",
&["a", "x", "b", "c"],
&[1.0, 2.0, 3.0, 4.0],
14.0,
);
}
#[test]
fn closure_is_send_sync() {
let atom = parse!("x^2+1").unwrap();
let f = atom.lambdify_borrowed_thread_safe(&["x"]);
let handle = std::thread::spawn(move || f(&[3.0]));
assert!((handle.join().unwrap() - 10.0).abs() < 1e-10);
}
#[test]
fn closure_shared_across_threads() {
use std::sync::Arc;
let atom = parse!("x^2").unwrap();
let f = Arc::new(atom.lambdify_borrowed_thread_safe(&["x"]));
let handles: Vec<_> = (1..=4)
.map(|i| {
let f = Arc::clone(&f);
std::thread::spawn(move || f(&[i as f64]))
})
.collect();
let results: Vec<f64> = handles.into_iter().map(|h| h.join().unwrap()).collect();
assert_eq!(results, vec![1.0, 4.0, 9.0, 16.0]);
}
#[test]
#[should_panic(expected = "lambdify: expected 1 argument(s), got 2")]
fn wrong_arity_panics() {
let atom = parse!("x").unwrap();
let f = atom.lambdify_borrowed_thread_safe(&["x"]);
f(&[1.0, 2.0]); }
#[test]
#[should_panic(expected = "lambdify: failed to prepare evaluator")]
fn unbound_variable_panics() {
let atom = parse!("x+y").unwrap();
let f = atom.lambdify_borrowed_thread_safe(&["x"]);
f(&[1.0]);
}
}