use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use crate::core::Envelope;
use crate::handler_pool::HandlerTask;
use crate::handler_set::RoutingTable;
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
pub(crate) enum RouterError {
#[error("handler channel {index} is closed")]
ChannelClosed {
index: usize,
},
}
pub(crate) struct Router<Tok> {
routing_table: RoutingTable,
senders: Vec<mpsc::Sender<HandlerTask<Tok>>>,
}
impl<Tok> Clone for Router<Tok> {
fn clone(&self) -> Self { Self { routing_table: self.routing_table.clone(), senders: self.senders.clone() } }
}
impl<Tok: Clone> Router<Tok> {
pub(crate) fn new(routing_table: RoutingTable, senders: Vec<mpsc::Sender<HandlerTask<Tok>>>) -> Self {
Self { routing_table, senders }
}
pub(crate) async fn route(
&self, envelope: Envelope, token: Tok, cancellation: &CancellationToken,
) -> Result<(), RouterError> {
let Some(&mask) = self.routing_table.get(&envelope.stream) else {
tracing::warn!(stream = %envelope.stream, "no handler subscribes to stream; dropping message");
return Ok(());
};
let target_count = mask.count_ones();
if target_count == 0 {
return Ok(());
}
if target_count == 1 {
let index = mask.trailing_zeros() as usize;
return self.send(index, HandlerTask { envelope, token }, cancellation).await;
}
let mut remaining = mask;
while remaining != 0 {
let index = remaining.trailing_zeros() as usize;
remaining &= remaining - 1;
let task = HandlerTask { envelope: envelope.clone(), token: token.clone() };
self.send(index, task, cancellation).await?;
}
Ok(())
}
async fn send(
&self, index: usize, task: HandlerTask<Tok>, cancellation: &CancellationToken,
) -> Result<(), RouterError> {
let Some(sender) = self.senders.get(index) else {
tracing::error!(index, "routing table referenced a missing handler channel");
return Ok(());
};
tokio::select! {
_ = cancellation.cancelled() => Ok(()),
result = sender.send(task) => result.map_err(|_| RouterError::ChannelClosed { index }),
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use super::Router;
use crate::handler_set::RoutingTable;
use crate::test_support::make_envelope;
#[tokio::test]
async fn routes_single_and_fanout() {
let (tx0, mut rx0) = mpsc::channel(8);
let (tx1, mut rx1) = mpsc::channel(8);
let mut table = RoutingTable::new();
table.insert(Arc::from("orders"), 1 << 0);
table.insert(Arc::from("shared"), (1 << 0) | (1 << 1));
let router = Router::new(table, vec![tx0, tx1]);
let cancel = CancellationToken::new();
router.route(make_envelope("orders", "1-1"), "t1".to_owned(), &cancel).await.unwrap();
router.route(make_envelope("shared", "1-2"), "t2".to_owned(), &cancel).await.unwrap();
assert_eq!(rx0.recv().await.unwrap().envelope.id.as_ref(), "1-1");
assert_eq!(rx0.recv().await.unwrap().envelope.id.as_ref(), "1-2");
assert_eq!(rx1.recv().await.unwrap().envelope.id.as_ref(), "1-2");
}
#[tokio::test]
async fn unknown_stream_is_dropped() {
let (tx0, mut rx0) = mpsc::channel(8);
let mut table = RoutingTable::new();
table.insert(Arc::from("orders"), 1 << 0);
let router = Router::new(table, vec![tx0]);
let cancel = CancellationToken::new();
router.route(make_envelope("unknown", "1-1"), "t".to_owned(), &cancel).await.unwrap();
assert!(rx0.try_recv().is_err());
}
}