use std::collections::BTreeSet;
use std::sync::Mutex;
use std::sync::atomic::{AtomicBool, Ordering};
use std::time::Duration;
use tokio::sync::Notify;
use crate::core::RunId;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct DrainReport {
pub at_start: Vec<RunId>,
pub unfinished: Vec<RunId>,
}
impl DrainReport {
#[must_use]
pub fn settled(&self) -> usize {
self.at_start
.iter()
.filter(|run| !self.unfinished.contains(run))
.count()
}
#[must_use]
pub fn is_complete(&self) -> bool {
self.unfinished.is_empty()
}
}
#[derive(Debug, Default)]
pub(crate) struct InFlight {
running: Mutex<BTreeSet<RunId>>,
idle: Notify,
draining: AtomicBool,
}
impl InFlight {
pub(crate) fn is_draining(&self) -> bool {
self.draining.load(Ordering::Acquire)
}
pub(crate) fn enter(self: &std::sync::Arc<Self>, run: RunId) -> Ticket {
crate::core::poison::recover(&self.running).insert(run);
Ticket {
inflight: std::sync::Arc::clone(self),
run,
}
}
fn leave(&self, run: RunId) {
let empty = {
let mut running = crate::core::poison::recover(&self.running);
running.remove(&run);
running.is_empty()
};
if empty {
self.idle.notify_waiters();
}
}
fn snapshot(&self) -> Vec<RunId> {
crate::core::poison::recover(&self.running)
.iter()
.copied()
.collect()
}
pub(crate) async fn drain(&self, grace: Duration) -> DrainReport {
self.draining.store(true, Ordering::Release);
let at_start = self.snapshot();
let deadline = tokio::time::Instant::now() + grace;
loop {
let waiter = self.idle.notified();
tokio::pin!(waiter);
waiter.as_mut().enable();
if crate::core::poison::recover(&self.running).is_empty() {
break;
}
tokio::select! {
() = waiter => {}
() = tokio::time::sleep_until(deadline) => break,
}
}
DrainReport {
at_start,
unfinished: self.snapshot(),
}
}
}
pub(crate) struct Ticket {
inflight: std::sync::Arc<InFlight>,
run: RunId,
}
impl Drop for Ticket {
fn drop(&mut self) {
self.inflight.leave(self.run);
}
}
#[cfg(test)]
mod tests {
use super::{DrainReport, InFlight};
use crate::core::RunId;
use std::sync::Arc;
use std::time::Duration;
#[tokio::test]
async fn a_drain_with_nothing_running_returns_at_once() {
let inflight = Arc::new(InFlight::default());
let report = inflight.drain(Duration::from_secs(30)).await;
assert!(report.is_complete());
assert_eq!(report.settled(), 0);
assert!(inflight.is_draining());
}
#[tokio::test]
async fn a_run_that_ends_during_the_drain_is_not_waited_out() {
let inflight = Arc::new(InFlight::default());
let run = RunId::generate();
let ticket = inflight.enter(run);
let waiting = tokio::spawn({
let inflight = Arc::clone(&inflight);
async move { inflight.drain(Duration::from_secs(600)).await }
});
tokio::task::yield_now().await;
drop(ticket);
let report = waiting.await.expect("the drain task");
assert_eq!(report.at_start, vec![run]);
assert!(report.is_complete());
}
#[tokio::test]
async fn a_run_that_outlasts_the_grace_period_is_named() {
let inflight = Arc::new(InFlight::default());
let run = RunId::generate();
let _ticket = inflight.enter(run);
let report = inflight.drain(Duration::from_millis(50)).await;
assert_eq!(report.unfinished, vec![run]);
assert_eq!(report.settled(), 0);
assert!(!report.is_complete());
}
#[test]
fn a_report_counts_what_settled_rather_than_what_started() {
let (settled, stuck) = (RunId::generate(), RunId::generate());
let report = DrainReport {
at_start: vec![settled, stuck],
unfinished: vec![stuck],
};
assert_eq!(report.settled(), 1);
assert!(!report.is_complete());
}
#[test]
fn a_run_that_appeared_after_the_drain_began_does_not_underflow_the_count() {
let report = DrainReport {
at_start: Vec::new(),
unfinished: vec![RunId::generate()],
};
assert_eq!(report.settled(), 0);
assert!(!report.is_complete());
}
}