[llvm-branch-commits] [mlir] [mlir] Add `OpFoldResults` definition (PR #228755)

Victor Perez via llvm-branch-commits llvm-branch-commits at lists.llvm.org
Sat Oct 3 11:44:09 PDT 2026


https://github.com/victor-eds created https://github.com/llvm/llvm-project/pull/228755

MLIR does not support partial folds for multi-result operations. This patch starts the enablement of this feature by adding the `OpFoldResults` new folder methods will return. Nothing uses the class yet.

The new class `OpFoldResults` in `mlir/IR/OpFoldResult.h` holds either no replacement or one replacement per result, and an in-place bit. For each replacement (of `OpFoldResult` type), a null replacement keeps its result. The class provides constructor for common constructs: `{}`, `nullptr`, and `failure()` are failures, and `success()` is an in-place change with no replacement. A helper `normalize` turns a replacement that is its own result into a null replacement, and it clears the replacements if no result is replaced. In assert builds, `normalize` also checks the type of each `Value` replacement. `normalize` won't be called from drivers, only by hooks, `fold` callers always get normalized `OpFoldResults`.

We include a `detail::convertLegacyFoldResults` migration utility while the legacy form is still allowed. This converts the result of a legacy vector fold to the new return type.

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

>From 51b7d3f7f45a36175c240d460fa23c1e81d3c5f7 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:36:38 -0700
Subject: [PATCH] [mlir] Add OpFoldResults
MIME-Version: 1.0
Content-Type: text/plain; charset=UTF-8
Content-Transfer-Encoding: 8bit

A fold of a multi-result op can only replace all results or change the
op in place. It cannot replace some results and keep the others. For
example, a loop that yields one of its iteration arguments unchanged
could replace the matching result with the initial value and keep the
other results. This patch adds the class that holds such a partial
fold. Nothing uses the class yet.

The new class `OpFoldResults` in `mlir/IR/OpFoldResult.h` holds either
no replacement or one replacement per result, and an in-place bit. The
plural name says that the object holds one `OpFoldResult` per result of
the op. A null replacement keeps its result. `{}`, `nullptr`, and
`failure()` are failures, and `success()` is an in-place change.
`normalize` turns a replacement that is its own result into a null
replacement, and it clears the replacements if no result is replaced. In
assert builds, `normalize` also checks the type of each `Value`
replacement.

`detail::convertLegacyFoldResults` converts the result of a legacy
vector fold with the strict legacy contract: failure, success with an
empty vector (in place), or success with one entry per result. A null
entry asserts.

The new unit tests in `OpFoldResultsTest.cpp` cover the constructors,
`normalize`, null safety, the free `succeeded` and `failed` helpers,
and the debug asserts of `normalize`.

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/OpFoldResult.h     | 105 +++++-
 mlir/lib/IR/OpFoldResult.cpp            | 125 +++++++
 mlir/unittests/IR/CMakeLists.txt        |   1 +
 mlir/unittests/IR/OpFoldResultsTest.cpp | 452 ++++++++++++++++++++++++
 4 files changed, 682 insertions(+), 1 deletion(-)
 create mode 100644 mlir/unittests/IR/OpFoldResultsTest.cpp

diff --git a/mlir/include/mlir/IR/OpFoldResult.h b/mlir/include/mlir/IR/OpFoldResult.h
index 0bda0e75403e2..f65b703206cdf 100644
--- a/mlir/include/mlir/IR/OpFoldResult.h
+++ b/mlir/include/mlir/IR/OpFoldResult.h
@@ -7,7 +7,8 @@
 //===----------------------------------------------------------------------===//
 //
 // This file defines OpFoldResult, the result of a fold of an operation with
-// exactly one result.
+// exactly one result, and OpFoldResults, the result of a fold of any
+// operation.
 //
 //===----------------------------------------------------------------------===//
 
@@ -18,9 +19,15 @@
 #include "mlir/IR/Value.h"
 #include "mlir/Support/LLVM.h"
 #include "llvm/ADT/PointerUnion.h"
+#include "llvm/ADT/STLExtras.h"
+#include "llvm/ADT/SmallVector.h"
 #include "llvm/Support/Compiler.h"
+#include <cstddef>
+#include <initializer_list>
+#include <type_traits>
 
 namespace mlir {
+class Operation;
 
 /// This class represents a single result from folding an operation.
 class OpFoldResult : public PointerUnion<Attribute, Value> {
@@ -60,6 +67,102 @@ namespace mlir {
 /// Allow printing to a stream.
 raw_ostream &operator<<(raw_ostream &os, OpFoldResult ofr);
 
+/// The result of a fold of any op. Replacement i is an Attribute (replace
+/// result i with a constant), a Value (replace result i with that value), or
+/// null / op->getResult(i) (keep result i). A separate bit records an in-place
+/// change of the op.
+class [[nodiscard]] OpFoldResults {
+public:
+  /// Failure: the fold did not apply and the IR is unchanged.
+  OpFoldResults() = default;
+  /// Failure, the same as the default constructor.
+  OpFoldResults(std::nullptr_t);
+  /// success(): the op changed in place. failure(): failure.
+  OpFoldResults(LogicalResult status);
+  /// One replacement. Valid only when the op has exactly one result at run
+  /// time. A null replacement, or the op's own result, keeps the result, so the
+  /// fold fails. Unlike `OpFoldResult fold`, the op's own result does not mean
+  /// in place; use success() or setModifiedInPlace() for an in-place change.
+  OpFoldResults(OpFoldResult replacement);
+  OpFoldResults(Value replacement);
+  OpFoldResults(Attribute replacement);
+  /// One replacement per result.
+  OpFoldResults(std::initializer_list<OpFoldResult> replacements);
+  /// One replacement per range element. An empty range is a failure.
+  template <typename RangeT,
+            typename = std::enable_if_t<
+                !std::is_convertible_v<RangeT, Attribute> &&
+                !std::is_convertible_v<RangeT, Value> &&
+                std::is_convertible_v<llvm::detail::ValueOfRange<RangeT>,
+                                      OpFoldResult>>>
+  OpFoldResults(RangeT &&range) {
+    llvm::append_range(replacements, range);
+  }
+  /// Incremental form: every result is kept and the op is not changed in
+  /// place.
+  explicit OpFoldResults(Operation *op);
+
+  /// Set the replacement of `result`, which must be a result of the op. If
+  /// another op owns `result`, or if normalize() already ran, the behavior is
+  /// undefined. The object must already hold one replacement per result, for
+  /// example from OpFoldResults(op). A null replacement, or
+  /// `replacement == result`, keeps the result. The last write wins.
+  void replace(Value result, OpFoldResult replacement);
+  /// Same as above, with the index of the result. normalize() maps a
+  /// replacement that is the op's own result to "keep".
+  void replace(unsigned resultIndex, OpFoldResult replacement);
+  /// Set whether the fold changed the op in place (operands, attributes,
+  /// properties, or regions).
+  void setModifiedInPlace(bool modified = true);
+
+  // Queries for drivers. They are valid after normalize().
+
+  /// Return true if the fold changed the op in place or replaces a result.
+  bool succeeded() const;
+  /// Return true if the fold did not apply.
+  bool failed() const;
+  /// Return true if the fold changed the op in place.
+  bool modifiedInPlace() const;
+  /// Return true if the fold replaces at least one result.
+  bool replacesAny() const;
+  /// Return true if the operation has at least one result and all the results
+  /// are replaced.
+  bool replacesAll() const;
+  /// Return the number of replacements: 0 if no replacement, or one per result
+  /// of the op.
+  unsigned size() const;
+  /// Return the replacement of result `i`; null means keep. If size() == 0,
+  /// every index reads as keep; otherwise `i` must be less than size().
+  OpFoldResult operator[](unsigned i) const;
+  /// Return the replacements. The returned range points into this object.
+  ArrayRef<OpFoldResult> getReplacements() const LLVM_LIFETIME_BOUND;
+
+  /// If replacement i matches op's result i, set replacement i to null. Clear
+  /// replacements if all end up null. Idempotent.
+  void normalize(Operation *op);
+
+private:
+  SmallVector<OpFoldResult, 2> replacements;
+  bool inPlace = false;
+};
+
+/// Return true if `result` changed the op in place or replaces a result.
+inline bool succeeded(const OpFoldResults &result) {
+  return result.succeeded();
+}
+
+/// Return true if `result` is a failure.
+inline bool failed(const OpFoldResults &result) { return result.failed(); }
+
+namespace detail {
+/// Convert the result of a legacy vector fold with the strict legacy contract:
+/// failure stays failure, success with an empty vector means "in place", and
+/// success with a full vector has one replacement per result. The result is
+/// then normalized.
+OpFoldResults convertLegacyFoldResults(LogicalResult status,
+                                       ArrayRef<OpFoldResult> results);
+} // namespace detail
+
 } // namespace mlir
 
 #endif // MLIR_IR_OPFOLDRESULT_H
diff --git a/mlir/lib/IR/OpFoldResult.cpp b/mlir/lib/IR/OpFoldResult.cpp
index 3a448d4b456ed..185d771e773b9 100644
--- a/mlir/lib/IR/OpFoldResult.cpp
+++ b/mlir/lib/IR/OpFoldResult.cpp
@@ -7,6 +7,8 @@
 //===----------------------------------------------------------------------===//
 
 #include "mlir/IR/OpFoldResult.h"
+#include "mlir/IR/Diagnostics.h"
+#include "mlir/IR/Operation.h"
 #include "llvm/Support/raw_ostream.h"
 
 using namespace mlir;
@@ -24,3 +26,126 @@ raw_ostream &mlir::operator<<(raw_ostream &os, OpFoldResult ofr) {
     llvm::dyn_cast_if_present<Attribute>(ofr).print(os);
   return os;
 }
+
+/// Return true if every replacement is null.
+static bool replacesNone(ArrayRef<OpFoldResult> replacements) {
+  return llvm::none_of(replacements, [](OpFoldResult replacement) {
+    return static_cast<bool>(replacement);
+  });
+}
+
+//===----------------------------------------------------------------------===//
+// OpFoldResults
+//===----------------------------------------------------------------------===//
+
+OpFoldResults::OpFoldResults(std::nullptr_t) {}
+
+OpFoldResults::OpFoldResults(LogicalResult status)
+    : inPlace(status.succeeded()) {}
+
+OpFoldResults::OpFoldResults(OpFoldResult replacement) {
+  if (replacement)
+    replacements.push_back(replacement);
+}
+
+OpFoldResults::OpFoldResults(Value replacement)
+    : OpFoldResults(OpFoldResult(replacement)) {}
+
+OpFoldResults::OpFoldResults(Attribute replacement)
+    : OpFoldResults(OpFoldResult(replacement)) {}
+
+OpFoldResults::OpFoldResults(std::initializer_list<OpFoldResult> replacements)
+    : replacements(replacements) {}
+
+OpFoldResults::OpFoldResults(Operation *op) {
+  assert(op && "expected a non-null operation");
+  replacements.resize(op->getNumResults());
+}
+
+void OpFoldResults::replace(Value result, OpFoldResult replacement) {
+  auto opResult = dyn_cast_if_present<OpResult>(result);
+  assert(opResult && "expected an OpResult");
+  if (dyn_cast_if_present<Value>(replacement) == result)
+    replacement = OpFoldResult();
+  replace(opResult.getResultNumber(), replacement);
+}
+
+void OpFoldResults::replace(unsigned resultIndex, OpFoldResult replacement) {
+  assert(resultIndex < replacements.size() && "result index out of range");
+  replacements[resultIndex] = replacement;
+}
+
+void OpFoldResults::setModifiedInPlace(bool modified) { inPlace = modified; }
+
+bool OpFoldResults::succeeded() const { return inPlace || replacesAny(); }
+
+bool OpFoldResults::failed() const { return !succeeded(); }
+
+bool OpFoldResults::modifiedInPlace() const { return inPlace; }
+
+bool OpFoldResults::replacesAny() const { return !replacesNone(replacements); }
+
+bool OpFoldResults::replacesAll() const {
+  return !replacements.empty() &&
+         llvm::all_of(replacements, [](OpFoldResult replacement) {
+           return static_cast<bool>(replacement);
+         });
+}
+
+unsigned OpFoldResults::size() const { return replacements.size(); }
+
+OpFoldResult OpFoldResults::operator[](unsigned i) const {
+  if (replacements.empty())
+    return OpFoldResult();
+  assert(i < replacements.size() && "result index out of range");
+  return replacements[i];
+}
+
+ArrayRef<OpFoldResult> OpFoldResults::getReplacements() const {
+  return replacements;
+}
+
+void OpFoldResults::normalize(Operation *op) {
+  assert(op && "expected a non-null operation");
+  if (!replacements.empty()) {
+    assert(replacements.size() == op->getNumResults() &&
+           "expected one replacement per operation result");
+    for (auto [replacement, result] :
+         llvm::zip_equal(replacements, op->getResults()))
+      if (dyn_cast_if_present<Value>(replacement) == result)
+        replacement = OpFoldResult();
+    if (replacesNone(replacements))
+      replacements.clear();
+  }
+
+#ifndef NDEBUG
+  for (auto [index, replacement] : llvm::enumerate(replacements)) {
+    auto value = dyn_cast_if_present<Value>(replacement);
+    if (!value)
+      continue;
+    Type expectedType = op->getResult(index).getType();
+    if (value.getType() != expectedType) {
+      op->emitOpError() << "folder produced a value of incorrect type: "
+                        << value.getType() << ", expected: " << expectedType;
+      assert(false && "incorrect fold result type");
+    }
+  }
+#endif // NDEBUG
+}
+
+//===----------------------------------------------------------------------===//
+// Legacy fold results
+//===----------------------------------------------------------------------===//
+
+OpFoldResults detail::convertLegacyFoldResults(LogicalResult status,
+                                               ArrayRef<OpFoldResult> results) {
+  if (failed(status))
+    return failure();
+  if (results.empty())
+    return success();
+  assert(llvm::all_of(
+             results,
+             [](OpFoldResult result) { return static_cast<bool>(result); }) &&
+         "legacy fold returned a null result");
+  return results;
+}
diff --git a/mlir/unittests/IR/CMakeLists.txt b/mlir/unittests/IR/CMakeLists.txt
index ce9b230222f62..e6858585aaf6c 100644
--- a/mlir/unittests/IR/CMakeLists.txt
+++ b/mlir/unittests/IR/CMakeLists.txt
@@ -14,6 +14,7 @@ add_mlir_unittest(MLIRIRTests
   LocationTest.cpp
   MLIRContextResetTest.cpp
   MemrefLayoutTest.cpp
+  OpFoldResultsTest.cpp
   OpImplementationTest.cpp
   OperationSupportTest.cpp
   PatternMatchTest.cpp
diff --git a/mlir/unittests/IR/OpFoldResultsTest.cpp b/mlir/unittests/IR/OpFoldResultsTest.cpp
new file mode 100644
index 0000000000000..8cb90ed967c73
--- /dev/null
+++ b/mlir/unittests/IR/OpFoldResultsTest.cpp
@@ -0,0 +1,452 @@
+//===- OpFoldResultsTest.cpp - OpFoldResults unit tests -------------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+
+#include "mlir/IR/Builders.h"
+#include "mlir/IR/BuiltinAttributes.h"
+#include "mlir/IR/BuiltinTypes.h"
+#include "mlir/IR/MLIRContext.h"
+#include "mlir/IR/OpFoldResult.h"
+#include "mlir/IR/Operation.h"
+#include "gtest/gtest.h"
+
+using namespace mlir;
+
+namespace {
+class OpFoldResultsTest : public ::testing::Test {
+protected:
+  OpFoldResultsTest() : builder(&context) {
+    context.allowUnregisteredDialects();
+    i32 = builder.getI32Type();
+    f32 = builder.getF32Type();
+  }
+
+  ~OpFoldResultsTest() override {
+    // Destroy users before the ops that define their operands.
+    for (Operation *op : llvm::reverse(ops))
+      op->destroy();
+  }
+
+  /// Create an op with the given result types and operands. The fixture
+  /// destroys it.
+  Operation *createOp(TypeRange resultTypes, StringRef name = "foo.bar",
+                      ValueRange operands = {}) {
+    OperationState state(UnknownLoc::get(&context), name);
+    state.addTypes(resultTypes);
+    state.addOperands(operands);
+    Operation *op = Operation::create(state);
+    ops.push_back(op);
+    return op;
+  }
+
+  MLIRContext context;
+  Builder builder;
+  Type i32;
+  Type f32;
+  SmallVector<Operation *> ops;
+};
+} // namespace
+
+static void expectOneReplacement(const OpFoldResults &result,
+                                 OpFoldResult expected) {
+  EXPECT_TRUE(result.succeeded());
+  EXPECT_FALSE(result.modifiedInPlace());
+  EXPECT_TRUE(result.replacesAny());
+  EXPECT_TRUE(result.replacesAll());
+  ASSERT_EQ(result.size(), 1u);
+  EXPECT_EQ(result[0], expected);
+}
+
+// The helpers below return the way a folder returns, so they use copy
+// initialization.
+static OpFoldResults foldToEmptyBraces() { return {}; }
+static OpFoldResults foldToTypedValue(TypedValue<IntegerType> value) {
+  return value;
+}
+static OpFoldResults foldToIntegerAttr(IntegerAttr attr) { return attr; }
+static OpFoldResults foldToArrayAttr(ArrayAttr attr) { return attr; }
+static OpFoldResults foldToList(OpFoldResult lhs, OpFoldResult rhs) {
+  return {lhs, rhs};
+}
+
+TEST_F(OpFoldResultsTest, DefaultIsFailure) {
+  OpFoldResults result;
+  EXPECT_TRUE(result.failed());
+  EXPECT_FALSE(result.succeeded());
+  EXPECT_FALSE(result.modifiedInPlace());
+  EXPECT_FALSE(result.replacesAny());
+  EXPECT_FALSE(result.replacesAll());
+  EXPECT_EQ(result.size(), 0u);
+  EXPECT_TRUE(result.getReplacements().empty());
+}
+
+TEST_F(OpFoldResultsTest, EmptyBracesAreFailure) {
+  OpFoldResults result = foldToEmptyBraces();
+  EXPECT_TRUE(result.failed());
+  EXPECT_FALSE(result.modifiedInPlace());
+  EXPECT_EQ(result.size(), 0u);
+}
+
+TEST_F(OpFoldResultsTest, NullptrIsFailure) {
+  OpFoldResults result = nullptr;
+  EXPECT_TRUE(result.failed());
+  EXPECT_FALSE(result.modifiedInPlace());
+  EXPECT_EQ(result.size(), 0u);
+}
+
+TEST_F(OpFoldResultsTest, LogicalResult) {
+  OpFoldResults inPlace = success();
+  EXPECT_TRUE(inPlace.succeeded());
+  EXPECT_TRUE(inPlace.modifiedInPlace());
+  EXPECT_FALSE(inPlace.replacesAny());
+  EXPECT_FALSE(inPlace.replacesAll());
+  EXPECT_EQ(inPlace.size(), 0u);
+
+  OpFoldResults failedResult = failure();
+  EXPECT_TRUE(failedResult.failed());
+  EXPECT_FALSE(failedResult.modifiedInPlace());
+  EXPECT_FALSE(failedResult.replacesAny());
+  EXPECT_EQ(failedResult.size(), 0u);
+}
+
+TEST_F(OpFoldResultsTest, SingleReplacement) {
+  Operation *producer = createOp({i32});
+  Operation *op = createOp({i32});
+  Value value = producer->getResult(0);
+  IntegerAttr intAttr = builder.getI32IntegerAttr(7);
+  Attribute attr = intAttr;
+
+  OpFoldResults fromOpFoldResult = OpFoldResult(attr);
+  expectOneReplacement(fromOpFoldResult, attr);
+
+  OpFoldResults fromValue = value;
+  expectOneReplacement(fromValue, value);
+
+  OpFoldResults fromAttr = attr;
+  expectOneReplacement(fromAttr, attr);
+
+  // TypedValue and IntegerAttr use the Value and Attribute constructors.
+  OpFoldResults fromTypedValue =
+      foldToTypedValue(cast<TypedValue<IntegerType>>(value));
+  expectOneReplacement(fromTypedValue, value);
+
+  OpFoldResults fromIntegerAttr = foldToIntegerAttr(intAttr);
+  expectOneReplacement(fromIntegerAttr, attr);
+
+  // The one replacement matches the single result.
+  fromValue.normalize(op);
+  expectOneReplacement(fromValue, value);
+  EXPECT_EQ(fromValue.size(), op->getNumResults());
+}
+
+TEST_F(OpFoldResultsTest, ArrayAttrFillsOneReplacement) {
+  Operation *op = createOp({i32});
+  ArrayAttr arrayAttr = builder.getI32ArrayAttr({1, 2, 3});
+  OpFoldResults result = foldToArrayAttr(arrayAttr);
+  expectOneReplacement(result, arrayAttr);
+  result.normalize(op);
+  expectOneReplacement(result, arrayAttr);
+}
+
+TEST_F(OpFoldResultsTest, InitializerList) {
+  Operation *producer = createOp({i32});
+  Operation *op = createOp({i32, f32});
+  Value value = producer->getResult(0);
+  Attribute attr = builder.getF32FloatAttr(1.0);
+
+  OpFoldResults result = {value, attr};
+  EXPECT_TRUE(result.succeeded());
+  ASSERT_EQ(result.size(), 2u);
+  EXPECT_EQ(result[0], OpFoldResult(value));
+  EXPECT_EQ(result[1], OpFoldResult(attr));
+  result.normalize(op);
+  EXPECT_TRUE(result.succeeded());
+  EXPECT_FALSE(result.modifiedInPlace());
+  EXPECT_TRUE(result.replacesAll());
+
+  OpFoldResults partial = foldToList(OpFoldResult(), attr);
+  partial.normalize(op);
+  EXPECT_TRUE(partial.succeeded());
+  EXPECT_TRUE(partial.replacesAny());
+  EXPECT_FALSE(partial.replacesAll());
+  ASSERT_EQ(partial.size(), 2u);
+  EXPECT_FALSE(partial[0]);
+  EXPECT_EQ(partial[1], OpFoldResult(attr));
+}
+
+TEST_F(OpFoldResultsTest, Range) {
+  Operation *producer = createOp({i32, f32});
+  Operation *op = createOp({i32, f32});
+  Attribute attr = builder.getI32IntegerAttr(3);
+
+  SmallVector<OpFoldResult> vector = {attr, OpFoldResult()};
+  OpFoldResults fromVector = vector;
+  EXPECT_TRUE(fromVector.succeeded());
+  ASSERT_EQ(fromVector.size(), 2u);
+  EXPECT_EQ(fromVector[0], OpFoldResult(attr));
+  EXPECT_FALSE(fromVector[1]);
+  fromVector.normalize(op);
+  EXPECT_TRUE(fromVector.replacesAny());
+  EXPECT_FALSE(fromVector.replacesAll());
+
+  ValueRange range = producer->getResults();
+  OpFoldResults fromRange = range;
+  EXPECT_TRUE(fromRange.succeeded());
+  ASSERT_EQ(fromRange.size(), 2u);
+  EXPECT_EQ(fromRange[0], OpFoldResult(producer->getResult(0)));
+  EXPECT_EQ(fromRange[1], OpFoldResult(producer->getResult(1)));
+  fromRange.normalize(op);
+  EXPECT_TRUE(fromRange.replacesAll());
+
+  OpFoldResults fromEmptyRange = ValueRange();
+  EXPECT_TRUE(fromEmptyRange.failed());
+  EXPECT_EQ(fromEmptyRange.size(), 0u);
+
+  OpFoldResults fromEmptyVector = SmallVector<OpFoldResult>();
+  EXPECT_TRUE(fromEmptyVector.failed());
+
+  SmallVector<OpFoldResult> nulls(2);
+  OpFoldResults fromNulls = nulls;
+  EXPECT_TRUE(fromNulls.failed());
+}
+
+TEST_F(OpFoldResultsTest, IncrementalForm) {
+  Operation *producer = createOp({i32});
+  Operation *op = createOp({i32, i32});
+  Value value = producer->getResult(0);
+  Attribute attr = builder.getI32IntegerAttr(5);
+
+  OpFoldResults result(op);
+  EXPECT_TRUE(result.failed());
+  EXPECT_FALSE(result.modifiedInPlace());
+  EXPECT_FALSE(result.replacesAny());
+  EXPECT_EQ(result.size(), 2u);
+
+  result.replace(op->getResult(1), attr);
+  EXPECT_TRUE(result.succeeded());
+  EXPECT_TRUE(result.replacesAny());
+  EXPECT_FALSE(result.replacesAll());
+  EXPECT_FALSE(result[0]);
+  EXPECT_EQ(result[1], OpFoldResult(attr));
+
+  // A literal 0 picks the unsigned overload.
+  result.replace(0, value);
+  EXPECT_TRUE(result.replacesAll());
+  EXPECT_EQ(result[0], OpFoldResult(value));
+
+  // The last write wins.
+  result.replace(0u, attr);
+  EXPECT_EQ(result[0], OpFoldResult(attr));
+
+  // A null replacement, or the result itself, keeps the result.
+  result.replace(op->getResult(0), OpFoldResult());
+  EXPECT_FALSE(result[0]);
+  result.replace(op->getResult(1), op->getResult(1));
+  EXPECT_FALSE(result[1]);
+  EXPECT_TRUE(result.failed());
+  EXPECT_FALSE(result.replacesAny());
+
+  result.setModifiedInPlace();
+  EXPECT_TRUE(result.succeeded());
+  EXPECT_TRUE(result.modifiedInPlace());
+  result.setModifiedInPlace(false);
+  EXPECT_TRUE(result.failed());
+  EXPECT_FALSE(result.modifiedInPlace());
+
+  // normalize() maps the op's own result, set by index, to "keep".
+  OpFoldResults ownResult(op);
+  ownResult.replace(1u, op->getResult(1));
+  ownResult.normalize(op);
+  EXPECT_TRUE(ownResult.failed());
+  EXPECT_EQ(ownResult.size(), 0u);
+
+  OpFoldResults partial(op);
+  partial.replace(1u, value);
+  partial.normalize(op);
+  EXPECT_TRUE(partial.succeeded());
+  ASSERT_EQ(partial.size(), 2u);
+  EXPECT_FALSE(partial[0]);
+  EXPECT_EQ(partial[1], OpFoldResult(value));
+  ArrayRef<OpFoldResult> replacements = partial.getReplacements();
+  ASSERT_EQ(replacements.size(), 2u);
+  EXPECT_FALSE(replacements[0]);
+  EXPECT_EQ(replacements[1], OpFoldResult(value));
+}
+
+TEST_F(OpFoldResultsTest, NullSafety) {
+  Operation *op = createOp({i32, i32});
+
+  EXPECT_TRUE(OpFoldResults(Value()).failed());
+  EXPECT_TRUE(OpFoldResults(Attribute()).failed());
+  EXPECT_TRUE(OpFoldResults(OpFoldResult()).failed());
+
+  OpFoldResults nullList = {OpFoldResult(), OpFoldResult()};
+  EXPECT_TRUE(nullList.failed());
+  nullList.normalize(op);
+  EXPECT_TRUE(nullList.failed());
+  EXPECT_EQ(nullList.size(), 0u);
+  EXPECT_FALSE(nullList[0]);
+  EXPECT_FALSE(nullList[1]);
+
+  Attribute attr = builder.getI32IntegerAttr(1);
+  OpFoldResults mixed = {OpFoldResult(), attr};
+  mixed.normalize(op);
+  EXPECT_TRUE(mixed.succeeded());
+  EXPECT_FALSE(mixed[0]);
+  EXPECT_EQ(mixed[1], OpFoldResult(attr));
+
+  OpFoldResults incremental(op);
+  incremental.replace(0u, OpFoldResult());
+  incremental.normalize(op);
+  EXPECT_TRUE(incremental.failed());
+  EXPECT_EQ(incremental.size(), 0u);
+}
+
+TEST_F(OpFoldResultsTest, NormalizeMapsOwnResultsToKeep) {
+  Operation *op = createOp({i32, i32});
+  Attribute attr = builder.getI32IntegerAttr(4);
+
+  OpFoldResults result = {op->getResult(0), attr};
+  result.normalize(op);
+  EXPECT_TRUE(result.succeeded());
+  EXPECT_TRUE(result.replacesAny());
+  EXPECT_FALSE(result.replacesAll());
+  ASSERT_EQ(result.size(), 2u);
+  EXPECT_FALSE(result[0]);
+  EXPECT_EQ(result[1], OpFoldResult(attr));
+
+  // A replacement may name another result of the op if that result is kept.
+  OpFoldResults forward = {op->getResult(1), OpFoldResult()};
+  forward.normalize(op);
+  EXPECT_TRUE(forward.succeeded());
+  ASSERT_EQ(forward.size(), 2u);
+  EXPECT_EQ(forward[0], OpFoldResult(op->getResult(1)));
+  EXPECT_FALSE(forward[1]);
+}
+
+TEST_F(OpFoldResultsTest, NormalizeCollapsesToFailure) {
+  Operation *op = createOp({i32, i32});
+  Operation *oneResultOp = createOp({i32});
+
+  OpFoldResults ownResults = {op->getResult(0), op->getResult(1)};
+  ownResults.normalize(op);
+  EXPECT_TRUE(ownResults.failed());
+  EXPECT_FALSE(ownResults.replacesAny());
+  EXPECT_EQ(ownResults.size(), 0u);
+  EXPECT_TRUE(ownResults.getReplacements().empty());
+
+  OpFoldResults ownRange = op->getResults();
+  ownRange.normalize(op);
+  EXPECT_TRUE(ownRange.failed());
+
+  OpFoldResults ownResult = oneResultOp->getResult(0);
+  ownResult.normalize(oneResultOp);
+  EXPECT_TRUE(ownResult.failed());
+
+  OpFoldResults nothingReplaced(op);
+  nothingReplaced.normalize(op);
+  EXPECT_TRUE(nothingReplaced.failed());
+  EXPECT_EQ(nothingReplaced.size(), 0u);
+
+  // An in-place change keeps the result successful.
+  OpFoldResults inPlace = {op->getResult(0), op->getResult(1)};
+  inPlace.setModifiedInPlace();
+  inPlace.normalize(op);
+  EXPECT_TRUE(inPlace.succeeded());
+  EXPECT_TRUE(inPlace.modifiedInPlace());
+  EXPECT_FALSE(inPlace.replacesAny());
+  EXPECT_FALSE(inPlace.replacesAll());
+  EXPECT_EQ(inPlace.size(), 0u);
+}
+
+TEST_F(OpFoldResultsTest, NormalizeIsIdempotent) {
+  Operation *op = createOp({i32, i32});
+  Attribute attr = builder.getI32IntegerAttr(6);
+
+  OpFoldResults result = {op->getResult(0), attr};
+  result.setModifiedInPlace();
+  result.normalize(op);
+  SmallVector<OpFoldResult> replacements(result.getReplacements());
+  result.normalize(op);
+  EXPECT_TRUE(result.succeeded());
+  EXPECT_TRUE(result.modifiedInPlace());
+  EXPECT_TRUE(result.replacesAny());
+  EXPECT_FALSE(result.replacesAll());
+  ASSERT_EQ(result.size(), replacements.size());
+  for (unsigned i = 0, e = replacements.size(); i < e; ++i)
+    EXPECT_EQ(result[i], replacements[i]);
+
+  OpFoldResults failedResult = failure();
+  failedResult.normalize(op);
+  failedResult.normalize(op);
+  EXPECT_TRUE(failedResult.failed());
+  EXPECT_EQ(failedResult.size(), 0u);
+
+  OpFoldResults inPlace = success();
+  inPlace.normalize(op);
+  inPlace.normalize(op);
+  EXPECT_TRUE(inPlace.succeeded());
+  EXPECT_TRUE(inPlace.modifiedInPlace());
+  EXPECT_EQ(inPlace.size(), 0u);
+}
+
+TEST_F(OpFoldResultsTest, ZeroResultInPlaceDoesNotReplaceAll) {
+  Operation *op = createOp({});
+
+  OpFoldResults inPlace = success();
+  inPlace.normalize(op);
+  EXPECT_TRUE(inPlace.succeeded());
+  EXPECT_TRUE(inPlace.modifiedInPlace());
+  EXPECT_FALSE(inPlace.replacesAny());
+  EXPECT_FALSE(inPlace.replacesAll());
+  EXPECT_EQ(inPlace.size(), 0u);
+
+  OpFoldResults failedResult = failure();
+  failedResult.normalize(op);
+  EXPECT_TRUE(failedResult.failed());
+  EXPECT_FALSE(failedResult.replacesAll());
+}
+
+TEST_F(OpFoldResultsTest, FreeHelpersMatchMembers) {
+  Operation *op = createOp({i32, i32});
+  OpFoldResults partial(op);
+  partial.replace(1u, builder.getI32IntegerAttr(1));
+  OpFoldResults inPlace = success();
+  OpFoldResults failedResult = failure();
+  for (const OpFoldResults *result : {&partial, &inPlace, &failedResult}) {
+    EXPECT_EQ(succeeded(*result), result->succeeded());
+    EXPECT_EQ(failed(*result), result->failed());
+  }
+  EXPECT_TRUE(succeeded(partial));
+  EXPECT_TRUE(succeeded(inPlace));
+  EXPECT_TRUE(failed(failedResult));
+}
+
+#ifdef GTEST_HAS_DEATH_TEST
+#ifndef NDEBUG
+namespace {
+class OpFoldResultsDeathTest : public OpFoldResultsTest {};
+} // namespace
+
+TEST_F(OpFoldResultsDeathTest, ValueReplacementOfIncorrectType) {
+  Operation *producer = createOp({f32});
+  Operation *op = createOp({i32, i32});
+  OpFoldResults result(op);
+  result.replace(0u, producer->getResult(0));
+  EXPECT_DEATH(result.normalize(op), "incorrect fold result type");
+}
+
+TEST_F(OpFoldResultsDeathTest, ReplacementCountMismatch) {
+  Operation *op = createOp({i32, i32});
+  Attribute attr = builder.getI32IntegerAttr(1);
+  OpFoldResults result = {attr, attr, attr};
+  EXPECT_DEATH(result.normalize(op),
+               "expected one replacement per operation result");
+}
+
+#endif // NDEBUG
+#endif // GTEST_HAS_DEATH_TEST



More information about the llvm-branch-commits mailing list