Skip to main content

kafrust_protocol/api/
fetch.rs

1use crate::codec::{DecodeLimits, Decoder, Encoder};
2use crate::error::{Error, Result};
3use crate::header::RequestHeader;
4use crate::record_batch::{decompress_record_batch_records_with_limit, RecordBatchCompression};
5
6pub const API_KEY: i16 = 1;
7
8#[derive(Debug, Clone, PartialEq, Eq)]
9pub struct FetchRequestV2 {
10    pub correlation_id: i32,
11    pub client_id: Option<String>,
12    pub replica_id: i32,
13    pub max_wait_ms: i32,
14    pub min_bytes: i32,
15    pub topics: Vec<FetchTopicV2>,
16}
17
18#[derive(Debug, Clone, PartialEq, Eq)]
19pub struct FetchRequestV4 {
20    pub correlation_id: i32,
21    pub client_id: Option<String>,
22    pub replica_id: i32,
23    pub max_wait_ms: i32,
24    pub min_bytes: i32,
25    pub max_bytes: i32,
26    pub isolation_level: i8,
27    pub topics: Vec<FetchTopicV2>,
28}
29
30impl FetchRequestV4 {
31    pub fn encode(&self) -> Result<Vec<u8>> {
32        let mut encoder = Encoder::new();
33        RequestHeader {
34            api_key: API_KEY,
35            api_version: 4,
36            correlation_id: self.correlation_id,
37            client_id: self.client_id.clone(),
38        }
39        .encode_v1(&mut encoder)?;
40        encoder.write_i32(self.replica_id);
41        encoder.write_i32(self.max_wait_ms);
42        encoder.write_i32(self.min_bytes);
43        encoder.write_i32(self.max_bytes);
44        encoder.write_i8(self.isolation_level);
45        encoder.write_array(Some(self.topics.as_slice()), |encoder, topic| {
46            topic.encode(encoder)
47        })?;
48        Ok(encoder.into_bytes())
49    }
50}
51
52impl FetchRequestV2 {
53    pub fn encode(&self) -> Result<Vec<u8>> {
54        let mut encoder = Encoder::new();
55        RequestHeader {
56            api_key: API_KEY,
57            api_version: 2,
58            correlation_id: self.correlation_id,
59            client_id: self.client_id.clone(),
60        }
61        .encode_v1(&mut encoder)?;
62        encoder.write_i32(self.replica_id);
63        encoder.write_i32(self.max_wait_ms);
64        encoder.write_i32(self.min_bytes);
65        encoder.write_array(Some(self.topics.as_slice()), |encoder, topic| {
66            topic.encode(encoder)
67        })?;
68        Ok(encoder.into_bytes())
69    }
70}
71
72#[derive(Debug, Clone, PartialEq, Eq)]
73pub struct FetchTopicV2 {
74    pub name: String,
75    pub partitions: Vec<FetchPartitionV2>,
76}
77
78impl FetchTopicV2 {
79    fn encode(&self, encoder: &mut Encoder) -> Result<()> {
80        encoder.write_string(&self.name)?;
81        encoder.write_array(Some(self.partitions.as_slice()), |encoder, partition| {
82            partition.encode(encoder)
83        })
84    }
85}
86
87#[derive(Debug, Clone, PartialEq, Eq)]
88pub struct FetchPartitionV2 {
89    pub partition_index: i32,
90    pub fetch_offset: i64,
91    pub max_bytes: i32,
92}
93
94impl FetchPartitionV2 {
95    fn encode(&self, encoder: &mut Encoder) -> Result<()> {
96        encoder.write_i32(self.partition_index);
97        encoder.write_i64(self.fetch_offset);
98        encoder.write_i32(self.max_bytes);
99        Ok(())
100    }
101}
102
103#[derive(Debug, Clone, PartialEq, Eq)]
104pub struct FetchResponseV2 {
105    pub throttle_time_ms: i32,
106    pub responses: Vec<FetchTopicResponseV2>,
107}
108
109#[derive(Debug, Clone, PartialEq, Eq)]
110pub struct FetchResponseV4 {
111    pub throttle_time_ms: i32,
112    pub responses: Vec<FetchTopicResponseV4>,
113}
114
115impl FetchResponseV4 {
116    pub fn decode_body(decoder: &mut Decoder<'_>) -> Result<Self> {
117        Ok(Self {
118            throttle_time_ms: decoder.read_i32()?,
119            responses: decoder
120                .read_array("fetch responses", FetchTopicResponseV4::decode)?
121                .unwrap_or_default(),
122        })
123    }
124}
125
126#[derive(Debug, Clone, PartialEq, Eq)]
127pub struct FetchTopicResponseV4 {
128    pub name: String,
129    pub partitions: Vec<FetchPartitionResponseV4>,
130}
131
132impl FetchTopicResponseV4 {
133    fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
134        Ok(Self {
135            name: decoder.read_string()?,
136            partitions: decoder
137                .read_array(
138                    "fetch partition responses",
139                    FetchPartitionResponseV4::decode,
140                )?
141                .unwrap_or_default(),
142        })
143    }
144}
145
146#[derive(Debug, Clone, PartialEq, Eq)]
147pub struct FetchPartitionResponseV4 {
148    pub partition_index: i32,
149    pub error_code: i16,
150    pub high_watermark: i64,
151    pub last_stable_offset: i64,
152    pub aborted_transactions: Vec<AbortedTransactionV4>,
153    pub records: Vec<MessageSetRecord>,
154}
155
156impl FetchPartitionResponseV4 {
157    fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
158        let limits = decoder.limits();
159        Ok(Self {
160            partition_index: decoder.read_i32()?,
161            error_code: decoder.read_i16()?,
162            high_watermark: decoder.read_i64()?,
163            last_stable_offset: decoder.read_i64()?,
164            aborted_transactions: decoder
165                .read_array("aborted transactions", AbortedTransactionV4::decode)?
166                .unwrap_or_default(),
167            records: decode_message_set(&decoder.read_bytes()?, limits)?,
168        })
169    }
170}
171
172#[derive(Debug, Clone, PartialEq, Eq)]
173pub struct AbortedTransactionV4 {
174    pub producer_id: i64,
175    pub first_offset: i64,
176}
177
178impl AbortedTransactionV4 {
179    fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
180        Ok(Self {
181            producer_id: decoder.read_i64()?,
182            first_offset: decoder.read_i64()?,
183        })
184    }
185}
186
187impl FetchResponseV2 {
188    pub fn decode_body(decoder: &mut Decoder<'_>) -> Result<Self> {
189        Ok(Self {
190            throttle_time_ms: decoder.read_i32()?,
191            responses: decoder
192                .read_array("fetch responses", FetchTopicResponseV2::decode)?
193                .unwrap_or_default(),
194        })
195    }
196}
197
198#[derive(Debug, Clone, PartialEq, Eq)]
199pub struct FetchTopicResponseV2 {
200    pub name: String,
201    pub partitions: Vec<FetchPartitionResponseV2>,
202}
203
204impl FetchTopicResponseV2 {
205    fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
206        Ok(Self {
207            name: decoder.read_string()?,
208            partitions: decoder
209                .read_array(
210                    "fetch partition responses",
211                    FetchPartitionResponseV2::decode,
212                )?
213                .unwrap_or_default(),
214        })
215    }
216}
217
218#[derive(Debug, Clone, PartialEq, Eq)]
219pub struct FetchPartitionResponseV2 {
220    pub partition_index: i32,
221    pub error_code: i16,
222    pub high_watermark: i64,
223    pub records: Vec<MessageSetRecord>,
224}
225
226impl FetchPartitionResponseV2 {
227    fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
228        let limits = decoder.limits();
229        Ok(Self {
230            partition_index: decoder.read_i32()?,
231            error_code: decoder.read_i16()?,
232            high_watermark: decoder.read_i64()?,
233            records: decode_message_set(&decoder.read_bytes()?, limits)?,
234        })
235    }
236}
237
238#[derive(Debug, Clone, PartialEq, Eq)]
239pub struct MessageSetRecord {
240    pub offset: i64,
241    pub timestamp_ms: i64,
242    pub key: Option<Vec<u8>>,
243    pub value: Option<Vec<u8>>,
244    pub producer_id: Option<i64>,
245    pub transactional: bool,
246    pub control: bool,
247}
248
249fn decode_message_set(bytes: &[u8], limits: DecodeLimits) -> Result<Vec<MessageSetRecord>> {
250    let mut decoder = Decoder::with_limits(bytes, limits);
251    let mut records = Vec::new();
252
253    while decoder.remaining() >= 12 {
254        let offset = decoder.read_i64()?;
255        let message_size = decoder.read_i32()?;
256        if message_size < 0 {
257            return Err(Error::NegativeLength {
258                kind: "message",
259                length: message_size,
260            });
261        }
262        let message_size =
263            usize::try_from(message_size).map_err(|_| Error::LengthOverflow("message"))?;
264        // Fetch responses may end with a partial trailing message set entry.
265        if decoder.remaining() < message_size {
266            break;
267        }
268        let message = decoder.read_exact(message_size)?;
269        let decoded = decode_message_or_batch(offset, message, limits)?;
270        let total = records
271            .len()
272            .checked_add(decoded.len())
273            .ok_or(Error::LengthOverflow("fetch records"))?;
274        decoder.ensure_collection_length("fetch records", total)?;
275        records.extend(decoded);
276    }
277
278    Ok(records)
279}
280
281fn decode_message_or_batch(
282    offset: i64,
283    bytes: &[u8],
284    limits: DecodeLimits,
285) -> Result<Vec<MessageSetRecord>> {
286    match bytes.get(4).copied() {
287        Some(2) => decode_record_batch(offset, bytes, limits),
288        _ => Ok(vec![decode_message(offset, bytes, limits)?]),
289    }
290}
291
292fn decode_message(offset: i64, bytes: &[u8], limits: DecodeLimits) -> Result<MessageSetRecord> {
293    let mut decoder = Decoder::with_limits(bytes, limits);
294    let _crc = decoder.read_i32()?;
295    let magic = decoder.read_i8()?;
296    let _attributes = decoder.read_i8()?;
297    let timestamp_ms = match magic {
298        0 => -1,
299        1 => decoder.read_i64()?,
300        _ => {
301            return Err(Error::UnsupportedVersion {
302                kind: "message magic",
303                version: i16::from(magic),
304            })
305        }
306    };
307    let key = decoder.read_nullable_bytes()?;
308    let value = decoder.read_nullable_bytes()?;
309
310    Ok(MessageSetRecord {
311        offset,
312        timestamp_ms,
313        key,
314        value,
315        producer_id: None,
316        transactional: false,
317        control: false,
318    })
319}
320
321fn decode_record_batch(
322    base_offset: i64,
323    bytes: &[u8],
324    limits: DecodeLimits,
325) -> Result<Vec<MessageSetRecord>> {
326    let mut decoder = Decoder::with_limits(bytes, limits);
327    let _partition_leader_epoch = decoder.read_i32()?;
328    let magic = decoder.read_i8()?;
329    if magic != 2 {
330        return Err(Error::UnsupportedVersion {
331            kind: "record batch magic",
332            version: i16::from(magic),
333        });
334    }
335    let _crc = decoder.read_i32()?;
336    let attributes = decoder.read_i16()?;
337    let compression = RecordBatchCompression::from_attributes(attributes)?;
338    let _last_offset_delta = decoder.read_i32()?;
339    let base_timestamp = decoder.read_i64()?;
340    let _max_timestamp = decoder.read_i64()?;
341    let producer_id = decoder.read_i64()?;
342    let _producer_epoch = decoder.read_i16()?;
343    let _base_sequence = decoder.read_i32()?;
344    let record_count = decoder.read_i32()?;
345    if record_count < 0 {
346        return Err(Error::NegativeLength {
347            kind: "record batch records",
348            length: record_count,
349        });
350    }
351
352    let record_count =
353        usize::try_from(record_count).map_err(|_| Error::LengthOverflow("record batch records"))?;
354    decoder.ensure_collection_length("record batch records", record_count)?;
355    let record_bytes = if compression.is_compressed() {
356        let compressed = decoder.read_exact(decoder.remaining())?;
357        decompress_record_batch_records_with_limit(
358            compression,
359            compressed,
360            limits.max_decompressed_record_bytes(),
361        )?
362    } else {
363        if decoder.remaining() > limits.max_decompressed_record_bytes() {
364            return Err(Error::LimitExceeded {
365                kind: "decompressed record batch bytes",
366                actual: decoder.remaining(),
367                max: limits.max_decompressed_record_bytes(),
368            });
369        }
370        decoder.read_exact(decoder.remaining())?.to_vec()
371    };
372    let mut record_decoder = Decoder::with_limits(&record_bytes, limits);
373    let mut records = Vec::with_capacity(record_count);
374    for _ in 0..record_count {
375        let record_length = record_decoder.read_varint()?;
376        if record_length < 0 {
377            return Err(Error::NegativeLength {
378                kind: "record",
379                length: record_length,
380            });
381        }
382        let record_length =
383            usize::try_from(record_length).map_err(|_| Error::LengthOverflow("record"))?;
384        let record_bytes = record_decoder.read_exact(record_length)?;
385        records.push(decode_record(
386            base_offset,
387            base_timestamp,
388            producer_id,
389            attributes,
390            record_bytes,
391            limits,
392        )?);
393    }
394
395    Ok(records)
396}
397
398fn decode_record(
399    base_offset: i64,
400    base_timestamp: i64,
401    producer_id: i64,
402    batch_attributes: i16,
403    bytes: &[u8],
404    limits: DecodeLimits,
405) -> Result<MessageSetRecord> {
406    let mut decoder = Decoder::with_limits(bytes, limits);
407    let _attributes = decoder.read_i8()?;
408    let timestamp_delta = decoder.read_varlong()?;
409    let offset_delta = decoder.read_varint()?;
410    let key = decoder.read_varint_nullable_bytes()?;
411    let value = decoder.read_varint_nullable_bytes()?;
412    let header_count = decoder.read_varint()?;
413    if header_count < 0 {
414        return Err(Error::NegativeLength {
415            kind: "record headers",
416            length: header_count,
417        });
418    }
419    let header_count =
420        usize::try_from(header_count).map_err(|_| Error::LengthOverflow("record headers"))?;
421    decoder.ensure_collection_length("record headers", header_count)?;
422    for _ in 0..header_count {
423        let _header_key = decoder.read_varint_bytes()?;
424        let _header_value = decoder.read_varint_nullable_bytes()?;
425    }
426
427    Ok(MessageSetRecord {
428        offset: base_offset.saturating_add(i64::from(offset_delta)),
429        timestamp_ms: base_timestamp.saturating_add(timestamp_delta),
430        key,
431        value,
432        producer_id: (producer_id >= 0).then_some(producer_id),
433        transactional: batch_attributes & 0x10 != 0,
434        control: batch_attributes & 0x20 != 0,
435    })
436}
437
438#[cfg(test)]
439#[allow(clippy::unwrap_used)]
440mod tests {
441    use super::{
442        FetchPartitionV2, FetchRequestV2, FetchRequestV4, FetchResponseV2, FetchResponseV4,
443        FetchTopicV2, MessageSetRecord,
444    };
445    use crate::codec::{Decoder, Encoder};
446
447    #[test]
448    fn encodes_fetch_request_v2() {
449        let request = FetchRequestV2 {
450            correlation_id: 7,
451            client_id: Some("kafrust".to_owned()),
452            replica_id: -1,
453            max_wait_ms: 500,
454            min_bytes: 1,
455            topics: vec![FetchTopicV2 {
456                name: "orders".to_owned(),
457                partitions: vec![FetchPartitionV2 {
458                    partition_index: 0,
459                    fetch_offset: 42,
460                    max_bytes: 1_048_576,
461                }],
462            }],
463        };
464
465        let bytes = request.encode().unwrap();
466        assert_eq!(&bytes[0..4], &[0, 1, 0, 2]);
467        assert!(bytes.len() > 40);
468    }
469
470    #[test]
471    fn encodes_fetch_request_v4() {
472        let request = FetchRequestV4 {
473            correlation_id: 8,
474            client_id: Some("kafrust".to_owned()),
475            replica_id: -1,
476            max_wait_ms: 500,
477            min_bytes: 1,
478            max_bytes: 1_048_576,
479            isolation_level: 0,
480            topics: vec![FetchTopicV2 {
481                name: "orders".to_owned(),
482                partitions: vec![FetchPartitionV2 {
483                    partition_index: 0,
484                    fetch_offset: 42,
485                    max_bytes: 1_048_576,
486                }],
487            }],
488        };
489
490        let bytes = request.encode().unwrap();
491        assert_eq!(&bytes[0..4], &[0, 1, 0, 4]);
492        assert_eq!(&bytes[4..8], &[0, 0, 0, 8]);
493        assert!(bytes.len() > 45);
494    }
495
496    #[test]
497    fn decodes_fetch_response_v4_with_aborted_transaction() {
498        let mut bytes = Encoder::new();
499        bytes.write_i32(0);
500        bytes.write_i32(1);
501        bytes.write_string("orders").unwrap();
502        bytes.write_i32(1);
503        bytes.write_i32(0);
504        bytes.write_i16(0);
505        bytes.write_i64(43);
506        bytes.write_i64(42);
507        bytes.write_i32(1);
508        bytes.write_i64(7);
509        bytes.write_i64(40);
510        bytes.write_bytes(&[]).unwrap();
511        let bytes = bytes.into_bytes();
512
513        let mut decoder = Decoder::new(&bytes);
514        let response = FetchResponseV4::decode_body(&mut decoder).unwrap();
515        let partition = &response.responses[0].partitions[0];
516
517        assert_eq!(partition.high_watermark, 43);
518        assert_eq!(partition.last_stable_offset, 42);
519        assert_eq!(partition.aborted_transactions[0].producer_id, 7);
520        assert_eq!(partition.aborted_transactions[0].first_offset, 40);
521        assert!(partition.records.is_empty());
522        assert!(decoder.is_empty());
523    }
524
525    #[test]
526    fn decodes_fetch_response_v2_with_message_set() {
527        let mut message = Encoder::new();
528        message.write_i32(0);
529        message.write_i8(1);
530        message.write_i8(0);
531        message.write_i64(123);
532        message.write_nullable_bytes(Some(b"order-1")).unwrap();
533        message.write_nullable_bytes(Some(b"created")).unwrap();
534        let message = message.into_bytes();
535
536        let mut set = Encoder::new();
537        set.write_i64(42);
538        set.write_i32(i32::try_from(message.len()).unwrap());
539        set.write_raw(&message);
540        let set = set.into_bytes();
541
542        let mut bytes = Encoder::new();
543        bytes.write_i32(0);
544        bytes.write_i32(1);
545        bytes.write_string("orders").unwrap();
546        bytes.write_i32(1);
547        bytes.write_i32(0);
548        bytes.write_i16(0);
549        bytes.write_i64(43);
550        bytes.write_bytes(&set).unwrap();
551        let bytes = bytes.into_bytes();
552
553        let mut decoder = Decoder::new(&bytes);
554        let response = FetchResponseV2::decode_body(&mut decoder).unwrap();
555        let record = MessageSetRecord {
556            offset: 42,
557            timestamp_ms: 123,
558            key: Some(b"order-1".to_vec()),
559            value: Some(b"created".to_vec()),
560            producer_id: None,
561            transactional: false,
562            control: false,
563        };
564
565        assert_eq!(response.throttle_time_ms, 0);
566        assert_eq!(response.responses[0].partitions[0].high_watermark, 43);
567        assert_eq!(response.responses[0].partitions[0].records, vec![record]);
568        assert!(decoder.is_empty());
569    }
570
571    #[test]
572    fn decodes_fetch_response_v2_ignores_partial_trailing_message_set_entry() {
573        let mut message = Encoder::new();
574        message.write_i32(0);
575        message.write_i8(1);
576        message.write_i8(0);
577        message.write_i64(123);
578        message.write_nullable_bytes(Some(b"order-1")).unwrap();
579        message.write_nullable_bytes(Some(b"created")).unwrap();
580        let message = message.into_bytes();
581
582        let mut set = Encoder::new();
583        set.write_i64(42);
584        set.write_i32(i32::try_from(message.len()).unwrap());
585        set.write_raw(&message);
586        set.write_i64(-1);
587        set.write_i32(61);
588        set.write_raw(&[0; 22]);
589        let set = set.into_bytes();
590
591        let mut bytes = Encoder::new();
592        bytes.write_i32(0);
593        bytes.write_i32(1);
594        bytes.write_string("orders").unwrap();
595        bytes.write_i32(1);
596        bytes.write_i32(0);
597        bytes.write_i16(0);
598        bytes.write_i64(43);
599        bytes.write_bytes(&set).unwrap();
600        let bytes = bytes.into_bytes();
601
602        let mut decoder = Decoder::new(&bytes);
603        let response = FetchResponseV2::decode_body(&mut decoder).unwrap();
604
605        assert_eq!(response.responses[0].partitions[0].records.len(), 1);
606        assert_eq!(response.responses[0].partitions[0].records[0].offset, 42);
607        assert!(decoder.is_empty());
608    }
609
610    #[test]
611    fn decodes_fetch_response_v2_with_record_batch() {
612        let mut record = Vec::new();
613        record.push(0);
614        write_varlong(&mut record, 5);
615        write_varint(&mut record, 0);
616        write_varint(&mut record, 7);
617        record.extend_from_slice(b"order-1");
618        write_varint(&mut record, 7);
619        record.extend_from_slice(b"created");
620        write_varint(&mut record, 0);
621
622        let mut batch = Encoder::new();
623        batch.write_i32(0);
624        batch.write_i8(2);
625        batch.write_i32(0);
626        batch.write_i16(0x10);
627        batch.write_i32(0);
628        batch.write_i64(1_000);
629        batch.write_i64(1_005);
630        batch.write_i64(7);
631        batch.write_i16(-1);
632        batch.write_i32(-1);
633        batch.write_i32(1);
634        let mut encoded_record = Vec::new();
635        write_varint(&mut encoded_record, i32::try_from(record.len()).unwrap());
636        encoded_record.extend_from_slice(&record);
637        batch.write_raw(&encoded_record);
638        let batch = batch.into_bytes();
639
640        let mut set = Encoder::new();
641        set.write_i64(42);
642        set.write_i32(i32::try_from(batch.len()).unwrap());
643        set.write_raw(&batch);
644        let set = set.into_bytes();
645
646        let mut bytes = Encoder::new();
647        bytes.write_i32(0);
648        bytes.write_i32(1);
649        bytes.write_string("orders").unwrap();
650        bytes.write_i32(1);
651        bytes.write_i32(0);
652        bytes.write_i16(0);
653        bytes.write_i64(43);
654        bytes.write_bytes(&set).unwrap();
655        let bytes = bytes.into_bytes();
656
657        let mut decoder = Decoder::new(&bytes);
658        let response = FetchResponseV2::decode_body(&mut decoder).unwrap();
659        let record = MessageSetRecord {
660            offset: 42,
661            timestamp_ms: 1_005,
662            key: Some(b"order-1".to_vec()),
663            value: Some(b"created".to_vec()),
664            producer_id: Some(7),
665            transactional: true,
666            control: false,
667        };
668
669        assert_eq!(response.responses[0].partitions[0].records, vec![record]);
670        assert!(decoder.is_empty());
671    }
672
673    fn write_varint(output: &mut Vec<u8>, value: i32) {
674        write_unsigned_varint(output, u64::from(((value << 1) ^ (value >> 31)) as u32));
675    }
676
677    fn write_varlong(output: &mut Vec<u8>, value: i64) {
678        write_unsigned_varint(output, ((value << 1) ^ (value >> 63)) as u64);
679    }
680
681    fn write_unsigned_varint(output: &mut Vec<u8>, mut value: u64) {
682        loop {
683            let mut byte = (value & 0x7f) as u8;
684            value >>= 7;
685            if value != 0 {
686                byte |= 0x80;
687            }
688            output.push(byte);
689            if value == 0 {
690                break;
691            }
692        }
693    }
694}