1use super::decode_capacity;
7use bytes::{Buf, BufMut, Bytes, BytesMut};
8
9use crate::error::{KrafkaError, ProtocolErrorKind, Result};
10use crate::util::{crc32c, varint};
11
12#[non_exhaustive]
23#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
24#[repr(u8)]
25pub enum Compression {
26 #[default]
28 None = 0,
29 Gzip = 1,
33 Snappy = 2,
37 Lz4 = 3,
41 Zstd = 4,
46}
47
48impl Compression {
49 #[inline]
51 #[must_use]
52 pub const fn from_i8(value: i8) -> Option<Self> {
53 match value {
54 0 => Some(Self::None),
55 1 => Some(Self::Gzip),
56 2 => Some(Self::Snappy),
57 3 => Some(Self::Lz4),
58 4 => Some(Self::Zstd),
59 _ => None,
60 }
61 }
62
63 #[inline]
70 #[must_use]
71 pub const fn from_u8(value: u8) -> Option<Self> {
72 match value & 0x07 {
73 0 => Some(Self::None),
74 1 => Some(Self::Gzip),
75 2 => Some(Self::Snappy),
76 3 => Some(Self::Lz4),
77 4 => Some(Self::Zstd),
78 _ => None,
79 }
80 }
81
82 #[inline]
95 #[must_use]
96 pub const fn is_available(&self) -> bool {
97 match self {
98 Self::None => true,
99 Self::Gzip => cfg!(feature = "gzip"),
100 Self::Snappy => cfg!(feature = "snappy"),
101 Self::Lz4 => cfg!(feature = "lz4"),
102 Self::Zstd => cfg!(feature = "zstd"),
103 }
104 }
105
106 #[inline]
110 #[must_use]
111 pub const fn required_feature(&self) -> Option<&'static str> {
112 match self {
113 Self::None => Option::None,
114 Self::Gzip => Option::Some("gzip"),
115 Self::Snappy => Option::Some("snappy"),
116 Self::Lz4 => Option::Some("lz4"),
117 Self::Zstd => Option::Some("zstd"),
118 }
119 }
120
121 #[must_use]
128 pub const fn supports_level(&self) -> bool {
129 matches!(self, Self::Gzip | Self::Zstd)
130 }
131
132 #[must_use]
140 pub fn level_range(&self) -> Option<std::ops::RangeInclusive<i32>> {
141 match self {
142 Self::Gzip => Some(0..=9),
143 #[cfg(feature = "zstd")]
144 Self::Zstd => Some(zstd::compression_level_range()),
145 #[cfg(not(feature = "zstd"))]
146 Self::Zstd => Some(-131_072..=22),
147 _ => None,
148 }
149 }
150
151 pub(crate) fn compress_with_level(&self, payload: &[u8], level: Option<i32>) -> Result<Bytes> {
160 let _ = level;
161 match self {
162 Self::None => Ok(Bytes::copy_from_slice(payload)),
163 #[cfg(feature = "gzip")]
164 Self::Gzip => {
165 use flate2::write::GzEncoder;
166 use std::io::Write;
167
168 let flate_level = match level {
169 Some(l) => flate2::Compression::new(l.clamp(0, 9) as u32),
173 None => flate2::Compression::default(),
174 };
175 let mut encoder = GzEncoder::new(Vec::new(), flate_level);
176 encoder
177 .write_all(payload)
178 .map_err(|e| KrafkaError::compression(e.to_string()))?;
179 let compressed = encoder
180 .finish()
181 .map_err(|e| KrafkaError::compression(e.to_string()))?;
182 Ok(Bytes::from(compressed))
183 }
184 #[cfg(not(feature = "gzip"))]
185 Self::Gzip => Err(KrafkaError::compression(
186 "gzip compression requires the `gzip` Cargo feature",
187 )),
188 #[cfg(feature = "snappy")]
189 Self::Snappy => {
190 let mut encoder = snap::raw::Encoder::new();
196 let compressed = encoder
197 .compress_vec(payload)
198 .map_err(|e| KrafkaError::compression(e.to_string()))?;
199 Ok(Bytes::from(compressed))
200 }
201 #[cfg(not(feature = "snappy"))]
202 Self::Snappy => Err(KrafkaError::compression(
203 "snappy compression requires the `snappy` Cargo feature",
204 )),
205 #[cfg(feature = "lz4")]
206 Self::Lz4 => {
207 use std::io::Write;
208
209 let mut compressed = Vec::new();
216 let mut encoder = lz4_flex::frame::FrameEncoder::new(&mut compressed);
217 encoder
218 .write_all(payload)
219 .map_err(|e| KrafkaError::compression(e.to_string()))?;
220 encoder
221 .finish()
222 .map_err(|e| KrafkaError::compression(e.to_string()))?;
223 Ok(Bytes::from(compressed))
224 }
225 #[cfg(not(feature = "lz4"))]
226 Self::Lz4 => Err(KrafkaError::compression(
227 "lz4 compression requires the `lz4` Cargo feature",
228 )),
229 #[cfg(feature = "zstd")]
230 Self::Zstd => {
231 let compressed = zstd::encode_all(payload, level.unwrap_or(3))
233 .map_err(|e| KrafkaError::compression(e.to_string()))?;
234 Ok(Bytes::from(compressed))
235 }
236 #[cfg(not(feature = "zstd"))]
237 Self::Zstd => Err(KrafkaError::compression(
238 "zstd compression requires the `zstd` Cargo feature",
239 )),
240 }
241 }
242}
243
244#[non_exhaustive]
246#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
247#[repr(u8)]
248pub enum TimestampType {
249 #[default]
251 CreateTime = 0,
252 LogAppendTime = 1,
254}
255
256impl TimestampType {
257 #[inline]
259 pub fn from_attributes(attributes: i16) -> Self {
260 if attributes & 0x08 != 0 {
261 Self::LogAppendTime
262 } else {
263 Self::CreateTime
264 }
265 }
266}
267
268#[derive(Debug, Clone, PartialEq, Eq)]
270pub struct RecordHeader {
271 pub key: Bytes,
277 pub value: Option<Bytes>,
279}
280
281impl RecordHeader {
282 pub fn new(key: impl Into<Bytes>, value: impl Into<Bytes>) -> Self {
284 Self {
285 key: key.into(),
286 value: Some(value.into()),
287 }
288 }
289
290 pub fn null(key: impl Into<Bytes>) -> Self {
296 Self {
297 key: key.into(),
298 value: None,
299 }
300 }
301
302 #[inline]
304 pub fn key_str(&self) -> Option<&str> {
305 std::str::from_utf8(&self.key).ok()
306 }
307
308 #[inline]
310 pub fn encode(&self, buf: &mut impl BufMut) -> Result<()> {
311 let key_len = i32::try_from(self.key.len()).map_err(|_| {
312 KrafkaError::protocol_kind(
313 ProtocolErrorKind::InvalidLength,
314 "record header key too large for i32 length",
315 )
316 })?;
317 varint::encode_signed_varint(key_len, buf);
318 buf.put_slice(&self.key);
319 match &self.value {
320 Some(v) => {
321 let val_len = i32::try_from(v.len()).map_err(|_| {
322 KrafkaError::protocol_kind(
323 ProtocolErrorKind::InvalidLength,
324 "record header value too large for i32 length",
325 )
326 })?;
327 varint::encode_signed_varint(val_len, buf);
328 buf.put_slice(v);
329 }
330 None => varint::encode_signed_varint(-1, buf),
331 }
332 Ok(())
333 }
334
335 #[inline]
337 pub fn decode(buf: &mut impl Buf) -> Result<Self> {
338 let key_len = varint::decode_signed_varint(buf)?;
339 if key_len < 0 || buf.remaining() < key_len as usize {
340 return Err(KrafkaError::protocol_kind(
341 ProtocolErrorKind::InvalidValue,
342 "invalid header key length",
343 ));
344 }
345 let key = buf.copy_to_bytes(key_len as usize);
346
347 let value_len = varint::decode_signed_varint(buf)?;
348 let value = if value_len < 0 {
349 None
350 } else {
351 if buf.remaining() < value_len as usize {
352 return Err(KrafkaError::protocol_kind(
353 ProtocolErrorKind::InvalidValue,
354 "invalid header value length",
355 ));
356 }
357 Some(buf.copy_to_bytes(value_len as usize))
358 };
359
360 Ok(Self { key, value })
361 }
362}
363
364#[must_use = "contains record key, value and headers"]
366#[derive(Debug, Clone, PartialEq, Eq)]
367pub struct Record {
368 pub attributes: i8,
370 pub timestamp_delta: i64,
372 pub offset_delta: i32,
374 pub key: Option<Bytes>,
376 pub value: Option<Bytes>,
378 pub headers: Vec<RecordHeader>,
380}
381
382impl Record {
383 pub fn new(key: Option<Bytes>, value: Option<Bytes>) -> Self {
385 Self {
386 attributes: 0,
387 timestamp_delta: 0,
388 offset_delta: 0,
389 key,
390 value,
391 headers: Vec::new(),
392 }
393 }
394
395 pub fn with_header(mut self, key: impl Into<Bytes>, value: impl Into<Bytes>) -> Self {
397 self.headers.push(RecordHeader::new(key, value));
398 self
399 }
400
401 pub fn with_timestamp_delta(mut self, delta: i64) -> Self {
403 self.timestamp_delta = delta;
404 self
405 }
406
407 pub fn with_offset_delta(mut self, delta: i32) -> Self {
409 self.offset_delta = delta;
410 self
411 }
412
413 #[inline]
419 pub fn encode(&self, buf: &mut impl BufMut) -> Result<()> {
420 let body_size = self.record_body_size()?;
421 let record_len = i32::try_from(body_size).map_err(|_| {
422 KrafkaError::protocol_kind(
423 ProtocolErrorKind::InvalidLength,
424 "record too large for i32 length prefix",
425 )
426 })?;
427 varint::encode_signed_varint(record_len, buf);
428 self.encode_body(buf)?;
429 Ok(())
430 }
431
432 #[inline]
436 pub fn record_body_size(&self) -> Result<usize> {
437 let mut size: usize = 0;
438 size += 1;
440 size += varint::signed_varlong_size(self.timestamp_delta);
442 size += varint::signed_varint_size(self.offset_delta);
444 match &self.key {
446 Some(k) => {
447 let key_len = i32::try_from(k.len()).map_err(|_| {
448 KrafkaError::protocol_kind(
449 ProtocolErrorKind::InvalidLength,
450 "record key too large for i32 length",
451 )
452 })?;
453 size += varint::signed_varint_size(key_len);
454 size += k.len();
455 }
456 None => {
457 size += varint::signed_varint_size(-1);
458 }
459 }
460 match &self.value {
462 Some(v) => {
463 let val_len = i32::try_from(v.len()).map_err(|_| {
464 KrafkaError::protocol_kind(
465 ProtocolErrorKind::InvalidLength,
466 "record value too large for i32 length",
467 )
468 })?;
469 size += varint::signed_varint_size(val_len);
470 size += v.len();
471 }
472 None => {
473 size += varint::signed_varint_size(-1);
474 }
475 }
476 let headers_len = i32::try_from(self.headers.len()).map_err(|_| {
478 KrafkaError::protocol_kind(
479 ProtocolErrorKind::InvalidLength,
480 "record headers count exceeds i32 limit",
481 )
482 })?;
483 size += varint::signed_varint_size(headers_len);
484 for header in &self.headers {
485 let key_len = i32::try_from(header.key.len()).map_err(|_| {
486 KrafkaError::protocol_kind(
487 ProtocolErrorKind::InvalidLength,
488 "record header key too large for i32 length",
489 )
490 })?;
491 size += varint::signed_varint_size(key_len);
492 size += header.key.len();
493 match &header.value {
494 Some(v) => {
495 let val_len = i32::try_from(v.len()).map_err(|_| {
496 KrafkaError::protocol_kind(
497 ProtocolErrorKind::InvalidLength,
498 "record header value too large for i32 length",
499 )
500 })?;
501 size += varint::signed_varint_size(val_len);
502 size += v.len();
503 }
504 None => {
505 size += varint::signed_varint_size(-1);
506 }
507 }
508 }
509 Ok(size)
510 }
511
512 #[inline]
513 fn encode_body(&self, buf: &mut impl BufMut) -> Result<()> {
514 buf.put_i8(self.attributes);
515 varint::encode_signed_varlong(self.timestamp_delta, buf);
516 varint::encode_signed_varint(self.offset_delta, buf);
517
518 match &self.key {
520 Some(k) => {
521 let key_len = i32::try_from(k.len()).map_err(|_| {
522 KrafkaError::protocol_kind(
523 ProtocolErrorKind::InvalidLength,
524 "record key too large for i32 length",
525 )
526 })?;
527 varint::encode_signed_varint(key_len, buf);
528 buf.put_slice(k);
529 }
530 None => varint::encode_signed_varint(-1, buf),
531 }
532
533 match &self.value {
535 Some(v) => {
536 let val_len = i32::try_from(v.len()).map_err(|_| {
537 KrafkaError::protocol_kind(
538 ProtocolErrorKind::InvalidLength,
539 "record value too large for i32 length",
540 )
541 })?;
542 varint::encode_signed_varint(val_len, buf);
543 buf.put_slice(v);
544 }
545 None => varint::encode_signed_varint(-1, buf),
546 }
547
548 let headers_len = i32::try_from(self.headers.len()).map_err(|_| {
550 KrafkaError::protocol_kind(
551 ProtocolErrorKind::InvalidLength,
552 "record headers count exceeds i32 limit",
553 )
554 })?;
555 varint::encode_signed_varint(headers_len, buf);
556 for header in &self.headers {
557 header.encode(buf)?;
558 }
559 Ok(())
560 }
561
562 #[inline]
564 pub fn decode(buf: &mut impl Buf) -> Result<Self> {
565 let length = varint::decode_signed_varint(buf)?;
566 if length < 0 {
567 return Err(KrafkaError::protocol_kind(
568 crate::error::ProtocolErrorKind::InvalidValue,
569 format!("invalid record length: {length}"),
570 ));
571 }
572 let length = usize::try_from(length).map_err(|_| {
573 KrafkaError::protocol_kind(
574 crate::error::ProtocolErrorKind::InvalidLength,
575 format!("record length {length} overflows usize on this target"),
576 )
577 })?;
578 if buf.remaining() < length {
579 return Err(KrafkaError::protocol_kind(
580 crate::error::ProtocolErrorKind::TruncatedFrame,
581 format!(
582 "record body truncated: need {length} bytes, have {}",
583 buf.remaining()
584 ),
585 ));
586 }
587
588 let mut rbuf = buf.copy_to_bytes(length);
593
594 let attributes = if rbuf.has_remaining() {
595 rbuf.get_i8()
596 } else {
597 return Err(KrafkaError::protocol_kind(
598 ProtocolErrorKind::Malformed,
599 "missing record attributes",
600 ));
601 };
602
603 let timestamp_delta = varint::decode_signed_varlong(&mut rbuf)?;
604 let offset_delta = varint::decode_signed_varint(&mut rbuf)?;
605
606 let key_len = varint::decode_signed_varint(&mut rbuf)?;
608 let key = if key_len < 0 {
609 None
610 } else {
611 if rbuf.remaining() < key_len as usize {
612 return Err(KrafkaError::protocol_kind(
613 ProtocolErrorKind::InvalidValue,
614 "invalid record key length",
615 ));
616 }
617 Some(rbuf.copy_to_bytes(key_len as usize))
618 };
619
620 let value_len = varint::decode_signed_varint(&mut rbuf)?;
622 let value = if value_len < 0 {
623 None
624 } else {
625 if rbuf.remaining() < value_len as usize {
626 return Err(KrafkaError::protocol_kind(
627 ProtocolErrorKind::InvalidValue,
628 "invalid record value length",
629 ));
630 }
631 Some(rbuf.copy_to_bytes(value_len as usize))
632 };
633
634 let header_count = varint::decode_signed_varint(&mut rbuf)?;
636 if header_count < 0 {
637 return Err(KrafkaError::protocol_kind(
638 ProtocolErrorKind::InvalidValue,
639 format!("negative header count {header_count} in record"),
640 ));
641 }
642 let header_count = header_count as usize;
643 if header_count > super::MAX_DECODE_ARRAY_LEN {
644 return Err(KrafkaError::protocol_kind(
645 ProtocolErrorKind::InvalidLength,
646 format!(
647 "header count {header_count} exceeds safety limit {}",
648 super::MAX_DECODE_ARRAY_LEN
649 ),
650 ));
651 }
652 let mut headers = Vec::with_capacity(decode_capacity(header_count, rbuf.remaining()));
660 for _ in 0..header_count {
661 headers.push(RecordHeader::decode(&mut rbuf)?);
662 }
663
664 Ok(Self {
665 attributes,
666 timestamp_delta,
667 offset_delta,
668 key,
669 value,
670 headers,
671 })
672 }
673}
674
675#[derive(Debug, Clone, Copy, Default)]
677pub struct RecordBatchAttributes {
678 pub compression: Compression,
680 pub timestamp_type: TimestampType,
682 pub is_transactional: bool,
684 pub is_control_batch: bool,
686}
687
688impl RecordBatchAttributes {
689 #[inline]
695 pub fn from_i16(value: i16) -> Result<Self> {
696 let compression_bits = (value & 0x07) as u8;
697 let compression = Compression::from_u8(compression_bits).ok_or_else(|| {
698 KrafkaError::protocol_kind(
699 crate::error::ProtocolErrorKind::InvalidValue,
700 format!("unknown compression codec discriminant: {compression_bits}"),
701 )
702 })?;
703 Ok(Self {
704 compression,
705 timestamp_type: TimestampType::from_attributes(value),
706 is_transactional: value & 0x10 != 0,
707 is_control_batch: value & 0x20 != 0,
708 })
709 }
710
711 #[inline]
713 pub fn to_i16(self) -> i16 {
714 let mut value = self.compression as i16;
715 if matches!(self.timestamp_type, TimestampType::LogAppendTime) {
716 value |= 0x08;
717 }
718 if self.is_transactional {
719 value |= 0x10;
720 }
721 if self.is_control_batch {
722 value |= 0x20;
723 }
724 value
725 }
726}
727
728#[derive(Debug, Clone)]
730pub struct RecordBatch {
731 pub base_offset: i64,
733 pub partition_leader_epoch: i32,
735 pub magic: i8,
737 pub attributes: RecordBatchAttributes,
739 pub last_offset_delta: i32,
741 pub base_timestamp: i64,
743 pub max_timestamp: i64,
745 pub producer_id: i64,
747 pub producer_epoch: i16,
749 pub base_sequence: i32,
751 pub records: Vec<Record>,
753 pub(crate) compression_level: Option<i32>,
760}
761
762impl RecordBatch {
763 pub fn new() -> Self {
765 Self {
766 base_offset: 0,
767 partition_leader_epoch: 0,
768 magic: 2,
769 attributes: RecordBatchAttributes::default(),
770 last_offset_delta: 0,
771 base_timestamp: 0,
772 max_timestamp: 0,
773 producer_id: -1,
774 producer_epoch: -1,
775 base_sequence: -1,
776 records: Vec::new(),
777 compression_level: None,
778 }
779 }
780
781 pub fn with_compression(mut self, compression: Compression) -> Self {
783 self.attributes.compression = compression;
784 self
785 }
786
787 pub fn add_record(&mut self, record: Record) {
789 self.records.push(record);
790 }
791
792 pub fn encode(&self) -> Result<Bytes> {
814 const BATCH_LENGTH_POS: usize = 8;
816 const CRC_POS: usize = 17;
817 const CRC_REGION_START: usize = 21;
818 const FIXED_OVERHEAD: usize = 49; let records_count = i32::try_from(self.records.len()).map_err(|_| {
823 KrafkaError::protocol_kind(
824 ProtocolErrorKind::InvalidLength,
825 "record batch record count exceeds i32 limit",
826 )
827 })?;
828
829 let estimated_records_size: usize = self
832 .records
833 .iter()
834 .map(|r| {
835 r.key.as_ref().map_or(0, Bytes::len) + r.value.as_ref().map_or(0, Bytes::len) + 25
836 })
837 .sum();
838
839 if matches!(self.attributes.compression, Compression::None) {
840 const HEADER_SIZE: usize = 12 + FIXED_OVERHEAD;
846
847 let mut buf = BytesMut::with_capacity(HEADER_SIZE + estimated_records_size);
848
849 buf.put_i64(self.base_offset);
850 buf.put_i32(0); buf.put_i32(self.partition_leader_epoch);
852 buf.put_i8(self.magic);
853 buf.put_u32(0); buf.put_i16(self.attributes.to_i16());
856 buf.put_i32(self.last_offset_delta);
857 buf.put_i64(self.base_timestamp);
858 buf.put_i64(self.max_timestamp);
859 buf.put_i64(self.producer_id);
860 buf.put_i16(self.producer_epoch);
861 buf.put_i32(self.base_sequence);
862 buf.put_i32(records_count);
863
864 for record in &self.records {
865 record.encode(&mut buf)?;
866 }
867
868 let batch_length = i32::try_from(buf.len() - 12).map_err(|_| {
870 KrafkaError::protocol_kind(
871 ProtocolErrorKind::InvalidLength,
872 "record batch too large for i32 length prefix",
873 )
874 })?;
875 buf[BATCH_LENGTH_POS..BATCH_LENGTH_POS + 4]
876 .copy_from_slice(&batch_length.to_be_bytes());
877
878 let crc = crc32c(&buf[CRC_REGION_START..]);
880 buf[CRC_POS..CRC_POS + 4].copy_from_slice(&crc.to_be_bytes());
881
882 Ok(buf.freeze())
883 } else {
884 let mut records_buf = BytesMut::with_capacity(estimated_records_size);
887 for record in &self.records {
888 record.encode(&mut records_buf)?;
889 }
890
891 let compressed_records = self.compress_records(&records_buf)?;
892
893 let batch_length =
894 i32::try_from(FIXED_OVERHEAD + compressed_records.len()).map_err(|_| {
895 KrafkaError::protocol_kind(
896 ProtocolErrorKind::InvalidLength,
897 "record batch too large for i32 length prefix",
898 )
899 })?;
900
901 let mut buf = BytesMut::with_capacity(12 + batch_length as usize);
902
903 buf.put_i64(self.base_offset);
904 buf.put_i32(batch_length);
905 buf.put_i32(self.partition_leader_epoch);
906 buf.put_i8(self.magic);
907 buf.put_u32(0); buf.put_i16(self.attributes.to_i16());
910 buf.put_i32(self.last_offset_delta);
911 buf.put_i64(self.base_timestamp);
912 buf.put_i64(self.max_timestamp);
913 buf.put_i64(self.producer_id);
914 buf.put_i16(self.producer_epoch);
915 buf.put_i32(self.base_sequence);
916 buf.put_i32(records_count);
917 buf.put_slice(&compressed_records);
918
919 let crc = crc32c(&buf[CRC_REGION_START..]);
920 buf[CRC_POS..CRC_POS + 4].copy_from_slice(&crc.to_be_bytes());
921
922 Ok(buf.freeze())
923 }
924 }
925
926 fn compress_records(&self, records: &[u8]) -> Result<Bytes> {
927 self.attributes
928 .compression
929 .compress_with_level(records, self.compression_level)
930 }
931
932 pub fn decode(buf: &mut impl Buf) -> Result<Self> {
938 Self::decode_with_limit(buf, Self::MAX_DECOMPRESSED_SIZE)
939 }
940
941 pub fn decode_with_limit(buf: &mut impl Buf, max_decompressed_size: usize) -> Result<Self> {
946 if buf.remaining() < 12 {
947 return Err(KrafkaError::protocol_kind(
948 ProtocolErrorKind::TruncatedFrame,
949 "not enough bytes for record batch header",
950 ));
951 }
952
953 let base_offset = buf.get_i64();
954 let batch_length_i32 = buf.get_i32();
955
956 if batch_length_i32 < 49 {
957 return Err(KrafkaError::protocol_kind(
958 ProtocolErrorKind::InvalidValue,
959 format!("invalid record batch length: {batch_length_i32}"),
960 ));
961 }
962
963 let batch_length = batch_length_i32 as usize;
964
965 if buf.remaining() < batch_length {
966 return Err(KrafkaError::protocol_kind(
967 ProtocolErrorKind::TruncatedFrame,
968 "not enough bytes for record batch",
969 ));
970 }
971
972 let partition_leader_epoch = buf.get_i32();
973 let magic = buf.get_i8();
974
975 if magic != 2 {
976 return Err(KrafkaError::protocol_kind(
977 ProtocolErrorKind::UnsupportedMagic,
978 format!("unsupported record batch magic: {magic}"),
979 ));
980 }
981
982 let crc = buf.get_u32();
983
984 let crc_covered_len = batch_length - 9;
1000 let crc_covered = buf.copy_to_bytes(crc_covered_len);
1001
1002 let computed_crc = crc32c(&crc_covered);
1003 if computed_crc != crc {
1004 return Err(KrafkaError::protocol_kind(
1005 ProtocolErrorKind::CrcMismatch,
1006 format!("CRC mismatch: expected {crc:08x}, got {computed_crc:08x}"),
1007 ));
1008 }
1009
1010 let mut cbuf = crc_covered;
1013 let attributes = RecordBatchAttributes::from_i16(cbuf.get_i16())?;
1014 let last_offset_delta = cbuf.get_i32();
1015 let base_timestamp = cbuf.get_i64();
1016 let max_timestamp = cbuf.get_i64();
1017 let producer_id = cbuf.get_i64();
1018 let producer_epoch = cbuf.get_i16();
1019 let base_sequence = cbuf.get_i32();
1020 let records_count = cbuf.get_i32();
1021
1022 if records_count < 0 {
1023 return Err(KrafkaError::protocol_kind(
1024 ProtocolErrorKind::InvalidValue,
1025 format!("invalid negative records count: {records_count}"),
1026 ));
1027 }
1028
1029 let compressed_records = cbuf;
1032
1033 let decompressed = Self::decompress_records(
1035 attributes.compression,
1036 &compressed_records,
1037 max_decompressed_size,
1038 )?;
1039 let mut records_buf = decompressed.as_ref();
1040
1041 let records_len = records_count as usize;
1044 if records_len > super::MAX_DECODE_ARRAY_LEN {
1045 return Err(KrafkaError::protocol_kind(
1046 ProtocolErrorKind::InvalidLength,
1047 format!(
1048 "records count {records_len} exceeds safety limit {}",
1049 super::MAX_DECODE_ARRAY_LEN
1050 ),
1051 ));
1052 }
1053 let mut records = Vec::with_capacity(decode_capacity(records_len, records_buf.len()));
1056 for _ in 0..records_len {
1057 records.push(Record::decode(&mut records_buf)?);
1058 }
1059
1060 Ok(Self {
1061 base_offset,
1062 partition_leader_epoch,
1063 magic,
1064 attributes,
1065 last_offset_delta,
1066 base_timestamp,
1067 max_timestamp,
1068 producer_id,
1069 producer_epoch,
1070 base_sequence,
1071 records,
1072 compression_level: None,
1073 })
1074 }
1075
1076 pub const MAX_DECOMPRESSED_SIZE: usize = 128 * 1024 * 1024;
1082
1083 fn decompress_records(
1090 compression: Compression,
1091 data: &Bytes,
1092 _max_decompressed_size: usize,
1093 ) -> Result<Bytes> {
1094 #[allow(unused_variables)]
1096 let compressed: &[u8] = data.as_ref();
1097 #[allow(unused_variables)]
1098 let result: Vec<u8> = match compression {
1099 Compression::None => return Ok(data.clone()),
1102 #[cfg(feature = "gzip")]
1103 Compression::Gzip => {
1104 use flate2::read::GzDecoder;
1105 use std::io::Read;
1106
1107 let decoder = GzDecoder::new(compressed);
1108 let mut limited = decoder.take(_max_decompressed_size as u64 + 1);
1109 let capacity = compressed
1110 .len()
1111 .saturating_mul(3)
1112 .min(_max_decompressed_size);
1113 let mut decompressed = Vec::with_capacity(capacity);
1114 limited
1115 .read_to_end(&mut decompressed)
1116 .map_err(|e| KrafkaError::compression(e.to_string()))?;
1117 decompressed
1118 }
1119 #[cfg(not(feature = "gzip"))]
1120 Compression::Gzip => {
1121 return Err(KrafkaError::compression(
1122 "gzip decompression requires the `gzip` Cargo feature",
1123 ));
1124 }
1125 #[cfg(feature = "snappy")]
1126 Compression::Snappy => {
1127 let declared_len = snap::raw::decompress_len(compressed)
1130 .map_err(|e| KrafkaError::compression(e.to_string()))?;
1131 if declared_len > _max_decompressed_size {
1132 return Err(KrafkaError::compression(format!(
1133 "snappy declared decompressed size {} exceeds maximum {} bytes (possible compression bomb)",
1134 declared_len, _max_decompressed_size
1135 )));
1136 }
1137 let mut decoder = snap::raw::Decoder::new();
1138 decoder
1139 .decompress_vec(compressed)
1140 .map_err(|e| KrafkaError::compression(e.to_string()))?
1141 }
1142 #[cfg(not(feature = "snappy"))]
1143 Compression::Snappy => {
1144 return Err(KrafkaError::compression(
1145 "snappy decompression requires the `snappy` Cargo feature",
1146 ));
1147 }
1148 #[cfg(feature = "lz4")]
1149 Compression::Lz4 => {
1150 use std::io::Read;
1151 let decoder = lz4_flex::frame::FrameDecoder::new(compressed);
1152 let mut limited = decoder.take(_max_decompressed_size as u64 + 1);
1153 let capacity = compressed
1154 .len()
1155 .saturating_mul(4)
1156 .min(_max_decompressed_size);
1157 let mut decompressed = Vec::with_capacity(capacity);
1158 limited
1159 .read_to_end(&mut decompressed)
1160 .map_err(|e| KrafkaError::compression(e.to_string()))?;
1161 decompressed
1162 }
1163 #[cfg(not(feature = "lz4"))]
1164 Compression::Lz4 => {
1165 return Err(KrafkaError::compression(
1166 "lz4 decompression requires the `lz4` Cargo feature",
1167 ));
1168 }
1169 #[cfg(feature = "zstd")]
1170 Compression::Zstd => {
1171 use std::io::Read;
1174 let decoder = zstd::Decoder::new(compressed)
1175 .map_err(|e| KrafkaError::compression(e.to_string()))?;
1176 let mut limited = decoder.take(_max_decompressed_size as u64 + 1);
1177 let capacity = compressed
1178 .len()
1179 .saturating_mul(3)
1180 .min(_max_decompressed_size);
1181 let mut decompressed = Vec::with_capacity(capacity);
1182 limited
1183 .read_to_end(&mut decompressed)
1184 .map_err(|e| KrafkaError::compression(e.to_string()))?;
1185 decompressed
1186 }
1187 #[cfg(not(feature = "zstd"))]
1188 Compression::Zstd => {
1189 return Err(KrafkaError::compression(
1190 "zstd decompression requires the `zstd` Cargo feature",
1191 ));
1192 }
1193 };
1194
1195 #[allow(unreachable_code)]
1198 {
1199 if result.len() > _max_decompressed_size {
1200 return Err(KrafkaError::compression(format!(
1201 "decompressed size {} exceeds maximum {} bytes (possible compression bomb)",
1202 result.len(),
1203 _max_decompressed_size
1204 )));
1205 }
1206
1207 Ok(Bytes::from(result))
1208 }
1209 }
1210}
1211
1212impl Default for RecordBatch {
1213 fn default() -> Self {
1214 Self::new()
1215 }
1216}
1217
1218#[must_use = "builders do nothing until .build() is called"]
1220#[derive(Debug, Default)]
1221pub struct RecordBatchBuilder {
1222 compression: Compression,
1223 compression_level: Option<i32>,
1224 records: Vec<Record>,
1225 base_timestamp: Option<i64>,
1226 producer_id: i64,
1227 producer_epoch: i16,
1228 base_sequence: i32,
1229 is_transactional: bool,
1230}
1231
1232impl RecordBatchBuilder {
1233 pub fn new() -> Self {
1235 Self {
1236 compression: Compression::None,
1237 compression_level: None,
1238 records: Vec::new(),
1239 base_timestamp: None,
1240 producer_id: -1,
1241 producer_epoch: -1,
1242 base_sequence: -1,
1243 is_transactional: false,
1244 }
1245 }
1246
1247 pub fn compression(mut self, compression: Compression) -> Self {
1249 self.compression = compression;
1250 self
1251 }
1252
1253 pub fn compression_level(mut self, level: Option<i32>) -> Self {
1259 self.compression_level = level;
1260 self
1261 }
1262
1263 pub fn producer(mut self, id: i64, epoch: i16, sequence: i32) -> Self {
1265 self.producer_id = id;
1266 self.producer_epoch = epoch;
1267 self.base_sequence = sequence;
1268 self
1269 }
1270
1271 pub fn transactional(mut self, is_transactional: bool) -> Self {
1276 self.is_transactional = is_transactional;
1277 self
1278 }
1279
1280 pub fn base_timestamp(mut self, timestamp: i64) -> Self {
1282 self.base_timestamp = Some(timestamp);
1283 self
1284 }
1285
1286 pub fn add_record(
1288 mut self,
1289 key: Option<impl Into<Bytes>>,
1290 value: Option<impl Into<Bytes>>,
1291 ) -> Self {
1292 debug_assert!(
1293 self.records.len() < i32::MAX as usize,
1294 "batch record count would overflow i32"
1295 );
1296 let offset_delta = self.records.len() as i32;
1297 let record =
1298 Record::new(key.map(Into::into), value.map(Into::into)).with_offset_delta(offset_delta);
1299 self.records.push(record);
1300 self
1301 }
1302
1303 pub fn add_record_with_headers(
1309 mut self,
1310 key: Option<impl Into<Bytes>>,
1311 value: Option<impl Into<Bytes>>,
1312 headers: Vec<(impl Into<Bytes>, Option<impl Into<Bytes>>)>,
1313 ) -> Self {
1314 debug_assert!(
1315 self.records.len() < i32::MAX as usize,
1316 "batch record count would overflow i32"
1317 );
1318 let offset_delta = self.records.len() as i32;
1319 let mut record =
1320 Record::new(key.map(Into::into), value.map(Into::into)).with_offset_delta(offset_delta);
1321 for (k, v) in headers {
1322 record.headers.push(RecordHeader {
1323 key: k.into(),
1324 value: v.map(Into::into),
1325 });
1326 }
1327 self.records.push(record);
1328 self
1329 }
1330
1331 pub fn build(self) -> RecordBatch {
1333 let now = std::time::SystemTime::now()
1334 .duration_since(std::time::UNIX_EPOCH)
1335 .map(|d| d.as_millis() as i64)
1336 .unwrap_or(0);
1337
1338 let base_timestamp = self.base_timestamp.unwrap_or(now);
1339 let last_offset_delta = self.records.len().saturating_sub(1) as i32;
1340
1341 RecordBatch {
1342 base_offset: 0,
1343 partition_leader_epoch: 0,
1344 magic: 2,
1345 attributes: RecordBatchAttributes {
1346 compression: self.compression,
1347 timestamp_type: TimestampType::CreateTime,
1348 is_transactional: self.is_transactional,
1349 is_control_batch: false,
1350 },
1351 last_offset_delta,
1352 base_timestamp,
1353 max_timestamp: base_timestamp,
1354 producer_id: self.producer_id,
1355 producer_epoch: self.producer_epoch,
1356 base_sequence: self.base_sequence,
1357 records: self.records,
1358 compression_level: self.compression_level,
1359 }
1360 }
1361}
1362
1363#[must_use = "contains lazily-decoded record batch data"]
1380#[derive(Debug, Clone)]
1381pub struct LazyRecordBatch {
1382 pub base_offset: i64,
1384 pub partition_leader_epoch: i32,
1386 pub attributes: RecordBatchAttributes,
1388 pub last_offset_delta: i32,
1390 pub base_timestamp: i64,
1392 pub max_timestamp: i64,
1394 pub producer_id: i64,
1396 pub producer_epoch: i16,
1398 pub base_sequence: i32,
1400 pub records_count: i32,
1402 raw_records: Bytes,
1404}
1405
1406impl LazyRecordBatch {
1407 pub fn decode(buf: &mut impl Buf) -> Result<Self> {
1413 Self::decode_with_limit(buf, RecordBatch::MAX_DECOMPRESSED_SIZE)
1414 }
1415
1416 pub fn decode_with_limit(buf: &mut impl Buf, max_decompressed_size: usize) -> Result<Self> {
1421 if buf.remaining() < 12 {
1422 return Err(KrafkaError::protocol_kind(
1423 ProtocolErrorKind::TruncatedFrame,
1424 "not enough bytes for record batch header",
1425 ));
1426 }
1427
1428 let base_offset = buf.get_i64();
1429 let batch_length_i32 = buf.get_i32();
1430
1431 if batch_length_i32 < 49 {
1432 return Err(KrafkaError::protocol_kind(
1433 ProtocolErrorKind::InvalidValue,
1434 format!("invalid record batch length: {batch_length_i32}"),
1435 ));
1436 }
1437
1438 let batch_length = batch_length_i32 as usize;
1439
1440 if buf.remaining() < batch_length {
1441 return Err(KrafkaError::protocol_kind(
1442 ProtocolErrorKind::TruncatedFrame,
1443 "not enough bytes for record batch",
1444 ));
1445 }
1446
1447 let partition_leader_epoch = buf.get_i32();
1448 let magic = buf.get_i8();
1449
1450 if magic != 2 {
1451 return Err(KrafkaError::protocol_kind(
1452 ProtocolErrorKind::UnsupportedMagic,
1453 format!("unsupported record batch magic: {magic}"),
1454 ));
1455 }
1456
1457 let crc = buf.get_u32();
1458
1459 let crc_covered_len = batch_length - 9;
1463 let crc_covered = buf.copy_to_bytes(crc_covered_len);
1464
1465 let computed_crc = crc32c(&crc_covered);
1466 if computed_crc != crc {
1467 return Err(KrafkaError::protocol_kind(
1468 ProtocolErrorKind::CrcMismatch,
1469 format!("CRC mismatch: expected {crc:08x}, got {computed_crc:08x}"),
1470 ));
1471 }
1472
1473 let mut cbuf = crc_covered;
1474 let attributes = RecordBatchAttributes::from_i16(cbuf.get_i16())?;
1475 let last_offset_delta = cbuf.get_i32();
1476 let base_timestamp = cbuf.get_i64();
1477 let max_timestamp = cbuf.get_i64();
1478 let producer_id = cbuf.get_i64();
1479 let producer_epoch = cbuf.get_i16();
1480 let base_sequence = cbuf.get_i32();
1481 let records_count = cbuf.get_i32();
1482
1483 if records_count < 0 {
1484 return Err(KrafkaError::protocol_kind(
1485 ProtocolErrorKind::InvalidValue,
1486 format!("invalid negative records count: {records_count}"),
1487 ));
1488 }
1489 if records_count as usize > super::MAX_DECODE_ARRAY_LEN {
1490 return Err(KrafkaError::protocol_kind(
1491 ProtocolErrorKind::InvalidLength,
1492 format!(
1493 "records count {} exceeds safety limit {}",
1494 records_count,
1495 super::MAX_DECODE_ARRAY_LEN
1496 ),
1497 ));
1498 }
1499
1500 let compressed_records = cbuf;
1502
1503 let raw_records = RecordBatch::decompress_records(
1505 attributes.compression,
1506 &compressed_records,
1507 max_decompressed_size,
1508 )?;
1509
1510 Ok(Self {
1511 base_offset,
1512 partition_leader_epoch,
1513 attributes,
1514 last_offset_delta,
1515 base_timestamp,
1516 max_timestamp,
1517 producer_id,
1518 producer_epoch,
1519 base_sequence,
1520 records_count,
1521 raw_records,
1522 })
1523 }
1524
1525 #[inline]
1527 pub fn len(&self) -> usize {
1528 self.records_count as usize
1529 }
1530
1531 #[inline]
1533 pub fn is_empty(&self) -> bool {
1534 self.records_count == 0
1535 }
1536
1537 #[inline]
1541 pub fn records(&self) -> LazyRecordIterator {
1542 LazyRecordIterator {
1543 buf: self.raw_records.clone(),
1544 remaining: self.records_count as usize,
1545 }
1546 }
1547
1548 pub fn decode_all(&self) -> Result<Vec<Record>> {
1555 let mut records = Vec::with_capacity(decode_capacity(
1558 (self.records_count as usize).min(super::MAX_DECODE_ARRAY_LEN),
1559 self.raw_records.len(),
1560 ));
1561 for result in self.records() {
1562 records.push(result?);
1563 }
1564 Ok(records)
1565 }
1566
1567 pub fn into_record_batch(self) -> Result<RecordBatch> {
1569 Ok(RecordBatch {
1570 base_offset: self.base_offset,
1571 partition_leader_epoch: self.partition_leader_epoch,
1572 magic: 2,
1573 attributes: self.attributes,
1574 last_offset_delta: self.last_offset_delta,
1575 base_timestamp: self.base_timestamp,
1576 max_timestamp: self.max_timestamp,
1577 producer_id: self.producer_id,
1578 producer_epoch: self.producer_epoch,
1579 base_sequence: self.base_sequence,
1580 records: self.decode_all()?,
1581 compression_level: None,
1582 })
1583 }
1584}
1585
1586#[must_use = "iterators are lazy and do nothing unless consumed"]
1588pub struct LazyRecordIterator {
1589 buf: Bytes,
1590 remaining: usize,
1591}
1592
1593impl Iterator for LazyRecordIterator {
1594 type Item = Result<Record>;
1595
1596 #[inline]
1605 fn next(&mut self) -> Option<Self::Item> {
1606 if self.remaining == 0 {
1607 return None;
1608 }
1609 if self.buf.is_empty() {
1610 self.remaining = 0;
1613 return Some(Err(KrafkaError::protocol_kind(
1614 ProtocolErrorKind::TruncatedFrame,
1615 "record batch declares more records than the buffer contains",
1616 )));
1617 }
1618 self.remaining -= 1;
1619 Some(Record::decode(&mut self.buf))
1620 }
1621
1622 fn size_hint(&self) -> (usize, Option<usize>) {
1629 (0, Some(self.remaining))
1630 }
1631}
1632
1633#[cfg(test)]
1634#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
1635mod tests {
1636
1637 #[cfg(any(feature = "zstd", feature = "gzip"))]
1645 fn compressible_payload() -> Vec<u8> {
1646 let mut out = String::new();
1647 let mut x: u64 = 0x2545_F491_4F6C_DD1D;
1648 for i in 0..2_000 {
1649 x ^= x << 13;
1650 x ^= x >> 7;
1651 x ^= x << 17;
1652 out.push_str(&format!(
1653 "{{\"id\":{i},\"user\":\"u{}\",\"ts\":{},\"evt\":\"click\",\"v\":{}}}\n",
1654 x % 100_000,
1655 1_700_000_000_000u64 + (x % 1_000_000),
1656 x % 997
1657 ));
1658 }
1659 out.into_bytes()
1660 }
1661
1662 #[cfg(feature = "zstd")]
1670 #[test]
1671 fn zstd_compression_level_changes_the_encoded_bytes() {
1672 let payload = compressible_payload();
1673
1674 let fast = Compression::Zstd
1675 .compress_with_level(&payload, Some(1))
1676 .expect("level 1 must encode");
1677 let dense = Compression::Zstd
1678 .compress_with_level(&payload, Some(19))
1679 .expect("level 19 must encode");
1680 let default = Compression::Zstd
1681 .compress_with_level(&payload, None)
1682 .expect("default must encode");
1683
1684 assert!(
1690 dense.len() < fast.len(),
1691 "level 19 must compress better than level 1, got {} vs {}",
1692 dense.len(),
1693 fast.len()
1694 );
1695 assert_eq!(
1696 default,
1697 Compression::Zstd
1698 .compress_with_level(&payload, Some(3))
1699 .expect("level 3 must encode"),
1700 "the documented default is 3; if that changes, the docs are wrong"
1701 );
1702 }
1703
1704 #[cfg(feature = "gzip")]
1705 #[test]
1706 fn gzip_compression_level_changes_the_encoded_bytes() {
1707 let payload = compressible_payload();
1708
1709 let none = Compression::Gzip
1710 .compress_with_level(&payload, Some(0))
1711 .expect("level 0 must encode");
1712 let best = Compression::Gzip
1713 .compress_with_level(&payload, Some(9))
1714 .expect("level 9 must encode");
1715
1716 assert!(
1717 best.len() < none.len(),
1718 "level 9 must compress better than level 0 (store), got {} vs {}",
1719 best.len(),
1720 none.len()
1721 );
1722 assert_eq!(
1723 Compression::Gzip
1724 .compress_with_level(&payload, None)
1725 .expect("default must encode"),
1726 Compression::Gzip
1727 .compress_with_level(&payload, Some(6))
1728 .expect("level 6 must encode"),
1729 "zlib's default is 6; if that changes, the docs are wrong"
1730 );
1731 }
1732
1733 #[cfg(feature = "zstd")]
1735 #[test]
1736 fn record_batch_carries_the_compression_level_to_the_wire() {
1737 let payload = compressible_payload();
1738
1739 let encode_at = |level: Option<i32>| {
1740 RecordBatchBuilder::new()
1741 .compression(Compression::Zstd)
1742 .compression_level(level)
1743 .add_record(None::<Bytes>, Some(Bytes::from(payload.clone())))
1744 .build()
1745 .encode()
1746 .expect("batch must encode")
1747 };
1748
1749 let fast = encode_at(Some(1));
1750 let dense = encode_at(Some(19));
1751 assert!(
1752 dense.len() < fast.len(),
1753 "the level must reach the codec through RecordBatch::encode, got {} vs {}",
1754 dense.len(),
1755 fast.len()
1756 );
1757 }
1758
1759 #[test]
1762 fn only_gzip_and_zstd_accept_a_level() {
1763 assert!(Compression::Gzip.supports_level());
1764 assert!(Compression::Zstd.supports_level());
1765 assert!(!Compression::Snappy.supports_level());
1766 assert!(!Compression::Lz4.supports_level());
1767 assert!(!Compression::None.supports_level());
1768
1769 assert_eq!(Compression::Gzip.level_range(), Some(0..=9));
1770 assert!(Compression::Snappy.level_range().is_none());
1771 }
1772 use super::*;
1773
1774 #[test]
1775 fn test_record_encode_decode() {
1776 let record = Record::new(Some(Bytes::from("key")), Some(Bytes::from("value")))
1777 .with_timestamp_delta(100)
1778 .with_offset_delta(0)
1779 .with_header("header1", Bytes::from("value1"));
1780
1781 let mut buf = BytesMut::new();
1782 record.encode(&mut buf).unwrap();
1783
1784 let decoded = Record::decode(&mut buf.freeze()).unwrap();
1785 assert_eq!(decoded.key, Some(Bytes::from("key")));
1786 assert_eq!(decoded.value, Some(Bytes::from("value")));
1787 assert_eq!(decoded.timestamp_delta, 100);
1788 assert_eq!(decoded.offset_delta, 0);
1789 assert_eq!(decoded.headers.len(), 1);
1790 assert_eq!(decoded.headers[0].key, "header1");
1791 }
1792
1793 #[test]
1794 fn test_record_null_key_value() {
1795 let record = Record::new(None, Some(Bytes::from("value")));
1796
1797 let mut buf = BytesMut::new();
1798 record.encode(&mut buf).unwrap();
1799
1800 let decoded = Record::decode(&mut buf.freeze()).unwrap();
1801 assert!(decoded.key.is_none());
1802 assert_eq!(decoded.value, Some(Bytes::from("value")));
1803 }
1804
1805 #[test]
1806 fn test_record_batch_builder() {
1807 let batch = RecordBatchBuilder::new()
1808 .compression(Compression::None)
1809 .add_record(Some("key1"), Some("value1"))
1810 .add_record(Some("key2"), Some("value2"))
1811 .build();
1812
1813 assert_eq!(batch.records.len(), 2);
1814 assert_eq!(batch.last_offset_delta, 1);
1815 }
1816
1817 #[test]
1818 fn test_record_batch_encode_decode() {
1819 let batch = RecordBatchBuilder::new()
1820 .base_timestamp(1234567890000)
1821 .add_record(Some("key"), Some("value"))
1822 .build();
1823
1824 let encoded = batch.encode().unwrap();
1825 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
1826
1827 assert_eq!(decoded.base_offset, 0);
1828 assert_eq!(decoded.base_timestamp, 1234567890000);
1829 assert_eq!(decoded.records.len(), 1);
1830 assert_eq!(decoded.records[0].key, Some(Bytes::from("key")));
1831 assert_eq!(decoded.records[0].value, Some(Bytes::from("value")));
1832 }
1833
1834 #[test]
1835 #[cfg(feature = "gzip")]
1836 fn test_record_batch_compression_gzip() {
1837 let batch = RecordBatchBuilder::new()
1838 .compression(Compression::Gzip)
1839 .base_timestamp(1234567890000)
1840 .add_record(Some("key"), Some("value"))
1841 .build();
1842
1843 let encoded = batch.encode().unwrap();
1844 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
1845
1846 assert_eq!(decoded.records.len(), 1);
1847 assert_eq!(decoded.records[0].key, Some(Bytes::from("key")));
1848 }
1849
1850 #[test]
1851 #[cfg(feature = "snappy")]
1852 fn test_record_batch_compression_snappy() {
1853 let batch = RecordBatchBuilder::new()
1854 .compression(Compression::Snappy)
1855 .base_timestamp(1234567890000)
1856 .add_record(Some("key"), Some("value"))
1857 .build();
1858
1859 let encoded = batch.encode().unwrap();
1860 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
1861
1862 assert_eq!(decoded.records.len(), 1);
1863 assert_eq!(decoded.records[0].key, Some(Bytes::from("key")));
1864 }
1865
1866 #[test]
1867 #[cfg(feature = "lz4")]
1868 fn test_record_batch_compression_lz4() {
1869 let batch = RecordBatchBuilder::new()
1870 .compression(Compression::Lz4)
1871 .base_timestamp(1234567890000)
1872 .add_record(Some("key"), Some("value"))
1873 .build();
1874
1875 let encoded = batch.encode().unwrap();
1876 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
1877
1878 assert_eq!(decoded.records.len(), 1);
1879 assert_eq!(decoded.records[0].key, Some(Bytes::from("key")));
1880 }
1881
1882 #[test]
1883 #[cfg(feature = "zstd")]
1884 fn test_record_batch_compression_zstd() {
1885 let batch = RecordBatchBuilder::new()
1886 .compression(Compression::Zstd)
1887 .base_timestamp(1234567890000)
1888 .add_record(Some("key"), Some("value"))
1889 .build();
1890
1891 let encoded = batch.encode().unwrap();
1892 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
1893
1894 assert_eq!(decoded.records.len(), 1);
1895 assert_eq!(decoded.records[0].key, Some(Bytes::from("key")));
1896 }
1897
1898 #[test]
1899 fn test_compression_is_available() {
1900 assert!(Compression::None.is_available());
1902
1903 assert_eq!(Compression::Gzip.is_available(), cfg!(feature = "gzip"));
1905 assert_eq!(Compression::Snappy.is_available(), cfg!(feature = "snappy"));
1906 assert_eq!(Compression::Lz4.is_available(), cfg!(feature = "lz4"));
1907 assert_eq!(Compression::Zstd.is_available(), cfg!(feature = "zstd"));
1908 }
1909
1910 #[test]
1911 fn test_compression_required_feature() {
1912 assert_eq!(Compression::None.required_feature(), None);
1913 assert_eq!(Compression::Gzip.required_feature(), Some("gzip"));
1914 assert_eq!(Compression::Snappy.required_feature(), Some("snappy"));
1915 assert_eq!(Compression::Lz4.required_feature(), Some("lz4"));
1916 assert_eq!(Compression::Zstd.required_feature(), Some("zstd"));
1917 }
1918
1919 #[test]
1920 fn test_disabled_codec_returns_error() {
1921 for compression in [
1924 Compression::Gzip,
1925 Compression::Snappy,
1926 Compression::Lz4,
1927 Compression::Zstd,
1928 ] {
1929 if compression.is_available() {
1930 continue;
1931 }
1932 let batch = RecordBatchBuilder::new()
1933 .compression(compression)
1934 .add_record(Some("k"), Some("v"))
1935 .build();
1936 let err = batch.encode().unwrap_err();
1937 let msg = err.to_string();
1938 let feature = compression.required_feature().unwrap();
1939 assert!(
1940 msg.contains(feature),
1941 "error for {compression:?} should mention feature `{feature}`, got: {msg}"
1942 );
1943 }
1944 }
1945
1946 #[test]
1947 fn test_compression_roundtrip() {
1948 #[allow(clippy::single_element_loop)]
1949 for compression in [
1950 Compression::None,
1951 #[cfg(feature = "gzip")]
1952 Compression::Gzip,
1953 #[cfg(feature = "snappy")]
1954 Compression::Snappy,
1955 #[cfg(feature = "lz4")]
1956 Compression::Lz4,
1957 #[cfg(feature = "zstd")]
1958 Compression::Zstd,
1959 ] {
1960 let batch = RecordBatchBuilder::new()
1961 .compression(compression)
1962 .base_timestamp(1234567890000)
1963 .add_record(Some("key1"), Some("value1"))
1964 .add_record(Some("key2"), Some("value2"))
1965 .add_record(Some("key3"), Some("value3"))
1966 .build();
1967
1968 let encoded = batch.encode().unwrap();
1969 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
1970
1971 assert_eq!(
1972 decoded.records.len(),
1973 3,
1974 "Failed for compression {compression:?}"
1975 );
1976 }
1977 }
1978
1979 #[test]
1980 fn test_record_batch_attributes() {
1981 let attrs = RecordBatchAttributes {
1982 compression: Compression::Lz4,
1983 timestamp_type: TimestampType::LogAppendTime,
1984 is_transactional: true,
1985 is_control_batch: false,
1986 };
1987
1988 let raw = attrs.to_i16();
1989 let decoded = RecordBatchAttributes::from_i16(raw).unwrap();
1990
1991 assert_eq!(decoded.compression, Compression::Lz4);
1992 assert_eq!(decoded.timestamp_type, TimestampType::LogAppendTime);
1993 assert!(decoded.is_transactional);
1994 assert!(!decoded.is_control_batch);
1995 }
1996
1997 #[test]
1998 fn test_record_batch_attributes_rejects_unknown_compression_discriminant() {
1999 let err = RecordBatchAttributes::from_i16(0x0005).unwrap_err();
2001 match err {
2002 KrafkaError::Protocol { kind, .. } => {
2003 assert_eq!(kind, crate::error::ProtocolErrorKind::InvalidValue)
2004 }
2005 other => panic!("expected protocol invalid-value error, got: {other}"),
2006 }
2007 }
2008
2009 #[test]
2010 fn test_lazy_record_batch_decode() {
2011 let batch = RecordBatchBuilder::new()
2012 .compression(Compression::None)
2013 .base_timestamp(1234567890000)
2014 .add_record(Some("key1"), Some("value1"))
2015 .add_record(Some("key2"), Some("value2"))
2016 .add_record(Some("key3"), Some("value3"))
2017 .build();
2018
2019 let encoded = batch.encode().unwrap();
2020 let lazy = LazyRecordBatch::decode(&mut encoded.clone()).unwrap();
2021
2022 assert_eq!(lazy.len(), 3);
2023 assert!(!lazy.is_empty());
2024 assert_eq!(lazy.base_timestamp, 1234567890000);
2025
2026 let records: Vec<Record> = lazy.records().map(|r| r.unwrap()).collect();
2028 assert_eq!(records.len(), 3);
2029 assert_eq!(records[0].key, Some(Bytes::from("key1")));
2030 assert_eq!(records[1].key, Some(Bytes::from("key2")));
2031 assert_eq!(records[2].key, Some(Bytes::from("key3")));
2032 }
2033
2034 #[test]
2035 #[cfg(feature = "lz4")]
2036 fn test_lazy_record_batch_into_eager() {
2037 let batch = RecordBatchBuilder::new()
2038 .compression(Compression::Lz4)
2039 .base_timestamp(1234567890000)
2040 .add_record(Some("key"), Some("value"))
2041 .build();
2042
2043 let encoded = batch.encode().unwrap();
2044 let lazy = LazyRecordBatch::decode(&mut encoded.clone()).unwrap();
2045 let eager = lazy.into_record_batch().unwrap();
2046
2047 assert_eq!(eager.records.len(), 1);
2048 assert_eq!(eager.records[0].key, Some(Bytes::from("key")));
2049 assert_eq!(eager.base_timestamp, 1234567890000);
2050 }
2051
2052 #[test]
2053 fn test_lazy_record_batch_with_compression() {
2054 #[allow(clippy::single_element_loop)]
2055 for compression in [
2056 Compression::None,
2057 #[cfg(feature = "gzip")]
2058 Compression::Gzip,
2059 #[cfg(feature = "snappy")]
2060 Compression::Snappy,
2061 #[cfg(feature = "lz4")]
2062 Compression::Lz4,
2063 #[cfg(feature = "zstd")]
2064 Compression::Zstd,
2065 ] {
2066 let batch = RecordBatchBuilder::new()
2067 .compression(compression)
2068 .base_timestamp(1234567890000)
2069 .add_record(Some("key1"), Some("value1"))
2070 .add_record(Some("key2"), Some("value2"))
2071 .build();
2072
2073 let encoded = batch.encode().unwrap();
2074 let lazy = LazyRecordBatch::decode(&mut encoded.clone()).unwrap();
2075
2076 assert_eq!(lazy.len(), 2, "Failed for compression {compression:?}");
2077
2078 let records: Result<Vec<_>> = lazy.records().collect();
2079 let records = records.unwrap();
2080 assert_eq!(records.len(), 2, "Failed for compression {compression:?}");
2081 }
2082 }
2083
2084 #[test]
2085 #[cfg(feature = "gzip")]
2086 fn test_decompress_normal_data_within_limit() {
2087 let batch = RecordBatchBuilder::new()
2089 .compression(Compression::Gzip)
2090 .add_record(Some("key"), Some("value"))
2091 .build();
2092
2093 let encoded = batch.encode().unwrap();
2094 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
2095 assert_eq!(decoded.records.len(), 1);
2096 }
2097
2098 #[test]
2099 fn test_max_decompressed_size_constant() {
2100 assert_eq!(RecordBatch::MAX_DECOMPRESSED_SIZE, 128 * 1024 * 1024);
2102 }
2103
2104 #[test]
2105 #[cfg(feature = "snappy")]
2106 fn test_snappy_decompression_bomb_rejected() {
2107 let huge_size: u64 = 256 * 1024 * 1024;
2111 let mut fake_snappy = Vec::new();
2113 let mut val = huge_size;
2114 while val >= 0x80 {
2115 fake_snappy.push((val as u8) | 0x80);
2116 val >>= 7;
2117 }
2118 fake_snappy.push(val as u8);
2119 fake_snappy.extend_from_slice(&[0u8; 16]);
2121
2122 let result = RecordBatch::decompress_records(
2123 Compression::Snappy,
2124 &Bytes::from(fake_snappy),
2125 RecordBatch::MAX_DECOMPRESSED_SIZE,
2126 );
2127 assert!(result.is_err());
2128 let err_msg = result.unwrap_err().to_string();
2129 assert!(
2130 err_msg.contains("compression bomb") || err_msg.contains("exceeds maximum"),
2131 "Error should mention size limit: {err_msg}"
2132 );
2133 }
2134
2135 #[test]
2136 #[cfg(feature = "zstd")]
2137 fn test_zstd_decompression_uses_streaming_limit() {
2138 let batch = RecordBatchBuilder::new()
2141 .compression(Compression::Zstd)
2142 .add_record(Some("key"), Some("value"))
2143 .build();
2144
2145 let encoded = batch.encode().unwrap();
2146 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
2147 assert_eq!(decoded.records.len(), 1);
2148 }
2149
2150 #[test]
2151 fn test_record_batch_builder_transactional_flag() {
2152 let batch = RecordBatchBuilder::new()
2153 .transactional(true)
2154 .add_record(Some("key"), Some("value"))
2155 .build();
2156
2157 assert!(batch.attributes.is_transactional);
2158
2159 let encoded = batch.encode().unwrap();
2161 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
2162 assert!(decoded.attributes.is_transactional);
2163 }
2164
2165 #[test]
2166 fn test_record_batch_builder_non_transactional_default() {
2167 let batch = RecordBatchBuilder::new()
2168 .add_record(Some("key"), Some("value"))
2169 .build();
2170
2171 assert!(!batch.attributes.is_transactional);
2172 }
2173
2174 #[test]
2175 fn test_record_batch_builder_producer_identity() {
2176 let batch = RecordBatchBuilder::new()
2177 .producer(12345, 7, 42)
2178 .transactional(true)
2179 .add_record(Some("key"), Some("value"))
2180 .build();
2181
2182 assert_eq!(batch.producer_id, 12345);
2183 assert_eq!(batch.producer_epoch, 7);
2184 assert_eq!(batch.base_sequence, 42);
2185 assert!(batch.attributes.is_transactional);
2186
2187 let encoded = batch.encode().unwrap();
2189 let decoded = RecordBatch::decode(&mut encoded.clone()).unwrap();
2190 assert_eq!(decoded.producer_id, 12345);
2191 assert_eq!(decoded.producer_epoch, 7);
2192 assert_eq!(decoded.base_sequence, 42);
2193 }
2194
2195 #[test]
2196 fn test_record_batch_attributes_transactional_bit() {
2197 let attrs = RecordBatchAttributes::from_i16(0x10).unwrap();
2199 assert!(attrs.is_transactional);
2200 assert!(!attrs.is_control_batch);
2201
2202 let raw = attrs.to_i16();
2203 assert_eq!(raw & 0x10, 0x10);
2204
2205 let attrs = RecordBatchAttributes::from_i16(0x00).unwrap();
2207 assert!(!attrs.is_transactional);
2208 }
2209
2210 #[test]
2211 fn test_record_batch_decode_rejects_negative_batch_length() {
2212 let mut buf = BytesMut::new();
2214 buf.put_i64(0); buf.put_i32(-1); let result = RecordBatch::decode(&mut buf.freeze());
2218 assert!(result.is_err(), "negative batch_length should be rejected");
2219 let err_msg = format!("{}", result.unwrap_err());
2220 assert!(
2221 err_msg.contains("invalid record batch length"),
2222 "error should mention invalid length: {err_msg}"
2223 );
2224 }
2225
2226 #[test]
2227 fn test_record_batch_decode_rejects_too_small_batch_length() {
2228 let mut buf = BytesMut::new();
2230 buf.put_i64(0); buf.put_i32(10); let result = RecordBatch::decode(&mut buf.freeze());
2234 assert!(result.is_err(), "batch_length < 49 should be rejected");
2235 }
2236
2237 #[test]
2238 fn test_lazy_record_batch_decode_rejects_negative_batch_length() {
2239 let mut buf = BytesMut::new();
2240 buf.put_i64(0); buf.put_i32(-100); let result = LazyRecordBatch::decode(&mut buf.freeze());
2244 assert!(result.is_err(), "negative batch_length should be rejected");
2245 }
2246
2247 #[test]
2248 fn test_record_batch_decode_rejects_negative_records_count() {
2249 let mut batch = RecordBatch::new();
2252 batch
2253 .records
2254 .push(Record::new(Some(Bytes::from("k")), Some(Bytes::from("v"))));
2255 let encoded = batch.encode().unwrap();
2256
2257 let mut tampered = BytesMut::from(encoded.as_ref());
2259 let rc_offset = 57;
2264 tampered[rc_offset..rc_offset + 4].copy_from_slice(&(-1i32).to_be_bytes());
2265
2266 let crc_data = &tampered[21..];
2269 let new_crc = crate::util::crc32c(crc_data);
2270 tampered[17..21].copy_from_slice(&new_crc.to_be_bytes());
2271
2272 let result = RecordBatch::decode(&mut tampered.freeze());
2273 assert!(result.is_err(), "negative records_count should be rejected");
2274 let err_msg = format!("{}", result.unwrap_err());
2275 assert!(
2276 err_msg.contains("negative records count"),
2277 "error should mention negative records count: {err_msg}"
2278 );
2279 }
2280
2281 #[test]
2282 fn test_lazy_record_batch_decode_rejects_negative_records_count() {
2283 let mut batch = RecordBatch::new();
2285 batch
2286 .records
2287 .push(Record::new(Some(Bytes::from("k")), Some(Bytes::from("v"))));
2288 let encoded = batch.encode().unwrap();
2289
2290 let mut tampered = BytesMut::from(encoded.as_ref());
2291 let rc_offset = 57;
2292 tampered[rc_offset..rc_offset + 4].copy_from_slice(&(-1i32).to_be_bytes());
2293
2294 let crc_data = &tampered[21..];
2296 let new_crc = crate::util::crc32c(crc_data);
2297 tampered[17..21].copy_from_slice(&new_crc.to_be_bytes());
2298
2299 let result = LazyRecordBatch::decode(&mut tampered.freeze());
2300 assert!(result.is_err(), "negative records_count should be rejected");
2301 }
2302
2303 #[test]
2304 fn test_kafka_bytes_encode_normal_size() {
2305 use crate::protocol::primitives::{KafkaBytes, TryEncode};
2307 let b = KafkaBytes::new(vec![1, 2, 3]);
2308 let mut buf = BytesMut::new();
2309 b.try_encode(&mut buf).unwrap();
2310 assert_eq!(buf.len(), 4 + 3); }
2312
2313 fn short_lazy_batch(declared: i32, actual: usize) -> LazyRecordBatch {
2318 let mut raw = BytesMut::new();
2319 for i in 0..actual {
2320 Record::new(Some(Bytes::from("k")), Some(Bytes::from("v")))
2321 .with_offset_delta(i as i32)
2322 .encode(&mut raw)
2323 .unwrap();
2324 }
2325 LazyRecordBatch {
2326 base_offset: 0,
2327 partition_leader_epoch: -1,
2328 attributes: RecordBatchAttributes::default(),
2329 last_offset_delta: 0,
2330 base_timestamp: 0,
2331 max_timestamp: 0,
2332 producer_id: -1,
2333 producer_epoch: -1,
2334 base_sequence: -1,
2335 records_count: declared,
2336 raw_records: raw.freeze(),
2337 }
2338 }
2339
2340 #[test]
2347 fn lazy_iterator_errors_on_truncated_records() {
2348 let lazy = short_lazy_batch(100, 3);
2349 let results: Vec<_> = lazy.records().collect();
2350
2351 assert_eq!(results.len(), 4, "expected 3 records + 1 error");
2353 for r in results.iter().take(3) {
2354 assert!(r.is_ok(), "first three records must decode");
2355 }
2356 let err = results[3].as_ref().unwrap_err();
2357 assert!(
2358 format!("{err}").contains("more records than"),
2359 "expected a truncated-frame error, got: {err}"
2360 );
2361 }
2362
2363 #[test]
2365 fn lazy_decode_all_errors_on_truncated_records() {
2366 assert!(short_lazy_batch(100, 3).decode_all().is_err());
2367 }
2368
2369 #[test]
2373 fn lazy_and_eager_agree_on_truncated_batch() {
2374 let lazy = short_lazy_batch(100, 3);
2375 let mut raw = lazy.raw_records.clone();
2376
2377 let mut eager_err = false;
2379 for _ in 0..100 {
2380 if Record::decode(&mut raw).is_err() {
2381 eager_err = true;
2382 break;
2383 }
2384 }
2385 assert!(eager_err, "eager path must reject the truncated batch");
2386 assert!(
2387 lazy.decode_all().is_err(),
2388 "lazy path must reject it too — the two must not disagree"
2389 );
2390 }
2391
2392 #[test]
2396 fn lazy_iterator_size_hint_lower_bound_is_zero() {
2397 let lazy = short_lazy_batch(100, 3);
2398 let it = lazy.records();
2399 assert_eq!(it.size_hint(), (0, Some(100)));
2400 }
2401
2402 #[test]
2404 fn lazy_iterator_exact_batch_has_no_error() {
2405 let lazy = short_lazy_batch(3, 3);
2406 let records: Vec<_> = lazy.records().collect();
2407 assert_eq!(records.len(), 3);
2408 assert!(records.iter().all(|r| r.is_ok()));
2409 assert_eq!(lazy.decode_all().unwrap().len(), 3);
2410 }
2411
2412 #[test]
2417 fn decompress_none_is_zero_copy() {
2418 let src = Bytes::from(vec![7u8; 4096]);
2419 let out = RecordBatch::decompress_records(
2420 Compression::None,
2421 &src,
2422 RecordBatch::MAX_DECOMPRESSED_SIZE,
2423 )
2424 .unwrap();
2425 assert_eq!(out, src);
2426 assert_eq!(
2427 out.as_ptr(),
2428 src.as_ptr(),
2429 "uncompressed path must not copy the record payload"
2430 );
2431 }
2432}