16#include <trantor/utils/NonCopyable.h>
17#include <trantor/net/EventLoop.h>
18#include <trantor/utils/Logger.h>
22#include <condition_variable>
35auto getAwaiterImpl(T &&value)
noexcept(
36 noexcept(
static_cast<T &&
>(value).operator
co_await()))
37 ->
decltype(
static_cast<T &&
>(value).
operator co_await())
39 return static_cast<T &&
>(value).
operator co_await();
43auto getAwaiterImpl(T &&value)
noexcept(
44 noexcept(
operator co_await(
static_cast<T &&
>(value))))
45 ->
decltype(
operator co_await(
static_cast<T &&
>(value)))
47 return operator co_await(
static_cast<T &&
>(value));
51auto getAwaiter(T &&value)
noexcept(
52 noexcept(getAwaiterImpl(
static_cast<T &&
>(value))))
53 ->
decltype(getAwaiterImpl(
static_cast<T &&
>(value)))
55 return getAwaiterImpl(
static_cast<T &&
>(value));
59using void_to_false_t =
60 std::conditional_t<std::is_same_v<T, void>, std::false_type, T>;
67 using awaiter_t =
decltype(internal::getAwaiter(std::declval<T>()));
68 using type =
decltype(std::declval<awaiter_t>().await_resume());
72using await_result_t =
typename await_result<T>::type;
74template <
typename T,
typename = std::
void_t<>>
82 std::void_t<decltype(internal::getAwaiter(std::declval<T>()))>>
97 bool await_ready()
noexcept
102 template <
typename T>
103 auto await_suspend(std::coroutine_handle<T> handle)
noexcept
105 return handle.promise().continuation_;
108 void await_resume()
noexcept
121template <
typename Promise>
124 using handle_type = std::coroutine_handle<Promise>;
127 explicit task_awaiter(handle_type coro) : coro_(coro)
131 bool await_ready()
noexcept
133 return !coro_ || coro_.done();
136 auto await_suspend(std::coroutine_handle<> handle)
noexcept
138 coro_.promise().setContinuation(handle);
144 if constexpr (std::is_void_v<
decltype(coro_.promise().result())>)
146 coro_.promise().result();
151 return std::move(coro_.promise().result());
159template <
typename T =
void>
160struct [[nodiscard]] Task
163 using handle_type = std::coroutine_handle<promise_type>;
165 Task(handle_type h) : coro_(h)
169 Task(
const Task &) =
delete;
171 Task(Task &&other)
noexcept
174 other.coro_ =
nullptr;
183 Task &operator=(
const Task &) =
delete;
185 Task &operator=(Task &&other)
noexcept
187 if (std::addressof(other) ==
this)
193 other.coro_ =
nullptr;
199 Task<T> get_return_object()
201 return Task<T>{handle_type::from_promise(*
this)};
204 std::suspend_always initial_suspend()
209 void return_value(
const T &v)
214 void return_value(T &&v)
216 value = std::move(v);
219 auto final_suspend()
noexcept
224 void unhandled_exception()
226 exception_ = std::current_exception();
231 if (exception_ !=
nullptr)
232 std::rethrow_exception(exception_);
233 assert(value.has_value() ==
true);
234 return std::move(value.value());
239 if (exception_ !=
nullptr)
240 std::rethrow_exception(exception_);
241 assert(value.has_value() ==
true);
242 return value.value();
245 void setContinuation(std::coroutine_handle<> handle)
247 continuation_ = handle;
250 std::optional<T> value;
251 std::exception_ptr exception_;
252 std::coroutine_handle<> continuation_{std::noop_coroutine()};
255 auto operator co_await()
const noexcept
264struct [[nodiscard]] Task<void>
267 using handle_type = std::coroutine_handle<promise_type>;
269 Task(handle_type handle) : coro_(handle)
273 Task(
const Task &) =
delete;
275 Task(Task &&other)
noexcept
278 other.coro_ =
nullptr;
287 Task &operator=(
const Task &) =
delete;
289 Task &operator=(Task &&other)
noexcept
291 if (std::addressof(other) ==
this)
297 other.coro_ =
nullptr;
303 Task<> get_return_object()
305 return Task<>{handle_type::from_promise(*
this)};
308 std::suspend_always initial_suspend()
317 auto final_suspend()
noexcept
322 void unhandled_exception()
324 exception_ = std::current_exception();
329 if (exception_ !=
nullptr)
330 std::rethrow_exception(exception_);
333 void setContinuation(std::coroutine_handle<> handle)
335 continuation_ = handle;
338 std::exception_ptr exception_;
339 std::coroutine_handle<> continuation_{std::noop_coroutine()};
342 auto operator co_await()
const noexcept
357 using handle_type = std::coroutine_handle<promise_type>;
359 AsyncTask() =
default;
361 AsyncTask(handle_type h) : coro_(h)
365 AsyncTask(
const AsyncTask &) =
delete;
367 AsyncTask(AsyncTask &&other)
noexcept
370 other.coro_ =
nullptr;
373 AsyncTask &operator=(
const AsyncTask &) =
delete;
375 AsyncTask &operator=(AsyncTask &&other)
noexcept
377 if (std::addressof(other) ==
this)
381 other.coro_ =
nullptr;
387 AsyncTask get_return_object()
noexcept
389 return {std::coroutine_handle<promise_type>::from_promise(*
this)};
392 std::suspend_never initial_suspend()
const noexcept
397 void unhandled_exception()
399 LOG_FATAL <<
"Exception escaping AsyncTask.";
403 void return_void()
noexcept
407 std::suspend_never final_suspend()
const noexcept
419template <
typename T =
void>
422 bool await_ready()
noexcept
427 bool hasException()
const noexcept
429 return exception_ !=
nullptr;
432 const T &await_resume()
const noexcept(
false)
437 assert(result_.has_value() ==
true || exception_ !=
nullptr);
440 std::rethrow_exception(exception_);
441 return result_.value();
448 std::optional<T> result_;
449 std::exception_ptr exception_{
nullptr};
452 void setException(
const std::exception_ptr &e)
457 void setValue(
const T &v)
464 result_.emplace(std::move(v));
471 bool await_ready()
noexcept
476 void await_resume()
noexcept(
false)
479 std::rethrow_exception(exception_);
482 bool hasException()
const noexcept
484 return exception_ !=
nullptr;
488 std::exception_ptr exception_{
nullptr};
491 void setException(
const std::exception_ptr &e)
499template <
typename Await>
500auto sync_wait(Await &&await)
502 static_assert(is_awaitable_v<std::decay_t<Await>>);
503 using value_type =
typename await_result<Await>::type;
504 std::condition_variable cv;
506 std::atomic<bool> flag =
false;
507 std::exception_ptr exception_ptr;
508 std::unique_lock lk(mtx);
510 if constexpr (std::is_same_v<value_type, void>)
519 exception_ptr = std::current_exception();
521 std::unique_lock lk(mtx);
526 std::thread thr([&]() { task(); });
527 cv.wait(lk, [&]() {
return (
bool)flag; });
530 std::rethrow_exception(exception_ptr);
534 std::optional<value_type> value;
538 value =
co_await await;
542 exception_ptr = std::current_exception();
544 std::unique_lock lk(mtx);
549 std::thread thr([&]() { task(); });
550 cv.wait(lk, [&]() {
return (
bool)flag; });
551 assert(value.has_value() ==
true || exception_ptr);
555 std::rethrow_exception(exception_ptr);
557 return std::move(value.value());
562template <
typename Await>
563inline auto co_future(Await &&await)
noexcept
564 -> std::future<await_result_t<Await>>
566 using Result = await_result_t<Await>;
567 std::promise<Result> prom;
568 auto fut = prom.get_future();
569 [](std::promise<Result> prom, Await await) ->
AsyncTask {
572 if constexpr (std::is_void_v<Result>)
574 co_await std::move(await);
578 prom.set_value(
co_await std::move(await));
582 prom.set_exception(std::current_exception());
584 }(std::move(prom), std::move(await));
592 TimerAwaiter(trantor::EventLoop *loop,
593 const std::chrono::duration<double> &delay)
594 : loop_(loop), delay_(delay.count())
598 TimerAwaiter(trantor::EventLoop *loop,
double delay)
599 : loop_(loop), delay_(delay)
603 void await_suspend(std::coroutine_handle<> handle)
605 loop_->runAfter(delay_, [handle]() { handle.resume(); });
609 trantor::EventLoop *loop_;
615 LoopAwaiter(trantor::EventLoop *workLoop,
616 std::function<
void()> &&taskFunc,
617 trantor::EventLoop *resumeLoop =
nullptr)
618 : workLoop_(workLoop),
619 resumeLoop_(resumeLoop),
620 taskFunc_(std::move(taskFunc))
625 void await_suspend(std::coroutine_handle<> handle)
627 workLoop_->queueInLoop([handle,
this]() {
631 if (resumeLoop_ && resumeLoop_ != workLoop_)
632 resumeLoop_->queueInLoop([handle]() { handle.resume(); });
638 setException(std::current_exception());
639 if (resumeLoop_ && resumeLoop_ != workLoop_)
640 resumeLoop_->queueInLoop([handle]() { handle.resume(); });
648 trantor::EventLoop *workLoop_{
nullptr};
649 trantor::EventLoop *resumeLoop_{
nullptr};
650 std::function<void()> taskFunc_;
655 explicit SwitchThreadAwaiter(trantor::EventLoop *loop) : loop_(loop)
659 void await_suspend(std::coroutine_handle<> handle)
661 loop_->runInLoop([handle]() { handle.resume(); });
665 trantor::EventLoop *loop_;
670 EndAwaiter(trantor::EventLoop *loop) : loop_(loop)
675 void await_suspend(std::coroutine_handle<> handle)
677 loop_->runOnQuit([handle]() { handle.resume(); });
681 trantor::EventLoop *loop_{
nullptr};
687 trantor::EventLoop *loop,
688 const std::chrono::duration<double> &delay)
noexcept
691 return {loop, delay};
695 double delay)
noexcept
698 return {loop, delay};
702 trantor::EventLoop *workLoop,
703 std::function<
void()> taskFunc,
704 trantor::EventLoop *resumeLoop =
nullptr)
707 return {workLoop, std::move(taskFunc), resumeLoop};
711 trantor::EventLoop *loop)
noexcept
723template <
typename T,
typename = std::
void_t<>>
731 std::void_t<decltype(internal::getAwaiter(std::declval<T>()))>>
748template <
typename Coro>
751 using CoroValueType = std::decay_t<Coro>;
752 auto functor = [](CoroValueType coro) ->
AsyncTask {
755 using FrameType = std::decay_t<
decltype(frame)>;
756 static_assert(is_awaitable_v<FrameType>);
761 functor(std::forward<Coro>(coro));
768template <
typename Coro>
771 return [coro = std::forward<Coro>(coro)]()
mutable {
781 EventLoopAwaiter(std::function<T()> &&task, trantor::EventLoop *loop)
782 : task_(std::move(task)), loop_(loop)
786 void await_suspend(std::coroutine_handle<> handle)
788 loop_->queueInLoop([
this, handle]() {
791 if constexpr (!std::is_same_v<T, void>)
793 this->setValue(task_());
802 catch (
const std::exception &err)
804 LOG_ERROR << err.what();
805 this->setException(std::current_exception());
812 std::function<T()> task_;
813 trantor::EventLoop *loop_;
816template <
typename... Tasks>
819 std::tuple<internal::void_to_false_t<await_result_t<Tasks>>...>>
821 WhenAllAwaiter(Tasks... tasks)
822 : tasks_(std::forward<Tasks>(tasks)...), counter_(
sizeof...(tasks))
826 void await_suspend(std::coroutine_handle<> handle)
834 await_suspend_impl(handle, std::index_sequence_for<Tasks...>{});
838 std::tuple<Tasks...> tasks_;
839 std::atomic<size_t> counter_;
840 std::tuple<internal::void_to_false_t<await_result_t<Tasks>>...> results_;
841 std::atomic_flag exceptionFlag_;
843 template <
size_t Idx>
844 void launch_task(std::coroutine_handle<> handle)
846 using Self = WhenAllAwaiter<Tasks...>;
847 [](Self *self, std::coroutine_handle<> handle) ->
AsyncTask {
850 using TaskType = std::tuple_element_t<
852 std::remove_cvref_t<
decltype(results_)>>;
853 if constexpr (std::is_same_v<TaskType, std::false_type>)
855 co_await std::get<Idx>(self->tasks_);
856 std::get<Idx>(self->results_) = std::false_type{};
860 std::get<Idx>(self->results_) =
861 co_await std::get<Idx>(self->tasks_);
866 if (self->exceptionFlag_.test_and_set() ==
false)
867 self->setException(std::current_exception());
870 if (self->counter_.fetch_sub(1, std::memory_order_acq_rel) == 1)
872 if (!self->hasException())
873 self->setValue(std::move(self->results_));
879 template <
size_t... Is>
880 void await_suspend_impl(std::coroutine_handle<> handle,
881 std::index_sequence<Is...>)
883 ((launch_task<Is>(handle)), ...);
888struct WhenAllAwaiter<
std::vector<Task<T>>>
891 WhenAllAwaiter(std::vector<
Task<T>> tasks)
892 : tasks_(std::move(tasks)),
893 counter_(tasks_.size()),
894 results_(tasks_.size())
898 void await_suspend(std::coroutine_handle<> handle)
902 this->setValue(std::vector<T>{});
907 const size_t count = tasks_.size();
908 for (
size_t i = 0; i < count; ++i)
910 [](WhenAllAwaiter *self,
911 std::coroutine_handle<> handle,
916 auto result =
co_await task;
917 self->results_[index] = std::move(result);
921 if (self->exceptionFlag_.test_and_set() ==
false)
922 self->setException(std::current_exception());
925 if (self->counter_.fetch_sub(1, std::memory_order_acq_rel) == 1)
927 if (!self->hasException())
929 self->setValue(std::move(self->results_));
933 }(
this, handle, std::move(tasks_[i]), i);
938 std::vector<Task<T>> tasks_;
939 std::atomic<size_t> counter_;
940 std::vector<T> results_;
941 std::atomic_flag exceptionFlag_;
948 : tasks_(std::move(t)), counter_(tasks_.size())
952 void await_suspend(std::coroutine_handle<> handle)
963 for (
size_t i = 0; i < count; ++i)
965 [](WhenAllAwaiter *self,
966 std::coroutine_handle<> handle,
974 if (self->exceptionFlag_.test_and_set() ==
false)
975 self->setException(std::current_exception());
977 if (self->counter_.fetch_sub(1, std::memory_order_acq_rel) == 1)
981 }(
this, handle, std::move(tasks_[i]));
985 std::vector<Task<void>> tasks_;
986 std::atomic<size_t> counter_;
987 std::atomic_flag exceptionFlag_;
997 std::function<T()> task)
1004 class ScopedCoroMutexAwaiter;
1005 class CoroMutexAwaiter;
1008 Mutex() noexcept : state_(unlockedValue()), waiters_(
nullptr)
1012 Mutex(
const Mutex &) =
delete;
1013 Mutex(Mutex &&) =
delete;
1014 Mutex &operator=(
const Mutex &) =
delete;
1015 Mutex &operator=(Mutex &&) =
delete;
1019 [[maybe_unused]]
auto state = state_.load(std::memory_order_relaxed);
1020 assert(state == unlockedValue() || state ==
nullptr);
1021 assert(waiters_ ==
nullptr);
1024 bool try_lock()
noexcept
1026 void *oldValue = unlockedValue();
1027 return state_.compare_exchange_strong(oldValue,
1029 std::memory_order_acquire,
1030 std::memory_order_relaxed);
1033 [[nodiscard]] ScopedCoroMutexAwaiter scoped_lock(
1034 trantor::EventLoop *loop =
1035 trantor::EventLoop::getEventLoopOfCurrentThread())
noexcept
1037 return ScopedCoroMutexAwaiter(*
this, loop);
1040 [[nodiscard]] CoroMutexAwaiter lock(
1041 trantor::EventLoop *loop =
1042 trantor::EventLoop::getEventLoopOfCurrentThread())
noexcept
1044 return CoroMutexAwaiter(*
this, loop);
1047 void unlock()
noexcept
1049 assert(state_.load(std::memory_order_relaxed) != unlockedValue());
1050 auto *waitersHead = waiters_;
1051 if (waitersHead ==
nullptr)
1053 void *currentState = state_.load(std::memory_order_relaxed);
1054 if (currentState ==
nullptr)
1056 const bool releasedLock =
1057 state_.compare_exchange_strong(currentState,
1059 std::memory_order_release,
1060 std::memory_order_relaxed);
1066 currentState = state_.exchange(
nullptr, std::memory_order_acquire);
1067 assert(currentState != unlockedValue());
1068 assert(currentState !=
nullptr);
1069 auto *waiter =
static_cast<CoroMutexAwaiter *
>(currentState);
1072 auto *temp = waiter->next_;
1073 waiter->next_ = waitersHead;
1074 waitersHead = waiter;
1076 }
while (waiter !=
nullptr);
1078 assert(waitersHead !=
nullptr);
1079 waiters_ = waitersHead->next_;
1080 if (waitersHead->loop_)
1082 auto handle = waitersHead->handle_;
1083 waitersHead->loop_->runInLoop([handle] { handle.resume(); });
1087 waitersHead->handle_.resume();
1092 class CoroMutexAwaiter
1095 CoroMutexAwaiter(Mutex &mutex, trantor::EventLoop *loop) noexcept
1096 : mutex_(mutex), loop_(loop)
1100 bool await_ready()
noexcept
1102 return mutex_.try_lock();
1105 bool await_suspend(std::coroutine_handle<> handle)
noexcept
1108 return mutex_.asynclockImpl(
this);
1111 void await_resume()
noexcept
1119 trantor::EventLoop *loop_;
1120 std::coroutine_handle<> handle_;
1121 CoroMutexAwaiter *next_;
1124 class ScopedCoroMutexAwaiter :
public CoroMutexAwaiter
1127 ScopedCoroMutexAwaiter(Mutex &mutex, trantor::EventLoop *loop)
1128 : CoroMutexAwaiter(mutex, loop)
1132 [[nodiscard]]
auto await_resume()
noexcept
1134 return std::unique_lock<Mutex>{mutex_, std::adopt_lock};
1138 bool asynclockImpl(CoroMutexAwaiter *awaiter)
1140 void *oldValue = state_.load(std::memory_order_relaxed);
1143 if (oldValue == unlockedValue())
1145 void *newValue =
nullptr;
1146 if (state_.compare_exchange_weak(oldValue,
1148 std::memory_order_acquire,
1149 std::memory_order_relaxed))
1156 void *newValue = awaiter;
1157 awaiter->next_ =
static_cast<CoroMutexAwaiter *
>(oldValue);
1158 if (state_.compare_exchange_weak(oldValue,
1160 std::memory_order_release,
1161 std::memory_order_relaxed))
1169 void *unlockedValue()
noexcept
1174 std::atomic<void *> state_;
1175 CoroMutexAwaiter *waiters_;
1178template <
typename... Tasks>
1184template <
typename T>
1185internal::WhenAllAwaiter<std::vector<Task<T>>> when_all(
1186 std::vector<Task<T>> tasks)
1188 return internal::WhenAllAwaiter(std::move(tasks));
Drogon Test is a minimal effort test framework developed because the major C++ test frameworks doesn'...
Definition Attribute.h:23
void async_run(Coro &&coro)
Runs a coroutine from a regular function.
Definition coroutine.h:749
std::function< void()> async_func(Coro &&coro)
returns a function that calls a coroutine
Definition coroutine.h:769
Definition coroutine.h:386
Definition coroutine.h:355
Definition coroutine.h:421
Definition coroutine.h:198
Definition coroutine.h:302
Definition coroutine.h:161
Definition coroutine.h:66
An awaiter for Task::promise_type::final_suspend(). Transfer execution back to the coroutine who is c...
Definition coroutine.h:96
Definition coroutine.h:669
Definition coroutine.h:780
Definition coroutine.h:614
Definition coroutine.h:654
Definition coroutine.h:591
Definition coroutine.h:820
Definition coroutine.h:76
Definition coroutine.h:725
Convert Task to an awaiter when it is co_awaited. Following things will happen:
Definition coroutine.h:123