use std::collections::HashMap;
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::config::Acks;
use super::idempotent::ProducerIdentity;
use super::partitioner::{Partitioner, UniformStickyPartitioner};
use super::record::{ProducerRecord, RecordMetadata, TopicHandle};
use super::retry::RetryPolicy;
use crate::barrier::InFlightBarrier;
use crate::consumer::ConsumerGroupMetadata;
use crate::metrics::ProducerMetrics;
use crate::serdes::Serializer;
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;
const TV3_MIN_INIT_PRODUCER_ID_VERSION: i16 = 6;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Default)]
#[non_exhaustive]
#[repr(u8)]
pub enum TransactionVersion {
#[default]
V1 = 1,
V2 = 2,
V3 = 3,
}
impl From<u8> for TransactionVersion {
fn from(v: u8) -> Self {
if v == Self::V3 as u8 {
Self::V3
} else 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 >= 3 {
Self::V3
} else if level >= 2 {
Self::V2
} else {
Self::V1
}
}
#[must_use]
#[inline]
pub fn is_v2(self) -> bool {
matches!(self, Self::V2 | Self::V3)
}
#[must_use]
#[inline]
pub fn supports_two_phase_commit(self) -> bool {
matches!(self, Self::V3)
}
}
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"),
Self::V3 => write!(f, "TV3"),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct BrokerTransactionSupport {
transaction_version_level: i16,
init_producer_id_max: Option<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)
{
if feature.supports_two_phase_commit()
&& supports(self.init_producer_id_max, TV3_MIN_INIT_PRODUCER_ID_VERSION)
{
TransactionVersion::V3
} else {
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,
Prepared = 8,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PreparedTxnState {
producer_id: i64,
producer_epoch: i16,
}
impl PreparedTxnState {
#[must_use]
pub const fn none() -> Self {
Self {
producer_id: -1,
producer_epoch: -1,
}
}
#[must_use]
pub const fn is_prepared(&self) -> bool {
self.producer_id >= 0
}
#[must_use]
pub const fn producer_id(&self) -> i64 {
self.producer_id
}
#[must_use]
pub const fn producer_epoch(&self) -> i16 {
self.producer_epoch
}
}
impl std::fmt::Display for PreparedTxnState {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}:{}", self.producer_id, self.producer_epoch)
}
}
impl std::str::FromStr for PreparedTxnState {
type Err = KrafkaError;
fn from_str(s: &str) -> Result<Self> {
let malformed = || {
KrafkaError::config(format!(
"malformed PreparedTxnState {s:?}; expected `producer_id:epoch`"
))
};
let (id, epoch) = s.split_once(':').ok_or_else(malformed)?;
Ok(Self {
producer_id: id.trim().parse().map_err(|_| malformed())?,
producer_epoch: epoch.trim().parse().map_err(|_| malformed())?,
})
}
}
#[non_exhaustive]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TransactionOutcome {
Committed,
Aborted,
}
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,
8 => Self::Prepared,
_ => {
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",
Self::Prepared => "Prepared",
})
}
}
#[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: Duration,
request_timeout: Duration,
connect_timeout: Duration,
delivery_timeout: Duration,
max_request_size: usize,
compression: Compression,
compression_level: Option<i32>,
topic_compression: HashMap<String, Compression>,
batch_size: usize,
linger: Duration,
buffer_memory: usize,
max_block: Duration,
metadata_max_age: Duration,
metadata_topic_cache_ttl: Option<Duration>,
metadata_recovery_strategy: crate::metadata::MetadataRecoveryStrategy,
metadata_recovery_rebootstrap_trigger: Duration,
auth: Option<AuthConfig>,
two_phase_commit: bool,
transport: crate::network::TransportConfig,
dead_letter_queue: Option<Arc<dyn crate::dlq::DeadLetterQueue>>,
}
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: Duration::from_secs(60),
two_phase_commit: false,
request_timeout: Duration::from_secs(30),
connect_timeout: crate::network::DEFAULT_CONNECT_TIMEOUT,
delivery_timeout: Duration::from_secs(120),
max_request_size: crate::protocol::MAX_MESSAGE_SIZE,
compression: Compression::None,
compression_level: None,
topic_compression: HashMap::new(),
batch_size: 16384,
linger: Duration::from_millis(5),
buffer_memory: 32 * 1024 * 1024,
max_block: Duration::from_secs(60),
metadata_max_age: Duration::from_secs(300),
metadata_topic_cache_ttl: Some(Duration::from_secs(300)),
metadata_recovery_strategy: crate::metadata::MetadataRecoveryStrategy::Rebootstrap,
metadata_recovery_rebootstrap_trigger: Duration::from_secs(300),
auth: None,
transport: crate::network::TransportConfig::default(),
dead_letter_queue: None,
}
}
}
impl TransactionalProducerConfig {
#[inline]
pub fn bootstrap_servers(&self) -> &str {
&self.bootstrap_servers
}
#[inline]
pub fn client_id(&self) -> &str {
&self.client_id
}
#[inline]
pub fn transactional_id(&self) -> &str {
&self.transactional_id
}
#[inline]
pub fn transaction_timeout(&self) -> Duration {
self.transaction_timeout
}
#[inline]
pub fn two_phase_commit(&self) -> bool {
self.two_phase_commit
}
#[inline]
pub fn request_timeout(&self) -> Duration {
self.request_timeout
}
#[inline]
pub fn connect_timeout(&self) -> Duration {
self.connect_timeout
}
#[inline]
pub fn delivery_timeout(&self) -> Duration {
self.delivery_timeout
}
#[inline]
pub fn max_request_size(&self) -> usize {
self.max_request_size
}
#[inline]
pub fn compression(&self) -> Compression {
self.compression
}
#[inline]
pub fn compression_level(&self) -> Option<i32> {
self.compression_level
}
#[inline]
pub fn compression_for(&self, topic: &str) -> Compression {
self.topic_compression
.get(topic)
.copied()
.unwrap_or(self.compression)
}
#[inline]
pub fn batch_size(&self) -> usize {
self.batch_size
}
#[inline]
pub fn linger(&self) -> Duration {
self.linger
}
#[inline]
pub fn buffer_memory(&self) -> usize {
self.buffer_memory
}
#[inline]
pub fn max_block(&self) -> Duration {
self.max_block
}
#[inline]
pub fn metadata_max_age(&self) -> Duration {
self.metadata_max_age
}
#[inline]
pub fn metadata_topic_cache_ttl(&self) -> Option<Duration> {
self.metadata_topic_cache_ttl
}
#[inline]
pub fn metadata_recovery_strategy(&self) -> crate::metadata::MetadataRecoveryStrategy {
self.metadata_recovery_strategy
}
#[inline]
pub fn metadata_recovery_rebootstrap_trigger(&self) -> Duration {
self.metadata_recovery_rebootstrap_trigger
}
#[inline]
pub fn auth(&self) -> Option<&AuthConfig> {
self.auth.as_ref()
}
#[inline]
pub fn acks(&self) -> Acks {
Acks::All
}
}
fn validate(
config: &TransactionalProducerConfig,
has_shared_pool: bool,
transaction_timeout_set: bool,
) -> Result<()> {
if !has_shared_pool && config.bootstrap_servers.is_empty() {
return Err(KrafkaError::config("bootstrap_servers is required"));
}
if config.transactional_id.is_empty() {
return Err(KrafkaError::config("transactional_id is required"));
}
const MAX_KAFKA_STRING_LEN: usize = i16::MAX as usize;
if 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}",
config.transactional_id.len()
)));
}
if 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}",
config.client_id.len()
)));
}
if config.transaction_timeout.is_zero() {
return Err(KrafkaError::config("transaction_timeout must be > 0"));
}
if config.two_phase_commit && transaction_timeout_set {
return Err(KrafkaError::config(
"two_phase_commit and transaction_timeout contradict each other: under \
KIP-939 the coordinator must hold a prepared transaction until an \
external coordinator decides, so `transaction.max.timeout.ms` is not \
applied and the timeout is sent as i32::MAX. Silently ignoring the value \
would leave an operator believing in a bound that does not exist. Drop \
one of the two.",
));
}
if config.max_request_size == 0 {
return Err(KrafkaError::config("max_request_size must be >= 1"));
}
if config.batch_size == 0 {
return Err(KrafkaError::config("batch_size must be >= 1"));
}
if config.delivery_timeout.is_zero() {
return Err(KrafkaError::config(
"delivery_timeout must be greater than zero",
));
}
if config.buffer_memory > 0 && config.batch_size > config.buffer_memory {
return Err(KrafkaError::config(format!(
"batch_size must not exceed buffer_memory (got batch_size={}, buffer_memory={})",
config.batch_size, config.buffer_memory
)));
}
if config.batch_size > config.max_request_size {
return Err(KrafkaError::config(format!(
"batch_size must not exceed max_request_size (got batch_size={}, max_request_size={})",
config.batch_size, config.max_request_size
)));
}
super::config::validate_compression(
config.compression,
config.compression_level,
&config.topic_compression,
)?;
let transaction_timeout = config.transaction_timeout;
if config.delivery_timeout > transaction_timeout {
warn!(
delivery_timeout_secs = config.delivery_timeout.as_secs_f64(),
transaction_timeout_secs = transaction_timeout.as_secs_f64(),
"delivery_timeout exceeds transaction_timeout; the coordinator aborts the \
transaction first, so the extra delivery budget is unreachable"
);
}
Ok(())
}
#[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>,
pool_owned: bool,
partitioner: Arc<dyn Partitioner>,
state: AtomicU8,
ongoing_prepared_txn: arc_swap::ArcSwap<PreparedTxnState>,
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_serializer: Option<Arc<dyn Serializer>>,
value_serializer: Option<Arc<dyn Serializer>>,
interceptor: Arc<dyn crate::interceptor::ProducerInterceptor>,
state_store: Option<Arc<dyn super::idempotent::ErasedProducerStateStore>>,
}
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,
init_producer_id_max: conn.negotiate_api_version(
ApiKey::InitProducerId,
versions::INIT_PRODUCER_ID_MAX,
versions::INIT_PRODUCER_ID_MIN,
),
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<()> {
self.init_transactions_inner(false).await
}
pub async fn init_transactions_keeping_prepared(&self) -> Result<Option<PreparedTxnState>> {
if !self.config.two_phase_commit {
return Err(KrafkaError::invalid_state(
"init_transactions_keeping_prepared() requires \
TransactionalProducer::builder().two_phase_commit(true); without it the \
coordinator was not told to hold prepared transactions and has already \
aborted anything this transactional.id left open",
));
}
self.init_transactions_inner(true).await?;
let ongoing = **self.ongoing_prepared_txn.load();
Ok(ongoing.is_prepared().then_some(ongoing))
}
async fn init_transactions_inner(&self, keep_prepared_txn: bool) -> 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);
if self.config.two_phase_commit && !version.supports_two_phase_commit() {
self.set_state(TransactionState::Uninitialized);
let cause = if versions::INIT_PRODUCER_ID_MAX < TV3_MIN_INIT_PRODUCER_ID_VERSION {
format!(
"this build of krafka negotiates InitProducerId up to \
v{}, and enable2Pc needs \
v{TV3_MIN_INIT_PRODUCER_ID_VERSION} — enable the \
`unstable-protocol` feature",
versions::INIT_PRODUCER_ID_MAX
)
} else {
format!(
"this cluster negotiated {version}; it must finalize \
transaction.version at 3 and every broker must serve \
InitProducerId v{TV3_MIN_INIT_PRODUCER_ID_VERSION}"
)
};
return Err(KrafkaError::invalid_state(format!(
"two_phase_commit (KIP-939) is not available: {cause}. The broker \
must also grant TWO_PHASE_COMMIT alongside WRITE on \
transactional_id '{}'.",
self.config.transactional_id
)));
}
let result = self.do_init_transactions(keep_prepared_txn).await;
if result.is_err() {
self.set_state(TransactionState::Uninitialized);
}
result
}
async fn do_init_transactions(&self, keep_prepared_txn: bool) -> 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 = if self.config.two_phase_commit {
InitProducerIdRequest::two_phase_commit(
&self.config.transactional_id,
keep_prepared_txn,
)
} else {
InitProducerIdRequest::transactional(
&self.config.transactional_id,
crate::util::duration_to_millis_i32(self.config.transaction_timeout),
)
};
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.ongoing_prepared_txn.store(Arc::new(PreparedTxnState {
producer_id: response.ongoing_txn_producer_id,
producer_epoch: response.ongoing_txn_producer_epoch,
}));
if let Some(ref store) = self.state_store {
match store.load_erased().await {
Ok(Some(snapshot))
if snapshot.producer_id == self.identity.producer_id()
&& snapshot.producer_epoch == self.identity.producer_epoch() =>
{
self.identity.restore_from_snapshot(&snapshot);
info!(
pid = self.identity.producer_id(),
epoch = self.identity.producer_epoch(),
partitions = snapshot.partition_sequences.len(),
"Transactional producer identity restored from state store"
);
}
Ok(Some(_)) => {
debug!(
"State store snapshot PID/epoch mismatch — sequences not \
restored; the coordinator assigned a new producer identity"
);
}
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"
),
}
}
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> {
self.enqueue(record).await?.await
}
pub async fn enqueue(&self, record: ProducerRecord) -> Result<TransactionalDeliveryHandle<'_>> {
let send_started_at = std::time::Instant::now();
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;
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 _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 enqueued = self
.accumulator
.enqueue_routed_with_guard(
topic,
record,
record_size,
partition,
operation_guard,
send_started_at,
)
.await;
Ok(TransactionalDeliveryHandle {
inner: self.classify_produce_result(enqueued)?,
producer: self,
})
}
fn classify_produce_result<T>(&self, result: Result<T>) -> Result<T> {
match self.classify_transaction_result(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 _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 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 prepare_transaction(&self) -> Result<PreparedTxnState> {
if !self.config.two_phase_commit {
return Err(KrafkaError::invalid_state(
"prepare_transaction() requires \
TransactionalProducer::builder().two_phase_commit(true); without it the \
coordinator applies transaction.max.timeout.ms and would abort the \
prepared transaction out from under the external coordinator",
));
}
self.ensure_transaction_can_continue("prepare transaction")?;
if let Err(actual) =
self.try_transition(TransactionState::InTransaction, TransactionState::Prepared)
{
return Err(KrafkaError::invalid_state(format!(
"cannot prepare in state {actual:?}; a transaction must be open"
)));
}
let target = self.in_flight_barrier.snapshot();
self.in_flight_barrier.wait_for(target).await;
if let Err(error) = self.accumulator.flush().await {
let _ =
self.try_transition(TransactionState::Prepared, TransactionState::InTransaction);
return Err(error);
}
let state = PreparedTxnState {
producer_id: self.identity.producer_id(),
producer_epoch: self.identity.producer_epoch(),
};
info!(
transactional_id = %self.config.transactional_id,
producer_id = state.producer_id,
producer_epoch = state.producer_epoch,
"Transaction prepared; awaiting the external coordinator's decision"
);
Ok(state)
}
pub async fn complete_transaction(
&self,
stored: PreparedTxnState,
) -> Result<TransactionOutcome> {
let ongoing = **self.ongoing_prepared_txn.load();
if !ongoing.is_prepared() {
return Err(KrafkaError::invalid_state(
"complete_transaction(): the coordinator is holding no prepared \
transaction for this transactional.id. Call \
init_transactions_keeping_prepared() first, and check its return \
value — `None` means there is nothing to resolve",
));
}
if stored == ongoing {
info!(
transactional_id = %self.config.transactional_id,
producer_id = ongoing.producer_id,
"Recovered prepared transaction matches the stored state; committing"
);
self.commit_prepared(ongoing).await?;
Ok(TransactionOutcome::Committed)
} else {
info!(
transactional_id = %self.config.transactional_id,
stored = %stored,
ongoing = %ongoing,
"Recovered prepared transaction does not match the stored state; \
the prepare was never recorded externally, so aborting"
);
self.abort_prepared(ongoing).await?;
Ok(TransactionOutcome::Aborted)
}
}
async fn commit_prepared(&self, ongoing: PreparedTxnState) -> Result<()> {
self.adopt_prepared(ongoing);
self.commit_transaction().await
}
async fn abort_prepared(&self, ongoing: PreparedTxnState) -> Result<()> {
self.adopt_prepared(ongoing);
self.abort_transaction().await
}
fn adopt_prepared(&self, ongoing: PreparedTxnState) {
self.identity
.initialize(ongoing.producer_id, ongoing.producer_epoch);
self.set_state(TransactionState::Prepared);
}
pub async fn commit_transaction(&self) -> Result<()> {
self.ensure_transaction_can_continue("commit transaction")?;
let entered_from = if self
.try_transition(
TransactionState::InTransaction,
TransactionState::Committing,
)
.is_ok()
{
TransactionState::InTransaction
} else if self
.try_transition(TransactionState::Prepared, TransactionState::Committing)
.is_ok()
{
TransactionState::Prepared
} else if self
.try_transition(
TransactionState::CommitIndeterminate,
TransactionState::Committing,
)
.is_ok()
{
TransactionState::CommitIndeterminate
} else {
return Err(KrafkaError::invalid_state(format!(
"cannot commit in state {:?}",
self.state()
)));
};
self.in_flight_barrier
.wait_for(self.in_flight_barrier.snapshot())
.await;
if let Err(error) = self.accumulator.flush().await {
let _ = self.try_transition(TransactionState::Committing, entered_from);
return Err(error);
}
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()) => {
let revert_to = if entered_from == TransactionState::Prepared {
TransactionState::Prepared
} else {
TransactionState::InTransaction
};
match self.try_transition(TransactionState::Committing, revert_to) {
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
|| entered_from == TransactionState::CommitIndeterminate
{
TransactionState::CommitIndeterminate
} else {
entered_from
};
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 entered_from = if self
.try_transition(TransactionState::InTransaction, TransactionState::Aborting)
.is_ok()
{
TransactionState::InTransaction
} else if self
.try_transition(TransactionState::Prepared, TransactionState::Aborting)
.is_ok()
{
TransactionState::Prepared
} else if self
.try_transition(TransactionState::Committing, TransactionState::Aborting)
.is_ok()
{
TransactionState::InTransaction
} else {
return Err(KrafkaError::invalid_state(format!(
"cannot abort in state {:?}",
self.state()
)));
};
self.in_flight_barrier
.wait_for(self.in_flight_barrier.snapshot())
.await;
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(false).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(false).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, entered_from) {
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)
}
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(())
}
#[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);
if self.pool_owned {
self.pool.close_all().await;
info!(
"TransactionalProducer closed: txn.id()={} (connection pool torn down)",
self.config.transactional_id
);
} else {
info!(
"TransactionalProducer closed: txn.id()={} (shared connection pool left open)",
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]
#[must_use]
pub fn owns_pool(&self) -> bool {
self.pool_owned
}
#[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
)
}
#[must_use = "a dropped handle discards the acknowledgement; the record is still sent"]
pub struct TransactionalDeliveryHandle<'a> {
inner: super::DeliveryHandle,
producer: &'a TransactionalProducer,
}
impl TransactionalDeliveryHandle<'_> {
#[inline]
#[must_use]
pub fn partition(&self) -> crate::PartitionId {
self.inner.partition()
}
}
impl std::future::Future for TransactionalDeliveryHandle<'_> {
type Output = Result<RecordMetadata>;
fn poll(
mut self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Self::Output> {
let producer = self.producer;
match std::pin::Pin::new(&mut self.inner).poll(cx) {
std::task::Poll::Ready(result) => {
std::task::Poll::Ready(producer.classify_produce_result(result))
}
std::task::Poll::Pending => std::task::Poll::Pending,
}
}
}
#[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_serializer: Option<Arc<dyn Serializer>>,
value_serializer: Option<Arc<dyn Serializer>>,
interceptors: Vec<Arc<dyn crate::interceptor::ProducerInterceptor>>,
shared: Option<(Arc<ConnectionPool>, Arc<ClusterMetadata>)>,
state_store: Option<Arc<dyn super::idempotent::ErasedProducerStateStore>>,
transaction_timeout_set: bool,
}
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 = timeout;
self.transaction_timeout_set = true;
self
}
pub fn two_phase_commit(mut self, enable: bool) -> Self {
self.config.two_phase_commit = enable;
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 delivery_timeout(mut self, timeout: Duration) -> Self {
self.config.delivery_timeout = timeout;
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 topic_compression(mut self, topic: impl Into<String>, compression: Compression) -> Self {
self.config
.topic_compression
.insert(topic.into(), compression);
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 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 state_store(mut self, store: impl super::ProducerStateStore + 'static) -> Self {
self.state_store = Some(Arc::new(store));
self
}
pub fn with_client(mut self, client: &crate::client::KrafkaClient) -> Self {
self.shared = Some((client.pool().clone(), client.metadata().clone()));
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 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.transport.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: 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_provider(
mut self,
provider: impl crate::auth::OAuthBearerTokenProvider + 'static,
) -> Self {
self.config.auth = Some(AuthConfig::sasl_oauthbearer_provider(provider));
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 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 fn build_config(self) -> Result<TransactionalProducerConfig> {
validate(
&self.config,
self.shared.is_some(),
self.transaction_timeout_set,
)?;
Ok(self.config)
}
pub async fn build(self) -> Result<TransactionalProducer> {
validate(
&self.config,
self.shared.is_some(),
self.transaction_timeout_set,
)?;
let pool_owned = self.shared.is_none();
let (pool, metadata) = if let Some((pool, metadata)) = self.shared.clone() {
(pool, metadata)
} else {
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());
}
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({
let mut meta = 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);
if let Some(ttl) = self.config.metadata_topic_cache_ttl {
meta = meta.with_topic_cache_ttl(ttl);
} else {
meta = meta.with_topic_cache_ttl_disabled();
}
meta
});
metadata.refresh().await?;
(pool, metadata)
};
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 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 retry_policy = self
.retry_policy
.with_delivery_timeout(Some(self.config.delivery_timeout));
let accumulator = RecordAccumulator::spawn(
AccumulatorConfig {
batch_size: self.config.batch_size,
linger: self.config.linger,
compression: self.config.compression,
compression_level: self.config.compression_level,
topic_compression: self.config.topic_compression.clone().into_iter().collect(),
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,
interceptor: interceptor.clone(),
identity: Some(identity.clone()),
partitioner: partitioner.clone(),
state_store: self.state_store.clone(),
transactional_id: Some(self.config.transactional_id.clone()),
dead_letter_queue: self.config.dead_letter_queue.clone(),
},
metadata.clone(),
retry_policy.clone(),
metrics.clone(),
in_flight_barrier.clone(),
);
Ok(TransactionalProducer {
config: self.config,
metadata,
pool,
partitioner,
state: AtomicU8::new(TransactionState::Uninitialized as u8),
ongoing_prepared_txn: arc_swap::ArcSwap::from_pointee(PreparedTxnState::none()),
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,
in_flight_barrier,
key_serializer: self.key_serializer,
value_serializer: self.value_serializer,
interceptor,
state_store: self.state_store,
pool_owned,
})
}
}
#[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, Duration::from_secs(60));
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),
ongoing_prepared_txn: arc_swap::ArcSwap::from_pointee(PreparedTxnState::none()),
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_serializer: None,
value_serializer: None,
interceptor: Arc::new(crate::interceptor::NoOpProducerInterceptor),
state_store: None,
pool_owned: true,
};
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),
ongoing_prepared_txn: arc_swap::ArcSwap::from_pointee(PreparedTxnState::none()),
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_serializer: None,
value_serializer: None,
interceptor: Arc::new(crate::interceptor::NoOpProducerInterceptor),
state_store: None,
pool_owned: true,
};
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),
ongoing_prepared_txn: arc_swap::ArcSwap::from_pointee(PreparedTxnState::none()),
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_serializer: None,
value_serializer: None,
interceptor: Arc::new(crate::interceptor::NoOpProducerInterceptor),
state_store: None,
pool_owned: true,
};
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());
}
fn valid_txn_builder() -> TransactionalProducerBuilder {
TransactionalProducer::builder()
.bootstrap_servers("localhost:9092")
.transactional_id("txn-1")
}
#[test]
fn build_config_returns_a_validated_config_without_connecting() {
let config = valid_txn_builder()
.client_id("checkout")
.delivery_timeout(Duration::from_secs(45))
.build_config()
.expect("a minimal transactional configuration is valid");
assert_eq!(config.transactional_id(), "txn-1");
assert_eq!(config.client_id(), "checkout");
assert_eq!(config.delivery_timeout(), Duration::from_secs(45));
assert_eq!(
config.acks(),
Acks::All,
"acks is fixed, not merely defaulted"
);
}
#[test]
fn build_config_rejects_a_missing_transactional_id() {
let err = TransactionalProducer::builder()
.bootstrap_servers("localhost:9092")
.build_config()
.expect_err("transactional_id is required")
.to_string();
assert!(err.contains("transactional_id"), "got: {err}");
}
#[test]
fn build_config_rejects_an_empty_bootstrap_list() {
let err = TransactionalProducer::builder()
.transactional_id("txn-1")
.build_config()
.expect_err("bootstrap servers are required without a shared client")
.to_string();
assert!(err.contains("bootstrap"), "got: {err}");
}
#[test]
fn build_config_rejects_zero_delivery_timeout() {
let err = valid_txn_builder()
.delivery_timeout(Duration::ZERO)
.build_config()
.expect_err("a zero delivery budget can never be met")
.to_string();
assert!(err.contains("delivery_timeout"), "got: {err}");
}
#[test]
fn build_config_rejects_a_batch_larger_than_the_buffer() {
let err = valid_txn_builder()
.batch_size(4096)
.buffer_memory(1024)
.build_config()
.expect_err("a batch that cannot fit in the buffer would deadlock")
.to_string();
assert!(err.contains("buffer_memory"), "got: {err}");
}
#[cfg(feature = "snappy")]
#[test]
fn build_config_rejects_a_level_on_a_levelless_codec() {
let err = valid_txn_builder()
.compression(Compression::Snappy)
.compression_level(Some(9))
.build_config()
.expect_err("Snappy takes no level")
.to_string();
assert!(
err.contains("takes no level"),
"the error must say the codec has no level, got: {err}"
);
}
#[cfg(feature = "gzip")]
#[test]
fn build_config_rejects_an_out_of_range_level() {
let err = valid_txn_builder()
.compression(Compression::Gzip)
.compression_level(Some(42))
.build_config()
.expect_err("gzip tops out at 9")
.to_string();
assert!(
err.contains("0..=9"),
"the error must name the range: {err}"
);
}
#[cfg(feature = "zstd")]
#[test]
fn a_valid_compression_level_reaches_the_config() {
let config = valid_txn_builder()
.compression(Compression::Zstd)
.compression_level(Some(1))
.build_config()
.expect("level 1 is valid for zstd");
assert_eq!(config.compression_level(), Some(1));
}
#[cfg(all(feature = "zstd", feature = "snappy"))]
#[test]
fn a_per_topic_codec_is_validated_against_the_level() {
let err = valid_txn_builder()
.compression(Compression::Zstd)
.compression_level(Some(1))
.topic_compression("events", Compression::Snappy)
.build_config()
.expect_err("the per-topic Snappy override takes no level")
.to_string();
assert!(
err.contains("events"),
"the error must name the topic: {err}"
);
}
#[test]
fn topic_compression_overrides_reach_the_config() {
let config = valid_txn_builder()
.topic_compression("high-volume", Compression::None)
.build_config()
.expect("an override to None is always available");
assert_eq!(config.compression_for("high-volume"), Compression::None);
assert_eq!(
config.compression_for("anything-else"),
config.compression()
);
}
#[test]
fn metadata_topic_cache_ttl_round_trips_and_can_be_disabled() {
let ttl = valid_txn_builder()
.metadata_topic_cache_ttl(Duration::from_secs(600))
.build_config()
.expect("valid");
assert_eq!(
ttl.metadata_topic_cache_ttl(),
Some(Duration::from_secs(600))
);
let disabled = valid_txn_builder()
.disable_metadata_topic_cache_ttl()
.build_config()
.expect("valid");
assert_eq!(disabled.metadata_topic_cache_ttl(), 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 prepared_txn_state_round_trips_through_a_string() {
let state = PreparedTxnState {
producer_id: 4242,
producer_epoch: 7,
};
assert_eq!(state.to_string(), "4242:7");
assert_eq!(
"4242:7".parse::<PreparedTxnState>().expect("valid"),
state,
"a state written to a database must read back identical"
);
assert_eq!(
" 4242 : 7 ".parse::<PreparedTxnState>().expect("valid"),
state
);
for malformed in ["", "4242", "4242:", ":7", "abc:7", "4242:xyz"] {
let err = malformed
.parse::<PreparedTxnState>()
.expect_err("malformed state must not silently become a valid one");
assert!(
err.to_string().contains("producer_id:epoch"),
"the error must show the expected shape, got: {err}"
);
}
}
#[test]
fn the_absent_prepared_state_is_distinguishable_from_a_real_one() {
assert!(!PreparedTxnState::none().is_prepared());
assert!(
PreparedTxnState {
producer_id: 0,
producer_epoch: 0
}
.is_prepared(),
"producer ID 0 is a real producer ID, not an absence"
);
}
#[test]
fn two_phase_commit_and_an_explicit_timeout_are_refused_together() {
let err = TransactionalProducer::builder()
.bootstrap_servers("localhost:9092")
.transactional_id("txn")
.two_phase_commit(true)
.transaction_timeout(Duration::from_secs(30))
.build_config()
.expect_err("the two settings contradict each other");
assert!(err.to_string().contains("two_phase_commit"), "got: {err}");
TransactionalProducer::builder()
.bootstrap_servers("localhost:9092")
.transactional_id("txn")
.two_phase_commit(true)
.build_config()
.expect("2PC without an explicit timeout is the normal configuration");
}
#[tokio::test]
async fn the_two_phase_entry_points_require_the_setting() {
let producer = test_producer(TransactionVersion::V2);
producer.set_state(TransactionState::InTransaction);
let err = producer
.prepare_transaction()
.await
.expect_err("prepare without 2PC must be refused");
assert!(err.to_string().contains("two_phase_commit"), "got: {err}");
let err = producer
.init_transactions_keeping_prepared()
.await
.expect_err("keeping prepared transactions without 2PC must be refused");
assert!(err.to_string().contains("two_phase_commit"), "got: {err}");
}
#[tokio::test]
async fn completing_without_a_prepared_transaction_is_an_error() {
let producer = test_producer(TransactionVersion::V2);
producer.set_state(TransactionState::Ready);
let err = producer
.complete_transaction(PreparedTxnState {
producer_id: 1,
producer_epoch: 0,
})
.await
.expect_err("there is nothing to complete");
assert!(
err.to_string()
.contains("init_transactions_keeping_prepared"),
"the error must name the call that was skipped, got: {err}"
);
}
#[test]
fn test_transaction_state_unknown_maps_to_fatal() {
for state in [
TransactionState::Uninitialized,
TransactionState::Ready,
TransactionState::InTransaction,
TransactionState::Committing,
TransactionState::Aborting,
TransactionState::FatalError,
TransactionState::Initializing,
TransactionState::CommitIndeterminate,
TransactionState::Prepared,
] {
assert_eq!(
TransactionState::from(state as u8),
state,
"{state} must survive the u8 round trip"
);
}
assert_eq!(TransactionState::from(9), 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
));
}
pub(super) fn test_producer(version: TransactionVersion) -> TransactionalProducer {
test_producer_at(version, "localhost:9092")
}
pub(super) fn test_producer_at(
version: TransactionVersion,
address: &str,
) -> TransactionalProducer {
let pool = Arc::new(ConnectionPool::new(ConnectionConfig::default()));
let metadata = Arc::new(ClusterMetadata::new(
vec![address.to_string()],
pool.clone(),
Duration::from_secs(300),
));
TransactionalProducer {
config: TransactionalProducerConfig {
bootstrap_servers: address.to_string(),
transactional_id: "txn-test".to_string(),
..TransactionalProducerConfig::default()
},
metadata,
pool,
partitioner: Arc::new(UniformStickyPartitioner::new()),
state: AtomicU8::new(TransactionState::InTransaction as u8),
ongoing_prepared_txn: arc_swap::ArcSwap::from_pointee(PreparedTxnState::none()),
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_serializer: None,
value_serializer: None,
interceptor: Arc::new(crate::interceptor::NoOpProducerInterceptor),
state_store: None,
pool_owned: true,
}
}
#[tokio::test]
async fn close_only_tears_down_a_pool_it_owns() {
for pool_owned in [true, false] {
let mut producer = test_producer(TransactionVersion::V2);
producer.pool_owned = pool_owned;
producer.set_state(TransactionState::Ready);
let pool = producer.pool.clone();
pool.start_idle_evictor();
assert!(pool.has_background_tasks());
assert_eq!(producer.owns_pool(), pool_owned);
producer.close().await;
assert!(producer.is_closed());
assert_eq!(
pool.has_background_tasks(),
!pool_owned,
"pool_owned={pool_owned}: the pool must be torn down only when owned"
);
}
}
fn support(
level: i16,
produce: i16,
txn_offset_commit: i16,
end_txn: i16,
) -> BrokerTransactionSupport {
BrokerTransactionSupport {
transaction_version_level: level,
init_producer_id_max: Some(TV3_MIN_INIT_PRODUCER_ID_VERSION),
produce_max: Some(produce),
txn_offset_commit_max: Some(txn_offset_commit),
end_txn_max: Some(end_txn),
}
}
#[test]
fn tv3_requires_an_init_producer_id_that_can_carry_enable_2pc() {
let mut broker = support(
3,
TV2_MIN_PRODUCE_VERSION,
TV2_MIN_TXN_OFFSET_COMMIT_VERSION,
TV2_MIN_END_TXN_VERSION,
);
assert_eq!(broker.version(), TransactionVersion::V3);
broker.init_producer_id_max = Some(TV3_MIN_INIT_PRODUCER_ID_VERSION - 1);
assert_eq!(
broker.version(),
TransactionVersion::V2,
"a broker that cannot encode enable2Pc is not a TV3 broker, whatever \
the feature level says"
);
broker.init_producer_id_max = None;
assert_eq!(broker.version(), TransactionVersion::V2);
let mut level_2 = support(
2,
TV2_MIN_PRODUCE_VERSION,
TV2_MIN_TXN_OFFSET_COMMIT_VERSION,
TV2_MIN_END_TXN_VERSION,
);
level_2.init_producer_id_max = Some(TV3_MIN_INIT_PRODUCER_ID_VERSION);
assert_eq!(level_2.version(), TransactionVersion::V2);
}
#[test]
fn a_single_lagging_broker_holds_the_cluster_below_tv3() {
let tv3 = support(
3,
TV2_MIN_PRODUCE_VERSION,
TV2_MIN_TXN_OFFSET_COMMIT_VERSION,
TV2_MIN_END_TXN_VERSION,
);
let mut lagging = tv3;
lagging.init_producer_id_max = Some(TV3_MIN_INIT_PRODUCER_ID_VERSION - 1);
assert_eq!(
negotiated_transaction_version(&[tv3, lagging]),
TransactionVersion::V2,
"a rolling upgrade must not enable 2PC before every broker can serve it"
);
}
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::V3
);
assert_eq!(
TransactionVersion::from_feature_level(4),
TransactionVersion::V3
);
assert!(TransactionVersion::V2.is_v2());
assert!(
TransactionVersion::V3.is_v2(),
"TV3 is a superset of TV2, not an alternative to it"
);
assert!(!TransactionVersion::V1.is_v2());
assert!(TransactionVersion::V3.supports_two_phase_commit());
assert!(!TransactionVersion::V2.supports_two_phase_commit());
assert!(!TransactionVersion::V1.supports_two_phase_commit());
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"
);
}
#[tokio::test]
async fn interceptors_run_on_the_transactional_send_path() {
use std::sync::atomic::AtomicUsize;
#[derive(Debug, Default)]
struct CountingInterceptor {
sends: AtomicUsize,
}
impl crate::interceptor::ProducerInterceptor for CountingInterceptor {
fn on_send(
&self,
record: &mut ProducerRecord,
) -> crate::interceptor::InterceptorResult {
self.sends.fetch_add(1, Ordering::SeqCst);
record
.headers
.push(("seen-by".to_string(), Bytes::from_static(b"interceptor")));
Ok(())
}
}
let interceptor = Arc::new(CountingInterceptor::default());
let mut producer = test_producer(TransactionVersion::V2);
producer.interceptor = interceptor.clone();
producer.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), producer.send_record(record)).await;
assert_eq!(
interceptor.sends.load(Ordering::SeqCst),
1,
"on_send must be invoked exactly once per transactional send"
);
}
#[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"
);
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod failed_completion_revert_tests {
use super::tests::{test_producer, test_producer_at};
use super::*;
#[tokio::test]
async fn a_failed_flush_returns_the_commit_to_the_state_it_entered_from() {
for entered_from in [
TransactionState::InTransaction,
TransactionState::Prepared,
TransactionState::CommitIndeterminate,
] {
let producer = test_producer(TransactionVersion::V2);
producer.set_state(entered_from);
producer
.accumulator
.shutdown()
.await
.expect("accumulator shutdown succeeds");
let err = producer
.commit_transaction()
.await
.expect_err("flush against a shut-down accumulator must fail");
assert!(
!err.is_retriable(),
"an accumulator-closed error is invalid_state: {err}"
);
assert_eq!(
producer.state(),
entered_from,
"a commit that never reached the coordinator must return to \
the state it was entered from"
);
}
}
#[tokio::test]
async fn an_indeterminate_commit_stays_abort_proof_across_a_failed_retry() {
let producer = test_producer(TransactionVersion::V2);
producer.set_state(TransactionState::CommitIndeterminate);
producer
.accumulator
.shutdown()
.await
.expect("accumulator shutdown succeeds");
let _ = producer
.commit_transaction()
.await
.expect_err("the retried commit fails on flush");
let err = producer
.abort_transaction()
.await
.expect_err("abort must still be refused after the failed retry");
assert!(
err.to_string().contains("KAFKA-17754"),
"the refusal must explain the hazard, got: {err}"
);
}
#[tokio::test]
async fn a_retriable_abort_failure_returns_a_prepared_transaction_to_prepared() {
let producer = test_producer_at(TransactionVersion::V1, "127.0.0.1:9");
producer.identity.initialize(7, 0);
producer.set_state(TransactionState::Prepared);
let err = producer
.abort_transaction()
.await
.expect_err("no coordinator is reachable");
assert!(
err.is_retriable(),
"a connection failure is retriable: {err}"
);
assert_eq!(
producer.state(),
TransactionState::Prepared,
"a failed abort must not reopen a prepared transaction to sends"
);
let err = producer
.abort_transaction()
.await
.expect_err("still unreachable");
assert!(err.is_retriable());
assert_eq!(producer.state(), TransactionState::Prepared);
}
}