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(®ion, &.{}, &.{});
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(®ion, &.{}, &.{}));
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(®ion, &.{ 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(®ion, &.{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(®ion, &.{ 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 }