[Mlir-commits] [mlir] [mlir-c] Add RewriterBase insertion point save/restore (PR #206531)
Maksim Levental
llvmlistbot at llvm.org
Mon Jun 29 11:51:30 PDT 2026
https://github.com/makslevental updated https://github.com/llvm/llvm-project/pull/206531
>From d24b49b2540b795a4680d6e98a3f7dd86ff2dfb3 Mon Sep 17 00:00:00 2001
From: makslevental <maksim.levental at gmail.com>
Date: Mon, 29 Jun 2026 10:04:44 -0700
Subject: [PATCH] [mlir-c] Add RewriterBase insertion point save/restore
---
mlir/include/mlir-c/Rewrite.h | 19 +++++++++
mlir/lib/CAPI/Transforms/Rewrite.cpp | 26 ++++++++++++
mlir/test/CAPI/rewrite.c | 63 ++++++++++++++++++++++++++++
3 files changed, 108 insertions(+)
diff --git a/mlir/include/mlir-c/Rewrite.h b/mlir/include/mlir-c/Rewrite.h
index 3356e6f445e47..ccdba963a446c 100644
--- a/mlir/include/mlir-c/Rewrite.h
+++ b/mlir/include/mlir-c/Rewrite.h
@@ -133,6 +133,25 @@ mlirRewriterBaseGetBlock(MlirRewriterBase rewriter);
MLIR_CAPI_EXPORTED MlirOperation
mlirRewriterBaseGetOperationAfterInsertion(MlirRewriterBase rewriter);
+/// A saved insertion point: a (block, operationAfter) pair. `operationAfter` is
+/// the operation that subsequent insertions go before. If `operationAfter` is
+/// null, the insertion point is at the end of `block`. If `block` is null, the
+/// insertion point is not set (cleared).
+typedef struct MlirRewriterBaseInsertPoint {
+ MlirBlock block;
+ MlirOperation operationAfter;
+} MlirRewriterBaseInsertPoint;
+
+/// Returns the current insertion point of the rewriter so that it can be
+/// restored later with mlirRewriterBaseRestoreInsertionPoint.
+MLIR_CAPI_EXPORTED MlirRewriterBaseInsertPoint
+mlirRewriterBaseSaveInsertionPoint(MlirRewriterBase rewriter);
+
+/// Restores a previously saved insertion point.
+MLIR_CAPI_EXPORTED void
+mlirRewriterBaseRestoreInsertionPoint(MlirRewriterBase rewriter,
+ MlirRewriterBaseInsertPoint insertPoint);
+
//===----------------------------------------------------------------------===//
/// Block and operation creation/insertion/cloning
//===----------------------------------------------------------------------===//
diff --git a/mlir/lib/CAPI/Transforms/Rewrite.cpp b/mlir/lib/CAPI/Transforms/Rewrite.cpp
index 083ed6f999ae3..d59a4118c6d9f 100644
--- a/mlir/lib/CAPI/Transforms/Rewrite.cpp
+++ b/mlir/lib/CAPI/Transforms/Rewrite.cpp
@@ -87,6 +87,32 @@ mlirRewriterBaseGetOperationAfterInsertion(MlirRewriterBase rewriter) {
return wrap(std::addressof(*it));
}
+MlirRewriterBaseInsertPoint
+mlirRewriterBaseSaveInsertionPoint(MlirRewriterBase rewriter) {
+ OpBuilder::InsertPoint ip = unwrap(rewriter)->saveInsertionPoint();
+ if (!ip.isSet())
+ return {{nullptr}, {nullptr}};
+ Block *block = ip.getBlock();
+ MlirOperation operationAfter = ip.getPoint() == block->end()
+ ? MlirOperation{nullptr}
+ : wrap(&*ip.getPoint());
+ return {wrap(block), operationAfter};
+}
+
+void mlirRewriterBaseRestoreInsertionPoint(
+ MlirRewriterBase rewriter, MlirRewriterBaseInsertPoint insertPoint) {
+ if (mlirBlockIsNull(insertPoint.block)) {
+ unwrap(rewriter)->clearInsertionPoint();
+ return;
+ }
+ Block *block = unwrap(insertPoint.block);
+ if (mlirOperationIsNull(insertPoint.operationAfter))
+ unwrap(rewriter)->setInsertionPointToEnd(block);
+ else
+ unwrap(rewriter)->setInsertionPoint(
+ block, Block::iterator(unwrap(insertPoint.operationAfter)));
+}
+
//===----------------------------------------------------------------------===//
/// Block and operation creation/insertion/cloning
//===----------------------------------------------------------------------===//
diff --git a/mlir/test/CAPI/rewrite.c b/mlir/test/CAPI/rewrite.c
index 439d1355af822..55cd046bc071d 100644
--- a/mlir/test/CAPI/rewrite.c
+++ b/mlir/test/CAPI/rewrite.c
@@ -623,6 +623,68 @@ void testCloneWithMapping(MlirContext ctx) {
fprintf(stderr, "testCloneWithMapping: PASSED\n");
}
+void testInsertionPointSaveRestore(MlirContext ctx) {
+ // CHECK-LABEL: @testInsertionPointSaveRestore
+ fprintf(stderr, "@testInsertionPointSaveRestore\n");
+
+ const char *moduleString = "\"dialect.op1\"() : () -> ()\n"
+ "\"dialect.op2\"() : () -> ()\n";
+ MlirModule module =
+ mlirModuleCreateParse(ctx, mlirStringRefCreateFromCString(moduleString));
+ MlirOperation op = mlirModuleGetOperation(module);
+ MlirBlock body = mlirModuleGetBody(module);
+ MlirOperation op1 = mlirBlockGetFirstOperation(body);
+ MlirOperation op2 = mlirOperationGetNextInBlock(op1);
+
+ MlirRewriterBase rewriter = mlirIRRewriterCreate(ctx);
+
+ // Save an insertion point that points right before op2.
+ mlirRewriterBaseSetInsertionPointBefore(rewriter, op2);
+ MlirRewriterBaseInsertPoint saved =
+ mlirRewriterBaseSaveInsertionPoint(rewriter);
+ assert(!mlirBlockIsNull(saved.block));
+ assert(mlirOperationEqual(saved.operationAfter, op2));
+
+ // Move the insertion point to the end of the block. An end-of-block insertion
+ // point round-trips with a null `operationAfter`.
+ mlirRewriterBaseSetInsertionPointToEnd(rewriter, body);
+ MlirRewriterBaseInsertPoint endIp =
+ mlirRewriterBaseSaveInsertionPoint(rewriter);
+ assert(!mlirBlockIsNull(endIp.block));
+ assert(mlirOperationIsNull(endIp.operationAfter));
+ MlirOperation opEnd = createOperationWithName(ctx, "dialect.op_end");
+ mlirRewriterBaseInsert(rewriter, opEnd);
+
+ // Restoring the first saved insertion point makes subsequent insertions land
+ // before op2 again, not at the end where we just were.
+ mlirRewriterBaseRestoreInsertionPoint(rewriter, saved);
+ MlirOperation opRestored =
+ createOperationWithName(ctx, "dialect.op_restored");
+ mlirRewriterBaseInsert(rewriter, opRestored);
+
+ // A cleared insertion point round-trips as a null block.
+ mlirRewriterBaseClearInsertionPoint(rewriter);
+ MlirRewriterBaseInsertPoint clearedIp =
+ mlirRewriterBaseSaveInsertionPoint(rewriter);
+ assert(mlirBlockIsNull(clearedIp.block));
+
+ mlirOperationDump(op);
+ // clang-format off
+ // CHECK: module {
+ // CHECK-NEXT: "dialect.op1"() : () -> ()
+ // CHECK-NEXT: %{{.*}} = "dialect.op_restored"() : () -> index
+ // CHECK-NEXT: "dialect.op2"() : () -> ()
+ // CHECK-NEXT: %{{.*}} = "dialect.op_end"() : () -> index
+ // CHECK-NEXT: }
+ // clang-format on
+
+ mlirIRRewriterDestroy(rewriter);
+ mlirModuleDestroy(module);
+
+ // CHECK: testInsertionPointSaveRestore: PASSED
+ fprintf(stderr, "testInsertionPointSaveRestore: PASSED\n");
+}
+
static MlirConversionTargetLegality dynamicLegalityAlwaysLegal(MlirOperation op,
void *userData) {
(void)op;
@@ -817,6 +879,7 @@ int main(void) {
testReplaceUses(ctx);
testGreedyRewriteDriverConfig(ctx);
testCloneWithMapping(ctx);
+ testInsertionPointSaveRestore(ctx);
testConversionTargetDynamicLegality(ctx);
mlirContextDestroy(ctx);
More information about the Mlir-commits
mailing list