diff --git a/src/callable/function.rs b/src/callable/function.rs index 73c1760..cae0533 100644 --- a/src/callable/function.rs +++ b/src/callable/function.rs @@ -1,3 +1,5 @@ +use std::fmt::{Debug, Formatter}; + use crate::{ ast::Stmt, callable::Callable, @@ -7,29 +9,53 @@ use crate::{ value::Value, }; -#[derive(Clone, Debug)] +#[derive(Clone)] pub struct LoxFunction { params: Vec, + closure: Env, body: Vec>, } impl LoxFunction { - pub fn new(params: Vec, body: Vec>) -> Self { - Self { params, body } + pub fn new(params: Vec, closure: Env, body: Vec>) -> Self { + Self { + params, + closure, + body, + } } } impl Callable for LoxFunction { - fn params(&self) -> Vec { - self.params.clone() + fn arity(&self) -> usize { + self.params.len() } - fn call(&self, interpreter: &mut Interpreter, env: &Env) -> Result { + fn call(&self, interpreter: &mut Interpreter, args: Vec) -> Result { + // Define our arguments + let env = self.closure.new_child(); + + // Define a new env for the call + for (param, arg) in self.params.clone().into_iter().zip(args) { + env.define(param, arg) + } + + // Execute the body for stmt in &self.body { - if let Some(return_value) = interpreter.execute_stmt(stmt, env)? { + if let Some(return_value) = interpreter.execute_stmt(stmt, &env)? { return Ok(return_value); } } + Ok(Value::Nil) } } + +impl Debug for LoxFunction { + /// Including `closure` in output would cause infinite recursion + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.debug_struct("LoxFunction") + .field("params", &self.params) + .finish_non_exhaustive() + } +} diff --git a/src/callable/mod.rs b/src/callable/mod.rs index 14a1de8..a46854c 100644 --- a/src/callable/mod.rs +++ b/src/callable/mod.rs @@ -3,14 +3,12 @@ mod native; use std::fmt::Debug; -use crate::{environment::Env, interpreter::Interpreter, interpreter::RuntimeError, value::Value}; +use crate::{interpreter::Interpreter, interpreter::RuntimeError, value::Value}; pub use function::LoxFunction; pub use native::NativeFunction; pub trait Callable: Debug { - fn params(&self) -> Vec { - Vec::new() - } - fn call(&self, interpreter: &mut Interpreter, env: &Env) -> Result; + fn arity(&self) -> usize; + fn call(&self, interpreter: &mut Interpreter, args: Vec) -> Result; } diff --git a/src/callable/native.rs b/src/callable/native.rs index d4c2a56..2bec998 100644 --- a/src/callable/native.rs +++ b/src/callable/native.rs @@ -36,8 +36,11 @@ impl NativeFunction { } impl Callable for NativeFunction { - fn call(&self, _: &mut Interpreter, env: &Env) -> Result { - (self.implementation)(env) + fn arity(&self) -> usize { + 0 + } + fn call(&self, _: &mut Interpreter, _: Vec) -> Result { + (self.implementation)(&Env::empty()) } } diff --git a/src/environment.rs b/src/environment.rs index 4b03140..db5bc4e 100644 --- a/src/environment.rs +++ b/src/environment.rs @@ -1,13 +1,18 @@ -use std::{cell::RefCell, collections::hash_map::HashMap, rc::Rc}; +use std::{ + cell::RefCell, + collections::hash_map::HashMap, + fmt::{Debug, Formatter}, + rc::Rc, +}; use crate::{callable::NativeFunction, value::Value}; -#[derive(Debug, Clone)] +#[derive(Clone)] pub struct Env(Rc>); impl Default for Env { fn default() -> Self { - let env = Self(Rc::new(RefCell::new(Inner::empty()))); + let env = Self::empty(); let clock = Value::Callable(Rc::new(NativeFunction::clock())); env.define("clock".to_string(), clock); env @@ -15,6 +20,10 @@ impl Default for Env { } impl Env { + pub fn empty() -> Self { + Self(Rc::new(RefCell::new(Inner::empty()))) + } + pub fn new_child(&self) -> Self { let parent = self.clone(); let inner = Inner::with_parent(parent); @@ -87,6 +96,20 @@ impl Inner { } } +impl Debug for Env { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + let inner = self.0.borrow(); + let mut d = f.debug_map(); + + d.entries(inner.values.iter()); + if let Some(parent) = &inner.parent { + d.entry(&"parent", parent); + } + + d.finish() + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/src/interpreter.rs b/src/interpreter.rs index 5c81ec9..b85980c 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -78,7 +78,7 @@ impl<'a> Interpreter<'a> { env.define(name.clone(), val); } Stmt::Function { name, params, body } => { - let callable = LoxFunction::new(params.clone(), body.clone()); + let callable = LoxFunction::new(params.clone(), env.clone(), body.clone()); env.define(name.clone(), Value::callable(callable)); } Stmt::Block(stmts) => { @@ -173,8 +173,7 @@ impl<'a> Interpreter<'a> { }; // Check the arity - let params = callable.params(); - if params.len() != arguments.len() { + if callable.arity() != arguments.len() { return Err(RuntimeError::incorrect_arity( callable.as_ref(), arguments.len(), @@ -188,14 +187,8 @@ impl<'a> Interpreter<'a> { .map(|arg| self.interpret_expr(arg, env)) .collect::, _>>()?; - // Define a new env for the call - let calling_env = env.new_child(); - for (param, arg) in params.into_iter().zip(evaluated_args) { - calling_env.define(param, arg) - } - // Call it - callable.call(self, &calling_env) + callable.call(self, evaluated_args) } } } @@ -421,7 +414,7 @@ impl RuntimeError { fn incorrect_arity(callable: &dyn Callable, args_len: usize, callee_span: Span) -> Self { Self::IncorrectArity { - expected: callable.params().len(), + expected: callable.arity(), actual: args_len, span: callee_span.into(), } @@ -564,6 +557,26 @@ mod test { assert!(out.is_empty()); } + #[test] + fn closures_close() { + let (_, out) = execute_stmts( + r#" + fun makeCounter() { + var i = 0; + fun count() { + i = i + 1; + print i; + } + return count; + } + var counter = makeCounter(); + counter(); + counter(); + "#, + ); + assert_eq!(out, "1\n2"); + } + #[test] fn interprets_logical_or() { for (input, expected_value) in [