lib/choir/src/backends/wasm/backend.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("../../core/root.zig");
3 const dialects = @import("../../dialects/root.zig");
4 const backends = @import("../root.zig");
5 const wasm = @import("root.zig");
6
7 const Allocator = std.mem.Allocator;
8 const BackendError = backends.interface.BackendError;
9 const BuiltinDialect = dialects.BuiltinDialect;
10 const FuncDialect = dialects.FuncDialect;
11
12 pub const Backend = struct {
13 allocator: Allocator,
14 ctx: *ir.Context,
15
16 pub fn init(allocator: Allocator, ctx: *ir.Context) Allocator.Error!Backend {
17 return .{
18 .allocator = allocator,
19 .ctx = ctx,
20 };
21 }
22
23 pub fn deinit(_: *Backend) void {}
24
25 pub fn verify(_: *Backend, module: *ir.Operation) BackendError!void {
26 return backends.contract.verifyModule("backend/wasm/verify", module);
27 }
28
29 pub fn lower(self: *Backend, module: *ir.Operation) BackendError!*ir.Operation {
30 try self.verify(module);
31 return module;
32 }
33
34 pub fn emit(
35 self: *Backend,
36 module: *ir.Operation,
37 options: backends.interface.EmitOptions,
38 writer: *std.Io.Writer,
39 ) BackendError!void {
40 const limits = wasm.ModuleEmitter.Limits.inspect(module, options) catch |err| {
41 return mapEmissionError(err);
42 };
43 var emitter = wasm.ModuleEmitter.init(self.allocator, limits) catch |err| {
44 return mapEmissionError(err);
45 };
46 defer emitter.deinit(self.allocator);
47 emitter.activate() catch |err| return mapEmissionError(err);
48 const encoded = emitter.emit() catch |err| return mapEmissionError(err);
49 writer.writeAll(encoded) catch return BackendError.CodeGenFailed;
50 }
51
52 pub fn compileModuleToArtifact(
53 self: *Backend,
54 module: *ir.Operation,
55 options: backends.interface.EmitOptions,
56 ) BackendError!backends.artifact.Artifact {
57 const limits = wasm.ModuleEmitter.Limits.inspect(module, options) catch |err| {
58 return mapEmissionError(err);
59 };
60 var emitter = wasm.ModuleEmitter.init(self.allocator, limits) catch |err| {
61 return mapEmissionError(err);
62 };
63 defer emitter.deinit(self.allocator);
64 emitter.activate() catch |err| return mapEmissionError(err);
65 const encoded = emitter.emit() catch |err| return mapEmissionError(err);
66
67 const name = options.entry orelse "module";
68 var artifact = backends.artifact.webassemblyModuleArtifact(
69 self.allocator,
70 name,
71 encoded,
72 ) catch return error.OutOfMemory;
73 errdefer artifact.deinit();
74
75 var symbols = wasm.emission.symbols.Iterator.init(module, options) catch |err| {
76 return mapEmissionError(err);
77 };
78 while (symbols.next() catch |err| return mapEmissionError(err)) |symbol| {
79 switch (symbol.role) {
80 .required => artifact.linkage.addRequired(.{
81 .name = symbol.name,
82 .kind = .function,
83 .binding = .external,
84 }) catch return error.OutOfMemory,
85 .provided => artifact.linkage.addProvided(.{
86 .name = symbol.name,
87 .kind = .function,
88 .binding = .external,
89 }) catch return error.OutOfMemory,
90 }
91 }
92 return artifact;
93 }
94 };
95
96 fn mapEmissionError(err: wasm.emission.Error) BackendError {
97 return switch (err) {
98 error.OutOfMemory => error.OutOfMemory,
99 error.FunctionNotFound => error.FunctionNotFound,
100 error.UnsupportedOperation => error.UnsupportedOperation,
101 else => error.CodeGenFailed,
102 };
103 }
104
105 pub fn initHandle(allocator: Allocator, ctx: *ir.Context) BackendError!backends.interface.BackendHandle {
106 return backends.interface.initHandle(
107 Backend,
108 allocator,
109 ctx,
110 backends.interface.BackendTarget.wasm,
111 "wasm",
112 .{
113 .artifact = .{ .webassembly_module = true },
114 },
115 );
116 }
117
118 test "wasm backend handle emits module bytes" {
119 const allocator = std.testing.allocator;
120 const ArithDialect = dialects.ArithDialect;
121
122 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
123 defer ctx.deinit(allocator);
124 try dialects.registerAllDialects(&ctx);
125
126 const loc = ir.Location.getUnknown();
127 const i32_type = try ArithDialect.getI32Type(&ctx);
128 const module = try BuiltinDialect.ModuleOp.create(&ctx, loc);
129 const block = module.getBodyBlock();
130
131 var func = try FuncDialect.FuncOp.create(&ctx, loc, "identity", &.{i32_type}, &.{i32_type});
132 try block.addOperation(func.op);
133 const ret = try FuncDialect.ReturnOp.create(&ctx, loc, &.{func.getArgument(0)});
134 try func.getEntryBlock().addOperation(ret.op);
135
136 var handle = try initHandle(allocator, &ctx);
137 defer handle.deinit();
138
139 var out = std.Io.Writer.Allocating.init(allocator);
140 defer out.deinit();
141 try handle.emit(module.op, .{ .entry = "identity" }, &out.writer);
142
143 const bytes = out.written();
144 try std.testing.expect(bytes.len >= 8);
145 try std.testing.expectEqualSlices(u8, &.{ 0x00, 0x61, 0x73, 0x6d }, bytes[0..4]);
146 }
147
148 test "wasm backend refuses overflow arithmetic by name" {
149 const allocator = std.testing.allocator;
150 const Arith = dialects.ArithDialect;
151 inline for (.{ Arith.AddoOp, Arith.SuboOp, Arith.MuloOp }) |Op| {
152 for ([_]bool{ false, true }) |used| {
153 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
154 defer ctx.deinit(allocator);
155 try dialects.registerAllDialects(&ctx);
156 const integer = try Arith.getScalarType(&ctx, .i64);
157 const module = try BuiltinDialect.ModuleOp.create(&ctx, .unknown);
158 const function = try FuncDialect.FuncOp.create(&ctx, .unknown, "overflow", &.{ integer, integer }, &.{integer});
159 try module.getBodyBlock().addOperation(function.op);
160 const checked = try Op.create(&ctx, .unknown, function.getArgument(0), function.getArgument(1));
161 try function.getEntryBlock().addOperation(checked.op);
162 const value = if (used) checked.getResult() else function.getArgument(0);
163 const ret = try FuncDialect.ReturnOp.create(&ctx, .unknown, &.{value});
164 try function.getEntryBlock().addOperation(ret.op);
165 var backend = try Backend.init(allocator, &ctx);
166 defer backend.deinit();
167 try std.testing.expectError(error.UnsupportedOperation, backend.compileModuleToArtifact(module.op, .{ .entry = "overflow" }));
168 }
169 }
170 }