commit 0957828a80842921dffcc8c7be965514eb49f4b2
parent 1c61c9d487badce74c4ab873ccdc170b774411ab
Author: Szymon Mikulicz <szymon.mikulicz@aptiv.com>
Date: Fri, 23 Aug 2024 10:19:46 +0200
Start closures
Diffstat:
12 files changed, 129 insertions(+), 58 deletions(-)
diff --git a/zlox/src/chunk.zig b/zlox/src/chunk.zig
@@ -34,6 +34,7 @@ pub const OP = enum(u8) {
SET_INDEX,
GET_INDEX,
CALL,
+ CLOSURE,
};
pub const Chunk = struct {
diff --git a/zlox/src/compiler.zig b/zlox/src/compiler.zig
@@ -79,23 +79,14 @@ pub fn Compiler(size: comptime_int) type {
v.* = switch (tok) {
// zig fmt: off
T.LEFT_PAREN => R(S.grouping, S.call, P.CALL ),
- T.RIGHT_PAREN => R(null, null, P.NONE ),
- T.LEFT_BRACE => R(null, null, P.NONE ),
- T.RIGHT_BRACE => R(null, null, P.NONE ),
T.LEFT_BRACKET => R(S.listTable,S.index, P.CALL ),
- T.RIGHT_BRACKET => R(null, null, P.NONE ),
- T.COMMA => R(null, null, P.NONE ),
- T.DOT => R(null, null, P.NONE ),
T.MINUS => R(S.unary, S.binary, P.TERM ),
T.PLUS => R(null, S.binary, P.TERM ),
- T.COLON => R(null, null, P.NONE ),
- T.SEMICOLON => R(null, null, P.NONE ),
T.SLASH => R(null, S.binary, P.FACTOR ),
T.STAR => R(null, S.binary, P.FACTOR ),
T.QUESTION => R(null, S.ternary, P.TERNARY ),
T.BANG => R(S.unary, null, P.NONE ),
T.BANG_EQUAL => R(null, S.binary, P.EQUALITY ),
- T.EQUAL => R(null, null, P.NONE ),
T.EQUAL_EQUAL => R(null, S.binary, P.EQUALITY ),
T.GREATER => R(null, S.binary, P.COMPARISON ),
T.GREATER_EQUAL => R(null, S.binary, P.COMPARISON ),
@@ -106,26 +97,11 @@ pub fn Compiler(size: comptime_int) type {
T.CHAR => R(S.char, null, P.NONE ),
T.NUMBER => R(S.number, null, P.NONE ),
T.AND => R(null, S._and, P.AND ),
- T.CLASS => R(null, null, P.NONE ),
- T.ELSE => R(null, null, P.NONE ),
T.FALSE => R(S.literal, null, P.NONE ),
- T.FOR => R(null, null, P.NONE ),
- T.FUN => R(null, null, P.NONE ),
- T.IF => R(null, null, P.NONE ),
T.NIL => R(S.literal, null, P.NONE ),
T.OR => R(null, S._or, P.OR ),
- T.PRINT => R(null, null, P.NONE ),
- T.RETURN => R(null, null, P.NONE ),
- T.SUPER => R(null, null, P.NONE ),
- T.THIS => R(null, null, P.NONE ),
T.TRUE => R(S.literal, null, P.NONE ),
- T.VAR => R(null, null, P.NONE ),
- T.CON => R(null, null, P.NONE ),
- T.WHILE => R(null, null, P.NONE ),
- T.SWITCH => R(null, null, P.NONE ),
- T.CASE => R(null, null, P.NONE ),
- T.DEFAULT => R(null, null, P.NONE ),
- T.EOF => R(null, null, P.NONE ),
+ else => R(null, null, P.NONE ),
// zig fmt: on
};
}
@@ -543,7 +519,7 @@ pub fn Compiler(size: comptime_int) type {
if (compiler.hadError) {
self.lastError = compiler.lastError;
} else {
- self.emit(OP.CONSTANT, self.makeConstant(Value.init(compiler.end().cast())));
+ self.emit(OP.CLOSURE, self.makeConstant(Value.init(compiler.end().cast())));
}
}
diff --git a/zlox/src/debug.zig b/zlox/src/debug.zig
@@ -55,6 +55,7 @@ pub fn disassembleInstruction(ch: *const chunk.Chunk, offset: usize) !usize {
@intFromEnum(OP.SET_INDEX) => simpleInstruction(name, offset),
@intFromEnum(OP.GET_INDEX) => simpleInstruction(name, offset),
@intFromEnum(OP.CALL) => simpleInstruction(name, offset),
+ @intFromEnum(OP.CLOSURE) => try closureInstruction(name, ch, offset),
else => blk: {
print("Unknown opcode {d} {s}\n", .{op, name});
break :blk offset + 1;
@@ -86,3 +87,9 @@ fn jumpInstruction(name: []const u8, sign: bool, ch: *const chunk.Chunk, offset:
print("{s:<32} {d:4} -> {d}\n", .{name, offset, if (sign) offset + 3 + jump else offset + 3 - jump});
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;
+}
diff --git a/zlox/src/obj.zig b/zlox/src/obj.zig
@@ -13,6 +13,7 @@ pub const Obj = packed struct {
pub const Table = @import("obj/table.zig").Table;
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 Type = enum(u8) {
String,
@@ -20,6 +21,7 @@ pub const Obj = packed struct {
Function,
Native,
List,
+ Closure,
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
@@ -0,0 +1,47 @@
+const std = @import("std");
+const utils = @import("../comptime_utils.zig");
+
+const Super = @import("../obj.zig").Obj;
+const Error = Super.Error;
+
+pub const Closure = packed struct {
+ const Self = @This();
+
+ pub const Arg = *const Super.Function;
+
+ obj: Super,
+ function: *const Super.Function,
+
+ 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.Closure,
+ },
+ .function = arg,
+ };
+ return self;
+ }
+
+ pub fn cast(self: anytype) utils.copy_const(@TypeOf(self), *Super) {
+ return @ptrCast(self);
+ }
+
+ pub fn format(self: *const Self, comptime _: []const u8, _: std.fmt.FormatOptions, writer: anytype) !void {
+ _ = try writer.write("<C: ");
+ if (self.function.name) |name| {
+ _ = try writer.write(name.slice());
+ } else {
+ _ = try writer.write("-");
+ }
+ _ = try writer.writeAll(">");
+ }
+
+ 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/obj/function.zig b/zlox/src/obj/function.zig
@@ -1,5 +1,6 @@
const std = @import("std");
const chunk = @import("../chunk.zig");
+const utils = @import("../comptime_utils.zig");
const Super = @import("../obj.zig").Obj;
const Error = Super.Error;
@@ -35,7 +36,7 @@ pub const Function = packed struct {
return self;
}
- pub fn cast(self: *Self) *Super {
+ pub fn cast(self: anytype) utils.copy_const(@TypeOf(self), *Super) {
return @ptrCast(self);
}
diff --git a/zlox/src/obj/list.zig b/zlox/src/obj/list.zig
@@ -35,7 +35,7 @@ pub const List = packed struct {
return self;
}
- pub fn cast(self: *Self) *Super {
+ pub fn cast(self: anytype) utils.copy_const(@TypeOf(self), *Super) {
return @ptrCast(self);
}
diff --git a/zlox/src/obj/native.zig b/zlox/src/obj/native.zig
@@ -1,5 +1,6 @@
const std = @import("std");
+const utils = @import("../comptime_utils.zig");
const GC = @import("../gc.zig").GC;
const Value = @import("../value.zig").Value;
const Super = @import("../obj.zig").Obj;
@@ -47,7 +48,7 @@ pub const Native = packed struct {
return self.fun(gc, args[0..argCount]);
}
- pub fn cast(self: *Self) *Super {
+ pub fn cast(self: anytype) utils.copy_const(@TypeOf(self), *Super) {
return @ptrCast(self);
}
diff --git a/zlox/src/obj/string.zig b/zlox/src/obj/string.zig
@@ -38,9 +38,11 @@ pub const String = packed struct {
pub fn slice(self: *const Self) []const u8 {
return self.data()[0..self.len];
}
- pub fn cast(self: *Self) *Super {
+
+ pub fn cast(self: anytype) utils.copy_const(@TypeOf(self), *Super) {
return @ptrCast(self);
}
+
pub fn format(self: *const Self, comptime _: []const u8, _: std.fmt.FormatOptions, writer: anytype) !void {
_ = try writer.writeAll(self.slice());
}
diff --git a/zlox/src/obj/table.zig b/zlox/src/obj/table.zig
@@ -28,7 +28,8 @@ pub const Table = packed struct {
self.table.* = Self.Table.init(allocator);
return self;
}
- pub fn cast(self: *Self) *Super {
+
+ pub fn cast(self: anytype) utils.copy_const(@TypeOf(self), *Super) {
return @ptrCast(self);
}
diff --git a/zlox/src/obj/template.zig b/zlox/src/obj/template.zig
@@ -1,5 +1,6 @@
const std = @import("std");
+const utils = @import("../comptime_utils.zig");
const Super = @import("../obj.zig").Obj;
const Error = Super.Error;
@@ -15,7 +16,7 @@ pub const Template = packed struct {
return error.OutOfMemory;
}
- pub fn cast(self: *Self) *Super {
+ pub fn cast(self: anytype) utils.copy_const(@TypeOf(self), *Super) {
return @ptrCast(self);
}
diff --git a/zlox/src/vm.zig b/zlox/src/vm.zig
@@ -47,15 +47,26 @@ pub const VM = struct {
const Globals = table.Table(*const Obj.String, Global, hash.hash_t(*const Obj.String), Obj.String.eql);
const CallFrame = struct {
- function: *const Obj.Function,
+ callee: *const Obj,
ip: [*]const u8,
slots: [*]Value,
-
- pub fn init(function: *const Obj.Function, slots: [*]Value) @This() {
- return @This() {
- .function = function,
- .ip = function.chunk.code.data.ptr,
- .slots = slots
+ chunk: *const Chunk,
+
+ pub fn init(comptime tp: Obj.Type, callee: *const tp.get(), slots: [*]Value) @This() {
+ return switch(tp) {
+ .Function => @This() {
+ .callee = callee.cast(),
+ .ip = callee.chunk.code.data.ptr,
+ .chunk = callee.chunk,
+ .slots = slots
+ },
+ .Closure => @This() {
+ .callee = callee.cast(),
+ .ip = callee.function.chunk.code.data.ptr,
+ .chunk = callee.function.chunk,
+ .slots = slots
+ },
+ else => @compileError("Invalid type")
};
}
};
@@ -107,7 +118,7 @@ pub const VM = struct {
.vm = vm
};
self.stackTop = &self.stack;
- self.frames[0] = CallFrame.init(function, self.stackTop);
+ self.frames[0] = CallFrame.init(.Function, function, self.stackTop);
self.push(Value.init(function.cast()));
try self.execute(dbg);
}
@@ -141,7 +152,7 @@ pub const VM = struct {
}
fn read_constant(self: *@This()) Value {
- return self.frame().function.chunk.constants.get(self.read_byte()) catch unreachable;
+ return self.frame().chunk.constants.get(self.read_byte()) catch unreachable;
}
fn read_string(self: *@This()) *const Obj.String {
@@ -164,24 +175,30 @@ pub const VM = struct {
fn callValue(self: *@This(), callee: Value, argCount: u8) !void {
if(callee.is(Obj.Type.Function)) {
- try self.call(callee.obj.cast(.Function) catch unreachable, argCount);
+ try self.callFunction(callee.obj.cast(.Function) catch unreachable, argCount);
+ } else if(callee.is(Obj.Type.Closure)) {
+ try self.callClosure(callee.obj.cast(.Closure) catch unreachable, argCount);
} else if(callee.is(Obj.Type.Native)) {
- const native = callee.obj.cast(.Native) catch unreachable;
- if (argCount < native.arity_min or argCount > native.arity_max) {
- self.runtimeError("Expected from {d} to {d} arguments but got {d}", .{native.arity_min, native.arity_max, argCount});
- return InterpreterError.RuntimeError;
- }
- const result = try native.call(&self.vm.objects, argCount, self.stackTop - argCount);
- self.stackTop -= argCount + 1;
- self.push(result);
- return;
+ try self.callNative(callee.obj.cast(.Native) catch unreachable, argCount);
} else {
self.runtimeError("Can only call functions and classes", .{});
return InterpreterError.RuntimeError;
}
}
- fn call(self: *@This(), callee: *Obj.Function, argCount: u8) !void {
+ fn callClosure(self: *@This(), callee: *Obj.Closure, argCount: u8) !void {
+ if (argCount != callee.function.arity) {
+ self.runtimeError("Expected {d} arguments but got {d}", .{callee.function.arity, argCount});
+ return InterpreterError.RuntimeError;
+ }
+ if (self.frameCount == callstack_size - 1)
+ return InterpreterError.StackOverflow;
+ self.frameCount += 1;
+ self.frames[self.frameCount - 1] = CallFrame.init(.Closure, callee, self.stackTop - argCount - 1);
+ }
+
+
+ fn callFunction(self: *@This(), callee: *Obj.Function, argCount: u8) !void {
if (argCount != callee.arity) {
self.runtimeError("Expected {d} arguments but got {d}", .{callee.arity, argCount});
return InterpreterError.RuntimeError;
@@ -189,7 +206,18 @@ pub const VM = struct {
if (self.frameCount == callstack_size - 1)
return InterpreterError.StackOverflow;
self.frameCount += 1;
- self.frames[self.frameCount - 1] = CallFrame.init(callee, self.stackTop - argCount - 1);
+ self.frames[self.frameCount - 1] = CallFrame.init(.Function, callee, self.stackTop - argCount - 1);
+ }
+
+ fn callNative(self: *@This(), native: *Obj.Native, argCount: u8) !void {
+ if (argCount < native.arity_min or argCount > native.arity_max) {
+ self.runtimeError("Expected from {d} to {d} arguments but got {d}", .{native.arity_min, native.arity_max, argCount});
+ return InterpreterError.RuntimeError;
+ }
+ const result = try native.call(&self.vm.objects, argCount, self.stackTop - argCount);
+ self.stackTop -= argCount + 1;
+ self.push(result);
+ return;
}
fn binary_op(self: *@This(), comptime in_tag: anytype, comptime out_tag: anytype, op: Callback.Type(in_tag, out_tag)) InterpreterError!void {
@@ -204,7 +232,7 @@ pub const VM = struct {
}
fn instruction_idx(self: *const @This()) usize {
- return @intFromPtr(self.ip()) - @intFromPtr(self.frame().function.chunk.code.data.ptr);
+ return @intFromPtr(self.ip()) - @intFromPtr(self.frame().chunk.code.data.ptr);
}
fn execute(self: *@This(), dbg: bool) !void {
@@ -216,7 +244,7 @@ pub const VM = struct {
std.debug.print("[{s}]", .{stackPtr[0]});
}
std.debug.print("\n", .{});
- _ = try debug.disassembleInstruction(self.frame().function.chunk, self.instruction_idx());
+ _ = try debug.disassembleInstruction(self.frame().chunk, self.instruction_idx());
}
const instruction: u8 = self.read_byte();
switch (instruction) {
@@ -295,7 +323,7 @@ pub const VM = struct {
var pushed = false;
if (obj.is(Value.obj)) {
switch(obj.obj.type) {
- .Function, .Native => {},
+ .Function, .Native, .Closure => {},
inline else => |tp| {
self.push((obj.obj.cast(tp) catch unreachable).get(key) catch Value.init({}));
pushed = true;
@@ -345,6 +373,10 @@ pub const VM = struct {
const argCount = self.read_byte();
try self.callValue(self.peek(argCount), argCount);
},
+ @intFromEnum(OP.CLOSURE) => {
+ const function = try self.read_constant().obj.cast(.Function);
+ self.push(Value.init(try self.vm.objects.emplace_cast(.Closure, function)));
+ },
@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())),
@intFromEnum(OP.SUBTRACT) => try self.binary_op(Value.number, Value.number, Callback.sub),
@@ -366,8 +398,8 @@ pub const VM = struct {
var i = self.frameCount - 1;
while (true) : (i -= 1) {
const fram = self.frames[i];
- const idx = @intFromPtr(fram.ip) - @intFromPtr(fram.function.chunk.code.data.ptr);
- std.debug.print("[line {d}] in {s}\n", .{fram.function.chunk.lines.get(idx) catch 1, fram.function});
+ const idx = @intFromPtr(fram.ip) - @intFromPtr(fram.chunk.code.data.ptr);
+ std.debug.print("[line {d}] in {s}\n", .{fram.chunk.lines.get(idx) catch 1, fram.callee});
if (i == 0) break;
}
std.debug.print(fmt ++ "\n", args);