Skip to main content

async_runtime/
runtime.rs

1use crate::error::{ShutdownError, ShutdownOutcome, SpawnError};
2use crate::priority::{Priority, PriorityWeights};
3use crate::scheduler::Scheduler;
4use crate::task::Task;
5use crate::worker;
6use async_task::Builder as TaskBuilder;
7use std::future::Future;
8use std::io;
9use std::num::NonZeroUsize;
10use std::sync::atomic::{AtomicUsize, Ordering};
11use std::sync::{Arc, Condvar, Mutex, Weak};
12use std::thread::{self, JoinHandle, ThreadId};
13use std::time::Duration;
14
15pub struct RuntimeBuilder {
16    worker_threads: NonZeroUsize,
17    weights: PriorityWeights,
18}
19
20impl RuntimeBuilder {
21    /// Creates a builder for an explicit, non-zero worker count.
22    pub fn new(worker_threads: NonZeroUsize) -> Self {
23        Self {
24            worker_threads,
25            weights: PriorityWeights::default(),
26        }
27    }
28    /// Replaces the default `8:4:1` priority weights.
29    #[must_use]
30    pub fn priority_weights(mut self, weights: PriorityWeights) -> Self {
31        self.weights = weights;
32        self
33    }
34    /// Starts all worker threads and returns the owning runtime handle.
35    ///
36    /// # Errors
37    ///
38    /// Returns an I/O error if any worker thread cannot be created. Workers
39    /// created earlier in the same build attempt are stopped and joined.
40    pub fn build(self) -> io::Result<Runtime> {
41        let (scheduler, worker_queues) = Scheduler::new(self.worker_threads.get());
42        let state = Arc::new(RuntimeState {
43            scheduler,
44            gate: Mutex::new(Gate::Running),
45            accepted_tasks: AtomicUsize::new(0),
46            drain_lock: Mutex::new(()),
47            drained: Condvar::new(),
48            worker_ids: Mutex::new(Vec::with_capacity(self.worker_threads.get())),
49            #[cfg(test)]
50            admission_pause: Mutex::new(None),
51            #[cfg(test)]
52            last_completion_pause: Mutex::new(None),
53        });
54        let mut workers = Vec::with_capacity(self.worker_threads.get());
55        for (number, queues) in worker_queues.into_iter().enumerate() {
56            let worker_state = Arc::clone(&state);
57            match thread::Builder::new()
58                .name(format!("async-runtime-{number}"))
59                .spawn(move || worker::run(&worker_state, number, queues, self.weights))
60            {
61                Ok(handle) => workers.push(handle),
62                Err(error) => {
63                    state.request_stop();
64                    for handle in workers {
65                        let _ = handle.join();
66                    }
67                    return Err(error);
68                }
69            }
70        }
71        Ok(Runtime {
72            state,
73            workers: Mutex::new(workers),
74            closed: false,
75        })
76    }
77}
78
79pub struct Runtime {
80    pub(crate) state: Arc<RuntimeState>,
81    workers: Mutex<Vec<JoinHandle<()>>>,
82    closed: bool,
83}
84
85impl Runtime {
86    /// Returns a weak capability for submitting work to this runtime.
87    pub fn spawner(&self) -> Spawner {
88        Spawner {
89            state: Arc::downgrade(&self.state),
90        }
91    }
92    /// Submits a task to one priority queue.
93    ///
94    /// # Errors
95    ///
96    /// Returns [`SpawnError::Closed`] after shutdown begins.
97    ///
98    /// A successful return means the task was admitted into the runtime's
99    /// lifecycle. A concurrent forced or timed shutdown may cancel it before
100    /// its first poll; graceful shutdown still waits for it to reach a terminal
101    /// state.
102    pub fn spawn<F, T>(&self, priority: Priority, future: F) -> Result<Task<T>, SpawnError>
103    where
104        F: Future<Output = T> + Send + 'static,
105        T: Send + 'static,
106    {
107        self.state.spawn(priority, future)
108    }
109
110    /// Returns a relaxed snapshot of scheduler observability counters.
111    #[cfg(feature = "stats")]
112    pub fn stats(&self) -> crate::RuntimeStats {
113        self.state.scheduler.stats()
114    }
115    /// Rejects new tasks, drains every accepted task, and joins all workers.
116    ///
117    /// # Errors
118    ///
119    /// Returns [`ShutdownError::CalledFromWorker`] when invoked by this
120    /// runtime's worker, or [`ShutdownError::WorkerPanicked`] if a joined
121    /// worker panicked.
122    pub fn shutdown_graceful(mut self) -> Result<(), ShutdownError> {
123        if self.state.is_current_worker() {
124            self.state.begin_close();
125            self.state.request_stop();
126            self.closed = true;
127            let _ = self.join_workers();
128            self.state.finish_close();
129            return Err(ShutdownError::CalledFromWorker);
130        }
131        self.state.begin_close();
132        self.state.wait_for_drain();
133        self.state.request_stop();
134        self.closed = true;
135        let result = self.join_workers();
136        self.state.finish_close();
137        result
138    }
139    /// Drains accepted tasks until `timeout`, then cancels the remainder.
140    ///
141    /// # Errors
142    ///
143    /// Returns [`ShutdownError::CalledFromWorker`] when invoked by this
144    /// runtime's worker, or [`ShutdownError::WorkerPanicked`] if a joined
145    /// worker panicked.
146    pub fn shutdown_timeout(mut self, timeout: Duration) -> Result<ShutdownOutcome, ShutdownError> {
147        if self.state.is_current_worker() {
148            self.state.begin_close();
149            self.state.request_stop();
150            self.closed = true;
151            let _ = self.join_workers();
152            self.state.finish_close();
153            return Err(ShutdownError::CalledFromWorker);
154        }
155        self.state.begin_close();
156        let outcome = if self.state.wait_for_drain_timeout(timeout) {
157            ShutdownOutcome::Completed
158        } else {
159            ShutdownOutcome::TimedOut {
160                remaining_tasks: self.state.accepted_tasks.load(Ordering::Acquire),
161            }
162        };
163        self.state.request_stop();
164        self.closed = true;
165        let result = self.join_workers();
166        self.state.finish_close();
167        result?;
168        Ok(outcome)
169    }
170    /// Cancels remaining work and joins all workers.
171    ///
172    /// # Errors
173    ///
174    /// Returns [`ShutdownError::CalledFromWorker`] when invoked by this
175    /// runtime's worker, or [`ShutdownError::WorkerPanicked`] if a joined
176    /// worker panicked.
177    pub fn shutdown_now(mut self) -> Result<(), ShutdownError> {
178        let called_from_worker = self.state.is_current_worker();
179        self.state.begin_close();
180        self.state.request_stop();
181        self.closed = true;
182        let result = self.join_workers();
183        self.state.finish_close();
184        if called_from_worker {
185            Err(ShutdownError::CalledFromWorker)
186        } else {
187            result
188        }
189    }
190    fn join_workers(&self) -> Result<(), ShutdownError> {
191        let current = thread::current().id();
192        let workers = {
193            let mut workers = self.workers.lock().expect("runtime worker list poisoned");
194            std::mem::take(&mut *workers)
195        };
196        let mut panicked = false;
197        for worker in workers {
198            if worker.thread().id() == current {
199                continue;
200            }
201            if worker.join().is_err() {
202                panicked = true;
203            }
204        }
205        if panicked {
206            Err(ShutdownError::WorkerPanicked)
207        } else {
208            Ok(())
209        }
210    }
211}
212
213impl Drop for Runtime {
214    fn drop(&mut self) {
215        if !self.closed {
216            self.state.begin_close();
217            self.state.request_stop();
218            let _ = self.join_workers();
219            self.state.finish_close();
220            self.closed = true;
221        }
222    }
223}
224
225#[derive(Clone)]
226pub struct Spawner {
227    state: Weak<RuntimeState>,
228}
229impl Spawner {
230    /// Submits a task to one priority queue.
231    ///
232    /// A successful return means the task was admitted into the runtime's
233    /// lifecycle. A concurrent forced or timed shutdown may cancel it before
234    /// its first poll; graceful shutdown still waits for it to reach a terminal
235    /// state.
236    ///
237    /// # Errors
238    ///
239    /// Returns [`SpawnError::Closed`] if the runtime no longer exists or has
240    /// begun shutting down.
241    pub fn spawn<F, T>(&self, priority: Priority, future: F) -> Result<Task<T>, SpawnError>
242    where
243        F: Future<Output = T> + Send + 'static,
244        T: Send + 'static,
245    {
246        self.state
247            .upgrade()
248            .ok_or(SpawnError::Closed)?
249            .spawn(priority, future)
250    }
251}
252
253#[derive(Clone, Copy, PartialEq, Eq)]
254enum Gate {
255    Running,
256    Closing,
257    Closed,
258}
259
260/// Shared worker state. Its gate serializes task acceptance with transition to closing.
261pub(crate) struct RuntimeState {
262    pub(crate) scheduler: Arc<Scheduler>,
263    gate: Mutex<Gate>,
264    pub(crate) accepted_tasks: AtomicUsize,
265    drain_lock: Mutex<()>,
266    drained: Condvar,
267    worker_ids: Mutex<Vec<ThreadId>>,
268    #[cfg(test)]
269    admission_pause: Mutex<Option<Arc<TestPause>>>,
270    #[cfg(test)]
271    last_completion_pause: Mutex<Option<Arc<TestPause>>>,
272}
273
274impl RuntimeState {
275    fn spawn<F, T>(self: &Arc<Self>, priority: Priority, future: F) -> Result<Task<T>, SpawnError>
276    where
277        F: Future<Output = T> + Send + 'static,
278        T: Send + 'static,
279    {
280        let gate = self.gate.lock().expect("runtime lifecycle gate poisoned");
281        if *gate != Gate::Running {
282            return Err(SpawnError::Closed);
283        }
284        self.accepted_tasks.fetch_add(1, Ordering::AcqRel);
285        // This token begins at admission, not at the first poll. It is moved
286        // into the task immediately so construction unwind, cancellation,
287        // forced shutdown, and normal completion all retire the admission
288        // exactly once.
289        let completion = CompletionGuard {
290            // Tasks are owned by the scheduler inside this state. Keeping only a weak
291            // reference here is essential: a queued task must not keep the scheduler
292            // (and therefore itself) alive during shutdown_now or Runtime::drop.
293            state: Arc::downgrade(self),
294        };
295        drop(gate);
296        #[cfg(test)]
297        self.pause_after_admission();
298        let tracked = async move {
299            let _completion = completion;
300            future.await
301        };
302        let scheduler = Arc::downgrade(&self.scheduler);
303        let schedule = move |runnable| {
304            if let Some(scheduler) = scheduler.upgrade() {
305                scheduler.schedule(priority, runnable);
306            }
307        };
308        let (runnable, task) = TaskBuilder::new()
309            .propagate_panic(true)
310            .spawn(|()| tracked, schedule);
311        self.scheduler
312            .record_spawn(self.scheduler.is_current_worker());
313        // A successful return means the runtime owns the queued task or its
314        // cancellation cleanup. Initial scheduling follows the same local vs
315        // global routing rule as every later wake.
316        runnable.schedule();
317        Ok(Task::direct(task))
318    }
319    pub(crate) fn register_worker(&self, id: ThreadId) {
320        self.worker_ids
321            .lock()
322            .expect("runtime worker id list poisoned")
323            .push(id);
324    }
325    fn is_current_worker(&self) -> bool {
326        self.worker_ids
327            .lock()
328            .expect("runtime worker id list poisoned")
329            .contains(&thread::current().id())
330    }
331    fn begin_close(&self) {
332        let mut gate = self.gate.lock().expect("runtime lifecycle gate poisoned");
333        if *gate == Gate::Running {
334            *gate = Gate::Closing;
335        }
336    }
337    fn finish_close(&self) {
338        *self.gate.lock().expect("runtime lifecycle gate poisoned") = Gate::Closed;
339    }
340    pub(crate) fn request_stop(&self) {
341        self.scheduler.stop();
342    }
343    fn wait_for_drain(&self) {
344        let guard = self.drain_lock.lock().expect("runtime drain lock poisoned");
345        drop(
346            self.drained
347                .wait_while(guard, |()| self.accepted_tasks.load(Ordering::Acquire) != 0)
348                .expect("runtime drain lock poisoned"),
349        );
350    }
351    fn wait_for_drain_timeout(&self, timeout: Duration) -> bool {
352        let guard = self.drain_lock.lock().expect("runtime drain lock poisoned");
353        let (guard, _) = self
354            .drained
355            .wait_timeout_while(guard, timeout, |()| {
356                self.accepted_tasks.load(Ordering::Acquire) != 0
357            })
358            .expect("runtime drain lock poisoned");
359        let drained = self.accepted_tasks.load(Ordering::Acquire) == 0;
360        drop(guard);
361        drained
362    }
363    fn complete_task(&self) {
364        let previous = self.accepted_tasks.fetch_sub(1, Ordering::AcqRel);
365        assert!(previous > 0, "runtime accepted task count underflow");
366        if previous == 1 {
367            #[cfg(test)]
368            self.pause_before_last_completion_notify();
369            // The waiter checks the zero predicate while holding this same
370            // mutex. Locking before notify closes both interleavings: either
371            // the waiter observes zero, or it has atomically begun waiting.
372            let _guard = self.drain_lock.lock().expect("runtime drain lock poisoned");
373            self.drained.notify_all();
374        }
375    }
376
377    #[cfg(test)]
378    fn pause_after_admission(&self) {
379        let pause = self
380            .admission_pause
381            .lock()
382            .expect("runtime test hook poisoned")
383            .clone();
384        if let Some(pause) = pause {
385            pause.pause();
386        }
387    }
388
389    #[cfg(test)]
390    fn pause_before_last_completion_notify(&self) {
391        let pause = self
392            .last_completion_pause
393            .lock()
394            .expect("runtime test hook poisoned")
395            .clone();
396        if let Some(pause) = pause {
397            pause.pause();
398        }
399    }
400}
401
402struct CompletionGuard {
403    state: Weak<RuntimeState>,
404}
405impl Drop for CompletionGuard {
406    fn drop(&mut self) {
407        // Graceful shutdown keeps the runtime state alive, so every accepted task
408        // contributes to its drain count. During forced teardown the state may
409        // already be gone; there is then no waiter left to notify.
410        if let Some(state) = self.state.upgrade() {
411            state.complete_task();
412        }
413    }
414}
415
416#[cfg(test)]
417struct TestPause {
418    reached: std::sync::Barrier,
419    released: std::sync::Barrier,
420}
421
422#[cfg(test)]
423impl TestPause {
424    fn new() -> Arc<Self> {
425        Arc::new(Self {
426            reached: std::sync::Barrier::new(2),
427            released: std::sync::Barrier::new(2),
428        })
429    }
430
431    fn pause(&self) {
432        self.reached.wait();
433        self.released.wait();
434    }
435
436    fn wait_until_reached(&self) {
437        self.reached.wait();
438    }
439
440    fn release(&self) {
441        self.released.wait();
442    }
443}
444
445#[cfg(test)]
446mod tests {
447    use super::{Gate, RuntimeBuilder, ShutdownOutcome, TestPause};
448    use crate::Priority;
449    use futures_lite::future;
450    use std::num::NonZeroUsize;
451    use std::sync::atomic::Ordering;
452    use std::sync::{mpsc, Arc};
453    use std::time::Duration;
454
455    fn runtime() -> super::Runtime {
456        RuntimeBuilder::new(NonZeroUsize::new(1).expect("non-zero worker count"))
457            .build()
458            .expect("runtime builds")
459    }
460
461    fn wait_until_not_running(state: &super::RuntimeState) {
462        for _ in 0..10_000 {
463            if *state.gate.lock().expect("runtime lifecycle gate poisoned") != Gate::Running {
464                return;
465            }
466            std::thread::yield_now();
467        }
468        panic!("shutdown did not close the admission gate");
469    }
470
471    #[test]
472    fn graceful_shutdown_waits_for_admitted_task_before_initial_schedule() {
473        let runtime = runtime();
474        let state = Arc::clone(&runtime.state);
475        let pause = TestPause::new();
476        *state.admission_pause.lock().expect("test hook poisoned") = Some(Arc::clone(&pause));
477        let spawner = runtime.spawner();
478        let (ran_tx, ran_rx) = mpsc::channel();
479        let producer = std::thread::spawn(move || {
480            spawner
481                .spawn(Priority::Normal, async move {
482                    ran_tx.send(()).expect("test receiver remains alive");
483                })
484                .expect("admitted spawn succeeds")
485                .detach();
486        });
487        pause.wait_until_reached();
488
489        let (shutdown_tx, shutdown_rx) = mpsc::channel();
490        let shutdown = std::thread::spawn(move || {
491            shutdown_tx
492                .send(runtime.shutdown_graceful())
493                .expect("test receiver remains alive");
494        });
495        wait_until_not_running(&state);
496        assert!(shutdown_rx.try_recv().is_err());
497
498        pause.release();
499        producer.join().expect("producer does not panic");
500        shutdown_rx
501            .recv_timeout(Duration::from_secs(1))
502            .expect("graceful shutdown finishes after scheduling")
503            .expect("graceful shutdown succeeds");
504        shutdown.join().expect("shutdown thread does not panic");
505        ran_rx
506            .recv_timeout(Duration::from_secs(1))
507            .expect("admitted task runs before graceful shutdown returns");
508        assert_eq!(state.accepted_tasks.load(Ordering::Acquire), 0);
509    }
510
511    #[test]
512    fn forced_shutdown_cancels_admitted_task_before_initial_schedule() {
513        let runtime = runtime();
514        let state = Arc::clone(&runtime.state);
515        let pause = TestPause::new();
516        *state.admission_pause.lock().expect("test hook poisoned") = Some(Arc::clone(&pause));
517        let spawner = runtime.spawner();
518        let (task_tx, task_rx) = mpsc::channel();
519        let producer = std::thread::spawn(move || {
520            let task = spawner
521                .spawn(Priority::Normal, async { 7_u8 })
522                .expect("spawn was admitted before shutdown");
523            task_tx.send(task).expect("test receiver remains alive");
524        });
525        pause.wait_until_reached();
526
527        runtime.shutdown_now().expect("forced shutdown succeeds");
528        pause.release();
529        producer.join().expect("producer does not panic");
530        let task = task_rx
531            .recv_timeout(Duration::from_secs(1))
532            .expect("spawn returns its cancelled task handle");
533        assert_eq!(future::block_on(task.fallible()), None);
534        assert_eq!(state.accepted_tasks.load(Ordering::Acquire), 0);
535    }
536
537    #[test]
538    fn timed_shutdown_cancels_admitted_task_before_initial_schedule() {
539        let runtime = runtime();
540        let state = Arc::clone(&runtime.state);
541        let pause = TestPause::new();
542        *state.admission_pause.lock().expect("test hook poisoned") = Some(Arc::clone(&pause));
543        let spawner = runtime.spawner();
544        let (task_tx, task_rx) = mpsc::channel();
545        let producer = std::thread::spawn(move || {
546            let task = spawner
547                .spawn(Priority::Normal, async { 9_u8 })
548                .expect("spawn was admitted before shutdown");
549            task_tx.send(task).expect("test receiver remains alive");
550        });
551        pause.wait_until_reached();
552
553        assert!(matches!(
554            runtime
555                .shutdown_timeout(Duration::ZERO)
556                .expect("timed shutdown succeeds"),
557            ShutdownOutcome::TimedOut { remaining_tasks: 1 }
558        ));
559        pause.release();
560        producer.join().expect("producer does not panic");
561        let task = task_rx
562            .recv_timeout(Duration::from_secs(1))
563            .expect("spawn returns its cancelled task handle");
564        assert_eq!(future::block_on(task.fallible()), None);
565        assert_eq!(state.accepted_tasks.load(Ordering::Acquire), 0);
566    }
567
568    #[test]
569    fn waiter_observes_zero_when_last_completion_precedes_notification() {
570        let runtime = runtime();
571        let state = Arc::clone(&runtime.state);
572        let pause = TestPause::new();
573        *state
574            .last_completion_pause
575            .lock()
576            .expect("test hook poisoned") = Some(Arc::clone(&pause));
577        let task = runtime
578            .spawn(Priority::Normal, async {})
579            .expect("spawn succeeds");
580        pause.wait_until_reached();
581        assert_eq!(state.accepted_tasks.load(Ordering::Acquire), 0);
582
583        let (wait_tx, wait_rx) = mpsc::channel();
584        let wait_state = Arc::clone(&state);
585        let waiter = std::thread::spawn(move || {
586            wait_state.wait_for_drain();
587            wait_tx.send(()).expect("test receiver remains alive");
588        });
589        wait_rx
590            .recv_timeout(Duration::from_secs(1))
591            .expect("waiter sees zero without needing the pending notification");
592
593        pause.release();
594        future::block_on(task);
595        waiter.join().expect("waiter does not panic");
596        runtime.shutdown_graceful().expect("shutdown succeeds");
597    }
598}