lib/choir/src/core/rewrite/pattern.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("../root.zig");
3 const PatternRewriter = ir.rewrite.PatternRewriter;
4
5 pub const PatternBenefit = u16;
6
7 pub const RewritePatternKind = enum {
8 rewrite,
9 fold,
10 };
11
12 pub const RewritePatternProducts = union(enum) {
13 unknown,
14 none,
15 operations: []const []const u8,
16 };
17
18 pub const RewritePatternSpec = struct {
19 name: []const u8,
20 root_op_name: []const u8,
21 benefit: PatternBenefit = 1,
22 kind: RewritePatternKind = .rewrite,
23 products: RewritePatternProducts = .unknown,
24 };
25
26 pub const RewritePattern = struct {
27 spec: RewritePatternSpec,
28
29 rewrite_fn: *const fn (
30 op: *ir.Operation,
31 rewriter: *PatternRewriter,
32 ) PatternResult,
33
34 match_fn: ?*const fn (op: *ir.Operation) bool,
35
36 order: usize,
37
38 pub fn init(
39 spec: RewritePatternSpec,
40 rewrite_fn: *const fn (*ir.Operation, *PatternRewriter) PatternResult,
41 ) RewritePattern {
42 return .{
43 .spec = spec,
44 .rewrite_fn = rewrite_fn,
45 .match_fn = null,
46 .order = 0,
47 };
48 }
49
50 pub fn initWithMatch(
51 spec: RewritePatternSpec,
52 match_fn: *const fn (*ir.Operation) bool,
53 rewrite_fn: *const fn (*ir.Operation, *PatternRewriter) PatternResult,
54 ) RewritePattern {
55 return .{
56 .spec = spec,
57 .rewrite_fn = rewrite_fn,
58 .match_fn = match_fn,
59 .order = 0,
60 };
61 }
62
63 pub fn matches(self: *const RewritePattern, op: *ir.Operation) bool {
64 if (!std.mem.eql(u8, op.name.name, self.spec.root_op_name)) {
65 return false;
66 }
67 return self.matchesAfterRoot(op);
68 }
69
70 pub fn matchesAfterRoot(self: *const RewritePattern, op: *ir.Operation) bool {
71 if (self.match_fn) |match| {
72 return match(op);
73 }
74 return true;
75 }
76
77 pub fn apply(self: *const RewritePattern, op: *ir.Operation, rewriter: *PatternRewriter) PatternResult {
78 return self.rewrite_fn(op, rewriter);
79 }
80
81 pub fn matchAndRewrite(self: *const RewritePattern, op: *ir.Operation, rewriter: *PatternRewriter) PatternResult {
82 if (!self.matches(op)) return .failure;
83 return self.apply(op, rewriter);
84 }
85
86 pub fn lessThan(_: void, a: RewritePattern, b: RewritePattern) bool {
87 if (a.spec.benefit != b.spec.benefit) return a.spec.benefit > b.spec.benefit;
88 return a.order < b.order;
89 }
90 };
91
92 pub const PatternResult = enum {
93 success,
94 failure,
95 };
96
97 test "RewritePattern stores inspectable declaration" {
98 const pattern = RewritePattern.init(.{
99 .name = "test-forward",
100 .root_op_name = "test.source",
101 .benefit = 9,
102 .products = .none,
103 }, rewriteNoopForPatternOrder);
104
105 try std.testing.expectEqualStrings("test-forward", pattern.spec.name);
106 try std.testing.expectEqualStrings("test.source", pattern.spec.root_op_name);
107 try std.testing.expectEqual(@as(PatternBenefit, 9), pattern.spec.benefit);
108 try std.testing.expectEqual(.rewrite, pattern.spec.kind);
109 switch (pattern.spec.products) {
110 .none => {},
111 else => return error.TestExpectedNoProducts,
112 }
113 }
114
115 fn rewriteNoopForPatternOrder(_: *ir.Operation, _: *PatternRewriter) PatternResult {
116 return .failure;
117 }