Skip to main content

s2_api/v1/stream/
s2s.rs

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
12/*
13  REGULAR MESSAGE:
14  ┌─────────────┬────────────┬─────────────────────────────┐
15  │   LENGTH    │   FLAGS    │        PAYLOAD DATA         │
16  │  (3 bytes)  │  (1 byte)  │     (variable length)       │
17  ├─────────────┼────────────┼─────────────────────────────┤
18  │ 0x00 00 XX  │ 0 CA RXXXX │  Compressed proto message   │
19  └─────────────┴────────────┴─────────────────────────────┘
20
21  TERMINAL MESSAGE:
22  ┌─────────────┬────────────┬─────────────┬───────────────┐
23  │   LENGTH    │   FLAGS    │ STATUS CODE │   JSON BODY   │
24  │  (3 bytes)  │  (1 byte)  │  (2 bytes)  │  (variable)   │
25  ├─────────────┼────────────┼─────────────┼───────────────┤
26  │ 0x00 00 XX  │ 1 CA RXXXX │   HTTP Code │   JSON data   │
27  └─────────────┴────────────┴─────────────┴───────────────┘
28
29  LENGTH = size of (FLAGS + PAYLOAD), does NOT include length header itself
30  Implemented limit: 2 MiB (smaller than 24-bit protocol maximum)
31*/
32
33const LENGTH_PREFIX_SIZE: usize = 3;
34const STATUS_CODE_SIZE: usize = 2;
35const COMPRESSION_THRESHOLD_BYTES: usize = 1024; // 1 KiB
36const MAX_FRAME_BYTES: usize = 2 * 1024 * 1024; // 2 MiB
37
38/*
39Flag byte layout:
40  ┌───┬───┬───┬───┬───┬───┬───┬───┐
41  │ 7 │ 6 │ 5 │ 4 │ 3 │ 2 │ 1 │ 0 │  Bit positions
42  ├───┼───┴───┼───┼───┴───┴───┴───┤
43  │ T │  C C  │ R │ Reserved (0s) │  Purpose
44  └───┴───────┴───┴───────────────┘
45
46  T = Terminal flag (1 bit)
47  C = Compression (2 bits, encodes 0-3)
48  R = Reconnect advised (1 bit, set by a server that is about to terminate)
49*/
50
51const FLAG_TOTAL_SIZE: usize = 1;
52// The frame length budget includes one flag byte, so payload bytes are capped at budget - flag.
53const 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
91/// Whether an `Accept-Encoding` parameter is a `q=0` weight, which marks the coding as
92/// "not acceptable" (RFC 9110 §12.4.2).
93fn 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    /// Whether this frame carried the reconnect-advised flag.
116    ///
117    /// A server sets the flag on responses when it is about to terminate,
118    /// signaling that the client should proactively reconnect.
119    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        // Decode at most `MAX_DECOMPRESSED_PAYLOAD_BYTES + 1` bytes
172        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
328/// Set the reconnect-advised flag on an already encoded frame.
329///
330/// Terminal frames are returned unchanged.
331pub 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}