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 }