diff --git a/include/exec/timed_thread_scheduler.hpp b/include/exec/timed_thread_scheduler.hpp index 45a8914d8..e43aada00 100644 --- a/include/exec/timed_thread_scheduler.hpp +++ b/include/exec/timed_thread_scheduler.hpp @@ -149,6 +149,12 @@ namespace experimental::execution using stop_type = _time_thrd_sched::timed_thread_stop_operation; using time_point = std::chrono::steady_clock::time_point; + [[nodiscard]] + auto is_running_thread() const noexcept -> bool + { + return std::this_thread::get_id() == run_thread_.get_id(); + } + void run() { while (true) @@ -211,6 +217,22 @@ namespace experimental::execution void schedule(command_type* op) { + // A stop requested on the context thread completes inline only if it finds the operation + // in the heap (see request_stop()), so on this thread insert straight into the heap + // rather than queueing a schedule command. run() rereads heap_.front() after every + // completion, so the operation is still seen before the thread next waits. + if (op->command_ == command_type::command_type::schedule && is_running_thread()) + { + auto* task = static_cast(op); + if (n_submissions_in_flight_.load(STDEXEC::__std::memory_order_relaxed) < 0) + { + task->set_stopped_(task); + return; + } + task->when_ = _time_thrd_sched::when_type{task->time_point_, submission_counter_++}; + heap_.insert(task); + return; + } std::ptrdiff_t n = n_submissions_in_flight_.fetch_add(1, STDEXEC::__std::memory_order_relaxed); if (n < 0) @@ -348,7 +370,19 @@ namespace experimental::execution constexpr void request_stop() noexcept { - if (ref_count_.fetch_add(1, STDEXEC::__std::memory_order_relaxed) == 1) + // On the context thread, an operation still waiting in the heap (ref_count_ == 1: + // started, no stop in flight) is removed and completed with set_stopped() inline. + // Queueing a stop command instead would hold the completion until run() next drains + // its queue. + if (context_.is_running_thread() + && ref_count_.load(STDEXEC::__std::memory_order_relaxed) == 1 + && context_.heap_.erase(this)) + { + ref_count_.store(0, STDEXEC::__std::memory_order_relaxed); + stop_callback_.reset(); + STDEXEC::set_stopped(std::move(rcvr_)); + } + else if (ref_count_.fetch_add(1, STDEXEC::__std::memory_order_relaxed) == 1) { context_.schedule(&stop_op_); } diff --git a/test/exec/test_timed_thread_scheduler.cpp b/test/exec/test_timed_thread_scheduler.cpp index cd2b9ca58..67777858d 100644 --- a/test/exec/test_timed_thread_scheduler.cpp +++ b/test/exec/test_timed_thread_scheduler.cpp @@ -336,6 +336,28 @@ namespace CHECK(n == 1); } + TEST_CASE("timed_thread_scheduler - stop on the context thread completes inline", + "[timed_thread_scheduler][when_any]") + { + exec::timed_thread_context context; + exec::timed_thread_scheduler scheduler = context.get_scheduler(); + exec::async_scope scope; + bool completed_before_start_returned = false; + STDEXEC::sync_wait( + STDEXEC::schedule(scheduler) + | STDEXEC::then( + [&] + { + bool completed = false; + scope.spawn(exec::when_any(exec::schedule_after(scheduler, std::chrono::seconds(5)), + STDEXEC::just()) + | STDEXEC::then([&completed] { completed = true; })); + completed_before_start_returned = completed; + })); + CHECK(completed_before_start_returned); + CHECK(STDEXEC::sync_wait(scope.on_empty())); + } + TEST_CASE("timed_thread_scheduler - many timers with async scope", "[timed_thread_scheduler][async_scope]") {