use crate::rdbms::RdbmsError;
use ecat_errors::{Error, ErrorCode};
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::Duration;
pub static TRANSACTIONS_LEAKED: AtomicU64 = AtomicU64::new(0);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BackendKind {
Rdbms = 0,
Cache = 1,
Search = 2,
Graph = 3,
Document = 4,
Storage = 5,
Tsdb = 6,
}
impl BackendKind {
pub fn slug(&self) -> &'static str {
match self {
BackendKind::Rdbms => "rdbms",
BackendKind::Cache => "cache",
BackendKind::Search => "search",
BackendKind::Graph => "graph",
BackendKind::Document => "document",
BackendKind::Storage => "storage",
BackendKind::Tsdb => "tsdb",
}
}
}
pub static TIMEOUTS: [AtomicU64; 7] = [const { AtomicU64::new(0) }; 7];
pub fn timeout_counter(kind: BackendKind) -> &'static AtomicU64 {
match kind {
BackendKind::Rdbms => &TIMEOUTS[0],
BackendKind::Cache => &TIMEOUTS[1],
BackendKind::Search => &TIMEOUTS[2],
BackendKind::Graph => &TIMEOUTS[3],
BackendKind::Document => &TIMEOUTS[4],
BackendKind::Storage => &TIMEOUTS[5],
BackendKind::Tsdb => &TIMEOUTS[6],
}
}
pub trait TimeoutError: Sized {
fn from_timeout(kind: BackendKind, d: Duration) -> Self;
}
impl TimeoutError for RdbmsError {
fn from_timeout(_kind: BackendKind, d: Duration) -> Self {
RdbmsError::Timeout(format!("query exceeded {d:?}"))
}
}
impl TimeoutError for Error {
fn from_timeout(kind: BackendKind, d: Duration) -> Self {
Error::new(
ErrorCode::DeadlineExceeded,
kind.slug(),
format!("call exceeded {d:?}"),
)
}
}
pub async fn run_with_timeout<F, T, E>(
kind: BackendKind,
timeout: Option<Duration>,
fut: F,
) -> Result<T, E>
where
F: std::future::Future<Output = Result<T, E>>,
E: TimeoutError,
{
match timeout {
None => fut.await,
Some(d) => match tokio::time::timeout(d, fut).await {
Ok(result) => result,
Err(_) => {
timeout_counter(kind).fetch_add(1, Ordering::Relaxed);
Err(E::from_timeout(kind, d))
}
},
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::rdbms::RdbmsError;
use std::sync::atomic::Ordering;
fn count(kind: BackendKind) -> u64 {
timeout_counter(kind).load(Ordering::SeqCst)
}
#[test]
fn timeout_counter_indexes_match_discriminants() {
for k in [
BackendKind::Rdbms,
BackendKind::Cache,
BackendKind::Search,
BackendKind::Graph,
BackendKind::Document,
BackendKind::Storage,
BackendKind::Tsdb,
] {
assert!(
std::ptr::eq(timeout_counter(k), &TIMEOUTS[k as usize]),
"timeout_counter({k:?}) 必须与 TIMEOUTS[{}] 是同一个槽",
k as usize
);
}
}
#[tokio::test]
async fn none_timeout_passes_result_through() {
let r: Result<u64, RdbmsError> =
run_with_timeout(BackendKind::Rdbms, 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(BackendKind::Rdbms, 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 = count(BackendKind::Rdbms);
let r: Result<(), RdbmsError> =
run_with_timeout(BackendKind::Rdbms, 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!(count(BackendKind::Rdbms), before + 1);
}
#[tokio::test]
async fn error_path_maps_to_deadline_exceeded() {
let r: Result<(), Error> =
run_with_timeout(BackendKind::Cache, Some(Duration::from_millis(10)), async {
tokio::time::sleep(Duration::from_millis(200)).await;
Ok(())
})
.await;
let err = r.unwrap_err();
assert_eq!(err.code, ErrorCode::DeadlineExceeded, "got: {err:?}");
assert_eq!(err.reason, "cache", "reason 应为组件标识,got: {err:?}");
}
#[tokio::test]
async fn timeout_counters_are_per_backend_kind() {
let witness_before = count(BackendKind::Storage);
let tsdb_before = count(BackendKind::Tsdb);
let slow = || async {
tokio::time::sleep(Duration::from_millis(200)).await;
Ok::<(), Error>(())
};
let _: Result<(), Error> =
run_with_timeout(BackendKind::Tsdb, Some(Duration::from_millis(10)), slow()).await;
assert_eq!(count(BackendKind::Tsdb), tsdb_before + 1);
assert_eq!(
count(BackendKind::Storage),
witness_before,
"Tsdb 的超时不得落到别的槽"
);
}
}