Oct 2, 2026

[coroutine] c9 idiom of cancellation.

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.