Skip to main content

kafrust_protocol/api/
fetch.rs

1use crate::codec::{Decoder, Encoder};
2use crate::error::{Error, Result};
3use crate::header::RequestHeader;
4
5pub const API_KEY: i16 = 1;
6
7#[derive(Debug, Clone, PartialEq, Eq)]
8pub struct FetchRequestV2 {
9    pub correlation_id: i32,
10    pub client_id: Option<String>,
11    pub replica_id: i32,
12    pub max_wait_ms: i32,
13    pub min_bytes: i32,
14    pub topics: Vec<FetchTopicV2>,
15}
16
17impl FetchRequestV2 {
18    pub fn encode(&self) -> Result<Vec<u8>> {
19        let mut encoder = Encoder::new();
20        RequestHeader {
21            api_key: API_KEY,
22            api_version: 2,
23            correlation_id: self.correlation_id,
24            client_id: self.client_id.clone(),
25        }
26        .encode_v1(&mut encoder)?;
27        encoder.write_i32(self.replica_id);
28        encoder.write_i32(self.max_wait_ms);
29        encoder.write_i32(self.min_bytes);
30        encoder.write_array(Some(self.topics.as_slice()), |encoder, topic| {
31            topic.encode(encoder)
32        })?;
33        Ok(encoder.into_bytes())
34    }
35}
36
37#[derive(Debug, Clone, PartialEq, Eq)]
38pub struct FetchTopicV2 {
39    pub name: String,
40    pub partitions: Vec<FetchPartitionV2>,
41}
42
43impl FetchTopicV2 {
44    fn encode(&self, encoder: &mut Encoder) -> Result<()> {
45        encoder.write_string(&self.name)?;
46        encoder.write_array(Some(self.partitions.as_slice()), |encoder, partition| {
47            partition.encode(encoder)
48        })
49    }
50}
51
52#[derive(Debug, Clone, PartialEq, Eq)]
53pub struct FetchPartitionV2 {
54    pub partition_index: i32,
55    pub fetch_offset: i64,
56    pub max_bytes: i32,
57}
58
59impl FetchPartitionV2 {
60    fn encode(&self, encoder: &mut Encoder) -> Result<()> {
61        encoder.write_i32(self.partition_index);
62        encoder.write_i64(self.fetch_offset);
63        encoder.write_i32(self.max_bytes);
64        Ok(())
65    }
66}
67
68#[derive(Debug, Clone, PartialEq, Eq)]
69pub struct FetchResponseV2 {
70    pub throttle_time_ms: i32,
71    pub responses: Vec<FetchTopicResponseV2>,
72}
73
74impl FetchResponseV2 {
75    pub fn decode_body(decoder: &mut Decoder<'_>) -> Result<Self> {
76        Ok(Self {
77            throttle_time_ms: decoder.read_i32()?,
78            responses: decoder
79                .read_array("fetch responses", FetchTopicResponseV2::decode)?
80                .unwrap_or_default(),
81        })
82    }
83}
84
85#[derive(Debug, Clone, PartialEq, Eq)]
86pub struct FetchTopicResponseV2 {
87    pub name: String,
88    pub partitions: Vec<FetchPartitionResponseV2>,
89}
90
91impl FetchTopicResponseV2 {
92    fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
93        Ok(Self {
94            name: decoder.read_string()?,
95            partitions: decoder
96                .read_array(
97                    "fetch partition responses",
98                    FetchPartitionResponseV2::decode,
99                )?
100                .unwrap_or_default(),
101        })
102    }
103}
104
105#[derive(Debug, Clone, PartialEq, Eq)]
106pub struct FetchPartitionResponseV2 {
107    pub partition_index: i32,
108    pub error_code: i16,
109    pub high_watermark: i64,
110    pub records: Vec<MessageSetRecord>,
111}
112
113impl FetchPartitionResponseV2 {
114    fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
115        Ok(Self {
116            partition_index: decoder.read_i32()?,
117            error_code: decoder.read_i16()?,
118            high_watermark: decoder.read_i64()?,
119            records: decode_message_set(&decoder.read_bytes()?)?,
120        })
121    }
122}
123
124#[derive(Debug, Clone, PartialEq, Eq)]
125pub struct MessageSetRecord {
126    pub offset: i64,
127    pub timestamp_ms: i64,
128    pub key: Option<Vec<u8>>,
129    pub value: Option<Vec<u8>>,
130}
131
132fn decode_message_set(bytes: &[u8]) -> Result<Vec<MessageSetRecord>> {
133    let mut decoder = Decoder::new(bytes);
134    let mut records = Vec::new();
135
136    while decoder.remaining() >= 12 {
137        let offset = decoder.read_i64()?;
138        let message_size = decoder.read_i32()?;
139        if message_size < 0 {
140            return Err(Error::NegativeLength {
141                kind: "message",
142                length: message_size,
143            });
144        }
145        let message_size =
146            usize::try_from(message_size).map_err(|_| Error::LengthOverflow("message"))?;
147        // Fetch responses may end with a partial trailing message set entry.
148        if decoder.remaining() < message_size {
149            break;
150        }
151        let message = decoder.read_exact(message_size)?;
152        records.extend(decode_message_or_batch(offset, message)?);
153    }
154
155    Ok(records)
156}
157
158fn decode_message_or_batch(offset: i64, bytes: &[u8]) -> Result<Vec<MessageSetRecord>> {
159    match bytes.get(4).copied() {
160        Some(2) => decode_record_batch(offset, bytes),
161        _ => Ok(vec![decode_message(offset, bytes)?]),
162    }
163}
164
165fn decode_message(offset: i64, bytes: &[u8]) -> Result<MessageSetRecord> {
166    let mut decoder = Decoder::new(bytes);
167    let _crc = decoder.read_i32()?;
168    let magic = decoder.read_i8()?;
169    let _attributes = decoder.read_i8()?;
170    let timestamp_ms = match magic {
171        0 => -1,
172        1 => decoder.read_i64()?,
173        _ => {
174            return Err(Error::UnsupportedVersion {
175                kind: "message magic",
176                version: i16::from(magic),
177            })
178        }
179    };
180    let key = decoder.read_nullable_bytes()?;
181    let value = decoder.read_nullable_bytes()?;
182
183    Ok(MessageSetRecord {
184        offset,
185        timestamp_ms,
186        key,
187        value,
188    })
189}
190
191fn decode_record_batch(base_offset: i64, bytes: &[u8]) -> Result<Vec<MessageSetRecord>> {
192    let mut decoder = Decoder::new(bytes);
193    let _partition_leader_epoch = decoder.read_i32()?;
194    let magic = decoder.read_i8()?;
195    if magic != 2 {
196        return Err(Error::UnsupportedVersion {
197            kind: "record batch magic",
198            version: i16::from(magic),
199        });
200    }
201    let _crc = decoder.read_i32()?;
202    let _attributes = decoder.read_i16()?;
203    let _last_offset_delta = decoder.read_i32()?;
204    let base_timestamp = decoder.read_i64()?;
205    let _max_timestamp = decoder.read_i64()?;
206    let _producer_id = decoder.read_i64()?;
207    let _producer_epoch = decoder.read_i16()?;
208    let _base_sequence = decoder.read_i32()?;
209    let record_count = decoder.read_i32()?;
210    if record_count < 0 {
211        return Err(Error::NegativeLength {
212            kind: "record batch records",
213            length: record_count,
214        });
215    }
216
217    let record_count =
218        usize::try_from(record_count).map_err(|_| Error::LengthOverflow("record batch records"))?;
219    let mut records = Vec::with_capacity(record_count);
220    for _ in 0..record_count {
221        let record_length = decoder.read_varint()?;
222        if record_length < 0 {
223            return Err(Error::NegativeLength {
224                kind: "record",
225                length: record_length,
226            });
227        }
228        let record_length =
229            usize::try_from(record_length).map_err(|_| Error::LengthOverflow("record"))?;
230        let record_bytes = decoder.read_exact(record_length)?;
231        records.push(decode_record(base_offset, base_timestamp, record_bytes)?);
232    }
233
234    Ok(records)
235}
236
237fn decode_record(base_offset: i64, base_timestamp: i64, bytes: &[u8]) -> Result<MessageSetRecord> {
238    let mut decoder = Decoder::new(bytes);
239    let _attributes = decoder.read_i8()?;
240    let timestamp_delta = decoder.read_varlong()?;
241    let offset_delta = decoder.read_varint()?;
242    let key = decoder.read_varint_nullable_bytes()?;
243    let value = decoder.read_varint_nullable_bytes()?;
244    let header_count = decoder.read_varint()?;
245    if header_count < 0 {
246        return Err(Error::NegativeLength {
247            kind: "record headers",
248            length: header_count,
249        });
250    }
251    for _ in 0..header_count {
252        let _header_key = decoder.read_varint_bytes()?;
253        let _header_value = decoder.read_varint_nullable_bytes()?;
254    }
255
256    Ok(MessageSetRecord {
257        offset: base_offset.saturating_add(i64::from(offset_delta)),
258        timestamp_ms: base_timestamp.saturating_add(timestamp_delta),
259        key,
260        value,
261    })
262}
263
264#[cfg(test)]
265#[allow(clippy::unwrap_used)]
266mod tests {
267    use super::{
268        FetchPartitionV2, FetchRequestV2, FetchResponseV2, FetchTopicV2, MessageSetRecord,
269    };
270    use crate::codec::{Decoder, Encoder};
271
272    #[test]
273    fn encodes_fetch_request_v2() {
274        let request = FetchRequestV2 {
275            correlation_id: 7,
276            client_id: Some("kafrust".to_owned()),
277            replica_id: -1,
278            max_wait_ms: 500,
279            min_bytes: 1,
280            topics: vec![FetchTopicV2 {
281                name: "orders".to_owned(),
282                partitions: vec![FetchPartitionV2 {
283                    partition_index: 0,
284                    fetch_offset: 42,
285                    max_bytes: 1_048_576,
286                }],
287            }],
288        };
289
290        let bytes = request.encode().unwrap();
291        assert_eq!(&bytes[0..4], &[0, 1, 0, 2]);
292        assert!(bytes.len() > 40);
293    }
294
295    #[test]
296    fn decodes_fetch_response_v2_with_message_set() {
297        let mut message = Encoder::new();
298        message.write_i32(0);
299        message.write_i8(1);
300        message.write_i8(0);
301        message.write_i64(123);
302        message.write_nullable_bytes(Some(b"order-1")).unwrap();
303        message.write_nullable_bytes(Some(b"created")).unwrap();
304        let message = message.into_bytes();
305
306        let mut set = Encoder::new();
307        set.write_i64(42);
308        set.write_i32(i32::try_from(message.len()).unwrap());
309        set.write_raw(&message);
310        let set = set.into_bytes();
311
312        let mut bytes = Encoder::new();
313        bytes.write_i32(0);
314        bytes.write_i32(1);
315        bytes.write_string("orders").unwrap();
316        bytes.write_i32(1);
317        bytes.write_i32(0);
318        bytes.write_i16(0);
319        bytes.write_i64(43);
320        bytes.write_bytes(&set).unwrap();
321        let bytes = bytes.into_bytes();
322
323        let mut decoder = Decoder::new(&bytes);
324        let response = FetchResponseV2::decode_body(&mut decoder).unwrap();
325        let record = MessageSetRecord {
326            offset: 42,
327            timestamp_ms: 123,
328            key: Some(b"order-1".to_vec()),
329            value: Some(b"created".to_vec()),
330        };
331
332        assert_eq!(response.throttle_time_ms, 0);
333        assert_eq!(response.responses[0].partitions[0].high_watermark, 43);
334        assert_eq!(response.responses[0].partitions[0].records, vec![record]);
335        assert!(decoder.is_empty());
336    }
337
338    #[test]
339    fn decodes_fetch_response_v2_ignores_partial_trailing_message_set_entry() {
340        let mut message = Encoder::new();
341        message.write_i32(0);
342        message.write_i8(1);
343        message.write_i8(0);
344        message.write_i64(123);
345        message.write_nullable_bytes(Some(b"order-1")).unwrap();
346        message.write_nullable_bytes(Some(b"created")).unwrap();
347        let message = message.into_bytes();
348
349        let mut set = Encoder::new();
350        set.write_i64(42);
351        set.write_i32(i32::try_from(message.len()).unwrap());
352        set.write_raw(&message);
353        set.write_i64(-1);
354        set.write_i32(61);
355        set.write_raw(&[0; 22]);
356        let set = set.into_bytes();
357
358        let mut bytes = Encoder::new();
359        bytes.write_i32(0);
360        bytes.write_i32(1);
361        bytes.write_string("orders").unwrap();
362        bytes.write_i32(1);
363        bytes.write_i32(0);
364        bytes.write_i16(0);
365        bytes.write_i64(43);
366        bytes.write_bytes(&set).unwrap();
367        let bytes = bytes.into_bytes();
368
369        let mut decoder = Decoder::new(&bytes);
370        let response = FetchResponseV2::decode_body(&mut decoder).unwrap();
371
372        assert_eq!(response.responses[0].partitions[0].records.len(), 1);
373        assert_eq!(response.responses[0].partitions[0].records[0].offset, 42);
374        assert!(decoder.is_empty());
375    }
376
377    #[test]
378    fn decodes_fetch_response_v2_with_record_batch() {
379        let mut record = Vec::new();
380        record.push(0);
381        write_varlong(&mut record, 5);
382        write_varint(&mut record, 0);
383        write_varint(&mut record, 7);
384        record.extend_from_slice(b"order-1");
385        write_varint(&mut record, 7);
386        record.extend_from_slice(b"created");
387        write_varint(&mut record, 0);
388
389        let mut batch = Encoder::new();
390        batch.write_i32(0);
391        batch.write_i8(2);
392        batch.write_i32(0);
393        batch.write_i16(0);
394        batch.write_i32(0);
395        batch.write_i64(1_000);
396        batch.write_i64(1_005);
397        batch.write_i64(-1);
398        batch.write_i16(-1);
399        batch.write_i32(-1);
400        batch.write_i32(1);
401        let mut encoded_record = Vec::new();
402        write_varint(&mut encoded_record, i32::try_from(record.len()).unwrap());
403        encoded_record.extend_from_slice(&record);
404        batch.write_raw(&encoded_record);
405        let batch = batch.into_bytes();
406
407        let mut set = Encoder::new();
408        set.write_i64(42);
409        set.write_i32(i32::try_from(batch.len()).unwrap());
410        set.write_raw(&batch);
411        let set = set.into_bytes();
412
413        let mut bytes = Encoder::new();
414        bytes.write_i32(0);
415        bytes.write_i32(1);
416        bytes.write_string("orders").unwrap();
417        bytes.write_i32(1);
418        bytes.write_i32(0);
419        bytes.write_i16(0);
420        bytes.write_i64(43);
421        bytes.write_bytes(&set).unwrap();
422        let bytes = bytes.into_bytes();
423
424        let mut decoder = Decoder::new(&bytes);
425        let response = FetchResponseV2::decode_body(&mut decoder).unwrap();
426        let record = MessageSetRecord {
427            offset: 42,
428            timestamp_ms: 1_005,
429            key: Some(b"order-1".to_vec()),
430            value: Some(b"created".to_vec()),
431        };
432
433        assert_eq!(response.responses[0].partitions[0].records, vec![record]);
434        assert!(decoder.is_empty());
435    }
436
437    fn write_varint(output: &mut Vec<u8>, value: i32) {
438        write_unsigned_varint(output, u64::from(((value << 1) ^ (value >> 31)) as u32));
439    }
440
441    fn write_varlong(output: &mut Vec<u8>, value: i64) {
442        write_unsigned_varint(output, ((value << 1) ^ (value >> 63)) as u64);
443    }
444
445    fn write_unsigned_varint(output: &mut Vec<u8>, mut value: u64) {
446        loop {
447            let mut byte = (value & 0x7f) as u8;
448            value >>= 7;
449            if value != 0 {
450                byte |= 0x80;
451            }
452            output.push(byte);
453            if value == 0 {
454                break;
455            }
456        }
457    }
458}