1use crate::rdbms::RdbmsError;
3use std::sync::atomic::{AtomicU64, Ordering};
4use std::time::Duration;
5
6pub static QUERY_TIMEOUTS: AtomicU64 = AtomicU64::new(0);
12
13pub static TRANSACTIONS_LEAKED: AtomicU64 = AtomicU64::new(0);
15
16pub async fn run_with_timeout<F, T>(timeout: Option<Duration>, fut: F) -> Result<T, RdbmsError>
24where
25 F: std::future::Future<Output = Result<T, RdbmsError>>,
26{
27 match timeout {
28 None => fut.await,
29 Some(d) => match tokio::time::timeout(d, fut).await {
30 Ok(result) => result,
31 Err(_) => {
32 QUERY_TIMEOUTS.fetch_add(1, Ordering::Relaxed);
33 Err(RdbmsError::Timeout(format!("query exceeded {d:?}")))
34 }
35 },
36 }
37}
38
39#[cfg(test)]
40mod tests {
41 use super::*;
42 use crate::rdbms::RdbmsError;
43 use std::sync::atomic::Ordering;
44
45 #[tokio::test]
46 async fn none_timeout_passes_result_through() {
47 let r: Result<u64, RdbmsError> = run_with_timeout(None, async { Ok(42) }).await;
48 assert_eq!(r.unwrap(), 42);
49 }
50
51 #[tokio::test]
52 async fn fast_future_completes_within_timeout() {
53 let r: Result<u64, RdbmsError> =
54 run_with_timeout(Some(Duration::from_secs(5)), async { Ok(1) }).await;
55 assert_eq!(r.unwrap(), 1);
56 }
57
58 #[tokio::test]
59 async fn slow_future_times_out_and_counts() {
60 let before = QUERY_TIMEOUTS.load(Ordering::SeqCst);
61 let r: Result<(), RdbmsError> = run_with_timeout(Some(Duration::from_millis(10)), async {
62 tokio::time::sleep(Duration::from_millis(200)).await;
63 Ok(())
64 })
65 .await;
66 let err = r.unwrap_err();
67 assert!(matches!(err, RdbmsError::Timeout(_)), "got: {err:?}");
68 assert_eq!(QUERY_TIMEOUTS.load(Ordering::SeqCst), before + 1);
69 }
70}