use std::collections::{BTreeSet, HashMap};
use std::sync::{Arc, Mutex};
use rdkafka::consumer::{BaseConsumer, ConsumerContext, Rebalance};
use rdkafka::{ClientContext, TopicPartitionList};
use tokio::sync::Notify;
use tokio::sync::futures::Notified;
#[derive(Debug)]
struct PartitionState {
outstanding: BTreeSet<i64>,
highest: i64,
stored: Option<i64>,
}
impl PartitionState {
fn starting_at(offset: i64) -> Self {
Self {
outstanding: BTreeSet::from([offset]),
highest: offset,
stored: None,
}
}
}
#[derive(Debug, Default)]
pub(crate) struct CommitTracker {
partitions: Mutex<HashMap<(String, i32), PartitionState>>,
advanced: Notify,
}
impl CommitTracker {
pub(crate) fn delivered(&self, topic: &str, partition: i32, offset: i64) {
let mut partitions = self
.partitions
.lock()
.expect("commit tracker mutex poisoned");
partitions
.entry((topic.to_owned(), partition))
.and_modify(|state| {
if offset <= state.highest {
*state = PartitionState::starting_at(offset);
} else {
state.highest = offset;
state.outstanding.insert(offset);
}
})
.or_insert_with(|| PartitionState::starting_at(offset));
}
#[allow(clippy::significant_drop_tightening)]
pub(crate) fn settle_with<E>(
&self,
topic: &str,
partition: i32,
offset: i64,
store: impl FnOnce(i64) -> Result<(), E>,
) -> Result<(), E> {
let mut partitions = self
.partitions
.lock()
.expect("commit tracker mutex poisoned");
let Some(state) = partitions.get_mut(&(topic.to_owned(), partition)) else {
return Ok(());
};
if !state.outstanding.remove(&offset) {
return Ok(());
}
let position = state
.outstanding
.first()
.map_or(state.highest, |lowest| lowest - 1);
if position < 0 || state.stored.is_some_and(|stored| position <= stored) {
return Ok(());
}
store(position)?;
state.stored = Some(position);
self.advanced.notify_waiters();
Ok(())
}
pub(crate) fn stored_position(&self, topic: &str, partition: i32) -> Option<i64> {
let partitions = self
.partitions
.lock()
.expect("commit tracker mutex poisoned");
partitions
.get(&(topic.to_owned(), partition))
.and_then(|state| state.stored)
}
pub(crate) fn covers(&self, topic: &str, partition: i32) -> bool {
let partitions = self
.partitions
.lock()
.expect("commit tracker mutex poisoned");
partitions.contains_key(&(topic.to_owned(), partition))
}
pub(crate) fn stored_positions(&self) -> Vec<((String, i32), i64)> {
let partitions = self
.partitions
.lock()
.expect("commit tracker mutex poisoned");
partitions
.iter()
.filter_map(|(key, state)| state.stored.map(|stored| (key.clone(), stored)))
.collect()
}
pub(crate) fn advance_waiter(&self) -> Notified<'_> {
self.advanced.notified()
}
fn clear(&self, revoked: &TopicPartitionList) {
let mut partitions = self
.partitions
.lock()
.expect("commit tracker mutex poisoned");
for element in revoked.elements() {
partitions.remove(&(element.topic().to_owned(), element.partition()));
}
}
}
pub(crate) struct TrackingContext {
tracker: Arc<CommitTracker>,
}
impl TrackingContext {
pub(crate) fn new(tracker: Arc<CommitTracker>) -> Self {
Self { tracker }
}
}
impl ClientContext for TrackingContext {}
impl ConsumerContext for TrackingContext {
fn pre_rebalance(&self, _consumer: &BaseConsumer<Self>, rebalance: &Rebalance<'_>) {
if let Rebalance::Revoke(revoked) = rebalance {
self.tracker.clear(revoked);
}
}
}
#[cfg(test)]
mod tests {
use std::convert::Infallible;
use super::*;
fn settle(tracker: &CommitTracker, offset: i64) -> Option<i64> {
let mut stored = None;
tracker
.settle_with("t", 0, offset, |position| {
stored = Some(position);
Ok::<(), Infallible>(())
})
.expect("infallible");
stored
}
#[test]
fn contiguous_acks_advance_the_position() {
let tracker = CommitTracker::default();
tracker.delivered("t", 0, 5);
tracker.delivered("t", 0, 6);
assert_eq!(settle(&tracker, 5), Some(5));
assert_eq!(settle(&tracker, 6), Some(6));
}
#[test]
fn offset_gaps_never_block_the_position() {
let tracker = CommitTracker::default();
tracker.delivered("t", 0, 0);
tracker.delivered("t", 0, 1);
tracker.delivered("t", 0, 3);
assert_eq!(settle(&tracker, 0), Some(0));
assert_eq!(settle(&tracker, 1), Some(2));
assert_eq!(settle(&tracker, 3), Some(3));
}
#[test]
fn out_of_order_acks_stay_bounded_by_the_lowest_outstanding() {
let tracker = CommitTracker::default();
for offset in 3..=5 {
tracker.delivered("t", 0, offset);
}
assert_eq!(settle(&tracker, 4), Some(2));
assert_eq!(settle(&tracker, 5), None);
assert_eq!(settle(&tracker, 3), Some(5));
}
#[test]
fn unsettled_delivery_blocks_the_position() {
let tracker = CommitTracker::default();
for offset in 0..3 {
tracker.delivered("t", 0, offset);
}
assert_eq!(settle(&tracker, 1), None);
assert_eq!(settle(&tracker, 2), None);
}
#[test]
fn partitions_are_tracked_independently() {
let tracker = CommitTracker::default();
tracker.delivered("t", 0, 10);
tracker.delivered("t", 1, 20);
let mut stored = None;
tracker
.settle_with("t", 1, 20, |position| {
stored = Some(position);
Ok::<(), Infallible>(())
})
.expect("infallible");
assert_eq!(stored, Some(20));
assert_eq!(settle(&tracker, 10), Some(10));
}
#[test]
fn replay_resets_the_partition_state() {
let tracker = CommitTracker::default();
tracker.delivered("t", 0, 10);
assert_eq!(settle(&tracker, 10), Some(10));
tracker.delivered("t", 0, 4);
assert_eq!(settle(&tracker, 4), Some(4));
}
#[test]
fn duplicate_settles_are_ignored() {
let tracker = CommitTracker::default();
tracker.delivered("t", 0, 0);
tracker.delivered("t", 0, 1);
assert_eq!(settle(&tracker, 0), Some(0));
assert_eq!(settle(&tracker, 0), None);
assert_eq!(settle(&tracker, 1), Some(1));
}
#[test]
fn store_failure_is_retried_by_the_next_settle() {
let tracker = CommitTracker::default();
tracker.delivered("t", 0, 0);
tracker.delivered("t", 0, 1);
let failed: Result<(), &str> = tracker.settle_with("t", 0, 0, |_| Err("store failed"));
assert!(failed.is_err());
assert_eq!(settle(&tracker, 1), Some(1));
}
}