Skip to main content

async_local_executor/
executor.rs

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