https://github.com/amanmaurya92 updated https://github.com/llvm/llvm-project/pull/227527
>From 412ca3f1aa48d9a08ebab1680738a7ed5d656b7e Mon Sep 17 00:00:00 2001 From: amanmaurya92 <[email protected]> Date: Wed, 30 Sep 2026 06:05:37 +0530 Subject: [PATCH] [ClangIR] Support bool-returning await_suspend Support bool-returning await_suspend in cir.await and CIRGen. Closes #227404 Assisted-by: Antigravity --- clang/include/clang/CIR/Dialect/IR/CIROps.td | 2 + clang/lib/CIR/CodeGen/CIRGenCoroutine.cpp | 23 ++++++---- clang/lib/CIR/Dialect/IR/CIRDialect.cpp | 39 ++++++++++++---- .../coro-await-suspend-bool.cpp | 44 +++++++++++++++++++ clang/test/CIR/IR/await.cir | 37 ++++++++++++++++ clang/test/CIR/IR/invalid-await.cir | 33 +++++++++++++- 6 files changed, 160 insertions(+), 18 deletions(-) create mode 100644 clang/test/CIR/CodeGenCoroutines/coro-await-suspend-bool.cpp diff --git a/clang/include/clang/CIR/Dialect/IR/CIROps.td b/clang/include/clang/CIR/Dialect/IR/CIROps.td index 89e93a716424d..b1db879126811 100644 --- a/clang/include/clang/CIR/Dialect/IR/CIROps.td +++ b/clang/include/clang/CIR/Dialect/IR/CIROps.td @@ -1129,6 +1129,8 @@ def CIR_ConditionOp : CIR_Op<"condition", [ if true, or exits it if false. - When in the `ready` region of a `cir.await`, it branches to the `resume` region when true, and to the `suspend` region when false. + - When in the `suspend` region of a `cir.await`, it suspends (exits `cir.await`) + when true, or branches to the `resume` region when false (veto suspension). Example: diff --git a/clang/lib/CIR/CodeGen/CIRGenCoroutine.cpp b/clang/lib/CIR/CodeGen/CIRGenCoroutine.cpp index 71bf3424c087a..560f1bca98fc7 100644 --- a/clang/lib/CIR/CodeGen/CIRGenCoroutine.cpp +++ b/clang/lib/CIR/CodeGen/CIRGenCoroutine.cpp @@ -662,16 +662,21 @@ emitSuspendExpression(CIRGenFunction &cgf, CGCoroData &coro, // and coro.suspend here, that should be done as part of lowering this // to LLVM dialect (or some other MLIR dialect) - // A invalid suspendRet indicates "void returning await_suspend" - mlir::Value suspendRet = cgf.emitScalarExpr(s.getSuspendExpr()); - - // Veto suspension if requested by bool returning await_suspend. - if (suspendRet) { - cgf.cgm.errorNYI("Veto await_suspend"); + if (s.getSuspendReturnType() == + CoroutineSuspendExpr::SuspendReturnType::SuspendBool) { + mlir::Value suspendRet = cgf.evaluateExprAsBool(s.getSuspendExpr()); + // Veto suspension if requested by bool returning await_suspend. + builder.createCondition(suspendRet); + } else if (s.getSuspendReturnType() == + CoroutineSuspendExpr::SuspendReturnType::SuspendVoid) { + cgf.emitScalarExpr(s.getSuspendExpr()); + // Signals the parent that execution flows to next region. + cir::CoroSuspendPoint::create(builder, loc); + } else { + cgf.cgm.errorNYI(s.getSourceRange(), + "await_suspend returning handle"); + cir::CoroSuspendPoint::create(builder, loc); } - - // Signals the parent that execution flows to next region. - cir::CoroSuspendPoint::create(builder, loc); }, /*resumeBuilder=*/ [&](mlir::OpBuilder &b, mlir::Location loc) { diff --git a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp index c4ed6697b869f..c14c38c00c37a 100644 --- a/clang/lib/CIR/Dialect/IR/CIRDialect.cpp +++ b/clang/lib/CIR/Dialect/IR/CIRDialect.cpp @@ -702,10 +702,21 @@ void cir::ConditionOp::getSuccessorRegions( return; } - // Parent is an await: condition may branch to resume or suspend regions. + // Parent is an await: condition in ready region branches to resume or + // suspend regions. Condition in suspend region branches to resume (veto) or + // exits to parent op (suspend). auto await = cast<AwaitOp>(getOperation()->getParentOp()); - regions.emplace_back(&await.getResume()); - regions.emplace_back(&await.getSuspend()); + mlir::Region *parentRegion = getOperation()->getBlock()->getParent(); + if (parentRegion == &await.getReady()) { + regions.emplace_back(&await.getResume()); + regions.emplace_back(&await.getSuspend()); + return; + } + if (parentRegion == &await.getSuspend()) { + regions.emplace_back(getOperation()); + regions.emplace_back(&await.getResume()); + return; + } } MutableOperandRange @@ -3575,17 +3586,29 @@ void cir::AwaitOp::getSuccessorRegions( return; } + // Branching from suspend: if terminated by cir.condition, it may branch to + // exit to parent op (suspend) or resume (veto). + if (&getSuspend() == parentRegion) { + if (isa<ConditionOp>(point.getTerminatorPredecessorOrNull())) { + regions.emplace_back(getOperation()); + regions.emplace_back(&getResume()); + return; + } + } + // Branching from suspend or resume: exit to the parent operation. regions.emplace_back(getOperation()); } LogicalResult cir::AwaitOp::verify() { - if (!isa<ConditionOp>(this->getReady().back().getTerminator())) + if (this->getReady().empty() || + !isa<ConditionOp>(this->getReady().back().getTerminator())) return emitOpError("ready region must end with cir.condition"); - if (this->getSuspend().empty()) - return emitOpError("suspend region must not be empty"); - if (!isa<CoroSuspendPoint>(this->getSuspend().back().getTerminator())) - return emitOpError("suspend region must end with cir.coro.suspend_point"); + if (this->getSuspend().empty() || + !isa<CoroSuspendPoint, ConditionOp>( + this->getSuspend().back().getTerminator())) + return emitOpError( + "suspend region must end with cir.coro.suspend_point or cir.condition"); return success(); } diff --git a/clang/test/CIR/CodeGenCoroutines/coro-await-suspend-bool.cpp b/clang/test/CIR/CodeGenCoroutines/coro-await-suspend-bool.cpp new file mode 100644 index 0000000000000..f88531206b0e9 --- /dev/null +++ b/clang/test/CIR/CodeGenCoroutines/coro-await-suspend-bool.cpp @@ -0,0 +1,44 @@ +// RUN: %clang_cc1 -std=c++20 -triple x86_64-unknown-linux-gnu -fclangir -Wno-coroutine-missing-unhandled-exception -emit-cir %s -o %t.cir +// RUN: FileCheck --input-file=%t.cir %s -check-prefix=CIR +// RUN: %clang_cc1 -std=c++20 -triple x86_64-unknown-linux-gnu -emit-llvm -disable-llvm-passes -Wno-coroutine-missing-unhandled-exception %s -o %t.ll +// RUN: FileCheck --input-file=%t.ll %s -check-prefix=OGCG + +#include "Inputs/coroutine.h" + +struct Task { + struct promise_type { + Task get_return_object() { return {}; } + std::suspend_never initial_suspend() noexcept { return {}; } + std::suspend_never final_suspend() noexcept { return {}; } + void return_void() {} + void unhandled_exception() {} + }; +}; + +struct BoolAwaiter { + bool await_ready() { return false; } + bool await_suspend(std::coroutine_handle<>) { return false; } + void await_resume() {} +}; + +// CIR-LABEL: cir.func coroutine {{.*}} @_Z15await_bool_vetov +// OGCG-LABEL: define dso_local void @_Z15await_bool_vetov +Task await_bool_veto() { + // CIR: cir.await(user, ready : { + // CIR: %[[READY:.*]] = cir.call @_ZN11BoolAwaiter11await_readyEv(%{{.*}}) : (!cir.ptr<!rec_BoolAwaiter>{{.*}}) -> (!cir.bool{{.*}}) + // CIR: cir.condition(%[[READY]]) + // CIR: }, suspend : { + // CIR: %[[SUSPEND_RET:.*]] = cir.call @_ZN11BoolAwaiter13await_suspendESt16coroutine_handleIvE(%{{.*}}) : (!cir.ptr<!rec_BoolAwaiter>{{.*}}) -> (!cir.bool{{.*}}) + // CIR: cir.condition(%[[SUSPEND_RET]]) + // CIR: }, resume : { + // CIR: cir.call @_ZN11BoolAwaiter12await_resumeEv(%{{.*}}) : (!cir.ptr<!rec_BoolAwaiter>{{.*}}) -> () + // CIR: cir.yield + // CIR: },) + + // OGCG: %[[READY_RES:.*]] = call noundef zeroext i1 @_ZN11BoolAwaiter11await_readyEv(ptr {{.*}}) + // OGCG: br i1 %[[READY_RES]], label %[[AWAIT_READY_DEST:.*]], label %[[AWAIT_SUSPEND:.*]] + // OGCG: [[AWAIT_SUSPEND]]: + // OGCG: %[[SUSP_RET:.*]] = call i1 @llvm.coro.await.suspend.bool(ptr {{.*}}, ptr {{.*}}, ptr {{.*}}) + // OGCG: br i1 %[[SUSP_RET]], label %{{.*}}, label %[[AWAIT_READY_DEST]] + co_await BoolAwaiter{}; +} diff --git a/clang/test/CIR/IR/await.cir b/clang/test/CIR/IR/await.cir index b5d3df6175b9c..e3f09df0d73f6 100644 --- a/clang/test/CIR/IR/await.cir +++ b/clang/test/CIR/IR/await.cir @@ -36,3 +36,40 @@ cir.func coroutine @checkPrintParse(%arg0 : !cir.bool) { // CHECK: }, resume : { // CHECK: cir.yield // CHECK: },) + +cir.func coroutine @checkPrintParseBoolSuspend(%arg0 : !cir.bool) { + cir.coroutine initialSuspend : { + cir.await(init, ready : { + cir.condition(%arg0) + }, suspend : { + cir.coro.suspend_point + }, resume : { + cir.yield + },) + cir.yield + }, body : { + cir.await(user, ready : { + cir.condition(%arg0) + }, suspend : { + cir.condition(%arg0) + }, resume : { + cir.yield + },) + cir.yield + }, finalSuspend : { + cir.yield + }, destroy : { + cir.yield + }, exit : { + cir.return + } + cir.trap +} + +// CHECK: cir.await(user, ready : { +// CHECK: cir.condition(%arg0) +// CHECK: }, suspend : { +// CHECK: cir.condition(%arg0) +// CHECK: }, resume : { +// CHECK: cir.yield +// CHECK: },) diff --git a/clang/test/CIR/IR/invalid-await.cir b/clang/test/CIR/IR/invalid-await.cir index 813f10b65b519..cbb85cd724873 100644 --- a/clang/test/CIR/IR/invalid-await.cir +++ b/clang/test/CIR/IR/invalid-await.cir @@ -31,7 +31,7 @@ cir.func coroutine @missing_condition() { cir.func coroutine @missing_suspend_point(%arg0 : !cir.bool) { cir.coroutine initialSuspend : { - cir.await(init, ready : { // expected-error {{suspend region must end with cir.coro.suspend_point}} + cir.await(init, ready : { // expected-error {{suspend region must end with cir.coro.suspend_point or cir.condition}} cir.condition(%arg0) }, suspend : { cir.yield @@ -50,3 +50,34 @@ cir.func coroutine @missing_suspend_point(%arg0 : !cir.bool) { } cir.trap } + +// ----- + +cir.func coroutine @invalid_suspend_terminator(%arg0 : !cir.bool) { + cir.coroutine initialSuspend : { + cir.await(init, ready : { + cir.condition(%arg0) + }, suspend : { + cir.coro.suspend_point + }, resume : { + cir.yield + },) + cir.yield + }, body : { + cir.await(user, ready : { // expected-error {{suspend region must end with cir.coro.suspend_point or cir.condition}} + cir.condition(%arg0) + }, suspend : { + cir.unreachable + }, resume : { + cir.yield + },) + cir.yield + }, finalSuspend : { + cir.yield + }, destroy : { + cir.yield + }, exit : { + cir.return + } + cir.trap +} _______________________________________________ cfe-commits mailing list [email protected] https://lists.llvm.org/cgi-bin/mailman/listinfo/cfe-commits
