use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU8, Ordering};
use std::time::Duration;
use bytes::Bytes;
use tokio::sync::{Notify, RwLock};
use tracing::{debug, info, warn};
use crate::auth::AuthConfig;
use crate::error::{ErrorCode, KrafkaError, ProtocolErrorKind, Result};
use crate::metadata::ClusterMetadata;
use crate::network::{BrokerConnection, ConnectionConfig, ConnectionPool};
use crate::protocol::{
AddOffsetsToTxnRequest, AddOffsetsToTxnResponse, AddPartitionsToTxnRequest,
AddPartitionsToTxnResponse, ApiKey, Compression, EndTxnRequest, EndTxnResponse,
FindCoordinatorRequest, FindCoordinatorResponse, InitProducerIdRequest, InitProducerIdResponse,
TxnOffsetCommitRequest, TxnOffsetCommitResponse, VersionedDecode, VersionedEncode, versions,
};
use crate::{Offset, PartitionId};
use super::accumulator::{AccumulatorConfig, RecordAccumulator, RecordAccumulatorHandle};
use super::barrier::InFlightBarrier;
use super::config::Acks;
use super::idempotent::ProducerIdentity;
use super::partitioner::{Partitioner, UniformStickyPartitioner};
use super::record::{ProducerRecord, RecordMetadata, TopicHandle};
use super::retry::RetryPolicy;
use crate::consumer::ConsumerGroupMetadata;
use crate::metrics::ProducerMetrics;
use crate::schema_registry::SchemaEncoder;
const TRANSACTION_VERSION_FEATURE: &str = "transaction.version";
const TV2_MIN_PRODUCE_VERSION: i16 = 12;
const TV2_MIN_TXN_OFFSET_COMMIT_VERSION: i16 = 5;
const TV2_MIN_END_TXN_VERSION: i16 = 4;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)]
#[non_exhaustive]
#[repr(u8)]
pub enum TransactionVersion {
#[default]
V1 = 1,
V2 = 2,
}
impl From<u8> for TransactionVersion {
fn from(v: u8) -> Self {
if v == Self::V2 as u8 {
Self::V2
} else {
Self::V1
}
}
}
impl TransactionVersion {
#[must_use]
pub fn from_feature_level(level: i16) -> Self {
if level >= 2 { Self::V2 } else { Self::V1 }
}
#[must_use]
#[inline]
pub fn is_v2(self) -> bool {
matches!(self, Self::V2)
}
}
impl std::fmt::Display for TransactionVersion {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::V1 => write!(f, "TV1"),
Self::V2 => write!(f, "TV2"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct BrokerTransactionSupport {
transaction_version_level: i16,
produce_max: Option<i16>,
txn_offset_commit_max: Option<i16>,
end_txn_max: Option<i16>,
}
impl BrokerTransactionSupport {
fn version(self) -> TransactionVersion {
let feature = TransactionVersion::from_feature_level(self.transaction_version_level);
if !feature.is_v2() {
return TransactionVersion::V1;
}
let supports =
|negotiated: Option<i16>, required: i16| negotiated.is_some_and(|v| v >= required);
if supports(self.produce_max, TV2_MIN_PRODUCE_VERSION)
&& supports(
self.txn_offset_commit_max,
TV2_MIN_TXN_OFFSET_COMMIT_VERSION,
)
&& supports(self.end_txn_max, TV2_MIN_END_TXN_VERSION)
{
TransactionVersion::V2
} else {
TransactionVersion::V1
}
}
}
fn negotiated_transaction_version(reports: &[BrokerTransactionSupport]) -> TransactionVersion {
reports
.iter()
.map(|r| r.version())
.min()
.unwrap_or(TransactionVersion::V1)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
#[repr(u8)]
pub enum TransactionState {
Uninitialized = 0,
Ready = 1,
InTransaction = 2,
Committing = 3,
Aborting = 4,
FatalError = 5,
Initializing = 6,
CommitIndeterminate = 7,
}
impl From<u8> for TransactionState {
fn from(v: u8) -> Self {
match v {
0 => Self::Uninitialized,
1 => Self::Ready,
2 => Self::InTransaction,
3 => Self::Committing,
4 => Self::Aborting,
5 => Self::FatalError,
6 => Self::Initializing,
7 => Self::CommitIndeterminate,
_ => {
warn!(
discriminant = v,
"unknown TransactionState discriminant — treating as FatalError"
);
Self::FatalError
}
}
}
}
impl std::fmt::Display for TransactionState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(match self {
Self::Uninitialized => "Uninitialized",
Self::Ready => "Ready",
Self::InTransaction => "InTransaction",
Self::Committing => "Committing",
Self::Aborting => "Aborting",
Self::FatalError => "FatalError",
Self::Initializing => "Initializing",
Self::CommitIndeterminate => "CommitIndeterminate",
})
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct TopicPartitionOffset {
pub topic: String,
pub partition: PartitionId,
pub next_offset: Offset,
}
impl TopicPartitionOffset {
pub fn new(topic: impl Into<String>, partition: PartitionId, next_offset: Offset) -> Self {
Self {
topic: topic.into(),
partition,
next_offset,
}
}
}
#[derive(Debug, Clone)]
pub struct TransactionalProducerConfig {
bootstrap_servers: String,
client_id: String,
transactional_id: String,
transaction_timeout_ms: i32,
request_timeout: Duration,
connect_timeout: Duration,
max_request_size: usize,
compression: Compression,
batch_size: usize,
linger: Duration,
buffer_memory: usize,
max_block: Duration,
max_in_flight: usize,
metadata_max_age: Duration,
metadata_recovery_strategy: crate::metadata::MetadataRecoveryStrategy,
metadata_recovery_rebootstrap_trigger: Duration,
auth: Option<AuthConfig>,
#[cfg(feature = "socks5")]
proxy: Option<crate::network::ProxyConfig>,
transport: crate::network::TransportConfig,
}
impl Default for TransactionalProducerConfig {
fn default() -> Self {
Self {
bootstrap_servers: String::new(),
client_id: "krafka-txn-producer".to_string(),
transactional_id: String::new(),
transaction_timeout_ms: 60000,
request_timeout: Duration::from_secs(30),
connect_timeout: crate::network::DEFAULT_CONNECT_TIMEOUT,
max_request_size: crate::protocol::MAX_MESSAGE_SIZE,
compression: Compression::None,
batch_size: 16384,
linger: Duration::from_millis(5),
buffer_memory: 32 * 1024 * 1024,
max_block: Duration::from_secs(60),
max_in_flight: 5,
metadata_max_age: Duration::from_secs(300),
metadata_recovery_strategy: crate::metadata::MetadataRecoveryStrategy::Rebootstrap,
metadata_recovery_rebootstrap_trigger: Duration::from_secs(300),
auth: None,
#[cfg(feature = "socks5")]
proxy: None,
transport: crate::network::TransportConfig::default(),
}
}
}
#[derive(Debug, Clone)]
enum PartitionAddState {
Pending(Arc<Notify>),
Added,
Failed(Arc<KrafkaError>),
}
#[cfg_attr(test, derive(Debug))]
enum BeginAddResult {
AlreadyAdded,
Wait(Arc<Notify>),
NeedAdd(Arc<Notify>),
Fatal(Arc<KrafkaError>),
}
#[derive(Debug, Default)]
struct TransactionPartitions {
partitions: std::collections::HashMap<
String,
std::collections::HashMap<PartitionId, PartitionAddState>,
>,
}
impl TransactionPartitions {
fn begin_add(&mut self, topic: &str, partition: PartitionId) -> BeginAddResult {
if let Some(topic_map) = self.partitions.get(topic) {
match topic_map.get(&partition) {
Some(PartitionAddState::Added) => return BeginAddResult::AlreadyAdded,
Some(PartitionAddState::Pending(notify)) => {
return BeginAddResult::Wait(notify.clone());
}
Some(PartitionAddState::Failed(err)) => {
return BeginAddResult::Fatal(err.clone());
}
None => {}
}
}
let notify = Arc::new(Notify::new());
self.partitions
.entry(topic.to_string())
.or_default()
.insert(partition, PartitionAddState::Pending(notify.clone()));
BeginAddResult::NeedAdd(notify)
}
fn confirm_add(&mut self, topic: &str, partition: PartitionId, notify: &Notify) {
self.partitions
.entry(topic.to_string())
.or_default()
.insert(partition, PartitionAddState::Added);
notify.notify_waiters();
}
fn cancel_add(&mut self, topic: &str, partition: PartitionId, notify: &Notify) {
if let Some(topic_map) = self.partitions.get_mut(topic) {
topic_map.remove(&partition);
if topic_map.is_empty() {
self.partitions.remove(topic);
}
}
notify.notify_waiters();
}
fn fail_add(
&mut self,
topic: &str,
partition: PartitionId,
error: Arc<KrafkaError>,
notify: &Notify,
) {
self.partitions
.entry(topic.to_string())
.or_default()
.insert(partition, PartitionAddState::Failed(error));
notify.notify_waiters();
}
fn clear(&mut self) {
self.partitions.clear();
}
#[cfg(test)]
fn is_empty(&self) -> bool {
self.partitions.is_empty()
}
}
struct PendingAddGuard {
txn_partitions: Arc<RwLock<TransactionPartitions>>,
topic: TopicHandle,
partition: PartitionId,
notify: Arc<Notify>,
defused: bool,
}
impl PendingAddGuard {
async fn confirm(mut self, topic: &str, partition: PartitionId) {
self.defused = true;
let mut txn_partitions = self.txn_partitions.write().await;
txn_partitions.confirm_add(topic, partition, &self.notify);
}
async fn cancel(mut self, topic: &str, partition: PartitionId) {
self.defused = true;
let mut txn_partitions = self.txn_partitions.write().await;
txn_partitions.cancel_add(topic, partition, &self.notify);
}
async fn fail(mut self, topic: &str, partition: PartitionId, error: Arc<KrafkaError>) {
self.defused = true;
let mut txn_partitions = self.txn_partitions.write().await;
txn_partitions.fail_add(topic, partition, error, &self.notify);
}
}
impl Drop for PendingAddGuard {
fn drop(&mut self) {
if !self.defused {
let topic = self.topic.clone();
let partition = self.partition;
let notify = self.notify.clone();
if let Ok(mut tp) = self.txn_partitions.try_write() {
tp.cancel_add(&topic, partition, ¬ify);
} else if let Ok(handle) = tokio::runtime::Handle::try_current() {
let txn_partitions = self.txn_partitions.clone();
handle.spawn(async move {
let mut tp = txn_partitions.write().await;
tp.cancel_add(&topic, partition, ¬ify);
});
} else {
let mut tp = self.txn_partitions.blocking_write();
tp.cancel_add(&topic, partition, ¬ify);
}
}
}
}
pub struct TransactionalProducer {
config: TransactionalProducerConfig,
metadata: Arc<ClusterMetadata>,
pool: Arc<ConnectionPool>,
partitioner: Arc<dyn Partitioner>,
state: AtomicU8,
transaction_version: AtomicU8,
abort_required: AtomicBool,
coordinator_id: RwLock<Option<i32>>,
txn_partitions: Arc<RwLock<TransactionPartitions>>,
identity: Arc<ProducerIdentity>,
accumulator: RecordAccumulatorHandle,
metrics: Arc<ProducerMetrics>,
retry_policy: RetryPolicy,
in_flight_barrier: Arc<InFlightBarrier>,
key_encoder: Option<Arc<dyn SchemaEncoder>>,
value_encoder: Option<Arc<dyn SchemaEncoder>>,
}
impl TransactionalProducer {
pub fn builder() -> TransactionalProducerBuilder {
TransactionalProducerBuilder::default()
}
#[inline]
pub fn state(&self) -> TransactionState {
TransactionState::from(self.state.load(Ordering::SeqCst))
}
#[inline]
pub fn transaction_version(&self) -> TransactionVersion {
TransactionVersion::from(self.transaction_version.load(Ordering::SeqCst))
}
#[inline]
fn requires_explicit_partition_registration(&self) -> bool {
!self.transaction_version().is_v2()
}
async fn detect_transaction_version(&self) -> TransactionVersion {
let brokers = self.metadata.brokers();
let mut reports = Vec::with_capacity(brokers.len());
for broker in &brokers {
match self.probe_broker_transaction_support(broker).await {
Ok(report) => reports.push(report),
Err(error) => {
debug!(
broker = broker.id(),
%error,
"Could not read transaction.version from broker; \
excluding it from the negotiated transaction version"
);
}
}
}
let version = negotiated_transaction_version(&reports);
info!(
%version,
brokers_probed = reports.len(),
"Negotiated KIP-890 transaction version"
);
version
}
async fn probe_broker_transaction_support(
&self,
broker: &crate::metadata::BrokerInfo,
) -> Result<BrokerTransactionSupport> {
let conn = self
.pool
.get_connection_by_id(broker.id(), broker.address())
.await?;
let av_version = conn
.negotiate_api_version(ApiKey::ApiVersions, versions::API_VERSIONS_MAX, 3)
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
"broker does not support ApiVersions v3+, so it cannot report finalized features",
)
})?;
let request = crate::protocol::ApiVersionsRequest::new()
.with_client_software("krafka", env!("CARGO_PKG_VERSION"));
let response_bytes = conn
.send_request(ApiKey::ApiVersions, av_version, |buf| {
if av_version >= 5 {
request.encode_v5(buf)
} else {
request.encode_v3(buf)
}
})
.await?;
let mut buf = response_bytes;
let response = crate::protocol::ApiVersionsResponse::decode_v3(&mut buf)?;
if response.error_code != 0 {
return Err(KrafkaError::broker(
ErrorCode::from(response.error_code),
"ApiVersions request failed while reading transaction.version",
));
}
let transaction_version_level = response
.get_finalized_feature(TRANSACTION_VERSION_FEATURE)
.map_or(0, |f| f.max_version_level);
Ok(BrokerTransactionSupport {
transaction_version_level,
produce_max: conn.negotiate_api_version(
ApiKey::Produce,
versions::PRODUCE_MAX,
versions::PRODUCE_MIN,
),
txn_offset_commit_max: conn.negotiate_api_version(
ApiKey::TxnOffsetCommit,
versions::TXN_OFFSET_COMMIT_MAX,
versions::TXN_OFFSET_COMMIT_MIN,
),
end_txn_max: conn.negotiate_api_version(
ApiKey::EndTxn,
versions::END_TXN_MAX,
versions::END_TXN_MIN,
),
})
}
fn checked_transactional_identity(&self) -> Result<(i64, i16)> {
let producer_id = self.identity.producer_id();
let producer_epoch = self.identity.producer_epoch();
if producer_id < 0 || producer_epoch < 0 {
return Err(KrafkaError::invalid_state(
"transactional producer identity not initialized",
));
}
debug_assert!(
producer_id >= 0 && producer_epoch >= 0,
"transactional producer identity must be initialized before sending"
);
Ok((producer_id, producer_epoch))
}
#[inline]
fn abort_required(&self) -> bool {
self.abort_required.load(Ordering::SeqCst)
}
fn ensure_transaction_can_continue(&self, operation: &str) -> Result<()> {
if self.abort_required() {
return Err(KrafkaError::broker(
ErrorCode::TransactionAbortable,
format!("cannot {operation}: abort_transaction() is required before continuing"),
));
}
Ok(())
}
fn mark_unknown_producer_id_abort_required(&self, operation: &str) -> KrafkaError {
self.abort_required.store(true, Ordering::SeqCst);
KrafkaError::broker(
ErrorCode::TransactionAbortable,
format!(
"{operation} failed with UnknownProducerId; abort_transaction() is required before continuing"
),
)
}
fn is_unknown_producer_id_error(error: &KrafkaError) -> bool {
matches!(
error,
KrafkaError::Broker {
code: ErrorCode::UnknownProducerId,
..
}
)
}
fn classify_transaction_result<T>(&self, result: Result<T>) -> Result<T> {
if let Err(KrafkaError::Broker { code, .. }) = &result
&& is_fatal_transaction_error(*code, self.transaction_version())
{
warn!(
error_code = ?code,
"Fatal transactional error from coordinator; producer is fenced and must be recreated"
);
self.set_state(TransactionState::FatalError);
return result;
}
if let Err(error) = &result
&& Self::is_abortable_transaction_error(error, self.transaction_version())
{
self.abort_required.store(true, Ordering::SeqCst);
}
result
}
fn is_abortable_transaction_error(error: &KrafkaError, version: TransactionVersion) -> bool {
let KrafkaError::Broker { code, .. } = error else {
return false;
};
match code {
ErrorCode::TransactionAbortable => true,
ErrorCode::InvalidProducerIdMapping => !version.is_v2(),
_ => false,
}
}
async fn coordinator_connection(&self, attempt: u32) -> Result<(i32, Arc<BrokerConnection>)> {
let coordinator_id = {
let cached = *self.coordinator_id.read().await;
match cached {
Some(id) => id,
None => {
let id = self.find_coordinator(attempt).await?;
*self.coordinator_id.write().await = Some(id);
debug!("Auto-discovered transaction coordinator: broker {}", id);
id
}
}
};
let brokers = self.metadata.brokers();
let broker = brokers
.iter()
.find(|b| b.id() == coordinator_id)
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
"coordinator not found in metadata",
)
})?;
let conn = self
.pool
.get_connection_by_id(broker.id(), broker.address())
.await?;
Ok((coordinator_id, conn))
}
fn needs_coordinator_refresh(err: &KrafkaError) -> bool {
match err {
KrafkaError::Broker { code, .. } => matches!(
code,
ErrorCode::NotCoordinator
| ErrorCode::CoordinatorNotAvailable
| ErrorCode::CoordinatorLoadInProgress
),
KrafkaError::Network(_) | KrafkaError::Timeout { .. } => true,
_ => false,
}
}
async fn invalidate_coordinator(&self) {
*self.coordinator_id.write().await = None;
}
async fn retry_with_coordinator<F, Fut>(&self, op_name: &str, op: F) -> Result<()>
where
F: Fn(u32) -> Fut,
Fut: Future<Output = Result<()>>,
{
let max_retries = self.retry_policy.max_retries;
for attempt in 0..=max_retries {
if attempt > 0 {
tokio::time::sleep(self.retry_policy.calculate_backoff(attempt)).await;
}
let result = op(attempt).await;
match &result {
Ok(()) => return Ok(()),
Err(e) if Self::is_unknown_producer_id_error(e) => return result,
Err(e) if Self::needs_coordinator_refresh(e) && attempt < max_retries => {
warn!(
attempt,
error = %e,
op_name,
"Coordinator error, refreshing and retrying"
);
self.invalidate_coordinator().await;
}
Err(e) if e.is_retriable() && attempt < max_retries => {
warn!(
attempt,
error = %e,
op_name,
"Retriable error, retrying"
);
}
Err(_) => return result,
}
}
Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!("{op_name} retry loop exhausted after {max_retries} retries"),
))
}
fn set_state(&self, state: TransactionState) {
self.state.store(state as u8, Ordering::SeqCst);
}
fn try_transition(
&self,
expected: TransactionState,
new: TransactionState,
) -> std::result::Result<(), TransactionState> {
self.state
.compare_exchange(
expected as u8,
new as u8,
Ordering::AcqRel,
Ordering::Acquire,
)
.map(|_| ())
.map_err(TransactionState::from)
}
pub async fn init_transactions(&self) -> Result<()> {
if let Err(actual) = self.try_transition(
TransactionState::Uninitialized,
TransactionState::Initializing,
) {
return Err(KrafkaError::invalid_state(format!(
"init_transactions can only be called once (state={:?})",
actual
)));
}
let version = self.detect_transaction_version().await;
self.transaction_version
.store(version as u8, Ordering::SeqCst);
let result = self.do_init_transactions().await;
if result.is_err() {
self.set_state(TransactionState::Uninitialized);
}
result
}
async fn do_init_transactions(&self) -> Result<()> {
self.retry_with_coordinator("InitProducerId", |attempt| async move {
let (_coordinator_id, conn) = self.coordinator_connection(attempt).await?;
let ip_version = conn
.negotiate_api_version(
ApiKey::InitProducerId,
versions::INIT_PRODUCER_ID_MAX,
versions::INIT_PRODUCER_ID_MIN,
)
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
"no mutually supported InitProducerId API version",
)
})?;
let request = InitProducerIdRequest::transactional(
&self.config.transactional_id,
self.config.transaction_timeout_ms,
);
let response_bytes = conn
.send_request(ApiKey::InitProducerId, ip_version, |buf| {
request.encode_versioned(ip_version, buf)
})
.await?;
let mut buf = response_bytes;
let response = InitProducerIdResponse::decode_versioned(ip_version, &mut buf)?;
if !response.is_ok() {
return Err(KrafkaError::broker(
response.error_code,
"failed to initialize producer ID",
));
}
self.identity
.initialize(response.producer_id, response.producer_epoch);
self.abort_required.store(false, Ordering::SeqCst);
self.set_state(TransactionState::Ready);
info!(
"Transactional producer initialized: PID={}, epoch={}",
response.producer_id, response.producer_epoch
);
Ok(())
})
.await
}
async fn find_coordinator(&self, attempt: u32) -> Result<i32> {
let brokers = self.metadata.brokers();
if brokers.is_empty() {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
"no brokers available",
));
}
let broker = &brokers[attempt as usize % brokers.len()];
let conn = self
.pool
.get_connection_by_id(broker.id(), broker.address())
.await?;
let request = FindCoordinatorRequest::for_transaction(&self.config.transactional_id);
let fc_version = conn
.negotiate_api_version(
ApiKey::FindCoordinator,
versions::FIND_COORDINATOR_MAX,
versions::FIND_COORDINATOR_MIN,
)
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
"no mutually supported FindCoordinator API version; \
transactional coordinator lookup requires v1+",
)
})?;
let response_bytes = conn
.send_request(ApiKey::FindCoordinator, fc_version, |buf| {
request.encode_versioned(fc_version, buf)
})
.await?;
let mut buf = response_bytes;
let response = FindCoordinatorResponse::decode_versioned(fc_version, &mut buf)?;
if !response.error_code.is_ok() {
return Err(KrafkaError::broker(
response.error_code,
"failed to find transaction coordinator",
));
}
debug!(
"Found transaction coordinator: broker {} at {}:{}",
response.node_id, response.host, response.port
);
Ok(response.node_id)
}
pub fn begin_transaction(&self) -> Result<()> {
if let Err(actual) =
self.try_transition(TransactionState::Ready, TransactionState::InTransaction)
{
return Err(KrafkaError::invalid_state(format!(
"cannot begin transaction in state {:?}",
actual
)));
}
debug!("Transaction started");
Ok(())
}
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_record(&self, record: ProducerRecord) -> Result<RecordMetadata> {
let operation_guard = self.in_flight_barrier.start("transactional producer")?;
let current = self.state();
if current != TransactionState::InTransaction {
return Err(KrafkaError::invalid_state(format!(
"cannot send in state {:?}",
current
)));
}
self.ensure_transaction_can_continue("send records")?;
let mut record = record;
if let Some(enc) = &self.value_encoder {
record.value = enc
.encode(
record.value.clone(),
&record.topic,
record.record_name.as_deref(),
false,
)
.await?;
}
if let Some(enc) = &self.key_encoder {
let key = record.key.clone().unwrap_or_default();
record.key = Some(
enc.encode(key, &record.topic, record.record_name.as_deref(), true)
.await?,
);
}
record.validate()?;
let _identity = self.checked_transactional_identity()?;
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)
}
};
if self.requires_explicit_partition_registration() {
self.add_partition_to_txn_if_needed(&topic, partition)
.await?;
}
let result = self
.accumulator
.append_routed_with_guard(topic, record, record_size, partition, operation_guard)
.await;
let result = self.classify_transaction_result(result);
match result {
Err(error) if Self::is_unknown_producer_id_error(&error) => {
Err(self.mark_unknown_producer_id_abort_required("transactional produce"))
}
other => other,
}
}
async fn add_partition_to_txn_if_needed(
&self,
topic: &Arc<str>,
partition: PartitionId,
) -> Result<()> {
loop {
let mut txn_partitions = self.txn_partitions.write().await;
match txn_partitions.begin_add(topic.as_ref(), partition) {
BeginAddResult::AlreadyAdded => break,
BeginAddResult::Fatal(err) => {
return Err((*err).clone());
}
BeginAddResult::Wait(notify) => {
let notified = notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
drop(txn_partitions);
notified.await;
}
BeginAddResult::NeedAdd(notify) => {
drop(txn_partitions);
let guard = PendingAddGuard {
txn_partitions: self.txn_partitions.clone(),
topic: topic.clone(),
partition,
notify,
defused: false,
};
match self.add_partition_to_txn(topic.as_ref(), partition).await {
Ok(()) => {
guard.confirm(topic.as_ref(), partition).await;
}
Err(e) if e.is_retriable() => {
guard.cancel(topic.as_ref(), partition).await;
return Err(e);
}
Err(e) => {
guard
.fail(topic.as_ref(), partition, Arc::new(e.clone()))
.await;
return Err(e);
}
}
break;
}
}
}
Ok(())
}
async fn add_partition_to_txn(&self, topic: &str, partition: PartitionId) -> Result<()> {
let result = self.retry_with_coordinator("AddPartitionsToTxn", |attempt| async move {
let (_coordinator_id, conn) = self.coordinator_connection(attempt).await?;
let (producer_id, producer_epoch) = self.checked_transactional_identity()?;
let apt_version = conn
.negotiate_api_version(
ApiKey::AddPartitionsToTxn,
versions::ADD_PARTITIONS_TO_TXN_MAX,
versions::ADD_PARTITIONS_TO_TXN_MIN,
)
.ok_or_else(|| {
KrafkaError::protocol_kind(ProtocolErrorKind::UnknownApiVersion, "no mutually supported AddPartitionsToTxn API version")
})?;
let request = AddPartitionsToTxnRequest::new(
&self.config.transactional_id,
producer_id,
producer_epoch,
)
.add_partition(topic, partition);
let response_bytes = conn
.send_request(ApiKey::AddPartitionsToTxn, apt_version, |buf| {
request.encode_versioned(apt_version, buf)
})
.await?;
let mut buf = response_bytes;
let response = AddPartitionsToTxnResponse::decode_versioned(apt_version, &mut buf)?;
if !response.is_ok() {
for topic_result in &response.results {
for partition_result in &topic_result.partitions {
if !partition_result.error_code.is_ok() {
return Err(KrafkaError::broker(
partition_result.error_code,
format!("failed to add {}-{} to transaction", topic, partition),
));
}
}
}
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!(
"failed to add {}-{} to transaction: response indicated error but no per-partition error found",
topic, partition
),
));
}
debug!("Added partition {}-{} to transaction", topic, partition);
Ok(())
})
.await;
match self.classify_transaction_result(result) {
Err(error) if Self::is_unknown_producer_id_error(&error) => {
Err(self.mark_unknown_producer_id_abort_required("AddPartitionsToTxn"))
}
other => other,
}
}
pub async fn send_offsets_to_transaction(
&self,
offsets: &[TopicPartitionOffset],
group_metadata: &ConsumerGroupMetadata,
) -> Result<()> {
let current = self.state();
if current != TransactionState::InTransaction {
return Err(KrafkaError::invalid_state(format!(
"cannot send offsets in state {:?}",
current
)));
}
self.ensure_transaction_can_continue("send offsets")?;
if !group_metadata.is_fenceable() {
self.abort_required.store(true, Ordering::SeqCst);
return Err(KrafkaError::broker(
ErrorCode::TransactionAbortable,
format!(
"consumer group metadata for '{}' carries no valid generation \
(generation_id={}, member_id={:?}); the offset commit could not be \
fenced against a zombie consumer. abort_transaction() is required.",
group_metadata.group_id(),
group_metadata.generation_id(),
group_metadata.member_id(),
),
));
}
let group_id = group_metadata.group_id();
let (producer_id, producer_epoch) = self.checked_transactional_identity()?;
if self.requires_explicit_partition_registration() {
self.add_offsets_to_txn(producer_id, producer_epoch, group_id)
.await?;
}
let commit_request = build_txn_offset_commit_request(
&self.config.transactional_id,
group_metadata,
producer_id,
producer_epoch,
offsets,
);
let toc_min_version = if self.transaction_version().is_v2() {
TV2_MIN_TXN_OFFSET_COMMIT_VERSION
} else {
versions::TXN_OFFSET_COMMIT_MIN
};
let max_retries = self.retry_policy.max_retries;
for attempt in 0..=max_retries {
if attempt > 0 {
tokio::time::sleep(self.retry_policy.calculate_backoff(attempt)).await;
}
let result: Result<()> = async {
let (group_node_id, group_host, group_port) =
self.find_group_coordinator(group_id, attempt).await?;
let group_addr = format!("{group_host}:{group_port}");
let group_conn = self
.pool
.get_connection_by_id(group_node_id, &group_addr)
.await?;
let toc_version = group_conn
.negotiate_api_version(
ApiKey::TxnOffsetCommit,
versions::TXN_OFFSET_COMMIT_MAX,
toc_min_version,
)
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
format!(
"no mutually supported TxnOffsetCommit API version (need v{toc_min_version}+)"
),
)
})?;
let response_bytes = group_conn
.send_request(ApiKey::TxnOffsetCommit, toc_version, |buf| {
commit_request.encode_versioned(toc_version, buf)
})
.await?;
let mut buf = response_bytes;
let commit_response =
TxnOffsetCommitResponse::decode_versioned(toc_version, &mut buf)?;
if !commit_response.is_ok() {
for topic_result in &commit_response.topics {
for part_result in &topic_result.partitions {
if !part_result.error_code.is_ok() {
return Err(KrafkaError::broker(
part_result.error_code,
format!(
"failed to commit offset for {}-{} in transaction",
topic_result.name, part_result.partition
),
));
}
}
}
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
"failed to commit offsets in transaction",
));
}
Ok(())
}
.await;
let result = self.classify_transaction_result(result);
if let Err(error) = &result
&& Self::is_unknown_producer_id_error(error)
{
return Err(self.mark_unknown_producer_id_abort_required("TxnOffsetCommit"));
}
if self.state() == TransactionState::FatalError {
return result;
}
match &result {
Ok(()) => {
debug!("Added offsets to transaction for group {}", group_id);
return Ok(());
}
Err(e) if Self::needs_coordinator_refresh(e) && attempt < max_retries => {
warn!(
attempt,
error = %e,
"TxnOffsetCommit group coordinator error, re-discovering and retrying"
);
}
Err(e) if e.is_retriable() && attempt < max_retries => {
warn!(
attempt,
error = %e,
"TxnOffsetCommit retriable error, retrying"
);
}
Err(_) => return result,
}
}
Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!("TxnOffsetCommit retry loop exhausted after {max_retries} retries"),
))
}
async fn add_offsets_to_txn(
&self,
producer_id: i64,
producer_epoch: i16,
group_id: &str,
) -> Result<()> {
let add_offsets_result = self
.retry_with_coordinator("AddOffsetsToTxn", |attempt| async move {
let (_coordinator_id, conn) = self.coordinator_connection(attempt).await?;
let add_request = AddOffsetsToTxnRequest::new(
&self.config.transactional_id,
producer_id,
producer_epoch,
group_id,
);
let aot_version = conn
.negotiate_api_version(
ApiKey::AddOffsetsToTxn,
versions::ADD_OFFSETS_TO_TXN_MAX,
versions::ADD_OFFSETS_TO_TXN_MIN,
)
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
"no mutually supported AddOffsetsToTxn API version",
)
})?;
let response_bytes = conn
.send_request(ApiKey::AddOffsetsToTxn, aot_version, |buf| {
add_request.encode_versioned(aot_version, buf)
})
.await?;
let mut buf = response_bytes;
let add_response =
AddOffsetsToTxnResponse::decode_versioned(aot_version, &mut buf)?;
if !add_response.is_ok() {
return Err(KrafkaError::broker(
add_response.error_code,
"failed to add offsets to transaction",
));
}
Ok(())
})
.await;
match self.classify_transaction_result(add_offsets_result) {
Err(error) if Self::is_unknown_producer_id_error(&error) => {
Err(self.mark_unknown_producer_id_abort_required("AddOffsetsToTxn"))
}
other => other,
}
}
async fn find_group_coordinator(
&self,
group_id: &str,
attempt: u32,
) -> Result<(i32, String, i32)> {
let brokers = self.metadata.brokers();
if brokers.is_empty() {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
"no brokers available",
));
}
let broker = &brokers[attempt as usize % brokers.len()];
let conn = self
.pool
.get_connection_by_id(broker.id(), broker.address())
.await?;
let request = FindCoordinatorRequest::for_group(group_id);
let fc_version = conn
.negotiate_api_version(
ApiKey::FindCoordinator,
versions::FIND_COORDINATOR_MAX,
versions::FIND_COORDINATOR_MIN,
)
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
"no mutually supported FindCoordinator API version",
)
})?;
let response_bytes = conn
.send_request(ApiKey::FindCoordinator, fc_version, |buf| {
request.encode_versioned(fc_version, buf)
})
.await?;
let mut buf = response_bytes;
let response = FindCoordinatorResponse::decode_versioned(fc_version, &mut buf)?;
if !response.error_code.is_ok() {
return Err(KrafkaError::broker(
response.error_code,
"failed to find group coordinator",
));
}
Ok((response.node_id, response.host, response.port))
}
pub async fn commit_transaction(&self) -> Result<()> {
self.ensure_transaction_can_continue("commit transaction")?;
self.accumulator.flush().await?;
if let Err(actual) = self
.try_transition(
TransactionState::InTransaction,
TransactionState::Committing,
)
.or_else(|_| {
self.try_transition(
TransactionState::CommitIndeterminate,
TransactionState::Committing,
)
})
{
return Err(KrafkaError::invalid_state(format!(
"cannot commit in state {:?}",
actual
)));
}
let result = match self.end_transaction(true).await {
Err(error) if Self::is_unknown_producer_id_error(&error) => {
Err(self.mark_unknown_producer_id_abort_required("commit_transaction"))
}
other => other,
};
match &result {
Ok(()) => {
self.set_state(TransactionState::Ready);
self.txn_partitions.write().await.clear();
info!("Transaction committed");
}
Err(e) if Self::is_abortable_transaction_error(e, self.transaction_version()) => {
match self.try_transition(
TransactionState::Committing,
TransactionState::InTransaction,
) {
Ok(()) => {
warn!("Transaction commit failed (abort required): {}", e);
}
Err(actual) => {
warn!(
"Transaction commit failed (abort required): {}; \
state is now {:?} (concurrent abort may be in progress)",
e, actual
);
}
}
}
Err(e) => {
if e.is_retriable() {
let outcome_unknown =
matches!(e, KrafkaError::Timeout { .. } | KrafkaError::Network(_));
let revert_to = if outcome_unknown {
TransactionState::CommitIndeterminate
} else {
TransactionState::InTransaction
};
match self.try_transition(TransactionState::Committing, revert_to) {
Ok(()) => {
if outcome_unknown {
warn!(
"Transaction commit outcome unknown ({e}); the coordinator \
may already have committed it. Retry commit_transaction() — \
aborting from here could tear the transaction (KAFKA-17754)."
);
} else {
warn!("Transaction commit failed (retriable): {}", e);
}
}
Err(actual) => {
warn!(
"Transaction commit failed (retriable): {}; \
state is now {:?} (concurrent abort may be in progress)",
e, actual
);
}
}
} else {
self.set_state(TransactionState::FatalError);
warn!("Transaction commit failed (fatal): {}", e);
}
}
}
result
}
pub async fn abort_transaction(&self) -> Result<()> {
if self.state() == TransactionState::CommitIndeterminate {
return Err(KrafkaError::invalid_state(
"cannot abort: a previous commit_transaction() timed out or lost its \
connection, so the coordinator may already have committed this \
transaction. Aborting now could be applied to a later transaction and \
tear it (KAFKA-17754). Retry commit_transaction() — EndTxn is idempotent \
for the same producer id and epoch — or drop this producer and let the \
coordinator resolve the transaction via transaction.timeout.ms."
.to_string(),
));
}
let transition = self
.try_transition(TransactionState::InTransaction, TransactionState::Aborting)
.or_else(|_| {
self.try_transition(TransactionState::Committing, TransactionState::Aborting)
});
if let Err(actual) = transition {
return Err(KrafkaError::invalid_state(format!(
"cannot abort in state {:?}",
actual
)));
}
if let Err(err) = self.accumulator.flush().await {
debug!(error = %err, "Accumulator flush during abort_transaction failed");
}
let needs_reinitialize = self.abort_required.swap(false, Ordering::SeqCst);
let result = if needs_reinitialize {
match self.end_transaction(false).await {
Ok(()) => self.do_init_transactions().await,
Err(error) if Self::is_unknown_producer_id_error(&error) => {
debug!(
"Abort observed UnknownProducerId after transactional error; reinitializing producer identity"
);
self.do_init_transactions().await
}
Err(error) => Err(error),
}
} else {
self.end_transaction(false).await
};
match &result {
Ok(()) => {
self.set_state(TransactionState::Ready);
self.txn_partitions.write().await.clear();
info!("Transaction aborted");
}
Err(e) if e.is_retriable() => {
match self
.try_transition(TransactionState::Aborting, TransactionState::InTransaction)
{
Ok(()) => {
if needs_reinitialize {
self.abort_required.store(true, Ordering::SeqCst);
}
warn!(
"Transaction abort failed (retriable), retry abort_transaction(): {e}"
);
}
Err(actual) => {
warn!(
"Transaction abort failed (retriable): {e}; state is now {actual:?} \
(concurrent operation may be in progress)"
);
}
}
}
Err(e) => {
self.set_state(TransactionState::FatalError);
warn!("Transaction abort failed (fatal): {e}; producer must be recreated");
}
}
result
}
async fn end_transaction(&self, commit: bool) -> Result<()> {
let is_v2 = self.transaction_version().is_v2();
let et_min_version = if is_v2 {
TV2_MIN_END_TXN_VERSION
} else {
versions::END_TXN_MIN
};
let result = self
.retry_with_coordinator("EndTxn", |attempt| async move {
let (_coordinator_id, conn) = self.coordinator_connection(attempt).await?;
let (producer_id, producer_epoch) = self.checked_transactional_identity()?;
let et_version = conn
.negotiate_api_version(ApiKey::EndTxn, versions::END_TXN_MAX, et_min_version)
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
format!(
"no mutually supported EndTxn API version (need v{et_min_version}+)"
),
)
})?;
let request = if commit {
EndTxnRequest::commit(
&self.config.transactional_id,
producer_id,
producer_epoch,
)
} else {
EndTxnRequest::abort(&self.config.transactional_id, producer_id, producer_epoch)
};
let response_bytes = conn
.send_request(ApiKey::EndTxn, et_version, |buf| {
request.encode_versioned(et_version, buf)
})
.await?;
let mut buf = response_bytes;
let response = EndTxnResponse::decode_versioned(et_version, &mut buf)?;
if !response.is_ok() {
return Err(KrafkaError::broker(
response.error_code,
if commit {
"failed to commit transaction"
} else {
"failed to abort transaction"
},
));
}
match (response.producer_id, response.producer_epoch) {
(Some(pid), Some(epoch)) if pid >= 0 && epoch >= 0 => {
debug!(
pid,
epoch,
transaction_version = %self.transaction_version(),
"Adopting broker-bumped producer identity from EndTxn response"
);
self.identity.bump_epoch(pid, epoch);
}
_ if is_v2 => {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
"EndTxn response omitted the bumped producer id/epoch that \
transaction version 2 requires",
));
}
_ => {}
}
Ok(())
})
.await;
self.classify_transaction_result(result)
}
#[inline]
pub fn transactional_id(&self) -> &str {
&self.config.transactional_id
}
#[inline]
pub fn producer_id(&self) -> i64 {
self.identity.producer_id()
}
#[inline]
pub fn producer_epoch(&self) -> i16 {
self.identity.producer_epoch()
}
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(err) = self.accumulator.shutdown().await {
warn!(error = %err, "Accumulator shutdown error during transactional close");
}
self.in_flight_barrier.wait_for(target).await;
let current = self.state();
if current == TransactionState::InTransaction {
warn!("Closing transactional producer with active transaction — aborting");
self.abort_transaction().await?;
} else if current == TransactionState::CommitIndeterminate {
warn!(
"Closing transactional producer after a commit whose outcome is unknown; \
leaving the transaction for the coordinator to resolve via \
transaction.timeout.ms rather than aborting a possibly-committed \
transaction (KAFKA-17754)"
);
}
Ok::<(), KrafkaError>(())
};
let close_result = if let Some(timeout) = timeout {
tokio::time::timeout(timeout, graceful_close)
.await
.map_err(|_| KrafkaError::timeout("transactional producer close"))?
} else {
graceful_close.await
};
self.set_state(TransactionState::FatalError);
self.pool.close_all().await;
info!(
"TransactionalProducer closed: txn.id()={}",
self.config.transactional_id
);
close_result
}
pub async fn refresh_tls(&self) -> Result<()> {
self.pool.refresh_tls().await
}
pub fn update_seed_brokers(&self, servers: Vec<String>) -> Result<()> {
self.metadata.update_seed_brokers(servers)
}
pub async fn rebootstrap(&self) {
self.metadata.rebootstrap().await;
}
#[inline]
pub fn connection_metrics(&self) -> Arc<crate::metrics::ConnectionMetrics> {
self.pool.metrics()
}
#[inline]
pub fn metrics(&self) -> crate::producer::ProducerMetricsSnapshot {
crate::producer::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<ProducerMetrics> {
self.metrics.clone()
}
#[inline]
pub fn is_closed(&self) -> bool {
self.in_flight_barrier.is_closing()
}
}
fn build_txn_offset_commit_request(
transactional_id: &str,
group_metadata: &ConsumerGroupMetadata,
producer_id: i64,
producer_epoch: i16,
offsets: &[TopicPartitionOffset],
) -> TxnOffsetCommitRequest {
let mut request = TxnOffsetCommitRequest::new(
transactional_id,
group_metadata.group_id(),
producer_id,
producer_epoch,
);
request.generation_id = group_metadata.generation_id();
request.member_id = group_metadata.member_id().to_string();
request.group_instance_id = group_metadata.group_instance_id().map(str::to_string);
for tpo in offsets {
request = request.add_offset(&tpo.topic, tpo.partition, tpo.next_offset, None);
}
request
}
fn is_fatal_transaction_error(error_code: ErrorCode, version: TransactionVersion) -> bool {
if error_code == ErrorCode::InvalidProducerIdMapping {
return version.is_v2();
}
matches!(
error_code,
ErrorCode::InvalidProducerEpoch
| ErrorCode::ProducerFenced
| ErrorCode::TransactionalIdAuthorizationFailed
| ErrorCode::InvalidTxnState
| ErrorCode::TransactionCoordinatorFenced
)
}
#[derive(Default)]
#[must_use = "builders do nothing until .build() is called"]
pub struct TransactionalProducerBuilder {
config: TransactionalProducerConfig,
retry_policy: RetryPolicy,
partitioner: Option<Arc<dyn Partitioner>>,
key_encoder: Option<Arc<dyn SchemaEncoder>>,
value_encoder: Option<Arc<dyn SchemaEncoder>>,
}
impl TransactionalProducerBuilder {
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 transactional_id(mut self, txn_id: impl Into<String>) -> Self {
self.config.transactional_id = txn_id.into();
self
}
pub fn transaction_timeout(mut self, timeout: Duration) -> Self {
self.config.transaction_timeout_ms = crate::util::duration_to_millis_i32(timeout);
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 max_request_size(mut self, bytes: usize) -> Self {
self.config.max_request_size = bytes;
self
}
pub fn batch_size(mut self, bytes: usize) -> Self {
self.config.batch_size = bytes;
self
}
pub fn linger(mut self, linger: Duration) -> Self {
self.config.linger = linger;
self
}
pub fn buffer_memory(mut self, bytes: usize) -> Self {
self.config.buffer_memory = bytes;
self
}
pub fn max_block(mut self, max_block: Duration) -> Self {
self.config.max_block = max_block;
self
}
pub fn max_in_flight(mut self, max: usize) -> Self {
self.config.max_in_flight = max;
self
}
pub fn compression(mut self, compression: Compression) -> Self {
self.config.compression = compression;
self
}
pub fn partitioner(mut self, partitioner: impl Partitioner + 'static) -> Self {
self.partitioner = Some(Arc::new(partitioner));
self
}
pub fn auth(mut self, auth: AuthConfig) -> Self {
self.config.auth = Some(auth);
self
}
#[cfg(feature = "socks5")]
pub fn proxy(mut self, proxy: crate::network::ProxyConfig) -> Self {
self.config.proxy = Some(proxy);
self
}
pub fn sasl_oauthbearer(mut self, token: impl Into<String>) -> Self {
self.config.auth = Some(crate::auth::AuthConfig::sasl_oauthbearer(token));
self
}
pub fn metadata_recovery_strategy(
mut self,
strategy: crate::metadata::MetadataRecoveryStrategy,
) -> Self {
self.config.metadata_recovery_strategy = strategy;
self
}
pub fn metadata_max_age(mut self, age: Duration) -> Self {
self.config.metadata_max_age = age;
self
}
pub fn metadata_recovery_rebootstrap_trigger(mut self, duration: Duration) -> Self {
self.config.metadata_recovery_rebootstrap_trigger = duration;
self
}
pub fn transport(mut self, transport: crate::network::TransportConfig) -> Self {
self.config.transport = transport;
self
}
pub fn sasl_plain(mut self, username: &str, password: &str) -> crate::Result<Self> {
self.config.auth = Some(AuthConfig::sasl_plain(username, password)?);
Ok(self)
}
pub fn sasl_scram_sha256(mut self, username: &str, password: &str) -> Self {
self.config.auth = Some(AuthConfig::sasl_scram_sha256(username, password));
self
}
pub fn sasl_scram_sha512(mut self, username: &str, password: &str) -> Self {
self.config.auth = Some(AuthConfig::sasl_scram_sha512(username, password));
self
}
pub fn key_encoder(mut self, encoder: Arc<dyn SchemaEncoder>) -> Self {
self.key_encoder = Some(encoder);
self
}
pub fn value_encoder(mut self, encoder: Arc<dyn SchemaEncoder>) -> Self {
self.value_encoder = Some(encoder);
self
}
pub fn retries(mut self, retries: u32) -> Self {
self.retry_policy = self.retry_policy.with_max_retries(retries);
self
}
pub fn retry_backoff(mut self, backoff: Duration) -> Self {
self.retry_policy = self.retry_policy.with_initial_backoff(backoff);
self
}
pub async fn build(self) -> Result<TransactionalProducer> {
if self.config.bootstrap_servers.is_empty() {
return Err(KrafkaError::config("bootstrap.servers is required"));
}
if self.config.transactional_id.is_empty() {
return Err(KrafkaError::config("transactional_id is required"));
}
const MAX_KAFKA_STRING_LEN: usize = i16::MAX as usize;
if self.config.transactional_id.len() > MAX_KAFKA_STRING_LEN {
return Err(KrafkaError::config(format!(
"transactional_id is {} bytes, exceeding the Kafka wire limit of {MAX_KAFKA_STRING_LEN}",
self.config.transactional_id.len()
)));
}
if self.config.client_id.len() > MAX_KAFKA_STRING_LEN {
return Err(KrafkaError::config(format!(
"client_id is {} bytes, exceeding the Kafka wire limit of {MAX_KAFKA_STRING_LEN}",
self.config.client_id.len()
)));
}
if self.config.transaction_timeout_ms <= 0 {
return Err(KrafkaError::config("transaction_timeout must be > 0"));
}
if self.config.max_request_size == 0 {
return Err(KrafkaError::config("max_request_size must be >= 1"));
}
let mut pool_config_builder = self.config.transport.apply(
ConnectionConfig::builder()
.client_id(&self.config.client_id)
.request_timeout(self.config.request_timeout)
.connect_timeout(self.config.connect_timeout),
);
if let Some(ref auth) = self.config.auth {
pool_config_builder = pool_config_builder.auth(auth.clone());
}
#[cfg(feature = "socks5")]
if let Some(ref proxy) = self.config.proxy {
pool_config_builder = pool_config_builder.proxy(proxy.clone());
}
let mut pool_config = pool_config_builder.build()?;
pool_config.init_tls().await?;
let pool = self.config.transport.build_pool(pool_config);
let bootstrap_servers =
crate::util::parse_bootstrap_servers(&self.config.bootstrap_servers)?;
let metadata = Arc::new(
ClusterMetadata::new(
bootstrap_servers,
pool.clone(),
self.config.metadata_max_age,
)
.with_recovery_strategy(self.config.metadata_recovery_strategy)
.with_rebootstrap_trigger(self.config.metadata_recovery_rebootstrap_trigger),
);
metadata.refresh().await?;
info!(
"TransactionalProducer created with transactional.id()={}",
self.config.transactional_id
);
let partitioner: Arc<dyn Partitioner> = self
.partitioner
.unwrap_or_else(|| Arc::new(UniformStickyPartitioner::new()));
let identity = Arc::new(ProducerIdentity::new());
let metrics = Arc::new(ProducerMetrics::default());
let in_flight_barrier = Arc::new(InFlightBarrier::new());
let accumulator = RecordAccumulator::spawn(
AccumulatorConfig {
batch_size: self.config.batch_size,
linger: self.config.linger,
compression: self.config.compression,
topic_compression: ahash::AHashMap::new(),
acks: Acks::All.to_i16(),
client_id: self.config.client_id.clone(),
request_timeout: self.config.request_timeout,
max_request_size: self.config.max_request_size,
buffer_memory: self.config.buffer_memory,
max_block_ms: self.config.max_block,
in_flight_semaphore: Arc::new(tokio::sync::Semaphore::new(
self.config.max_in_flight.max(1),
)),
interceptor: Arc::new(crate::interceptor::NoOpProducerInterceptor),
identity: Some(identity.clone()),
partitioner: partitioner.clone(),
state_store: None,
transactional_id: Some(self.config.transactional_id.clone()),
},
metadata.clone(),
self.retry_policy.clone(),
metrics.clone(),
in_flight_barrier.clone(),
);
Ok(TransactionalProducer {
config: self.config,
metadata,
pool,
partitioner,
state: AtomicU8::new(TransactionState::Uninitialized as u8),
transaction_version: AtomicU8::new(TransactionVersion::V1 as u8),
abort_required: AtomicBool::new(false),
coordinator_id: RwLock::new(None),
txn_partitions: Arc::new(RwLock::new(TransactionPartitions::default())),
identity,
accumulator,
metrics,
retry_policy: self.retry_policy,
in_flight_barrier,
key_encoder: self.key_encoder,
value_encoder: self.value_encoder,
})
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
use crate::metadata::ClusterMetadata;
use crate::network::ConnectionPool;
fn test_accumulator() -> RecordAccumulatorHandle {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
let metadata = Arc::new(ClusterMetadata::new(
vec!["localhost:9092".to_string()],
pool,
Duration::from_secs(300),
));
RecordAccumulator::spawn(
AccumulatorConfig::default(),
metadata,
RetryPolicy::default(),
Arc::new(ProducerMetrics::default()),
Arc::new(InFlightBarrier::new()),
)
}
#[test]
fn test_transaction_state() {
assert_eq!(TransactionState::from(0), TransactionState::Uninitialized);
assert_eq!(TransactionState::from(1), TransactionState::Ready);
assert_eq!(TransactionState::from(2), TransactionState::InTransaction);
assert_eq!(TransactionState::from(3), TransactionState::Committing);
assert_eq!(TransactionState::from(4), TransactionState::Aborting);
assert_eq!(TransactionState::from(5), TransactionState::FatalError);
assert_eq!(TransactionState::from(99), TransactionState::FatalError);
}
#[test]
fn test_transactional_producer_config_default() {
let config = TransactionalProducerConfig::default();
assert_eq!(config.client_id, "krafka-txn-producer");
assert_eq!(config.transaction_timeout_ms, 60000);
assert_eq!(config.max_request_size, crate::protocol::MAX_MESSAGE_SIZE);
}
#[test]
fn test_transaction_partitions() {
let mut partitions = TransactionPartitions::default();
assert!(partitions.is_empty());
let result = partitions.begin_add("topic1", 0);
let notify = match result {
BeginAddResult::NeedAdd(n) => n,
_ => panic!("expected NeedAdd"),
};
assert!(!partitions.is_empty());
assert!(matches!(
partitions.begin_add("topic1", 0),
BeginAddResult::Wait(_)
));
partitions.confirm_add("topic1", 0, ¬ify);
assert!(matches!(
partitions.begin_add("topic1", 0),
BeginAddResult::AlreadyAdded
));
assert!(matches!(
partitions.begin_add("topic1", 1),
BeginAddResult::NeedAdd(_)
));
partitions.clear();
assert!(partitions.is_empty());
}
#[test]
fn test_is_fatal_transaction_error() {
for version in [TransactionVersion::V1, TransactionVersion::V2] {
assert!(is_fatal_transaction_error(
ErrorCode::InvalidProducerEpoch,
version
));
assert!(is_fatal_transaction_error(
ErrorCode::ProducerFenced,
version
));
assert!(is_fatal_transaction_error(
ErrorCode::TransactionCoordinatorFenced,
version
));
assert!(is_fatal_transaction_error(
ErrorCode::TransactionalIdAuthorizationFailed,
version
));
assert!(is_fatal_transaction_error(
ErrorCode::InvalidTxnState,
version
));
assert!(!is_fatal_transaction_error(ErrorCode::None, version));
assert!(!is_fatal_transaction_error(
ErrorCode::UnknownServerError,
version
));
}
}
#[test]
fn test_needs_coordinator_refresh() {
assert!(TransactionalProducer::needs_coordinator_refresh(
&KrafkaError::broker(ErrorCode::NotCoordinator, "test")
));
assert!(TransactionalProducer::needs_coordinator_refresh(
&KrafkaError::broker(ErrorCode::CoordinatorNotAvailable, "test")
));
assert!(TransactionalProducer::needs_coordinator_refresh(
&KrafkaError::broker(ErrorCode::CoordinatorLoadInProgress, "test")
));
assert!(TransactionalProducer::needs_coordinator_refresh(
&KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionRefused,
"refused"
))
));
assert!(TransactionalProducer::needs_coordinator_refresh(
&KrafkaError::timeout("test operation")
));
assert!(!TransactionalProducer::needs_coordinator_refresh(
&KrafkaError::broker(ErrorCode::InvalidProducerEpoch, "test")
));
assert!(!TransactionalProducer::needs_coordinator_refresh(
&KrafkaError::broker(ErrorCode::TransactionCoordinatorFenced, "test")
));
assert!(!TransactionalProducer::needs_coordinator_refresh(
&KrafkaError::protocol_kind(ProtocolErrorKind::Other, "test")
));
assert!(!TransactionalProducer::needs_coordinator_refresh(
&KrafkaError::invalid_state("test")
));
}
#[tokio::test]
async fn test_builder_missing_bootstrap() {
let result = TransactionalProducer::builder()
.transactional_id("my-txn")
.build()
.await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_send_record_requires_initialized_transactional_identity() {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
let metadata = Arc::new(ClusterMetadata::new(
vec!["localhost:9092".to_string()],
pool.clone(),
Duration::from_secs(300),
));
let producer = TransactionalProducer {
config: TransactionalProducerConfig {
bootstrap_servers: "localhost:9092".to_string(),
transactional_id: "txn-test".to_string(),
..TransactionalProducerConfig::default()
},
metadata,
pool,
partitioner: Arc::new(UniformStickyPartitioner::new()),
state: AtomicU8::new(TransactionState::InTransaction as u8),
transaction_version: AtomicU8::new(TransactionVersion::V1 as u8),
abort_required: AtomicBool::new(false),
coordinator_id: RwLock::new(None),
txn_partitions: Arc::new(RwLock::new(TransactionPartitions::default())),
identity: Arc::new(ProducerIdentity::new()),
accumulator: test_accumulator(),
metrics: Arc::new(ProducerMetrics::default()),
retry_policy: RetryPolicy::default(),
in_flight_barrier: Arc::new(InFlightBarrier::new()),
key_encoder: None,
value_encoder: None,
};
let record = ProducerRecord::new("topic", Bytes::from_static(b"value")).with_partition(0);
let err = producer.send_record(record).await.unwrap_err();
assert!(
err.to_string()
.contains("transactional producer identity not initialized"),
"expected invalid identity guard, got: {err}"
);
}
#[tokio::test]
async fn test_builder_missing_txn_id() {
let result = TransactionalProducer::builder()
.bootstrap_servers("localhost:9092")
.build()
.await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_mark_unknown_producer_id_requires_abort() {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
let metadata = Arc::new(ClusterMetadata::new(
vec!["localhost:9092".to_string()],
pool.clone(),
Duration::from_secs(300),
));
let producer = TransactionalProducer {
config: TransactionalProducerConfig {
bootstrap_servers: "localhost:9092".to_string(),
transactional_id: "txn-test".to_string(),
..TransactionalProducerConfig::default()
},
metadata,
pool,
partitioner: Arc::new(UniformStickyPartitioner::new()),
state: AtomicU8::new(TransactionState::InTransaction as u8),
transaction_version: AtomicU8::new(TransactionVersion::V1 as u8),
abort_required: AtomicBool::new(false),
coordinator_id: RwLock::new(None),
txn_partitions: Arc::new(RwLock::new(TransactionPartitions::default())),
identity: Arc::new(ProducerIdentity::new()),
accumulator: test_accumulator(),
metrics: Arc::new(ProducerMetrics::default()),
retry_policy: RetryPolicy::default(),
in_flight_barrier: Arc::new(InFlightBarrier::new()),
key_encoder: None,
value_encoder: None,
};
let error = producer.mark_unknown_producer_id_abort_required("transactional produce");
assert!(matches!(
error,
KrafkaError::Broker {
code: ErrorCode::TransactionAbortable,
..
}
));
assert!(producer.abort_required());
let gate_error = producer
.ensure_transaction_can_continue("commit transaction")
.unwrap_err();
assert!(matches!(
gate_error,
KrafkaError::Broker {
code: ErrorCode::TransactionAbortable,
..
}
));
}
#[tokio::test]
async fn test_commit_transaction_rejects_abort_required() {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
let metadata = Arc::new(ClusterMetadata::new(
vec!["localhost:9092".to_string()],
pool.clone(),
Duration::from_secs(300),
));
let producer = TransactionalProducer {
config: TransactionalProducerConfig {
bootstrap_servers: "localhost:9092".to_string(),
transactional_id: "txn-test".to_string(),
..TransactionalProducerConfig::default()
},
metadata,
pool,
partitioner: Arc::new(UniformStickyPartitioner::new()),
state: AtomicU8::new(TransactionState::InTransaction as u8),
transaction_version: AtomicU8::new(TransactionVersion::V1 as u8),
abort_required: AtomicBool::new(true),
coordinator_id: RwLock::new(None),
txn_partitions: Arc::new(RwLock::new(TransactionPartitions::default())),
identity: Arc::new(ProducerIdentity::new()),
accumulator: test_accumulator(),
metrics: Arc::new(ProducerMetrics::default()),
retry_policy: RetryPolicy::default(),
in_flight_barrier: Arc::new(InFlightBarrier::new()),
key_encoder: None,
value_encoder: None,
};
let error = producer.commit_transaction().await.unwrap_err();
assert!(matches!(
error,
KrafkaError::Broker {
code: ErrorCode::TransactionAbortable,
..
}
));
assert_eq!(producer.state(), TransactionState::InTransaction);
}
#[test]
fn test_try_transition_success() {
let state = AtomicU8::new(TransactionState::Ready as u8);
let result = state.compare_exchange(
TransactionState::Ready as u8,
TransactionState::InTransaction as u8,
Ordering::SeqCst,
Ordering::SeqCst,
);
assert!(result.is_ok());
assert_eq!(
TransactionState::from(state.load(Ordering::SeqCst)),
TransactionState::InTransaction
);
}
#[test]
fn test_try_transition_failure() {
let state = AtomicU8::new(TransactionState::Uninitialized as u8);
let result = state.compare_exchange(
TransactionState::Ready as u8,
TransactionState::InTransaction as u8,
Ordering::SeqCst,
Ordering::SeqCst,
);
assert!(result.is_err());
assert_eq!(
TransactionState::from(state.load(Ordering::SeqCst)),
TransactionState::Uninitialized
);
}
#[test]
fn test_txn_builder_no_auth_by_default() {
let builder = TransactionalProducer::builder()
.bootstrap_servers("broker:9092")
.transactional_id("txn-1");
assert!(builder.config.auth.is_none());
}
#[test]
fn test_txn_builder_sets_max_request_size() {
let builder = TransactionalProducer::builder()
.bootstrap_servers("broker:9092")
.transactional_id("txn-1")
.max_request_size(65_536);
assert_eq!(builder.config.max_request_size, 65_536);
}
#[test]
fn test_txn_builder_sasl_plain() {
let builder = TransactionalProducer::builder()
.bootstrap_servers("broker:9093")
.transactional_id("txn-1")
.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_txn_builder_sasl_scram_sha256() {
let builder = TransactionalProducer::builder()
.bootstrap_servers("broker:9093")
.transactional_id("txn-1")
.sasl_scram_sha256("user", "pass");
let auth = builder.config.auth.as_ref().unwrap();
assert!(auth.requires_sasl());
assert!(auth.scram_credentials.is_some());
}
#[test]
fn test_txn_builder_sasl_scram_sha512() {
let builder = TransactionalProducer::builder()
.bootstrap_servers("broker:9093")
.transactional_id("txn-1")
.sasl_scram_sha512("user", "pass");
let auth = builder.config.auth.as_ref().unwrap();
assert!(auth.requires_sasl());
assert!(auth.scram_credentials.is_some());
}
#[test]
fn test_txn_builder_auth_config() {
use crate::auth::AuthConfig;
let auth = AuthConfig::sasl_scram_sha256("admin", "secret");
let builder = TransactionalProducer::builder()
.bootstrap_servers("broker:9093")
.transactional_id("txn-1")
.auth(auth);
let auth = builder.config.auth.as_ref().unwrap();
assert!(auth.requires_sasl());
assert!(auth.scram_credentials.is_some());
}
#[test]
fn test_txn_builder_initializes_producer_identity() {
let builder = TransactionalProducer::builder()
.bootstrap_servers("broker:9092")
.transactional_id("txn-test");
assert_eq!(builder.config.transactional_id, "txn-test");
}
#[test]
fn test_txn_builder_requires_transactional_id() {
let builder = TransactionalProducer::builder().bootstrap_servers("broker:9092");
assert!(builder.config.transactional_id.is_empty());
}
#[tokio::test]
async fn test_txn_builder_rejects_zero_timeout() {
let result = TransactionalProducer::builder()
.bootstrap_servers("localhost:9092")
.transactional_id("txn-1")
.transaction_timeout(Duration::ZERO)
.build()
.await;
match result {
Err(e) => assert!(e.to_string().contains("transaction_timeout")),
Ok(_) => panic!("expected error for transaction_timeout=0"),
}
}
#[tokio::test]
async fn test_txn_builder_rejects_zero_max_request_size() {
let result = TransactionalProducer::builder()
.bootstrap_servers("localhost:9092")
.transactional_id("txn-1")
.max_request_size(0)
.build()
.await;
match result {
Err(e) => assert!(e.to_string().contains("max_request_size")),
Ok(_) => panic!("expected error for max_request_size=0"),
}
}
#[tokio::test]
async fn test_txn_builder_rejects_negative_timeout() {
let result = TransactionalProducer::builder()
.bootstrap_servers("localhost:9092")
.transactional_id("txn-1")
.transaction_timeout(Duration::ZERO)
.build()
.await;
assert!(result.is_err());
}
#[test]
fn test_transaction_state_initializing_from_u8() {
assert_eq!(TransactionState::from(6), TransactionState::Initializing);
}
#[test]
fn test_transaction_state_initializing_value() {
assert_eq!(TransactionState::Initializing as u8, 6);
}
#[test]
fn test_transaction_state_initializing_round_trip() {
let state = TransactionState::Initializing;
let val = state as u8;
assert_eq!(TransactionState::from(val), TransactionState::Initializing);
}
#[test]
fn test_transaction_state_unknown_maps_to_fatal() {
assert_eq!(TransactionState::from(8), TransactionState::FatalError);
assert_eq!(TransactionState::from(255), TransactionState::FatalError);
}
#[test]
fn test_try_transition_uninitialized_to_initializing() {
let state = AtomicU8::new(TransactionState::Uninitialized as u8);
let result = state.compare_exchange(
TransactionState::Uninitialized as u8,
TransactionState::Initializing as u8,
Ordering::SeqCst,
Ordering::SeqCst,
);
assert!(result.is_ok());
assert_eq!(
TransactionState::from(state.load(Ordering::SeqCst)),
TransactionState::Initializing
);
}
#[test]
fn test_try_transition_initializing_blocks_second_init() {
let state = AtomicU8::new(TransactionState::Initializing as u8);
let result = state.compare_exchange(
TransactionState::Uninitialized as u8,
TransactionState::Initializing as u8,
Ordering::SeqCst,
Ordering::SeqCst,
);
assert!(result.is_err());
assert_eq!(
TransactionState::from(state.load(Ordering::SeqCst)),
TransactionState::Initializing
);
}
#[test]
fn test_commit_fatal_error_state_machine() {
let state = AtomicU8::new(TransactionState::Committing as u8);
let error = KrafkaError::broker(ErrorCode::InvalidProducerEpoch, "epoch fenced");
assert!(!error.is_retriable());
if error.is_retriable() {
state.store(TransactionState::InTransaction as u8, Ordering::SeqCst);
} else {
state.store(TransactionState::FatalError as u8, Ordering::SeqCst);
}
assert_eq!(
TransactionState::from(state.load(Ordering::SeqCst)),
TransactionState::FatalError
);
}
#[test]
fn test_commit_retriable_error_reverts_to_in_transaction() {
let state = AtomicU8::new(TransactionState::Committing as u8);
let error = KrafkaError::broker(ErrorCode::CoordinatorNotAvailable, "coordinator down");
assert!(error.is_retriable());
if error.is_retriable() {
state.store(TransactionState::InTransaction as u8, Ordering::SeqCst);
} else {
state.store(TransactionState::FatalError as u8, Ordering::SeqCst);
}
assert_eq!(
TransactionState::from(state.load(Ordering::SeqCst)),
TransactionState::InTransaction
);
}
#[test]
fn test_txn_close_sets_fatal_error_state() {
let state = AtomicU8::new(TransactionState::Ready as u8);
state.store(TransactionState::FatalError as u8, Ordering::SeqCst);
assert_eq!(
TransactionState::from(state.load(Ordering::SeqCst)),
TransactionState::FatalError
);
}
#[test]
fn test_out_of_order_sequence_is_retriable() {
let error = KrafkaError::broker(ErrorCode::OutOfOrderSequenceNumber, "sequence mismatch");
assert!(error.is_retriable());
}
#[test]
fn test_producer_record_with_timestamp() {
use crate::producer::ProducerRecord;
let record = ProducerRecord::new("topic", b"value".to_vec()).with_timestamp(1234567890);
assert_eq!(record.timestamp, Some(1234567890));
}
#[test]
fn test_transaction_partitions_state_machine() {
let mut tp = TransactionPartitions::default();
let result = tp.begin_add("topic", 0);
let notify = match result {
BeginAddResult::NeedAdd(n) => n,
_ => panic!("expected NeedAdd"),
};
let result2 = tp.begin_add("topic", 0);
assert!(matches!(result2, BeginAddResult::Wait(_)));
tp.confirm_add("topic", 0, ¬ify);
assert!(matches!(
tp.begin_add("topic", 0),
BeginAddResult::AlreadyAdded
));
let result3 = tp.begin_add("topic", 1);
let notify2 = match result3 {
BeginAddResult::NeedAdd(n) => n,
_ => panic!("expected NeedAdd"),
};
tp.cancel_add("topic", 1, ¬ify2);
assert!(matches!(
tp.begin_add("topic", 1),
BeginAddResult::NeedAdd(_)
));
tp.clear();
assert!(tp.is_empty());
}
#[test]
fn test_transaction_partitions_fail_add_propagates_as_fatal() {
let mut tp = TransactionPartitions::default();
let notify = match tp.begin_add("t", 0) {
BeginAddResult::NeedAdd(n) => n,
other => panic!("expected NeedAdd, got {other:?}"),
};
assert!(matches!(tp.begin_add("t", 0), BeginAddResult::Wait(_)));
let err = Arc::new(KrafkaError::invalid_state("fatal"));
tp.fail_add("t", 0, err.clone(), ¬ify);
assert!(
matches!(tp.begin_add("t", 0), BeginAddResult::Fatal(_)),
"expected Fatal after fail_add"
);
match tp.begin_add("t", 0) {
BeginAddResult::Fatal(stored) => {
assert_eq!(stored.to_string(), err.to_string());
}
other => panic!("expected Fatal, got {other:?}"),
}
}
#[test]
fn test_transactional_producer_is_send_sync() {
fn assert_send_sync<T: Send + Sync>() {}
assert_send_sync::<TransactionalProducer>();
}
#[test]
fn test_txn_offset_commit_carries_group_metadata() {
let metadata =
ConsumerGroupMetadata::new("my-group", 42, "member-7", Some("instance-3".to_string()));
let offsets = vec![
TopicPartitionOffset::new("orders", 0, 101),
TopicPartitionOffset::new("orders", 1, 55),
];
let request = build_txn_offset_commit_request("txn-1", &metadata, 12345, 4, &offsets);
assert_eq!(request.transactional_id, "txn-1");
assert_eq!(request.group_id, "my-group");
assert_eq!(request.producer_id, 12345);
assert_eq!(request.producer_epoch, 4);
assert_eq!(request.generation_id, 42, "KIP-447 generation must be sent");
assert_eq!(
request.member_id, "member-7",
"KIP-447 member_id must be sent"
);
assert_eq!(
request.group_instance_id.as_deref(),
Some("instance-3"),
"KIP-345 static instance id must be sent"
);
assert_eq!(request.topics.len(), 1);
assert_eq!(request.topics[0].name, "orders");
assert_eq!(request.topics[0].partitions.len(), 2);
assert_eq!(request.topics[0].partitions[0].committed_offset, 101);
assert_eq!(request.topics[0].partitions[1].committed_offset, 55);
}
#[test]
fn test_txn_offset_commit_without_static_membership() {
let metadata = ConsumerGroupMetadata::new("g", 3, "m", None);
let request = build_txn_offset_commit_request("txn", &metadata, 1, 0, &[]);
assert_eq!(request.generation_id, 3);
assert_eq!(request.member_id, "m");
assert!(request.group_instance_id.is_none());
}
#[test]
fn test_unfenceable_group_metadata_is_rejected() {
assert!(!ConsumerGroupMetadata::new("g", -1, "", None).is_fenceable());
assert!(!ConsumerGroupMetadata::new("g", 5, "", None).is_fenceable());
assert!(!ConsumerGroupMetadata::new("g", -1, "m", None).is_fenceable());
assert!(ConsumerGroupMetadata::new("g", 0, "m", None).is_fenceable());
}
#[test]
fn test_fencing_error_codes_are_fatal() {
for code in [
ErrorCode::InvalidProducerEpoch,
ErrorCode::ProducerFenced,
ErrorCode::TransactionalIdAuthorizationFailed,
ErrorCode::InvalidTxnState,
ErrorCode::TransactionCoordinatorFenced,
] {
for version in [TransactionVersion::V1, TransactionVersion::V2] {
assert!(
is_fatal_transaction_error(code, version),
"{code:?} must be classified as a fatal transaction error under {version}"
);
}
}
assert!(!is_fatal_transaction_error(
ErrorCode::NotCoordinator,
TransactionVersion::V1
));
assert!(!is_fatal_transaction_error(
ErrorCode::None,
TransactionVersion::V1
));
}
fn test_producer(version: TransactionVersion) -> TransactionalProducer {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
let metadata = Arc::new(ClusterMetadata::new(
vec!["localhost:9092".to_string()],
pool.clone(),
Duration::from_secs(300),
));
TransactionalProducer {
config: TransactionalProducerConfig {
bootstrap_servers: "localhost:9092".to_string(),
transactional_id: "txn-test".to_string(),
..TransactionalProducerConfig::default()
},
metadata,
pool,
partitioner: Arc::new(UniformStickyPartitioner::new()),
state: AtomicU8::new(TransactionState::InTransaction as u8),
transaction_version: AtomicU8::new(version as u8),
abort_required: AtomicBool::new(false),
coordinator_id: RwLock::new(None),
txn_partitions: Arc::new(RwLock::new(TransactionPartitions::default())),
identity: Arc::new(ProducerIdentity::new()),
accumulator: test_accumulator(),
metrics: Arc::new(ProducerMetrics::default()),
retry_policy: RetryPolicy::no_retries(),
in_flight_barrier: Arc::new(InFlightBarrier::new()),
key_encoder: None,
value_encoder: None,
}
}
fn support(
level: i16,
produce: i16,
txn_offset_commit: i16,
end_txn: i16,
) -> BrokerTransactionSupport {
BrokerTransactionSupport {
transaction_version_level: level,
produce_max: Some(produce),
txn_offset_commit_max: Some(txn_offset_commit),
end_txn_max: Some(end_txn),
}
}
fn tv2_broker() -> BrokerTransactionSupport {
support(
2,
versions::PRODUCE_MAX,
versions::TXN_OFFSET_COMMIT_MAX,
versions::END_TXN_MAX,
)
}
#[test]
fn test_transaction_version_from_feature_level() {
assert_eq!(
TransactionVersion::from_feature_level(0),
TransactionVersion::V1
);
assert_eq!(
TransactionVersion::from_feature_level(1),
TransactionVersion::V1
);
assert_eq!(
TransactionVersion::from_feature_level(2),
TransactionVersion::V2
);
assert_eq!(
TransactionVersion::from_feature_level(3),
TransactionVersion::V2
);
assert_eq!(
TransactionVersion::from_feature_level(-1),
TransactionVersion::V1
);
}
#[test]
fn test_transaction_version_defaults_to_v1() {
assert_eq!(TransactionVersion::default(), TransactionVersion::V1);
assert!(!TransactionVersion::V1.is_v2());
assert!(TransactionVersion::V2.is_v2());
assert_eq!(
TransactionVersion::from(TransactionVersion::V2 as u8),
TransactionVersion::V2
);
assert_eq!(
TransactionVersion::from(TransactionVersion::V1 as u8),
TransactionVersion::V1
);
assert_eq!(TransactionVersion::from(99), TransactionVersion::V1);
}
#[test]
fn test_negotiated_version_uniform_tv2_cluster() {
let cluster = [tv2_broker(), tv2_broker(), tv2_broker()];
assert_eq!(
negotiated_transaction_version(&cluster),
TransactionVersion::V2
);
}
#[test]
fn test_negotiated_version_takes_minimum_across_mixed_cluster() {
let mixed_with_v1 = [
tv2_broker(),
tv2_broker(),
support(
1,
versions::PRODUCE_MAX,
versions::TXN_OFFSET_COMMIT_MAX,
versions::END_TXN_MAX,
),
];
assert_eq!(
negotiated_transaction_version(&mixed_with_v1),
TransactionVersion::V1,
"one level-1 broker must downgrade the entire cluster to TV1"
);
let mixed_with_feature_absent = [
tv2_broker(),
support(
0,
versions::PRODUCE_MAX,
versions::TXN_OFFSET_COMMIT_MAX,
versions::END_TXN_MAX,
),
];
assert_eq!(
negotiated_transaction_version(&mixed_with_feature_absent),
TransactionVersion::V1
);
let laggard_first = [
support(
0,
versions::PRODUCE_MAX,
versions::TXN_OFFSET_COMMIT_MAX,
versions::END_TXN_MAX,
),
tv2_broker(),
];
assert_eq!(
negotiated_transaction_version(&laggard_first),
TransactionVersion::V1
);
}
#[test]
fn test_negotiated_version_empty_cluster_is_v1() {
assert_eq!(negotiated_transaction_version(&[]), TransactionVersion::V1);
}
#[test]
fn test_negotiated_version_requires_the_tv2_api_versions() {
let produce_too_old = support(
2,
TV2_MIN_PRODUCE_VERSION - 1,
versions::TXN_OFFSET_COMMIT_MAX,
versions::END_TXN_MAX,
);
assert_eq!(
negotiated_transaction_version(&[produce_too_old]),
TransactionVersion::V1,
"TV2 needs Produce v{TV2_MIN_PRODUCE_VERSION}+ to add partitions implicitly"
);
let txn_offset_commit_too_old = support(
2,
versions::PRODUCE_MAX,
TV2_MIN_TXN_OFFSET_COMMIT_VERSION - 1,
versions::END_TXN_MAX,
);
assert_eq!(
negotiated_transaction_version(&[txn_offset_commit_too_old]),
TransactionVersion::V1,
"TV2 needs TxnOffsetCommit v{TV2_MIN_TXN_OFFSET_COMMIT_VERSION}+"
);
let end_txn_too_old = support(
2,
versions::PRODUCE_MAX,
versions::TXN_OFFSET_COMMIT_MAX,
TV2_MIN_END_TXN_VERSION - 1,
);
assert_eq!(
negotiated_transaction_version(&[end_txn_too_old]),
TransactionVersion::V1,
"TV2 needs EndTxn v{TV2_MIN_END_TXN_VERSION}+ to receive the bumped epoch"
);
let at_floor = support(
2,
TV2_MIN_PRODUCE_VERSION,
TV2_MIN_TXN_OFFSET_COMMIT_VERSION,
TV2_MIN_END_TXN_VERSION,
);
assert_eq!(
negotiated_transaction_version(&[at_floor]),
TransactionVersion::V2
);
}
#[test]
fn test_negotiated_version_unnegotiable_api_is_v1() {
let no_produce = BrokerTransactionSupport {
produce_max: None,
..tv2_broker()
};
assert_eq!(
negotiated_transaction_version(&[no_produce]),
TransactionVersion::V1
);
}
#[test]
fn test_crate_supports_the_tv2_api_versions() {
const {
assert!(versions::PRODUCE_MAX >= TV2_MIN_PRODUCE_VERSION);
assert!(versions::TXN_OFFSET_COMMIT_MAX >= TV2_MIN_TXN_OFFSET_COMMIT_VERSION);
assert!(versions::END_TXN_MAX >= TV2_MIN_END_TXN_VERSION);
}
}
#[tokio::test]
async fn test_tv2_skips_explicit_partition_registration() {
let tv1 = test_producer(TransactionVersion::V1);
assert!(
tv1.requires_explicit_partition_registration(),
"TV1 must send AddPartitionsToTxn before the first write to a partition"
);
let tv2 = test_producer(TransactionVersion::V2);
assert!(
!tv2.requires_explicit_partition_registration(),
"TV2 adds partitions implicitly via Produce; AddPartitionsToTxn must be skipped"
);
}
#[tokio::test]
async fn test_tv2_produce_path_does_not_contact_the_coordinator() {
let tv2 = test_producer(TransactionVersion::V2);
tv2.identity.initialize(7, 3);
let record = ProducerRecord::new("topic", Bytes::from_static(b"value")).with_partition(0);
let _ = tokio::time::timeout(Duration::from_secs(2), tv2.send_record(record)).await;
assert!(
tv2.txn_partitions.read().await.is_empty(),
"TV2 must not record per-partition registration state"
);
assert!(
tv2.coordinator_id.read().await.is_none(),
"TV2 must not perform coordinator discovery on the produce path"
);
}
#[test]
fn test_invalid_producer_id_mapping_is_fatal_only_under_tv2() {
assert!(
!is_fatal_transaction_error(
ErrorCode::InvalidProducerIdMapping,
TransactionVersion::V1
),
"under TV1 the producer aborts and re-initializes"
);
assert!(
is_fatal_transaction_error(ErrorCode::InvalidProducerIdMapping, TransactionVersion::V2),
"under TV2 recovering in place could break exactly-once, so it is fatal"
);
let error = KrafkaError::broker(ErrorCode::InvalidProducerIdMapping, "test");
assert!(
TransactionalProducer::is_abortable_transaction_error(&error, TransactionVersion::V1),
"the TV1 classification must be abortable, not merely non-fatal"
);
assert!(
!TransactionalProducer::is_abortable_transaction_error(&error, TransactionVersion::V2),
"fatal and abortable must stay mutually exclusive"
);
}
#[test]
fn test_transaction_abortable_is_abortable_not_fatal() {
let error = KrafkaError::broker(ErrorCode::TransactionAbortable, "test");
for version in [TransactionVersion::V1, TransactionVersion::V2] {
assert!(
!is_fatal_transaction_error(ErrorCode::TransactionAbortable, version),
"TRANSACTION_ABORTABLE must not be fatal under {version}"
);
assert!(
TransactionalProducer::is_abortable_transaction_error(&error, version),
"TRANSACTION_ABORTABLE must be abortable under {version}"
);
}
assert!(!error.is_retriable());
}
#[tokio::test]
async fn test_classify_transaction_result_by_version() {
let tv2 = test_producer(TransactionVersion::V2);
let result: Result<()> = Err(KrafkaError::broker(
ErrorCode::InvalidProducerIdMapping,
"test",
));
assert!(tv2.classify_transaction_result(result).is_err());
assert_eq!(tv2.state(), TransactionState::FatalError);
assert!(
!tv2.abort_required(),
"a fatal error is unrecoverable; abort_transaction() cannot help"
);
let tv1 = test_producer(TransactionVersion::V1);
let result: Result<()> = Err(KrafkaError::broker(
ErrorCode::InvalidProducerIdMapping,
"test",
));
assert!(tv1.classify_transaction_result(result).is_err());
assert_eq!(
tv1.state(),
TransactionState::InTransaction,
"TV1 must not fence the producer over a PID-mapping mismatch"
);
assert!(
tv1.abort_required(),
"the transaction is over; the caller must abort before continuing"
);
}
#[tokio::test]
async fn test_transaction_version_accessor_defaults_to_v1() {
let producer = test_producer(TransactionVersion::V1);
assert_eq!(producer.transaction_version(), TransactionVersion::V1);
let producer = test_producer(TransactionVersion::V2);
assert_eq!(producer.transaction_version(), TransactionVersion::V2);
}
#[test]
fn test_endtxn_epoch_bump_resets_sequences() {
let identity = ProducerIdentity::new();
identity.initialize(42, 0);
for _ in 0..5 {
identity.next_sequence("orders", 0).expect("allocate");
}
identity.next_sequence("payments", 1).expect("allocate");
assert_eq!(identity.peek_sequence("orders", 0), 5);
assert_eq!(identity.peek_sequence("payments", 1), 1);
identity.bump_epoch(42, 1);
assert_eq!(identity.producer_id(), 42);
assert_eq!(identity.producer_epoch(), 1);
assert_eq!(
identity.peek_sequence("orders", 0),
0,
"a bumped epoch starts a fresh sequence space"
);
assert_eq!(identity.peek_sequence("payments", 1), 0);
}
#[test]
fn test_endtxn_bump_adopts_new_producer_id_on_epoch_overflow() {
let identity = ProducerIdentity::new();
identity.initialize(42, i16::MAX);
identity.next_sequence("orders", 0).expect("allocate");
identity.bump_epoch(1000, 0);
assert_eq!(identity.producer_id(), 1000);
assert_eq!(identity.producer_epoch(), 0);
assert_eq!(identity.peek_sequence("orders", 0), 0);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod commit_indeterminate_tests {
use super::*;
#[test]
fn a_timed_out_commit_is_not_reported_as_still_in_transaction() {
for error in [
KrafkaError::timeout("EndTxn"),
KrafkaError::network(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
"connection closed",
)),
] {
assert!(
matches!(error, KrafkaError::Timeout { .. } | KrafkaError::Network(_)),
"{error} must classify as outcome-unknown"
);
assert!(
error.is_retriable(),
"{error} must be retriable, or it would take the fatal path instead"
);
}
}
#[test]
fn a_broker_rejection_is_a_definite_answer_not_an_unknown_outcome() {
let error = KrafkaError::broker(
ErrorCode::CoordinatorNotAvailable,
"coordinator moved".to_string(),
);
assert!(error.is_retriable());
assert!(
!matches!(error, KrafkaError::Timeout { .. } | KrafkaError::Network(_)),
"a broker error must not be treated as an unknown outcome"
);
}
#[test]
fn commit_indeterminate_round_trips_through_its_discriminant() {
assert_eq!(
TransactionState::from(TransactionState::CommitIndeterminate as u8),
TransactionState::CommitIndeterminate
);
assert_eq!(
TransactionState::CommitIndeterminate.to_string(),
"CommitIndeterminate"
);
}
}