use std::future::Future;
use std::pin::Pin;
use std::task::{Context, Poll};
use std::time::Duration;
use tracing::Instrument;
use crate::error::MytheclipseError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TimeoutError {
Elapsed,
}
impl std::fmt::Display for TimeoutError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Elapsed => write!(f, "deadline elapsed"),
}
}
}
impl std::error::Error for TimeoutError {}
pub async fn with_timeout<T, F>(dur: Duration, future: F) -> Result<T, TimeoutError>
where
F: Future<Output = T>,
{
let span = tracing::info_span!("mytheclipse_timeout_task");
match tokio::time::timeout(dur, future.instrument(span)).await {
Ok(value) => Ok(value),
Err(_) => Err(TimeoutError::Elapsed),
}
}
pub struct Timeout<T> {
inner: Pin<Box<dyn Future<Output = Result<T, TimeoutError>> + Send>>,
}
impl<T: Send + 'static> Timeout<T> {
pub fn new<F>(dur: Duration, future: F) -> Self
where
F: Future<Output = T> + Send + 'static,
{
let future = async move {
match tokio::time::timeout(dur, future).await {
Ok(value) => Ok(value),
Err(_) => Err(TimeoutError::Elapsed),
}
};
Self {
inner: Box::pin(future),
}
}
}
impl<T> Future for Timeout<T> {
type Output = Result<T, TimeoutError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.get_mut().inner.as_mut().poll(cx)
}
}
pub async fn timeout<T, F>(dur: Duration, future: F) -> Result<T, MytheclipseError>
where
F: Future<Output = T>,
{
with_timeout(dur, future)
.await
.map_err(|_| MytheclipseError::Timeout)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn completes_within_bound_returns_value() {
let value = with_timeout(Duration::from_secs(5), async { 42u32 }).await;
assert_eq!(value.unwrap(), 42);
}
#[tokio::test]
async fn exceeding_bound_yields_elapsed() {
let outcome = with_timeout(Duration::from_millis(20), async {
tokio::time::sleep(Duration::from_secs(5)).await;
42u32
})
.await;
assert_eq!(outcome, Err(TimeoutError::Elapsed));
}
#[tokio::test]
async fn timeout_wrapper_maps_to_shared_error() {
let outcome = timeout(Duration::from_millis(10), async {
tokio::time::sleep(Duration::from_secs(5)).await;
})
.await;
assert_eq!(outcome, Err(MytheclipseError::Timeout));
}
#[tokio::test]
async fn timeout_future_is_spawnable() {
let bounded = Timeout::new(Duration::from_secs(5), async { 7u32 });
assert_eq!(tokio::spawn(bounded).await.unwrap().unwrap(), 7);
}
}