use std::{
future::Future,
io, thread,
time::{Duration, Instant},
};
use crate::{BatchError, Channel, Receiver, Sender, sync};
pub fn spawn<
T: Channel + Send + 'static,
F: Future<Output = Result<(), BatchError<T>>> + Send + 'static,
>(
thread_name: impl Into<String>,
receiver: Receiver<T>,
on_batch: impl FnMut(T) -> F + Send + 'static,
) -> io::Result<thread::JoinHandle<()>>
where
T::Item: Send + 'static,
{
let receive = async move {
receiver
.exec(|delay| tokio::time::sleep(delay), on_batch)
.await
};
thread::Builder::new()
.name(thread_name.into())
.spawn(move || {
tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.unwrap()
.block_on(receive);
})
}
pub async fn exec<T: Channel, F: Future<Output = Result<(), BatchError<T>>>>(
receiver: Receiver<T>,
on_batch: impl FnMut(T) -> F,
) {
receiver
.exec(|delay| tokio::time::sleep(delay), on_batch)
.await
}
pub fn blocking_flush<T: Channel>(sender: &Sender<T>, timeout: Duration) -> bool {
match tokio::runtime::Handle::try_current() {
Ok(handle) => handle.block_on(flush(sender, timeout)),
Err(_) => sync::blocking_flush(sender, timeout),
}
}
pub async fn flush<T: Channel>(sender: &Sender<T>, timeout: Duration) -> bool {
let (notifier, notified) = tokio::sync::oneshot::channel();
sender.when_flushed(move || {
let _ = notifier.send(());
});
wait(notified, timeout).await
}
pub fn blocking_send<T: Channel>(
sender: &Sender<T>,
msg: T::Item,
timeout: Duration,
) -> Result<(), BatchError<T::Item>> {
match tokio::runtime::Handle::try_current() {
Ok(handle) => handle.block_on(send(sender, msg, timeout)),
Err(_) => sync::blocking_send(sender, msg, timeout),
}
}
pub async fn send<T: Channel>(
sender: &Sender<T>,
msg: T::Item,
timeout: Duration,
) -> Result<(), BatchError<T::Item>> {
let start = Instant::now();
sender
.send_or_wait(
msg,
timeout,
|| start.elapsed(),
|sender, timeout| async move {
let (notifier, notified) = tokio::sync::oneshot::channel();
sender.when_empty(move || {
let _ = notifier.send(());
});
wait(notified, timeout).await;
},
)
.await
}
async fn wait(mut notified: tokio::sync::oneshot::Receiver<()>, timeout: Duration) -> bool {
if notified.try_recv().is_ok() {
return true;
}
if timeout == Duration::ZERO {
return false;
}
match tokio::time::timeout(timeout, notified).await {
Ok(Ok(())) => true,
Ok(Err(_)) => true,
Err(_) => false,
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{Arc, Mutex};
use tokio::sync::{Barrier, Semaphore, broadcast};
use crate::TestBarriers;
fn barrier() -> Arc<Barrier> {
Arc::new(Barrier::new(2))
}
#[tokio::test]
async fn async_send_recv_flush() {
let received = Arc::new(Mutex::new(0));
let (sender, receiver) = crate::bounded::<Vec<()>>(10);
let _ = spawn("test_receiver", receiver, {
let received = received.clone();
move |batch| {
let received = received.clone();
async move {
*received.lock().unwrap() += batch.len();
Ok(())
}
}
})
.unwrap();
for _ in 0..100 {
send(&sender, (), Duration::from_secs(1))
.await
.map_err(|_| "failed to send")
.unwrap();
}
flush(&sender, Duration::from_secs(1)).await;
assert_eq!(100, *received.lock().unwrap());
}
#[tokio::test]
async fn send_full_capacity() {
let received = Arc::new(Mutex::new(Vec::new()));
let post_process_barrier = barrier();
let (sender, mut receiver) = crate::bounded::<Vec<i32>>(5);
for i in 0..10 {
sender.send(i);
}
receiver.test_barriers = TestBarriers {
post_process: Some(post_process_barrier.clone()),
..Default::default()
};
let _ = spawn("test_receiver", receiver, {
let received = received.clone();
move |batch| {
let received = received.clone();
async move {
received.lock().unwrap().extend(batch);
Ok(())
}
}
})
.unwrap();
post_process_barrier.wait().await;
assert_eq!(vec![5, 6, 7, 8, 9], *received.lock().unwrap());
}
#[tokio::test]
async fn async_send_full_capacity() {
let received = Arc::new(Mutex::new(0));
let (sender, receiver) = crate::bounded::<Vec<()>>(5);
let _ = spawn("test_receiver", receiver, {
let received = received.clone();
move |batch| {
let received = received.clone();
async move {
*received.lock().unwrap() += batch.len();
Ok(())
}
}
})
.unwrap();
for _ in 0..10 {
send(&sender, (), Duration::from_secs(1)).await.unwrap();
}
flush(&sender, Duration::from_secs(1)).await;
assert_eq!(10, *received.lock().unwrap());
}
#[tokio::test]
async fn async_send_timeout() {
let (receiver_ready_tx, mut receiver_ready_rx) = broadcast::channel::<()>(100);
let blocker = Arc::new(Semaphore::new(0));
let (sender, receiver) = crate::bounded::<Vec<i32>>(5);
let receiver_task = tokio::task::spawn(async move {
exec(receiver, {
let receiver_ready_tx = receiver_ready_tx.clone();
let blocker = blocker.clone();
move |_batch| {
let receiver_ready_tx = receiver_ready_tx.clone();
let blocker = blocker.clone();
async move {
let _ = receiver_ready_tx.send(());
let _ = blocker.acquire().await;
Ok(())
}
}
})
.await
});
for i in 0..5 {
sender.send(i);
}
receiver_ready_rx.recv().await.ok();
for i in 0..5 {
sender.send(i);
}
let result = send(&sender, 99, Duration::from_millis(10)).await;
assert!(result.is_err());
receiver_task.abort();
let _ = receiver_task.await;
}
#[tokio::test]
async fn flush_empty() {
let (sender, receiver) = crate::bounded::<Vec<()>>(10);
let _ = spawn("test_receiver", receiver, |batch| async move {
let _ = batch;
Ok(())
})
.unwrap();
assert!(flush(&sender, Duration::ZERO).await);
}
#[tokio::test]
async fn flush_active() {
let batch_count = Arc::new(Mutex::new(0));
let (receiver_ready_tx, mut receiver_ready_rx) = broadcast::channel::<()>(100);
let (sender, receiver) = crate::bounded::<Vec<i32>>(10);
let _ = spawn("test_receiver", receiver, {
let batch_count = batch_count.clone();
let receiver_ready_tx = receiver_ready_tx.clone();
move |_batch| {
let batch_count = batch_count.clone();
let receiver_ready_tx = receiver_ready_tx.clone();
async move {
*batch_count.lock().unwrap() += 1;
let _ = receiver_ready_tx.send(());
Ok(())
}
}
})
.unwrap();
for i in 0..3 {
sender.send(i);
}
receiver_ready_rx.recv().await.ok();
for i in 3..6 {
sender.send(i);
}
let flushed = flush(&sender, Duration::from_secs(1)).await;
assert!(flushed);
assert_eq!(2, *batch_count.lock().unwrap());
}
#[tokio::test]
async fn retry_on_batch_failure() {
let (receiver_processed_tx, mut receiver_processed_rx) = broadcast::channel::<()>(100);
let attempt_count = Arc::new(Mutex::new(0));
let received = Arc::new(Mutex::new(false));
let (sender, receiver) = crate::bounded::<Vec<i32>>(10);
let _ = spawn("test_receiver", receiver, {
let attempt_count = attempt_count.clone();
let received = received.clone();
move |batch| {
let attempt_count = attempt_count.clone();
let received = received.clone();
let receiver_processed_tx = receiver_processed_tx.clone();
async move {
let mut count = attempt_count.lock().unwrap();
*count += 1;
if *count < 3 {
Err(BatchError::retry(
std::io::Error::new(std::io::ErrorKind::Other, "temporary failure"),
batch,
))
} else {
*received.lock().unwrap() = true;
receiver_processed_tx.send(()).unwrap();
Ok(())
}
}
}
})
.unwrap();
sender.send(42);
receiver_processed_rx.recv().await.ok();
assert!(*received.lock().unwrap());
}
#[tokio::test]
async fn processes_remaining_after_drop() {
let (receiver_processed_tx, mut receiver_processed_rx) = broadcast::channel::<()>(100);
let received = Arc::new(Mutex::new(Vec::new()));
let post_process_barrier = barrier();
let (sender, mut receiver) = crate::bounded::<Vec<i32>>(10);
receiver.test_barriers = TestBarriers {
post_process: Some(post_process_barrier.clone()),
..Default::default()
};
let _ = spawn("test_receiver", receiver, {
let received = received.clone();
move |batch| {
let received = received.clone();
let receiver_processed_tx = receiver_processed_tx.clone();
async move {
received.lock().unwrap().extend(batch);
receiver_processed_tx.send(()).unwrap();
Ok(())
}
}
})
.unwrap();
for i in 0..5 {
sender.send(i);
}
drop(sender);
post_process_barrier.wait().await;
receiver_processed_rx.recv().await.ok();
assert_eq!(vec![0, 1, 2, 3, 4], *received.lock().unwrap());
}
#[tokio::test]
async fn try_send_behavior() {
let pre_take_barrier = barrier();
let post_process_barrier = barrier();
let (sender, mut receiver) = crate::bounded::<Vec<i32>>(3);
receiver.test_barriers = TestBarriers {
pre_take: Some(pre_take_barrier.clone()),
post_process: Some(post_process_barrier.clone()),
..Default::default()
};
let _ = spawn("test_receiver", receiver, |batch| async move {
let _ = batch;
Ok(())
})
.unwrap();
sender.try_send(1).unwrap();
sender.try_send(2).unwrap();
sender.try_send(3).unwrap();
let result = sender.try_send(4);
assert!(result.is_err());
pre_take_barrier.wait().await;
post_process_barrier.wait().await;
sender.try_send(4).unwrap();
}
#[tokio::test]
async fn try_send_on_closed_channel() {
let (sender, receiver) = crate::bounded::<Vec<i32>>(10);
drop(receiver);
let result = sender.try_send(1);
assert!(result.is_err());
let err = result.err().unwrap();
assert!(err.into_retryable().is_none());
}
#[tokio::test]
async fn when_empty_callback() {
let callback_fired = Arc::new(Mutex::new(false));
let pre_take_barrier = barrier();
let post_take_barrier = barrier();
let (sender, mut receiver) = crate::bounded::<Vec<i32>>(10);
receiver.test_barriers = TestBarriers {
pre_take: Some(pre_take_barrier.clone()),
post_take: Some(post_take_barrier.clone()),
..Default::default()
};
let _ = spawn("test_receiver", receiver, |_batch| async move { Ok(()) }).unwrap();
sender.send(1);
sender.when_empty({
let callback_fired = callback_fired.clone();
move || {
*callback_fired.lock().unwrap() = true;
}
});
assert!(!*callback_fired.lock().unwrap());
pre_take_barrier.wait().await;
post_take_barrier.wait().await;
assert!(*callback_fired.lock().unwrap());
}
#[tokio::test]
async fn when_flushed_callback() {
let (when_flushed_tx, mut when_flushed_rx) = broadcast::channel::<()>(100);
let callback_fired = Arc::new(Mutex::new(false));
let post_process_barrier = barrier();
let (sender, mut receiver) = crate::bounded::<Vec<i32>>(10);
receiver.test_barriers = TestBarriers {
post_process: Some(post_process_barrier.clone()),
..Default::default()
};
let _ = spawn("test_receiver", receiver, |_batch| async move { Ok(()) }).unwrap();
sender.send(1);
let callback_fired_clone = callback_fired.clone();
sender.when_flushed(move || {
when_flushed_tx.send(()).unwrap();
*callback_fired_clone.lock().unwrap() = true;
});
assert!(!*callback_fired.lock().unwrap());
post_process_barrier.wait().await;
when_flushed_rx.recv().await.ok();
assert!(*callback_fired.lock().unwrap());
}
}