Skip to main content

matrix_sdk_common/
executor.rs

1// Copyright 2021 The Matrix.org Foundation C.I.C.
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15//! Abstraction over an executor so we can spawn tasks under Wasm the same way
16//! we do usually.
17//!
18//! On non Wasm platforms, this re-exports parts of tokio directly.  For Wasm,
19//! we provide a single-threaded solution that matches the interface that tokio
20//! provides as a drop in replacement.
21
22use std::{
23    future::Future,
24    pin::Pin,
25    task::{Context, Poll},
26};
27
28#[cfg(not(target_family = "wasm"))]
29mod sys {
30    pub use tokio::{
31        runtime::{Handle, Runtime},
32        task::{AbortHandle, JoinError, JoinHandle, spawn},
33    };
34}
35
36#[cfg(target_family = "wasm")]
37mod sys {
38    use std::{
39        future::Future,
40        pin::Pin,
41        task::{Context, Poll},
42    };
43
44    pub use futures_util::future::AbortHandle;
45    use futures_util::{
46        FutureExt,
47        future::{Abortable, RemoteHandle},
48    };
49
50    /// A Wasm specific version of `tokio::task::JoinError` designed to work in
51    /// the single-threaded environment available in Wasm environments.
52    #[derive(Debug)]
53    pub enum JoinError {
54        Cancelled,
55        Panic,
56    }
57
58    impl JoinError {
59        /// Returns true if the error was caused by the task being cancelled.
60        ///
61        /// See [the module level docs] for more information on cancellation.
62        ///
63        /// [the module level docs]: crate::task#cancellation
64        pub fn is_cancelled(&self) -> bool {
65            matches!(self, JoinError::Cancelled)
66        }
67    }
68
69    impl std::fmt::Display for JoinError {
70        fn fmt(&self, fmt: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
71            match &self {
72                JoinError::Cancelled => write!(fmt, "task was cancelled"),
73                JoinError::Panic => write!(fmt, "task panicked"),
74            }
75        }
76    }
77
78    /// A Wasm specific version of `tokio::task::JoinHandle` that holds handles
79    /// to locally executing futures.
80    #[derive(Debug)]
81    pub struct JoinHandle<T> {
82        remote_handle: Option<RemoteHandle<T>>,
83        abort_handle: AbortHandle,
84    }
85
86    impl<T> JoinHandle<T> {
87        /// Aborts the spawned future, preventing it from being polled again.
88        pub fn abort(&self) {
89            self.abort_handle.abort();
90        }
91
92        /// Returns the handle to the `AbortHandle` that can be used to abort
93        /// the spawned future.
94        pub fn abort_handle(&self) -> AbortHandle {
95            self.abort_handle.clone()
96        }
97
98        /// Returns true if the spawned future has been aborted.
99        pub fn is_finished(&self) -> bool {
100            self.abort_handle.is_aborted()
101        }
102    }
103
104    impl<T> Drop for JoinHandle<T> {
105        fn drop(&mut self) {
106            // don't abort the spawned future
107            if let Some(h) = self.remote_handle.take() {
108                h.forget();
109            }
110        }
111    }
112
113    impl<T: 'static> Future for JoinHandle<T> {
114        type Output = Result<T, JoinError>;
115
116        fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
117            if self.abort_handle.is_aborted() {
118                // The future has been aborted. It is not possible to poll it
119                // again.
120                Poll::Ready(Err(JoinError::Cancelled))
121            } else if let Some(handle) = self.remote_handle.as_mut() {
122                Pin::new(handle).poll(cx).map(Ok)
123            } else {
124                Poll::Ready(Err(JoinError::Panic))
125            }
126        }
127    }
128
129    /// A Wasm specific version of `tokio::task::spawn` that utilizes
130    /// wasm_bindgen_futures to spawn futures on the local executor.
131    pub fn spawn<F, T>(future: F) -> JoinHandle<T>
132    where
133        F: Future<Output = T> + 'static,
134    {
135        let (future, remote_handle) = future.remote_handle();
136        let (abort_handle, abort_registration) = AbortHandle::new_pair();
137        let future = Abortable::new(future, abort_registration);
138
139        wasm_bindgen_futures::spawn_local(async {
140            // Poll the future, and ignore the result (either it's `Ok(())`, or
141            // it's `Err(Aborted)`).
142            let _ = future.await;
143        });
144
145        JoinHandle { remote_handle: Some(remote_handle), abort_handle }
146    }
147}
148
149pub use sys::*;
150
151/// A type ensuring a task is aborted on drop.
152#[derive(Debug)]
153pub struct AbortOnDrop<T>(JoinHandle<T>);
154
155impl<T> AbortOnDrop<T> {
156    pub fn new(join_handle: JoinHandle<T>) -> Self {
157        Self(join_handle)
158    }
159}
160
161impl<T> Drop for AbortOnDrop<T> {
162    fn drop(&mut self) {
163        self.0.abort();
164    }
165}
166
167impl<T: 'static> Future for AbortOnDrop<T> {
168    type Output = Result<T, JoinError>;
169
170    fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
171        Pin::new(&mut self.0).poll(context)
172    }
173}
174
175/// Trait to create an [`AbortOnDrop`] from a [`JoinHandle`].
176pub trait JoinHandleExt<T> {
177    fn abort_on_drop(self) -> AbortOnDrop<T>;
178}
179
180impl<T> JoinHandleExt<T> for JoinHandle<T> {
181    fn abort_on_drop(self) -> AbortOnDrop<T> {
182        AbortOnDrop::new(self)
183    }
184}
185
186#[cfg(test)]
187mod tests {
188    use assert_matches::assert_matches;
189    use matrix_sdk_test_macros::async_test;
190
191    use super::spawn;
192
193    #[async_test]
194    async fn test_spawn() {
195        let future = async { 42 };
196        let join_handle = spawn(future);
197
198        assert_matches!(join_handle.await, Ok(42));
199    }
200
201    #[async_test]
202    async fn test_abort() {
203        let future = async { 42 };
204        let join_handle = spawn(future);
205
206        join_handle.abort();
207
208        assert!(join_handle.await.is_err());
209    }
210}