lib/choir/src/passes/pass/operation.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("../../core/root.zig");
3 const subject = @import("root.zig");
4
5 const Pass = subject.Pass;
6 const PassContext = subject.PassContext;
7 const PassMutationScope = subject.PassMutationScope;
8 const PassResult = subject.PassResult;
9
10 pub fn OperationPass(
11 comptime OpType: type,
12 comptime run_on_op: fn (*OpType, *PassContext) PassResult,
13 ) type {
14 return struct {
15 base: Pass,
16
17 const Self = @This();
18
19 pub fn init(
20 name: []const u8,
21 description: []const u8,
22 ) Self {
23 return initWithMutationScope(name, description, .whole_module);
24 }
25
26 pub fn initWithMutationScope(
27 name: []const u8,
28 description: []const u8,
29 mutation_scope: PassMutationScope,
30 ) Self {
31 return .{
32 .base = .{
33 .name = name,
34 .description = description,
35 .run_fn = &runImpl,
36 .mutation_scope = mutation_scope,
37 },
38 };
39 }
40
41 fn runImpl(ctx: *PassContext) PassResult {
42 return walkAndApply(ctx.op, ctx);
43 }
44
45 fn walkAndApply(op: *ir.Operation, ctx: *PassContext) PassResult {
46 if (std.mem.eql(u8, op.name.name, OpType.operation_name)) {
47 var typed_op = OpType{ .op = op };
48 const result = run_on_op(&typed_op, ctx);
49 if (result == .failure) {
50 return .failure;
51 }
52 }
53
54 for (op.regions.items) |*region| {
55 var block_iter = region.getBlocks();
56 while (block_iter.next()) |block| {
57 var current_op: ?*ir.Operation = @ptrCast(@alignCast(block.operations.head));
58 while (current_op) |nested_op| {
59 const result = walkAndApply(nested_op, ctx);
60 if (result == .failure) {
61 return .failure;
62 }
63 current_op = nested_op.next_op;
64 }
65 }
66 }
67
68 return .success;
69 }
70 };
71 }