1use crate::{Instant, Priority, RunnableMeta, Scheduler, SessionId, Timer};
2use async_task::Runnable;
3use std::{
4 any::Any,
5 future::Future,
6 marker::PhantomData,
7 mem::ManuallyDrop,
8 panic::Location,
9 pin::Pin,
10 rc::Rc,
11 sync::Arc,
12 task::{Context, Poll, Waker},
13 thread::{self, ThreadId},
14 time::Duration,
15};
16
17#[derive(Clone)]
22pub struct LocalExecutor {
23 session_id: SessionId,
24 scheduler: Arc<dyn Scheduler>,
25 dispatch: Arc<dyn Fn(Runnable<RunnableMeta>) + Send + Sync>,
29 not_send: PhantomData<Rc<()>>,
30}
31
32impl LocalExecutor {
33 pub fn new(
42 session_id: SessionId,
43 scheduler: Arc<dyn Scheduler>,
44 dispatch: impl Fn(Runnable<RunnableMeta>) + Send + Sync + 'static,
45 ) -> Self {
46 Self {
47 session_id,
48 scheduler,
49 dispatch: Arc::new(dispatch),
50 not_send: PhantomData,
51 }
52 }
53
54 pub fn session_id(&self) -> SessionId {
55 self.session_id
56 }
57
58 pub fn scheduler(&self) -> &Arc<dyn Scheduler> {
59 &self.scheduler
60 }
61
62 pub fn is_test(&self) -> bool {
64 self.scheduler.as_test().is_some()
65 }
66
67 #[track_caller]
68 pub fn spawn<F>(&self, future: F) -> Task<F::Output>
69 where
70 F: Future + 'static,
71 F::Output: 'static,
72 {
73 let schedule = self.schedule();
74 let location = Location::caller();
75 let (runnable, task) = spawn_local_with_source_location(
76 future,
77 schedule,
78 RunnableMeta {
79 location,
80 spawned: crate::SpawnTime(Instant::now()),
81 },
82 );
83 runnable.schedule();
84 Task(TaskState::Spawned(task))
85 }
86
87 #[track_caller]
92 pub fn spawn_with_dispatch<F>(
93 &self,
94 future: F,
95 dispatch: impl Fn(Runnable<RunnableMeta>) + Send + Sync + 'static,
96 ) -> Task<F::Output>
97 where
98 F: Future + 'static,
99 F::Output: 'static,
100 {
101 let location = Location::caller();
102 let (runnable, task) = spawn_local_with_source_location(
103 future,
104 dispatch,
105 RunnableMeta {
106 location,
107 spawned: crate::SpawnTime(Instant::now()),
108 },
109 );
110 runnable.schedule();
111 Task(TaskState::Spawned(task))
112 }
113
114 #[cfg(not(target_family = "wasm"))]
115 pub fn block_on<Fut: Future>(&self, future: Fut) -> Fut::Output {
116 use std::cell::Cell;
117
118 let output = Cell::new(None);
119 let future = async {
120 output.set(Some(future.await));
121 };
122 let mut future = std::pin::pin!(future);
123
124 self.scheduler
125 .block(Some(self.session_id), future.as_mut(), None);
126
127 output.take().expect("block_on future did not complete")
128 }
129
130 #[cfg(not(target_family = "wasm"))]
133 pub fn block_with_timeout<Fut: Future>(
134 &self,
135 timeout: Duration,
136 future: Fut,
137 ) -> Result<Fut::Output, impl Future<Output = Fut::Output> + use<Fut>> {
138 use std::cell::Cell;
139
140 let output = Cell::new(None);
141 let mut future = Box::pin(future);
142
143 {
144 let future_ref = &mut future;
145 let wrapper = async {
146 output.set(Some(future_ref.await));
147 };
148 let mut wrapper = std::pin::pin!(wrapper);
149
150 self.scheduler
151 .block(Some(self.session_id), wrapper.as_mut(), Some(timeout));
152 }
153
154 match output.take() {
155 Some(value) => Ok(value),
156 None => Err(future),
157 }
158 }
159
160 #[track_caller]
161 pub fn timer(&self, duration: Duration) -> Timer {
162 self.scheduler.timer(duration)
163 }
164
165 pub fn now(&self) -> Instant {
166 self.scheduler.clock().now()
167 }
168
169 #[track_caller]
178 pub fn spawn_dedicated<F, Fut>(&self, f: F) -> Task<Fut::Output>
179 where
180 F: FnOnce(LocalExecutor) -> Fut + Send + 'static,
181 Fut: Future + 'static,
182 Fut::Output: Send + Sync + 'static,
183 {
184 self.scheduler
185 .clone()
186 .spawn_dedicated(box_dedicated(f))
187 .downcast::<Fut::Output>()
188 }
189
190 fn schedule(&self) -> impl Fn(Runnable<RunnableMeta>) + Send + Sync + 'static {
191 let dispatch = self.dispatch.clone();
192 move |runnable| dispatch(runnable)
193 }
194}
195
196fn box_dedicated<F, Fut>(
201 f: F,
202) -> Box<
203 dyn FnOnce(LocalExecutor) -> Pin<Box<dyn Future<Output = Box<dyn Any + Send + Sync>> + 'static>>
204 + Send
205 + 'static,
206>
207where
208 F: FnOnce(LocalExecutor) -> Fut + Send + 'static,
209 Fut: Future + 'static,
210 Fut::Output: Send + Sync + 'static,
211{
212 Box::new(move |executor| {
213 Box::pin(async move { Box::new(f(executor).await) as Box<dyn Any + Send + Sync> })
214 })
215}
216
217#[derive(Clone)]
218pub struct BackgroundExecutor {
219 scheduler: Arc<dyn Scheduler>,
220}
221
222impl BackgroundExecutor {
223 pub fn new(scheduler: Arc<dyn Scheduler>) -> Self {
224 Self { scheduler }
225 }
226
227 #[track_caller]
228 pub fn spawn<F>(&self, future: F) -> Task<F::Output>
229 where
230 F: Future + Send + 'static,
231 F::Output: Send + 'static,
232 {
233 self.spawn_with_priority(Priority::default(), future)
234 }
235
236 #[track_caller]
237 pub fn spawn_with_priority<F>(&self, priority: Priority, future: F) -> Task<F::Output>
238 where
239 F: Future + Send + 'static,
240 F::Output: Send + 'static,
241 {
242 let schedule = self.schedule_with_priority(priority);
243 let location = Location::caller();
244 let (runnable, task) = async_task::Builder::new()
245 .metadata(RunnableMeta {
246 location,
247 spawned: crate::SpawnTime(Instant::now()),
248 })
249 .spawn(move |_| future, schedule);
250 runnable.schedule();
251 Task(TaskState::Spawned(task))
252 }
253
254 #[track_caller]
256 pub fn spawn_realtime<F>(&self, future: F) -> Task<F::Output>
257 where
258 F: Future + Send + 'static,
259 F::Output: Send + 'static,
260 {
261 let location = Location::caller();
262 let (tx, rx) = flume::bounded::<async_task::Runnable<RunnableMeta>>(1);
263
264 self.scheduler.spawn_realtime(Box::new(move || {
265 while let Ok(runnable) = rx.recv() {
266 runnable.run();
267 }
268 }));
269
270 let (runnable, task) = async_task::Builder::new()
271 .metadata(RunnableMeta {
272 location,
273 spawned: crate::SpawnTime(Instant::now()),
274 })
275 .spawn(
276 move |_| future,
277 move |runnable| {
278 let _ = tx.send(runnable);
279 },
280 );
281 runnable.schedule();
282 Task(TaskState::Spawned(task))
283 }
284
285 #[track_caller]
286 pub fn timer(&self, duration: Duration) -> Timer {
287 self.scheduler.timer(duration)
288 }
289
290 pub fn now(&self) -> Instant {
291 self.scheduler.clock().now()
292 }
293
294 pub fn scheduler(&self) -> &Arc<dyn Scheduler> {
295 &self.scheduler
296 }
297
298 pub fn is_test(&self) -> bool {
300 self.scheduler.as_test().is_some()
301 }
302
303 #[track_caller]
312 pub fn spawn_dedicated<F, Fut>(&self, f: F) -> Task<Fut::Output>
313 where
314 F: FnOnce(LocalExecutor) -> Fut + Send + 'static,
315 Fut: Future + 'static,
316 Fut::Output: Send + Sync + 'static,
317 {
318 self.scheduler
319 .clone()
320 .spawn_dedicated(box_dedicated(f))
321 .downcast::<Fut::Output>()
322 }
323
324 fn schedule_with_priority(
325 &self,
326 priority: Priority,
327 ) -> impl Fn(Runnable<RunnableMeta>) + Send + Sync + 'static {
328 let scheduler = Arc::downgrade(&self.scheduler);
329 move |runnable| {
330 if let Some(scheduler) = scheduler.upgrade() {
331 scheduler.schedule_background_with_priority(runnable, priority);
332 }
333 }
334 }
335}
336
337pub struct DedicatedExecutor {
348 sender: flume::Sender<Runnable<RunnableMeta>>,
349 _session: Task<()>,
350}
351
352impl DedicatedExecutor {
353 #[track_caller]
357 pub fn new(executor: &BackgroundExecutor) -> Self {
358 let (sender, receiver) = flume::unbounded::<Runnable<RunnableMeta>>();
359 let session = executor.spawn_dedicated(move |_executor| async move {
360 while let Ok(runnable) = receiver.recv_async().await {
361 runnable.run();
362 }
363 });
364 Self {
365 sender,
366 _session: session,
367 }
368 }
369
370 #[track_caller]
376 pub fn spawn<F>(&self, future: F) -> Task<F::Output>
377 where
378 F: Future + Send + 'static,
379 F::Output: Send + 'static,
380 {
381 let sender = self.sender.clone();
382 let (runnable, task) = async_task::Builder::new()
383 .metadata(RunnableMeta::new_with_callers_location())
384 .spawn(
385 move |_| future,
386 move |runnable| {
387 let _ = sender.send(runnable);
388 },
389 );
390 runnable.schedule();
391 Task(TaskState::Spawned(task))
392 }
393}
394
395#[must_use]
402pub struct Task<T>(TaskState<T>);
403
404enum TaskState<T> {
405 Ready(Option<T>),
407
408 Spawned(async_task::Task<T, RunnableMeta>),
410
411 Downcast {
415 inner: Box<Task<Box<dyn Any + Send + Sync>>>,
416 marker: PhantomData<fn() -> T>,
417 },
418
419 Rendezvous(RendezvousReceiver<T>),
423}
424
425enum RendezvousState<T> {
427 Pending(Option<Waker>),
429 Delivered(Task<T>),
431 Cancelled,
433 Detached,
435 Taken,
437}
438
439pub(crate) struct RendezvousReceiver<T> {
440 shared: Arc<parking_lot::Mutex<RendezvousState<T>>>,
441}
442
443impl<T> RendezvousReceiver<T> {
444 fn poll_take(&self, cx: &mut Context) -> Poll<Task<T>> {
445 let mut state = self.shared.lock();
446 match std::mem::replace(&mut *state, RendezvousState::Taken) {
447 RendezvousState::Delivered(task) => Poll::Ready(task),
448 RendezvousState::Pending(_) => {
449 *state = RendezvousState::Pending(Some(cx.waker().clone()));
450 Poll::Pending
451 }
452 RendezvousState::Cancelled | RendezvousState::Detached | RendezvousState::Taken => {
453 unreachable!("a rendezvous task was polled after its receiver was consumed")
454 }
455 }
456 }
457
458 fn is_ready(&self) -> bool {
459 match &*self.shared.lock() {
460 RendezvousState::Delivered(task) => task.is_ready(),
461 _ => false,
462 }
463 }
464
465 fn detach(self) {
466 let mut state = self.shared.lock();
467 match std::mem::replace(&mut *state, RendezvousState::Detached) {
468 RendezvousState::Delivered(task) => {
469 *state = RendezvousState::Taken;
470 drop(state);
471 task.detach();
472 }
473 RendezvousState::Pending(_) => {}
475 RendezvousState::Cancelled | RendezvousState::Detached | RendezvousState::Taken => {
476 unreachable!("a rendezvous task was detached after its receiver was consumed")
477 }
478 }
479 }
480}
481
482impl<T> Drop for RendezvousReceiver<T> {
483 fn drop(&mut self) {
484 let mut state = self.shared.lock();
485 match &*state {
486 RendezvousState::Pending(_) => *state = RendezvousState::Cancelled,
487 RendezvousState::Delivered(_) => {
488 let delivered = std::mem::replace(&mut *state, RendezvousState::Cancelled);
489 drop(state);
490 drop(delivered);
491 }
492 _ => {}
496 }
497 }
498}
499
500pub(crate) struct RendezvousSender<T> {
501 shared: Arc<parking_lot::Mutex<RendezvousState<T>>>,
502}
503
504impl<T> RendezvousSender<T> {
505 pub(crate) fn deliver(self, task: Task<T>) {
508 let mut state = self.shared.lock();
509 match std::mem::replace(&mut *state, RendezvousState::Delivered(task)) {
510 RendezvousState::Pending(waker) => {
511 drop(state);
512 if let Some(waker) = waker {
513 waker.wake();
514 }
515 }
516 RendezvousState::Cancelled => {
517 let delivered = std::mem::replace(&mut *state, RendezvousState::Cancelled);
518 drop(state);
519 drop(delivered);
520 }
521 RendezvousState::Detached => {
522 let RendezvousState::Delivered(task) =
523 std::mem::replace(&mut *state, RendezvousState::Detached)
524 else {
525 unreachable!("the delivered task was just stored");
526 };
527 drop(state);
528 task.detach();
529 }
530 RendezvousState::Delivered(_) | RendezvousState::Taken => {
531 unreachable!("a rendezvous task was delivered twice")
532 }
533 }
534 }
535}
536
537impl<T> Task<T> {
538 pub fn ready(val: T) -> Self {
540 Task(TaskState::Ready(Some(val)))
541 }
542
543 pub(crate) fn rendezvous() -> (Self, RendezvousSender<T>) {
548 let shared = Arc::new(parking_lot::Mutex::new(RendezvousState::Pending(None)));
549 (
550 Task(TaskState::Rendezvous(RendezvousReceiver {
551 shared: shared.clone(),
552 })),
553 RendezvousSender { shared },
554 )
555 }
556
557 pub fn from_async_task(task: async_task::Task<T, RunnableMeta>) -> Self {
559 Task(TaskState::Spawned(task))
560 }
561
562 pub fn is_ready(&self) -> bool {
563 match &self.0 {
564 TaskState::Ready(_) => true,
565 TaskState::Spawned(task) => task.is_finished(),
566 TaskState::Downcast { inner, .. } => inner.is_ready(),
567 TaskState::Rendezvous(receiver) => receiver.is_ready(),
568 }
569 }
570
571 pub fn detach(self) {
573 match self {
574 Task(TaskState::Ready(_)) => {}
575 Task(TaskState::Spawned(task)) => task.detach(),
576 Task(TaskState::Downcast { inner, .. }) => inner.detach(),
577 Task(TaskState::Rendezvous(receiver)) => receiver.detach(),
578 }
579 }
580
581 pub fn fallible(self) -> FallibleTask<T> {
583 FallibleTask(match self.0 {
584 TaskState::Ready(val) => FallibleTaskState::Ready(val),
585 TaskState::Spawned(task) => FallibleTaskState::Spawned(task.fallible()),
586 TaskState::Downcast { inner, .. } => FallibleTaskState::Downcast {
587 inner: Box::new(inner.fallible()),
588 marker: PhantomData,
589 },
590 TaskState::Rendezvous(receiver) => FallibleTaskState::Rendezvous(receiver),
591 })
592 }
593}
594
595impl Task<Box<dyn Any + Send + Sync>> {
596 pub fn downcast<T: Send + Sync + 'static>(self) -> Task<T> {
605 Task(TaskState::Downcast {
606 inner: Box::new(self),
607 marker: PhantomData,
608 })
609 }
610}
611
612impl<T> std::fmt::Debug for Task<T> {
613 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
614 match &self.0 {
615 TaskState::Ready(_) => f.debug_tuple("Task::Ready").finish(),
616 TaskState::Spawned(task) => f.debug_tuple("Task::Spawned").field(task).finish(),
617 TaskState::Downcast { inner, .. } => {
618 f.debug_tuple("Task::Downcast").field(inner).finish()
619 }
620 TaskState::Rendezvous(_) => f.debug_tuple("Task::Rendezvous").finish(),
621 }
622 }
623}
624
625#[must_use]
627pub struct FallibleTask<T>(FallibleTaskState<T>);
628
629enum FallibleTaskState<T> {
630 Ready(Option<T>),
632
633 Spawned(async_task::FallibleTask<T, RunnableMeta>),
635
636 Downcast {
638 inner: Box<FallibleTask<Box<dyn Any + Send + Sync>>>,
639 marker: PhantomData<fn() -> T>,
640 },
641
642 Rendezvous(RendezvousReceiver<T>),
644}
645
646impl<T> FallibleTask<T> {
647 pub fn ready(val: T) -> Self {
649 FallibleTask(FallibleTaskState::Ready(Some(val)))
650 }
651
652 pub fn detach(self) {
654 match self.0 {
655 FallibleTaskState::Ready(_) => {}
656 FallibleTaskState::Spawned(task) => task.detach(),
657 FallibleTaskState::Downcast { inner, .. } => inner.detach(),
658 FallibleTaskState::Rendezvous(receiver) => receiver.detach(),
659 }
660 }
661}
662
663impl<T: 'static> Future for FallibleTask<T> {
664 type Output = Option<T>;
665
666 fn poll(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
667 let this = unsafe { self.get_unchecked_mut() };
668 loop {
669 match &mut this.0 {
670 FallibleTaskState::Ready(val) => return Poll::Ready(val.take()),
671 FallibleTaskState::Spawned(task) => return Pin::new(task).poll(cx),
672 FallibleTaskState::Downcast { inner, .. } => {
673 return match Pin::new(inner.as_mut()).poll(cx) {
674 Poll::Ready(Some(boxed_any)) => Poll::Ready(Some(
675 *boxed_any
676 .downcast::<T>()
677 .expect("FallibleTask::poll: downcast type mismatch"),
678 )),
679 Poll::Ready(None) => Poll::Ready(None),
680 Poll::Pending => Poll::Pending,
681 };
682 }
683 FallibleTaskState::Rendezvous(receiver) => match receiver.poll_take(cx) {
684 Poll::Ready(task) => {
685 this.0 = task.fallible().0;
686 continue;
687 }
688 Poll::Pending => return Poll::Pending,
689 },
690 }
691 }
692 }
693}
694
695impl<T> std::fmt::Debug for FallibleTask<T> {
696 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
697 match &self.0 {
698 FallibleTaskState::Ready(_) => f.debug_tuple("FallibleTask::Ready").finish(),
699 FallibleTaskState::Spawned(task) => {
700 f.debug_tuple("FallibleTask::Spawned").field(task).finish()
701 }
702 FallibleTaskState::Downcast { inner, .. } => f
703 .debug_tuple("FallibleTask::Downcast")
704 .field(inner)
705 .finish(),
706 FallibleTaskState::Rendezvous(_) => f.debug_tuple("FallibleTask::Rendezvous").finish(),
707 }
708 }
709}
710
711impl<T: 'static> Future for Task<T> {
712 type Output = T;
713
714 fn poll(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
715 let this = unsafe { self.get_unchecked_mut() };
716 loop {
717 match &mut this.0 {
718 TaskState::Ready(val) => return Poll::Ready(val.take().unwrap()),
719 TaskState::Spawned(task) => return Pin::new(task).poll(cx),
720 TaskState::Downcast { inner, .. } => {
721 return match Pin::new(inner.as_mut()).poll(cx) {
722 Poll::Ready(boxed_any) => Poll::Ready(
723 *boxed_any
724 .downcast::<T>()
725 .expect("Task::poll: downcast type mismatch"),
726 ),
727 Poll::Pending => Poll::Pending,
728 };
729 }
730 TaskState::Rendezvous(receiver) => match receiver.poll_take(cx) {
731 Poll::Ready(task) => {
732 this.0 = task.0;
733 continue;
734 }
735 Poll::Pending => return Poll::Pending,
736 },
737 }
738 }
739 }
740}
741
742#[track_caller]
744fn spawn_local_with_source_location<Fut, S>(
745 future: Fut,
746 schedule: S,
747 metadata: RunnableMeta,
748) -> (
749 async_task::Runnable<RunnableMeta>,
750 async_task::Task<Fut::Output, RunnableMeta>,
751)
752where
753 Fut: Future + 'static,
754 Fut::Output: 'static,
755 S: async_task::Schedule<RunnableMeta> + Send + Sync + 'static,
756{
757 #[inline]
758 fn thread_id() -> ThreadId {
759 std::thread_local! {
760 static ID: ThreadId = thread::current().id();
761 }
762 ID.try_with(|id| *id)
763 .unwrap_or_else(|_| thread::current().id())
764 }
765
766 struct Checked<F> {
767 id: ThreadId,
768 inner: ManuallyDrop<F>,
769 location: &'static Location<'static>,
770 }
771
772 impl<F> Drop for Checked<F> {
773 fn drop(&mut self) {
774 assert_eq!(
775 self.id,
776 thread_id(),
777 "local task dropped by a thread that didn't spawn it. Task spawned at {}",
778 self.location
779 );
780 unsafe { ManuallyDrop::drop(&mut self.inner) };
784 }
785 }
786
787 impl<F: Future> Future for Checked<F> {
788 type Output = F::Output;
789
790 fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
791 let this = unsafe { self.get_unchecked_mut() };
795 assert!(
796 this.id == thread_id(),
797 "local task polled by a thread that didn't spawn it. Task spawned at {}",
798 this.location
799 );
800 unsafe { Pin::new_unchecked(&mut *this.inner).poll(cx) }
805 }
806 }
807
808 let location = metadata.location;
809
810 let future = move |_| Checked {
811 id: thread_id(),
812 inner: ManuallyDrop::new(future),
813 location,
814 };
815
816 let builder = async_task::Builder::new().metadata(metadata);
817 unsafe { builder.spawn_unchecked(future, schedule) }
821}