use crate::native::{KafkaClientError, KafkaClientResult};
#[derive(Debug, Default, Clone)]
pub(crate) struct Encoder {
buf: Vec<u8>,
}
impl Encoder {
#[cfg(test)]
pub(crate) fn new() -> Self {
Self { buf: Vec::new() }
}
pub(crate) fn with_capacity(capacity: usize) -> Self {
Self {
buf: Vec::with_capacity(capacity),
}
}
pub(crate) fn into_inner(self) -> Vec<u8> {
self.buf
}
pub(crate) fn put_i8(&mut self, value: i8) {
self.buf.push(value as u8);
}
pub(crate) fn put_i16(&mut self, value: i16) {
self.buf.extend_from_slice(&value.to_be_bytes());
}
pub(crate) fn put_i32(&mut self, value: i32) {
self.buf.extend_from_slice(&value.to_be_bytes());
}
pub(crate) fn put_i64(&mut self, value: i64) {
self.buf.extend_from_slice(&value.to_be_bytes());
}
pub(crate) fn put_bool(&mut self, value: bool) {
self.buf.push(u8::from(value));
}
pub(crate) fn put_raw(&mut self, value: &[u8]) {
self.buf.extend_from_slice(value);
}
pub(crate) fn put_bytes(&mut self, value: &[u8]) -> KafkaClientResult<()> {
let len = i32::try_from(value.len())
.map_err(|_| KafkaClientError::protocol("Kafka bytes exceed int32 length"))?;
self.put_i32(len);
self.buf.extend_from_slice(value);
Ok(())
}
pub(crate) fn put_nullable_bytes(&mut self, value: Option<&[u8]>) -> KafkaClientResult<()> {
match value {
Some(value) => self.put_bytes(value),
None => {
self.put_i32(-1);
Ok(())
}
}
}
pub(crate) fn put_string(&mut self, value: &str) -> KafkaClientResult<()> {
let len = i16::try_from(value.len())
.map_err(|_| KafkaClientError::protocol("Kafka string exceeds int16 length"))?;
self.put_i16(len);
self.buf.extend_from_slice(value.as_bytes());
Ok(())
}
pub(crate) fn put_nullable_string(&mut self, value: Option<&str>) -> KafkaClientResult<()> {
match value {
Some(value) => self.put_string(value),
None => {
self.put_i16(-1);
Ok(())
}
}
}
pub(crate) fn put_compact_string(&mut self, value: &str) -> KafkaClientResult<()> {
let len = u32::try_from(value.len())
.map_err(|_| KafkaClientError::protocol("Kafka compact string exceeds u32 length"))?;
self.put_unsigned_varint(len + 1);
self.buf.extend_from_slice(value.as_bytes());
Ok(())
}
pub(crate) fn put_compact_nullable_string(
&mut self,
value: Option<&str>,
) -> KafkaClientResult<()> {
match value {
Some(value) => self.put_compact_string(value),
None => {
self.put_unsigned_varint(0);
Ok(())
}
}
}
pub(crate) fn put_compact_bytes(&mut self, value: &[u8]) -> KafkaClientResult<()> {
let len = u32::try_from(value.len())
.map_err(|_| KafkaClientError::protocol("Kafka compact bytes exceed u32 length"))?;
self.put_unsigned_varint(len + 1);
self.buf.extend_from_slice(value);
Ok(())
}
pub(crate) fn put_array_len(&mut self, len: usize, flexible: bool) -> KafkaClientResult<()> {
if flexible {
let len = u32::try_from(len)
.map_err(|_| KafkaClientError::protocol("Kafka array exceeds u32 length"))?;
self.put_unsigned_varint(len + 1);
} else {
let len = i32::try_from(len)
.map_err(|_| KafkaClientError::protocol("Kafka array exceeds int32 length"))?;
self.put_i32(len);
}
Ok(())
}
pub(crate) fn put_empty_tags(&mut self) {
self.put_unsigned_varint(0);
}
pub(crate) fn put_unsigned_varint(&mut self, mut value: u32) {
while (value & !0x7f) != 0 {
self.buf.push(((value & 0x7f) as u8) | 0x80);
value >>= 7;
}
self.buf.push(value as u8);
}
#[cfg(test)]
pub(crate) fn put_varint_i32(&mut self, value: i32) {
self.put_unsigned_varint(((value << 1) ^ (value >> 31)) as u32);
}
#[cfg(test)]
pub(crate) fn put_varint_i64(&mut self, value: i64) {
let mut value = ((value << 1) ^ (value >> 63)) as u64;
while (value & !0x7f) != 0 {
self.buf.push(((value & 0x7f) as u8) | 0x80);
value >>= 7;
}
self.buf.push(value as u8);
}
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct Decoder<'a> {
input: &'a [u8],
pos: usize,
}
impl<'a> Decoder<'a> {
pub(crate) fn new(input: &'a [u8]) -> Self {
Self { input, pos: 0 }
}
pub(crate) fn position(&self) -> usize {
self.pos
}
pub(crate) fn remaining(&self) -> usize {
self.input.len().saturating_sub(self.pos)
}
pub(crate) fn is_done(&self) -> bool {
self.remaining() == 0
}
pub(crate) fn take(&mut self, len: usize) -> KafkaClientResult<&'a [u8]> {
let end = self
.pos
.checked_add(len)
.ok_or_else(|| KafkaClientError::protocol("Kafka decoder length overflow"))?;
if end > self.input.len() {
return Err(KafkaClientError::protocol(format!(
"short Kafka frame: need {len} bytes, have {}",
self.remaining()
)));
}
let slice = &self.input[self.pos..end];
self.pos = end;
Ok(slice)
}
pub(crate) fn get_i8(&mut self) -> KafkaClientResult<i8> {
Ok(self.take(1)?[0] as i8)
}
pub(crate) fn get_i16(&mut self) -> KafkaClientResult<i16> {
let bytes: [u8; 2] = self.take(2)?.try_into().expect("slice length checked");
Ok(i16::from_be_bytes(bytes))
}
pub(crate) fn get_i32(&mut self) -> KafkaClientResult<i32> {
let bytes: [u8; 4] = self.take(4)?.try_into().expect("slice length checked");
Ok(i32::from_be_bytes(bytes))
}
pub(crate) fn get_i64(&mut self) -> KafkaClientResult<i64> {
let bytes: [u8; 8] = self.take(8)?.try_into().expect("slice length checked");
Ok(i64::from_be_bytes(bytes))
}
pub(crate) fn get_bool(&mut self) -> KafkaClientResult<bool> {
match self.take(1)?[0] {
0 => Ok(false),
1 => Ok(true),
value => Err(KafkaClientError::protocol(format!(
"invalid Kafka boolean value {value}"
))),
}
}
pub(crate) fn get_string(&mut self) -> KafkaClientResult<String> {
let len = self.get_i16()?;
if len < 0 {
return Err(KafkaClientError::protocol(
"non-null Kafka string had null length",
));
}
self.string_of_len(len as usize)
}
pub(crate) fn get_nullable_string(&mut self) -> KafkaClientResult<Option<String>> {
let len = self.get_i16()?;
if len < 0 {
return Ok(None);
}
self.string_of_len(len as usize).map(Some)
}
pub(crate) fn get_compact_string(&mut self) -> KafkaClientResult<String> {
let len = self.get_unsigned_varint()?;
if len == 0 {
return Err(KafkaClientError::protocol(
"non-null compact Kafka string had null length",
));
}
self.string_of_len((len - 1) as usize)
}
pub(crate) fn get_compact_nullable_string(&mut self) -> KafkaClientResult<Option<String>> {
let len = self.get_unsigned_varint()?;
if len == 0 {
return Ok(None);
}
self.string_of_len((len - 1) as usize).map(Some)
}
fn string_of_len(&mut self, len: usize) -> KafkaClientResult<String> {
let bytes = self.take(len)?;
std::str::from_utf8(bytes)
.map(str::to_owned)
.map_err(|_| KafkaClientError::protocol("Kafka string was not valid UTF-8"))
}
pub(crate) fn get_compact_nullable_bytes(&mut self) -> KafkaClientResult<Option<&'a [u8]>> {
let len = self.get_unsigned_varint()?;
if len == 0 {
return Ok(None);
}
self.take((len - 1) as usize).map(Some)
}
pub(crate) fn get_compact_bytes(&mut self) -> KafkaClientResult<&'a [u8]> {
let len = self.get_unsigned_varint()?;
if len == 0 {
return Err(KafkaClientError::protocol(
"non-null compact Kafka bytes had null length",
));
}
self.take((len - 1) as usize)
}
pub(crate) fn get_nullable_bytes(&mut self) -> KafkaClientResult<Option<&'a [u8]>> {
let len = self.get_i32()?;
if len < 0 {
return Ok(None);
}
self.take(len as usize).map(Some)
}
pub(crate) fn get_array_len(&mut self, flexible: bool) -> KafkaClientResult<usize> {
if flexible {
let len = self.get_unsigned_varint()?;
if len == 0 {
return Err(KafkaClientError::protocol(
"non-null compact Kafka array had null length",
));
}
Ok((len - 1) as usize)
} else {
let len = self.get_i32()?;
if len < 0 {
return Err(KafkaClientError::protocol(
"non-null Kafka array had null length",
));
}
Ok(len as usize)
}
}
pub(crate) fn get_nullable_array_len(
&mut self,
flexible: bool,
) -> KafkaClientResult<Option<usize>> {
if flexible {
let len = self.get_unsigned_varint()?;
if len == 0 {
return Ok(None);
}
Ok(Some((len - 1) as usize))
} else {
let len = self.get_i32()?;
if len < 0 {
return Ok(None);
}
Ok(Some(len as usize))
}
}
pub(crate) fn skip_tags(&mut self) -> KafkaClientResult<()> {
let count = self.get_unsigned_varint()?;
let mut previous = None;
for _ in 0..count {
let tag = self.get_unsigned_varint()?;
if previous.is_some_and(|previous| tag <= previous) {
return Err(KafkaClientError::protocol(
"Kafka tagged fields were not strictly increasing",
));
}
previous = Some(tag);
let size = self.get_unsigned_varint()? as usize;
self.take(size)?;
}
Ok(())
}
pub(crate) fn get_unsigned_varint(&mut self) -> KafkaClientResult<u32> {
let mut value = 0_u32;
for shift in (0..35).step_by(7) {
let byte = self.take(1)?[0];
value |= ((byte & 0x7f) as u32) << shift;
if (byte & 0x80) == 0 {
return Ok(value);
}
}
Err(KafkaClientError::protocol("Kafka unsigned varint overflow"))
}
pub(crate) fn skip_varint(&mut self) -> KafkaClientResult<()> {
for _ in 0..10 {
let byte = self.take(1)?[0];
if (byte & 0x80) == 0 {
return Ok(());
}
}
Err(KafkaClientError::protocol("Kafka varint overflow"))
}
pub(crate) fn get_varint_i32(&mut self) -> KafkaClientResult<i32> {
let value = self.get_unsigned_varint()?;
Ok(((value >> 1) as i32) ^ (-((value & 1) as i32)))
}
pub(crate) fn get_varint_i64(&mut self) -> KafkaClientResult<i64> {
let mut value = 0_u64;
for shift in (0..70).step_by(7) {
let byte = self.take(1)?[0];
value |= ((byte & 0x7f) as u64) << shift;
if (byte & 0x80) == 0 {
return Ok(((value >> 1) as i64) ^ (-((value & 1) as i64)));
}
}
Err(KafkaClientError::protocol("Kafka varlong overflow"))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn unsigned_varint_round_trips_many_values() {
let values = [0, 1, 2, 63, 64, 127, 128, 16_384, 1_048_576, u32::MAX];
for value in values {
let mut encoder = Encoder::new();
encoder.put_unsigned_varint(value);
let bytes = encoder.into_inner();
let mut decoder = Decoder::new(&bytes);
assert_eq!(decoder.get_unsigned_varint().expect("decode"), value);
assert!(decoder.is_done());
}
}
#[test]
fn signed_varints_round_trip_many_values() {
let values = [i32::MIN, -123456, -1, 0, 1, 123456, i32::MAX];
for value in values {
let mut encoder = Encoder::new();
encoder.put_varint_i32(value);
let bytes = encoder.into_inner();
let mut decoder = Decoder::new(&bytes);
assert_eq!(decoder.get_varint_i32().expect("decode"), value);
assert!(decoder.is_done());
}
}
#[test]
fn unknown_tagged_fields_are_skipped() {
let mut encoder = Encoder::new();
encoder.put_unsigned_varint(2);
encoder.put_unsigned_varint(1);
encoder.put_unsigned_varint(3);
encoder.put_i8(1);
encoder.put_i8(2);
encoder.put_i8(3);
encoder.put_unsigned_varint(5);
encoder.put_unsigned_varint(1);
encoder.put_i8(9);
let bytes = encoder.into_inner();
let mut decoder = Decoder::new(&bytes);
decoder.skip_tags().expect("skip tags");
assert!(decoder.is_done());
}
#[test]
fn short_frame_reports_error() {
let mut decoder = Decoder::new(&[1, 2, 3]);
assert!(decoder.get_i32().is_err());
}
}