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 }