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 }