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 = &registry,
459     };
460     try registry.installTestUtilities(&utility_ctx);
461     try std.testing.expectEqual(@as(usize, 1), fake_test_utilities_installed);
462 }