matrix_sdk_common/
executor.rs1use 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 #[derive(Debug)]
53 pub enum JoinError {
54 Cancelled,
55 Panic,
56 }
57
58 impl JoinError {
59 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 #[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 pub fn abort(&self) {
89 self.abort_handle.abort();
90 }
91
92 pub fn abort_handle(&self) -> AbortHandle {
95 self.abort_handle.clone()
96 }
97
98 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 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 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 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 let _ = future.await;
143 });
144
145 JoinHandle { remote_handle: Some(remote_handle), abort_handle }
146 }
147}
148
149pub use sys::*;
150
151#[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
175pub 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}