use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
pub fn new<T>(size: usize) -> (SchedulerSender<T>, SchedulerReceiver<T>) {
let inner = Arc::new(Inner::new(size));
let tx = SchedulerSender {
inner: inner.clone(),
};
let rx = SchedulerReceiver { inner };
(tx, rx)
}
pub struct SchedulerSender<T> {
inner: Arc<Inner<T>>,
}
impl<T> Clone for SchedulerSender<T> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
impl<T> SchedulerSender<T> {
pub fn send(&self, mut value: T) -> Result<(), T> {
match self.try_send_spin(value) {
Ok(()) => return Ok(()),
Err(TrySendError::Disconnected(value)) => return Err(value),
Err(TrySendError::Full(v)) => {
value = v;
},
}
futures_executor::block_on(self.send_async(value))
}
pub async fn send_async(&self, mut value: T) -> Result<(), T> {
match self.try_send_spin(value) {
Ok(()) => return Ok(()),
Err(TrySendError::Disconnected(value)) => return Err(value),
Err(TrySendError::Full(v)) => {
value = v;
},
}
let notify_fut = self.inner.notify.notified();
tokio::pin!(notify_fut);
loop {
notify_fut.as_mut().enable();
match self.inner.try_send(value) {
Err(TrySendError::Disconnected(value)) => return Err(value),
Err(TrySendError::Full(v)) => {
value = v;
notify_fut.as_mut().await;
notify_fut.set(self.inner.notify.notified());
},
Ok(()) => return Ok(()),
}
}
}
fn try_send_spin(&self, mut value: T) -> Result<(), TrySendError<T>> {
match self.inner.try_send(value) {
Err(TrySendError::Disconnected(value)) => {
return Err(TrySendError::Disconnected(value));
},
Err(TrySendError::Full(v)) => {
value = v;
},
Ok(()) => return Ok(()),
}
for _ in 0..10 {
match self.inner.try_send(value) {
Err(TrySendError::Disconnected(value)) => {
return Err(TrySendError::Disconnected(value));
},
Err(TrySendError::Full(v)) => {
value = v;
},
Ok(()) => return Ok(()),
}
}
Err(TrySendError::Full(value))
}
}
pub struct SchedulerReceiver<T> {
inner: Arc<Inner<T>>,
}
impl<T> SchedulerReceiver<T> {
pub fn pop(&self) -> Option<T> {
self.inner.queue.pop()
}
pub fn wake_n(&self, n: usize) {
for _ in 0..n {
self.inner.notify.notify_one();
}
}
pub fn wake_all(&self) {
self.inner.notify.notify_waiters();
}
pub fn is_empty(&self) -> bool {
self.inner.queue.is_empty()
}
pub fn len(&self) -> usize {
self.inner.queue.len()
}
pub fn is_disconnected(&self) -> bool {
Arc::strong_count(&self.inner) == 1
}
}
impl<T> Drop for SchedulerReceiver<T> {
fn drop(&mut self) {
self.inner
.receiver_disconnected
.store(true, Ordering::Release);
}
}
struct Inner<T> {
queue: crossbeam_queue::ArrayQueue<T>,
notify: tokio::sync::Notify,
receiver_disconnected: AtomicBool,
}
impl<T> Inner<T> {
fn new(size: usize) -> Self {
Self {
queue: crossbeam_queue::ArrayQueue::new(size),
notify: tokio::sync::Notify::new(),
receiver_disconnected: AtomicBool::new(false),
}
}
fn try_send(&self, value: T) -> Result<(), TrySendError<T>> {
if self.receiver_disconnected.load(Ordering::Acquire) {
return Err(TrySendError::Disconnected(value));
}
if let Err(v) = self.queue.push(value) {
Err(TrySendError::Full(v))
} else {
Ok(())
}
}
}
enum TrySendError<T> {
Disconnected(T),
Full(T),
}
#[cfg(test)]
mod tests {
#[test]
fn test_queue_disconnect() {
let (tx, rx) = super::new::<()>(9);
assert!(!rx.is_disconnected());
drop(tx);
assert!(rx.is_disconnected());
}
#[test]
fn test_queue_sync_handling() {
let _ = tracing_subscriber::fmt::try_init();
let (tx, rx) = super::new(9);
for i in 0..8 {
tracing::info!(i);
assert!(tx.send(i).is_ok());
}
tracing::info!("spawn thread");
let handle = std::thread::spawn(move || tx.send(8));
assert_eq!(rx.pop(), Some(0));
assert_eq!(rx.pop(), Some(1));
assert_eq!(rx.pop(), Some(2));
rx.wake_all();
let result = handle.join().unwrap();
assert!(result.is_ok());
}
#[tokio::test]
async fn test_queue_async_handling() {
let _ = tracing_subscriber::fmt::try_init();
let (tx, rx) = super::new(9);
for i in 0..8 {
tracing::info!(i);
assert!(tx.send_async(i).await.is_ok());
}
tracing::info!("spawn thread");
let handle = tokio::spawn(async move { tx.send_async(8).await });
assert_eq!(rx.pop(), Some(0));
assert_eq!(rx.pop(), Some(1));
assert_eq!(rx.pop(), Some(2));
rx.wake_all();
let result = handle.await.unwrap();
assert!(result.is_ok());
}
}