use std::sync::Arc;
use tokio::time::Instant;
use ahash::AHashMap as HashMap;
use super::commit::CommitWaiter;
use crate::consumer::TopicPartition;
use crate::error::{ErrorCode, KrafkaError};
use crate::protocol::ShareAcknowledgementBatch;
use crate::{BrokerId, Offset};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum AckType {
Gap = 0,
Accept = 1,
Release = 2,
Reject = 3,
Renew = 4,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct AckRange {
pub first: Offset,
pub last: Offset,
pub kind: AckType,
}
impl AckRange {
pub(crate) fn one(offset: Offset, kind: AckType) -> Self {
Self {
first: offset,
last: offset,
kind,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum AckFailure {
Retry,
NotLeader,
Permanent,
}
pub(crate) fn classify(error: &KrafkaError) -> AckFailure {
match error {
KrafkaError::Network(_) | KrafkaError::Timeout { .. } => AckFailure::Retry,
KrafkaError::Broker { code, .. } => classify_code(*code),
_ => AckFailure::Permanent,
}
}
pub(crate) fn classify_code(code: ErrorCode) -> AckFailure {
if is_session_error(code) || code == ErrorCode::RequestTimedOut {
AckFailure::Retry
} else if is_leader_error(code) {
AckFailure::NotLeader
} else {
AckFailure::Permanent
}
}
pub(crate) fn is_session_error(code: ErrorCode) -> bool {
matches!(
code,
ErrorCode::ShareSessionNotFound
| ErrorCode::InvalidShareSessionEpoch
| ErrorCode::ShareSessionLimitReached
)
}
pub(crate) fn is_leader_error(code: ErrorCode) -> bool {
matches!(
code,
ErrorCode::NotLeaderForPartition
| ErrorCode::FencedLeaderEpoch
| ErrorCode::UnknownLeaderEpoch
| ErrorCode::UnknownTopicOrPartition
| ErrorCode::UnknownTopicId
)
}
#[derive(Debug)]
pub(crate) struct AckEntry {
pub topic_id: [u8; 16],
pub pending: Vec<AckRange>,
pub in_flight: Vec<AckRange>,
pub queued_at: Instant,
pub waiters: Vec<Arc<CommitWaiter>>,
}
impl AckEntry {
fn new(topic_id: [u8; 16]) -> Self {
Self {
topic_id,
pending: Vec::new(),
in_flight: Vec::new(),
queued_at: Instant::now(),
waiters: Vec::new(),
}
}
fn is_empty(&self) -> bool {
self.pending.is_empty() && self.in_flight.is_empty()
}
}
pub(crate) type AckKey = (BrokerId, TopicPartition);
#[derive(Debug)]
pub(crate) struct Resolved {
pub partition: TopicPartition,
pub ranges: Vec<AckRange>,
pub result: Result<(), KrafkaError>,
}
#[derive(Debug, Default)]
pub(crate) struct AckBook {
entries: HashMap<AckKey, AckEntry>,
}
impl AckBook {
pub(crate) fn add(
&mut self,
node: BrokerId,
partition: TopicPartition,
topic_id: [u8; 16],
range: AckRange,
) {
let entry = self
.entries
.entry((node, partition))
.or_insert_with(|| AckEntry::new(topic_id));
if entry.pending.is_empty() {
entry.queued_at = Instant::now();
}
insert_range(&mut entry.pending, range);
}
pub(crate) fn has_pending(&self, node: BrokerId) -> bool {
self.entries
.iter()
.any(|((n, _), e)| *n == node && !e.pending.is_empty())
}
pub(crate) fn nodes(&self) -> Vec<BrokerId> {
let mut nodes: Vec<BrokerId> = self.entries.keys().map(|(n, _)| *n).collect();
nodes.sort_unstable();
nodes.dedup();
nodes
}
pub(crate) fn pending_for(
&self,
node: BrokerId,
) -> impl Iterator<Item = (&TopicPartition, &AckEntry)> {
self.entries
.iter()
.filter(move |((n, _), e)| *n == node && !e.pending.is_empty())
.map(|((_, tp), e)| (tp, e))
}
pub(crate) fn take(
&mut self,
node: BrokerId,
partitions: Option<&[TopicPartition]>,
) -> Vec<(TopicPartition, [u8; 16], Vec<AckRange>)> {
let mut out = Vec::new();
for ((n, tp), entry) in &mut self.entries {
if *n != node || entry.pending.is_empty() || !entry.in_flight.is_empty() {
continue;
}
if partitions.is_some_and(|ps| !ps.contains(tp)) {
continue;
}
entry.in_flight = std::mem::take(&mut entry.pending);
out.push((tp.clone(), entry.topic_id, entry.in_flight.clone()));
}
out.sort_by(|a, b| (&a.0.topic, a.0.partition).cmp(&(&b.0.topic, b.0.partition)));
out
}
pub(crate) fn settle(
&mut self,
node: BrokerId,
partition: &TopicPartition,
result: Result<(), KrafkaError>,
retry_cutoff: Option<Instant>,
) -> Option<Resolved> {
let key = (node, partition.clone());
let entry = self.entries.get_mut(&key)?;
let in_flight = std::mem::take(&mut entry.in_flight);
if in_flight.is_empty() {
return None;
}
let result = match result {
Err(error) if classify(&error) == AckFailure::Retry => {
if retry_cutoff.is_some_and(|cutoff| entry.queued_at >= cutoff) {
let mut retried = in_flight;
for range in std::mem::take(&mut entry.pending) {
insert_range(&mut retried, range);
}
entry.pending = retried;
return None;
}
Err(KrafkaError::timeout(format!(
"acknowledging {}-{} on node {node} ({error})",
partition.topic, partition.partition
)))
}
other => other,
};
let resolved = Resolved {
partition: partition.clone(),
ranges: in_flight,
result,
};
self.finish(&key, &resolved);
Some(resolved)
}
pub(crate) fn fail_pending(
&mut self,
mut select: impl FnMut(&AckKey, &AckEntry) -> bool,
error: impl Fn(&AckKey) -> KrafkaError,
) -> Vec<Resolved> {
let keys: Vec<AckKey> = self
.entries
.iter()
.filter(|(k, e)| !e.pending.is_empty() && select(k, e))
.map(|(k, _)| k.clone())
.collect();
let mut out = Vec::new();
for key in keys {
let Some(entry) = self.entries.get_mut(&key) else {
continue;
};
let resolved = Resolved {
partition: key.1.clone(),
ranges: std::mem::take(&mut entry.pending),
result: Err(error(&key)),
};
self.finish(&key, &resolved);
out.push(resolved);
}
out
}
pub(crate) fn attach(&mut self, waiter: &Arc<CommitWaiter>) -> usize {
for entry in self.entries.values_mut() {
entry.waiters.push(Arc::clone(waiter));
}
self.entries.len()
}
pub(crate) fn detach(&mut self, waiter: &Arc<CommitWaiter>) {
for entry in self.entries.values_mut() {
entry.waiters.retain(|w| !Arc::ptr_eq(w, waiter));
}
}
pub(crate) fn waited_by(&self, waiter: &Arc<CommitWaiter>) -> Vec<TopicPartition> {
self.entries
.iter()
.filter(|(_, e)| e.waiters.iter().any(|w| Arc::ptr_eq(w, waiter)))
.map(|((_, tp), _)| tp.clone())
.collect()
}
pub(crate) fn has_waiters(&self, node: BrokerId) -> bool {
self.pending_for(node).any(|(_, e)| !e.waiters.is_empty())
}
fn finish(&mut self, key: &AckKey, resolved: &Resolved) {
let Some(entry) = self.entries.get(key) else {
return;
};
for waiter in &entry.waiters {
waiter.record(&resolved.partition, &resolved.result);
}
if entry.is_empty()
&& let Some(entry) = self.entries.remove(key)
{
for waiter in entry.waiters {
waiter.entry_done();
}
}
}
}
fn insert_range(ranges: &mut Vec<AckRange>, range: AckRange) {
if let Some(last) = ranges.last_mut() {
if last.last < range.first {
if last.kind == range.kind && last.last + 1 == range.first {
last.last = range.last;
} else {
ranges.push(range);
}
return;
}
} else {
ranges.push(range);
return;
}
let mut out = Vec::with_capacity(ranges.len() + 1);
for existing in ranges.drain(..) {
if existing.last < range.first || existing.first > range.last {
out.push(existing);
continue;
}
if existing.first < range.first {
out.push(AckRange {
last: range.first - 1,
..existing
});
}
if existing.last > range.last {
out.push(AckRange {
first: range.last + 1,
..existing
});
}
}
out.push(range);
out.sort_unstable_by_key(|r| r.first);
*ranges = out;
}
pub(crate) fn to_batches(ranges: &[AckRange]) -> Vec<ShareAcknowledgementBatch> {
let mut sorted = ranges.to_vec();
sorted.sort_unstable_by_key(|r| r.first);
let mut batches: Vec<ShareAcknowledgementBatch> = Vec::with_capacity(sorted.len());
for range in sorted {
let kind = range.kind as i8;
match batches.last_mut() {
Some(prev)
if prev.last_offset.checked_add(1) == Some(range.first)
&& prev.acknowledge_types == [kind] =>
{
prev.last_offset = range.last;
}
_ => batches.push(ShareAcknowledgementBatch {
first_offset: range.first,
last_offset: range.last,
acknowledge_types: vec![kind],
}),
}
}
batches
}
pub(crate) fn has_renew(ranges: &[AckRange]) -> bool {
ranges.iter().any(|r| r.kind == AckType::Renew)
}