[llvm-branch-commits] [mlir] [mlir] Return `OpFoldResults` from the type-erased fold hooks (PR #228762)

Victor Perez via llvm-branch-commits llvm-branch-commits at lists.llvm.org
Sat Oct 3 15:39:02 PDT 2026


https://github.com/victor-eds updated https://github.com/llvm/llvm-project/pull/228762

>From 5b06a3b83ef39bdb8b9ca31efbda594fadbaf79f Mon Sep 17 00:00:00 2001
From: =?UTF-8?q?V=C3=ADctor=20P=C3=A9rez=20Carrasco?=
 <victor.pc.upm at gmail.com>
Date: Fri, 2 Oct 2026 14:45:09 -0700
Subject: [PATCH] [mlir] Return OpFoldResults from the type-erased fold hooks
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit

The type-erased fold hooks return `LogicalResult` and fill a vector of
`OpFoldResult`. This form cannot express a fold that replaces only some
results of an op. This patch changes the hooks to return
`OpFoldResults`. The drivers do not use partial folds yet.

The patch keeps every existing fold signature and wraps each one in an
adapter that returns `OpFoldResults`:
- `OpFoldResult fold(FoldAdaptor)` and the single-result `foldTrait`: a
  null result is a failure, the op's own result is an in-place change,
  and any other value replaces the result.
- `LogicalResult fold(FoldAdaptor, SmallVectorImpl<OpFoldResult> &)` and
  the general `foldTrait`: `detail::convertLegacyFoldResults` applies
  the strict legacy contract. Failure stays failure, success with an
  empty vector is an in-place change, and success with one entry per
  result replaces every result.
- The legacy fold hooks of `DynamicOpDefinition` and
  `DialectFoldInterface` use the same strict conversion.

As before, the traits fold only when the op's own fold replaces no
result, and an op without a `fold` method folds only its traits.

The new `Operation::fold` overloads return a normalized
`OpFoldResults`. The legacy `Operation::fold` overloads return the full
vector only when the fold replaces all results. Otherwise they return
the in-place state and an empty vector. For folds and fold traits that
follow the documented contract, the legacy overloads return the same
result as before. The type check in `OpFoldResults::normalize` replaces
`checkFoldResultTypes`.

Downstream impact:
- Code that overrides `OperationName::Impl::foldHook`, calls
  `OperationName::foldHook`, or names `OperationName::FoldHookFn` must
  use the new signature
  `OpFoldResults(Operation *, ArrayRef<Attribute>)`.
  A dynamic op can keep a legacy fold hook through
  `DynamicOpDefinition::LegacyFoldHookFn`.
- `DynamicOpDefinition::get(...)` with `nullptr` or `{}` as the fold
  hook is ambiguous now. `setFoldHookFn(nullptr)` and
  `setFoldHookFn({})` remove the fold hook, so that folding the op
  fails.
- A legacy fold or fold trait that breaks the contract can behave
  differently. If the replacement of result i is result i itself, the
  result is now kept. If the fold keeps every result, it fails instead
  of replacing the op with itself. If it keeps only some results, the
  legacy `Operation::fold` overloads drop the partial fold. A fold that
  returns `failure()` after it writes results no longer blocks the later
  traits.

This patch also fixes a crash. In a graph region,
`builtin.unrealized_conversion_cast` can forward its own result:
`%0 = builtin.unrealized_conversion_cast %0 : i32 to i32`. The legacy
fold returns that result, so the greedy driver replaced the op with
itself and then erased the op while it still had uses. Now the fold
keeps the result and fails, and the cast stays. Without this patch, the
new test in `Builtin/canonicalize.mlir` aborts:

  $ mlir-opt mlir/test/Dialect/Builtin/canonicalize.mlir \
      -canonicalize="test-convergence"
  mlir-opt: mlir/lib/IR/PatternMatch.cpp:156: virtual void
  mlir::RewriterBase::eraseOp(Operation *): Assertion
  `op->use_empty() && "expected 'op' to have no uses"' failed.

The new unit tests cover the legacy adapters, the dynamic fold hooks,
`Operation::fold`, the `DialectFoldInterface` fallback, an unregistered
op, and the assert on a null legacy entry.

RFC: https://discourse.llvm.org/t/rfc-partial-folding-for-multi-result-ops/91954

Signed-off-by: Víctor Pérez Carrasco <victor.pc.upm at gmail.com>
---
 mlir/include/mlir/IR/ExtensibleDialect.h    |  23 +-
 mlir/include/mlir/IR/OpDefinition.h         | 148 +++++----
 mlir/include/mlir/IR/Operation.h            |  21 ++
 mlir/include/mlir/IR/OperationSupport.h     |  53 ++--
 mlir/lib/IR/ExtensibleDialect.cpp           |  41 ++-
 mlir/lib/IR/MLIRContext.cpp                 |   5 +-
 mlir/lib/IR/Operation.cpp                   |  76 ++---
 mlir/test/Dialect/Builtin/canonicalize.mlir |  13 +
 mlir/unittests/IR/OpFoldResultsTest.cpp     | 330 ++++++++++++++++++++
 9 files changed, 569 insertions(+), 141 deletions(-)

diff --git a/mlir/include/mlir/IR/ExtensibleDialect.h b/mlir/include/mlir/IR/ExtensibleDialect.h
index 79c07cbefcc5ce..7a1ee4b71e7c7f 100644
--- a/mlir/include/mlir/IR/ExtensibleDialect.h
+++ b/mlir/include/mlir/IR/ExtensibleDialect.h
@@ -440,6 +440,12 @@ class DynamicOpDefinition : public OperationName::Impl {
 public:
   using GetCanonicalizationPatternsFn =
       llvm::unique_function<void(RewritePatternSet &, MLIRContext *) const>;
+  /// The legacy fold hook signature. Kept for backwards compatibility until the
+  /// switch to the signature returning OpFoldResults is completed. It has the
+  /// strict legacy contract: failure, success with an empty vector (in place),
+  /// or success with one entry per result.
+  using LegacyFoldHookFn = llvm::unique_function<LogicalResult(
+      Operation *, ArrayRef<Attribute>, SmallVectorImpl<OpFoldResult> &) const>;
 
   /// Create a new op at runtime. The op is registered only after passing it to
   /// the dialect using registerDynamicOp.
@@ -462,6 +468,14 @@ class DynamicOpDefinition : public OperationName::Impl {
       OperationName::FoldHookFn &&foldHookFn,
       GetCanonicalizationPatternsFn &&getCanonicalizationPatternsFn,
       OperationName::PopulateDefaultAttrsFn &&populateDefaultAttrsFn);
+  static std::unique_ptr<DynamicOpDefinition>
+  get(StringRef name, ExtensibleDialect *dialect,
+      OperationName::VerifyInvariantsFn &&verifyFn,
+      OperationName::VerifyRegionInvariantsFn &&verifyRegionFn,
+      OperationName::ParseAssemblyFn &&parseFn,
+      OperationName::PrintAssemblyFn &&printFn, LegacyFoldHookFn &&foldHookFn,
+      GetCanonicalizationPatternsFn &&getCanonicalizationPatternsFn,
+      OperationName::PopulateDefaultAttrsFn &&populateDefaultAttrsFn);
 
   /// Returns the op typeID.
   TypeID getTypeID() { return typeID; }
@@ -495,6 +509,10 @@ class DynamicOpDefinition : public OperationName::Impl {
   void setFoldHookFn(OperationName::FoldHookFn &&foldHook) {
     foldHookFn = std::move(foldHook);
   }
+  /// Same as above, but with a legacy fold hook.
+  void setFoldHookFn(LegacyFoldHookFn &&foldHook);
+  /// Remove the fold hook, so that folding the op always fails.
+  void setFoldHookFn(std::nullptr_t);
 
   /// Set the hook returning any canonicalization pattern rewrites that the op
   /// supports, for use by the canonicalization pass.
@@ -514,9 +532,8 @@ class DynamicOpDefinition : public OperationName::Impl {
     return traits.insert(std::move(trait));
   }
 
-  LogicalResult foldHook(Operation *op, ArrayRef<Attribute> attrs,
-                         SmallVectorImpl<OpFoldResult> &results) final {
-    return foldHookFn(op, attrs, results);
+  OpFoldResults foldHook(Operation *op, ArrayRef<Attribute> attrs) final {
+    return foldHookFn(op, attrs);
   }
   void getCanonicalizationPatterns(RewritePatternSet &set,
                                    MLIRContext *context) final {
diff --git a/mlir/include/mlir/IR/OpDefinition.h b/mlir/include/mlir/IR/OpDefinition.h
index 783cd1359e9706..79ccbab2bc99be 100644
--- a/mlir/include/mlir/IR/OpDefinition.h
+++ b/mlir/include/mlir/IR/OpDefinition.h
@@ -1560,46 +1560,70 @@ using detect_has_any_fold_trait =
 /// that is specialized for operations that have a single result.
 template <typename Trait>
 std::enable_if_t<detect_has_single_result_fold_trait<Trait>::value,
-                 LogicalResult>
-foldTrait(Operation *op, ArrayRef<Attribute> operands,
-          SmallVectorImpl<OpFoldResult> &results) {
+                 OpFoldResults>
+foldTrait(Operation *op, ArrayRef<Attribute> operands) {
   assert(op->hasTrait<OpTrait::OneResult>() &&
          "expected trait on non single-result operation to implement the "
          "general `foldTrait` method");
-  // If a previous trait has already been folded and replaced this operation, we
-  // fail to fold this trait.
-  if (!results.empty())
+  OpFoldResult result = Trait::foldTrait(op, operands);
+  if (!result)
     return failure();
-
-  if (OpFoldResult result = Trait::foldTrait(op, operands)) {
-    if (llvm::dyn_cast_if_present<Value>(result) != op->getResult(0))
-      results.push_back(result);
+  // The op's own result means that the trait changed the op in place.
+  if (llvm::dyn_cast_if_present<Value>(result) == op->getResult(0))
     return success();
-  }
-  return failure();
+  return result;
 }
-/// Returns the result of folding a trait that implements a generalized
-/// `foldTrait` function that is supports any operation type.
+/// Returns the result of folding a trait that implements the generalized
+/// `foldTrait(op, operands, results)` function, with the strict legacy
+/// contract.
 template <typename Trait>
-std::enable_if_t<detect_has_fold_trait<Trait>::value, LogicalResult>
-foldTrait(Operation *op, ArrayRef<Attribute> operands,
-          SmallVectorImpl<OpFoldResult> &results) {
-  // If a previous trait has already been folded and replaced this operation, we
-  // fail to fold this trait.
-  return results.empty() ? Trait::foldTrait(op, operands, results) : failure();
+std::enable_if_t<detect_has_fold_trait<Trait>::value, OpFoldResults>
+foldTrait(Operation *op, ArrayRef<Attribute> operands) {
+  SmallVector<OpFoldResult, 2> results;
+  LogicalResult status = Trait::foldTrait(op, operands, results);
+  return ::mlir::detail::convertLegacyFoldResults(status, results);
 }
 template <typename Trait>
-inline std::enable_if_t<!detect_has_any_fold_trait<Trait>::value, LogicalResult>
-foldTrait(Operation *, ArrayRef<Attribute>, SmallVectorImpl<OpFoldResult> &) {
+inline std::enable_if_t<!detect_has_any_fold_trait<Trait>::value, OpFoldResults>
+foldTrait(Operation *, ArrayRef<Attribute>) {
   return failure();
 }
 
 /// Given a tuple type containing a set of traits, return the result of folding
-/// the given operation.
+/// the given operation. `own` is the normalized result of the op's own fold.
+/// The traits run only when `own` replaced no result, and folding stops at the
+/// first trait that does not fail.
 template <typename... Ts>
-LogicalResult foldTraits(Operation *op, ArrayRef<Attribute> operands,
-                         SmallVectorImpl<OpFoldResult> &results) {
-  return success((succeeded(foldTrait<Ts>(op, operands, results)) || ...));
+std::enable_if_t<std::disjunction_v<detect_has_any_fold_trait<Ts>...>,
+                 OpFoldResults>
+foldTraits(Operation *op, ArrayRef<Attribute> operands,
+           OpFoldResults own = {}) {
+  if (own.replacesAny())
+    return own;
+  OpFoldResults trait;
+  // Stop at the first trait whose normalized result does not fail.
+  auto tryTrait = [&](OpFoldResults result) {
+    result.normalize(op);
+    if (result.failed())
+      return false;
+    trait = std::move(result);
+    return true;
+  };
+  (void)((detect_has_any_fold_trait<Ts>::value &&
+          tryTrait(foldTrait<Ts>(op, operands))) ||
+         ...);
+  bool inPlace = own.modifiedInPlace() || trait.modifiedInPlace();
+  OpFoldResults result =
+      trait.replacesAny() ? std::move(trait) : std::move(own);
+  result.setModifiedInPlace(inPlace);
+  return result;
+}
+/// Same as above, for a set of traits in which no trait can fold.
+template <typename... Ts>
+std::enable_if_t<!std::disjunction_v<detect_has_any_fold_trait<Ts>...>,
+                 OpFoldResults>
+foldTraits(Operation *, ArrayRef<Attribute>, OpFoldResults own = {}) {
+  return own;
 }
 
 //===----------------------------------------------------------------------===//
@@ -1889,8 +1913,7 @@ class Op : public OpState, public Traits<ConcreteType>... {
     return detail::InterfaceMap::template get<Traits<ConcreteType>...>();
   }
 
-  using FoldHookFn = LogicalResult (*)(Operation *, ArrayRef<Attribute>,
-                                       SmallVectorImpl<OpFoldResult> &);
+  using FoldHookFn = OpFoldResults (*)(Operation *, ArrayRef<Attribute>);
   using HasTraitFn = bool (*)(TypeID);
   using PrintAssemblyFn = void (*)(Operation *, OpAsmPrinter &, StringRef);
   using VerifyInvariantsFn = LogicalResult (*)(Operation *);
@@ -1909,30 +1932,26 @@ class Op : public OpState, public Traits<ConcreteType>... {
                   "SmallVectorImpl<OpFoldResult> &)`");
     // If the operation is single result and defines a `fold` method.
     if constexpr (hasOneResult && hasSingleResultFold)
-      return [](Operation *op, ArrayRef<Attribute> operands,
-                SmallVectorImpl<OpFoldResult> &results) {
-        return foldSingleResultHook<ConcreteType>(op, operands, results);
+      return [](Operation *op, ArrayRef<Attribute> operands) {
+        return foldSingleResultHook<ConcreteType>(op, operands);
       };
     // The operation is not single result and defines a `fold` method.
     if constexpr (has_fold_v<ConcreteType> || has_fold_adaptor_v<ConcreteType>)
-      return [](Operation *op, ArrayRef<Attribute> operands,
-                SmallVectorImpl<OpFoldResult> &results) {
-        return foldHook<ConcreteType>(op, operands, results);
+      return [](Operation *op, ArrayRef<Attribute> operands) {
+        return foldHook<ConcreteType>(op, operands);
       };
     // The operation does not define a `fold` method.
-    return [](Operation *op, ArrayRef<Attribute> operands,
-              SmallVectorImpl<OpFoldResult> &results) {
+    return [](Operation *op, ArrayRef<Attribute> operands) {
       // In this case, we only need to fold the traits of the operation.
-      return op_definition_impl::foldTraits<Traits<ConcreteType>...>(
-          op, operands, results);
+      return op_definition_impl::foldTraits<Traits<ConcreteType>...>(op,
+                                                                     operands);
     };
   }
   /// Return the result of folding a single result operation that defines a
   /// `fold` method.
   template <typename ConcreteOpT>
-  static LogicalResult
-  foldSingleResultHook(Operation *op, ArrayRef<Attribute> operands,
-                       SmallVectorImpl<OpFoldResult> &results) {
+  static OpFoldResults foldSingleResultHook(Operation *op,
+                                            ArrayRef<Attribute> operands) {
     OpFoldResult result;
     if constexpr (has_fold_adaptor_single_result_v<ConcreteOpT>) {
       result = cast<ConcreteOpT>(op).fold(
@@ -1941,39 +1960,34 @@ class Op : public OpState, public Traits<ConcreteType>... {
       result = cast<ConcreteOpT>(op).fold(operands);
     }
 
-    // If the fold failed or was in-place, try to fold the traits of the
-    // operation.
-    if (!result ||
-        llvm::dyn_cast_if_present<Value>(result) == op->getResult(0)) {
-      if (succeeded(op_definition_impl::foldTraits<Traits<ConcreteType>...>(
-              op, operands, results)))
-        return success();
-      return success(static_cast<bool>(result));
-    }
-    results.push_back(result);
-    return success();
-  }
-  /// Return the result of folding an operation that defines a `fold` method.
+    // A null result is a failure. The op's own result means that the fold
+    // changed the op in place.
+    OpFoldResults own =
+        llvm::dyn_cast_if_present<Value>(result) == op->getResult(0)
+            ? OpFoldResults(success())
+            : OpFoldResults(result);
+    return op_definition_impl::foldTraits<Traits<ConcreteType>...>(
+        op, operands, std::move(own));
+  }
+  /// Return the result of folding an operation that defines a legacy
+  /// multi-result `fold` method, with the strict legacy contract.
   template <typename ConcreteOpT>
-  static LogicalResult foldHook(Operation *op, ArrayRef<Attribute> operands,
-                                SmallVectorImpl<OpFoldResult> &results) {
-    auto result = LogicalResult::failure();
+  static OpFoldResults foldHook(Operation *op, ArrayRef<Attribute> operands) {
+    SmallVector<OpFoldResult, 2> results;
+    auto status = LogicalResult::failure();
     if constexpr (has_fold_adaptor_v<ConcreteOpT>) {
-      result = cast<ConcreteOpT>(op).fold(
+      status = cast<ConcreteOpT>(op).fold(
           typename ConcreteOpT::FoldAdaptor(operands, cast<ConcreteOpT>(op)),
           results);
     } else {
-      result = cast<ConcreteOpT>(op).fold(operands, results);
+      status = cast<ConcreteOpT>(op).fold(operands, results);
     }
 
-    // If the fold failed or was in-place, try to fold the traits of the
-    // operation.
-    if (failed(result) || results.empty()) {
-      if (succeeded(op_definition_impl::foldTraits<Traits<ConcreteType>...>(
-              op, operands, results)))
-        return success();
-    }
-    return result;
+    OpFoldResults own =
+        ::mlir::detail::convertLegacyFoldResults(status, results);
+    own.normalize(op);
+    return op_definition_impl::foldTraits<Traits<ConcreteType>...>(
+        op, operands, std::move(own));
   }
 
   static constexpr HasTraitFn getHasTraitFn() {
diff --git a/mlir/include/mlir/IR/Operation.h b/mlir/include/mlir/IR/Operation.h
index 3ea9ad8ca0cfc4..16b79a11465b8f 100644
--- a/mlir/include/mlir/IR/Operation.h
+++ b/mlir/include/mlir/IR/Operation.h
@@ -767,6 +767,21 @@ class alignas(8) Operation final
   // Accessors for various properties of operations
   //===--------------------------------------------------------------------===//
 
+  /// Attempt to fold this operation with the specified constant operand values
+  /// - the elements in "operands" will correspond directly to the operands of
+  /// the operation, but may be null if non-constant.
+  ///
+  /// The result may replace some, all, or none of the results of this
+  /// operation, and may also record that this operation was modified in place.
+  /// The result is normalized: it has either no replacement or one replacement
+  /// per result, and a null replacement means "keep the result". See
+  /// `OperationName::foldHook` for the contract.
+  OpFoldResults fold(ArrayRef<Attribute> operands);
+
+  /// Attempt to fold this operation. Same as above, but computes the constant
+  /// operand values.
+  OpFoldResults fold();
+
   /// Attempt to fold this operation with the specified constant operand values
   /// - the elements in "operands" will correspond directly to the operands of
   /// the operation, but may be null if non-constant.
@@ -776,6 +791,9 @@ class alignas(8) Operation final
   ///   `results` is empty.
   /// * Otherwise, `results` is filled with the folded results.
   /// If folding was unsuccessful, this function returns "failure".
+  /// If the fold replaces only some results, those replacements are dropped:
+  /// this function returns "success" with an empty `results` if this operation
+  /// was modified in place, and "failure" otherwise.
   LogicalResult fold(ArrayRef<Attribute> operands,
                      SmallVectorImpl<OpFoldResult> &results);
 
@@ -786,6 +804,9 @@ class alignas(8) Operation final
   ///   `results` is empty.
   /// * Otherwise, `results` is filled with the folded results.
   /// If folding was unsuccessful, this function returns "failure".
+  /// If the fold replaces only some results, those replacements are dropped:
+  /// this function returns "success" with an empty `results` if this operation
+  /// was modified in place, and "failure" otherwise.
   LogicalResult fold(SmallVectorImpl<OpFoldResult> &results);
 
   /// Returns true if `InterfaceT` has been promised by the dialect or
diff --git a/mlir/include/mlir/IR/OperationSupport.h b/mlir/include/mlir/IR/OperationSupport.h
index 991d9847672e2e..7248903860788a 100644
--- a/mlir/include/mlir/IR/OperationSupport.h
+++ b/mlir/include/mlir/IR/OperationSupport.h
@@ -20,6 +20,7 @@
 #include "mlir/IR/Diagnostics.h"
 #include "mlir/IR/DialectRegistry.h"
 #include "mlir/IR/Location.h"
+#include "mlir/IR/OpFoldResult.h"
 #include "mlir/IR/TypeRange.h"
 #include "mlir/IR/Types.h"
 #include "mlir/IR/Value.h"
@@ -52,7 +53,6 @@ class OpAsmParser;
 class OpAsmPrinter;
 class OperandRange;
 class OperandRangeRange;
-class OpFoldResult;
 class Pattern;
 class Region;
 class ResultRange;
@@ -131,8 +131,8 @@ class PropertyRef {
 
 class OperationName {
 public:
-  using FoldHookFn = llvm::unique_function<LogicalResult(
-      Operation *, ArrayRef<Attribute>, SmallVectorImpl<OpFoldResult> &) const>;
+  using FoldHookFn = llvm::unique_function<OpFoldResults(
+      Operation *, ArrayRef<Attribute>) const>;
   using HasTraitFn = llvm::unique_function<bool(TypeID) const>;
   using ParseAssemblyFn =
       llvm::unique_function<ParseResult(OpAsmParser &, OperationState &)>;
@@ -154,8 +154,7 @@ class OperationName {
   /// may not be populated.
   struct InterfaceConcept {
     virtual ~InterfaceConcept() = default;
-    virtual LogicalResult foldHook(Operation *, ArrayRef<Attribute>,
-                                   SmallVectorImpl<OpFoldResult> &) = 0;
+    virtual OpFoldResults foldHook(Operation *, ArrayRef<Attribute>) = 0;
     virtual void getCanonicalizationPatterns(RewritePatternSet &,
                                              MLIRContext *) = 0;
     virtual bool hasTrait(TypeID) = 0;
@@ -247,8 +246,7 @@ class OperationName {
         : Impl(name, dialect, typeID, std::move(interfaceMap)) {
       propertiesTypeID = TypeID::get<Attribute>();
     }
-    LogicalResult foldHook(Operation *, ArrayRef<Attribute>,
-                           SmallVectorImpl<OpFoldResult> &) final;
+    OpFoldResults foldHook(Operation *, ArrayRef<Attribute>) final;
     void getCanonicalizationPatterns(RewritePatternSet &, MLIRContext *) final;
     bool hasTrait(TypeID) final;
     OperationName::ParseAssemblyFn getParseAssemblyFn() final;
@@ -297,24 +295,34 @@ class OperationName {
   /// can implement this to provide simplifications rules that are applied by
   /// the Builder::createOrFold API and the canonicalization pass.
   ///
-  /// This is an intentionally limited interface - implementations of this
-  /// hook can only perform the following changes to the operation:
+  /// This is an intentionally limited interface. The returned OpFoldResults
+  /// holds either no replacement or one replacement per result of the
+  /// operation, and an in-place bit:
   ///
-  ///  1. They can leave the operation alone and without changing the IR, and
-  ///     return failure.
-  ///  2. They can mutate the operation in place, without changing anything
-  ///     else in the IR. In this case, return success.
-  ///  3. They can return a list of existing values that can be used instead
-  ///     of the operation. In this case, fill in the results list and return
-  ///     success. The caller will remove the operation and use those results
-  ///     instead.
+  ///  1. A replacement is an Attribute (replace the result with a constant),
+  ///     a Value (replace the result with that value), or null or the result
+  ///     itself (keep the result). Replacement i may be result j only if
+  ///     result j is kept.
+  ///  2. A failure replaces no result and has no in-place mark. The IR must be
+  ///     unchanged.
+  ///  3. The hook can mutate the operation in place, without changing anything
+  ///     else in the IR. In this case, it marks the result as modified in
+  ///     place. The operation must still verify.
+  ///  4. The hook can replace some but not all results and also mutate the
+  ///     operation in place. The operation stays.
+  ///  5. If the operation has results and every result is replaced, the
+  ///     caller removes the operation and uses the replacements instead, even
+  ///     if the hook also mutated it in place.
+  ///
+  /// The hook creates no operations and changes no IR outside the operation. A
+  /// replacement Value must exist before the fold, must have the type of the
+  /// result that it replaces, and must dominate the operation.
   ///
   /// This allows expression of some simple in-place canonicalizations (e.g.
   /// "x+0 -> x", "min(x,y,x,z) -> min(x,y,z)", "x+y-x -> y", etc), as well as
   /// generalized constant folding.
-  LogicalResult foldHook(Operation *op, ArrayRef<Attribute> operands,
-                         SmallVectorImpl<OpFoldResult> &results) const {
-    return getImpl()->foldHook(op, operands, results);
+  OpFoldResults foldHook(Operation *op, ArrayRef<Attribute> operands) const {
+    return getImpl()->foldHook(op, operands);
   }
 
   /// This hook returns any canonicalization pattern rewrites that the
@@ -600,9 +608,8 @@ class RegisteredOperationName : public OperationName {
                TypeID::get<ConcreteOp>(), ConcreteOp::getInterfaceMap()) {
       propertiesTypeID = TypeID::get<Properties>();
     }
-    LogicalResult foldHook(Operation *op, ArrayRef<Attribute> attrs,
-                           SmallVectorImpl<OpFoldResult> &results) final {
-      return ConcreteOp::getFoldHookFn()(op, attrs, results);
+    OpFoldResults foldHook(Operation *op, ArrayRef<Attribute> attrs) final {
+      return ConcreteOp::getFoldHookFn()(op, attrs);
     }
     void getCanonicalizationPatterns(RewritePatternSet &set,
                                      MLIRContext *context) final {
diff --git a/mlir/lib/IR/ExtensibleDialect.cpp b/mlir/lib/IR/ExtensibleDialect.cpp
index b51d3301606468..c6b36a765973b5 100644
--- a/mlir/lib/IR/ExtensibleDialect.cpp
+++ b/mlir/lib/IR/ExtensibleDialect.cpp
@@ -285,6 +285,17 @@ void DynamicAttr::print(AsmPrinter &printer) {
 // Dynamic operation
 //===----------------------------------------------------------------------===//
 
+/// Wrap a legacy fold hook into a fold hook that returns OpFoldResults.
+static OperationName::FoldHookFn
+adaptLegacyFoldHookFn(DynamicOpDefinition::LegacyFoldHookFn &&foldHookFn) {
+  return [foldHookFn = std::move(foldHookFn)](
+             Operation *op, ArrayRef<Attribute> operands) -> OpFoldResults {
+    SmallVector<OpFoldResult> results;
+    LogicalResult status = foldHookFn(op, operands, results);
+    return detail::convertLegacyFoldResults(status, results);
+  };
+}
+
 DynamicOpDefinition::DynamicOpDefinition(
     StringRef name, ExtensibleDialect *dialect,
     OperationName::VerifyInvariantsFn &&verifyFn,
@@ -333,8 +344,8 @@ std::unique_ptr<DynamicOpDefinition> DynamicOpDefinition::get(
     OperationName::VerifyRegionInvariantsFn &&verifyRegionFn,
     OperationName::ParseAssemblyFn &&parseFn,
     OperationName::PrintAssemblyFn &&printFn) {
-  auto foldHookFn = [](Operation *op, ArrayRef<Attribute> operands,
-                       SmallVectorImpl<OpFoldResult> &results) {
+  auto foldHookFn = [](Operation *op,
+                       ArrayRef<Attribute> operands) -> OpFoldResults {
     return failure();
   };
 
@@ -366,6 +377,32 @@ std::unique_ptr<DynamicOpDefinition> DynamicOpDefinition::get(
       std::move(populateDefaultAttrsFn)));
 }
 
+std::unique_ptr<DynamicOpDefinition> DynamicOpDefinition::get(
+    StringRef name, ExtensibleDialect *dialect,
+    OperationName::VerifyInvariantsFn &&verifyFn,
+    OperationName::VerifyRegionInvariantsFn &&verifyRegionFn,
+    OperationName::ParseAssemblyFn &&parseFn,
+    OperationName::PrintAssemblyFn &&printFn, LegacyFoldHookFn &&foldHookFn,
+    GetCanonicalizationPatternsFn &&getCanonicalizationPatternsFn,
+    OperationName::PopulateDefaultAttrsFn &&populateDefaultAttrsFn) {
+  return DynamicOpDefinition::get(name, dialect, std::move(verifyFn),
+                                  std::move(verifyRegionFn), std::move(parseFn),
+                                  std::move(printFn),
+                                  adaptLegacyFoldHookFn(std::move(foldHookFn)),
+                                  std::move(getCanonicalizationPatternsFn),
+                                  std::move(populateDefaultAttrsFn));
+}
+
+void DynamicOpDefinition::setFoldHookFn(LegacyFoldHookFn &&foldHook) {
+  foldHookFn = adaptLegacyFoldHookFn(std::move(foldHook));
+}
+
+void DynamicOpDefinition::setFoldHookFn(std::nullptr_t) {
+  foldHookFn = [](Operation *, ArrayRef<Attribute>) -> OpFoldResults {
+    return failure();
+  };
+}
+
 //===----------------------------------------------------------------------===//
 // Extensible dialect
 //===----------------------------------------------------------------------===//
diff --git a/mlir/lib/IR/MLIRContext.cpp b/mlir/lib/IR/MLIRContext.cpp
index 94cf420070131c..f64240bb0a583a 100644
--- a/mlir/lib/IR/MLIRContext.cpp
+++ b/mlir/lib/IR/MLIRContext.cpp
@@ -949,9 +949,8 @@ StringRef OperationName::getDialectNamespace() const {
   return getStringRef().split('.').first;
 }
 
-LogicalResult
-OperationName::UnregisteredOpModel::foldHook(Operation *, ArrayRef<Attribute>,
-                                             SmallVectorImpl<OpFoldResult> &) {
+OpFoldResults
+OperationName::UnregisteredOpModel::foldHook(Operation *, ArrayRef<Attribute>) {
   return failure();
 }
 void OperationName::UnregisteredOpModel::getCanonicalizationPatterns(
diff --git a/mlir/lib/IR/Operation.cpp b/mlir/lib/IR/Operation.cpp
index d262e844603a1a..e771d3f194bcc4 100644
--- a/mlir/lib/IR/Operation.cpp
+++ b/mlir/lib/IR/Operation.cpp
@@ -600,63 +600,53 @@ void Operation::setSuccessor(Block *block, unsigned index) {
   getBlockOperands()[index].set(block);
 }
 
-#ifndef NDEBUG
-/// Assert that the folded results (in case of values) have the same type as
-/// the results of the given op.
-static void checkFoldResultTypes(Operation *op,
-                                 SmallVectorImpl<OpFoldResult> &results) {
-  if (results.empty())
-    return;
-
-  for (auto [ofr, opResult] : llvm::zip_equal(results, op->getResults())) {
-    if (auto value = dyn_cast<Value>(ofr)) {
-      if (value.getType() != opResult.getType()) {
-        op->emitOpError() << "folder produced a value of incorrect type: "
-                          << value.getType()
-                          << ", expected: " << opResult.getType();
-        assert(false && "incorrect fold result type");
-      }
-    }
-  }
-}
-#endif // NDEBUG
-
 /// Attempt to fold this operation using the Op's registered foldHook.
-LogicalResult Operation::fold(ArrayRef<Attribute> operands,
-                              SmallVectorImpl<OpFoldResult> &results) {
+OpFoldResults Operation::fold(ArrayRef<Attribute> operands) {
   // If we have a registered operation definition matching this one, use it to
   // try to constant fold the operation.
-  if (succeeded(name.foldHook(this, operands, results))) {
-#ifndef NDEBUG
-    checkFoldResultTypes(this, results);
-#endif // NDEBUG
-    return success();
-  }
+  OpFoldResults results = name.foldHook(this, operands);
+  results.normalize(this);
+  if (results.succeeded())
+    return results;
 
   // Otherwise, fall back on the dialect hook to handle it.
   Dialect *dialect = getDialect();
   if (!dialect)
-    return failure();
+    return results;
 
   auto *interface = dyn_cast<DialectFoldInterface>(dialect);
   if (!interface)
-    return failure();
+    return results;
 
-  LogicalResult status = interface->fold(this, operands, results);
-#ifndef NDEBUG
-  if (succeeded(status))
-    checkFoldResultTypes(this, results);
-#endif // NDEBUG
-  return status;
+  SmallVector<OpFoldResult> legacyResults;
+  LogicalResult status = interface->fold(this, operands, legacyResults);
+  results = detail::convertLegacyFoldResults(status, legacyResults);
+  results.normalize(this);
+  return results;
+}
+
+/// Compute the constant operand values of `op`.
+static SmallVector<Attribute> getConstantOperands(Operation *op) {
+  SmallVector<Attribute> constants(op->getNumOperands(), Attribute());
+  for (unsigned i = 0, e = op->getNumOperands(); i != e; ++i)
+    matchPattern(op->getOperand(i), m_Constant(&constants[i]));
+  return constants;
+}
+
+OpFoldResults Operation::fold() { return fold(getConstantOperands(this)); }
+
+LogicalResult Operation::fold(ArrayRef<Attribute> operands,
+                              SmallVectorImpl<OpFoldResult> &results) {
+  OpFoldResults foldResults = fold(operands);
+  if (foldResults.replacesAll()) {
+    llvm::append_range(results, foldResults.getReplacements());
+    return success();
+  }
+  return success(foldResults.modifiedInPlace());
 }
 
 LogicalResult Operation::fold(SmallVectorImpl<OpFoldResult> &results) {
-  // Check if any operands are constants.
-  SmallVector<Attribute> constants;
-  constants.assign(getNumOperands(), Attribute());
-  for (unsigned i = 0, e = getNumOperands(); i != e; ++i)
-    matchPattern(getOperand(i), m_Constant(&constants[i]));
-  return fold(constants, results);
+  return fold(getConstantOperands(this), results);
 }
 
 /// Emit an error with the op name prefixed, like "'dim' op " which is
diff --git a/mlir/test/Dialect/Builtin/canonicalize.mlir b/mlir/test/Dialect/Builtin/canonicalize.mlir
index 2e36b7ee371c32..6e38d469fb91dd 100644
--- a/mlir/test/Dialect/Builtin/canonicalize.mlir
+++ b/mlir/test/Dialect/Builtin/canonicalize.mlir
@@ -23,3 +23,16 @@ func.func @multiple_conversion_casts_failure(%arg0: i32, %arg1: i32, %arg2: i64)
   %outputs:2 = builtin.unrealized_conversion_cast %arg2, %inputs#1 : i64, i64 to i32, i32
   return %outputs#0, %outputs#1 : i32, i32
 }
+
+// In a graph region, a cast can forward its own result. The fold keeps that
+// result, so the cast stays.
+// CHECK-LABEL: func @graph_region_self_forward
+//       CHECK:   test.graph_region
+//       CHECK:     %[[CAST:.+]] = builtin.unrealized_conversion_cast %[[CAST]] : i32 to i32
+func.func @graph_region_self_forward() {
+  test.graph_region {
+    %0 = builtin.unrealized_conversion_cast %0 : i32 to i32
+    "test.use"(%0) : (i32) -> ()
+  }
+  return
+}
diff --git a/mlir/unittests/IR/OpFoldResultsTest.cpp b/mlir/unittests/IR/OpFoldResultsTest.cpp
index 8cb90ed967c739..bd118d54a2738a 100644
--- a/mlir/unittests/IR/OpFoldResultsTest.cpp
+++ b/mlir/unittests/IR/OpFoldResultsTest.cpp
@@ -9,26 +9,96 @@
 #include "mlir/IR/Builders.h"
 #include "mlir/IR/BuiltinAttributes.h"
 #include "mlir/IR/BuiltinTypes.h"
+#include "mlir/IR/Dialect.h"
+#include "mlir/IR/ExtensibleDialect.h"
 #include "mlir/IR/MLIRContext.h"
+#include "mlir/IR/OpDefinition.h"
 #include "mlir/IR/OpFoldResult.h"
 #include "mlir/IR/Operation.h"
+#include "mlir/Interfaces/FoldInterfaces.h"
 #include "gtest/gtest.h"
 
+#include <functional>
+
 using namespace mlir;
 
+// The fallback TypeID resolver rejects a trait template in an anonymous
+// namespace, so the test dialect lives in a named namespace.
+namespace op_fold_results_test {
+using LegacyFoldFn =
+    std::function<LogicalResult(Operation *, SmallVectorImpl<OpFoldResult> &)>;
+
+/// Per-test behavior of the ops, the trait, and the dialect interface below.
+struct FoldState {
+  LegacyFoldFn traitFoldFn;
+  LegacyFoldFn dialectFoldFn;
+  unsigned traitCalls = 0;
+  unsigned dialectCalls = 0;
+};
+
+/// The state of the running test. The test fixture owns it.
+static FoldState *foldState = nullptr;
+
+template <typename ConcreteType>
+struct LegacyFoldTrait
+    : public OpTrait::TraitBase<ConcreteType, LegacyFoldTrait> {
+  static LogicalResult foldTrait(Operation *op, ArrayRef<Attribute>,
+                                 SmallVectorImpl<OpFoldResult> &results) {
+    ++foldState->traitCalls;
+    return foldState->traitFoldFn ? foldState->traitFoldFn(op, results)
+                                  : failure();
+  }
+};
+
+/// An op with two results and no fold method. Only its trait folds.
+struct PartialFoldOp : public Op<PartialFoldOp, OpTrait::NResults<2>::Impl,
+                                 OpTrait::VariadicOperands, LegacyFoldTrait> {
+  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(PartialFoldOp)
+  using Op::Op;
+  static ArrayRef<StringRef> getAttributeNames() { return {}; }
+  static StringRef getOperationName() { return "fold_test.partial"; }
+};
+
+struct TestFoldInterface : public DialectFoldInterface {
+  using DialectFoldInterface::DialectFoldInterface;
+  LogicalResult fold(Operation *op, ArrayRef<Attribute>,
+                     SmallVectorImpl<OpFoldResult> &results) const final {
+    ++foldState->dialectCalls;
+    return foldState->dialectFoldFn ? foldState->dialectFoldFn(op, results)
+                                    : failure();
+  }
+};
+
+struct FoldTestDialect : public Dialect {
+  MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(FoldTestDialect)
+  static constexpr StringLiteral getDialectNamespace() { return "fold_test"; }
+  explicit FoldTestDialect(MLIRContext *context)
+      : Dialect(getDialectNamespace(), context,
+                TypeID::get<FoldTestDialect>()) {
+    addOperations<PartialFoldOp>();
+    addInterfaces<TestFoldInterface>();
+  }
+};
+} // namespace op_fold_results_test
+
+using namespace op_fold_results_test;
+
 namespace {
 class OpFoldResultsTest : public ::testing::Test {
 protected:
   OpFoldResultsTest() : builder(&context) {
     context.allowUnregisteredDialects();
+    context.loadDialect<FoldTestDialect>();
     i32 = builder.getI32Type();
     f32 = builder.getF32Type();
+    foldState = &state;
   }
 
   ~OpFoldResultsTest() override {
     // Destroy users before the ops that define their operands.
     for (Operation *op : llvm::reverse(ops))
       op->destroy();
+    foldState = nullptr;
   }
 
   /// Create an op with the given result types and operands. The fixture
@@ -43,11 +113,60 @@ class OpFoldResultsTest : public ::testing::Test {
     return op;
   }
 
+  /// Load a dynamic dialect. The fold hook of `test_fold.op` calls `foldFn`.
+  /// The legacy fold hooks of `test_fold.legacy_op` and `test_fold.get_op`
+  /// call `legacyFoldFn`; `test_fold.get_op` uses the legacy overload of
+  /// DynamicOpDefinition::get. `test_fold.no_fold_op` has no fold hook.
+  void loadDynamicDialect() {
+    context.getOrLoadDynamicDialect("test_fold", [&](DynamicDialect *dialect) {
+      auto verify = [](Operation *) { return success(); };
+      std::unique_ptr<DynamicOpDefinition> opDef =
+          DynamicOpDefinition::get("op", dialect, verify, verify);
+      opDef->setFoldHookFn(
+          [this](Operation *op, ArrayRef<Attribute>) { return foldFn(op); });
+      dialect->registerDynamicOp(std::move(opDef));
+
+      std::unique_ptr<DynamicOpDefinition> legacyOpDef =
+          DynamicOpDefinition::get("legacy_op", dialect, verify, verify);
+      legacyOpDef->setFoldHookFn(
+          [this](Operation *op, ArrayRef<Attribute>,
+                 SmallVectorImpl<OpFoldResult> &results) {
+            return legacyFoldFn(op, results);
+          });
+      dialect->registerDynamicOp(std::move(legacyOpDef));
+
+      DynamicOpDefinition::LegacyFoldHookFn legacyFold =
+          [this](Operation *op, ArrayRef<Attribute>,
+                 SmallVectorImpl<OpFoldResult> &results) {
+            return legacyFoldFn(op, results);
+          };
+      dialect->registerDynamicOp(DynamicOpDefinition::get(
+          "get_op", dialect, verify, verify,
+          [](OpAsmParser &, OperationState &) -> ParseResult {
+            return failure();
+          },
+          [](Operation *, OpAsmPrinter &, StringRef) {}, std::move(legacyFold),
+          [](RewritePatternSet &, MLIRContext *) {},
+          [](const OperationName &, NamedAttrList &) {}));
+
+      std::unique_ptr<DynamicOpDefinition> noFoldOpDef =
+          DynamicOpDefinition::get("no_fold_op", dialect, verify, verify);
+      noFoldOpDef->setFoldHookFn(
+          [this](Operation *op, ArrayRef<Attribute>) { return foldFn(op); });
+      noFoldOpDef->setFoldHookFn(nullptr);
+      dialect->registerDynamicOp(std::move(noFoldOpDef));
+    });
+  }
+
   MLIRContext context;
   Builder builder;
   Type i32;
   Type f32;
   SmallVector<Operation *> ops;
+  std::function<OpFoldResults(Operation *)> foldFn;
+  std::function<LogicalResult(Operation *, SmallVectorImpl<OpFoldResult> &)>
+      legacyFoldFn;
+  FoldState state;
 };
 } // namespace
 
@@ -411,6 +530,174 @@ TEST_F(OpFoldResultsTest, ZeroResultInPlaceDoesNotReplaceAll) {
   EXPECT_FALSE(failedResult.replacesAll());
 }
 
+TEST_F(OpFoldResultsTest, UnregisteredOpFoldFails) {
+  Operation *op = createOp({i32});
+  EXPECT_TRUE(op->fold().failed());
+  SmallVector<OpFoldResult> results;
+  EXPECT_TRUE(failed(op->fold(results)));
+  EXPECT_TRUE(results.empty());
+}
+
+TEST_F(OpFoldResultsTest, LegacyOperationFoldKeepsStrictContract) {
+  loadDynamicDialect();
+  Operation *producer = createOp({i32, i32});
+  Operation *op = createOp({i32, i32}, "test_fold.op");
+  Attribute attr = builder.getI32IntegerAttr(1);
+  SmallVector<OpFoldResult> results;
+
+  foldFn = [](Operation *) -> OpFoldResults { return failure(); };
+  EXPECT_TRUE(failed(op->fold(results)));
+  EXPECT_TRUE(results.empty());
+  EXPECT_TRUE(op->fold().failed());
+
+  foldFn = [](Operation *) -> OpFoldResults { return success(); };
+  EXPECT_TRUE(succeeded(op->fold(results)));
+  EXPECT_TRUE(results.empty());
+
+  foldFn = [&](Operation *) -> OpFoldResults {
+    return {attr, producer->getResult(1)};
+  };
+  EXPECT_TRUE(succeeded(op->fold(results)));
+  ASSERT_EQ(results.size(), 2u);
+  EXPECT_EQ(results[0], OpFoldResult(attr));
+  EXPECT_EQ(results[1], OpFoldResult(producer->getResult(1)));
+  results.clear();
+
+  // The legacy overloads do not apply a partial fold. Without an in-place
+  // change, the partial fold is a failure.
+  foldFn = [&](Operation *foldedOp) {
+    OpFoldResults partial(foldedOp);
+    partial.replace(1u, attr);
+    return partial;
+  };
+  EXPECT_TRUE(failed(op->fold(results)));
+  EXPECT_TRUE(results.empty());
+  OpFoldResults partialResult = op->fold();
+  EXPECT_TRUE(partialResult.succeeded());
+  EXPECT_FALSE(partialResult.modifiedInPlace());
+  ASSERT_EQ(partialResult.size(), 2u);
+  EXPECT_FALSE(partialResult[0]);
+  EXPECT_EQ(partialResult[1], OpFoldResult(attr));
+
+  // With an in-place change, the partial fold is reported as in place.
+  foldFn = [&](Operation *foldedOp) {
+    OpFoldResults partial(foldedOp);
+    partial.replace(1u, attr);
+    partial.setModifiedInPlace();
+    return partial;
+  };
+  EXPECT_TRUE(succeeded(op->fold(results)));
+  EXPECT_TRUE(results.empty());
+
+  // A fold that keeps every result and is not in place is a failure.
+  foldFn = [](Operation *foldedOp) -> OpFoldResults {
+    return foldedOp->getResults();
+  };
+  EXPECT_TRUE(failed(op->fold(results)));
+  EXPECT_TRUE(results.empty());
+  EXPECT_TRUE(op->fold().failed());
+}
+
+TEST_F(OpFoldResultsTest, LegacyDynamicFoldHookKeepsStrictContract) {
+  loadDynamicDialect();
+  Operation *producer = createOp({i32, i32});
+  Operation *op = createOp({i32, i32}, "test_fold.legacy_op");
+  Attribute attr = builder.getI32IntegerAttr(1);
+
+  legacyFoldFn = [](Operation *, SmallVectorImpl<OpFoldResult> &) {
+    return failure();
+  };
+  EXPECT_TRUE(op->fold().failed());
+
+  legacyFoldFn = [](Operation *, SmallVectorImpl<OpFoldResult> &) {
+    return success();
+  };
+  OpFoldResults inPlace = op->fold();
+  EXPECT_TRUE(inPlace.succeeded());
+  EXPECT_TRUE(inPlace.modifiedInPlace());
+  EXPECT_FALSE(inPlace.replacesAny());
+
+  legacyFoldFn = [&](Operation *, SmallVectorImpl<OpFoldResult> &results) {
+    results.push_back(attr);
+    results.push_back(producer->getResult(0));
+    return success();
+  };
+  OpFoldResults all = op->fold();
+  EXPECT_TRUE(all.succeeded());
+  EXPECT_FALSE(all.modifiedInPlace());
+  EXPECT_TRUE(all.replacesAll());
+  ASSERT_EQ(all.size(), 2u);
+  EXPECT_EQ(all[0], OpFoldResult(attr));
+  EXPECT_EQ(all[1], OpFoldResult(producer->getResult(0)));
+
+  // A legacy fold that forwards the op's own results keeps every result, so
+  // it is a failure.
+  legacyFoldFn = [](Operation *foldedOp,
+                    SmallVectorImpl<OpFoldResult> &results) {
+    llvm::append_range(results, foldedOp->getResults());
+    return success();
+  };
+  EXPECT_TRUE(op->fold().failed());
+  SmallVector<OpFoldResult> results;
+  EXPECT_TRUE(failed(op->fold(results)));
+  EXPECT_TRUE(results.empty());
+}
+
+TEST_F(OpFoldResultsTest, LegacyDynamicFoldHookMayForwardAnotherResult) {
+  loadDynamicDialect();
+  Operation *op = createOp({i32, i32}, "test_fold.legacy_op");
+  Attribute attr = builder.getI32IntegerAttr(1);
+
+  // Replacement 0 is result 1, which the fold also replaces. The legacy
+  // adapter does not reject this.
+  legacyFoldFn = [&](Operation *foldedOp,
+                     SmallVectorImpl<OpFoldResult> &results) {
+    results.push_back(foldedOp->getResult(1));
+    results.push_back(attr);
+    return success();
+  };
+  OpFoldResults result = op->fold();
+  EXPECT_TRUE(result.replacesAll());
+  ASSERT_EQ(result.size(), 2u);
+  EXPECT_EQ(result[0], OpFoldResult(op->getResult(1)));
+  EXPECT_EQ(result[1], OpFoldResult(attr));
+
+  SmallVector<OpFoldResult> results;
+  EXPECT_TRUE(succeeded(op->fold(results)));
+  ASSERT_EQ(results.size(), 2u);
+  EXPECT_EQ(results[0], OpFoldResult(op->getResult(1)));
+  EXPECT_EQ(results[1], OpFoldResult(attr));
+}
+
+TEST_F(OpFoldResultsTest, LegacyDynamicOpDefinitionGet) {
+  loadDynamicDialect();
+  Operation *op = createOp({i32}, "test_fold.get_op");
+  Attribute attr = builder.getI32IntegerAttr(1);
+
+  legacyFoldFn = [&](Operation *, SmallVectorImpl<OpFoldResult> &results) {
+    results.push_back(attr);
+    return success();
+  };
+  OpFoldResults result = op->fold();
+  expectOneReplacement(result, attr);
+
+  legacyFoldFn = [](Operation *, SmallVectorImpl<OpFoldResult> &) {
+    return failure();
+  };
+  EXPECT_TRUE(op->fold().failed());
+}
+
+TEST_F(OpFoldResultsTest, NullFoldHookFails) {
+  loadDynamicDialect();
+  Operation *op = createOp({i32}, "test_fold.no_fold_op");
+  // The removed hook would report an in-place fold.
+  foldFn = [](Operation *) -> OpFoldResults { return success(); };
+  EXPECT_TRUE(op->fold().failed());
+  SmallVector<OpFoldResult> results;
+  EXPECT_TRUE(failed(op->fold(results)));
+  EXPECT_TRUE(results.empty());
+}
+
 TEST_F(OpFoldResultsTest, FreeHelpersMatchMembers) {
   Operation *op = createOp({i32, i32});
   OpFoldResults partial(op);
@@ -426,6 +713,38 @@ TEST_F(OpFoldResultsTest, FreeHelpersMatchMembers) {
   EXPECT_TRUE(failed(failedResult));
 }
 
+TEST_F(OpFoldResultsTest, DialectFoldInterfaceFallback) {
+  Operation *op = createOp({i32, i32}, "fold_test.partial");
+  Attribute attr = builder.getI32IntegerAttr(1);
+  foldState->dialectFoldFn = [&](Operation *,
+                                 SmallVectorImpl<OpFoldResult> &results) {
+    results.append(2, attr);
+    return success();
+  };
+
+  // The traits fail.
+  OpFoldResults result = op->fold();
+  EXPECT_TRUE(result.replacesAll());
+  ASSERT_EQ(result.size(), 2u);
+  EXPECT_EQ(result[0], OpFoldResult(attr));
+  EXPECT_EQ(foldState->traitCalls, 1u);
+  EXPECT_EQ(foldState->dialectCalls, 1u);
+  SmallVector<OpFoldResult> results;
+  EXPECT_TRUE(succeeded(op->fold(results)));
+  EXPECT_EQ(results.size(), 2u);
+  EXPECT_EQ(foldState->dialectCalls, 2u);
+
+  // The fallback result is normalized.
+  foldState->dialectCalls = 0;
+  foldState->dialectFoldFn = [](Operation *foldedOp,
+                                SmallVectorImpl<OpFoldResult> &results) {
+    llvm::append_range(results, foldedOp->getResults());
+    return success();
+  };
+  EXPECT_TRUE(op->fold().failed());
+  EXPECT_EQ(foldState->dialectCalls, 1u);
+}
+
 #ifdef GTEST_HAS_DEATH_TEST
 #ifndef NDEBUG
 namespace {
@@ -448,5 +767,16 @@ TEST_F(OpFoldResultsDeathTest, ReplacementCountMismatch) {
                "expected one replacement per operation result");
 }
 
+TEST_F(OpFoldResultsDeathTest, LegacyNullResult) {
+  loadDynamicDialect();
+  Operation *op = createOp({i32, i32}, "test_fold.legacy_op");
+  Attribute attr = builder.getI32IntegerAttr(1);
+  legacyFoldFn = [&](Operation *, SmallVectorImpl<OpFoldResult> &results) {
+    results.push_back(attr);
+    results.push_back(OpFoldResult());
+    return success();
+  };
+  EXPECT_DEATH((void)op->fold(), "legacy fold returned a null result");
+}
 #endif // NDEBUG
 #endif // GTEST_HAS_DEATH_TEST



More information about the llvm-branch-commits mailing list