use std::future::Future;
use std::sync::Arc;
use std::sync::OnceLock;
use std::time::Duration;
use tokio::time::Instant;
use tokio_util::sync::CancellationToken;
use crate::error::Error;
use crate::timeout::TimeoutScope;
#[derive(Debug, Clone)]
pub(crate) enum CancelReason {
Timeout {
scope: TimeoutScope,
elapsed: Duration,
},
}
#[derive(Debug, Clone)]
pub(crate) struct CallCancellation {
token: CancellationToken,
reason: Arc<OnceLock<CancelReason>>,
}
impl CallCancellation {
pub(crate) fn new(parent: &CancellationToken) -> Self {
Self {
token: parent.child_token(),
reason: Arc::new(OnceLock::new()),
}
}
pub(crate) fn child(&self) -> Self {
Self {
token: self.token.child_token(),
reason: Arc::clone(&self.reason),
}
}
pub(crate) fn token(&self) -> &CancellationToken {
&self.token
}
pub(crate) fn is_cancelled(&self) -> bool {
self.token.is_cancelled()
}
pub(crate) fn error(&self) -> Error {
match self.reason.get() {
Some(CancelReason::Timeout { scope, elapsed }) => Error::Timeout {
scope: scope.clone(),
elapsed: *elapsed,
},
None => Error::Cancelled,
}
}
pub(crate) fn map_error(&self, error: Error) -> Error {
match error {
Error::Cancelled => self.error(),
other => other,
}
}
pub(crate) fn cancel_for_timeout(&self, scope: TimeoutScope, elapsed: Duration) {
let _ = self.reason.set(CancelReason::Timeout { scope, elapsed });
self.token.cancel();
}
pub(crate) async fn with_timeout<T>(
&self,
scope: TimeoutScope,
duration: Option<Duration>,
future: impl Future<Output = Result<T, Error>>,
) -> Result<T, Error> {
let Some(duration) = duration else {
return future.await;
};
let start = Instant::now();
match tokio::time::timeout(duration, future).await {
Ok(result) => result,
Err(_) => {
let elapsed = start.elapsed();
self.cancel_for_timeout(scope.clone(), elapsed);
Err(Error::Timeout { scope, elapsed })
}
}
}
}