use super::{
ASTVisitor, Add, ArithmeticExpression, DiceExpression, Div, Exp,
Expression, Mod, Mul, Sub
};
#[derive(Copy, Clone)]
pub(crate) enum Node<'a, 'src>
{
Expression(&'a Expression<'src>),
Dice(&'a DiceExpression<'src>),
Arithmetic(&'a ArithmeticExpression<'src>)
}
#[derive(Copy, Clone)]
pub(crate) enum Event<'a, 'src>
{
Enter(Node<'a, 'src>),
Leave(Node<'a, 'src>)
}
pub(crate) struct Walk<'a, 'src>
{
stack: Vec<Event<'a, 'src>>
}
impl<'a, 'src> Walk<'a, 'src>
{
pub(crate) fn new(root: Node<'a, 'src>) -> Self
{
Walk {
stack: vec![Event::Enter(root)]
}
}
fn expression(&mut self, expression: &'a Expression<'src>)
{
self.stack.push(Event::Enter(Node::Expression(expression)));
}
fn schedule_children(&mut self, node: Node<'a, 'src>)
{
match node
{
Node::Expression(expression) => match expression
{
Expression::Group(group) => self.expression(&group.expression),
Expression::Constant(_) | Expression::Variable(_) =>
{},
Expression::Binding(binding) =>
{
self.expression(&binding.expression)
},
Expression::Range(range) =>
{
self.expression(&range.end);
self.expression(&range.start);
},
Expression::Dice(dice) => self.schedule_dice_children(dice),
Expression::Arithmetic(arithmetic) =>
{
self.schedule_arithmetic_children(arithmetic)
},
},
Node::Dice(dice) => self.schedule_dice_children(dice),
Node::Arithmetic(arithmetic) =>
{
self.schedule_arithmetic_children(arithmetic)
},
}
}
fn schedule_dice_children(&mut self, dice: &'a DiceExpression<'src>)
{
let (dice, drop) = match dice
{
DiceExpression::Standard(standard) =>
{
self.expression(&standard.faces);
self.expression(&standard.count);
return
},
DiceExpression::Custom(custom) =>
{
self.expression(&custom.count);
return
},
DiceExpression::DropLowest(clause) => (&clause.dice, &clause.drop),
DiceExpression::DropHighest(clause) => (&clause.dice, &clause.drop)
};
if let Some(drop) = drop
{
self.expression(drop);
}
self.stack.push(Event::Enter(Node::Dice(dice)));
}
fn schedule_arithmetic_children(
&mut self,
arithmetic: &'a ArithmeticExpression<'src>
)
{
match arithmetic
{
ArithmeticExpression::Add(Add { left, right, .. })
| ArithmeticExpression::Sub(Sub { left, right, .. })
| ArithmeticExpression::Mul(Mul { left, right, .. })
| ArithmeticExpression::Div(Div { left, right, .. })
| ArithmeticExpression::Mod(Mod { left, right, .. })
| ArithmeticExpression::Exp(Exp { left, right, .. }) =>
{
self.expression(right);
self.expression(left);
},
ArithmeticExpression::Neg(neg) => self.expression(&neg.operand)
}
}
}
impl<'a, 'src> Iterator for Walk<'a, 'src>
{
type Item = Event<'a, 'src>;
fn next(&mut self) -> Option<Self::Item>
{
let event = self.stack.pop()?;
if let Event::Enter(node) = event
{
self.stack.push(Event::Leave(node));
self.schedule_children(node);
}
Some(event)
}
}
pub(super) fn fold<'a, 'src, V: ASTVisitor<'a, 'src>>(
root: Node<'a, 'src>,
visitor: &mut V
) -> Result<V::Output, V::Error>
{
let mut outputs = Vec::new();
for event in Walk::new(root)
{
match event
{
Event::Enter(node) => enter(node, visitor)?,
Event::Leave(node) =>
{
let output = leave(node, visitor, &mut outputs)?;
outputs.push(output);
}
}
}
let output = pop(&mut outputs);
debug_assert!(outputs.is_empty(), "a walk must consume every output");
Ok(output)
}
fn enter<'a, 'src, V: ASTVisitor<'a, 'src>>(
node: Node<'a, 'src>,
visitor: &mut V
) -> Result<(), V::Error>
{
match node
{
Node::Expression(expression) => match expression
{
Expression::Group(group) => visitor.enter_group(group),
Expression::Constant(constant) => visitor.enter_constant(constant),
Expression::Variable(variable) => visitor.enter_variable(variable),
Expression::Binding(binding) => visitor.enter_binding(binding),
Expression::Range(range) => visitor.enter_range(range),
Expression::Dice(dice) => enter_dice(dice, visitor),
Expression::Arithmetic(arithmetic) =>
{
enter_arithmetic(arithmetic, visitor)
},
},
Node::Dice(dice) => enter_dice(dice, visitor),
Node::Arithmetic(arithmetic) => enter_arithmetic(arithmetic, visitor)
}
}
fn enter_dice<'a, 'src, V: ASTVisitor<'a, 'src>>(
dice: &'a DiceExpression<'src>,
visitor: &mut V
) -> Result<(), V::Error>
{
match dice
{
DiceExpression::Standard(standard) =>
{
visitor.enter_standard_dice(standard)
},
DiceExpression::Custom(custom) => visitor.enter_custom_dice(custom),
DiceExpression::DropLowest(clause) => visitor.enter_drop_lowest(clause),
DiceExpression::DropHighest(clause) =>
{
visitor.enter_drop_highest(clause)
},
}
}
fn enter_arithmetic<'a, 'src, V: ASTVisitor<'a, 'src>>(
arithmetic: &'a ArithmeticExpression<'src>,
visitor: &mut V
) -> Result<(), V::Error>
{
match arithmetic
{
ArithmeticExpression::Add(add) => visitor.enter_add(add),
ArithmeticExpression::Sub(sub) => visitor.enter_sub(sub),
ArithmeticExpression::Mul(mul) => visitor.enter_mul(mul),
ArithmeticExpression::Div(div) => visitor.enter_div(div),
ArithmeticExpression::Mod(modulo) => visitor.enter_mod(modulo),
ArithmeticExpression::Exp(exp) => visitor.enter_exp(exp),
ArithmeticExpression::Neg(neg) => visitor.enter_neg(neg)
}
}
fn leave<'a, 'src, V: ASTVisitor<'a, 'src>>(
node: Node<'a, 'src>,
visitor: &mut V,
outputs: &mut Vec<V::Output>
) -> Result<V::Output, V::Error>
{
match node
{
Node::Expression(expression) =>
{
let output = match expression
{
Expression::Group(group) =>
{
let inner = pop(outputs);
visitor.visit_group(group, inner)
},
Expression::Constant(constant) =>
{
visitor.visit_constant(constant)
},
Expression::Variable(variable) =>
{
visitor.visit_variable(variable)
},
Expression::Binding(binding) =>
{
let inner = pop(outputs);
visitor.visit_binding(binding, inner)
},
Expression::Range(range) =>
{
let (start, end) = pop_pair(outputs);
visitor.visit_range(range, start, end)
},
Expression::Dice(dice) => leave_dice(dice, visitor, outputs),
Expression::Arithmetic(arithmetic) =>
{
leave_arithmetic(arithmetic, visitor, outputs)
},
}?;
visitor.visit_expression(expression, output)
},
Node::Dice(dice) => leave_dice(dice, visitor, outputs),
Node::Arithmetic(arithmetic) =>
{
leave_arithmetic(arithmetic, visitor, outputs)
},
}
}
fn leave_dice<'a, 'src, V: ASTVisitor<'a, 'src>>(
dice: &'a DiceExpression<'src>,
visitor: &mut V,
outputs: &mut Vec<V::Output>
) -> Result<V::Output, V::Error>
{
match dice
{
DiceExpression::Standard(standard) =>
{
let (count, faces) = pop_pair(outputs);
visitor.visit_standard_dice(standard, count, faces)
},
DiceExpression::Custom(custom) =>
{
let count = pop(outputs);
visitor.visit_custom_dice(custom, count)
},
DiceExpression::DropLowest(clause) =>
{
let drop = clause.drop.as_ref().map(|_| pop(outputs));
let dice = pop(outputs);
visitor.visit_drop_lowest(clause, dice, drop)
},
DiceExpression::DropHighest(clause) =>
{
let drop = clause.drop.as_ref().map(|_| pop(outputs));
let dice = pop(outputs);
visitor.visit_drop_highest(clause, dice, drop)
}
}
}
fn leave_arithmetic<'a, 'src, V: ASTVisitor<'a, 'src>>(
arithmetic: &'a ArithmeticExpression<'src>,
visitor: &mut V,
outputs: &mut Vec<V::Output>
) -> Result<V::Output, V::Error>
{
if let ArithmeticExpression::Neg(neg) = arithmetic
{
let operand = pop(outputs);
return visitor.visit_neg(neg, operand)
}
let (left, right) = pop_pair(outputs);
match arithmetic
{
ArithmeticExpression::Add(add) => visitor.visit_add(add, left, right),
ArithmeticExpression::Sub(sub) => visitor.visit_sub(sub, left, right),
ArithmeticExpression::Mul(mul) => visitor.visit_mul(mul, left, right),
ArithmeticExpression::Div(div) => visitor.visit_div(div, left, right),
ArithmeticExpression::Mod(modulo) =>
{
visitor.visit_mod(modulo, left, right)
},
ArithmeticExpression::Exp(exp) => visitor.visit_exp(exp, left, right),
ArithmeticExpression::Neg(_) => unreachable!()
}
}
fn pop<T>(outputs: &mut Vec<T>) -> T
{
outputs
.pop()
.expect("a child's output must precede its parent's visit")
}
fn pop_pair<T>(outputs: &mut Vec<T>) -> (T, T)
{
let second = pop(outputs);
let first = pop(outputs);
(first, second)
}