1use std::future::Future;
29use std::marker::PhantomData;
30use std::sync::Arc;
31use std::sync::Mutex;
32
33use futures::StreamExt;
34use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc};
35
36use crate::error::{Error, Result};
37use crate::event::Event;
38use crate::io::acker::NackContext;
39use crate::io::stream::SpawnedStream;
40use crate::io::{Acker, Message, Reader};
41use crate::payload::Payload;
42
43const CHANNEL_BUFFER: usize = 64;
44
45pub struct BufferEntry<C, Id, P = Payload> {
46 pub id: Id,
47 pub event: Event<P>,
48 pub cursor: C,
49}
50
51pub trait BufferStore<C, P = Payload>: Clone + Send + Sync + 'static
52where
53 P: Send + Sync,
54{
55 type Id: Clone + Send + Sync;
56
57 fn push(&self, event: &Event<P>, cursor: &C) -> impl Future<Output = Result<Self::Id>> + Send;
58
59 fn pending(&self) -> impl Future<Output = Result<Vec<BufferEntry<C, Self::Id, P>>>> + Send;
63
64 fn ack(&self, id: &Self::Id) -> impl Future<Output = Result<()>> + Send;
65
66 fn nack(&self, id: &Self::Id) -> impl Future<Output = Result<()>> + Send;
67}
68
69pub struct BufferedReaderConfig {
70 pub max_pending: usize,
71}
72
73impl Default for BufferedReaderConfig {
74 fn default() -> Self {
75 Self { max_pending: 1024 }
76 }
77}
78
79pub struct BufferAcker<S: BufferStore<C, P>, C, P = Payload>
80where
81 P: Send + Sync,
82{
83 store: S,
84 id: S::Id,
85 permit: Arc<Mutex<Option<OwnedSemaphorePermit>>>,
86 _cursor: PhantomData<fn(C, P)>,
87}
88
89impl<S, C, P> BufferAcker<S, C, P>
90where
91 S: BufferStore<C, P>,
92 P: Send + Sync,
93{
94 fn new(store: S, id: <S as BufferStore<C, P>>::Id, permit: OwnedSemaphorePermit) -> Self {
95 Self {
96 store,
97 id,
98 permit: Arc::new(Mutex::new(Some(permit))),
99 _cursor: PhantomData,
100 }
101 }
102
103 fn release_slot(&self) {
104 self.permit.lock().unwrap().take();
105 }
106}
107
108impl<S, C, P> Acker for BufferAcker<S, C, P>
109where
110 S: BufferStore<C, P> + 'static,
111 C: Send + Sync + 'static,
112 P: Send + Sync + 'static,
113{
114 async fn ack(&self) -> Result<()> {
115 self.store.ack(&self.id).await?;
116 self.release_slot();
117 Ok(())
118 }
119
120 async fn nack(&self) -> Result<()> {
121 self.store.nack(&self.id).await?;
122 self.release_slot();
123 Ok(())
124 }
125
126 async fn nack_with(&self, _context: NackContext) -> Result<()> {
127 self.store.nack(&self.id).await?;
128 self.release_slot();
129 Ok(())
130 }
131}
132
133impl<S, C, P> Clone for BufferAcker<S, C, P>
134where
135 S: BufferStore<C, P> + Clone,
136 P: Send + Sync,
137{
138 fn clone(&self) -> Self {
139 Self {
140 store: self.store.clone(),
141 id: self.id.clone(),
142 permit: Arc::clone(&self.permit),
143 _cursor: PhantomData,
144 }
145 }
146}
147
148impl<S, C, P> Drop for BufferAcker<S, C, P>
149where
150 S: BufferStore<C, P>,
151 P: Send + Sync,
152{
153 fn drop(&mut self) {
154 self.release_slot();
155 }
156}
157
158pub struct BufferedReader<R, S> {
159 inner: R,
160 store: S,
161 config: BufferedReaderConfig,
162}
163
164impl<R, S> BufferedReader<R, S> {
165 pub fn new(inner: R, store: S) -> Self {
166 Self {
167 inner,
168 store,
169 config: BufferedReaderConfig::default(),
170 }
171 }
172
173 pub fn with_config(inner: R, store: S, config: BufferedReaderConfig) -> Self {
174 Self {
175 inner,
176 store,
177 config,
178 }
179 }
180}
181
182impl<R, S, P> Reader<P> for BufferedReader<R, S>
183where
184 R: Reader<P> + Send + Sync + 'static,
185 R::Cursor: Clone + Send + Sync + 'static,
186 R::Subscription: Send + 'static,
187 R::Acker: Acker + 'static,
188 R::Stream: Send + 'static,
189 S: BufferStore<R::Cursor, P> + 'static,
190 P: Send + Sync + 'static,
191{
192 type Subscription = R::Subscription;
193 type Acker = BufferAcker<S, R::Cursor, P>;
194 type Cursor = R::Cursor;
195 type Stream = SpawnedStream<BufferAcker<S, R::Cursor, P>, R::Cursor, P>;
196
197 async fn read(&self, subscription: Self::Subscription) -> Result<Self::Stream> {
198 let store = self.store.clone();
199 let (tx, rx) = mpsc::channel::<Result<Message<BufferAcker<S, R::Cursor, P>, R::Cursor, P>>>(
200 CHANNEL_BUFFER,
201 );
202
203 let pending_entries = store.pending().await?;
204 let inner = self.inner.read(subscription).await?;
205 let semaphore = Arc::new(Semaphore::new(self.config.max_pending));
206
207 let handle = tokio::spawn(async move {
208 let mut inner = Box::pin(inner);
209
210 for entry in pending_entries {
211 let permit = match Arc::clone(&semaphore).acquire_owned().await {
212 Ok(p) => p,
213 Err(_) => return,
214 };
215 let acker = BufferAcker::new(store.clone(), entry.id, permit);
216 let msg = Message::new(entry.event, acker, entry.cursor);
217 if tx.send(Ok(msg)).await.is_err() {
218 return;
219 }
220 }
221
222 loop {
223 let permit = match Arc::clone(&semaphore).acquire_owned().await {
224 Ok(p) => p,
225 Err(_) => return,
226 };
227
228 let item = inner.next().await;
229 let msg = match item {
230 Some(Ok(m)) => m,
231 Some(Err(e)) => {
232 let _ = tx.send(Err(e)).await;
233 return;
234 }
235 None => return,
236 };
237
238 let id = match store.push(msg.event(), msg.cursor()).await {
239 Ok(id) => id,
240 Err(e) => {
241 let _ = tx.send(Err(e)).await;
242 return;
243 }
244 };
245
246 let (event, inner_acker, cursor) = msg.into_parts();
247
248 if let Err(e) = inner_acker.ack().await {
249 let _ = tx
250 .send(Err(Error::Store(format!(
251 "buffer reader: inner ack failed: {e}"
252 ))))
253 .await;
254 return;
255 }
256
257 let acker = BufferAcker::new(store.clone(), id, permit);
258 let out = Message::new(event, acker, cursor);
259
260 if tx.send(Ok(out)).await.is_err() {
261 return;
262 }
263 }
264 });
265
266 Ok(SpawnedStream::new(rx, handle))
267 }
268}
269
270#[cfg(test)]
271mod tests {
272 use std::collections::HashMap;
273 use std::pin::Pin;
274 use std::sync::Mutex;
275 use std::sync::atomic::{AtomicUsize, Ordering};
276 use std::time::Duration;
277
278 use futures::{Stream, StreamExt, stream};
279
280 use super::*;
281 use crate::error::Error;
282 use crate::io::acker::NoopAcker;
283 use crate::io::{Message, Reader};
284 use crate::payload::Payload;
285
286 fn ev(key: &str) -> Event {
287 Event::builder(
288 "acme",
289 "/x",
290 "thing.happened",
291 key,
292 Payload::from_string("p"),
293 )
294 .unwrap()
295 .build()
296 .expect("valid event")
297 }
298
299 #[derive(Debug, Clone, Copy, Eq, PartialEq, Hash)]
300 struct TestId(u64);
301
302 struct TestState<C> {
303 entries: HashMap<TestId, (Event, C)>,
304 next_id: TestId,
305 }
306
307 #[derive(Clone)]
308 struct TestBufferStore<C> {
309 state: Arc<Mutex<TestState<C>>>,
310 }
311
312 impl<C: Clone + Send + Sync + 'static> TestBufferStore<C> {
313 fn new() -> Self {
314 Self {
315 state: Arc::new(Mutex::new(TestState {
316 entries: HashMap::new(),
317 next_id: TestId(0),
318 })),
319 }
320 }
321
322 fn pending_count(&self) -> usize {
323 self.state.lock().unwrap().entries.len()
324 }
325 }
326
327 impl<C: Clone + Send + Sync + 'static> BufferStore<C> for TestBufferStore<C> {
328 type Id = TestId;
329
330 async fn push(&self, event: &Event, cursor: &C) -> Result<Self::Id> {
331 let mut state = self.state.lock().unwrap();
332 let id = state.next_id;
333 state.next_id = TestId(id.0 + 1);
334 state.entries.insert(id, (event.clone(), cursor.clone()));
335 Ok(id)
336 }
337
338 async fn pending(&self) -> Result<Vec<BufferEntry<C, Self::Id>>> {
339 let state = self.state.lock().unwrap();
340 let mut entries: Vec<BufferEntry<C, Self::Id>> = state
341 .entries
342 .iter()
343 .map(|(id, (e, c))| BufferEntry {
344 id: *id,
345 event: e.clone(),
346 cursor: c.clone(),
347 })
348 .collect();
349 entries.sort_by_key(|e| e.id.0);
350 Ok(entries)
351 }
352
353 async fn ack(&self, id: &Self::Id) -> Result<()> {
354 self.state.lock().unwrap().entries.remove(id);
355 Ok(())
356 }
357
358 async fn nack(&self, _id: &Self::Id) -> Result<()> {
359 Ok(())
360 }
361 }
362
363 #[derive(Debug, Clone, Copy, Eq, PartialEq)]
364 struct TestCursor(u64);
365
366 struct VecReader {
367 events: Mutex<Option<Vec<Event>>>,
368 }
369
370 impl Reader for VecReader {
371 type Subscription = ();
372 type Acker = NoopAcker;
373 type Cursor = TestCursor;
374 type Stream = Pin<Box<dyn Stream<Item = Result<Message<NoopAcker, TestCursor>>> + Send>>;
375
376 async fn read(&self, _: ()) -> Result<Self::Stream> {
377 let events = self.events.lock().unwrap().take().unwrap_or_default();
378 let iter = events
379 .into_iter()
380 .enumerate()
381 .map(|(i, e)| Ok(Message::new(e, NoopAcker, TestCursor(i as u64 + 1))));
382 Ok(Box::pin(stream::iter(iter)))
383 }
384 }
385
386 #[tokio::test]
387 async fn delivers_events_and_acks_store() {
388 let events: Vec<Event> = (0..3).map(|i| ev(&format!("k{i}"))).collect();
389 let store = TestBufferStore::<TestCursor>::new();
390 let reader = VecReader {
391 events: Mutex::new(Some(events)),
392 };
393 let buffered = BufferedReader::new(reader, store.clone());
394 let mut stream = buffered.read(()).await.unwrap();
395
396 for i in 0..3 {
397 let msg = tokio::time::timeout(Duration::from_secs(2), stream.next())
398 .await
399 .unwrap()
400 .unwrap()
401 .unwrap();
402 assert_eq!(msg.event().key().as_str(), &format!("k{i}"));
403 msg.ack().await.unwrap();
404 }
405
406 assert_eq!(store.pending_count(), 0);
407 }
408
409 #[tokio::test]
410 async fn drain_on_restart_replays_unacked_events() {
411 let events: Vec<Event> = (0..3).map(|i| ev(&format!("k{i}"))).collect();
412 let store = TestBufferStore::<TestCursor>::new();
413 let reader = VecReader {
414 events: Mutex::new(Some(events)),
415 };
416 let buffered = BufferedReader::new(reader, store.clone());
417 let mut stream = buffered.read(()).await.unwrap();
418
419 let msg0 = tokio::time::timeout(Duration::from_secs(2), stream.next())
420 .await
421 .unwrap()
422 .unwrap()
423 .unwrap();
424 msg0.ack().await.unwrap();
425
426 let msg1 = tokio::time::timeout(Duration::from_secs(2), stream.next())
427 .await
428 .unwrap()
429 .unwrap()
430 .unwrap();
431 msg1.nack().await.unwrap();
432
433 let msg2 = tokio::time::timeout(Duration::from_secs(2), stream.next())
434 .await
435 .unwrap()
436 .unwrap()
437 .unwrap();
438 drop(msg2);
439
440 assert!(stream.next().await.is_none());
441
442 let store2 = store.clone();
443 let reader2 = VecReader {
444 events: Mutex::new(Some(vec![])),
445 };
446 let buffered2 = BufferedReader::new(reader2, store2);
447 let mut stream2 = buffered2.read(()).await.unwrap();
448
449 let replayed1 = tokio::time::timeout(Duration::from_secs(2), stream2.next())
450 .await
451 .unwrap()
452 .unwrap()
453 .unwrap();
454 assert_eq!(replayed1.event().key().as_str(), "k1");
455 replayed1.ack().await.unwrap();
456
457 let replayed2 = tokio::time::timeout(Duration::from_secs(2), stream2.next())
458 .await
459 .unwrap()
460 .unwrap()
461 .unwrap();
462 assert_eq!(replayed2.event().key().as_str(), "k2");
463 replayed2.ack().await.unwrap();
464
465 assert_eq!(store.pending_count(), 0);
466 }
467
468 #[tokio::test]
469 async fn inner_acker_called_after_persist() {
470 #[derive(Clone, Default)]
471 struct CountingAcker {
472 count: Arc<AtomicUsize>,
473 }
474
475 impl Acker for CountingAcker {
476 async fn ack(&self) -> Result<()> {
477 self.count.fetch_add(1, Ordering::SeqCst);
478 Ok(())
479 }
480 async fn nack(&self) -> Result<()> {
481 Ok(())
482 }
483 }
484
485 struct CountingReader {
486 events: Mutex<Option<Vec<Event>>>,
487 acker: CountingAcker,
488 }
489
490 impl Reader for CountingReader {
491 type Subscription = ();
492 type Acker = CountingAcker;
493 type Cursor = TestCursor;
494 type Stream =
495 Pin<Box<dyn Stream<Item = Result<Message<CountingAcker, TestCursor>>> + Send>>;
496
497 async fn read(&self, _: ()) -> Result<Self::Stream> {
498 let events = self.events.lock().unwrap().take().unwrap_or_default();
499 let acker = self.acker.clone();
500 let iter = events.into_iter().enumerate().map(move |(i, e)| {
501 Ok(Message::new(e, acker.clone(), TestCursor(i as u64 + 1)))
502 });
503 Ok(Box::pin(stream::iter(iter)))
504 }
505 }
506
507 let acker = CountingAcker::default();
508 let store = TestBufferStore::<TestCursor>::new();
509 let reader = CountingReader {
510 events: Mutex::new(Some(vec![ev("k0")])),
511 acker: acker.clone(),
512 };
513 let buffered = BufferedReader::new(reader, store.clone());
514 let mut stream = buffered.read(()).await.unwrap();
515
516 let msg = tokio::time::timeout(Duration::from_secs(2), stream.next())
517 .await
518 .unwrap()
519 .unwrap()
520 .unwrap();
521
522 assert_eq!(acker.count.load(Ordering::SeqCst), 1);
523 assert_eq!(store.pending_count(), 1);
524
525 msg.ack().await.unwrap();
526 assert_eq!(store.pending_count(), 0);
527 }
528
529 #[tokio::test]
530 async fn inner_read_error_propagates() {
531 struct FailingReader;
532
533 impl Reader for FailingReader {
534 type Subscription = ();
535 type Acker = NoopAcker;
536 type Cursor = TestCursor;
537 type Stream =
538 Pin<Box<dyn Stream<Item = Result<Message<NoopAcker, TestCursor>>> + Send>>;
539
540 async fn read(&self, _: ()) -> Result<Self::Stream> {
541 Err(Error::Store("read failed".into()))
542 }
543 }
544
545 let store = TestBufferStore::<TestCursor>::new();
546 let reader = FailingReader;
547 let buffered = BufferedReader::new(reader, store);
548 let result = buffered.read(()).await;
549 assert!(result.is_err());
550 }
551
552 #[tokio::test]
553 async fn backpressure_blocks_intake_when_buffer_full() {
554 let events: Vec<Event> = (0..5).map(|i| ev(&format!("k{i}"))).collect();
555 let store = TestBufferStore::<TestCursor>::new();
556 let reader = VecReader {
557 events: Mutex::new(Some(events)),
558 };
559 let buffered =
560 BufferedReader::with_config(reader, store, BufferedReaderConfig { max_pending: 2 });
561 let mut stream = buffered.read(()).await.unwrap();
562
563 let msg0 = tokio::time::timeout(Duration::from_secs(2), stream.next())
564 .await
565 .unwrap()
566 .unwrap()
567 .unwrap();
568
569 let _msg1 = tokio::time::timeout(Duration::from_secs(2), stream.next())
570 .await
571 .unwrap()
572 .unwrap()
573 .unwrap();
574
575 let blocked = tokio::time::timeout(Duration::from_millis(200), stream.next()).await;
576 assert!(
577 blocked.is_err(),
578 "3rd message should be blocked by max_pending=2"
579 );
580
581 msg0.ack().await.unwrap();
582
583 let msg2 = tokio::time::timeout(Duration::from_secs(2), stream.next())
584 .await
585 .unwrap()
586 .unwrap()
587 .unwrap();
588 assert_eq!(msg2.event().key().as_str(), "k2");
589 }
590}