use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::time::Duration;
use tokio::sync::watch;
use tokio::time::Instant;
static NEXT_OPERATION_ID: AtomicU64 = AtomicU64::new(1);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct OperationId(u64);
#[derive(Clone)]
pub struct OperationContext {
id: OperationId,
deadline: Instant,
canceled: watch::Receiver<bool>,
interrupted: Arc<AtomicBool>,
}
#[derive(Clone)]
pub struct OperationCancellation {
id: OperationId,
canceled: watch::Sender<bool>,
interrupted: Arc<AtomicBool>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OperationFailure {
Canceled,
TimedOut,
}
impl OperationContext {
pub fn new(timeout: Duration) -> (Self, OperationCancellation) {
let id = OperationId(NEXT_OPERATION_ID.fetch_add(1, Ordering::Relaxed));
let (canceled, receiver) = watch::channel(false);
let interrupted = Arc::new(AtomicBool::new(false));
(
Self {
id,
deadline: Instant::now() + timeout,
canceled: receiver,
interrupted: Arc::clone(&interrupted),
},
OperationCancellation {
id,
canceled,
interrupted,
},
)
}
pub fn id(&self) -> OperationId {
self.id
}
pub fn is_canceled(&self) -> bool {
self.interrupted.load(Ordering::Acquire)
}
pub fn interruption_flag(&self) -> Arc<AtomicBool> {
Arc::clone(&self.interrupted)
}
pub async fn run<F, T>(&mut self, future: F) -> Result<T, OperationFailure>
where
F: Future<Output = T>,
{
if self.is_canceled() {
return Err(OperationFailure::Canceled);
}
tokio::select! {
biased;
changed = self.canceled.changed() => {
let _ = changed;
self.interrupted.store(true, Ordering::Release);
Err(OperationFailure::Canceled)
}
() = tokio::time::sleep_until(self.deadline) => {
self.interrupted.store(true, Ordering::Release);
Err(OperationFailure::TimedOut)
},
value = future => Ok(value),
}
}
}
impl OperationCancellation {
pub fn id(&self) -> OperationId {
self.id
}
pub fn cancel(&self) {
self.interrupted.store(true, Ordering::Release);
self.canceled.send_replace(true);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn cancellation_and_timeout_have_distinct_outcomes() {
let (mut canceled_context, cancellation) = OperationContext::new(Duration::from_secs(1));
let canceled_flag = canceled_context.interruption_flag();
cancellation.cancel();
assert_eq!(
canceled_context.run(std::future::pending::<()>()).await,
Err(OperationFailure::Canceled)
);
assert!(canceled_flag.load(Ordering::Acquire));
let (mut timed_context, _timeout_cancellation) =
OperationContext::new(Duration::from_millis(1));
let timed_flag = timed_context.interruption_flag();
assert_eq!(
timed_context.run(std::future::pending::<()>()).await,
Err(OperationFailure::TimedOut)
);
assert!(timed_flag.load(Ordering::Acquire));
assert!(canceled_flag.load(Ordering::Acquire));
}
#[test]
fn operation_ids_are_unique_and_shared_with_the_cancellation_handle() {
let (first, first_cancellation) = OperationContext::new(Duration::from_secs(1));
let (second, _) = OperationContext::new(Duration::from_secs(1));
assert_eq!(first.id(), first_cancellation.id());
assert_ne!(first.id(), second.id());
}
}