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 };