//! A reusable, read-only AST visitor with default traversal. //! //! The traversal is split into two layers so that many different static //! analyses can share a single definition of "how to walk the tree": //! //! * The `walk_*` free functions contain the canonical recursion. For each //! node they call back into the visitor on every child. This is the only //! place that needs to know the shape of the AST. //! * The [`Visitor`] trait's `visit_*` methods are the overridable hooks. Each //! one defaults to calling the matching `walk_*` function, so a visitor that //! overrides nothing still performs a complete traversal. //! //! To implement an analysis, implement [`Visitor`] and override only the hooks //! you care about. Inside an override, call the corresponding `walk_*` function //! whenever you want the default "descend into children" behaviour to happen. //! Accumulate results in your own struct fields (errors, scope tables, type //! information, ...) rather than through the return value, which is only used //! to short-circuit on error. //! //! # Example //! //! ```ignore //! struct IdentifierCounter { count: usize } //! //! impl Visitor for IdentifierCounter { //! fn visit_expr(&mut self, expr: &AstNode) -> LoxResult<()> { //! if let Expr::Identifier { .. } = &expr.node { //! self.count += 1; //! } //! walk_expr(self, expr) // keep descending into children //! } //! } //! ``` use crate::common::{ ast::{AstNode, Expr}, base_value::{BaseValue, LoxFunction}, lox_result::LoxResult, }; /// A read-only visitor over the AST. /// /// Every hook has a default implementation that performs the standard /// recursive traversal, so implementors only override the cases they need. pub trait Visitor: Sized { /// Visit an expression node. Defaults to [`walk_expr`]. fn visit_expr(&mut self, expr: &AstNode) -> LoxResult<()> { walk_expr(self, expr) } /// Visit a function literal (parameters, optional guard and body). /// /// Defaults to [`walk_function`], which walks the guard (if any) and the /// body. Override this to manage a parameter scope before descending. fn visit_function(&mut self, function: &LoxFunction) -> LoxResult<()> { walk_function(self, function) } } /// Recurse into the children of `expr`, calling back into `visitor`. pub fn walk_expr(visitor: &mut V, expr: &AstNode) -> LoxResult<()> { match &expr.node { // A function literal carries an entire sub-tree (its body), so it is // not a leaf: hand it to the dedicated function hook. Expr::Literal { value } => { if let BaseValue::Function(function) = &**value { visitor.visit_function(function)?; } Ok(()) } Expr::Identifier { .. } => Ok(()), Expr::Binary { left, right, .. } => { visitor.visit_expr(left)?; visitor.visit_expr(right) } Expr::Unary { operand, .. } => visitor.visit_expr(operand), Expr::Assign { value, .. } => visitor.visit_expr(value), Expr::Grouping { expression } => visitor.visit_expr(expression), Expr::Call { callee, arguments, .. } => { visitor.visit_expr(callee)?; for argument in arguments.iter() { visitor.visit_expr(argument)?; } Ok(()) } Expr::Print { expression, .. } => visitor.visit_expr(expression), Expr::VarDeclaration { initializer, .. } => { if let Some(initializer) = initializer { visitor.visit_expr(initializer)?; } Ok(()) } Expr::Return { expression, .. } => visitor.visit_expr(expression), Expr::Block { statements, .. } => { for statement in statements.iter() { visitor.visit_expr(statement)?; } Ok(()) } Expr::If { condition, then_branch, elif_branches, else_branch, .. } => { visitor.visit_expr(condition)?; visitor.visit_expr(then_branch)?; for (elif_condition, elif_body) in elif_branches.iter() { visitor.visit_expr(elif_condition)?; visitor.visit_expr(elif_body)?; } if let Some(else_body) = else_branch { visitor.visit_expr(else_body)?; } Ok(()) } Expr::While { condition, body, .. } => { visitor.visit_expr(condition)?; visitor.visit_expr(body) } Expr::For { variable, condition, increment, body, .. } => { visitor.visit_expr(variable)?; visitor.visit_expr(condition)?; visitor.visit_expr(increment)?; visitor.visit_expr(body) } // Struct declarations carry only type information, no child expressions. Expr::Struct { .. } => Ok(()), } } /// Walk the guard (if present) and body of a function literal. pub fn walk_function(visitor: &mut V, function: &LoxFunction) -> LoxResult<()> { if let Some(guard) = &function.guard { visitor.visit_expr(guard)?; } visitor.visit_expr(&function.body) } #[cfg(test)] mod tests { use super::*; use crate::common::lox_result::runtime_error; use crate::frontend::lexer::Lexer; use crate::frontend::parser::Parser; fn parse(src: &str) -> Vec { let tokens = Lexer::new(src.to_string(), 0) .scans_tokens() .expect("source should lex"); Parser::new(tokens).parse().expect("source should parse") } /// A visitor that relies entirely on the default traversal and just counts /// how many nodes it sees. #[derive(Default)] struct Counter { nodes: usize, } impl Visitor for Counter { fn visit_expr(&mut self, node: &AstNode) -> LoxResult<()> { self.nodes += 1; walk_expr(self, node) } } fn count(src: &str) -> Counter { let mut counter = Counter::default(); for stmt in parse(src).iter() { counter.visit_expr(stmt).unwrap(); } counter } #[test] fn counts_every_node_via_default_traversal() { // 1 + 2 * 3 => Binary(+){ Literal, Binary(*){ Literal, Literal } } let counter = count("1 + 2 * 3;"); assert_eq!(counter.nodes, 5); } #[test] fn descends_into_function_bodies() { // The function body must be traversed through `visit_function`, so the // `return a;` inside it should contribute to the counts. let counter = count("f :: fn (a): Number do return a; end;"); // VarDeclaration + Function literal + Block + Return + identifier `a` assert_eq!(counter.nodes, 5); } /// A visitor that aborts as soon as it sees an identifier, used to check /// that errors short-circuit the traversal. struct FailOnIdentifier; impl Visitor for FailOnIdentifier { fn visit_expr(&mut self, expr: &AstNode) -> LoxResult<()> { if let Expr::Identifier { .. } = &expr.node { return runtime_error(expr.source_slice.clone(), "found an identifier"); } walk_expr(self, expr) } } #[test] fn errors_propagate_through_traversal() { let stmts = parse("var x: Int = 1; x;"); let mut visitor = FailOnIdentifier; let mut result = Ok(()); for stmt in stmts.iter() { result = visitor.visit_expr(stmt); if result.is_err() { break; } } assert!(result.is_err()); } }