From 5072e03e03be6fe282083887b07448c5ead23130 Mon Sep 17 00:00:00 2001 From: Alex van de Sandt Date: Wed, 10 Jul 2024 18:22:27 -0400 Subject: [PATCH] Implement custom function execution --- src/ast/stmt.rs | 19 ++++--------------- src/callable.rs | 32 ++++++++++++++++++++++++++++++-- src/interpreter.rs | 22 +++++++++++++++++++--- src/value.rs | 4 ++++ 4 files changed, 57 insertions(+), 20 deletions(-) diff --git a/src/ast/stmt.rs b/src/ast/stmt.rs index db4aa62..35e108a 100644 --- a/src/ast/stmt.rs +++ b/src/ast/stmt.rs @@ -17,8 +17,8 @@ pub enum Stmt { }, Function { name: String, - _params: Vec, - _body: Vec>, + params: Vec, + body: Vec>, }, Block(Vec>), If { @@ -72,14 +72,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: params, - _body: body, - }, - span, - ) + Spanned::new(Self::Function { name, params, body }, span) } pub fn block( @@ -243,11 +236,7 @@ impl Stmt { pub fn unwrap_fun_decl(&self) -> (&String, &[String], &[Spanned]) { match self { - Self::Function { - name, - _params: params, - _body: body, - } => (name, params, body), + Self::Function { name, params, body } => (name, params, body), other => panic!("expected function decl, found {:?}", other.as_str()), } } diff --git a/src/callable.rs b/src/callable.rs index 159f6e2..84eb9d3 100644 --- a/src/callable.rs +++ b/src/callable.rs @@ -3,6 +3,9 @@ use std::{ time::SystemTime, }; +use crate::ast::Stmt; +use crate::interpreter::Interpreter; +use crate::span::Spanned; use crate::{environment::Env, interpreter::RuntimeError, value::Value}; type CallableOutput = Result; @@ -11,7 +14,32 @@ pub trait Callable: Debug { fn params(&self) -> Vec { Vec::new() } - fn call(&self, env: &Env) -> CallableOutput; + fn call(&self, interpreter: &mut Interpreter, env: &Env) -> CallableOutput; +} + +#[derive(Clone, Debug)] +pub struct LoxFunction { + params: Vec, + body: Vec>, +} + +impl LoxFunction { + pub fn new(params: Vec, body: Vec>) -> Self { + Self { params, body } + } +} + +impl Callable for LoxFunction { + fn params(&self) -> Vec { + self.params.clone() + } + + fn call(&self, interpreter: &mut Interpreter, env: &Env) -> CallableOutput { + for stmt in &self.body { + interpreter.execute_stmt(stmt, env)?; + } + Ok(Value::Nil) + } } pub struct NativeFunction { @@ -40,7 +68,7 @@ impl NativeFunction { } impl Callable for NativeFunction { - fn call(&self, env: &Env) -> Result { + fn call(&self, _: &mut Interpreter, env: &Env) -> Result { (self.implementation)(env) } } diff --git a/src/interpreter.rs b/src/interpreter.rs index 987198e..76a34e2 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -2,6 +2,7 @@ use std::io::Write; use miette::SourceSpan; +use crate::callable::LoxFunction; use crate::{ ast::{Ast, BinaryOp, Expr, LogicalOp, Lval, Stmt, UnaryOp}, callable::Callable, @@ -59,7 +60,7 @@ impl<'a> Interpreter<'a> { } #[tracing::instrument(name = "stmt", skip_all)] - fn execute_stmt(&mut self, stmt: &Spanned, env: &Env) -> Result<()> { + pub fn execute_stmt(&mut self, stmt: &Spanned, env: &Env) -> Result<()> { match stmt.as_ref() { Stmt::Expr(e) => { self.interpret_expr(e, env)?; @@ -76,7 +77,10 @@ impl<'a> Interpreter<'a> { .unwrap_or(Value::Nil); env.define(name.clone(), val); } - Stmt::Function { .. } => todo!("function delcaration execution"), + Stmt::Function { name, params, body } => { + let callable = LoxFunction::new(params.clone(), body.clone()); + env.define(name.clone(), Value::callable(callable)); + } Stmt::Block(stmts) => { let child_env = env.new_child(); for stmt in stmts { @@ -175,7 +179,7 @@ impl<'a> Interpreter<'a> { } // Call it - callable.call(&calling_env) + callable.call(self, &calling_env) } } } @@ -526,6 +530,18 @@ mod test { assert_eq!(out, "1"); } + #[test] + fn interprets_function_decl() { + let (e, out) = execute_stmts("fun foo() { print 1; } foo();"); + assert_matches!(e.get("foo"), Some(Value::Callable(_))); + assert_eq!(out, "1"); + + let (e, out) = + execute_stmts(r#"fun greet(name) { print "Hello, " + name; } greet("Ferris");"#); + assert_matches!(e.get("greet"), Some(Value::Callable(_))); + assert_eq!(out, "Hello, Ferris"); + } + #[test] fn interprets_logical_or() { for (input, expected_value) in [ diff --git a/src/value.rs b/src/value.rs index d5a0881..6f31b5e 100644 --- a/src/value.rs +++ b/src/value.rs @@ -13,6 +13,10 @@ pub enum Value { } impl Value { + pub fn callable(c: impl Callable + 'static) -> Self { + Self::Callable(Rc::new(c)) + } + pub fn is_truthy(&self) -> bool { match self { Value::Bool(b) => *b, -- 2.51.2