Skip to main content

srt_runtime/
caller.rs

1//! Caller-side handshake engine — `draft-sharabayko-srt-01` §4.3.1
2//! (Caller-Listener Handshake), Caller role.
3//!
4//! [`CallerHandshake`] is a driveable, sans-IO state machine:
5//! [`CallerHandshake::start`] returns the INDUCTION handshake bytes to send;
6//! [`CallerHandshake::feed`] consumes an inbound (already-parsed)
7//! [`ControlPacket`] and returns the [`HandshakeOutput`]s produced — further
8//! bytes to send, the negotiated parameters on success, or a rejection. No
9//! sockets, no clock: retransmit timing is driven by
10//! [`CallerHandshake::tick`] calls from the caller.
11//!
12//! Flow implemented (§4.3.1.1 / §4.3.1.2):
13//!
14//! 1. `start()`: send INDUCTION (Version 4, Encryption Field 0, Extension
15//!    Field 2, SYN Cookie 0).
16//! 2. `feed()` the Listener's INDUCTION response: validate Version 5 and the
17//!    SRT magic code `0x4A17`, capture the cookie and Listener's Socket ID,
18//!    then send CONCLUSION (Version 5, the captured cookie, HSREQ + optional
19//!    Stream ID / Group extensions).
20//! 3. `feed()` the Listener's CONCLUSION response: validate it, decode its
21//!    Handshake Extension Message, and reach [`CallerHandshakeState::Connected`]
22//!    with [`NegotiatedParams`] — or [`CallerHandshakeState::Rejected`] on any
23//!    validation failure or explicit peer rejection.
24
25use alloc::vec;
26use alloc::vec::Vec;
27
28use crate::error::{Error, Result};
29use crate::handshake_sm::{
30    self, HANDSHAKE_VERSION_4, HANDSHAKE_VERSION_5, HandshakeConfig, HandshakeOutput,
31    INDUCTION_LEGACY_SOCKET_TYPE, NegotiatedParams, RejectionReason, SRT_MAGIC_CODE,
32};
33use crate::packet::{
34    ControlPacket, EncryptionField, ExtensionType, HandshakeExtensionFlags, HandshakeExtensions,
35    HandshakePacket, HandshakeType, HsExtMessage,
36};
37
38/// Caller-side handshake lifecycle state (`draft-sharabayko-srt-01` §4.3.1).
39#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
40#[cfg_attr(feature = "serde", derive(serde::Serialize))]
41#[non_exhaustive]
42pub enum CallerHandshakeState {
43    /// No handshake message sent yet.
44    Idle,
45    /// INDUCTION sent; awaiting the Listener's INDUCTION response.
46    AwaitingInductionResponse,
47    /// CONCLUSION sent; awaiting the Listener's CONCLUSION response.
48    AwaitingConclusionResponse,
49    /// The handshake completed; [`NegotiatedParams`] are available.
50    Connected,
51    /// The handshake was rejected (locally, or by the peer).
52    Rejected,
53    /// No response arrived after the configured retry budget.
54    TimedOut,
55}
56
57impl CallerHandshakeState {
58    /// A short label for this state.
59    pub fn name(&self) -> &'static str {
60        match self {
61            CallerHandshakeState::Idle => "Idle",
62            CallerHandshakeState::AwaitingInductionResponse => "AwaitingInductionResponse",
63            CallerHandshakeState::AwaitingConclusionResponse => "AwaitingConclusionResponse",
64            CallerHandshakeState::Connected => "Connected",
65            CallerHandshakeState::Rejected => "Rejected",
66            CallerHandshakeState::TimedOut => "TimedOut",
67        }
68    }
69}
70
71impl core::fmt::Display for CallerHandshakeState {
72    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
73        f.write_str(self.name())
74    }
75}
76
77/// A driveable Caller-side SRT handshake (`draft-sharabayko-srt-01` §4.3.1).
78#[derive(Debug)]
79pub struct CallerHandshake {
80    own_socket_id: u32,
81    config: HandshakeConfig,
82    state: CallerHandshakeState,
83    peer_socket_id: u32,
84    syn_cookie: u32,
85    last_sent: Option<Vec<u8>>,
86    ticks_since_send: u32,
87    retries: u32,
88    negotiated: Option<NegotiatedParams>,
89}
90
91impl CallerHandshake {
92    /// Creates a fresh Caller handshake in [`CallerHandshakeState::Idle`].
93    pub fn new(own_socket_id: u32, config: HandshakeConfig) -> Self {
94        CallerHandshake {
95            own_socket_id,
96            config,
97            state: CallerHandshakeState::Idle,
98            peer_socket_id: 0,
99            syn_cookie: 0,
100            last_sent: None,
101            ticks_since_send: 0,
102            retries: 0,
103            negotiated: None,
104        }
105    }
106
107    /// The current state.
108    pub fn state(&self) -> CallerHandshakeState {
109        self.state
110    }
111
112    /// The negotiated parameters, once [`CallerHandshakeState::Connected`].
113    pub fn negotiated(&self) -> Option<&NegotiatedParams> {
114        self.negotiated.as_ref()
115    }
116
117    /// Builds the initial INDUCTION handshake (§4.3.1.1) and transitions to
118    /// [`CallerHandshakeState::AwaitingInductionResponse`].
119    ///
120    /// # Errors
121    /// [`Error::HandshakeOutOfSequence`] if called more than once.
122    pub fn start(&mut self) -> Result<Vec<u8>> {
123        if self.state != CallerHandshakeState::Idle {
124            return Err(Error::HandshakeOutOfSequence {
125                state: self.state.name(),
126                reason: "start() called after the handshake already began",
127            });
128        }
129        let hp = HandshakePacket {
130            timestamp: 0,
131            dest_socket_id: 0, // §4.3.1.1: 0 is interpreted as a connection request.
132            version: HANDSHAKE_VERSION_4,
133            encryption_field: EncryptionField::NoEncryption,
134            extension_field: HandshakeExtensionFlags(INDUCTION_LEGACY_SOCKET_TYPE),
135            initial_seq_number: self.config.initial_seq_number,
136            mtu: self.config.mtu,
137            max_flow_window_size: self.config.max_flow_window_size,
138            handshake_type: HandshakeType::Induction,
139            srt_socket_id: self.own_socket_id,
140            syn_cookie: 0,
141            peer_ip: self.config.local_ip,
142            extensions: HandshakeExtensions(&[]),
143        };
144        let bytes = handshake_sm::build_bytes(hp)?;
145        self.last_sent = Some(bytes.clone());
146        self.ticks_since_send = 0;
147        self.state = CallerHandshakeState::AwaitingInductionResponse;
148        Ok(bytes)
149    }
150
151    /// Feeds an inbound control packet. Only meaningful while awaiting the
152    /// Listener's INDUCTION or CONCLUSION response.
153    ///
154    /// # Errors
155    /// [`Error::UnexpectedControlPacket`] if `packet` is not a Handshake
156    /// packet; [`Error::HandshakeOutOfSequence`] if fed outside those two
157    /// states (a driver bug, not a peer failure).
158    pub fn feed(&mut self, packet: &ControlPacket<'_>) -> Result<Vec<HandshakeOutput>> {
159        let hp = match packet {
160            ControlPacket::Handshake(hp) => hp,
161            other => {
162                return Err(Error::UnexpectedControlPacket {
163                    actual: other.control_type().name(),
164                });
165            }
166        };
167        match self.state {
168            CallerHandshakeState::AwaitingInductionResponse => self.on_induction_response(hp),
169            CallerHandshakeState::AwaitingConclusionResponse => self.on_conclusion_response(hp),
170            _ => Err(Error::HandshakeOutOfSequence {
171                state: self.state.name(),
172                reason: "not awaiting a handshake response",
173            }),
174        }
175    }
176
177    /// Convenience wrapper: parses `bytes` as an [`ControlPacket`] then
178    /// [`Self::feed`]s it.
179    pub fn feed_bytes(&mut self, bytes: &[u8]) -> Result<Vec<HandshakeOutput>> {
180        let packet = ControlPacket::parse(bytes)?;
181        self.feed(&packet)
182    }
183
184    /// Advances retransmit timing by one caller-defined tick. If the
185    /// handshake is waiting on a peer response and
186    /// [`HandshakeConfig::retransmit_after_ticks`] have elapsed with none
187    /// arriving, re-emits the last sent packet; after
188    /// [`HandshakeConfig::max_retries`] such retransmissions, transitions to
189    /// [`CallerHandshakeState::TimedOut`].
190    pub fn tick(&mut self) -> Vec<HandshakeOutput> {
191        if !matches!(
192            self.state,
193            CallerHandshakeState::AwaitingInductionResponse
194                | CallerHandshakeState::AwaitingConclusionResponse
195        ) {
196            return Vec::new();
197        }
198        self.ticks_since_send += 1;
199        if self.ticks_since_send < self.config.retransmit_after_ticks {
200            return Vec::new();
201        }
202        self.ticks_since_send = 0;
203        self.retries += 1;
204        if self.retries > self.config.max_retries {
205            self.state = CallerHandshakeState::TimedOut;
206            return vec![HandshakeOutput::TimedOut];
207        }
208        match self.last_sent.clone() {
209            Some(bytes) => vec![HandshakeOutput::Send(bytes)],
210            None => Vec::new(),
211        }
212    }
213
214    fn on_induction_response(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
215        if let Some(reason) = RejectionReason::from_handshake_type(hp.handshake_type) {
216            self.state = CallerHandshakeState::Rejected;
217            return Ok(vec![HandshakeOutput::Rejected(reason)]);
218        }
219        if hp.handshake_type != HandshakeType::Induction {
220            self.state = CallerHandshakeState::Rejected;
221            return Ok(vec![HandshakeOutput::Rejected(RejectionReason::Rogue)]);
222        }
223        if hp.version != HANDSHAKE_VERSION_5 {
224            // Not an SRT party (or an incompatible version) — §4.3.1.1.
225            self.state = CallerHandshakeState::Rejected;
226            return Ok(vec![HandshakeOutput::Rejected(RejectionReason::Version)]);
227        }
228        if hp.extension_field.0 != SRT_MAGIC_CODE {
229            // §4.3.1.1: "whether the Extension Flags contains the magic
230            // value 0x4A17; otherwise the connection is rejected."
231            self.state = CallerHandshakeState::Rejected;
232            return Ok(vec![HandshakeOutput::Rejected(RejectionReason::Rogue)]);
233        }
234
235        self.peer_socket_id = hp.srt_socket_id;
236        self.syn_cookie = hp.syn_cookie;
237
238        let hs_msg = HsExtMessage {
239            srt_version: self.config.srt_version,
240            srt_flags: self.config.flags,
241            receiver_tsbpd_delay_ms: self.config.latency_ms,
242            sender_tsbpd_delay_ms: self.config.latency_ms,
243        };
244        let (ext_bytes, ext_flags) = handshake_sm::build_conclusion_extensions(
245            ExtensionType::HsReq,
246            &hs_msg,
247            self.config.stream_id.as_deref(),
248            self.config.group,
249        )?;
250        // §6.1.5: the connection initiator (Caller) is the one that sends
251        // the Key Material request, piggybacked on this same CONCLUSION.
252        #[cfg(feature = "crypto")]
253        let (ext_bytes, ext_flags) = {
254            let mut ext_bytes = ext_bytes;
255            let mut ext_flags = ext_flags;
256            if let Some(crypto) = &self.config.crypto {
257                let km_ext = handshake_sm::build_key_material_extension(crypto)?;
258                ext_bytes.extend(km_ext);
259                ext_flags |= crate::packet::handshake::HS_EXT_FLAG_KMREQ;
260            }
261            (ext_bytes, ext_flags)
262        };
263
264        let hp_out = HandshakePacket {
265            timestamp: 0,
266            // §4.3.1.2: the socket ID previously received in the induction phase.
267            dest_socket_id: self.peer_socket_id,
268            version: HANDSHAKE_VERSION_5,
269            encryption_field: self.config.encryption_field,
270            extension_field: HandshakeExtensionFlags(ext_flags),
271            initial_seq_number: self.config.initial_seq_number,
272            mtu: self.config.mtu,
273            max_flow_window_size: self.config.max_flow_window_size,
274            handshake_type: HandshakeType::Conclusion,
275            srt_socket_id: self.own_socket_id,
276            syn_cookie: self.syn_cookie,
277            peer_ip: self.config.local_ip,
278            extensions: HandshakeExtensions(&ext_bytes),
279        };
280        let bytes = handshake_sm::build_bytes(hp_out)?;
281        self.last_sent = Some(bytes.clone());
282        self.ticks_since_send = 0;
283        self.retries = 0;
284        self.state = CallerHandshakeState::AwaitingConclusionResponse;
285        Ok(vec![HandshakeOutput::Send(bytes)])
286    }
287
288    fn on_conclusion_response(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
289        if let Some(reason) = RejectionReason::from_handshake_type(hp.handshake_type) {
290            self.state = CallerHandshakeState::Rejected;
291            return Ok(vec![HandshakeOutput::Rejected(reason)]);
292        }
293        if hp.handshake_type != HandshakeType::Conclusion {
294            self.state = CallerHandshakeState::Rejected;
295            return Ok(vec![HandshakeOutput::Rejected(RejectionReason::Rogue)]);
296        }
297        if hp.version != HANDSHAKE_VERSION_5 {
298            self.state = CallerHandshakeState::Rejected;
299            return Ok(vec![HandshakeOutput::Rejected(RejectionReason::Version)]);
300        }
301
302        let parsed = match handshake_sm::parse_peer_extensions(hp) {
303            Ok(p) => p,
304            Err(_) => {
305                self.state = CallerHandshakeState::Rejected;
306                return Ok(vec![HandshakeOutput::Rejected(RejectionReason::Rogue)]);
307            }
308        };
309        let peer_msg = match parsed.hs_msg {
310            Some(m) => m,
311            None => {
312                self.state = CallerHandshakeState::Rejected;
313                return Ok(vec![HandshakeOutput::Rejected(RejectionReason::Rogue)]);
314            }
315        };
316
317        // §6.1.5: this Caller (the initiator) already holds the plaintext
318        // SEK/Salt it generated itself — it does not need to unwrap
319        // anything. It only needs the Listener's echo to *confirm* the
320        // Listener derived the same SEK (§6.1.5, "the responder echoes the
321        // same KM message back to prove it derived the same SEK"); a
322        // missing or mismatched echo means the Listener does not have (or
323        // never confirmed) the SEK.
324        #[cfg(feature = "crypto")]
325        let crypto_cfg = self.config.crypto.clone();
326        #[cfg(feature = "crypto")]
327        let crypto_result: Option<handshake_sm::RecoveredSek> = match &crypto_cfg {
328            Some(crypto) => match &parsed.km {
329                Some(echoed) if handshake_sm::verify_km_echo(crypto, echoed) => {
330                    Some((crypto.sek.clone(), crypto.salt))
331                }
332                _ => {
333                    self.state = CallerHandshakeState::Rejected;
334                    return Ok(vec![HandshakeOutput::Rejected(RejectionReason::BadSecret)]);
335                }
336            },
337            None => None,
338        };
339
340        let negotiated = NegotiatedParams {
341            version: HANDSHAKE_VERSION_5,
342            flags: crate::packet::HandshakeExtensionMessageFlags(
343                self.config.flags.0 & peer_msg.srt_flags.0,
344            ),
345            latency_ms: handshake_sm::negotiate_latency_ms(self.config.latency_ms, &peer_msg),
346            own_socket_id: self.own_socket_id,
347            peer_socket_id: self.peer_socket_id,
348            stream_id: self.config.stream_id.clone(),
349            group: self.config.group,
350            #[cfg(feature = "crypto")]
351            sek: crypto_result.as_ref().map(|(sek, _)| sek.clone()),
352            #[cfg(feature = "crypto")]
353            salt: crypto_result.as_ref().map(|(_, salt)| *salt),
354        };
355        self.negotiated = Some(negotiated.clone());
356        self.state = CallerHandshakeState::Connected;
357        Ok(vec![HandshakeOutput::Connected(negotiated)])
358    }
359}
360
361#[cfg(test)]
362mod tests {
363    use super::*;
364    use crate::packet::handshake::{HANDSHAKE_TYPE_INDUCTION, HS_EXT_FLAG_HSREQ};
365
366    #[test]
367    fn start_is_idempotent_guard() {
368        let mut c = CallerHandshake::new(1, HandshakeConfig::default());
369        assert!(c.start().is_ok());
370        assert!(c.start().is_err());
371    }
372
373    #[test]
374    fn induction_wire_values_match_draft_4_3_1_1() {
375        let mut c = CallerHandshake::new(0xAAAA_BBBB, HandshakeConfig::default());
376        let bytes = c.start().unwrap();
377        let pkt = ControlPacket::parse(&bytes).unwrap();
378        match pkt {
379            ControlPacket::Handshake(hp) => {
380                assert_eq!(hp.version, HANDSHAKE_VERSION_4);
381                assert_eq!(hp.encryption_field, EncryptionField::NoEncryption);
382                assert_eq!(hp.extension_field.0, INDUCTION_LEGACY_SOCKET_TYPE);
383                assert_eq!(hp.handshake_type.to_bits(), HANDSHAKE_TYPE_INDUCTION);
384                assert_eq!(hp.srt_socket_id, 0xAAAA_BBBB);
385                assert_eq!(hp.syn_cookie, 0);
386                assert_eq!(hp.dest_socket_id, 0);
387            }
388            _ => panic!("expected handshake"),
389        }
390        assert_eq!(c.state(), CallerHandshakeState::AwaitingInductionResponse);
391    }
392
393    fn induction_response(cookie: u32, listener_id: u32) -> ControlPacket<'static> {
394        ControlPacket::Handshake(HandshakePacket {
395            timestamp: 0,
396            dest_socket_id: 0xAAAA_BBBB,
397            version: HANDSHAKE_VERSION_5,
398            encryption_field: EncryptionField::NoEncryption,
399            extension_field: HandshakeExtensionFlags(SRT_MAGIC_CODE),
400            initial_seq_number: 0,
401            mtu: 1500,
402            max_flow_window_size: 8192,
403            handshake_type: HandshakeType::Induction,
404            srt_socket_id: listener_id,
405            syn_cookie: cookie,
406            peer_ip: [0; 4],
407            extensions: HandshakeExtensions(&[]),
408        })
409    }
410
411    #[test]
412    fn conclusion_carries_the_captured_cookie_and_hsreq() {
413        let mut c = CallerHandshake::new(0xAAAA_BBBB, HandshakeConfig::default());
414        c.start().unwrap();
415        let outputs = c
416            .feed(&induction_response(0xC0FF_EE00, 0x1111_2222))
417            .unwrap();
418        assert_eq!(outputs.len(), 1);
419        let bytes = match &outputs[0] {
420            HandshakeOutput::Send(b) => b.clone(),
421            other => panic!("expected Send, got {other:?}"),
422        };
423        let pkt = ControlPacket::parse(&bytes).unwrap();
424        match pkt {
425            ControlPacket::Handshake(hp) => {
426                assert_eq!(hp.version, HANDSHAKE_VERSION_5);
427                assert_eq!(hp.handshake_type, HandshakeType::Conclusion);
428                assert_eq!(hp.syn_cookie, 0xC0FF_EE00);
429                assert_eq!(hp.dest_socket_id, 0x1111_2222);
430                assert_eq!(hp.extension_field.0 & HS_EXT_FLAG_HSREQ, HS_EXT_FLAG_HSREQ);
431                let blocks: Vec<_> = hp.extensions.iter().map(|b| b.unwrap()).collect();
432                assert_eq!(blocks.len(), 1);
433                assert_eq!(blocks[0].ext_type, ExtensionType::HsReq);
434            }
435            _ => panic!("expected handshake"),
436        }
437        assert_eq!(c.state(), CallerHandshakeState::AwaitingConclusionResponse);
438    }
439
440    #[test]
441    fn induction_response_bad_magic_is_rejected() {
442        let mut c = CallerHandshake::new(1, HandshakeConfig::default());
443        c.start().unwrap();
444        let mut bad = induction_response(1, 2);
445        if let ControlPacket::Handshake(hp) = &mut bad {
446            hp.extension_field = HandshakeExtensionFlags(0x0000);
447        }
448        let outputs = c.feed(&bad).unwrap();
449        assert_eq!(
450            outputs,
451            vec![HandshakeOutput::Rejected(RejectionReason::Rogue)]
452        );
453        assert_eq!(c.state(), CallerHandshakeState::Rejected);
454    }
455
456    #[test]
457    fn induction_response_bad_version_is_rejected() {
458        let mut c = CallerHandshake::new(1, HandshakeConfig::default());
459        c.start().unwrap();
460        let mut bad = induction_response(1, 2);
461        if let ControlPacket::Handshake(hp) = &mut bad {
462            hp.version = HANDSHAKE_VERSION_4;
463        }
464        let outputs = c.feed(&bad).unwrap();
465        assert_eq!(
466            outputs,
467            vec![HandshakeOutput::Rejected(RejectionReason::Version)]
468        );
469        assert_eq!(c.state(), CallerHandshakeState::Rejected);
470    }
471
472    #[test]
473    fn explicit_peer_rejection_is_surfaced() {
474        let mut c = CallerHandshake::new(1, HandshakeConfig::default());
475        c.start().unwrap();
476        let mut rejected = induction_response(1, 2);
477        if let ControlPacket::Handshake(hp) = &mut rejected {
478            hp.handshake_type = RejectionReason::Backlog.to_handshake_type();
479        }
480        let outputs = c.feed(&rejected).unwrap();
481        assert_eq!(
482            outputs,
483            vec![HandshakeOutput::Rejected(RejectionReason::Backlog)]
484        );
485        assert_eq!(c.state(), CallerHandshakeState::Rejected);
486    }
487
488    #[test]
489    fn feed_before_start_is_out_of_sequence() {
490        let mut c = CallerHandshake::new(1, HandshakeConfig::default());
491        let resp = induction_response(1, 2);
492        assert!(matches!(
493            c.feed(&resp),
494            Err(Error::HandshakeOutOfSequence { .. })
495        ));
496    }
497
498    #[test]
499    fn feed_rejects_non_handshake_packets() {
500        use crate::packet::misc::KeepAlivePacket;
501        let mut c = CallerHandshake::new(1, HandshakeConfig::default());
502        c.start().unwrap();
503        let ka = ControlPacket::KeepAlive(KeepAlivePacket {
504            timestamp: 0,
505            dest_socket_id: 0,
506        });
507        assert!(matches!(
508            c.feed(&ka),
509            Err(Error::UnexpectedControlPacket { .. })
510        ));
511    }
512
513    #[test]
514    fn tick_retransmits_then_times_out() {
515        let config = HandshakeConfig {
516            retransmit_after_ticks: 2,
517            max_retries: 1,
518            ..HandshakeConfig::default()
519        };
520        let mut c = CallerHandshake::new(1, config);
521        let first = c.start().unwrap();
522
523        assert_eq!(c.tick(), Vec::new()); // 1 tick, threshold 2: nothing yet
524        let out = c.tick(); // 2 ticks: retransmit #1
525        assert_eq!(out, vec![HandshakeOutput::Send(first.clone())]);
526
527        assert_eq!(c.tick(), Vec::new());
528        let out = c.tick(); // retransmit #2 exceeds max_retries=1
529        assert_eq!(out, vec![HandshakeOutput::TimedOut]);
530        assert_eq!(c.state(), CallerHandshakeState::TimedOut);
531    }
532}