matrix_sdk_common/
timeout.rs1use std::{error::Error, fmt, time::Duration};
16
17use futures_core::Future;
18#[cfg(target_family = "wasm")]
19use futures_util::future::{Either, select};
20#[cfg(target_family = "wasm")]
21use gloo_timers::future::TimeoutFuture;
22#[cfg(not(target_family = "wasm"))]
23use tokio::time::timeout as tokio_timeout;
24
25#[derive(Clone, Copy, Debug, PartialEq, Eq)]
27pub struct ElapsedError();
28
29impl fmt::Display for ElapsedError {
30 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
31 write!(f, "time waiting for future has elapsed!")
32 }
33}
34
35impl Error for ElapsedError {}
36
37pub async fn timeout<F, T>(future: F, duration: Duration) -> Result<T, ElapsedError>
42where
43 F: Future<Output = T>,
44{
45 #[cfg(not(target_family = "wasm"))]
46 return tokio_timeout(duration, future).await.map_err(|_| ElapsedError());
47
48 #[cfg(target_family = "wasm")]
49 {
50 let timeout_future =
51 TimeoutFuture::new(u32::try_from(duration.as_millis()).expect("Overlong duration"));
52
53 match select(std::pin::pin!(future), timeout_future).await {
54 Either::Left((res, _)) => Ok(res),
55 Either::Right((_, _)) => Err(ElapsedError()),
56 }
57 }
58}
59
60#[cfg(test)]
61pub(crate) mod tests {
62 use std::{future, time::Duration};
63
64 use matrix_sdk_test_macros::async_test;
65
66 use super::timeout;
67
68 #[cfg(target_family = "wasm")]
69 wasm_bindgen_test::wasm_bindgen_test_configure!(run_in_browser);
70
71 #[async_test]
72 async fn test_without_timeout() {
73 timeout(future::ready(()), Duration::from_millis(100))
74 .await
75 .expect("future should have completed without ElapsedError");
76 }
77
78 #[async_test]
79 async fn test_with_timeout() {
80 timeout(future::pending::<()>(), Duration::from_millis(100))
81 .await
82 .expect_err("future should return an ElapsedError");
83 }
84}