use std::collections::{HashMap, HashSet};
use crate::{
CompilationError, SourceSpan,
ast::{
self, ASTVisitor, Add, Binding, Constant, CustomDice, Div, DropHighest,
DropLowest, Event, Exp, Expression, Group, Mod, Mul, Neg, Node, Range,
StandardDice, Sub, Variable, Walk
}
};
#[derive(Copy, Clone, Debug, Default)]
pub struct Validator;
impl Validator
{
#[inline]
pub const fn new() -> Self { Self }
pub fn validate<'src>(
ast: &ast::Function<'src>
) -> Result<(), CompilationError<'src>>
{
check_duplicate_parameters(ast)?;
let bindings = collect_bindings_and_check_collisions(ast)?;
check_use_before_bind(&ast.body, &bindings)
}
}
fn check_duplicate_parameters<'src>(
ast: &ast::Function<'src>
) -> Result<(), CompilationError<'src>>
{
if let Some(ref parameters) = ast.parameters
{
let mut seen: HashMap<&str, SourceSpan> =
HashMap::with_capacity(parameters.len());
for param in parameters
{
if let Some(&first) = seen.get(&*param.name)
{
return Err(CompilationError::DuplicateParameter {
name: param.name.clone(),
first,
duplicate: param.span
});
}
seen.insert(¶m.name, param.span);
}
}
Ok(())
}
fn collect_bindings_and_check_collisions<'a, 'src>(
ast: &'a ast::Function<'src>
) -> Result<HashMap<&'a str, SourceSpan>, CompilationError<'src>>
{
let parameter_spans = match ast.parameters
{
Some(ref parameters) => parameters
.iter()
.map(|p| (&*p.name, p.span))
.collect::<HashMap<_, _>>(),
None => HashMap::new()
};
let mut bindings: HashMap<&'a str, SourceSpan> = HashMap::new();
for event in Walk::new(Node::Expression(&ast.body))
{
if let Event::Enter(Node::Expression(Expression::Binding(b))) = event
{
if let Some(¶meter) = parameter_spans.get(&*b.name)
{
return Err(CompilationError::BindingCollidesWithParameter {
name: b.name.clone(),
parameter,
binding: b.name_span
});
}
if let Some(&first) = bindings.get(&*b.name)
{
return Err(CompilationError::DuplicateBinding {
name: b.name.clone(),
first,
duplicate: b.name_span
});
}
bindings.insert(&b.name, b.name_span);
}
}
Ok(bindings)
}
fn check_use_before_bind<'a, 'src>(
body: &'a Expression<'src>,
bindings: &HashMap<&'a str, SourceSpan>
) -> Result<(), CompilationError<'src>>
{
let mut seen: HashSet<&'a str> = HashSet::new();
for event in Walk::new(Node::Expression(body))
{
match event
{
Event::Enter(Node::Expression(Expression::Variable(v))) =>
{
if let Some(&binding_span) = bindings.get(&*v.name)
&& !seen.contains(&*v.name)
{
return Err(CompilationError::UseBeforeBind {
name: v.name.clone(),
reference: v.span,
binding: binding_span
});
}
},
Event::Leave(Node::Expression(Expression::Binding(b))) =>
{
seen.insert(&b.name);
},
_ =>
{}
}
}
Ok(())
}
impl<'a, 'src: 'a> ASTVisitor<'a, 'src> for Validator
{
type Error = CompilationError<'src>;
type Output = ();
fn enter_function(
&mut self,
node: &'a ast::Function<'src>
) -> Result<(), Self::Error>
{
check_duplicate_parameters(node)
}
fn visit_function(
&mut self,
_node: &'a ast::Function<'src>,
_body: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_group(
&mut self,
_node: &'a Group<'src>,
_expression: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_constant(&mut self, _node: &'a Constant)
-> Result<(), Self::Error>
{
Ok(())
}
fn visit_variable(
&mut self,
_node: &'a Variable<'src>
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_binding(
&mut self,
_node: &'a Binding<'src>,
_expression: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_range(
&mut self,
_node: &'a Range<'src>,
_start: (),
_end: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_standard_dice(
&mut self,
_node: &'a StandardDice<'src>,
_count: (),
_faces: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_custom_dice(
&mut self,
_node: &'a CustomDice<'src>,
_count: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_drop_lowest(
&mut self,
_node: &'a DropLowest<'src>,
_dice: (),
_drop: Option<()>
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_drop_highest(
&mut self,
_node: &'a DropHighest<'src>,
_dice: (),
_drop: Option<()>
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_add(
&mut self,
_node: &'a Add<'src>,
_left: (),
_right: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_sub(
&mut self,
_node: &'a Sub<'src>,
_left: (),
_right: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_mul(
&mut self,
_node: &'a Mul<'src>,
_left: (),
_right: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_div(
&mut self,
_node: &'a Div<'src>,
_left: (),
_right: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_mod(
&mut self,
_node: &'a Mod<'src>,
_left: (),
_right: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_exp(
&mut self,
_node: &'a Exp<'src>,
_left: (),
_right: ()
) -> Result<(), Self::Error>
{
Ok(())
}
fn visit_neg(
&mut self,
_node: &'a Neg<'src>,
_operand: ()
) -> Result<(), Self::Error>
{
Ok(())
}
}