use std::num::NonZeroUsize;
use std::time::Duration;
use futures::{Stream, StreamExt};
use tokio::time::sleep;
use crate::{BatchSubscriber, ConnectedBroker, Seekable, Subscriber, SubscriptionSource};
const DEFAULT_MAX_WAIT: Duration = Duration::from_millis(10);
#[derive(Debug, Clone)]
pub struct Buffered<S> {
source: S,
max_wait: Duration,
}
impl<S> Buffered<S> {
#[must_use]
pub fn new(source: S) -> Self {
Self {
source,
max_wait: DEFAULT_MAX_WAIT,
}
}
#[must_use]
pub fn max_wait(mut self, max_wait: Duration) -> Self {
self.max_wait = max_wait;
self
}
}
impl<C, S> SubscriptionSource<C> for Buffered<S>
where
C: ConnectedBroker,
S: SubscriptionSource<C> + Send,
S::Subscriber: Send,
{
type Subscriber = BufferedSubscriber<S::Subscriber>;
fn name(&self) -> &str {
self.source.name()
}
async fn subscribe(self, connected: &C) -> Result<Self::Subscriber, C::Error> {
Ok(BufferedSubscriber {
inner: self.source.subscribe(connected).await?,
max_wait: self.max_wait,
})
}
}
#[derive(Debug)]
pub struct BufferedSubscriber<S> {
inner: S,
max_wait: Duration,
}
impl<S> BufferedSubscriber<S> {
#[must_use]
pub fn new(inner: S) -> Self {
Self {
inner,
max_wait: DEFAULT_MAX_WAIT,
}
}
#[must_use]
pub fn max_wait(mut self, max_wait: Duration) -> Self {
self.max_wait = max_wait;
self
}
}
impl<S: Subscriber> Subscriber for BufferedSubscriber<S> {
type Message = S::Message;
type Error = S::Error;
fn stream(&mut self) -> impl Stream<Item = Result<Self::Message, Self::Error>> + Send + '_ {
self.inner.stream()
}
}
impl<S: Seekable> Seekable for BufferedSubscriber<S> {
type Seeker = S::Seeker;
fn seeker(&self) -> S::Seeker {
self.inner.seeker()
}
}
enum Carry<E> {
Nothing,
Error(E),
Ended,
}
impl<S: Subscriber> BatchSubscriber for BufferedSubscriber<S> {
type Batch = Vec<S::Message>;
fn batches(
&mut self,
size: NonZeroUsize,
) -> impl Stream<Item = Result<Self::Batch, <Self as Subscriber>::Error>> + Send + '_ {
let max_size = size.get();
let max_wait = self.max_wait;
let inner = Box::pin(self.inner.stream());
futures::stream::unfold(
(inner, Carry::Nothing),
move |(mut stream, carry)| async move {
match carry {
Carry::Error(err) => return Some((Err(err), (stream, Carry::Nothing))),
Carry::Ended => return None,
Carry::Nothing => {}
}
let first = match stream.next().await? {
Ok(msg) => msg,
Err(err) => return Some((Err(err), (stream, Carry::Nothing))),
};
let mut batch = Vec::with_capacity(max_size.min(64));
batch.push(first);
let mut carry = Carry::Nothing;
if max_size > 1 {
let deadline = sleep(max_wait);
tokio::pin!(deadline);
loop {
tokio::select! {
() = &mut deadline => break,
next = stream.next() => match next {
Some(Ok(msg)) => {
batch.push(msg);
if batch.len() >= max_size {
break;
}
}
Some(Err(err)) => {
carry = Carry::Error(err);
break;
}
None => {
carry = Carry::Ended;
break;
}
}
}
}
}
Some((Ok(batch), (stream, carry)))
},
)
}
}
#[cfg(all(test, feature = "memory"))]
mod tests {
use std::future::ready;
use futures::StreamExt;
use super::*;
use crate::memory::{MemoryBroker, MemorySubscriber};
use crate::{Broker, IncomingMessage, Name, OutgoingMessage, Publisher};
fn batch(size: usize) -> NonZeroUsize {
NonZeroUsize::new(size).expect("test sizes are nonzero")
}
async fn buffered(
broker: &MemoryBroker,
max_wait: Duration,
) -> BufferedSubscriber<MemorySubscriber> {
let connected = broker
.clone()
.connect()
.await
.expect("memory connect is infallible");
Buffered::new(Name::new("buffered"))
.max_wait(max_wait)
.subscribe(&connected)
.await
.unwrap()
}
#[derive(Debug, thiserror::Error)]
#[error("subscriber stream failed")]
struct StreamFault;
struct Frame(Vec<u8>);
impl IncomingMessage for Frame {
fn payload(&self) -> &[u8] {
&self.0
}
fn headers(&self) -> &crate::HeaderMap {
static EMPTY: std::sync::LazyLock<crate::HeaderMap> =
std::sync::LazyLock::new(crate::HeaderMap::new);
&EMPTY
}
fn ack(self) -> impl Future<Output = Result<(), crate::AckError>> {
ready(Ok(()))
}
fn nack(self, _requeue: bool) -> impl Future<Output = Result<(), crate::AckError>> {
ready(Ok(()))
}
}
struct ScriptedSubscriber(Vec<Result<Frame, StreamFault>>);
impl Subscriber for ScriptedSubscriber {
type Message = Frame;
type Error = StreamFault;
fn stream(&mut self) -> impl Stream<Item = Result<Self::Message, Self::Error>> + Send + '_ {
futures::stream::iter(std::mem::take(&mut self.0))
}
}
fn scripted(script: Vec<Result<Frame, StreamFault>>) -> BufferedSubscriber<ScriptedSubscriber> {
BufferedSubscriber::new(ScriptedSubscriber(script)).max_wait(Duration::from_secs(60))
}
#[tokio::test(start_paused = true)]
async fn a_fault_mid_batch_is_carried_until_after_the_batch_it_interrupted() {
let mut sub = scripted(vec![
Ok(Frame(b"a".to_vec())),
Err(StreamFault),
Ok(Frame(b"b".to_vec())),
]);
let mut stream = std::pin::pin!(sub.batches(batch(8)));
let first = stream.next().await.unwrap().unwrap();
assert_eq!(first.len(), 1);
assert_eq!(first[0].payload(), b"a");
assert!(stream.next().await.unwrap().is_err());
let resumed = stream.next().await.unwrap().unwrap();
assert_eq!(resumed[0].payload(), b"b");
assert!(stream.next().await.is_none());
}
#[tokio::test(start_paused = true)]
async fn a_fault_as_the_first_item_is_yielded_without_a_batch() {
let mut sub = scripted(vec![Err(StreamFault), Ok(Frame(b"a".to_vec()))]);
let mut stream = std::pin::pin!(sub.batches(batch(8)));
assert!(stream.next().await.unwrap().is_err());
let recovered = stream.next().await.unwrap().unwrap();
assert_eq!(recovered[0].payload(), b"a");
}
#[tokio::test(start_paused = true)]
async fn a_stream_ending_mid_batch_flushes_it_before_terminating() {
let mut sub = scripted(vec![Ok(Frame(b"a".to_vec())), Ok(Frame(b"b".to_vec()))]);
let mut stream = std::pin::pin!(sub.batches(batch(8)));
let flushed = stream.next().await.unwrap().unwrap();
assert_eq!(flushed.len(), 2);
assert!(stream.next().await.is_none());
}
#[tokio::test]
async fn the_batch_size_closes_a_full_batch() {
let broker = MemoryBroker::new();
let mut sub = buffered(&broker, Duration::from_secs(60)).await;
let publisher = broker.publisher();
for i in 0..4u8 {
publisher
.publish(OutgoingMessage::new("buffered", &[i]))
.await
.unwrap();
}
let mut stream = std::pin::pin!(sub.batches(batch(2)));
let first = stream.next().await.unwrap().unwrap();
assert_eq!(first.len(), 2);
let second = stream.next().await.unwrap().unwrap();
assert_eq!(second.len(), 2);
for msg in first.into_iter().chain(second) {
msg.ack().await.unwrap();
}
}
#[tokio::test(start_paused = true)]
async fn deadline_flushes_a_partial_batch() {
let broker = MemoryBroker::new();
let mut sub = buffered(&broker, Duration::from_millis(10)).await;
let publisher = broker.publisher();
publisher
.publish(OutgoingMessage::new("buffered", b"only"))
.await
.unwrap();
let mut stream = std::pin::pin!(sub.batches(batch(64)));
let batch = stream.next().await.unwrap().unwrap();
assert_eq!(batch.len(), 1);
assert_eq!(batch[0].payload(), b"only");
for msg in batch {
msg.ack().await.unwrap();
}
}
#[tokio::test]
async fn the_seeker_reaches_through_the_buffer() {
use crate::memory::MemoryPosition;
use crate::{Seekable, Seeker};
let broker = MemoryBroker::new();
let publisher = broker.publisher();
for i in 0..2u8 {
publisher
.publish(OutgoingMessage::new("buffered", &[i]))
.await
.unwrap();
}
let mut sub = buffered(&broker, Duration::from_millis(10)).await;
sub.seeker()
.seek(MemoryPosition::start())
.await
.expect("the in-memory log replays from the start");
let mut stream = std::pin::pin!(sub.batches(batch(8)));
let replayed = stream.next().await.unwrap().unwrap();
let payloads: Vec<&[u8]> = replayed.iter().map(IncomingMessage::payload).collect();
assert_eq!(payloads, [[0].as_slice(), [1].as_slice()]);
for msg in replayed {
msg.ack().await.unwrap();
}
}
#[tokio::test]
async fn plain_stream_passes_through() {
let broker = MemoryBroker::new();
let mut sub = buffered(&broker, Duration::from_millis(10)).await;
let publisher = broker.publisher();
publisher
.publish(OutgoingMessage::new("buffered", b"single"))
.await
.unwrap();
let mut stream = std::pin::pin!(sub.stream());
let msg = stream.next().await.unwrap().unwrap();
assert_eq!(msg.payload(), b"single");
msg.ack().await.unwrap();
}
}