use core::fmt::Display;
use core::ops::Deref;
use alloc::boxed::Box;
use alloc::string::ToString;
use alloc::{vec, vec::Vec};
use num_traits::{FromPrimitive, Zero};
use rust_decimal::{Decimal, MathematicalOps};
use crate::Number;
use crate::error::MathsError;
use crate::node::common;
use crate::render::{Glyph, LayoutBlock, Layoutable, Renderer, LayoutComputationProperties};
use crate::nav::NavPathNavigator;
use super::function::Function;
use super::simplified::{Simplifiable, SimplifiedNode};
#[derive(PartialEq, Eq, Debug, Clone)]
pub enum StructuredNode {
Number(Number),
Variable(char),
Sqrt(Box<StructuredNode>),
Power(Box<StructuredNode>, Box<StructuredNode>),
Add(Box<StructuredNode>, Box<StructuredNode>),
Subtract(Box<StructuredNode>, Box<StructuredNode>),
Multiply(Box<StructuredNode>, Box<StructuredNode>),
Divide(Box<StructuredNode>, Box<StructuredNode>),
Parentheses(Box<StructuredNode>),
FunctionCall(Function, Vec<StructuredNode>),
}
#[derive(PartialEq, Eq, Debug, Copy, Clone)]
pub enum AngleUnit {
Degree,
Radian,
}
impl Default for AngleUnit {
fn default() -> Self {
Self::Degree
}
}
impl Display for AngleUnit {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str(match self {
Self::Degree => "Degree",
Self::Radian => "Radian",
})
}
}
#[derive(PartialEq, Eq, Debug, Clone, Default)]
pub struct EvaluationSettings {
pub angle_unit: AngleUnit,
}
impl StructuredNode {
pub fn add_or_sub(&self) -> bool {
matches!(&self, StructuredNode::Add(_, _) | StructuredNode::Subtract(_, _))
}
pub fn mul_or_div(&self) -> bool {
matches!(&self, StructuredNode::Multiply(_, _) | StructuredNode::Divide(_, _))
}
pub fn in_parentheses(&self) -> StructuredNode {
StructuredNode::Parentheses(box self.clone())
}
pub fn in_parentheses_or_clone(&self, parens: bool) -> StructuredNode {
if parens {
self.in_parentheses()
} else {
self.clone()
}
}
pub fn disambiguate(&self) -> Result<StructuredNode, MathsError> {
Ok(match self {
StructuredNode::Multiply(l, r) => {
let l = l.in_parentheses_or_clone(l.add_or_sub());
let r = r.in_parentheses_or_clone(r.add_or_sub() || r.mul_or_div());
StructuredNode::Multiply(box l, box r)
}
StructuredNode::Divide(l, r) => {
let l = l.in_parentheses_or_clone(l.add_or_sub());
let r = r.in_parentheses_or_clone(r.add_or_sub() || r.mul_or_div());
StructuredNode::Divide(box l, box r)
}
StructuredNode::Add(l, r) => {
let r = r.in_parentheses_or_clone(r.add_or_sub());
StructuredNode::Add(l.clone(), box r)
}
StructuredNode::Subtract(l, r) => {
let r = r.in_parentheses_or_clone(r.add_or_sub());
StructuredNode::Subtract(l.clone(), box r)
}
StructuredNode::Number(_) | StructuredNode::Sqrt(_) | StructuredNode::Parentheses(_) | StructuredNode::Variable(_) | StructuredNode::Power(_, _) | StructuredNode::FunctionCall(_, _)
=> self.clone(),
})
}
pub fn evaluate(&self, settings: &EvaluationSettings) -> Result<Number, MathsError> {
match self {
StructuredNode::Number(n) => Ok((*n).into()),
StructuredNode::Variable(_) => Err(MathsError::MissingVariable),
StructuredNode::Sqrt(inner) =>
inner.evaluate(settings)?.to_decimal().sqrt().map(|x| x.into()).ok_or(MathsError::InvalidSqrt),
StructuredNode::Power(b, e) => b.evaluate(settings)?.checked_pow(e.evaluate(settings)?),
StructuredNode::Add(a, b) => a.evaluate(settings)?.checked_add(b.evaluate(settings)?),
StructuredNode::Subtract(a, b) => a.evaluate(settings)?.checked_sub(b.evaluate(settings)?),
StructuredNode::Multiply(a, b) => a.evaluate(settings)?.checked_mul(b.evaluate(settings)?),
StructuredNode::Divide(a, b) => a.evaluate(settings)?.checked_div(b.evaluate(settings)?),
StructuredNode::Parentheses(inner) => inner.evaluate(settings),
StructuredNode::FunctionCall(func, args) => {
let args = args.iter().map(|n| n.evaluate(settings)).collect::<Result<Vec<_>, _>>()?;
func.evaluate(&args, settings)
}
}
}
pub fn walk(&self, func: &impl Fn(&StructuredNode)) {
func(self);
match self {
StructuredNode::Add(l, r)
| StructuredNode::Subtract(l, r)
| StructuredNode::Multiply(l, r)
| StructuredNode::Divide(l, r) => {
l.walk(func);
r.walk(func);
},
StructuredNode::Sqrt(inner) | StructuredNode::Parentheses(inner) => {
inner.walk(func);
},
StructuredNode::Power(b, e) => {
b.walk(func);
e.walk(func);
}
StructuredNode::FunctionCall(_, args) => {
for arg in args {
arg.walk(func);
}
}
StructuredNode::Number(_) | StructuredNode::Variable(_) => (),
}
}
pub fn walk_mut(&mut self, func: &mut impl FnMut(&mut StructuredNode)) {
func(self);
match self {
StructuredNode::Add(l, r)
| StructuredNode::Subtract(l, r)
| StructuredNode::Multiply(l, r)
| StructuredNode::Divide(l, r) => {
l.walk_mut(func);
r.walk_mut(func);
},
StructuredNode::Sqrt(inner) | StructuredNode::Parentheses(inner) => {
inner.walk_mut(func);
},
StructuredNode::Power(b, e) => {
b.walk_mut(func);
e.walk_mut(func);
}
StructuredNode::FunctionCall(_, args) => {
for arg in args {
arg.walk_mut(func);
}
}
StructuredNode::Number(_) | StructuredNode::Variable(_) => (),
}
}
pub fn substitute_variable(&self, var_name: char, subst: &StructuredNode) -> StructuredNode {
let mut clone = self.clone();
clone.walk_mut(&mut |n| {
if let StructuredNode::Variable(actual_var_name) = n {
if *actual_var_name == var_name {
*n = subst.clone();
}
}
});
clone
}
}
fn layout_binop(renderer: &mut impl Renderer, glyph: Glyph, properties: LayoutComputationProperties, left: &StructuredNode, right: &StructuredNode) -> LayoutBlock {
let left_layout = left.layout(renderer, None, properties);
let binop_layout = LayoutBlock::from_glyph(renderer, glyph, properties)
.move_right_of_other(&left_layout);
let right_layout = right.layout(renderer, None, properties)
.move_right_of_other(&binop_layout);
left_layout
.merge_along_baseline(&binop_layout)
.merge_along_baseline(&right_layout)
}
impl Layoutable for StructuredNode {
fn layout(&self, renderer: &mut impl Renderer, path: Option<&mut NavPathNavigator>, properties: LayoutComputationProperties) -> LayoutBlock {
match self {
StructuredNode::Number(Number::Decimal(mut number)) => {
let negative = number < Decimal::zero();
if negative {
number = -number;
}
let mut glyph_layouts = number
.to_string()
.chars()
.map(|c|
if c == '.' {
Glyph::Point
} else {
Glyph::Digit { number: c.to_digit(10).unwrap() as u8 }
}
)
.map(|g| LayoutBlock::from_glyph(renderer, g, properties))
.collect::<Vec<_>>();
if negative {
glyph_layouts.insert(
0,
LayoutBlock::from_glyph(renderer, Glyph::Subtract, properties)
)
}
LayoutBlock::layout_horizontal(&glyph_layouts[..])
},
StructuredNode::Number(Number::Rational(numer, denom)) => {
if *denom == 1 {
StructuredNode::Number(Number::Decimal(Decimal::from_i64(*numer).unwrap())).layout(renderer, path, properties)
} else {
common::layout_fraction(
&StructuredNode::Number(Number::Decimal(Decimal::from_i64(*numer).unwrap())),
&StructuredNode::Number(Number::Decimal(Decimal::from_i64(*denom).unwrap())),
renderer,
None,
properties,
)
}
},
StructuredNode::Variable(v) => LayoutBlock::from_glyph(renderer, Glyph::Variable { name: *v }, properties),
StructuredNode::Add(left, right) => layout_binop(renderer, Glyph::Add, properties, left, right),
StructuredNode::Subtract(left, right) => layout_binop(renderer, Glyph::Subtract, properties, left, right),
StructuredNode::Multiply(left, right) => layout_binop(renderer, Glyph::Multiply, properties, left, right),
StructuredNode::Divide(top, bottom)
=> common::layout_fraction(top.deref(), bottom.deref(), renderer, path, properties),
StructuredNode::Sqrt(inner)
=> common::layout_sqrt(inner.deref(), renderer, path, properties),
StructuredNode::Parentheses(inner)
=> common::layout_parentheses(inner.deref(), renderer, path, properties),
StructuredNode::Power(base, exp)
=> common::layout_power(Some(base.deref()), exp.deref(), renderer, path, properties),
StructuredNode::FunctionCall(func, args)
=> common::layout_function_call(*func, args, renderer, path, properties),
}
}
}
impl Simplifiable for StructuredNode {
fn simplify(&self) -> SimplifiedNode {
match self {
&Self::Number(n) => SimplifiedNode::Number(n),
&Self::Variable(n) => SimplifiedNode::Variable(n),
&Self::Add(ref l, ref r) => SimplifiedNode::Add(vec![
l.simplify(),
r.simplify(),
]),
&Self::Subtract(ref l, ref r) => SimplifiedNode::Add(vec![
l.simplify(),
r.simplify().negate(),
]),
&Self::Multiply(ref l, ref r) => SimplifiedNode::Multiply(vec![
l.simplify(),
r.simplify(),
]),
&Self::Divide(ref l, ref r) => SimplifiedNode::Multiply(vec![
l.simplify(),
r.simplify().reciprocal(),
]),
Self::Sqrt(n) => SimplifiedNode::Power(
box n.simplify(),
box SimplifiedNode::Number(Number::Rational(1, 2)),
),
Self::Power(b, e) => SimplifiedNode::Power(
box b.simplify(),
box e.simplify(),
),
Self::FunctionCall(func, args) => SimplifiedNode::FunctionCall(
*func,
args.iter().map(|n| n.simplify()).collect(),
),
Self::Parentheses(n) => n.simplify(),
}
}
}