From 9f2b929e0e13dab159a3a1fe6a9882b1a0a3243c Mon Sep 17 00:00:00 2001 From: Simon Hofmann Date: Mon, 31 Aug 2026 17:54:50 +0200 Subject: [PATCH] =?UTF-8?q?=F0=9F=90=9B=20Stop=20QCO=20wire=20traversal=20?= =?UTF-8?q?at=20unknown=20carriers=20(#2255)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Assisted-by: GPT-5.6 Sol via Codex --- mlir/lib/Dialect/QCO/Utils/WireIterator.cpp | 17 ++++++----- .../Dialect/QCO/Utils/test_wireiterator.cpp | 30 +++++++++++++++++++ 2 files changed, 39 insertions(+), 8 deletions(-) diff --git a/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp b/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp index b8e883b422..bf019e4aa4 100644 --- a/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp +++ b/mlir/lib/Dialect/QCO/Utils/WireIterator.cpp @@ -90,10 +90,7 @@ void WireIterator::forward() { .Case([&](IndexSwitchOp op) { qubit_ = op.getTiedResult(&(*qubit_.use_begin())); }) - .Default([&](Operation* op) { - llvm::reportFatalInternalError("unknown op in def-use chain: " + - op->getName().getStringRef()); - }); + .Default([&](Operation*) { isFinal_ = true; }); } } @@ -125,6 +122,7 @@ void WireIterator::backward() { } // Find the input from the output qubit SSA value. + bool reachedBoundary = false; TypeSwitch(op_) .Case( [&](UnitaryOpInterface op) { qubit_ = op.getInputForOutput(qubit_); }) @@ -163,10 +161,13 @@ void WireIterator::backward() { } llvm::reportFatalInternalError("expected result lookup"); }) - .Default([&](Operation* op) { - llvm::reportFatalInternalError("unknown op in def-use chain: " + - op->getName().getStringRef()); - }); + .Default([&](Operation*) { reachedBoundary = true; }); + + if (reachedBoundary) { + op_ = nullptr; + isFinal_ = false; + return; + } // Get the operation that produces the qubit value. // If the current qubit SSA value is a BlockArgument (no defining op), the diff --git a/mlir/unittests/Dialect/QCO/Utils/test_wireiterator.cpp b/mlir/unittests/Dialect/QCO/Utils/test_wireiterator.cpp index 2d9e0cbc03..7d894456d3 100644 --- a/mlir/unittests/Dialect/QCO/Utils/test_wireiterator.cpp +++ b/mlir/unittests/Dialect/QCO/Utils/test_wireiterator.cpp @@ -17,6 +17,7 @@ #include #include #include +#include #include #include #include @@ -264,6 +265,35 @@ TEST_P(WireIteratorTest, FunctionReturnTerminatesTraversal) { ASSERT_EQ(it.qubit(), output); } +TEST_P(WireIteratorTest, UnknownCarrierTerminatesTraversal) { + OpBuilder builder(context.get()); + const auto location = builder.getUnknownLoc(); + auto module = ModuleOp::create(location); + builder.setInsertionPointToStart(module.getBody()); + auto function = func::FuncOp::create(builder, location, "main", + builder.getFunctionType({}, {})); + Block* body = function.addEntryBlock(); + builder.setInsertionPointToStart(body); + + auto source = qco::AllocOp::create(builder, location).getResult(); + auto carrier = UnrealizedConversionCastOp::create( + builder, location, TypeRange{source.getType()}, ValueRange{source}); + auto carried = carrier.getResult(0); + qco::SinkOp::create(builder, location, carried); + func::ReturnOp::create(builder, location); + + qco::WireIterator forward(source); + ++forward; + EXPECT_EQ(forward.operation(), carrier.getOperation()); + ++forward; + EXPECT_EQ(forward, std::default_sentinel); + + qco::WireIterator backward(carried); + --backward; + EXPECT_EQ(backward.operation(), nullptr); + EXPECT_EQ(backward.qubit(), carried); +} + INSTANTIATE_TEST_SUITE_P(DynamicAndStatic, WireIteratorTest, ::testing::Bool(), [](const ::testing::TestParamInfo& info) { return info.param ? "Dynamic" : "Static";