use alloc::string::String;
use alloc::vec::Vec;
use super::key_material::KeyMaterial;
use super::{Error, Result, be16, be32, put_be16, put_be32};
pub const HANDSHAKE_CIF_FIXED_LEN: usize = 48;
pub const ENCRYPTION_FIELD_NONE: u16 = 0;
pub const ENCRYPTION_FIELD_AES_128: u16 = 2;
pub const ENCRYPTION_FIELD_AES_192: u16 = 3;
pub const ENCRYPTION_FIELD_AES_256: u16 = 4;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[non_exhaustive]
pub enum EncryptionField {
NoEncryption,
Aes128,
Aes192,
Aes256,
Reserved(u16),
}
impl EncryptionField {
pub fn from_bits(v: u16) -> Self {
match v {
ENCRYPTION_FIELD_NONE => EncryptionField::NoEncryption,
ENCRYPTION_FIELD_AES_128 => EncryptionField::Aes128,
ENCRYPTION_FIELD_AES_192 => EncryptionField::Aes192,
ENCRYPTION_FIELD_AES_256 => EncryptionField::Aes256,
other => EncryptionField::Reserved(other),
}
}
pub fn to_bits(self) -> u16 {
match self {
EncryptionField::NoEncryption => ENCRYPTION_FIELD_NONE,
EncryptionField::Aes128 => ENCRYPTION_FIELD_AES_128,
EncryptionField::Aes192 => ENCRYPTION_FIELD_AES_192,
EncryptionField::Aes256 => ENCRYPTION_FIELD_AES_256,
EncryptionField::Reserved(v) => v,
}
}
pub fn name(&self) -> &'static str {
match self {
EncryptionField::NoEncryption => "no encryption advertised",
EncryptionField::Aes128 => "AES-128",
EncryptionField::Aes192 => "AES-192",
EncryptionField::Aes256 => "AES-256",
EncryptionField::Reserved(_) => "reserved",
}
}
}
broadcast_common::impl_spec_display!(EncryptionField, Reserved);
pub const HANDSHAKE_TYPE_DONE: u32 = 0xFFFF_FFFD;
pub const HANDSHAKE_TYPE_AGREEMENT: u32 = 0xFFFF_FFFE;
pub const HANDSHAKE_TYPE_CONCLUSION: u32 = 0xFFFF_FFFF;
pub const HANDSHAKE_TYPE_WAVEHAND: u32 = 0x0000_0000;
pub const HANDSHAKE_TYPE_INDUCTION: u32 = 0x0000_0001;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[non_exhaustive]
pub enum HandshakeType {
Done,
Agreement,
Conclusion,
Wavehand,
Induction,
Reserved(u32),
}
impl HandshakeType {
pub fn from_bits(v: u32) -> Self {
match v {
HANDSHAKE_TYPE_DONE => HandshakeType::Done,
HANDSHAKE_TYPE_AGREEMENT => HandshakeType::Agreement,
HANDSHAKE_TYPE_CONCLUSION => HandshakeType::Conclusion,
HANDSHAKE_TYPE_WAVEHAND => HandshakeType::Wavehand,
HANDSHAKE_TYPE_INDUCTION => HandshakeType::Induction,
other => HandshakeType::Reserved(other),
}
}
pub fn to_bits(self) -> u32 {
match self {
HandshakeType::Done => HANDSHAKE_TYPE_DONE,
HandshakeType::Agreement => HANDSHAKE_TYPE_AGREEMENT,
HandshakeType::Conclusion => HANDSHAKE_TYPE_CONCLUSION,
HandshakeType::Wavehand => HANDSHAKE_TYPE_WAVEHAND,
HandshakeType::Induction => HANDSHAKE_TYPE_INDUCTION,
HandshakeType::Reserved(v) => v,
}
}
pub fn name(&self) -> &'static str {
match self {
HandshakeType::Done => "DONE",
HandshakeType::Agreement => "AGREEMENT",
HandshakeType::Conclusion => "CONCLUSION",
HandshakeType::Wavehand => "WAVEHAND",
HandshakeType::Induction => "INDUCTION",
HandshakeType::Reserved(_) => "reserved",
}
}
}
broadcast_common::impl_spec_display!(HandshakeType, Reserved);
pub const HS_EXT_FLAG_HSREQ: u16 = 0x0001;
pub const HS_EXT_FLAG_KMREQ: u16 = 0x0002;
pub const HS_EXT_FLAG_CONFIG: u16 = 0x0004;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct HandshakeExtensionFlags(pub u16);
impl HandshakeExtensionFlags {
pub fn hsreq(self) -> bool {
self.0 & HS_EXT_FLAG_HSREQ != 0
}
pub fn kmreq(self) -> bool {
self.0 & HS_EXT_FLAG_KMREQ != 0
}
pub fn config(self) -> bool {
self.0 & HS_EXT_FLAG_CONFIG != 0
}
}
pub const EXT_TYPE_HSREQ: u16 = 1;
pub const EXT_TYPE_HSRSP: u16 = 2;
pub const EXT_TYPE_KMREQ: u16 = 3;
pub const EXT_TYPE_KMRSP: u16 = 4;
pub const EXT_TYPE_SID: u16 = 5;
pub const EXT_TYPE_CONGESTION: u16 = 6;
pub const EXT_TYPE_FILTER: u16 = 7;
pub const EXT_TYPE_GROUP: u16 = 8;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[non_exhaustive]
pub enum ExtensionType {
HsReq,
HsRsp,
KmReq,
KmRsp,
Sid,
Congestion,
Filter,
Group,
Reserved(u16),
}
impl ExtensionType {
pub fn from_bits(v: u16) -> Self {
match v {
EXT_TYPE_HSREQ => ExtensionType::HsReq,
EXT_TYPE_HSRSP => ExtensionType::HsRsp,
EXT_TYPE_KMREQ => ExtensionType::KmReq,
EXT_TYPE_KMRSP => ExtensionType::KmRsp,
EXT_TYPE_SID => ExtensionType::Sid,
EXT_TYPE_CONGESTION => ExtensionType::Congestion,
EXT_TYPE_FILTER => ExtensionType::Filter,
EXT_TYPE_GROUP => ExtensionType::Group,
other => ExtensionType::Reserved(other),
}
}
pub fn to_bits(self) -> u16 {
match self {
ExtensionType::HsReq => EXT_TYPE_HSREQ,
ExtensionType::HsRsp => EXT_TYPE_HSRSP,
ExtensionType::KmReq => EXT_TYPE_KMREQ,
ExtensionType::KmRsp => EXT_TYPE_KMRSP,
ExtensionType::Sid => EXT_TYPE_SID,
ExtensionType::Congestion => EXT_TYPE_CONGESTION,
ExtensionType::Filter => EXT_TYPE_FILTER,
ExtensionType::Group => EXT_TYPE_GROUP,
ExtensionType::Reserved(v) => v,
}
}
pub fn name(&self) -> &'static str {
match self {
ExtensionType::HsReq => "SRT_CMD_HSREQ",
ExtensionType::HsRsp => "SRT_CMD_HSRSP",
ExtensionType::KmReq => "SRT_CMD_KMREQ",
ExtensionType::KmRsp => "SRT_CMD_KMRSP",
ExtensionType::Sid => "SRT_CMD_SID",
ExtensionType::Congestion => "SRT_CMD_CONGESTION",
ExtensionType::Filter => "SRT_CMD_FILTER",
ExtensionType::Group => "SRT_CMD_GROUP",
ExtensionType::Reserved(_) => "reserved",
}
}
}
broadcast_common::impl_spec_display!(ExtensionType, Reserved);
pub const HS_MSG_FLAG_TSBPDSND: u32 = 0x0000_0001;
pub const HS_MSG_FLAG_TSBPDRCV: u32 = 0x0000_0002;
pub const HS_MSG_FLAG_CRYPT: u32 = 0x0000_0004;
pub const HS_MSG_FLAG_TLPKTDROP: u32 = 0x0000_0008;
pub const HS_MSG_FLAG_PERIODICNAK: u32 = 0x0000_0010;
pub const HS_MSG_FLAG_REXMITFLG: u32 = 0x0000_0020;
pub const HS_MSG_FLAG_STREAM: u32 = 0x0000_0040;
pub const HS_MSG_FLAG_PACKET_FILTER: u32 = 0x0000_0080;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct HandshakeExtensionMessageFlags(pub u32);
impl HandshakeExtensionMessageFlags {
pub fn tsbpdsnd(self) -> bool {
self.0 & HS_MSG_FLAG_TSBPDSND != 0
}
pub fn tsbpdrcv(self) -> bool {
self.0 & HS_MSG_FLAG_TSBPDRCV != 0
}
pub fn crypt(self) -> bool {
self.0 & HS_MSG_FLAG_CRYPT != 0
}
pub fn tlpktdrop(self) -> bool {
self.0 & HS_MSG_FLAG_TLPKTDROP != 0
}
pub fn periodicnak(self) -> bool {
self.0 & HS_MSG_FLAG_PERIODICNAK != 0
}
pub fn rexmitflg(self) -> bool {
self.0 & HS_MSG_FLAG_REXMITFLG != 0
}
pub fn stream(self) -> bool {
self.0 & HS_MSG_FLAG_STREAM != 0
}
pub fn packet_filter(self) -> bool {
self.0 & HS_MSG_FLAG_PACKET_FILTER != 0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct HsExtMessage {
pub srt_version: u32,
pub srt_flags: HandshakeExtensionMessageFlags,
pub receiver_tsbpd_delay_ms: u16,
pub sender_tsbpd_delay_ms: u16,
}
pub const HS_EXT_MESSAGE_LEN: usize = 12;
impl HsExtMessage {
pub fn parse(bytes: &[u8]) -> Result<Self> {
if bytes.len() != HS_EXT_MESSAGE_LEN {
return Err(Error::BufferTooShort {
need: HS_EXT_MESSAGE_LEN,
have: bytes.len(),
what: "handshake extension message",
});
}
Ok(HsExtMessage {
srt_version: be32(bytes, 0),
srt_flags: HandshakeExtensionMessageFlags(be32(bytes, 4)),
receiver_tsbpd_delay_ms: be16(bytes, 8),
sender_tsbpd_delay_ms: be16(bytes, 10),
})
}
pub fn to_bytes(&self) -> [u8; HS_EXT_MESSAGE_LEN] {
let mut buf = [0u8; HS_EXT_MESSAGE_LEN];
put_be32(&mut buf, 0, self.srt_version);
put_be32(&mut buf, 4, self.srt_flags.0);
put_be16(&mut buf, 8, self.receiver_tsbpd_delay_ms);
put_be16(&mut buf, 10, self.sender_tsbpd_delay_ms);
buf
}
}
pub const GROUP_TYPE_UNDEFINED: u8 = 0;
pub const GROUP_TYPE_BROADCAST: u8 = 1;
pub const GROUP_TYPE_MAIN_BACKUP: u8 = 2;
pub const GROUP_TYPE_BALANCING: u8 = 3;
pub const GROUP_TYPE_MULTICAST: u8 = 4;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[non_exhaustive]
pub enum GroupType {
Undefined,
Broadcast,
MainBackup,
Balancing,
Multicast,
Reserved(u8),
}
impl GroupType {
pub fn from_bits(v: u8) -> Self {
match v {
GROUP_TYPE_UNDEFINED => GroupType::Undefined,
GROUP_TYPE_BROADCAST => GroupType::Broadcast,
GROUP_TYPE_MAIN_BACKUP => GroupType::MainBackup,
GROUP_TYPE_BALANCING => GroupType::Balancing,
GROUP_TYPE_MULTICAST => GroupType::Multicast,
other => GroupType::Reserved(other),
}
}
pub fn to_bits(self) -> u8 {
match self {
GroupType::Undefined => GROUP_TYPE_UNDEFINED,
GroupType::Broadcast => GROUP_TYPE_BROADCAST,
GroupType::MainBackup => GROUP_TYPE_MAIN_BACKUP,
GroupType::Balancing => GROUP_TYPE_BALANCING,
GroupType::Multicast => GROUP_TYPE_MULTICAST,
GroupType::Reserved(v) => v,
}
}
pub fn name(&self) -> &'static str {
match self {
GroupType::Undefined => "undefined",
GroupType::Broadcast => "broadcast",
GroupType::MainBackup => "main/backup",
GroupType::Balancing => "balancing",
GroupType::Multicast => "multicast",
GroupType::Reserved(_) => "reserved",
}
}
}
broadcast_common::impl_spec_display!(GroupType, Reserved);
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct GroupFlags(pub u8);
impl GroupFlags {
const M_BIT: u8 = 0x01;
pub fn message_number_sync(self) -> bool {
self.0 & Self::M_BIT != 0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct GroupMembershipExtension {
pub group_id: u32,
pub group_type: GroupType,
pub flags: GroupFlags,
pub weight: u16,
}
pub const GROUP_MEMBERSHIP_EXT_LEN: usize = 8;
impl GroupMembershipExtension {
pub fn parse(bytes: &[u8]) -> Result<Self> {
if bytes.len() != GROUP_MEMBERSHIP_EXT_LEN {
return Err(Error::BufferTooShort {
need: GROUP_MEMBERSHIP_EXT_LEN,
have: bytes.len(),
what: "group membership extension",
});
}
let group_id = be32(bytes, 0);
let word1 = be32(bytes, 4);
Ok(GroupMembershipExtension {
group_id,
group_type: GroupType::from_bits((word1 >> 24) as u8),
flags: GroupFlags((word1 >> 16) as u8),
weight: (word1 & 0xFFFF) as u16,
})
}
pub fn to_bytes(&self) -> [u8; GROUP_MEMBERSHIP_EXT_LEN] {
let mut buf = [0u8; GROUP_MEMBERSHIP_EXT_LEN];
let word1 = (u32::from(self.group_type.to_bits()) << 24)
| (u32::from(self.flags.0) << 16)
| u32::from(self.weight);
put_be32(&mut buf, 0, self.group_id);
put_be32(&mut buf, 4, word1);
buf
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct HandshakeExtensionBlock<'a> {
pub ext_type: ExtensionType,
pub contents: &'a [u8],
}
impl<'a> HandshakeExtensionBlock<'a> {
pub fn as_hs_ext_message(&self) -> Result<HsExtMessage> {
HsExtMessage::parse(self.contents)
}
pub fn as_key_material(&self) -> Result<KeyMaterial<'a>> {
KeyMaterial::parse(self.contents)
}
pub fn as_stream_id(&self) -> Result<String> {
if !self.contents.len().is_multiple_of(4) {
return Err(Error::BufferTooShort {
need: self.contents.len().div_ceil(4) * 4,
have: self.contents.len(),
what: "stream ID extension (not a whole number of 4-byte words)",
});
}
let mut bytes = Vec::with_capacity(self.contents.len());
for word in self.contents.chunks_exact(4) {
bytes.extend(word.iter().rev());
}
while bytes.last() == Some(&0u8) {
bytes.pop();
}
String::from_utf8(bytes).map_err(|_| Error::InvalidStreamIdUtf8)
}
pub fn as_group_membership(&self) -> Result<GroupMembershipExtension> {
GroupMembershipExtension::parse(self.contents)
}
}
pub fn encode_stream_id(id: &str) -> Vec<u8> {
let mut bytes = Vec::from(id.as_bytes());
while bytes.len() % 4 != 0 {
bytes.push(0);
}
let mut out = Vec::with_capacity(bytes.len());
for word in bytes.chunks_exact(4) {
out.extend(word.iter().rev());
}
out
}
pub fn build_extension_block(ext_type: ExtensionType, contents: &[u8]) -> Result<Vec<u8>> {
if !contents.len().is_multiple_of(4) {
return Err(Error::InvalidField {
what: "Extension Contents",
reason: "length must be a whole number of 4-byte words",
});
}
let words = contents.len() / 4;
let words_u16 = u16::try_from(words).map_err(|_| Error::FieldTooWide {
what: "Extension Length",
value: words as u64,
bits: 16,
})?;
let mut out = Vec::with_capacity(4 + contents.len());
out.extend_from_slice(&ext_type.to_bits().to_be_bytes());
out.extend_from_slice(&words_u16.to_be_bytes());
out.extend_from_slice(contents);
Ok(out)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct HandshakeExtensions<'a>(pub &'a [u8]);
impl<'a> HandshakeExtensions<'a> {
pub fn iter(&self) -> HandshakeExtensionIter<'a> {
HandshakeExtensionIter { rest: self.0 }
}
}
#[derive(Debug, Clone)]
pub struct HandshakeExtensionIter<'a> {
rest: &'a [u8],
}
impl<'a> Iterator for HandshakeExtensionIter<'a> {
type Item = Result<HandshakeExtensionBlock<'a>>;
fn next(&mut self) -> Option<Self::Item> {
if self.rest.is_empty() {
return None;
}
if self.rest.len() < 4 {
self.rest = &[];
return Some(Err(Error::BufferTooShort {
need: 4,
have: self.rest.len(),
what: "handshake extension block header",
}));
}
let ext_type_bits = be16(self.rest, 0);
let ext_len_words = be16(self.rest, 2);
let ext_len_bytes = usize::from(ext_len_words) * 4;
if self.rest.len() < 4 + ext_len_bytes {
let remaining = self.rest.len() - 4;
self.rest = &[];
return Some(Err(Error::ExtensionOverrun {
declared: ext_len_words,
remaining,
}));
}
let contents = &self.rest[4..4 + ext_len_bytes];
self.rest = &self.rest[4 + ext_len_bytes..];
Some(Ok(HandshakeExtensionBlock {
ext_type: ExtensionType::from_bits(ext_type_bits),
contents,
}))
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct HandshakePacket<'a> {
pub timestamp: u32,
pub dest_socket_id: u32,
pub version: u32,
pub encryption_field: EncryptionField,
pub extension_field: HandshakeExtensionFlags,
pub initial_seq_number: u32,
pub mtu: u32,
pub max_flow_window_size: u32,
pub handshake_type: HandshakeType,
pub srt_socket_id: u32,
pub syn_cookie: u32,
pub peer_ip: [u32; 4],
pub extensions: HandshakeExtensions<'a>,
}
impl<'a> HandshakePacket<'a> {
pub(crate) fn parse_cif(timestamp: u32, dest_socket_id: u32, cif: &'a [u8]) -> Result<Self> {
if cif.len() < HANDSHAKE_CIF_FIXED_LEN {
return Err(Error::BufferTooShort {
need: HANDSHAKE_CIF_FIXED_LEN,
have: cif.len(),
what: "handshake CIF",
});
}
let version = be32(cif, 0);
let word1 = be32(cif, 4);
let encryption_field = EncryptionField::from_bits((word1 >> 16) as u16);
let extension_field = HandshakeExtensionFlags((word1 & 0xFFFF) as u16);
let initial_seq_number = be32(cif, 8);
let mtu = be32(cif, 12);
let max_flow_window_size = be32(cif, 16);
let handshake_type = HandshakeType::from_bits(be32(cif, 20));
let srt_socket_id = be32(cif, 24);
let syn_cookie = be32(cif, 28);
let peer_ip = [be32(cif, 32), be32(cif, 36), be32(cif, 40), be32(cif, 44)];
let extensions = HandshakeExtensions(&cif[HANDSHAKE_CIF_FIXED_LEN..]);
Ok(HandshakePacket {
timestamp,
dest_socket_id,
version,
encryption_field,
extension_field,
initial_seq_number,
mtu,
max_flow_window_size,
handshake_type,
srt_socket_id,
syn_cookie,
peer_ip,
extensions,
})
}
pub(crate) fn cif_len(&self) -> usize {
HANDSHAKE_CIF_FIXED_LEN + self.extensions.0.len()
}
pub(crate) fn write_cif(&self, buf: &mut [u8]) -> Result<()> {
let word1 =
(u32::from(self.encryption_field.to_bits()) << 16) | u32::from(self.extension_field.0);
put_be32(buf, 0, self.version);
put_be32(buf, 4, word1);
put_be32(buf, 8, self.initial_seq_number);
put_be32(buf, 12, self.mtu);
put_be32(buf, 16, self.max_flow_window_size);
put_be32(buf, 20, self.handshake_type.to_bits());
put_be32(buf, 24, self.srt_socket_id);
put_be32(buf, 28, self.syn_cookie);
put_be32(buf, 32, self.peer_ip[0]);
put_be32(buf, 36, self.peer_ip[1]);
put_be32(buf, 40, self.peer_ip[2]);
put_be32(buf, 44, self.peer_ip[3]);
buf[HANDSHAKE_CIF_FIXED_LEN..].copy_from_slice(self.extensions.0);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::super::control::ControlPacket;
use super::*;
fn sample_no_ext() -> HandshakePacket<'static> {
HandshakePacket {
timestamp: 111,
dest_socket_id: 222,
version: 5,
encryption_field: EncryptionField::Aes128,
extension_field: HandshakeExtensionFlags(0),
initial_seq_number: 1000,
mtu: 1500,
max_flow_window_size: 8192,
handshake_type: HandshakeType::Induction,
srt_socket_id: 0xABCD_EF01,
syn_cookie: 0x1234_5678,
peer_ip: [0x0A00_0001, 0, 0, 0],
extensions: HandshakeExtensions(&[]),
}
}
#[test]
fn round_trips_hand_computed_bytes_no_extensions() {
let pkt = ControlPacket::Handshake(sample_no_ext());
let mut buf = [0u8; 16 + HANDSHAKE_CIF_FIXED_LEN];
let n = pkt.serialize_into(&mut buf).unwrap();
assert_eq!(n, buf.len());
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;
assert_eq!(&buf[20..24], &expected_word1.to_be_bytes());
assert_eq!(&buf[36..40], &HANDSHAKE_TYPE_INDUCTION.to_be_bytes());
assert_eq!(ControlPacket::parse(&buf).unwrap(), pkt);
}
#[test]
fn round_trips_with_hsreq_and_sid_extensions() {
let hs_msg = HsExtMessage {
srt_version: 0x0105_0000,
srt_flags: HandshakeExtensionMessageFlags(
HS_MSG_FLAG_TSBPDSND | HS_MSG_FLAG_TSBPDRCV | HS_MSG_FLAG_CRYPT,
),
receiver_tsbpd_delay_ms: 120,
sender_tsbpd_delay_ms: 120,
};
let hsreq_block = build_extension_block(ExtensionType::HsReq, &hs_msg.to_bytes()).unwrap();
let sid_contents = encode_stream_id("live/stream1");
let sid_block = build_extension_block(ExtensionType::Sid, &sid_contents).unwrap();
let mut ext_bytes = Vec::new();
ext_bytes.extend_from_slice(&hsreq_block);
ext_bytes.extend_from_slice(&sid_block);
let mut hp = sample_no_ext();
hp.handshake_type = HandshakeType::Conclusion;
hp.extension_field = HandshakeExtensionFlags(HS_EXT_FLAG_HSREQ);
hp.extensions = HandshakeExtensions(&ext_bytes);
let pkt = ControlPacket::Handshake(hp.clone());
let mut buf = alloc::vec![0u8; pkt.serialized_len()];
pkt.serialize_into(&mut buf).unwrap();
let parsed = ControlPacket::parse(&buf).unwrap();
assert_eq!(parsed, pkt);
if let ControlPacket::Handshake(h) = parsed {
let blocks: Vec<_> = h.extensions.iter().map(|b| b.unwrap()).collect();
assert_eq!(blocks.len(), 2);
assert_eq!(blocks[0].ext_type, ExtensionType::HsReq);
assert_eq!(blocks[0].as_hs_ext_message().unwrap(), hs_msg);
assert_eq!(blocks[1].ext_type, ExtensionType::Sid);
assert_eq!(blocks[1].as_stream_id().unwrap(), "live/stream1");
} else {
panic!("expected handshake");
}
}
#[test]
fn stream_id_padding_is_trimmed() {
let contents = encode_stream_id("abc");
assert_eq!(contents.len(), 4);
let block = HandshakeExtensionBlock {
ext_type: ExtensionType::Sid,
contents: &contents,
};
assert_eq!(block.as_stream_id().unwrap(), "abc");
}
#[test]
fn group_membership_extension_round_trips() {
let g = GroupMembershipExtension {
group_id: 42,
group_type: GroupType::MainBackup,
flags: GroupFlags(0x01),
weight: 7,
};
let bytes = g.to_bytes();
assert_eq!(GroupMembershipExtension::parse(&bytes).unwrap(), g);
assert!(g.flags.message_number_sync());
}
#[test]
fn extension_overrun_does_not_panic() {
let bytes = [0x00, 0x05, 0xFF, 0xFF];
let exts = HandshakeExtensions(&bytes);
let mut it = exts.iter();
assert!(matches!(
it.next(),
Some(Err(Error::ExtensionOverrun { .. }))
));
assert!(it.next().is_none());
}
#[test]
fn all_encryption_fields_and_types_round_trip() {
for e in [
EncryptionField::NoEncryption,
EncryptionField::Aes128,
EncryptionField::Aes192,
EncryptionField::Aes256,
] {
assert_eq!(EncryptionField::from_bits(e.to_bits()), e);
}
for t in [
HandshakeType::Done,
HandshakeType::Agreement,
HandshakeType::Conclusion,
HandshakeType::Wavehand,
HandshakeType::Induction,
] {
assert_eq!(HandshakeType::from_bits(t.to_bits()), t);
}
}
}