use super::ack::AckTx;
use super::tracker::{PartitionTracker, ResolveOutcome};
use super::{AckMsg, AckRef, BatchId};
use crate::error::FatalError;
use crate::record::PartitionId;
use std::collections::{HashMap, HashSet};
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use std::time::Instant;
#[derive(Clone, Copy, Debug)]
struct Registration {
id: BatchId,
last_offset: i64,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct DrainStats {
pub applied: usize,
pub stale_epoch: usize,
pub duplicates: usize,
pub unknown: usize,
}
#[derive(Debug)]
pub struct AckIssuer {
ack_tx: crossbeam_channel::Sender<AckMsg>,
reg_tx: crossbeam_channel::Sender<Registration>,
shared_epoch: Arc<AtomicU32>,
local_epoch: u32,
seqs: HashMap<PartitionId, u64>,
}
impl Clone for AckIssuer {
fn clone(&self) -> Self {
AckIssuer {
ack_tx: self.ack_tx.clone(),
reg_tx: self.reg_tx.clone(),
shared_epoch: Arc::clone(&self.shared_epoch),
local_epoch: self.local_epoch,
seqs: HashMap::new(),
}
}
}
impl AckIssuer {
pub fn issue(&mut self, partition: PartitionId, last_offset: i64) -> AckRef {
let epoch = self.shared_epoch.load(Ordering::Acquire);
if epoch != self.local_epoch {
self.local_epoch = epoch;
self.seqs.clear();
}
let seq_slot = self.seqs.entry(partition).or_insert(0);
let seq = *seq_slot;
*seq_slot += 1;
let id = BatchId {
partition,
epoch,
seq,
};
let _ = self.reg_tx.send(Registration { id, last_offset });
AckRef::new(id, last_offset, AckTx::Channel(self.ack_tx.clone()))
}
}
#[derive(Debug)]
pub struct Checkpointer {
ack_tx: crossbeam_channel::Sender<AckMsg>,
ack_rx: crossbeam_channel::Receiver<AckMsg>,
reg_tx: crossbeam_channel::Sender<Registration>,
reg_rx: crossbeam_channel::Receiver<Registration>,
shared_epoch: Arc<AtomicU32>,
epoch: u32,
trackers: HashMap<PartitionId, PartitionTracker>,
admitted: HashSet<PartitionId>,
}
impl Default for Checkpointer {
fn default() -> Self {
Self::new()
}
}
impl Checkpointer {
#[must_use]
pub fn new() -> Self {
let (ack_tx, ack_rx) = crossbeam_channel::unbounded();
let (reg_tx, reg_rx) = crossbeam_channel::unbounded();
Checkpointer {
ack_tx,
ack_rx,
reg_tx,
reg_rx,
shared_epoch: Arc::new(AtomicU32::new(0)),
epoch: 0,
trackers: HashMap::new(),
admitted: HashSet::new(),
}
}
#[must_use]
pub fn handle(&self) -> AckIssuer {
AckIssuer {
ack_tx: self.ack_tx.clone(),
reg_tx: self.reg_tx.clone(),
shared_epoch: Arc::clone(&self.shared_epoch),
local_epoch: self.shared_epoch.load(Ordering::Acquire),
seqs: HashMap::new(),
}
}
pub fn begin_epoch(&mut self, partitions: &[PartitionId], epoch: u32) {
assert!(
epoch > self.epoch || (self.epoch == 0 && self.trackers.is_empty()),
"assignment epochs must be strictly increasing: {} -> {epoch}",
self.epoch
);
self.epoch = epoch;
self.trackers = partitions
.iter()
.map(|&p| (p, PartitionTracker::new()))
.collect();
self.admitted = partitions.iter().copied().collect();
self.shared_epoch.store(epoch, Ordering::Release);
}
pub fn extend_epoch(&mut self, partitions: &[PartitionId]) -> Result<(), FatalError> {
for &p in partitions {
if self.admitted.contains(&p) {
let live = if self.trackers.contains_key(&p) {
"is already tracked"
} else {
"was revoked earlier in this epoch"
};
return Err(FatalError {
component: "checkpoint".into(),
reason: format!(
"additive assignment reused partition {} which {live}; every \
added lane must carry a partition never seen in this epoch \
(a returning partition needs a new epoch)",
p.0
),
});
}
}
for &p in partitions {
self.trackers.insert(p, PartitionTracker::new());
self.admitted.insert(p);
}
Ok(())
}
pub fn revoke(&mut self, partitions: &[PartitionId]) {
for p in partitions {
self.trackers.remove(p);
}
}
pub fn drain(&mut self) -> DrainStats {
let mut stats = DrainStats::default();
self.drain_registrations(&mut stats);
let mut deferred = Vec::new();
while let Ok(msg) = self.ack_rx.try_recv() {
self.apply(msg, &mut stats, Some(&mut deferred));
}
if !deferred.is_empty() {
self.drain_registrations(&mut stats);
for msg in deferred {
self.apply(msg, &mut stats, None);
}
}
stats
}
fn drain_registrations(&mut self, stats: &mut DrainStats) {
while let Ok(reg) = self.reg_rx.try_recv() {
if reg.id.epoch != self.epoch {
stats.stale_epoch += 1;
continue;
}
match self.trackers.get_mut(®.id.partition) {
Some(tracker) => tracker.register(reg.id.seq, reg.last_offset),
None => stats.stale_epoch += 1,
}
}
}
fn apply(&mut self, msg: AckMsg, stats: &mut DrainStats, defer: Option<&mut Vec<AckMsg>>) {
if msg.id.epoch != self.epoch {
stats.stale_epoch += 1;
return;
}
let Some(tracker) = self.trackers.get_mut(&msg.id.partition) else {
stats.stale_epoch += 1;
return;
};
match tracker.resolve(msg.id.seq, msg.status) {
ResolveOutcome::Applied => stats.applied += 1,
ResolveOutcome::Duplicate | ResolveOutcome::AlreadyAdvanced => stats.duplicates += 1,
ResolveOutcome::Unregistered => match defer {
Some(deferred) => deferred.push(msg),
None => {
debug_assert!(false, "resolution without registration: {:?}", msg.id);
stats.unknown += 1;
}
},
}
}
#[must_use]
pub fn take_watermarks(&mut self) -> Vec<(PartitionId, i64)> {
let mut out: Vec<_> = self
.trackers
.iter_mut()
.filter_map(|(&p, t)| t.advance().map(|w| (p, w)))
.collect();
out.sort_unstable_by_key(|&(p, _)| p);
out
}
#[must_use]
pub fn pending(&self, partition: PartitionId) -> usize {
self.trackers
.get(&partition)
.map_or(0, PartitionTracker::pending)
}
#[must_use]
pub fn max_pending(&self) -> usize {
self.trackers
.values()
.map(PartitionTracker::pending)
.max()
.unwrap_or(0)
}
#[must_use]
pub fn stalled_partitions(&self) -> Vec<(PartitionId, Instant)> {
let mut out: Vec<_> = self
.trackers
.iter()
.filter_map(|(&p, t)| t.stalled_since().map(|since| (p, since)))
.collect();
out.sort_unstable_by_key(|&(p, _)| p);
out
}
}
#[cfg(all(test, not(loom)))]
mod tests {
use super::*;
const P0: PartitionId = PartitionId(0);
const P1: PartitionId = PartitionId(1);
fn checkpointer(partitions: &[PartitionId]) -> (Checkpointer, AckIssuer) {
let mut cp = Checkpointer::new();
cp.begin_epoch(partitions, 1);
let issuer = cp.handle();
(cp, issuer)
}
#[test]
fn issue_drain_take_happy_path() {
let (mut cp, mut issuer) = checkpointer(&[P0]);
drop(issuer.issue(P0, 99));
drop(issuer.issue(P0, 199));
let stats = cp.drain();
assert_eq!(stats.applied, 2);
assert_eq!(
stats,
DrainStats {
applied: 2,
..Default::default()
}
);
assert_eq!(cp.take_watermarks(), vec![(P0, 200)]);
}
#[test]
fn extend_epoch_adds_partitions_without_disturbing_inflight_acks() {
let (mut cp, mut issuer) = checkpointer(&[P0]);
let ack = issuer.issue(P0, 99);
cp.extend_epoch(&[P1]).unwrap();
drop(ack);
drop(issuer.issue(P1, 9));
let stats = cp.drain();
assert_eq!(stats.applied, 2);
assert_eq!(stats.stale_epoch, 0);
assert_eq!(cp.take_watermarks(), vec![(P0, 100), (P1, 10)]);
}
#[test]
fn extend_epoch_rejects_a_live_partition() {
let (mut cp, _issuer) = checkpointer(&[P0]);
let err = cp.extend_epoch(&[P0]).unwrap_err();
assert!(err.reason.contains("already tracked"), "{err}");
}
#[test]
fn extend_epoch_rejects_a_partition_revoked_earlier_in_the_epoch() {
let (mut cp, mut issuer) = checkpointer(&[P0]);
drop(issuer.issue(P0, 9));
cp.drain();
cp.revoke(&[P0]);
let err = cp.extend_epoch(&[P0]).unwrap_err();
assert_eq!(err.component, "checkpoint");
assert!(
err.reason.contains("revoked earlier in this epoch"),
"{err}"
);
drop(issuer.issue(P0, 19));
let stats = cp.drain();
assert_eq!(stats.applied, 0);
assert_eq!(stats.stale_epoch, 2);
cp.begin_epoch(&[P0], 2);
cp.extend_epoch(&[P1]).unwrap();
}
#[test]
fn take_watermarks_is_empty_until_new_progress() {
let (mut cp, mut issuer) = checkpointer(&[P0]);
drop(issuer.issue(P0, 9));
cp.drain();
assert_eq!(cp.take_watermarks(), vec![(P0, 10)]);
assert_eq!(cp.take_watermarks(), vec![]);
drop(issuer.issue(P0, 19));
cp.drain();
assert_eq!(cp.take_watermarks(), vec![(P0, 20)]);
}
#[test]
fn out_of_order_acks_across_partitions() {
let (mut cp, mut issuer) = checkpointer(&[P0, P1]);
let a0 = issuer.issue(P0, 9);
let a1 = issuer.issue(P0, 19);
let b0 = issuer.issue(P1, 99);
drop(a1);
drop(b0);
cp.drain();
assert_eq!(cp.take_watermarks(), vec![(P1, 100)]);
assert_eq!(cp.pending(P0), 2);
drop(a0);
cp.drain();
assert_eq!(cp.take_watermarks(), vec![(P0, 20)]);
}
#[test]
fn failed_batch_stalls_partition_and_reports() {
let (mut cp, mut issuer) = checkpointer(&[P0]);
let bad = issuer.issue(P0, 9);
bad.fail();
drop(bad);
drop(issuer.issue(P0, 19));
cp.drain();
assert_eq!(cp.take_watermarks(), vec![]);
let stalled = cp.stalled_partitions();
assert_eq!(stalled.len(), 1);
assert_eq!(stalled[0].0, P0);
}
#[test]
fn stale_epoch_acks_are_discarded() {
let (mut cp, mut issuer) = checkpointer(&[P0]);
let old = issuer.issue(P0, 9);
cp.begin_epoch(&[P0], 2);
drop(old); let stats = cp.drain();
assert_eq!(stats.applied, 0);
assert_eq!(stats.stale_epoch, 2);
assert_eq!(cp.take_watermarks(), vec![]);
drop(issuer.issue(P0, 49));
let stats = cp.drain();
assert_eq!(stats.applied, 1);
assert_eq!(cp.take_watermarks(), vec![(P0, 50)]);
}
#[test]
fn revoke_mid_flight_discards_later_acks() {
let (mut cp, mut issuer) = checkpointer(&[P0, P1]);
let in_flight = issuer.issue(P1, 9);
cp.drain(); cp.revoke(&[P1]);
drop(in_flight);
let stats = cp.drain();
assert_eq!(stats.stale_epoch, 1);
assert_eq!(cp.take_watermarks(), vec![]);
assert_eq!(cp.pending(P1), 0);
}
#[test]
fn registration_and_ack_in_same_drain() {
let (mut cp, mut issuer) = checkpointer(&[P0]);
drop(issuer.issue(P0, 9));
let stats = cp.drain();
assert_eq!(stats.applied, 1);
assert_eq!(stats.unknown, 0);
assert_eq!(cp.take_watermarks(), vec![(P0, 10)]);
}
#[test]
fn cross_thread_issue_and_resolve() {
let (mut cp, issuer) = checkpointer(&[P0, P1]);
let handles: Vec<_> = [P0, P1]
.into_iter()
.map(|p| {
let mut issuer = issuer.clone();
std::thread::spawn(move || {
for i in 0..100i64 {
drop(issuer.issue(p, (i + 1) * 10 - 1));
}
})
})
.collect();
for h in handles {
h.join().unwrap();
}
let stats = cp.drain();
assert_eq!(stats.applied, 200);
assert_eq!(stats.unknown, 0);
assert_eq!(cp.take_watermarks(), vec![(P0, 1000), (P1, 1000)]);
}
#[test]
fn pending_counts_feed_backpressure() {
let (mut cp, mut issuer) = checkpointer(&[P0, P1]);
let held: Vec<_> = (0..5).map(|i| issuer.issue(P0, i)).collect();
drop(issuer.issue(P1, 9));
cp.drain();
assert_eq!(cp.pending(P0), 5);
assert_eq!(cp.max_pending(), 5);
drop(held);
cp.drain();
let _ = cp.take_watermarks();
assert_eq!(cp.max_pending(), 0);
}
#[test]
#[should_panic(expected = "strictly increasing")]
fn epoch_regression_panics() {
let mut cp = Checkpointer::new();
cp.begin_epoch(&[P0], 5);
cp.begin_epoch(&[P0], 5);
}
}
#[cfg(all(test, not(loom)))]
mod proptests {
use super::*;
use proptest::prelude::*;
#[derive(Clone, Debug)]
enum Op {
Issue { partition: u8, fail: bool },
ResolveOldest,
Rebalance { partitions: Vec<u8> },
DrainAndTake,
}
fn ops() -> impl Strategy<Value = Vec<Op>> {
prop::collection::vec(
prop_oneof![
(0..3u8, any::<bool>()).prop_map(|(partition, fail)| Op::Issue { partition, fail }),
Just(Op::ResolveOldest),
prop::collection::vec(0..3u8, 1..3)
.prop_map(|partitions| Op::Rebalance { partitions }),
Just(Op::DrainAndTake),
],
0..120,
)
}
proptest! {
#[test]
fn epoch_churn_never_leaks_stale_acks(ops in ops()) {
let mut cp = Checkpointer::new();
let mut epoch = 1u32;
let mut assigned: Vec<PartitionId> = vec![PartitionId(0), PartitionId(1), PartitionId(2)];
cp.begin_epoch(&assigned, epoch);
let mut issuer = cp.handle();
let mut offsets: std::collections::HashMap<PartitionId, i64> =
std::collections::HashMap::new();
let mut held: std::collections::VecDeque<(AckRef, u32, bool)> =
std::collections::VecDeque::new();
let mut last_watermark: std::collections::HashMap<PartitionId, i64> =
std::collections::HashMap::new();
for op in ops {
match op {
Op::Issue { partition, fail } => {
let p = PartitionId(u32::from(partition));
if !assigned.contains(&p) {
continue;
}
let next = offsets.entry(p).or_insert(0);
*next += 10;
let ack = issuer.issue(p, *next - 1);
if fail {
ack.fail();
}
held.push_back((ack, epoch, fail));
}
Op::ResolveOldest => {
held.pop_front(); }
Op::Rebalance { partitions } => {
epoch += 1;
assigned = partitions
.into_iter()
.map(|p| PartitionId(u32::from(p)))
.collect::<std::collections::BTreeSet<_>>()
.into_iter()
.collect();
cp.begin_epoch(&assigned, epoch);
offsets.clear();
last_watermark.clear();
}
Op::DrainAndTake => {
cp.drain();
for (p, w) in cp.take_watermarks() {
prop_assert!(
assigned.contains(&p),
"watermark for unassigned partition {p:?}"
);
if let Some(&prev) = last_watermark.get(&p) {
prop_assert!(w > prev, "watermark not monotonic for {p:?}");
}
last_watermark.insert(p, w);
}
}
}
}
let stale_epochs: Vec<u32> =
held.iter().map(|&(_, e, _)| e).filter(|&e| e != epoch).collect();
held.clear();
let stats = cp.drain();
prop_assert!(stats.unknown == 0, "driver-bug resolutions: {stats:?}");
for (p, w) in cp.take_watermarks() {
prop_assert!(assigned.contains(&p));
if let Some(&prev) = last_watermark.get(&p) {
prop_assert!(w > prev);
}
}
if !stale_epochs.is_empty() {
prop_assert!(stats.stale_epoch > 0);
}
}
}
}