use std::collections::{BTreeMap, BTreeSet};
use std::fmt;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex as StdMutex};
use std::time::Duration;
use kafrust_protocol::api::consumer_group_heartbeat::{
ConsumerGroupHeartbeatResponseV0, ConsumerGroupHeartbeatTopicPartitions,
};
use kafrust_protocol::api::find_coordinator::FindCoordinatorResponseV1;
use kafrust_protocol::api::join_group::JoinGroupMember;
use kafrust_protocol::api::leave_group::{LeaveGroupMemberIdentity, LeaveGroupResponseV3};
use kafrust_protocol::api::list_offsets::{
ListOffsetsPartitionV1, ListOffsetsResponseV1, ListOffsetsTopicV1,
};
use kafrust_protocol::api::metadata::{
MetadataRequestTopicV12, MetadataResponseV1, MetadataResponseV12,
};
use kafrust_protocol::api::offset_commit::{
OffsetCommitPartition, OffsetCommitPartitionV7, OffsetCommitPartitionV9, OffsetCommitTopic,
OffsetCommitTopicResponse, OffsetCommitTopicV7, OffsetCommitTopicV9,
};
use kafrust_protocol::api::offset_fetch::{
OffsetFetchPartitionResponse, OffsetFetchTopic, OffsetFetchTopicResponse, OffsetFetchTopicV9,
};
use kafrust_protocol::api::sync_group::SyncGroupAssignment;
use kafrust_protocol::codec::{Decoder, Encoder};
use kafrust_protocol::consumer_group::{
ConsumerProtocolAssignmentV0, ConsumerProtocolSubscriptionV0, ConsumerProtocolSubscriptionV1,
ConsumerProtocolTopicAssignment,
};
use crate::client::Client;
use crate::config::{ClientConfig, OAuthBearerTokenProvider, SecurityProtocol};
pub use crate::consumer::OffsetResetPolicy;
use crate::consumer::{
Consumer, ConsumerAssignment, ConsumerConfig, ConsumerPartitionQueue, ConsumerRecord,
IsolationLevel, LeaderEpochOffset, PartitionWatermarks,
};
use crate::error::{BrokerErrorKind, Error, Result};
use crate::metrics::ClientMetrics;
use regex::Regex;
use tokio::sync::{mpsc, oneshot, Mutex, Notify};
use tokio::task::JoinHandle;
use tokio::time::{self, MissedTickBehavior};
use tracing::{debug, Instrument};
const PROTOCOL_TYPE: &str = "consumer";
const RANGE_PROTOCOL: &str = "range";
const ROUND_ROBIN_PROTOCOL: &str = "roundrobin";
const STICKY_PROTOCOL: &str = "sticky";
const COOPERATIVE_STICKY_PROTOCOL: &str = "cooperative-sticky";
const GROUP_JOIN_RETRY_BACKOFF: Duration = Duration::from_millis(50);
const GROUP_JOIN_MAX_RETRY_BACKOFF: Duration = Duration::from_secs(1);
const DEFAULT_GROUP_MAX_RETRIES: u32 = 5;
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum ConsumerGroupProtocol {
#[default]
Classic,
Consumer,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RebalancePhase {
Before,
After,
}
#[derive(Debug, Clone, Copy)]
pub struct RebalanceEvent<'a> {
phase: RebalancePhase,
group_id: &'a str,
member_id: &'a str,
generation_id: i32,
protocol: ConsumerGroupProtocol,
assignments: &'a [ConsumerAssignment],
}
impl RebalanceEvent<'_> {
pub fn phase(&self) -> RebalancePhase {
self.phase
}
pub fn group_id(&self) -> &str {
self.group_id
}
pub fn member_id(&self) -> &str {
self.member_id
}
pub fn generation_id(&self) -> i32 {
self.generation_id
}
pub fn protocol(&self) -> ConsumerGroupProtocol {
self.protocol
}
pub fn assignments(&self) -> &[ConsumerAssignment] {
self.assignments
}
}
pub trait RebalanceListener: Send + Sync {
fn on_rebalance(&self, event: RebalanceEvent<'_>);
}
impl<F> RebalanceListener for F
where
F: for<'a> Fn(RebalanceEvent<'a>) + Send + Sync,
{
fn on_rebalance(&self, event: RebalanceEvent<'_>) {
self(event);
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum ConsumerGroupAssignmentStrategy {
#[default]
Range,
RoundRobin,
Sticky,
CooperativeSticky,
}
impl ConsumerGroupAssignmentStrategy {
fn protocol_name(self) -> &'static str {
match self {
Self::Range => RANGE_PROTOCOL,
Self::RoundRobin => ROUND_ROBIN_PROTOCOL,
Self::Sticky => STICKY_PROTOCOL,
Self::CooperativeSticky => COOPERATIVE_STICKY_PROTOCOL,
}
}
}
struct JoinedGroup {
error_code: i16,
generation_id: i32,
protocol_name: String,
leader: String,
member_id: String,
members: Vec<JoinGroupMember>,
}
#[derive(Clone)]
pub struct ConsumerGroupConfig {
client: ClientConfig,
group_id: String,
group_instance_id: Option<String>,
topics: Vec<String>,
topic_pattern: Option<String>,
session_timeout_ms: i32,
rebalance_timeout_ms: i32,
retention_time_ms: i64,
auto_commit: bool,
auto_commit_interval: Duration,
offset_reset_policy: OffsetResetPolicy,
max_wait_ms: i32,
min_bytes: i32,
max_partition_bytes: i32,
max_retries: u32,
max_poll_records: usize,
partition_queue_capacity: usize,
isolation_level: IsolationLevel,
assignment_strategy: ConsumerGroupAssignmentStrategy,
group_protocol: ConsumerGroupProtocol,
server_assignor: Option<String>,
rebalance_listener: Option<Arc<dyn RebalanceListener>>,
}
#[derive(Debug, Clone)]
struct ConsumerProtocolHeartbeatState {
member_id: String,
member_epoch: i32,
owned_partitions: Option<Vec<ConsumerGroupHeartbeatTopicPartitions>>,
assignment_version: u64,
heartbeat_interval: Duration,
}
#[derive(Debug, Clone)]
struct ConsumerProtocolHeartbeatConfig {
group_instance_id: Option<String>,
topics: Vec<String>,
server_assignor: Option<String>,
rebalance_timeout_ms: i32,
}
impl ConsumerGroupConfig {
pub fn new(
bootstrap_servers: impl IntoIterator<Item = impl Into<String>>,
group_id: impl Into<String>,
) -> Self {
Self {
client: ClientConfig::new(bootstrap_servers),
group_id: group_id.into(),
group_instance_id: None,
topics: Vec::new(),
topic_pattern: None,
session_timeout_ms: 10_000,
rebalance_timeout_ms: 30_000,
retention_time_ms: 86_400_000,
auto_commit: false,
auto_commit_interval: Duration::from_millis(5_000),
offset_reset_policy: OffsetResetPolicy::Offset(0),
max_wait_ms: 500,
min_bytes: 1,
max_partition_bytes: 1_048_576,
max_retries: DEFAULT_GROUP_MAX_RETRIES,
max_poll_records: 500,
partition_queue_capacity: 1024,
isolation_level: IsolationLevel::ReadUncommitted,
assignment_strategy: ConsumerGroupAssignmentStrategy::Range,
group_protocol: ConsumerGroupProtocol::Classic,
server_assignor: None,
rebalance_listener: None,
}
}
pub fn client_id(mut self, client_id: impl Into<String>) -> Self {
self.client = self.client.client_id(client_id);
self
}
pub fn client_rack(mut self, client_rack: impl Into<String>) -> Self {
self.client = self.client.client_rack(client_rack);
self
}
pub fn request_timeout_ms(mut self, request_timeout_ms: u64) -> Self {
self.client = self.client.request_timeout_ms(request_timeout_ms);
self
}
pub fn max_response_bytes(mut self, max_response_bytes: usize) -> Self {
self.client = self.client.max_response_bytes(max_response_bytes);
self
}
pub fn max_decode_array_elements(mut self, max: usize) -> Self {
self.client = self.client.max_decode_array_elements(max);
self
}
pub fn max_decompressed_record_bytes(mut self, max: usize) -> Self {
self.client = self.client.max_decompressed_record_bytes(max);
self
}
pub fn metrics(mut self, metrics: ClientMetrics) -> Self {
self.client = self.client.metrics(metrics);
self
}
pub fn security_protocol(mut self, security_protocol: SecurityProtocol) -> Self {
self.client = self.client.security_protocol(security_protocol);
self
}
pub fn tls_server_name(mut self, server_name: impl Into<String>) -> Self {
self.client = self.client.tls_server_name(server_name);
self
}
pub fn tls_root_certificate_der(mut self, certificate: impl Into<Vec<u8>>) -> Self {
self.client = self.client.tls_root_certificate_der(certificate);
self
}
pub fn sasl_plain(mut self, username: impl Into<String>, password: impl Into<String>) -> Self {
self.client = self.client.sasl_plain(username, password);
self
}
pub fn sasl_scram_sha_256(
mut self,
username: impl Into<String>,
password: impl Into<String>,
) -> Self {
self.client = self.client.sasl_scram_sha_256(username, password);
self
}
pub fn sasl_scram_sha_512(
mut self,
username: impl Into<String>,
password: impl Into<String>,
) -> Self {
self.client = self.client.sasl_scram_sha_512(username, password);
self
}
pub fn sasl_oauthbearer(mut self, token: impl Into<String>) -> Self {
self.client = self.client.sasl_oauthbearer(token);
self
}
pub fn sasl_oauthbearer_with_username(
mut self,
username: impl Into<String>,
token: impl Into<String>,
) -> Self {
self.client = self.client.sasl_oauthbearer_with_username(username, token);
self
}
pub fn sasl_oauthbearer_provider<P>(mut self, provider: P) -> Self
where
P: OAuthBearerTokenProvider + 'static,
{
self.client = self.client.sasl_oauthbearer_provider(provider);
self
}
pub fn sasl_oauthbearer_with_username_and_provider<P>(
mut self,
username: impl Into<String>,
provider: P,
) -> Self
where
P: OAuthBearerTokenProvider + 'static,
{
self.client = self
.client
.sasl_oauthbearer_with_username_and_provider(username, provider);
self
}
pub fn subscribe(mut self, topic: impl Into<String>) -> Self {
self.topic_pattern = None;
self.topics.push(topic.into());
self
}
pub fn subscribe_pattern(mut self, pattern: impl Into<String>) -> Self {
self.topics.clear();
self.topic_pattern = Some(pattern.into());
self
}
pub fn group_instance_id(mut self, group_instance_id: impl Into<String>) -> Self {
self.group_instance_id = Some(group_instance_id.into());
self
}
pub fn group_instance_id_ref(&self) -> Option<&str> {
self.group_instance_id.as_deref()
}
pub fn session_timeout_ms(mut self, session_timeout_ms: i32) -> Self {
self.session_timeout_ms = session_timeout_ms;
self
}
pub fn rebalance_timeout_ms(mut self, rebalance_timeout_ms: i32) -> Self {
self.rebalance_timeout_ms = rebalance_timeout_ms;
self
}
pub fn retention_time_ms(mut self, retention_time_ms: i64) -> Self {
self.retention_time_ms = retention_time_ms;
self
}
pub fn enable_auto_commit(mut self, enabled: bool) -> Self {
self.auto_commit = enabled;
self
}
pub fn auto_commit_interval_ms(mut self, interval_ms: u64) -> Self {
self.auto_commit_interval = Duration::from_millis(interval_ms);
self
}
pub fn auto_commit_enabled(&self) -> bool {
self.auto_commit
}
pub fn auto_commit_interval(&self) -> Duration {
self.auto_commit_interval
}
pub fn start_offset(mut self, start_offset: i64) -> Self {
self.offset_reset_policy = OffsetResetPolicy::Offset(start_offset);
self
}
pub fn offset_reset_policy(mut self, offset_reset_policy: OffsetResetPolicy) -> Self {
self.offset_reset_policy = offset_reset_policy;
self
}
pub fn max_wait_ms(mut self, max_wait_ms: i32) -> Self {
self.max_wait_ms = max_wait_ms;
self
}
pub fn min_bytes(mut self, min_bytes: i32) -> Self {
self.min_bytes = min_bytes;
self
}
pub fn max_partition_bytes(mut self, max_partition_bytes: i32) -> Self {
self.max_partition_bytes = max_partition_bytes;
self
}
pub fn max_retries(mut self, max_retries: u32) -> Self {
self.max_retries = max_retries;
self
}
pub fn max_poll_records(mut self, max_poll_records: usize) -> Self {
self.max_poll_records = max_poll_records;
self
}
pub fn partition_queue_capacity(mut self, partition_queue_capacity: usize) -> Self {
self.partition_queue_capacity = partition_queue_capacity.max(1);
self
}
pub fn partition_queue_capacity_ref(&self) -> usize {
self.partition_queue_capacity
}
pub fn isolation_level(mut self, isolation_level: IsolationLevel) -> Self {
self.isolation_level = isolation_level;
self
}
pub fn assignment_strategy(
mut self,
assignment_strategy: ConsumerGroupAssignmentStrategy,
) -> Self {
self.assignment_strategy = assignment_strategy;
self
}
pub fn assignment_strategy_ref(&self) -> ConsumerGroupAssignmentStrategy {
self.assignment_strategy
}
pub fn group_protocol(mut self, group_protocol: ConsumerGroupProtocol) -> Self {
self.group_protocol = group_protocol;
self
}
pub fn group_protocol_ref(&self) -> ConsumerGroupProtocol {
self.group_protocol
}
pub fn server_assignor(mut self, server_assignor: impl Into<String>) -> Self {
self.server_assignor = Some(server_assignor.into());
self
}
pub fn server_assignor_ref(&self) -> Option<&str> {
self.server_assignor.as_deref()
}
pub fn rebalance_listener<F>(mut self, listener: F) -> Self
where
F: for<'a> Fn(RebalanceEvent<'a>) + Send + Sync + 'static,
{
self.rebalance_listener = Some(Arc::new(listener));
self
}
pub fn rebalance_listener_handler<L>(mut self, listener: L) -> Self
where
L: RebalanceListener + 'static,
{
self.rebalance_listener = Some(Arc::new(listener));
self
}
pub fn has_rebalance_listener(&self) -> bool {
self.rebalance_listener.is_some()
}
pub fn offset_reset_policy_ref(&self) -> OffsetResetPolicy {
self.offset_reset_policy
}
pub fn isolation_level_ref(&self) -> IsolationLevel {
self.isolation_level
}
pub fn group_id(&self) -> &str {
&self.group_id
}
pub fn topics(&self) -> &[String] {
&self.topics
}
pub fn topic_pattern_ref(&self) -> Option<&str> {
self.topic_pattern.as_deref()
}
fn has_subscription(&self) -> bool {
!self.topics.is_empty() || self.topic_pattern.is_some()
}
async fn resolve_subscription_topics(&self, bootstrap: &mut Client) -> Result<Vec<String>> {
if !self.topics.is_empty() {
return Ok(self.topics.clone());
}
let pattern = self
.topic_pattern
.as_deref()
.ok_or(Error::Unsupported("consumer group without subscriptions"))?;
let metadata = bootstrap.metadata(None).await?;
resolve_topic_pattern(pattern, &metadata)
}
fn consumer_config(&self) -> ConsumerConfig {
ConsumerConfig::from_client_config(self.client.clone())
.max_wait_ms(self.max_wait_ms)
.min_bytes(self.min_bytes)
.max_partition_bytes(self.max_partition_bytes)
.max_retries(self.max_retries)
.max_poll_records(self.max_poll_records)
.partition_queue_capacity(self.partition_queue_capacity)
.isolation_level(self.isolation_level)
.offset_reset_policy(self.offset_reset_policy)
}
#[tracing::instrument(
level = "debug",
name = "kafka.consumer_group.join",
skip_all,
fields(group_id = self.group_id.as_str(), topic_count = self.topics.len()),
err
)]
pub async fn join(self) -> Result<ConsumerGroup> {
self.validate()?;
let mut group = if self.group_protocol == ConsumerGroupProtocol::Consumer {
if self.assignment_strategy != ConsumerGroupAssignmentStrategy::Range {
return Err(Error::Unsupported(
"KIP-848 groups use a broker-side assignor; configure ConsumerGroupAssignmentStrategy::Range or select the classic protocol",
));
}
self.join_consumer(None, 0, None).await?
} else {
self.join_with_owned_partitions(Vec::new()).await?
};
group.start_auto_commit_if_enabled().await?;
Ok(group)
}
async fn join_with_owned_partitions(
self,
owned_partitions: Vec<ConsumerProtocolTopicAssignment>,
) -> Result<ConsumerGroup> {
self.join_with_member_id(owned_partitions, None, None).await
}
async fn join_with_member_id(
self,
owned_partitions: Vec<ConsumerProtocolTopicAssignment>,
member_id: Option<String>,
previous_generation: Option<i32>,
) -> Result<ConsumerGroup> {
if !self.has_subscription() {
return Err(Error::Unsupported("consumer group without subscriptions"));
}
if self.group_instance_id.as_deref() == Some("") {
return Err(Error::InvalidGroupInstanceId);
}
let mut member_id = member_id;
let mut retry_attempt = 0;
loop {
debug!(
group_id = self.group_id.as_str(),
topic_count = self.topics.len(),
retry_attempt,
"joining kafka consumer group"
);
let config = self.clone();
let mut bootstrap = self.client.clone().connect().await?;
let topics = self.resolve_subscription_topics(&mut bootstrap).await?;
let coordinator = find_group_coordinator_with_retry(
&mut bootstrap,
&self.client,
&self.group_id,
self.max_retries,
)
.await?;
let coordinator_addr = coordinator_addr(&coordinator);
let mut coordinator_client = self.client.connect_broker(coordinator_addr).await?;
let is_rejoin = member_id.is_some();
let subscription = match self.assignment_strategy {
ConsumerGroupAssignmentStrategy::Sticky => ConsumerProtocolSubscriptionV0 {
topics: topics.clone(),
user_data: if is_rejoin {
Some(encode_sticky_user_data(
&owned_partitions,
previous_generation,
)?)
} else {
None
},
}
.encode()?,
ConsumerGroupAssignmentStrategy::CooperativeSticky => {
ConsumerProtocolSubscriptionV1 {
topics: topics.clone(),
user_data: None,
owned_partitions: owned_partitions.clone(),
}
.encode()?
}
ConsumerGroupAssignmentStrategy::Range
| ConsumerGroupAssignmentStrategy::RoundRobin => ConsumerProtocolSubscriptionV0 {
topics: topics.clone(),
user_data: None,
}
.encode()?,
};
let protocols = vec![kafrust_protocol::api::join_group::JoinGroupProtocol {
name: self.assignment_strategy.protocol_name().to_owned(),
metadata: subscription,
}];
let requested_member_id = member_id.take().unwrap_or_default();
let joined = if let Some(group_instance_id) = &self.group_instance_id {
let response = coordinator_client
.join_group_v5(
self.group_id.clone(),
self.session_timeout_ms,
self.rebalance_timeout_ms,
requested_member_id.clone(),
Some(group_instance_id.clone()),
PROTOCOL_TYPE,
protocols,
)
.await?;
JoinedGroup {
error_code: response.error_code,
generation_id: response.generation_id,
protocol_name: response.protocol_name,
leader: response.leader,
member_id: response.member_id,
members: response
.members
.into_iter()
.map(|member| JoinGroupMember {
member_id: member.member_id,
metadata: member.metadata,
})
.collect(),
}
} else {
let response = coordinator_client
.join_group_v2(
self.group_id.clone(),
self.session_timeout_ms,
self.rebalance_timeout_ms,
requested_member_id.clone(),
PROTOCOL_TYPE,
protocols,
)
.await?;
JoinedGroup {
error_code: response.error_code,
generation_id: response.generation_id,
protocol_name: response.protocol_name,
leader: response.leader,
member_id: response.member_id,
members: response.members,
}
};
if joined.error_code != 0 {
let error = self
.client
.broker_error(joined.error_code, format!("join group {}", self.group_id));
if retry_attempt < self.max_retries && should_rejoin_group(&error) {
retry_attempt += 1;
self.client.record_retry();
member_id = member_id_after_join_error(&error, requested_member_id);
debug!(
group_id = self.group_id.as_str(),
retry_attempt,
error = %error,
"retrying kafka consumer group join after transient join error"
);
time::sleep(group_retry_backoff(retry_attempt)).await;
continue;
}
return Err(error);
}
let assignments = if joined.member_id == joined.leader {
let metadata = bootstrap.metadata(Some(topics.clone())).await?;
assignments_for_strategy(&joined.protocol_name, &joined.members, &metadata)?
} else {
Vec::new()
};
let synced = if let Some(group_instance_id) = &self.group_instance_id {
coordinator_client
.sync_group_v3(
self.group_id.clone(),
joined.generation_id,
joined.member_id.clone(),
Some(group_instance_id.clone()),
assignments,
)
.await?
} else {
coordinator_client
.sync_group_v2(
self.group_id.clone(),
joined.generation_id,
joined.member_id.clone(),
assignments,
)
.await?
};
if synced.error_code != 0 {
let error = self
.client
.broker_error(synced.error_code, format!("sync group {}", self.group_id));
if retry_attempt < self.max_retries && should_rejoin_group(&error) {
retry_attempt += 1;
self.client.record_retry();
member_id = Some(joined.member_id.clone());
debug!(
group_id = self.group_id.as_str(),
retry_attempt,
error = %error,
"retrying kafka consumer group join after transient sync error"
);
time::sleep(group_retry_backoff(retry_attempt)).await;
continue;
}
return Err(error);
}
let assignment = ConsumerProtocolAssignmentV0::decode(&synced.assignment)?;
let cooperative_rejoin_required = self.assignment_strategy
== ConsumerGroupAssignmentStrategy::CooperativeSticky
&& cooperative_rejoin_required_for_assignment(
&joined.members,
&owned_partitions,
&assignment,
is_rejoin,
)?;
let consumer_assignments = assignments_from_protocol(
&mut coordinator_client,
&mut bootstrap,
&self.client,
&self.group_id,
self.offset_reset_policy,
&assignment,
None,
)
.await?;
let consumer_config = self.consumer_config();
let consumer_client = self.client.clone().connect().await?;
debug!(
group_id = self.group_id.as_str(),
member_id = joined.member_id.as_str(),
generation_id = joined.generation_id,
assignment_count = consumer_assignments.len(),
"joined kafka consumer group"
);
let group = ConsumerGroup {
config,
group_id: self.group_id,
generation_id: joined.generation_id,
member_id: joined.member_id,
retention_time_ms: self.retention_time_ms,
coordinator: coordinator_client,
cooperative_rejoin_required,
protocol: ConsumerGroupProtocol::Classic,
resolved_topics: topics,
consumer_topic_ids: BTreeMap::new(),
consumer_session: None,
consumer_heartbeat_state: None,
consumer_heartbeat_assignment_version: 0,
consumer_owned_partitions: None,
consumer: Consumer::from_assignments(
consumer_client,
consumer_config,
consumer_assignments,
),
pending_commit_offsets: BTreeMap::new(),
commit_worker: None,
auto_commit_worker: None,
};
group.notify_rebalance(RebalancePhase::After, group.consumer.assignments());
return Ok(group);
}
}
pub fn validate(&self) -> Result<()> {
self.client.validate()?;
if self.group_id.trim().is_empty() {
return Err(Error::InvalidConfiguration {
field: "group_id",
reason: "must not be empty",
});
}
if self.topics.iter().any(|topic| topic.trim().is_empty()) {
return Err(Error::InvalidConfiguration {
field: "topics",
reason: "entries must not be empty",
});
}
if let Some(pattern) = self.topic_pattern.as_deref() {
if pattern.is_empty() {
return Err(Error::InvalidConfiguration {
field: "topic_pattern",
reason: "must not be empty",
});
}
Regex::new(pattern).map_err(|error| Error::InvalidTopicPattern {
pattern: pattern.to_owned(),
reason: error.to_string(),
})?;
}
if !self.has_subscription() {
return Err(Error::InvalidConfiguration {
field: "subscription",
reason: "at least one topic or topic pattern is required",
});
}
if self.group_instance_id.as_deref() == Some("") {
return Err(Error::InvalidGroupInstanceId);
}
if self.session_timeout_ms <= 0 {
return Err(Error::InvalidConfiguration {
field: "session_timeout_ms",
reason: "must be greater than zero",
});
}
if self.rebalance_timeout_ms <= 0 {
return Err(Error::InvalidConfiguration {
field: "rebalance_timeout_ms",
reason: "must be greater than zero",
});
}
if self.max_wait_ms < 0 {
return Err(Error::InvalidConfiguration {
field: "max_wait_ms",
reason: "must not be negative",
});
}
if self.min_bytes < 0 {
return Err(Error::InvalidConfiguration {
field: "min_bytes",
reason: "must not be negative",
});
}
if self.max_partition_bytes <= 0 {
return Err(Error::InvalidConfiguration {
field: "max_partition_bytes",
reason: "must be greater than zero",
});
}
if self.max_poll_records == 0 {
return Err(Error::InvalidConfiguration {
field: "max_poll_records",
reason: "must be greater than zero",
});
}
if self.auto_commit {
validate_commit_worker_interval(self.auto_commit_interval)?;
}
Ok(())
}
async fn join_consumer(
self,
member_id: Option<String>,
member_epoch: i32,
owned_partitions: Option<Vec<ConsumerGroupHeartbeatTopicPartitions>>,
) -> Result<ConsumerGroup> {
let mut retry_attempt = 0;
loop {
match self
.clone()
.join_consumer_once(member_id.clone(), member_epoch, owned_partitions.clone())
.await
{
Ok(group) => return Ok(group),
Err(error)
if retry_attempt < self.max_retries
&& should_retry_consumer_join_transport(&error) =>
{
retry_attempt += 1;
self.client.record_retry();
debug!(
group_id = self.group_id.as_str(),
retry_attempt,
error = %error,
"retrying KIP-848 consumer group join after transport failure"
);
time::sleep(group_retry_backoff(retry_attempt)).await;
}
Err(error) => return Err(error),
}
}
}
async fn join_consumer_once(
self,
member_id: Option<String>,
member_epoch: i32,
owned_partitions: Option<Vec<ConsumerGroupHeartbeatTopicPartitions>>,
) -> Result<ConsumerGroup> {
if !self.has_subscription() {
return Err(Error::Unsupported("consumer group without subscriptions"));
}
if self.group_instance_id.as_deref() == Some("") {
return Err(Error::InvalidGroupInstanceId);
}
let config = self.clone();
let mut bootstrap = self.client.clone().connect().await?;
let topics = self.resolve_subscription_topics(&mut bootstrap).await?;
let topic_partitions = match owned_partitions {
Some(owned_partitions) => Some(owned_partitions),
None if member_epoch == 0 => Some(Vec::new()),
None => None,
};
let mut coordinator = find_group_coordinator_with_retry(
&mut bootstrap,
&self.client,
&self.group_id,
self.max_retries,
)
.await?;
let metadata = bootstrap
.metadata_v12(Some(
topics
.iter()
.map(|name| MetadataRequestTopicV12 {
topic_id: [0; 16],
name: Some(name.clone()),
})
.collect(),
))
.await?;
let has_assignable_partitions = metadata
.topics
.iter()
.any(|topic| topic.error_code == 0 && !topic.partitions.is_empty());
let topic_ids = topic_ids_for_names(&metadata, &topics).map_err(|error| match error {
Error::Broker { code, context } => self.client.broker_error(code, context),
error => error,
})?;
let requested_member_id = member_id.clone().unwrap_or_default();
let mut heartbeat_retry_attempt = 0;
let (mut coordinator_client, response) = loop {
let mut coordinator_client = self
.client
.connect_broker(coordinator_addr(&coordinator))
.await?;
let response = coordinator_client
.consumer_group_heartbeat_v0(
self.group_id.clone(),
requested_member_id.clone(),
member_epoch,
self.group_instance_id.clone(),
None,
self.rebalance_timeout_ms,
Some(topics.clone()),
self.server_assignor.clone(),
topic_partitions.clone(),
)
.await?;
if response.error_code == 0 {
break (coordinator_client, response);
}
let error = consumer_group_heartbeat_error(&self.client, &self.group_id, &response);
if heartbeat_retry_attempt >= self.max_retries || !should_rejoin_group(&error) {
return Err(error);
}
heartbeat_retry_attempt += 1;
self.client.record_retry();
debug!(
group_id = self.group_id.as_str(),
retry_attempt = heartbeat_retry_attempt,
error = %error,
"retrying kafka KIP-848 consumer group heartbeat"
);
time::sleep(group_retry_backoff(heartbeat_retry_attempt)).await;
coordinator = find_group_coordinator_with_retry(
&mut bootstrap,
&self.client,
&self.group_id,
self.max_retries,
)
.await?;
};
let member_id = response.member_id.or(member_id).ok_or(Error::Unsupported(
"consumer group heartbeat omitted member ID",
))?;
let assignment = response.assignment.clone().unwrap_or_default();
let protocol_assignment = consumer_protocol_assignment(&assignment, &topic_ids)?;
let consumer_assignments = assignments_from_protocol(
&mut coordinator_client,
&mut bootstrap,
&self.client,
&self.group_id,
self.offset_reset_policy,
&protocol_assignment,
Some((member_id.as_str(), response.member_epoch)),
)
.await?;
let consumer_config = self.consumer_config();
let consumer_client = self.client.clone().connect().await?;
debug!(
group_id = self.group_id.as_str(),
member_id = member_id.as_str(),
member_epoch = response.member_epoch,
assignment_count = consumer_assignments.len(),
"joined kafka KIP-848 consumer group"
);
let mut group = ConsumerGroup {
config,
group_id: self.group_id,
generation_id: response.member_epoch,
member_id,
retention_time_ms: self.retention_time_ms,
coordinator: coordinator_client,
cooperative_rejoin_required: false,
protocol: ConsumerGroupProtocol::Consumer,
resolved_topics: topics,
consumer_topic_ids: topic_ids,
consumer_session: Some(Arc::new(())),
consumer_heartbeat_state: None,
consumer_heartbeat_assignment_version: 0,
consumer_owned_partitions: response.assignment,
consumer: Consumer::from_assignments(
consumer_client,
consumer_config,
consumer_assignments,
),
pending_commit_offsets: BTreeMap::new(),
commit_worker: None,
auto_commit_worker: None,
};
if !consumer_assignment_ready(
group.consumer_owned_partitions.as_deref(),
has_assignable_partitions,
) {
group
.wait_for_consumer_assignment(has_assignable_partitions)
.await?;
}
group.notify_rebalance(RebalancePhase::After, group.consumer.assignments());
Ok(group)
}
}
fn compile_topic_pattern(pattern: &str) -> Result<Regex> {
Regex::new(pattern).map_err(|error| Error::InvalidTopicPattern {
pattern: pattern.to_owned(),
reason: error.to_string(),
})
}
fn resolve_topic_pattern(pattern: &str, metadata: &MetadataResponseV1) -> Result<Vec<String>> {
let regex = compile_topic_pattern(pattern)?;
let mut topics = metadata
.topics
.iter()
.filter(|topic| topic.error_code == 0 && regex.is_match(&topic.name))
.map(|topic| topic.name.clone())
.collect::<Vec<_>>();
topics.sort();
topics.dedup();
if topics.is_empty() {
return Err(Error::Unsupported(
"consumer group topic pattern matched no topics",
));
}
Ok(topics)
}
fn queue_commit_offset(
pending: &mut BTreeMap<(String, i32), i64>,
topic: &str,
partition: i32,
offset: i64,
) -> Result<()> {
let next_offset = offset
.checked_add(1)
.ok_or(Error::Unsupported("consumer record offset overflow"))?;
let queued = pending
.entry((topic.to_owned(), partition))
.or_insert(next_offset);
*queued = (*queued).max(next_offset);
Ok(())
}
fn retain_pending_commit_offsets_for(
assignments: &[ConsumerAssignment],
pending: &mut BTreeMap<(String, i32), i64>,
) {
let assigned = assignments
.iter()
.map(|assignment| (assignment.topic().to_owned(), assignment.partition()))
.collect::<BTreeSet<_>>();
pending.retain(|key, _| assigned.contains(key));
}
fn pending_commit_assignments(
assignments: &[ConsumerAssignment],
pending: &BTreeMap<(String, i32), i64>,
) -> Vec<ConsumerAssignment> {
assignments
.iter()
.filter_map(|assignment| {
pending
.get(&(assignment.topic().to_owned(), assignment.partition()))
.map(|offset| {
ConsumerAssignment::new(
assignment.topic().to_owned(),
assignment.partition(),
*offset,
)
})
})
.collect()
}
impl fmt::Debug for ConsumerGroupConfig {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("ConsumerGroupConfig")
.field("client", &self.client)
.field("group_id", &self.group_id)
.field("group_instance_id", &self.group_instance_id)
.field("topics", &self.topics)
.field("topic_pattern", &self.topic_pattern)
.field("session_timeout_ms", &self.session_timeout_ms)
.field("rebalance_timeout_ms", &self.rebalance_timeout_ms)
.field("retention_time_ms", &self.retention_time_ms)
.field("auto_commit", &self.auto_commit)
.field("auto_commit_interval", &self.auto_commit_interval)
.field("offset_reset_policy", &self.offset_reset_policy)
.field("max_wait_ms", &self.max_wait_ms)
.field("min_bytes", &self.min_bytes)
.field("max_partition_bytes", &self.max_partition_bytes)
.field("max_retries", &self.max_retries)
.field("max_poll_records", &self.max_poll_records)
.field("partition_queue_capacity", &self.partition_queue_capacity)
.field("isolation_level", &self.isolation_level)
.field("assignment_strategy", &self.assignment_strategy)
.field("group_protocol", &self.group_protocol)
.field("server_assignor", &self.server_assignor)
.field("has_rebalance_listener", &self.has_rebalance_listener())
.finish()
}
}
impl PartialEq for ConsumerGroupConfig {
fn eq(&self, other: &Self) -> bool {
self.client == other.client
&& self.group_id == other.group_id
&& self.group_instance_id == other.group_instance_id
&& self.topics == other.topics
&& self.topic_pattern == other.topic_pattern
&& self.session_timeout_ms == other.session_timeout_ms
&& self.rebalance_timeout_ms == other.rebalance_timeout_ms
&& self.retention_time_ms == other.retention_time_ms
&& self.auto_commit == other.auto_commit
&& self.auto_commit_interval == other.auto_commit_interval
&& self.offset_reset_policy == other.offset_reset_policy
&& self.max_wait_ms == other.max_wait_ms
&& self.min_bytes == other.min_bytes
&& self.max_partition_bytes == other.max_partition_bytes
&& self.max_retries == other.max_retries
&& self.max_poll_records == other.max_poll_records
&& self.partition_queue_capacity == other.partition_queue_capacity
&& self.isolation_level == other.isolation_level
&& self.assignment_strategy == other.assignment_strategy
&& self.group_protocol == other.group_protocol
&& self.server_assignor == other.server_assignor
&& match (&self.rebalance_listener, &other.rebalance_listener) {
(None, None) => true,
(Some(left), Some(right)) => Arc::ptr_eq(left, right),
_ => false,
}
}
}
impl Eq for ConsumerGroupConfig {}
#[derive(Debug)]
pub struct ConsumerGroup {
config: ConsumerGroupConfig,
group_id: String,
generation_id: i32,
member_id: String,
retention_time_ms: i64,
coordinator: Client,
cooperative_rejoin_required: bool,
protocol: ConsumerGroupProtocol,
resolved_topics: Vec<String>,
consumer_topic_ids: BTreeMap<String, [u8; 16]>,
consumer_session: Option<Arc<()>>,
consumer_heartbeat_state: Option<Arc<Mutex<ConsumerProtocolHeartbeatState>>>,
consumer_heartbeat_assignment_version: u64,
consumer_owned_partitions: Option<Vec<ConsumerGroupHeartbeatTopicPartitions>>,
consumer: Consumer,
pending_commit_offsets: BTreeMap<(String, i32), i64>,
commit_worker: Option<CommitWorkerLink>,
auto_commit_worker: Option<ConsumerGroupCommitWorker>,
}
#[derive(Debug, Clone)]
struct CommitWorkerState {
group_id: String,
generation_id: i32,
member_id: String,
group_instance_id: Option<String>,
protocol: ConsumerGroupProtocol,
retention_time_ms: i64,
assignments: Vec<ConsumerAssignment>,
pending_commit_offsets: BTreeMap<(String, i32), i64>,
}
#[derive(Debug, Clone)]
struct CommitWorkerMembership {
group_id: String,
generation_id: i32,
member_id: String,
group_instance_id: Option<String>,
protocol: ConsumerGroupProtocol,
retention_time_ms: i64,
assignments: Vec<ConsumerAssignment>,
}
#[derive(Debug, Clone)]
struct CommitWorkerLink {
state: Arc<StdMutex<CommitWorkerState>>,
flush_tx: mpsc::Sender<CommitWorkerCommand>,
shutdown: Arc<AtomicBool>,
shutdown_notify: Arc<Notify>,
finished: Arc<AtomicBool>,
finished_notify: Arc<Notify>,
}
#[derive(Debug)]
enum CommitWorkerCommand {
Flush(oneshot::Sender<Result<()>>),
}
impl CommitWorkerLink {
fn lock_state(&self) -> Result<std::sync::MutexGuard<'_, CommitWorkerState>> {
self.state
.lock()
.map_err(|_| Error::Unsupported("consumer group commit worker state was poisoned"))
}
fn queue_record(&self, topic: &str, partition: i32, offset: i64) -> Result<()> {
let mut state = self.lock_state()?;
queue_commit_offset(&mut state.pending_commit_offsets, topic, partition, offset)?;
drop(state);
Ok(())
}
fn queue_assignments(&self, assignments: &[ConsumerAssignment]) -> Result<()> {
let mut state = self.lock_state()?;
for assignment in assignments {
state.pending_commit_offsets.insert(
(assignment.topic().to_owned(), assignment.partition()),
assignment.next_offset(),
);
}
drop(state);
Ok(())
}
fn pending_count(&self) -> usize {
match self.state.lock() {
Ok(state) => state.pending_commit_offsets.len(),
Err(poisoned) => poisoned.into_inner().pending_commit_offsets.len(),
}
}
fn snapshot_pending(&self) -> Result<BTreeMap<(String, i32), i64>> {
Ok(self.lock_state()?.pending_commit_offsets.clone())
}
fn update_membership(
&self,
membership: CommitWorkerMembership,
pending_commit_offsets: BTreeMap<(String, i32), i64>,
) -> Result<()> {
let mut state = self.lock_state()?;
state.group_id = membership.group_id;
state.generation_id = membership.generation_id;
state.member_id = membership.member_id;
state.group_instance_id = membership.group_instance_id;
state.protocol = membership.protocol;
state.retention_time_ms = membership.retention_time_ms;
state.assignments = membership.assignments;
state.pending_commit_offsets = pending_commit_offsets;
let assignments = state.assignments.clone();
retain_pending_commit_offsets_for(&assignments, &mut state.pending_commit_offsets);
drop(state);
Ok(())
}
fn retain_assigned(&self) -> Result<()> {
let mut state = self.lock_state()?;
let assignments = state.assignments.clone();
retain_pending_commit_offsets_for(&assignments, &mut state.pending_commit_offsets);
Ok(())
}
fn clear_committed(
&self,
snapshot: &CommitWorkerState,
assignments: &[ConsumerAssignment],
) -> Result<()> {
let mut state = self.lock_state()?;
if state.generation_id != snapshot.generation_id
|| state.member_id != snapshot.member_id
|| state.protocol != snapshot.protocol
{
return Ok(());
}
for assignment in assignments {
let key = (assignment.topic().to_owned(), assignment.partition());
let committed = state
.pending_commit_offsets
.get(&key)
.is_some_and(|pending| *pending <= assignment.next_offset());
if committed {
state.pending_commit_offsets.remove(&key);
}
}
Ok(())
}
async fn flush(&self) -> Result<()> {
let (ack_tx, ack_rx) = oneshot::channel();
self.flush_tx
.send(CommitWorkerCommand::Flush(ack_tx))
.await
.map_err(|_| Error::Unsupported("consumer group commit worker stopped"))?;
ack_rx
.await
.map_err(|_| Error::Unsupported("consumer group commit worker stopped"))?
}
fn signal_shutdown(&self) {
self.shutdown.store(true, Ordering::Release);
self.shutdown_notify.notify_waiters();
}
fn mark_finished(&self) {
self.finished.store(true, Ordering::Release);
self.finished_notify.notify_waiters();
}
async fn wait_finished(&self) {
while !self.finished.load(Ordering::Acquire) {
self.finished_notify.notified().await;
}
}
}
#[derive(Debug)]
pub struct ConsumerGroupCommitWorker {
group_id: String,
interval: Duration,
link: CommitWorkerLink,
handle: Option<JoinHandle<Result<()>>>,
}
impl ConsumerGroupCommitWorker {
pub fn group_id(&self) -> &str {
&self.group_id
}
pub fn interval(&self) -> Duration {
self.interval
}
pub fn is_finished(&self) -> bool {
match &self.handle {
Some(handle) => handle.is_finished(),
None => true,
}
}
pub async fn try_wait(&mut self) -> Result<Option<()>> {
let Some(handle) = &self.handle else {
return Ok(Some(()));
};
if !handle.is_finished() {
return Ok(None);
}
let Some(handle) = self.handle.take() else {
return Ok(Some(()));
};
handle.await?.map(Some)
}
pub async fn stop(mut self) -> Result<()> {
self.link.signal_shutdown();
let Some(handle) = self.handle.take() else {
return Ok(());
};
handle.await?
}
}
impl Drop for ConsumerGroupCommitWorker {
fn drop(&mut self) {
self.link.signal_shutdown();
self.link.mark_finished();
if let Some(handle) = &self.handle {
handle.abort();
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ConsumerGroupMetadata {
group_id: String,
generation_id: i32,
member_id: String,
group_instance_id: Option<String>,
}
impl ConsumerGroupMetadata {
pub fn group_id(&self) -> &str {
&self.group_id
}
pub fn generation_id(&self) -> i32 {
self.generation_id
}
pub fn member_id(&self) -> &str {
&self.member_id
}
pub fn group_instance_id(&self) -> Option<&str> {
self.group_instance_id.as_deref()
}
}
fn consumer_assignment_ready(
assignments: Option<&[ConsumerGroupHeartbeatTopicPartitions]>,
require_partitions: bool,
) -> bool {
assignments.is_some_and(|assignments| {
!require_partitions
|| assignments
.iter()
.any(|assignment| !assignment.partitions.is_empty())
})
}
impl ConsumerGroup {
pub fn group_id(&self) -> &str {
&self.group_id
}
pub fn group_protocol(&self) -> ConsumerGroupProtocol {
self.protocol
}
pub fn member_id(&self) -> &str {
&self.member_id
}
pub fn generation_id(&self) -> i32 {
self.generation_id
}
pub fn metadata(&self) -> ConsumerGroupMetadata {
ConsumerGroupMetadata {
group_id: self.group_id.clone(),
generation_id: self.generation_id,
member_id: self.member_id.clone(),
group_instance_id: self.config.group_instance_id.clone(),
}
}
pub fn assignments(&self) -> &[ConsumerAssignment] {
self.consumer.assignments()
}
async fn start_auto_commit_if_enabled(&mut self) -> Result<()> {
if !self.config.auto_commit {
return Ok(());
}
let interval = self.config.auto_commit_interval;
let worker = self.spawn_commit_worker(interval).await?;
self.auto_commit_worker = Some(worker);
Ok(())
}
async fn wait_for_consumer_assignment(&mut self, require_partitions: bool) -> Result<()> {
let timeout_ms = self.config.rebalance_timeout_ms.max(1) as u64;
let deadline = time::Instant::now() + Duration::from_millis(timeout_ms);
while !consumer_assignment_ready(
self.consumer_owned_partitions.as_deref(),
require_partitions,
) {
self.heartbeat_consumer().await?;
if consumer_assignment_ready(
self.consumer_owned_partitions.as_deref(),
require_partitions,
) {
break;
}
if time::Instant::now() >= deadline {
return Err(Error::Unsupported(
"consumer group assignment was not delivered before the rebalance timeout",
));
}
time::sleep(Duration::from_millis(50)).await;
}
Ok(())
}
fn queue_auto_commit_offsets(&self) -> Result<()> {
if !self.config.auto_commit {
return Ok(());
}
let worker = self.commit_worker.as_ref().ok_or(Error::Unsupported(
"automatic consumer group commit worker is not running",
))?;
worker.queue_assignments(self.consumer.assignments())
}
async fn observe_auto_commit_worker(&mut self) -> Result<()> {
if let Some(worker) = &mut self.auto_commit_worker {
if worker.try_wait().await?.is_some() {
return Err(Error::Unsupported(
"automatic consumer group commit worker stopped",
));
}
}
Ok(())
}
fn synchronize_commit_worker_membership(&self) -> Result<()> {
let Some(worker) = &self.commit_worker else {
return Ok(());
};
let pending_commit_offsets = worker.snapshot_pending()?;
worker.update_membership(
CommitWorkerMembership {
group_id: self.group_id.clone(),
generation_id: self.generation_id,
member_id: self.member_id.clone(),
group_instance_id: self.config.group_instance_id.clone(),
protocol: self.protocol,
retention_time_ms: self.retention_time_ms,
assignments: self.consumer.assignments().to_vec(),
},
pending_commit_offsets,
)
}
pub fn split_partition_queue(
&mut self,
topic: impl Into<String>,
partition: i32,
) -> Result<ConsumerPartitionQueue> {
self.consumer.split_partition_queue(topic, partition)
}
pub fn commit_record(&mut self, record: &ConsumerRecord) -> Result<()> {
if self
.consumer
.position(record.topic(), record.partition())
.is_none()
{
return Err(Error::UnassignedTopicPartition {
topic: record.topic().to_owned(),
partition: record.partition(),
});
}
if let Some(worker) = &self.commit_worker {
worker.queue_record(record.topic(), record.partition(), record.offset())
} else {
queue_commit_offset(
&mut self.pending_commit_offsets,
record.topic(),
record.partition(),
record.offset(),
)
}
}
pub fn pending_commit_count(&self) -> usize {
self.commit_worker.as_ref().map_or_else(
|| self.pending_commit_offsets.len(),
CommitWorkerLink::pending_count,
)
}
pub fn position(&self, topic: &str, partition: i32) -> Option<i64> {
self.consumer.position(topic, partition)
}
pub fn seek(&mut self, topic: &str, partition: i32, offset: i64) -> Result<()> {
self.consumer.seek(topic, partition, offset)
}
pub fn pause(&mut self, topic: &str, partition: i32) -> Result<()> {
self.consumer.pause(topic, partition)
}
pub fn resume(&mut self, topic: &str, partition: i32) -> Result<()> {
self.consumer.resume(topic, partition)
}
pub async fn fetch_watermarks(
&mut self,
topic: impl Into<String>,
partition: i32,
) -> Result<PartitionWatermarks> {
self.consumer.fetch_watermarks(topic, partition).await
}
pub async fn offset_for_leader_epoch(
&mut self,
topic: impl Into<String>,
partition: i32,
current_leader_epoch: i32,
leader_epoch: i32,
) -> Result<LeaderEpochOffset> {
self.consumer
.offset_for_leader_epoch(topic, partition, current_leader_epoch, leader_epoch)
.await
}
#[tracing::instrument(
level = "debug",
name = "kafka.consumer_group.leave",
skip_all,
fields(group_id = self.group_id.as_str(), member_id = self.member_id.as_str(), generation_id = self.generation_id),
err
)]
pub async fn leave(mut self) -> Result<()> {
if let Some(worker) = self.auto_commit_worker.take() {
worker.stop().await?;
}
if let Some(worker) = &self.commit_worker {
worker.signal_shutdown();
worker.wait_finished().await;
}
if self.protocol == ConsumerGroupProtocol::Consumer {
return self.leave_consumer().await;
}
let response = self
.coordinator
.leave_group_v3(
self.group_id.clone(),
vec![LeaveGroupMemberIdentity {
member_id: self.member_id.clone(),
group_instance_id: self.config.group_instance_id.clone(),
}],
)
.await?;
if let Some(error) = leave_group_response_error(&self.group_id, &response) {
return Err(match error {
Error::Broker { code, context } => self.config.client.broker_error(code, context),
error => error,
});
}
Ok(())
}
async fn leave_consumer(mut self) -> Result<()> {
let response = self
.coordinator
.consumer_group_heartbeat_v0(
self.group_id.clone(),
self.member_id.clone(),
-1,
self.config.group_instance_id.clone(),
None,
-1,
None,
None,
None,
)
.await?;
if response.error_code != 0 {
return Err(consumer_group_heartbeat_error(
&self.config.client,
&self.group_id,
&response,
));
}
Ok(())
}
#[tracing::instrument(
level = "debug",
name = "kafka.consumer_group.poll",
skip_all,
fields(group_id = self.group_id.as_str(), member_id = self.member_id.as_str(), generation_id = self.generation_id),
err
)]
pub async fn poll(&mut self) -> Result<Vec<ConsumerRecord>> {
self.observe_auto_commit_worker().await?;
if self.cooperative_rejoin_required {
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
"completing cooperative consumer group ownership transfer"
);
self.rejoin().await?;
}
match self.heartbeat().await {
Ok(()) => {}
Err(error) if should_rejoin_group(&error) => {
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
error = %error,
"rejoining kafka consumer group after heartbeat"
);
self.config.client.record_retry();
self.rejoin().await?
}
Err(error) => return Err(error),
}
let records = self.consumer.poll().await?;
self.queue_auto_commit_offsets()?;
self.observe_auto_commit_worker().await?;
Ok(records)
}
#[tracing::instrument(
level = "debug",
name = "kafka.consumer_group.poll_with_heartbeat",
skip_all,
fields(group_id = self.group_id.as_str(), member_id = self.member_id.as_str(), generation_id = self.generation_id),
err
)]
pub async fn poll_with_heartbeat(
&mut self,
heartbeat: &mut ConsumerGroupHeartbeat,
) -> Result<Vec<ConsumerRecord>> {
if self.protocol == ConsumerGroupProtocol::Consumer {
return self.poll_with_consumer_heartbeat(heartbeat).await;
}
let interval = heartbeat.interval;
let mut restart_heartbeat = false;
match heartbeat.state_for(&self.group_id, &self.member_id, self.generation_id) {
HeartbeatHandleState::Current => {
if should_rejoin_after_background_heartbeat(heartbeat).await? {
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
"rejoining kafka consumer group after background heartbeat"
);
self.config.client.record_retry();
self.rejoin().await?;
restart_heartbeat = true;
} else if heartbeat.is_finished() {
restart_heartbeat = true;
}
}
HeartbeatHandleState::StaleGeneration => {
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
heartbeat_member_id = heartbeat.member_id(),
heartbeat_generation_id = heartbeat.generation_id(),
"stopping stale kafka consumer group heartbeat task"
);
heartbeat.stop_stale_generation().await?;
restart_heartbeat = true;
}
HeartbeatHandleState::DifferentGroup => {
return Err(Error::Unsupported(
"background heartbeat handle belongs to a different consumer group",
));
}
}
if restart_heartbeat {
*heartbeat = self.spawn_heartbeat_task(interval).await?;
}
let result = self.poll().await;
match heartbeat.state_for(&self.group_id, &self.member_id, self.generation_id) {
HeartbeatHandleState::Current => {}
HeartbeatHandleState::StaleGeneration => {
heartbeat.stop_stale_generation().await?;
*heartbeat = self.spawn_heartbeat_task(interval).await?;
}
HeartbeatHandleState::DifferentGroup => {
return Err(Error::Unsupported(
"background heartbeat handle belongs to a different consumer group",
));
}
}
result
}
async fn poll_with_consumer_heartbeat(
&mut self,
heartbeat: &mut ConsumerGroupHeartbeat,
) -> Result<Vec<ConsumerRecord>> {
self.observe_auto_commit_worker().await?;
let interval = heartbeat.interval;
match heartbeat.state_for_consumer(&self.group_id, self.consumer_session.as_ref()) {
HeartbeatHandleState::Current => {}
HeartbeatHandleState::StaleGeneration => {
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
"stopping stale KIP-848 consumer group heartbeat task"
);
heartbeat.stop_stale_generation().await?;
*heartbeat = self.spawn_heartbeat_task(interval).await?;
}
HeartbeatHandleState::DifferentGroup => {
return Err(Error::Unsupported(
"background heartbeat handle belongs to a different consumer group",
));
}
}
self.consumer_heartbeat_state = heartbeat.consumer_state_handle();
match heartbeat.try_wait().await {
Ok(None) => {}
Ok(Some(())) => {
self.apply_consumer_heartbeat_state(heartbeat).await?;
*heartbeat = self.spawn_heartbeat_task(interval).await?;
self.consumer_heartbeat_state = heartbeat.consumer_state_handle();
}
Err(error) if should_rejoin_group(&error) => {
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
error = %error,
"rejoining KIP-848 consumer group after background heartbeat"
);
self.config.client.record_retry();
self.rejoin().await?;
*heartbeat = self.spawn_heartbeat_task(interval).await?;
self.consumer_heartbeat_state = heartbeat.consumer_state_handle();
}
Err(error) => return Err(error),
}
self.apply_consumer_heartbeat_state(heartbeat).await?;
let result = self.consumer.poll().await;
if result.is_ok() {
self.queue_auto_commit_offsets()?;
self.observe_auto_commit_worker().await?;
match heartbeat.try_wait().await {
Ok(None) => {}
Ok(Some(())) => {
self.apply_consumer_heartbeat_state(heartbeat).await?;
*heartbeat = self.spawn_heartbeat_task(interval).await?;
self.consumer_heartbeat_state = heartbeat.consumer_state_handle();
}
Err(error) if should_rejoin_group(&error) => {
self.config.client.record_retry();
self.rejoin().await?;
*heartbeat = self.spawn_heartbeat_task(interval).await?;
self.consumer_heartbeat_state = heartbeat.consumer_state_handle();
return Err(error);
}
Err(error) => return Err(error),
}
self.apply_consumer_heartbeat_state(heartbeat).await?;
}
result
}
async fn apply_consumer_heartbeat_state(
&mut self,
heartbeat: &mut ConsumerGroupHeartbeat,
) -> Result<()> {
let Some(state) = heartbeat.consumer_state_snapshot().await else {
return Ok(());
};
let member_id = state.member_id.clone();
let member_epoch = state.member_epoch;
self.member_id = member_id.clone();
self.generation_id = member_epoch;
if let Some(assignment) = state
.owned_partitions
.clone()
.filter(|_| state.assignment_version > self.consumer_heartbeat_assignment_version)
{
let previous_assignments = self.consumer.assignments().to_vec();
self.notify_rebalance(RebalancePhase::Before, &previous_assignments);
let protocol_assignment =
consumer_protocol_assignment(&assignment, &self.consumer_topic_ids)?;
let mut bootstrap = self.config.client.clone().connect().await?;
let assignments = assignments_from_protocol(
&mut self.coordinator,
&mut bootstrap,
&self.config.client,
&self.group_id,
self.config.offset_reset_policy,
&protocol_assignment,
Some((member_id.as_str(), member_epoch)),
)
.await?;
self.consumer.replace_assignments(assignments);
self.consumer_owned_partitions = Some(assignment);
self.consumer_heartbeat_assignment_version = state.assignment_version;
let current_assignments = self.consumer.assignments().to_vec();
self.notify_rebalance(RebalancePhase::After, ¤t_assignments);
}
heartbeat.member_id = member_id;
heartbeat.generation_id = member_epoch;
self.synchronize_commit_worker_membership()?;
Ok(())
}
#[tracing::instrument(
level = "debug",
name = "kafka.consumer_group.spawn_commit_worker",
skip_all,
fields(group_id = self.group_id.as_str(), member_id = self.member_id.as_str(), generation_id = self.generation_id, interval_ms = duration_millis(interval)),
err
)]
pub async fn spawn_commit_worker(
&mut self,
interval: Duration,
) -> Result<ConsumerGroupCommitWorker> {
validate_commit_worker_interval(interval)?;
if self.commit_worker.is_some() {
return Err(Error::Unsupported(
"consumer group commit worker is already running",
));
}
let coordinator = connect_group_coordinator_with_retry(
&self.config.client,
&self.group_id,
self.config.max_retries,
)
.await?;
let state = Arc::new(StdMutex::new(CommitWorkerState {
group_id: self.group_id.clone(),
generation_id: self.generation_id,
member_id: self.member_id.clone(),
group_instance_id: self.config.group_instance_id.clone(),
protocol: self.protocol,
retention_time_ms: self.retention_time_ms,
assignments: self.consumer.assignments().to_vec(),
pending_commit_offsets: self.pending_commit_offsets.clone(),
}));
let shutdown = Arc::new(AtomicBool::new(false));
let shutdown_notify = Arc::new(Notify::new());
let finished = Arc::new(AtomicBool::new(false));
let finished_notify = Arc::new(Notify::new());
let (flush_tx, flush_rx) = mpsc::channel(8);
let link = CommitWorkerLink {
state: state.clone(),
flush_tx,
shutdown: shutdown.clone(),
shutdown_notify: shutdown_notify.clone(),
finished: finished.clone(),
finished_notify: finished_notify.clone(),
};
let client_config = self.config.client.clone();
let group_id = self.group_id.clone();
let max_retries = self.config.max_retries;
let worker_span = tracing::debug_span!(
"kafka.consumer_group.background_commit",
group_id = group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
interval_ms = duration_millis(interval),
);
let worker_link = link.clone();
let handle = tokio::spawn(
async move {
let result = run_background_commit_worker(
coordinator,
worker_link.clone(),
flush_rx,
client_config,
group_id,
max_retries,
interval,
)
.await;
worker_link.mark_finished();
result
}
.instrument(worker_span),
);
self.pending_commit_offsets.clear();
self.commit_worker = Some(link.clone());
Ok(ConsumerGroupCommitWorker {
group_id: self.group_id.clone(),
interval,
link,
handle: Some(handle),
})
}
#[tracing::instrument(
level = "debug",
name = "kafka.consumer_group.spawn_heartbeat",
skip_all,
fields(group_id = self.group_id.as_str(), member_id = self.member_id.as_str(), generation_id = self.generation_id, interval_ms = duration_millis(interval)),
err
)]
pub async fn spawn_heartbeat_task(&self, interval: Duration) -> Result<ConsumerGroupHeartbeat> {
validate_heartbeat_interval(interval)?;
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
interval_ms = duration_millis(interval),
"starting kafka consumer group heartbeat task"
);
let mut bootstrap = self.config.client.clone().connect().await?;
let coordinator = find_group_coordinator_with_retry(
&mut bootstrap,
&self.config.client,
&self.group_id,
self.config.max_retries,
)
.await?;
let mut coordinator = self
.config
.client
.connect_broker(coordinator_addr(&coordinator))
.await?;
let group_id = self.group_id.clone();
let generation_id = self.generation_id;
let member_id = self.member_id.clone();
let group_instance_id = self.config.group_instance_id.clone();
let consumer_session = self.consumer_session.clone();
let consumer_state = (self.protocol == ConsumerGroupProtocol::Consumer).then(|| {
Arc::new(Mutex::new(ConsumerProtocolHeartbeatState {
member_id: self.member_id.clone(),
member_epoch: self.generation_id,
owned_partitions: self.consumer_owned_partitions.clone(),
assignment_version: 0,
heartbeat_interval: interval,
}))
});
let consumer_state_for_task = consumer_state.clone();
let topics = self.resolved_topics.clone();
let server_assignor = self.config.server_assignor.clone();
let rebalance_timeout_ms = self.config.rebalance_timeout_ms;
let (shutdown, shutdown_rx) = oneshot::channel();
let heartbeat_span = tracing::debug_span!(
"kafka.consumer_group.background_heartbeat",
group_id = group_id.as_str(),
member_id = member_id.as_str(),
generation_id,
interval_ms = duration_millis(interval),
);
let handle = tokio::spawn(
async move {
if let Some(consumer_state) = consumer_state_for_task {
run_background_consumer_heartbeat(
&mut coordinator,
group_id,
consumer_state,
ConsumerProtocolHeartbeatConfig {
group_instance_id,
topics,
server_assignor,
rebalance_timeout_ms,
},
shutdown_rx,
)
.await
} else {
run_background_heartbeat(
&mut coordinator,
group_id,
generation_id,
member_id,
group_instance_id,
interval,
shutdown_rx,
)
.await
}
}
.instrument(heartbeat_span),
);
Ok(ConsumerGroupHeartbeat {
group_id: self.group_id.clone(),
generation_id: self.generation_id,
member_id: self.member_id.clone(),
interval,
shutdown: Some(shutdown),
handle: Some(handle),
consumer_session,
consumer_state,
})
}
pub async fn rejoin(&mut self) -> Result<()> {
let previous_assignments = self.consumer.assignments().to_vec();
self.notify_rebalance(RebalancePhase::Before, &previous_assignments);
let paused = self
.consumer
.assignments()
.iter()
.filter(|assignment| assignment.is_paused())
.map(|assignment| (assignment.topic().to_owned(), assignment.partition()))
.collect::<Vec<_>>();
let owned_partitions = owned_partitions_from_assignments(self.consumer.assignments());
let commit_worker = self.commit_worker.clone();
let pending_commit_offsets = match &commit_worker {
Some(worker) => worker.snapshot_pending()?,
None => self.pending_commit_offsets.clone(),
};
if self.protocol == ConsumerGroupProtocol::Consumer {
let owned_partitions = consumer_owned_partitions_from_assignments(
self.consumer.assignments(),
&self.consumer_topic_ids,
);
let mut joined = self
.config
.clone()
.join_consumer(
Some(self.member_id.clone()),
self.generation_id,
Some(owned_partitions),
)
.await?;
for (topic, partition) in paused {
if joined.consumer.position(&topic, partition).is_some() {
joined.consumer.pause(&topic, partition)?;
}
}
joined.pending_commit_offsets = if commit_worker.is_some() {
BTreeMap::new()
} else {
pending_commit_offsets.clone()
};
joined.commit_worker = commit_worker.clone();
joined.auto_commit_worker = self.auto_commit_worker.take();
if commit_worker.is_none() {
joined.retain_pending_commit_offsets();
}
let current_assignments = joined.consumer.assignments().to_vec();
*self = joined;
if let Some(worker) = &self.commit_worker {
worker.update_membership(
CommitWorkerMembership {
group_id: self.group_id.clone(),
generation_id: self.generation_id,
member_id: self.member_id.clone(),
group_instance_id: self.config.group_instance_id.clone(),
protocol: self.protocol,
retention_time_ms: self.retention_time_ms,
assignments: self.consumer.assignments().to_vec(),
},
pending_commit_offsets,
)?;
}
self.notify_rebalance(RebalancePhase::After, ¤t_assignments);
return Ok(());
}
let mut joined = self
.config
.clone()
.join_with_member_id(
owned_partitions,
Some(self.member_id.clone()),
Some(self.generation_id),
)
.await?;
for (topic, partition) in paused {
if joined.consumer.position(&topic, partition).is_some() {
joined.consumer.pause(&topic, partition)?;
}
}
joined.pending_commit_offsets = if commit_worker.is_some() {
BTreeMap::new()
} else {
pending_commit_offsets.clone()
};
joined.commit_worker = commit_worker.clone();
joined.auto_commit_worker = self.auto_commit_worker.take();
if commit_worker.is_none() {
joined.retain_pending_commit_offsets();
}
let current_assignments = joined.consumer.assignments().to_vec();
*self = joined;
if let Some(worker) = &self.commit_worker {
worker.update_membership(
CommitWorkerMembership {
group_id: self.group_id.clone(),
generation_id: self.generation_id,
member_id: self.member_id.clone(),
group_instance_id: self.config.group_instance_id.clone(),
protocol: self.protocol,
retention_time_ms: self.retention_time_ms,
assignments: self.consumer.assignments().to_vec(),
},
pending_commit_offsets,
)?;
}
self.notify_rebalance(RebalancePhase::After, ¤t_assignments);
Ok(())
}
fn notify_rebalance(&self, phase: RebalancePhase, assignments: &[ConsumerAssignment]) {
if let Some(listener) = &self.config.rebalance_listener {
listener.on_rebalance(RebalanceEvent {
phase,
group_id: &self.group_id,
member_id: &self.member_id,
generation_id: self.generation_id,
protocol: self.protocol,
assignments,
});
}
}
fn retain_pending_commit_offsets(&mut self) {
let assigned = self
.consumer
.assignments()
.iter()
.map(|assignment| (assignment.topic().to_owned(), assignment.partition()))
.collect::<BTreeSet<_>>();
self.pending_commit_offsets
.retain(|key, _| assigned.contains(key));
}
#[tracing::instrument(
level = "debug",
name = "kafka.consumer_group.heartbeat",
skip_all,
fields(group_id = self.group_id.as_str(), member_id = self.member_id.as_str(), generation_id = self.generation_id),
err
)]
pub async fn heartbeat(&mut self) -> Result<()> {
if self.protocol == ConsumerGroupProtocol::Consumer {
return self.heartbeat_consumer().await;
}
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
"sending kafka consumer group heartbeat"
);
let response = if let Some(group_instance_id) = &self.config.group_instance_id {
self.coordinator
.heartbeat_v3(
self.group_id.clone(),
self.generation_id,
self.member_id.clone(),
Some(group_instance_id.clone()),
)
.await?
} else {
self.coordinator
.heartbeat_v2(
self.group_id.clone(),
self.generation_id,
self.member_id.clone(),
)
.await?
};
if response.error_code != 0 {
return Err(self.config.client.broker_error(
response.error_code,
format!("heartbeat group {}", self.group_id),
));
}
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
"sent kafka consumer group heartbeat"
);
Ok(())
}
async fn heartbeat_consumer(&mut self) -> Result<()> {
let owned_partitions = self.consumer_owned_partitions.clone();
let mut retry_attempt = 0;
let response = loop {
let response = match self
.coordinator
.consumer_group_heartbeat_v0(
self.group_id.clone(),
self.member_id.clone(),
self.generation_id,
self.config.group_instance_id.clone(),
None,
self.config.rebalance_timeout_ms,
Some(self.resolved_topics.clone()),
self.config.server_assignor.clone(),
owned_partitions.clone(),
)
.await
{
Ok(response) => response,
Err(error)
if retry_attempt < self.config.max_retries && should_rejoin_group(&error) =>
{
retry_attempt += 1;
self.config.client.record_retry();
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
member_epoch = self.generation_id,
retry_attempt,
error = %error,
"retrying kafka KIP-848 consumer heartbeat after transport failure"
);
time::sleep(group_retry_backoff(retry_attempt)).await;
self.coordinator = connect_group_coordinator_with_retry(
&self.config.client,
&self.group_id,
self.config.max_retries.saturating_sub(retry_attempt),
)
.await?;
continue;
}
Err(error) => return Err(error),
};
if response.error_code == 0 {
break response;
}
let error =
consumer_group_heartbeat_error(&self.config.client, &self.group_id, &response);
if retry_attempt >= self.config.max_retries || !should_rejoin_group(&error) {
return Err(error);
}
retry_attempt += 1;
self.config.client.record_retry();
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
member_epoch = self.generation_id,
retry_attempt,
error = %error,
"retrying kafka KIP-848 consumer group heartbeat"
);
time::sleep(group_retry_backoff(retry_attempt)).await;
self.coordinator = connect_group_coordinator_with_retry(
&self.config.client,
&self.group_id,
self.config.max_retries.saturating_sub(retry_attempt),
)
.await?;
};
if let Some(member_id) = response.member_id.clone() {
self.member_id = member_id;
}
self.generation_id = response.member_epoch;
if let Some(assignment) = response.assignment {
let previous_assignments = self.consumer.assignments().to_vec();
self.notify_rebalance(RebalancePhase::Before, &previous_assignments);
let protocol_assignment =
consumer_protocol_assignment(&assignment, &self.consumer_topic_ids)?;
let mut bootstrap = self.config.client.clone().connect().await?;
let assignments = assignments_from_protocol(
&mut self.coordinator,
&mut bootstrap,
&self.config.client,
&self.group_id,
self.config.offset_reset_policy,
&protocol_assignment,
Some((self.member_id.as_str(), self.generation_id)),
)
.await?;
self.consumer.replace_assignments(assignments);
self.consumer_owned_partitions = Some(assignment);
let current_assignments = self.consumer.assignments().to_vec();
self.notify_rebalance(RebalancePhase::After, ¤t_assignments);
}
self.synchronize_commit_worker_membership()?;
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
member_epoch = self.generation_id,
assignment_count = self.consumer.assignments().len(),
"sent kafka KIP-848 consumer group heartbeat"
);
Ok(())
}
#[tracing::instrument(
level = "debug",
name = "kafka.consumer_group.commit_offsets",
skip_all,
fields(group_id = self.group_id.as_str(), member_id = self.member_id.as_str(), generation_id = self.generation_id, assignment_count = self.consumer.assignments().len()),
err
)]
pub async fn commit_offsets(&mut self) -> Result<()> {
let assignments = self.consumer.assignments().to_vec();
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
topic_count = offset_commit_topics(&assignments).len(),
"committing kafka consumer group offsets"
);
if let Some(worker) = &self.commit_worker {
worker.queue_assignments(&assignments)?;
return worker.flush().await;
}
self.commit_assignment_offsets(&assignments).await?;
self.clear_pending_commit_offsets(&assignments);
Ok(())
}
#[tracing::instrument(
level = "debug",
name = "kafka.consumer_group.commit_queued_offsets",
skip_all,
fields(group_id = self.group_id.as_str(), member_id = self.member_id.as_str(), generation_id = self.generation_id),
err
)]
pub async fn commit_queued_offsets(&mut self) -> Result<()> {
if let Some(worker) = &self.commit_worker {
worker.retain_assigned()?;
return worker.flush().await;
}
self.retain_pending_commit_offsets();
let assignments =
pending_commit_assignments(self.consumer.assignments(), &self.pending_commit_offsets);
if assignments.is_empty() {
return Ok(());
}
self.commit_assignment_offsets(&assignments).await?;
self.clear_pending_commit_offsets(&assignments);
Ok(())
}
async fn commit_assignment_offsets(
&mut self,
assignments: &[ConsumerAssignment],
) -> Result<()> {
let response_topics = match if self.protocol == ConsumerGroupProtocol::Consumer {
self.coordinator
.offset_commit_v9(
self.group_id.clone(),
self.generation_id,
self.member_id.clone(),
self.config.group_instance_id.clone(),
offset_commit_topics_v9(assignments),
)
.await
.map(|response| response.topics)
} else if let Some(group_instance_id) = &self.config.group_instance_id {
self.coordinator
.offset_commit_v7(
self.group_id.clone(),
self.generation_id,
self.member_id.clone(),
Some(group_instance_id.clone()),
offset_commit_topics_v7(assignments),
)
.await
.map(|response| response.topics)
} else {
self.coordinator
.offset_commit_v2(
self.group_id.clone(),
self.generation_id,
self.member_id.clone(),
self.retention_time_ms,
offset_commit_topics(assignments),
)
.await
.map(|response| response.topics)
} {
Ok(topics) => topics,
Err(error) if should_rejoin_group(&error) => {
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
error = %error,
"rejoining kafka consumer group after offset commit request"
);
self.config.client.record_retry();
self.rejoin().await?;
return Err(error);
}
Err(error) => return Err(error),
};
if let Some(error) = offset_commit_response_error(&self.group_id, &response_topics) {
let error = match error {
Error::Broker { code, context } => self.config.client.broker_error(code, context),
error => error,
};
if should_rejoin_group(&error) {
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
error = %error,
"rejoining kafka consumer group after offset commit"
);
self.config.client.record_retry();
self.rejoin().await?;
}
return Err(error);
}
debug!(
group_id = self.group_id.as_str(),
member_id = self.member_id.as_str(),
generation_id = self.generation_id,
"committed kafka consumer group offsets"
);
Ok(())
}
fn clear_pending_commit_offsets(&mut self, assignments: &[ConsumerAssignment]) {
for assignment in assignments {
let key = (assignment.topic().to_owned(), assignment.partition());
let committed = self
.pending_commit_offsets
.get(&key)
.is_some_and(|pending| *pending <= assignment.next_offset());
if committed {
self.pending_commit_offsets.remove(&key);
}
}
}
}
#[derive(Debug)]
pub struct ConsumerGroupHeartbeat {
group_id: String,
generation_id: i32,
member_id: String,
interval: Duration,
shutdown: Option<oneshot::Sender<()>>,
handle: Option<JoinHandle<Result<()>>>,
consumer_session: Option<Arc<()>>,
consumer_state: Option<Arc<Mutex<ConsumerProtocolHeartbeatState>>>,
}
impl ConsumerGroupHeartbeat {
pub fn group_id(&self) -> &str {
&self.group_id
}
pub fn generation_id(&self) -> i32 {
self.generation_id
}
pub fn member_id(&self) -> &str {
&self.member_id
}
pub fn interval(&self) -> Duration {
self.interval
}
pub fn is_finished(&self) -> bool {
match &self.handle {
Some(handle) => handle.is_finished(),
None => true,
}
}
pub async fn try_wait(&mut self) -> Result<Option<()>> {
let Some(handle) = &self.handle else {
return Ok(Some(()));
};
if !handle.is_finished() {
return Ok(None);
}
let Some(handle) = self.handle.take() else {
return Ok(Some(()));
};
handle.await?.map(Some)
}
pub async fn stop(mut self) -> Result<()> {
self.signal_shutdown();
let Some(handle) = self.handle.take() else {
return Ok(());
};
handle.await?
}
fn state_for(
&self,
group_id: &str,
member_id: &str,
generation_id: i32,
) -> HeartbeatHandleState {
if self.group_id != group_id {
return HeartbeatHandleState::DifferentGroup;
}
if self.member_id != member_id || self.generation_id != generation_id {
return HeartbeatHandleState::StaleGeneration;
}
HeartbeatHandleState::Current
}
fn state_for_consumer(
&self,
group_id: &str,
group_session: Option<&Arc<()>>,
) -> HeartbeatHandleState {
if self.group_id != group_id {
return HeartbeatHandleState::DifferentGroup;
}
if self
.consumer_session
.as_ref()
.zip(group_session)
.is_some_and(|(handle, group)| Arc::ptr_eq(handle, group))
{
HeartbeatHandleState::Current
} else {
HeartbeatHandleState::StaleGeneration
}
}
fn consumer_state_handle(&self) -> Option<Arc<Mutex<ConsumerProtocolHeartbeatState>>> {
self.consumer_state.clone()
}
async fn consumer_state_snapshot(&self) -> Option<ConsumerProtocolHeartbeatState> {
let state = self.consumer_state.as_ref()?;
Some(state.lock().await.clone())
}
async fn stop_stale_generation(&mut self) -> Result<()> {
self.signal_shutdown();
let Some(handle) = self.handle.take() else {
return Ok(());
};
match handle.await? {
Ok(()) => Ok(()),
Err(error) if should_rejoin_group(&error) => Ok(()),
Err(error) => Err(error),
}
}
fn signal_shutdown(&mut self) {
if let Some(shutdown) = self.shutdown.take() {
let _ = shutdown.send(());
}
}
}
impl Drop for ConsumerGroupHeartbeat {
fn drop(&mut self) {
self.signal_shutdown();
if let Some(handle) = &self.handle {
handle.abort();
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum HeartbeatHandleState {
Current,
StaleGeneration,
DifferentGroup,
}
fn assignments_for_strategy(
protocol_name: &str,
members: &[JoinGroupMember],
metadata: &MetadataResponseV1,
) -> Result<Vec<SyncGroupAssignment>> {
match protocol_name {
RANGE_PROTOCOL => range_assignments(members, metadata),
ROUND_ROBIN_PROTOCOL => round_robin_assignments(members, metadata),
STICKY_PROTOCOL => sticky_assignments(members, metadata),
COOPERATIVE_STICKY_PROTOCOL => cooperative_sticky_assignments(members, metadata),
_ => Err(Error::Unsupported(
"consumer group selected an unsupported assignment strategy",
)),
}
}
#[derive(Debug, Default)]
struct ClassicSubscription {
topics: Vec<String>,
user_data: Option<Vec<u8>>,
owned_partitions: Vec<ConsumerProtocolTopicAssignment>,
generation: Option<i32>,
}
fn decode_classic_subscription(bytes: &[u8]) -> Result<ClassicSubscription> {
let mut decoder = Decoder::new(bytes);
let version = decoder.read_i16()?;
if version < 0 {
return Err(Error::Unsupported(
"consumer group subscription has an invalid version",
));
}
let topics = decoder
.read_array("consumer group subscription topics", Decoder::read_string)?
.unwrap_or_default();
let user_data = decoder.read_nullable_bytes()?;
let owned_partitions = if version >= 1 {
decoder
.read_array("consumer group subscription owned partitions", |decoder| {
let topic = decoder.read_string()?;
let partitions = decoder
.read_array("consumer group subscription partitions", Decoder::read_i32)?
.unwrap_or_default();
Ok(ConsumerProtocolTopicAssignment { topic, partitions })
})?
.unwrap_or_default()
} else {
Vec::new()
};
let generation = if version >= 2 {
let generation = decoder.read_i32()?;
(generation >= 0).then_some(generation)
} else {
None
};
if version >= 3 {
let _rack_id = decoder.read_nullable_string()?;
}
Ok(ClassicSubscription {
topics,
user_data,
owned_partitions,
generation,
})
}
pub(crate) fn range_assignments(
members: &[JoinGroupMember],
metadata: &MetadataResponseV1,
) -> Result<Vec<SyncGroupAssignment>> {
let mut subscriptions_by_topic = BTreeMap::<String, Vec<String>>::new();
for member in members {
let subscription = decode_classic_subscription(&member.metadata)?;
for topic in &subscription.topics {
subscriptions_by_topic
.entry(topic.clone())
.or_default()
.push(member.member_id.clone());
}
}
let mut assigned_by_member = BTreeMap::<String, BTreeMap<String, Vec<i32>>>::new();
for member in members {
assigned_by_member.insert(member.member_id.clone(), BTreeMap::new());
}
for (topic, mut topic_members) in subscriptions_by_topic {
topic_members.sort();
topic_members.dedup();
let partitions = partitions_for(metadata, &topic)?;
for (member_id, partitions) in range_for_topic(&topic_members, &partitions) {
assigned_by_member
.entry(member_id)
.or_default()
.insert(topic.clone(), partitions);
}
}
encode_member_assignments(assigned_by_member)
}
pub(crate) fn round_robin_assignments(
members: &[JoinGroupMember],
metadata: &MetadataResponseV1,
) -> Result<Vec<SyncGroupAssignment>> {
let mut subscriptions = BTreeMap::<String, BTreeSet<String>>::new();
let mut subscribed_topics = BTreeSet::new();
for member in members {
let subscription = decode_classic_subscription(&member.metadata)?;
let topics = subscription.topics.into_iter().collect::<BTreeSet<_>>();
subscribed_topics.extend(topics.iter().cloned());
subscriptions.insert(member.member_id.clone(), topics);
}
let member_ids = subscriptions.keys().cloned().collect::<Vec<_>>();
let mut assigned_by_member = member_ids
.iter()
.map(|member_id| (member_id.clone(), BTreeMap::new()))
.collect::<BTreeMap<String, BTreeMap<String, Vec<i32>>>>();
if member_ids.is_empty() {
return Ok(Vec::new());
}
let mut next_member = 0;
for topic in subscribed_topics {
for partition in partitions_for(metadata, &topic)? {
let member_index = (0..member_ids.len())
.map(|offset| (next_member + offset) % member_ids.len())
.find(|index| {
subscriptions
.get(&member_ids[*index])
.is_some_and(|topics| topics.contains(&topic))
})
.ok_or(Error::Unsupported(
"round-robin assignment found no subscribed member",
))?;
assigned_by_member
.entry(member_ids[member_index].clone())
.or_default()
.entry(topic.clone())
.or_default()
.push(partition);
next_member = (member_index + 1) % member_ids.len();
}
}
encode_member_assignments(assigned_by_member)
}
pub(crate) fn sticky_assignments(
members: &[JoinGroupMember],
metadata: &MetadataResponseV1,
) -> Result<Vec<SyncGroupAssignment>> {
let mut subscriptions = BTreeMap::<String, BTreeSet<String>>::new();
let mut owned_by_member = BTreeMap::<String, BTreeSet<(String, i32)>>::new();
let mut generations = BTreeMap::<String, i32>::new();
for member in members {
let subscription = decode_classic_subscription(&member.metadata)?;
let member_data = decode_sticky_user_data(subscription.user_data.as_deref());
subscriptions.insert(
member.member_id.clone(),
subscription.topics.into_iter().collect(),
);
owned_by_member.insert(
member.member_id.clone(),
member_data
.partitions
.into_iter()
.flat_map(|assignment| {
assignment
.partitions
.into_iter()
.map(move |partition| (assignment.topic.clone(), partition))
})
.collect(),
);
generations.insert(
member.member_id.clone(),
member_data.generation.unwrap_or(-1),
);
}
sticky_assignments_from_owned(subscriptions, owned_by_member, generations, metadata, false)
}
pub(crate) fn cooperative_sticky_assignments(
members: &[JoinGroupMember],
metadata: &MetadataResponseV1,
) -> Result<Vec<SyncGroupAssignment>> {
let mut subscriptions = BTreeMap::<String, BTreeSet<String>>::new();
let mut owned_by_member = BTreeMap::<String, BTreeSet<(String, i32)>>::new();
let mut generations = BTreeMap::<String, i32>::new();
for member in members {
let subscription = decode_classic_subscription(&member.metadata)?;
subscriptions.insert(
member.member_id.clone(),
subscription.topics.into_iter().collect(),
);
owned_by_member.insert(
member.member_id.clone(),
subscription
.owned_partitions
.into_iter()
.flat_map(|assignment| {
assignment
.partitions
.into_iter()
.map(move |partition| (assignment.topic.clone(), partition))
})
.collect(),
);
generations.insert(
member.member_id.clone(),
subscription.generation.unwrap_or(-1),
);
}
sticky_assignments_from_owned(subscriptions, owned_by_member, generations, metadata, true)
}
fn sticky_assignments_from_owned(
subscriptions: BTreeMap<String, BTreeSet<String>>,
owned_by_member: BTreeMap<String, BTreeSet<(String, i32)>>,
generations: BTreeMap<String, i32>,
metadata: &MetadataResponseV1,
cooperative: bool,
) -> Result<Vec<SyncGroupAssignment>> {
let mut assigned_by_member = subscriptions
.keys()
.map(|member_id| (member_id.clone(), BTreeMap::<String, Vec<i32>>::new()))
.collect::<BTreeMap<_, _>>();
let mut owner_by_partition = BTreeMap::<(String, i32), String>::new();
let mut invalidated_previous = BTreeSet::<(String, i32)>::new();
let mut available = BTreeSet::<(String, i32)>::new();
let mut assignment_order = Vec::<(String, i32)>::new();
let topics = subscriptions
.values()
.flat_map(|topics| topics.iter())
.cloned()
.collect::<BTreeSet<_>>();
let mut topics_by_candidate_count = topics
.into_iter()
.map(|topic| {
let candidate_count = subscriptions
.values()
.filter(|topics| topics.contains(&topic))
.count();
(candidate_count, topic)
})
.collect::<Vec<_>>();
topics_by_candidate_count.sort();
for (_, topic) in topics_by_candidate_count {
for partition in partitions_for(metadata, &topic)? {
let partition = (topic.clone(), partition);
available.insert(partition.clone());
assignment_order.push(partition);
}
}
for (member_id, owned) in &owned_by_member {
for partition in owned {
if !available.contains(partition) {
continue;
}
if invalidated_previous.contains(partition) {
continue;
}
if let Some(current_owner) = owner_by_partition.get(partition).cloned() {
let current_generation = generations.get(¤t_owner).copied().unwrap_or(-1);
let member_generation = generations.get(member_id).copied().unwrap_or(-1);
if member_generation == current_generation {
debug!(
topic = partition.0.as_str(),
partition = partition.1,
first_member_id = current_owner.as_str(),
second_member_id = member_id.as_str(),
generation = member_generation,
"invalidating sticky previous ownership claimed by multiple members"
);
remove_assignment(&mut assigned_by_member, ¤t_owner, partition);
owner_by_partition.remove(partition);
invalidated_previous.insert(partition.clone());
continue;
}
if member_generation <= current_generation {
continue;
}
remove_assignment(&mut assigned_by_member, ¤t_owner, partition);
}
owner_by_partition.insert(partition.clone(), member_id.clone());
add_assignment(&mut assigned_by_member, member_id, partition);
}
}
for partition in assignment_order {
if owner_by_partition.contains_key(&partition) {
continue;
}
let topic = &partition.0;
let member_id = subscriptions
.iter()
.filter(|(_, topics)| topics.contains(topic))
.min_by_key(|(member_id, _)| {
(assignment_count(&assigned_by_member, member_id), *member_id)
})
.map(|(member_id, _)| member_id.clone())
.ok_or(Error::Unsupported(
"sticky assignor found no subscribed member",
))?;
owner_by_partition.insert(partition.clone(), member_id.clone());
add_assignment(&mut assigned_by_member, &member_id, &partition);
}
balance_owned_assignments(
&mut assigned_by_member,
&mut owner_by_partition,
&subscriptions,
);
if cooperative {
let mut transfers = BTreeSet::<(String, i32)>::new();
for (partition, new_owner) in &owner_by_partition {
if let Some(old_owner) = owned_by_member
.iter()
.find_map(|(member_id, owned)| owned.contains(partition).then_some(member_id))
{
if old_owner != new_owner && subscriptions.contains_key(old_owner) {
transfers.insert(partition.clone());
}
}
}
for partition in transfers {
if let Some(new_owner) = owner_by_partition.get(&partition) {
remove_assignment(&mut assigned_by_member, new_owner, &partition);
}
}
}
encode_member_assignments(assigned_by_member)
}
#[derive(Debug, Default)]
struct StickyMemberData {
partitions: Vec<ConsumerProtocolTopicAssignment>,
generation: Option<i32>,
}
fn encode_sticky_user_data(
owned_partitions: &[ConsumerProtocolTopicAssignment],
generation: Option<i32>,
) -> Result<Vec<u8>> {
let mut encoder = Encoder::new();
encoder.write_array(Some(owned_partitions), |encoder, assignment| {
encoder.write_string(&assignment.topic)?;
encoder.write_array(
Some(assignment.partitions.as_slice()),
|encoder, partition| {
encoder.write_i32(*partition);
Ok(())
},
)
})?;
if let Some(generation) = generation {
encoder.write_i32(generation);
}
Ok(encoder.into_bytes())
}
fn decode_sticky_user_data(user_data: Option<&[u8]>) -> StickyMemberData {
let Some(user_data) = user_data else {
return StickyMemberData::default();
};
let decoded = (|| -> Result<StickyMemberData> {
let mut decoder = Decoder::new(user_data);
let partitions = decoder
.read_array("sticky assignor topics", |decoder| {
let topic = decoder.read_string()?;
let partitions = decoder
.read_array("sticky assignor partitions", Decoder::read_i32)?
.unwrap_or_default();
Ok(ConsumerProtocolTopicAssignment { topic, partitions })
})?
.unwrap_or_default();
let generation = if decoder.is_empty() {
None
} else {
Some(decoder.read_i32()?)
};
if !decoder.is_empty() {
return Err(Error::Unsupported("invalid sticky assignor user data"));
}
Ok(StickyMemberData {
partitions,
generation,
})
})();
match decoded {
Ok(member_data) => member_data,
Err(error) => {
debug!(error = %error, "ignoring malformed sticky assignor user data");
StickyMemberData::default()
}
}
}
fn balance_owned_assignments(
assigned_by_member: &mut BTreeMap<String, BTreeMap<String, Vec<i32>>>,
owner_by_partition: &mut BTreeMap<(String, i32), String>,
subscriptions: &BTreeMap<String, BTreeSet<String>>,
) {
let member_ids = subscriptions.keys().cloned().collect::<Vec<_>>();
loop {
let Some(source) = member_ids
.iter()
.max_by_key(|member_id| assignment_count(assigned_by_member, member_id))
else {
return;
};
let source_count = assignment_count(assigned_by_member, source);
let source_partitions = assigned_by_member
.get(source)
.into_iter()
.flat_map(|topics| {
topics.iter().flat_map(|(topic, partitions)| {
partitions
.iter()
.map(|partition| (topic.clone(), *partition))
})
})
.collect::<Vec<_>>();
let mut candidate = None;
for partition in source_partitions {
for target in &member_ids {
if target == source
|| !subscriptions
.get(target)
.is_some_and(|topics| topics.contains(&partition.0))
{
continue;
}
let target_count = assignment_count(assigned_by_member, target);
if target_count + 1 >= source_count {
continue;
}
let should_replace = candidate.as_ref().map_or(true, |(_, _, _, current_count)| {
target_count < *current_count
});
if should_replace {
candidate = Some((
partition.0.clone(),
partition.1,
target.clone(),
target_count,
));
}
}
}
let Some((topic, partition, target, _)) = candidate else {
return;
};
let partition = (topic, partition);
remove_assignment(assigned_by_member, source, &partition);
add_assignment(assigned_by_member, &target, &partition);
owner_by_partition.insert(partition, target);
}
}
fn add_assignment(
assigned_by_member: &mut BTreeMap<String, BTreeMap<String, Vec<i32>>>,
member_id: &str,
partition: &(String, i32),
) {
assigned_by_member
.entry(member_id.to_owned())
.or_default()
.entry(partition.0.clone())
.or_default()
.push(partition.1);
}
fn remove_assignment(
assigned_by_member: &mut BTreeMap<String, BTreeMap<String, Vec<i32>>>,
member_id: &str,
partition: &(String, i32),
) {
let Some(topics) = assigned_by_member.get_mut(member_id) else {
return;
};
let Some(partitions) = topics.get_mut(&partition.0) else {
return;
};
partitions.retain(|candidate| *candidate != partition.1);
if partitions.is_empty() {
topics.remove(&partition.0);
}
}
fn assignment_count(
assigned_by_member: &BTreeMap<String, BTreeMap<String, Vec<i32>>>,
member_id: &str,
) -> usize {
assigned_by_member
.get(member_id)
.into_iter()
.flat_map(|topics| topics.values())
.map(Vec::len)
.sum()
}
fn owned_partitions_from_assignments(
assignments: &[ConsumerAssignment],
) -> Vec<ConsumerProtocolTopicAssignment> {
let mut owned = BTreeMap::<String, Vec<i32>>::new();
for assignment in assignments {
owned
.entry(assignment.topic().to_owned())
.or_default()
.push(assignment.partition());
}
owned
.into_iter()
.map(|(topic, mut partitions)| {
partitions.sort_unstable();
ConsumerProtocolTopicAssignment { topic, partitions }
})
.collect()
}
fn consumer_owned_partitions_from_assignments(
assignments: &[ConsumerAssignment],
topic_ids: &BTreeMap<String, [u8; 16]>,
) -> Vec<ConsumerGroupHeartbeatTopicPartitions> {
let mut by_topic = BTreeMap::<[u8; 16], Vec<i32>>::new();
for assignment in assignments {
if let Some(topic_id) = topic_ids.get(assignment.topic()) {
by_topic
.entry(*topic_id)
.or_default()
.push(assignment.partition());
}
}
by_topic
.into_iter()
.map(|(topic_id, mut partitions)| {
partitions.sort_unstable();
ConsumerGroupHeartbeatTopicPartitions {
topic_id,
partitions,
}
})
.collect()
}
fn cooperative_assignment_requires_rejoin(
members: &[JoinGroupMember],
assignment: &ConsumerProtocolAssignmentV0,
) -> Result<bool> {
let mut owned_partitions = BTreeSet::<(String, i32)>::new();
for member in members {
let subscription = decode_classic_subscription(&member.metadata)?;
for owned in subscription.owned_partitions {
for partition in owned.partitions {
owned_partitions.insert((owned.topic.clone(), partition));
}
}
}
let assigned_partitions = assignment
.assignments
.iter()
.flat_map(|assigned| {
assigned
.partitions
.iter()
.map(move |partition| (assigned.topic.clone(), *partition))
})
.collect::<BTreeSet<_>>();
Ok(owned_partitions
.iter()
.any(|partition| !assigned_partitions.contains(partition)))
}
fn cooperative_rejoin_required_for_assignment(
members: &[JoinGroupMember],
owned_partitions: &[ConsumerProtocolTopicAssignment],
assignment: &ConsumerProtocolAssignmentV0,
is_rejoin: bool,
) -> Result<bool> {
if !members.is_empty() {
return cooperative_assignment_requires_rejoin(members, assignment);
}
Ok(
owned_partitions_require_rejoin(owned_partitions, assignment)
|| (!is_rejoin && assignment.assignments.is_empty()),
)
}
fn owned_partitions_require_rejoin(
owned_partitions: &[ConsumerProtocolTopicAssignment],
assignment: &ConsumerProtocolAssignmentV0,
) -> bool {
owned_partitions.iter().any(|owned| {
owned.partitions.iter().any(|partition| {
!assignment.assignments.iter().any(|assigned| {
assigned.topic == owned.topic && assigned.partitions.contains(partition)
})
})
})
}
fn encode_member_assignments(
assigned_by_member: BTreeMap<String, BTreeMap<String, Vec<i32>>>,
) -> Result<Vec<SyncGroupAssignment>> {
assigned_by_member
.into_iter()
.map(|(member_id, topics)| {
let assignments = topics
.into_iter()
.map(|(topic, mut partitions)| {
partitions.sort_unstable();
ConsumerProtocolTopicAssignment { topic, partitions }
})
.collect();
Ok(SyncGroupAssignment {
member_id,
assignment: ConsumerProtocolAssignmentV0 {
assignments,
user_data: None,
}
.encode()?,
})
})
.collect()
}
async fn assignments_from_protocol(
coordinator: &mut Client,
bootstrap: &mut Client,
client_config: &ClientConfig,
group_id: &str,
offset_reset_policy: OffsetResetPolicy,
assignment: &ConsumerProtocolAssignmentV0,
consumer_member: Option<(&str, i32)>,
) -> Result<Vec<ConsumerAssignment>> {
let mut assignments = Vec::new();
let mut reset_partitions = Vec::new();
let leader_epochs = assignment_leader_epochs(bootstrap, assignment).await?;
let (offset_topics, offset_error_code) =
if let Some((member_id, member_epoch)) = consumer_member {
let offsets = coordinator
.offset_fetch_v9(
group_id.to_owned(),
Some(member_id.to_owned()),
member_epoch,
Some(offset_fetch_topics_v9(assignment)),
)
.await?;
let group = offsets
.groups
.into_iter()
.find(|group| group.group_id == group_id)
.ok_or(Error::Unsupported(
"offset fetch response omitted requested group",
))?;
(group.topics, group.error_code)
} else {
let offsets = coordinator
.offset_fetch_v2(group_id.to_owned(), Some(offset_fetch_topics(assignment)))
.await?;
(offsets.topics, offsets.error_code)
};
if offset_error_code != 0 {
return Err(
coordinator.broker_error(offset_error_code, format!("offset fetch group {group_id}"))
);
}
for topic in &assignment.assignments {
for partition in &topic.partitions {
let committed = committed_offset(group_id, &offset_topics, &topic.topic, *partition)
.map_err(|error| match error {
Error::Broker { code, context } => coordinator.broker_error(code, context),
error => error,
})?;
match committed.filter(|offset| *offset >= 0) {
Some(next_offset) => assignments.push(ConsumerAssignment::new(
topic.topic.clone(),
*partition,
next_offset,
)),
None => match offset_reset_policy {
OffsetResetPolicy::Offset(next_offset) => {
assignments.push(ConsumerAssignment::new(
topic.topic.clone(),
*partition,
next_offset,
));
}
OffsetResetPolicy::Earliest | OffsetResetPolicy::Latest => {
reset_partitions.push((topic.topic.clone(), *partition));
}
},
}
}
}
if let Some(timestamp) = offset_reset_policy.timestamp() {
assignments.extend(
resolve_reset_offsets(bootstrap, client_config, &reset_partitions, timestamp).await?,
);
}
for assignment in &mut assignments {
if let Some(leader_epoch) =
leader_epochs.get(&(assignment.topic().to_owned(), assignment.partition()))
{
assignment.set_leader_epoch(*leader_epoch);
}
}
assignments.sort_by(|left, right| {
left.topic()
.cmp(right.topic())
.then_with(|| left.partition().cmp(&right.partition()))
});
Ok(assignments)
}
async fn assignment_leader_epochs(
bootstrap: &mut Client,
assignment: &ConsumerProtocolAssignmentV0,
) -> Result<BTreeMap<(String, i32), i32>> {
if assignment.assignments.is_empty() || !bootstrap.supports_metadata_v12().await? {
return Ok(BTreeMap::new());
}
let requests = assignment
.assignments
.iter()
.map(|topic| MetadataRequestTopicV12 {
topic_id: [0; 16],
name: Some(topic.topic.clone()),
})
.collect();
let metadata = match bootstrap.metadata_v12(Some(requests)).await {
Ok(metadata) => metadata,
Err(error) => {
debug!(error = %error, "initial consumer-group leader epoch metadata unavailable");
return Ok(BTreeMap::new());
}
};
Ok(assignment_leader_epochs_from_metadata(
&metadata, assignment,
))
}
fn assignment_leader_epochs_from_metadata(
metadata: &MetadataResponseV12,
assignment: &ConsumerProtocolAssignmentV0,
) -> BTreeMap<(String, i32), i32> {
let assigned = assignment
.assignments
.iter()
.flat_map(|topic| {
topic
.partitions
.iter()
.map(|partition| (topic.topic.as_str(), *partition))
})
.collect::<BTreeSet<_>>();
let mut epochs = BTreeMap::new();
for topic in &metadata.topics {
let Some(topic_name) = topic.name.as_deref() else {
continue;
};
if topic.error_code != 0 {
continue;
}
for partition in &topic.partitions {
if partition.error_code == 0
&& partition.leader_epoch >= 0
&& assigned.contains(&(topic_name, partition.partition_index))
{
epochs.insert(
(topic_name.to_owned(), partition.partition_index),
partition.leader_epoch,
);
}
}
}
epochs
}
async fn resolve_reset_offsets(
bootstrap: &mut Client,
client_config: &ClientConfig,
partitions: &[(String, i32)],
timestamp: i64,
) -> Result<Vec<ConsumerAssignment>> {
if partitions.is_empty() {
return Ok(Vec::new());
}
let topic_names = partitions
.iter()
.map(|(topic, _)| topic.clone())
.collect::<BTreeSet<_>>()
.into_iter()
.collect();
let metadata = bootstrap.metadata(Some(topic_names)).await?;
let mut requests_by_broker = BTreeMap::<i32, BTreeMap<String, Vec<i32>>>::new();
for (topic, partition) in partitions {
let leader_id =
offset_leader_for(&metadata, topic, *partition).map_err(|error| match error {
Error::Broker { code, context } => client_config.broker_error(code, context),
error => error,
})?;
requests_by_broker
.entry(leader_id)
.or_default()
.entry(topic.clone())
.or_default()
.push(*partition);
}
let mut resolved = BTreeMap::new();
for (leader_id, topics) in requests_by_broker {
let broker_addr = offset_broker_addr_for(&metadata, leader_id)?;
let mut leader = client_config.connect_broker(broker_addr).await?;
let response = leader
.list_offsets_v1(list_offsets_topics(topics, timestamp))
.await?;
for (topic, partition) in partitions {
let partition_leader =
offset_leader_for(&metadata, topic, *partition).map_err(|error| match error {
Error::Broker { code, context } => client_config.broker_error(code, context),
error => error,
})?;
if partition_leader != leader_id {
continue;
}
let offset =
list_offset(&response, topic, *partition).map_err(|error| match error {
Error::Broker { code, context } => client_config.broker_error(code, context),
error => error,
})?;
resolved.insert((topic.clone(), *partition), offset);
}
}
Ok(resolved
.into_iter()
.map(|((topic, partition), offset)| ConsumerAssignment::new(topic, partition, offset))
.collect())
}
fn offset_leader_for(
metadata: &MetadataResponseV1,
topic_name: &str,
partition_index: i32,
) -> Result<i32> {
let topic = metadata
.topics
.iter()
.find(|topic| topic.name == topic_name)
.ok_or_else(|| Error::UnknownTopicOrPartition {
topic: topic_name.to_owned(),
partition: partition_index,
})?;
if topic.error_code != 0 {
return Err(Error::Broker {
code: topic.error_code,
context: format!("metadata topic {topic_name}"),
});
}
let partition = topic
.partitions
.iter()
.find(|partition| partition.partition_index == partition_index)
.ok_or_else(|| Error::UnknownTopicOrPartition {
topic: topic_name.to_owned(),
partition: partition_index,
})?;
if partition.error_code != 0 {
return Err(Error::Broker {
code: partition.error_code,
context: format!("metadata {topic_name}-{partition_index}"),
});
}
(partition.leader_id >= 0)
.then_some(partition.leader_id)
.ok_or_else(|| Error::MissingLeader {
topic: topic_name.to_owned(),
partition: partition_index,
})
}
fn offset_broker_addr_for(metadata: &MetadataResponseV1, node_id: i32) -> Result<String> {
metadata
.brokers
.iter()
.find(|broker| broker.node_id == node_id)
.map(|broker| format!("{}:{}", broker.host, broker.port))
.ok_or(Error::MissingBroker { node_id })
}
fn list_offsets_topics(
topics: BTreeMap<String, Vec<i32>>,
timestamp: i64,
) -> Vec<ListOffsetsTopicV1> {
topics
.into_iter()
.map(|(name, partitions)| ListOffsetsTopicV1 {
name,
partitions: partitions
.into_iter()
.map(|partition_index| ListOffsetsPartitionV1 {
partition_index,
timestamp,
})
.collect(),
})
.collect()
}
fn list_offset(
response: &ListOffsetsResponseV1,
topic_name: &str,
partition_index: i32,
) -> Result<i64> {
let partition = response
.topics
.iter()
.find(|topic| topic.name == topic_name)
.and_then(|topic| {
topic
.partitions
.iter()
.find(|partition| partition.partition_index == partition_index)
})
.ok_or_else(|| Error::UnknownTopicOrPartition {
topic: topic_name.to_owned(),
partition: partition_index,
})?;
if partition.error_code != 0 {
return Err(Error::Broker {
code: partition.error_code,
context: format!("list offsets {topic_name}-{partition_index}"),
});
}
Ok(partition.offset)
}
fn offset_fetch_topics(assignment: &ConsumerProtocolAssignmentV0) -> Vec<OffsetFetchTopic> {
assignment
.assignments
.iter()
.map(|topic| OffsetFetchTopic {
name: topic.topic.clone(),
partition_indexes: topic.partitions.clone(),
})
.collect()
}
fn offset_fetch_topics_v9(assignment: &ConsumerProtocolAssignmentV0) -> Vec<OffsetFetchTopicV9> {
offset_fetch_topics(assignment)
.into_iter()
.map(|topic| OffsetFetchTopicV9 {
name: topic.name,
partition_indexes: topic.partition_indexes,
})
.collect()
}
fn committed_offset(
group_id: &str,
topics: &[OffsetFetchTopicResponse],
topic_name: &str,
partition_index: i32,
) -> Result<Option<i64>> {
let Some(partition) = partition_response(topics, topic_name, partition_index) else {
return Ok(None);
};
if partition.error_code != 0 {
return Err(Error::Broker {
code: partition.error_code,
context: format!("offset fetch group {group_id} {topic_name}-{partition_index}"),
});
}
Ok(Some(partition.committed_offset))
}
fn partition_response<'a>(
topics: &'a [OffsetFetchTopicResponse],
topic_name: &str,
partition_index: i32,
) -> Option<&'a OffsetFetchPartitionResponse> {
topics
.iter()
.find(|topic| topic.name == topic_name)
.and_then(|topic| {
topic
.partitions
.iter()
.find(|partition| partition.partition_index == partition_index)
})
}
fn offset_commit_topics(assignments: &[ConsumerAssignment]) -> Vec<OffsetCommitTopic> {
let mut topics = BTreeMap::<String, Vec<OffsetCommitPartition>>::new();
for assignment in assignments {
topics
.entry(assignment.topic().to_owned())
.or_default()
.push(OffsetCommitPartition {
partition_index: assignment.partition(),
committed_offset: assignment.next_offset(),
committed_metadata: None,
});
}
topics
.into_iter()
.map(|(name, mut partitions)| {
partitions.sort_by_key(|partition| partition.partition_index);
OffsetCommitTopic { name, partitions }
})
.collect()
}
fn offset_commit_topics_v7(assignments: &[ConsumerAssignment]) -> Vec<OffsetCommitTopicV7> {
let mut topics = BTreeMap::<String, Vec<OffsetCommitPartitionV7>>::new();
for assignment in assignments {
topics
.entry(assignment.topic().to_owned())
.or_default()
.push(OffsetCommitPartitionV7 {
partition_index: assignment.partition(),
committed_offset: assignment.next_offset(),
committed_leader_epoch: assignment.leader_epoch(),
committed_metadata: None,
});
}
topics
.into_iter()
.map(|(name, mut partitions)| {
partitions.sort_by_key(|partition| partition.partition_index);
OffsetCommitTopicV7 { name, partitions }
})
.collect()
}
fn offset_commit_topics_v9(assignments: &[ConsumerAssignment]) -> Vec<OffsetCommitTopicV9> {
let mut topics = BTreeMap::<String, Vec<OffsetCommitPartitionV9>>::new();
for assignment in assignments {
topics
.entry(assignment.topic().to_owned())
.or_default()
.push(OffsetCommitPartitionV9 {
partition_index: assignment.partition(),
committed_offset: assignment.next_offset(),
committed_leader_epoch: assignment.leader_epoch(),
committed_metadata: None,
});
}
topics
.into_iter()
.map(|(name, mut partitions)| {
partitions.sort_by_key(|partition| partition.partition_index);
OffsetCommitTopicV9 { name, partitions }
})
.collect()
}
fn offset_commit_response_error(
group_id: &str,
topics: &[OffsetCommitTopicResponse],
) -> Option<Error> {
for topic in topics {
for partition in &topic.partitions {
if partition.error_code != 0 {
return Some(Error::Broker {
code: partition.error_code,
context: format!(
"offset commit group {group_id} {}-{}",
topic.name, partition.partition_index
),
});
}
}
}
None
}
fn leave_group_response_error(group_id: &str, response: &LeaveGroupResponseV3) -> Option<Error> {
if response.error_code != 0 {
return Some(Error::Broker {
code: response.error_code,
context: format!("leave group {group_id}"),
});
}
response
.members
.iter()
.find(|member| member.error_code != 0)
.map(|member| Error::Broker {
code: member.error_code,
context: format!("leave group {group_id} member {}", member.member_id),
})
}
fn should_rejoin_group(error: &Error) -> bool {
if matches!(error, Error::Io(_) | Error::RequestTimedOut { .. }) {
return true;
}
matches!(
error.broker_error_kind(),
Some(
BrokerErrorKind::CoordinatorLoadInProgress
| BrokerErrorKind::CoordinatorNotAvailable
| BrokerErrorKind::NotCoordinator
| BrokerErrorKind::IllegalGeneration
| BrokerErrorKind::UnknownMemberId
| BrokerErrorKind::RebalanceInProgress
| BrokerErrorKind::FencedMemberEpoch
| BrokerErrorKind::StaleMemberEpoch
)
)
}
fn should_retry_consumer_join_transport(error: &Error) -> bool {
matches!(error, Error::Io(_) | Error::RequestTimedOut { .. })
}
fn group_retry_backoff(retry_attempt: u32) -> Duration {
let exponent = retry_attempt.saturating_sub(1).min(5);
let multiplier = 1u64 << exponent;
let milliseconds = (GROUP_JOIN_RETRY_BACKOFF.as_millis() as u64).saturating_mul(multiplier);
Duration::from_millis(milliseconds).min(GROUP_JOIN_MAX_RETRY_BACKOFF)
}
fn member_id_after_join_error(error: &Error, requested_member_id: String) -> Option<String> {
if matches!(
error.broker_error_kind(),
Some(BrokerErrorKind::UnknownMemberId)
) {
None
} else {
Some(requested_member_id)
}
}
async fn should_rejoin_after_background_heartbeat(
heartbeat: &mut ConsumerGroupHeartbeat,
) -> Result<bool> {
match heartbeat.try_wait().await {
Ok(None | Some(())) => Ok(false),
Err(error) if should_rejoin_group(&error) => Ok(true),
Err(error) => Err(error),
}
}
fn validate_heartbeat_interval(interval: Duration) -> Result<()> {
if interval.is_zero() {
return Err(Error::InvalidConfiguration {
field: "heartbeat_interval",
reason: "must be greater than zero",
});
}
Ok(())
}
fn validate_commit_worker_interval(interval: Duration) -> Result<()> {
if interval.is_zero() {
return Err(Error::InvalidConfiguration {
field: "auto_commit_interval_ms",
reason: "must be greater than zero",
});
}
Ok(())
}
async fn run_background_commit_worker(
mut coordinator: Client,
link: CommitWorkerLink,
mut flush_rx: mpsc::Receiver<CommitWorkerCommand>,
client_config: ClientConfig,
group_id: String,
max_retries: u32,
interval: Duration,
) -> Result<()> {
let mut ticker = time::interval(interval);
ticker.set_missed_tick_behavior(MissedTickBehavior::Delay);
ticker.tick().await;
loop {
if link.shutdown.load(Ordering::Acquire) {
return Ok(());
}
tokio::select! {
_ = link.shutdown_notify.notified() => return Ok(()),
command = flush_rx.recv() => {
let Some(CommitWorkerCommand::Flush(ack)) = command else {
return Ok(());
};
match flush_commit_worker(
&mut coordinator,
&link,
&client_config,
&group_id,
max_retries,
)
.await
{
Ok(()) => {
let _ = ack.send(Ok(()));
}
Err(error) => {
let _ = ack.send(Err(Error::Unsupported(
"background commit worker stopped after a commit failure",
)));
link.signal_shutdown();
return Err(error);
}
}
}
_ = ticker.tick() => {
if let Err(error) = flush_commit_worker(
&mut coordinator,
&link,
&client_config,
&group_id,
max_retries,
)
.await
{
link.signal_shutdown();
return Err(error);
}
}
}
}
}
async fn flush_commit_worker(
coordinator: &mut Client,
link: &CommitWorkerLink,
client_config: &ClientConfig,
group_id: &str,
max_retries: u32,
) -> Result<()> {
let snapshot = link.lock_state()?.clone();
let assignments =
pending_commit_assignments(&snapshot.assignments, &snapshot.pending_commit_offsets);
if assignments.is_empty() {
return Ok(());
}
let mut retry_attempt = 0;
loop {
match commit_worker_request(coordinator, &snapshot, &assignments).await {
Ok(()) => {
link.clear_committed(&snapshot, &assignments)?;
return Ok(());
}
Err(error) if retry_attempt < max_retries && should_retry_commit_worker(&error) => {
retry_attempt += 1;
client_config.record_retry();
debug!(
group_id,
retry_attempt,
error = %error,
"retrying background kafka consumer group offset commit"
);
*coordinator =
connect_group_coordinator_with_retry(client_config, group_id, max_retries)
.await?;
time::sleep(group_retry_backoff(retry_attempt)).await;
}
Err(error) => return Err(error),
}
}
}
async fn commit_worker_request(
coordinator: &mut Client,
state: &CommitWorkerState,
assignments: &[ConsumerAssignment],
) -> Result<()> {
let response_topics = match state.protocol {
ConsumerGroupProtocol::Consumer => coordinator
.offset_commit_v9(
state.group_id.clone(),
state.generation_id,
state.member_id.clone(),
state.group_instance_id.clone(),
offset_commit_topics_v9(assignments),
)
.await
.map(|response| response.topics),
ConsumerGroupProtocol::Classic if state.group_instance_id.is_some() => coordinator
.offset_commit_v7(
state.group_id.clone(),
state.generation_id,
state.member_id.clone(),
state.group_instance_id.clone(),
offset_commit_topics_v7(assignments),
)
.await
.map(|response| response.topics),
ConsumerGroupProtocol::Classic => coordinator
.offset_commit_v2(
state.group_id.clone(),
state.generation_id,
state.member_id.clone(),
state.retention_time_ms,
offset_commit_topics(assignments),
)
.await
.map(|response| response.topics),
}?;
if let Some(error) = offset_commit_response_error(&state.group_id, &response_topics) {
return Err(match error {
Error::Broker { code, context } => coordinator.broker_error(code, context),
error => error,
});
}
Ok(())
}
fn should_retry_commit_worker(error: &Error) -> bool {
if matches!(error, Error::Io(_) | Error::RequestTimedOut { .. }) {
return true;
}
matches!(
error.broker_error_kind(),
Some(
BrokerErrorKind::CoordinatorLoadInProgress
| BrokerErrorKind::CoordinatorNotAvailable
| BrokerErrorKind::NotCoordinator
)
)
}
async fn run_background_heartbeat(
coordinator: &mut Client,
group_id: String,
generation_id: i32,
member_id: String,
group_instance_id: Option<String>,
interval: Duration,
mut shutdown: oneshot::Receiver<()>,
) -> Result<()> {
let mut heartbeat = time::interval(interval);
heartbeat.set_missed_tick_behavior(MissedTickBehavior::Delay);
heartbeat.tick().await;
loop {
tokio::select! {
_ = &mut shutdown => return Ok(()),
_ = heartbeat.tick() => {
debug!(
group_id = group_id.as_str(),
member_id = member_id.as_str(),
generation_id,
"sending background kafka consumer group heartbeat"
);
let response = if let Some(group_instance_id) = &group_instance_id {
coordinator
.heartbeat_v3(
group_id.clone(),
generation_id,
member_id.clone(),
Some(group_instance_id.clone()),
)
.await?
} else {
coordinator
.heartbeat_v2(group_id.clone(), generation_id, member_id.clone())
.await?
};
if response.error_code != 0 {
return Err(coordinator.broker_error(
response.error_code,
format!("background heartbeat group {group_id}"),
));
}
}
}
}
}
async fn run_background_consumer_heartbeat(
coordinator: &mut Client,
group_id: String,
state: Arc<Mutex<ConsumerProtocolHeartbeatState>>,
config: ConsumerProtocolHeartbeatConfig,
mut shutdown: oneshot::Receiver<()>,
) -> Result<()> {
let mut first_heartbeat = true;
loop {
if !first_heartbeat {
let interval = state.lock().await.heartbeat_interval;
tokio::select! {
_ = &mut shutdown => return Ok(()),
_ = time::sleep(interval) => {}
}
}
first_heartbeat = false;
let (member_id, member_epoch, owned_partitions) = {
let state = state.lock().await;
(
state.member_id.clone(),
state.member_epoch,
state.owned_partitions.clone(),
)
};
debug!(
group_id = group_id.as_str(),
member_id = member_id.as_str(),
member_epoch,
"sending background KIP-848 consumer group heartbeat"
);
let response = coordinator
.consumer_group_heartbeat_v0(
group_id.clone(),
member_id,
member_epoch,
config.group_instance_id.clone(),
None,
config.rebalance_timeout_ms,
Some(config.topics.clone()),
config.server_assignor.clone(),
owned_partitions,
)
.await?;
if response.error_code != 0 {
let message = response
.error_message
.as_deref()
.unwrap_or("broker returned a consumer-group heartbeat error");
return Err(coordinator.broker_error(
response.error_code,
format!("consumer group heartbeat {group_id}: {message}"),
));
}
record_consumer_heartbeat_response(&state, response).await;
}
}
async fn record_consumer_heartbeat_response(
state: &Arc<Mutex<ConsumerProtocolHeartbeatState>>,
response: ConsumerGroupHeartbeatResponseV0,
) {
let mut state = state.lock().await;
if let Some(member_id) = response.member_id {
state.member_id = member_id;
}
state.member_epoch = response.member_epoch;
if response.heartbeat_interval_ms > 0 {
state.heartbeat_interval = Duration::from_millis(
u64::try_from(response.heartbeat_interval_ms).unwrap_or(u64::MAX),
);
}
if let Some(assignment) = response.assignment {
state.owned_partitions = Some(assignment);
state.assignment_version = state.assignment_version.wrapping_add(1);
}
}
fn coordinator_addr(coordinator: &FindCoordinatorResponseV1) -> String {
format!("{}:{}", coordinator.host, coordinator.port)
}
async fn connect_group_coordinator_with_retry(
client: &ClientConfig,
group_id: &str,
max_retries: u32,
) -> Result<Client> {
let mut retry_attempt = 0;
loop {
let mut bootstrap = match client.clone().connect().await {
Ok(bootstrap) => bootstrap,
Err(error) if retry_attempt < max_retries && should_rejoin_group(&error) => {
retry_attempt += 1;
client.record_retry();
time::sleep(group_retry_backoff(retry_attempt)).await;
continue;
}
Err(error) => return Err(error),
};
let coordinator =
find_group_coordinator_with_retry(&mut bootstrap, client, group_id, max_retries)
.await?;
match client.connect_broker(coordinator_addr(&coordinator)).await {
Ok(coordinator) => return Ok(coordinator),
Err(error) if retry_attempt < max_retries && should_rejoin_group(&error) => {
retry_attempt += 1;
client.record_retry();
debug!(
group_id,
retry_attempt,
error = %error,
"retrying kafka consumer group coordinator connection"
);
time::sleep(group_retry_backoff(retry_attempt)).await;
}
Err(error) => return Err(error),
}
}
}
async fn find_group_coordinator_with_retry(
bootstrap: &mut Client,
client: &ClientConfig,
group_id: &str,
max_retries: u32,
) -> Result<FindCoordinatorResponseV1> {
let mut retry_attempt = 0;
loop {
let response = match bootstrap.find_group_coordinator(group_id.to_owned()).await {
Ok(response) => response,
Err(error) if retry_attempt < max_retries && should_rejoin_group(&error) => {
retry_attempt += 1;
client.record_retry();
time::sleep(group_retry_backoff(retry_attempt)).await;
continue;
}
Err(error) => return Err(error),
};
if response.error_code == 0 {
return Ok(response);
}
let error = client.broker_error(
response.error_code,
format!("find group coordinator {group_id}"),
);
if retry_attempt < max_retries && should_rejoin_group(&error) {
retry_attempt += 1;
client.record_retry();
debug!(
group_id,
retry_attempt,
error = %error,
"retrying kafka consumer group coordinator lookup"
);
time::sleep(group_retry_backoff(retry_attempt)).await;
continue;
}
return Err(error);
}
}
fn topic_ids_for_names(
metadata: &MetadataResponseV12,
topic_names: &[String],
) -> Result<BTreeMap<String, [u8; 16]>> {
let mut topic_ids = BTreeMap::new();
for topic_name in topic_names {
let topic = metadata
.topics
.iter()
.find(|topic| topic.name.as_deref() == Some(topic_name.as_str()))
.ok_or_else(|| Error::UnknownTopicOrPartition {
topic: topic_name.clone(),
partition: -1,
})?;
if topic.error_code != 0 {
return Err(Error::Broker {
code: topic.error_code,
context: format!("metadata topic {topic_name}"),
});
}
topic_ids.insert(topic_name.clone(), topic.topic_id);
}
Ok(topic_ids)
}
fn consumer_protocol_assignment(
assignment: &[ConsumerGroupHeartbeatTopicPartitions],
topic_ids: &BTreeMap<String, [u8; 16]>,
) -> Result<ConsumerProtocolAssignmentV0> {
let names_by_id = topic_ids
.iter()
.map(|(name, topic_id)| (*topic_id, name.clone()))
.collect::<BTreeMap<_, _>>();
let assignments = assignment
.iter()
.map(|topic| {
let topic_name =
names_by_id
.get(&topic.topic_id)
.ok_or_else(|| Error::UnknownTopicOrPartition {
topic: format!("topic-id-{:02x?}", topic.topic_id),
partition: -1,
})?;
Ok(ConsumerProtocolTopicAssignment {
topic: topic_name.clone(),
partitions: topic.partitions.clone(),
})
})
.collect::<Result<Vec<_>>>()?;
Ok(ConsumerProtocolAssignmentV0 {
assignments,
user_data: None,
})
}
fn consumer_group_heartbeat_error(
client: &ClientConfig,
group_id: &str,
response: &ConsumerGroupHeartbeatResponseV0,
) -> Error {
let message = response
.error_message
.as_deref()
.unwrap_or("broker returned a consumer-group heartbeat error");
client.broker_error(
response.error_code,
format!("consumer group heartbeat {group_id}: {message}"),
)
}
fn duration_millis(duration: std::time::Duration) -> u64 {
u64::try_from(duration.as_millis()).unwrap_or(u64::MAX)
}
fn partitions_for(metadata: &MetadataResponseV1, topic_name: &str) -> Result<Vec<i32>> {
let topic = metadata
.topics
.iter()
.find(|topic| topic.name == topic_name)
.ok_or_else(|| Error::UnknownTopicOrPartition {
topic: topic_name.to_owned(),
partition: -1,
})?;
let mut partitions = topic
.partitions
.iter()
.map(|partition| partition.partition_index)
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
partitions.sort();
Ok(partitions)
}
fn range_for_topic(members: &[String], partitions: &[i32]) -> Vec<(String, Vec<i32>)> {
if members.is_empty() {
return Vec::new();
}
let partition_count = partitions.len();
let member_count = members.len();
let partitions_per_member = partition_count / member_count;
let extra_partitions = partition_count % member_count;
members
.iter()
.enumerate()
.map(|(index, member)| {
let start = partitions_per_member * index + extra_partitions.min(index);
let length = partitions_per_member + usize::from(index < extra_partitions);
let end = start + length;
(member.clone(), partitions[start..end].to_vec())
})
.collect()
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use std::collections::BTreeMap;
use std::time::Duration;
use super::{
assignment_leader_epochs_from_metadata, assignments_for_strategy, committed_offset,
consumer_assignment_ready, cooperative_assignment_requires_rejoin,
cooperative_rejoin_required_for_assignment, cooperative_sticky_assignments,
decode_classic_subscription, decode_sticky_user_data, encode_sticky_user_data,
group_retry_backoff, leave_group_response_error, list_offset, list_offsets_topics,
member_id_after_join_error, offset_commit_response_error, offset_commit_topics,
offset_commit_topics_v7, offset_commit_topics_v9, offset_fetch_topics,
pending_commit_assignments, queue_commit_offset, range_assignments,
record_consumer_heartbeat_response, round_robin_assignments,
should_rejoin_after_background_heartbeat, should_rejoin_group, should_retry_commit_worker,
should_retry_consumer_join_transport, sticky_assignments, validate_commit_worker_interval,
validate_heartbeat_interval, CommitWorkerLink, CommitWorkerMembership, CommitWorkerState,
ConsumerGroupAssignmentStrategy, ConsumerGroupConfig, ConsumerGroupHeartbeat,
ConsumerGroupHeartbeatTopicPartitions, ConsumerGroupProtocol,
ConsumerProtocolHeartbeatState as ConsumerGroupHeartbeatState, HeartbeatHandleState,
IsolationLevel, OffsetResetPolicy, RebalanceEvent, RebalancePhase, SecurityProtocol,
DEFAULT_GROUP_MAX_RETRIES, GROUP_JOIN_MAX_RETRY_BACKOFF,
};
use crate::consumer::ConsumerAssignment;
use crate::Error;
use kafrust_protocol::api::consumer_group_heartbeat::ConsumerGroupHeartbeatResponseV0;
use kafrust_protocol::api::join_group::JoinGroupMember;
use kafrust_protocol::api::leave_group::{LeaveGroupMemberResponse, LeaveGroupResponseV3};
use kafrust_protocol::api::list_offsets::{
ListOffsetsPartitionResponseV1, ListOffsetsResponseV1, ListOffsetsTopicResponseV1,
EARLIEST_TIMESTAMP,
};
use kafrust_protocol::api::metadata::{
MetadataPartitionV12, MetadataResponseV1, MetadataResponseV12, MetadataTopicV12,
PartitionMetadata, TopicMetadata,
};
use kafrust_protocol::api::offset_commit::{
OffsetCommitPartitionResponse, OffsetCommitTopicResponse,
};
use kafrust_protocol::api::offset_fetch::{
OffsetFetchPartitionResponse, OffsetFetchTopicResponse,
};
use kafrust_protocol::codec::Encoder;
use kafrust_protocol::consumer_group::{
ConsumerProtocolAssignmentV0, ConsumerProtocolSubscriptionV0,
ConsumerProtocolSubscriptionV1, ConsumerProtocolTopicAssignment,
};
use std::sync::atomic::AtomicBool;
use std::sync::Arc;
use tokio::sync::{mpsc, Notify};
#[test]
fn builds_consumer_group_config() {
let config = ConsumerGroupConfig::new(["localhost:9092"], "orders-group")
.client_id("orders-reader")
.request_timeout_ms(5_000)
.security_protocol(SecurityProtocol::SaslPlaintext)
.tls_server_name("broker.example.com")
.tls_root_certificate_der([1, 2, 3])
.sasl_plain("alice", "secret-password")
.subscribe("orders")
.group_instance_id("orders-reader-1")
.session_timeout_ms(8_000)
.rebalance_timeout_ms(20_000)
.retention_time_ms(60_000)
.enable_auto_commit(true)
.auto_commit_interval_ms(250)
.start_offset(5)
.max_wait_ms(250)
.min_bytes(10)
.max_partition_bytes(1024)
.max_retries(3)
.max_poll_records(10)
.partition_queue_capacity(7)
.isolation_level(IsolationLevel::ReadCommitted)
.assignment_strategy(ConsumerGroupAssignmentStrategy::RoundRobin);
assert_eq!(config.group_id(), "orders-group");
assert_eq!(config.group_instance_id_ref(), Some("orders-reader-1"));
assert_eq!(config.topics(), &["orders".to_owned()]);
assert_eq!(
config.client.security_protocol_ref(),
SecurityProtocol::SaslPlaintext
);
assert_eq!(config.partition_queue_capacity_ref(), 7);
assert_eq!(
config.client.tls_server_name_ref(),
Some("broker.example.com")
);
assert_eq!(config.client.tls_root_certificates_der(), &[vec![1, 2, 3]]);
assert_eq!(
config.client.sasl_credentials_ref().unwrap().username(),
"alice"
);
assert_eq!(config.isolation_level_ref(), IsolationLevel::ReadCommitted);
assert!(config.auto_commit_enabled());
assert_eq!(config.auto_commit_interval(), Duration::from_millis(250));
assert_eq!(
config.assignment_strategy_ref(),
ConsumerGroupAssignmentStrategy::RoundRobin
);
assert_eq!(
config.offset_reset_policy_ref(),
OffsetResetPolicy::Offset(5)
);
}
#[test]
fn configures_regex_topic_subscription_and_validates_pattern() {
let config = ConsumerGroupConfig::new(["localhost:9092"], "orders-group")
.subscribe_pattern(r"^orders-[0-9]+$");
assert!(config.has_subscription());
assert!(config.topics().is_empty());
assert_eq!(config.topic_pattern_ref(), Some(r"^orders-[0-9]+$"));
let error = super::compile_topic_pattern("[").err();
assert!(matches!(error, Some(Error::InvalidTopicPattern { .. })));
}
#[test]
fn extracts_leader_epochs_for_assigned_partitions() {
let assignment = ConsumerProtocolAssignmentV0 {
assignments: vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![1],
}],
user_data: None,
};
let metadata = MetadataResponseV12 {
throttle_time_ms: 0,
brokers: Vec::new(),
cluster_id: None,
controller_id: 1,
topics: vec![MetadataTopicV12 {
error_code: 0,
name: Some("orders".to_owned()),
topic_id: [0; 16],
is_internal: false,
partitions: vec![
MetadataPartitionV12 {
error_code: 0,
partition_index: 0,
leader_id: 1,
leader_epoch: 4,
replica_nodes: vec![1],
isr_nodes: vec![1],
offline_replicas: Vec::new(),
},
MetadataPartitionV12 {
error_code: 0,
partition_index: 1,
leader_id: 2,
leader_epoch: 7,
replica_nodes: vec![2],
isr_nodes: vec![2],
offline_replicas: Vec::new(),
},
],
topic_authorized_operations: -1,
}],
};
let epochs = assignment_leader_epochs_from_metadata(&metadata, &assignment);
assert_eq!(epochs.get(&(String::from("orders"), 1)), Some(&7));
assert!(!epochs.contains_key(&(String::from("orders"), 0)));
}
#[tokio::test]
async fn rejects_invalid_group_configuration_before_connecting() {
let cases = [
(
ConsumerGroupConfig::new(["127.0.0.1:1"], "")
.subscribe("orders")
.join()
.await
.unwrap_err(),
"group_id",
),
(
ConsumerGroupConfig::new(["127.0.0.1:1"], "orders-group")
.join()
.await
.unwrap_err(),
"subscription",
),
(
ConsumerGroupConfig::new(["127.0.0.1:1"], "orders-group")
.subscribe("orders")
.session_timeout_ms(0)
.join()
.await
.unwrap_err(),
"session_timeout_ms",
),
(
ConsumerGroupConfig::new(["127.0.0.1:1"], "orders-group")
.subscribe("orders")
.max_poll_records(0)
.join()
.await
.unwrap_err(),
"max_poll_records",
),
];
for (error, field) in cases {
assert!(matches!(
error,
Error::InvalidConfiguration {
field: actual,
..
} if actual == field
));
}
assert!(ConsumerGroupConfig::new(["127.0.0.1:1"], "orders-group")
.subscribe("orders")
.validate()
.is_ok());
}
#[test]
fn public_group_validation_checks_client_and_worker_settings() {
assert!(matches!(
ConsumerGroupConfig::new(std::iter::empty::<&str>(), "orders-group")
.subscribe("orders")
.validate(),
Err(Error::MissingBootstrapServer)
));
assert!(matches!(
ConsumerGroupConfig::new(["localhost:9092"], "orders-group")
.subscribe("orders")
.enable_auto_commit(true)
.auto_commit_interval_ms(0)
.validate(),
Err(Error::InvalidConfiguration {
field: "auto_commit_interval_ms",
reason: "must be greater than zero"
})
));
}
#[test]
fn resolves_regex_topics_from_metadata_in_sorted_order() {
let mut inaccessible = topic_metadata("orders-3", &[0]);
inaccessible.error_code = 29;
let metadata = MetadataResponseV1 {
brokers: Vec::new(),
controller_id: 1,
topics: vec![
topic_metadata("orders-2", &[0]),
topic_metadata("payments", &[0]),
topic_metadata("orders-1", &[0]),
inaccessible,
],
};
let topics = super::resolve_topic_pattern(r"^orders-[0-9]+$", &metadata).ok();
assert_eq!(
topics,
Some(vec!["orders-1".to_owned(), "orders-2".to_owned()])
);
}
#[test]
fn rejects_regex_subscription_without_matching_topics() {
let metadata = metadata_fixture("payments", &[0]);
let error = super::resolve_topic_pattern(r"^orders-", &metadata).err();
assert!(matches!(
error,
Some(Error::Unsupported(
"consumer group topic pattern matched no topics"
))
));
}
#[test]
fn coalesces_queued_offsets_and_filters_unassigned_partitions() {
let mut pending = BTreeMap::new();
assert!(queue_commit_offset(&mut pending, "orders", 0, 4).is_ok());
assert!(queue_commit_offset(&mut pending, "orders", 0, 7).is_ok());
assert!(queue_commit_offset(&mut pending, "orders", 0, 5).is_ok());
let assignments = vec![
ConsumerAssignment::new("orders".to_owned(), 0, 8),
ConsumerAssignment::new("payments".to_owned(), 0, 3),
];
assert_eq!(
pending_commit_assignments(&assignments, &pending),
vec![ConsumerAssignment::new("orders".to_owned(), 0, 8)]
);
assert_eq!(pending.get(&("orders".to_owned(), 0)), Some(&8));
}
#[test]
fn rejects_queued_offset_overflow() {
let mut pending = BTreeMap::new();
let error = queue_commit_offset(&mut pending, "orders", 0, i64::MAX).err();
assert!(matches!(
error,
Some(Error::Unsupported("consumer record offset overflow"))
));
}
#[test]
fn commit_worker_state_coalesces_and_tracks_current_assignments() {
let (flush_tx, _flush_rx) = mpsc::channel(1);
let link = CommitWorkerLink {
state: Arc::new(std::sync::Mutex::new(CommitWorkerState {
group_id: "orders-group".to_owned(),
generation_id: 4,
member_id: "member-a".to_owned(),
group_instance_id: None,
protocol: ConsumerGroupProtocol::Classic,
retention_time_ms: -1,
assignments: vec![ConsumerAssignment::new("orders".to_owned(), 0, 0)],
pending_commit_offsets: BTreeMap::new(),
})),
flush_tx,
shutdown: Arc::new(AtomicBool::new(false)),
shutdown_notify: Arc::new(Notify::new()),
finished: Arc::new(AtomicBool::new(false)),
finished_notify: Arc::new(Notify::new()),
};
link.queue_record("orders", 0, 4).unwrap();
link.queue_record("orders", 0, 7).unwrap();
link.queue_record("orders", 0, 5).unwrap();
link.queue_record("payments", 0, 2).unwrap();
assert_eq!(link.pending_count(), 2);
assert_eq!(
link.snapshot_pending().unwrap()[&(String::from("orders"), 0)],
8
);
link.update_membership(
CommitWorkerMembership {
group_id: "orders-group".to_owned(),
generation_id: 5,
member_id: "member-b".to_owned(),
group_instance_id: None,
protocol: ConsumerGroupProtocol::Classic,
retention_time_ms: -1,
assignments: vec![ConsumerAssignment::new("payments".to_owned(), 0, 3)],
},
link.snapshot_pending().unwrap(),
)
.unwrap();
assert_eq!(link.pending_count(), 1);
assert_eq!(
link.snapshot_pending().unwrap()[&(String::from("payments"), 0)],
3
);
}
#[test]
fn classifies_commit_worker_retryable_errors() {
assert!(should_retry_commit_worker(&Error::Io(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
"connection reset",
))));
assert!(should_retry_commit_worker(&Error::RequestTimedOut {
timeout_ms: 1_000,
}));
assert!(should_retry_commit_worker(&Error::Broker {
code: 14,
context: "coordinator loading".to_owned(),
}));
assert!(should_retry_commit_worker(&Error::Broker {
code: 15,
context: "coordinator unavailable".to_owned(),
}));
assert!(should_retry_commit_worker(&Error::Broker {
code: 16,
context: "not coordinator".to_owned(),
}));
assert!(!should_retry_commit_worker(&Error::Broker {
code: 22,
context: "illegal generation".to_owned(),
}));
}
#[test]
fn rejects_zero_commit_worker_interval() {
assert!(matches!(
validate_commit_worker_interval(Duration::ZERO),
Err(Error::InvalidConfiguration {
field: "auto_commit_interval_ms",
reason: "must be greater than zero"
})
));
assert!(validate_commit_worker_interval(Duration::from_millis(1)).is_ok());
}
#[test]
fn configures_kip_848_consumer_protocol_and_server_assignor() {
let config = ConsumerGroupConfig::new(["localhost:9092"], "orders-group")
.group_protocol(ConsumerGroupProtocol::Consumer)
.server_assignor("uniform")
.subscribe("orders");
assert_eq!(config.group_protocol_ref(), ConsumerGroupProtocol::Consumer);
assert_eq!(config.server_assignor_ref(), Some("uniform"));
}
#[test]
fn waits_for_non_empty_kip_848_assignment_when_partitions_exist() {
let empty = [ConsumerGroupHeartbeatTopicPartitions {
topic_id: [1; 16],
partitions: Vec::new(),
}];
let assigned = [ConsumerGroupHeartbeatTopicPartitions {
topic_id: [1; 16],
partitions: vec![0],
}];
assert!(!consumer_assignment_ready(Some(&empty), true));
assert!(consumer_assignment_ready(Some(&empty), false));
assert!(consumer_assignment_ready(Some(&assigned), true));
assert!(!consumer_assignment_ready(None, true));
}
#[test]
fn configures_and_invokes_rebalance_listener() {
use std::sync::atomic::{AtomicUsize, Ordering};
let calls = Arc::new(AtomicUsize::new(0));
let callback_calls = calls.clone();
let config = ConsumerGroupConfig::new(["localhost:9092"], "orders-group")
.rebalance_listener(move |event| {
assert_eq!(event.phase(), RebalancePhase::After);
assert_eq!(event.group_id(), "orders-group");
assert_eq!(event.member_id(), "member-a");
assert_eq!(event.generation_id(), 4);
assert_eq!(event.protocol(), ConsumerGroupProtocol::Classic);
assert!(event.assignments().is_empty());
callback_calls.fetch_add(1, Ordering::SeqCst);
});
assert!(config.has_rebalance_listener());
config
.rebalance_listener
.as_ref()
.unwrap()
.on_rebalance(RebalanceEvent {
phase: RebalancePhase::After,
group_id: "orders-group",
member_id: "member-a",
generation_id: 4,
protocol: ConsumerGroupProtocol::Classic,
assignments: &[],
});
assert_eq!(calls.load(Ordering::SeqCst), 1);
}
#[test]
fn configures_offset_reset_policy_and_preserves_offset_zero_default() {
let default = ConsumerGroupConfig::new(["localhost:9092"], "orders-group");
assert_eq!(
default.offset_reset_policy_ref(),
OffsetResetPolicy::Offset(0)
);
assert_eq!(default.max_retries, DEFAULT_GROUP_MAX_RETRIES);
assert!(!default.auto_commit_enabled());
assert_eq!(default.auto_commit_interval(), Duration::from_secs(5));
let latest = default
.start_offset(12)
.offset_reset_policy(OffsetResetPolicy::Latest);
assert_eq!(latest.offset_reset_policy_ref(), OffsetResetPolicy::Latest);
}
#[tokio::test]
async fn rejects_empty_static_group_instance_id_before_connecting() {
let error = ConsumerGroupConfig::new(["localhost:9092"], "orders-group")
.group_instance_id("")
.subscribe("orders")
.join()
.await
.unwrap_err();
assert!(matches!(error, Error::InvalidGroupInstanceId));
}
#[tokio::test]
async fn rejects_zero_auto_commit_interval_before_connecting() {
let error = ConsumerGroupConfig::new(["localhost:9092"], "orders-group")
.subscribe("orders")
.enable_auto_commit(true)
.auto_commit_interval_ms(0)
.join()
.await
.unwrap_err();
assert!(matches!(
error,
Error::InvalidConfiguration {
field: "auto_commit_interval_ms",
reason: "must be greater than zero"
}
));
}
#[tokio::test]
async fn rejects_classic_assignment_strategy_for_kip_848_before_connecting() {
let error = ConsumerGroupConfig::new(["localhost:9092"], "orders-group")
.group_protocol(ConsumerGroupProtocol::Consumer)
.assignment_strategy(ConsumerGroupAssignmentStrategy::RoundRobin)
.subscribe("orders")
.join()
.await
.unwrap_err();
assert!(matches!(
error,
Error::Unsupported(
"KIP-848 groups use a broker-side assignor; configure ConsumerGroupAssignmentStrategy::Range or select the classic protocol"
)
));
}
#[test]
fn group_fetch_consumer_preserves_client_security_config() {
let config = ConsumerGroupConfig::new(["localhost:9092"], "orders-group")
.client_id("orders-reader")
.request_timeout_ms(5_000)
.security_protocol(SecurityProtocol::SaslPlaintext)
.tls_server_name("broker.example.com")
.tls_root_certificate_der([1, 2, 3])
.sasl_plain("alice", "secret-password")
.max_retries(3)
.max_poll_records(10)
.isolation_level(IsolationLevel::ReadCommitted);
let consumer_config = config.consumer_config();
let client_config = consumer_config.client_config();
assert_eq!(client_config.client_id_ref(), Some("orders-reader"));
assert_eq!(client_config.request_timeout().as_millis(), 5_000);
assert_eq!(
client_config.security_protocol_ref(),
SecurityProtocol::SaslPlaintext
);
assert_eq!(
client_config.tls_server_name_ref(),
Some("broker.example.com")
);
assert_eq!(client_config.tls_root_certificates_der(), &[vec![1, 2, 3]]);
assert_eq!(
client_config.sasl_credentials_ref().unwrap().username(),
"alice"
);
assert_eq!(consumer_config.max_retries_ref(), 3);
assert_eq!(consumer_config.max_poll_records_ref(), 10);
assert_eq!(
consumer_config.isolation_level_ref(),
IsolationLevel::ReadCommitted
);
let recovery_config = config.offset_reset_policy(OffsetResetPolicy::Earliest);
assert_eq!(
recovery_config.consumer_config().offset_reset_policy_ref(),
OffsetResetPolicy::Earliest
);
}
#[test]
fn assigns_partitions_with_range_strategy() {
let members = vec![
member("member-b", &["orders"]),
member("member-a", &["orders"]),
];
let metadata = metadata_fixture("orders", &[0, 1, 2]);
let assignments = range_assignments(&members, &metadata).unwrap();
let member_a = decode_assignment(&assignments[0].assignment);
let member_b = decode_assignment(&assignments[1].assignment);
assert_eq!(assignments[0].member_id, "member-a");
assert_eq!(member_a.assignments[0].topic, "orders");
assert_eq!(member_a.assignments[0].partitions, vec![0, 1]);
assert_eq!(assignments[1].member_id, "member-b");
assert_eq!(member_b.assignments[0].topic, "orders");
assert_eq!(member_b.assignments[0].partitions, vec![2]);
}
#[test]
fn assigns_only_subscribed_members_per_topic() {
let members = vec![
member("member-a", &["orders", "payments"]),
member("member-b", &["orders"]),
];
let mut metadata = metadata_fixture("orders", &[0, 1]);
metadata.topics.push(topic_metadata("payments", &[0, 1, 2]));
let assignments = range_assignments(&members, &metadata).unwrap();
let member_a = decode_assignment(&assignments[0].assignment);
assert_eq!(member_a.assignments[0].topic, "orders");
assert_eq!(member_a.assignments[0].partitions, vec![0]);
assert_eq!(member_a.assignments[1].topic, "payments");
assert_eq!(member_a.assignments[1].partitions, vec![0, 1, 2]);
}
#[test]
fn assigns_partitions_with_round_robin_strategy() {
let members = vec![
member("member-b", &["orders", "payments"]),
member("member-a", &["orders", "payments"]),
];
let mut metadata = metadata_fixture("orders", &[0, 1, 2]);
metadata.topics.push(topic_metadata("payments", &[0, 1]));
let assignments = round_robin_assignments(&members, &metadata).unwrap();
let member_a = decode_assignment(&assignments[0].assignment);
let member_b = decode_assignment(&assignments[1].assignment);
assert_eq!(member_a.assignments[0].partitions, vec![0, 2]);
assert_eq!(member_a.assignments[1].partitions, vec![1]);
assert_eq!(member_b.assignments[0].partitions, vec![1]);
assert_eq!(member_b.assignments[1].partitions, vec![0]);
}
#[test]
fn assigns_partitions_with_eager_sticky_strategy() {
let members = vec![
sticky_member(
"member-a",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0, 1],
}],
Some(4),
),
sticky_member(
"member-b",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![2],
}],
Some(4),
),
sticky_member("member-c", &["orders"], Vec::new(), Some(4)),
];
let metadata = metadata_fixture("orders", &[0, 1, 2]);
let assignments = sticky_assignments(&members, &metadata).unwrap();
assert_eq!(assignments[0].member_id, "member-a");
assert_eq!(assignments[1].member_id, "member-b");
assert_eq!(assignments[2].member_id, "member-c");
let counts = assignments
.iter()
.map(|assignment| {
decode_assignment(&assignment.assignment).assignments[0]
.partitions
.len()
})
.collect::<Vec<_>>();
assert_eq!(counts, vec![1, 1, 1]);
}
#[test]
fn eager_sticky_assignment_transfers_partitions_immediately() {
let members = vec![
sticky_member(
"member-a",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0, 1],
}],
Some(8),
),
sticky_member(
"member-c",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![2, 3],
}],
Some(8),
),
];
let metadata = metadata_fixture("orders", &[0, 1, 2, 3]);
let assignments = sticky_assignments(&members, &metadata).unwrap();
let member_a = decode_assignment(&assignments[0].assignment);
let member_c = decode_assignment(&assignments[1].assignment);
assert_eq!(member_a.assignments[0].partitions, vec![0, 1]);
assert_eq!(member_c.assignments[0].partitions, vec![2, 3]);
}
#[test]
fn sticky_invalidates_duplicate_previous_owners_in_the_same_generation() {
let members = vec![
sticky_member(
"member-a",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0, 1],
}],
Some(9),
),
sticky_member(
"member-b",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0, 2],
}],
Some(9),
),
sticky_member("member-c", &["orders"], Vec::new(), Some(9)),
];
let metadata = metadata_fixture("orders", &[0, 1, 2, 3]);
let assignments = sticky_assignments(&members, &metadata).unwrap();
let member_a = decode_assignment(&assignments[0].assignment);
let member_b = decode_assignment(&assignments[1].assignment);
let member_c = decode_assignment(&assignments[2].assignment);
assert_eq!(member_a.assignments[0].partitions, vec![1, 3]);
assert_eq!(member_b.assignments[0].partitions, vec![2]);
assert_eq!(member_c.assignments[0].partitions, vec![0]);
}
#[test]
fn sticky_balances_mixed_topic_subscriptions_by_topic_candidate_count() {
let members = vec![
sticky_member("member-a", &["orders"], Vec::new(), Some(1)),
sticky_member("member-b", &["orders", "payments"], Vec::new(), Some(1)),
sticky_member("member-c", &["orders"], Vec::new(), Some(1)),
];
let mut metadata = metadata_fixture("orders", &[0, 1, 2]);
metadata.topics.push(topic_metadata("payments", &[0, 1]));
let assignments = sticky_assignments(&members, &metadata).unwrap();
let member_a = decode_assignment(&assignments[0].assignment);
let member_b = decode_assignment(&assignments[1].assignment);
let member_c = decode_assignment(&assignments[2].assignment);
assert_eq!(member_a.assignments[0].topic, "orders");
assert_eq!(member_a.assignments[0].partitions, vec![0, 2]);
assert_eq!(member_b.assignments[0].topic, "payments");
assert_eq!(member_b.assignments[0].partitions, vec![0, 1]);
assert_eq!(member_c.assignments[0].topic, "orders");
assert_eq!(member_c.assignments[0].partitions, vec![1]);
}
#[test]
fn sticky_user_data_round_trips_with_optional_generation() {
let owned = vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![1, 3],
}];
let encoded = encode_sticky_user_data(&owned, Some(12)).unwrap();
let decoded = decode_sticky_user_data(Some(&encoded));
assert_eq!(decoded.partitions, owned);
assert_eq!(decoded.generation, Some(12));
assert_eq!(
encode_sticky_user_data(&owned, None).unwrap(),
[
0, 0, 0, 1, 0, 6, b'o', b'r', b'd', b'e', b'r', b's', 0, 0, 0, 2, 0, 0, 0, 1, 0, 0, 0, 3, ]
);
}
#[test]
fn decodes_classic_subscription_versions_with_append_only_fields() {
let mut encoder = Encoder::new();
encoder.write_i16(3);
encoder
.write_array(Some(&["orders".to_owned()]), |encoder, topic| {
encoder.write_string(topic)
})
.unwrap();
encoder.write_nullable_bytes(Some(&[7, 8])).unwrap();
encoder
.write_array(
Some(&[ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![1, 3],
}]),
|encoder, assignment| {
encoder.write_string(&assignment.topic)?;
encoder.write_array(
Some(assignment.partitions.as_slice()),
|encoder, partition| {
encoder.write_i32(*partition);
Ok(())
},
)
},
)
.unwrap();
encoder.write_i32(12);
encoder.write_nullable_string(Some("rack-a")).unwrap();
let decoded = decode_classic_subscription(&encoder.into_bytes()).unwrap();
assert_eq!(decoded.topics, vec!["orders"]);
assert_eq!(decoded.user_data, Some(vec![7, 8]));
assert_eq!(decoded.owned_partitions[0].partitions, vec![1, 3]);
assert_eq!(decoded.generation, Some(12));
}
#[test]
fn assigns_cooperatively_and_preserves_existing_ownership() {
let members = vec![
cooperative_member(
"member-a",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0, 1],
}],
),
cooperative_member(
"member-b",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![2],
}],
),
cooperative_member("member-c", &["orders"], Vec::new()),
];
let metadata = metadata_fixture("orders", &[0, 1, 2]);
let assignments = cooperative_sticky_assignments(&members, &metadata).unwrap();
let member_a = decode_assignment(&assignments[0].assignment);
let member_b = decode_assignment(&assignments[1].assignment);
let member_c = decode_assignment(&assignments[2].assignment);
assert_eq!(assignments[0].member_id, "member-a");
assert_eq!(member_a.assignments[0].partitions, vec![1]);
assert_eq!(member_b.assignments[0].partitions, vec![2]);
assert!(member_c.assignments.is_empty());
}
#[test]
fn assigns_cooperative_new_members_without_immediate_partition_transfer() {
let members = vec![
cooperative_member("member-a", &["orders"], Vec::new()),
cooperative_member("member-b", &["orders"], Vec::new()),
];
let metadata = metadata_fixture("orders", &[0, 1, 2, 3]);
let assignments = cooperative_sticky_assignments(&members, &metadata).unwrap();
let member_a = decode_assignment(&assignments[0].assignment);
let member_b = decode_assignment(&assignments[1].assignment);
assert_eq!(member_a.assignments[0].partitions, vec![0, 2]);
assert_eq!(member_b.assignments[0].partitions, vec![1, 3]);
}
#[test]
fn requests_cooperative_rejoin_for_staged_transfer() {
let members = vec![
cooperative_member(
"member-a",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0, 1, 2, 3],
}],
),
cooperative_member("member-b", &["orders"], Vec::new()),
];
let assignment = ConsumerProtocolAssignmentV0 {
assignments: vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0, 2],
}],
user_data: None,
};
assert!(cooperative_assignment_requires_rejoin(&members, &assignment).unwrap());
}
#[test]
fn does_not_request_cooperative_rejoin_after_member_loss() {
let members = vec![cooperative_member(
"member-a",
&["orders"],
vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0],
}],
)];
let assignment = ConsumerProtocolAssignmentV0 {
assignments: vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0, 1],
}],
user_data: None,
};
assert!(!cooperative_assignment_requires_rejoin(&members, &assignment).unwrap());
}
#[test]
fn requests_one_rejoin_for_new_non_leader_with_empty_assignment() {
let assignment = ConsumerProtocolAssignmentV0 {
assignments: Vec::new(),
user_data: None,
};
assert!(cooperative_rejoin_required_for_assignment(&[], &[], &assignment, false,).unwrap());
assert!(!cooperative_rejoin_required_for_assignment(&[], &[], &assignment, true,).unwrap());
}
#[test]
fn round_robin_skips_members_not_subscribed_to_a_topic() {
let members = vec![
member("member-a", &["orders"]),
member("member-b", &["payments"]),
member("member-c", &["orders", "payments"]),
];
let mut metadata = metadata_fixture("orders", &[0, 1]);
metadata.topics.push(topic_metadata("payments", &[0, 1]));
let assignments = round_robin_assignments(&members, &metadata).unwrap();
let member_a = decode_assignment(&assignments[0].assignment);
let member_b = decode_assignment(&assignments[1].assignment);
let member_c = decode_assignment(&assignments[2].assignment);
assert_eq!(member_a.assignments[0].partitions, vec![0]);
assert_eq!(member_b.assignments[0].partitions, vec![0]);
assert_eq!(member_c.assignments[0].partitions, vec![1]);
assert_eq!(member_c.assignments[1].partitions, vec![1]);
}
#[test]
fn rejects_unknown_selected_assignment_strategy() {
let error = assignments_for_strategy("unsupported", &[], &metadata_fixture("orders", &[0]))
.unwrap_err();
assert!(matches!(
error,
Error::Unsupported("consumer group selected an unsupported assignment strategy")
));
}
#[test]
fn builds_offset_fetch_topics_from_assignment() {
let assignment = ConsumerProtocolAssignmentV0 {
assignments: vec![ConsumerProtocolTopicAssignment {
topic: "orders".to_owned(),
partitions: vec![0, 2],
}],
user_data: None,
};
let topics = offset_fetch_topics(&assignment);
assert_eq!(topics.len(), 1);
assert_eq!(topics[0].name, "orders");
assert_eq!(topics[0].partition_indexes, vec![0, 2]);
}
#[test]
fn reads_committed_offset_and_surfaces_partition_error() {
let topics = vec![OffsetFetchTopicResponse {
name: "orders".to_owned(),
partitions: vec![
OffsetFetchPartitionResponse {
partition_index: 0,
committed_offset: 42,
metadata: None,
error_code: 0,
},
OffsetFetchPartitionResponse {
partition_index: 1,
committed_offset: -1,
metadata: None,
error_code: 25,
},
],
}];
assert_eq!(
committed_offset("orders-group", &topics, "orders", 0).unwrap(),
Some(42)
);
assert_eq!(
committed_offset("orders-group", &topics, "orders", 2).unwrap(),
None
);
assert!(matches!(
committed_offset("orders-group", &topics, "orders", 1).unwrap_err(),
Error::Broker { code: 25, .. }
));
}
#[test]
fn builds_list_offsets_topics_and_reads_partition_offset() {
let topics = BTreeMap::from([
("orders".to_owned(), vec![0, 2]),
("payments".to_owned(), vec![1]),
]);
let request = list_offsets_topics(topics, EARLIEST_TIMESTAMP);
assert_eq!(request[0].name, "orders");
assert_eq!(request[0].partitions[0].partition_index, 0);
assert_eq!(request[0].partitions[0].timestamp, EARLIEST_TIMESTAMP);
assert_eq!(request[1].name, "payments");
let response = ListOffsetsResponseV1 {
topics: vec![ListOffsetsTopicResponseV1 {
name: "orders".to_owned(),
partitions: vec![ListOffsetsPartitionResponseV1 {
partition_index: 2,
error_code: 0,
timestamp: -1,
offset: 41,
}],
}],
};
assert_eq!(list_offset(&response, "orders", 2).unwrap(), 41);
assert!(matches!(
list_offset(&response, "orders", 0).unwrap_err(),
Error::UnknownTopicOrPartition { .. }
));
}
#[test]
fn surfaces_list_offsets_partition_error() {
let response = ListOffsetsResponseV1 {
topics: vec![ListOffsetsTopicResponseV1 {
name: "orders".to_owned(),
partitions: vec![ListOffsetsPartitionResponseV1 {
partition_index: 0,
error_code: 3,
timestamp: -1,
offset: -1,
}],
}],
};
assert!(matches!(
list_offset(&response, "orders", 0).unwrap_err(),
Error::Broker {
code: 3,
context
} if context == "list offsets orders-0"
));
}
#[test]
fn builds_offset_commit_topics_from_current_assignment_offsets() {
let mut assignments = vec![
ConsumerAssignment::new("orders".to_owned(), 1, 43),
ConsumerAssignment::new("orders".to_owned(), 0, 11),
ConsumerAssignment::new("payments".to_owned(), 0, 7),
];
assignments[1].set_leader_epoch(4);
let topics = offset_commit_topics(&assignments);
assert_eq!(topics[0].name, "orders");
assert_eq!(topics[0].partitions[0].partition_index, 0);
assert_eq!(topics[0].partitions[0].committed_offset, 11);
assert_eq!(topics[0].partitions[1].partition_index, 1);
let static_topics = offset_commit_topics_v7(&assignments);
assert_eq!(static_topics[0].partitions[0].committed_offset, 11);
assert_eq!(static_topics[0].partitions[0].committed_leader_epoch, 4);
let consumer_protocol_topics = offset_commit_topics_v9(&assignments);
assert_eq!(
consumer_protocol_topics[0].partitions[0].committed_leader_epoch,
4
);
assert_eq!(topics[1].name, "payments");
}
#[test]
fn surfaces_offset_commit_partition_error() {
let response = vec![OffsetCommitTopicResponse {
name: "orders".to_owned(),
partitions: vec![OffsetCommitPartitionResponse {
partition_index: 0,
error_code: 25,
}],
}];
let error = offset_commit_response_error("orders-group", &response).unwrap();
assert!(matches!(error, Error::Broker { code: 25, .. }));
}
#[test]
fn surfaces_leave_group_member_error() {
let response = LeaveGroupResponseV3 {
throttle_time_ms: 0,
error_code: 0,
members: vec![LeaveGroupMemberResponse {
member_id: "member-a".to_owned(),
group_instance_id: Some("orders-reader-1".to_owned()),
error_code: 82,
}],
};
let error = leave_group_response_error("orders-group", &response).unwrap();
assert!(matches!(error, Error::Broker { code: 82, .. }));
}
#[test]
fn classifies_offset_commit_rejoin_error() {
let response = vec![OffsetCommitTopicResponse {
name: "orders".to_owned(),
partitions: vec![OffsetCommitPartitionResponse {
partition_index: 0,
error_code: 27,
}],
}];
let error = offset_commit_response_error("orders-group", &response).unwrap();
assert!(should_rejoin_group(&error));
assert!(matches!(error, Error::Broker { code: 27, .. }));
}
#[test]
fn classifies_offset_commit_non_rejoin_error() {
let response = vec![OffsetCommitTopicResponse {
name: "orders".to_owned(),
partitions: vec![OffsetCommitPartitionResponse {
partition_index: 0,
error_code: 7,
}],
}];
let error = offset_commit_response_error("orders-group", &response).unwrap();
assert!(!should_rejoin_group(&error));
assert!(matches!(error, Error::Broker { code: 7, .. }));
}
#[test]
fn classifies_group_errors_that_require_rejoin() {
for code in [14, 15, 16, 22, 25, 27] {
assert!(should_rejoin_group(&Error::Broker {
code,
context: "heartbeat group orders-group".to_owned(),
}));
}
assert!(!should_rejoin_group(&Error::Broker {
code: 7,
context: "heartbeat group orders-group".to_owned(),
}));
assert!(!should_rejoin_group(&Error::Broker {
code: 82,
context: "heartbeat group orders-group".to_owned(),
}));
assert!(should_rejoin_group(&Error::RequestTimedOut {
timeout_ms: 5
}));
assert!(should_rejoin_group(&Error::Io(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
"reset",
))));
assert!(!should_rejoin_group(&Error::Unsupported("group feature")));
}
#[test]
fn retries_kip_848_join_transport_failures() {
assert!(should_retry_consumer_join_transport(
&Error::RequestTimedOut { timeout_ms: 1_000 }
));
assert!(should_retry_consumer_join_transport(&Error::Io(
std::io::Error::new(std::io::ErrorKind::ConnectionReset, "reset"),
)));
assert!(!should_retry_consumer_join_transport(&Error::Broker {
code: 15,
context: "consumer group coordinator".to_owned(),
}));
}
#[test]
fn group_retry_backoff_is_bounded_and_exponential() {
assert_eq!(group_retry_backoff(1), Duration::from_millis(50));
assert_eq!(group_retry_backoff(2), Duration::from_millis(100));
assert_eq!(group_retry_backoff(5), Duration::from_millis(800));
assert_eq!(group_retry_backoff(6), GROUP_JOIN_MAX_RETRY_BACKOFF);
assert_eq!(group_retry_backoff(u32::MAX), GROUP_JOIN_MAX_RETRY_BACKOFF);
}
#[test]
fn clears_unknown_member_id_before_join_retry() {
let error = Error::Broker {
code: 25,
context: "join group orders-group".to_owned(),
};
assert_eq!(
member_id_after_join_error(&error, "member-a".to_owned()),
None
);
}
#[test]
fn preserves_member_id_for_other_transient_join_errors() {
let error = Error::Broker {
code: 27,
context: "join group orders-group".to_owned(),
};
assert_eq!(
member_id_after_join_error(&error, "member-a".to_owned()),
Some("member-a".to_owned())
);
}
#[test]
fn rejects_zero_background_heartbeat_interval() {
assert!(matches!(
validate_heartbeat_interval(std::time::Duration::ZERO).unwrap_err(),
Error::InvalidConfiguration {
field: "heartbeat_interval",
reason: "must be greater than zero"
}
));
validate_heartbeat_interval(std::time::Duration::from_millis(1)).unwrap();
}
#[tokio::test]
async fn kip_848_background_state_preserves_null_assignment_and_tracks_updates() {
let initial_assignment = vec![ConsumerGroupHeartbeatTopicPartitions {
topic_id: [1; 16],
partitions: vec![0, 1],
}];
let state = std::sync::Arc::new(tokio::sync::Mutex::new(ConsumerGroupHeartbeatState {
member_id: "member-a".to_owned(),
member_epoch: 4,
owned_partitions: Some(initial_assignment.clone()),
assignment_version: 0,
heartbeat_interval: std::time::Duration::from_secs(1),
}));
record_consumer_heartbeat_response(
&state,
ConsumerGroupHeartbeatResponseV0 {
throttle_time_ms: 0,
error_code: 0,
error_message: None,
member_id: None,
member_epoch: 5,
heartbeat_interval_ms: 2_500,
assignment: None,
},
)
.await;
{
let state = state.lock().await;
assert_eq!(state.member_epoch, 5);
assert_eq!(state.owned_partitions, Some(initial_assignment));
assert_eq!(state.assignment_version, 0);
assert_eq!(
state.heartbeat_interval,
std::time::Duration::from_millis(2_500)
);
}
let updated_assignment = vec![ConsumerGroupHeartbeatTopicPartitions {
topic_id: [2; 16],
partitions: vec![2],
}];
record_consumer_heartbeat_response(
&state,
ConsumerGroupHeartbeatResponseV0 {
throttle_time_ms: 0,
error_code: 0,
error_message: None,
member_id: Some("member-b".to_owned()),
member_epoch: 6,
heartbeat_interval_ms: 2_500,
assignment: Some(updated_assignment.clone()),
},
)
.await;
let state = state.lock().await;
assert_eq!(state.member_id, "member-b");
assert_eq!(state.member_epoch, 6);
assert_eq!(state.owned_partitions, Some(updated_assignment));
assert_eq!(state.assignment_version, 1);
}
#[tokio::test]
async fn try_wait_reports_running_heartbeat_task() {
let (shutdown, shutdown_rx) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async move {
let _ = shutdown_rx.await;
Ok(())
});
let mut heartbeat = heartbeat_handle(shutdown, handle, "orders-group", "member-a", 1);
assert!(!heartbeat.is_finished());
assert_eq!(heartbeat.try_wait().await.unwrap(), None);
heartbeat.stop().await.unwrap();
}
#[tokio::test]
async fn try_wait_surfaces_finished_heartbeat_error() {
let (shutdown, _shutdown_rx) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async {
Err(Error::Broker {
code: 27,
context: "background heartbeat group orders-group".to_owned(),
})
});
let mut heartbeat = heartbeat_handle(shutdown, handle, "orders-group", "member-a", 1);
while !heartbeat.is_finished() {
tokio::task::yield_now().await;
}
assert!(matches!(
heartbeat.try_wait().await,
Err(Error::Broker { code: 27, .. })
));
assert!(heartbeat.is_finished());
}
#[tokio::test]
async fn background_heartbeat_observation_ignores_running_task() {
let (shutdown, shutdown_rx) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async move {
let _ = shutdown_rx.await;
Ok(())
});
let mut heartbeat = heartbeat_handle(shutdown, handle, "orders-group", "member-a", 1);
assert!(!should_rejoin_after_background_heartbeat(&mut heartbeat)
.await
.unwrap());
heartbeat.stop().await.unwrap();
}
#[tokio::test]
async fn background_heartbeat_observation_requests_rejoin_for_group_error() {
let (shutdown, _shutdown_rx) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async {
Err(Error::Broker {
code: 27,
context: "background heartbeat group orders-group".to_owned(),
})
});
let mut heartbeat = heartbeat_handle(shutdown, handle, "orders-group", "member-a", 1);
while !heartbeat.is_finished() {
tokio::task::yield_now().await;
}
assert!(should_rejoin_after_background_heartbeat(&mut heartbeat)
.await
.unwrap());
assert!(heartbeat.is_finished());
}
#[tokio::test]
async fn background_heartbeat_observation_surfaces_non_rejoin_error() {
let (shutdown, _shutdown_rx) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async {
Err(Error::Broker {
code: 7,
context: "background heartbeat group orders-group".to_owned(),
})
});
let mut heartbeat = heartbeat_handle(shutdown, handle, "orders-group", "member-a", 1);
while !heartbeat.is_finished() {
tokio::task::yield_now().await;
}
assert!(matches!(
should_rejoin_after_background_heartbeat(&mut heartbeat).await,
Err(Error::Broker { code: 7, .. })
));
assert!(heartbeat.is_finished());
}
#[tokio::test]
async fn heartbeat_handle_exposes_group_identity() {
let (shutdown, shutdown_rx) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async move {
let _ = shutdown_rx.await;
Ok(())
});
let heartbeat = heartbeat_handle(shutdown, handle, "orders-group", "member-a", 7);
assert_eq!(heartbeat.group_id(), "orders-group");
assert_eq!(heartbeat.member_id(), "member-a");
assert_eq!(heartbeat.generation_id(), 7);
assert_eq!(heartbeat.interval(), std::time::Duration::from_millis(100));
assert_eq!(
heartbeat.state_for("orders-group", "member-a", 7),
HeartbeatHandleState::Current
);
assert_eq!(
heartbeat.state_for("orders-group", "member-b", 8),
HeartbeatHandleState::StaleGeneration
);
assert_eq!(
heartbeat.state_for("payments-group", "member-a", 7),
HeartbeatHandleState::DifferentGroup
);
heartbeat.stop().await.unwrap();
}
#[tokio::test]
async fn stale_heartbeat_shutdown_stops_running_task() {
let (shutdown, shutdown_rx) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async move {
let _ = shutdown_rx.await;
Ok(())
});
let mut heartbeat = heartbeat_handle(shutdown, handle, "orders-group", "member-a", 1);
heartbeat.stop_stale_generation().await.unwrap();
assert!(heartbeat.is_finished());
assert_eq!(heartbeat.try_wait().await.unwrap(), Some(()));
}
#[tokio::test]
async fn stale_heartbeat_shutdown_ignores_rejoinable_group_error() {
let (shutdown, _shutdown_rx) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async {
Err(Error::Broker {
code: 27,
context: "background heartbeat group orders-group".to_owned(),
})
});
let mut heartbeat = heartbeat_handle(shutdown, handle, "orders-group", "member-a", 1);
while !heartbeat.is_finished() {
tokio::task::yield_now().await;
}
heartbeat.stop_stale_generation().await.unwrap();
assert!(heartbeat.is_finished());
}
#[tokio::test]
async fn stale_heartbeat_shutdown_surfaces_non_rejoin_error() {
let (shutdown, _shutdown_rx) = tokio::sync::oneshot::channel();
let handle = tokio::spawn(async {
Err(Error::Broker {
code: 7,
context: "background heartbeat group orders-group".to_owned(),
})
});
let mut heartbeat = heartbeat_handle(shutdown, handle, "orders-group", "member-a", 1);
while !heartbeat.is_finished() {
tokio::task::yield_now().await;
}
assert!(matches!(
heartbeat.stop_stale_generation().await,
Err(Error::Broker { code: 7, .. })
));
assert!(heartbeat.is_finished());
}
fn member(member_id: &str, topics: &[&str]) -> JoinGroupMember {
JoinGroupMember {
member_id: member_id.to_owned(),
metadata: ConsumerProtocolSubscriptionV0 {
topics: topics.iter().map(|topic| (*topic).to_owned()).collect(),
user_data: None,
}
.encode()
.unwrap(),
}
}
fn sticky_member(
member_id: &str,
topics: &[&str],
owned_partitions: Vec<ConsumerProtocolTopicAssignment>,
generation: Option<i32>,
) -> JoinGroupMember {
JoinGroupMember {
member_id: member_id.to_owned(),
metadata: ConsumerProtocolSubscriptionV0 {
topics: topics.iter().map(|topic| (*topic).to_owned()).collect(),
user_data: if owned_partitions.is_empty() && generation.is_none() {
None
} else {
Some(encode_sticky_user_data(&owned_partitions, generation).unwrap())
},
}
.encode()
.unwrap(),
}
}
fn cooperative_member(
member_id: &str,
topics: &[&str],
owned_partitions: Vec<ConsumerProtocolTopicAssignment>,
) -> JoinGroupMember {
JoinGroupMember {
member_id: member_id.to_owned(),
metadata: ConsumerProtocolSubscriptionV1 {
topics: topics.iter().map(|topic| (*topic).to_owned()).collect(),
user_data: None,
owned_partitions,
}
.encode()
.unwrap(),
}
}
fn heartbeat_handle(
shutdown: tokio::sync::oneshot::Sender<()>,
handle: tokio::task::JoinHandle<crate::Result<()>>,
group_id: &str,
member_id: &str,
generation_id: i32,
) -> ConsumerGroupHeartbeat {
ConsumerGroupHeartbeat {
group_id: group_id.to_owned(),
generation_id,
member_id: member_id.to_owned(),
interval: std::time::Duration::from_millis(100),
shutdown: Some(shutdown),
handle: Some(handle),
consumer_session: None,
consumer_state: None,
}
}
fn decode_assignment(bytes: &[u8]) -> ConsumerProtocolAssignmentV0 {
ConsumerProtocolAssignmentV0::decode(bytes).unwrap()
}
fn metadata_fixture(topic: &str, partitions: &[i32]) -> MetadataResponseV1 {
MetadataResponseV1 {
brokers: Vec::new(),
controller_id: 1,
topics: vec![topic_metadata(topic, partitions)],
}
}
fn topic_metadata(topic: &str, partitions: &[i32]) -> TopicMetadata {
TopicMetadata {
error_code: 0,
name: topic.to_owned(),
is_internal: false,
partitions: partitions
.iter()
.map(|partition| PartitionMetadata {
error_code: 0,
partition_index: *partition,
leader_id: 1,
replica_nodes: vec![1],
isr_nodes: vec![1],
})
.collect(),
}
}
}