Skip to main content

srt_runtime/packet/
handshake.rs

1//! Handshake control packet — `draft-sharabayko-srt-01` §3.2.1, Figure 5, and
2//! its extension messages: Handshake Extension Message (§3.2.1.1 / Figure 6,
3//! flags in §3.2.1.1.1 / Table 6), Key Material Extension (§3.2.1.2, content
4//! is a [`super::KeyMaterial`] — §3.2.2), Stream ID Extension (§3.2.1.3 /
5//! Figure 7), and Group Membership Extension (§3.2.1.4 / Figures 8-9).
6//!
7//! ```text
8//! word0..0   Version (32)
9//! word1      Encryption Field (16) | Extension Field (16)
10//! word2      Initial Packet Sequence Number (32)
11//! word3      Maximum Transmission Unit Size (32)
12//! word4      Maximum Flow Window Size (32)
13//! word5      Handshake Type (32)
14//! word6      SRT Socket ID (32)
15//! word7      SYN Cookie (32)
16//! word8..11  Peer IP Address (128)
17//! ── repeated ──
18//!            Extension Type (16) | Extension Length (16, in 4-byte blocks)
19//!            Extension Contents (Extension Length * 4 bytes)
20//! ```
21//!
22//! This module covers *structure* only: parsing the fixed CIF fields and
23//! walking/decoding extension blocks. The handshake *exchange* (caller /
24//! listener / rendezvous state machine, §4.3) is an explicit follow-up.
25
26use alloc::string::String;
27use alloc::vec::Vec;
28
29use super::key_material::KeyMaterial;
30use super::{Error, Result, be16, be32, put_be16, put_be32};
31
32/// Length in bytes of the fixed Handshake CIF core (Version through Peer IP
33/// Address, §3.2.1 Figure 5) — 12 header words.
34pub const HANDSHAKE_CIF_FIXED_LEN: usize = 48;
35
36/// `Encryption Field` wire values (§3.2.1, Table 2).
37pub const ENCRYPTION_FIELD_NONE: u16 = 0;
38/// AES-128.
39pub const ENCRYPTION_FIELD_AES_128: u16 = 2;
40/// AES-192.
41pub const ENCRYPTION_FIELD_AES_192: u16 = 3;
42/// AES-256.
43pub const ENCRYPTION_FIELD_AES_256: u16 = 4;
44
45/// `Encryption Field`: block cipher family and key size advertised in the
46/// handshake (`draft-sharabayko-srt-01` §3.2.1, Table 2).
47#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
48#[cfg_attr(feature = "serde", derive(serde::Serialize))]
49#[non_exhaustive]
50pub enum EncryptionField {
51    /// `0`: no encryption advertised.
52    NoEncryption,
53    /// `2`: AES-128 (the default).
54    Aes128,
55    /// `3`: AES-192.
56    Aes192,
57    /// `4`: AES-256.
58    Aes256,
59    /// A value Table 2 does not define (includes `1`).
60    Reserved(u16),
61}
62
63impl EncryptionField {
64    /// Decode the 16-bit `Encryption Field`.
65    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    /// The wire value.
76    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    /// Spec label.
87    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
100/// `Handshake Type` wire values (§3.2.1, Table 4).
101pub const HANDSHAKE_TYPE_DONE: u32 = 0xFFFF_FFFD;
102/// AGREEMENT.
103pub const HANDSHAKE_TYPE_AGREEMENT: u32 = 0xFFFF_FFFE;
104/// CONCLUSION.
105pub const HANDSHAKE_TYPE_CONCLUSION: u32 = 0xFFFF_FFFF;
106/// WAVEHAND.
107pub const HANDSHAKE_TYPE_WAVEHAND: u32 = 0x0000_0000;
108/// INDUCTION.
109pub const HANDSHAKE_TYPE_INDUCTION: u32 = 0x0000_0001;
110
111/// `Handshake Type` (`draft-sharabayko-srt-01` §3.2.1, Table 4).
112#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
113#[cfg_attr(feature = "serde", derive(serde::Serialize))]
114#[non_exhaustive]
115pub enum HandshakeType {
116    /// `0xFFFFFFFD`.
117    Done,
118    /// `0xFFFFFFFE`.
119    Agreement,
120    /// `0xFFFFFFFF`.
121    Conclusion,
122    /// `0x00000000`.
123    Wavehand,
124    /// `0x00000001`.
125    Induction,
126    /// A value Table 4 does not define (real SRT implementations also use
127    /// this range for `REJ_*` rejection-reason codes, not standardised here).
128    Reserved(u32),
129}
130
131impl HandshakeType {
132    /// Decode the 32-bit `Handshake Type`.
133    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    /// The wire value.
145    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    /// Spec label.
157    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
171/// `Extension Field` bitmask values (§3.2.1, Table 3). Only meaningful on a
172/// CONCLUSION handshake; on an INDUCTION response this 16-bit field is
173/// instead echoed back opaquely by the Listener.
174pub const HS_EXT_FLAG_HSREQ: u16 = 0x0001;
175/// KMREQ.
176pub const HS_EXT_FLAG_KMREQ: u16 = 0x0002;
177/// CONFIG.
178pub const HS_EXT_FLAG_CONFIG: u16 = 0x0004;
179
180/// The `Extension Field` (§3.2.1, Table 3) — a 16-bit bitmask on a CONCLUSION
181/// handshake, or an opaque echoed value on an INDUCTION response.
182#[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    /// `HSREQ` bit set.
188    pub fn hsreq(self) -> bool {
189        self.0 & HS_EXT_FLAG_HSREQ != 0
190    }
191    /// `KMREQ` bit set.
192    pub fn kmreq(self) -> bool {
193        self.0 & HS_EXT_FLAG_KMREQ != 0
194    }
195    /// `CONFIG` bit set.
196    pub fn config(self) -> bool {
197        self.0 & HS_EXT_FLAG_CONFIG != 0
198    }
199}
200
201/// `Extension Type` wire values (§3.2.1, Table 5).
202pub const EXT_TYPE_HSREQ: u16 = 1;
203/// `SRT_CMD_HSRSP`.
204pub const EXT_TYPE_HSRSP: u16 = 2;
205/// `SRT_CMD_KMREQ` — also used as the control-packet `Subtype` for a Key
206/// Material message delivered as a User-Defined control packet (§3.2.2).
207pub const EXT_TYPE_KMREQ: u16 = 3;
208/// `SRT_CMD_KMRSP`.
209pub const EXT_TYPE_KMRSP: u16 = 4;
210/// `SRT_CMD_SID`.
211pub const EXT_TYPE_SID: u16 = 5;
212/// `SRT_CMD_CONGESTION`.
213pub const EXT_TYPE_CONGESTION: u16 = 6;
214/// `SRT_CMD_FILTER`.
215pub const EXT_TYPE_FILTER: u16 = 7;
216/// `SRT_CMD_GROUP`.
217pub const EXT_TYPE_GROUP: u16 = 8;
218
219/// Handshake Extension `Extension Type` (`draft-sharabayko-srt-01` §3.2.1,
220/// Table 5).
221#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
222#[cfg_attr(feature = "serde", derive(serde::Serialize))]
223#[non_exhaustive]
224pub enum ExtensionType {
225    /// `1`: `SRT_CMD_HSREQ` — Handshake Extension request (§3.2.1.1).
226    HsReq,
227    /// `2`: `SRT_CMD_HSRSP` — Handshake Extension response (§3.2.1.1).
228    HsRsp,
229    /// `3`: `SRT_CMD_KMREQ` — Key Material request (§3.2.1.2).
230    KmReq,
231    /// `4`: `SRT_CMD_KMRSP` — Key Material response (§3.2.1.2).
232    KmRsp,
233    /// `5`: `SRT_CMD_SID` — Stream ID (§3.2.1.3).
234    Sid,
235    /// `6`: `SRT_CMD_CONGESTION`.
236    Congestion,
237    /// `7`: `SRT_CMD_FILTER`.
238    Filter,
239    /// `8`: `SRT_CMD_GROUP` — Group Membership (§3.2.1.4).
240    Group,
241    /// A value Table 5 does not define.
242    Reserved(u16),
243}
244
245impl ExtensionType {
246    /// Decode the 16-bit `Extension Type`.
247    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    /// The wire value.
262    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    /// Spec label.
277    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
294/// `SRT Flags` bitmask values of the Handshake Extension Message (§3.2.1.1.1,
295/// Table 6).
296pub const HS_MSG_FLAG_TSBPDSND: u32 = 0x0000_0001;
297/// TSBPDRCV.
298pub const HS_MSG_FLAG_TSBPDRCV: u32 = 0x0000_0002;
299/// CRYPT.
300pub const HS_MSG_FLAG_CRYPT: u32 = 0x0000_0004;
301/// TLPKTDROP.
302pub const HS_MSG_FLAG_TLPKTDROP: u32 = 0x0000_0008;
303/// PERIODICNAK.
304pub const HS_MSG_FLAG_PERIODICNAK: u32 = 0x0000_0010;
305/// REXMITFLG.
306pub const HS_MSG_FLAG_REXMITFLG: u32 = 0x0000_0020;
307/// STREAM.
308pub const HS_MSG_FLAG_STREAM: u32 = 0x0000_0040;
309/// PACKET_FILTER.
310pub const HS_MSG_FLAG_PACKET_FILTER: u32 = 0x0000_0080;
311
312/// `SRT Flags` of a Handshake Extension Message (§3.2.1.1.1, Table 6).
313#[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    /// TSBPD used for sending.
319    pub fn tsbpdsnd(self) -> bool {
320        self.0 & HS_MSG_FLAG_TSBPDSND != 0
321    }
322    /// TSBPD used for receiving.
323    pub fn tsbpdrcv(self) -> bool {
324        self.0 & HS_MSG_FLAG_TSBPDRCV != 0
325    }
326    /// Legacy flag: peer understands the data packet `KK` field. MUST be set.
327    pub fn crypt(self) -> bool {
328        self.0 & HS_MSG_FLAG_CRYPT != 0
329    }
330    /// Too-late packet drop will be used.
331    pub fn tlpktdrop(self) -> bool {
332        self.0 & HS_MSG_FLAG_TLPKTDROP != 0
333    }
334    /// Peer will send periodic NAK packets.
335    pub fn periodicnak(self) -> bool {
336        self.0 & HS_MSG_FLAG_PERIODICNAK != 0
337    }
338    /// Legacy flag: peer understands the data packet `R` field. MUST be set.
339    pub fn rexmitflg(self) -> bool {
340        self.0 & HS_MSG_FLAG_REXMITFLG != 0
341    }
342    /// Buffer mode (`true`) vs message mode (`false`).
343    pub fn stream(self) -> bool {
344        self.0 & HS_MSG_FLAG_STREAM != 0
345    }
346    /// Peer supports packet filter.
347    pub fn packet_filter(self) -> bool {
348        self.0 & HS_MSG_FLAG_PACKET_FILTER != 0
349    }
350}
351
352/// Handshake Extension Message contents (§3.2.1.1, Figure 6) — the payload of
353/// an [`ExtensionType::HsReq`] / [`ExtensionType::HsRsp`] block.
354#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
355#[cfg_attr(feature = "serde", derive(serde::Serialize))]
356pub struct HsExtMessage {
357    /// SRT library version: `major * 0x10000 + minor * 0x100 + patch`.
358    pub srt_version: u32,
359    /// SRT configuration flags.
360    pub srt_flags: HandshakeExtensionMessageFlags,
361    /// Receiver TSBPD Delay, in milliseconds.
362    pub receiver_tsbpd_delay_ms: u16,
363    /// Sender TSBPD Delay, in milliseconds.
364    pub sender_tsbpd_delay_ms: u16,
365}
366
367/// Wire length of a Handshake Extension Message (Figure 6).
368pub const HS_EXT_MESSAGE_LEN: usize = 12;
369
370impl HsExtMessage {
371    /// Parse a Handshake Extension Message from an extension block's
372    /// contents (must be exactly 12 bytes).
373    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    /// Serialize into a fresh 12-byte buffer.
390    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
400/// `Type` (Group Type) wire values of the Group Membership Extension
401/// (§3.2.1.4).
402pub const GROUP_TYPE_UNDEFINED: u8 = 0;
403/// Broadcast.
404pub const GROUP_TYPE_BROADCAST: u8 = 1;
405/// Main/backup.
406pub const GROUP_TYPE_MAIN_BACKUP: u8 = 2;
407/// Balancing.
408pub const GROUP_TYPE_BALANCING: u8 = 3;
409/// Multicast (reserved for future use).
410pub const GROUP_TYPE_MULTICAST: u8 = 4;
411
412/// Group Membership Extension `Type` (`draft-sharabayko-srt-01` §3.2.1.4,
413/// `SRT_GTYPE_*`).
414#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
415#[cfg_attr(feature = "serde", derive(serde::Serialize))]
416#[non_exhaustive]
417pub enum GroupType {
418    /// `0`: undefined group type.
419    Undefined,
420    /// `1`: broadcast group type.
421    Broadcast,
422    /// `2`: main/backup group type.
423    MainBackup,
424    /// `3`: balancing group type.
425    Balancing,
426    /// `4`: multicast group type (reserved for future use).
427    Multicast,
428    /// A value not defined above.
429    Reserved(u8),
430}
431
432impl GroupType {
433    /// Decode the 8-bit `Type`.
434    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    /// The wire value.
446    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    /// Spec label.
458    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/// Group Membership Extension `Flags` (§3.2.1.4, Figure 9): only the `M` bit
473/// is defined.
474#[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    /// `M`: synchronize on message numbers (`true`) vs sequence numbers
482    /// (`false`).
483    pub fn message_number_sync(self) -> bool {
484        self.0 & Self::M_BIT != 0
485    }
486}
487
488/// Group Membership Extension (§3.2.1.4, Figure 8) — the payload of an
489/// [`ExtensionType::Group`] block.
490#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
491#[cfg_attr(feature = "serde", derive(serde::Serialize))]
492pub struct GroupMembershipExtension {
493    /// Group ID.
494    pub group_id: u32,
495    /// Group Type.
496    pub group_type: GroupType,
497    /// Flags.
498    pub flags: GroupFlags,
499    /// Link priority (main/backup) or otherwise reserved.
500    pub weight: u16,
501}
502
503/// Wire length of a Group Membership Extension (Figure 8).
504pub const GROUP_MEMBERSHIP_EXT_LEN: usize = 8;
505
506impl GroupMembershipExtension {
507    /// Parse from an extension block's contents (must be exactly 8 bytes).
508    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    /// Serialize into a fresh 8-byte buffer.
527    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/// One decoded Handshake Extension block (§3.2.1): `Extension Type` /
539/// `Extension Length` / `Extension Contents`. Data-carrying (borrows the raw
540/// contents) — decode the contents with [`Self::as_hs_ext_message`],
541/// [`Self::as_key_material`], [`Self::as_stream_id`], or
542/// [`Self::as_group_membership`] as appropriate for [`Self::ext_type`].
543#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
544#[cfg_attr(feature = "serde", derive(serde::Serialize))]
545pub struct HandshakeExtensionBlock<'a> {
546    /// The extension's type.
547    pub ext_type: ExtensionType,
548    /// The raw contents (`Extension Length * 4` bytes).
549    pub contents: &'a [u8],
550}
551
552impl<'a> HandshakeExtensionBlock<'a> {
553    /// Decode [`Self::contents`] as a Handshake Extension Message (§3.2.1.1)
554    /// — valid for [`ExtensionType::HsReq`] / [`ExtensionType::HsRsp`].
555    pub fn as_hs_ext_message(&self) -> Result<HsExtMessage> {
556        HsExtMessage::parse(self.contents)
557    }
558
559    /// Decode [`Self::contents`] as a Key Material message (§3.2.2) — valid
560    /// for [`ExtensionType::KmReq`] / [`ExtensionType::KmRsp`].
561    pub fn as_key_material(&self) -> Result<KeyMaterial<'a>> {
562        KeyMaterial::parse(self.contents)
563    }
564
565    /// Decode [`Self::contents`] as the Stream ID extension (§3.2.1.3) —
566    /// valid for [`ExtensionType::Sid`].
567    ///
568    /// Per §3.2.1.3 "The content is stored as 32-bit little endian words":
569    /// each 4-byte word of `contents` is byte-reversed before concatenation,
570    /// trailing NUL padding is trimmed, and the result is validated as UTF-8.
571    pub fn as_stream_id(&self) -> Result<String> {
572        if self.contents.len() % 4 != 0 {
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    /// Decode [`Self::contents`] as a Group Membership Extension (§3.2.1.4) —
590    /// valid for [`ExtensionType::Group`].
591    pub fn as_group_membership(&self) -> Result<GroupMembershipExtension> {
592        GroupMembershipExtension::parse(self.contents)
593    }
594}
595
596/// Encode the Stream ID extension's 32-bit-little-endian-word storage
597/// (§3.2.1.3) from a plain UTF-8 stream id string, padding with `0x00` up to
598/// the next 4-byte boundary. Inverse of [`HandshakeExtensionBlock::as_stream_id`].
599pub 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
611/// Build one Handshake Extension block's raw bytes (`Extension Type` +
612/// `Extension Length` + `Extension Contents`, §3.2.1) — append the result of
613/// repeated calls to chain multiple blocks.
614///
615/// # Errors
616/// [`Error::FieldTooWide`] if `contents` is not a whole number of 4-byte
617/// words, or has more than `0xFFFF` such words.
618pub fn build_extension_block(ext_type: ExtensionType, contents: &[u8]) -> Result<Vec<u8>> {
619    if contents.len() % 4 != 0 {
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/// A borrowed, lazily-walked list of Handshake Extension blocks — the same
639/// convention `dvb-si` uses for descriptor loops. Iterate with [`Self::iter`].
640#[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    /// Walk the extension blocks.
646    pub fn iter(&self) -> HandshakeExtensionIter<'a> {
647        HandshakeExtensionIter { rest: self.0 }
648    }
649}
650
651/// Iterator over [`HandshakeExtensions`]. See [`HandshakeExtensions::iter`].
652#[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/// Handshake control packet (§3.2.1, Figure 5).
693#[derive(Debug, Clone, PartialEq, Eq, Hash)]
694#[cfg_attr(feature = "serde", derive(serde::Serialize))]
695pub struct HandshakePacket<'a> {
696    /// Timestamp (§3).
697    pub timestamp: u32,
698    /// Destination Socket ID (§3).
699    pub dest_socket_id: u32,
700    /// A base protocol version number (`4` or `5`; `>5` reserved).
701    pub version: u32,
702    /// Block cipher family and key size (Table 2).
703    pub encryption_field: EncryptionField,
704    /// Message-specific extension flags/echo (Table 3, or opaque on
705    /// INDUCTION).
706    pub extension_field: HandshakeExtensionFlags,
707    /// The sequence number of the very first data packet to be sent.
708    pub initial_seq_number: u32,
709    /// Typically `1500` (Ethernet MTU) or less.
710    pub mtu: u32,
711    /// Maximum number of data packets allowed in flight.
712    pub max_flow_window_size: u32,
713    /// The handshake packet type (Table 4).
714    pub handshake_type: HandshakeType,
715    /// The ID of the source SRT socket issuing this handshake packet.
716    pub srt_socket_id: u32,
717    /// Randomized value for processing the handshake.
718    pub syn_cookie: u32,
719    /// IPv4 or IPv6 address of the packet's sender, as 4 wire words (IPv4:
720    /// only the first word is non-zero).
721    pub peer_ip: [u32; 4],
722    /// The trailing Handshake Extension blocks (possibly empty).
723    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()); // control type 0
818        assert_eq!(&buf[16..20], &5u32.to_be_bytes()); // Version
819        // Extension Field is 0 (no extensions on this sample).
820        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        // "abc" is not a multiple of 4 bytes; encode_stream_id pads with a
871        // NUL, and as_stream_id must trim it back off.
872        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        // Declares 0xFFFF words (huge) but supplies none.
897        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}