lib/choir/src/core/rewrite/canonicalization.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

 1 const std = @import("std");
 2 const ir = @import("../root.zig");
 3 const rewrite = ir.rewrite;
 4 
 5 pub const builtin_pattern_benefit: rewrite.PatternBenefit = 10;
 6 pub const fold_pattern_benefit: rewrite.PatternBenefit = 100;
 7 
 8 pub const DialectCanonicalizationInterface = struct {
 9     pub const interface_name = "choir.pass.dialect_canonicalization";
10     pub const id: ir.InterfaceId = ir.interfaceId(interface_name);
11 
12     pub const VTable = struct {
13         patterns: []const rewrite.RewritePattern,
14     };
15 
16     fn entry(vtable: *const VTable) ir.InterfaceEntry {
17         return .{ .id = id, .vtable = vtable };
18     }
19 
20     pub fn vtableFor(
21         comptime dialect_name: []const u8,
22         comptime patterns: []const rewrite.RewritePattern,
23     ) *const VTable {
24         for (patterns, 0..) |pattern, index| {
25             const separator = std.mem.indexOfScalar(u8, pattern.spec.root_op_name, '.') orelse
26                 @compileError("dialect canonicalization pattern root requires a dialect namespace");
27             if (!std.mem.eql(u8, pattern.spec.root_op_name[0..separator], dialect_name)) {
28                 @compileError("dialect canonicalization pattern root must belong to its dialect");
29             }
30             if (pattern.spec.benefit <= builtin_pattern_benefit or
31                 pattern.spec.benefit >= fold_pattern_benefit)
32             {
33                 @compileError("dialect canonicalization pattern benefit must remain between builtin and fold tiers");
34             }
35             for (patterns[0..index]) |previous| {
36                 if (!std.mem.eql(u8, previous.spec.root_op_name, pattern.spec.root_op_name)) continue;
37                 if (previous.spec.benefit < pattern.spec.benefit) {
38                     @compileError("dialect canonicalization patterns must descend by benefit per root");
39                 }
40             }
41         }
42         return &.{ .patterns = patterns };
43     }
44 
45     pub fn entryFor(
46         comptime dialect_name: []const u8,
47         comptime patterns: []const rewrite.RewritePattern,
48     ) ir.InterfaceEntry {
49         return entry(vtableFor(dialect_name, patterns));
50     }
51 
52     pub fn fromOpaque(ptr: *const anyopaque) *const VTable {
53         return @ptrCast(@alignCast(ptr));
54     }
55 };