如何让C++20栈式协程的co_await支持任意T返回类型?
解决C++20栈式协程跨
task<T>类型co_await的问题 问题背景
基于Stack Overflow帖子实现C++20栈式协程系统时,原代码仅支持嵌套协程为相同task<T>类型,现需扩展让co_await能接受任意T类型的task,即等待方与被等待方的task<T>的T可以不同。当前代码在await_suspend函数处编译失败,原因是std::coroutine_handle<>未暴露promise()成员;若改用原代码的std::coroutine_handle<promise_type>,又会回到仅支持同类型T的限制。
当前代码
template<class T> class task { public: class promise_type { protected: T value; std::coroutine_handle<> innerHandler{}; std::coroutine_handle<> outerHandler{}; friend class task; public: task get_return_object() { return task(std::coroutine_handle<promise_type>::from_promise(*this)); } auto initial_suspend() { return std::suspend_never{}; } auto final_suspend() noexcept { return std::suspend_always{}; } auto return_value(T v) { this->value = std::move(v); return std::suspend_always{}; } void unhandled_exception() { abort("Unhandled exception in Coroutine"); } }; explicit task(std::coroutine_handle<promise_type> handle) : handle(handle) {} task(const task&) = delete; task(task&& c) noexcept : handle(std::exchange(c.handle, nullptr)) {} task& operator=(const task&) = delete; task& operator=(task&& c) noexcept { this->handle = std::exchange(c.handle, nullptr); return *this; } ~task() { if (this->handle) { this->handle.destroy(); } } constexpr bool await_ready() const noexcept { return !this->handle || this->handle.done(); } bool await_suspend(std::coroutine_handle<> h) { h.promise().innerHandler = this->handle; this->handle.promise().outerHandler = h; return true; } constexpr T await_resume() const noexcept { return std::move(this->handle.promise().value); } bool next() { auto cur = this->handle; while (cur) { if (!cur.promise().innerHandler) { while (!cur.done()) { cur.resume(); if (!cur.done()) { return true; } if (cur.promise().outerHandler) { cur = cur.promise().outerHandler; cur.promise().innerHandler = nullptr; } else { return false; } } break; } cur = cur.promise().innerHandler; } return !cur.done(); } private: std::coroutine_handle<promise_type> handle; };
问题分析
编译失败的核心原因:std::coroutine_handle<>(无类型参数版本)不提供promise()成员函数,只有带具体promise类型的std::coroutine_handle<P>才能访问该成员。原代码将外部协程句柄限制为同类型promise_type,导致无法支持不同T的task嵌套。
解决方案
通过给所有task<T>的promise_type定义公共基类,统一管理协程句柄的嵌套关系,摆脱具体task<T>类型的限制。
步骤1:定义公共promise基类
创建非模板基类,包含所有promise_type共享的成员和访问接口:
class task_promise_base { protected: std::coroutine_handle<> innerHandler{}; std::coroutine_handle<> outerHandler{}; public: virtual ~task_promise_base() = default; // 公共访问接口 void set_inner_handler(std::coroutine_handle<> h) { innerHandler = h; } std::coroutine_handle<> get_inner_handler() const { return innerHandler; } void set_outer_handler(std::coroutine_handle<> h) { outerHandler = h; } std::coroutine_handle<> get_outer_handler() const { return outerHandler; } // 从无类型协程句柄转换为基类指针 static task_promise_base* from_coroutine_handle(std::coroutine_handle<> h) { return static_cast<task_promise_base*>(h.address()); } };
步骤2:修改task<T>::promise_type继承基类
让每个task<T>的promise_type公有继承自task_promise_base,继承公共成员:
template<class T> class task { public: class promise_type : public task_promise_base { protected: T value; friend class task; public: // 原有成员函数保持不变 task get_return_object() { return task(std::coroutine_handle<promise_type>::from_promise(*this)); } auto initial_suspend() { return std::suspend_never{}; } auto final_suspend() noexcept { return std::suspend_always{}; } auto return_value(T v) { this->value = std::move(v); return std::suspend_always{}; } void unhandled_exception() { abort("Unhandled exception in Coroutine"); } }; // 其他成员保持不变 };
步骤3:修复await_suspend函数
通过公共基类访问外部协程的成员,替代原有的h.promise()调用:
bool await_suspend(std::coroutine_handle<> h) { // 获取外部协程的promise基类指针 auto outer_promise = task_promise_base::from_coroutine_handle(h); outer_promise->set_inner_handler(this->handle); // 当前task的promise继承自基类,直接调用接口设置外部句柄 this->handle.promise().set_outer_handler(h); return true; }
步骤4:修复next函数中的promise访问
将原cur.promise().xxx的调用替换为基类接口:
bool next() { auto cur = this->handle; while (cur) { auto cur_promise = task_promise_base::from_coroutine_handle(cur); if (!cur_promise->get_inner_handler()) { while (!cur.done()) { cur.resume(); if (!cur.done()) { return true; } auto outer_h = cur_promise->get_outer_handler(); if (outer_h) { cur = outer_h; auto outer_promise = task_promise_base::from_coroutine_handle(cur); outer_promise->set_inner_handler(nullptr); } else { return false; } } break; } cur = cur_promise->get_inner_handler(); } return !cur.done(); }
说明
- 公共基类统一管理协程句柄的嵌套关系,实现了任意
task<T>之间的co_await支持。 - 无需使用
await_transform,逻辑直观,适合现有代码的改造。 std::coroutine_handle<>::address()返回promise对象的起始地址,由于promise_type公有继承自基类,静态转换符合C++对象布局规则,是安全的。
内容的提问来源于stack exchange,提问作者debevv
相关产品推荐
相关产品推荐

