use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
pub type DepthGauge = fn(f64);
#[derive(Debug)]
pub enum Rejected<T> {
Full(T),
Closed(T),
}
#[derive(Debug)]
pub enum Blocked<T> {
Closed(T),
TimedOut,
}
pub struct BoundedWorker<T> {
senders: Vec<mpsc::Sender<T>>,
next: Arc<AtomicUsize>,
pending: Arc<AtomicUsize>,
gauge: DepthGauge,
}
impl<T> Clone for BoundedWorker<T> {
fn clone(&self) -> Self {
Self {
senders: self.senders.clone(),
next: self.next.clone(),
pending: self.pending.clone(),
gauge: self.gauge,
}
}
}
impl<T> BoundedWorker<T> {
pub fn new(
shards: usize,
capacity_per_shard: usize,
gauge: DepthGauge,
) -> (Self, Vec<WorkerReceiver<T>>) {
let shards = shards.max(1);
let capacity_per_shard = capacity_per_shard.max(1);
let pending = Arc::new(AtomicUsize::new(0));
let mut senders = Vec::with_capacity(shards);
let mut receivers = Vec::with_capacity(shards);
for _ in 0..shards {
let (tx, rx) = mpsc::channel::<T>(capacity_per_shard);
senders.push(tx);
receivers.push(WorkerReceiver {
rx,
pending: pending.clone(),
gauge,
});
}
let worker = Self {
senders,
next: Arc::new(AtomicUsize::new(0)),
pending,
gauge,
};
(worker, receivers)
}
pub fn disabled(gauge: DepthGauge) -> Self {
Self {
senders: Vec::new(),
next: Arc::new(AtomicUsize::new(0)),
pending: Arc::new(AtomicUsize::new(0)),
gauge,
}
}
pub fn is_disabled(&self) -> bool {
self.senders.is_empty()
}
pub fn depth(&self) -> usize {
self.pending.load(Ordering::Acquire)
}
pub fn try_submit(&self, item: T) -> Result<(), Rejected<T>> {
if self.senders.is_empty() {
return Err(Rejected::Closed(item));
}
self.pending.fetch_add(1, Ordering::AcqRel);
let start = self.next.fetch_add(1, Ordering::Relaxed);
let mut item = item;
let mut any_open = false;
for i in 0..self.senders.len() {
let shard = &self.senders[(start.wrapping_add(i)) % self.senders.len()];
match shard.try_send(item) {
Ok(()) => {
(self.gauge)(self.depth() as f64);
return Ok(());
}
Err(mpsc::error::TrySendError::Full(returned)) => {
any_open = true;
item = returned;
}
Err(mpsc::error::TrySendError::Closed(returned)) => {
item = returned;
}
}
}
self.pending.fetch_sub(1, Ordering::AcqRel);
Err(if any_open {
Rejected::Full(item)
} else {
Rejected::Closed(item)
})
}
pub async fn submit_blocking(&self, item: T, timeout: Duration) -> Result<(), Blocked<T>> {
if self.senders.is_empty() {
return Err(Blocked::Closed(item));
}
self.pending.fetch_add(1, Ordering::AcqRel);
let start = self.next.fetch_add(1, Ordering::Relaxed);
let shard = &self.senders[start % self.senders.len()];
match tokio::time::timeout(timeout, shard.send(item)).await {
Ok(Ok(())) => {
(self.gauge)(self.depth() as f64);
Ok(())
}
Ok(Err(mpsc::error::SendError(returned))) => {
self.pending.fetch_sub(1, Ordering::AcqRel);
Err(Blocked::Closed(returned))
}
Err(_elapsed) => {
self.pending.fetch_sub(1, Ordering::AcqRel);
Err(Blocked::TimedOut)
}
}
}
pub fn drain_handle(&self) -> DrainHandle<T> {
DrainHandle {
senders: self.senders.clone(),
pending: self.pending.clone(),
}
}
}
pub struct WorkerReceiver<T> {
rx: mpsc::Receiver<T>,
pending: Arc<AtomicUsize>,
gauge: DepthGauge,
}
#[derive(Debug)]
pub enum Recv<T> {
Item(T),
Closed,
Elapsed,
}
impl<T> WorkerReceiver<T> {
pub async fn recv(&mut self) -> Option<T> {
let item = self.rx.recv().await?;
self.release();
Some(item)
}
pub async fn recv_leased(&mut self) -> Option<Leased<T>> {
let item = self.rx.recv().await?;
Some(Leased {
item,
pending: self.pending.clone(),
gauge: self.gauge,
})
}
pub async fn recv_timeout(&mut self, within: Duration) -> Recv<T> {
match tokio::time::timeout(within, self.rx.recv()).await {
Ok(Some(item)) => {
self.release();
Recv::Item(item)
}
Ok(None) => Recv::Closed,
Err(_elapsed) => Recv::Elapsed,
}
}
#[cfg(test)]
pub(crate) fn try_recv(&mut self) -> Option<T> {
let item = self.rx.try_recv().ok()?;
self.release();
Some(item)
}
fn release(&self) {
let depth = self
.pending
.fetch_sub(1, Ordering::AcqRel)
.saturating_sub(1);
(self.gauge)(depth as f64);
}
}
pub struct Leased<T> {
item: T,
pending: Arc<AtomicUsize>,
gauge: DepthGauge,
}
impl<T> std::ops::Deref for Leased<T> {
type Target = T;
fn deref(&self) -> &T {
&self.item
}
}
impl<T> Drop for Leased<T> {
fn drop(&mut self) {
let depth = self
.pending
.fetch_sub(1, Ordering::Release)
.saturating_sub(1);
(self.gauge)(depth as f64);
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DrainWitness {
TasksExit,
QueueEmpty,
}
#[derive(Debug, PartialEq, Eq)]
pub enum DrainOutcome {
Drained,
WorkerPanicked,
TimedOut {
lost: usize,
},
}
pub struct DrainHandle<T> {
senders: Vec<mpsc::Sender<T>>,
pending: Arc<AtomicUsize>,
}
const DRAIN_POLL_INTERVAL: Duration = Duration::from_millis(2);
impl<T> DrainHandle<T> {
pub fn depth(&self) -> usize {
self.pending.load(Ordering::Acquire)
}
pub async fn drain(
self,
mut joins: Vec<JoinHandle<()>>,
witness: DrainWitness,
deadline: Duration,
) -> DrainOutcome {
drop(self.senders);
if joins.is_empty() {
return DrainOutcome::Drained;
}
let pending = self.pending.clone();
let finished = tokio::time::timeout(deadline, async {
let all_joined = async {
let mut panicked = false;
for join in joins.iter_mut() {
if join.await.is_err() {
panicked = true;
}
}
panicked
};
match witness {
DrainWitness::TasksExit => all_joined.await,
DrainWitness::QueueEmpty => {
tokio::select! {
panicked = all_joined => panicked,
() = async {
while pending.load(Ordering::Acquire) > 0 {
tokio::time::sleep(DRAIN_POLL_INTERVAL).await;
}
} => false,
}
}
}
})
.await;
for join in &joins {
join.abort();
}
match finished {
Ok(false) => DrainOutcome::Drained,
Ok(true) => DrainOutcome::WorkerPanicked,
Err(_elapsed) => DrainOutcome::TimedOut {
lost: self.pending.load(Ordering::Acquire),
},
}
}
}
#[cfg(test)]
mod tests {
#![allow(clippy::panic)]
use super::*;
fn no_gauge(_: f64) {}
#[tokio::test]
async fn a_full_queue_sheds_immediately_and_hands_the_item_back() {
let (worker, _rx) = BoundedWorker::<u32>::new(1, 1, no_gauge);
worker.try_submit(1).expect("first fits");
match worker.try_submit(2) {
Err(Rejected::Full(item)) => assert_eq!(item, 2),
other => panic!("expected Full(2), got {other:?}"),
}
}
#[tokio::test]
async fn a_closed_queue_is_distinct_from_a_full_one() {
let (worker, rx) = BoundedWorker::<u32>::new(1, 1, no_gauge);
drop(rx);
match worker.try_submit(1) {
Err(Rejected::Closed(item)) => assert_eq!(item, 1),
other => panic!("expected Closed(1), got {other:?}"),
}
}
#[tokio::test]
async fn a_rejection_releases_its_reservation() {
let (worker, _rx) = BoundedWorker::<u32>::new(1, 1, no_gauge);
worker.try_submit(1).expect("first fits");
assert_eq!(worker.depth(), 1);
for _ in 0..10 {
assert!(worker.try_submit(2).is_err());
}
assert_eq!(worker.depth(), 1, "shed submissions must not accumulate");
}
#[tokio::test]
async fn a_full_shard_falls_through_to_its_siblings() {
let (worker, mut receivers) = BoundedWorker::<u32>::new(2, 1, no_gauge);
worker.try_submit(1).expect("shard a");
worker.try_submit(2).expect("shard b");
assert!(worker.try_submit(3).is_err(), "both shards are full");
receivers[0].recv().await.expect("an item");
worker
.try_submit(4)
.expect("the freed shard must be found by fallthrough");
}
#[tokio::test]
async fn one_dead_shard_among_several_reports_full_not_closed() {
let (worker, mut receivers) = BoundedWorker::<u32>::new(2, 1, no_gauge);
let dead = receivers.remove(0);
drop(dead);
worker.try_submit(1).expect("the live shard");
match worker.try_submit(2) {
Err(Rejected::Full(_)) => {}
other => panic!("expected Full, got {other:?}"),
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn the_depth_counter_never_exceeds_capacity_under_interleaving() {
const SHARDS: usize = 2;
const CAPACITY: usize = 4;
const PRODUCERS: usize = 4;
const PER_PRODUCER: usize = 500;
let (worker, receivers) = BoundedWorker::<u32>::new(SHARDS, CAPACITY, no_gauge);
let ceiling = SHARDS * CAPACITY + PRODUCERS;
let consumers: Vec<_> = receivers
.into_iter()
.map(|mut rx| {
tokio::spawn(async move {
let mut taken = 0usize;
while rx.recv().await.is_some() {
taken += 1;
tokio::task::yield_now().await;
}
taken
})
})
.collect();
let observer = {
let worker = worker.clone();
tokio::spawn(async move {
let mut worst = 0usize;
for _ in 0..20_000 {
worst = worst.max(worker.depth());
tokio::task::yield_now().await;
}
worst
})
};
let producers: Vec<_> = (0..PRODUCERS)
.map(|_| {
let worker = worker.clone();
tokio::spawn(async move {
let mut accepted = 0usize;
for i in 0..PER_PRODUCER {
if worker.try_submit(i as u32).is_ok() {
accepted += 1;
}
tokio::task::yield_now().await;
}
accepted
})
})
.collect();
let mut accepted = 0usize;
for p in producers {
accepted += p.await.expect("producer");
}
let worst = observer.await.expect("observer");
drop(worker);
let mut taken = 0usize;
for c in consumers {
taken += c.await.expect("consumer");
}
assert!(
worst <= ceiling,
"depth reached {worst}, above the {ceiling} the shards hold plus one \
reservation per producer — the counter wrapped"
);
assert_eq!(
accepted, taken,
"every accepted item must be delivered once"
);
}
#[tokio::test]
async fn a_drain_waits_for_its_workers_to_exit() {
let (worker, mut receivers) = BoundedWorker::<u32>::new(1, 4, no_gauge);
let mut rx = receivers.pop().expect("one receiver");
let drained = Arc::new(AtomicUsize::new(0));
let join = {
let drained = drained.clone();
tokio::spawn(async move {
while rx.recv().await.is_some() {
drained.fetch_add(1, Ordering::Release);
}
})
};
for i in 0..4 {
worker.try_submit(i).expect("fits");
}
let handle = worker.drain_handle();
drop(worker);
let outcome = handle
.drain(vec![join], DrainWitness::TasksExit, Duration::from_secs(5))
.await;
assert_eq!(outcome, DrainOutcome::Drained);
assert_eq!(drained.load(Ordering::Acquire), 4);
}
#[tokio::test]
async fn queue_empty_finishes_while_a_producer_clone_is_still_held() {
let (worker, mut receivers) = BoundedWorker::<u32>::new(1, 4, no_gauge);
let mut rx = receivers.pop().expect("one receiver");
let join = tokio::spawn(async move { while rx.recv().await.is_some() {} });
worker.try_submit(1).expect("fits");
let handle = worker.drain_handle();
let outcome = handle
.drain(vec![join], DrainWitness::QueueEmpty, Duration::from_secs(5))
.await;
assert_eq!(outcome, DrainOutcome::Drained);
assert_eq!(worker.depth(), 0);
}
#[tokio::test]
async fn a_drain_that_times_out_reports_what_it_abandoned() {
let (worker, receivers) = BoundedWorker::<u32>::new(1, 4, no_gauge);
let _receivers = receivers;
let join = tokio::spawn(async { std::future::pending::<()>().await });
for i in 0..3 {
worker.try_submit(i).expect("fits");
}
let handle = worker.drain_handle();
drop(worker);
let outcome = handle
.drain(
vec![join],
DrainWitness::TasksExit,
Duration::from_millis(50),
)
.await;
assert_eq!(outcome, DrainOutcome::TimedOut { lost: 3 });
}
#[tokio::test]
async fn a_panicking_worker_is_reported_rather_than_read_as_a_clean_drain() {
let (worker, receivers) = BoundedWorker::<u32>::new(1, 4, no_gauge);
let _receivers = receivers;
let join = tokio::spawn(async { panic!("worker exploded") });
let handle = worker.drain_handle();
drop(worker);
let outcome = handle
.drain(vec![join], DrainWitness::TasksExit, Duration::from_secs(5))
.await;
assert_eq!(outcome, DrainOutcome::WorkerPanicked);
}
#[tokio::test]
async fn blocking_submission_gives_up_after_its_timeout() {
let (worker, _receivers) = BoundedWorker::<u32>::new(1, 1, no_gauge);
worker.try_submit(1).expect("first fits");
match worker.submit_blocking(2, Duration::from_millis(25)).await {
Err(Blocked::TimedOut) => {}
other => panic!("expected TimedOut, got {other:?}"),
}
assert_eq!(
worker.depth(),
1,
"the abandoned send must release its slot"
);
}
#[tokio::test]
async fn a_disabled_queue_accepts_nothing_and_counts_nothing() {
let worker = BoundedWorker::<u32>::disabled(no_gauge);
assert!(worker.is_disabled());
assert!(matches!(worker.try_submit(1), Err(Rejected::Closed(1))));
assert_eq!(worker.depth(), 0);
}
}