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 pub fn new(worker_threads: NonZeroUsize) -> Self {
23 Self {
24 worker_threads,
25 weights: PriorityWeights::default(),
26 }
27 }
28 #[must_use]
30 pub fn priority_weights(mut self, weights: PriorityWeights) -> Self {
31 self.weights = weights;
32 self
33 }
34 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 pub fn spawner(&self) -> Spawner {
88 Spawner {
89 state: Arc::downgrade(&self.state),
90 }
91 }
92 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 #[cfg(feature = "stats")]
112 pub fn stats(&self) -> crate::RuntimeStats {
113 self.state.scheduler.stats()
114 }
115 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 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 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 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
260pub(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 let completion = CompletionGuard {
290 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 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 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 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}