[Mlir-commits] [mlir] [mlir-c] Add TypeConverter source and target materialization (PR #208934)

Maksim Levental llvmlistbot at llvm.org
Sat Jul 18 11:44:07 PDT 2026


https://github.com/makslevental updated https://github.com/llvm/llvm-project/pull/208934

>From 496a4a2d6b35cb97e7493d1d1612e6cb9e7b6781 Mon Sep 17 00:00:00 2001
From: makslevental <maksim.levental at gmail.com>
Date: Mon, 29 Jun 2026 09:57:46 -0700
Subject: [PATCH 1/5] [mlir-c] Add TypeConverter source materialization

---
 mlir/include/mlir-c/Rewrite.h        |  16 ++++
 mlir/lib/CAPI/Transforms/Rewrite.cpp |  32 +++++++
 mlir/test/CAPI/rewrite.c             | 125 +++++++++++++++++++++++++++
 3 files changed, 173 insertions(+)

diff --git a/mlir/include/mlir-c/Rewrite.h b/mlir/include/mlir-c/Rewrite.h
index 3356e6f445e47..2d7ed3e4b7111 100644
--- a/mlir/include/mlir-c/Rewrite.h
+++ b/mlir/include/mlir-c/Rewrite.h
@@ -605,6 +605,22 @@ mlirTypeConverterAddConversion(MlirTypeConverter typeConverter,
 MLIR_CAPI_EXPORTED MlirType
 mlirTypeConverterConvertType(MlirTypeConverter typeConverter, MlirType type);
 
+/// Callback type for type materializations. Given a builder (passed as a
+/// rewriter), the desired output type, the input values, and a location, the
+/// callback must build a cast-like operation that produces a single value of
+/// `outputType` and return it. Returning a null MlirValue indicates failure, in
+/// which case another registered materialization may be attempted.
+typedef MlirValue (*MlirTypeConverterMaterializationCallback)(
+    MlirRewriterBase rewriter, MlirType outputType, intptr_t nInputs,
+    MlirValue *inputs, MlirLocation loc, void *userData);
+
+/// Register a source materialization with the given TypeConverter. This is
+/// invoked when a replacement value must be converted back to its original
+/// source type because some uses persist beyond the main conversion.
+MLIR_CAPI_EXPORTED void mlirTypeConverterAddSourceMaterialization(
+    MlirTypeConverter typeConverter,
+    MlirTypeConverterMaterializationCallback callback, void *userData);
+
 //===----------------------------------------------------------------------===//
 /// ConversionPattern API
 //===----------------------------------------------------------------------===//
diff --git a/mlir/lib/CAPI/Transforms/Rewrite.cpp b/mlir/lib/CAPI/Transforms/Rewrite.cpp
index 083ed6f999ae3..2fb8eff5ad79d 100644
--- a/mlir/lib/CAPI/Transforms/Rewrite.cpp
+++ b/mlir/lib/CAPI/Transforms/Rewrite.cpp
@@ -669,6 +669,38 @@ MlirType mlirTypeConverterConvertType(MlirTypeConverter typeConverter,
   return wrap(unwrap(typeConverter)->convertType(unwrap(type)));
 }
 
+namespace {
+/// Wraps a C materialization callback as a C++ materialization callback of the
+/// form `Value(OpBuilder &, Type, ValueRange, Location)`, shared by both source
+/// and target materializations. The builder is always a RewriterBase in the
+/// conversion driver, so it is safe to expose it as an MlirRewriterBase.
+std::function<Value(OpBuilder &, Type, ValueRange, Location)>
+wrapMaterializationCallback(MlirTypeConverterMaterializationCallback callback,
+                            void *userData) {
+  return [callback, userData](OpBuilder &builder, Type type, ValueRange inputs,
+                              Location loc) -> Value {
+    SmallVector<MlirValue> wrappedInputs;
+    wrappedInputs.reserve(inputs.size());
+    for (Value v : inputs)
+      wrappedInputs.push_back(wrap(v));
+    MlirValue result =
+        callback(wrap(static_cast<RewriterBase *>(&builder)), wrap(type),
+                 static_cast<intptr_t>(wrappedInputs.size()),
+                 wrappedInputs.data(), wrap(loc), userData);
+    return mlirValueIsNull(result) ? Value() : unwrap(result);
+  };
+}
+} // namespace
+
+void mlirTypeConverterAddSourceMaterialization(
+    MlirTypeConverter typeConverter,
+    MlirTypeConverterMaterializationCallback callback, void *userData) {
+  assert(callback && "expected non-null materialization callback");
+  unwrap(typeConverter)
+      ->addSourceMaterialization(
+          wrapMaterializationCallback(callback, userData));
+}
+
 //===----------------------------------------------------------------------===//
 /// ConversionPattern API
 //===----------------------------------------------------------------------===//
diff --git a/mlir/test/CAPI/rewrite.c b/mlir/test/CAPI/rewrite.c
index 439d1355af822..8e5b98abdb883 100644
--- a/mlir/test/CAPI/rewrite.c
+++ b/mlir/test/CAPI/rewrite.c
@@ -802,6 +802,130 @@ void testConversionTargetDynamicLegality(MlirContext ctx) {
   fprintf(stderr, "testConversionTargetDynamicLegality: PASSED\n");
 }
 
+// Type conversion callback: maps i32 -> i64 and leaves every other type
+// unchanged (identity). Used by the materialization tests below.
+static MlirLogicalResult widenI32ToI64(MlirType type, MlirType *result,
+                                       void *userData) {
+  (void)userData;
+  if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32)
+    *result = mlirIntegerTypeGet(mlirTypeGetContext(type), 64);
+  else
+    *result = type;
+  return mlirLogicalResultSuccess();
+}
+
+// Materialization callback: builds a `test.cast` op that produces a single
+// value of `outputType` from the given inputs, and records that it ran by
+// bumping the counter passed as userData.
+static MlirValue buildCastMaterialization(MlirRewriterBase rewriter,
+                                          MlirType outputType, intptr_t nInputs,
+                                          MlirValue *inputs, MlirLocation loc,
+                                          void *userData) {
+  intptr_t *counter = (intptr_t *)userData;
+  if (counter)
+    (*counter)++;
+  MlirOperationState state =
+      mlirOperationStateGet(mlirStringRefCreateFromCString("test.cast"), loc);
+  mlirOperationStateAddOperands(&state, nInputs, inputs);
+  mlirOperationStateAddResults(&state, 1, &outputType);
+  MlirOperation castOp = mlirOperationCreate(&state);
+  mlirRewriterBaseInsert(rewriter, castOp);
+  return mlirOperationGetResult(castOp, 0);
+}
+
+// Conversion pattern for `test.source`: replaces it with a `test.source_i64`
+// op whose result has the widened (i64) type. Because the original result type
+// (i32) differs from the replacement type (i64), persisting uses force the
+// framework to insert a source materialization.
+static MlirLogicalResult convertSource(MlirConversionPattern pattern,
+                                       MlirOperation op, intptr_t nOperands,
+                                       MlirValue *operands,
+                                       MlirConversionPatternRewriter rewriter,
+                                       void *userData) {
+  (void)pattern;
+  (void)nOperands;
+  (void)operands;
+  (void)userData;
+  MlirContext ctx = mlirOperationGetContext(op);
+  MlirLocation loc = mlirOperationGetLocation(op);
+  MlirType i64 = mlirIntegerTypeGet(ctx, 64);
+  MlirOperationState state = mlirOperationStateGet(
+      mlirStringRefCreateFromCString("test.source_i64"), loc);
+  mlirOperationStateAddResults(&state, 1, &i64);
+  MlirOperation newOp = mlirOperationCreate(&state);
+
+  MlirRewriterBase base = mlirPatternRewriterAsBase(
+      mlirConversionPatternRewriterAsPatternRewriter(rewriter));
+  mlirRewriterBaseInsert(base, newOp);
+  MlirValue newVal = mlirOperationGetResult(newOp, 0);
+  mlirRewriterBaseReplaceOpWithValues(base, op, 1, &newVal);
+  return mlirLogicalResultSuccess();
+}
+
+void testTypeConverterSourceMaterialization(MlirContext ctx) {
+  // CHECK-LABEL: @testTypeConverterSourceMaterialization
+  fprintf(stderr, "@testTypeConverterSourceMaterialization\n");
+
+  // `test.source` produces an i32 that is consumed by the (legal) `test.user`.
+  // Converting `test.source` to an i64-producing op leaves `test.user` wanting
+  // the original i32, which triggers a source materialization back to i32.
+  const char *moduleString = "%0 = \"test.source\"() : () -> i32\n"
+                             "\"test.user\"(%0) : (i32) -> ()\n";
+  MlirModule module =
+      mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
+  MlirOperation moduleOp = mlirModuleGetOperation(module);
+
+  MlirTypeConverter converter = mlirTypeConverterCreate();
+  mlirTypeConverterAddConversion(converter, widenI32ToI64, NULL);
+  intptr_t materializationCounter = 0;
+  mlirTypeConverterAddSourceMaterialization(converter, buildCastMaterialization,
+                                            &materializationCounter);
+
+  MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
+  MlirConversionPatternCallbacks callbacks = {NULL, NULL, convertSource};
+  MlirConversionPattern pattern = mlirOpConversionPatternCreate(
+      mlirStringRefCreateFromCString("test.source"), 1, ctx, converter,
+      callbacks, NULL, 0, NULL);
+  mlirRewritePatternSetAdd(patterns,
+                           mlirConversionPatternAsRewritePattern(pattern));
+  MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
+
+  MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+  mlirConversionTargetAddIllegalOp(
+      target, mlirStringRefCreateFromCString("test.source"));
+  mlirConversionTargetAddLegalOp(
+      target, mlirStringRefCreateFromCString("test.source_i64"));
+  mlirConversionTargetAddLegalOp(target,
+                                 mlirStringRefCreateFromCString("test.cast"));
+  mlirConversionTargetAddLegalOp(target,
+                                 mlirStringRefCreateFromCString("test.user"));
+  mlirConversionTargetAddLegalOp(
+      target, mlirStringRefCreateFromCString("builtin.module"));
+
+  MlirConversionConfig config = mlirConversionConfigCreate();
+  MlirLogicalResult result =
+      mlirApplyPartialConversion(moduleOp, target, frozen, config);
+  assert(mlirLogicalResultIsSuccess(result));
+  assert(materializationCounter > 0 &&
+         "source materialization callback must be invoked");
+
+  mlirOperationDump(moduleOp);
+  // clang-format off
+  // CHECK: %[[v:.*]] = "test.source_i64"() : () -> i64
+  // CHECK: %[[c:.*]] = "test.cast"(%[[v]]) : (i64) -> i32
+  // CHECK: "test.user"(%[[c]]) : (i32) -> ()
+  // clang-format on
+
+  mlirConversionConfigDestroy(config);
+  mlirConversionTargetDestroy(target);
+  mlirFrozenRewritePatternSetDestroy(frozen);
+  mlirTypeConverterDestroy(converter);
+  mlirModuleDestroy(module);
+
+  // CHECK: testTypeConverterSourceMaterialization: PASSED
+  fprintf(stderr, "testTypeConverterSourceMaterialization: PASSED\n");
+}
+
 int main(void) {
   MlirContext ctx = mlirContextCreate();
   mlirContextSetAllowUnregisteredDialects(ctx, true);
@@ -818,6 +942,7 @@ int main(void) {
   testGreedyRewriteDriverConfig(ctx);
   testCloneWithMapping(ctx);
   testConversionTargetDynamicLegality(ctx);
+  testTypeConverterSourceMaterialization(ctx);
 
   mlirContextDestroy(ctx);
   return 0;

>From 12609605057480049675bea96fc6de538c16141d Mon Sep 17 00:00:00 2001
From: makslevental <maksim.levental at gmail.com>
Date: Mon, 29 Jun 2026 09:59:12 -0700
Subject: [PATCH 2/5] [mlir-c] Add TypeConverter target materialization

---
 mlir/include/mlir-c/Rewrite.h        |  7 +++
 mlir/lib/CAPI/Transforms/Rewrite.cpp |  9 +++
 mlir/test/CAPI/rewrite.c             | 90 ++++++++++++++++++++++++++++
 3 files changed, 106 insertions(+)

diff --git a/mlir/include/mlir-c/Rewrite.h b/mlir/include/mlir-c/Rewrite.h
index 2d7ed3e4b7111..e074fc96f2aeb 100644
--- a/mlir/include/mlir-c/Rewrite.h
+++ b/mlir/include/mlir-c/Rewrite.h
@@ -621,6 +621,13 @@ MLIR_CAPI_EXPORTED void mlirTypeConverterAddSourceMaterialization(
     MlirTypeConverter typeConverter,
     MlirTypeConverterMaterializationCallback callback, void *userData);
 
+/// Register a target materialization with the given TypeConverter. This is
+/// invoked when a value must be converted to a target type according to a
+/// pattern's type converter.
+MLIR_CAPI_EXPORTED void mlirTypeConverterAddTargetMaterialization(
+    MlirTypeConverter typeConverter,
+    MlirTypeConverterMaterializationCallback callback, void *userData);
+
 //===----------------------------------------------------------------------===//
 /// ConversionPattern API
 //===----------------------------------------------------------------------===//
diff --git a/mlir/lib/CAPI/Transforms/Rewrite.cpp b/mlir/lib/CAPI/Transforms/Rewrite.cpp
index 2fb8eff5ad79d..92456d6d8a435 100644
--- a/mlir/lib/CAPI/Transforms/Rewrite.cpp
+++ b/mlir/lib/CAPI/Transforms/Rewrite.cpp
@@ -701,6 +701,15 @@ void mlirTypeConverterAddSourceMaterialization(
           wrapMaterializationCallback(callback, userData));
 }
 
+void mlirTypeConverterAddTargetMaterialization(
+    MlirTypeConverter typeConverter,
+    MlirTypeConverterMaterializationCallback callback, void *userData) {
+  assert(callback && "expected non-null materialization callback");
+  unwrap(typeConverter)
+      ->addTargetMaterialization(
+          wrapMaterializationCallback(callback, userData));
+}
+
 //===----------------------------------------------------------------------===//
 /// ConversionPattern API
 //===----------------------------------------------------------------------===//
diff --git a/mlir/test/CAPI/rewrite.c b/mlir/test/CAPI/rewrite.c
index 8e5b98abdb883..c1624c54eb461 100644
--- a/mlir/test/CAPI/rewrite.c
+++ b/mlir/test/CAPI/rewrite.c
@@ -926,6 +926,95 @@ void testTypeConverterSourceMaterialization(MlirContext ctx) {
   fprintf(stderr, "testTypeConverterSourceMaterialization: PASSED\n");
 }
 
+// Conversion pattern for `test.consumer`: replaces it with a
+// `test.consumer_legal` op that consumes the (already remapped) operands. The
+// operand of the original op has type i32 but its producer is not converted, so
+// the framework inserts a target materialization to i64 before invoking this
+// pattern -- the remapped `operands` are therefore the i64 cast results.
+static MlirLogicalResult convertConsumer(MlirConversionPattern pattern,
+                                         MlirOperation op, intptr_t nOperands,
+                                         MlirValue *operands,
+                                         MlirConversionPatternRewriter rewriter,
+                                         void *userData) {
+  (void)pattern;
+  (void)userData;
+  MlirLocation loc = mlirOperationGetLocation(op);
+  MlirOperationState state = mlirOperationStateGet(
+      mlirStringRefCreateFromCString("test.consumer_legal"), loc);
+  mlirOperationStateAddOperands(&state, nOperands, operands);
+  MlirOperation newOp = mlirOperationCreate(&state);
+
+  MlirRewriterBase base = mlirPatternRewriterAsBase(
+      mlirConversionPatternRewriterAsPatternRewriter(rewriter));
+  mlirRewriterBaseInsert(base, newOp);
+  mlirRewriterBaseEraseOp(base, op);
+  return mlirLogicalResultSuccess();
+}
+
+void testTypeConverterTargetMaterialization(MlirContext ctx) {
+  // CHECK-LABEL: @testTypeConverterTargetMaterialization
+  fprintf(stderr, "@testTypeConverterTargetMaterialization\n");
+
+  // `test.consumer` takes an i32 from the (legal, unconverted) `test.producer`.
+  // Converting `test.consumer` requires its operand as i64, which triggers a
+  // target materialization from i32 to i64.
+  const char *moduleString = "%0 = \"test.producer\"() : () -> i32\n"
+                             "\"test.consumer\"(%0) : (i32) -> ()\n";
+  MlirModule module =
+      mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
+  MlirOperation moduleOp = mlirModuleGetOperation(module);
+
+  MlirTypeConverter converter = mlirTypeConverterCreate();
+  mlirTypeConverterAddConversion(converter, widenI32ToI64, NULL);
+  intptr_t materializationCounter = 0;
+  mlirTypeConverterAddTargetMaterialization(converter, buildCastMaterialization,
+                                            &materializationCounter);
+
+  MlirRewritePatternSet patterns = mlirRewritePatternSetCreate(ctx);
+  MlirConversionPatternCallbacks callbacks = {NULL, NULL, convertConsumer};
+  MlirConversionPattern pattern = mlirOpConversionPatternCreate(
+      mlirStringRefCreateFromCString("test.consumer"), 1, ctx, converter,
+      callbacks, NULL, 0, NULL);
+  mlirRewritePatternSetAdd(patterns,
+                           mlirConversionPatternAsRewritePattern(pattern));
+  MlirFrozenRewritePatternSet frozen = mlirFreezeRewritePattern(patterns);
+
+  MlirConversionTarget target = mlirConversionTargetCreate(ctx);
+  mlirConversionTargetAddIllegalOp(
+      target, mlirStringRefCreateFromCString("test.consumer"));
+  mlirConversionTargetAddLegalOp(
+      target, mlirStringRefCreateFromCString("test.producer"));
+  mlirConversionTargetAddLegalOp(
+      target, mlirStringRefCreateFromCString("test.consumer_legal"));
+  mlirConversionTargetAddLegalOp(target,
+                                 mlirStringRefCreateFromCString("test.cast"));
+  mlirConversionTargetAddLegalOp(
+      target, mlirStringRefCreateFromCString("builtin.module"));
+
+  MlirConversionConfig config = mlirConversionConfigCreate();
+  MlirLogicalResult result =
+      mlirApplyPartialConversion(moduleOp, target, frozen, config);
+  assert(mlirLogicalResultIsSuccess(result));
+  assert(materializationCounter > 0 &&
+         "target materialization callback must be invoked");
+
+  mlirOperationDump(moduleOp);
+  // clang-format off
+  // CHECK: %[[v:.*]] = "test.producer"() : () -> i32
+  // CHECK: %[[c:.*]] = "test.cast"(%[[v]]) : (i32) -> i64
+  // CHECK: "test.consumer_legal"(%[[c]]) : (i64) -> ()
+  // clang-format on
+
+  mlirConversionConfigDestroy(config);
+  mlirConversionTargetDestroy(target);
+  mlirFrozenRewritePatternSetDestroy(frozen);
+  mlirTypeConverterDestroy(converter);
+  mlirModuleDestroy(module);
+
+  // CHECK: testTypeConverterTargetMaterialization: PASSED
+  fprintf(stderr, "testTypeConverterTargetMaterialization: PASSED\n");
+}
+
 int main(void) {
   MlirContext ctx = mlirContextCreate();
   mlirContextSetAllowUnregisteredDialects(ctx, true);
@@ -943,6 +1032,7 @@ int main(void) {
   testCloneWithMapping(ctx);
   testConversionTargetDynamicLegality(ctx);
   testTypeConverterSourceMaterialization(ctx);
+  testTypeConverterTargetMaterialization(ctx);
 
   mlirContextDestroy(ctx);
   return 0;

>From 7fb9b326d964f82016e7553d911d4c479798e1f4 Mon Sep 17 00:00:00 2001
From: makslevental <maksim.levental at gmail.com>
Date: Sat, 11 Jul 2026 11:55:10 -0700
Subject: [PATCH 3/5] [mlir-c] Tighten materialization test CHECKs to pin full
 module body

Use CHECK-NEXT to match the entire module body (module { ... }) rather
than loose CHECK lines, so the tests also assert the absence of any
stray ops -- e.g. a leftover builtin.unrealized_conversion_cast -- that
a loose CHECK would silently allow.
---
 mlir/test/CAPI/rewrite.c | 16 ++++++++++------
 1 file changed, 10 insertions(+), 6 deletions(-)

diff --git a/mlir/test/CAPI/rewrite.c b/mlir/test/CAPI/rewrite.c
index c1624c54eb461..6d87d0236061b 100644
--- a/mlir/test/CAPI/rewrite.c
+++ b/mlir/test/CAPI/rewrite.c
@@ -911,9 +911,11 @@ void testTypeConverterSourceMaterialization(MlirContext ctx) {
 
   mlirOperationDump(moduleOp);
   // clang-format off
-  // CHECK: %[[v:.*]] = "test.source_i64"() : () -> i64
-  // CHECK: %[[c:.*]] = "test.cast"(%[[v]]) : (i64) -> i32
-  // CHECK: "test.user"(%[[c]]) : (i32) -> ()
+  // CHECK:      module {
+  // CHECK-NEXT:   %[[v:.*]] = "test.source_i64"() : () -> i64
+  // CHECK-NEXT:   %[[c:.*]] = "test.cast"(%[[v]]) : (i64) -> i32
+  // CHECK-NEXT:   "test.user"(%[[c]]) : (i32) -> ()
+  // CHECK-NEXT: }
   // clang-format on
 
   mlirConversionConfigDestroy(config);
@@ -1000,9 +1002,11 @@ void testTypeConverterTargetMaterialization(MlirContext ctx) {
 
   mlirOperationDump(moduleOp);
   // clang-format off
-  // CHECK: %[[v:.*]] = "test.producer"() : () -> i32
-  // CHECK: %[[c:.*]] = "test.cast"(%[[v]]) : (i32) -> i64
-  // CHECK: "test.consumer_legal"(%[[c]]) : (i64) -> ()
+  // CHECK:      module {
+  // CHECK-NEXT:   %[[v:.*]] = "test.producer"() : () -> i32
+  // CHECK-NEXT:   %[[c:.*]] = "test.cast"(%[[v]]) : (i32) -> i64
+  // CHECK-NEXT:   "test.consumer_legal"(%[[c]]) : (i64) -> ()
+  // CHECK-NEXT: }
   // clang-format on
 
   mlirConversionConfigDestroy(config);

>From 0747c994b3badf2322f7af5362ceb9556e3f42ec Mon Sep 17 00:00:00 2001
From: makslevental <maksim.levental at gmail.com>
Date: Fri, 17 Jul 2026 12:22:08 -0700
Subject: [PATCH 4/5] address comment

---
 mlir/test/CAPI/rewrite.c | 8 ++++----
 1 file changed, 4 insertions(+), 4 deletions(-)

diff --git a/mlir/test/CAPI/rewrite.c b/mlir/test/CAPI/rewrite.c
index 6d87d0236061b..b10afbd8b812c 100644
--- a/mlir/test/CAPI/rewrite.c
+++ b/mlir/test/CAPI/rewrite.c
@@ -627,7 +627,7 @@ static MlirConversionTargetLegality dynamicLegalityAlwaysLegal(MlirOperation op,
                                                                void *userData) {
   (void)op;
   intptr_t *counter = (intptr_t *)userData;
-  (*counter)++;
+  ++(*counter);
   return MLIR_CONVERSION_TARGET_LEGALITY_LEGAL;
 }
 
@@ -635,7 +635,7 @@ static MlirConversionTargetLegality
 dynamicLegalityAlwaysIllegal(MlirOperation op, void *userData) {
   (void)op;
   intptr_t *counter = (intptr_t *)userData;
-  (*counter)++;
+  ++(*counter);
   return MLIR_CONVERSION_TARGET_LEGALITY_ILLEGAL;
 }
 
@@ -643,7 +643,7 @@ static MlirConversionTargetLegality dynamicLegalityNoOpinion(MlirOperation op,
                                                              void *userData) {
   (void)op;
   intptr_t *counter = (intptr_t *)userData;
-  (*counter)++;
+  ++(*counter);
   return MLIR_CONVERSION_TARGET_LEGALITY_NO_OPINION;
 }
 
@@ -823,7 +823,7 @@ static MlirValue buildCastMaterialization(MlirRewriterBase rewriter,
                                           void *userData) {
   intptr_t *counter = (intptr_t *)userData;
   if (counter)
-    (*counter)++;
+    ++(*counter);
   MlirOperationState state =
       mlirOperationStateGet(mlirStringRefCreateFromCString("test.cast"), loc);
   mlirOperationStateAddOperands(&state, nInputs, inputs);

>From 0e10d1f261c6670fbcf1f45afce2f86c9fb29fe6 Mon Sep 17 00:00:00 2001
From: makslevental <maksim.levental at gmail.com>
Date: Sat, 18 Jul 2026 11:37:08 -0700
Subject: [PATCH 5/5] [mlir-c] Use a status enum for the type conversion
 callback

The 1:1 conversion callback returned MlirLogicalResult and encoded the
three C++ conversion states implicitly: returning failure meant "try
another conversion", while returning success with a null out-parameter
meant a hard failure. This dual encoding was easy to misuse and the doc
comment conflated the two.

Return a MlirTypeConverterConversionStatus enum with explicit Success,
Failure (do not try another), and Declined (try another) states, and add
a test covering the decline-fallback and hard-failure paths.
---
 mlir/include/mlir-c/Rewrite.h        | 29 ++++++++---
 mlir/lib/CAPI/Transforms/Rewrite.cpp | 18 +++++--
 mlir/test/CAPI/rewrite.c             | 77 ++++++++++++++++++++++++++--
 3 files changed, 108 insertions(+), 16 deletions(-)

diff --git a/mlir/include/mlir-c/Rewrite.h b/mlir/include/mlir-c/Rewrite.h
index e074fc96f2aeb..cc98bf7ba9a63 100644
--- a/mlir/include/mlir-c/Rewrite.h
+++ b/mlir/include/mlir-c/Rewrite.h
@@ -588,12 +588,29 @@ MLIR_CAPI_EXPORTED MlirTypeConverter mlirTypeConverterCreate(void);
 MLIR_CAPI_EXPORTED void
 mlirTypeConverterDestroy(MlirTypeConverter typeConverter);
 
-/// Callback type for type conversion functions.
-/// Returns failure or sets convertedType to MlirType{NULL} to indicate failure.
-/// If failure is returned, the converter is allowed to try another
-/// conversion function to perform the conversion.
-typedef MlirLogicalResult (*MlirTypeConverterConversionCallback)(
-    MlirType type, MlirType *convertedType, void *userData);
+/// Outcome of a type conversion callback (MlirTypeConverterConversionCallback).
+/// Mirrors the three states of the underlying C++ conversion function.
+typedef enum MlirTypeConverterConversionStatus {
+  /// The type was converted; the callback set `*convertedType` to the result.
+  MlirTypeConverterConversionStatusSuccess = 0,
+  /// The conversion failed hard; no further conversion function will be tried.
+  MlirTypeConverterConversionStatusFailure = 1,
+  /// The conversion was declined; another registered conversion function may be
+  /// tried.
+  MlirTypeConverterConversionStatusDeclined = 2,
+} MlirTypeConverterConversionStatus;
+
+/// Callback type for type conversion functions. On
+/// MlirTypeConverterConversionStatusSuccess the callback must set
+/// `*convertedType` to the converted type.
+/// MlirTypeConverterConversionStatusDeclined leaves the type unconverted and
+/// allows another conversion function to be tried, whereas
+/// MlirTypeConverterConversionStatusFailure fails the conversion outright
+/// without trying any further conversion function.
+typedef MlirTypeConverterConversionStatus (
+    *MlirTypeConverterConversionCallback)(MlirType type,
+                                          MlirType *convertedType,
+                                          void *userData);
 
 /// Add a type conversion function to the given TypeConverter.
 MLIR_CAPI_EXPORTED void
diff --git a/mlir/lib/CAPI/Transforms/Rewrite.cpp b/mlir/lib/CAPI/Transforms/Rewrite.cpp
index 92456d6d8a435..e324a848691e9 100644
--- a/mlir/lib/CAPI/Transforms/Rewrite.cpp
+++ b/mlir/lib/CAPI/Transforms/Rewrite.cpp
@@ -654,13 +654,21 @@ void mlirTypeConverterAddConversion(
       ->addConversion(
           [convertType, userData](Type type) -> std::optional<Type> {
             MlirType converted{nullptr};
-            MlirLogicalResult result =
+            MlirTypeConverterConversionStatus status =
                 convertType(wrap(type), &converted, userData);
-            if (mlirLogicalResultIsFailure(result))
-              return std::nullopt; // allowed to try another conversion function
-            if (mlirTypeIsNull(converted))
+            switch (status) {
+            case MlirTypeConverterConversionStatusSuccess:
+              assert(!mlirTypeIsNull(converted) &&
+                     "a successful conversion must set convertedType");
+              return unwrap(converted);
+            case MlirTypeConverterConversionStatusFailure:
+              // A null result maps to success(false), i.e. a hard failure: the
+              // driver will not try another conversion function.
               return nullptr;
-            return unwrap(converted);
+            case MlirTypeConverterConversionStatusDeclined:
+              return std::nullopt; // allowed to try another conversion function
+            }
+            llvm_unreachable("unknown MlirTypeConverterConversionStatus");
           });
 }
 
diff --git a/mlir/test/CAPI/rewrite.c b/mlir/test/CAPI/rewrite.c
index b10afbd8b812c..6615364825e95 100644
--- a/mlir/test/CAPI/rewrite.c
+++ b/mlir/test/CAPI/rewrite.c
@@ -804,18 +804,45 @@ void testConversionTargetDynamicLegality(MlirContext ctx) {
 
 // Type conversion callback: maps i32 -> i64 and leaves every other type
 // unchanged (identity). Used by the materialization tests below.
-static MlirLogicalResult widenI32ToI64(MlirType type, MlirType *result,
-                                       void *userData) {
+static MlirTypeConverterConversionStatus
+widenI32ToI64(MlirType type, MlirType *result, void *userData) {
   (void)userData;
   if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32)
     *result = mlirIntegerTypeGet(mlirTypeGetContext(type), 64);
   else
     *result = type;
-  return mlirLogicalResultSuccess();
+  return MlirTypeConverterConversionStatusSuccess;
 }
 
-// Materialization callback: builds a `test.cast` op that produces a single
-// value of `outputType` from the given inputs, and records that it ran by
+// 1:1 conversion callback that declines i32 (returns Declined) so another
+// conversion function is tried, and is the identity on everything else. Bumps
+// the counter passed as userData when it declines.
+static MlirTypeConverterConversionStatus
+declineI32(MlirType type, MlirType *result, void *userData) {
+  intptr_t *counter = (intptr_t *)userData;
+  if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32) {
+    if (counter)
+      ++(*counter);
+    return MlirTypeConverterConversionStatusDeclined;
+  }
+  *result = type;
+  return MlirTypeConverterConversionStatusSuccess;
+}
+
+// 1:1 conversion callback that fails hard on i32 (returns Failure), which must
+// abort the conversion without trying another conversion function. Bumps the
+// counter passed as userData when it fails.
+static MlirTypeConverterConversionStatus
+failI32(MlirType type, MlirType *result, void *userData) {
+  intptr_t *counter = (intptr_t *)userData;
+  if (mlirTypeIsAInteger(type) && mlirIntegerTypeGetWidth(type) == 32) {
+    if (counter)
+      ++(*counter);
+    return MlirTypeConverterConversionStatusFailure;
+  }
+  *result = type;
+  return MlirTypeConverterConversionStatusSuccess;
+}
 // bumping the counter passed as userData.
 static MlirValue buildCastMaterialization(MlirRewriterBase rewriter,
                                           MlirType outputType, intptr_t nInputs,
@@ -862,6 +889,45 @@ static MlirLogicalResult convertSource(MlirConversionPattern pattern,
   return mlirLogicalResultSuccess();
 }
 
+void testTypeConverterConversionStates(MlirContext ctx) {
+  // CHECK-LABEL: @testTypeConverterConversionStates
+  fprintf(stderr, "@testTypeConverterConversionStates\n");
+
+  MlirType i32 = mlirIntegerTypeGet(ctx, 32);
+  MlirType i64 = mlirIntegerTypeGet(ctx, 64);
+  (void)i64;
+
+  // Declined: the i32 decliner is registered last (tried first) and defers to
+  // the i32 -> i64 conversion registered before it, so i32 still converts.
+  MlirTypeConverter declineConverter = mlirTypeConverterCreate();
+  mlirTypeConverterAddConversion(declineConverter, widenI32ToI64, NULL);
+  intptr_t declineCounter = 0;
+  mlirTypeConverterAddConversion(declineConverter, declineI32, &declineCounter);
+  MlirType declined = mlirTypeConverterConvertType(declineConverter, i32);
+  (void)declined;
+  assert(declineCounter == 1 && "declining conversion must be consulted");
+  assert(mlirTypeEqual(declined, i64) &&
+         "decline must fall back to the i32 -> i64 conversion");
+  mlirTypeConverterDestroy(declineConverter);
+
+  // Failure: the i32 hard-failure is registered last (tried first) and must
+  // abort the conversion without falling back to the i32 -> i64 conversion, so
+  // the type converts to null.
+  MlirTypeConverter failConverter = mlirTypeConverterCreate();
+  mlirTypeConverterAddConversion(failConverter, widenI32ToI64, NULL);
+  intptr_t failCounter = 0;
+  mlirTypeConverterAddConversion(failConverter, failI32, &failCounter);
+  MlirType failed = mlirTypeConverterConvertType(failConverter, i32);
+  (void)failed;
+  assert(failCounter == 1 && "failing conversion must be consulted");
+  assert(mlirTypeIsNull(failed) &&
+         "hard failure must not fall back to another conversion");
+  mlirTypeConverterDestroy(failConverter);
+
+  // CHECK: testTypeConverterConversionStates: PASSED
+  fprintf(stderr, "testTypeConverterConversionStates: PASSED\n");
+}
+
 void testTypeConverterSourceMaterialization(MlirContext ctx) {
   // CHECK-LABEL: @testTypeConverterSourceMaterialization
   fprintf(stderr, "@testTypeConverterSourceMaterialization\n");
@@ -1035,6 +1101,7 @@ int main(void) {
   testGreedyRewriteDriverConfig(ctx);
   testCloneWithMapping(ctx);
   testConversionTargetDynamicLegality(ctx);
+  testTypeConverterConversionStates(ctx);
   testTypeConverterSourceMaterialization(ctx);
   testTypeConverterTargetMaterialization(ctx);
 



More information about the Mlir-commits mailing list