1use std::{
2 io::{Read, Write},
3 pin::Pin,
4 task::{Context, Poll},
5};
6
7use bytes::{Buf, BufMut, Bytes, BytesMut};
8use flate2::{Compression, read::GzDecoder, write::GzEncoder};
9use futures_core::Stream;
10use strum::FromRepr;
11
12const LENGTH_PREFIX_SIZE: usize = 3;
34const STATUS_CODE_SIZE: usize = 2;
35const COMPRESSION_THRESHOLD_BYTES: usize = 1024; const MAX_FRAME_BYTES: usize = 2 * 1024 * 1024; const FLAG_TOTAL_SIZE: usize = 1;
52const MAX_FRAME_PAYLOAD_BYTES: usize = MAX_FRAME_BYTES - FLAG_TOTAL_SIZE;
54const MAX_DECOMPRESSED_PAYLOAD_BYTES: usize = MAX_FRAME_PAYLOAD_BYTES;
55const FLAG_TERMINAL: u8 = 0b1000_0000;
56const FLAG_COMPRESSION_MASK: u8 = 0b0110_0000;
57const FLAG_COMPRESSION_SHIFT: u8 = 5;
58const FLAG_RECONNECT_ADVISED: u8 = 0b0001_0000;
59
60#[derive(Debug, Clone, Copy, PartialEq, Eq, FromRepr)]
61#[repr(u8)]
62pub enum CompressionAlgorithm {
63 None = 0,
64 Zstd = 1,
65 Gzip = 2,
66}
67
68impl CompressionAlgorithm {
69 pub fn from_accept_encoding(headers: &http::HeaderMap) -> Self {
70 let mut gzip = false;
71 for header_value in headers.get_all(http::header::ACCEPT_ENCODING) {
72 if let Ok(value) = header_value.to_str() {
73 for encoding in value.split(',') {
74 let mut parts = encoding.split(';');
75 let encoding = parts.next().unwrap_or("").trim();
76 if parts.any(is_zero_qvalue) {
77 continue;
78 }
79 if encoding.eq_ignore_ascii_case("zstd") {
80 return Self::Zstd;
81 } else if encoding.eq_ignore_ascii_case("gzip") {
82 gzip = true;
83 }
84 }
85 }
86 }
87 if gzip { Self::Gzip } else { Self::None }
88 }
89}
90
91fn is_zero_qvalue(param: &str) -> bool {
94 let Some((name, value)) = param.split_once('=') else {
95 return false;
96 };
97 name.trim().eq_ignore_ascii_case("q") && value.trim().parse::<f32>().is_ok_and(|q| q == 0.0)
98}
99
100#[derive(Debug, Clone, PartialEq, Eq)]
101pub struct CompressedData {
102 compression: CompressionAlgorithm,
103 reconnect_advised: bool,
104 payload: Bytes,
105}
106
107impl CompressedData {
108 pub fn for_proto(
109 compression: CompressionAlgorithm,
110 proto: &impl prost::Message,
111 ) -> std::io::Result<Self> {
112 Self::compress(compression, proto.encode_to_vec())
113 }
114
115 pub fn reconnect_advised(&self) -> bool {
120 self.reconnect_advised
121 }
122
123 fn compress(compression: CompressionAlgorithm, data: Vec<u8>) -> std::io::Result<Self> {
124 if data.len() > MAX_DECOMPRESSED_PAYLOAD_BYTES {
125 return Err(std::io::Error::new(
126 std::io::ErrorKind::InvalidInput,
127 "payload exceeds decompressed limit",
128 ));
129 }
130
131 if compression == CompressionAlgorithm::None || data.len() < COMPRESSION_THRESHOLD_BYTES {
132 return Ok(Self {
133 compression: CompressionAlgorithm::None,
134 reconnect_advised: false,
135 payload: data.into(),
136 });
137 }
138 let mut buf = Vec::with_capacity(data.len());
139 match compression {
140 CompressionAlgorithm::Gzip => {
141 let mut encoder = GzEncoder::new(buf, Compression::default());
142 encoder.write_all(data.as_slice())?;
143 buf = encoder.finish()?;
144 }
145 CompressionAlgorithm::Zstd => {
146 zstd::stream::copy_encode(data.as_slice(), &mut buf, 0)?;
147 }
148 CompressionAlgorithm::None => unreachable!("handled above"),
149 };
150 let payload = Bytes::from(buf.into_boxed_slice());
151 if payload.len() > MAX_FRAME_PAYLOAD_BYTES {
152 return Err(std::io::Error::new(
153 std::io::ErrorKind::InvalidInput,
154 "compressed payload exceeds frame limit",
155 ));
156 }
157 Ok(Self {
158 compression,
159 reconnect_advised: false,
160 payload,
161 })
162 }
163
164 fn decompressed(self) -> std::io::Result<Bytes> {
165 let initial_capacity = self
166 .payload
167 .len()
168 .saturating_mul(2)
169 .clamp(COMPRESSION_THRESHOLD_BYTES, MAX_DECOMPRESSED_PAYLOAD_BYTES);
170
171 fn read_to_end_limited(
173 mut reader: impl Read,
174 initial_capacity: usize,
175 ) -> std::io::Result<Bytes> {
176 let mut limited = reader
177 .by_ref()
178 .take((MAX_DECOMPRESSED_PAYLOAD_BYTES + 1) as u64);
179 let mut buf = Vec::with_capacity(initial_capacity);
180 limited.read_to_end(&mut buf)?;
181 if buf.len() > MAX_DECOMPRESSED_PAYLOAD_BYTES {
182 return Err(std::io::Error::new(
183 std::io::ErrorKind::InvalidData,
184 "decompressed payload exceeds limit",
185 ));
186 }
187 Ok(Bytes::from(buf.into_boxed_slice()))
188 }
189
190 match self.compression {
191 CompressionAlgorithm::None => {
192 if self.payload.len() > MAX_DECOMPRESSED_PAYLOAD_BYTES {
193 return Err(std::io::Error::new(
194 std::io::ErrorKind::InvalidData,
195 "decompressed payload exceeds limit",
196 ));
197 }
198 Ok(self.payload)
199 }
200 CompressionAlgorithm::Gzip => {
201 let mut decoder = GzDecoder::new(&self.payload[..]);
202 read_to_end_limited(&mut decoder, initial_capacity)
203 }
204 CompressionAlgorithm::Zstd => {
205 let mut decoder = zstd::stream::Decoder::new(&self.payload[..])?;
206 read_to_end_limited(&mut decoder, initial_capacity)
207 }
208 }
209 }
210
211 pub fn try_into_proto<P: prost::Message + Default>(self) -> std::io::Result<P> {
212 let payload = self.decompressed()?;
213 P::decode(payload.as_ref())
214 .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e))
215 }
216}
217
218#[derive(Debug, Clone, PartialEq, Eq)]
219pub struct TerminalMessage {
220 pub status: u16,
221 pub body: String,
222}
223
224#[derive(Debug, Clone, PartialEq, Eq)]
225pub enum SessionMessage {
226 Regular(CompressedData),
227 Terminal(TerminalMessage),
228}
229
230impl From<CompressedData> for SessionMessage {
231 fn from(data: CompressedData) -> Self {
232 Self::Regular(data)
233 }
234}
235
236impl From<TerminalMessage> for SessionMessage {
237 fn from(msg: TerminalMessage) -> Self {
238 Self::Terminal(msg)
239 }
240}
241
242impl SessionMessage {
243 pub fn regular(
244 compression: CompressionAlgorithm,
245 proto: &impl prost::Message,
246 ) -> std::io::Result<Self> {
247 Ok(Self::Regular(CompressedData::for_proto(
248 compression,
249 proto,
250 )?))
251 }
252
253 pub fn encode(&self) -> Bytes {
254 let encoded_size = FLAG_TOTAL_SIZE + self.payload_size();
255 assert!(
256 encoded_size <= MAX_FRAME_BYTES,
257 "payload exceeds encoder limit"
258 );
259 let mut buf = BytesMut::with_capacity(LENGTH_PREFIX_SIZE + encoded_size);
260 buf.put_uint(encoded_size as u64, 3);
261 match self {
262 Self::Regular(msg) => {
263 let mut flag =
264 ((msg.compression as u8) << FLAG_COMPRESSION_SHIFT) & FLAG_COMPRESSION_MASK;
265 if msg.reconnect_advised {
266 flag |= FLAG_RECONNECT_ADVISED;
267 }
268 buf.put_u8(flag);
269 buf.extend_from_slice(&msg.payload);
270 }
271 Self::Terminal(msg) => {
272 buf.put_u8(FLAG_TERMINAL);
273 buf.put_u16(msg.status);
274 buf.extend_from_slice(msg.body.as_bytes());
275 }
276 }
277 buf.freeze()
278 }
279
280 fn decode_message(mut buf: Bytes) -> std::io::Result<Self> {
281 if buf.is_empty() {
282 return Err(std::io::Error::new(
283 std::io::ErrorKind::UnexpectedEof,
284 "empty frame payload",
285 ));
286 }
287 let flag = buf.get_u8();
288
289 let is_terminal = (flag & FLAG_TERMINAL) != 0;
290 if is_terminal {
291 if buf.len() < STATUS_CODE_SIZE {
292 return Err(std::io::Error::new(
293 std::io::ErrorKind::InvalidData,
294 "terminal message missing status code",
295 ));
296 }
297 let status = buf.get_u16();
298 let body = String::from_utf8(buf.into()).map_err(|_| {
299 std::io::Error::new(std::io::ErrorKind::InvalidData, "invalid utf-8")
300 })?;
301 return Ok(TerminalMessage { status, body }.into());
302 }
303
304 let compression_bits = (flag & FLAG_COMPRESSION_MASK) >> FLAG_COMPRESSION_SHIFT;
305 let Some(compression) = CompressionAlgorithm::from_repr(compression_bits) else {
306 return Err(std::io::Error::new(
307 std::io::ErrorKind::InvalidData,
308 "unknown compression algorithm",
309 ));
310 };
311
312 Ok(CompressedData {
313 compression,
314 reconnect_advised: (flag & FLAG_RECONNECT_ADVISED) != 0,
315 payload: buf,
316 }
317 .into())
318 }
319
320 fn payload_size(&self) -> usize {
321 match self {
322 Self::Regular(msg) => msg.payload.len(),
323 Self::Terminal(msg) => STATUS_CODE_SIZE + msg.body.len(),
324 }
325 }
326}
327
328pub fn advise_reconnect(frame: Bytes) -> Bytes {
332 if frame
333 .get(LENGTH_PREFIX_SIZE)
334 .is_none_or(|flag| flag & FLAG_TERMINAL != 0)
335 {
336 return frame;
337 }
338
339 let mut frame = frame
340 .try_into_mut()
341 .unwrap_or_else(|frame| BytesMut::from(frame.as_ref()));
342 frame[LENGTH_PREFIX_SIZE] |= FLAG_RECONNECT_ADVISED;
343 frame.freeze()
344}
345
346pub struct FramedMessageStream<S> {
347 inner: S,
348 compression: CompressionAlgorithm,
349 terminated: bool,
350}
351
352impl<S> FramedMessageStream<S> {
353 pub fn new(compression: CompressionAlgorithm, inner: S) -> Self {
354 Self {
355 inner,
356 compression,
357 terminated: false,
358 }
359 }
360}
361
362impl<S, P, E> Stream for FramedMessageStream<S>
363where
364 S: Stream<Item = Result<P, E>> + Unpin,
365 P: prost::Message,
366 E: Into<TerminalMessage>,
367{
368 type Item = std::io::Result<Bytes>;
369
370 fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
371 if self.terminated {
372 return Poll::Ready(None);
373 }
374
375 match Pin::new(&mut self.inner).poll_next(cx) {
376 Poll::Ready(Some(Ok(item))) => match SessionMessage::regular(self.compression, &item) {
377 Ok(msg) => Poll::Ready(Some(Ok(msg.encode()))),
378 Err(err) => {
379 self.terminated = true;
380 Poll::Ready(Some(Err(err)))
381 }
382 },
383 Poll::Ready(Some(Err(e))) => {
384 self.terminated = true;
385 let bytes = SessionMessage::Terminal(e.into()).encode();
386 Poll::Ready(Some(Ok(bytes)))
387 }
388 Poll::Ready(None) => {
389 self.terminated = true;
390 Poll::Ready(None)
391 }
392 Poll::Pending => Poll::Pending,
393 }
394 }
395}
396
397pub struct FrameDecoder;
398
399impl tokio_util::codec::Decoder for FrameDecoder {
400 type Item = SessionMessage;
401 type Error = std::io::Error;
402
403 fn decode(&mut self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
404 if src.len() < LENGTH_PREFIX_SIZE {
405 return Ok(None);
406 }
407
408 let length = ((src[0] as usize) << 16) | ((src[1] as usize) << 8) | (src[2] as usize);
409
410 if length > MAX_FRAME_BYTES {
411 return Err(std::io::Error::new(
412 std::io::ErrorKind::InvalidInput,
413 "frame exceeds decode limit",
414 ));
415 }
416
417 let total_size = LENGTH_PREFIX_SIZE + length;
418 if src.len() < total_size {
419 return Ok(None);
420 }
421
422 src.advance(LENGTH_PREFIX_SIZE);
423 let frame_bytes = src.split_to(length).freeze();
424 Ok(Some(SessionMessage::decode_message(frame_bytes)?))
425 }
426}
427
428#[cfg(test)]
429mod test {
430 use std::{
431 io,
432 pin::Pin,
433 task::{Context, Poll},
434 };
435
436 use bytes::BytesMut;
437 use futures::StreamExt;
438 use http::HeaderValue;
439 use proptest::{collection::vec, prelude::*};
440 use prost::Message;
441 use tokio_util::codec::Decoder;
442
443 use super::*;
444
445 #[derive(Clone, PartialEq, prost::Message)]
446 struct TestProto {
447 #[prost(bytes, tag = "1")]
448 payload: Vec<u8>,
449 }
450
451 impl TestProto {
452 fn new(payload: Vec<u8>) -> Self {
453 Self { payload }
454 }
455 }
456
457 #[derive(Debug, Clone)]
458 struct TestError {
459 status: u16,
460 body: &'static str,
461 }
462
463 impl From<TestError> for TerminalMessage {
464 fn from(val: TestError) -> Self {
465 TerminalMessage {
466 status: val.status,
467 body: val.body.to_string(),
468 }
469 }
470 }
471
472 fn decode_once(bytes: &Bytes) -> io::Result<SessionMessage> {
473 let mut decoder = FrameDecoder;
474 let mut buf = BytesMut::from(bytes.as_ref());
475 decoder
476 .decode(&mut buf)?
477 .ok_or_else(|| io::Error::new(io::ErrorKind::UnexpectedEof, "frame incomplete"))
478 }
479
480 fn compression_strategy() -> impl proptest::strategy::Strategy<Value = CompressionAlgorithm> {
481 prop_oneof![
482 Just(CompressionAlgorithm::None),
483 Just(CompressionAlgorithm::Gzip),
484 Just(CompressionAlgorithm::Zstd),
485 ]
486 }
487
488 fn chunk_bytes(data: &Bytes, pattern: &[usize]) -> Vec<Bytes> {
489 let mut chunks = Vec::new();
490 let mut offset = 0;
491 for &hint in pattern {
492 if offset >= data.len() {
493 break;
494 }
495 let remaining = data.len() - offset;
496 let take = (hint % remaining).saturating_add(1).min(remaining);
497 chunks.push(data.slice(offset..offset + take));
498 offset += take;
499 }
500 if offset < data.len() {
501 chunks.push(data.slice(offset..));
502 }
503 if chunks.is_empty() {
504 chunks.push(data.clone());
505 }
506 chunks
507 }
508
509 proptest! {
510 #[test]
511 fn regular_session_message_round_trips_proptest(
512 algo in compression_strategy(),
513 payload in vec(any::<u8>(), 0..=COMPRESSION_THRESHOLD_BYTES * 4)
514 ) {
515 let proto = TestProto::new(payload.clone());
516 let msg = SessionMessage::regular(algo, &proto).unwrap();
517 let encoded = msg.encode();
518 let decoded = decode_once(&encoded).unwrap();
519
520 prop_assert!(matches!(decoded, SessionMessage::Regular(_)));
521 let SessionMessage::Regular(data) = decoded else { unreachable!() };
522
523 let expected_compression = if algo == CompressionAlgorithm::None || proto.encoded_len() < COMPRESSION_THRESHOLD_BYTES {
524 CompressionAlgorithm::None
525 } else {
526 algo
527 };
528 let actual_compression = data.compression;
529
530 let restored = data.try_into_proto::<TestProto>().unwrap();
531 prop_assert_eq!(restored.payload, payload);
532 prop_assert_eq!(actual_compression, expected_compression);
533 }
534
535 #[test]
536 fn frame_decoder_handles_chunked_frames(
537 algo in compression_strategy(),
538 payload in vec(any::<u8>(), 0..=COMPRESSION_THRESHOLD_BYTES * 4),
539 chunk_pattern in vec(0usize..=16, 0..=16)
540 ) {
541 let proto = TestProto::new(payload);
542 let msg = SessionMessage::regular(algo, &proto).unwrap();
543 let encoded = msg.encode();
544 let expected = decode_once(&encoded).unwrap();
545
546 let chunks = chunk_bytes(&encoded, &chunk_pattern);
547 prop_assert_eq!(chunks.iter().map(|c| c.len()).sum::<usize>(), encoded.len());
548
549 let mut decoder = FrameDecoder;
550 let mut buf = BytesMut::new();
551 let mut decoded = None;
552
553 for (idx, chunk) in chunks.iter().enumerate() {
554 buf.extend_from_slice(chunk.as_ref());
555 let result = decoder.decode(&mut buf).expect("decode invocation failed");
556 if idx < chunks.len() - 1 {
557 prop_assert!(result.is_none());
558 } else {
559 let message = result.expect("final chunk should produce frame");
560 prop_assert!(buf.is_empty());
561 decoded = Some(message);
562 }
563 }
564
565 let decoded = decoded.expect("decoder never emitted frame");
566 prop_assert_eq!(decoded, expected);
567 }
568 }
569
570 #[test]
571 fn from_accept_encoding_prefers_zstd() {
572 let mut headers = http::HeaderMap::new();
573 headers.insert(
574 http::header::ACCEPT_ENCODING,
575 HeaderValue::from_static("gzip, zstd, br"),
576 );
577
578 let algo = CompressionAlgorithm::from_accept_encoding(&headers);
579 assert_eq!(algo, CompressionAlgorithm::Zstd);
580 }
581
582 #[test]
583 fn from_accept_encoding_falls_back_to_gzip() {
584 let mut headers = http::HeaderMap::new();
585 headers.insert(
586 http::header::ACCEPT_ENCODING,
587 HeaderValue::from_static("gzip;q=0.8, deflate"),
588 );
589
590 let algo = CompressionAlgorithm::from_accept_encoding(&headers);
591 assert_eq!(algo, CompressionAlgorithm::Gzip);
592 }
593
594 #[rstest::rstest]
595 #[case("zstd;q=0, gzip", CompressionAlgorithm::Gzip)]
596 #[case("zstd; q=0.0, gzip;q=0.5", CompressionAlgorithm::Gzip)]
597 #[case("gzip;q=0", CompressionAlgorithm::None)]
598 #[case("gzip;Q=0.000, zstd;q=0", CompressionAlgorithm::None)]
599 #[case("zstd;q=0.001", CompressionAlgorithm::Zstd)]
600 fn from_accept_encoding_skips_refused_codings(
601 #[case] accept_encoding: &'static str,
602 #[case] expected: CompressionAlgorithm,
603 ) {
604 let mut headers = http::HeaderMap::new();
605 headers.insert(
606 http::header::ACCEPT_ENCODING,
607 HeaderValue::from_static(accept_encoding),
608 );
609
610 let algo = CompressionAlgorithm::from_accept_encoding(&headers);
611 assert_eq!(algo, expected);
612 }
613
614 #[test]
615 fn from_accept_encoding_defaults_to_none() {
616 let headers = http::HeaderMap::new();
617 let algo = CompressionAlgorithm::from_accept_encoding(&headers);
618 assert_eq!(algo, CompressionAlgorithm::None);
619 }
620
621 #[test]
622 fn regular_session_message_round_trips() {
623 let proto = TestProto::new(vec![1, 2, 3, 4]);
624 let msg = SessionMessage::regular(CompressionAlgorithm::None, &proto).unwrap();
625 let encoded = msg.encode();
626 let decoded = decode_once(&encoded).unwrap();
627
628 match decoded {
629 SessionMessage::Regular(data) => {
630 assert_eq!(data.compression, CompressionAlgorithm::None);
631 let restored = data.try_into_proto::<TestProto>().unwrap();
632 assert_eq!(restored, proto);
633 }
634 SessionMessage::Terminal(_) => panic!("expected regular message"),
635 }
636 }
637
638 #[test]
639 fn terminal_session_message_round_trips() {
640 let terminal = TerminalMessage {
641 status: 418,
642 body: "short-circuit".to_string(),
643 };
644 let msg = SessionMessage::from(terminal.clone());
645 let encoded = msg.encode();
646 let decoded = decode_once(&encoded).unwrap();
647
648 match decoded {
649 SessionMessage::Regular(_) => panic!("expected terminal message"),
650 SessionMessage::Terminal(decoded_terminal) => {
651 assert_eq!(decoded_terminal, terminal);
652 }
653 }
654 }
655
656 #[test]
657 fn frame_decoder_waits_for_complete_frame() {
658 let proto = TestProto::new(vec![9, 9, 9]);
659 let msg = SessionMessage::regular(CompressionAlgorithm::None, &proto).unwrap();
660 let encoded = msg.encode();
661 let mut decoder = FrameDecoder;
662
663 let split_idx = encoded.len() - 1;
664 let mut buf = BytesMut::from(&encoded[..split_idx]);
665 assert!(decoder.decode(&mut buf).unwrap().is_none());
666 buf.extend_from_slice(&encoded[split_idx..]);
667 let decoded = decoder.decode(&mut buf).unwrap().unwrap();
668
669 match decoded {
670 SessionMessage::Regular(data) => {
671 let restored = data.try_into_proto::<TestProto>().unwrap();
672 assert_eq!(restored, proto);
673 }
674 SessionMessage::Terminal(_) => panic!("expected regular message"),
675 }
676 assert!(buf.is_empty());
677 }
678
679 #[test]
680 fn frame_decoder_rejects_frames_exceeding_decode_limit() {
681 let length = MAX_FRAME_BYTES + 1;
682 let prefix = [
683 ((length >> 16) & 0xFF) as u8,
684 ((length >> 8) & 0xFF) as u8,
685 (length & 0xFF) as u8,
686 ];
687 let mut buf = BytesMut::from(prefix.as_slice());
688 let mut decoder = FrameDecoder;
689 let err = decoder.decode(&mut buf).unwrap_err();
690 assert_eq!(err.kind(), std::io::ErrorKind::InvalidInput);
691 }
692
693 #[test]
694 #[should_panic(expected = "encoder limit")]
695 fn session_message_encode_rejects_frames_over_limit() {
696 let data = CompressedData {
697 compression: CompressionAlgorithm::None,
698 reconnect_advised: false,
699 payload: Bytes::from(vec![0u8; MAX_FRAME_BYTES]),
700 };
701 let msg = SessionMessage::from(data);
702 let _ = msg.encode();
703 }
704
705 #[test]
706 fn frame_decoder_rejects_unknown_compression() {
707 let mut raw = vec![0, 0, 1];
708 raw.push(0x60);
709 let mut decoder = FrameDecoder;
710 let mut buf = BytesMut::from(raw.as_slice());
711 let err = decoder.decode(&mut buf).unwrap_err();
712 assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
713 }
714
715 #[test]
716 fn frame_decoder_rejects_terminal_without_status() {
717 let mut raw = vec![0, 0, 1];
718 raw.push(FLAG_TERMINAL);
719 let mut decoder = FrameDecoder;
720 let mut buf = BytesMut::from(raw.as_slice());
721 let err = decoder.decode(&mut buf).unwrap_err();
722 assert_eq!(err.kind(), std::io::ErrorKind::InvalidData);
723 }
724
725 #[test]
726 fn frame_decoder_handles_empty_payload() {
727 let raw = vec![0, 0, 0];
728 let mut decoder = FrameDecoder;
729 let mut buf = BytesMut::from(raw.as_slice());
730 let err = decoder.decode(&mut buf).unwrap_err();
731 assert_eq!(err.kind(), std::io::ErrorKind::UnexpectedEof);
732 }
733
734 #[test]
735 fn compressed_data_round_trip_gzip() {
736 let payload = vec![42; 1_200_000];
737 let proto = TestProto::new(payload.clone());
738 let msg = SessionMessage::regular(CompressionAlgorithm::Gzip, &proto).unwrap();
739 let encoded = msg.encode();
740 let decoded = decode_once(&encoded).unwrap();
741
742 match decoded {
743 SessionMessage::Regular(data) => {
744 assert_eq!(data.compression, CompressionAlgorithm::Gzip);
745 assert!(data.payload.len() < proto.encode_to_vec().len());
746 let restored = data.try_into_proto::<TestProto>().unwrap();
747 assert_eq!(restored.payload, payload);
748 }
749 SessionMessage::Terminal(_) => panic!("expected regular message"),
750 }
751 }
752
753 #[test]
754 fn compressed_data_round_trip_zstd() {
755 let payload = vec![7; 1_100_000];
756 let proto = TestProto::new(payload.clone());
757 let msg = SessionMessage::regular(CompressionAlgorithm::Zstd, &proto).unwrap();
758 let encoded = msg.encode();
759 let decoded = decode_once(&encoded).unwrap();
760
761 match decoded {
762 SessionMessage::Regular(data) => {
763 assert_eq!(data.compression, CompressionAlgorithm::Zstd);
764 assert!(data.payload.len() < proto.encode_to_vec().len());
765 let restored = data.try_into_proto::<TestProto>().unwrap();
766 assert_eq!(restored.payload, payload);
767 }
768 SessionMessage::Terminal(_) => panic!("expected regular message"),
769 }
770 }
771
772 #[test]
773 fn decompression_rejects_payloads_exceeding_limit() {
774 let payload = vec![0; MAX_DECOMPRESSED_PAYLOAD_BYTES + 1];
775 let proto = TestProto::new(payload);
776 let encoded = proto.encode_to_vec();
777
778 for algo in [CompressionAlgorithm::Gzip, CompressionAlgorithm::Zstd] {
779 let compressed = match algo {
780 CompressionAlgorithm::Gzip => {
781 let mut out = Vec::new();
782 let mut encoder = GzEncoder::new(&mut out, Compression::default());
783 encoder.write_all(encoded.as_slice()).unwrap();
784 encoder.finish().unwrap();
785 out
786 }
787 CompressionAlgorithm::Zstd => {
788 let mut out = Vec::new();
789 zstd::stream::copy_encode(encoded.as_slice(), &mut out, 0).unwrap();
790 out
791 }
792 CompressionAlgorithm::None => unreachable!("explicitly excluded in test"),
793 };
794
795 let data = CompressedData {
796 compression: algo,
797 reconnect_advised: false,
798 payload: Bytes::from(compressed),
799 };
800 assert!(data.payload.len() <= MAX_FRAME_PAYLOAD_BYTES);
801
802 let err = data.try_into_proto::<TestProto>().expect_err("should fail");
803 assert_eq!(err.kind(), io::ErrorKind::InvalidData);
804 assert!(
805 err.to_string()
806 .contains("decompressed payload exceeds limit")
807 );
808 }
809 }
810
811 #[test]
812 fn compress_rejects_payloads_exceeding_decompressed_limit() {
813 let payload = vec![0; MAX_DECOMPRESSED_PAYLOAD_BYTES + 1];
814 let proto = TestProto::new(payload);
815
816 let err = CompressedData::compress(CompressionAlgorithm::Gzip, proto.encode_to_vec())
817 .expect_err("should fail");
818 assert_eq!(err.kind(), io::ErrorKind::InvalidInput);
819 assert!(
820 err.to_string()
821 .contains("payload exceeds decompressed limit")
822 );
823 }
824
825 #[test]
826 fn compress_allows_payload_at_exact_limit_without_encode_panic() {
827 let payload = vec![0; MAX_DECOMPRESSED_PAYLOAD_BYTES];
828 let data = CompressedData::compress(CompressionAlgorithm::None, payload).unwrap();
829 let encoded = SessionMessage::from(data).encode();
830 assert_eq!(encoded.len(), LENGTH_PREFIX_SIZE + MAX_FRAME_BYTES);
831 }
832
833 #[test]
834 fn compress_rejects_incompressible_payload_that_exceeds_frame_limit_after_compression() {
835 let mut payload = vec![0u8; MAX_DECOMPRESSED_PAYLOAD_BYTES];
836 let mut x = 0x1234_5678u32;
837 for byte in &mut payload {
838 x ^= x << 13;
839 x ^= x >> 17;
840 x ^= x << 5;
841 *byte = (x & 0xFF) as u8;
842 }
843
844 for algo in [CompressionAlgorithm::Gzip, CompressionAlgorithm::Zstd] {
845 let err = CompressedData::compress(algo, payload.clone()).expect_err("should fail");
846 assert_eq!(err.kind(), io::ErrorKind::InvalidInput);
847 assert!(
848 err.to_string()
849 .contains("compressed payload exceeds frame limit")
850 );
851 }
852 }
853
854 #[test]
855 fn framed_message_stream_yields_terminal_on_error() {
856 let proto = TestProto::new(vec![1, 2, 3]);
857 let items = vec![
858 Ok(proto.clone()),
859 Err(TestError {
860 status: 500,
861 body: "boom",
862 }),
863 Ok(proto.clone()),
864 ];
865
866 let stream = futures::stream::iter(items);
867 let framed = FramedMessageStream::new(CompressionAlgorithm::None, stream);
868 let outputs = futures::executor::block_on(async {
869 framed.collect::<Vec<std::io::Result<Bytes>>>().await
870 });
871
872 assert_eq!(outputs.len(), 2);
873
874 let first = outputs[0].as_ref().expect("first frame ok");
875 match decode_once(first).unwrap() {
876 SessionMessage::Regular(data) => {
877 let restored = data.try_into_proto::<TestProto>().unwrap();
878 assert_eq!(restored, proto);
879 }
880 SessionMessage::Terminal(_) => panic!("expected regular message"),
881 }
882
883 let second = outputs[1].as_ref().expect("second frame ok");
884 match decode_once(second).unwrap() {
885 SessionMessage::Regular(_) => panic!("expected terminal message"),
886 SessionMessage::Terminal(term) => {
887 assert_eq!(term.status, 500);
888 assert_eq!(term.body, "boom");
889 }
890 }
891 }
892
893 #[test]
894 fn framed_message_stream_stops_after_termination() {
895 let mut stream = FramedMessageStream::new(
896 CompressionAlgorithm::None,
897 futures::stream::iter(vec![
898 Ok(TestProto::new(vec![0])),
899 Err(TestError {
900 status: 400,
901 body: "bad",
902 }),
903 ]),
904 );
905
906 let mut cx = Context::from_waker(futures::task::noop_waker_ref());
907
908 match Pin::new(&mut stream).poll_next(&mut cx) {
909 Poll::Ready(Some(Ok(bytes))) => match decode_once(&bytes).unwrap() {
910 SessionMessage::Regular(_) => {}
911 SessionMessage::Terminal(_) => panic!("expected regular message"),
912 },
913 other => panic!("unexpected poll result: {other:?}"),
914 }
915
916 match Pin::new(&mut stream).poll_next(&mut cx) {
917 Poll::Ready(Some(Ok(bytes))) => match decode_once(&bytes).unwrap() {
918 SessionMessage::Terminal(term) => {
919 assert_eq!(term.status, 400);
920 assert_eq!(term.body, "bad");
921 }
922 SessionMessage::Regular(_) => panic!("expected terminal message"),
923 },
924 other => panic!("unexpected poll result: {other:?}"),
925 }
926
927 match Pin::new(&mut stream).poll_next(&mut cx) {
928 Poll::Ready(None) => {}
929 other => panic!("expected stream to terminate, got {other:?}"),
930 }
931 }
932
933 #[test]
934 fn framed_message_stream_terminates_after_encoding_error() {
935 let oversized = MAX_DECOMPRESSED_PAYLOAD_BYTES + 1;
936 let items: Vec<Result<TestProto, TestError>> = vec![
937 Ok(TestProto::new(vec![0u8; oversized])),
938 Ok(TestProto::new(vec![1u8; oversized])),
939 ];
940 let mut stream =
941 FramedMessageStream::new(CompressionAlgorithm::None, futures::stream::iter(items));
942
943 let mut cx = Context::from_waker(futures::task::noop_waker_ref());
944
945 match Pin::new(&mut stream).poll_next(&mut cx) {
946 Poll::Ready(Some(Err(err))) => {
947 assert_eq!(err.kind(), io::ErrorKind::InvalidInput);
948 assert!(
949 err.to_string()
950 .contains("payload exceeds decompressed limit")
951 );
952 }
953 other => panic!("expected encoding error, got {other:?}"),
954 }
955
956 match Pin::new(&mut stream).poll_next(&mut cx) {
957 Poll::Ready(None) => {}
958 other => panic!("expected stream to terminate after encoding error, got {other:?}"),
959 }
960 }
961}