1. Who creates the coroutine?
↓
2. Who owns the coroutine_handle?
↓
3. Who resumes it?
↓
4. What event causes the resume?
↓
5. Who destroys the frame?
// Minimal single-threaded model of C9-style cancellation.
//
// Key difference from the stop_token version: a cancelled coroutine is NEVER
// resumed. The cancel callback destroys every frame of the logical thread
// (leaf -> entry), running destructors, and then resumes an "on_exit" handle
// so the owner learns the thread is gone. No code checks a token.
//
// Name map to google3/util/c9/internal:
// PromiseLink -> CoroutinePromiseLink (waiter link, co_thread)
// CoThread -> CoThread (cancelled_ AsyncNotification, on_exit)
// SuspensionToken -> SuspensionToken (BecameReady/Cancelled/AlreadyCancelled,
// DestroyCallStack)
// Scheduler -> ThreadState ready queue + Schedule/CancelWaiter
// Yield / CondVar -> leaf awaitables like c9::Event::Awaitable
#include <algorithm>
#include <coroutine>
#include <deque>
#include <exception>
#include <functional>
#include <iostream>
#include <optional>
#include <type_traits>
#include <utility>
struct FrameCounter {
static inline int live = 0;
FrameCounter() { ++live; }
~FrameCounter() { --live; }
};
class CoThread;
// Base of every Co<T> promise. The `waiter` pointers form the intrusive
// "coroutine call stack" that cancellation walks and destroys.
struct PromiseLink {
std::coroutine_handle<> self; // this frame
PromiseLink* waiter = nullptr; // caller frame; nullptr for the entry frame
CoThread* co_thread = nullptr; // logical thread this frame belongs to
};
// One logical thread of execution = one chain of nested co_awaits.
class CoThread {
public:
explicit CoThread(std::coroutine_handle<> on_exit) : on_exit_(on_exit) {}
std::coroutine_handle<> on_exit() const { return on_exit_; }
// Called by the leaf awaitable while suspending. False = already cancelled,
// so the leaf must not park; it must destroy the stack instead.
bool Register(std::function<void()> on_cancel) {
if (cancelled_) return false;
on_cancel_ = std::move(on_cancel);
return true;
}
void Unregister() { on_cancel_ = nullptr; }
// Request cancellation. If a leaf is parked, its callback runs now and
// destroys the whole chain. If the chain is currently running, nothing
// happens here; it dies at its next suspension point (Register -> false).
void Cancel() {
cancelled_ = true;
if (on_cancel_) std::exchange(on_cancel_, nullptr)();
}
private:
std::coroutine_handle<> on_exit_;
bool cancelled_ = false;
std::function<void()> on_cancel_;
};
// Handed to a leaf awaitable at suspension. Exactly one of BecameReady,
// Cancelled, AlreadyCancelled is called for each suspension.
class SuspensionToken {
public:
explicit SuspensionToken(PromiseLink* leaf) : leaf_(leaf) {}
bool operator==(const SuspensionToken&) const = default;
bool RegisterForCancellation(std::function<void()> on_cancel) {
return co_thread().Register(std::move(on_cancel));
}
// Normal wakeup: stop listening for cancel; caller resumes the leaf.
std::coroutine_handle<> BecameReady() {
co_thread().Unregister();
return leaf_->self;
}
// Cancel callback won the race: destroy the chain, tell the owner. The leaf
// coroutine is never resumed.
void Cancelled() {
std::coroutine_handle<> on_exit = co_thread().on_exit(); // read first
DestroyCallStack(); // frees leaf_
on_exit.resume();
}
// Register failed (cancelled before we could park): destroy the chain and
// return on_exit for the caller to symmetric-transfer into.
std::coroutine_handle<> AlreadyCancelled() {
std::coroutine_handle<> on_exit = co_thread().on_exit();
DestroyCallStack();
return on_exit;
}
private:
CoThread& co_thread() { return *leaf_->co_thread; }
// Callee before caller, iteratively, so OS stack depth stays O(1).
void DestroyCallStack() {
for (PromiseLink* p = leaf_; p != nullptr;) {
PromiseLink* next = p->waiter;
p->self.destroy(); // runs this frame's destructors (RAII)
p = next;
}
}
PromiseLink* leaf_;
};
// Ready queue of wakeups. Holds tokens, not raw handles, so BecameReady runs
// (unregistering the cancel callback) right before resume.
class Scheduler {
public:
void Schedule(SuspensionToken t) { ready_.push_back(t); }
// Pull a scheduled wakeup back out so the frame can be destroyed instead.
// False = already dequeued and running; cancellation then takes effect at
// the coroutine's next suspension point.
bool CancelWaiter(const SuspensionToken& t) {
return std::erase(ready_, t) > 0;
}
void Run() {
while (!ready_.empty()) {
SuspensionToken t = ready_.front();
ready_.pop_front();
t.BecameReady().resume();
}
}
private:
std::deque<SuspensionToken> ready_;
};
// Awaitable child coroutine. Lazy start, symmetric transfer, frame destroys
// itself at final_suspend (result is written into the parent's frame first).
template <typename T>
class Co {
static constexpr bool kVoid = std::is_void_v<T>;
struct Empty {};
using Slot = std::conditional_t<kVoid, Empty, std::optional<T>>;
template <typename U>
struct ReturnValue {
std::optional<U>* result = nullptr; // lives in the awaiting parent
void return_value(U v) { result->emplace(std::move(v)); }
};
struct ReturnVoid {
void return_void() {}
};
public:
struct promise_type
: PromiseLink,
FrameCounter,
std::conditional_t<kVoid, ReturnVoid, ReturnValue<T>> {
Co get_return_object() {
auto h = std::coroutine_handle<promise_type>::from_promise(*this);
this->self = h;
return Co(h);
}
std::suspend_always initial_suspend() noexcept { return {}; }
auto final_suspend() noexcept {
struct Final {
bool await_ready() noexcept { return false; }
std::coroutine_handle<> await_suspend(
std::coroutine_handle<promise_type> h) noexcept {
PromiseLink* waiter = h.promise().waiter;
std::coroutine_handle<> on_exit = h.promise().co_thread->on_exit();
h.destroy(); // frame gone; only locals used from here on
return waiter != nullptr ? waiter->self : on_exit;
}
void await_resume() noexcept {}
};
return Final{};
}
void unhandled_exception() { std::terminate(); }
};
explicit Co(std::coroutine_handle<promise_type> h) : h_(h) {}
Co(Co&& o) noexcept : h_(std::exchange(o.h_, {})) {}
~Co() {
if (h_) h_.destroy(); // only if never started
}
bool await_ready() const noexcept { return false; }
template <typename P>
std::coroutine_handle<> await_suspend(std::coroutine_handle<P> parent) {
promise_type& p = h_.promise();
p.waiter = &parent.promise();
p.co_thread = parent.promise().co_thread;
if constexpr (!kVoid) p.result = &result_;
return std::exchange(h_, {}); // frame now owns itself
}
T await_resume() {
if constexpr (!kVoid) return std::move(*result_);
}
// For bridges only: hand the entry frame to a new CoThread.
std::coroutine_handle<promise_type> Release() {
return std::exchange(h_, {});
}
private:
std::coroutine_handle<promise_type> h_;
Slot result_;
};
// Tiny coroutine used as a CoThread's on_exit handle.
struct OnExit {
struct promise_type : FrameCounter {
OnExit get_return_object() {
return {std::coroutine_handle<promise_type>::from_promise(*this)};
}
std::suspend_always initial_suspend() noexcept { return {}; }
std::suspend_never final_suspend() noexcept { return {}; }
void return_void() {}
void unhandled_exception() { std::terminate(); }
};
std::coroutine_handle<> handle;
};
OnExit ReportExit(const char* name) {
std::cout << name << ": thread exited\n";
co_return;
}
// Bridge: start `co` as the entry frame of a fresh CoThread (what C9 bridges,
// RunConcurrently and WithDeadline do for their children).
void Spawn(Scheduler& sched, CoThread& ct, Co<void> co) {
auto h = co.Release();
h.promise().co_thread = &ct; // waiter stays nullptr: entry frame
sched.Schedule(SuspensionToken(&h.promise()));
}
// ---- Leaf awaitables (the only place cancellation is "implemented") ----
// Re-queue the current coroutine (simulates work). Cancellable.
class Yield {
public:
explicit Yield(Scheduler& s) : sched_(s) {}
bool await_ready() const noexcept { return false; }
template <typename P>
std::coroutine_handle<> await_suspend(std::coroutine_handle<P> h) {
token_.emplace(&h.promise());
if (!token_->RegisterForCancellation([this] { Cancel(); })) {
return token_->AlreadyCancelled();
}
sched_.Schedule(*token_);
return std::noop_coroutine();
}
void await_resume() const noexcept {}
private:
void Cancel() {
// Copy out: Cancelled() destroys the frame that holds *this.
Scheduler& sched = sched_;
SuspensionToken token = *token_;
if (sched.CancelWaiter(token)) token.Cancelled();
}
Scheduler& sched_;
std::optional<SuspensionToken> token_;
};
// Condition variable. Note Wait() takes no token: cancellation is implicit.
class CondVar {
public:
explicit CondVar(Scheduler& sched) : sched_(sched) {}
class Awaiter {
public:
explicit Awaiter(CondVar& cv) : cv_(cv) {}
bool await_ready() const noexcept { return false; }
template <typename P>
std::coroutine_handle<> await_suspend(std::coroutine_handle<P> h) {
token_.emplace(&h.promise());
if (!token_->RegisterForCancellation([this] { Cancel(); })) {
return token_->AlreadyCancelled();
}
cv_.waiters_.push_back(*token_);
return std::noop_coroutine();
}
void await_resume() const noexcept {}
private:
void Cancel() {
CondVar& cv = cv_;
SuspensionToken token = *token_;
// Either still parked here, or NotifyOne already moved it to the ready
// queue. Pull it back from wherever it is, then destroy.
if (std::erase(cv.waiters_, token) > 0 || cv.sched_.CancelWaiter(token)) {
token.Cancelled();
}
}
CondVar& cv_;
std::optional<SuspensionToken> token_;
};
Awaiter Wait() { return Awaiter(*this); }
void NotifyOne() {
if (waiters_.empty()) return;
sched_.Schedule(waiters_.front());
waiters_.pop_front();
}
private:
Scheduler& sched_;
std::deque<SuspensionToken> waiters_;
};
// ---- User code: no token anywhere ----
class Channel {
public:
explicit Channel(Scheduler& sched) : cv_(sched) {}
void Push(int v) {
items_.push_back(v);
cv_.NotifyOne();
}
Co<int> Pop() {
while (items_.empty()) co_await cv_.Wait();
int v = items_.front();
items_.pop_front();
co_return v;
}
private:
std::deque<int> items_;
CondVar cv_;
};
Co<void> Consumer(Scheduler& sched, Channel& ch, int id) {
struct Cleanup {
int id;
~Cleanup() { std::cout << "consumer " << id << " cleanup (RAII)\n"; }
} cleanup{id};
for (;;) { // no exit path: cancellation destroys this frame instead
int item = co_await ch.Pop();
std::cout << "consumer " << id << " processes " << item << "\n";
co_await Yield(sched);
}
}
Co<void> Producer(Scheduler& sched, Channel& ch) {
for (int i = 0;; ++i) {
std::cout << "producer pushes " << i << "\n";
ch.Push(i);
co_await Yield(sched);
}
}
Co<void> StopAfter(Scheduler& sched, std::deque<CoThread*> threads, int ticks) {
for (int i = 0; i < ticks; ++i) co_await Yield(sched);
std::cout << "=== cancel requested ===\n";
for (CoThread* t : threads) t->Cancel(); // callbacks destroy the stacks
}
int main() {
Scheduler sched;
Channel ch(sched);
CoThread c1(ReportExit("consumer 1").handle);
CoThread c2(ReportExit("consumer 2").handle);
CoThread prod(ReportExit("producer").handle);
CoThread stopper(ReportExit("stopper").handle);
Spawn(sched, c1, Consumer(sched, ch, 1));
Spawn(sched, c2, Consumer(sched, ch, 2));
Spawn(sched, prod, Producer(sched, ch));
Spawn(sched, stopper, StopAfter(sched, {&c1, &c2, &prod}, 3));
sched.Run();
std::cout << "live coroutine frames after Run(): " << FrameCounter::live
<< "\n";
return FrameCounter::live == 0 ? 0 : 1;
}
No comments:
Post a Comment
Note: Only a member of this blog may post a comment.