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