[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