commit 7f9410d7da372919957061986600b2e10624898c
parent 81e72bdef0fdc23b2151b2773a1564f9eb2f5d6a
Author: Szymon Mikulicz <szymon.mikulicz@posteo.net>
Date: Sun, 9 Oct 2022 00:42:37 +0200
Add resolving variables
Diffstat:
5 files changed, 230 insertions(+), 33 deletions(-)
diff --git a/source/app.d b/source/app.d
@@ -12,6 +12,7 @@ import ast;
import parser;
import error;
import interpreter;
+import resolver;
int main(string[] args) {
Lox lox = new Lox;
@@ -32,10 +33,12 @@ class Lox {
static bool hadError = false;
static bool hadRuntimeError = false;
+ Resolver resolver;
Interpreter interpreter;
this() {
interpreter = new Interpreter();
+ resolver = new Resolver(interpreter);
}
extern(C) static void completion(const char *buf, linenoiseCompletions *lc) {
@@ -83,6 +86,9 @@ class Lox {
auto expression = parser.parse();
if(hadError) return;
+ resolver.resolve(expression);
+ if(hadError) return;
+
auto result = interpreter.interpret(expression);
if(hadError) return;
diff --git a/source/environment.d b/source/environment.d
@@ -28,6 +28,15 @@ class Environment {
return value;
}
+ Environment ancestor(size_t distance) {
+ Environment environment = this;
+ foreach(_; 0..distance) {
+ environment = environment.enclosing;
+ }
+
+ return environment;
+ }
+
Variant get(TokenI name) {
return *find(name);
}
@@ -35,4 +44,13 @@ class Environment {
void assign(TokenI name, Variant new_value) {
*find(name) = new_value;
}
+
+ Variant getAt(TokenI name, size_t distance) {
+ return ancestor(distance).values[name.lexeme];
+ }
+
+ void assignAt(TokenI name, Variant value, size_t distance) {
+ ancestor(distance).values[name.lexeme] = value;
+ }
+
}
diff --git a/source/interpreter.d b/source/interpreter.d
@@ -15,9 +15,10 @@ import callable;
import fun;
class Interpreter : StmtVisitor, ExprVisitor {
- Variant value;
- Environment environment;
- Environment globals;
+ private Variant value;
+ private Environment environment;
+ private Environment globals;
+ private size_t[Expr] locals;
class BreakCalled : Exception {
this() {
@@ -36,6 +37,7 @@ class Interpreter : StmtVisitor, ExprVisitor {
this() {
globals = new Environment();
environment = globals;
+ locals = null;
globals.define("clock", Variant(new class Callable {
ulong arity() {
@@ -47,7 +49,7 @@ class Interpreter : StmtVisitor, ExprVisitor {
}));
}
- string interpret(Array!Stmt statements) {
+ string interpret(Stmt[] statements) {
try {
foreach(statement; statements) {
execute(statement);
@@ -66,6 +68,10 @@ class Interpreter : StmtVisitor, ExprVisitor {
return str;
}
+ void resolve(Expr expr, size_t depth) {
+ locals[expr] = depth;
+ }
+
private void execute(Stmt stmt) {
stmt.accept(this);
}
@@ -156,15 +162,24 @@ class Interpreter : StmtVisitor, ExprVisitor {
void visit(Assign expr) {
Variant variant = evaluate(expr.value);
- environment.assign(expr.name, variant);
+ if(size_t* distance = expr in locals) {
+ environment.assignAt(expr.name, value, *distance);
+ } else {
+ globals.assign(expr.name, variant);
+ }
value = variant;
}
void visit(Variable expr) {
- Variant var = environment.get(expr.name);
- if (!var.hasValue())
- throw new RuntimeError(expr.name, "Variable is not initialized.");
- value = var;
+ value = lookUpVariable(expr.name, expr);
+ }
+
+ private Variant lookUpVariable(TokenI name, Expr expr) {
+ if(size_t* distance = expr in locals) {
+ return environment.getAt(name, *distance);
+ } else {
+ return globals.get(name);
+ }
}
void visit(Logical expr) {
diff --git a/source/parser.d b/source/parser.d
@@ -61,17 +61,15 @@ class Parser {
private Array!TokenI tokens;
private int current = 0;
- private int breakAllowed = 0;
- private int returnAllowed = 0;
this(Array!TokenI tokens) {
this.tokens = tokens;
}
- Array!Stmt parse() {
- Array!Stmt statements = Array!Stmt();
+ Stmt[] parse() {
+ Stmt[] statements;
while (!isAtEnd()) {
- statements.insert(declaration());
+ statements ~= declaration();
}
return statements;
}
@@ -106,9 +104,7 @@ class Parser {
}
consume(TokenType.RIGHT_PAREN, "Expect ')' after parameters.");
consume(TokenType.LEFT_BRACE, "Expect '{' before function body.");
- returnAllowed++;
Stmt[] bod = block();
- returnAllowed--;
return new Function(parameters, bod);
}
@@ -138,24 +134,16 @@ class Parser {
}
private Stmt breakStatement() {
- if (breakAllowed) {
- return statement!(Break)(previous());
- } else {
- throw error(previous(), "Break not allowed here.");
- }
+ return statement!(Break)(previous());
}
private Stmt returnStatement() {
TokenI keyword = previous();
- if (returnAllowed) {
- Expr value = null;
- if (!check(TokenType.SEMICOLON) && !check(TokenType.RIGHT_BRACE)) {
- value = expression();
- }
- return statement!(Return)(keyword, value);
- } else {
- throw error(keyword, "Return not allowed here.");
+ Expr value = null;
+ if (!isAtEnd() && !check(TokenType.SEMICOLON) && !check(TokenType.RIGHT_BRACE)) {
+ value = expression();
}
+ return statement!(Return)(keyword, value);
}
private Stmt forStatement() {
@@ -183,9 +171,7 @@ class Parser {
}
consume(TokenType.RIGHT_PAREN, "Expect ')' after for clauses.");
- breakAllowed++;
Stmt bod = matchStatement();
- breakAllowed--;
if (increment !is null) {
bod = new Block([bod, new Expression(increment)]);
@@ -221,9 +207,7 @@ class Parser {
consume(TokenType.LEFT_PAREN, "Expect '(' after 'while'.");
Expr condition = expression();
consume(TokenType.RIGHT_PAREN, "Expect ')' after condition.");
- breakAllowed++;
Stmt bod = matchStatement();
- breakAllowed--;
return statement!(While)(condition, bod);
}
diff --git a/source/resolver.d b/source/resolver.d
@@ -0,0 +1,174 @@
+import interpreter;
+import ast;
+import std.container;
+import std.range;
+import token;
+import app;
+
+class Resolver : StmtVisitor, ExprVisitor {
+ private Interpreter interpreter;
+ private SList!(bool[string]) scopes;
+ private FunctionType currentFunction = FunctionType.NONE;
+ private LoopType currentLoop = LoopType.NONE;
+
+ private enum FunctionType {
+ NONE,
+ FUN
+ }
+
+ private enum LoopType {
+ NONE,
+ WHILE
+ }
+
+ this(Interpreter interpreter) {
+ this.interpreter = interpreter;
+ this.scopes = SList!(bool[string])();
+ }
+
+ void resolve(T)(T[] statements...) {
+ foreach(statement; statements) {
+ statement.accept(this);
+ }
+ }
+
+ void beginScope() {
+ scopes.insertFront(null);
+ }
+
+ private void endScope() {
+ scopes.removeFront();
+ }
+
+ private void declare(TokenI name) {
+ if(scopes.empty()) return;
+
+ auto sco = scopes.front();
+ if(name.lexeme in sco) {
+ Lox.error(name, "Already a variable with this name in this scope.");
+ }
+ sco[name.lexeme] = false;
+ }
+
+ private void define(TokenI name) {
+ if(scopes.empty()) return;
+
+ scopes.front()[name.lexeme] = true;
+ }
+
+ private void resolveLocal(Expr expr, TokenI name) {
+ foreach(i, sco; scopes[].enumerate()) {
+ if(name.lexeme in sco) {
+ interpreter.resolve(expr, i);
+ }
+ }
+ }
+
+ void visit(Print _print) {
+ resolve(_print.expression);
+ }
+
+ void visit(Expression _expression) {
+ resolve(_expression.expression);
+ }
+
+ void visit(Var stmt) {
+ declare(stmt.name);
+ if (stmt.initializer !is null) {
+ resolve(stmt.initializer);
+ }
+ define(stmt.name);
+ }
+
+ void visit(Block stmt) {
+ beginScope();
+ resolve(stmt.statements);
+ endScope();
+ }
+
+ void visit(If _if) {
+ resolve(_if.condition);
+ resolve(_if.thenBranch);
+ if (_if.elseBranch !is null) resolve(_if.elseBranch);
+ }
+
+ void visit(While _while) {
+ resolve(_while.condition);
+ LoopType enclosingLoop = currentLoop;
+ currentLoop = LoopType.WHILE;
+ resolve(_while.bod);
+ currentLoop = enclosingLoop;
+ }
+
+ void visit(Break _break) {
+ if (currentLoop == LoopType.NONE) {
+ Lox.error(_break.keyword, "Can't break outside a loop.");
+ }
+ }
+
+ void visit(Return _return) {
+ if (currentFunction == FunctionType.NONE) {
+ Lox.error(_return.keyword, "Can't return from top-level code.");
+ }
+ if (_return.value !is null) resolve(_return.value);
+ }
+
+ void visit(Ternary _ternary) {
+ resolve(_ternary.left);
+ resolve(_ternary.middle);
+ resolve(_ternary.right);
+ }
+
+ void visit(Binary _binary) {
+ resolve(_binary.left);
+ resolve(_binary.right);
+ }
+
+ void visit(Grouping _grouping) {
+ resolve(_grouping.expression);
+ }
+
+ void visit(Literal _) {}
+
+ void visit(Unary _unary) {
+ resolve(_unary.right);
+ }
+
+ void visit(Variable expr) {
+ if (!scopes.empty() && !scopes.front().get(expr.name.lexeme, true)) {
+ Lox.error(expr.name, "Can't read local variable in its own initializer.");
+ }
+
+ resolveLocal(expr, expr.name);
+ }
+ void visit(Assign expr) {
+ resolve(expr.value);
+ resolveLocal(expr, expr.name);
+ }
+
+ void visit(Logical _logical) {
+ resolve(_logical.left);
+ resolve(_logical.right);
+ }
+
+ void visit(Call _call) {
+ resolve(_call.callee);
+
+ foreach(arg; _call.arguments) {
+ resolve(arg);
+ }
+ }
+
+ void visit(Function expr) {
+ FunctionType enclosingFunction = currentFunction;
+ currentFunction = FunctionType.FUN;
+ beginScope();
+ foreach (param; expr.params) {
+ declare(param);
+ define(param);
+ }
+ resolve(expr.body);
+ endScope();
+ currentFunction = enclosingFunction;
+ }
+}