Skip to main content

async_local_executor/
executor.rs

1use 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}