1use std::collections::VecDeque;
17use std::future::Future;
18use std::pin::Pin;
19use std::sync::atomic::{AtomicBool, Ordering};
20use std::sync::{Arc, Condvar, Mutex};
21use std::task::{Context, Waker};
22use std::thread::JoinHandle;
23
24struct Task {
29 future: Mutex<Option<Pin<Box<dyn Future<Output = ()> + Send>>>>,
30 ready: Arc<ReadyQueue>,
31}
32
33impl std::task::Wake for Task {
34 fn wake(self: Arc<Self>) {
35 self.ready.clone().push(self);
36 }
37 fn wake_by_ref(self: &Arc<Self>) {
38 self.ready.clone().push(self.clone());
39 }
40}
41
42struct ReadyQueue {
43 queue: Mutex<VecDeque<Arc<Task>>>,
44 signal: Condvar,
45 shutdown: AtomicBool,
46}
47
48impl ReadyQueue {
49 fn push(&self, task: Arc<Task>) {
50 self.queue.lock().unwrap().push_back(task);
51 self.signal.notify_one();
52 }
53
54 fn pop(&self) -> Option<Arc<Task>> {
56 let mut q = self.queue.lock().unwrap();
57 loop {
58 if let Some(task) = q.pop_front() {
59 return Some(task);
60 }
61 if self.shutdown.load(Ordering::Acquire) {
62 return None;
63 }
64 q = self.signal.wait(q).unwrap();
65 }
66 }
67}
68
69pub struct TaskPool {
72 ready: Arc<ReadyQueue>,
73 workers: Vec<JoinHandle<()>>,
74}
75
76impl TaskPool {
77 pub fn new(n_workers: usize) -> Self {
79 let n = n_workers.max(1);
80 let ready = Arc::new(ReadyQueue {
81 queue: Mutex::new(VecDeque::new()),
82 signal: Condvar::new(),
83 shutdown: AtomicBool::new(false),
84 });
85 let workers = (0..n)
86 .map(|_| {
87 let ready = Arc::clone(&ready);
88 std::thread::spawn(move || worker_loop(ready))
89 })
90 .collect();
91 Self { ready, workers }
92 }
93
94 pub fn worker_count(&self) -> usize {
96 self.workers.len()
97 }
98
99 pub fn spawn(&self, future: impl Future<Output = ()> + Send + 'static) {
102 let task = Arc::new(Task {
103 future: Mutex::new(Some(Box::pin(future))),
104 ready: Arc::clone(&self.ready),
105 });
106 self.ready.push(task);
107 }
108
109 pub fn shutdown(self) {
112 self.ready.shutdown.store(true, Ordering::Release);
113 self.ready.signal.notify_all();
114 for w in self.workers {
115 w.join().ok();
116 }
117 }
118}
119
120fn worker_loop(ready: Arc<ReadyQueue>) {
121 while let Some(task) = ready.pop() {
122 let mut guard = task.future.lock().unwrap();
123 if let Some(fut) = guard.as_mut() {
124 let waker = Waker::from(Arc::clone(&task));
125 let mut cx = Context::from_waker(&waker);
126 if fut.as_mut().poll(&mut cx).is_ready() {
127 *guard = None;
130 }
131 }
132 }
133}
134
135#[cfg(test)]
136mod tests {
137 use super::*;
138 use std::sync::atomic::AtomicU64;
139
140 #[test]
141 fn runs_many_tasks_on_few_threads_with_yields() {
142 let pool = TaskPool::new(2);
145 let done = Arc::new(AtomicU64::new(0));
146 let n = 5_000u64;
147 for _ in 0..n {
148 let done = Arc::clone(&done);
149 pool.spawn(async move {
150 YieldOnce::default().await;
151 done.fetch_add(1, Ordering::AcqRel);
152 });
153 }
154 let start = std::time::Instant::now();
157 while done.load(Ordering::Acquire) < n {
158 if start.elapsed() > std::time::Duration::from_secs(10) {
159 panic!("only {} of {n} tasks finished", done.load(Ordering::Acquire));
160 }
161 std::hint::spin_loop();
162 }
163 assert_eq!(pool.worker_count(), 2);
164 pool.shutdown();
165 }
166
167 #[derive(Default)]
168 struct YieldOnce {
169 yielded: bool,
170 }
171 impl Future for YieldOnce {
172 type Output = ();
173 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> std::task::Poll<()> {
174 if self.yielded {
175 std::task::Poll::Ready(())
176 } else {
177 self.yielded = true;
178 cx.waker().wake_by_ref();
179 std::task::Poll::Pending
180 }
181 }
182 }
183}