mod accumulator;
mod config;
mod idempotent;
mod partitioner;
mod record;
mod retry;
mod transaction;
pub use accumulator::{
AccumulatorConfig, DeliveryHandle, RecordAccumulator, RecordAccumulatorHandle,
};
pub use config::{Acks, ProducerConfig};
pub use idempotent::{
PartitionSequenceSnapshot, ProducerIdentity, ProducerIdentitySnapshot, ProducerStateStore,
RollbackOutcome,
};
pub use partitioner::{
DefaultPartitioner, HashPartitioner, Partitioner, RoundRobinPartitioner, StickyPartitioner,
UniformStickyPartitioner, murmur2,
};
pub use record::{DeliveryConfirmation, ProducerRecord, RecordMetadata};
pub use retry::{RetryContext, RetryPolicy};
pub use transaction::{
PreparedTxnState, TopicPartitionOffset, TransactionOutcome, TransactionState,
TransactionVersion, TransactionalProducer, TransactionalProducerBuilder,
TransactionalProducerConfig,
};
use std::sync::Arc;
use std::time::{Duration, Instant};
use bytes::Bytes;
use tracing::{debug, info, warn};
use crate::PartitionId;
use crate::auth::AuthConfig;
use crate::error::{ErrorCode, KrafkaError, ProtocolErrorKind, Result};
use crate::metadata::ClusterMetadata;
use crate::metrics::{ConnectionMetrics, ProducerMetrics as ProducerMetricsInner};
use crate::network::{ConnectionConfig, ConnectionPool};
use crate::protocol::{
ApiKey, Compression, InitProducerIdRequest, InitProducerIdResponse, ProduceRequest,
ProduceResponse, VersionedDecode, VersionedEncode, versions,
};
use crate::serdes::Serializer;
use self::idempotent::ErasedProducerStateStore;
use crate::barrier::InFlightBarrier;
pub struct Producer {
config: ProducerConfig,
metadata: Arc<ClusterMetadata>,
pool: Arc<ConnectionPool>,
partitioner: Arc<dyn Partitioner>,
accumulator: RecordAccumulatorHandle,
in_flight_barrier: Arc<InFlightBarrier>,
metrics: Arc<ProducerMetricsInner>,
interceptor: Arc<dyn crate::interceptor::ProducerInterceptor>,
identity: Option<Arc<ProducerIdentity>>,
key_serializer: Option<Arc<dyn Serializer>>,
value_serializer: Option<Arc<dyn Serializer>>,
pool_owned: bool,
}
async fn init_idempotent_producer_id(
identity: &ProducerIdentity,
metadata: &ClusterMetadata,
retry_policy: &RetryPolicy,
) -> Result<()> {
let started_at = Instant::now();
for attempt in 0..=retry_policy.max_retries {
if let Some(deadline) = retry_policy.delivery_timeout
&& started_at.elapsed() >= deadline
{
return Err(KrafkaError::timeout("InitProducerId"));
}
if attempt > 0 {
let mut backoff = retry_policy.calculate_backoff(attempt);
if let Some(deadline) = retry_policy.delivery_timeout {
let elapsed = started_at.elapsed();
if elapsed >= deadline {
return Err(KrafkaError::timeout("InitProducerId"));
}
backoff = backoff.min(deadline.saturating_sub(elapsed));
}
if !backoff.is_zero() {
tokio::time::sleep(backoff).await;
}
}
if let Some(deadline) = retry_policy.delivery_timeout
&& started_at.elapsed() >= deadline
{
return Err(KrafkaError::timeout("InitProducerId"));
}
let brokers = metadata.brokers();
if brokers.is_empty() {
if attempt < retry_policy.max_retries {
warn!(attempt, "No brokers available for InitProducerId, retrying");
continue;
}
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
"no brokers available for InitProducerId",
));
}
let broker = &brokers[attempt as usize % brokers.len()];
let conn = match metadata.get_broker_connection(broker.id()).await {
Ok(connection) => connection,
Err(error) if error.is_retriable() && attempt < retry_policy.max_retries => {
warn!(
attempt,
error = %error,
"Connection failed for InitProducerId, retrying"
);
continue;
}
Err(error) => return Err(error),
};
let ip_version = match conn.negotiate_api_version(
ApiKey::InitProducerId,
versions::INIT_PRODUCER_ID_MAX,
versions::INIT_PRODUCER_ID_MIN,
) {
Some(version) => version,
None => {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
"no mutually supported InitProducerId API version",
));
}
};
let request = InitProducerIdRequest::idempotent();
let response_bytes = match conn
.send_request(ApiKey::InitProducerId, ip_version, |buf| {
request.encode_versioned(ip_version, buf)
})
.await
{
Ok(bytes) => bytes,
Err(error) if error.is_retriable() && attempt < retry_policy.max_retries => {
warn!(
attempt,
error = %error,
"InitProducerId request failed, retrying"
);
continue;
}
Err(error) => return Err(error),
};
let mut buf = response_bytes;
let response = InitProducerIdResponse::decode_versioned(ip_version, &mut buf)?;
if response.is_ok() {
identity.initialize(response.producer_id, response.producer_epoch);
info!(
"Idempotent producer initialized: PID={}, epoch={}",
response.producer_id, response.producer_epoch
);
return Ok(());
}
if response.error_code.is_retriable() && attempt < retry_policy.max_retries {
warn!(
error_code = ?response.error_code,
attempt,
"InitProducerId returned retriable error, retrying"
);
} else {
return Err(KrafkaError::broker(
response.error_code,
"failed to initialize producer ID",
));
}
}
Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!(
"InitProducerId retry loop exhausted after {} retries",
retry_policy.max_retries
),
))
}
async fn ensure_idempotent_producer_id_initialized(
identity: &ProducerIdentity,
metadata: &ClusterMetadata,
retry_policy: &RetryPolicy,
) -> Result<()> {
if identity.take_reinit_request() {
warn!(
"Producer identity invalidated; obtaining a fresh producer ID and \
restarting all partition sequences"
);
} else if identity.is_initialized() {
return Ok(());
}
init_idempotent_producer_id(identity, metadata, retry_policy).await
}
async fn recover_unknown_producer_id(
identity: &ProducerIdentity,
metadata: &ClusterMetadata,
retry_policy: &RetryPolicy,
topic: &str,
partition: PartitionId,
base_sequence: i32,
record_count: i32,
) -> Result<i32> {
if identity.needs_reinit() {
return Err(KrafkaError::broker(
ErrorCode::UnknownProducerId,
"producer identity is awaiting re-initialisation; retry the send",
));
}
if !identity.check_and_reset_if_retryable(topic, partition, base_sequence, record_count)? {
warn!(
topic,
partition,
"UnknownProducerId while newer batches for this partition are still in flight; \
requesting a producer-ID re-init instead of failing permanently"
);
identity.request_reinit();
return Err(KrafkaError::broker(
ErrorCode::UnknownProducerId,
format!(
"UnknownProducerId for {topic}-{partition} could not be resolved in place \
while newer batches were in flight; the producer ID will be re-initialised \
and this send should be retried"
),
));
}
init_idempotent_producer_id(identity, metadata, retry_policy).await?;
identity.allocate_sequence(topic, partition, record_count)
}
pub(crate) fn out_of_order_data_loss_error(
topic: &str,
partition: PartitionId,
base_sequence: i32,
) -> KrafkaError {
KrafkaError::invalid_state(format!(
"fatal OUT_OF_ORDER_SEQUENCE_NUMBER for {topic}-{partition} at base sequence \
{base_sequence}: the broker expected a different sequence, which means an earlier \
batch was never durably stored (log truncation or unclean leader election). \
Retrying would silently write this batch into the resulting gap, so the send is \
failed instead. Recreate the producer to resume."
))
}
fn apply_produce_leader_hint(
metadata: &ClusterMetadata,
topic: &str,
partition: PartitionId,
response: &ProduceResponse,
partition_response: &crate::protocol::ProducePartitionResponse,
) -> bool {
if !matches!(
partition_response.error_code,
ErrorCode::NotLeaderForPartition | ErrorCode::FencedLeaderEpoch
) {
return false;
}
let Some(leader) = partition_response.current_leader else {
return false;
};
let applied = metadata.apply_leader_hint(
topic,
partition,
leader.leader_id,
leader.leader_epoch,
crate::metadata::broker_info_for_node(&response.node_endpoints, leader.leader_id),
);
if applied {
debug!(
topic,
partition,
leader_id = leader.leader_id,
leader_epoch = leader.leader_epoch,
"broker named a new leader; retrying there without a metadata refresh (KIP-951)"
);
}
applied
}
fn request_header_size(api_key: ApiKey, api_version: i16, client_id: &str) -> Result<usize> {
if client_id.len() > i16::MAX as usize {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::InvalidLength,
format!(
"client_id length {} exceeds protocol limit of {}",
client_id.len(),
i16::MAX
),
));
}
let base = 2 + 2 + 4 + 2 + client_id.len();
match crate::protocol::RequestHeader::header_version(api_key, api_version) {
1 => Ok(base),
2 => Ok(base + 1), version => Err(KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
format!("unsupported request header version {version}"),
)),
}
}
fn encode_and_validate_produce_request(
client_id: &str,
max_request_size: usize,
api_version: i16,
request: &ProduceRequest,
) -> Result<Bytes> {
let mut body = bytes::BytesMut::new();
request.encode_versioned(api_version, &mut body)?;
let frame_size = 4 + request_header_size(ApiKey::Produce, api_version, client_id)? + body.len();
if frame_size > max_request_size {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::FrameTooLarge,
format!(
"produce request size {frame_size} exceeds max_request_size {max_request_size}"
),
));
}
Ok(body.freeze())
}
pub(crate) fn fill_produce_topic_ids(
request: &mut ProduceRequest,
metadata: &ClusterMetadata,
) -> bool {
let mut all_resolved = true;
for topic_data in &mut request.topic_data {
if topic_data.topic_id.is_none() {
if let Some(id) = metadata.topic_id_for_name(&topic_data.name) {
topic_data.topic_id = Some(id);
} else {
all_resolved = false;
}
}
}
all_resolved
}
impl std::fmt::Debug for Producer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Producer")
.field("client_id", &self.config.client_id)
.field("idempotent", &self.identity.is_some())
.field("connections", &self.pool.len())
.field("owns_pool", &self.owns_pool())
.finish_non_exhaustive()
}
}
impl Producer {
pub fn builder() -> ProducerBuilder {
ProducerBuilder::default()
}
async fn new(
config: ProducerConfig,
interceptor: Arc<dyn crate::interceptor::ProducerInterceptor>,
partitioner: Option<Arc<dyn Partitioner>>,
key_serializer: Option<Arc<dyn Serializer>>,
value_serializer: Option<Arc<dyn Serializer>>,
shared: Option<(Arc<ConnectionPool>, Arc<crate::metadata::ClusterMetadata>)>,
state_store: Option<Arc<dyn ErasedProducerStateStore>>,
) -> Result<Self> {
let pool_owned = shared.is_none();
let (pool, metadata) = if let Some((pool, metadata)) = shared {
(pool, metadata)
} else {
let mut pool_config_builder = config.transport.apply(
ConnectionConfig::builder()
.client_id(&config.client_id)
.request_timeout(config.request_timeout)
.connect_timeout(config.connect_timeout),
);
if let Some(ref auth) = config.auth {
pool_config_builder = pool_config_builder.auth(auth.clone());
}
let mut pool_config = pool_config_builder.build()?;
pool_config.init_tls().await?;
let pool = config.transport.build_pool(pool_config);
let bootstrap_servers =
crate::util::parse_bootstrap_servers(&config.bootstrap_servers)?;
let metadata = Arc::new({
let mut meta =
ClusterMetadata::new(bootstrap_servers, pool.clone(), config.metadata_max_age)
.with_recovery_strategy(config.metadata_recovery_strategy)
.with_rebootstrap_trigger(config.metadata_recovery_rebootstrap_trigger);
if let Some(ttl) = config.metadata_topic_cache_ttl {
meta = meta.with_topic_cache_ttl(ttl);
} else {
meta = meta.with_topic_cache_ttl_disabled();
}
meta
});
metadata.refresh().await?;
info!(
"Producer initialized with {} brokers",
metadata.brokers().len()
);
(pool, metadata)
};
let init_retry_policy = RetryPolicy::new()
.with_max_retries(config.retries)
.with_initial_backoff(config.retry_backoff)
.with_max_backoff(Duration::from_secs(10))
.with_delivery_timeout(Some(config.delivery_timeout));
let identity = if config.idempotent {
let identity = Arc::new(ProducerIdentity::new());
init_idempotent_producer_id(&identity, &metadata, &init_retry_policy).await?;
if let Some(ref store) = state_store {
match store.load_erased().await {
Ok(Some(snapshot))
if snapshot.producer_id == identity.producer_id()
&& snapshot.producer_epoch == identity.producer_epoch() =>
{
identity.restore_from_snapshot(&snapshot);
info!(
pid = identity.producer_id(),
epoch = identity.producer_epoch(),
partitions = snapshot.partition_sequences.len(),
"Producer identity restored from state store"
);
}
Ok(Some(_)) => {
debug!(
"State store snapshot PID/epoch mismatch — sequences not restored \
(expected for new transactional sessions or plain idempotent producers)"
);
}
Ok(None) => {
debug!("No previous producer state found in state store");
}
Err(err) => {
warn!(error = %err, "Failed to load producer state from store; continuing with fresh state");
}
}
}
Some(identity)
} else {
None
};
let partitioner: Arc<dyn Partitioner> =
partitioner.unwrap_or_else(|| Arc::new(UniformStickyPartitioner::new()));
let retry_policy = RetryPolicy::new()
.with_max_retries(config.retries)
.with_initial_backoff(config.retry_backoff)
.with_max_backoff(Duration::from_secs(30))
.with_delivery_timeout(Some(config.delivery_timeout));
let metrics = Arc::new(ProducerMetricsInner::default());
if config.buffer_memory == 0 {
warn!(
"buffer_memory=0 disables producer backpressure; \
memory usage is unbounded. Not recommended for production."
);
}
let in_flight_barrier = Arc::new(InFlightBarrier::new());
let acc_config = accumulator::AccumulatorConfig {
batch_size: config.batch_size,
linger: config.linger,
compression: config.compression,
compression_level: config.compression_level,
topic_compression: config.topic_compression.clone().into_iter().collect(),
acks: config.acks.to_i16(),
client_id: config.client_id.clone(),
request_timeout: config.request_timeout,
max_request_size: config.max_request_size,
buffer_memory: config.buffer_memory,
max_block_ms: config.max_block,
interceptor: interceptor.clone(),
identity: identity.clone(),
partitioner: partitioner.clone(),
state_store: state_store.clone(),
transactional_id: None,
dead_letter_queue: config.dead_letter_queue.clone(),
};
let accumulator = accumulator::RecordAccumulator::spawn(
acc_config,
metadata.clone(),
retry_policy.clone(),
metrics.clone(),
in_flight_barrier.clone(),
);
Ok(Self {
config: config.clone(),
metadata,
pool,
partitioner,
accumulator,
in_flight_barrier,
metrics,
interceptor,
identity,
key_serializer,
value_serializer,
pool_owned,
})
}
pub async fn send(
&self,
topic: &str,
key: Option<&[u8]>,
value: &[u8],
) -> Result<RecordMetadata> {
let mut record = ProducerRecord::new(topic, Bytes::copy_from_slice(value));
if let Some(k) = key {
record = record.with_key(Bytes::copy_from_slice(k));
}
self.send_record(record).await
}
pub async fn send_with_headers(
&self,
topic: &str,
key: Option<&[u8]>,
value: &[u8],
headers: Vec<(String, Bytes)>,
) -> Result<RecordMetadata> {
let mut record = ProducerRecord::new(topic, Bytes::copy_from_slice(value));
if let Some(k) = key {
record = record.with_key(Bytes::copy_from_slice(k));
}
record.headers = headers;
self.send_record(record).await
}
pub async fn send_record(&self, record: ProducerRecord) -> Result<RecordMetadata> {
self.enqueue(record).await?.await
}
pub async fn enqueue(&self, record: ProducerRecord) -> Result<DeliveryHandle> {
let send_started_at = Instant::now();
let operation_guard = self.in_flight_barrier.start("producer")?;
let mut record = record;
crate::interceptor::safe_on_send(&*self.interceptor, &mut record);
if let Some(enc) = &self.value_serializer {
record.value = enc
.serialize(
record.value.clone(),
&record.topic,
record.record_name.as_deref(),
false,
)
.await?;
}
if let Some(enc) = &self.key_serializer {
let key = record.key.clone().unwrap_or_default();
record.key = Some(
enc.serialize(key, &record.topic, record.record_name.as_deref(), true)
.await?,
);
}
record.validate()?;
let record_size = record.estimated_size();
let routed = record.into_routed_parts();
let topic = routed.topic;
let record = routed.record;
let partition = match routed.partition {
Some(p) => p,
None => {
let partition_count = self
.metadata
.partition_count(topic.as_ref())
.ok_or_else(|| KrafkaError::invalid_state(format!("unknown topic: {topic}")))?;
self.partitioner
.partition(topic.as_ref(), record.key_bytes(), partition_count)
}
};
self.accumulator
.enqueue_routed_with_guard(
topic,
record,
record_size,
partition,
operation_guard,
send_started_at,
)
.await
}
pub async fn flush(&self) -> Result<()> {
let target = self.in_flight_barrier.snapshot();
self.accumulator.flush().await?;
self.in_flight_barrier.wait_for(target).await;
Ok(())
}
pub fn update_seed_brokers(&self, servers: Vec<String>) -> Result<()> {
self.metadata.update_seed_brokers(servers)
}
pub async fn refresh_tls(&self) -> Result<()> {
self.pool.refresh_tls().await
}
pub async fn rebootstrap(&self) {
self.metadata.rebootstrap().await;
}
pub async fn close(&self) {
let _ = self.close_inner(None).await;
}
pub async fn close_with_timeout(&self, timeout: Duration) -> Result<()> {
self.close_inner(Some(timeout)).await
}
async fn close_inner(&self, timeout: Option<Duration>) -> Result<()> {
let Some(target) = self.in_flight_barrier.begin_close() else {
return Ok(());
};
let graceful_close = async {
if let Err(e) = self.accumulator.shutdown().await {
warn!("Accumulator shutdown error during close: {e}");
}
self.in_flight_barrier.wait_for(target).await;
};
let close_result = if let Some(timeout) = timeout {
tokio::time::timeout(timeout, graceful_close)
.await
.map_err(|_| {
warn!("Producer close timed out; batches in retry backoff may be lost");
KrafkaError::timeout("producer close")
})
} else {
graceful_close.await;
Ok(())
};
crate::interceptor::safe_producer_close(&*self.interceptor);
if self.pool_owned {
self.pool.close_all().await;
info!("Producer closed (connection pool torn down)");
} else {
info!("Producer closed (shared connection pool left open)");
}
close_result
}
#[inline]
#[must_use]
pub fn owns_pool(&self) -> bool {
self.pool_owned
}
#[inline]
pub fn is_closed(&self) -> bool {
self.in_flight_barrier.is_closing()
}
pub fn metrics(&self) -> ProducerMetricsSnapshot {
ProducerMetricsSnapshot {
connections: self.pool.len(),
records_sent: self.metrics.records_sent.get(),
bytes_sent: self.metrics.bytes_sent.get(),
errors: self.metrics.errors.get(),
retries: self.metrics.retries.get(),
buffered_records: self.metrics.buffered_records.get(),
}
}
#[inline]
pub fn metrics_handle(&self) -> Arc<ProducerMetricsInner> {
self.metrics.clone()
}
#[inline]
pub fn connection_metrics(&self) -> Arc<ConnectionMetrics> {
self.pool.metrics()
}
}
impl Drop for Producer {
fn drop(&mut self) {
if self.in_flight_barrier.is_closing() {
return;
}
if let Ok(runtime) = tokio::runtime::Handle::try_current() {
let accumulator = self.accumulator.clone();
drop(runtime.spawn(async move {
if let Err(err) = accumulator.shutdown().await {
warn!(error = %err, "Best-effort flush on Producer drop failed");
}
}));
}
if !std::thread::panicking() {
warn!(
"Producer dropped without close(); buffered batches are being flushed on a \
detached task with no completion guarantee and may still be lost. Call \
`Producer::close()` (or `close_with_timeout`) before drop."
);
}
}
}
#[derive(Debug, Clone)]
#[non_exhaustive]
pub struct ProducerMetricsSnapshot {
pub connections: usize,
pub records_sent: u64,
pub bytes_sent: u64,
pub errors: u64,
pub retries: u64,
pub buffered_records: u64,
}
#[derive(Default)]
#[must_use = "builders do nothing until .build() is called"]
pub struct ProducerBuilder {
config: ProducerConfig,
interceptors: Vec<Arc<dyn crate::interceptor::ProducerInterceptor>>,
partitioner: Option<Arc<dyn Partitioner>>,
key_serializer: Option<Arc<dyn Serializer>>,
value_serializer: Option<Arc<dyn Serializer>>,
shared: Option<(Arc<ConnectionPool>, Arc<crate::metadata::ClusterMetadata>)>,
state_store: Option<Arc<dyn ErasedProducerStateStore>>,
}
impl ProducerBuilder {
pub fn bootstrap_servers(mut self, servers: impl Into<String>) -> Self {
self.config.bootstrap_servers = servers.into();
self
}
pub fn client_id(mut self, client_id: impl Into<String>) -> Self {
self.config.client_id = client_id.into();
self
}
pub fn acks(mut self, acks: Acks) -> Self {
self.config.acks = acks;
self
}
pub fn topic_compression(mut self, topic: impl Into<String>, compression: Compression) -> Self {
self.config
.topic_compression
.insert(topic.into(), compression);
self
}
pub fn buffer_memory(mut self, bytes: usize) -> Self {
self.config.buffer_memory = bytes;
self
}
pub fn max_block(mut self, duration: Duration) -> Self {
self.config.max_block = duration;
self
}
pub fn dead_letter_queue(mut self, dlq: Arc<dyn crate::dlq::DeadLetterQueue>) -> Self {
self.config.dead_letter_queue = Some(dlq);
self
}
pub fn metadata_recovery_strategy(
mut self,
strategy: crate::metadata::MetadataRecoveryStrategy,
) -> Self {
self.config.metadata_recovery_strategy = strategy;
self
}
pub fn metadata_recovery_rebootstrap_trigger(mut self, duration: Duration) -> Self {
self.config.metadata_recovery_rebootstrap_trigger = duration;
self
}
#[cfg(feature = "socks5")]
pub fn proxy(mut self, proxy: crate::network::ProxyConfig) -> Self {
self.config.transport.proxy = Some(proxy);
self
}
pub fn transport(mut self, transport: crate::network::TransportConfig) -> Self {
self.config.transport = transport;
self
}
pub fn compression(mut self, compression: Compression) -> Self {
self.config.compression = compression;
self
}
pub fn compression_level(mut self, level: Option<i32>) -> Self {
self.config.compression_level = level;
self
}
pub fn batch_size(mut self, size: usize) -> Self {
self.config.batch_size = size;
self
}
pub fn linger(mut self, duration: Duration) -> Self {
self.config.linger = duration;
self
}
pub fn request_timeout(mut self, timeout: Duration) -> Self {
self.config.request_timeout = timeout;
self
}
pub fn connect_timeout(mut self, timeout: Duration) -> Self {
self.config.connect_timeout = timeout;
self
}
pub fn delivery_timeout(mut self, timeout: Duration) -> Self {
self.config.delivery_timeout = timeout;
self
}
pub fn retries(mut self, retries: u32) -> Self {
self.config.retries = retries;
self
}
pub fn retry_backoff(mut self, backoff: Duration) -> Self {
self.config.retry_backoff = backoff;
self
}
pub fn max_request_size(mut self, bytes: usize) -> Self {
self.config.max_request_size = bytes;
self
}
pub fn metadata_max_age(mut self, duration: Duration) -> Self {
self.config.metadata_max_age = duration;
self
}
pub fn metadata_topic_cache_ttl(mut self, ttl: Duration) -> Self {
self.config.metadata_topic_cache_ttl = Some(ttl);
self
}
pub fn disable_metadata_topic_cache_ttl(mut self) -> Self {
self.config.metadata_topic_cache_ttl = None;
self
}
pub fn idempotent(mut self, enable: bool) -> Self {
self.config.idempotent = enable;
self
}
pub fn auth(mut self, auth: AuthConfig) -> Self {
self.config.auth = Some(auth);
self
}
pub fn sasl_plain(
mut self,
username: impl Into<String>,
password: impl Into<String>,
) -> crate::Result<Self> {
self.config.auth = Some(AuthConfig::sasl_plain(username, password)?);
Ok(self)
}
pub fn sasl_scram_sha256(
mut self,
username: impl Into<String>,
password: impl Into<String>,
) -> Self {
self.config.auth = Some(AuthConfig::sasl_scram_sha256(username, password));
self
}
pub fn sasl_scram_sha512(
mut self,
username: impl Into<String>,
password: impl Into<String>,
) -> Self {
self.config.auth = Some(AuthConfig::sasl_scram_sha512(username, password));
self
}
pub fn sasl_oauthbearer(mut self, token: impl Into<String>) -> Self {
self.config.auth = Some(AuthConfig::sasl_oauthbearer(token));
self
}
pub fn sasl_oauthbearer_provider(
mut self,
provider: impl crate::auth::OAuthBearerTokenProvider + 'static,
) -> Self {
self.config.auth = Some(AuthConfig::sasl_oauthbearer_provider(provider));
self
}
pub fn partitioner(mut self, partitioner: impl Partitioner + 'static) -> Self {
self.partitioner = Some(Arc::new(partitioner));
self
}
pub fn interceptor(
mut self,
interceptor: Arc<dyn crate::interceptor::ProducerInterceptor>,
) -> Self {
self.interceptors = vec![interceptor];
self
}
pub fn add_interceptor(
mut self,
interceptor: Arc<dyn crate::interceptor::ProducerInterceptor>,
) -> Self {
self.interceptors.push(interceptor);
self
}
pub fn key_serializer(mut self, encoder: Arc<dyn Serializer>) -> Self {
self.key_serializer = Some(encoder);
self
}
pub fn value_serializer(mut self, encoder: Arc<dyn Serializer>) -> Self {
self.value_serializer = Some(encoder);
self
}
pub fn with_client(mut self, client: &crate::client::KrafkaClient) -> Self {
self.shared = Some((client.pool().clone(), client.metadata().clone()));
self
}
pub fn state_store(
mut self,
store: impl crate::producer::ProducerStateStore + 'static,
) -> Self {
self.state_store = Some(Arc::new(store));
self
}
pub fn build_config(self) -> Result<ProducerConfig> {
let has_shared_pool = self.shared.is_some();
let mut config = self.config;
config::validate(&mut config, has_shared_pool)?;
Ok(config)
}
pub async fn build(mut self) -> Result<Producer> {
config::validate(&mut self.config, self.shared.is_some())?;
let interceptor: Arc<dyn crate::interceptor::ProducerInterceptor> =
if self.interceptors.is_empty() {
Arc::new(crate::interceptor::NoOpProducerInterceptor)
} else if self.interceptors.len() == 1 {
let Some(single) = self.interceptors.into_iter().next() else {
unreachable!("len == 1 verified above");
};
single
} else {
Arc::new(crate::interceptor::ProducerInterceptorChain::new(
self.interceptors,
))
};
let producer = Producer::new(
self.config,
interceptor,
self.partitioner,
self.key_serializer,
self.value_serializer,
self.shared,
self.state_store,
)
.await?;
Ok(producer)
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
use crate::metadata::ClusterMetadata;
use crate::network::{ConnectionConfig, ConnectionPool};
use crate::protocol::{ProducePartitionData, ProduceTopicData};
#[test]
fn test_producer_builder() {
let builder = Producer::builder()
.bootstrap_servers("localhost:9092")
.client_id("test")
.acks(Acks::All)
.compression(Compression::Gzip)
.batch_size(32768)
.max_request_size(65536)
.linger(Duration::from_millis(10));
assert_eq!(builder.config.bootstrap_servers, "localhost:9092");
assert_eq!(builder.config.client_id, "test");
assert_eq!(builder.config.acks, Acks::All);
assert_eq!(builder.config.compression, Compression::Gzip);
assert_eq!(builder.config.batch_size, 32768);
assert_eq!(builder.config.max_request_size, 65536);
assert_eq!(builder.config.linger, Duration::from_millis(10));
assert!(builder.config.auth.is_none());
}
#[test]
fn test_producer_builder_with_auth() {
let builder = Producer::builder()
.bootstrap_servers("broker:9093")
.auth(AuthConfig::sasl_plain("user", "pass").unwrap());
let auth = builder.config.auth.as_ref().unwrap();
assert!(auth.requires_sasl());
assert!(!auth.requires_tls());
assert_eq!(
auth.security_protocol,
crate::auth::SecurityProtocol::SaslPlaintext
);
assert_eq!(auth.sasl_mechanism, Some(crate::auth::SaslMechanism::Plain));
}
#[test]
fn test_producer_builder_aws_msk_iam() {
let auth = AuthConfig::aws_msk_iam("AKID", "secret", "us-east-1");
let builder = Producer::builder()
.bootstrap_servers("broker:9094")
.auth(auth);
let auth = builder.config.auth.as_ref().unwrap();
assert!(auth.requires_tls());
assert!(auth.requires_sasl());
assert_eq!(
auth.sasl_mechanism,
Some(crate::auth::SaslMechanism::AwsMskIam)
);
assert!(auth.aws_msk_iam_credentials.is_some());
assert!(auth.tls_config.is_some());
}
#[test]
fn test_producer_builder_no_auth_by_default() {
let builder = Producer::builder().bootstrap_servers("broker:9092");
assert!(builder.config.auth.is_none());
}
#[test]
fn test_producer_builder_sasl_plain() {
let builder = Producer::builder()
.bootstrap_servers("broker:9093")
.sasl_plain("user", "pass")
.unwrap();
let auth = builder.config.auth.as_ref().unwrap();
assert!(auth.requires_sasl());
assert!(auth.plain_credentials.is_some());
}
#[test]
fn test_producer_builder_sasl_scram() {
let builder = Producer::builder()
.bootstrap_servers("broker:9093")
.sasl_scram_sha256("user", "pass");
let auth = builder.config.auth.as_ref().unwrap();
assert!(auth.requires_sasl());
assert!(auth.scram_credentials.is_some());
let builder = Producer::builder()
.bootstrap_servers("broker:9093")
.sasl_scram_sha512("user", "pass");
let auth = builder.config.auth.as_ref().unwrap();
assert!(auth.requires_sasl());
assert!(auth.scram_credentials.is_some());
}
#[tokio::test]
async fn test_producer_builder_no_servers() {
let result = Producer::builder().build().await;
assert!(result.is_err());
}
#[test]
fn test_producer_builder_retry_config() {
let builder = Producer::builder()
.bootstrap_servers("localhost:9092")
.retries(5)
.retry_backoff(Duration::from_millis(200));
assert_eq!(builder.config.retries, 5);
assert_eq!(builder.config.retry_backoff, Duration::from_millis(200));
}
#[test]
fn test_validate_produce_request_size_rejects_oversized_frame() {
let request = ProduceRequest {
transactional_id: None,
acks: Acks::All.to_i16(),
timeout_ms: 30_000,
topic_data: vec![ProduceTopicData {
name: "topic".to_string(),
topic_id: None,
partition_data: vec![ProducePartitionData {
index: 0,
records: Bytes::from(vec![0; 512]),
}],
}],
};
let error =
encode_and_validate_produce_request("client", 128, versions::PRODUCE_MIN, &request)
.expect_err("oversized frame should be rejected");
assert!(error.to_string().contains("max_request_size"));
}
#[test]
fn test_validate_produce_request_size_uses_exact_flexible_encoding_size() {
let request = ProduceRequest {
transactional_id: Some("txn-123".to_string()),
acks: Acks::All.to_i16(),
timeout_ms: 30_000,
topic_data: vec![ProduceTopicData {
name: "topic".to_string(),
topic_id: Some([0u8; 16]),
partition_data: vec![ProducePartitionData {
index: 0,
records: Bytes::from(vec![1; 32]),
}],
}],
};
let encoded = encode_and_validate_produce_request(
"client",
usize::MAX,
versions::PRODUCE_MAX,
&request,
)
.unwrap();
let exact_size = 4
+ request_header_size(ApiKey::Produce, versions::PRODUCE_MAX, "client").unwrap()
+ encoded.len();
encode_and_validate_produce_request("client", exact_size, versions::PRODUCE_MAX, &request)
.unwrap();
let error = encode_and_validate_produce_request(
"client",
exact_size.saturating_sub(1),
versions::PRODUCE_MAX,
&request,
)
.unwrap_err();
assert!(error.to_string().contains("max_request_size"));
}
#[test]
fn test_validate_produce_request_size_v13_requires_topic_id() {
let request = ProduceRequest {
transactional_id: None,
acks: Acks::All.to_i16(),
timeout_ms: 30_000,
topic_data: vec![ProduceTopicData {
name: "topic".to_string(),
topic_id: None,
partition_data: vec![ProducePartitionData {
index: 0,
records: Bytes::from_static(b"payload"),
}],
}],
};
let error = encode_and_validate_produce_request("client", 1024, 13, &request).unwrap_err();
assert!(error.to_string().contains("topic_id is required"));
}
#[tokio::test]
async fn test_recover_unknown_producer_id_requests_reinit_instead_of_poisoning() {
let identity = ProducerIdentity::new();
identity.initialize(7, 1);
assert_eq!(identity.allocate_sequence("topic", 0, 2).unwrap(), 0);
assert_eq!(identity.allocate_sequence("topic", 0, 1).unwrap(), 2);
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
let metadata = ClusterMetadata::new(
vec!["localhost:9092".to_string()],
pool,
Duration::from_secs(300),
);
let retry_policy = RetryPolicy::default();
let error =
recover_unknown_producer_id(&identity, &metadata, &retry_policy, "topic", 0, 0, 2)
.await
.unwrap_err();
assert!(
matches!(
error,
KrafkaError::Broker {
code: ErrorCode::UnknownProducerId,
..
}
),
"expected a retriable broker error, got: {error:?}"
);
assert!(
!error.to_string().contains("poisoned"),
"the producer must not be described as permanently poisoned"
);
assert!(identity.needs_reinit());
assert_eq!(identity.producer_id(), 7);
assert_eq!(identity.peek_sequence("topic", 0), 3);
assert!(identity.take_reinit_request());
assert!(!identity.is_initialized());
assert_eq!(identity.peek_sequence("topic", 0), 0);
identity.initialize(8, 0);
assert!(identity.is_initialized());
assert_eq!(identity.allocate_sequence("topic", 0, 1).unwrap(), 0);
}
#[test]
fn test_producer_metrics_snapshot() {
let snapshot = ProducerMetricsSnapshot {
connections: 3,
records_sent: 100,
bytes_sent: 50000,
errors: 2,
retries: 5,
buffered_records: 7,
};
assert_eq!(snapshot.connections, 3);
assert_eq!(snapshot.records_sent, 100);
assert_eq!(snapshot.bytes_sent, 50000);
assert_eq!(snapshot.errors, 2);
assert_eq!(snapshot.retries, 5);
assert_eq!(snapshot.buffered_records, 7);
}
#[test]
fn a_record_larger_than_buffer_memory_is_rejected_up_front() {
let err = accumulator::check_record_admission(1024, 16, usize::MAX)
.expect_err("a record larger than buffer_memory must be rejected");
assert!(
err.to_string().contains("buffer_memory"),
"the error must name the setting to raise, got: {err}"
);
}
#[test]
fn test_retry_policy_from_config() {
let policy = RetryPolicy::new()
.with_max_retries(10)
.with_initial_backoff(Duration::from_millis(50))
.with_max_backoff(Duration::from_secs(30));
assert_eq!(policy.max_retries, 10);
assert_eq!(policy.initial_backoff(), Duration::from_millis(50));
assert_eq!(policy.max_backoff(), Duration::from_secs(30));
}
#[test]
fn test_acks_none_returns_fire_and_forget_metadata() {
let builder = Producer::builder()
.bootstrap_servers("localhost:9092")
.acks(Acks::None);
assert_eq!(builder.config.acks, Acks::None);
assert_eq!(builder.config.acks.to_i16(), 0);
}
#[test]
fn test_idempotent_builder() {
let builder = Producer::builder().bootstrap_servers("broker:9092");
assert!(builder.config.idempotent);
let builder = Producer::builder()
.bootstrap_servers("broker:9092")
.idempotent(false);
assert!(!builder.config.idempotent);
}
#[tokio::test]
async fn test_idempotent_requires_acks_all() {
let builder = Producer::builder()
.bootstrap_servers("localhost:9092")
.acks(Acks::Leader)
.idempotent(true);
let result = builder.build().await;
match result {
Err(e) => assert!(e.to_string().contains("acks")),
Ok(_) => panic!("expected config error for idempotent with acks != All"),
}
}
#[tokio::test]
async fn test_producer_builder_rejects_zero_batch_size() {
let mut builder = Producer::builder().bootstrap_servers("localhost:9092");
builder.config.batch_size = 0;
let result = builder.build().await;
match result {
Err(e) => assert!(e.to_string().contains("batch_size")),
Ok(_) => panic!("expected error for batch_size=0"),
}
}
#[test]
fn test_producer_builder_interceptor() {
use crate::interceptor::{InterceptorResult, ProducerInterceptor};
#[derive(Debug)]
struct TestInterceptor;
impl ProducerInterceptor for TestInterceptor {
fn on_send(&self, _record: &mut ProducerRecord) -> InterceptorResult {
Ok(())
}
}
let builder = Producer::builder()
.bootstrap_servers("localhost:9092")
.interceptor(Arc::new(TestInterceptor));
assert_eq!(builder.interceptors.len(), 1);
}
#[test]
fn test_producer_builder_add_interceptor() {
use crate::interceptor::ProducerInterceptor;
#[derive(Debug)]
struct A;
impl ProducerInterceptor for A {}
#[derive(Debug)]
struct B;
impl ProducerInterceptor for B {}
let builder = Producer::builder()
.bootstrap_servers("localhost:9092")
.add_interceptor(Arc::new(A))
.add_interceptor(Arc::new(B));
assert_eq!(builder.interceptors.len(), 2);
}
#[test]
fn test_producer_builder_interceptor_replaces_chain() {
use crate::interceptor::ProducerInterceptor;
#[derive(Debug)]
struct A;
impl ProducerInterceptor for A {}
#[derive(Debug)]
struct B;
impl ProducerInterceptor for B {}
let builder = Producer::builder()
.bootstrap_servers("localhost:9092")
.add_interceptor(Arc::new(A))
.add_interceptor(Arc::new(A))
.interceptor(Arc::new(B));
assert_eq!(builder.interceptors.len(), 1);
}
#[test]
fn test_producer_is_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<Producer>();
}
#[tokio::test]
async fn test_a_sub_ten_second_request_timeout_is_reachable() {
let err = Producer::builder()
.bootstrap_servers("127.0.0.1:1")
.request_timeout(Duration::from_secs(2))
.build()
.await
.expect_err("request_timeout below the default connect_timeout must be rejected");
assert!(
err.to_string().contains("connect_timeout"),
"the error should name the setter to change: {err}"
);
let err = Producer::builder()
.bootstrap_servers("127.0.0.1:1")
.request_timeout(Duration::from_secs(2))
.connect_timeout(Duration::from_secs(2))
.build()
.await
.expect_err("no broker is listening on port 1");
assert!(
!err.to_string().contains("connect_timeout"),
"config validation should have passed, got {err}"
);
}
#[test]
fn test_transaction_version_is_publicly_nameable() {
let version: crate::producer::TransactionVersion = TransactionVersion::V2;
assert_eq!(version, TransactionVersion::V2);
}
}