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

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const ir = @import("../root.zig");
  3 const alloc_arena = @import("alloc_arena");
  4 const RewritePattern = ir.rewrite.RewritePattern;
  5 const PatternResult = ir.rewrite.PatternResult;
  6 const TypeConverter = ir.rewrite.TypeConverter;
  7 
  8 pub const PatternRewriteAttempt = struct {
  9     rewriter: *PatternRewriter,
 10     created_before: usize,
 11     guard: PatternRewriter.InsertionGuard,
 12 
 13     pub inline fn init(op: *ir.Operation, rewriter: *PatternRewriter) PatternRewriteAttempt {
 14         const attempt = PatternRewriteAttempt{
 15             .rewriter = rewriter,
 16             .created_before = rewriter.created_ops.items.len,
 17             .guard = rewriter.insertionGuard(),
 18         };
 19         rewriter.setInsertionPointBefore(op);
 20         return attempt;
 21     }
 22 
 23     pub inline fn finish(self: *PatternRewriteAttempt, result: PatternResult) bool {
 24         if (result == .failure and self.rewriter.created_ops.items.len != self.created_before) {
 25             self.rewriter.created_ops.items.len = self.created_before;
 26         }
 27         self.guard.deinit();
 28         self.* = undefined;
 29         return result == .success;
 30     }
 31 };
 32 
 33 pub inline fn tryApplyRewritePattern(
 34     pattern: *const RewritePattern,
 35     op: *ir.Operation,
 36     rewriter: *PatternRewriter,
 37 ) bool {
 38     if (!pattern.matchesAfterRoot(op)) return false;
 39     var attempt = PatternRewriteAttempt.init(op, rewriter);
 40     return attempt.finish(pattern.apply(op, rewriter));
 41 }
 42 
 43 fn rewrite_listener_target(comptime TargetPtr: type) type {
 44     return struct {
 45         fn notify_event(context: ?*anyopaque, event: PatternRewriter.RewriteEvent) void {
 46             const typed: TargetPtr = @ptrCast(@alignCast(context.?));
 47             typed.notifyRewriteEvent(event);
 48         }
 49     };
 50 }
 51 
 52 pub const PatternRewriter = struct {
 53     pub const ReplacementMode = enum {
 54         immediate,
 55         deferred,
 56     };
 57 
 58     allocator: std.mem.Allocator,
 59     ir_ctx: *ir.Context,
 60     type_converter: ?*const TypeConverter,
 61     replacement_mode: ReplacementMode,
 62 
 63     created_ops: std.ArrayListUnmanaged(*ir.Operation),
 64 
 65     ops_to_erase: std.ArrayListUnmanaged(*ir.Operation),
 66 
 67     replacements: std.AutoHashMap(*ir.Value, *ir.Value),
 68 
 69     builder: ir.OperationBuilder,
 70 
 71     listener: ?Listener,
 72 
 73     pub const InsertPoint: type = ir.OperationBuilder.InsertPoint;
 74     pub const InsertionGuard: type = ir.OperationBuilder.InsertionGuard;
 75 
 76     pub const OperationReplacement = struct {
 77         op: *ir.Operation,
 78         new_values: []const *ir.Value,
 79     };
 80 
 81     pub const DetachedOperation = struct {
 82         op: *ir.Operation,
 83         block: *ir.Block,
 84         next: ?*ir.Operation,
 85     };
 86 
 87     pub const RewriteEvent = union(enum) {
 88         modified: *ir.Operation,
 89         replaced: OperationReplacement,
 90         erased: *ir.Operation,
 91     };
 92 
 93     pub const Listener = struct {
 94         context: ?*anyopaque = null,
 95         notify_event: ?*const fn (?*anyopaque, RewriteEvent) void = null,
 96 
 97         pub fn bind(target: anytype) Listener {
 98             const TargetPtr = @TypeOf(target);
 99             comptime {
100                 const pointer_info = switch (@typeInfo(TargetPtr)) {
101                     .pointer => |info| info,
102                     else => @compileError("PatternRewriter.Listener.bind expects a mutable pointer"),
103                 };
104                 if (pointer_info.size != .one) {
105                     @compileError("PatternRewriter.Listener.bind expects a single-item pointer");
106                 }
107                 if (pointer_info.attrs.@"const") {
108                     @compileError("PatternRewriter.Listener.bind expects a mutable pointer");
109                 }
110                 if (!@hasDecl(pointer_info.child, "notifyRewriteEvent")) {
111                     @compileError("PatternRewriter listener target must define notifyRewriteEvent");
112                 }
113             }
114 
115             return .{
116                 .context = target,
117                 .notify_event = rewrite_listener_target(TargetPtr).notify_event,
118             };
119         }
120 
121         pub fn notify(self: Listener, event: RewriteEvent) void {
122             if (self.notify_event) |notify_event| notify_event(self.context, event);
123         }
124     };
125 
126     pub fn init(allocator: std.mem.Allocator, ir_ctx: *ir.Context) PatternRewriter {
127         return initWithMode(allocator, ir_ctx, .immediate);
128     }
129 
130     pub fn initDeferred(allocator: std.mem.Allocator, ir_ctx: *ir.Context) PatternRewriter {
131         return initWithMode(allocator, ir_ctx, .deferred);
132     }
133 
134     fn initWithMode(
135         allocator: std.mem.Allocator,
136         ir_ctx: *ir.Context,
137         replacement_mode: ReplacementMode,
138     ) PatternRewriter {
139         return .{
140             .allocator = allocator,
141             .ir_ctx = ir_ctx,
142             .type_converter = null,
143             .replacement_mode = replacement_mode,
144             .created_ops = .empty,
145             .ops_to_erase = .empty,
146             .replacements = std.AutoHashMap(*ir.Value, *ir.Value).init(allocator),
147             .builder = ir.OperationBuilder.init(ir_ctx),
148             .listener = null,
149         };
150     }
151 
152     pub fn initWithTypeConverter(
153         allocator: std.mem.Allocator,
154         ir_ctx: *ir.Context,
155         type_converter: ?*const TypeConverter,
156     ) PatternRewriter {
157         var rewriter = PatternRewriter.init(allocator, ir_ctx);
158         rewriter.type_converter = type_converter;
159         return rewriter;
160     }
161 
162     pub fn deinit(self: *PatternRewriter) void {
163         self.created_ops.deinit(self.allocator);
164         self.ops_to_erase.deinit(self.allocator);
165         self.replacements.deinit();
166     }
167 
168     pub fn setTypeConverter(self: *PatternRewriter, type_converter: ?*const TypeConverter) void {
169         self.type_converter = type_converter;
170     }
171 
172     pub fn getTypeConverter(self: *const PatternRewriter) ?*const TypeConverter {
173         return self.type_converter;
174     }
175 
176     pub fn setListener(self: *PatternRewriter, listener: ?Listener) void {
177         self.listener = listener;
178     }
179 
180     pub fn getListener(self: *const PatternRewriter) ?Listener {
181         return self.listener;
182     }
183 
184     pub fn materializeConversion(self: *PatternRewriter, value: *ir.Value, target_type: ir.Type) ?*ir.Value {
185         const converter = self.type_converter orelse return null;
186         return converter.materialize(self, value, target_type);
187     }
188 
189     pub fn clearInsertionPoint(self: *PatternRewriter) void {
190         self.builder.clearInsertionPoint();
191     }
192 
193     pub fn hasInsertionPoint(self: *const PatternRewriter) bool {
194         return self.builder.hasInsertionPoint();
195     }
196 
197     pub fn saveInsertionPoint(self: *const PatternRewriter) InsertPoint {
198         return self.builder.saveInsertionPoint();
199     }
200 
201     pub fn restoreInsertionPoint(self: *PatternRewriter, insert_point: InsertPoint) void {
202         self.builder.restoreInsertionPoint(insert_point);
203     }
204 
205     pub fn insertionGuard(self: *PatternRewriter) InsertionGuard {
206         return self.builder.insertionGuard();
207     }
208 
209     pub fn setInsertionPoint(self: *PatternRewriter, block: *ir.Block) void {
210         self.builder.setInsertionPoint(block);
211     }
212 
213     pub fn setInsertionPointBefore(self: *PatternRewriter, op: *ir.Operation) void {
214         self.builder.setInsertionPointBefore(op);
215     }
216 
217     pub fn setInsertionPointAfter(self: *PatternRewriter, op: *ir.Operation) void {
218         self.builder.setInsertionPointAfter(op);
219     }
220 
221     fn ensureBuilderListener(self: *PatternRewriter) void {
222         self.builder.setListener(.{
223             .context = self,
224             .notify_operation_inserted = recordInsertedOperation,
225         });
226     }
227 
228     fn recordInsertedOperation(context: ?*anyopaque, op: *ir.Operation, previous: ir.OperationBuilder.InsertPoint) !void {
229         _ = previous;
230         const self: *PatternRewriter = @ptrCast(@alignCast(context.?));
231         try self.created_ops.append(self.allocator, op);
232     }
233 
234     pub fn insert(self: *PatternRewriter, op: *ir.Operation) !*ir.Operation {
235         self.ensureBuilderListener();
236         return try self.builder.insert(op);
237     }
238 
239     pub fn create(self: *PatternRewriter, state: ir.Operation.State) !*ir.Operation {
240         self.ensureBuilderListener();
241         return try self.builder.create(state);
242     }
243 
244     fn notifyRewriteEvent(self: *const PatternRewriter, event: RewriteEvent) void {
245         if (self.listener) |listener| listener.notify(event);
246     }
247 
248     fn replaceUsesImmediately(self: *PatternRewriter, old_value: *ir.Value, new_value: *ir.Value) void {
249         if (old_value == new_value) return;
250         var uses = old_value.useIterator();
251         while (uses.next()) |operand| {
252             operand.setValue(new_value);
253             const owner: *ir.Operation = @ptrCast(@alignCast(operand.owner));
254             self.notifyRewriteEvent(.{ .modified = owner });
255         }
256     }
257 
258     fn ensureReplacementCapacity(self: *PatternRewriter, count: usize) !void {
259         if (self.replacement_mode == .deferred) {
260             try self.replacements.ensureUnusedCapacity(@intCast(count));
261         }
262     }
263 
264     fn ensureEraseCapacity(self: *PatternRewriter) !void {
265         try self.ops_to_erase.ensureUnusedCapacity(self.allocator, 1);
266     }
267 
268     pub fn setOperandValue(self: *PatternRewriter, op: *ir.Operation, index: usize, new_value: *ir.Value) void {
269         op.setOperandValue(index, new_value);
270         self.notifyRewriteEvent(.{ .modified = op });
271     }
272 
273     pub fn setAttr(self: *PatternRewriter, op: *ir.Operation, attr_name: []const u8, value: ir.Attribute) !void {
274         try op.setAttr(attr_name, value);
275         self.notifyRewriteEvent(.{ .modified = op });
276     }
277 
278     pub fn replaceAllUsesWith(self: *PatternRewriter, old_value: *ir.Value, new_value: *ir.Value) !void {
279         switch (self.replacement_mode) {
280             .immediate => self.replaceUsesImmediately(old_value, new_value),
281             .deferred => try self.replacements.put(old_value, new_value),
282         }
283     }
284 
285     pub fn replaceAllOpUsesWith(self: *PatternRewriter, op: *ir.Operation, new_values: []const *ir.Value) !void {
286         if (new_values.len != op.getNumResults()) return error.ReplacementResultCountMismatch;
287         try self.ensureReplacementCapacity(new_values.len);
288         self.notifyRewriteEvent(.{ .replaced = .{ .op = op, .new_values = new_values } });
289         switch (self.replacement_mode) {
290             .immediate => for (op.results.items, new_values) |*result, new_value| {
291                 self.replaceUsesImmediately(result, new_value);
292             },
293             .deferred => for (op.results.items, new_values) |*result, new_value| {
294                 self.replacements.putAssumeCapacity(result, new_value);
295             },
296         }
297     }
298 
299     pub fn replaceAllOpUsesWithOperation(self: *PatternRewriter, op: *ir.Operation, replacement: *ir.Operation) !void {
300         const result_count = op.getNumResults();
301         if (replacement.getNumResults() != result_count) return error.ReplacementResultCountMismatch;
302         if (result_count == 0) {
303             try self.replaceAllOpUsesWith(op, &.{});
304             return;
305         }
306         if (result_count == 1) {
307             const value = replacement.getResult(0) orelse return error.ReplacementResultCountMismatch;
308             try self.replaceAllOpUsesWith(op, &.{value});
309             return;
310         }
311 
312         const values = try self.allocator.alloc(*ir.Value, result_count);
313         defer self.allocator.free(values);
314         for (values, 0..) |*value, index| {
315             value.* = replacement.getResult(index) orelse return error.ReplacementResultCountMismatch;
316         }
317 
318         try self.replaceAllOpUsesWith(op, values);
319     }
320 
321     pub fn replaceOp(self: *PatternRewriter, op: *ir.Operation, new_values: []const *ir.Value) !void {
322         if (new_values.len != op.getNumResults()) return error.ReplacementResultCountMismatch;
323         try self.ensureEraseCapacity();
324         try self.replaceAllOpUsesWith(op, new_values);
325         self.ops_to_erase.appendAssumeCapacity(op);
326     }
327 
328     pub fn replaceOpWithOperation(self: *PatternRewriter, op: *ir.Operation, replacement: *ir.Operation) !void {
329         if (replacement.getNumResults() != op.getNumResults()) return error.ReplacementResultCountMismatch;
330         try self.ensureEraseCapacity();
331         try self.replaceAllOpUsesWithOperation(op, replacement);
332         self.ops_to_erase.appendAssumeCapacity(op);
333     }
334 
335     pub fn replaceOpWithNewOp(self: *PatternRewriter, op: *ir.Operation, state: ir.Operation.State) !*ir.Operation {
336         const result_count = op.getNumResults();
337         if (state.result_types.len != result_count) return error.ReplacementResultCountMismatch;
338         try self.ensureEraseCapacity();
339         try self.ensureReplacementCapacity(result_count);
340 
341         const new_op = try self.create(state);
342         try self.replaceOpWithOperation(op, new_op);
343         return new_op;
344     }
345 
346     pub fn replaceOpWithValue(self: *PatternRewriter, op: *ir.Operation, new_value: *ir.Value) !void {
347         try self.replaceOp(op, &.{new_value});
348     }
349 
350     pub fn eraseOp(self: *PatternRewriter, op: *ir.Operation) !void {
351         try self.ops_to_erase.append(self.allocator, op);
352     }
353 
354     pub fn isScheduledForErase(self: *const PatternRewriter, op: *const ir.Operation) bool {
355         for (self.ops_to_erase.items) |erased_op| {
356             if (erased_op == op) return true;
357         }
358         return false;
359     }
360 
361     pub fn getRemappedValue(self: *const PatternRewriter, value: *ir.Value) *ir.Value {
362         if (self.replacement_mode == .immediate) return value;
363         var current = value;
364         while (self.replacements.get(current)) |replacement| {
365             current = replacement;
366         }
367         return current;
368     }
369 
370     pub fn finalize(self: *PatternRewriter, root_op: *ir.Operation) void {
371         self.applyPendingReplacements(root_op);
372 
373         for (self.ops_to_erase.items) |op| {
374             self.notifyRewriteEvent(.{ .erased = op });
375             op.erase();
376         }
377     }
378 
379     pub fn applyPendingReplacements(self: *PatternRewriter, root_op: *ir.Operation) void {
380         if (self.replacement_mode == .immediate) return;
381         if (self.replacements.count() != 0) self.applyReplacementsToOp(root_op);
382     }
383 
384     pub fn detachScheduledOperations(
385         self: *PatternRewriter,
386         storage: []DetachedOperation,
387     ) []DetachedOperation {
388         std.debug.assert(storage.len >= self.ops_to_erase.items.len);
389         for (self.ops_to_erase.items, 0..) |op, index| {
390             const block = op.getBlock() orelse @panic("scheduled erased operation is detached");
391             storage[index] = .{
392                 .op = op,
393                 .block = block,
394                 .next = op.next_op,
395             };
396             self.notifyRewriteEvent(.{ .erased = op });
397             block.detachOperation(op);
398         }
399         return storage[0..self.ops_to_erase.items.len];
400     }
401 
402     pub fn restoreDetachedOperations(detached: []const DetachedOperation) void {
403         var index = detached.len;
404         while (index > 0) {
405             index -= 1;
406             const entry = detached[index];
407             if (entry.next) |next| {
408                 entry.block.insertBefore(entry.op, next) catch unreachable;
409             } else {
410                 entry.block.addOperation(entry.op) catch unreachable;
411             }
412         }
413     }
414 
415     pub fn destroyDetachedOperations(detached: []const DetachedOperation) void {
416         for (detached) |entry| entry.op.dropAllReferences();
417         var index = detached.len;
418         while (index > 0) {
419             index -= 1;
420             detached[index].op.erase();
421         }
422     }
423 
424     pub fn rollbackCreatedOperations(self: *PatternRewriter) void {
425         var index = self.created_ops.items.len;
426         while (index > 0) {
427             index -= 1;
428             self.created_ops.items[index].erase();
429         }
430         self.created_ops.clearRetainingCapacity();
431     }
432 
433     fn applyReplacementsToOp(self: *PatternRewriter, op: *ir.Operation) void {
434         var context = ReplacementWalk{ .rewriter = self };
435         _ = op.walk(.{ .order = .pre_order }, &context, ReplacementWalk.visit) catch unreachable;
436     }
437 };
438 
439 const ReplacementWalk = struct {
440     rewriter: *PatternRewriter,
441 
442     fn visit(self: *ReplacementWalk, op: *ir.Operation) void {
443         for (op.operands.items, 0..) |*operand, i| {
444             const new_value = self.rewriter.getRemappedValue(operand.value);
445             if (new_value != operand.value) {
446                 self.rewriter.setOperandValue(op, i, new_value);
447             }
448         }
449     }
450 };
451 
452 fn expectRewriterBlockNames(block: *ir.Block, expected: []const []const u8) !void {
453     var iter = block.getOperations();
454     for (expected) |name| {
455         const op = iter.next() orelse return error.TestExpectedOperation;
456         try std.testing.expectEqualStrings(name, op.name.name);
457     }
458     try std.testing.expect(iter.next() == null);
459 }
460 
461 const PatternRewriterEventRecorder = struct {
462     allocator: std.mem.Allocator,
463     events: std.ArrayListUnmanaged(Event) = .empty,
464 
465     const EventKind = enum {
466         modified,
467         replaced,
468         erased,
469     };
470 
471     const Event = struct {
472         kind: EventKind,
473         op: *ir.Operation,
474         value_count: usize = 0,
475     };
476 
477     fn init(allocator: std.mem.Allocator) PatternRewriterEventRecorder {
478         return .{ .allocator = allocator };
479     }
480 
481     fn deinit(self: *PatternRewriterEventRecorder) void {
482         self.events.deinit(self.allocator);
483     }
484 
485     pub fn notifyRewriteEvent(self: *PatternRewriterEventRecorder, event: PatternRewriter.RewriteEvent) void {
486         const recorded: Event = switch (event) {
487             .modified => |op| .{ .kind = .modified, .op = op },
488             .replaced => |replacement| .{
489                 .kind = .replaced,
490                 .op = replacement.op,
491                 .value_count = replacement.new_values.len,
492             },
493             .erased => |op| .{ .kind = .erased, .op = op },
494         };
495         self.events.append(self.allocator, recorded) catch unreachable;
496     }
497 };
498 
499 test "PatternRewriter delegates insertion guard to operation builder" {
500     const testing = std.testing;
501 
502     var ctx = try ir.Context.init(testing.allocator, ir.Context.Limits.testing);
503     defer ctx.deinit(testing.allocator);
504     try ctx.allowUnregistered();
505 
506     var block = ir.Block.init(testing.allocator);
507     defer block.deinit();
508 
509     const anchor = try ctx.createOperation(ir.Operation.State.init("test.rewriter.anchor", .unknown));
510     try block.addOperation(anchor);
511 
512     var rewriter = PatternRewriter.init(testing.allocator, &ctx);
513     defer rewriter.deinit();
514 
515     rewriter.setInsertionPointBefore(anchor);
516     {
517         var guard = rewriter.insertionGuard();
518         defer guard.deinit();
519         rewriter.setInsertionPoint(&block);
520         _ = try rewriter.create(ir.Operation.State.init("test.rewriter.tail", .unknown));
521     }
522 
523     _ = try rewriter.create(ir.Operation.State.init("test.rewriter.before_anchor", .unknown));
524 
525     try expectRewriterBlockNames(&block, &.{
526         "test.rewriter.before_anchor",
527         "test.rewriter.anchor",
528         "test.rewriter.tail",
529     });
530 }
531 
532 test "PatternRewriter tracks inserted ops through builder listener" {
533     const testing = std.testing;
534 
535     var arena = alloc_arena.Arena.init(std.testing.allocator);
536     defer arena.deinit();
537     const allocator = arena.allocator();
538 
539     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
540     defer ctx.deinit(allocator);
541     try ctx.allowUnregistered();
542 
543     var rewriter = PatternRewriter.init(allocator, &ctx);
544     defer rewriter.deinit();
545 
546     const detached = try rewriter.create(ir.Operation.State.init("test.detached", ir.Location.getUnknown()));
547     try testing.expect(detached.parent_block == null);
548     try testing.expectEqual(@as(usize, 0), rewriter.created_ops.items.len);
549 
550     var block = ir.Block.init(allocator);
551     defer block.deinit();
552 
553     rewriter.setInsertionPoint(&block);
554     const created = try rewriter.create(ir.Operation.State.init("test.created", ir.Location.getUnknown()));
555     try testing.expectEqual(@as(usize, 1), rewriter.created_ops.items.len);
556     try testing.expect(rewriter.created_ops.items[0] == created);
557 
558     const inserted = try ctx.createOperation(ir.Operation.State.init("test.inserted", ir.Location.getUnknown()));
559     _ = try rewriter.insert(inserted);
560     try testing.expectEqual(@as(usize, 2), rewriter.created_ops.items.len);
561     try testing.expect(rewriter.created_ops.items[1] == inserted);
562     try expectRewriterBlockNames(&block, &.{ "test.created", "test.inserted" });
563 }
564 
565 test "PatternRewriter listener binds typed rewrite events" {
566     const testing = std.testing;
567 
568     var arena = alloc_arena.Arena.init(std.testing.allocator);
569     defer arena.deinit();
570     const allocator = arena.allocator();
571 
572     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
573     defer ctx.deinit(allocator);
574     try ctx.allowUnregistered();
575 
576     const op = try ctx.createOperation(ir.Operation.State.init("test.attr_target", .unknown));
577     const attr = try ctx.getStringAttr("tracked");
578 
579     var recorder = PatternRewriterEventRecorder.init(allocator);
580     defer recorder.deinit();
581 
582     var rewriter = PatternRewriter.init(allocator, &ctx);
583     defer rewriter.deinit();
584     rewriter.setListener(PatternRewriter.Listener.bind(&recorder));
585 
586     try rewriter.setAttr(op, "debug.note", attr);
587 
588     try testing.expect(op.getAttr("debug.note") != null);
589     try testing.expectEqual(@as(usize, 1), recorder.events.items.len);
590     try testing.expectEqual(PatternRewriterEventRecorder.EventKind.modified, recorder.events.items[0].kind);
591     try testing.expect(recorder.events.items[0].op == op);
592 }
593 
594 test "PatternRewriter replaceOp rejects result count mismatches" {
595     const testing = std.testing;
596 
597     var arena = alloc_arena.Arena.init(testing.allocator);
598     defer arena.deinit();
599     const allocator = arena.allocator();
600 
601     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
602     defer ctx.deinit(allocator);
603     try ctx.allowUnregistered();
604 
605     const value_type = try ctx.getDialectTypeFromName("test.value");
606     var old_state = ir.Operation.State.init("test.old", .unknown);
607     old_state.addTypes(&.{value_type});
608     const old_op = try ctx.createOperation(old_state);
609 
610     var rewriter = PatternRewriter.init(allocator, &ctx);
611     defer rewriter.deinit();
612 
613     try testing.expectError(error.ReplacementResultCountMismatch, rewriter.replaceOp(old_op, &.{}));
614     try testing.expectEqual(@as(usize, 0), rewriter.replacements.count());
615     try testing.expect(!rewriter.isScheduledForErase(old_op));
616 }
617 
618 test "PatternRewriter immediate replacement uses no allocator or finalize walk" {
619     const testing = std.testing;
620 
621     var arena = alloc_arena.Arena.init(testing.allocator);
622     defer arena.deinit();
623     const allocator = arena.allocator();
624 
625     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
626     defer ctx.deinit(allocator);
627     try ctx.allowUnregistered();
628 
629     const old_type = try ctx.getDialectTypeFromName("test.old_type");
630     const new_type = try ctx.getDialectTypeFromName("test.new_type");
631 
632     var replacement_state = ir.Operation.State.init("test.replacement", .unknown);
633     replacement_state.addTypes(&.{new_type});
634     const replacement_op = try ctx.createOperation(replacement_state);
635     const replacement_value = replacement_op.getResult(0).?;
636 
637     var old_state = ir.Operation.State.init("test.old_value", .unknown);
638     old_state.addTypes(&.{old_type});
639     const old_op = try ctx.createOperation(old_state);
640     const old_value = old_op.getResult(0).?;
641 
642     var user_state = ir.Operation.State.init("test.user", .unknown);
643     user_state.addOperands(&.{old_value});
644     const user = try ctx.createOperation(user_state);
645 
646     var body = ir.context.initRegion(&ctx);
647     defer body.deinit();
648     const block = try body.addBlock();
649     try block.addOperation(replacement_op);
650     try block.addOperation(old_op);
651     try block.addOperation(user);
652 
653     var root_state = ir.Operation.State.init("test.root", .unknown);
654     root_state.addRegionBodies(&.{&body});
655     const root = try ctx.createOperation(root_state);
656 
657     var failing = std.testing.FailingAllocator.init(testing.allocator, .{ .fail_index = 0 });
658     var rewriter = PatternRewriter.init(failing.allocator(), &ctx);
659     defer rewriter.deinit();
660 
661     try rewriter.replaceAllOpUsesWithOperation(old_op, replacement_op);
662     try testing.expect(!rewriter.isScheduledForErase(old_op));
663     try testing.expect(user.operands.items[0].value == replacement_value);
664     try testing.expectEqual(@as(usize, 0), failing.alloc_index);
665 
666     rewriter.finalize(root);
667 
668     try testing.expect(user.operands.items[0].value == replacement_value);
669     try testing.expectEqual(@as(usize, 0), failing.alloc_index);
670     const root_block = root.getRegion(0).?.getEntryBlock().?;
671     try expectRewriterBlockNames(root_block, &.{ "test.replacement", "test.old_value", "test.user" });
672 }
673 
674 test "PatternRewriter listener records replacement modification and erase" {
675     const testing = std.testing;
676 
677     var arena = alloc_arena.Arena.init(testing.allocator);
678     defer arena.deinit();
679     const allocator = arena.allocator();
680 
681     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
682     defer ctx.deinit(allocator);
683     try ctx.allowUnregistered();
684 
685     const old_type = try ctx.getDialectTypeFromName("test.old_type");
686     const new_type = try ctx.getDialectTypeFromName("test.new_type");
687 
688     var replacement_state = ir.Operation.State.init("test.replacement", .unknown);
689     replacement_state.addTypes(&.{new_type});
690     const replacement_op = try ctx.createOperation(replacement_state);
691     const replacement_value = replacement_op.getResult(0).?;
692 
693     var old_state = ir.Operation.State.init("test.old_value", .unknown);
694     old_state.addTypes(&.{old_type});
695     const old_op = try ctx.createOperation(old_state);
696     const old_value = old_op.getResult(0).?;
697 
698     var user_state = ir.Operation.State.init("test.user", .unknown);
699     user_state.addOperands(&.{old_value});
700     const user = try ctx.createOperation(user_state);
701 
702     var body = ir.context.initRegion(&ctx);
703     defer body.deinit();
704     const block = try body.addBlock();
705     try block.addOperation(replacement_op);
706     try block.addOperation(old_op);
707     try block.addOperation(user);
708 
709     var root_state = ir.Operation.State.init("test.root", .unknown);
710     root_state.addRegionBodies(&.{&body});
711     const root = try ctx.createOperation(root_state);
712 
713     var recorder = PatternRewriterEventRecorder.init(allocator);
714     defer recorder.deinit();
715 
716     var rewriter = PatternRewriter.init(allocator, &ctx);
717     defer rewriter.deinit();
718     rewriter.setListener(PatternRewriter.Listener.bind(&recorder));
719 
720     try rewriter.replaceOpWithOperation(old_op, replacement_op);
721     try testing.expectEqual(@as(usize, 2), recorder.events.items.len);
722     try testing.expectEqual(PatternRewriterEventRecorder.EventKind.replaced, recorder.events.items[0].kind);
723     try testing.expect(recorder.events.items[0].op == old_op);
724     try testing.expectEqual(@as(usize, 1), recorder.events.items[0].value_count);
725 
726     rewriter.finalize(root);
727 
728     try testing.expect(user.operands.items[0].value == replacement_value);
729     try testing.expectEqual(@as(usize, 3), recorder.events.items.len);
730     try testing.expectEqual(PatternRewriterEventRecorder.EventKind.modified, recorder.events.items[1].kind);
731     try testing.expect(recorder.events.items[1].op == user);
732     try testing.expectEqual(PatternRewriterEventRecorder.EventKind.erased, recorder.events.items[2].kind);
733     try testing.expect(recorder.events.items[2].op == old_op);
734     const root_block = root.getRegion(0).?.getEntryBlock().?;
735     try expectRewriterBlockNames(root_block, &.{ "test.replacement", "test.user" });
736 }
737 
738 test "PatternRewriter replaceOpWithOperation remaps existing operation results" {
739     const testing = std.testing;
740 
741     var arena = alloc_arena.Arena.init(testing.allocator);
742     defer arena.deinit();
743     const allocator = arena.allocator();
744 
745     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
746     defer ctx.deinit(allocator);
747     try ctx.allowUnregistered();
748 
749     const old_type = try ctx.getDialectTypeFromName("test.old_type");
750     const new_type = try ctx.getDialectTypeFromName("test.new_type");
751 
752     var replacement_state = ir.Operation.State.init("test.replacement", .unknown);
753     replacement_state.addTypes(&.{new_type});
754     const replacement_op = try ctx.createOperation(replacement_state);
755     const replacement_value = replacement_op.getResult(0).?;
756 
757     var old_state = ir.Operation.State.init("test.old_value", .unknown);
758     old_state.addTypes(&.{old_type});
759     const old_op = try ctx.createOperation(old_state);
760     const old_value = old_op.getResult(0).?;
761 
762     var user_state = ir.Operation.State.init("test.user", .unknown);
763     user_state.addOperands(&.{old_value});
764     const user = try ctx.createOperation(user_state);
765 
766     var body = ir.context.initRegion(&ctx);
767     defer body.deinit();
768     const block = try body.addBlock();
769     try block.addOperation(replacement_op);
770     try block.addOperation(old_op);
771     try block.addOperation(user);
772 
773     var root_state = ir.Operation.State.init("test.root", .unknown);
774     root_state.addRegionBodies(&.{&body});
775     const root = try ctx.createOperation(root_state);
776 
777     var rewriter = PatternRewriter.init(allocator, &ctx);
778     defer rewriter.deinit();
779 
780     try rewriter.replaceOpWithOperation(old_op, replacement_op);
781     try testing.expectEqual(@as(usize, 0), rewriter.created_ops.items.len);
782     try testing.expect(rewriter.isScheduledForErase(old_op));
783 
784     rewriter.finalize(root);
785 
786     try testing.expect(user.operands.items[0].value == replacement_value);
787     const root_block = root.getRegion(0).?.getEntryBlock().?;
788     try expectRewriterBlockNames(root_block, &.{ "test.replacement", "test.user" });
789 }
790 
791 test "PatternRewriter detached erasures can be restored or committed" {
792     const testing = std.testing;
793 
794     var arena = alloc_arena.Arena.init(testing.allocator);
795     defer arena.deinit();
796     const allocator = arena.allocator();
797 
798     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
799     defer ctx.deinit(allocator);
800     try ctx.allowUnregistered();
801 
802     const value_type = try ctx.getDialectTypeFromName("test.value");
803     var replacement_state = ir.Operation.State.init("test.replacement", .unknown);
804     replacement_state.addTypes(&.{value_type});
805     const replacement = try ctx.createOperation(replacement_state);
806     var old_state = ir.Operation.State.init("test.old", .unknown);
807     old_state.addTypes(&.{value_type});
808     const old = try ctx.createOperation(old_state);
809     var user_state = ir.Operation.State.init("test.user", .unknown);
810     user_state.addOperands(&.{old.getResult(0).?});
811     const user = try ctx.createOperation(user_state);
812 
813     var body = ir.context.initRegion(&ctx);
814     defer body.deinit();
815     const block = try body.addBlock();
816     try block.addOperation(replacement);
817     try block.addOperation(old);
818     try block.addOperation(user);
819     var root_state = ir.Operation.State.init("test.root", .unknown);
820     root_state.addRegionBodies(&.{&body});
821     const root = try ctx.createOperation(root_state);
822 
823     var rewriter = PatternRewriter.initDeferred(allocator, &ctx);
824     defer rewriter.deinit();
825     try rewriter.replaceOpWithOperation(old, replacement);
826     try testing.expectEqual(old.getResult(0).?, user.getOperand(0).?);
827 
828     var detached_storage: [1]PatternRewriter.DetachedOperation = undefined;
829     rewriter.applyPendingReplacements(root);
830     const detached = rewriter.detachScheduledOperations(&detached_storage);
831     try testing.expectEqual(replacement.getResult(0).?, user.getOperand(0).?);
832     try expectRewriterBlockNames(block, &.{ "test.replacement", "test.user" });
833 
834     user.setOperandValue(0, old.getResult(0).?);
835     PatternRewriter.restoreDetachedOperations(detached);
836     try expectRewriterBlockNames(block, &.{ "test.replacement", "test.old", "test.user" });
837 
838     rewriter.applyPendingReplacements(root);
839     _ = rewriter.detachScheduledOperations(&detached_storage);
840     PatternRewriter.destroyDetachedOperations(&detached_storage);
841     try testing.expectEqual(replacement.getResult(0).?, user.getOperand(0).?);
842     try expectRewriterBlockNames(block, &.{ "test.replacement", "test.user" });
843 }
844 
845 test "PatternRewriter replaceOpWithNewOp remaps replacement results" {
846     const testing = std.testing;
847 
848     var arena = alloc_arena.Arena.init(testing.allocator);
849     defer arena.deinit();
850     const allocator = arena.allocator();
851 
852     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
853     defer ctx.deinit(allocator);
854     try ctx.allowUnregistered();
855 
856     const old_type = try ctx.getDialectTypeFromName("test.old_type");
857     const new_type = try ctx.getDialectTypeFromName("test.new_type");
858 
859     var old_state = ir.Operation.State.init("test.old_value", .unknown);
860     old_state.addTypes(&.{old_type});
861     const old_op = try ctx.createOperation(old_state);
862     const old_value = old_op.getResult(0).?;
863 
864     var user_state = ir.Operation.State.init("test.user", .unknown);
865     user_state.addOperands(&.{old_value});
866     const user = try ctx.createOperation(user_state);
867 
868     var body = ir.context.initRegion(&ctx);
869     defer body.deinit();
870     const block = try body.addBlock();
871     try block.addOperation(old_op);
872     try block.addOperation(user);
873 
874     var root_state = ir.Operation.State.init("test.root", .unknown);
875     root_state.addRegionBodies(&.{&body});
876     const root = try ctx.createOperation(root_state);
877 
878     var rewriter = PatternRewriter.init(allocator, &ctx);
879     defer rewriter.deinit();
880 
881     rewriter.setInsertionPointBefore(old_op);
882     var new_state = ir.Operation.State.init("test.new_value", .unknown);
883     new_state.addTypes(&.{new_type});
884     const new_op = try rewriter.replaceOpWithNewOp(old_op, new_state);
885     const new_value = new_op.getResult(0).?;
886 
887     try testing.expectEqual(@as(usize, 1), rewriter.created_ops.items.len);
888     try testing.expect(rewriter.created_ops.items[0] == new_op);
889     try testing.expect(rewriter.isScheduledForErase(old_op));
890 
891     rewriter.finalize(root);
892 
893     try testing.expect(user.operands.items[0].value == new_value);
894     const root_block = root.getRegion(0).?.getEntryBlock().?;
895     try expectRewriterBlockNames(root_block, &.{ "test.new_value", "test.user" });
896 }
897 
898 test "PatternRewriter replaceOpWithNewOp rejects result count mismatches before insertion" {
899     const testing = std.testing;
900 
901     var arena = alloc_arena.Arena.init(testing.allocator);
902     defer arena.deinit();
903     const allocator = arena.allocator();
904 
905     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
906     defer ctx.deinit(allocator);
907     try ctx.allowUnregistered();
908 
909     const i32_type = try ctx.getDialectTypeFromName("test.i32");
910     var block = ir.Block.init(allocator);
911     defer block.deinit();
912 
913     var old_state = ir.Operation.State.init("test.old", .unknown);
914     old_state.addTypes(&.{i32_type});
915     const old_op = try ctx.createOperation(old_state);
916     try block.addOperation(old_op);
917 
918     var rewriter = PatternRewriter.init(allocator, &ctx);
919     defer rewriter.deinit();
920     rewriter.setInsertionPointBefore(old_op);
921 
922     const no_result_state = ir.Operation.State.init("test.no_result", .unknown);
923     try testing.expectError(
924         error.ReplacementResultCountMismatch,
925         rewriter.replaceOpWithNewOp(old_op, no_result_state),
926     );
927 
928     try testing.expectEqual(@as(usize, 0), rewriter.created_ops.items.len);
929     try testing.expect(!rewriter.isScheduledForErase(old_op));
930     try expectRewriterBlockNames(&block, &.{"test.old"});
931 }
932 
933 test "PatternRewriter finalize remaps nested operands through operation walk" {
934     const testing = std.testing;
935 
936     var arena = alloc_arena.Arena.init(std.testing.allocator);
937     defer arena.deinit();
938     const allocator = arena.allocator();
939 
940     var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
941     defer ctx.deinit(allocator);
942     try ctx.allowUnregistered();
943 
944     const value_type = try ctx.getDialectTypeFromName("test.value");
945 
946     var old_state = ir.Operation.State.init("test.old_value", .unknown);
947     old_state.addTypes(&.{value_type});
948     const old_op = try ctx.createOperation(old_state);
949     const old_value = old_op.getResult(0).?;
950 
951     var new_state = ir.Operation.State.init("test.new_value", .unknown);
952     new_state.addTypes(&.{value_type});
953     const new_op = try ctx.createOperation(new_state);
954     const new_value = new_op.getResult(0).?;
955 
956     var body = ir.context.initRegion(&ctx);
957     defer body.deinit();
958     const block = try body.addBlock();
959 
960     var user_state = ir.Operation.State.init("test.nested_user", .unknown);
961     user_state.addOperands(&.{old_value});
962     const user = try ctx.createOperation(user_state);
963     try block.addOperation(user);
964 
965     var root_state = ir.Operation.State.init("test.rewrite_root", .unknown);
966     root_state.addRegionBodies(&.{&body});
967     const root = try ctx.createOperation(root_state);
968 
969     var rewriter = PatternRewriter.init(allocator, &ctx);
970     defer rewriter.deinit();
971 
972     try rewriter.replaceAllUsesWith(old_value, new_value);
973     rewriter.finalize(root);
974 
975     try testing.expect(user.operands.items[0].value == new_value);
976 }