diff --git a/src/ast/mod.rs b/src/ast/mod.rs index e57cc74..2a8601c 100644 --- a/src/ast/mod.rs +++ b/src/ast/mod.rs @@ -16,7 +16,11 @@ impl Ast { .map(|s| match s.as_ref() { Stmt::Expr(e) => e.print_rpn(), Stmt::Print(e) => format!("print {}", e.print_rpn()), - _ => todo!("var decl rpn"), + Stmt::VarDecl { + name, + initializer: Some(e), + } => format!("set {name} {}", e.print_rpn()), + Stmt::VarDecl { name, .. } => format!("set {name}"), }) .collect::>() .join("\n") diff --git a/src/environment.rs b/src/environment.rs new file mode 100644 index 0000000..c8788db --- /dev/null +++ b/src/environment.rs @@ -0,0 +1,18 @@ +use crate::value::Value; +use std::collections::HashMap; + +pub struct Env(HashMap); + +impl Env { + pub fn new() -> Self { + Self(HashMap::new()) + } + + pub fn define(&mut self, var: String, val: Value) { + self.0.insert(var, val); + } + + pub fn get(&self, var: &str) -> Option { + self.0.get(var).cloned() + } +} diff --git a/src/interpreter.rs b/src/interpreter.rs index e1f7220..2304730 100644 --- a/src/interpreter.rs +++ b/src/interpreter.rs @@ -4,6 +4,7 @@ use miette::SourceSpan; use crate::{ ast::{Ast, BinaryOp, Expr, Stmt, UnaryOp}, + environment::Env, span::{Span, Spanned}, value::Value, }; @@ -19,13 +20,15 @@ pub fn interpret_expr(expr: &Spanned) -> Result { Interpreter::new().interpret_expr(expr) } -struct Interpreter { +pub struct Interpreter { + env: Env, output: Output, } impl Interpreter { - fn new() -> Self { + pub fn new() -> Self { Self { + env: Env::new(), output: std::io::stdout(), } } @@ -34,16 +37,19 @@ impl Interpreter { #[cfg(test)] impl Interpreter { fn with_output(output: Output) -> Self { - Self { output } + Self { + env: Env::new(), + output, + } } - fn into_output(self) -> Output { - self.output + fn decompose(self) -> (Env, Output) { + (self.env, self.output) } } impl Interpreter { - fn interpret(&mut self, ast: &Ast) -> Result<()> { + pub fn interpret(&mut self, ast: &Ast) -> Result<()> { for stmt in &ast.0 { self.execute_stmt(stmt)?; } @@ -61,7 +67,14 @@ impl Interpreter { let value = self.interpret_expr(e)?; writeln!(&mut self.output, "{value}")?; } - Stmt::VarDecl { .. } => todo!("var decl execution"), + Stmt::VarDecl { name, initializer } => { + let val = initializer + .as_ref() + .map(|e| self.interpret_expr(e)) + .transpose()? + .unwrap_or(Value::Nil); + self.env.define(name.clone(), val); + } } Ok(()) @@ -77,7 +90,10 @@ impl Interpreter { Ok(Value::from(lit.as_ref())) } Expr::Unary { op, expr } => self.interpret_unary(op, expr), - Expr::Var { .. } => todo!("var access"), + Expr::Var { name } => self + .env + .get(name) + .ok_or_else(|| RuntimeError::undefined_var(name)), } } @@ -217,6 +233,12 @@ impl Interpreter { #[derive(Debug, thiserror::Error, miette::Diagnostic)] pub enum RuntimeError { + #[error("Attempted to access undefined variable `{name}`")] + UndefinedVar { + name: String, + #[label("this variable is undefined")] + span: SourceSpan, + }, #[error("Only strings can be concatenated")] NonStringConcat { actual_type: &'static str, @@ -245,6 +267,13 @@ pub enum RuntimeError { } impl RuntimeError { + fn undefined_var(name: &Spanned) -> Self { + Self::UndefinedVar { + name: name.as_ref().to_string(), + span: name.span().into(), + } + } + fn non_string_concat(value: &Value, value_span: Span, op_span: Span) -> Self { Self::NonStringConcat { actual_type: value.type_str(), @@ -269,9 +298,9 @@ mod test { parser::{parse, parse_expr}, scanner::scan, }; - use claims::assert_matches; + use claims::{assert_matches, assert_some_eq}; - fn execute_stmts(input: &str) -> String { + fn execute_stmts(input: &str) -> (Env, String) { let tokens = scan(input).unwrap_or_else(|e| panic!("input `{input}` should scan. error: {e:?}")); let ast = @@ -279,7 +308,8 @@ mod test { let mut i = Interpreter::with_output(Vec::new()); i.interpret(&ast).unwrap(); - String::from_utf8(i.into_output()).unwrap() + let (env, output) = i.decompose(); + (env, String::from_utf8(output).unwrap()) } fn interpret_expr_to_value(input: &str) -> Value { @@ -301,18 +331,31 @@ mod test { #[test] fn executes_expr_stmt() { - assert_eq!(execute_stmts("1 + 1;"), ""); + assert_eq!(execute_stmts("1 + 1;").1, ""); } #[test] fn executes_print_stmt() { - assert_eq!(execute_stmts("print 1;"), "1\n"); - assert_eq!(execute_stmts("print 1 > 2;"), "false\n"); + assert_eq!(execute_stmts("print 1;").1, "1\n"); + assert_eq!(execute_stmts("print 1 > 2;").1, "false\n"); assert_eq!( - execute_stmts("print \"hello\"; print true;"), + execute_stmts("print \"hello\"; print true;").1, "hello\ntrue\n" ); } + #[test] + fn defines_var() { + let (e, _) = execute_stmts("var foo;"); + assert_some_eq!(e.get("foo"), Value::Nil); + + let (e, _) = execute_stmts("var foo = 1;"); + assert_some_eq!(e.get("foo"), Value::Number(1.0)); + + let (e, _) = execute_stmts("var foo = 1; var bar = foo;"); + assert_some_eq!(e.get("foo"), Value::Number(1.0)); + assert_some_eq!(e.get("bar"), Value::Number(1.0)); + } + #[test] fn interprets_arithmetic() { for (input, expected_value) in [ @@ -410,6 +453,14 @@ mod test { } } + #[test] + fn errs_on_undefined_var() { + assert_matches!( + interpret_expr_to_err("foo"), + RuntimeError::UndefinedVar { .. } + ); + } + #[test] fn errs_on_illegal_arithmetic() { for input in [ diff --git a/src/lib.rs b/src/lib.rs index e8c72f5..a259e89 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -1,4 +1,5 @@ mod ast; +mod environment; mod interpreter; mod logging; mod match_token; diff --git a/src/runners/repl.rs b/src/runners/repl.rs index ea5a6a9..f57ddba 100644 --- a/src/runners/repl.rs +++ b/src/runners/repl.rs @@ -2,7 +2,7 @@ use std::io::{BufRead, Write}; use miette::{IntoDiagnostic, Report, Result, WrapErr}; -use crate::interpreter::interpret; +use crate::interpreter::{interpret, Interpreter}; use crate::{parser::parse, scanner::scan}; #[tracing::instrument] @@ -19,6 +19,8 @@ pub fn run_repl() -> Result<()> { Ok(()) }; + let mut interpreter = Interpreter::new(); + let stdin = std::io::stdin(); prompt()?; @@ -60,7 +62,10 @@ pub fn run_repl() -> Result<()> { }; println!("{}", ast.print_rpn()); - interpret(&ast).map_err(|e| Report::new(e).with_source_code(line.clone()))?; + interpreter + .interpret(&ast) + .map_err(|e| Report::new(e).with_source_code(line.clone()))?; + prompt()?; }