use crate::error::{MathError, Result};
use crate::expr::Expr;
use crate::simplify::simplify;
pub fn differentiate(expr: &Expr, var: &str) -> Result<Expr> {
Ok(simplify(&diff(expr, var)?))
}
fn diff(expr: &Expr, var: &str) -> Result<Expr> {
match expr {
Expr::Num(_) => Ok(Expr::num(0.0)),
Expr::Var(v) => {
if v == var {
Ok(Expr::num(1.0))
} else {
Ok(Expr::num(0.0))
}
}
Expr::Neg(e) => Ok(Expr::neg(diff(e, var)?)),
Expr::Add(a, b) => Ok(Expr::add(diff(a, var)?, diff(b, var)?)),
Expr::Sub(a, b) => Ok(Expr::sub(diff(a, var)?, diff(b, var)?)),
Expr::Mul(a, b) => {
let da = diff(a, var)?;
let db = diff(b, var)?;
Ok(Expr::add(
Expr::mul(da, (**b).clone()),
Expr::mul((**a).clone(), db),
))
}
Expr::Div(a, b) => {
let da = diff(a, var)?;
let db = diff(b, var)?;
Ok(Expr::div(
Expr::sub(
Expr::mul(da, (**b).clone()),
Expr::mul((**a).clone(), db),
),
Expr::pow((**b).clone(), Expr::num(2.0)),
))
}
Expr::Pow(base, exp) => match (base.as_ref(), exp.as_ref()) {
(_, Expr::Num(c)) => {
let c_val = *c;
let inner = Expr::pow((**base).clone(), Expr::num(c_val - 1.0));
Ok(Expr::mul(
Expr::mul(Expr::num(c_val), inner),
diff(base, var)?,
))
}
(Expr::Num(_), _) => Ok(Expr::mul(
expr.clone(),
Expr::mul(
Expr::func("ln", vec![(**base).clone()]),
diff(exp, var)?,
),
)),
_ => Ok(Expr::mul(
expr.clone(),
Expr::add(
Expr::mul(diff(exp, var)?, Expr::func("ln", vec![(**base).clone()])),
Expr::mul(
Expr::div((**exp).clone(), (**base).clone()),
diff(base, var)?,
),
),
)),
},
Expr::Func(name, args) => {
if args.len() != 1 {
return Err(MathError::Eval(format!(
"cannot differentiate multi-arg function {}",
name
)));
}
let arg = &args[0];
let arg_diff = diff(arg, var)?;
let deriv = derivative_of_builtin(name, arg)?;
Ok(Expr::mul(deriv, arg_diff))
}
}
}
fn derivative_of_builtin(name: &str, arg: &Expr) -> Result<Expr> {
let x = arg.clone();
Ok(match name {
"sin" => Expr::func("cos", vec![x]),
"cos" => Expr::neg(Expr::func("sin", vec![x])),
"tan" => Expr::div(
Expr::num(1.0),
Expr::pow(Expr::func("cos", vec![x]), Expr::num(2.0)),
),
"asin" => Expr::div(
Expr::num(1.0),
Expr::func(
"sqrt",
vec![Expr::sub(Expr::num(1.0), Expr::pow(x, Expr::num(2.0)))],
),
),
"acos" => Expr::neg(Expr::div(
Expr::num(1.0),
Expr::func(
"sqrt",
vec![Expr::sub(Expr::num(1.0), Expr::pow(x, Expr::num(2.0)))],
),
)),
"atan" => Expr::div(
Expr::num(1.0),
Expr::add(Expr::num(1.0), Expr::pow(x, Expr::num(2.0))),
),
"sinh" => Expr::func("cosh", vec![x]),
"cosh" => Expr::func("sinh", vec![x]),
"tanh" => Expr::div(
Expr::num(1.0),
Expr::pow(Expr::func("cosh", vec![x]), Expr::num(2.0)),
),
"exp" => Expr::func("exp", vec![x]),
"ln" | "log" => Expr::div(Expr::num(1.0), x),
"log2" => Expr::div(
Expr::num(1.0),
Expr::mul(x, Expr::func("ln", vec![Expr::num(2.0)])),
),
"log10" => Expr::div(
Expr::num(1.0),
Expr::mul(x, Expr::func("ln", vec![Expr::num(10.0)])),
),
"sqrt" => Expr::div(
Expr::num(1.0),
Expr::mul(Expr::num(2.0), Expr::func("sqrt", vec![x.clone()])),
),
"abs" => Expr::div(x.clone(), Expr::func("abs", vec![x])),
"floor" | "ceil" | "round" | "sign" | "fract" => Expr::num(0.0),
_ => {
return Err(MathError::Eval(format!(
"no symbolic derivative for function '{}'",
name
)))
}
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::eval::{eval, Context};
use crate::parser::Parser;
fn agrees(got: &Expr, want: &Expr, xs: &[f64]) {
let mut ctx = Context::standard();
for &x in xs {
ctx.set("x", x);
let g = eval(got, &ctx).unwrap();
let w = eval(want, &ctx).unwrap();
assert!(
(g - w).abs() < 1e-9,
"disagree at x={}: got {} want {}",
x,
g,
w
);
}
}
fn d(s: &str) -> Expr {
let e = Parser::parse(s).unwrap();
differentiate(&e, "x").unwrap()
}
#[test]
fn polynomial() {
let got = d("x^3 + 2*x^2 + x + 5");
let want = Parser::parse("3*x^2 + 4*x + 1").unwrap();
agrees(&got, &want, &[0.0, 1.0, -2.0, 3.5, 0.7]);
}
#[test]
fn product_rule() {
let got = d("x*sin(x)");
let want = Parser::parse("sin(x) + x*cos(x)").unwrap();
agrees(&got, &want, &[0.1, 0.5, 1.0, 2.0, -0.7]);
}
#[test]
fn quotient_rule() {
let got = d("x/(x+1)");
let want = Parser::parse("1/(x+1)^2").unwrap();
agrees(&got, &want, &[0.5, 1.5, 2.0, -0.5]);
}
#[test]
fn chain_rule() {
let got = d("sin(x^2)");
let want = Parser::parse("2*x*cos(x^2)").unwrap();
agrees(&got, &want, &[0.1, 0.5, 1.0, -0.7]);
}
#[test]
fn exp_ln() {
let got = d("exp(x)");
let want = Parser::parse("exp(x)").unwrap();
agrees(&got, &want, &[0.0, 1.0, -2.0]);
}
}