Skip to main content

ecat_data/
timeout.rs

1// Copyright (c) 2026 erik <erik@erik.xyz> — https://erik.xyz
2use crate::rdbms::RdbmsError;
3use std::sync::atomic::{AtomicU64, Ordering};
4use std::time::Duration;
5
6/// 查询超时累计次数。`metrics` feature 开启时由后端注册为
7/// `ecat_rdbms_query_timeout_total`。
8///
9/// 用进程级 `AtomicU64` 而非依赖 `ecat-metrics`:本 crate 保持零外部依赖,
10/// 指标的读取方按需接入。
11pub static QUERY_TIMEOUTS: AtomicU64 = AtomicU64::new(0);
12
13/// 未提交即 Drop 的事务累计数(`Transaction` 的 Drop guard 递增)。
14pub static TRANSACTIONS_LEAKED: AtomicU64 = AtomicU64::new(0);
15
16/// 给数据库调用套一层超时。
17///
18/// `None` 表示禁用超时,直接透传结果。超时发生时递增 [`QUERY_TIMEOUTS`]
19/// 并返回 [`RdbmsError::Timeout`]。
20///
21/// 这是池耗尽的头号防线:`acquire_timeout` 只约束「等连接」,
22/// 拿到连接后卡死的查询会一直占着它。
23pub 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}