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 }