use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU8, AtomicUsize, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
pub enum DeliveryStatus {
Delivered = 0,
Dropped = 1,
Rejected = 2,
Errored = 3,
}
impl DeliveryStatus {
fn from_code(code: u8) -> Self {
match code {
0 => Self::Delivered,
1 => Self::Dropped,
2 => Self::Rejected,
_ => Self::Errored,
}
}
#[must_use]
pub fn should_commit(self) -> bool {
self != Self::Errored
}
}
struct Shared {
worst: AtomicU8,
outstanding: AtomicUsize,
on_finalize: Mutex<Option<Box<dyn FnOnce(DeliveryStatus) + Send>>>,
}
impl Shared {
fn merge(&self, status: DeliveryStatus) {
self.worst.fetch_max(status as u8, Ordering::AcqRel);
}
fn release(&self) {
if self.outstanding.fetch_sub(1, Ordering::AcqRel) == 1 {
let status = DeliveryStatus::from_code(self.worst.load(Ordering::Acquire));
if let Some(cb) = self
.on_finalize
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take()
{
cb(status);
}
}
}
}
pub struct BatchFinalizer {
shared: Arc<Shared>,
}
impl BatchFinalizer {
#[must_use]
pub fn new<F>(on_finalize: F) -> Self
where
F: FnOnce(DeliveryStatus) + Send + 'static,
{
Self {
shared: Arc::new(Shared {
worst: AtomicU8::new(DeliveryStatus::Delivered as u8),
outstanding: AtomicUsize::new(1),
on_finalize: Mutex::new(Some(Box::new(on_finalize))),
}),
}
}
#[must_use]
pub fn piece(&self) -> PieceFinalizer {
self.shared.outstanding.fetch_add(1, Ordering::AcqRel);
PieceFinalizer {
shared: Arc::clone(&self.shared),
reported: false,
}
}
pub fn seal(self) {
self.shared.release();
}
}
pub struct PieceFinalizer {
shared: Arc<Shared>,
reported: bool,
}
impl PieceFinalizer {
pub fn report(mut self, status: DeliveryStatus) {
self.shared.merge(status);
self.reported = true;
}
}
impl Drop for PieceFinalizer {
fn drop(&mut self) {
if !self.reported {
self.shared.merge(DeliveryStatus::Errored);
}
self.shared.release();
}
}
impl std::fmt::Debug for BatchFinalizer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BatchFinalizer")
.field(
"outstanding",
&self.shared.outstanding.load(Ordering::Relaxed),
)
.finish_non_exhaustive()
}
}
impl std::fmt::Debug for PieceFinalizer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PieceFinalizer")
.field("reported", &self.reported)
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::mpsc;
fn capturing() -> (BatchFinalizer, mpsc::Receiver<DeliveryStatus>) {
let (tx, rx) = mpsc::channel();
let fin = BatchFinalizer::new(move |status| {
tx.send(status).unwrap();
});
(fin, rx)
}
#[test]
fn status_ordering_and_commit_rule() {
assert!(DeliveryStatus::Delivered < DeliveryStatus::Dropped);
assert!(DeliveryStatus::Dropped < DeliveryStatus::Rejected);
assert!(DeliveryStatus::Rejected < DeliveryStatus::Errored);
assert!(DeliveryStatus::Delivered.should_commit());
assert!(DeliveryStatus::Dropped.should_commit());
assert!(DeliveryStatus::Rejected.should_commit());
assert!(!DeliveryStatus::Errored.should_commit());
}
#[test]
fn all_delivered_acks_delivered() {
let (fin, rx) = capturing();
let pieces: Vec<_> = (0..4).map(|_| fin.piece()).collect();
for p in pieces {
p.report(DeliveryStatus::Delivered);
}
assert!(rx.try_recv().is_err(), "ack must wait for seal");
fin.seal();
assert_eq!(rx.recv().unwrap(), DeliveryStatus::Delivered);
}
#[test]
fn one_errored_piece_blocks_commit() {
let (fin, rx) = capturing();
let p1 = fin.piece();
let p2 = fin.piece();
let p3 = fin.piece();
p1.report(DeliveryStatus::Delivered);
p2.report(DeliveryStatus::Errored); p3.report(DeliveryStatus::Delivered);
fin.seal();
let status = rx.recv().unwrap();
assert_eq!(status, DeliveryStatus::Errored, "worst wins");
assert!(
!status.should_commit(),
"an errored piece must withhold the ack"
);
}
#[test]
fn rejected_still_allows_commit() {
let (fin, rx) = capturing();
let p1 = fin.piece();
let p2 = fin.piece();
p1.report(DeliveryStatus::Delivered);
p2.report(DeliveryStatus::Rejected);
fin.seal();
let status = rx.recv().unwrap();
assert_eq!(status, DeliveryStatus::Rejected);
assert!(
status.should_commit(),
"DLQ'd record lets the source advance"
);
}
#[test]
fn dropped_piece_without_report_counts_as_errored() {
let (fin, rx) = capturing();
let p1 = fin.piece();
let p2 = fin.piece();
p1.report(DeliveryStatus::Delivered);
drop(p2); fin.seal();
let status = rx.recv().unwrap();
assert_eq!(status, DeliveryStatus::Errored, "a lost piece is Errored");
assert!(!status.should_commit(), "lost piece must block the ack");
}
#[test]
fn seal_before_pieces_resolve_defers_ack() {
let (fin, rx) = capturing();
let p = fin.piece();
fin.seal();
assert!(rx.try_recv().is_err(), "outstanding piece defers the ack");
p.report(DeliveryStatus::Delivered);
assert_eq!(rx.recv().unwrap(), DeliveryStatus::Delivered);
}
#[test]
fn empty_batch_acks_delivered_on_seal() {
let (fin, rx) = capturing();
fin.seal();
assert_eq!(rx.recv().unwrap(), DeliveryStatus::Delivered);
}
use proptest::prelude::*;
proptest! {
#[test]
fn finalizer_acks_the_worst_status(actions in prop::collection::vec(0u8..=4, 0..64)) {
let (tx, rx) = mpsc::channel();
let fin = BatchFinalizer::new(move |s| {
tx.send(s).unwrap();
});
let mut expected = DeliveryStatus::Delivered; for &code in &actions {
let p = fin.piece();
if code <= 3 {
let st = DeliveryStatus::from_code(code);
p.report(st);
expected = expected.max(st);
} else {
drop(p);
expected = expected.max(DeliveryStatus::Errored);
}
}
fin.seal();
let got = rx.recv().expect("ack must fire once");
prop_assert_eq!(got, expected);
prop_assert!(rx.try_recv().is_err(), "ack must fire exactly once");
prop_assert_eq!(got.should_commit(), expected != DeliveryStatus::Errored);
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn concurrent_pieces_fire_ack_exactly_once() {
use std::sync::atomic::AtomicUsize;
let fired = Arc::new(AtomicUsize::new(0));
let fired2 = Arc::clone(&fired);
let fin = BatchFinalizer::new(move |_status| {
fired2.fetch_add(1, Ordering::SeqCst);
});
let mut handles = Vec::new();
for i in 0..64 {
let p = fin.piece();
handles.push(tokio::spawn(async move {
let status = if i % 2 == 0 {
DeliveryStatus::Delivered
} else {
DeliveryStatus::Dropped
};
p.report(status);
}));
}
fin.seal();
for h in handles {
h.await.unwrap();
}
assert_eq!(
fired.load(Ordering::SeqCst),
1,
"the ack callback must fire exactly once"
);
}
}