Skip to main content

kafrust_protocol/api/
produce.rs

1use crate::codec::{Decoder, Encoder};
2use crate::error::{Error, Result};
3use crate::header::RequestHeader;
4
5pub const API_KEY: i16 = 0;
6
7#[derive(Debug, Clone, PartialEq, Eq)]
8pub struct ProduceRequestV2 {
9    pub correlation_id: i32,
10    pub client_id: Option<String>,
11    pub acks: i16,
12    pub timeout_ms: i32,
13    pub topics: Vec<ProduceTopicV2>,
14}
15
16#[derive(Debug, Clone, PartialEq, Eq)]
17pub struct ProduceRequestV3 {
18    pub correlation_id: i32,
19    pub client_id: Option<String>,
20    pub transactional_id: Option<String>,
21    pub acks: i16,
22    pub timeout_ms: i32,
23    pub topics: Vec<ProduceTopicV3>,
24}
25
26impl ProduceRequestV3 {
27    pub fn encode(&self) -> Result<Vec<u8>> {
28        let mut encoder = Encoder::new();
29        RequestHeader {
30            api_key: API_KEY,
31            api_version: 3,
32            correlation_id: self.correlation_id,
33            client_id: self.client_id.clone(),
34        }
35        .encode_v1(&mut encoder)?;
36        encoder.write_nullable_string(self.transactional_id.as_deref())?;
37        encoder.write_i16(self.acks);
38        encoder.write_i32(self.timeout_ms);
39        encoder.write_array(Some(self.topics.as_slice()), |encoder, topic| {
40            topic.encode(encoder)
41        })?;
42        Ok(encoder.into_bytes())
43    }
44}
45
46impl ProduceRequestV2 {
47    pub fn encode(&self) -> Result<Vec<u8>> {
48        let mut encoder = Encoder::new();
49        RequestHeader {
50            api_key: API_KEY,
51            api_version: 2,
52            correlation_id: self.correlation_id,
53            client_id: self.client_id.clone(),
54        }
55        .encode_v1(&mut encoder)?;
56        encoder.write_i16(self.acks);
57        encoder.write_i32(self.timeout_ms);
58        encoder.write_array(Some(self.topics.as_slice()), |encoder, topic| {
59            topic.encode(encoder)
60        })?;
61        Ok(encoder.into_bytes())
62    }
63}
64
65#[derive(Debug, Clone, PartialEq, Eq)]
66pub struct ProduceTopicV2 {
67    pub name: String,
68    pub partitions: Vec<ProducePartitionV2>,
69}
70
71#[derive(Debug, Clone, PartialEq, Eq)]
72pub struct ProduceTopicV3 {
73    pub name: String,
74    pub partitions: Vec<ProducePartitionV3>,
75}
76
77impl ProduceTopicV3 {
78    fn encode(&self, encoder: &mut Encoder) -> Result<()> {
79        encoder.write_string(&self.name)?;
80        encoder.write_array(Some(self.partitions.as_slice()), |encoder, partition| {
81            partition.encode(encoder)
82        })
83    }
84}
85
86impl ProduceTopicV2 {
87    fn encode(&self, encoder: &mut Encoder) -> Result<()> {
88        encoder.write_string(&self.name)?;
89        encoder.write_array(Some(self.partitions.as_slice()), |encoder, partition| {
90            partition.encode(encoder)
91        })
92    }
93}
94
95#[derive(Debug, Clone, PartialEq, Eq)]
96pub struct ProducePartitionV2 {
97    pub partition_index: i32,
98    pub records: Vec<MessageSetMessage>,
99}
100
101#[derive(Debug, Clone, PartialEq, Eq)]
102pub struct ProducePartitionV3 {
103    pub partition_index: i32,
104    pub records: Vec<RecordBatchMessage>,
105}
106
107impl ProducePartitionV3 {
108    fn encode(&self, encoder: &mut Encoder) -> Result<()> {
109        encoder.write_i32(self.partition_index);
110        let record_set = encode_record_batch_set(&self.records)?;
111        encoder.write_bytes(&record_set)
112    }
113}
114
115impl ProducePartitionV2 {
116    fn encode(&self, encoder: &mut Encoder) -> Result<()> {
117        encoder.write_i32(self.partition_index);
118        let record_set = encode_message_set(&self.records)?;
119        encoder.write_bytes(&record_set)
120    }
121}
122
123/// Returns the encoded byte length of a Produce v2 message set.
124pub fn encoded_message_set_len(records: &[MessageSetMessage]) -> Result<usize> {
125    Ok(encode_message_set(records)?.len())
126}
127
128/// Returns the encoded byte length of a Produce v3 record batch set.
129pub fn encoded_record_batch_set_len(records: &[RecordBatchMessage]) -> Result<usize> {
130    Ok(encode_record_batch_set(records)?.len())
131}
132
133#[derive(Debug, Clone, PartialEq, Eq)]
134pub struct MessageSetMessage {
135    pub key: Option<Vec<u8>>,
136    pub value: Option<Vec<u8>>,
137    pub timestamp_ms: i64,
138}
139
140impl MessageSetMessage {
141    pub fn new(key: Option<Vec<u8>>, value: Option<Vec<u8>>, timestamp_ms: i64) -> Self {
142        Self {
143            key,
144            value,
145            timestamp_ms,
146        }
147    }
148}
149
150#[derive(Debug, Clone, PartialEq, Eq)]
151pub struct RecordBatchHeader {
152    pub key: String,
153    pub value: Option<Vec<u8>>,
154}
155
156impl RecordBatchHeader {
157    pub fn new(key: impl Into<String>, value: Option<Vec<u8>>) -> Self {
158        Self {
159            key: key.into(),
160            value,
161        }
162    }
163}
164
165#[derive(Debug, Clone, PartialEq, Eq)]
166pub struct RecordBatchMessage {
167    pub key: Option<Vec<u8>>,
168    pub value: Option<Vec<u8>>,
169    pub timestamp_ms: i64,
170    pub headers: Vec<RecordBatchHeader>,
171}
172
173impl RecordBatchMessage {
174    pub fn new(key: Option<Vec<u8>>, value: Option<Vec<u8>>, timestamp_ms: i64) -> Self {
175        Self {
176            key,
177            value,
178            timestamp_ms,
179            headers: Vec::new(),
180        }
181    }
182
183    pub fn header(mut self, key: impl Into<String>, value: Option<Vec<u8>>) -> Self {
184        self.headers.push(RecordBatchHeader::new(key, value));
185        self
186    }
187}
188
189#[derive(Debug, Clone, PartialEq, Eq)]
190pub struct ProduceResponseV2 {
191    pub responses: Vec<ProduceTopicResponseV2>,
192    pub throttle_time_ms: i32,
193}
194
195impl ProduceResponseV2 {
196    pub fn decode_body(decoder: &mut Decoder<'_>) -> Result<Self> {
197        Ok(Self {
198            responses: decoder
199                .read_array("produce responses", ProduceTopicResponseV2::decode)?
200                .unwrap_or_default(),
201            throttle_time_ms: decoder.read_i32()?,
202        })
203    }
204}
205
206#[derive(Debug, Clone, PartialEq, Eq)]
207pub struct ProduceTopicResponseV2 {
208    pub name: String,
209    pub partitions: Vec<ProducePartitionResponseV2>,
210}
211
212impl ProduceTopicResponseV2 {
213    fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
214        Ok(Self {
215            name: decoder.read_string()?,
216            partitions: decoder
217                .read_array(
218                    "produce partition responses",
219                    ProducePartitionResponseV2::decode,
220                )?
221                .unwrap_or_default(),
222        })
223    }
224}
225
226#[derive(Debug, Clone, PartialEq, Eq)]
227pub struct ProducePartitionResponseV2 {
228    pub partition_index: i32,
229    pub error_code: i16,
230    pub base_offset: i64,
231    pub log_append_time_ms: i64,
232}
233
234impl ProducePartitionResponseV2 {
235    fn decode(decoder: &mut Decoder<'_>) -> Result<Self> {
236        Ok(Self {
237            partition_index: decoder.read_i32()?,
238            error_code: decoder.read_i16()?,
239            base_offset: decoder.read_i64()?,
240            log_append_time_ms: decoder.read_i64()?,
241        })
242    }
243}
244
245fn encode_message_set(records: &[MessageSetMessage]) -> Result<Vec<u8>> {
246    let mut set = Encoder::new();
247    for record in records {
248        let message = encode_message(record)?;
249        set.write_i64(0);
250        set.write_i32(i32::try_from(message.len()).map_err(|_| Error::LengthOverflow("message"))?);
251        set.write_raw(&message);
252    }
253    Ok(set.into_bytes())
254}
255
256fn encode_message(record: &MessageSetMessage) -> Result<Vec<u8>> {
257    let mut body = Encoder::new();
258    body.write_i8(1);
259    body.write_i8(0);
260    body.write_i64(record.timestamp_ms);
261    body.write_nullable_bytes(record.key.as_deref())?;
262    body.write_nullable_bytes(record.value.as_deref())?;
263    let body = body.into_bytes();
264
265    let mut message = Encoder::new();
266    message.write_i32(crc32_ieee(&body) as i32);
267    message.write_raw(&body);
268    Ok(message.into_bytes())
269}
270
271fn encode_record_batch_set(records: &[RecordBatchMessage]) -> Result<Vec<u8>> {
272    let base_timestamp = records
273        .first()
274        .map(|record| record.timestamp_ms)
275        .unwrap_or_default();
276    let max_timestamp = records
277        .iter()
278        .map(|record| record.timestamp_ms)
279        .max()
280        .unwrap_or(base_timestamp);
281    let last_offset_delta = records
282        .len()
283        .checked_sub(1)
284        .map(|delta| i32::try_from(delta).map_err(|_| Error::LengthOverflow("record batch")))
285        .transpose()?
286        .unwrap_or_default();
287
288    let mut record_bytes = Encoder::new();
289    record_bytes.write_i32(
290        i32::try_from(records.len()).map_err(|_| Error::LengthOverflow("record batch records"))?,
291    );
292    for (offset_delta, record) in records.iter().enumerate() {
293        let encoded = encode_record(record, base_timestamp, offset_delta)?;
294        record_bytes.write_varint(
295            i32::try_from(encoded.len()).map_err(|_| Error::LengthOverflow("record"))?,
296        );
297        record_bytes.write_raw(&encoded);
298    }
299
300    let mut crc_payload = Encoder::new();
301    crc_payload.write_i16(0);
302    crc_payload.write_i32(last_offset_delta);
303    crc_payload.write_i64(base_timestamp);
304    crc_payload.write_i64(max_timestamp);
305    crc_payload.write_i64(-1);
306    crc_payload.write_i16(-1);
307    crc_payload.write_i32(-1);
308    crc_payload.write_raw(&record_bytes.into_bytes());
309    let crc_payload = crc_payload.into_bytes();
310
311    let mut batch = Encoder::new();
312    batch.write_i32(0);
313    batch.write_i8(2);
314    batch.write_i32(crc32c(&crc_payload) as i32);
315    batch.write_raw(&crc_payload);
316    let batch = batch.into_bytes();
317
318    let mut set = Encoder::new();
319    set.write_i64(0);
320    set.write_i32(i32::try_from(batch.len()).map_err(|_| Error::LengthOverflow("record batch"))?);
321    set.write_raw(&batch);
322    Ok(set.into_bytes())
323}
324
325fn encode_record(
326    record: &RecordBatchMessage,
327    base_timestamp: i64,
328    offset_delta: usize,
329) -> Result<Vec<u8>> {
330    let mut encoder = Encoder::new();
331    encoder.write_i8(0);
332    encoder.write_varlong(record.timestamp_ms.saturating_sub(base_timestamp));
333    encoder.write_varint(
334        i32::try_from(offset_delta).map_err(|_| Error::LengthOverflow("record offset delta"))?,
335    );
336    encoder.write_varint_nullable_bytes(record.key.as_deref())?;
337    encoder.write_varint_nullable_bytes(record.value.as_deref())?;
338    encoder.write_varint(
339        i32::try_from(record.headers.len()).map_err(|_| Error::LengthOverflow("record headers"))?,
340    );
341    for header in &record.headers {
342        encoder.write_varint_bytes(header.key.as_bytes())?;
343        encoder.write_varint_nullable_bytes(header.value.as_deref())?;
344    }
345    Ok(encoder.into_bytes())
346}
347
348fn crc32_ieee(bytes: &[u8]) -> u32 {
349    let mut crc = 0xffff_ffffu32;
350    for byte in bytes {
351        crc ^= u32::from(*byte);
352        for _ in 0..8 {
353            let mask = 0u32.wrapping_sub(crc & 1);
354            crc = (crc >> 1) ^ (0xedb8_8320 & mask);
355        }
356    }
357    !crc
358}
359
360fn crc32c(bytes: &[u8]) -> u32 {
361    let mut crc = 0xffff_ffffu32;
362    for byte in bytes {
363        crc ^= u32::from(*byte);
364        for _ in 0..8 {
365            let mask = 0u32.wrapping_sub(crc & 1);
366            crc = (crc >> 1) ^ (0x82f6_3b78 & mask);
367        }
368    }
369    !crc
370}
371
372#[cfg(test)]
373#[allow(clippy::unwrap_used)]
374mod tests {
375    use super::{
376        encode_message_set, encode_record_batch_set, encoded_message_set_len,
377        encoded_record_batch_set_len, MessageSetMessage, ProducePartitionV2, ProducePartitionV3,
378        ProduceRequestV2, ProduceRequestV3, ProduceResponseV2, ProduceTopicV2, ProduceTopicV3,
379        RecordBatchMessage,
380    };
381    use crate::codec::Decoder;
382    use crate::{api::fetch::FetchResponseV2, codec::Encoder};
383
384    #[test]
385    fn encodes_produce_request_v2() {
386        let request = ProduceRequestV2 {
387            correlation_id: 5,
388            client_id: Some("kafrust".to_owned()),
389            acks: 1,
390            timeout_ms: 30_000,
391            topics: vec![ProduceTopicV2 {
392                name: "orders".to_owned(),
393                partitions: vec![ProducePartitionV2 {
394                    partition_index: 0,
395                    records: vec![MessageSetMessage::new(
396                        Some(b"order-1".to_vec()),
397                        Some(b"created".to_vec()),
398                        0,
399                    )],
400                }],
401            }],
402        };
403
404        let bytes = request.encode().unwrap();
405        assert_eq!(
406            &bytes[0..17],
407            &[0, 0, 0, 2, 0, 0, 0, 5, 0, 7, b'k', b'a', b'f', b'r', b'u', b's', b't',]
408        );
409        assert!(bytes.len() > 60);
410    }
411
412    #[test]
413    fn encodes_produce_request_v3_with_record_batch() {
414        let request = ProduceRequestV3 {
415            correlation_id: 5,
416            client_id: Some("kafrust".to_owned()),
417            transactional_id: None,
418            acks: 1,
419            timeout_ms: 30_000,
420            topics: vec![ProduceTopicV3 {
421                name: "orders".to_owned(),
422                partitions: vec![ProducePartitionV3 {
423                    partition_index: 0,
424                    records: vec![RecordBatchMessage::new(
425                        Some(b"order-1".to_vec()),
426                        Some(b"created".to_vec()),
427                        1_000,
428                    )
429                    .header("source", Some(b"checkout".to_vec()))],
430                }],
431            }],
432        };
433
434        let bytes = request.encode().unwrap();
435
436        assert_eq!(&bytes[0..4], &[0, 0, 0, 3]);
437        assert!(bytes.len() > 80);
438    }
439
440    #[test]
441    fn record_batch_encoding_roundtrips_through_fetch_decoder() {
442        let record_set = encode_record_batch_set(&[RecordBatchMessage::new(
443            Some(b"order-1".to_vec()),
444            Some(b"created".to_vec()),
445            1_000,
446        )
447        .header("source", Some(b"checkout".to_vec()))])
448        .unwrap();
449
450        let mut bytes = Encoder::new();
451        bytes.write_i32(0);
452        bytes.write_i32(1);
453        bytes.write_string("orders").unwrap();
454        bytes.write_i32(1);
455        bytes.write_i32(0);
456        bytes.write_i16(0);
457        bytes.write_i64(43);
458        bytes.write_bytes(&record_set).unwrap();
459        let bytes = bytes.into_bytes();
460
461        let mut decoder = Decoder::new(&bytes);
462        let response = FetchResponseV2::decode_body(&mut decoder).unwrap();
463        let record = &response.responses[0].partitions[0].records[0];
464
465        assert_eq!(record.offset, 0);
466        assert_eq!(record.timestamp_ms, 1_000);
467        assert_eq!(record.key.as_deref(), Some(&b"order-1"[..]));
468        assert_eq!(record.value.as_deref(), Some(&b"created"[..]));
469        assert!(decoder.is_empty());
470    }
471
472    #[test]
473    fn reports_message_set_encoded_len() {
474        let records = [MessageSetMessage::new(
475            Some(b"order-1".to_vec()),
476            Some(b"created".to_vec()),
477            1_000,
478        )];
479
480        assert_eq!(
481            encoded_message_set_len(&records).unwrap(),
482            encode_message_set(&records).unwrap().len()
483        );
484    }
485
486    #[test]
487    fn reports_record_batch_set_encoded_len() {
488        let records =
489            [
490                RecordBatchMessage::new(
491                    Some(b"order-1".to_vec()),
492                    Some(b"created".to_vec()),
493                    1_000,
494                )
495                .header("source", Some(b"checkout".to_vec())),
496            ];
497
498        assert_eq!(
499            encoded_record_batch_set_len(&records).unwrap(),
500            encode_record_batch_set(&records).unwrap().len()
501        );
502    }
503
504    #[test]
505    fn decodes_produce_response_v2() {
506        let bytes = [
507            0, 0, 0, 1, // topic response count
508            0, 6, b'o', b'r', b'd', b'e', b'r', b's', // topic
509            0, 0, 0, 1, // partition response count
510            0, 0, 0, 0, // partition
511            0, 0, // error code
512            0, 0, 0, 0, 0, 0, 0, 42, // base offset
513            0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, // log append time -1
514            0, 0, 0, 0, // throttle time
515        ];
516        let mut decoder = Decoder::new(&bytes);
517        let response = ProduceResponseV2::decode_body(&mut decoder).unwrap();
518
519        assert_eq!(response.throttle_time_ms, 0);
520        assert_eq!(response.responses[0].name, "orders");
521        assert_eq!(response.responses[0].partitions[0].partition_index, 0);
522        assert_eq!(response.responses[0].partitions[0].error_code, 0);
523        assert_eq!(response.responses[0].partitions[0].base_offset, 42);
524        assert!(decoder.is_empty());
525    }
526}