1use bytes::{BufMut, Bytes, BytesMut};
4
5use crate::{
6 KafkaString,
7 error::Result,
8 frame::{FrameError, FrameErrorKind, MAX_FRAME_LENGTH},
9 generated::{ApiKey, RequestHeaderData, ResponseHeaderData},
10 version::{request_header_version, response_header_version},
11};
12
13#[derive(Debug, Clone, Copy)]
15pub struct RequestFrameSpec<'a> {
16 pub api_key: ApiKey,
18 pub api_version: i16,
20 pub correlation_id: i32,
22 pub client_id: &'a str,
24 pub capacity_hint: usize,
26}
27
28#[derive(Debug, Clone, PartialEq, Eq)]
30pub struct ResponseEnvelope {
31 pub correlation_id: i32,
33 pub body: Bytes,
35}
36
37pub fn encode_request_frame<F>(spec: RequestFrameSpec<'_>, write_body: F) -> Result<BytesMut>
39where
40 F: FnOnce(&mut BytesMut) -> Result<()>,
41{
42 let mut frame = BytesMut::with_capacity(spec.capacity_hint.max(4));
43 encode_request_frame_with_buffer(&mut frame, spec, write_body)?;
44 Ok(frame)
45}
46
47pub fn encode_request_frame_with_buffer<F>(
49 frame: &mut BytesMut,
50 spec: RequestFrameSpec<'_>,
51 write_body: F,
52) -> Result<()>
53where
54 F: FnOnce(&mut BytesMut) -> Result<()>,
55{
56 frame.clear();
57 frame.reserve(spec.capacity_hint.max(4));
58 frame.put_i32(0);
59 write_request_header(frame, spec)?;
60 write_body(frame)?;
61 finish_request_frame(frame)
62}
63
64pub fn encode_request_frame_body(spec: RequestFrameSpec<'_>, body: &[u8]) -> Result<BytesMut> {
66 encode_request_frame(spec, |frame| {
67 frame.extend_from_slice(body);
68 Ok(())
69 })
70}
71
72pub fn encode_request_frame_body_with_buffer(
75 frame: &mut BytesMut,
76 spec: RequestFrameSpec<'_>,
77 body: &[u8],
78) -> Result<()> {
79 encode_request_frame_with_buffer(frame, spec, |frame| {
80 frame.extend_from_slice(body);
81 Ok(())
82 })
83}
84
85pub fn request_frame_capacity_hint(spec: RequestFrameSpec<'_>, body_len: usize) -> Result<usize> {
90 let header_len = request_header_len(spec)?;
91 let payload_len = header_len.checked_add(body_len).ok_or_else(|| {
92 crate::ProtocolError::Frame(FrameError::from(FrameErrorKind::TooLarge {
93 length: i32::MAX,
94 max: MAX_FRAME_LENGTH,
95 }))
96 })?;
97 validate_payload_len(payload_len)?;
98 payload_len.checked_add(4).ok_or_else(|| {
99 crate::ProtocolError::Frame(FrameError::from(FrameErrorKind::TooLarge {
100 length: i32::MAX,
101 max: MAX_FRAME_LENGTH,
102 }))
103 })
104}
105
106pub fn decode_response_envelope(
108 api_key: ApiKey,
109 api_version: i16,
110 mut bytes: Bytes,
111) -> Result<ResponseEnvelope> {
112 let header_version = response_header_version(api_key as i16, api_version);
113 let header = ResponseHeaderData::read(&mut bytes, header_version)?;
114 Ok(ResponseEnvelope {
115 correlation_id: header.correlation_id,
116 body: bytes,
117 })
118}
119
120fn write_request_header(buf: &mut BytesMut, spec: RequestFrameSpec<'_>) -> Result<()> {
121 request_header(spec).write(
122 buf,
123 request_header_version(spec.api_key as i16, spec.api_version),
124 )?;
125 Ok(())
126}
127
128fn request_header_len(spec: RequestFrameSpec<'_>) -> Result<usize> {
129 let header_version = request_header_version(spec.api_key as i16, spec.api_version);
130 request_header(spec).encoded_len(header_version)
131}
132
133fn request_header(spec: RequestFrameSpec<'_>) -> RequestHeaderData {
134 RequestHeaderData {
135 request_api_key: spec.api_key as i16,
136 request_api_version: spec.api_version,
137 correlation_id: spec.correlation_id,
138 client_id: Some(KafkaString::from(spec.client_id.to_owned())),
139 _unknown_tagged_fields: Vec::new(),
140 }
141}
142
143fn finish_request_frame(request_frame: &mut BytesMut) -> Result<()> {
144 let payload_len = request_frame.len().checked_sub(4).ok_or_else(|| {
145 crate::ProtocolError::Frame(FrameError::from(FrameErrorKind::Truncated {
146 needed: 4,
147 available: request_frame.len(),
148 }))
149 })?;
150 validate_payload_len(payload_len)?;
151 let payload_len = i32::try_from(payload_len).map_err(|_error| {
152 crate::ProtocolError::Frame(FrameError::from(FrameErrorKind::TooLarge {
153 length: i32::MAX,
154 max: MAX_FRAME_LENGTH,
155 }))
156 })?;
157 let Some(length_slot) = request_frame.get_mut(..4) else {
158 return Err(crate::ProtocolError::Frame(FrameError::from(
159 FrameErrorKind::Truncated {
160 needed: 4,
161 available: request_frame.len(),
162 },
163 )));
164 };
165 length_slot.copy_from_slice(&payload_len.to_be_bytes());
166 Ok(())
167}
168
169fn validate_payload_len(payload_len: usize) -> Result<()> {
170 let max_frame_length = usize::try_from(MAX_FRAME_LENGTH).unwrap_or(usize::MAX);
171 if payload_len > max_frame_length {
172 return Err(crate::ProtocolError::Frame(FrameError::from(
173 FrameErrorKind::TooLarge {
174 length: i32::try_from(payload_len).unwrap_or(i32::MAX),
175 max: MAX_FRAME_LENGTH,
176 },
177 )));
178 }
179 Ok(())
180}
181
182#[cfg(test)]
183mod tests {
184 use bytes::{Buf, BytesMut};
185
186 use super::{
187 RequestFrameSpec, decode_response_envelope, encode_request_frame,
188 request_frame_capacity_hint,
189 };
190 use crate::{
191 KafkaString,
192 generated::{
193 ApiKey, ApiVersionsRequestData, ApiVersionsResponseData, RequestHeaderData,
194 ResponseHeaderData,
195 },
196 version::{request_header_version, response_header_version},
197 };
198
199 #[test]
200 fn encode_request_frame_writes_length_header_and_body_in_one_buffer() {
201 let request = ApiVersionsRequestData {
202 client_software_name: KafkaString::from("kacrab".to_owned()),
203 client_software_version: KafkaString::from("0.0.1".to_owned()),
204 _unknown_tagged_fields: Vec::new(),
205 };
206
207 let frame = encode_request_frame(
208 RequestFrameSpec {
209 api_key: ApiKey::ApiVersions,
210 api_version: 3,
211 correlation_id: 42,
212 client_id: "client-a",
213 capacity_hint: 64,
214 },
215 |buf| request.write(buf, 3),
216 )
217 .expect("request frame");
218
219 let mut payload = frame.freeze();
220 let frame_len = payload.get_i32();
221 assert_eq!(usize::try_from(frame_len).unwrap(), payload.remaining());
222
223 let header_version = request_header_version(ApiKey::ApiVersions as i16, 3);
224 let header = RequestHeaderData::read(&mut payload, header_version).expect("header");
225 assert_eq!(header.request_api_key, ApiKey::ApiVersions as i16);
226 assert_eq!(header.request_api_version, 3);
227 assert_eq!(header.correlation_id, 42);
228 assert_eq!(
229 header.client_id.as_ref().map(KafkaString::as_str),
230 Some("client-a")
231 );
232
233 let decoded = ApiVersionsRequestData::read(&mut payload, 3).expect("body");
234 assert_eq!(decoded.client_software_name.as_str(), "kacrab");
235 assert_eq!(decoded.client_software_version.as_str(), "0.0.1");
236 assert_eq!(payload.remaining(), 0);
237 }
238
239 #[test]
240 fn request_frame_capacity_hint_matches_encoded_frame_len() {
241 let request = ApiVersionsRequestData {
242 client_software_name: KafkaString::from("kacrab".to_owned()),
243 client_software_version: KafkaString::from("0.0.1".to_owned()),
244 _unknown_tagged_fields: Vec::new(),
245 };
246 let body_len = request.encoded_len(3).expect("body length");
247 let spec = RequestFrameSpec {
248 api_key: ApiKey::ApiVersions,
249 api_version: 3,
250 correlation_id: 42,
251 client_id: "client-a",
252 capacity_hint: 0,
253 };
254 let capacity_hint =
255 request_frame_capacity_hint(spec, body_len).expect("request capacity hint");
256
257 let frame = encode_request_frame(
258 RequestFrameSpec {
259 capacity_hint,
260 ..spec
261 },
262 |buf| request.write(buf, 3),
263 )
264 .expect("request frame");
265
266 assert_eq!(capacity_hint, frame.len());
267 assert_eq!(capacity_hint, frame.capacity());
268 }
269
270 #[test]
271 fn decode_response_envelope_returns_correlation_and_body_bytes() {
272 let response = ApiVersionsResponseData::default();
273 let mut body = BytesMut::new();
274 response.write(&mut body, 3).expect("response body");
275
276 let mut payload = BytesMut::new();
277 ResponseHeaderData {
278 correlation_id: 7,
279 _unknown_tagged_fields: Vec::new(),
280 }
281 .write(
282 &mut payload,
283 response_header_version(ApiKey::ApiVersions as i16, 3),
284 )
285 .expect("response header");
286 payload.extend_from_slice(&body);
287
288 let mut envelope = decode_response_envelope(ApiKey::ApiVersions, 3, payload.freeze())
289 .expect("response envelope");
290
291 assert_eq!(envelope.correlation_id, 7);
292 let decoded = ApiVersionsResponseData::read(&mut envelope.body, 3).expect("response");
293 assert_eq!(decoded.error_code, 0);
294 assert_eq!(envelope.body.remaining(), 0);
295 }
296}