[Mlir-commits] [mlir] [mlir-c] Add ConversionTarget dynamic legality C API (PR #206161)
Maksim Levental
llvmlistbot at llvm.org
Sun Jun 28 21:20:37 PDT 2026
https://github.com/makslevental updated https://github.com/llvm/llvm-project/pull/206161
>From 94e26aa37f1aa914fff0d7ee088de3db04d7687f Mon Sep 17 00:00:00 2001
From: makslevental <maksim.levental at gmail.com>
Date: Fri, 26 Jun 2026 12:20:59 -0700
Subject: [PATCH] [mlir-c] Add ConversionTarget dynamic legality C API
Add mlirConversionTargetAddDynamicallyLegalOp,
mlirConversionTargetAddDynamicallyLegalDialect,
mlirConversionTargetMarkOpRecursivelyLegal, and
mlirConversionTargetMarkUnknownOpDynamicallyLegal to enable
per-instance legality callbacks from C.
---
mlir/include/mlir-c/Rewrite.h | 31 ++++++
mlir/lib/CAPI/Transforms/Rewrite.cpp | 49 +++++++++
mlir/test/CAPI/rewrite.c | 149 +++++++++++++++++++++++++++
3 files changed, 229 insertions(+)
diff --git a/mlir/include/mlir-c/Rewrite.h b/mlir/include/mlir-c/Rewrite.h
index ac243a4c9d8f9..322e7a3e4654f 100644
--- a/mlir/include/mlir-c/Rewrite.h
+++ b/mlir/include/mlir-c/Rewrite.h
@@ -533,6 +533,37 @@ MLIR_CAPI_EXPORTED void
mlirConversionTargetAddIllegalDialect(MlirConversionTarget target,
MlirStringRef dialectName);
+/// Callback for dynamic legality checks. Return true if the operation is
+/// legal, false if illegal.
+typedef bool (*MlirConversionTargetDynamicLegalityCallback)(MlirOperation op,
+ void *userData);
+
+/// Register the given operation as dynamically legal, with a callback to
+/// determine per-instance legality. The callback must not be NULL.
+MLIR_CAPI_EXPORTED void mlirConversionTargetAddDynamicallyLegalOp(
+ MlirConversionTarget target, MlirStringRef opName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData);
+
+/// Register the given dialect as dynamically legal, with a callback to
+/// determine per-instance legality for all operations in the dialect. The
+/// callback must not be NULL.
+MLIR_CAPI_EXPORTED void mlirConversionTargetAddDynamicallyLegalDialect(
+ MlirConversionTarget target, MlirStringRef dialectName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData);
+
+/// Mark the given operation as recursively legal. The optional callback (may
+/// be NULL) determines whether a specific instance is recursively legal; a NULL
+/// callback marks the operation as unconditionally recursively legal.
+MLIR_CAPI_EXPORTED void mlirConversionTargetMarkOpRecursivelyLegal(
+ MlirConversionTarget target, MlirStringRef opName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData);
+
+/// Mark unknown operations as dynamically legal, with a callback. The callback
+/// must not be NULL.
+MLIR_CAPI_EXPORTED void mlirConversionTargetMarkUnknownOpDynamicallyLegal(
+ MlirConversionTarget target,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData);
+
//===----------------------------------------------------------------------===//
/// TypeConverter API
//===----------------------------------------------------------------------===//
diff --git a/mlir/lib/CAPI/Transforms/Rewrite.cpp b/mlir/lib/CAPI/Transforms/Rewrite.cpp
index 56ce9212f4811..a38d2329f801e 100644
--- a/mlir/lib/CAPI/Transforms/Rewrite.cpp
+++ b/mlir/lib/CAPI/Transforms/Rewrite.cpp
@@ -23,6 +23,8 @@
#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
#include "mlir/Transforms/WalkPatternRewriteDriver.h"
+#include <cassert>
+
using namespace mlir;
//===----------------------------------------------------------------------===//
@@ -575,6 +577,53 @@ void mlirConversionTargetAddIllegalDialect(MlirConversionTarget target,
unwrap(target)->addIllegalDialect(unwrap(dialectName));
}
+void mlirConversionTargetAddDynamicallyLegalOp(
+ MlirConversionTarget target, MlirStringRef opName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
+ assert(callback && "expected non-null legality callback");
+ MLIRContext *ctx = &unwrap(target)->getContext();
+ OperationName name(unwrap(opName), ctx);
+ unwrap(target)->addDynamicallyLegalOp(
+ name, [callback, userData](Operation *op) -> std::optional<bool> {
+ return callback(wrap(op), userData);
+ });
+}
+
+void mlirConversionTargetAddDynamicallyLegalDialect(
+ MlirConversionTarget target, MlirStringRef dialectName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
+ assert(callback && "expected non-null legality callback");
+ unwrap(target)->addDynamicallyLegalDialect(
+ [callback, userData](Operation *op) -> std::optional<bool> {
+ return callback(wrap(op), userData);
+ },
+ unwrap(dialectName));
+}
+
+void mlirConversionTargetMarkOpRecursivelyLegal(
+ MlirConversionTarget target, MlirStringRef opName,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
+ MLIRContext *ctx = &unwrap(target)->getContext();
+ OperationName name(unwrap(opName), ctx);
+ ConversionTarget::DynamicLegalityCallbackFn fn;
+ if (callback) {
+ fn = [callback, userData](Operation *op) -> std::optional<bool> {
+ return callback(wrap(op), userData);
+ };
+ }
+ unwrap(target)->markOpRecursivelyLegal(name, fn);
+}
+
+void mlirConversionTargetMarkUnknownOpDynamicallyLegal(
+ MlirConversionTarget target,
+ MlirConversionTargetDynamicLegalityCallback callback, void *userData) {
+ assert(callback && "expected non-null legality callback");
+ unwrap(target)->markUnknownOpDynamicallyLegal(
+ [callback, userData](Operation *op) -> std::optional<bool> {
+ return callback(wrap(op), userData);
+ });
+}
+
//===----------------------------------------------------------------------===//
/// TypeConverter API
//===----------------------------------------------------------------------===//
diff --git a/mlir/test/CAPI/rewrite.c b/mlir/test/CAPI/rewrite.c
index 3809dd6a7843f..8c1103b8e6119 100644
--- a/mlir/test/CAPI/rewrite.c
+++ b/mlir/test/CAPI/rewrite.c
@@ -623,6 +623,154 @@ void testCloneWithMapping(MlirContext ctx) {
fprintf(stderr, "testCloneWithMapping: PASSED\n");
}
+static bool dynamicLegalityAlwaysLegal(MlirOperation op, void *userData) {
+ (void)op;
+ intptr_t *counter = (intptr_t *)userData;
+ (*counter)++;
+ return true;
+}
+
+static bool dynamicLegalityAlwaysIllegal(MlirOperation op, void *userData) {
+ (void)op;
+ intptr_t *counter = (intptr_t *)userData;
+ (*counter)++;
+ return false;
+}
+
+// Runs a partial conversion of `moduleString` against `target` with an empty
+// pattern set and returns whether it succeeded. This is what actually drives
+// the registered dynamic-legality callbacks.
+static bool runPartialConversion(MlirContext ctx, const char *moduleString,
+ MlirConversionTarget target) {
+ MlirModule module =
+ mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
+ assert(!mlirModuleIsNull(module) && "expected module to parse");
+ MlirOperation moduleOp = mlirModuleGetOperation(module);
+
+ MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
+ MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
+ MlirConversionConfig config = mlirConversionConfigCreate();
+
+ MlirLogicalResult result =
+ mlirApplyPartialConversion(moduleOp, target, frozen, config);
+
+ mlirConversionConfigDestroy(config);
+ mlirFrozenRewritePatternSetDestroy(frozen);
+ mlirModuleDestroy(module);
+
+ return mlirLogicalResultIsSuccess(result);
+}
+
+void testConversionTargetDynamicLegality(MlirContext ctx) {
+ // CHECK-LABEL: @testConversionTargetDynamicLegality
+ fprintf(stderr, "@testConversionTargetDynamicLegality\n");
+
+ const char *opModule = "\"dialect.op1\"() : () -> ()\n";
+
+ // addDynamicallyLegalOp: callback returning true makes the op legal, so the
+ // (pattern-free) partial conversion succeeds and the callback is invoked.
+ {
+ MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+ intptr_t counter = 0;
+ mlirConversionTargetAddDynamicallyLegalOp(
+ target, mlirStringRefCreateFromCString("dialect.op1"),
+ dynamicLegalityAlwaysLegal, &counter);
+ assert(runPartialConversion(ctx, opModule, target));
+ assert(counter > 0 && "legality callback must be invoked");
+ mlirConversionTargetDestroy(target);
+ }
+
+ // addDynamicallyLegalOp: callback returning false makes the op illegal. With
+ // no pattern to legalize it, the partial conversion fails -- proving the
+ // callback's return value actually drives the result.
+ {
+ MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+ intptr_t counter = 0;
+ mlirConversionTargetAddDynamicallyLegalOp(
+ target, mlirStringRefCreateFromCString("dialect.op1"),
+ dynamicLegalityAlwaysIllegal, &counter);
+ assert(!runPartialConversion(ctx, opModule, target));
+ assert(counter > 0 && "legality callback must be invoked");
+ mlirConversionTargetDestroy(target);
+ }
+
+ // addDynamicallyLegalDialect: the callback applies to every op in the
+ // dialect. Returning true keeps `dialect.op1` legal -> success.
+ {
+ MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+ intptr_t counter = 0;
+ mlirConversionTargetAddDynamicallyLegalDialect(
+ target, mlirStringRefCreateFromCString("dialect"),
+ dynamicLegalityAlwaysLegal, &counter);
+ assert(runPartialConversion(ctx, opModule, target));
+ assert(counter > 0 && "dialect legality callback must be invoked");
+ mlirConversionTargetDestroy(target);
+ }
+
+ // markUnknownOpDynamicallyLegal: `dialect.op1` is unregistered and otherwise
+ // unmarked, so the unknown-op callback decides its legality.
+ {
+ MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+ intptr_t counter = 0;
+ mlirConversionTargetMarkUnknownOpDynamicallyLegal(
+ target, dynamicLegalityAlwaysLegal, &counter);
+ assert(runPartialConversion(ctx, opModule, target));
+ assert(counter > 0 && "unknown-op legality callback must be invoked");
+ mlirConversionTargetDestroy(target);
+ }
+
+ // markOpRecursivelyLegal: an op marked recursively legal short-circuits the
+ // walk so nested ops are never checked. Here `dialect.inner` is illegal, but
+ // because `dialect.outer` is recursively legal the conversion still succeeds
+ // and the inner op's (illegal) callback is never invoked.
+ {
+ const char *nestedModule = "\"dialect.outer\"() ({\n"
+ " \"dialect.inner\"() : () -> ()\n"
+ "}) : () -> ()\n";
+ MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+ intptr_t innerCounter = 0;
+ intptr_t recursiveCounter = 0;
+ mlirConversionTargetAddDynamicallyLegalOp(
+ target, mlirStringRefCreateFromCString("dialect.inner"),
+ dynamicLegalityAlwaysIllegal, &innerCounter);
+ mlirConversionTargetAddLegalOp(
+ target, mlirStringRefCreateFromCString("dialect.outer"));
+ mlirConversionTargetMarkOpRecursivelyLegal(
+ target, mlirStringRefCreateFromCString("dialect.outer"),
+ dynamicLegalityAlwaysLegal, &recursiveCounter);
+ assert(runPartialConversion(ctx, nestedModule, target));
+ assert(recursiveCounter > 0 && "recursive legality callback must run");
+ assert(innerCounter == 0 &&
+ "nested op must not be visited under recursive legality");
+ mlirConversionTargetDestroy(target);
+ }
+
+ // markOpRecursivelyLegal with a NULL callback: the op is unconditionally
+ // recursively legal (no per-instance check), so the nested illegal op is
+ // still skipped and the conversion succeeds.
+ {
+ const char *nestedModule = "\"dialect.outer\"() ({\n"
+ " \"dialect.inner\"() : () -> ()\n"
+ "}) : () -> ()\n";
+ MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+ intptr_t innerCounter = 0;
+ mlirConversionTargetAddDynamicallyLegalOp(
+ target, mlirStringRefCreateFromCString("dialect.inner"),
+ dynamicLegalityAlwaysIllegal, &innerCounter);
+ mlirConversionTargetAddLegalOp(
+ target, mlirStringRefCreateFromCString("dialect.outer"));
+ mlirConversionTargetMarkOpRecursivelyLegal(
+ target, mlirStringRefCreateFromCString("dialect.outer"), NULL, NULL);
+ assert(runPartialConversion(ctx, nestedModule, target));
+ assert(innerCounter == 0 &&
+ "nested op must not be visited under recursive legality");
+ mlirConversionTargetDestroy(target);
+ }
+
+ // CHECK: testConversionTargetDynamicLegality: PASSED
+ fprintf(stderr, "testConversionTargetDynamicLegality: PASSED\n");
+}
+
int main(void) {
MlirContext ctx = mlirContextCreate();
mlirContextSetAllowUnregisteredDialects(ctx, true);
@@ -638,6 +786,7 @@ int main(void) {
testReplaceUses(ctx);
testGreedyRewriteDriverConfig(ctx);
testCloneWithMapping(ctx);
+ testConversionTargetDynamicLegality(ctx);
mlirContextDestroy(ctx);
return 0;
More information about the Mlir-commits
mailing list