async_local_executor/
executor.rs1use crate::tls::{executor, try_executor};
2use async_local_channel::oneshot;
3use crossbeam_queue::SegQueue;
4use slotmap::new_key_type;
5use slotmap::{Key, SlotMap};
6use std::collections::VecDeque;
7use std::fmt::Debug;
8use std::pin::{Pin, pin};
9use std::sync::Arc;
10use std::task::Wake;
11use std::task::Waker;
12use std::task::{Context, Poll};
13
14new_key_type! { pub struct TaskId; }
15
16type WakeFn = Arc<dyn Fn() + Send + Sync>;
17
18struct WakerData {
19 task_id: TaskId,
20 queue: Arc<SegQueue<TaskId>>,
21 wake_fn: WakeFn,
22}
23
24impl Wake for WakerData {
25 fn wake(self: Arc<Self>) {
26 self.wake_by_ref();
27 }
28
29 fn wake_by_ref(self: &Arc<Self>) {
30 if let Some(mut executor) = try_executor() {
31 return executor.local_queue.push_back(self.task_id);
32 }
33 self.queue.push(self.task_id);
34 (self.wake_fn)();
35 }
36}
37
38fn create_waker(
39 task_id: TaskId,
40 tx: Arc<SegQueue<TaskId>>,
41 wake_fn: Arc<dyn Fn() + Send + Sync>,
42) -> Waker {
43 Waker::from(Arc::new(WakerData {
44 task_id,
45 queue: tx,
46 wake_fn,
47 }))
48}
49
50type LocalFuture = Pin<Box<dyn Future<Output = ()>>>;
51
52struct Task {
53 future: LocalFuture,
54 waker: Waker,
55}
56
57#[must_use]
58pub struct Ticker {
59 task_id: TaskId,
60 task: Task,
61}
62
63impl Ticker {
64 #[inline]
65 pub fn tick(mut self) {
66 let mut context = Context::from_waker(&self.task.waker);
67 if pin!(&mut self.task.future).poll(&mut context).is_ready() {
68 executor().task_completed(self.task_id);
69 } else {
70 executor().return_poller(self);
71 }
72 }
73}
74
75pub struct JoinHandle<T> {
76 task_id: TaskId,
77 result: oneshot::Receiver<T>,
78 cancel: bool,
79}
80
81impl<T> Debug for JoinHandle<T> {
82 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
83 f.debug_tuple("JoinHandle")
84 .field(&self.task_id.data())
85 .finish()
86 }
87}
88
89impl<T> JoinHandle<T> {
90 const fn new(task_id: TaskId, result: oneshot::Receiver<T>) -> Self {
91 Self {
92 task_id,
93 result,
94 cancel: false,
95 }
96 }
97
98 #[inline]
99 pub fn cancel(mut self) {
100 self.cancel = true;
101 }
102
103 #[inline]
104 #[must_use]
105 pub fn result(&self) -> Option<T> {
106 self.result.try_recv()
107 }
108}
109
110impl<T: 'static> Future for JoinHandle<T> {
111 type Output = T;
112
113 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
114 match pin!(self.result.recv()).poll(cx) {
115 Poll::Ready(value) => Poll::Ready(value.unwrap()),
116 Poll::Pending => Poll::Pending,
117 }
118 }
119}
120
121impl<T> Drop for JoinHandle<T> {
122 #[inline]
123 fn drop(&mut self) {
124 if self.cancel
125 && let Some(mut executor) = try_executor()
126 {
127 executor.task_completed(self.task_id);
128 }
129 }
130}
131
132pub struct Executor {
133 tasks: SlotMap<TaskId, Option<Task>>,
134 wake_fn: WakeFn,
135 local_queue: VecDeque<TaskId>,
136 queue: Arc<SegQueue<TaskId>>,
137}
138
139impl Debug for Executor {
140 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
141 f.debug_tuple("Executor").field(&self.tasks.len()).finish()
142 }
143}
144
145impl Executor {
146 #[inline]
147 pub fn new<F: Fn() + Send + Sync + 'static>(f: F) -> Self {
148 let tasks = SlotMap::with_key();
149 let wake_fn = Arc::new(f);
150 let queue = Arc::new(SegQueue::new());
151 let local_queue = VecDeque::new();
152 Self {
153 tasks,
154 wake_fn,
155 local_queue,
156 queue,
157 }
158 }
159
160 pub(crate) fn spawn_local<F>(&mut self, future: F) -> JoinHandle<F::Output>
161 where
162 F: Future + 'static,
163 F::Output: 'static,
164 {
165 let (tx, rx) = oneshot::channel();
166 let future = async {
167 let res = future.await;
168 let _ = tx.send(res);
169 };
170 let future = Box::pin(future);
171 let task_id = self.tasks.insert_with_key(|id| {
172 let waker = create_waker(id, self.queue.clone(), self.wake_fn.clone());
173 Some(Task { future, waker })
174 });
175 self.local_queue.push_back(task_id);
176 JoinHandle::new(task_id, rx.activate())
177 }
178
179 pub(crate) fn exit(&mut self) {
180 self.tasks.clear();
181 }
182
183 fn next_task_id(&mut self) -> Option<TaskId> {
184 if let Some(task_id) = self.local_queue.pop_front() {
185 return Some(task_id);
186 }
187 while let Some(task_id) = self.queue.pop() {
188 self.local_queue.push_back(task_id);
189 }
190 self.local_queue.pop_front()
191 }
192
193 pub(crate) fn ticker(&mut self) -> Option<Ticker> {
194 let task_id = self.next_task_id()?;
195 let task = self.tasks.get_mut(task_id)?.take()?;
196 Some(Ticker { task_id, task })
197 }
198
199 fn task_completed(&mut self, task_id: TaskId) {
200 self.tasks.remove(task_id);
201 }
202
203 fn return_poller(&mut self, ticker: Ticker) {
204 let Some(task) = self.tasks.get_mut(ticker.task_id) else {
205 return;
206 };
207 *task = Some(ticker.task);
208 }
209}