use std::{
future::Future,
io,
sync::Mutex,
thread,
time::{Duration, Instant},
};
use crate::{BatchError, Channel, Receiver, Sender, Wait, 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 = exec(receiver, on_batch);
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,
) {
let shared = receiver.shared.clone();
receiver
.exec_inner(
move |wait, delay| {
let shared = shared.clone();
async move {
match wait {
Wait::Idle => {
shared.receiver_notifier.tokio.wait_timeout(delay).await;
}
Wait::Retry => tokio::time::sleep(delay).await,
}
}
},
on_batch,
)
.await
}
pub(crate) struct Trigger(tokio::sync::Notify, Mutex<Option<bool>>);
impl Trigger {
pub fn new() -> Self {
Trigger(tokio::sync::Notify::new(), Mutex::new(None))
}
pub fn trigger(&self, value: bool) {
*self.1.lock().unwrap() = Some(value);
self.0.notify_one()
}
pub async fn wait_timeout(&self, timeout: Duration) -> bool {
let notified = self.0.notified();
tokio::pin!(notified);
notified.as_mut().enable();
match tokio::time::timeout(timeout, notified).await {
Ok(()) => self.1.lock().unwrap().take().unwrap_or(false),
Err::<(), tokio::time::error::Elapsed>(_) => {
self.1.lock().unwrap().take().unwrap_or(false)
}
}
}
}
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_inner(move |flushed| {
let _ = notifier.send(flushed);
});
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(true);
});
wait(notified, timeout).await;
},
)
.await
}
async fn wait(mut notified: tokio::sync::oneshot::Receiver<bool>, timeout: Duration) -> bool {
if let Ok(value) = notified.try_recv() {
return value;
}
if timeout == Duration::ZERO {
return false;
}
match tokio::time::timeout(timeout, notified).await {
Ok(Ok(value)) => value,
Ok(Err(_)) => false,
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_reports_failed_batch() {
let (sender, receiver) = crate::bounded::<Vec<i32>>(10);
let _ = spawn("test_receiver", receiver, |_| async {
Err(BatchError::no_retry(std::io::Error::new(
std::io::ErrorKind::Other,
"explicit failure",
)))
})
.unwrap();
sender.send(1);
assert!(!flush(&sender, Duration::from_secs(5)).await);
}
#[tokio::test]
async fn flush_wakes_idle_receiver() {
let received = Arc::new(Mutex::new(0));
let (sender, receiver) = crate::bounded(10);
let _ = spawn("test_receiver", receiver, {
let received = received.clone();
move |batch: Vec<()>| {
let received = received.clone();
async move {
*received.lock().unwrap() += batch.len();
Ok(())
}
}
})
.unwrap();
for _ in 0..3 {
tokio::time::sleep(Duration::from_millis(550)).await;
sender.send(());
assert!(flush(&sender, Duration::from_millis(200)).await);
}
assert_eq!(3, *received.lock().unwrap());
}
#[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 || {
*callback_fired_clone.lock().unwrap() = true;
when_flushed_tx.send(()).unwrap();
});
assert!(!*callback_fired.lock().unwrap());
post_process_barrier.wait().await;
when_flushed_rx.recv().await.ok();
assert!(*callback_fired.lock().unwrap());
}
}