use ahash::AHashMap;
use bytes::BufMut as _;
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::task::{Context, Poll};
use std::time::{Duration, Instant};
use tokio::sync::{Semaphore, mpsc, oneshot};
use tracing::{debug, trace, warn};
use crate::barrier::{InFlightBarrier, InFlightOpGuard};
use super::record::{
DeliveryConfirmation, ProducerRecord, RecordMetadata, RoutedRecord, TopicHandle,
};
use super::retry::{RetryContext, RetryPolicy};
use crate::PartitionId;
use crate::error::{ErrorCode, KrafkaError, ProtocolErrorKind, Result};
use crate::interceptor::ProducerInterceptor;
use crate::metadata::ClusterMetadata;
use crate::metrics::ProducerMetrics;
use crate::protocol::{
ApiKey, Compression, ProducePartitionData, ProduceRequest, ProduceResponse, ProduceTopicData,
RecordBatchBuilder, VersionedDecode, versions,
};
const MAX_CONCURRENT_BATCH_SENDS: usize = 64;
const MAX_BATCH_SPLIT_DEPTH: u8 = 2;
const PARTITION_INFLIGHT_PRUNE_THRESHOLD: usize = 1024;
const IDLE_TICK: Duration = Duration::from_secs(1);
#[derive(Debug)]
struct PartitionInFlight {
next_ticket: AtomicU64,
now_serving: AtomicU64,
advance: tokio::sync::Notify,
dispatch: Arc<tokio::sync::Notify>,
}
impl PartitionInFlight {
fn new(dispatch: Arc<tokio::sync::Notify>) -> Self {
Self {
next_ticket: AtomicU64::new(0),
now_serving: AtomicU64::new(0),
advance: tokio::sync::Notify::new(),
dispatch,
}
}
fn release(&self) {
self.now_serving.fetch_add(1, Ordering::AcqRel);
self.advance.notify_waiters();
self.dispatch.notify_one();
}
fn take_ticket(self: &Arc<Self>) -> PartitionTicket {
let ticket = self.next_ticket.fetch_add(1, Ordering::AcqRel);
PartitionTicket {
slot: Arc::clone(self),
ticket,
acquired: false,
}
}
fn is_idle(&self) -> bool {
self.now_serving.load(Ordering::Acquire) == self.next_ticket.load(Ordering::Acquire)
}
}
#[derive(Debug)]
struct PartitionTicket {
slot: Arc<PartitionInFlight>,
ticket: u64,
acquired: bool,
}
impl PartitionTicket {
async fn acquire(mut self) -> PartitionTurn {
while self.slot.now_serving.load(Ordering::Acquire) != self.ticket {
let notified = self.slot.advance.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if self.slot.now_serving.load(Ordering::Acquire) == self.ticket {
break;
}
notified.await;
}
self.acquired = true;
PartitionTurn {
slot: Arc::clone(&self.slot),
}
}
}
impl Drop for PartitionTicket {
fn drop(&mut self) {
if !self.acquired {
self.slot.release();
}
}
}
#[derive(Debug)]
struct PartitionTurn {
slot: Arc<PartitionInFlight>,
}
impl Drop for PartitionTurn {
fn drop(&mut self) {
self.slot.release();
}
}
struct ExtractedBatch {
batch: AccumulatorBatch,
guard: InFlightGuard,
ticket: PartitionTicket,
}
type ExtractedBatches = Vec<((TopicHandle, PartitionId), ExtractedBatch)>;
fn is_batch_too_large(error: &KrafkaError) -> bool {
match error {
KrafkaError::Broker { code, .. } => matches!(
code,
ErrorCode::MessageTooLarge | ErrorCode::RecordListTooLarge
),
KrafkaError::Protocol { kind, .. } => *kind == ProtocolErrorKind::FrameTooLarge,
_ => false,
}
}
#[inline]
fn is_cpu_heavy_compression(compression: Compression) -> bool {
matches!(compression, Compression::Gzip | Compression::Zstd)
}
fn max_record_semaphore_permits() -> usize {
Semaphore::MAX_PERMITS.min(u32::MAX as usize)
}
pub(crate) fn check_record_admission(
record_size: usize,
memory_capacity: usize,
max_request_size: usize,
) -> Result<()> {
let semaphore_limit = max_record_semaphore_permits();
if record_size > semaphore_limit {
return Err(KrafkaError::config(format!(
"record size {record_size} B exceeds the semaphore \
permit-count limit ({} B; min(u32::MAX, \
Semaphore::MAX_PERMITS)); Kafka records must be \
smaller",
semaphore_limit
)));
}
if max_request_size > 0 && record_size > max_request_size {
return Err(KrafkaError::config(format!(
"record size {record_size} B exceeds max_request_size \
({max_request_size} B); the broker will reject the record \
with MESSAGE_TOO_LARGE — raise ProducerConfig::max_request_size \
or shrink the record",
)));
}
if record_size > memory_capacity {
return Err(KrafkaError::config(format!(
"record size {record_size} B exceeds producer buffer_memory \
capacity ({} B); raise ProducerConfig::buffer_memory or \
shrink the record",
memory_capacity
)));
}
Ok(())
}
pub(crate) fn effective_memory_capacity(buffer_memory: usize) -> usize {
if buffer_memory > 0 {
if buffer_memory > Semaphore::MAX_PERMITS {
warn!(
requested = buffer_memory,
effective = Semaphore::MAX_PERMITS,
"buffer_memory exceeds Semaphore::MAX_PERMITS; clamping effective \
producer memory capacity"
);
Semaphore::MAX_PERMITS
} else {
buffer_memory
}
} else {
Semaphore::MAX_PERMITS
}
}
#[derive(Debug)]
pub(crate) struct BufferedRecordGuard {
buffered_records: Arc<AtomicUsize>,
metrics: Arc<ProducerMetrics>,
}
impl BufferedRecordGuard {
pub(crate) fn new(buffered_records: Arc<AtomicUsize>, metrics: Arc<ProducerMetrics>) -> Self {
buffered_records.fetch_add(1, Ordering::Relaxed);
metrics.buffered_records.inc();
Self {
buffered_records,
metrics,
}
}
}
impl Drop for BufferedRecordGuard {
fn drop(&mut self) {
self.buffered_records.fetch_sub(1, Ordering::Relaxed);
self.metrics.buffered_records.dec();
}
}
#[derive(Debug)]
#[must_use = "a dropped DeliveryHandle discards the acknowledgement; the record is still sent"]
pub struct DeliveryHandle {
response_rx: oneshot::Receiver<AppendResponse>,
partition: PartitionId,
}
impl DeliveryHandle {
#[inline]
#[must_use]
pub fn partition(&self) -> PartitionId {
self.partition
}
}
impl Future for DeliveryHandle {
type Output = Result<RecordMetadata>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
match Pin::new(&mut self.response_rx).poll(cx) {
Poll::Ready(Ok(AppendResponse::Done(result))) => Poll::Ready(result),
Poll::Ready(Err(_)) => Poll::Ready(Err(KrafkaError::invalid_state(
"accumulator response dropped",
))),
Poll::Pending => Poll::Pending,
}
}
}
#[derive(Debug)]
enum AppendResponse {
Done(Result<RecordMetadata>),
}
#[derive(Debug)]
struct AppendCommand {
topic: TopicHandle,
record: RoutedRecord,
partition: PartitionId,
record_size: usize,
send_started_at: Instant,
response_tx: oneshot::Sender<AppendResponse>,
operation_guard: InFlightOpGuard,
_buffered_record_guard: BufferedRecordGuard,
permit_reservation: PermitReservation,
}
#[derive(Debug)]
enum AccumulatorMessage {
Append(AppendCommand),
Flush {
response_tx: oneshot::Sender<Result<()>>,
},
Shutdown { response_tx: oneshot::Sender<()> },
}
struct PermitReservation {
bytes: usize,
memory_permits: Arc<Semaphore>,
}
impl PermitReservation {
fn forget(mut self) {
self.bytes = 0;
}
}
impl Drop for PermitReservation {
fn drop(&mut self) {
self.memory_permits.add_permits(self.bytes);
}
}
impl std::fmt::Debug for PermitReservation {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PermitReservation")
.field("bytes", &self.bytes)
.finish()
}
}
#[derive(Clone)]
pub struct RecordAccumulatorHandle {
sender: mpsc::Sender<AccumulatorMessage>,
memory_permits: Arc<Semaphore>,
memory_capacity: usize,
max_request_size: usize,
max_block_ms: Duration,
in_flight_barrier: Arc<InFlightBarrier>,
buffered_records: Arc<AtomicUsize>,
metrics: Arc<ProducerMetrics>,
}
impl RecordAccumulatorHandle {
pub async fn append(
&self,
record: ProducerRecord,
partition: PartitionId,
) -> Result<RecordMetadata> {
let send_started_at = Instant::now();
let operation_guard = self.in_flight_barrier.start("producer")?;
self.append_with_guard(record, partition, operation_guard, send_started_at)
.await
}
pub(crate) async fn append_with_guard(
&self,
record: ProducerRecord,
partition: PartitionId,
operation_guard: InFlightOpGuard,
send_started_at: Instant,
) -> Result<RecordMetadata> {
let record_size = record.estimated_size();
let routed = record.into_routed_parts();
self.append_routed_with_guard(
routed.topic,
routed.record,
record_size,
partition,
operation_guard,
send_started_at,
)
.await
}
pub(crate) async fn append_routed_with_guard(
&self,
topic: TopicHandle,
record: RoutedRecord,
record_size: usize,
partition: PartitionId,
operation_guard: InFlightOpGuard,
send_started_at: Instant,
) -> Result<RecordMetadata> {
self.enqueue_routed_with_guard(
topic,
record,
record_size,
partition,
operation_guard,
send_started_at,
)
.await?
.await
}
pub(crate) async fn enqueue_routed_with_guard(
&self,
topic: TopicHandle,
record: RoutedRecord,
record_size: usize,
partition: PartitionId,
operation_guard: InFlightOpGuard,
send_started_at: Instant,
) -> Result<DeliveryHandle> {
let deadline = tokio::time::Instant::now() + self.max_block_ms;
check_record_admission(record_size, self.memory_capacity, self.max_request_size)?;
let permit = match tokio::time::timeout(
deadline.saturating_duration_since(tokio::time::Instant::now()),
self.memory_permits.acquire_many(record_size as u32),
)
.await
{
Ok(Ok(p)) => p,
Ok(Err(_)) => return Err(KrafkaError::invalid_state("accumulator closed")),
Err(_) => {
return Err(KrafkaError::timeout(
"producer append: max_block exceeded while waiting for buffer \
memory (ProducerConfig::max_block / AccumulatorConfig::max_block_ms)",
));
}
};
let permit_reservation = PermitReservation {
bytes: record_size,
memory_permits: self.memory_permits.clone(),
};
permit.forget();
let (response_tx, response_rx) = oneshot::channel();
let buffered_record_guard =
BufferedRecordGuard::new(self.buffered_records.clone(), self.metrics.clone());
let remaining = deadline.saturating_duration_since(tokio::time::Instant::now());
match tokio::time::timeout(
remaining,
self.sender.send(AccumulatorMessage::Append(AppendCommand {
topic,
record,
partition,
record_size,
send_started_at,
response_tx,
operation_guard,
_buffered_record_guard: buffered_record_guard,
permit_reservation,
})),
)
.await
{
Ok(Ok(())) => {}
Ok(Err(_)) => return Err(KrafkaError::invalid_state("accumulator closed")),
Err(_) => {
return Err(KrafkaError::timeout(
"producer append: max_block exceeded while sending to accumulator",
));
}
}
Ok(DeliveryHandle {
response_rx,
partition,
})
}
pub async fn flush(&self) -> Result<()> {
let (response_tx, response_rx) = oneshot::channel();
self.sender
.send(AccumulatorMessage::Flush { response_tx })
.await
.map_err(|_| KrafkaError::invalid_state("accumulator closed"))?;
response_rx
.await
.map_err(|_| KrafkaError::invalid_state("accumulator response dropped"))?
}
pub async fn shutdown(&self) -> Result<()> {
let (response_tx, response_rx) = oneshot::channel();
self.sender
.send(AccumulatorMessage::Shutdown { response_tx })
.await
.map_err(|_| {
warn!("Accumulator shutdown failed: task already exited");
KrafkaError::invalid_state("accumulator already shut down")
})?;
response_rx.await.map_err(|_| {
warn!("Accumulator shutdown: response channel dropped before completion");
KrafkaError::invalid_state("accumulator shutdown interrupted")
})?;
Ok(())
}
}
pub struct AccumulatorConfig {
pub batch_size: usize,
pub linger: Duration,
pub compression: Compression,
pub compression_level: Option<i32>,
pub topic_compression: AHashMap<String, Compression>,
pub acks: i16,
pub client_id: String,
pub request_timeout: Duration,
pub max_request_size: usize,
pub buffer_memory: usize,
pub max_block_ms: Duration,
pub interceptor: Arc<dyn ProducerInterceptor>,
pub identity: Option<Arc<super::idempotent::ProducerIdentity>>,
pub partitioner: Arc<dyn super::partitioner::Partitioner>,
pub(crate) state_store: Option<Arc<dyn super::idempotent::ErasedProducerStateStore>>,
pub transactional_id: Option<String>,
pub(crate) dead_letter_queue: Option<Arc<dyn crate::dlq::DeadLetterQueue>>,
}
impl Clone for AccumulatorConfig {
fn clone(&self) -> Self {
Self {
batch_size: self.batch_size,
linger: self.linger,
compression: self.compression,
compression_level: self.compression_level,
topic_compression: self.topic_compression.clone(),
acks: self.acks,
client_id: self.client_id.clone(),
request_timeout: self.request_timeout,
max_request_size: self.max_request_size,
buffer_memory: self.buffer_memory,
max_block_ms: self.max_block_ms,
interceptor: self.interceptor.clone(),
identity: self.identity.clone(),
partitioner: self.partitioner.clone(),
state_store: self.state_store.clone(),
transactional_id: self.transactional_id.clone(),
dead_letter_queue: self.dead_letter_queue.clone(),
}
}
}
impl fmt::Debug for AccumulatorConfig {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("AccumulatorConfig")
.field("batch_size", &self.batch_size)
.field("linger", &self.linger)
.field("compression", &self.compression)
.field("acks", &self.acks)
.field("client_id", &self.client_id)
.field("request_timeout", &self.request_timeout)
.field("max_request_size", &self.max_request_size)
.field("buffer_memory", &self.buffer_memory)
.field("max_block_ms", &self.max_block_ms)
.field("interceptor", &self.interceptor)
.field("partitioner", &"<dyn Partitioner>")
.finish()
}
}
impl Default for AccumulatorConfig {
fn default() -> Self {
Self {
batch_size: 16384,
linger: Duration::ZERO,
compression: Compression::None,
compression_level: None,
topic_compression: AHashMap::new(),
acks: -1,
client_id: "krafka".to_string(),
request_timeout: Duration::from_secs(30),
max_request_size: crate::protocol::MAX_MESSAGE_SIZE,
buffer_memory: 32 * 1024 * 1024, max_block_ms: Duration::from_secs(60), interceptor: Arc::new(crate::interceptor::NoOpProducerInterceptor),
identity: None,
partitioner: Arc::new(super::partitioner::UniformStickyPartitioner::new()),
state_store: None,
transactional_id: None,
dead_letter_queue: None,
}
}
}
struct PendingRecord {
record: RoutedRecord,
response_tx: oneshot::Sender<AppendResponse>,
offset_in_batch: i64,
estimated_size: usize,
_buffered_record_guard: BufferedRecordGuard,
_operation_guard: InFlightOpGuard,
}
struct InFlightGuard {
bytes: usize,
in_flight_memory: Arc<AtomicUsize>,
memory_permits: Arc<Semaphore>,
}
impl Drop for InFlightGuard {
fn drop(&mut self) {
self.in_flight_memory
.fetch_sub(self.bytes, Ordering::Relaxed);
self.memory_permits.add_permits(self.bytes);
}
}
struct AccumulatorBatch {
current_size: usize,
max_size: usize,
pending: Vec<PendingRecord>,
opened_at: Instant,
deadline_from: Instant,
}
impl AccumulatorBatch {
fn new(max_size: usize) -> Self {
let now = Instant::now();
Self {
current_size: 0,
max_size,
pending: Vec::new(),
opened_at: now,
deadline_from: now,
}
}
#[inline]
fn is_empty(&self) -> bool {
self.pending.is_empty()
}
#[inline]
fn len(&self) -> usize {
self.pending.len()
}
#[inline]
fn is_full(&self) -> bool {
self.current_size >= self.max_size
}
#[inline]
fn would_fit(&self, record_size: usize) -> bool {
self.is_empty() || self.current_size + record_size <= self.max_size
}
#[inline]
fn track(&mut self, record_size: usize) {
self.current_size += record_size;
}
fn age(&self) -> Duration {
self.opened_at.elapsed()
}
fn linger_deadline(&self, linger: Duration) -> Instant {
self.opened_at + linger
}
fn charge_from(&mut self, started_at: Instant) {
if started_at < self.deadline_from {
self.deadline_from = started_at;
}
}
}
pub struct RecordAccumulator {
config: AccumulatorConfig,
batches: AHashMap<(TopicHandle, PartitionId), AccumulatorBatch>,
metadata: Arc<ClusterMetadata>,
send_semaphore: Arc<Semaphore>,
in_flight_memory: Arc<AtomicUsize>,
retry_policy: RetryPolicy,
metrics: Arc<ProducerMetrics>,
memory_permits: Arc<Semaphore>,
partitioner: Arc<dyn super::partitioner::Partitioner>,
partition_inflight: AHashMap<(TopicHandle, PartitionId), Arc<PartitionInFlight>>,
dispatch_notify: Arc<tokio::sync::Notify>,
}
impl RecordAccumulator {
pub(crate) fn spawn(
config: AccumulatorConfig,
metadata: Arc<ClusterMetadata>,
retry_policy: RetryPolicy,
metrics: Arc<ProducerMetrics>,
in_flight_barrier: Arc<InFlightBarrier>,
) -> RecordAccumulatorHandle {
let channel_capacity = if config.buffer_memory > 0 {
let batch = config.batch_size.max(1);
(config.buffer_memory / 10 / batch).clamp(1, 256)
} else {
64
};
let (sender, receiver) = mpsc::channel(channel_capacity);
let memory_capacity = effective_memory_capacity(config.buffer_memory);
let memory_permits = Arc::new(Semaphore::new(memory_capacity));
let in_flight_memory = Arc::new(AtomicUsize::new(0));
let buffered_records = Arc::new(AtomicUsize::new(0));
let handle_buffered_records = buffered_records.clone();
let handle_metrics = metrics.clone();
let max_block_ms = config.max_block_ms;
let max_request_size = config.max_request_size;
let accumulator_partitioner = config.partitioner.clone();
let accumulator = Self {
config,
batches: AHashMap::new(),
metadata,
send_semaphore: Arc::new(Semaphore::new(MAX_CONCURRENT_BATCH_SENDS)),
in_flight_memory,
retry_policy,
metrics,
memory_permits: memory_permits.clone(),
partitioner: accumulator_partitioner,
partition_inflight: AHashMap::new(),
dispatch_notify: Arc::new(tokio::sync::Notify::new()),
};
let memory_permits_panic = memory_permits.clone();
tokio::spawn(async move {
let join_handle = tokio::spawn(accumulator.run(receiver));
if let Err(join_err) = join_handle.await {
if join_err.is_panic() {
tracing::error!("Accumulator task panicked: {join_err}");
} else {
tracing::error!("Accumulator task cancelled: {join_err}");
}
memory_permits_panic.close();
}
});
RecordAccumulatorHandle {
sender,
memory_permits,
memory_capacity,
max_request_size,
max_block_ms,
in_flight_barrier,
buffered_records: handle_buffered_records,
metrics: handle_metrics,
}
}
async fn run(mut self, mut receiver: mpsc::Receiver<AccumulatorMessage>) {
let dispatch_notify = self.dispatch_notify.clone();
loop {
let wake_at = self.next_wake_deadline();
tokio::select! {
() = dispatch_notify.notified() => {
self.dispatch_unblocked_partitions();
}
msg = receiver.recv() => {
match msg {
Some(AccumulatorMessage::Append(append)) => {
self.handle_append(append);
}
Some(AccumulatorMessage::Flush { response_tx }) => {
let result = self.flush_all().await;
let _ = response_tx.send(result);
}
Some(AccumulatorMessage::Shutdown { response_tx }) => {
debug!("Accumulator shutting down, flushing remaining batches");
let _ = self.flush_all().await;
let _ = response_tx.send(());
break;
}
None => {
debug!("Accumulator channel closed, flushing remaining batches");
let _ = self.flush_all().await;
break;
}
}
}
() = tokio::time::sleep_until(wake_at.into()) => {
self.check_linger_expiry();
}
}
}
debug!("Accumulator shutdown complete");
}
fn next_wake_deadline(&self) -> Instant {
let idle = Instant::now() + IDLE_TICK;
if self.config.linger.is_zero() {
return idle;
}
self.batches
.values()
.filter(|batch| !batch.is_empty())
.map(|batch| batch.linger_deadline(self.config.linger))
.min()
.map_or(idle, |deadline| deadline.min(idle))
}
fn handle_append(&mut self, append: AppendCommand) {
let AppendCommand {
topic,
record,
partition,
record_size,
send_started_at,
response_tx,
operation_guard,
_buffered_record_guard: buffered_record_guard,
permit_reservation,
} = append;
let key = (topic, partition);
let batch_size = self.config.batch_size;
let accumulator_batch = self
.batches
.entry(key.clone())
.or_insert_with(|| AccumulatorBatch::new(batch_size));
accumulator_batch.charge_from(send_started_at);
let offset = accumulator_batch.len() as i64;
if accumulator_batch.would_fit(record_size) {
accumulator_batch.track(record_size);
accumulator_batch.pending.push(PendingRecord {
record,
response_tx,
offset_in_batch: offset,
estimated_size: record_size,
_buffered_record_guard: buffered_record_guard,
_operation_guard: operation_guard,
});
permit_reservation.forget();
if accumulator_batch.is_full() {
trace!("Batch full for {}-{}, flushing", key.0, partition);
let partition_count = self
.metadata
.partition_count(key.0.as_ref())
.unwrap_or(partition as usize + 1);
self.partitioner
.on_new_batch(key.0.as_ref(), partition, partition_count);
self.flush_batch(&key);
} else if self.config.linger.is_zero() && self.partition_is_idle(&key) {
trace!("Linger=0 for {}-{}, flushing immediately", key.0, partition);
self.flush_batch(&key);
}
} else {
let partition_count = self
.metadata
.partition_count(key.0.as_ref())
.unwrap_or(partition as usize + 1);
self.partitioner
.on_new_batch(key.0.as_ref(), partition, partition_count);
self.flush_batch(&key);
let mut new_batch = AccumulatorBatch::new(batch_size);
new_batch.charge_from(send_started_at);
new_batch.track(record_size);
new_batch.pending.push(PendingRecord {
record,
response_tx,
offset_in_batch: 0,
estimated_size: record_size,
_buffered_record_guard: buffered_record_guard,
_operation_guard: operation_guard,
});
self.batches.insert(key, new_batch);
permit_reservation.forget();
}
}
fn check_linger_expiry(&mut self) {
self.prune_idle_partition_inflight();
if self.config.linger.is_zero() {
self.dispatch_unblocked_partitions();
return;
}
let keys_to_flush: Vec<_> = self
.batches
.iter()
.filter(|(_, batch)| !batch.is_empty() && batch.age() >= self.config.linger)
.map(|(key, _)| key.clone())
.collect();
if keys_to_flush.is_empty() {
return;
}
let mut extracted = Vec::with_capacity(keys_to_flush.len());
for key in keys_to_flush {
trace!("Linger expired for {:?}, flushing", key);
if let Some(item) = self.extract_batch(&key) {
extracted.push((key, item));
}
}
Self::spawn_batches_detached(
extracted,
&self.metadata,
&self.config,
&self.retry_policy,
&self.metrics,
self.send_semaphore.clone(),
);
}
fn partition_is_idle(&self, key: &(TopicHandle, PartitionId)) -> bool {
self.partition_inflight
.get(key)
.is_none_or(|slot| slot.is_idle())
}
fn dispatch_unblocked_partitions(&mut self) {
let keys_to_flush: Vec<_> = self
.batches
.iter()
.filter(|(_, batch)| !batch.is_empty())
.map(|(key, _)| key.clone())
.filter(|key| self.partition_is_idle(key))
.collect();
if keys_to_flush.is_empty() {
return;
}
let mut extracted = Vec::with_capacity(keys_to_flush.len());
for key in keys_to_flush {
if let Some(item) = self.extract_batch(&key) {
extracted.push((key, item));
}
}
Self::spawn_batches_detached(
extracted,
&self.metadata,
&self.config,
&self.retry_policy,
&self.metrics,
self.send_semaphore.clone(),
);
}
async fn spawn_batches_bounded(
extracted: ExtractedBatches,
metadata: &Arc<ClusterMetadata>,
config: &AccumulatorConfig,
retry_policy: &RetryPolicy,
metrics: &Arc<ProducerMetrics>,
send_semaphore: Arc<Semaphore>,
) {
let mut join_set = tokio::task::JoinSet::new();
for (
(topic, partition),
ExtractedBatch {
batch,
guard,
ticket,
},
) in extracted
{
let metadata = metadata.clone();
let config = config.clone();
let retry_policy = retry_policy.clone();
let metrics = metrics.clone();
let send_semaphore = send_semaphore.clone();
join_set.spawn(async move {
Self::send_extracted_batch(
send_semaphore,
topic,
partition,
batch.pending,
batch.deadline_from,
guard,
ticket,
metadata,
config,
retry_policy,
metrics,
)
.await;
});
}
while let Some(result) = join_set.join_next().await {
if let Err(e) = result
&& e.is_panic()
{
warn!("send_extracted_batch task panicked: {e}");
}
}
}
fn spawn_batches_detached(
extracted: ExtractedBatches,
metadata: &Arc<ClusterMetadata>,
config: &AccumulatorConfig,
retry_policy: &RetryPolicy,
metrics: &Arc<ProducerMetrics>,
send_semaphore: Arc<Semaphore>,
) {
if extracted.is_empty() {
return;
}
let metadata = metadata.clone();
let config = config.clone();
let retry_policy = retry_policy.clone();
let metrics = metrics.clone();
drop(tokio::spawn(async move {
Self::spawn_batches_bounded(
extracted,
&metadata,
&config,
&retry_policy,
&metrics,
send_semaphore,
)
.await;
}));
}
fn extract_batch(&mut self, key: &(TopicHandle, PartitionId)) -> Option<ExtractedBatch> {
let batch = self.batches.remove(key)?;
if batch.is_empty() {
return None;
}
let batch_memory: usize = batch.pending.iter().map(|p| p.estimated_size).sum();
self.in_flight_memory
.fetch_add(batch_memory, Ordering::Relaxed);
let guard = InFlightGuard {
bytes: batch_memory,
in_flight_memory: self.in_flight_memory.clone(),
memory_permits: self.memory_permits.clone(),
};
let dispatch_notify = &self.dispatch_notify;
let ticket = self
.partition_inflight
.entry(key.clone())
.or_insert_with(|| Arc::new(PartitionInFlight::new(dispatch_notify.clone())))
.take_ticket();
Some(ExtractedBatch {
batch,
guard,
ticket,
})
}
fn prune_idle_partition_inflight(&mut self) {
if self.partition_inflight.len() <= PARTITION_INFLIGHT_PRUNE_THRESHOLD {
return;
}
self.partition_inflight
.retain(|_, slot| Arc::strong_count(slot) > 1 || !slot.is_idle());
}
fn flush_batch(&mut self, key: &(TopicHandle, PartitionId)) {
if let Some(item) = self.extract_batch(key) {
Self::spawn_batches_detached(
vec![(key.clone(), item)],
&self.metadata,
&self.config,
&self.retry_policy,
&self.metrics,
self.send_semaphore.clone(),
);
}
}
#[allow(clippy::too_many_arguments)]
async fn send_extracted_batch(
send_semaphore: Arc<Semaphore>,
topic: TopicHandle,
partition: PartitionId,
pending: Vec<PendingRecord>,
enqueued_at: Instant,
_in_flight_guard: InFlightGuard,
ticket: PartitionTicket,
metadata: Arc<ClusterMetadata>,
config: AccumulatorConfig,
retry_policy: RetryPolicy,
metrics: Arc<ProducerMetrics>,
) {
let _turn = ticket.acquire().await;
let Ok(_send_slot) = send_semaphore.acquire_owned().await else {
return;
};
if let Some(identity) = config.identity.as_ref() {
let init_result = if config.transactional_id.is_some() {
if identity.is_initialized() {
Ok(())
} else {
Err(KrafkaError::invalid_state(
"transactional producer identity not initialized; \
call init_transactions() before sending",
))
}
} else {
super::ensure_idempotent_producer_id_initialized(identity, &metadata, &retry_policy)
.await
};
if let Err(error) = init_result {
metrics.record_error_for_topic(topic.as_ref());
Self::fail_pending(&topic, partition, pending, &error, &config).await;
return;
}
}
let _timer = metrics.send_latency.start();
Self::produce_pending(
&topic,
partition,
pending,
enqueued_at,
&metadata,
&config,
&retry_policy,
&metrics,
0,
)
.await;
}
async fn allocate_sequence_range(
config: &AccumulatorConfig,
metadata: &Arc<ClusterMetadata>,
retry_policy: &RetryPolicy,
topic: &TopicHandle,
partition: PartitionId,
record_count: i32,
) -> Result<Option<i32>> {
let Some(identity) = config.identity.as_ref() else {
return Ok(None);
};
if let Some(base) =
identity.checked_allocate_sequence(topic.as_ref(), partition, record_count)?
{
return Ok(Some(base));
}
if config.transactional_id.is_some() {
return Err(KrafkaError::invalid_state(
"transactional producer identity was cleared before the batch could allocate \
a sequence range; the transaction must be aborted and restarted",
));
}
super::ensure_idempotent_producer_id_initialized(identity, metadata, retry_policy).await?;
identity
.checked_allocate_sequence(topic.as_ref(), partition, record_count)?
.map(Some)
.ok_or_else(|| {
KrafkaError::broker(
ErrorCode::UnknownProducerId,
format!(
"producer identity was reset while allocating sequences for \
{topic}-{partition}; retry the send"
),
)
})
}
async fn fail_pending(
topic: &TopicHandle,
partition: PartitionId,
pending: Vec<PendingRecord>,
error: &KrafkaError,
config: &AccumulatorConfig,
) {
let topic_owned = topic.to_string();
for p in pending {
let meta = RecordMetadata {
topic: topic_owned.clone(),
partition,
offset: -1,
timestamp: 0,
delivery: DeliveryConfirmation::Failed,
};
crate::interceptor::safe_on_acknowledgement(&*config.interceptor, &meta, Some(error));
if let Some(dlq) = config.dead_letter_queue.as_ref() {
let dlq_record = ProducerRecord {
topic: topic_owned.clone(),
partition: Some(partition),
key: p.record.key.clone(),
value: p.record.value.clone(),
timestamp: p.record.timestamp,
headers: p.record.headers.clone(),
record_name: None,
};
dlq.send(dlq_record, error.to_string()).await;
}
let _ = p.response_tx.send(AppendResponse::Done(Err(error.clone())));
}
}
async fn encode_batch_request(
topic: &TopicHandle,
partition: PartitionId,
pending: &[PendingRecord],
sequence: Option<i32>,
config: &AccumulatorConfig,
) -> Result<(ProduceRequest, u64, u64)> {
let effective_compression = config
.topic_compression
.get(topic.as_ref())
.copied()
.unwrap_or(config.compression);
let mut batch_builder = RecordBatchBuilder::new()
.compression(effective_compression)
.compression_level(config.compression_level);
if let (Some(identity), Some(s)) = (&config.identity, sequence) {
batch_builder =
batch_builder.producer(identity.producer_id(), identity.producer_epoch(), s);
}
if config.transactional_id.is_some() {
batch_builder = batch_builder.transactional(true);
}
let uncompressed_len: u64 = pending.iter().map(|p| p.estimated_size as u64).sum();
for p in pending {
batch_builder = p.record.append_to_batch_builder(batch_builder);
}
let batch = batch_builder.build();
let batch_bytes = if is_cpu_heavy_compression(effective_compression) {
tokio::task::spawn_blocking(move || batch.encode())
.await
.map_err(|join_err| {
KrafkaError::invalid_state(format!(
"record-batch compression task failed: {join_err}"
))
})??
} else {
batch.encode()?
};
let compressed_len = batch_bytes.len() as u64;
Ok((
ProduceRequest {
transactional_id: config.transactional_id.clone(),
acks: config.acks,
timeout_ms: crate::util::duration_to_millis_i32(config.request_timeout),
topic_data: vec![ProduceTopicData {
name: topic.to_string(),
topic_id: None,
partition_data: vec![ProducePartitionData {
index: partition,
records: batch_bytes,
}],
}],
},
compressed_len,
uncompressed_len,
))
}
#[allow(clippy::too_many_arguments)]
fn produce_pending<'a>(
topic: &'a TopicHandle,
partition: PartitionId,
pending: Vec<PendingRecord>,
enqueued_at: Instant,
metadata: &'a Arc<ClusterMetadata>,
config: &'a AccumulatorConfig,
retry_policy: &'a RetryPolicy,
metrics: &'a Arc<ProducerMetrics>,
split_depth: u8,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send + 'a>> {
Box::pin(Self::produce_pending_inner(
topic,
partition,
pending,
enqueued_at,
metadata,
config,
retry_policy,
metrics,
split_depth,
))
}
#[allow(clippy::too_many_arguments)]
async fn produce_pending_inner(
topic: &TopicHandle,
partition: PartitionId,
mut pending: Vec<PendingRecord>,
enqueued_at: Instant,
metadata: &Arc<ClusterMetadata>,
config: &AccumulatorConfig,
retry_policy: &RetryPolicy,
metrics: &Arc<ProducerMetrics>,
split_depth: u8,
) {
let record_count = pending.len() as i32;
if record_count == 0 {
return;
}
let mut sequence: Option<i32> = match Self::allocate_sequence_range(
config,
metadata,
retry_policy,
topic,
partition,
record_count,
)
.await
{
Ok(s) => s,
Err(e) => {
Self::fail_pending(topic, partition, pending, &e, config).await;
return;
}
};
let (mut request, compressed_len, uncompressed_len) =
match Self::encode_batch_request(topic, partition, &pending, sequence, config).await {
Ok(r) => r,
Err(e) => {
if let (Some(identity), Some(base)) = (config.identity.as_ref(), sequence) {
Self::release_failed_sequence_range(
identity,
topic,
partition,
base,
record_count,
);
}
Self::fail_pending(topic, partition, pending, &e, config).await;
return;
}
};
if config.compression != Compression::None {
metrics.record_compression(compressed_len, uncompressed_len);
}
let mut retry_ctx = RetryContext::new_with_start(
retry_policy.clone(),
format!("batch({topic}-{partition})"),
enqueued_at,
);
let result: std::result::Result<(i64, i64), KrafkaError> = loop {
let conn = match metadata
.get_leader_connection(topic.as_ref(), partition)
.await
{
Ok(c) => c,
Err(e) => {
if e.is_retriable() {
debug!(
topic = %topic,
partition = partition,
error = %e,
"Batch connection error, refreshing metadata"
);
if let Err(refresh_err) = metadata
.refresh_for_topics_forced(Some(&[topic.as_ref()]))
.await
{
debug!(error = %refresh_err, "Metadata refresh failed during batch retry");
}
}
if let Some(backoff) = retry_ctx.record_failure(&e) {
metrics.record_retry();
retry_ctx.wait(backoff).await;
continue;
}
break Err(e);
}
};
conn.await_throttle().await;
let mut produce_version = match conn.negotiate_api_version(
ApiKey::Produce,
versions::PRODUCE_MAX,
versions::PRODUCE_MIN,
) {
Some(v) => v,
None => {
let e = KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
"no mutually supported Produce API version",
);
debug!(
topic = %topic,
partition = partition,
"Produce version negotiation failed, refreshing metadata"
);
if let Err(refresh_err) = metadata
.refresh_for_topics_forced(Some(&[topic.as_ref()]))
.await
{
debug!(
error = %refresh_err,
"Metadata refresh failed during batch retry"
);
}
if let Some(backoff) = retry_ctx.record_failure(&e) {
metrics.record_retry();
retry_ctx.wait(backoff).await;
continue;
}
break Err(e);
}
};
if produce_version >= 13 && !super::fill_produce_topic_ids(&mut request, metadata) {
produce_version = 12;
}
let encoded_body = match super::encode_and_validate_produce_request(
&config.client_id,
config.max_request_size,
produce_version,
&request,
) {
Ok(b) => b,
Err(error) => break Err(error),
};
if config.acks == 0 {
match conn
.send_fire_and_forget(ApiKey::Produce, produce_version, |buf| {
buf.put_slice(&encoded_body);
Ok(())
})
.await
{
Ok(()) => {
retry_ctx.record_success();
break Ok((-1, -1));
}
Err(e) => {
if let Some(backoff) = retry_ctx.record_failure(&e) {
metrics.record_retry();
retry_ctx.wait(backoff).await;
continue;
}
break Err(e);
}
}
}
let response_result = conn
.send_request(ApiKey::Produce, produce_version, |buf| {
buf.put_slice(&encoded_body);
Ok(())
})
.await;
match response_result {
Ok(mut response_buf) => {
match ProduceResponse::decode_versioned(produce_version, &mut response_buf) {
Ok(produce_response) => {
conn.notify_throttle(produce_response.throttle_time_ms);
let pr = produce_response
.responses
.iter()
.find(|r| {
if produce_version >= 13 {
r.topic_id.as_ref().is_some_and(|id| {
metadata.topic_name_for_id(id).as_deref()
== Some(topic.as_ref())
})
} else {
r.name == topic.as_ref()
}
})
.and_then(|r| {
r.partition_responses.iter().find(|p| p.index == partition)
});
match pr {
Some(pr) if pr.error_code.is_ok() => {
retry_ctx.record_success();
break Ok((pr.base_offset, pr.log_append_time_ms));
}
Some(pr)
if pr.error_code == ErrorCode::DuplicateSequenceNumber
&& config.identity.is_some() =>
{
debug!(
topic = %topic,
partition = partition,
"DuplicateSequenceNumber in batch — dedup confirmed"
);
retry_ctx.record_success();
break Ok((-1, -1));
}
Some(pr) => {
let err = KrafkaError::broker(
pr.error_code,
format!("batch produce failed for {topic}-{partition}"),
);
if pr.error_code == ErrorCode::UnknownProducerId
&& let (Some(identity), Some(current_sequence)) =
(config.identity.as_ref(), sequence)
{
warn!(
topic = %topic,
partition = partition,
"UnknownProducerId in batch, reinitializing idempotent producer state"
);
let new_sequence = match super::recover_unknown_producer_id(
identity,
metadata,
retry_policy,
topic.as_ref(),
partition,
current_sequence,
record_count,
)
.await
{
Ok(new_sequence) => new_sequence,
Err(recovery_error) => break Err(recovery_error),
};
sequence = Some(new_sequence);
match Self::encode_batch_request(
topic, partition, &pending, sequence, config,
)
.await
{
Ok((new_request, ..)) => request = new_request,
Err(encode_err) => break Err(encode_err),
}
} else if pr.error_code == ErrorCode::OutOfOrderSequenceNumber
&& let (Some(identity), Some(base)) =
(config.identity.as_ref(), sequence)
{
match identity.can_reset_after_out_of_order(
topic.as_ref(),
partition,
base,
record_count,
) {
Ok(true) => {
warn!(
topic = %topic,
partition = partition,
base_sequence = base,
"OutOfOrderSequenceNumber for head-of-line batch, \
resetting sequence and retrying"
);
let new_seq = match identity.reset_and_allocate(
topic.as_ref(),
partition,
record_count,
) {
Ok(s) => s,
Err(e) => break Err(e),
};
sequence = Some(new_seq);
match Self::encode_batch_request(
topic, partition, &pending, sequence, config,
)
.await
{
Ok((r, ..)) => request = r,
Err(encode_err) => break Err(encode_err),
}
}
Ok(false) => {
break Err(super::out_of_order_data_loss_error(
topic.as_ref(),
partition,
base,
));
}
Err(e) => break Err(e),
}
} else if err.is_retriable()
&& !super::apply_produce_leader_hint(
metadata,
topic.as_ref(),
partition,
&produce_response,
pr,
)
&& let Err(refresh_err) = metadata
.refresh_for_topics_forced(Some(&[topic.as_ref()]))
.await
{
debug!(error = %refresh_err, "Metadata refresh failed during batch retry");
}
if let Some(backoff) = retry_ctx.record_failure(&err) {
metrics.record_retry();
retry_ctx.wait(backoff).await;
continue;
}
break Err(err);
}
None => {
break Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
"partition not found in response",
));
}
}
}
Err(e) => {
if let Some(backoff) = retry_ctx.record_failure(&e) {
metrics.record_retry();
retry_ctx.wait(backoff).await;
continue;
}
break Err(e);
}
}
}
Err(e) => {
if e.is_retriable() {
debug!(
topic = %topic,
partition = partition,
error = %e,
"Batch send error, refreshing metadata"
);
if let Err(refresh_err) = metadata
.refresh_for_topics_forced(Some(&[topic.as_ref()]))
.await
{
debug!(error = %refresh_err, "Metadata refresh failed during batch retry");
}
}
if let Some(backoff) = retry_ctx.record_failure(&e) {
metrics.record_retry();
retry_ctx.wait(backoff).await;
continue;
}
break Err(e);
}
}
};
match result {
Ok((base_offset, timestamp)) => {
if let (Some(identity), Some(seq)) = (&config.identity, sequence)
&& let Ok(last_seq) =
super::idempotent::last_sequence_of_batch(seq, record_count)
{
identity.acknowledge(topic.as_ref(), partition, last_seq);
if let Some(ref store) = config.state_store {
let snapshot = identity.snapshot();
let store = Arc::clone(store);
tokio::spawn(async move {
if let Err(err) = store.store_erased(&snapshot).await {
tracing::warn!(error = %err, "Failed to persist producer state snapshot");
}
});
}
}
let batch_bytes_total: u64 = pending.iter().map(|p| p.estimated_size as u64).sum();
metrics.record_batch_for_topic(
topic.as_ref(),
pending.len() as u64,
batch_bytes_total,
);
let topic_owned = topic.to_string();
for p in pending {
let meta = RecordMetadata {
topic: topic_owned.clone(),
partition,
offset: if base_offset >= 0 {
base_offset + p.offset_in_batch
} else {
-1
},
timestamp,
delivery: if base_offset >= 0 {
DeliveryConfirmation::Offset
} else if config.acks == 0 {
DeliveryConfirmation::Unacknowledged
} else {
DeliveryConfirmation::Deduplicated
},
};
crate::interceptor::safe_on_acknowledgement(&*config.interceptor, &meta, None);
let _ = p.response_tx.send(AppendResponse::Done(Ok(meta)));
}
}
Err(e) => {
let sequence_space_intact =
if let (Some(identity), Some(base)) = (config.identity.as_ref(), sequence) {
Self::release_failed_sequence_range(
identity,
topic,
partition,
base,
record_count,
)
} else {
true
};
if sequence_space_intact
&& is_batch_too_large(&e)
&& pending.len() > 1
&& split_depth < MAX_BATCH_SPLIT_DEPTH
{
let mid = pending.len() / 2;
let second_half = pending.split_off(mid);
warn!(
topic = %topic,
partition = partition,
records = pending.len() + second_half.len(),
error = %e,
"Batch rejected as too large; splitting and resubmitting both halves"
);
metrics.record_retry();
for half in [pending, second_half] {
Self::produce_pending(
topic,
partition,
half,
enqueued_at,
metadata,
config,
retry_policy,
metrics,
split_depth + 1,
)
.await;
}
return;
}
metrics.record_error_for_topic(topic.as_ref());
Self::fail_pending(topic, partition, pending, &e, config).await;
}
}
}
fn release_failed_sequence_range(
identity: &super::idempotent::ProducerIdentity,
topic: &TopicHandle,
partition: PartitionId,
base_sequence: i32,
record_count: i32,
) -> bool {
match identity.rollback_sequence_range(
topic.as_ref(),
partition,
base_sequence,
record_count,
) {
Ok(super::idempotent::RollbackOutcome::RolledBack) => true,
Ok(super::idempotent::RollbackOutcome::NotTail) | Err(_) => {
warn!(
topic = %topic,
partition = partition,
base_sequence,
"Failed batch no longer owns the tail of its sequence range; \
resetting partition sequences and requesting a producer-ID re-init \
instead of rewinding into a newer allocation"
);
identity.reset_partition_sequences(topic.as_ref(), partition);
identity.request_reinit();
false
}
}
}
async fn flush_all(&mut self) -> Result<()> {
let keys: Vec<_> = self
.batches
.iter()
.filter(|(_, batch)| !batch.is_empty())
.map(|(key, _)| key.clone())
.collect();
let mut extracted = Vec::with_capacity(keys.len());
for key in keys {
if let Some(item) = self.extract_batch(&key) {
extracted.push((key, item));
}
}
Self::spawn_batches_bounded(
extracted,
&self.metadata,
&self.config,
&self.retry_policy,
&self.metrics,
self.send_semaphore.clone(),
)
.await;
Ok(())
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
#[test]
fn test_accumulator_config_default() {
let config = AccumulatorConfig::default();
assert_eq!(config.batch_size, 16384);
assert_eq!(config.linger, Duration::ZERO);
assert_eq!(config.acks, -1);
}
#[test]
fn test_accumulator_batch_age() {
let batch = AccumulatorBatch::new(16384);
std::thread::sleep(Duration::from_millis(10));
assert!(batch.age() >= Duration::from_millis(10));
}
#[test]
fn backpressure_charges_the_delivery_clock_but_not_the_linger_clock() {
const LINGER: Duration = Duration::from_millis(50);
let mut batch = AccumulatorBatch::new(16384);
let opened_at = batch.opened_at;
let waited_since = opened_at - Duration::from_secs(60);
batch.charge_from(waited_since);
assert_eq!(
batch.deadline_from, waited_since,
"the delivery budget must cover the backpressure wait"
);
assert_eq!(
batch.opened_at, opened_at,
"the linger window must not be retroactively expired by it"
);
assert!(
batch.linger_deadline(LINGER) > Instant::now(),
"a freshly opened batch must still have its full linger window"
);
batch.charge_from(opened_at + Duration::from_secs(1));
assert_eq!(batch.deadline_from, waited_since);
}
#[tokio::test]
async fn an_idle_accumulator_sleeps_instead_of_polling() {
let mut accumulator = test_accumulator(Duration::from_millis(5));
let now = Instant::now();
assert!(
accumulator.next_wake_deadline() >= now + IDLE_TICK - Duration::from_millis(1),
"an empty accumulator has nothing due before the housekeeping tick"
);
accumulator
.batches
.insert((Arc::from("orders"), 0), non_empty_batch());
let wake = accumulator.next_wake_deadline();
assert!(
wake < now + IDLE_TICK,
"an open batch must wake the loop at its linger deadline, not the idle tick"
);
let mut immediate = test_accumulator(Duration::ZERO);
immediate
.batches
.insert((Arc::from("orders"), 0), non_empty_batch());
assert!(
immediate.next_wake_deadline() >= now + IDLE_TICK - Duration::from_millis(1),
"linger = 0 dispatch is event-driven; the timer is housekeeping only"
);
}
#[test]
fn test_accumulator_batch_new() {
let batch = AccumulatorBatch::new(32768);
assert!(batch.is_empty());
assert!(batch.pending.is_empty());
}
#[test]
fn test_accumulator_config_custom() {
let config = AccumulatorConfig {
batch_size: 65536,
linger: Duration::from_millis(50),
compression: Compression::Snappy,
acks: 1,
client_id: "test-client".to_string(),
request_timeout: Duration::from_secs(10),
max_request_size: 131072,
buffer_memory: 64 * 1024 * 1024,
max_block_ms: Duration::from_secs(30),
interceptor: Arc::new(crate::interceptor::NoOpProducerInterceptor),
identity: None,
partitioner: Arc::new(crate::producer::partitioner::DefaultPartitioner::new()),
state_store: None,
topic_compression: AHashMap::new(),
transactional_id: None,
compression_level: None,
dead_letter_queue: None,
};
assert_eq!(config.batch_size, 65536);
assert_eq!(config.linger, Duration::from_millis(50));
assert_eq!(config.acks, 1);
assert_eq!(config.client_id, "test-client");
assert_eq!(config.max_request_size, 131072);
assert_eq!(config.buffer_memory, 64 * 1024 * 1024);
}
fn test_pending(value: bytes::Bytes) -> PendingRecord {
let (response_tx, _response_rx) = oneshot::channel();
let estimated_size = value.len();
PendingRecord {
record: RoutedRecord {
key: None,
value,
timestamp: None,
headers: Vec::new(),
},
response_tx,
offset_in_batch: 0,
estimated_size,
_buffered_record_guard: BufferedRecordGuard::new(
Arc::new(AtomicUsize::new(0)),
Arc::new(ProducerMetrics::default()),
),
_operation_guard: Arc::new(InFlightBarrier::new())
.start("test")
.expect("a fresh barrier is open"),
}
}
fn test_accumulator(linger: Duration) -> RecordAccumulator {
let pool = Arc::new(crate::network::ConnectionPool::new(
crate::network::ConnectionConfig::default(),
));
RecordAccumulator {
config: AccumulatorConfig {
linger,
..AccumulatorConfig::default()
},
batches: AHashMap::new(),
metadata: Arc::new(ClusterMetadata::new(
vec!["127.0.0.1:9092".to_string()],
pool,
Duration::from_secs(300),
)),
send_semaphore: Arc::new(Semaphore::new(MAX_CONCURRENT_BATCH_SENDS)),
in_flight_memory: Arc::new(AtomicUsize::new(0)),
retry_policy: RetryPolicy::default(),
metrics: Arc::new(ProducerMetrics::default()),
memory_permits: Arc::new(Semaphore::new(1024)),
partitioner: Arc::new(crate::producer::DefaultPartitioner::new()),
partition_inflight: AHashMap::new(),
dispatch_notify: Arc::new(tokio::sync::Notify::new()),
}
}
fn non_empty_batch() -> AccumulatorBatch {
let mut batch = AccumulatorBatch::new(16384);
batch
.pending
.push(test_pending(bytes::Bytes::from_static(b"v")));
batch
}
#[cfg(any(feature = "zstd", feature = "gzip"))]
fn compressible_payload() -> bytes::Bytes {
let mut s = String::with_capacity(96 * 1024);
for i in 0..2048 {
s.push_str(&format!(
"{{\"id\":{i},\"user\":\"user-{}\",\"score\":{},\"tag\":\"{}\"}}\n",
i % 97,
i * 7 % 1013,
if i % 3 == 0 { "alpha" } else { "beta" }
));
}
bytes::Bytes::from(s)
}
#[cfg(any(feature = "zstd", feature = "gzip"))]
async fn encoded_size(codec: Compression, level: Option<i32>, payload: bytes::Bytes) -> usize {
let topic: TopicHandle = Arc::from("levels");
let config = AccumulatorConfig {
compression: codec,
compression_level: level,
..AccumulatorConfig::default()
};
let pending = vec![test_pending(payload)];
let (request, compressed, _uncompressed) =
RecordAccumulator::encode_batch_request(&topic, 0, &pending, None, &config)
.await
.expect("encoding one batch cannot fail");
assert_eq!(
compressed as usize,
request.topic_data[0].partition_data[0].records.len(),
"the reported compressed length must match the encoded record set"
);
compressed as usize
}
#[cfg(feature = "zstd")]
#[tokio::test]
async fn zstd_compression_level_reaches_the_batched_encoder() {
let payload = compressible_payload();
let fast = encoded_size(Compression::Zstd, Some(1), payload.clone()).await;
let dense = encoded_size(Compression::Zstd, Some(19), payload).await;
assert!(
dense < fast,
"zstd level 19 must compress harder than level 1 (got {dense} vs {fast} bytes); \
equal sizes mean the level never reached the encoder"
);
}
#[cfg(feature = "gzip")]
#[tokio::test]
async fn gzip_compression_level_reaches_the_batched_encoder() {
let payload = compressible_payload();
let fast = encoded_size(Compression::Gzip, Some(1), payload.clone()).await;
let dense = encoded_size(Compression::Gzip, Some(9), payload).await;
assert!(
dense < fast,
"gzip level 9 must compress harder than level 1 (got {dense} vs {fast} bytes)"
);
}
#[cfg(all(feature = "zstd", feature = "gzip"))]
#[tokio::test]
async fn per_topic_codec_beats_the_producer_wide_setting() {
let topic: TopicHandle = Arc::from("overridden");
let mut topic_compression = AHashMap::new();
topic_compression.insert("overridden".to_string(), Compression::Gzip);
let config = AccumulatorConfig {
compression: Compression::Zstd,
compression_level: Some(1),
topic_compression,
..AccumulatorConfig::default()
};
let pending = vec![test_pending(compressible_payload())];
let (request, _c, _u) =
RecordAccumulator::encode_batch_request(&topic, 0, &pending, None, &config)
.await
.expect("gzip level 1 is valid");
let records = &request.topic_data[0].partition_data[0].records;
let attributes = i16::from_be_bytes([records[21], records[22]]);
assert_eq!(
attributes & 0x07,
1,
"the per-topic gzip override must beat the producer-wide zstd setting"
);
}
#[derive(Debug, Default)]
struct RecordingDlq {
received: parking_lot::Mutex<Vec<(String, String)>>,
}
impl crate::dlq::DeadLetterQueue for RecordingDlq {
fn send<'a>(
&'a self,
record: ProducerRecord,
error: String,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send + 'a>> {
let value = String::from_utf8_lossy(&record.value).into_owned();
Box::pin(async move {
self.received.lock().push((value, error));
})
}
}
#[tokio::test]
async fn a_permanently_failed_batch_reaches_the_dead_letter_queue() {
let dlq = Arc::new(RecordingDlq::default());
let config = AccumulatorConfig {
dead_letter_queue: Some(dlq.clone()),
..AccumulatorConfig::default()
};
let topic: TopicHandle = Arc::from("orders");
let pending = vec![
test_pending(bytes::Bytes::from_static(b"first")),
test_pending(bytes::Bytes::from_static(b"second")),
];
let error = KrafkaError::broker(ErrorCode::RecordListTooLarge, "batch rejected");
RecordAccumulator::fail_pending(&topic, 3, pending, &error, &config).await;
let received = dlq.received.lock().clone();
assert_eq!(
received.len(),
2,
"every record in the failed batch must be routed, got {received:?}"
);
assert_eq!(received[0].0, "first");
assert_eq!(received[1].0, "second");
assert!(
received[0].1.contains("batch rejected"),
"the DLQ must be told why, got: {}",
received[0].1
);
}
#[tokio::test]
async fn failing_a_batch_without_a_dlq_still_answers_every_caller() {
let config = AccumulatorConfig::default();
let topic: TopicHandle = Arc::from("orders");
let (tx_a, rx_a) = oneshot::channel();
let (tx_b, rx_b) = oneshot::channel();
let mut a = test_pending(bytes::Bytes::from_static(b"a"));
let mut b = test_pending(bytes::Bytes::from_static(b"b"));
a.response_tx = tx_a;
b.response_tx = tx_b;
let error = KrafkaError::broker(ErrorCode::RecordListTooLarge, "batch rejected");
RecordAccumulator::fail_pending(&topic, 0, vec![a, b], &error, &config).await;
for rx in [rx_a, rx_b] {
match rx.await {
Ok(AppendResponse::Done(Err(e))) => assert!(e.to_string().contains("rejected")),
other => panic!("expected a delivered failure, got {other:?}"),
}
}
}
#[test]
fn test_estimate_record_size() {
let record = ProducerRecord::new("test-topic", b"value".to_vec());
let size = record.estimated_size();
assert!(size >= 5);
assert!(size > 64);
let record_with_key =
ProducerRecord::new("test-topic", b"value".to_vec()).with_key(b"key".to_vec());
let size_with_key = record_with_key.estimated_size();
assert!(size_with_key > size);
}
#[test]
fn test_linger_zero_check_interval() {
let linger = Duration::ZERO;
let check_interval = Duration::from_millis(1).max(linger / 10);
assert_eq!(check_interval, Duration::from_millis(1));
}
#[test]
fn test_linger_zero_is_zero() {
let config = AccumulatorConfig {
linger: Duration::ZERO,
..Default::default()
};
assert!(config.linger.is_zero());
}
#[test]
fn test_accumulator_handle_is_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<RecordAccumulatorHandle>();
assert_send_sync::<RecordAccumulator>();
}
#[tokio::test]
async fn test_backpressure_timeout_returns_timeout_error() {
let (sender, _receiver) = mpsc::channel::<AccumulatorMessage>(16);
let handle = RecordAccumulatorHandle {
sender,
memory_permits: Arc::new(Semaphore::new(0)),
memory_capacity: 1024 * 1024, max_request_size: 0,
max_block_ms: Duration::from_millis(50),
in_flight_barrier: Arc::new(InFlightBarrier::new()),
buffered_records: Arc::new(AtomicUsize::new(0)),
metrics: Arc::new(ProducerMetrics::default()),
};
let record = ProducerRecord::new("topic", b"value".to_vec());
let result = handle.append(record, 0).await;
assert!(result.is_err());
let err = result.unwrap_err();
let err_msg = err.to_string();
assert!(
err_msg.contains("max_block"),
"expected max_block in error, got: {err_msg}"
);
assert!(
matches!(err, KrafkaError::Timeout { .. }),
"expected Timeout variant, got: {err:?}"
);
}
#[tokio::test]
async fn test_backpressure_unblocks_on_permit_release() {
let sem = Arc::new(Semaphore::new(0));
let s = sem.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(10)).await;
s.add_permits(128);
});
let result = tokio::time::timeout(Duration::from_secs(2), sem.acquire_many(64)).await;
assert!(result.is_ok(), "acquire_many should have completed");
assert!(
result.unwrap().is_ok(),
"acquire_many should have succeeded"
);
}
#[tokio::test]
async fn test_oversize_record_rejected_immediately() {
let (sender, _receiver) = mpsc::channel::<AccumulatorMessage>(16);
let handle = RecordAccumulatorHandle {
sender,
memory_permits: Arc::new(Semaphore::new(16)),
memory_capacity: 16, max_request_size: 0,
max_block_ms: Duration::from_secs(60),
in_flight_barrier: Arc::new(InFlightBarrier::new()),
buffered_records: Arc::new(AtomicUsize::new(0)),
metrics: Arc::new(ProducerMetrics::default()),
};
let record = ProducerRecord::new("topic", vec![0u8; 1024]);
let start = std::time::Instant::now();
let result = handle.append(record, 0).await;
assert!(start.elapsed() < Duration::from_secs(1));
let err = result.expect_err("oversize record must be rejected");
assert!(
err.to_string().contains("buffer_memory"),
"expected buffer_memory error, got: {err}"
);
}
#[tokio::test]
async fn test_closed_semaphore_unblocks_waiters() {
let (sender, _receiver) = mpsc::channel::<AccumulatorMessage>(16);
let sem = Arc::new(Semaphore::new(0));
let handle = RecordAccumulatorHandle {
sender,
memory_permits: sem.clone(),
memory_capacity: 1024 * 1024,
max_request_size: 0,
max_block_ms: Duration::from_secs(60),
in_flight_barrier: Arc::new(InFlightBarrier::new()),
buffered_records: Arc::new(AtomicUsize::new(0)),
metrics: Arc::new(ProducerMetrics::default()),
};
let sem_close = sem.clone();
tokio::spawn(async move {
tokio::time::sleep(Duration::from_millis(10)).await;
sem_close.close();
});
let record = ProducerRecord::new("topic", b"value".to_vec());
let start = std::time::Instant::now();
let result = handle.append(record, 0).await;
assert!(
start.elapsed() < Duration::from_secs(1),
"must unblock on close, not on max_block timeout"
);
let err = result.expect_err("closed semaphore must surface as error");
assert!(
matches!(err, KrafkaError::InvalidState { .. }),
"expected InvalidState variant, got: {err:?}"
);
}
#[tokio::test]
async fn test_permits_released_when_append_message_dropped() {
let (sender, mut receiver) = mpsc::channel::<AccumulatorMessage>(16);
let sem = Arc::new(Semaphore::new(1024));
let metrics = Arc::new(ProducerMetrics::default());
let buffered_records = Arc::new(AtomicUsize::new(0));
let handle = RecordAccumulatorHandle {
sender,
memory_permits: sem.clone(),
memory_capacity: 1024,
max_request_size: 0,
max_block_ms: Duration::from_millis(500),
in_flight_barrier: Arc::new(InFlightBarrier::new()),
buffered_records: buffered_records.clone(),
metrics: metrics.clone(),
};
let record = ProducerRecord::new("topic", vec![0u8; 256]);
let append_fut = tokio::spawn(async move { handle.append(record, 0).await });
let msg = tokio::time::timeout(Duration::from_secs(2), receiver.recv())
.await
.expect("timed out waiting for Append message to arrive in channel")
.expect("channel closed before message arrived");
assert_eq!(metrics.buffered_records.get(), 1);
assert_eq!(buffered_records.load(Ordering::Relaxed), 1);
drop(msg);
drop(receiver);
let _ = append_fut.await;
assert_eq!(
sem.available_permits(),
1024,
"permits leaked when the Append message was dropped"
);
assert_eq!(metrics.buffered_records.get(), 0);
assert_eq!(buffered_records.load(Ordering::Relaxed), 0);
}
#[test]
fn test_check_record_admission_rejects_oversized_for_buffer() {
let err = check_record_admission(1024, 16, 0).expect_err("must reject");
let msg = err.to_string();
assert!(
msg.contains("buffer_memory"),
"error must cite buffer_memory, got: {msg}"
);
assert!(
!msg.contains("u32::MAX"),
"must not cite u32::MAX for a buffer_memory rejection, got: {msg}"
);
}
#[test]
fn test_check_record_admission_rejects_oversized_for_semaphore_limit() {
let oversized = max_record_semaphore_permits() + 1;
let err = check_record_admission(oversized, usize::MAX, 0).expect_err("must reject");
let msg = err.to_string();
assert!(
msg.contains("Semaphore::MAX_PERMITS"),
"error must cite the effective semaphore limit, got: {msg}"
);
assert!(
!msg.contains("buffer_memory"),
"must not cite buffer_memory for a semaphore-limit rejection, got: {msg}"
);
}
#[test]
fn test_buffered_record_guard_updates_metric() {
let metrics = Arc::new(ProducerMetrics::default());
let buffered_records = Arc::new(AtomicUsize::new(0));
{
let _guard = BufferedRecordGuard::new(buffered_records.clone(), metrics.clone());
assert_eq!(metrics.buffered_records.get(), 1);
}
assert_eq!(metrics.buffered_records.get(), 0);
}
#[test]
fn test_effective_memory_capacity_zero_returns_max() {
assert_eq!(effective_memory_capacity(0), Semaphore::MAX_PERMITS);
}
#[test]
fn test_effective_memory_capacity_clamps_over_limit() {
let over = Semaphore::MAX_PERMITS + 1;
assert_eq!(effective_memory_capacity(over), Semaphore::MAX_PERMITS);
}
#[test]
fn test_effective_memory_capacity_passthrough() {
let within = Semaphore::MAX_PERMITS / 2;
assert_eq!(effective_memory_capacity(within), within);
}
#[tokio::test]
async fn test_same_partition_batches_dispatch_in_seal_order() {
let slot = Arc::new(PartitionInFlight::new(Arc::new(tokio::sync::Notify::new())));
let tickets: Vec<PartitionTicket> = (0..8).map(|_| slot.take_ticket()).collect();
let order = Arc::new(parking_lot::Mutex::new(Vec::new()));
let mut handles = Vec::new();
for (i, ticket) in tickets.into_iter().enumerate().rev() {
let order = order.clone();
handles.push(tokio::spawn(async move {
let turn = ticket.acquire().await;
order.lock().push(i);
tokio::time::sleep(Duration::from_millis(2)).await;
drop(turn);
}));
tokio::task::yield_now().await;
}
for h in handles {
h.await.expect("task panicked");
}
assert_eq!(
*order.lock(),
(0..8).collect::<Vec<usize>>(),
"same-partition batches must dispatch in seal order"
);
}
#[tokio::test]
async fn test_sequences_are_monotonic_and_gapless_under_partition_fifo() {
use super::super::idempotent::ProducerIdentity;
const BATCHES: usize = 16;
const RECORDS_PER_BATCH: i32 = 3;
let identity = Arc::new(ProducerIdentity::new());
identity.initialize(1, 0);
let slot = Arc::new(PartitionInFlight::new(Arc::new(tokio::sync::Notify::new())));
let tickets: Vec<PartitionTicket> = (0..BATCHES).map(|_| slot.take_ticket()).collect();
let observed = Arc::new(parking_lot::Mutex::new(Vec::new()));
let mut handles = Vec::new();
for ticket in tickets.into_iter().rev() {
let identity = identity.clone();
let observed = observed.clone();
handles.push(tokio::spawn(async move {
let _turn = ticket.acquire().await;
let base = identity
.allocate_sequence("t", 0, RECORDS_PER_BATCH)
.expect("allocate");
observed.lock().push(base);
}));
tokio::task::yield_now().await;
}
for h in handles {
h.await.expect("task panicked");
}
let expected: Vec<i32> = (0..BATCHES as i32).map(|i| i * RECORDS_PER_BATCH).collect();
assert_eq!(
*observed.lock(),
expected,
"sequence allocation order must follow dispatch order with no gaps"
);
}
#[tokio::test]
async fn test_different_partitions_are_not_serialized() {
let a = Arc::new(PartitionInFlight::new(Arc::new(tokio::sync::Notify::new())));
let b = Arc::new(PartitionInFlight::new(Arc::new(tokio::sync::Notify::new())));
let turn_a = a.take_ticket().acquire().await;
let turn_b = tokio::time::timeout(Duration::from_millis(500), b.take_ticket().acquire())
.await
.expect("cross-partition sends must not serialize");
drop(turn_a);
drop(turn_b);
}
#[tokio::test]
async fn test_dropped_ticket_does_not_stall_the_partition() {
let slot = Arc::new(PartitionInFlight::new(Arc::new(tokio::sync::Notify::new())));
let abandoned = slot.take_ticket();
let next = slot.take_ticket();
drop(abandoned);
let turn = tokio::time::timeout(Duration::from_millis(500), next.acquire())
.await
.expect("a dropped ticket must not stall the partition FIFO");
drop(turn);
assert!(slot.is_idle());
}
#[tokio::test]
async fn a_completed_batch_wakes_the_dispatch_loop() {
let dispatch = Arc::new(tokio::sync::Notify::new());
let slot = Arc::new(PartitionInFlight::new(dispatch.clone()));
let turn = slot.take_ticket().acquire().await;
drop(turn);
tokio::time::timeout(Duration::from_millis(500), dispatch.notified())
.await
.expect("a batch completion must wake the dispatch loop");
}
#[tokio::test]
async fn test_partition_inflight_idle_tracking() {
let slot = Arc::new(PartitionInFlight::new(Arc::new(tokio::sync::Notify::new())));
assert!(slot.is_idle());
let ticket = slot.take_ticket();
assert!(!slot.is_idle(), "an outstanding ticket means not idle");
let turn = ticket.acquire().await;
assert!(!slot.is_idle());
drop(turn);
assert!(slot.is_idle());
}
#[test]
fn test_is_batch_too_large_classification() {
assert!(is_batch_too_large(&KrafkaError::broker(
ErrorCode::MessageTooLarge,
"too big"
)));
assert!(is_batch_too_large(&KrafkaError::broker(
ErrorCode::RecordListTooLarge,
"too big"
)));
assert!(is_batch_too_large(&KrafkaError::protocol_kind(
ProtocolErrorKind::FrameTooLarge,
"produce request size 2000000 exceeds max_request_size 1000000"
)));
}
#[test]
fn test_is_batch_too_large_ignores_unrelated_errors() {
assert!(!is_batch_too_large(&KrafkaError::broker(
ErrorCode::NotLeaderForPartition,
"leader moved"
)));
assert!(!is_batch_too_large(&KrafkaError::broker(
ErrorCode::OutOfOrderSequenceNumber,
"oosn"
)));
assert!(!is_batch_too_large(&KrafkaError::timeout("produce")));
assert!(!is_batch_too_large(&KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
"bad response"
)));
}
#[test]
fn test_is_batch_too_large_ignores_response_decode_errors() {
assert!(!is_batch_too_large(&KrafkaError::protocol_kind(
ProtocolErrorKind::InvalidLength,
"array length 9999999 exceeds MAX_DECODE_ARRAY_LEN"
)));
assert!(!is_batch_too_large(&KrafkaError::protocol_kind(
ProtocolErrorKind::TruncatedFrame,
"buffer exhausted decoding ProduceResponse"
)));
}
#[test]
fn test_batch_split_is_bounded_and_lossless() {
fn split(records: usize, depth: u8, leaves: &mut Vec<usize>) {
if records > 1 && depth < MAX_BATCH_SPLIT_DEPTH {
let mid = records / 2;
split(mid, depth + 1, leaves);
split(records - mid, depth + 1, leaves);
} else {
leaves.push(records);
}
}
let mut leaves = Vec::new();
split(100, 0, &mut leaves);
assert_eq!(leaves.iter().sum::<usize>(), 100, "no records lost");
assert!(
leaves.len() <= 1 << MAX_BATCH_SPLIT_DEPTH,
"recursion must stay bounded, got {} leaves",
leaves.len()
);
let mut single = Vec::new();
split(1, 0, &mut single);
assert_eq!(single, vec![1]);
}
#[test]
fn test_cpu_heavy_compression_selection() {
assert!(is_cpu_heavy_compression(Compression::Gzip));
assert!(is_cpu_heavy_compression(Compression::Zstd));
assert!(!is_cpu_heavy_compression(Compression::None));
assert!(!is_cpu_heavy_compression(Compression::Snappy));
assert!(!is_cpu_heavy_compression(Compression::Lz4));
}
}