1use 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#[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 .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
109pub trait MarshallMessage: fmt::Debug {
111 type Input;
113
114 fn marshall(&self, input: Self::Input) -> Result<Message, Error>;
115}
116
117#[derive(Debug)]
119pub enum UnmarshalledMessage<T, E> {
120 Event(T),
121 Error(E),
122}
123
124pub trait UnmarshallMessage: fmt::Debug {
126 type Output;
128 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
238fn 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
258fn 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
269pub 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
277fn read_prelude_from<B: Buf>(mut buffer: B) -> Result<(u32, u32), Error> {
279 let mut crc_buffer = CrcBuf::new(&mut buffer);
280
281 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 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 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
300pub 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 let mut crc_buffer = CrcBuf::new(&mut buffer);
309 let (total_len, header_len) = read_prelude_from(&mut crc_buffer)?;
310
311 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 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 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
344pub 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 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 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#[derive(Debug)]
564pub enum DecodedFrame {
565 Incomplete,
567 Complete(Message),
569}
570
571#[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 pub fn new() -> Self {
582 Default::default()
583 }
584
585 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 fn reset(&mut self) {
607 self.prelude_read = false;
608 self.prelude = [0u8; PRELUDE_LENGTH_BYTES_USIZE];
609 }
610
611 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}