use super::operation::{
associative_operation, unary_operation, AssociativeOperator, UnaryOperator,
};
use super::*;
use crate::{CoefficientError, LinearMonomial, Monomial, MonomialDyn, QuadraticMonomial};
use std::ops::{Add, Neg};
impl Function {
pub fn zero() -> Self {
Function::Zero
}
pub fn is_zero(&self) -> bool {
matches!(self, Function::Zero)
}
pub(crate) fn normalize(self) -> Self {
fn constant_term<M: crate::Monomial>(p: &crate::PolynomialBase<M>) -> Coefficient {
p.get(&M::default())
.expect("non-zero degree-0 polynomial has a constant term")
}
match self {
Function::Zero | Function::Constant(_) => self,
Function::Linear(l) => {
if l.is_zero() {
Function::Zero
} else if l.degree() == 0 {
Function::Constant(constant_term(&l))
} else {
Function::Linear(l)
}
}
Function::Quadratic(q) => {
if q.is_zero() {
return Function::Zero;
}
match q.degree().into_inner() {
2 => Function::Quadratic(q),
1 => Function::Linear(
Linear::try_from(&q).expect("degree-1 polynomial is linear"),
),
_ => Function::Constant(constant_term(&q)),
}
}
Function::Polynomial(p) => {
if p.is_zero() {
return Function::Zero;
}
let mut degree = Degree::from(0);
for monomial in p.keys() {
degree = degree.max(monomial.degree());
if degree > 2 {
break;
}
}
if degree > 2 {
return Function::Polynomial(p);
}
match degree.into_inner() {
2 => Function::Quadratic(
Quadratic::try_from(&p).expect("degree-2 polynomial is quadratic"),
),
1 => Function::Linear(
Linear::try_from(&p).expect("degree-1 polynomial is linear"),
),
_ => Function::Constant(constant_term(&p)),
}
}
Function::Expression(_) => self,
}
}
pub fn try_add_assign_in_place(&mut self, rhs: Self) -> Result<(), CoefficientError> {
let lhs = std::mem::take(self);
*self = match (lhs, rhs) {
(Function::Zero, rhs) => rhs,
(lhs, Function::Zero) => lhs,
(Function::Constant(lhs), Function::Constant(rhs)) => {
if let Some(coefficient) = (lhs + rhs)? {
Function::Constant(coefficient)
} else {
Function::Zero
}
}
(Function::Constant(c), Function::Linear(mut l))
| (Function::Linear(mut l), Function::Constant(c)) => {
l.add_term(LinearMonomial::Constant, c)?;
Function::Linear(l)
}
(Function::Constant(c), Function::Quadratic(mut q))
| (Function::Quadratic(mut q), Function::Constant(c)) => {
q.add_term(QuadraticMonomial::Constant, c)?;
Function::Quadratic(q)
}
(Function::Constant(c), Function::Polynomial(mut p))
| (Function::Polynomial(mut p), Function::Constant(c)) => {
p.add_term(MonomialDyn::default(), c)?;
Function::Polynomial(p)
}
(Function::Linear(lhs), Function::Linear(rhs)) => Function::Linear((lhs + rhs)?),
(Function::Linear(l), Function::Quadratic(q))
| (Function::Quadratic(q), Function::Linear(l)) => Function::Quadratic((q + &l)?),
(Function::Linear(l), Function::Polynomial(p))
| (Function::Polynomial(p), Function::Linear(l)) => Function::Polynomial((p + &l)?),
(Function::Quadratic(lhs), Function::Quadratic(rhs)) => {
Function::Quadratic((lhs + rhs)?)
}
(Function::Quadratic(q), Function::Polynomial(p))
| (Function::Polynomial(p), Function::Quadratic(q)) => Function::Polynomial((p + &q)?),
(Function::Polynomial(lhs), Function::Polynomial(rhs)) => {
Function::Polynomial((lhs + rhs)?)
}
(lhs, rhs) => associative_operation(AssociativeOperator::Add, lhs, rhs),
};
Ok(())
}
fn try_add_assign_ref_in_place(&mut self, rhs: &Self) -> Result<(), CoefficientError> {
let lhs = std::mem::take(self);
*self = match (lhs, rhs) {
(Function::Zero, rhs) => rhs.clone(),
(lhs, Function::Zero) => lhs,
(Function::Constant(lhs), Function::Constant(rhs)) => {
if let Some(coefficient) = (lhs + *rhs)? {
Function::Constant(coefficient)
} else {
Function::Zero
}
}
(Function::Constant(c), Function::Linear(l)) => Function::Linear((l + c)?),
(Function::Linear(l), Function::Constant(c)) => Function::Linear((l + *c)?),
(Function::Constant(c), Function::Quadratic(q)) => Function::Quadratic((q + c)?),
(Function::Quadratic(q), Function::Constant(c)) => Function::Quadratic((q + *c)?),
(Function::Constant(c), Function::Polynomial(p)) => Function::Polynomial((p + c)?),
(Function::Polynomial(p), Function::Constant(c)) => Function::Polynomial((p + *c)?),
(Function::Linear(lhs), Function::Linear(rhs)) => Function::Linear((lhs + rhs)?),
(Function::Linear(l), Function::Quadratic(q)) => Function::Quadratic((q + &l)?),
(Function::Quadratic(q), Function::Linear(l)) => Function::Quadratic((q + l)?),
(Function::Linear(l), Function::Polynomial(p)) => Function::Polynomial((p + &l)?),
(Function::Polynomial(p), Function::Linear(l)) => Function::Polynomial((p + l)?),
(Function::Quadratic(lhs), Function::Quadratic(rhs)) => {
Function::Quadratic((lhs + rhs)?)
}
(Function::Quadratic(q), Function::Polynomial(p)) => Function::Polynomial((p + &q)?),
(Function::Polynomial(p), Function::Quadratic(q)) => Function::Polynomial((p + q)?),
(Function::Polynomial(lhs), Function::Polynomial(rhs)) => {
Function::Polynomial((lhs + rhs)?)
}
(lhs, rhs) => associative_operation(AssociativeOperator::Add, lhs, rhs.clone()),
};
Ok(())
}
fn try_add_refs(lhs: &Self, rhs: &Self) -> Result<Self, CoefficientError> {
Ok(match (lhs, rhs) {
(Function::Zero, rhs) => rhs.clone(),
(lhs, Function::Zero) => lhs.clone(),
(Function::Constant(lhs), Function::Constant(rhs)) => {
if let Some(coefficient) = (*lhs + *rhs)? {
Function::Constant(coefficient)
} else {
Function::Zero
}
}
(Function::Constant(c), Function::Linear(l))
| (Function::Linear(l), Function::Constant(c)) => Function::Linear((l + *c)?),
(Function::Constant(c), Function::Quadratic(q))
| (Function::Quadratic(q), Function::Constant(c)) => Function::Quadratic((q + *c)?),
(Function::Constant(c), Function::Polynomial(p))
| (Function::Polynomial(p), Function::Constant(c)) => Function::Polynomial((p + *c)?),
(Function::Linear(lhs), Function::Linear(rhs)) => Function::Linear((lhs + rhs)?),
(Function::Linear(l), Function::Quadratic(q))
| (Function::Quadratic(q), Function::Linear(l)) => Function::Quadratic((q + l)?),
(Function::Linear(l), Function::Polynomial(p))
| (Function::Polynomial(p), Function::Linear(l)) => Function::Polynomial((p + l)?),
(Function::Quadratic(lhs), Function::Quadratic(rhs)) => {
Function::Quadratic((lhs + rhs)?)
}
(Function::Quadratic(q), Function::Polynomial(p))
| (Function::Polynomial(p), Function::Quadratic(q)) => Function::Polynomial((p + q)?),
(Function::Polynomial(lhs), Function::Polynomial(rhs)) => {
Function::Polynomial((lhs + rhs)?)
}
(lhs, rhs) => associative_operation(AssociativeOperator::Add, lhs.clone(), rhs.clone()),
})
}
fn try_add_linear_ref(mut self, rhs: &Linear) -> Result<Self, CoefficientError> {
self = match self {
Function::Zero => Function::Linear(rhs.clone()),
Function::Constant(c) => Function::Linear((rhs + c)?),
Function::Linear(lhs) => Function::Linear((lhs + rhs)?),
Function::Quadratic(lhs) => Function::Quadratic((lhs + rhs)?),
Function::Polynomial(lhs) => Function::Polynomial((lhs + rhs)?),
lhs @ Function::Expression(_) => {
associative_operation(AssociativeOperator::Add, lhs, Function::Linear(rhs.clone()))
}
};
Ok(self)
}
fn try_add_quadratic_ref(mut self, rhs: &Quadratic) -> Result<Self, CoefficientError> {
self = match self {
Function::Zero => Function::Quadratic(rhs.clone()),
Function::Constant(c) => Function::Quadratic((rhs + c)?),
Function::Linear(lhs) => Function::Quadratic((rhs + &lhs)?),
Function::Quadratic(lhs) => Function::Quadratic((lhs + rhs)?),
Function::Polynomial(lhs) => Function::Polynomial((lhs + rhs)?),
lhs @ Function::Expression(_) => associative_operation(
AssociativeOperator::Add,
lhs,
Function::Quadratic(rhs.clone()),
),
};
Ok(self)
}
fn try_add_polynomial_ref(mut self, rhs: &Polynomial) -> Result<Self, CoefficientError> {
self = match self {
Function::Zero => Function::Polynomial(rhs.clone()),
Function::Constant(c) => Function::Polynomial((rhs + c)?),
Function::Linear(lhs) => Function::Polynomial((rhs + &lhs)?),
Function::Quadratic(lhs) => Function::Polynomial((rhs + &lhs)?),
Function::Polynomial(lhs) => Function::Polynomial((lhs + rhs)?),
lhs @ Function::Expression(_) => associative_operation(
AssociativeOperator::Add,
lhs,
Function::Polynomial(rhs.clone()),
),
};
Ok(self)
}
fn try_add_linear_refs(lhs: &Self, rhs: &Linear) -> Result<Self, CoefficientError> {
Ok(match lhs {
Function::Zero => Function::Linear(rhs.clone()),
Function::Constant(c) => Function::Linear((rhs + *c)?),
Function::Linear(lhs) => Function::Linear((lhs + rhs)?),
Function::Quadratic(lhs) => Function::Quadratic((lhs + rhs)?),
Function::Polynomial(lhs) => Function::Polynomial((lhs + rhs)?),
lhs @ Function::Expression(_) => associative_operation(
AssociativeOperator::Add,
lhs.clone(),
Function::Linear(rhs.clone()),
),
})
}
fn try_add_quadratic_refs(lhs: &Self, rhs: &Quadratic) -> Result<Self, CoefficientError> {
Ok(match lhs {
Function::Zero => Function::Quadratic(rhs.clone()),
Function::Constant(c) => Function::Quadratic((rhs + *c)?),
Function::Linear(lhs) => Function::Quadratic((rhs + lhs)?),
Function::Quadratic(lhs) => Function::Quadratic((lhs + rhs)?),
Function::Polynomial(lhs) => Function::Polynomial((lhs + rhs)?),
lhs @ Function::Expression(_) => associative_operation(
AssociativeOperator::Add,
lhs.clone(),
Function::Quadratic(rhs.clone()),
),
})
}
fn try_add_polynomial_refs(lhs: &Self, rhs: &Polynomial) -> Result<Self, CoefficientError> {
Ok(match lhs {
Function::Zero => Function::Polynomial(rhs.clone()),
Function::Constant(c) => Function::Polynomial((rhs + *c)?),
Function::Linear(lhs) => Function::Polynomial((rhs + lhs)?),
Function::Quadratic(lhs) => Function::Polynomial((rhs + lhs)?),
Function::Polynomial(lhs) => Function::Polynomial((lhs + rhs)?),
lhs @ Function::Expression(_) => associative_operation(
AssociativeOperator::Add,
lhs.clone(),
Function::Polynomial(rhs.clone()),
),
})
}
}
impl Add for Function {
type Output = Result<Self, CoefficientError>;
fn add(self, rhs: Self) -> Self::Output {
let mut out = self;
out.try_add_assign_in_place(rhs)?;
Ok(out.normalize())
}
}
impl Add for &Function {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: Self) -> Self::Output {
Ok(Function::try_add_refs(self, rhs)?.normalize())
}
}
impl Add<Function> for &Function {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: Function) -> Self::Output {
if self.is_polynomial() && rhs.is_polynomial() {
rhs + self
} else {
self.clone() + rhs
}
}
}
impl Add<&Function> for Function {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: &Function) -> Self::Output {
let mut out = self;
out.try_add_assign_ref_in_place(rhs)?;
Ok(out.normalize())
}
}
impl Add<Coefficient> for Function {
type Output = Result<Self, CoefficientError>;
fn add(mut self, rhs: Coefficient) -> Self::Output {
self.try_add_assign_in_place(Function::Constant(rhs))?;
Ok(self.normalize())
}
}
impl Add<Coefficient> for &Function {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: Coefficient) -> Self::Output {
self + Function::Constant(rhs)
}
}
impl Add<&Coefficient> for Function {
type Output = Result<Self, CoefficientError>;
fn add(self, rhs: &Coefficient) -> Self::Output {
self + *rhs
}
}
impl Add<&Coefficient> for &Function {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: &Coefficient) -> Self::Output {
self + Function::Constant(*rhs)
}
}
impl Add<Function> for Coefficient {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: Function) -> Self::Output {
Function::Constant(self) + rhs
}
}
impl Add<&Function> for Coefficient {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: &Function) -> Self::Output {
Function::Constant(self) + rhs
}
}
macro_rules! impl_add_polynomial_rhs {
($rhs:ty, $owned_ref_method:ident, $refs_method:ident) => {
impl Add<&$rhs> for Function {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: &$rhs) -> Self::Output {
Ok(self.$owned_ref_method(rhs)?.normalize())
}
}
impl Add<&$rhs> for &Function {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: &$rhs) -> Self::Output {
Ok(Function::$refs_method(self, rhs)?.normalize())
}
}
impl Add<$rhs> for Function {
type Output = Result<Function, CoefficientError>;
fn add(mut self, rhs: $rhs) -> Self::Output {
self.try_add_assign_in_place(Function::from(rhs))?;
Ok(self.normalize())
}
}
impl Add<$rhs> for &Function {
type Output = Result<Function, CoefficientError>;
fn add(self, rhs: $rhs) -> Self::Output {
self + Function::from(rhs)
}
}
};
}
impl_add_polynomial_rhs!(Linear, try_add_linear_ref, try_add_linear_refs);
impl_add_polynomial_rhs!(Quadratic, try_add_quadratic_ref, try_add_quadratic_refs);
impl_add_polynomial_rhs!(Polynomial, try_add_polynomial_ref, try_add_polynomial_refs);
impl Neg for Function {
type Output = Self;
fn neg(mut self) -> Self::Output {
if self.is_polynomial() {
self.values_mut()
.expect("polynomial function has coefficient values")
.for_each(|v| *v = -(*v));
return self;
}
unary_operation(UnaryOperator::Neg, self)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{coeff, linear, FunctionParameters, PolynomialParameters};
use ::approx::assert_abs_diff_eq;
use proptest::prelude::*;
fn polynomial_function() -> BoxedStrategy<Function> {
Function::arbitrary_with(FunctionParameters::polynomial_only(
PolynomialParameters::default(),
))
}
proptest! {
#[test]
fn add_ref(a in any::<Function>(), b in any::<Function>()) {
let ans = (a.clone() + b.clone()).unwrap();
assert_abs_diff_eq!((&a + &b).unwrap(), ans);
assert_abs_diff_eq!((&a + b.clone()).unwrap(), ans);
assert_abs_diff_eq!((a + &b).unwrap(), ans);
}
#[test]
fn zero(a in any::<Function>()) {
assert_abs_diff_eq!((&a + Function::zero()).unwrap(), a.clone());
assert_abs_diff_eq!((Function::zero() + &a).unwrap(), a.clone());
}
#[test]
fn add_commutative(a in polynomial_function(), b in polynomial_function()) {
assert_abs_diff_eq!((&a + &b).unwrap(), (&b + &a).unwrap());
}
#[test]
fn add_nary(
a in polynomial_function(),
b in polynomial_function(),
c in polynomial_function(),
) {
assert_abs_diff_eq!((&a + (&b + &c).unwrap()).unwrap(), ((&a + &b).unwrap() + &c).unwrap());
}
#[test]
fn fixed_rhs_refs_match_function_dispatch(
lhs in any::<Function>(),
linear in any::<Linear>(),
quadratic in any::<Quadratic>(),
polynomial in any::<Polynomial>(),
) {
prop_assert_eq!(lhs.clone() + &linear, lhs.clone() + Function::Linear(linear.clone()));
prop_assert_eq!(&lhs + &linear, &lhs + Function::Linear(linear));
prop_assert_eq!(lhs.clone() + &quadratic, lhs.clone() + Function::Quadratic(quadratic.clone()));
prop_assert_eq!(&lhs + &quadratic, &lhs + Function::Quadratic(quadratic));
prop_assert_eq!(lhs.clone() + &polynomial, lhs.clone() + Function::Polynomial(polynomial.clone()));
prop_assert_eq!(&lhs + &polynomial, &lhs + Function::Polynomial(polynomial));
}
}
#[test]
fn arithmetic_normalizes_polynomials_but_keeps_division_composed() {
let linear = Function::from(linear!(1));
let constant = Function::from(coeff!(2.0));
assert!(matches!(
(linear.clone() + constant).unwrap(),
Function::Linear(_)
));
assert!(matches!(
(linear.clone() * Function::from(linear!(2))).unwrap(),
Function::Quadratic(_)
));
assert!(matches!(
(linear.clone() / coeff!(2.0)).unwrap(),
Function::Expression(_)
));
assert!(matches!((linear.clone() - linear).unwrap(), Function::Zero));
}
#[test]
fn borrowed_addition_preserves_coefficient_error() {
let huge = Function::Linear(Linear::single_term(linear!(1), coeff!(f64::MAX)));
assert!(matches!(&huge + &huge, Err(CoefficientError::Infinite)));
}
#[test]
fn fixed_borrowed_addition_preserves_expression_order_and_coefficient_error() {
let lhs = Function::from(linear!(1)).abs();
let rhs = Linear::from(linear!(2));
let expected = (lhs.clone() + Function::Linear(rhs.clone())).unwrap();
assert!((lhs.clone() + &rhs).unwrap() == expected);
assert!((&lhs + &rhs).unwrap() == expected);
let huge = Function::Linear(Linear::single_term(linear!(1), coeff!(f64::MAX)));
let huge_rhs = Linear::single_term(linear!(1), coeff!(f64::MAX));
assert!(matches!(
huge.clone() + &huge_rhs,
Err(CoefficientError::Infinite)
));
assert!(matches!(&huge + &huge_rhs, Err(CoefficientError::Infinite)));
}
}