1use bytes::Bytes;
12use tokio_util::{
13 bytes::{Buf, BufMut, BytesMut},
14 codec::{Decoder, Encoder, LengthDelimitedCodec},
15};
16
17mod two_part;
18pub mod zero_copy_decoder;
19
20pub use two_part::{TwoPartCodec, TwoPartMessage, TwoPartMessageType};
21pub use zero_copy_decoder::{TcpRequestMessageZeroCopy, ZeroCopyTcpDecoder};
22
23const TCP_REQUEST_ENDPOINT_LEN_WIDTH: usize = 2;
24const TCP_REQUEST_HEADERS_LEN_WIDTH: usize = 2;
25const TCP_REQUEST_PAYLOAD_LEN_WIDTH: usize = 4;
26
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
28struct TcpRequestWireHeader {
29 endpoint_len: usize,
30 headers_len: usize,
31 payload_len: usize,
32 header_size: usize,
33 total_len: usize,
34}
35
36impl TcpRequestWireHeader {
37 fn endpoint_start(&self) -> usize {
38 TCP_REQUEST_ENDPOINT_LEN_WIDTH
39 }
40
41 fn endpoint_end(&self) -> usize {
42 self.endpoint_start() + self.endpoint_len
43 }
44
45 fn headers_start(&self) -> usize {
46 self.endpoint_end() + TCP_REQUEST_HEADERS_LEN_WIDTH
47 }
48
49 fn headers_end(&self) -> usize {
50 self.headers_start() + self.headers_len
51 }
52
53 fn payload_start(&self) -> usize {
54 self.header_size
55 }
56}
57
58fn tcp_request_header_size(endpoint_len: usize, headers_len: usize) -> usize {
59 TCP_REQUEST_ENDPOINT_LEN_WIDTH
60 + endpoint_len
61 + TCP_REQUEST_HEADERS_LEN_WIDTH
62 + headers_len
63 + TCP_REQUEST_PAYLOAD_LEN_WIDTH
64}
65
66fn tcp_request_total_len(
67 endpoint_len: usize,
68 headers_len: usize,
69 payload_len: usize,
70) -> Result<TcpRequestWireHeader, std::io::Error> {
71 let header_size = tcp_request_header_size(endpoint_len, headers_len);
72 let total_len = header_size.checked_add(payload_len).ok_or_else(|| {
73 std::io::Error::new(
74 std::io::ErrorKind::InvalidData,
75 "TCP request message length overflow",
76 )
77 })?;
78
79 Ok(TcpRequestWireHeader {
80 endpoint_len,
81 headers_len,
82 payload_len,
83 header_size,
84 total_len,
85 })
86}
87
88fn validate_tcp_request_encode_lengths(
89 endpoint_len: usize,
90 headers_len: usize,
91 payload_len: usize,
92) -> Result<TcpRequestWireHeader, std::io::Error> {
93 if endpoint_len > u16::MAX as usize {
94 return Err(std::io::Error::new(
95 std::io::ErrorKind::InvalidInput,
96 format!("Endpoint path too long: {} bytes", endpoint_len),
97 ));
98 }
99
100 if headers_len > u16::MAX as usize {
101 return Err(std::io::Error::new(
102 std::io::ErrorKind::InvalidInput,
103 format!("Headers too large: {} bytes", headers_len),
104 ));
105 }
106
107 if payload_len > u32::MAX as usize {
108 return Err(std::io::Error::new(
109 std::io::ErrorKind::InvalidInput,
110 format!("Payload too large: {} bytes", payload_len),
111 ));
112 }
113
114 tcp_request_total_len(endpoint_len, headers_len, payload_len)
115}
116
117fn tcp_request_endpoint_len(bytes: &[u8]) -> Result<usize, std::io::Error> {
118 if bytes.len() < TCP_REQUEST_ENDPOINT_LEN_WIDTH {
119 return Err(std::io::Error::new(
120 std::io::ErrorKind::UnexpectedEof,
121 "Not enough bytes for endpoint path length",
122 ));
123 }
124
125 Ok(u16::from_be_bytes([bytes[0], bytes[1]]) as usize)
126}
127
128fn tcp_request_headers_len(bytes: &[u8], endpoint_len: usize) -> Result<usize, std::io::Error> {
129 let endpoint_end = TCP_REQUEST_ENDPOINT_LEN_WIDTH + endpoint_len;
130 if bytes.len() < endpoint_end {
131 return Err(std::io::Error::new(
132 std::io::ErrorKind::UnexpectedEof,
133 "Not enough bytes for endpoint path",
134 ));
135 }
136
137 if bytes.len() < endpoint_end + TCP_REQUEST_HEADERS_LEN_WIDTH {
138 return Err(std::io::Error::new(
139 std::io::ErrorKind::UnexpectedEof,
140 "Not enough bytes for headers length",
141 ));
142 }
143
144 Ok(u16::from_be_bytes([bytes[endpoint_end], bytes[endpoint_end + 1]]) as usize)
145}
146
147fn parse_tcp_request_frame_header(bytes: &[u8]) -> Result<TcpRequestWireHeader, std::io::Error> {
148 let endpoint_len = tcp_request_endpoint_len(bytes)?;
149 let headers_len = tcp_request_headers_len(bytes, endpoint_len)?;
150
151 let headers_end =
152 TCP_REQUEST_ENDPOINT_LEN_WIDTH + endpoint_len + TCP_REQUEST_HEADERS_LEN_WIDTH + headers_len;
153 if bytes.len() < headers_end {
154 return Err(std::io::Error::new(
155 std::io::ErrorKind::UnexpectedEof,
156 "Not enough bytes for headers",
157 ));
158 }
159
160 if bytes.len() < headers_end + TCP_REQUEST_PAYLOAD_LEN_WIDTH {
161 return Err(std::io::Error::new(
162 std::io::ErrorKind::UnexpectedEof,
163 "Not enough bytes for payload length",
164 ));
165 }
166
167 let payload_len = u32::from_be_bytes([
168 bytes[headers_end],
169 bytes[headers_end + 1],
170 bytes[headers_end + 2],
171 bytes[headers_end + 3],
172 ]) as usize;
173
174 tcp_request_total_len(endpoint_len, headers_len, payload_len)
175}
176
177fn parse_tcp_request_frame(bytes: &[u8]) -> Result<TcpRequestWireHeader, std::io::Error> {
178 let parsed = parse_tcp_request_frame_header(bytes)?;
179 if bytes.len() < parsed.total_len {
180 return Err(std::io::Error::new(
181 std::io::ErrorKind::UnexpectedEof,
182 format!(
183 "Not enough bytes for payload: expected {}, got {}",
184 parsed.payload_len,
185 bytes.len().saturating_sub(parsed.payload_start())
186 ),
187 ));
188 }
189
190 Ok(parsed)
191}
192
193fn check_tcp_request_max_message_size(
194 total_len: usize,
195 max_message_size: usize,
196) -> Result<(), std::io::Error> {
197 if total_len > max_message_size {
198 return Err(std::io::Error::new(
199 std::io::ErrorKind::InvalidData,
200 format!(
201 "message too large: {} bytes (max: {} bytes)",
202 total_len, max_message_size
203 ),
204 ));
205 }
206
207 Ok(())
208}
209
210#[derive(Debug, Clone, PartialEq, Eq)]
220pub struct TcpRequestMessage {
221 pub endpoint_path: String,
222 pub headers: std::collections::HashMap<String, String>,
223 pub payload: Bytes,
224}
225
226#[derive(Debug, Clone, PartialEq, Eq)]
231pub struct TcpRequestFrame {
232 pub header: Bytes,
233 pub payload: Bytes,
234}
235
236impl TcpRequestFrame {
237 pub fn encoded_len(&self) -> usize {
238 self.header.len() + self.payload.len()
239 }
240}
241
242impl TcpRequestMessage {
243 pub fn new(endpoint_path: String, payload: Bytes) -> Self {
244 Self {
245 endpoint_path,
246 headers: std::collections::HashMap::new(),
247 payload,
248 }
249 }
250
251 pub fn with_headers(
252 endpoint_path: String,
253 headers: std::collections::HashMap<String, String>,
254 payload: Bytes,
255 ) -> Self {
256 Self {
257 endpoint_path,
258 headers,
259 payload,
260 }
261 }
262
263 pub fn encode(&self) -> Result<Bytes, std::io::Error> {
265 let endpoint_bytes = self.endpoint_path.as_bytes();
266 let endpoint_len = endpoint_bytes.len();
267
268 let headers_json = serde_json::to_vec(&self.headers).map_err(|e| {
270 std::io::Error::new(
271 std::io::ErrorKind::InvalidInput,
272 format!("Failed to encode headers: {}", e),
273 )
274 })?;
275 let headers_len = headers_json.len();
276
277 let parsed =
278 validate_tcp_request_encode_lengths(endpoint_len, headers_len, self.payload.len())?;
279
280 let mut buf = BytesMut::with_capacity(parsed.total_len);
282
283 buf.put_u16(endpoint_len as u16);
285
286 buf.put_slice(endpoint_bytes);
288
289 buf.put_u16(headers_len as u16);
291
292 buf.put_slice(&headers_json);
294
295 buf.put_u32(self.payload.len() as u32);
297
298 buf.put_slice(&self.payload);
300
301 Ok(buf.freeze())
303 }
304
305 pub fn into_frame(self) -> Result<TcpRequestFrame, std::io::Error> {
309 let endpoint_bytes = self.endpoint_path.as_bytes();
310 let endpoint_len = endpoint_bytes.len();
311
312 let headers_json = serde_json::to_vec(&self.headers).map_err(|e| {
313 std::io::Error::new(
314 std::io::ErrorKind::InvalidInput,
315 format!("Failed to encode headers: {}", e),
316 )
317 })?;
318 let headers_len = headers_json.len();
319 let payload_len = self.payload.len();
320
321 let parsed = validate_tcp_request_encode_lengths(endpoint_len, headers_len, payload_len)?;
322 let mut header = BytesMut::with_capacity(parsed.header_size);
323
324 header.put_u16(endpoint_len as u16);
325 header.put_slice(endpoint_bytes);
326 header.put_u16(headers_len as u16);
327 header.put_slice(&headers_json);
328 header.put_u32(payload_len as u32);
329
330 Ok(TcpRequestFrame {
331 header: header.freeze(),
332 payload: self.payload,
333 })
334 }
335
336 pub fn decode(bytes: &Bytes) -> Result<Self, std::io::Error> {
338 let parsed = parse_tcp_request_frame(bytes)?;
339
340 let endpoint_path =
342 String::from_utf8(bytes[parsed.endpoint_start()..parsed.endpoint_end()].to_vec())
343 .map_err(|e| {
344 std::io::Error::new(
345 std::io::ErrorKind::InvalidData,
346 format!("Invalid UTF-8 in endpoint path: {}", e),
347 )
348 })?;
349
350 let headers: std::collections::HashMap<String, String> = serde_json::from_slice(
352 &bytes[parsed.headers_start()..parsed.headers_end()],
353 )
354 .map_err(|e| {
355 std::io::Error::new(
356 std::io::ErrorKind::InvalidData,
357 format!("Invalid JSON in headers: {}", e),
358 )
359 })?;
360
361 let payload = bytes.slice(parsed.payload_start()..parsed.total_len);
363
364 Ok(Self {
365 endpoint_path,
366 headers,
367 payload,
368 })
369 }
370}
371
372#[derive(Debug, Clone, PartialEq, Eq)]
378pub struct TcpResponseMessage {
379 pub data: Bytes,
380}
381
382impl TcpResponseMessage {
383 pub fn new(data: Bytes) -> Self {
384 Self { data }
385 }
386
387 pub fn empty() -> Self {
388 Self { data: Bytes::new() }
389 }
390
391 pub fn encode(&self) -> Result<Bytes, std::io::Error> {
393 if self.data.len() > u32::MAX as usize {
394 return Err(std::io::Error::new(
395 std::io::ErrorKind::InvalidInput,
396 format!("Response too large: {} bytes", self.data.len()),
397 ));
398 }
399
400 let mut buf = BytesMut::with_capacity(4 + self.data.len());
401 buf.put_u32(self.data.len() as u32);
402 buf.put_slice(&self.data);
403 Ok(buf.freeze())
404 }
405
406 pub fn decode(bytes: &Bytes) -> Result<Self, std::io::Error> {
408 if bytes.len() < 4 {
409 return Err(std::io::Error::new(
410 std::io::ErrorKind::UnexpectedEof,
411 "Not enough bytes for response length",
412 ));
413 }
414
415 let len = u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]) as usize;
417
418 if bytes.len() < 4 + len {
419 return Err(std::io::Error::new(
420 std::io::ErrorKind::UnexpectedEof,
421 format!(
422 "Not enough bytes for response: expected {}, got {}",
423 len,
424 bytes.len() - 4
425 ),
426 ));
427 }
428
429 let data = bytes.slice(4..4 + len);
431
432 Ok(Self { data })
433 }
434}
435
436const RESPONSE_LENGTH_WIDTH: usize = std::mem::size_of::<u32>();
437
438fn response_payload_limit(max_message_size: Option<usize>) -> usize {
439 max_message_size
440 .map(|max| max.saturating_sub(RESPONSE_LENGTH_WIDTH))
441 .unwrap_or(u32::MAX as usize)
442 .min(u32::MAX as usize)
443}
444
445#[derive(Clone)]
448pub struct TcpResponseCodec {
449 decoder: LengthDelimitedCodec,
450 reject_all_frames: bool,
451}
452
453impl TcpResponseCodec {
454 pub fn new(max_message_size: Option<usize>) -> Self {
455 let reject_all_frames = max_message_size.is_some_and(|max| max < RESPONSE_LENGTH_WIDTH);
456 let decoder = LengthDelimitedCodec::builder()
457 .length_field_type::<u32>()
458 .big_endian()
459 .length_adjustment(RESPONSE_LENGTH_WIDTH as isize)
460 .num_skip(0)
461 .max_frame_length(response_payload_limit(max_message_size))
462 .new_codec();
463
464 Self {
465 decoder,
466 reject_all_frames,
467 }
468 }
469}
470
471impl Default for TcpResponseCodec {
472 fn default() -> Self {
473 Self::new(None)
474 }
475}
476
477impl Decoder for TcpResponseCodec {
478 type Item = TcpResponseMessage;
479 type Error = std::io::Error;
480
481 fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
482 if self.reject_all_frames && src.len() >= RESPONSE_LENGTH_WIDTH {
483 return Err(std::io::ErrorKind::InvalidData.into());
484 }
485
486 self.decoder.decode(src).map(|frame| {
487 frame.map(|mut frame| {
488 frame.advance(RESPONSE_LENGTH_WIDTH);
489 TcpResponseMessage {
490 data: frame.freeze(),
491 }
492 })
493 })
494 }
495}
496
497impl Encoder<TcpResponseMessage> for TcpResponseCodec {
498 type Error = std::io::Error;
499
500 fn encode(&mut self, item: TcpResponseMessage, dst: &mut BytesMut) -> Result<(), Self::Error> {
501 if self.reject_all_frames {
502 return Err(std::io::ErrorKind::InvalidInput.into());
503 }
504
505 LengthDelimitedCodec::builder()
506 .length_field_type::<u32>()
507 .big_endian()
508 .max_frame_length(self.decoder.max_frame_length())
509 .new_codec()
510 .encode(item.data, dst)
511 }
512}
513
514#[cfg(test)]
515mod tests {
516 use super::*;
517
518 #[test]
519 fn test_tcp_request_encode_decode() {
520 let msg = TcpRequestMessage::new(
521 "test.endpoint".to_string(),
522 Bytes::from(vec![1, 2, 3, 4, 5]),
523 );
524
525 let encoded = msg.encode().unwrap();
526 let decoded = TcpRequestMessage::decode(&encoded).unwrap();
527
528 assert_eq!(decoded, msg);
529 }
530
531 #[test]
532 fn test_tcp_request_empty_payload() {
533 let msg = TcpRequestMessage::new("test".to_string(), Bytes::new());
534
535 let encoded = msg.encode().unwrap();
536 let decoded = TcpRequestMessage::decode(&encoded).unwrap();
537
538 assert_eq!(decoded, msg);
539 }
540
541 #[test]
542 fn test_tcp_request_into_frame_matches_encode() {
543 let mut basic_headers = std::collections::HashMap::new();
544 basic_headers.insert("request-id".to_string(), "abc-123".to_string());
545
546 let mut multibyte_headers = std::collections::HashMap::new();
547 multibyte_headers.insert("trace".to_string(), "snowman-โ".to_string());
548 multibyte_headers.insert("emoji".to_string(), "rocket-๐".to_string());
549
550 let mut large_headers = std::collections::HashMap::new();
551 large_headers.insert("x-long".to_string(), "v".repeat(4096));
552
553 let cases = [
554 (
555 "test.endpoint".to_string(),
556 basic_headers,
557 Bytes::from_static(b"payload-body"),
558 ),
559 (
560 "empty.payload".to_string(),
561 std::collections::HashMap::new(),
562 Bytes::new(),
563 ),
564 (
565 "unicode.endpoint".to_string(),
566 multibyte_headers,
567 Bytes::from("ใใใซใกใฏ"),
568 ),
569 (
570 "large.payload".to_string(),
571 large_headers,
572 Bytes::from(vec![42u8; 64 * 1024]),
573 ),
574 ];
575
576 for (endpoint, headers, payload) in cases {
577 let msg = TcpRequestMessage::with_headers(endpoint, headers, payload.clone());
578 let encoded = msg.clone().encode().unwrap();
579 let frame = msg.into_frame().unwrap();
580
581 assert_eq!(frame.encoded_len(), encoded.len());
582 if !payload.is_empty() {
583 assert_eq!(frame.payload.as_ptr(), payload.as_ptr());
584 }
585
586 let mut combined = BytesMut::with_capacity(frame.encoded_len());
587 combined.put_slice(&frame.header);
588 combined.put_slice(&frame.payload);
589 assert_eq!(combined.freeze(), encoded);
590 }
591 }
592
593 #[test]
594 fn test_tcp_request_large_payload() {
595 let payload = Bytes::from(vec![42u8; 1024 * 1024]); let msg = TcpRequestMessage::new("large".to_string(), payload);
597
598 let encoded = msg.encode().unwrap();
599 let decoded = TcpRequestMessage::decode(&encoded).unwrap();
600
601 assert_eq!(decoded, msg);
602 }
603
604 #[test]
605 fn test_tcp_request_decode_truncated() {
606 let msg = TcpRequestMessage::new("test".to_string(), Bytes::from(vec![1, 2, 3, 4, 5]));
607 let encoded = msg.encode().unwrap();
608
609 let truncated = encoded.slice(..encoded.len() - 2);
611 let result = TcpRequestMessage::decode(&truncated);
612
613 assert!(result.is_err());
614 }
615
616 #[test]
617 fn test_tcp_request_decode_invalid_endpoint_utf8() {
618 let mut encoded = BytesMut::new();
619 encoded.put_u16(2);
620 encoded.put_slice(&[0xff, 0xff]);
621 encoded.put_u16(2);
622 encoded.put_slice(b"{}");
623 encoded.put_u32(0);
624
625 let result = TcpRequestMessage::decode(&encoded.freeze());
626
627 assert!(result.is_err());
628 let err = result.unwrap_err();
629 assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
630 assert!(err.to_string().contains("Invalid UTF-8"));
631 }
632
633 #[test]
634 fn test_tcp_request_decode_invalid_headers_json() {
635 let mut encoded = BytesMut::new();
636 encoded.put_u16(4);
637 encoded.put_slice(b"test");
638 encoded.put_u16(1);
639 encoded.put_slice(b"{");
640 encoded.put_u32(0);
641
642 let result = TcpRequestMessage::decode(&encoded.freeze());
643
644 assert!(result.is_err());
645 let err = result.unwrap_err();
646 assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
647 assert!(err.to_string().contains("Invalid JSON"));
648 }
649
650 #[test]
651 fn test_tcp_request_empty_endpoint_path() {
652 let msg = TcpRequestMessage::new(String::new(), Bytes::from_static(b"payload"));
653
654 let encoded = msg.encode().unwrap();
655 let decoded = TcpRequestMessage::decode(&encoded).unwrap();
656
657 assert_eq!(decoded, msg);
658 }
659
660 #[test]
661 fn test_tcp_response_encode_decode() {
662 let msg = TcpResponseMessage::new(Bytes::from(vec![1, 2, 3, 4, 5]));
663
664 let encoded = msg.encode().unwrap();
665 let decoded = TcpResponseMessage::decode(&encoded).unwrap();
666
667 assert_eq!(decoded, msg);
668 }
669
670 #[test]
671 fn test_tcp_response_empty() {
672 let msg = TcpResponseMessage::empty();
673
674 let encoded = msg.encode().unwrap();
675 let decoded = TcpResponseMessage::decode(&encoded).unwrap();
676
677 assert_eq!(decoded, msg);
678 assert_eq!(decoded.data.len(), 0);
679 }
680
681 #[test]
682 fn test_tcp_response_decode_truncated() {
683 let msg = TcpResponseMessage::new(Bytes::from(vec![1, 2, 3, 4, 5]));
684 let encoded = msg.encode().unwrap();
685
686 let truncated = encoded.slice(..3);
688 let result = TcpResponseMessage::decode(&truncated);
689
690 assert!(result.is_err());
691 }
692
693 #[test]
694 fn test_tcp_request_unicode_endpoint() {
695 let msg = TcpRequestMessage::new("ัะตัั.็ซฏ็น".to_string(), Bytes::from(vec![1, 2, 3]));
696
697 let encoded = msg.encode().unwrap();
698 let decoded = TcpRequestMessage::decode(&encoded).unwrap();
699
700 assert_eq!(decoded, msg);
701 }
702
703 #[test]
704 fn test_tcp_response_codec() {
705 use tokio_util::codec::{Decoder, Encoder};
706
707 let msg = TcpResponseMessage::new(Bytes::from(vec![1, 2, 3, 4, 5]));
708
709 let mut codec = TcpResponseCodec::new(None);
710 let mut buf = BytesMut::new();
711
712 codec.encode(msg.clone(), &mut buf).unwrap();
714
715 let decoded = codec.decode(&mut buf).unwrap().unwrap();
717 assert_eq!(decoded, msg);
718 }
719
720 #[test]
721 fn test_tcp_response_codec_partial() {
722 use tokio_util::codec::Decoder;
723
724 let msg = TcpResponseMessage::new(Bytes::from(vec![1, 2, 3, 4, 5]));
725
726 let encoded = msg.encode().unwrap();
727 let mut codec = TcpResponseCodec::new(None);
728
729 let mut buf = BytesMut::from(&encoded[..3]);
731 assert!(codec.decode(&mut buf).unwrap().is_none());
732
733 buf.extend_from_slice(&encoded[3..]);
735 let decoded = codec.decode(&mut buf).unwrap().unwrap();
736 assert_eq!(decoded, msg);
737 }
738
739 #[test]
740 fn test_tcp_response_codec_max_size() {
741 use tokio_util::codec::Encoder;
742
743 let msg = TcpResponseMessage::new(Bytes::from(vec![1, 2, 3, 4, 5]));
744
745 let mut codec = TcpResponseCodec::new(Some(5)); let mut buf = BytesMut::new();
747
748 let result = codec.encode(msg, &mut buf);
749 assert!(result.is_err());
750 }
751}