use super::Cas;
use super::pattern::{Call2, Guard, Pattern, Rule, RuleSet};
use laddu_cas_macros::cas_rules;
pub(super) fn standard(cas: &Cas) -> RuleSet {
let equations = cas_rules! {
x + 0 => x;
x * 1 => x;
x * 0 => 0;
x - x => 0;
x - 0 => x;
0 - x => -x;
x / 1 => x;
x / x => 1 if nonzero(x);
-(-x) => x;
x ^ 0 => 1;
x ^ 1 => x;
sqrt(x) ^ 2 => x;
conj(conj(x)) => x;
conj(complex(re, im)) => complex(re, -im) if real(re) && real(im);
conj(exp(x)) => exp(conj(x));
conj(-x) => -conj(x);
conj(x + y) => conj(x) + conj(y);
conj(x * y) => conj(x) * conj(y);
real(real(x)) => real(x);
imag(real(x)) => 0;
real(complex(re, im)) => re;
imag(complex(re, im)) => im;
real(conj(x)) => real(x);
imag(conj(x)) => -imag(x);
conj(x) => x if real(x);
real(x) => x if real(x);
imag(x) => 0 if real(x);
norm_sqr(conj(x)) => norm_sqr(x);
norm_sqr(x) => x ^ 2 if real(x);
norm_sqr(exp(I * x) * y) => norm_sqr(y) if real(x);
cos(-x) => cos(x);
sin(-x) => -sin(x);
sin(x)^2 + cos(x)^2 => 1;
1 - sin(x)^2 => cos(x)^2;
1 - cos(x)^2 => sin(x)^2;
1 + -sin(x)^2 => cos(x)^2;
1 + -cos(x)^2 => sin(x)^2;
sin(0.5 * x) ^ 2 => 0.5 * (1 - cos(x));
cos(0.5 * x) ^ 2 => 0.5 * (1 + cos(x));
exp(0) => 1;
exp(x) * exp(y) <=> exp(x + y);
cos(x) + I * sin(x) <=> exp(I * x) if real(x);
cos(x) - I * sin(x) => exp(-(I * x)) if real(x);
-(x + y) => -x + -y;
c * cos(x) + (c * I) * sin(x) => c * exp(I * x) if real(x);
x * a + x * b <=> x * (a + b);
x^2 - y^2 <=> (x - y) * (x + y);
matmul(identity, a) => a if identity(identity);
matmul(a, identity) => a if identity(identity);
matvec(identity, v) => v if identity(identity);
dot(zero, v) => 0 if zero(zero);
solve(identity, rhs) => rhs if identity(identity);
x + (-x) => 0;
x + x => 2 * x;
x - (-y) => x + y;
-(x - y) => y - x;
0 / x => 0 if nonzero(x);
(-x) / y => -(x / y);
x / (-y) => -(x / y);
(x * y) / x => y if nonzero(x);
(x / y) * y => x if nonzero(y);
sqrt(0) => 0;
sqrt(1) => 1;
I^2 => -1;
conj(I) => -I;
real(I) => 0;
imag(I) => 1;
complex(a, 0) => a if real(a);
complex(a, b) + complex(c, d) => complex(a + c, b + d) if real(a) && real(b) && real(c) && real(d);
conj(x / y) => conj(x) / conj(y) if nonzero(y);
real(x + y) => real(x) + real(y);
imag(x + y) => imag(x) + imag(y);
real(-x) => -real(x);
imag(-x) => -imag(x);
norm_sqr(0) => 0;
norm_sqr(1) => 1;
norm_sqr(-x) => norm_sqr(x);
norm_sqr(x / y) => norm_sqr(x) / norm_sqr(y) if nonzero(y);
matmul(zero, m) => zero if zero(zero);
matvec(zero, v) => zero if zero(zero);
matvec(a, zero) => zero if zero(zero);
dot(a, zero) => 0 if zero(zero);
cos(0) => 1;
sin(0) => 0;
};
RuleSet::new(cas, equations)
}