use core::num::NonZeroUsize;
use std::collections::HashMap;
use anyhow::Context as _;
use tokio::sync::mpsc;
use crate::bus::Wire;
use crate::bus::adapter::{AckHandle, Delivery, Source};
use crate::lifecycle::{DrainReason, Shutdown};
use crate::topology::partition_of_subject;
pub struct PartitionSource<T, H: AckHandle> {
receiver: mpsc::Receiver<Delivery<T, H>>,
}
pub type PartitionRoutes<T, H> = HashMap<u16, mpsc::Sender<Delivery<T, H>>>;
pub type PartitionSources<T, H> = HashMap<u16, PartitionSource<T, H>>;
impl<T: Wire + Send, H: AckHandle> Source<T> for PartitionSource<T, H> {
type Handle = H;
async fn next(&mut self) -> Option<anyhow::Result<Delivery<T, Self::Handle>>> {
self.receiver.recv().await.map(Ok)
}
}
pub fn partition_routes<T, H: AckHandle>(
partitions: impl IntoIterator<Item = u16>,
capacity: NonZeroUsize,
) -> (PartitionRoutes<T, H>, PartitionSources<T, H>) {
let mut senders = HashMap::new();
let mut sources = HashMap::new();
for partition in partitions {
let (sender, receiver) = mpsc::channel(capacity.get());
assert!(
senders.insert(partition, sender).is_none(),
"duplicate partition route"
);
assert!(
sources
.insert(partition, PartitionSource { receiver })
.is_none(),
"duplicate partition source"
);
}
(senders, sources)
}
pub async fn route<T, Src>(
mut source: Src,
routes: PartitionRoutes<T, Src::Handle>,
shutdown: Shutdown,
) -> anyhow::Result<()>
where
T: Wire + Send,
Src: Source<T>,
{
loop {
let next = tokio::select! {
_ = shutdown.triggered() => return Ok(()),
next = source.next() => next,
};
let Some(delivery) = next else {
if shutdown.is_triggered() {
return Ok(());
}
shutdown.trigger(DrainReason::Fatal);
anyhow::bail!("sharded source closed before shutdown");
};
let delivery = match delivery {
Ok(delivery) => delivery,
Err(error) => {
shutdown.trigger(DrainReason::Fatal);
return Err(error).context("sharded JetStream source failed");
}
};
let Some(sender) = partition_of_subject(&delivery.subject).and_then(|p| routes.get(&p))
else {
let subject = delivery.subject.clone();
let _ = delivery.handle.nak(None).await;
shutdown.trigger(DrainReason::Fatal);
anyhow::bail!("shared consumer delivered unroutable subject {subject}");
};
tokio::select! {
_ = shutdown.triggered() => {
return Ok(());
}
sent = sender.send(delivery) => match sent {
Ok(()) => {}
Err(error) => {
let subject = error.0.subject.clone();
let _ = error.0.handle.nak(None).await;
shutdown.trigger(DrainReason::Fatal);
anyhow::bail!("partition route for {subject} closed");
}
}
}
}
}
#[must_use]
pub fn raw_replay_start(frontiers: impl IntoIterator<Item = Option<u64>>) -> Option<u64> {
let mut earliest = None;
for frontier in frontiers {
let next = frontier?.saturating_add(1);
earliest = Some(earliest.map_or(next, |current: u64| current.min(next)));
}
earliest
}
#[cfg(test)]
mod tests {
use alloc::collections::VecDeque;
use alloc::sync::Arc;
use core::time::Duration;
use std::sync::Mutex;
use async_nats::HeaderMap;
use super::*;
use crate::bus::adapter::AckHandle;
#[derive(Clone, Debug, PartialEq, Eq)]
struct TestWire(u8);
impl Wire for TestWire {
fn encode(&self) -> anyhow::Result<Vec<u8>> {
Ok(vec![self.0])
}
fn decode(bytes: &[u8]) -> anyhow::Result<Self> {
Ok(Self(bytes[0]))
}
}
#[derive(Clone)]
struct TestAck {
sequence: u64,
operations: Arc<Mutex<Vec<&'static str>>>,
}
impl AckHandle for TestAck {
async fn ack(self) -> anyhow::Result<()> {
self.operations.lock().unwrap().push("ack");
Ok(())
}
async fn nak(self, _: Option<Duration>) -> anyhow::Result<()> {
self.operations.lock().unwrap().push("nak");
Ok(())
}
fn sequence(&self) -> u64 {
self.sequence
}
fn deliveries(&self) -> u32 {
1
}
}
struct TestSource {
deliveries: VecDeque<Delivery<TestWire, TestAck>>,
}
impl Source<TestWire> for TestSource {
type Handle = TestAck;
async fn next(&mut self) -> Option<anyhow::Result<Delivery<TestWire, Self::Handle>>> {
self.deliveries.pop_front().map(Ok)
}
}
fn delivery(partition: u16, value: u8, ack: TestAck) -> Delivery<TestWire, TestAck> {
Delivery {
item: TestWire(value),
handle: ack,
subject: format!("events.raw.p.{partition}"),
msg_id: None,
headers: HeaderMap::new(),
sent_at: None,
redelivered: false,
}
}
#[tokio::test]
async fn routing_is_partition_owned_and_keeps_each_queue_ordered() {
let ops = Arc::new(Mutex::new(Vec::new()));
let (routes, mut sources) = partition_routes([2, 3], NonZeroUsize::new(4).unwrap());
let source = TestSource {
deliveries: [
delivery(
2,
1,
TestAck {
sequence: 1,
operations: ops.clone(),
},
),
delivery(
3,
9,
TestAck {
sequence: 2,
operations: ops.clone(),
},
),
delivery(
2,
2,
TestAck {
sequence: 3,
operations: ops.clone(),
},
),
]
.into(),
};
let shutdown = Shutdown::new();
assert!(route(source, routes, shutdown.clone()).await.is_err());
assert_eq!(shutdown.reason(), Some(DrainReason::Fatal));
let two = sources.get_mut(&2).unwrap();
assert_eq!(two.next().await.unwrap().unwrap().item, TestWire(1));
assert_eq!(two.next().await.unwrap().unwrap().item, TestWire(2));
assert_eq!(
sources
.get_mut(&3)
.unwrap()
.next()
.await
.unwrap()
.unwrap()
.item,
TestWire(9)
);
assert!(ops.lock().unwrap().is_empty());
}
#[tokio::test]
async fn closed_owner_naks_instead_of_losing_the_delivery() {
let ops = Arc::new(Mutex::new(Vec::new()));
let (mut routes, sources) = partition_routes([2], NonZeroUsize::new(1).unwrap());
drop(sources);
let source = TestSource {
deliveries: [delivery(
2,
1,
TestAck {
sequence: 1,
operations: ops.clone(),
},
)]
.into(),
};
assert!(
route(source, routes.clone(), Shutdown::new())
.await
.is_err()
);
assert_eq!(*ops.lock().unwrap(), ["nak"]);
routes.clear();
}
#[test]
fn raw_replay_starts_at_the_earliest_member_frontier() {
assert_eq!(raw_replay_start([Some(20), Some(7), Some(30)]), Some(8));
assert_eq!(raw_replay_start([Some(u64::MAX)]), Some(u64::MAX));
assert_eq!(raw_replay_start([Some(7), None, Some(20)]), None);
assert_eq!(raw_replay_start([]), None);
}
#[tokio::test]
async fn source_closure_forces_a_fatal_drain() {
let (routes, _sources) = partition_routes([2], NonZeroUsize::new(1).unwrap());
let shutdown = Shutdown::new();
let source = TestSource {
deliveries: VecDeque::new(),
};
assert!(route(source, routes, shutdown.clone()).await.is_err());
assert_eq!(shutdown.reason(), Some(DrainReason::Fatal));
}
}