#[cfg(feature = "shuttle-test")]
use shuttle::sync::atomic::AtomicUsize;
#[cfg(not(feature = "shuttle-test"))]
use std::sync::atomic::AtomicUsize;
use {
crossbeam_channel::{RecvError, SendError, TryRecvError},
crossbeam_utils::Backoff,
log::error,
std::{
mem,
sync::{
Arc,
atomic::{AtomicU32, Ordering, fence},
},
thread::{self, JoinHandle},
},
};
#[derive(Default)]
struct WakeEvent {
cookie: AtomicU32,
waiters: AtomicUsize,
}
impl WakeEvent {
fn register_waiter(&self) -> WakeWaiter<'_> {
let cookie = self.cookie.load(Ordering::Relaxed);
self.waiters.fetch_add(1, Ordering::Relaxed);
fence(Ordering::SeqCst);
WakeWaiter {
event: self,
cookie,
}
}
fn wake_one(&self) {
fence(Ordering::SeqCst);
if self.waiters.load(Ordering::Relaxed) != 0 {
self.cookie.fetch_add(1, Ordering::Relaxed);
atomic_wait::wake_one(&self.cookie);
}
}
fn wake_all(&self) {
fence(Ordering::SeqCst);
if self.waiters.load(Ordering::Relaxed) != 0 {
self.cookie.fetch_add(1, Ordering::Relaxed);
atomic_wait::wake_all(&self.cookie);
}
}
}
struct WakeWaiter<'a> {
event: &'a WakeEvent,
cookie: u32,
}
impl WakeWaiter<'_> {
fn wait(self) {
if self.event.cookie.load(Ordering::Relaxed) == self.cookie {
atomic_wait::wait(&self.event.cookie, self.cookie);
}
}
}
impl Drop for WakeWaiter<'_> {
fn drop(&mut self) {
self.event.waiters.fetch_sub(1, Ordering::Relaxed);
}
}
struct Shared {
wake_event: WakeEvent,
num_senders: AtomicUsize,
}
struct Sender<T> {
inner: crossbeam_channel::Sender<T>,
shared: Arc<Shared>,
}
impl<T> Sender<T> {
fn send(&self, value: T) -> Result<(), SendError<T>> {
self.inner.send(value)?;
self.shared.wake_event.wake_one();
Ok(())
}
}
impl<T> Clone for Sender<T> {
fn clone(&self) -> Self {
self.shared.num_senders.fetch_add(1, Ordering::Relaxed);
Self {
inner: self.inner.clone(),
shared: Arc::clone(&self.shared),
}
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
let (replacement, _) = crossbeam_channel::bounded(0);
drop(mem::replace(&mut self.inner, replacement));
if self.shared.num_senders.fetch_sub(1, Ordering::AcqRel) == 1 {
self.shared.wake_event.wake_all();
}
}
}
struct Receiver<T> {
inner: crossbeam_channel::Receiver<T>,
shared: Arc<Shared>,
}
impl<T> Receiver<T> {
fn recv(&self) -> Result<T, RecvError> {
loop {
let backoff = Backoff::new();
loop {
match self.inner.try_recv() {
Ok(value) => return Ok(value),
Err(TryRecvError::Disconnected) => return Err(RecvError),
Err(TryRecvError::Empty) if backoff.is_completed() => break,
Err(TryRecvError::Empty) => backoff.snooze(),
}
}
let waiter = self.shared.wake_event.register_waiter();
match self.inner.try_recv() {
Ok(value) => return Ok(value),
Err(TryRecvError::Disconnected) => return Err(RecvError),
Err(TryRecvError::Empty) => waiter.wait(),
}
}
}
}
impl<T> Clone for Receiver<T> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
shared: Arc::clone(&self.shared),
}
}
}
fn bounded<T>(capacity: usize) -> (Sender<T>, Receiver<T>) {
assert_ne!(capacity, 0, "channel capacity must be nonzero");
let (sender, receiver) = crossbeam_channel::bounded(capacity);
let shared = Arc::new(Shared {
wake_event: WakeEvent::default(),
num_senders: AtomicUsize::new(1),
});
(
Sender {
inner: sender,
shared: Arc::clone(&shared),
},
Receiver {
inner: receiver,
shared,
},
)
}
pub(crate) trait WorkerJob: Send + 'static {
fn run(self);
}
pub(crate) struct WorkerPool<J: WorkerJob> {
job_sender: Sender<J>,
worker_handles: Vec<JoinHandle<()>>,
}
impl<J: WorkerJob> WorkerPool<J> {
pub(crate) fn new(
thread_name_prefix: &str,
num_workers: usize,
job_queue_capacity: usize,
) -> Self {
assert_ne!(num_workers, 0, "worker pool must have at least one worker");
let (job_sender, job_receiver) = bounded::<J>(job_queue_capacity);
let worker_handles = (0..num_workers)
.map(|index| {
let job_receiver = job_receiver.clone();
thread::Builder::new()
.name(format!("{thread_name_prefix}{index:02}"))
.stack_size(2 * 1024 * 1024)
.spawn(move || {
while let Ok(job) = job_receiver.recv() {
job.run();
}
})
.expect("failed to spawn worker thread")
})
.collect();
Self {
job_sender,
worker_handles,
}
}
pub(crate) fn send(&self, job: J) {
self.job_sender
.send(job)
.expect("worker threads exited unexpectedly");
}
pub(crate) fn num_workers(&self) -> usize {
self.worker_handles.len()
}
}
impl<J: WorkerJob> Drop for WorkerPool<J> {
fn drop(&mut self) {
let (tmp, _) = bounded(1);
drop(mem::replace(&mut self.job_sender, tmp));
for worker_handle in self.worker_handles.drain(..) {
if let Err(err) = worker_handle.join() {
error!("worker thread failed: {err:?}");
}
}
}
}
#[cfg(all(test, not(feature = "shuttle-test")))]
mod tests {
use {
super::*,
std::{sync::Barrier, thread},
};
#[test]
fn test_wake_before_wait() {
let event = WakeEvent::default();
let waiter = event.register_waiter();
event.wake_one();
waiter.wait();
}
#[test]
fn test_wake_receivers_and_disconnect() {
const NUM_RECEIVERS: usize = 4;
let (sender, receiver) = bounded(NUM_RECEIVERS);
let sender1 = sender.clone();
drop(sender);
let barrier = Arc::new(Barrier::new(NUM_RECEIVERS + 1));
let handles = (0..NUM_RECEIVERS)
.map(|_| {
let receiver = receiver.clone();
let barrier = Arc::clone(&barrier);
thread::spawn(move || {
receiver.recv().unwrap();
barrier.wait();
assert!(receiver.recv().is_err());
})
})
.collect::<Vec<_>>();
while receiver.shared.wake_event.waiters.load(Ordering::Relaxed) != NUM_RECEIVERS {
thread::yield_now();
}
for _ in 0..NUM_RECEIVERS {
sender1.send(()).unwrap();
}
barrier.wait();
drop(sender1);
for handle in handles {
handle.join().unwrap();
}
}
}
#[cfg(all(test, feature = "shuttle-test"))]
mod shuttle_tests {
use {super::*, shuttle::thread};
#[test]
fn test_disconnect_is_visible_before_wake() {
shuttle::check_dfs(
|| {
let (sender, receiver) = bounded::<()>(1);
let sender1 = sender.clone();
let observer_receiver = receiver.clone();
let _waiter = receiver.shared.wake_event.register_waiter();
let sender_drop = thread::spawn(move || drop(sender));
let sender1_drop = thread::spawn(move || drop(sender1));
let observer = thread::spawn(move || {
if observer_receiver
.shared
.wake_event
.cookie
.load(Ordering::Relaxed)
!= 0
{
assert_eq!(
observer_receiver.inner.try_recv(),
Err(TryRecvError::Disconnected),
);
}
});
sender_drop.join().unwrap();
sender1_drop.join().unwrap();
observer.join().unwrap();
assert_ne!(receiver.shared.wake_event.cookie.load(Ordering::Relaxed), 0);
assert_eq!(receiver.inner.try_recv(), Err(TryRecvError::Disconnected));
},
None,
);
}
}