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 }