use anyhow::Result;
use async_trait::async_trait;
use flume::Sender;
use tokio::task::JoinHandle;
use tracing::{debug, error};
#[async_trait]
pub trait Producer {
type Item: Send;
async fn run(self, sender: Sender<Self::Item>);
}
#[async_trait]
pub trait Processor {
type Item: Send;
type Context: Send + Sync + Clone;
async fn process(&self, item: Self::Item, context: &Self::Context);
}
pub struct WorkDispatcher<P: Producer, C: Processor<Item = P::Item>> {
producer: P,
processor: C,
context: C::Context,
num_workers: usize,
channel_buffer: usize,
}
impl<P, C> WorkDispatcher<P, C>
where
P: Producer + Send + 'static,
P::Item: Send + 'static,
C: Processor<Item = P::Item> + Clone + Send + Sync + 'static,
{
pub fn new(producer: P, processor: C, context: C::Context) -> Self {
let default_workers = num_cpus::get().max(1); Self {
producer,
processor,
context,
num_workers: default_workers,
channel_buffer: default_workers * 100,
}
}
pub fn workers(mut self, count: usize) -> Self {
self.num_workers = count.max(1); self
}
pub fn buffer(mut self, size: usize) -> Self {
self.channel_buffer = size;
self
}
pub async fn run(self) -> Result<()> {
debug!(
"Starting WorkDispatcher with {} workers and a channel buffer of {}...",
self.num_workers, self.channel_buffer
);
let (sender, receiver) = flume::bounded::<P::Item>(self.channel_buffer);
let producer_handle = self.producer.run(sender);
let mut worker_handles = Vec::with_capacity(self.num_workers);
for i in 0..self.num_workers {
let receiver_clone = receiver.clone();
let processor_clone = self.processor.clone();
let context_clone = self.context.clone();
let handle: JoinHandle<()> = tokio::spawn(async move {
while let Ok(item) = receiver_clone.recv_async().await {
processor_clone.process(item, &context_clone).await;
}
debug!("[Worker {}] Channel closed, shutting down.", i);
});
worker_handles.push(handle);
}
producer_handle.await;
debug!("Producer has finished successfully.");
for (i, handle) in worker_handles.into_iter().enumerate() {
if let Err(e) = handle.await {
error!("[Worker {}] Panicked: {}", i, e);
}
}
debug!("All workers have finished.");
Ok(())
}
}