use crate::rdbms::RdbmsError;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
pub static QUERY_TIMEOUTS: AtomicU64 = AtomicU64::new(0);
pub static TRANSACTIONS_LEAKED: AtomicU64 = AtomicU64::new(0);
pub async fn run_with_timeout<F, T>(timeout: Option<Duration>, fut: F) -> Result<T, RdbmsError>
where
F: std::future::Future<Output = Result<T, RdbmsError>>,
{
match timeout {
None => fut.await,
Some(d) => match tokio::time::timeout(d, fut).await {
Ok(result) => result,
Err(_) => {
QUERY_TIMEOUTS.fetch_add(1, Ordering::Relaxed);
Err(RdbmsError::Timeout(format!("query exceeded {d:?}")))
}
},
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::rdbms::RdbmsError;
use std::sync::atomic::Ordering;
#[tokio::test]
async fn none_timeout_passes_result_through() {
let r: Result<u64, RdbmsError> = run_with_timeout(None, async { Ok(42) }).await;
assert_eq!(r.unwrap(), 42);
}
#[tokio::test]
async fn fast_future_completes_within_timeout() {
let r: Result<u64, RdbmsError> =
run_with_timeout(Some(Duration::from_secs(5)), async { Ok(1) }).await;
assert_eq!(r.unwrap(), 1);
}
#[tokio::test]
async fn slow_future_times_out_and_counts() {
let before = QUERY_TIMEOUTS.load(Ordering::SeqCst);
let r: Result<(), RdbmsError> = run_with_timeout(Some(Duration::from_millis(10)), async {
tokio::time::sleep(Duration::from_millis(200)).await;
Ok(())
})
.await;
let err = r.unwrap_err();
assert!(matches!(err, RdbmsError::Timeout(_)), "got: {err:?}");
assert_eq!(QUERY_TIMEOUTS.load(Ordering::SeqCst), before + 1);
}
}