Skip to main content

aws_smithy_eventstream/
frame.rs

1/*
2 * Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
3 * SPDX-License-Identifier: Apache-2.0
4 */
5
6//! Event Stream message frame types and serialization/deserialization logic.
7
8use crate::buf::count::CountBuf;
9use crate::buf::crc::{CrcBuf, CrcBufMut};
10use crate::error::{Error, ErrorKind};
11pub use aws_smithy_types::event_stream::DeferredSignerSender;
12use aws_smithy_types::event_stream::{DeferredSignerReceiver, Header, HeaderValue, Message};
13use aws_smithy_types::str_bytes::StrBytes;
14use aws_smithy_types::DateTime;
15use bytes::{Buf, BufMut};
16use std::fmt;
17use std::mem::size_of;
18
19const PRELUDE_LENGTH_BYTES: u32 = 3 * size_of::<u32>() as u32;
20const PRELUDE_LENGTH_BYTES_USIZE: usize = PRELUDE_LENGTH_BYTES as usize;
21const MESSAGE_CRC_LENGTH_BYTES: u32 = size_of::<u32>() as u32;
22const MAX_HEADER_NAME_LEN: usize = 255;
23const MIN_HEADER_LEN: usize = 2;
24
25pub(crate) const TYPE_TRUE: u8 = 0;
26pub(crate) const TYPE_FALSE: u8 = 1;
27pub(crate) const TYPE_BYTE: u8 = 2;
28pub(crate) const TYPE_INT16: u8 = 3;
29pub(crate) const TYPE_INT32: u8 = 4;
30pub(crate) const TYPE_INT64: u8 = 5;
31pub(crate) const TYPE_BYTE_ARRAY: u8 = 6;
32pub(crate) const TYPE_STRING: u8 = 7;
33pub(crate) const TYPE_TIMESTAMP: u8 = 8;
34pub(crate) const TYPE_UUID: u8 = 9;
35
36pub use aws_smithy_types::event_stream::{SignMessage, SignMessageError};
37
38/// Deferred event stream signer to allow a signer to be wired up later.
39///
40/// HTTP request signing takes place after serialization, and the event stream
41/// message stream body is established during serialization. Since event stream
42/// signing may need context from the initial HTTP signing operation, this
43/// [`DeferredSigner`] is needed to wire up the signer later in the request lifecycle.
44///
45/// This signer basically just establishes a MPSC channel so that the sender can
46/// be placed in the request's config. Then the HTTP signer implementation can
47/// retrieve the sender from that config and send an actual signing implementation
48/// with all the context needed.
49///
50/// When an event stream implementation needs to sign a message, the first call to
51/// sign will acquire a signing implementation off of the channel and cache it
52/// for the remainder of the operation.
53#[derive(Debug)]
54pub struct DeferredSigner {
55    rx: Option<DeferredSignerReceiver>,
56    signer: Option<Box<dyn SignMessage + Send + Sync>>,
57}
58
59impl DeferredSigner {
60    pub fn new() -> (Self, DeferredSignerSender) {
61        let (rx, sender) = DeferredSignerSender::new();
62        (
63            Self {
64                rx: Some(rx),
65                signer: None,
66            },
67            sender,
68        )
69    }
70
71    fn acquire(&mut self) -> &mut (dyn SignMessage + Send + Sync) {
72        if self.signer.is_none() {
73            self.signer = Some(
74                self.rx
75                    .take()
76                    .expect("only taken once")
77                    .recv::<Box<dyn SignMessage + Send + Sync>>()
78                    // Fall back to NoOpSigner when no signer is sent (e.g., server-side
79                    // event streams or tests that don't configure signing).
80                    .unwrap_or_else(|_| Box::new(NoOpSigner {}) as _),
81            );
82        }
83        self.signer.as_mut().unwrap().as_mut()
84    }
85}
86
87impl SignMessage for DeferredSigner {
88    fn sign(&mut self, message: Message) -> Result<Message, SignMessageError> {
89        self.acquire().sign(message)
90    }
91
92    fn sign_empty(&mut self) -> Option<Result<Message, SignMessageError>> {
93        self.acquire().sign_empty()
94    }
95}
96
97#[derive(Debug)]
98pub struct NoOpSigner {}
99impl SignMessage for NoOpSigner {
100    fn sign(&mut self, message: Message) -> Result<Message, SignMessageError> {
101        Ok(message)
102    }
103
104    fn sign_empty(&mut self) -> Option<Result<Message, SignMessageError>> {
105        None
106    }
107}
108
109/// Converts a Smithy modeled Event Stream type into a [`Message`].
110pub trait MarshallMessage: fmt::Debug {
111    /// Smithy modeled input type to convert from.
112    type Input;
113
114    fn marshall(&self, input: Self::Input) -> Result<Message, Error>;
115}
116
117/// A successfully unmarshalled message that is either an `Event` or an `Error`.
118#[derive(Debug)]
119pub enum UnmarshalledMessage<T, E> {
120    Event(T),
121    Error(E),
122}
123
124/// Converts an Event Stream [`Message`] into a Smithy modeled type.
125pub trait UnmarshallMessage: fmt::Debug {
126    /// Smithy modeled type to convert into.
127    type Output;
128    /// Smithy modeled error to convert into.
129    type Error;
130
131    fn unmarshall(
132        &self,
133        message: &Message,
134    ) -> Result<UnmarshalledMessage<Self::Output, Self::Error>, Error>;
135}
136
137macro_rules! read_value {
138    ($buf:ident, $typ:ident, $size_typ:ident, $read_fn:ident) => {
139        if $buf.remaining() >= size_of::<$size_typ>() {
140            Ok(HeaderValue::$typ($buf.$read_fn()))
141        } else {
142            Err(ErrorKind::InvalidHeaderValue.into())
143        }
144    };
145}
146
147fn read_header_value_from<B: Buf>(mut buffer: B) -> Result<HeaderValue, Error> {
148    let value_type = buffer.get_u8();
149    match value_type {
150        TYPE_TRUE => Ok(HeaderValue::Bool(true)),
151        TYPE_FALSE => Ok(HeaderValue::Bool(false)),
152        TYPE_BYTE => read_value!(buffer, Byte, i8, get_i8),
153        TYPE_INT16 => read_value!(buffer, Int16, i16, get_i16),
154        TYPE_INT32 => read_value!(buffer, Int32, i32, get_i32),
155        TYPE_INT64 => read_value!(buffer, Int64, i64, get_i64),
156        TYPE_BYTE_ARRAY | TYPE_STRING => {
157            if buffer.remaining() > size_of::<u16>() {
158                let len = buffer.get_u16() as usize;
159                if buffer.remaining() < len {
160                    return Err(ErrorKind::InvalidHeaderValue.into());
161                }
162                let bytes = buffer.copy_to_bytes(len);
163                if value_type == TYPE_STRING {
164                    Ok(HeaderValue::String(
165                        bytes.try_into().map_err(|_| ErrorKind::InvalidUtf8String)?,
166                    ))
167                } else {
168                    Ok(HeaderValue::ByteArray(bytes))
169                }
170            } else {
171                Err(ErrorKind::InvalidHeaderValue.into())
172            }
173        }
174        TYPE_TIMESTAMP => {
175            if buffer.remaining() >= size_of::<i64>() {
176                let epoch_millis = buffer.get_i64();
177                Ok(HeaderValue::Timestamp(DateTime::from_millis(epoch_millis)))
178            } else {
179                Err(ErrorKind::InvalidHeaderValue.into())
180            }
181        }
182        TYPE_UUID => read_value!(buffer, Uuid, u128, get_u128),
183        _ => Err(ErrorKind::InvalidHeaderValueType(value_type).into()),
184    }
185}
186
187fn write_header_value_to<B: BufMut>(value: &HeaderValue, mut buffer: B) -> Result<(), Error> {
188    use HeaderValue::*;
189    match value {
190        Bool(val) => buffer.put_u8(if *val { TYPE_TRUE } else { TYPE_FALSE }),
191        Byte(val) => {
192            buffer.put_u8(TYPE_BYTE);
193            buffer.put_i8(*val);
194        }
195        Int16(val) => {
196            buffer.put_u8(TYPE_INT16);
197            buffer.put_i16(*val);
198        }
199        Int32(val) => {
200            buffer.put_u8(TYPE_INT32);
201            buffer.put_i32(*val);
202        }
203        Int64(val) => {
204            buffer.put_u8(TYPE_INT64);
205            buffer.put_i64(*val);
206        }
207        ByteArray(val) => {
208            buffer.put_u8(TYPE_BYTE_ARRAY);
209            buffer.put_u16(checked(val.len(), ErrorKind::HeaderValueTooLong.into())?);
210            buffer.put_slice(&val[..]);
211        }
212        String(val) => {
213            buffer.put_u8(TYPE_STRING);
214            buffer.put_u16(checked(
215                val.as_bytes().len(),
216                ErrorKind::HeaderValueTooLong.into(),
217            )?);
218            buffer.put_slice(&val.as_bytes()[..]);
219        }
220        Timestamp(time) => {
221            buffer.put_u8(TYPE_TIMESTAMP);
222            buffer.put_i64(
223                time.to_millis()
224                    .map_err(|_| ErrorKind::TimestampValueTooLarge(*time))?,
225            );
226        }
227        Uuid(val) => {
228            buffer.put_u8(TYPE_UUID);
229            buffer.put_u128(*val);
230        }
231        _ => {
232            panic!("matched on unexpected variant in `aws_smithy_types::event_stream::HeaderValue`")
233        }
234    }
235    Ok(())
236}
237
238/// Reads a header from the given `buffer`.
239fn read_header_from<B: Buf>(mut buffer: B) -> Result<(Header, usize), Error> {
240    if buffer.remaining() < MIN_HEADER_LEN {
241        return Err(ErrorKind::InvalidHeadersLength.into());
242    }
243
244    let mut counting_buf = CountBuf::new(&mut buffer);
245    let name_len = counting_buf.get_u8();
246    if name_len as usize >= counting_buf.remaining() {
247        return Err(ErrorKind::InvalidHeaderNameLength.into());
248    }
249
250    let name: StrBytes = counting_buf
251        .copy_to_bytes(name_len as usize)
252        .try_into()
253        .map_err(|_| ErrorKind::InvalidUtf8String)?;
254    let value = read_header_value_from(&mut counting_buf)?;
255    Ok((Header::new(name, value), counting_buf.into_count()))
256}
257
258/// Writes the header to the given `buffer`.
259fn write_header_to<B: BufMut>(header: &Header, mut buffer: B) -> Result<(), Error> {
260    if header.name().as_bytes().len() > MAX_HEADER_NAME_LEN {
261        return Err(ErrorKind::InvalidHeaderNameLength.into());
262    }
263
264    buffer.put_u8(u8::try_from(header.name().as_bytes().len()).expect("bounds check above"));
265    buffer.put_slice(&header.name().as_bytes()[..]);
266    write_header_value_to(header.value(), buffer)
267}
268
269/// Writes the given `headers` to a `buffer`.
270pub fn write_headers_to<B: BufMut>(headers: &[Header], mut buffer: B) -> Result<(), Error> {
271    for header in headers {
272        write_header_to(header, &mut buffer)?;
273    }
274    Ok(())
275}
276
277// Returns (total_len, header_len)
278fn read_prelude_from<B: Buf>(mut buffer: B) -> Result<(u32, u32), Error> {
279    let mut crc_buffer = CrcBuf::new(&mut buffer);
280
281    // If the buffer doesn't have the entire, then error
282    let total_len = crc_buffer.get_u32();
283    if crc_buffer.remaining() + size_of::<u32>() < total_len as usize {
284        return Err(ErrorKind::InvalidMessageLength.into());
285    }
286
287    // Validate the prelude
288    let header_len = crc_buffer.get_u32();
289    let (expected_crc, prelude_crc) = (crc_buffer.into_crc(), buffer.get_u32());
290    if expected_crc != prelude_crc {
291        return Err(ErrorKind::PreludeChecksumMismatch(expected_crc, prelude_crc).into());
292    }
293    // The header length can be 0 or >= 2, but must fit within the frame size
294    if header_len == 1 || header_len > max_header_len(total_len)? {
295        return Err(ErrorKind::InvalidHeadersLength.into());
296    }
297    Ok((total_len, header_len))
298}
299
300/// Reads a message from the given `buffer`. For streaming use cases, use
301/// the [`MessageFrameDecoder`] instead of this.
302pub fn read_message_from<B: Buf>(mut buffer: B) -> Result<Message, Error> {
303    if buffer.remaining() < PRELUDE_LENGTH_BYTES_USIZE {
304        return Err(ErrorKind::InvalidMessageLength.into());
305    }
306
307    // Calculate a CRC as we go and read the prelude
308    let mut crc_buffer = CrcBuf::new(&mut buffer);
309    let (total_len, header_len) = read_prelude_from(&mut crc_buffer)?;
310
311    // Verify we have the full frame before continuing
312    let remaining_len = total_len
313        .checked_sub(PRELUDE_LENGTH_BYTES)
314        .ok_or_else(|| Error::from(ErrorKind::InvalidMessageLength))?;
315    if crc_buffer.remaining() < remaining_len as usize {
316        return Err(ErrorKind::InvalidMessageLength.into());
317    }
318
319    // Read headers
320    let mut header_bytes_read = 0;
321    let mut headers = Vec::new();
322    while header_bytes_read < header_len as usize {
323        let (header, bytes_read) = read_header_from(&mut crc_buffer)?;
324        header_bytes_read += bytes_read;
325        if header_bytes_read > header_len as usize {
326            return Err(ErrorKind::InvalidHeaderValue.into());
327        }
328        headers.push(header);
329    }
330
331    // Read payload
332    let payload_len = payload_len(total_len, header_len)?;
333    let payload = crc_buffer.copy_to_bytes(payload_len as usize);
334
335    let expected_crc = crc_buffer.into_crc();
336    let message_crc = buffer.get_u32();
337    if expected_crc != message_crc {
338        return Err(ErrorKind::MessageChecksumMismatch(expected_crc, message_crc).into());
339    }
340
341    Ok(Message::new_from_parts(headers, payload))
342}
343
344/// Writes the `message` to the given `buffer`.
345pub fn write_message_to(message: &Message, buffer: &mut dyn BufMut) -> Result<(), Error> {
346    let mut headers = Vec::new();
347    for header in message.headers() {
348        write_header_to(header, &mut headers)?;
349    }
350
351    let headers_len = checked(headers.len(), ErrorKind::HeadersTooLong.into())?;
352    let payload_len = checked(message.payload().len(), ErrorKind::PayloadTooLong.into())?;
353    let message_len = [
354        PRELUDE_LENGTH_BYTES,
355        headers_len,
356        payload_len,
357        MESSAGE_CRC_LENGTH_BYTES,
358    ]
359    .iter()
360    .try_fold(0u32, |acc, v| {
361        acc.checked_add(*v)
362            .ok_or_else(|| Error::from(ErrorKind::MessageTooLong))
363    })?;
364
365    let mut crc_buffer = CrcBufMut::new(buffer);
366    crc_buffer.put_u32(message_len);
367    crc_buffer.put_u32(headers_len);
368    crc_buffer.put_crc();
369    crc_buffer.put(&headers[..]);
370    crc_buffer.put(&message.payload()[..]);
371    crc_buffer.put_crc();
372    Ok(())
373}
374
375fn checked<T: TryFrom<U>, U>(from: U, err: Error) -> Result<T, Error> {
376    T::try_from(from).map_err(|_| err)
377}
378
379fn max_header_len(total_len: u32) -> Result<u32, Error> {
380    total_len
381        .checked_sub(PRELUDE_LENGTH_BYTES + MESSAGE_CRC_LENGTH_BYTES)
382        .ok_or_else(|| Error::from(ErrorKind::InvalidMessageLength))
383}
384
385fn payload_len(total_len: u32, header_len: u32) -> Result<u32, Error> {
386    total_len
387        .checked_sub(
388            header_len
389                .checked_add(PRELUDE_LENGTH_BYTES + MESSAGE_CRC_LENGTH_BYTES)
390                .ok_or_else(|| Error::from(ErrorKind::InvalidHeadersLength))?,
391        )
392        .ok_or_else(|| Error::from(ErrorKind::InvalidMessageLength))
393}
394
395#[cfg(test)]
396mod message_tests {
397    use super::read_message_from;
398    use crate::error::ErrorKind;
399    use crate::frame::{write_message_to, Header, HeaderValue, Message};
400    use aws_smithy_types::DateTime;
401    use bytes::Bytes;
402
403    macro_rules! read_message_expect_err {
404        ($bytes:expr, $err:pat) => {
405            let result = read_message_from(&mut Bytes::from_static($bytes));
406            let result = result.as_ref();
407            assert!(result.is_err(), "Expected error, got {:?}", result);
408            assert!(
409                matches!(result.err().unwrap().kind(), $err),
410                "Expected {}, got {:?}",
411                stringify!($err),
412                result
413            );
414        };
415    }
416
417    #[test]
418    fn invalid_messages() {
419        read_message_expect_err!(
420            include_bytes!("../test_data/invalid_header_string_value_length"),
421            ErrorKind::InvalidHeaderValue
422        );
423        read_message_expect_err!(
424            include_bytes!("../test_data/invalid_header_string_length_cut_off"),
425            ErrorKind::InvalidHeaderValue
426        );
427        read_message_expect_err!(
428            include_bytes!("../test_data/invalid_header_value_type"),
429            ErrorKind::InvalidHeaderValueType(0x60)
430        );
431        read_message_expect_err!(
432            include_bytes!("../test_data/invalid_header_name_length"),
433            ErrorKind::InvalidHeaderNameLength
434        );
435        read_message_expect_err!(
436            include_bytes!("../test_data/invalid_headers_length"),
437            ErrorKind::InvalidHeadersLength
438        );
439        read_message_expect_err!(
440            include_bytes!("../test_data/invalid_prelude_checksum"),
441            ErrorKind::PreludeChecksumMismatch(0x8BB495FB, 0xDEADBEEF)
442        );
443        read_message_expect_err!(
444            include_bytes!("../test_data/invalid_message_checksum"),
445            ErrorKind::MessageChecksumMismatch(0x01a05860, 0xDEADBEEF)
446        );
447        read_message_expect_err!(
448            include_bytes!("../test_data/invalid_header_name_length_too_long"),
449            ErrorKind::InvalidUtf8String
450        );
451    }
452
453    #[test]
454    fn read_message_no_headers() {
455        // Test message taken from the CRT:
456        // https://github.com/awslabs/aws-c-event-stream/blob/main/tests/message_deserializer_test.c
457        let data: &'static [u8] = &[
458            0x00, 0x00, 0x00, 0x1D, 0x00, 0x00, 0x00, 0x00, 0xfd, 0x52, 0x8c, 0x5a, 0x7b, 0x27,
459            0x66, 0x6f, 0x6f, 0x27, 0x3a, 0x27, 0x62, 0x61, 0x72, 0x27, 0x7d, 0xc3, 0x65, 0x39,
460            0x36,
461        ];
462
463        let result = read_message_from(&mut Bytes::from_static(data)).unwrap();
464        assert_eq!(result.headers(), Vec::new());
465
466        let expected_payload = b"{'foo':'bar'}";
467        assert_eq!(expected_payload, result.payload().as_ref());
468    }
469
470    #[test]
471    fn read_message_one_header() {
472        // Test message taken from the CRT:
473        // https://github.com/awslabs/aws-c-event-stream/blob/main/tests/message_deserializer_test.c
474        let data: &'static [u8] = &[
475            0x00, 0x00, 0x00, 0x3D, 0x00, 0x00, 0x00, 0x20, 0x07, 0xFD, 0x83, 0x96, 0x0C, b'c',
476            b'o', b'n', b't', b'e', b'n', b't', b'-', b't', b'y', b'p', b'e', 0x07, 0x00, 0x10,
477            b'a', b'p', b'p', b'l', b'i', b'c', b'a', b't', b'i', b'o', b'n', b'/', b'j', b's',
478            b'o', b'n', 0x7b, 0x27, 0x66, 0x6f, 0x6f, 0x27, 0x3a, 0x27, 0x62, 0x61, 0x72, 0x27,
479            0x7d, 0x8D, 0x9C, 0x08, 0xB1,
480        ];
481
482        let result = read_message_from(&mut Bytes::from_static(data)).unwrap();
483        assert_eq!(
484            result.headers(),
485            vec![Header::new(
486                "content-type",
487                HeaderValue::String("application/json".into())
488            )]
489        );
490
491        let expected_payload = b"{'foo':'bar'}";
492        assert_eq!(expected_payload, result.payload().as_ref());
493    }
494
495    #[test]
496    fn read_all_headers_and_payload() {
497        let message = include_bytes!("../test_data/valid_with_all_headers_and_payload");
498        let result = read_message_from(&mut Bytes::from_static(message)).unwrap();
499        assert_eq!(
500            result.headers(),
501            vec![
502                Header::new("true", HeaderValue::Bool(true)),
503                Header::new("false", HeaderValue::Bool(false)),
504                Header::new("byte", HeaderValue::Byte(50)),
505                Header::new("short", HeaderValue::Int16(20_000)),
506                Header::new("int", HeaderValue::Int32(500_000)),
507                Header::new("long", HeaderValue::Int64(50_000_000_000)),
508                Header::new(
509                    "bytes",
510                    HeaderValue::ByteArray(Bytes::from(&b"some bytes"[..]))
511                ),
512                Header::new("str", HeaderValue::String("some str".into())),
513                Header::new(
514                    "time",
515                    HeaderValue::Timestamp(DateTime::from_secs(5_000_000))
516                ),
517                Header::new(
518                    "uuid",
519                    HeaderValue::Uuid(0xb79bc914_de21_4e13_b8b2_bc47e85b7f0b)
520                ),
521            ]
522        );
523
524        assert_eq!(b"some payload", result.payload().as_ref());
525    }
526
527    #[test]
528    fn round_trip_all_headers_payload() {
529        let message = Message::new(&b"some payload"[..])
530            .add_header(Header::new("true", HeaderValue::Bool(true)))
531            .add_header(Header::new("false", HeaderValue::Bool(false)))
532            .add_header(Header::new("byte", HeaderValue::Byte(50)))
533            .add_header(Header::new("short", HeaderValue::Int16(20_000)))
534            .add_header(Header::new("int", HeaderValue::Int32(500_000)))
535            .add_header(Header::new("long", HeaderValue::Int64(50_000_000_000)))
536            .add_header(Header::new(
537                "bytes",
538                HeaderValue::ByteArray((&b"some bytes"[..]).into()),
539            ))
540            .add_header(Header::new("str", HeaderValue::String("some str".into())))
541            .add_header(Header::new(
542                "time",
543                HeaderValue::Timestamp(DateTime::from_secs(5_000_000)),
544            ))
545            .add_header(Header::new(
546                "uuid",
547                HeaderValue::Uuid(0xb79bc914_de21_4e13_b8b2_bc47e85b7f0b),
548            ));
549
550        let mut actual = Vec::new();
551        write_message_to(&message, &mut actual).unwrap();
552
553        let expected = include_bytes!("../test_data/valid_with_all_headers_and_payload").to_vec();
554        assert_eq!(expected, actual);
555
556        let result = read_message_from(&mut Bytes::from(actual)).unwrap();
557        assert_eq!(message.headers(), result.headers());
558        assert_eq!(message.payload().as_ref(), result.payload().as_ref());
559    }
560}
561
562/// Return value from [`MessageFrameDecoder`].
563#[derive(Debug)]
564pub enum DecodedFrame {
565    /// There wasn't enough data in the buffer to decode a full message.
566    Incomplete,
567    /// There was enough data in the buffer to decode.
568    Complete(Message),
569}
570
571/// Streaming decoder for decoding a [`Message`] from a stream.
572#[non_exhaustive]
573#[derive(Default, Debug)]
574pub struct MessageFrameDecoder {
575    prelude: [u8; PRELUDE_LENGTH_BYTES_USIZE],
576    prelude_read: bool,
577}
578
579impl MessageFrameDecoder {
580    /// Returns a new `MessageFrameDecoder`.
581    pub fn new() -> Self {
582        Default::default()
583    }
584
585    /// Determines if the `buffer` has enough data in it to read a full frame.
586    /// Returns `Ok(None)` if there's not enough data, or `Some(remaining)` where
587    /// `remaining` is the number of bytes after the prelude that belong to the
588    /// message that's in the buffer.
589    fn remaining_bytes_if_frame_available<B: Buf>(
590        &self,
591        buffer: &B,
592    ) -> Result<Option<usize>, Error> {
593        if self.prelude_read {
594            let remaining_len = (&self.prelude[..])
595                .get_u32()
596                .checked_sub(PRELUDE_LENGTH_BYTES)
597                .ok_or_else(|| Error::from(ErrorKind::InvalidMessageLength))?;
598            if buffer.remaining() >= remaining_len as usize {
599                return Ok(Some(remaining_len as usize));
600            }
601        }
602        Ok(None)
603    }
604
605    /// Resets the decoder.
606    fn reset(&mut self) {
607        self.prelude_read = false;
608        self.prelude = [0u8; PRELUDE_LENGTH_BYTES_USIZE];
609    }
610
611    /// Attempts to decode a [`Message`] from the given `buffer`. This function expects
612    /// to be called over and over again with more data in the buffer each time its called.
613    /// When there's not enough data to decode a message, it returns `Ok(None)`.
614    ///
615    /// Once there is enough data to read a message prelude, then it will mutate the `Buf`
616    /// position. The state from the reading of the prelude is stored in the decoder so that
617    /// the next call will be able to decode the entire message, even though the prelude
618    /// is no longer available in the `Buf`.
619    pub fn decode_frame<B: Buf>(&mut self, mut buffer: B) -> Result<DecodedFrame, Error> {
620        if !self.prelude_read && buffer.remaining() >= PRELUDE_LENGTH_BYTES_USIZE {
621            buffer.copy_to_slice(&mut self.prelude);
622            self.prelude_read = true;
623        }
624
625        if let Some(remaining_len) = self.remaining_bytes_if_frame_available(&buffer)? {
626            let mut message_buf = (&self.prelude[..]).chain(buffer.take(remaining_len));
627            let result = read_message_from(&mut message_buf).map(DecodedFrame::Complete);
628            self.reset();
629            return result;
630        }
631
632        Ok(DecodedFrame::Incomplete)
633    }
634}
635
636#[cfg(test)]
637mod message_frame_decoder_tests {
638    use super::{DecodedFrame, MessageFrameDecoder};
639    use crate::frame::read_message_from;
640    use bytes::Bytes;
641    use bytes_utils::SegmentedBuf;
642
643    #[test]
644    fn single_streaming_message() {
645        let message = include_bytes!("../test_data/valid_with_all_headers_and_payload");
646
647        let mut decoder = MessageFrameDecoder::new();
648        let mut segmented = SegmentedBuf::new();
649        for i in 0..(message.len() - 1) {
650            segmented.push(&message[i..(i + 1)]);
651            if let DecodedFrame::Complete(_) = decoder.decode_frame(&mut segmented).unwrap() {
652                panic!("incomplete frame shouldn't result in message");
653            }
654        }
655
656        segmented.push(&message[(message.len() - 1)..]);
657        match decoder.decode_frame(&mut segmented).unwrap() {
658            DecodedFrame::Incomplete => panic!("frame should be complete now"),
659            DecodedFrame::Complete(actual) => {
660                let expected = read_message_from(&mut Bytes::from_static(message)).unwrap();
661                assert_eq!(expected, actual);
662            }
663        }
664    }
665
666    fn multiple_streaming_messages_chunk_size(chunk_size: usize) {
667        let message1 = include_bytes!("../test_data/valid_with_all_headers_and_payload");
668        let message2 = include_bytes!("../test_data/valid_empty_payload");
669        let message3 = include_bytes!("../test_data/valid_no_headers");
670        let mut repeated = message1.to_vec();
671        repeated.extend_from_slice(message2);
672        repeated.extend_from_slice(message3);
673
674        let mut decoder = MessageFrameDecoder::new();
675        let mut segmented = SegmentedBuf::new();
676        let mut decoded = Vec::new();
677        for window in repeated.chunks(chunk_size) {
678            segmented.push(window);
679            match dbg!(decoder.decode_frame(&mut segmented)).unwrap() {
680                DecodedFrame::Incomplete => {}
681                DecodedFrame::Complete(message) => {
682                    decoded.push(message);
683                }
684            }
685        }
686
687        let expected1 = read_message_from(&mut Bytes::from_static(message1)).unwrap();
688        let expected2 = read_message_from(&mut Bytes::from_static(message2)).unwrap();
689        let expected3 = read_message_from(&mut Bytes::from_static(message3)).unwrap();
690        assert_eq!(3, decoded.len());
691        assert_eq!(expected1, decoded[0]);
692        assert_eq!(expected2, decoded[1]);
693        assert_eq!(expected3, decoded[2]);
694    }
695
696    #[test]
697    fn multiple_streaming_messages() {
698        for chunk_size in 1..=11 {
699            println!("chunk size: {chunk_size}");
700            multiple_streaming_messages_chunk_size(chunk_size);
701        }
702    }
703}
704
705#[cfg(test)]
706mod deferred_signer_tests {
707    use crate::frame::{DeferredSigner, Header, HeaderValue, Message, SignMessage};
708    use bytes::Bytes;
709
710    fn check_send_sync<T: Send + Sync>(value: T) -> T {
711        value
712    }
713
714    #[test]
715    fn deferred_signer() {
716        #[derive(Default, Debug)]
717        struct TestSigner {
718            call_num: i32,
719        }
720        impl SignMessage for TestSigner {
721            fn sign(
722                &mut self,
723                message: Message,
724            ) -> Result<Message, crate::frame::SignMessageError> {
725                self.call_num += 1;
726                Ok(message.add_header(Header::new("call_num", HeaderValue::Int32(self.call_num))))
727            }
728
729            fn sign_empty(&mut self) -> Option<Result<Message, crate::frame::SignMessageError>> {
730                None
731            }
732        }
733
734        let (mut signer, sender) = check_send_sync(DeferredSigner::new());
735
736        sender
737            .send(Box::<TestSigner>::default() as Box<dyn SignMessage + Send + Sync>)
738            .expect("success");
739
740        let message = signer.sign(Message::new(Bytes::new())).expect("success");
741        assert_eq!(1, message.headers()[0].value().as_int32().unwrap());
742
743        let message = signer.sign(Message::new(Bytes::new())).expect("success");
744        assert_eq!(2, message.headers()[0].value().as_int32().unwrap());
745
746        assert!(signer.sign_empty().is_none());
747    }
748
749    #[test]
750    fn deferred_signer_defaults_to_noop_signer() {
751        let (mut signer, _sender) = DeferredSigner::new();
752        assert_eq!(
753            Message::new(Bytes::new()),
754            signer.sign(Message::new(Bytes::new())).unwrap()
755        );
756        assert!(signer.sign_empty().is_none());
757    }
758}