commit 2a1da9dbbc9718c63924cae9604dbad3918c63cb
parent 0957828a80842921dffcc8c7be965514eb49f4b2
Author: Szymon Mikulicz <szymon.mikulicz@aptiv.com>
Date: Sat, 24 Aug 2024 19:08:12 +0200
Closures part 1
Diffstat:
9 files changed, 199 insertions(+), 32 deletions(-)
diff --git a/zlox/src/chunk.zig b/zlox/src/chunk.zig
@@ -27,6 +27,8 @@ pub const OP = enum(u8) {
SET_GLOBAL,
GET_LOCAL,
SET_LOCAL,
+ GET_UPVALUE,
+ SET_UPVALUE,
JUMP_IF_FALSE,
JUMP,
JUMP_POP,
diff --git a/zlox/src/compiler.zig b/zlox/src/compiler.zig
@@ -48,9 +48,17 @@ pub fn Compiler(size: comptime_int) type {
locals: [size]Local,
localCount: usize,
scopeDepth: usize,
+ enclosing: ?*Self,
+ upvalues: [upvalues_size]Upvalue,
const Self = @This();
+ pub const Upvalue = struct {
+ index: u8,
+ isLocal: bool,
+ };
+ const upvalues_size = std.math.maxInt(u8);
+
const Local = struct {
name: scanner.Token,
depth: ?usize,
@@ -335,27 +343,60 @@ pub fn Compiler(size: comptime_int) type {
}
fn namedVariable(self: *Self, tok: scanner.Token, canAssign: bool) void {
- var getOP = OP.GET_LOCAL;
- var setOP = OP.SET_LOCAL;
- const arg = self.resolveLocal(tok) catch blk: {
- getOP = OP.GET_GLOBAL;
- setOP = OP.SET_GLOBAL;
- break :blk self.identifierConstant(tok) catch return;
- };
+ const OPs: struct {get: OP, set: OP, arg: u8} = if (self.resolveLocal(tok)) |arg|
+ .{.get = OP.GET_LOCAL, .set = OP.SET_LOCAL, .arg = arg}
+ else if (self.resolveUpvalue(tok)) |arg|
+ .{.get = OP.GET_UPVALUE, .set = OP.SET_UPVALUE, .arg = arg}
+ else
+ .{.get = OP.GET_GLOBAL, .set = OP.SET_GLOBAL, .arg = self.identifierConstant(tok) catch return};
if (canAssign and self.match(Token.EQUAL)) {
- if(setOP == OP.SET_LOCAL and self.locals[arg].con) {
+ if(OPs.get == OP.GET_LOCAL and self.locals[OPs.arg].con) {
self.errorAtPrevious("Cannot assign to a constant");
return;
}
self.expression();
- self.emit(setOP, arg);
+ self.emit(OPs.set, OPs.arg);
} else {
- self.emit(getOP, arg);
+ self.emit(OPs.get, OPs.arg);
+ }
+ }
+
+ fn resolveUpvalue(self: *Self, name: scanner.Token) ?u8 {
+ if (self.enclosing) |enclosing| {
+ std.debug.print("Resolving '{s}' in '{s}'\n", .{name.lexeme, enclosing.currentFunction});
+ if (enclosing.resolveLocal(name)) |local| {
+ return self.addUpvalue(local, true) catch null;
+ } else if (enclosing.resolveUpvalue(name)) |upvalue| {
+ return self.addUpvalue(upvalue, false) catch null;
+ }
+ }
+ return null;
+ }
+
+ fn addUpvalue(self: *Self, idx: u8, isLocal: bool) !u8 {
+ const count = self.currentFunction.upvalue_count;
+
+ for (self.upvalues[0..count], 0..) |upvalue, i| {
+ if (upvalue.index == idx and upvalue.isLocal == isLocal) {
+ std.debug.print("Found existing at {d}\n", .{i});
+ return @intCast(i);
+ }
}
+
+ if (count == upvalues_size) {
+ self.errorAtPrevious("Too many upvalues");
+ self.lastError = error.OutOfMemory;
+ return self.lastError;
+ }
+
+ std.debug.print("New {s} upvalue in {s} idx {d} at {d} \n", .{if (isLocal) "local" else "upvalue", self.currentFunction, idx, count});
+ self.upvalues[count] = .{.index = idx, .isLocal = isLocal};
+ self.currentFunction.upvalue_count += 1;
+ return count;
}
- fn resolveLocal(self: *Self, name: scanner.Token) !u8 {
+ fn resolveLocal(self: *Self, name: scanner.Token) ?u8 {
var i = self.localCount;
while (i > 0) : (i -= 1) {
if (identifiersEql(self.locals[i-1].name, name)) {
@@ -366,7 +407,7 @@ pub fn Compiler(size: comptime_int) type {
}
}
}
- return error.NotFound;
+ return null;
}
fn emitConstant(self: *Self, val: Value) void {
@@ -493,9 +534,7 @@ pub fn Compiler(size: comptime_int) type {
return;
};
- var compiler = Self.init(self.scanner, self.objects, fun);
- compiler.current = self.current;
- compiler.beginScope();
+ var compiler = Self.init_enclosed(self, fun);
compiler.consume(Token.LEFT_PAREN, "Expect '(' after function name");
if (!compiler.check(Token.RIGHT_PAREN)) {
@@ -518,8 +557,15 @@ pub fn Compiler(size: comptime_int) type {
if (compiler.hadError) {
self.lastError = compiler.lastError;
+ } else if (compiler.currentFunction.upvalue_count == 0) {
+ self.emit(OP.CONSTANT, self.makeConstant(Value.init(compiler.end().cast())));
} else {
self.emit(OP.CLOSURE, self.makeConstant(Value.init(compiler.end().cast())));
+
+ for (compiler.upvalues[0..compiler.currentFunction.upvalue_count]) |upvalue| {
+ self.emitByte(if (upvalue.isLocal) 1 else 0);
+ self.emitByte(upvalue.index);
+ }
}
}
@@ -906,9 +952,19 @@ pub fn Compiler(size: comptime_int) type {
.locals = [_]Local{Local{.name = scanner.Token.Empty, .depth = 0, .con = true}} ** size,
.localCount = 1,
.scopeDepth = 0,
+ .enclosing = null,
+ .upvalues = [_]Upvalue{Upvalue{.index = 0, .isLocal = false}} ** upvalues_size,
};
}
+ fn init_enclosed(enclosing: *Self, fun: *Obj.Function) Self {
+ var enclosed = Self.init(enclosing.scanner, enclosing.objects, fun);
+ enclosed.current = enclosing.current;
+ enclosed.enclosing = enclosing;
+ enclosed.beginScope();
+ return enclosed;
+ }
+
pub fn compile(source: []const u8, objects: *GC) CompilerError!*Obj.Function {
var scan = try scanner.Scanner.init(source);
const fun = try objects.emplace(Obj.Type.Function, Obj.Function.Type.Script);
diff --git a/zlox/src/debug.zig b/zlox/src/debug.zig
@@ -1,27 +1,35 @@
const std = @import("std");
const chunk = @import("chunk.zig");
const value = @import("value.zig");
+const Obj = @import("obj.zig").Obj;
const print = std.debug.print;
-pub fn disassembleChunk(ch: *const chunk.Chunk, name: []const u8) !void {
- print("== {s} ==\n", .{name});
+const Error = error{OutOfMemory, KeyError, IllegalCastError, NotFound, IndexOutOfBounds};
+
+pub fn disassembleChunk(ch: *const chunk.Chunk, name: []const u8) Error!void {
+ print("/= {s} =\\\n", .{name});
var offset: usize = 0;
while (offset < ch.code.len) {
offset = try disassembleInstruction(ch, offset);
}
+ print("\\= {s} =/\n", .{name});
}
-pub fn disassembleInstruction(ch: *const chunk.Chunk, offset: usize) !usize {
- const OP = chunk.OP;
+pub fn print_offset(ch: *const chunk.Chunk, offset: usize) !void {
print("{d:0>4} ", .{offset});
if (offset > 0 and (try ch.lines.get(offset)) == (try ch.lines.get(offset - 1))) {
print(" | ", .{});
} else {
print("{d:4} ", .{try ch.lines.get(offset)});
}
+}
+pub fn disassembleInstruction(ch: *const chunk.Chunk, offset: usize) Error!usize {
+ try print_offset(ch, offset);
+
+ const OP = chunk.OP;
const op = try ch.code.get(offset);
const name = @tagName(@as(OP, @enumFromInt(op)));
@@ -48,13 +56,15 @@ pub fn disassembleInstruction(ch: *const chunk.Chunk, offset: usize) !usize {
@intFromEnum(OP.POP) => simpleInstruction(name, offset),
@intFromEnum(OP.GET_LOCAL) => try byteInstruction(name, ch, offset),
@intFromEnum(OP.SET_LOCAL) => try byteInstruction(name, ch, offset),
+ @intFromEnum(OP.GET_UPVALUE) => try byteInstruction(name, ch, offset),
+ @intFromEnum(OP.SET_UPVALUE) => try byteInstruction(name, ch, offset),
@intFromEnum(OP.JUMP_IF_FALSE) => try jumpInstruction(name, true, ch, offset),
@intFromEnum(OP.JUMP_POP) => simpleInstruction(name, offset),
@intFromEnum(OP.JUMP) => try jumpInstruction(name, true, ch, offset),
@intFromEnum(OP.LOOP) => try jumpInstruction(name, false, ch, offset),
@intFromEnum(OP.SET_INDEX) => simpleInstruction(name, offset),
@intFromEnum(OP.GET_INDEX) => simpleInstruction(name, offset),
- @intFromEnum(OP.CALL) => simpleInstruction(name, offset),
+ @intFromEnum(OP.CALL) => try byteInstruction(name, ch, offset),
@intFromEnum(OP.CLOSURE) => try closureInstruction(name, ch, offset),
else => blk: {
print("Unknown opcode {d} {s}\n", .{op, name});
@@ -68,13 +78,23 @@ fn simpleInstruction(name: []const u8, offset: usize) usize {
return offset + 1;
}
-fn constantInstruction(name: []const u8, ch: *const chunk.Chunk, offset: usize) !usize {
+fn constantInstruction(name: []const u8, ch: *const chunk.Chunk, offset: usize) Error!usize {
const constant = try ch.code.get(offset + 1);
- print("{s:<32} {d:4} '{s}'\n", .{ name, constant, try ch.constants.get(constant)});
+ const constval = try ch.constants.get(constant);
+ print("{s:<32} {d:4} '{s}'\n", .{ name, constant, constval});
+ if (constval.is(Obj.Type.Function)) {
+ const function = constval.obj.cast(.Function) catch unreachable;
+ if (function.name) |str| {
+ try disassembleChunk(function.chunk, str.slice());
+ } else {
+ try disassembleChunk(function.chunk, "<anon>");
+ }
+
+ }
return offset + 2;
}
-fn byteInstruction(name: []const u8, ch:*const chunk.Chunk, offset: usize) !usize {
+fn byteInstruction(name: []const u8, ch:*const chunk.Chunk, offset: usize) Error!usize {
print("{s:<32} {d:4}\n", .{name, try ch.code.get(offset+1)});
return offset + 2;
}
@@ -88,8 +108,24 @@ fn jumpInstruction(name: []const u8, sign: bool, ch: *const chunk.Chunk, offset:
return offset + 3;
}
-fn closureInstruction(name: []const u8, ch: *const chunk.Chunk, offset: usize) !usize {
- const constant = try ch.code.get(offset + 1);
- print("{s:<32} {d:4} '{s}'\n", .{ name, constant, try ch.constants.get(constant)});
- return offset + 2;
+fn closureInstruction(name: []const u8, ch: *const chunk.Chunk, offset: usize) Error!usize {
+ var off = offset + 1;
+ const constant = try ch.code.get(off);
+ const val = try ch.constants.get(constant);
+ const function = try val.obj.cast(.Function);
+ print("{s:<32} {d:4} '{s}'\n", .{ name, constant, function});
+ for (0..function.upvalue_count) |_| {
+ const isLocal = try ch.code.get(off + 1);
+ const idx = try ch.code.get(off + 2);
+ try print_offset(ch, off + 1);
+ print("{s:<38}|-> {s} {d}\n", .{ "", if (isLocal == 1) "local" else "upvalue", idx});
+ off += 2;
+ }
+ if (function.name) |str| {
+ try disassembleChunk(function.chunk, str.slice());
+ } else {
+ try disassembleChunk(function.chunk, "<anon>");
+ }
+
+ return off + 1;
}
diff --git a/zlox/src/main.zig b/zlox/src/main.zig
@@ -39,9 +39,7 @@ pub fn runFile(allocator: std.mem.Allocator, path: []const u8) anyerror!void {
const text = try file.reader().readAllAlloc(allocator, 999999);
defer allocator.free(text);
- VM.interpret(text, false) catch |err| {
- std.debug.print("Error: {}\n", .{err});
- };
+ try VM.interpret(text, true);
}
pub fn repl(allocator: std.mem.Allocator, dbg: bool) anyerror!void {
diff --git a/zlox/src/obj.zig b/zlox/src/obj.zig
@@ -14,6 +14,7 @@ pub const Obj = packed struct {
pub const Function = @import("obj/function.zig").Function;
pub const Native = @import("obj/native.zig").Native;
pub const Closure = @import("obj/closure.zig").Closure;
+ pub const Upvalue = @import("obj/upvalue.zig").Upvalue;
pub const Type = enum(u8) {
String,
@@ -22,6 +23,7 @@ pub const Obj = packed struct {
Native,
List,
Closure,
+ Upvalue,
pub fn get(comptime self: @This()) type {
return @field(Super, @tagName(self));
diff --git a/zlox/src/obj/closure.zig b/zlox/src/obj/closure.zig
@@ -11,6 +11,8 @@ pub const Closure = packed struct {
obj: Super,
function: *const Super.Function,
+ upvalues: [*]?*Super.Upvalue,
+ upvalues_len: u8,
pub fn init(arg: Arg, allocator: std.mem.Allocator) Error!*Self {
const self: *Self = try allocator.create(Self);
@@ -18,8 +20,12 @@ pub const Closure = packed struct {
.obj = Super{
.type = Super.Type.Closure,
},
+ .upvalues = (try allocator.alloc(?*Super.Upvalue, arg.upvalue_count)).ptr,
+ .upvalues_len = arg.upvalue_count,
.function = arg,
};
+ for(self.upvalues[0..self.upvalues_len])
+ |*upvalue| upvalue.* = null;
return self;
}
@@ -42,6 +48,7 @@ pub const Closure = packed struct {
}
pub fn free(self: *const Self, allocator: std.mem.Allocator) void {
+ allocator.free(self.upvalues[0..self.upvalues_len]);
allocator.destroy(self);
}
};
diff --git a/zlox/src/obj/function.zig b/zlox/src/obj/function.zig
@@ -20,6 +20,7 @@ pub const Function = packed struct {
chunk: *chunk.Chunk,
name: ?*const String,
type: Type,
+ upvalue_count: u8,
pub fn init(tp: Arg, allocator: std.mem.Allocator) Error!*Self {
const self: *Self = try allocator.create(Self);
@@ -31,6 +32,7 @@ pub const Function = packed struct {
.arity = 0,
.name = null,
.type = tp,
+ .upvalue_count = 0,
};
self.chunk.* = try chunk.Chunk.init(allocator);
return self;
diff --git a/zlox/src/obj/upvalue.zig b/zlox/src/obj/upvalue.zig
@@ -0,0 +1,43 @@
+const std = @import("std");
+
+const utils = @import("../comptime_utils.zig");
+const Value = @import("../value.zig").Value;
+const Super = @import("../obj.zig").Obj;
+const Error = Super.Error;
+
+pub const Upvalue = packed struct {
+ const Self = @This();
+
+ pub const Arg = *Value;
+
+ obj: Super,
+ location: *Value,
+
+ pub fn init(arg: Arg, allocator: std.mem.Allocator) Error!*Self {
+ const self: *Self = try allocator.create(Self);
+ self.* = Self{
+ .obj = Super{
+ .type = Super.Type.Upvalue,
+ },
+ .location = arg,
+ };
+ return self;
+ }
+
+ pub fn cast(self: anytype) utils.copy_const(@TypeOf(self), *Super) {
+ return @ptrCast(self);
+ }
+
+ pub fn format(_: *const Self, comptime _: []const u8, _: std.fmt.FormatOptions, writer: anytype) !void {
+ _ = try writer.writeAll("<upvalue>");
+ }
+
+ pub fn eql(_: *const Self, _: *const Self) bool {
+ return false;
+ }
+
+ pub fn free(self: *const Self, allocator: std.mem.Allocator) void {
+ allocator.destroy(self);
+ }
+};
+
diff --git a/zlox/src/vm.zig b/zlox/src/vm.zig
@@ -220,6 +220,10 @@ pub const VM = struct {
return;
}
+ fn captureUpvalue(self: *@This(), local: *Value) !*Obj.Upvalue {
+ return self.vm.objects.emplace(.Upvalue, local);
+ }
+
fn binary_op(self: *@This(), comptime in_tag: anytype, comptime out_tag: anytype, op: Callback.Type(in_tag, out_tag)) InterpreterError!void {
const b = self.pop();
const a = self.pop();
@@ -317,13 +321,21 @@ pub const VM = struct {
return InterpreterError.RuntimeError;
}
},
+ @intFromEnum(OP.GET_UPVALUE) => {
+ const closure = try self.frame().callee.cast(.Closure);
+ self.push(closure.upvalues[self.read_byte()].?.location.*);
+ },
+ @intFromEnum(OP.SET_UPVALUE) => {
+ const closure = try self.frame().callee.cast(.Closure);
+ closure.upvalues[self.read_byte()].?.location.* = self.peek(0);
+ },
@intFromEnum(OP.GET_INDEX) => {
const key = self.pop();
const obj = self.pop();
var pushed = false;
if (obj.is(Value.obj)) {
switch(obj.obj.type) {
- .Function, .Native, .Closure => {},
+ .Function, .Native, .Closure, .Upvalue => {},
inline else => |tp| {
self.push((obj.obj.cast(tp) catch unreachable).get(key) catch Value.init({}));
pushed = true;
@@ -375,7 +387,16 @@ pub const VM = struct {
},
@intFromEnum(OP.CLOSURE) => {
const function = try self.read_constant().obj.cast(.Function);
- self.push(Value.init(try self.vm.objects.emplace_cast(.Closure, function)));
+ const closure = try self.vm.objects.emplace(.Closure, function);
+ for(closure.upvalues[0..closure.upvalues_len]) |*upvalue| {
+ if(self.read_byte() == 1) {
+ upvalue.* = try self.captureUpvalue(&self.frame().slots[self.read_byte()]);
+ } else {
+ const callee = try self.frame().callee.cast(.Closure);
+ upvalue.* = callee.upvalues[self.read_byte()];
+ }
+ }
+ self.push(Value.init(closure.cast()));
},
@intFromEnum(OP.DEFINE_GLOBAL) => _ = try self.vm.globals.set(self.read_string(), Global.make_var(self.pop())),
@intFromEnum(OP.DEFINE_GLOBAL_CONSTANT) => _ = try self.vm.globals.set(self.read_string(), Global.make_con(self.pop())),