lib/choir/src/core/builder.zig

daab053ee43316e1809a84551d573ddd1e5bf3d2

  1 const std = @import("std");
  2 const Operation = @import("operation/root.zig").Operation;
  3 const Block = @import("block.zig").Block;
  4 const Region = @import("region.zig").Region;
  5 const Type = @import("type.zig").Type;
  6 const Location = @import("location.zig").Location;
  7 const context_mod = @import("context/root.zig");
  8 const Context = context_mod.Context;
  9 
 10 pub const OperationBuilder = struct {
 11     ir_ctx: *Context,
 12     insertion_block: ?*Block,
 13     insertion_before: ?*Operation,
 14     listener: ?Listener,
 15 
 16     pub const Listener = struct {
 17         context: ?*anyopaque = null,
 18         notify_operation_inserted: ?*const fn (?*anyopaque, *Operation, InsertPoint) anyerror!void = null,
 19         notify_block_inserted: ?*const fn (?*anyopaque, *Block) anyerror!void = null,
 20 
 21         pub fn notifyOperationInserted(self: Listener, op: *Operation, previous: InsertPoint) !void {
 22             if (self.notify_operation_inserted) |notify| {
 23                 try notify(self.context, op, previous);
 24             }
 25         }
 26 
 27         pub fn notifyBlockInserted(self: Listener, block: *Block) !void {
 28             if (self.notify_block_inserted) |notify| {
 29                 try notify(self.context, block);
 30             }
 31         }
 32     };
 33 
 34     pub const InsertPoint = struct {
 35         block: ?*Block = null,
 36         before: ?*Operation = null,
 37 
 38         pub fn isSet(self: InsertPoint) bool {
 39             return self.block != null;
 40         }
 41     };
 42 
 43     pub const InsertionGuard = struct {
 44         builder: ?*OperationBuilder,
 45         insert_point: InsertPoint,
 46 
 47         pub fn init(builder: *OperationBuilder) InsertionGuard {
 48             return .{
 49                 .builder = builder,
 50                 .insert_point = builder.saveInsertionPoint(),
 51             };
 52         }
 53 
 54         pub fn deinit(self: *InsertionGuard) void {
 55             const builder = self.builder orelse return;
 56             builder.restoreInsertionPoint(self.insert_point);
 57             self.builder = null;
 58         }
 59     };
 60 
 61     pub fn init(ir_ctx: *Context) OperationBuilder {
 62         return .{
 63             .ir_ctx = ir_ctx,
 64             .insertion_block = null,
 65             .insertion_before = null,
 66             .listener = null,
 67         };
 68     }
 69 
 70     pub fn setListener(self: *OperationBuilder, listener: ?Listener) void {
 71         self.listener = listener;
 72     }
 73 
 74     pub fn getListener(self: *const OperationBuilder) ?Listener {
 75         return self.listener;
 76     }
 77 
 78     pub fn clearInsertionPoint(self: *OperationBuilder) void {
 79         self.insertion_block = null;
 80         self.insertion_before = null;
 81     }
 82 
 83     pub fn hasInsertionPoint(self: *const OperationBuilder) bool {
 84         return self.insertion_block != null;
 85     }
 86 
 87     pub fn saveInsertionPoint(self: *const OperationBuilder) InsertPoint {
 88         return .{
 89             .block = self.insertion_block,
 90             .before = self.insertion_before,
 91         };
 92     }
 93 
 94     pub fn restoreInsertionPoint(self: *OperationBuilder, insert_point: InsertPoint) void {
 95         if (insert_point.isSet()) {
 96             self.insertion_block = insert_point.block;
 97             self.insertion_before = insert_point.before;
 98             return;
 99         }
100         self.clearInsertionPoint();
101     }
102 
103     pub fn insertionGuard(self: *OperationBuilder) InsertionGuard {
104         return InsertionGuard.init(self);
105     }
106 
107     pub fn setInsertionPoint(self: *OperationBuilder, block: *Block) void {
108         self.setInsertionPointToEnd(block);
109     }
110 
111     pub fn setInsertionPointToEnd(self: *OperationBuilder, block: *Block) void {
112         self.insertion_block = block;
113         self.insertion_before = null;
114     }
115 
116     pub fn setInsertionPointBefore(self: *OperationBuilder, op: *Operation) void {
117         self.insertion_block = op.getBlock();
118         self.insertion_before = op;
119     }
120 
121     pub fn setInsertionPointAfter(self: *OperationBuilder, op: *Operation) void {
122         if (op.next_op) |next| {
123             self.setInsertionPointBefore(next);
124             return;
125         }
126         if (op.getBlock()) |block| {
127             self.setInsertionPointToEnd(block);
128             return;
129         }
130         self.clearInsertionPoint();
131     }
132 
133     pub fn getInsertionBlock(self: *const OperationBuilder) ?*Block {
134         return self.insertion_block;
135     }
136 
137     pub fn getInsertionBefore(self: *const OperationBuilder) ?*Operation {
138         return self.insertion_before;
139     }
140 
141     pub fn createBlock(self: *OperationBuilder, region: *Region, arg_types: []const Type, locs: []const Location) !*Block {
142         if (arg_types.len != locs.len) return error.BlockArgumentLocationMismatch;
143 
144         const previous = self.saveInsertionPoint();
145         errdefer self.restoreInsertionPoint(previous);
146         const block = try region.addBlock();
147         errdefer std.debug.assert(region.eraseBlock(block));
148 
149         for (arg_types, locs) |arg_type, loc| {
150             _ = try block.addArgument(arg_type, loc);
151         }
152 
153         self.setInsertionPointToEnd(block);
154         if (self.listener) |listener| {
155             try listener.notifyBlockInserted(block);
156         }
157         return block;
158     }
159 
160     pub fn createBlockWithLoc(self: *OperationBuilder, region: *Region, arg_types: []const Type, loc: Location) !*Block {
161         const previous = self.saveInsertionPoint();
162         errdefer self.restoreInsertionPoint(previous);
163         const block = try region.addBlock();
164         errdefer std.debug.assert(region.eraseBlock(block));
165 
166         for (arg_types) |arg_type| {
167             _ = try block.addArgument(arg_type, loc);
168         }
169 
170         self.setInsertionPointToEnd(block);
171         if (self.listener) |listener| {
172             try listener.notifyBlockInserted(block);
173         }
174         return block;
175     }
176 
177     pub fn insert(self: *OperationBuilder, op: *Operation) !*Operation {
178         var inserted = false;
179         if (self.insertion_block) |block| {
180             if (self.insertion_before) |before| {
181                 if (before.getBlock() == block) {
182                     try block.insertBefore(op, before);
183                 } else {
184                     try block.addOperation(op);
185                 }
186             } else {
187                 try block.addOperation(op);
188             }
189             inserted = true;
190         }
191 
192         if (inserted) {
193             if (self.listener) |listener| {
194                 listener.notifyOperationInserted(op, .{}) catch |err| {
195                     if (op.getBlock()) |block| block.detachOperation(op);
196                     return err;
197                 };
198             }
199         }
200 
201         return op;
202     }
203 
204     pub fn create(self: *OperationBuilder, state: Operation.State) !*Operation {
205         const op = try self.ir_ctx.createOperation(state);
206         errdefer if (op.getBlock() == null) op.erase();
207         return try self.insert(op);
208     }
209 };
210 
211 fn expectBlockOperationNames(block: *Block, expected: []const []const u8) !void {
212     var iter = block.getOperations();
213     for (expected) |name| {
214         const op = iter.next() orelse return error.TestExpectedOperation;
215         try std.testing.expectEqualStrings(name, op.name.name);
216     }
217     try std.testing.expect(iter.next() == null);
218 }
219 
220 const BuilderListenerRecorder = struct {
221     inserted_ops: usize = 0,
222     inserted_blocks: usize = 0,
223     last_op: ?*Operation = null,
224     last_block: ?*Block = null,
225     last_previous: OperationBuilder.InsertPoint = .{},
226 
227     fn notifyOperationInserted(listener_context: ?*anyopaque, op: *Operation, previous: OperationBuilder.InsertPoint) !void {
228         const recorder = try contextAsRecorder(listener_context);
229         recorder.inserted_ops += 1;
230         recorder.last_op = op;
231         recorder.last_previous = previous;
232     }
233 
234     fn notifyBlockInserted(listener_context: ?*anyopaque, block: *Block) !void {
235         const recorder = try contextAsRecorder(listener_context);
236         recorder.inserted_blocks += 1;
237         recorder.last_block = block;
238     }
239 
240     fn contextAsRecorder(listener_context: ?*anyopaque) !*BuilderListenerRecorder {
241         return @ptrCast(@alignCast(listener_context orelse return error.MissingListenerContext));
242     }
243 };
244 
245 const RejectingBuilderListener = struct {
246     fn notifyOperationInserted(_: ?*anyopaque, _: *Operation, _: OperationBuilder.InsertPoint) !void {
247         return error.InsertionRejected;
248     }
249 
250     fn notifyBlockInserted(_: ?*anyopaque, _: *Block) !void {
251         return error.InsertionRejected;
252     }
253 };
254 
255 test "OperationBuilder insertion guard restores saved point" {
256     const testing = std.testing;
257 
258     var ctx = try Context.init(testing.allocator, Context.Limits.testing);
259     defer ctx.deinit(testing.allocator);
260     try ctx.allowUnregistered();
261 
262     var block = Block.init(testing.allocator);
263     defer block.deinit();
264 
265     var builder = OperationBuilder.init(&ctx);
266     builder.setInsertionPointToEnd(&block);
267     _ = try builder.create(Operation.State.init("test.builder.first", .unknown));
268     const tail = try builder.create(Operation.State.init("test.builder.tail", .unknown));
269 
270     builder.setInsertionPointBefore(tail);
271     _ = try builder.create(Operation.State.init("test.builder.before_tail", .unknown));
272 
273     {
274         var guard = builder.insertionGuard();
275         defer guard.deinit();
276         builder.setInsertionPointBefore(tail.prev_op.?);
277         _ = try builder.create(Operation.State.init("test.builder.guarded", .unknown));
278     }
279 
280     _ = try builder.create(Operation.State.init("test.builder.restored", .unknown));
281 
282     try expectBlockOperationNames(&block, &.{
283         "test.builder.first",
284         "test.builder.guarded",
285         "test.builder.before_tail",
286         "test.builder.restored",
287         "test.builder.tail",
288     });
289 }
290 
291 test "OperationBuilder listener records inserted operations and blocks" {
292     const testing = std.testing;
293 
294     var ctx = try Context.init(testing.allocator, Context.Limits.testing);
295     defer ctx.deinit(testing.allocator);
296     try ctx.allowUnregistered();
297 
298     var region = context_mod.initRegion(&ctx);
299     defer region.deinit();
300 
301     var recorder = BuilderListenerRecorder{};
302     var builder = OperationBuilder.init(&ctx);
303     builder.setListener(.{
304         .context = &recorder,
305         .notify_operation_inserted = BuilderListenerRecorder.notifyOperationInserted,
306         .notify_block_inserted = BuilderListenerRecorder.notifyBlockInserted,
307     });
308 
309     const detached = try builder.create(Operation.State.init("test.builder.detached", .unknown));
310     try testing.expect(detached.parent_block == null);
311     try testing.expectEqual(@as(usize, 0), recorder.inserted_ops);
312     try testing.expectEqual(@as(usize, 0), recorder.inserted_blocks);
313 
314     const block = try builder.createBlock(&region, &.{}, &.{});
315     try testing.expectEqual(@as(usize, 1), recorder.inserted_blocks);
316     try testing.expect(recorder.last_block == block);
317 
318     const inserted = try builder.create(Operation.State.init("test.builder.inserted", .unknown));
319     try testing.expect(inserted.parent_block == block);
320     try testing.expectEqual(@as(usize, 1), recorder.inserted_ops);
321     try testing.expect(recorder.last_op == inserted);
322     try testing.expect(!recorder.last_previous.isSet());
323 
324     builder.setListener(null);
325     _ = try builder.create(Operation.State.init("test.builder.unobserved", .unknown));
326     try testing.expectEqual(@as(usize, 1), recorder.inserted_ops);
327 }
328 
329 test "OperationBuilder listener failure rolls insertion back transactionally" {
330     const testing = std.testing;
331 
332     var ctx = try Context.init(testing.allocator, Context.Limits.testing);
333     defer ctx.deinit(testing.allocator);
334     try ctx.allowUnregistered();
335 
336     var block = Block.init(testing.allocator);
337     defer block.deinit();
338 
339     var builder = OperationBuilder.init(&ctx);
340     builder.setInsertionPointToEnd(&block);
341     builder.setListener(.{
342         .notify_operation_inserted = RejectingBuilderListener.notifyOperationInserted,
343     });
344 
345     const detached = try ctx.createOperation(Operation.State.init("test.builder.rejected_insert", .unknown));
346     defer detached.erase();
347     try testing.expectError(error.InsertionRejected, builder.insert(detached));
348     try testing.expect(detached.getBlock() == null);
349     try testing.expect(block.operations.head == null);
350     try testing.expect(block.operations.tail == null);
351 
352     const operations = ctx.operationCount();
353     try testing.expectError(
354         error.InsertionRejected,
355         builder.create(Operation.State.init("test.builder.rejected_create", .unknown)),
356     );
357     try testing.expectEqual(operations, ctx.operationCount());
358     try testing.expect(block.operations.head == null);
359     try testing.expect(block.operations.tail == null);
360 }
361 
362 test "OperationBuilder block listener failure restores its insertion point" {
363     const testing = std.testing;
364 
365     var ctx = try Context.init(testing.allocator, Context.Limits.testing);
366     defer ctx.deinit(testing.allocator);
367     try ctx.allowUnregistered();
368 
369     var original = Block.init(testing.allocator);
370     defer original.deinit();
371     var region = context_mod.initRegion(&ctx);
372     defer region.deinit();
373 
374     var builder = OperationBuilder.init(&ctx);
375     builder.setInsertionPointToEnd(&original);
376     builder.setListener(.{
377         .notify_block_inserted = RejectingBuilderListener.notifyBlockInserted,
378     });
379 
380     try testing.expectError(error.InsertionRejected, builder.createBlock(&region, &.{}, &.{}));
381     try testing.expect(region.empty());
382     try testing.expectEqual(&original, builder.getInsertionBlock().?);
383     try testing.expect(builder.getInsertionBefore() == null);
384 }
385 
386 test "OperationBuilder createBlock adds arguments and sets insertion point" {
387     const testing = std.testing;
388 
389     var ctx = try Context.init(testing.allocator, Context.Limits.testing);
390     defer ctx.deinit(testing.allocator);
391     try ctx.allowUnregistered();
392 
393     var region = context_mod.initRegion(&ctx);
394     defer region.deinit();
395 
396     const i32_type = try ctx.getDialectTypeFromName("test.i32");
397     const i64_type = try ctx.getDialectTypeFromName("test.i64");
398     var builder = OperationBuilder.init(&ctx);
399 
400     const block = try builder.createBlock(&region, &.{ i32_type, i64_type }, &.{ .unknown, .unknown });
401     try testing.expect(region.getEntryBlock() == block);
402     try testing.expect(builder.getInsertionBlock() == block);
403     try testing.expect(builder.getInsertionBefore() == null);
404     try testing.expectEqual(@as(usize, 2), block.getNumArguments());
405     try testing.expect(block.getArgument(0).?.type.eql(i32_type));
406     try testing.expect(block.getArgument(1).?.type.eql(i64_type));
407 
408     const op = try builder.create(Operation.State.init("test.builder.created_in_block", .unknown));
409     try testing.expect(op.parent_block == block);
410     try testing.expect(block.operations.head == @as(?*anyopaque, op));
411 }
412 
413 test "OperationBuilder createBlock rejects mismatched argument locations" {
414     const testing = std.testing;
415 
416     var ctx = try Context.init(testing.allocator, Context.Limits.testing);
417     defer ctx.deinit(testing.allocator);
418     try ctx.allowUnregistered();
419 
420     var region = context_mod.initRegion(&ctx);
421     defer region.deinit();
422 
423     const i32_type = try ctx.getDialectTypeFromName("test.i32");
424     var builder = OperationBuilder.init(&ctx);
425 
426     try testing.expectError(
427         error.BlockArgumentLocationMismatch,
428         builder.createBlock(&region, &.{i32_type}, &.{}),
429     );
430     try testing.expect(region.empty());
431     try testing.expect(!builder.hasInsertionPoint());
432 }
433 
434 test "OperationBuilder createBlockWithLoc adds uniform-location arguments" {
435     const testing = std.testing;
436 
437     var ctx = try Context.init(testing.allocator, Context.Limits.testing);
438     defer ctx.deinit(testing.allocator);
439     try ctx.allowUnregistered();
440 
441     var region = context_mod.initRegion(&ctx);
442     defer region.deinit();
443 
444     const i1_type = try ctx.getDialectTypeFromName("test.i1");
445     const i8_type = try ctx.getDialectTypeFromName("test.i8");
446     var builder = OperationBuilder.init(&ctx);
447 
448     const block = try builder.createBlockWithLoc(&region, &.{ i1_type, i8_type }, .unknown);
449     try testing.expect(region.getEntryBlock() == block);
450     try testing.expect(builder.getInsertionBlock() == block);
451     try testing.expectEqual(@as(usize, 2), block.getNumArguments());
452     try testing.expect(block.getArgument(0).?.type.eql(i1_type));
453     try testing.expect(block.getArgument(1).?.type.eql(i8_type));
454 }
455 
456 test "OperationBuilder save restore and after insertion point" {
457     const testing = std.testing;
458 
459     var ctx = try Context.init(testing.allocator, Context.Limits.testing);
460     defer ctx.deinit(testing.allocator);
461     try ctx.allowUnregistered();
462 
463     var block = Block.init(testing.allocator);
464     defer block.deinit();
465 
466     var builder = OperationBuilder.init(&ctx);
467     try testing.expect(!builder.hasInsertionPoint());
468 
469     builder.setInsertionPointToEnd(&block);
470     const first = try builder.create(Operation.State.init("test.builder.after_first", .unknown));
471     const last = try builder.create(Operation.State.init("test.builder.after_last", .unknown));
472     const end_point = builder.saveInsertionPoint();
473 
474     builder.setInsertionPointAfter(first);
475     _ = try builder.create(Operation.State.init("test.builder.after_middle", .unknown));
476 
477     builder.restoreInsertionPoint(end_point);
478     _ = try builder.create(Operation.State.init("test.builder.after_end", .unknown));
479 
480     builder.clearInsertionPoint();
481     try testing.expect(!builder.hasInsertionPoint());
482     builder.restoreInsertionPoint(end_point);
483     try testing.expect(builder.hasInsertionPoint());
484 
485     _ = last;
486     try expectBlockOperationNames(&block, &.{
487         "test.builder.after_first",
488         "test.builder.after_middle",
489         "test.builder.after_last",
490         "test.builder.after_end",
491     });
492 }