lib/choir/src/backends/wasm/emission/module/section/export.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const wasm = @import("../../../root.zig");
  3 const emission = @import("../../root.zig");
  4 const common = @import("../root.zig").common;
  5 
  6 const memory_export_name = "memory";
  7 
  8 pub fn write(
  9     out: anytype,
 10     plan: *const emission.Plan,
 11     module: *common.ir.Operation,
 12     options: emission.Options,
 13 ) emission.Error!void {
 14     try common.writeSectionHeader(out, .export_, plan.sections.export_);
 15     try writeBody(out, plan, module, options);
 16 }
 17 
 18 pub fn countBody(
 19     plan: *const emission.Plan,
 20     module: *common.ir.Operation,
 21     options: emission.Options,
 22 ) emission.Error!usize {
 23     var count = wasm.binary.Count{};
 24     try writeBody(&count, plan, module, options);
 25     return common.requireU32(count.len);
 26 }
 27 
 28 pub fn entryCount(
 29     module: *common.ir.Operation,
 30     options: emission.Options,
 31     needs_memory: bool,
 32 ) emission.Error!usize {
 33     var count: usize = if (needs_memory) 1 else 0;
 34     var symbols = try emission.symbols.Iterator.init(module, options);
 35     while (try symbols.next()) |symbol| {
 36         if (symbol.role != .provided) continue;
 37         if (needs_memory and std.mem.eql(u8, symbol.name, memory_export_name)) {
 38             return error.CodeGenFailed;
 39         }
 40         count = std.math.add(usize, count, 1) catch return error.CapacityOverflow;
 41     }
 42     return count;
 43 }
 44 
 45 fn writeBody(
 46     out: anytype,
 47     plan: *const emission.Plan,
 48     module: *common.ir.Operation,
 49     options: emission.Options,
 50 ) emission.Error!void {
 51     const count = try entryCount(module, options, plan.needsMemory());
 52     _ = try common.requireU32(count);
 53     try wasm.binary.writeUleb(out, count);
 54     var symbols = try emission.symbols.Iterator.init(module, options);
 55     while (try symbols.next()) |symbol| {
 56         if (symbol.role != .provided) continue;
 57         const function_index = plan.functionIndex(symbol.name) orelse
 58             return error.FunctionNotFound;
 59         try common.writeName(out, symbol.name);
 60         try out.writeByte(@backingInt(wasm.binary.ExternalKind.function));
 61         try wasm.binary.writeUleb(out, function_index);
 62     }
 63     if (plan.needsMemory()) {
 64         try common.writeName(out, memory_export_name);
 65         try out.writeByte(@backingInt(wasm.binary.ExternalKind.memory));
 66         try wasm.binary.writeUleb(out, @as(u32, 0));
 67     }
 68 }
 69 
 70 test "WASM export rejects a function that collides with linear memory" {
 71     const dialects = @import("../../../../../dialects/root.zig");
 72     const allocator = std.testing.allocator;
 73     var context = try common.ir.Context.init(allocator, common.ir.Context.Limits.testing);
 74     defer context.deinit(allocator);
 75     try dialects.registerAllDialects(&context);
 76 
 77     const location = common.ir.Location.getUnknown();
 78     const f32_type = try dialects.ArithDialect.getScalarType(&context, .f32);
 79     const memref_type = try dialects.MemrefDialect.getMemrefTypeDynamic(
 80         &context,
 81         f32_type,
 82         .host,
 83     );
 84     const source_module = try common.BuiltinDialect.ModuleOp.create(&context, location);
 85     var function = try common.FuncDialect.FuncOp.create(
 86         &context,
 87         location,
 88         memory_export_name,
 89         &.{memref_type},
 90         &.{},
 91     );
 92     try source_module.getBodyBlock().addOperation(function.op);
 93     const return_op = try common.FuncDialect.ReturnOp.create(&context, location, &.{});
 94     try function.getEntryBlock().addOperation(return_op.op);
 95 
 96     const limits = try emission.Limits.inspect(source_module.op, .{});
 97     try std.testing.expectError(
 98         error.CodeGenFailed,
 99         emission.ModuleEmitter.init(allocator, limits),
100     );
101 }