1use alloc::string::String;
27use alloc::vec::Vec;
28
29use super::key_material::KeyMaterial;
30use super::{Error, Result, be16, be32, put_be16, put_be32};
31
32pub const HANDSHAKE_CIF_FIXED_LEN: usize = 48;
35
36pub const ENCRYPTION_FIELD_NONE: u16 = 0;
38pub const ENCRYPTION_FIELD_AES_128: u16 = 2;
40pub const ENCRYPTION_FIELD_AES_192: u16 = 3;
42pub const ENCRYPTION_FIELD_AES_256: u16 = 4;
44
45#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
48#[cfg_attr(feature = "serde", derive(serde::Serialize))]
49#[non_exhaustive]
50pub enum EncryptionField {
51 NoEncryption,
53 Aes128,
55 Aes192,
57 Aes256,
59 Reserved(u16),
61}
62
63impl EncryptionField {
64 pub fn from_bits(v: u16) -> Self {
66 match v {
67 ENCRYPTION_FIELD_NONE => EncryptionField::NoEncryption,
68 ENCRYPTION_FIELD_AES_128 => EncryptionField::Aes128,
69 ENCRYPTION_FIELD_AES_192 => EncryptionField::Aes192,
70 ENCRYPTION_FIELD_AES_256 => EncryptionField::Aes256,
71 other => EncryptionField::Reserved(other),
72 }
73 }
74
75 pub fn to_bits(self) -> u16 {
77 match self {
78 EncryptionField::NoEncryption => ENCRYPTION_FIELD_NONE,
79 EncryptionField::Aes128 => ENCRYPTION_FIELD_AES_128,
80 EncryptionField::Aes192 => ENCRYPTION_FIELD_AES_192,
81 EncryptionField::Aes256 => ENCRYPTION_FIELD_AES_256,
82 EncryptionField::Reserved(v) => v,
83 }
84 }
85
86 pub fn name(&self) -> &'static str {
88 match self {
89 EncryptionField::NoEncryption => "no encryption advertised",
90 EncryptionField::Aes128 => "AES-128",
91 EncryptionField::Aes192 => "AES-192",
92 EncryptionField::Aes256 => "AES-256",
93 EncryptionField::Reserved(_) => "reserved",
94 }
95 }
96}
97
98broadcast_common::impl_spec_display!(EncryptionField, Reserved);
99
100pub const HANDSHAKE_TYPE_DONE: u32 = 0xFFFF_FFFD;
102pub const HANDSHAKE_TYPE_AGREEMENT: u32 = 0xFFFF_FFFE;
104pub const HANDSHAKE_TYPE_CONCLUSION: u32 = 0xFFFF_FFFF;
106pub const HANDSHAKE_TYPE_WAVEHAND: u32 = 0x0000_0000;
108pub const HANDSHAKE_TYPE_INDUCTION: u32 = 0x0000_0001;
110
111#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
113#[cfg_attr(feature = "serde", derive(serde::Serialize))]
114#[non_exhaustive]
115pub enum HandshakeType {
116 Done,
118 Agreement,
120 Conclusion,
122 Wavehand,
124 Induction,
126 Reserved(u32),
129}
130
131impl HandshakeType {
132 pub fn from_bits(v: u32) -> Self {
134 match v {
135 HANDSHAKE_TYPE_DONE => HandshakeType::Done,
136 HANDSHAKE_TYPE_AGREEMENT => HandshakeType::Agreement,
137 HANDSHAKE_TYPE_CONCLUSION => HandshakeType::Conclusion,
138 HANDSHAKE_TYPE_WAVEHAND => HandshakeType::Wavehand,
139 HANDSHAKE_TYPE_INDUCTION => HandshakeType::Induction,
140 other => HandshakeType::Reserved(other),
141 }
142 }
143
144 pub fn to_bits(self) -> u32 {
146 match self {
147 HandshakeType::Done => HANDSHAKE_TYPE_DONE,
148 HandshakeType::Agreement => HANDSHAKE_TYPE_AGREEMENT,
149 HandshakeType::Conclusion => HANDSHAKE_TYPE_CONCLUSION,
150 HandshakeType::Wavehand => HANDSHAKE_TYPE_WAVEHAND,
151 HandshakeType::Induction => HANDSHAKE_TYPE_INDUCTION,
152 HandshakeType::Reserved(v) => v,
153 }
154 }
155
156 pub fn name(&self) -> &'static str {
158 match self {
159 HandshakeType::Done => "DONE",
160 HandshakeType::Agreement => "AGREEMENT",
161 HandshakeType::Conclusion => "CONCLUSION",
162 HandshakeType::Wavehand => "WAVEHAND",
163 HandshakeType::Induction => "INDUCTION",
164 HandshakeType::Reserved(_) => "reserved",
165 }
166 }
167}
168
169broadcast_common::impl_spec_display!(HandshakeType, Reserved);
170
171pub const HS_EXT_FLAG_HSREQ: u16 = 0x0001;
175pub const HS_EXT_FLAG_KMREQ: u16 = 0x0002;
177pub const HS_EXT_FLAG_CONFIG: u16 = 0x0004;
179
180#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
183#[cfg_attr(feature = "serde", derive(serde::Serialize))]
184pub struct HandshakeExtensionFlags(pub u16);
185
186impl HandshakeExtensionFlags {
187 pub fn hsreq(self) -> bool {
189 self.0 & HS_EXT_FLAG_HSREQ != 0
190 }
191 pub fn kmreq(self) -> bool {
193 self.0 & HS_EXT_FLAG_KMREQ != 0
194 }
195 pub fn config(self) -> bool {
197 self.0 & HS_EXT_FLAG_CONFIG != 0
198 }
199}
200
201pub const EXT_TYPE_HSREQ: u16 = 1;
203pub const EXT_TYPE_HSRSP: u16 = 2;
205pub const EXT_TYPE_KMREQ: u16 = 3;
208pub const EXT_TYPE_KMRSP: u16 = 4;
210pub const EXT_TYPE_SID: u16 = 5;
212pub const EXT_TYPE_CONGESTION: u16 = 6;
214pub const EXT_TYPE_FILTER: u16 = 7;
216pub const EXT_TYPE_GROUP: u16 = 8;
218
219#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
222#[cfg_attr(feature = "serde", derive(serde::Serialize))]
223#[non_exhaustive]
224pub enum ExtensionType {
225 HsReq,
227 HsRsp,
229 KmReq,
231 KmRsp,
233 Sid,
235 Congestion,
237 Filter,
239 Group,
241 Reserved(u16),
243}
244
245impl ExtensionType {
246 pub fn from_bits(v: u16) -> Self {
248 match v {
249 EXT_TYPE_HSREQ => ExtensionType::HsReq,
250 EXT_TYPE_HSRSP => ExtensionType::HsRsp,
251 EXT_TYPE_KMREQ => ExtensionType::KmReq,
252 EXT_TYPE_KMRSP => ExtensionType::KmRsp,
253 EXT_TYPE_SID => ExtensionType::Sid,
254 EXT_TYPE_CONGESTION => ExtensionType::Congestion,
255 EXT_TYPE_FILTER => ExtensionType::Filter,
256 EXT_TYPE_GROUP => ExtensionType::Group,
257 other => ExtensionType::Reserved(other),
258 }
259 }
260
261 pub fn to_bits(self) -> u16 {
263 match self {
264 ExtensionType::HsReq => EXT_TYPE_HSREQ,
265 ExtensionType::HsRsp => EXT_TYPE_HSRSP,
266 ExtensionType::KmReq => EXT_TYPE_KMREQ,
267 ExtensionType::KmRsp => EXT_TYPE_KMRSP,
268 ExtensionType::Sid => EXT_TYPE_SID,
269 ExtensionType::Congestion => EXT_TYPE_CONGESTION,
270 ExtensionType::Filter => EXT_TYPE_FILTER,
271 ExtensionType::Group => EXT_TYPE_GROUP,
272 ExtensionType::Reserved(v) => v,
273 }
274 }
275
276 pub fn name(&self) -> &'static str {
278 match self {
279 ExtensionType::HsReq => "SRT_CMD_HSREQ",
280 ExtensionType::HsRsp => "SRT_CMD_HSRSP",
281 ExtensionType::KmReq => "SRT_CMD_KMREQ",
282 ExtensionType::KmRsp => "SRT_CMD_KMRSP",
283 ExtensionType::Sid => "SRT_CMD_SID",
284 ExtensionType::Congestion => "SRT_CMD_CONGESTION",
285 ExtensionType::Filter => "SRT_CMD_FILTER",
286 ExtensionType::Group => "SRT_CMD_GROUP",
287 ExtensionType::Reserved(_) => "reserved",
288 }
289 }
290}
291
292broadcast_common::impl_spec_display!(ExtensionType, Reserved);
293
294pub const HS_MSG_FLAG_TSBPDSND: u32 = 0x0000_0001;
297pub const HS_MSG_FLAG_TSBPDRCV: u32 = 0x0000_0002;
299pub const HS_MSG_FLAG_CRYPT: u32 = 0x0000_0004;
301pub const HS_MSG_FLAG_TLPKTDROP: u32 = 0x0000_0008;
303pub const HS_MSG_FLAG_PERIODICNAK: u32 = 0x0000_0010;
305pub const HS_MSG_FLAG_REXMITFLG: u32 = 0x0000_0020;
307pub const HS_MSG_FLAG_STREAM: u32 = 0x0000_0040;
309pub const HS_MSG_FLAG_PACKET_FILTER: u32 = 0x0000_0080;
311
312#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
314#[cfg_attr(feature = "serde", derive(serde::Serialize))]
315pub struct HandshakeExtensionMessageFlags(pub u32);
316
317impl HandshakeExtensionMessageFlags {
318 pub fn tsbpdsnd(self) -> bool {
320 self.0 & HS_MSG_FLAG_TSBPDSND != 0
321 }
322 pub fn tsbpdrcv(self) -> bool {
324 self.0 & HS_MSG_FLAG_TSBPDRCV != 0
325 }
326 pub fn crypt(self) -> bool {
328 self.0 & HS_MSG_FLAG_CRYPT != 0
329 }
330 pub fn tlpktdrop(self) -> bool {
332 self.0 & HS_MSG_FLAG_TLPKTDROP != 0
333 }
334 pub fn periodicnak(self) -> bool {
336 self.0 & HS_MSG_FLAG_PERIODICNAK != 0
337 }
338 pub fn rexmitflg(self) -> bool {
340 self.0 & HS_MSG_FLAG_REXMITFLG != 0
341 }
342 pub fn stream(self) -> bool {
344 self.0 & HS_MSG_FLAG_STREAM != 0
345 }
346 pub fn packet_filter(self) -> bool {
348 self.0 & HS_MSG_FLAG_PACKET_FILTER != 0
349 }
350}
351
352#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
355#[cfg_attr(feature = "serde", derive(serde::Serialize))]
356pub struct HsExtMessage {
357 pub srt_version: u32,
359 pub srt_flags: HandshakeExtensionMessageFlags,
361 pub receiver_tsbpd_delay_ms: u16,
363 pub sender_tsbpd_delay_ms: u16,
365}
366
367pub const HS_EXT_MESSAGE_LEN: usize = 12;
369
370impl HsExtMessage {
371 pub fn parse(bytes: &[u8]) -> Result<Self> {
374 if bytes.len() != HS_EXT_MESSAGE_LEN {
375 return Err(Error::BufferTooShort {
376 need: HS_EXT_MESSAGE_LEN,
377 have: bytes.len(),
378 what: "handshake extension message",
379 });
380 }
381 Ok(HsExtMessage {
382 srt_version: be32(bytes, 0),
383 srt_flags: HandshakeExtensionMessageFlags(be32(bytes, 4)),
384 receiver_tsbpd_delay_ms: be16(bytes, 8),
385 sender_tsbpd_delay_ms: be16(bytes, 10),
386 })
387 }
388
389 pub fn to_bytes(&self) -> [u8; HS_EXT_MESSAGE_LEN] {
391 let mut buf = [0u8; HS_EXT_MESSAGE_LEN];
392 put_be32(&mut buf, 0, self.srt_version);
393 put_be32(&mut buf, 4, self.srt_flags.0);
394 put_be16(&mut buf, 8, self.receiver_tsbpd_delay_ms);
395 put_be16(&mut buf, 10, self.sender_tsbpd_delay_ms);
396 buf
397 }
398}
399
400pub const GROUP_TYPE_UNDEFINED: u8 = 0;
403pub const GROUP_TYPE_BROADCAST: u8 = 1;
405pub const GROUP_TYPE_MAIN_BACKUP: u8 = 2;
407pub const GROUP_TYPE_BALANCING: u8 = 3;
409pub const GROUP_TYPE_MULTICAST: u8 = 4;
411
412#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
415#[cfg_attr(feature = "serde", derive(serde::Serialize))]
416#[non_exhaustive]
417pub enum GroupType {
418 Undefined,
420 Broadcast,
422 MainBackup,
424 Balancing,
426 Multicast,
428 Reserved(u8),
430}
431
432impl GroupType {
433 pub fn from_bits(v: u8) -> Self {
435 match v {
436 GROUP_TYPE_UNDEFINED => GroupType::Undefined,
437 GROUP_TYPE_BROADCAST => GroupType::Broadcast,
438 GROUP_TYPE_MAIN_BACKUP => GroupType::MainBackup,
439 GROUP_TYPE_BALANCING => GroupType::Balancing,
440 GROUP_TYPE_MULTICAST => GroupType::Multicast,
441 other => GroupType::Reserved(other),
442 }
443 }
444
445 pub fn to_bits(self) -> u8 {
447 match self {
448 GroupType::Undefined => GROUP_TYPE_UNDEFINED,
449 GroupType::Broadcast => GROUP_TYPE_BROADCAST,
450 GroupType::MainBackup => GROUP_TYPE_MAIN_BACKUP,
451 GroupType::Balancing => GROUP_TYPE_BALANCING,
452 GroupType::Multicast => GROUP_TYPE_MULTICAST,
453 GroupType::Reserved(v) => v,
454 }
455 }
456
457 pub fn name(&self) -> &'static str {
459 match self {
460 GroupType::Undefined => "undefined",
461 GroupType::Broadcast => "broadcast",
462 GroupType::MainBackup => "main/backup",
463 GroupType::Balancing => "balancing",
464 GroupType::Multicast => "multicast",
465 GroupType::Reserved(_) => "reserved",
466 }
467 }
468}
469
470broadcast_common::impl_spec_display!(GroupType, Reserved);
471
472#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
475#[cfg_attr(feature = "serde", derive(serde::Serialize))]
476pub struct GroupFlags(pub u8);
477
478impl GroupFlags {
479 const M_BIT: u8 = 0x01;
480
481 pub fn message_number_sync(self) -> bool {
484 self.0 & Self::M_BIT != 0
485 }
486}
487
488#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
491#[cfg_attr(feature = "serde", derive(serde::Serialize))]
492pub struct GroupMembershipExtension {
493 pub group_id: u32,
495 pub group_type: GroupType,
497 pub flags: GroupFlags,
499 pub weight: u16,
501}
502
503pub const GROUP_MEMBERSHIP_EXT_LEN: usize = 8;
505
506impl GroupMembershipExtension {
507 pub fn parse(bytes: &[u8]) -> Result<Self> {
509 if bytes.len() != GROUP_MEMBERSHIP_EXT_LEN {
510 return Err(Error::BufferTooShort {
511 need: GROUP_MEMBERSHIP_EXT_LEN,
512 have: bytes.len(),
513 what: "group membership extension",
514 });
515 }
516 let group_id = be32(bytes, 0);
517 let word1 = be32(bytes, 4);
518 Ok(GroupMembershipExtension {
519 group_id,
520 group_type: GroupType::from_bits((word1 >> 24) as u8),
521 flags: GroupFlags((word1 >> 16) as u8),
522 weight: (word1 & 0xFFFF) as u16,
523 })
524 }
525
526 pub fn to_bytes(&self) -> [u8; GROUP_MEMBERSHIP_EXT_LEN] {
528 let mut buf = [0u8; GROUP_MEMBERSHIP_EXT_LEN];
529 let word1 = (u32::from(self.group_type.to_bits()) << 24)
530 | (u32::from(self.flags.0) << 16)
531 | u32::from(self.weight);
532 put_be32(&mut buf, 0, self.group_id);
533 put_be32(&mut buf, 4, word1);
534 buf
535 }
536}
537
538#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
544#[cfg_attr(feature = "serde", derive(serde::Serialize))]
545pub struct HandshakeExtensionBlock<'a> {
546 pub ext_type: ExtensionType,
548 pub contents: &'a [u8],
550}
551
552impl<'a> HandshakeExtensionBlock<'a> {
553 pub fn as_hs_ext_message(&self) -> Result<HsExtMessage> {
556 HsExtMessage::parse(self.contents)
557 }
558
559 pub fn as_key_material(&self) -> Result<KeyMaterial<'a>> {
562 KeyMaterial::parse(self.contents)
563 }
564
565 pub fn as_stream_id(&self) -> Result<String> {
572 if !self.contents.len().is_multiple_of(4) {
573 return Err(Error::BufferTooShort {
574 need: self.contents.len().div_ceil(4) * 4,
575 have: self.contents.len(),
576 what: "stream ID extension (not a whole number of 4-byte words)",
577 });
578 }
579 let mut bytes = Vec::with_capacity(self.contents.len());
580 for word in self.contents.chunks_exact(4) {
581 bytes.extend(word.iter().rev());
582 }
583 while bytes.last() == Some(&0u8) {
584 bytes.pop();
585 }
586 String::from_utf8(bytes).map_err(|_| Error::InvalidStreamIdUtf8)
587 }
588
589 pub fn as_group_membership(&self) -> Result<GroupMembershipExtension> {
592 GroupMembershipExtension::parse(self.contents)
593 }
594}
595
596pub fn encode_stream_id(id: &str) -> Vec<u8> {
600 let mut bytes = Vec::from(id.as_bytes());
601 while bytes.len() % 4 != 0 {
602 bytes.push(0);
603 }
604 let mut out = Vec::with_capacity(bytes.len());
605 for word in bytes.chunks_exact(4) {
606 out.extend(word.iter().rev());
607 }
608 out
609}
610
611pub fn build_extension_block(ext_type: ExtensionType, contents: &[u8]) -> Result<Vec<u8>> {
619 if !contents.len().is_multiple_of(4) {
620 return Err(Error::InvalidField {
621 what: "Extension Contents",
622 reason: "length must be a whole number of 4-byte words",
623 });
624 }
625 let words = contents.len() / 4;
626 let words_u16 = u16::try_from(words).map_err(|_| Error::FieldTooWide {
627 what: "Extension Length",
628 value: words as u64,
629 bits: 16,
630 })?;
631 let mut out = Vec::with_capacity(4 + contents.len());
632 out.extend_from_slice(&ext_type.to_bits().to_be_bytes());
633 out.extend_from_slice(&words_u16.to_be_bytes());
634 out.extend_from_slice(contents);
635 Ok(out)
636}
637
638#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
641#[cfg_attr(feature = "serde", derive(serde::Serialize))]
642pub struct HandshakeExtensions<'a>(pub &'a [u8]);
643
644impl<'a> HandshakeExtensions<'a> {
645 pub fn iter(&self) -> HandshakeExtensionIter<'a> {
647 HandshakeExtensionIter { rest: self.0 }
648 }
649}
650
651#[derive(Debug, Clone)]
653pub struct HandshakeExtensionIter<'a> {
654 rest: &'a [u8],
655}
656
657impl<'a> Iterator for HandshakeExtensionIter<'a> {
658 type Item = Result<HandshakeExtensionBlock<'a>>;
659
660 fn next(&mut self) -> Option<Self::Item> {
661 if self.rest.is_empty() {
662 return None;
663 }
664 if self.rest.len() < 4 {
665 self.rest = &[];
666 return Some(Err(Error::BufferTooShort {
667 need: 4,
668 have: self.rest.len(),
669 what: "handshake extension block header",
670 }));
671 }
672 let ext_type_bits = be16(self.rest, 0);
673 let ext_len_words = be16(self.rest, 2);
674 let ext_len_bytes = usize::from(ext_len_words) * 4;
675 if self.rest.len() < 4 + ext_len_bytes {
676 let remaining = self.rest.len() - 4;
677 self.rest = &[];
678 return Some(Err(Error::ExtensionOverrun {
679 declared: ext_len_words,
680 remaining,
681 }));
682 }
683 let contents = &self.rest[4..4 + ext_len_bytes];
684 self.rest = &self.rest[4 + ext_len_bytes..];
685 Some(Ok(HandshakeExtensionBlock {
686 ext_type: ExtensionType::from_bits(ext_type_bits),
687 contents,
688 }))
689 }
690}
691
692#[derive(Debug, Clone, PartialEq, Eq, Hash)]
694#[cfg_attr(feature = "serde", derive(serde::Serialize))]
695pub struct HandshakePacket<'a> {
696 pub timestamp: u32,
698 pub dest_socket_id: u32,
700 pub version: u32,
702 pub encryption_field: EncryptionField,
704 pub extension_field: HandshakeExtensionFlags,
707 pub initial_seq_number: u32,
709 pub mtu: u32,
711 pub max_flow_window_size: u32,
713 pub handshake_type: HandshakeType,
715 pub srt_socket_id: u32,
717 pub syn_cookie: u32,
719 pub peer_ip: [u32; 4],
722 pub extensions: HandshakeExtensions<'a>,
724}
725
726impl<'a> HandshakePacket<'a> {
727 pub(crate) fn parse_cif(timestamp: u32, dest_socket_id: u32, cif: &'a [u8]) -> Result<Self> {
728 if cif.len() < HANDSHAKE_CIF_FIXED_LEN {
729 return Err(Error::BufferTooShort {
730 need: HANDSHAKE_CIF_FIXED_LEN,
731 have: cif.len(),
732 what: "handshake CIF",
733 });
734 }
735 let version = be32(cif, 0);
736 let word1 = be32(cif, 4);
737 let encryption_field = EncryptionField::from_bits((word1 >> 16) as u16);
738 let extension_field = HandshakeExtensionFlags((word1 & 0xFFFF) as u16);
739 let initial_seq_number = be32(cif, 8);
740 let mtu = be32(cif, 12);
741 let max_flow_window_size = be32(cif, 16);
742 let handshake_type = HandshakeType::from_bits(be32(cif, 20));
743 let srt_socket_id = be32(cif, 24);
744 let syn_cookie = be32(cif, 28);
745 let peer_ip = [be32(cif, 32), be32(cif, 36), be32(cif, 40), be32(cif, 44)];
746 let extensions = HandshakeExtensions(&cif[HANDSHAKE_CIF_FIXED_LEN..]);
747 Ok(HandshakePacket {
748 timestamp,
749 dest_socket_id,
750 version,
751 encryption_field,
752 extension_field,
753 initial_seq_number,
754 mtu,
755 max_flow_window_size,
756 handshake_type,
757 srt_socket_id,
758 syn_cookie,
759 peer_ip,
760 extensions,
761 })
762 }
763
764 pub(crate) fn cif_len(&self) -> usize {
765 HANDSHAKE_CIF_FIXED_LEN + self.extensions.0.len()
766 }
767
768 pub(crate) fn write_cif(&self, buf: &mut [u8]) -> Result<()> {
769 let word1 =
770 (u32::from(self.encryption_field.to_bits()) << 16) | u32::from(self.extension_field.0);
771 put_be32(buf, 0, self.version);
772 put_be32(buf, 4, word1);
773 put_be32(buf, 8, self.initial_seq_number);
774 put_be32(buf, 12, self.mtu);
775 put_be32(buf, 16, self.max_flow_window_size);
776 put_be32(buf, 20, self.handshake_type.to_bits());
777 put_be32(buf, 24, self.srt_socket_id);
778 put_be32(buf, 28, self.syn_cookie);
779 put_be32(buf, 32, self.peer_ip[0]);
780 put_be32(buf, 36, self.peer_ip[1]);
781 put_be32(buf, 40, self.peer_ip[2]);
782 put_be32(buf, 44, self.peer_ip[3]);
783 buf[HANDSHAKE_CIF_FIXED_LEN..].copy_from_slice(self.extensions.0);
784 Ok(())
785 }
786}
787
788#[cfg(test)]
789mod tests {
790 use super::super::control::ControlPacket;
791 use super::*;
792
793 fn sample_no_ext() -> HandshakePacket<'static> {
794 HandshakePacket {
795 timestamp: 111,
796 dest_socket_id: 222,
797 version: 5,
798 encryption_field: EncryptionField::Aes128,
799 extension_field: HandshakeExtensionFlags(0),
800 initial_seq_number: 1000,
801 mtu: 1500,
802 max_flow_window_size: 8192,
803 handshake_type: HandshakeType::Induction,
804 srt_socket_id: 0xABCD_EF01,
805 syn_cookie: 0x1234_5678,
806 peer_ip: [0x0A00_0001, 0, 0, 0],
807 extensions: HandshakeExtensions(&[]),
808 }
809 }
810
811 #[test]
812 fn round_trips_hand_computed_bytes_no_extensions() {
813 let pkt = ControlPacket::Handshake(sample_no_ext());
814 let mut buf = [0u8; 16 + HANDSHAKE_CIF_FIXED_LEN];
815 let n = pkt.serialize_into(&mut buf).unwrap();
816 assert_eq!(n, buf.len());
817 assert_eq!(&buf[0..4], &0x8000_0000u32.to_be_bytes()); assert_eq!(&buf[16..20], &5u32.to_be_bytes()); let expected_word1 = u32::from(ENCRYPTION_FIELD_AES_128) << 16;
821 assert_eq!(&buf[20..24], &expected_word1.to_be_bytes());
822 assert_eq!(&buf[36..40], &HANDSHAKE_TYPE_INDUCTION.to_be_bytes());
823 assert_eq!(ControlPacket::parse(&buf).unwrap(), pkt);
824 }
825
826 #[test]
827 fn round_trips_with_hsreq_and_sid_extensions() {
828 let hs_msg = HsExtMessage {
829 srt_version: 0x0105_0000,
830 srt_flags: HandshakeExtensionMessageFlags(
831 HS_MSG_FLAG_TSBPDSND | HS_MSG_FLAG_TSBPDRCV | HS_MSG_FLAG_CRYPT,
832 ),
833 receiver_tsbpd_delay_ms: 120,
834 sender_tsbpd_delay_ms: 120,
835 };
836 let hsreq_block = build_extension_block(ExtensionType::HsReq, &hs_msg.to_bytes()).unwrap();
837
838 let sid_contents = encode_stream_id("live/stream1");
839 let sid_block = build_extension_block(ExtensionType::Sid, &sid_contents).unwrap();
840
841 let mut ext_bytes = Vec::new();
842 ext_bytes.extend_from_slice(&hsreq_block);
843 ext_bytes.extend_from_slice(&sid_block);
844
845 let mut hp = sample_no_ext();
846 hp.handshake_type = HandshakeType::Conclusion;
847 hp.extension_field = HandshakeExtensionFlags(HS_EXT_FLAG_HSREQ);
848 hp.extensions = HandshakeExtensions(&ext_bytes);
849
850 let pkt = ControlPacket::Handshake(hp.clone());
851 let mut buf = alloc::vec![0u8; pkt.serialized_len()];
852 pkt.serialize_into(&mut buf).unwrap();
853 let parsed = ControlPacket::parse(&buf).unwrap();
854 assert_eq!(parsed, pkt);
855
856 if let ControlPacket::Handshake(h) = parsed {
857 let blocks: Vec<_> = h.extensions.iter().map(|b| b.unwrap()).collect();
858 assert_eq!(blocks.len(), 2);
859 assert_eq!(blocks[0].ext_type, ExtensionType::HsReq);
860 assert_eq!(blocks[0].as_hs_ext_message().unwrap(), hs_msg);
861 assert_eq!(blocks[1].ext_type, ExtensionType::Sid);
862 assert_eq!(blocks[1].as_stream_id().unwrap(), "live/stream1");
863 } else {
864 panic!("expected handshake");
865 }
866 }
867
868 #[test]
869 fn stream_id_padding_is_trimmed() {
870 let contents = encode_stream_id("abc");
873 assert_eq!(contents.len(), 4);
874 let block = HandshakeExtensionBlock {
875 ext_type: ExtensionType::Sid,
876 contents: &contents,
877 };
878 assert_eq!(block.as_stream_id().unwrap(), "abc");
879 }
880
881 #[test]
882 fn group_membership_extension_round_trips() {
883 let g = GroupMembershipExtension {
884 group_id: 42,
885 group_type: GroupType::MainBackup,
886 flags: GroupFlags(0x01),
887 weight: 7,
888 };
889 let bytes = g.to_bytes();
890 assert_eq!(GroupMembershipExtension::parse(&bytes).unwrap(), g);
891 assert!(g.flags.message_number_sync());
892 }
893
894 #[test]
895 fn extension_overrun_does_not_panic() {
896 let bytes = [0x00, 0x05, 0xFF, 0xFF];
898 let exts = HandshakeExtensions(&bytes);
899 let mut it = exts.iter();
900 assert!(matches!(
901 it.next(),
902 Some(Err(Error::ExtensionOverrun { .. }))
903 ));
904 assert!(it.next().is_none());
905 }
906
907 #[test]
908 fn all_encryption_fields_and_types_round_trip() {
909 for e in [
910 EncryptionField::NoEncryption,
911 EncryptionField::Aes128,
912 EncryptionField::Aes192,
913 EncryptionField::Aes256,
914 ] {
915 assert_eq!(EncryptionField::from_bits(e.to_bits()), e);
916 }
917 for t in [
918 HandshakeType::Done,
919 HandshakeType::Agreement,
920 HandshakeType::Conclusion,
921 HandshakeType::Wavehand,
922 HandshakeType::Induction,
923 ] {
924 assert_eq!(HandshakeType::from_bits(t.to_bits()), t);
925 }
926 }
927}