use bytes::{Buf, BufMut, Bytes};
use crate::error::{KrafkaError, ProtocolErrorKind, Result};
use crate::protocol::primitives::{Decode, Encode, KafkaBytes, KafkaString, TryEncode};
use crate::protocol::{array_len_i32, check_decode_array_len, decode_capacity};
pub const CONSUMER_PROTOCOL_TYPE: &str = "consumer";
pub const CONSUMER_PROTOCOL_MAX_VERSION: i16 = 3;
pub const CONSUMER_PROTOCOL_V1: i16 = 1;
pub const CONSUMER_PROTOCOL_V2: i16 = 2;
pub const CONSUMER_PROTOCOL_V3: i16 = 3;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ConsumerProtocolTopicPartitions {
pub topic: String,
pub partitions: Vec<i32>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ConsumerProtocolSubscription {
pub version: i16,
pub topics: Vec<String>,
pub user_data: Option<Bytes>,
pub owned_partitions: Vec<ConsumerProtocolTopicPartitions>,
pub generation_id: i32,
pub rack_id: Option<String>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ConsumerProtocolAssignment {
pub version: i16,
pub assigned_partitions: Vec<ConsumerProtocolTopicPartitions>,
pub user_data: Option<Bytes>,
}
impl ConsumerProtocolSubscription {
pub fn new(topics: Vec<String>) -> Self {
Self {
version: 0,
topics,
user_data: None,
owned_partitions: Vec::new(),
generation_id: -1,
rack_id: None,
}
}
#[must_use]
pub fn with_owned_partitions(
mut self,
owned_partitions: Vec<ConsumerProtocolTopicPartitions>,
) -> Self {
self.owned_partitions = owned_partitions;
self.version = self.version.max(CONSUMER_PROTOCOL_V1);
self
}
#[must_use]
pub fn with_generation_id(mut self, generation_id: i32) -> Self {
self.generation_id = generation_id;
self.version = self.version.max(CONSUMER_PROTOCOL_V2);
self
}
#[must_use]
pub fn with_rack_id(mut self, rack_id: impl Into<String>) -> Self {
self.rack_id = Some(rack_id.into());
self.version = self.version.max(CONSUMER_PROTOCOL_V3);
self
}
}
impl ConsumerProtocolAssignment {
pub fn new(version: i16, assigned_partitions: Vec<ConsumerProtocolTopicPartitions>) -> Self {
Self {
version,
assigned_partitions,
user_data: None,
}
}
}
pub fn decode_consumer_protocol_subscription(data: &Bytes) -> Result<ConsumerProtocolSubscription> {
if data.is_empty() {
return Ok(ConsumerProtocolSubscription::default());
}
let mut buf = data.clone();
let version = decode_version(&mut buf, "ConsumerProtocolSubscription")?;
let effective = version.min(CONSUMER_PROTOCOL_MAX_VERSION);
let topic_count = check_decode_array_len(i32::decode(&mut buf)?)?;
let mut topics = Vec::with_capacity(decode_capacity(topic_count, buf.remaining()));
for _ in 0..topic_count {
topics.push(non_null_string(&mut buf, "subscription topic")?);
}
let user_data = KafkaBytes::decode(&mut buf)?.0;
let owned_partitions = if effective >= CONSUMER_PROTOCOL_V1 {
decode_topic_partitions(&mut buf, "owned partitions")?
} else {
Vec::new()
};
let generation_id = if effective >= CONSUMER_PROTOCOL_V2 {
i32::decode(&mut buf)?
} else {
-1
};
let rack_id = if effective >= CONSUMER_PROTOCOL_V3 {
KafkaString::decode(&mut buf)?.0
} else {
None
};
Ok(ConsumerProtocolSubscription {
version,
topics,
user_data,
owned_partitions,
generation_id,
rack_id,
})
}
pub fn decode_consumer_protocol_assignment(data: &Bytes) -> Result<ConsumerProtocolAssignment> {
if data.is_empty() {
return Ok(ConsumerProtocolAssignment::default());
}
let mut buf = data.clone();
let version = decode_version(&mut buf, "ConsumerProtocolAssignment")?;
let assigned_partitions = decode_topic_partitions(&mut buf, "assigned partitions")?;
let user_data = KafkaBytes::decode(&mut buf)?.0;
Ok(ConsumerProtocolAssignment {
version,
assigned_partitions,
user_data,
})
}
pub fn encode_consumer_protocol_subscription(
subscription: &ConsumerProtocolSubscription,
buf: &mut impl BufMut,
) -> Result<()> {
let version = subscription.version;
if version < 0 {
return Err(invalid_length(format!(
"ConsumerProtocolSubscription version {version} is negative"
)));
}
if version < CONSUMER_PROTOCOL_V1 && !subscription.owned_partitions.is_empty() {
return Err(invalid_length(
"owned partitions need ConsumerProtocolSubscription v1 or newer",
));
}
if version < CONSUMER_PROTOCOL_V2 && subscription.generation_id != -1 {
return Err(invalid_length(
"generation_id needs ConsumerProtocolSubscription v2 or newer",
));
}
if version < CONSUMER_PROTOCOL_V3 && subscription.rack_id.is_some() {
return Err(invalid_length(
"rack_id needs ConsumerProtocolSubscription v3 or newer",
));
}
version.encode(buf);
buf.put_i32(array_len_i32(subscription.topics.len())?);
for topic in &subscription.topics {
KafkaString::new(topic).try_encode(buf)?;
}
encode_nullable_bytes(subscription.user_data.as_ref(), buf);
if version >= CONSUMER_PROTOCOL_V1 {
encode_topic_partitions(&subscription.owned_partitions, buf)?;
}
if version >= CONSUMER_PROTOCOL_V2 {
subscription.generation_id.encode(buf);
}
if version >= CONSUMER_PROTOCOL_V3 {
match &subscription.rack_id {
Some(rack) => KafkaString::new(rack).try_encode(buf)?,
None => KafkaString::null().try_encode(buf)?,
}
}
Ok(())
}
pub fn encode_consumer_protocol_assignment(
assignment: &ConsumerProtocolAssignment,
buf: &mut impl BufMut,
) -> Result<()> {
if assignment.version < 0 {
return Err(invalid_length(format!(
"ConsumerProtocolAssignment version {} is negative",
assignment.version
)));
}
assignment.version.encode(buf);
encode_topic_partitions(&assignment.assigned_partitions, buf)?;
encode_nullable_bytes(assignment.user_data.as_ref(), buf);
Ok(())
}
fn decode_version(buf: &mut impl Buf, what: &str) -> Result<i16> {
let version = i16::decode(buf)?;
if version < 0 {
return Err(KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!("{what} has negative version {version}"),
));
}
Ok(version)
}
fn decode_topic_partitions(
buf: &mut impl Buf,
what: &str,
) -> Result<Vec<ConsumerProtocolTopicPartitions>> {
let topic_count = check_decode_array_len(i32::decode(buf)?)?;
let mut topics = Vec::with_capacity(decode_capacity(topic_count, buf.remaining()));
for _ in 0..topic_count {
let topic = non_null_string(buf, what)?;
let partition_count = check_decode_array_len(i32::decode(buf)?)?;
let mut partitions = Vec::with_capacity(decode_capacity(partition_count, buf.remaining()));
for _ in 0..partition_count {
partitions.push(i32::decode(buf)?);
}
topics.push(ConsumerProtocolTopicPartitions { topic, partitions });
}
Ok(topics)
}
fn encode_topic_partitions(
topics: &[ConsumerProtocolTopicPartitions],
buf: &mut impl BufMut,
) -> Result<()> {
buf.put_i32(array_len_i32(topics.len())?);
for entry in topics {
KafkaString::new(&entry.topic).try_encode(buf)?;
buf.put_i32(array_len_i32(entry.partitions.len())?);
for &partition in &entry.partitions {
partition.encode(buf);
}
}
Ok(())
}
fn encode_nullable_bytes(value: Option<&Bytes>, buf: &mut impl BufMut) {
match value {
Some(bytes) => match i32::try_from(bytes.len()) {
Ok(len) => {
buf.put_i32(len);
buf.put_slice(bytes);
}
Err(_) => buf.put_i32(-1),
},
None => buf.put_i32(-1),
}
}
fn non_null_string(buf: &mut impl Buf, what: &str) -> Result<String> {
KafkaString::decode(buf)?.0.ok_or_else(|| {
KrafkaError::protocol_kind(
ProtocolErrorKind::Malformed,
format!("null topic name in {what}"),
)
})
}
fn invalid_length(message: impl Into<String>) -> KrafkaError {
KrafkaError::protocol_kind(ProtocolErrorKind::InvalidLength, message)
}