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