diff --git a/src/callable.rs b/src/callable.rs index 84eb9d3..51b6414 100644 --- a/src/callable.rs +++ b/src/callable.rs @@ -36,7 +36,9 @@ impl Callable for LoxFunction { fn call(&self, interpreter: &mut Interpreter, env: &Env) -> CallableOutput { for stmt in &self.body { - interpreter.execute_stmt(stmt, env)?; + if let Some(return_value) = interpreter.execute_stmt(stmt, env)? { + return Ok(return_value); + } } Ok(Value::Nil) } diff --git a/src/interpreter.rs b/src/interpreter.rs index dc2e716..5c81ec9 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -60,7 +60,7 @@ impl<'a> Interpreter<'a> { } #[tracing::instrument(name = "stmt", skip_all)] - pub 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)?; @@ -84,7 +84,9 @@ impl<'a> Interpreter<'a> { Stmt::Block(stmts) => { let child_env = env.new_child(); for stmt in stmts { - self.execute_stmt(stmt, &child_env)?; + if let Some(return_value) = self.execute_stmt(stmt, &child_env)? { + return Ok(Some(return_value)); + } } } Stmt::If { @@ -94,22 +96,35 @@ impl<'a> Interpreter<'a> { } => { let condition_val = self.interpret_expr(condition, env)?; if condition_val.is_truthy() { - self.execute_stmt(then, env)?; + if let Some(return_value) = self.execute_stmt(then, env)? { + return Ok(Some(return_value)); + } } else if let Some(otherwise) = otherwise { - self.execute_stmt(otherwise, env)?; + if let Some(return_value) = self.execute_stmt(otherwise, env)? { + return Ok(Some(return_value)); + } } } Stmt::While { condition, body } => { let mut condition_val = self.interpret_expr(condition, env)?; while condition_val.is_truthy() { - self.execute_stmt(body, env)?; + if let Some(return_value) = self.execute_stmt(body, env)? { + return Ok(Some(return_value)); + } condition_val = self.interpret_expr(condition, env)?; } } - Stmt::Return(_) => todo!("return statement execution"), + Stmt::Return(return_expr) => { + let value = return_expr + .as_ref() + .map(|e| self.interpret_expr(e, env)) + .transpose()? + .unwrap_or(Value::Nil); + return Ok(Some(value)); + } } - Ok(()) + Ok(None) } #[tracing::instrument(name = "expr", skip_all)] @@ -543,6 +558,12 @@ mod test { assert_eq!(out, "Hello, Ferris"); } + #[test] + fn returns_early() { + let (_, out) = execute_stmts(r#"fun foo() { return 1; print "bad"; }"#); + assert!(out.is_empty()); + } + #[test] fn interprets_logical_or() { for (input, expected_value) in [