From b04b938b58edc6be4505cf1c35f6f33ac8684ed3 Mon Sep 17 00:00:00 2001 From: Alex van de Sandt Date: Thu, 11 Jul 2024 23:10:48 -0400 Subject: [PATCH] Store function AST data in separate struct --- src/ast/mod.rs | 2 +- src/ast/stmt.rs | 26 +++++++++++++++++--------- src/interpreter.rs | 4 ++-- src/parser.rs | 12 ++++++------ src/resolver.rs | 4 ++-- 5 files changed, 28 insertions(+), 20 deletions(-) diff --git a/src/ast/mod.rs b/src/ast/mod.rs index 5320e27..6504a52 100644 --- a/src/ast/mod.rs +++ b/src/ast/mod.rs @@ -4,7 +4,7 @@ mod stmt; use crate::span::Spanned; pub use expr::{BinaryOp, Expr, Literal, LogicalOp, Lval, UnaryOp}; -pub use stmt::Stmt; +pub use stmt::{Function, Stmt}; #[derive(Clone, Debug)] pub struct Ast(pub Vec>); diff --git a/src/ast/stmt.rs b/src/ast/stmt.rs index 5f7fe92..8a562e0 100644 --- a/src/ast/stmt.rs +++ b/src/ast/stmt.rs @@ -15,11 +15,7 @@ pub enum Stmt { name: String, // spanned? initializer: Option>, }, - Function { - name: String, - params: Vec, - body: Vec>, - }, + Function(Function), Block(Vec>), If { condition: Spanned, @@ -30,7 +26,6 @@ pub enum Stmt { condition: Spanned, body: Box>, }, - Return(#[allow(dead_code)] Option>), } @@ -74,7 +69,7 @@ impl Stmt { debug_assert_eq!(closing_brace.kind, TokenKind::RightBrace); let span = fun_token.span.join(&closing_brace.span); - Spanned::new(Self::Function { name, params, body }, span) + Spanned::new(Self::Function(Function::new(name, params, body)), span) } pub fn block( @@ -248,9 +243,9 @@ impl Stmt { } } - pub fn unwrap_fun_decl(&self) -> (&String, &[String], &[Spanned]) { + pub fn unwrap_fun_decl(&self) -> &Function { match self { - Self::Function { name, params, body } => (name, params, body), + Self::Function(f) => f, other => panic!("expected function decl, found {:?}", other.as_str()), } } @@ -329,3 +324,16 @@ impl Stmt { } } } + +#[derive(Clone, Debug)] +pub struct Function { + pub name: String, + pub params: Vec, + pub body: Vec>, +} + +impl Function { + pub fn new(name: String, params: Vec, body: Vec>) -> Self { + Function { name, params, body } + } +} diff --git a/src/interpreter.rs b/src/interpreter.rs index ca69ad8..ccfb668 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -3,7 +3,7 @@ use std::io::Write; use miette::SourceSpan; use crate::{ - ast::{Ast, BinaryOp, Expr, LogicalOp, Lval, Stmt, UnaryOp}, + ast::{Ast, BinaryOp, Expr, Function, LogicalOp, Lval, Stmt, UnaryOp}, callable::Callable, callable::LoxFunction, environment::Env, @@ -77,7 +77,7 @@ impl<'a> Interpreter<'a> { .unwrap_or(Value::Nil); env.define(name.clone(), val); } - Stmt::Function { name, params, body } => { + Stmt::Function(Function { name, params, body }) => { let callable = LoxFunction::new(params.clone(), env.clone(), body.clone()); env.define(name.clone(), Value::callable(callable)); } diff --git a/src/parser.rs b/src/parser.rs index 83c0507..c9448e3 100644 --- a/src/parser.rs +++ b/src/parser.rs @@ -783,21 +783,21 @@ mod tests { #[test] fn parses_fun_decl() { let stmt = parse_stmt("fun no_args() {}"); - let (name, args, body) = stmt.unwrap_fun_decl(); + let Function { name, params, body } = stmt.unwrap_fun_decl(); assert_eq!(name, "no_args"); - assert!(args.is_empty()); + assert!(params.is_empty()); assert!(body.is_empty()); let stmt = parse_stmt("fun one_arg_one_stmt(bar) { print bar; }"); - let (name, args, body) = stmt.unwrap_fun_decl(); + let Function { name, params, body } = stmt.unwrap_fun_decl(); assert_eq!(name, "one_arg_one_stmt"); - assert_eq!(args.len(), 1); + assert_eq!(params.len(), 1); assert_eq!(body.len(), 1); let stmt = parse_stmt("fun two_args_one_stmt(bar, baz) { print bar + baz; }"); - let (name, args, body) = stmt.unwrap_fun_decl(); + let Function { name, params, body } = stmt.unwrap_fun_decl(); assert_eq!(name, "two_args_one_stmt"); - assert_eq!(args.len(), 2); + assert_eq!(params.len(), 2); assert_eq!(body.len(), 1); } diff --git a/src/resolver.rs b/src/resolver.rs index 25eec12..9d6cdc1 100644 --- a/src/resolver.rs +++ b/src/resolver.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use crate::{ - ast::{Ast, Expr, Lval, Stmt}, + ast::{Ast, Expr, Function, Lval, Stmt}, span::{Span, Spanned}, }; @@ -80,7 +80,7 @@ impl Resolver { } self.end_scope(); } - Stmt::Function { name, params, body } => { + Stmt::Function(Function { name, params, body }) => { self.declare(name.clone()); self.define(name.clone()); -- 2.51.2