use std::future::Future;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU8, Ordering};
use tokio::sync::{Notify, RwLock};
use tracing::{debug, info, warn};
use crate::client::CloseOptions;
use crate::error::{ErrorCode, KrafkaError, ProtocolErrorKind, Result};
use crate::metadata::ClusterMetadata;
use crate::network::{BrokerConnection, ConnectionPool};
use crate::protocol::{
AddOffsetsToTxnRequest, AddOffsetsToTxnResponse, AddPartitionsToTxnRequest,
AddPartitionsToTxnResponse, ApiKey, EndTxnRequest, EndTxnResponse, FindCoordinatorRequest,
FindCoordinatorResponse, InitProducerIdRequest, InitProducerIdResponse, TxnOffsetCommitRequest,
TxnOffsetCommitResponse, VersionedDecode, VersionedEncode, versions,
};
use crate::{Offset, PartitionId};
use super::Producer;
use super::accumulator::DeliveryHandle;
use super::gate::TxnGate;
use super::record::{Record, RecordMetadata, TopicHandle};
use super::retry::{self, Backoff};
use crate::consumer::ConsumerGroupMetadata;
use crate::metrics::Metrics;
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 = 5;
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]
pub enum TransactionState {
Uninitialized,
Initializing,
Ready,
Open,
Committing,
CommitUnknown,
Aborting,
Prepared,
Fatal,
}
#[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 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::Initializing => "Initializing",
Self::Ready => "Ready",
Self::Open => "Open",
Self::Committing => "Committing",
Self::CommitUnknown => "CommitUnknown",
Self::Aborting => "Aborting",
Self::Prepared => "Prepared",
Self::Fatal => "Fatal",
})
}
}
#[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)]
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 {
producer: Producer,
transactional_id: String,
gate: Arc<TxnGate>,
identity: parking_lot::Mutex<(i64, i16)>,
ongoing_prepared_txn: arc_swap::ArcSwap<PreparedTxnState>,
transaction_version: AtomicU8,
end_txn_unanswered: AtomicBool,
reinit_required: AtomicBool,
reinit_lock: tokio::sync::Mutex<()>,
coordinator_id: RwLock<Option<i32>>,
txn_partitions: Arc<RwLock<TransactionPartitions>>,
backoff: Backoff,
}
impl std::fmt::Debug for TransactionalProducer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TransactionalProducer")
.field("transactional_id", &self.transactional_id)
.field("state", &self.state())
.finish_non_exhaustive()
}
}
impl TransactionalProducer {
pub(crate) async fn start(
producer: Producer,
transactional_id: String,
gate: Arc<TxnGate>,
) -> Result<Self> {
let backoff = Backoff::new(producer.config.retry_backoff);
let keep_prepared = producer.config.two_phase_commit;
let txn = Self {
producer,
transactional_id,
gate,
identity: parking_lot::Mutex::new((-1, -1)),
ongoing_prepared_txn: arc_swap::ArcSwap::from_pointee(PreparedTxnState::none()),
transaction_version: AtomicU8::new(TransactionVersion::V1 as u8),
end_txn_unanswered: AtomicBool::new(false),
reinit_required: AtomicBool::new(false),
reinit_lock: tokio::sync::Mutex::new(()),
coordinator_id: RwLock::new(None),
txn_partitions: Arc::new(RwLock::new(TransactionPartitions::default())),
backoff,
};
if let Err(error) = txn.init_transactions(keep_prepared).await {
let _ = txn.producer.close().await;
return Err(error);
}
info!(transactional_id = %txn.transactional_id, "transactional producer started");
Ok(txn)
}
fn metadata(&self) -> &ClusterMetadata {
self.producer.kafka.metadata()
}
fn pool(&self) -> &ConnectionPool {
self.producer.kafka.pool()
}
#[inline]
pub fn state(&self) -> TransactionState {
self.gate.state()
}
#[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, producer_epoch) = *self.identity.lock();
if producer_id < 0 || producer_epoch < 0 {
return Err(KrafkaError::illegal_state(
"transactional producer identity not initialized",
));
}
Ok((producer_id, producer_epoch))
}
fn adopt_identity(&self, producer_id: i64, producer_epoch: i16) {
*self.identity.lock() = (producer_id, producer_epoch);
self.producer.accumulator.set_identity(
producer_id,
producer_epoch,
self.transaction_version().is_v2(),
);
}
fn deadline(&self) -> tokio::time::Instant {
tokio::time::Instant::now() + self.producer.config.max_block
}
fn classify_transaction_result<T>(&self, result: Result<T>) -> Result<T> {
let Err(error) = &result else {
return result;
};
if let KrafkaError::Broker { code, message } = error
&& is_fatal_transaction_error(*code, self.transaction_version())
{
warn!(
error_code = ?code,
"Fatal transactional error from coordinator; producer must be recreated"
);
self.gate.set(TransactionState::Fatal);
if is_fencing_code(*code) {
return Err(KrafkaError::fenced(format!("{code:?}: {message}")));
}
return result;
}
if Self::is_abortable_transaction_error(error, self.transaction_version()) {
self.gate.fail(error.clone());
}
result
}
fn is_abortable_transaction_error(error: &KrafkaError, version: TransactionVersion) -> bool {
let code = match error {
KrafkaError::TransactionAbortable { .. } => return true,
KrafkaError::Broker { code, .. } => code,
_ => return false,
};
match code {
ErrorCode::TransactionAbortable | ErrorCode::UnknownProducerId => true,
ErrorCode::InvalidProducerIdMapping => !version.is_v2(),
_ => false,
}
}
async fn coordinator_connection(&self, attempt: u32) -> Result<Arc<BrokerConnection>> {
let cached = *self.coordinator_id.read().await;
let coordinator_id = match cached {
Some(id) => id,
None => {
let id = self.find_coordinator(attempt).await?;
*self.coordinator_id.write().await = Some(id);
debug!("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::broker(
ErrorCode::CoordinatorNotAvailable,
"the transaction coordinator is not in the metadata",
)
})?;
self.pool()
.get_connection_by_id(broker.id(), broker.address())
.await
}
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 with_coordinator<T, F, Fut>(&self, what: &str, attempt: F) -> Result<T>
where
F: Fn(u32) -> Fut,
Fut: Future<Output = Result<T>>,
{
retry::until_deadline(&self.backoff, self.deadline(), what, |n| {
let call = attempt(n);
async move {
let result = call.await;
if let Err(error) = &result
&& Self::needs_coordinator_refresh(error)
{
*self.coordinator_id.write().await = None;
}
result
}
})
.await
}
async fn init_transactions(&self, keep_prepared_txn: bool) -> Result<()> {
if let Err(actual) = self.gate.transition(
&[TransactionState::Uninitialized],
TransactionState::Initializing,
) {
return Err(KrafkaError::illegal_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.producer.config.two_phase_commit && !version.supports_two_phase_commit() {
self.gate.set(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::illegal_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.transactional_id
)));
}
match self.init_producer_id(keep_prepared_txn, None).await {
Ok(()) => {
self.gate.set(TransactionState::Ready);
Ok(())
}
Err(error) => {
self.gate.set(TransactionState::Uninitialized);
Err(error)
}
}
}
async fn init_producer_id(
&self,
keep_prepared_txn: bool,
current: Option<(i64, i16)>,
) -> Result<()> {
let result = self
.with_coordinator("InitProducerId", |attempt| async move {
let conn = self.coordinator_connection(attempt).await?;
let version = conn
.negotiate_api_version(
ApiKey::InitProducerId,
versions::INIT_PRODUCER_ID_MAX,
versions::INIT_PRODUCER_ID_MIN,
)
.ok_or_else(|| KrafkaError::transactions_unsupported(ApiKey::InitProducerId))?;
conn.negotiate_api_version(
ApiKey::EndTxn,
versions::END_TXN_MAX,
versions::END_TXN_MIN,
)
.ok_or_else(|| KrafkaError::transactions_unsupported(ApiKey::EndTxn))?;
let mut request = if self.producer.config.two_phase_commit {
InitProducerIdRequest::two_phase_commit(
&self.transactional_id,
keep_prepared_txn,
)
} else {
InitProducerIdRequest::transactional(
&self.transactional_id,
crate::util::duration_to_millis_i32(
self.producer.config.transaction_timeout(),
),
)
};
if let Some((producer_id, epoch)) = current {
request.producer_id = producer_id;
request.producer_epoch = epoch;
}
let mut bytes = conn
.send_request(ApiKey::InitProducerId, version, |buf| {
request.encode_versioned(version, buf)
})
.await?;
let response = InitProducerIdResponse::decode_versioned(version, &mut bytes)?;
if !response.is_ok() {
return Err(KrafkaError::broker(
response.error_code,
"failed to initialize producer ID",
));
}
Ok(response)
})
.await;
let response = self.classify_transaction_result(result)?;
self.adopt_identity(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,
}));
self.reinit_required.store(false, Ordering::SeqCst);
self.end_txn_unanswered.store(false, Ordering::SeqCst);
info!(
"Transactional producer initialized: PID={}, epoch={}",
response.producer_id, response.producer_epoch
);
Ok(())
}
async fn ensure_reinitialised(&self) -> Result<()> {
if !self.reinit_required.load(Ordering::SeqCst) {
return Ok(());
}
let _guard = self.reinit_lock.lock().await;
if !self.reinit_required.load(Ordering::SeqCst) {
return Ok(());
}
let current = self.checked_transactional_identity()?;
self.init_producer_id(false, Some(current)).await
}
async fn find_coordinator(&self, attempt: u32) -> Result<i32> {
let brokers = self.metadata().brokers();
if brokers.is_empty() {
return Err(KrafkaError::broker(
ErrorCode::CoordinatorNotAvailable,
"no brokers known",
));
}
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.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(&self) -> Result<()> {
if self.producer.barrier.is_closing() {
return Err(KrafkaError::closed("transactional producer is closed"));
}
self.gate.begin().map_err(|actual| {
KrafkaError::illegal_state(format!("cannot begin a transaction in state {actual}"))
})?;
debug!("Transaction started");
Ok(())
}
pub async fn send(&self, record: Record) -> Result<RecordMetadata> {
self.enqueue(record).await?.await
}
pub async fn enqueue(&self, record: Record) -> Result<DeliveryHandle> {
if let Err(error) = self.gate.check_open("send") {
return Err(self.refuse_send(record, error));
}
if let Err(error) = self.ensure_reinitialised().await {
return Err(self.refuse_send(record, error));
}
let explicit = self.requires_explicit_partition_registration();
let gate = Arc::clone(&self.gate);
super::enqueue_record(
&self.producer.accumulator,
self.metadata(),
&self.producer.partitioning,
&*self.producer.interceptor,
self.producer.kafka.client_id(),
self.producer.config.max_block,
record,
|topic, partition| async move {
if explicit {
self.add_partition_to_txn_if_needed(&topic, partition).await
} else {
Ok(())
}
},
move || gate.admit().map(Some),
)
.await
}
fn refuse_send(&self, mut record: Record, error: KrafkaError) -> KrafkaError {
match super::SendObligation::on_send(
&*self.producer.interceptor,
&mut record,
self.producer.kafka.client_id(),
) {
Ok(mut obligation) => obligation.fail(super::UNKNOWN_PARTITION, &record.headers, error),
Err(_) => error,
}
}
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
.with_coordinator("AddPartitionsToTxn", |attempt| async move {
let conn = self.coordinator_connection(attempt).await?;
let (producer_id, producer_epoch) = self.checked_transactional_identity()?;
let version = conn
.negotiate_api_version(
ApiKey::AddPartitionsToTxn,
versions::ADD_PARTITIONS_TO_TXN_MAX,
versions::ADD_PARTITIONS_TO_TXN_MIN,
)
.ok_or_else(|| {
KrafkaError::transactions_unsupported(ApiKey::AddPartitionsToTxn)
})?;
let request = AddPartitionsToTxnRequest::new(
&self.transactional_id,
producer_id,
producer_epoch,
)
.add_partition(topic, partition);
let mut bytes = conn
.send_request(ApiKey::AddPartitionsToTxn, version, |buf| {
request.encode_versioned(version, buf)
})
.await?;
let response = AddPartitionsToTxnResponse::decode_versioned(version, &mut bytes)?;
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 {topic}-{partition} to transaction"),
));
}
}
}
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!(
"failed to add {topic}-{partition} to transaction: the response \
reported an error but no per-partition error"
),
));
}
debug!("Added partition {}-{} to transaction", topic, partition);
Ok(())
})
.await;
self.classify_transaction_result(result)
}
pub async fn send_offsets(
&self,
offsets: &[TopicPartitionOffset],
group_metadata: &ConsumerGroupMetadata,
) -> Result<()> {
self.gate.check_open("send offsets")?;
self.ensure_reinitialised().await?;
let ticket = self.gate.admit()?;
let result = self.send_offsets_inner(offsets, group_metadata).await;
ticket.complete(result.as_ref().map(|_| ()));
result
}
async fn send_offsets_inner(
&self,
offsets: &[TopicPartitionOffset],
group_metadata: &ConsumerGroupMetadata,
) -> Result<()> {
if !group_metadata.is_fenceable() {
return Err(KrafkaError::transaction_abortable(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() 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.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 commit_request = &commit_request;
let result = retry::until_deadline(
&self.backoff,
self.deadline(),
"TxnOffsetCommit",
|attempt| async move {
let (group_node_id, group_host, group_port) =
self.find_group_coordinator(group_id, attempt).await?;
let group_conn = self
.pool()
.get_connection_by_id(group_node_id, &format!("{group_host}:{group_port}"))
.await?;
let 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 mut bytes = group_conn
.send_request(ApiKey::TxnOffsetCommit, version, |buf| {
commit_request.encode_versioned(version, buf)
})
.await?;
let response = TxnOffsetCommitResponse::decode_versioned(version, &mut bytes)?;
if !response.is_ok() {
for topic_result in &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 result.is_ok() {
debug!("Added offsets to transaction for group {}", group_id);
}
result
}
async fn add_offsets_to_txn(
&self,
producer_id: i64,
producer_epoch: i16,
group_id: &str,
) -> Result<()> {
let result = self
.with_coordinator("AddOffsetsToTxn", |attempt| async move {
let conn = self.coordinator_connection(attempt).await?;
let request = AddOffsetsToTxnRequest::new(
&self.transactional_id,
producer_id,
producer_epoch,
group_id,
);
let version = conn
.negotiate_api_version(
ApiKey::AddOffsetsToTxn,
versions::ADD_OFFSETS_TO_TXN_MAX,
versions::ADD_OFFSETS_TO_TXN_MIN,
)
.ok_or_else(|| {
KrafkaError::transactions_unsupported(ApiKey::AddOffsetsToTxn)
})?;
let mut bytes = conn
.send_request(ApiKey::AddOffsetsToTxn, version, |buf| {
request.encode_versioned(version, buf)
})
.await?;
let response = AddOffsetsToTxnResponse::decode_versioned(version, &mut bytes)?;
if !response.is_ok() {
return Err(KrafkaError::broker(
response.error_code,
"failed to add offsets to transaction",
));
}
Ok(())
})
.await;
self.classify_transaction_result(result)
}
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::broker(
ErrorCode::CoordinatorNotAvailable,
"no brokers known",
));
}
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))
}
async fn drain_transaction(&self) -> Option<KrafkaError> {
let generation = self.producer.barrier.snapshot();
self.producer.accumulator.flush(generation);
self.gate.drained().await;
self.gate.failure("commit")
}
pub async fn prepare(&self) -> Result<PreparedTxnState> {
if !self.producer.config.two_phase_commit {
return Err(KrafkaError::illegal_state(
"prepare() requires ProducerBuilder::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",
));
}
if let Some(error) = self.gate.failure("prepare") {
return Err(error);
}
self.gate
.transition(&[TransactionState::Open], TransactionState::Prepared)
.map_err(|actual| {
KrafkaError::illegal_state(format!(
"cannot prepare in state {actual}; a transaction must be open"
))
})?;
if let Some(error) = self.drain_transaction().await {
let _ = self
.gate
.transition(&[TransactionState::Prepared], TransactionState::Open);
return Err(error);
}
let (producer_id, producer_epoch) = *self.identity.lock();
let state = PreparedTxnState {
producer_id,
producer_epoch,
};
info!(
transactional_id = %self.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(&self, stored: PreparedTxnState) -> Result<TransactionOutcome> {
let ongoing = **self.ongoing_prepared_txn.load();
if !ongoing.is_prepared() {
return Err(KrafkaError::illegal_state(
"complete(): the coordinator is holding no prepared transaction for this \
transactional.id; prepared_transaction() returned None",
));
}
self.adopt_prepared(ongoing);
if stored == ongoing {
info!(
transactional_id = %self.transactional_id,
producer_id = ongoing.producer_id,
"Recovered prepared transaction matches the stored state; committing"
);
self.commit().await?;
Ok(TransactionOutcome::Committed)
} else {
info!(
transactional_id = %self.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().await?;
Ok(TransactionOutcome::Aborted)
}
}
fn adopt_prepared(&self, ongoing: PreparedTxnState) {
self.adopt_identity(ongoing.producer_id, ongoing.producer_epoch);
self.gate.set(TransactionState::Prepared);
}
pub async fn commit(&self) -> Result<()> {
if self.state() == TransactionState::Open
&& let Some(error) = self.gate.failure("commit")
{
return Err(error);
}
let entered = self
.gate
.transition(
&[
TransactionState::Open,
TransactionState::Prepared,
TransactionState::CommitUnknown,
],
TransactionState::Committing,
)
.map_err(|actual| {
KrafkaError::illegal_state(format!("cannot commit in state {actual}"))
})?;
let mut ending = Ending::new(self, TransactionState::Committing, entered);
if entered != TransactionState::CommitUnknown {
if let Some(error) = self.drain_transaction().await {
let _ = self
.gate
.transition(&[TransactionState::Committing], entered);
return Err(error);
}
if let Err(error) = self.ensure_reinitialised().await {
let _ = self
.gate
.transition(&[TransactionState::Committing], entered);
return Err(error);
}
}
let unanswered = Arc::new(AtomicBool::new(false));
ending.sending();
let result = self.end_transaction(true, &unanswered).await;
let unanswered = unanswered.load(Ordering::SeqCst);
if unanswered {
self.end_txn_unanswered.store(true, Ordering::SeqCst);
}
match result {
Ok(()) => {
ending.ended();
self.txn_partitions.write().await.clear();
self.finish_transaction(false).await;
ending.disarm();
info!("Transaction committed");
Ok(())
}
Err(error) => {
ending.disarm();
if self.state() == TransactionState::Fatal {
warn!("Transaction commit failed (fatal): {error}");
} else if unanswered || entered == TransactionState::CommitUnknown {
let _ = self.gate.transition(
&[TransactionState::Committing],
TransactionState::CommitUnknown,
);
warn!(
"Transaction commit outcome unknown ({error}); the coordinator may already \
have committed it. Retry commit() — aborting could tear the \
transaction (KAFKA-17754)."
);
} else {
let _ = self
.gate
.transition(&[TransactionState::Committing], entered);
warn!("Transaction commit failed: {error}");
}
Err(error)
}
}
}
async fn finish_transaction(&self, send_failed: bool) {
let needs_bump = !self.transaction_version().is_v2()
&& (self.end_txn_unanswered.load(Ordering::SeqCst) || send_failed);
if needs_bump {
self.reinit_required.store(true, Ordering::SeqCst);
if let Err(error) = self.ensure_reinitialised().await {
warn!(
%error,
"could not bump the producer epoch after the transaction; retrying before \
the next transaction writes"
);
}
}
self.gate.set(TransactionState::Ready);
}
pub async fn abort(&self) -> Result<()> {
let entered = self
.gate
.transition(
&[TransactionState::Open, TransactionState::Prepared],
TransactionState::Aborting,
)
.map_err(|actual| match actual {
TransactionState::CommitUnknown => KrafkaError::illegal_state(
"cannot abort: a previous commit() may already have been \
applied by the coordinator. Aborting now could be applied to a later \
transaction and tear it (KAFKA-17754). Retry commit(), or \
drop this producer and let the coordinator resolve the transaction via \
transaction.timeout.ms.",
),
actual => KrafkaError::illegal_state(format!("cannot abort in state {actual}")),
})?;
let mut ending = Ending::new(self, TransactionState::Aborting, entered);
let send_failed = self.gate.has_failed();
self.producer
.accumulator
.fail_unsent(KrafkaError::transaction_abortable(
"the transaction was aborted before this record was sent",
))
.await;
self.gate.drained().await;
let unanswered = Arc::new(AtomicBool::new(false));
ending.sending();
let result = self.end_transaction(false, &unanswered).await;
if unanswered.load(Ordering::SeqCst) {
self.end_txn_unanswered.store(true, Ordering::SeqCst);
}
match result {
Ok(()) => {
ending.ended();
self.txn_partitions.write().await.clear();
self.finish_transaction(send_failed).await;
ending.disarm();
info!("Transaction aborted");
Ok(())
}
Err(error) => {
ending.disarm();
if self.state() != TransactionState::Fatal {
let _ = self.gate.transition(&[TransactionState::Aborting], entered);
warn!("Transaction abort failed, retry abort(): {error}");
}
Err(error)
}
}
}
async fn end_transaction(&self, commit: bool, unanswered: &Arc<AtomicBool>) -> Result<()> {
let is_v2 = self.transaction_version().is_v2();
let (min_version, max_version) = if is_v2 {
(TV2_MIN_END_TXN_VERSION, versions::END_TXN_MAX)
} else {
(
versions::END_TXN_MIN,
versions::END_TXN_MAX.min(TV2_MIN_END_TXN_VERSION - 1),
)
};
let result = self
.with_coordinator("EndTxn", |attempt| {
let unanswered = Arc::clone(unanswered);
async move {
let conn = self.coordinator_connection(attempt).await?;
let (producer_id, producer_epoch) = self.checked_transactional_identity()?;
let version = conn
.negotiate_api_version(ApiKey::EndTxn, max_version, min_version)
.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::UnknownApiVersion,
format!(
"no mutually supported EndTxn API version in \
v{min_version}..=v{max_version}"
),
)
})?;
let request = if commit {
EndTxnRequest::commit(&self.transactional_id, producer_id, producer_epoch)
} else {
EndTxnRequest::abort(&self.transactional_id, producer_id, producer_epoch)
};
let pending = Unanswered::arm(unanswered);
let mut bytes = conn
.send_request(ApiKey::EndTxn, version, |buf| {
request.encode_versioned(version, buf)
})
.await?;
let response = EndTxnResponse::decode_versioned(version, &mut bytes)?;
pending.disarm();
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 && is_v2 => {
debug!(pid, epoch, "Adopting the producer epoch EndTxn bumped");
self.adopt_identity(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 partitions_for(&self, topic: &str) -> Result<Vec<crate::PartitionInfo>> {
self.producer.partitions_for(topic).await
}
pub async fn flush(&self) -> Result<()> {
self.producer.flush().await
}
#[inline]
pub fn transactional_id(&self) -> &str {
&self.transactional_id
}
#[inline]
pub fn producer_id(&self) -> i64 {
self.identity.lock().0
}
#[inline]
pub fn producer_epoch(&self) -> i16 {
self.identity.lock().1
}
pub fn prepared_transaction(&self) -> Option<PreparedTxnState> {
let ongoing = **self.ongoing_prepared_txn.load();
ongoing.is_prepared().then_some(ongoing)
}
pub async fn close(&self) -> Result<()> {
self.close_with(CloseOptions::new()).await
}
pub async fn close_with(&self, options: CloseOptions) -> Result<()> {
let Some(generation) = self.producer.barrier.begin_close() else {
return Ok(());
};
let started = tokio::time::Instant::now();
let settle = async {
match self.state() {
TransactionState::Open => {
warn!("closing a transactional producer with an open transaction; aborting it");
self.abort().await
}
TransactionState::CommitUnknown => {
warn!(
"closing after a commit whose outcome is unknown; leaving the transaction \
for the coordinator to resolve rather than aborting a possibly-committed \
transaction (KAFKA-17754)"
);
Ok(())
}
_ => Ok(()),
}
};
let settled = match options.timeout {
Some(timeout) => tokio::time::timeout(timeout, settle)
.await
.unwrap_or_else(|_| Err(KrafkaError::timeout("transactional producer close"))),
None => settle.await,
};
let remaining = options
.timeout
.map(|timeout| timeout.saturating_sub(started.elapsed()));
let closed = self.producer.finish_close(generation, remaining).await;
settled.and(closed)
}
#[inline]
pub fn is_closed(&self) -> bool {
self.producer.is_closed()
}
#[inline]
pub fn metrics(&self) -> Metrics {
self.producer.metrics()
}
pub async fn client_instance_id(
&self,
timeout: std::time::Duration,
) -> Result<Option<crate::metrics::ClientInstanceId>> {
self.producer.client_instance_id(timeout).await
}
}
struct Ending<'a> {
producer: &'a TransactionalProducer,
during: TransactionState,
on_drop: TransactionState,
sending: bool,
armed: bool,
}
impl<'a> Ending<'a> {
fn new(
producer: &'a TransactionalProducer,
during: TransactionState,
entered: TransactionState,
) -> Self {
Self {
producer,
during,
on_drop: entered,
sending: false,
armed: true,
}
}
fn sending(&mut self) {
self.sending = true;
if self.during == TransactionState::Committing {
self.on_drop = TransactionState::CommitUnknown;
}
}
fn ended(&mut self) {
self.sending = false;
self.on_drop = TransactionState::Ready;
}
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for Ending<'_> {
fn drop(&mut self) {
if !self.armed {
return;
}
if self.sending {
self.producer
.end_txn_unanswered
.store(true, Ordering::SeqCst);
}
if self.on_drop == TransactionState::Ready {
if !self.producer.transaction_version().is_v2() {
self.producer.reinit_required.store(true, Ordering::SeqCst);
}
if let Ok(mut partitions) = self.producer.txn_partitions.try_write() {
partitions.clear();
}
}
let _ = self.producer.gate.transition(&[self.during], self.on_drop);
}
}
struct Unanswered(Option<Arc<AtomicBool>>);
impl Unanswered {
fn arm(flag: Arc<AtomicBool>) -> Self {
Self(Some(flag))
}
fn disarm(mut self) {
self.0 = None;
}
}
impl Drop for Unanswered {
fn drop(&mut self) {
if let Some(flag) = self.0.take() {
flag.store(true, Ordering::SeqCst);
}
}
}
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
)
}
fn is_fencing_code(error_code: ErrorCode) -> bool {
matches!(
error_code,
ErrorCode::ProducerFenced
| ErrorCode::InvalidProducerEpoch
| ErrorCode::TransactionCoordinatorFenced
| ErrorCode::InvalidProducerIdMapping
)
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
use std::time::Duration;
#[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::illegal_state("test")
));
}
#[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 test_out_of_order_sequence_is_retriable() {
let error = KrafkaError::broker(ErrorCode::OutOfOrderSequenceNumber, "sequence mismatch");
assert!(error.is_retriable());
}
#[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::illegal_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());
}
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);
}
}
#[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"
);
}
fn test_producer(version: TransactionVersion) -> TransactionalProducer {
let kafka = crate::Kafka::detached();
let config = super::super::ProducerConfig {
max_block: Duration::from_millis(200),
..super::super::ProducerConfig::default()
};
let metrics = Arc::new(crate::metrics::ProducerRecorder::default());
let metrics_source = crate::metrics::MetricsSource::producer(&kafka, Arc::clone(&metrics));
let barrier = Arc::new(crate::barrier::InFlightBarrier::new());
let gate = Arc::new(TxnGate::new());
let interceptor: Arc<dyn crate::interceptor::ProducerInterceptor> =
Arc::new(crate::interceptor::NoOpProducerInterceptor);
let backoff = Backoff::new(Duration::from_millis(1));
let accumulator = super::super::accumulator::Accumulator::spawn(
super::super::engine::EngineConfig {
batch_size: config.batch_size,
linger: config.linger,
delivery_timeout: config.delivery_timeout,
request_timeout: kafka.request_timeout(),
max_request_size: config.max_request_size,
acks: -1,
compression: crate::protocol::Compression::None,
compression_level: None,
topic_compression: ahash::AHashMap::new(),
client_id: "test".to_string(),
transactional_id: Some("txn-test".to_string()),
backoff: backoff.clone(),
interceptor: Arc::clone(&interceptor),
mode: super::super::identity::Mode::Transactional {
tv2: version.is_v2(),
},
identity: None,
gate: Some(Arc::clone(&gate)),
},
config.buffer_memory,
config.max_block,
Arc::clone(kafka.metadata()),
Arc::clone(&metrics),
Arc::clone(&barrier),
);
let producer = Producer {
partitioning: Arc::new(super::super::partitioner::Partitioning::new(
None,
config.batch_size,
None,
)),
kafka,
config,
accumulator,
barrier,
metrics_source,
telemetry: crate::telemetry::Telemetry::disabled(),
interceptor,
};
TransactionalProducer {
producer,
transactional_id: "txn-test".to_string(),
gate,
identity: parking_lot::Mutex::new((7, 0)),
ongoing_prepared_txn: arc_swap::ArcSwap::from_pointee(PreparedTxnState::none()),
transaction_version: AtomicU8::new(version as u8),
end_txn_unanswered: AtomicBool::new(false),
reinit_required: AtomicBool::new(false),
reinit_lock: tokio::sync::Mutex::new(()),
coordinator_id: RwLock::new(None),
txn_partitions: Arc::new(RwLock::new(TransactionPartitions::default())),
backoff,
}
}
#[tokio::test]
async fn the_two_phase_entry_points_require_the_setting() {
let producer = test_producer(TransactionVersion::V2);
producer.gate.set(TransactionState::Open);
let err = producer.prepare().await.unwrap_err();
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.gate.set(TransactionState::Ready);
let err = producer
.complete(PreparedTxnState {
producer_id: 1,
producer_epoch: 0,
})
.await
.unwrap_err();
assert!(
err.to_string().contains("prepared_transaction"),
"got: {err}"
);
}
#[tokio::test]
async fn an_unknown_commit_refuses_abort() {
let producer = test_producer(TransactionVersion::V1);
producer.gate.set(TransactionState::CommitUnknown);
let err = producer.abort().await.unwrap_err();
assert!(err.to_string().contains("KAFKA-17754"), "got: {err}");
assert_eq!(producer.state(), TransactionState::CommitUnknown);
}
#[tokio::test]
async fn a_failed_send_refuses_commit() {
let producer = test_producer(TransactionVersion::V1);
producer.gate.set(TransactionState::Ready);
producer.begin().unwrap();
producer
.gate
.fail(KrafkaError::broker(ErrorCode::InvalidRecord, "r2"));
let err = producer.commit().await.unwrap_err();
assert!(err.requires_abort(), "{err}");
assert_eq!(producer.state(), TransactionState::Open);
}
#[test]
fn test_transaction_state_display() {
assert_eq!(TransactionState::CommitUnknown.to_string(), "CommitUnknown");
assert_eq!(TransactionState::Open.to_string(), "Open");
}
#[test]
fn tv1_end_txn_stays_below_the_tv2_version() {
assert_eq!(TV2_MIN_END_TXN_VERSION, 5);
assert!(versions::END_TXN_MAX.min(TV2_MIN_END_TXN_VERSION - 1) < TV2_MIN_END_TXN_VERSION);
}
}