lib/choir/src/extensions.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const ir = @import("core/root.zig");
3 const rewrite = ir.rewrite;
4 const core_dialects = @import("core/root.zig").dialects;
5 const passes = @import("passes/root.zig");
6 const backends = @import("backends/root.zig");
7 const translation_mod = backends.translation;
8
9 const Allocator = std.mem.Allocator;
10
11 pub const RegisterError = Allocator.Error || ir.Context.Error || error{
12 DialectLoadCycle,
13 };
14
15 pub const DialectRegistration = struct {
16 name: []const u8,
17 load: core_dialects.DialectLoaderFn,
18 is_backend: bool = false,
19 };
20
21 pub const DialectExtensionRegistration = struct {
22 dialect_name: []const u8,
23 extend: core_dialects.DialectExtensionFn,
24 };
25
26 pub const PassRegistration = passes.PassRegistration;
27
28 pub const ConfigureConversionTargetFn = *const fn (*passes.ConversionTarget) anyerror!void;
29 pub const PopulateRewritePatternsFn = *const fn (*rewrite.RewritePatternSet) anyerror!void;
30 pub const ConfigureTypeConverterFn = *const fn (*rewrite.TypeConverter) anyerror!void;
31
32 pub const ConversionRegistration = struct {
33 name: []const u8,
34 configure_target: ?ConfigureConversionTargetFn = null,
35 populate_patterns: ?PopulateRewritePatternsFn = null,
36 configure_type_converter: ?ConfigureTypeConverterFn = null,
37 };
38
39 pub const BackendTranslationRegistration = translation_mod.Registration;
40
41 pub const TestUtilityContext = struct {
42 allocator: Allocator,
43 ir_ctx: *ir.Context,
44 registry: *ExtensionRegistry,
45 };
46
47 pub const InstallTestUtilityFn = *const fn (*TestUtilityContext) anyerror!void;
48
49 pub const TestUtilityRegistration = struct {
50 name: []const u8,
51 install: InstallTestUtilityFn,
52 };
53
54 pub const PackageExtension = struct {
55 name: []const u8,
56 dialects: []const DialectRegistration = &.{},
57 dialect_extensions: []const DialectExtensionRegistration = &.{},
58 passes: []const PassRegistration = &.{},
59 pipelines: []const passes.PipelineRegistration = &.{},
60 conversions: []const ConversionRegistration = &.{},
61 instrumentations: []const passes.PassInstrumentation = &.{},
62 backend_translations: []const BackendTranslationRegistration = &.{},
63 test_utilities: []const TestUtilityRegistration = &.{},
64
65 pub fn registerContext(self: PackageExtension, ctx: *ir.Context) anyerror!void {
66 for (self.dialects) |dialect| {
67 try ctx.registerDialectLoader(dialect.name, dialect.load);
68 if (dialect.is_backend) {
69 try ctx.registerBackendDialect(dialect.name);
70 }
71 }
72
73 for (self.dialect_extensions) |extension| {
74 try ctx.addDialectExtension(extension.dialect_name, extension.extend);
75 }
76 }
77 };
78
79 pub const choir_optimization_package_extension = PackageExtension{
80 .name = "choir-optimization",
81 .passes = &passes.optimization_pass_registrations,
82 .pipelines = &passes.optimization_pipeline_registrations,
83 };
84
85 pub const ExtensionRegistry = struct {
86 allocator: Allocator,
87 packages: std.ArrayListUnmanaged([]const u8) = .empty,
88 passes: std.ArrayListUnmanaged(PassRegistration) = .empty,
89 pipelines: std.ArrayListUnmanaged(passes.PipelineRegistration) = .empty,
90 conversions: std.ArrayListUnmanaged(ConversionRegistration) = .empty,
91 instrumentations: std.ArrayListUnmanaged(passes.PassInstrumentation) = .empty,
92 backend_translations: std.ArrayListUnmanaged(BackendTranslationRegistration) = .empty,
93 test_utilities: std.ArrayListUnmanaged(TestUtilityRegistration) = .empty,
94
95 pub fn init(allocator: Allocator) ExtensionRegistry {
96 return .{ .allocator = allocator };
97 }
98
99 pub fn deinit(self: *ExtensionRegistry) void {
100 self.packages.deinit(self.allocator);
101 self.passes.deinit(self.allocator);
102 self.pipelines.deinit(self.allocator);
103 self.conversions.deinit(self.allocator);
104 self.instrumentations.deinit(self.allocator);
105 self.backend_translations.deinit(self.allocator);
106 self.test_utilities.deinit(self.allocator);
107 }
108
109 pub fn registerPackage(self: *ExtensionRegistry, ctx: *ir.Context, extension: PackageExtension) anyerror!void {
110 try self.reserve(extension);
111 try extension.registerContext(ctx);
112
113 self.packages.appendAssumeCapacity(extension.name);
114 appendSliceAssumeCapacity(PassRegistration, &self.passes, extension.passes);
115 appendSliceAssumeCapacity(passes.PipelineRegistration, &self.pipelines, extension.pipelines);
116 appendSliceAssumeCapacity(ConversionRegistration, &self.conversions, extension.conversions);
117 appendSliceAssumeCapacity(passes.PassInstrumentation, &self.instrumentations, extension.instrumentations);
118 appendSliceAssumeCapacity(BackendTranslationRegistration, &self.backend_translations, extension.backend_translations);
119 appendSliceAssumeCapacity(TestUtilityRegistration, &self.test_utilities, extension.test_utilities);
120 }
121
122 pub fn addPassesTo(self: *ExtensionRegistry, manager: *passes.PassManager) anyerror!void {
123 for (self.passes.items) |registration| {
124 if (registration.target_op_name) |target| {
125 const nested = try manager.nest(target);
126 try nested.addPass(registration.pass);
127 } else {
128 try manager.addPass(registration.pass);
129 }
130 }
131 }
132
133 pub fn registerPassEntriesTo(self: *ExtensionRegistry, pass_registry: *passes.PassRegistry) anyerror!void {
134 for (self.passes.items) |registration| {
135 try pass_registry.registerPass(registration);
136 }
137 for (self.pipelines.items) |registration| {
138 try pass_registry.registerPipeline(registration);
139 }
140 }
141
142 pub fn registerPipelinesTo(self: *ExtensionRegistry, pipeline_registry: *passes.PipelineRegistry) anyerror!void {
143 for (self.pipelines.items) |registration| {
144 try pipeline_registry.registerPipeline(registration);
145 }
146 }
147
148 pub fn addPipelineTo(
149 self: *ExtensionRegistry,
150 name: []const u8,
151 manager: *passes.PassManager,
152 ) anyerror!void {
153 const registration = self.lookupPipeline(name) orelse return error.UnknownPipeline;
154 try registration.addTo(&manager.root);
155 }
156
157 pub fn addInstrumentationsTo(self: *ExtensionRegistry, manager: *passes.PassManager) anyerror!void {
158 for (self.instrumentations.items) |instrumentation| {
159 try manager.addInstrumentation(instrumentation);
160 }
161 }
162
163 pub fn configureConversion(
164 self: *ExtensionRegistry,
165 target: *passes.ConversionTarget,
166 patterns: *rewrite.RewritePatternSet,
167 type_converter: ?*rewrite.TypeConverter,
168 ) anyerror!void {
169 for (self.conversions.items) |registration| {
170 if (registration.configure_target) |configure| {
171 try configure(target);
172 }
173 if (registration.populate_patterns) |populate| {
174 try populate(patterns);
175 }
176 if (type_converter) |converter| {
177 if (registration.configure_type_converter) |configure| {
178 try configure(converter);
179 }
180 }
181 }
182 }
183
184 pub fn registerBackendTranslationsTo(self: *ExtensionRegistry, registry: *translation_mod.Registry) anyerror!void {
185 try registry.registerAll(self.backend_translations.items);
186 }
187
188 pub fn installTestUtilities(self: *ExtensionRegistry, context: *TestUtilityContext) anyerror!void {
189 for (self.test_utilities.items) |registration| {
190 try registration.install(context);
191 }
192 }
193
194 fn lookupPipeline(self: *const ExtensionRegistry, name: []const u8) ?passes.PipelineRegistration {
195 for (self.pipelines.items) |registration| {
196 if (std.mem.eql(u8, registration.name, name)) return registration;
197 }
198 return null;
199 }
200
201 fn reserve(self: *ExtensionRegistry, extension: PackageExtension) Allocator.Error!void {
202 try self.packages.ensureTotalCapacity(self.allocator, self.packages.items.len + 1);
203 try self.passes.ensureTotalCapacity(self.allocator, self.passes.items.len + extension.passes.len);
204 try self.pipelines.ensureTotalCapacity(self.allocator, self.pipelines.items.len + extension.pipelines.len);
205 try self.conversions.ensureTotalCapacity(self.allocator, self.conversions.items.len + extension.conversions.len);
206 try self.instrumentations.ensureTotalCapacity(self.allocator, self.instrumentations.items.len + extension.instrumentations.len);
207 try self.backend_translations.ensureTotalCapacity(self.allocator, self.backend_translations.items.len + extension.backend_translations.len);
208 try self.test_utilities.ensureTotalCapacity(self.allocator, self.test_utilities.items.len + extension.test_utilities.len);
209 }
210 };
211
212 fn appendSliceAssumeCapacity(comptime T: type, list: *std.ArrayListUnmanaged(T), items: []const T) void {
213 for (items) |item| {
214 list.appendAssumeCapacity(item);
215 }
216 }
217
218 const fake_dialect_name = "external_fake";
219 const fake_op_name = "external_fake.op";
220
221 var fake_pass_runs: usize = 0;
222 var fake_pipeline_pass_runs: usize = 0;
223 var fake_before_pass_runs: usize = 0;
224 var fake_verify_runs: usize = 0;
225 var fake_test_utilities_installed: usize = 0;
226 var fake_translation_runs: usize = 0;
227
228 fn resetFakeCounters() void {
229 fake_pass_runs = 0;
230 fake_pipeline_pass_runs = 0;
231 fake_before_pass_runs = 0;
232 fake_verify_runs = 0;
233 fake_test_utilities_installed = 0;
234 fake_translation_runs = 0;
235 }
236
237 test "Choir optimization package extension registers pass entries" {
238 const allocator = std.testing.allocator;
239
240 var ctx = try ir.Context.init(allocator, ir.Context.Limits.testing);
241 defer ctx.deinit(allocator);
242
243 var registry = ExtensionRegistry.init(allocator);
244 defer registry.deinit();
245 try registry.registerPackage(&ctx, choir_optimization_package_extension);
246
247 try std.testing.expectEqual(@as(usize, 1), registry.packages.items.len);
248 try std.testing.expectEqual(@as(usize, passes.optimization_pass_registrations.len), registry.passes.items.len);
249 try std.testing.expectEqual(@as(usize, passes.optimization_pipeline_registrations.len), registry.pipelines.items.len);
250
251 var pass_registry = passes.PassRegistry.init(allocator);
252 defer pass_registry.deinit();
253 try registry.registerPassEntriesTo(&pass_registry);
254
255 try std.testing.expect(pass_registry.lookupPass(passes.canonicalization_pass_name) != null);
256 try std.testing.expect(pass_registry.lookupPass(passes.sparse_conditional_constant_propagation_pass_name) != null);
257 try std.testing.expect(pass_registry.lookupPass(passes.common_subexpression_elimination_pass_name) != null);
258 try std.testing.expect(pass_registry.lookupPass(passes.dead_store_elimination_pass_name) != null);
259 try std.testing.expect(pass_registry.lookupPipeline(passes.default_optimization_pipeline_name) != null);
260
261 var manager = passes.PassManager.init(allocator);
262 defer manager.deinit();
263 try passes.parsePassPipeline(&pass_registry, passes.default_optimization_pipeline_name, &manager);
264 try std.testing.expectEqual(@as(usize, 11), manager.root.pipeline.items.len);
265 }
266
267 fn loadFakeDialect(ctx: *ir.Context) !void {
268 _ = try ctx.registerOperation(fake_op_name, .{});
269 _ = try ctx.registerType("external_fake.type");
270 }
271
272 fn verifyFakeOp(_: *const anyopaque) anyerror!void {
273 fake_verify_runs += 1;
274 }
275
276 fn extendFakeDialect(ctx: *ir.Context) !void {
277 try ctx.registerOperationInterfaceExternal(fake_op_name, ir.VerifyOpInterface.entryFor(verifyFakeOp));
278 }
279
280 fn runFakePass(ctx: *passes.PassContext) passes.PassResult {
281 fake_pass_runs += 1;
282 ctx.markModified();
283 return .success;
284 }
285
286 fn runFakePipelinePass(ctx: *passes.PassContext) passes.PassResult {
287 fake_pipeline_pass_runs += 1;
288 ctx.preserveAllAnalyses();
289 return .success;
290 }
291
292 fn buildFakePipeline(manager: *passes.OpPassManager) anyerror!void {
293 try manager.addPass(.{
294 .name = "external-fake-pipeline-pass",
295 .description = "fake external package pipeline pass",
296 .run_fn = runFakePipelinePass,
297 });
298 }
299
300 fn beforeFakePass(_: ?*anyopaque, _: passes.PassInfo) void {
301 fake_before_pass_runs += 1;
302 }
303
304 fn configureFakeTarget(target: *passes.ConversionTarget) anyerror!void {
305 try target.addLegalDialect(fake_dialect_name);
306 }
307
308 fn rewriteFake(_: *ir.Operation, _: *rewrite.PatternRewriter) rewrite.PatternResult {
309 return .success;
310 }
311
312 fn populateFakePatterns(patterns: *rewrite.RewritePatternSet) anyerror!void {
313 try patterns.add(rewrite.RewritePattern.init(.{
314 .name = "external-fake-rewrite",
315 .root_op_name = fake_op_name,
316 .benefit = 3,
317 .products = .none,
318 }, rewriteFake));
319 }
320
321 fn installFakeTestUtility(_: *TestUtilityContext) anyerror!void {
322 fake_test_utilities_installed += 1;
323 }
324
325 const FakeBackendTranslation = struct {
326 pub fn emitBytes(result_allocator: Allocator, request: translation_mod.Request) anyerror![]u8 {
327 _ = request;
328 fake_translation_runs += 1;
329 return result_allocator.dupe(u8, "external-fake-bytes");
330 }
331 };
332
333 const fake_backend_translation = translation_mod.registration(FakeBackendTranslation, .{
334 .name = "external-fake-binary",
335 .target_dialect_name = fake_dialect_name,
336 .format_name = "fake-binary",
337 });
338
339 test "external package registers dialects passes conversions translations and test utilities" {
340 resetFakeCounters();
341
342 var ctx = try ir.Context.init(std.testing.allocator, ir.Context.Limits.testing);
343 defer ctx.deinit(std.testing.allocator);
344 try ctx.requireRegistered();
345
346 var registry = ExtensionRegistry.init(std.testing.allocator);
347 defer registry.deinit();
348
349 const fake_pass = passes.Pass{
350 .name = "external-fake-pass",
351 .description = "fake external package pass",
352 .run_fn = runFakePass,
353 };
354
355 const fake_extension = PackageExtension{
356 .name = "external-fake-package",
357 .dialects = &.{
358 .{ .name = fake_dialect_name, .load = loadFakeDialect },
359 },
360 .dialect_extensions = &.{
361 .{ .dialect_name = fake_dialect_name, .extend = extendFakeDialect },
362 },
363 .passes = &.{
364 .{
365 .name = "external-fake-pass",
366 .description = "fake external package pass",
367 .pass = fake_pass,
368 },
369 },
370 .pipelines = &.{
371 .{
372 .name = "external-fake-pipeline",
373 .description = "fake external package pipeline",
374 .build = buildFakePipeline,
375 },
376 },
377 .conversions = &.{
378 .{
379 .name = "external-fake-conversion",
380 .configure_target = configureFakeTarget,
381 .populate_patterns = populateFakePatterns,
382 },
383 },
384 .instrumentations = &.{
385 .{ .runBeforePass = beforeFakePass },
386 },
387 .backend_translations = &.{
388 fake_backend_translation,
389 },
390 .test_utilities = &.{
391 .{ .name = "external-fake-test-utility", .install = installFakeTestUtility },
392 },
393 };
394
395 try registry.registerPackage(&ctx, fake_extension);
396 try std.testing.expectEqual(@as(usize, 1), registry.packages.items.len);
397 try std.testing.expectEqual(@as(usize, 1), registry.pipelines.items.len);
398 try std.testing.expectEqual(@as(usize, 1), registry.backend_translations.items.len);
399 try std.testing.expect(ctx.lookupOperation(fake_op_name) == null);
400
401 _ = try ctx.getOrLoadDialect(fake_dialect_name);
402 try std.testing.expect(ctx.lookupOperation(fake_op_name) != null);
403 try std.testing.expect(ctx.lookupType("external_fake.type") != null);
404
405 const op = try ctx.createOperation(ir.Operation.State.init(fake_op_name, .unknown));
406 try ir.verifyOperation(op, ir.verify.default_options);
407 try std.testing.expectEqual(@as(usize, 1), fake_verify_runs);
408
409 var manager = passes.PassManager.init(std.testing.allocator);
410 defer manager.deinit();
411 try registry.addInstrumentationsTo(&manager);
412 try registry.addPassesTo(&manager);
413
414 try std.testing.expectEqual(passes.PassResult.success, manager.run(op, &ctx));
415 try std.testing.expectEqual(@as(usize, 1), fake_before_pass_runs);
416 try std.testing.expectEqual(@as(usize, 1), fake_pass_runs);
417
418 var pipeline_registry = passes.PipelineRegistry.init(std.testing.allocator);
419 defer pipeline_registry.deinit();
420 try registry.registerPipelinesTo(&pipeline_registry);
421 try std.testing.expect(pipeline_registry.lookup("external-fake-pipeline") != null);
422
423 var pass_registry = passes.PassRegistry.init(std.testing.allocator);
424 defer pass_registry.deinit();
425 try registry.registerPassEntriesTo(&pass_registry);
426 try std.testing.expect(pass_registry.lookupPass("external-fake-pass") != null);
427 try std.testing.expect(pass_registry.lookupPipeline("external-fake-pipeline") != null);
428
429 var pipeline_manager = passes.PassManager.init(std.testing.allocator);
430 defer pipeline_manager.deinit();
431 try registry.addPipelineTo("external-fake-pipeline", &pipeline_manager);
432 try std.testing.expectEqual(passes.PassResult.success, pipeline_manager.run(op, &ctx));
433 try std.testing.expectEqual(@as(usize, 1), fake_pipeline_pass_runs);
434
435 var target = passes.ConversionTarget.init(std.testing.allocator);
436 defer target.deinit();
437 var patterns = rewrite.RewritePatternSet.init(std.testing.allocator);
438 defer patterns.deinit();
439 var type_converter = rewrite.TypeConverter.init(std.testing.allocator);
440 defer type_converter.deinit();
441 try registry.configureConversion(&target, &patterns, &type_converter);
442 try std.testing.expect(target.isLegal(op));
443 try std.testing.expectEqual(@as(usize, 1), patterns.patterns.items.len);
444
445 var translation_registry = translation_mod.Registry.init(std.testing.allocator);
446 defer translation_registry.deinit();
447 try registry.registerBackendTranslationsTo(&translation_registry);
448 try std.testing.expect(translation_registry.lookup("external-fake-binary") != null);
449
450 const translated = try translation_registry.emitBytes(std.testing.allocator, fake_dialect_name, "fake-binary", .{ .module = op });
451 defer std.testing.allocator.free(translated);
452 try std.testing.expectEqualStrings("external-fake-bytes", translated);
453 try std.testing.expectEqual(@as(usize, 1), fake_translation_runs);
454
455 var utility_ctx = TestUtilityContext{
456 .allocator = std.testing.allocator,
457 .ir_ctx = &ctx,
458 .registry = ®istry,
459 };
460 try registry.installTestUtilities(&utility_ctx);
461 try std.testing.expectEqual(@as(usize, 1), fake_test_utilities_installed);
462 }