use std::future::Future;
use std::pin::Pin;
use futures::channel::oneshot;
use tokio::{spawn, task::JoinHandle};
pub struct JobQueue<I, O> {
job_sender: async_channel::Sender<(I, oneshot::Sender<O>)>,
completion_sender: async_channel::Sender<Completion<O>>,
sync_receiver: async_channel::Receiver<()>,
workers: Vec<JoinHandle<()>>,
consumer: JoinHandle<()>,
}
enum Completion<O> {
Completion(oneshot::Receiver<O>),
Sync,
}
impl<I, O> JobQueue<I, O>
where
I: Send + 'static,
O: Send + 'static,
{
pub fn new(
num_workers: usize,
worker_builder: impl Fn() -> Box<dyn FnMut(I) -> Pin<Box<dyn Future<Output = O> + Send>> + Send>,
mut consumer_func: impl FnMut(O) + Send + 'static,
) -> Self {
assert_ne!(num_workers, 0);
let (job_sender, job_receiver) =
async_channel::bounded::<(I, oneshot::Sender<O>)>(num_workers);
let (completion_sender, completion_receiver) = async_channel::bounded(2 * num_workers + 1);
let (sync_sender, sync_receiver) = async_channel::bounded(1);
let workers = (0..num_workers)
.map(move |_| {
let mut worker_fn = worker_builder();
let job_receiver = job_receiver.clone();
spawn(async move {
loop {
let Ok((input, completion_sender)) = job_receiver.recv().await else {
return;
};
let result = worker_fn(input).await;
if completion_sender.send(result).is_err() {
return;
};
}
})
})
.collect();
let consumer = spawn(async move {
loop {
match completion_receiver.recv().await {
Err(_) => {
return;
}
Ok(Completion::Completion(receiver)) => {
let Ok(v) = receiver.await else {
continue;
};
consumer_func(v);
}
Ok(Completion::Sync) => {
if sync_sender.send(()).await.is_err() {
return;
}
}
}
}
});
Self {
job_sender,
completion_sender,
sync_receiver,
workers,
consumer,
}
}
pub async fn push_job(&self, job: I) {
let (completion_sender, completion_receiver) = oneshot::channel();
let _ = self.job_sender.send((job, completion_sender)).await;
let _ = self
.completion_sender
.send(Completion::Completion(completion_receiver))
.await;
}
pub async fn flush(&self) {
let _ = self.completion_sender.send(Completion::Sync).await;
let _ = self.sync_receiver.recv().await;
}
}
impl<I, O> Drop for JobQueue<I, O> {
fn drop(&mut self) {
self.consumer.abort();
for worker in self.workers.drain(..) {
worker.abort();
}
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn test_job_queue() {
let result = Arc::new(Mutex::new(Vec::new()));
let result_clone = result.clone();
let job_queue = super::JobQueue::new(
6,
|| Box::new(|i: u32| Box::pin(async move { i })),
move |i| result_clone.lock().unwrap().push(i),
);
for i in 0..1000000 {
job_queue.push_job(i).await;
}
job_queue.flush().await;
let expected: Vec<u32> = (0..1000000).collect();
assert_eq!(&*result.lock().unwrap(), &*expected);
}
}