1use alloc::vec;
25use alloc::vec::Vec;
26
27use crate::error::{Error, Result};
28use crate::handshake_sm::{
29 self, HANDSHAKE_VERSION_5, HandshakeConfig, HandshakeOutput, NegotiatedParams, RejectionReason,
30 SRT_MAGIC_CODE,
31};
32use crate::packet::{
33 ControlPacket, EncryptionField, ExtensionType, HandshakeExtensionFlags, HandshakeExtensions,
34 HandshakePacket, HandshakeType, HsExtMessage,
35};
36
37#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
39#[cfg_attr(feature = "serde", derive(serde::Serialize))]
40#[non_exhaustive]
41pub enum ListenerHandshakeState {
42 Idle,
44 AwaitingConclusion,
46 Connected,
48 Rejected,
50 TimedOut,
52}
53
54impl ListenerHandshakeState {
55 pub fn name(&self) -> &'static str {
57 match self {
58 ListenerHandshakeState::Idle => "Idle",
59 ListenerHandshakeState::AwaitingConclusion => "AwaitingConclusion",
60 ListenerHandshakeState::Connected => "Connected",
61 ListenerHandshakeState::Rejected => "Rejected",
62 ListenerHandshakeState::TimedOut => "TimedOut",
63 }
64 }
65}
66
67impl core::fmt::Display for ListenerHandshakeState {
68 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
69 f.write_str(self.name())
70 }
71}
72
73#[derive(Debug)]
79pub struct ListenerHandshake {
80 own_socket_id: u32,
81 syn_cookie: u32,
82 config: HandshakeConfig,
83 state: ListenerHandshakeState,
84 peer_socket_id: u32,
85 last_sent: Option<Vec<u8>>,
86 ticks_since_send: u32,
87 retries: u32,
88 negotiated: Option<NegotiatedParams>,
89}
90
91impl ListenerHandshake {
92 pub fn new(own_socket_id: u32, syn_cookie: u32, config: HandshakeConfig) -> Self {
96 ListenerHandshake {
97 own_socket_id,
98 syn_cookie,
99 config,
100 state: ListenerHandshakeState::Idle,
101 peer_socket_id: 0,
102 last_sent: None,
103 ticks_since_send: 0,
104 retries: 0,
105 negotiated: None,
106 }
107 }
108
109 pub fn state(&self) -> ListenerHandshakeState {
111 self.state
112 }
113
114 pub fn negotiated(&self) -> Option<&NegotiatedParams> {
116 self.negotiated.as_ref()
117 }
118
119 pub fn feed(&mut self, packet: &ControlPacket<'_>) -> Result<Vec<HandshakeOutput>> {
127 let hp = match packet {
128 ControlPacket::Handshake(hp) => hp,
129 other => {
130 return Err(Error::UnexpectedControlPacket {
131 actual: other.control_type().name(),
132 });
133 }
134 };
135 match self.state {
136 ListenerHandshakeState::Idle => self.on_induction(hp),
137 ListenerHandshakeState::AwaitingConclusion => self.on_conclusion(hp),
138 _ => Err(Error::HandshakeOutOfSequence {
139 state: self.state.name(),
140 reason: "not awaiting an induction or conclusion",
141 }),
142 }
143 }
144
145 pub fn feed_bytes(&mut self, bytes: &[u8]) -> Result<Vec<HandshakeOutput>> {
148 let packet = ControlPacket::parse(bytes)?;
149 self.feed(&packet)
150 }
151
152 pub fn tick(&mut self) -> Vec<HandshakeOutput> {
155 if self.state != ListenerHandshakeState::AwaitingConclusion {
156 return Vec::new();
157 }
158 self.ticks_since_send += 1;
159 if self.ticks_since_send < self.config.retransmit_after_ticks {
160 return Vec::new();
161 }
162 self.ticks_since_send = 0;
163 self.retries += 1;
164 if self.retries > self.config.max_retries {
165 self.state = ListenerHandshakeState::TimedOut;
166 return vec![HandshakeOutput::TimedOut];
167 }
168 match self.last_sent.clone() {
169 Some(bytes) => vec![HandshakeOutput::Send(bytes)],
170 None => Vec::new(),
171 }
172 }
173
174 fn on_induction(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
175 if hp.handshake_type != HandshakeType::Induction {
176 return Err(Error::HandshakeOutOfSequence {
177 state: self.state.name(),
178 reason: "expected an INDUCTION handshake",
179 });
180 }
181 self.peer_socket_id = hp.srt_socket_id;
184
185 let hp_out = HandshakePacket {
186 timestamp: 0,
187 dest_socket_id: self.peer_socket_id,
188 version: HANDSHAKE_VERSION_5,
189 encryption_field: self.config.encryption_field,
190 extension_field: HandshakeExtensionFlags(SRT_MAGIC_CODE),
191 initial_seq_number: self.config.initial_seq_number,
192 mtu: self.config.mtu,
193 max_flow_window_size: self.config.max_flow_window_size,
194 handshake_type: HandshakeType::Induction,
195 srt_socket_id: self.own_socket_id,
196 syn_cookie: self.syn_cookie,
197 peer_ip: self.config.local_ip,
198 extensions: HandshakeExtensions(&[]),
199 };
200 let bytes = handshake_sm::build_bytes(hp_out)?;
201 self.last_sent = Some(bytes.clone());
202 self.ticks_since_send = 0;
203 self.state = ListenerHandshakeState::AwaitingConclusion;
204 Ok(vec![HandshakeOutput::Send(bytes)])
205 }
206
207 fn on_conclusion(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
208 if hp.handshake_type != HandshakeType::Conclusion {
209 return self.reject(RejectionReason::Rogue, hp);
210 }
211 if hp.version != HANDSHAKE_VERSION_5 {
212 return self.reject(RejectionReason::Version, hp);
213 }
214 if hp.syn_cookie != self.syn_cookie {
215 return self.reject(RejectionReason::Rogue, hp);
218 }
219
220 let parsed = match handshake_sm::parse_peer_extensions(hp) {
221 Ok(p) => p,
222 Err(_) => return self.reject(RejectionReason::Rogue, hp),
223 };
224 let peer_msg = match parsed.hs_msg {
225 Some(m) => m,
226 None => return self.reject(RejectionReason::Rogue, hp),
227 };
228
229 self.peer_socket_id = hp.srt_socket_id;
230
231 #[cfg(feature = "crypto")]
239 let crypto_cfg = self.config.crypto.clone();
240 #[cfg(feature = "crypto")]
241 type CryptoConclusionResult = (
242 Option<handshake_sm::RecoveredSek>,
243 Option<alloc::vec::Vec<u8>>,
244 );
245 #[cfg(feature = "crypto")]
246 let (crypto_result, km_echo_ext): CryptoConclusionResult = match (&crypto_cfg, &parsed.km) {
247 (Some(crypto), Some(km_req)) => {
248 match handshake_sm::recover_sek(&crypto.passphrase, km_req) {
249 Ok((sek, salt)) => {
250 let echo = match handshake_sm::echo_key_material_as_response(km_req) {
251 Ok(e) => e,
252 Err(_) => return self.reject(RejectionReason::Rogue, hp),
253 };
254 (Some((sek, salt)), Some(echo))
255 }
256 Err(_) => return self.reject(RejectionReason::BadSecret, hp),
257 }
258 }
259 (Some(_), None) | (None, Some(_)) => {
260 return self.reject(RejectionReason::Unsecure, hp);
261 }
262 (None, None) => (None, None),
263 };
264
265 let negotiated = NegotiatedParams {
266 version: HANDSHAKE_VERSION_5,
267 flags: crate::packet::HandshakeExtensionMessageFlags(
268 self.config.flags.0 & peer_msg.srt_flags.0,
269 ),
270 latency_ms: handshake_sm::negotiate_latency_ms(self.config.latency_ms, &peer_msg),
271 own_socket_id: self.own_socket_id,
272 peer_socket_id: self.peer_socket_id,
273 stream_id: parsed.stream_id,
274 group: parsed.group,
275 #[cfg(feature = "crypto")]
276 sek: crypto_result.as_ref().map(|(sek, _)| sek.clone()),
277 #[cfg(feature = "crypto")]
278 salt: crypto_result.as_ref().map(|(_, salt)| *salt),
279 };
280
281 let hs_msg = HsExtMessage {
282 srt_version: self.config.srt_version,
283 srt_flags: self.config.flags,
284 receiver_tsbpd_delay_ms: self.config.latency_ms,
285 sender_tsbpd_delay_ms: self.config.latency_ms,
286 };
287 let (ext_bytes, ext_flags) = handshake_sm::build_conclusion_extensions(
292 ExtensionType::HsRsp,
293 &hs_msg,
294 None,
295 self.config.group,
296 )?;
297 #[cfg(feature = "crypto")]
299 let (ext_bytes, ext_flags) = {
300 let mut ext_bytes = ext_bytes;
301 let mut ext_flags = ext_flags;
302 if let Some(echo) = km_echo_ext {
303 ext_bytes.extend(echo);
304 ext_flags |= crate::packet::handshake::HS_EXT_FLAG_KMREQ;
305 }
306 (ext_bytes, ext_flags)
307 };
308
309 let hp_out = HandshakePacket {
310 timestamp: 0,
311 dest_socket_id: self.peer_socket_id,
312 version: HANDSHAKE_VERSION_5,
313 encryption_field: self.config.encryption_field,
314 extension_field: HandshakeExtensionFlags(ext_flags),
315 initial_seq_number: self.config.initial_seq_number,
316 mtu: self.config.mtu,
317 max_flow_window_size: self.config.max_flow_window_size,
318 handshake_type: HandshakeType::Conclusion,
319 srt_socket_id: self.own_socket_id,
320 syn_cookie: 0, peer_ip: self.config.local_ip,
322 extensions: HandshakeExtensions(&ext_bytes),
323 };
324 let bytes = handshake_sm::build_bytes(hp_out)?;
325 self.last_sent = Some(bytes.clone());
326 self.negotiated = Some(negotiated.clone());
327 self.state = ListenerHandshakeState::Connected;
328 Ok(vec![
329 HandshakeOutput::Send(bytes),
330 HandshakeOutput::Connected(negotiated),
331 ])
332 }
333
334 fn reject(
337 &mut self,
338 reason: RejectionReason,
339 hp: &HandshakePacket<'_>,
340 ) -> Result<Vec<HandshakeOutput>> {
341 self.state = ListenerHandshakeState::Rejected;
342 let hp_out = HandshakePacket {
343 timestamp: 0,
344 dest_socket_id: hp.srt_socket_id,
345 version: HANDSHAKE_VERSION_5,
346 encryption_field: EncryptionField::NoEncryption,
347 extension_field: HandshakeExtensionFlags(0),
348 initial_seq_number: 0,
349 mtu: self.config.mtu,
350 max_flow_window_size: self.config.max_flow_window_size,
351 handshake_type: reason.to_handshake_type(),
352 srt_socket_id: self.own_socket_id,
353 syn_cookie: 0,
354 peer_ip: self.config.local_ip,
355 extensions: HandshakeExtensions(&[]),
356 };
357 let bytes = handshake_sm::build_bytes(hp_out)?;
358 self.last_sent = Some(bytes.clone());
359 Ok(vec![
360 HandshakeOutput::Send(bytes),
361 HandshakeOutput::Rejected(reason),
362 ])
363 }
364}
365
366#[cfg(test)]
367mod tests {
368 use super::*;
369 use crate::packet::handshake::{HANDSHAKE_CIF_FIXED_LEN, HS_EXT_FLAG_HSREQ};
370
371 fn caller_induction(caller_id: u32) -> ControlPacket<'static> {
372 ControlPacket::Handshake(HandshakePacket {
373 timestamp: 0,
374 dest_socket_id: 0,
375 version: crate::handshake_sm::HANDSHAKE_VERSION_4,
376 encryption_field: EncryptionField::NoEncryption,
377 extension_field: HandshakeExtensionFlags(2),
378 initial_seq_number: 0,
379 mtu: 1500,
380 max_flow_window_size: 8192,
381 handshake_type: HandshakeType::Induction,
382 srt_socket_id: caller_id,
383 syn_cookie: 0,
384 peer_ip: [0; 4],
385 extensions: HandshakeExtensions(&[]),
386 })
387 }
388
389 #[test]
390 fn induction_response_wire_values_match_draft_4_3_1_1() {
391 let mut l = ListenerHandshake::new(0x9999, 0xC0FF_EE00, HandshakeConfig::default());
392 let outputs = l.feed(&caller_induction(0x1234)).unwrap();
393 assert_eq!(outputs.len(), 1);
394 let bytes = match &outputs[0] {
395 HandshakeOutput::Send(b) => b.clone(),
396 other => panic!("expected Send, got {other:?}"),
397 };
398 let pkt = ControlPacket::parse(&bytes).unwrap();
399 match pkt {
400 ControlPacket::Handshake(hp) => {
401 assert_eq!(hp.version, HANDSHAKE_VERSION_5);
402 assert_eq!(hp.extension_field.0, SRT_MAGIC_CODE);
403 assert_eq!(hp.handshake_type, HandshakeType::Induction);
404 assert_eq!(hp.srt_socket_id, 0x9999);
405 assert_eq!(hp.syn_cookie, 0xC0FF_EE00);
406 assert_eq!(hp.dest_socket_id, 0x1234);
407 }
408 _ => panic!("expected handshake"),
409 }
410 assert_eq!(l.state(), ListenerHandshakeState::AwaitingConclusion);
411 }
412
413 fn caller_conclusion(
414 caller_id: u32,
415 listener_id: u32,
416 cookie: u32,
417 latency_ms: u16,
418 ) -> ControlPacket<'static> {
419 let hs_msg = HsExtMessage {
420 srt_version: 0x0105_0000,
421 srt_flags: crate::packet::HandshakeExtensionMessageFlags(0x6F),
422 receiver_tsbpd_delay_ms: latency_ms,
423 sender_tsbpd_delay_ms: latency_ms,
424 };
425 let ext = crate::packet::handshake::build_extension_block(
426 ExtensionType::HsReq,
427 &hs_msg.to_bytes(),
428 )
429 .unwrap();
430 let ext: &'static [u8] = Vec::leak(ext);
431 ControlPacket::Handshake(HandshakePacket {
432 timestamp: 0,
433 dest_socket_id: listener_id,
434 version: HANDSHAKE_VERSION_5,
435 encryption_field: EncryptionField::NoEncryption,
436 extension_field: HandshakeExtensionFlags(HS_EXT_FLAG_HSREQ),
437 initial_seq_number: 0,
438 mtu: 1500,
439 max_flow_window_size: 8192,
440 handshake_type: HandshakeType::Conclusion,
441 srt_socket_id: caller_id,
442 syn_cookie: cookie,
443 peer_ip: [0; 4],
444 extensions: HandshakeExtensions(ext),
445 })
446 }
447
448 #[test]
449 fn conclusion_bad_cookie_is_rejected() {
450 let mut l = ListenerHandshake::new(1, 0xC0FF_EE00, HandshakeConfig::default());
451 l.feed(&caller_induction(2)).unwrap();
452 let outputs = l.feed(&caller_conclusion(2, 1, 0xBAD_C00C, 120)).unwrap();
453 assert_eq!(outputs.len(), 2);
454 assert_eq!(
455 outputs[1],
456 HandshakeOutput::Rejected(RejectionReason::Rogue)
457 );
458 let bytes = match &outputs[0] {
459 HandshakeOutput::Send(b) => b,
460 other => panic!("expected Send, got {other:?}"),
461 };
462 let pkt = ControlPacket::parse(bytes).unwrap();
463 if let ControlPacket::Handshake(hp) = pkt {
464 assert_eq!(
465 hp.handshake_type,
466 RejectionReason::Rogue.to_handshake_type()
467 );
468 } else {
469 panic!("expected handshake");
470 }
471 assert_eq!(l.state(), ListenerHandshakeState::Rejected);
472 }
473
474 #[test]
475 fn conclusion_version_mismatch_is_rejected() {
476 let mut l = ListenerHandshake::new(1, 0xC0FF_EE00, HandshakeConfig::default());
477 l.feed(&caller_induction(2)).unwrap();
478 let mut bad = caller_conclusion(2, 1, 0xC0FF_EE00, 120);
479 if let ControlPacket::Handshake(hp) = &mut bad {
480 hp.version = 4;
481 }
482 let outputs = l.feed(&bad).unwrap();
483 assert_eq!(
484 outputs[1],
485 HandshakeOutput::Rejected(RejectionReason::Version)
486 );
487 assert_eq!(l.state(), ListenerHandshakeState::Rejected);
488 }
489
490 #[test]
491 fn conclusion_malformed_extension_is_rejected_not_panicking() {
492 let mut l = ListenerHandshake::new(1, 0xC0FF_EE00, HandshakeConfig::default());
493 l.feed(&caller_induction(2)).unwrap();
494 let bad_ext: &'static [u8] = &[0x00, 0x01, 0xFF, 0xFF];
497 let bad = ControlPacket::Handshake(HandshakePacket {
498 timestamp: 0,
499 dest_socket_id: 1,
500 version: HANDSHAKE_VERSION_5,
501 encryption_field: EncryptionField::NoEncryption,
502 extension_field: HandshakeExtensionFlags(HS_EXT_FLAG_HSREQ),
503 initial_seq_number: 0,
504 mtu: 1500,
505 max_flow_window_size: 8192,
506 handshake_type: HandshakeType::Conclusion,
507 srt_socket_id: 2,
508 syn_cookie: 0xC0FF_EE00,
509 peer_ip: [0; 4],
510 extensions: HandshakeExtensions(bad_ext),
511 });
512 let outputs = l.feed(&bad).unwrap();
513 assert_eq!(
514 outputs[1],
515 HandshakeOutput::Rejected(RejectionReason::Rogue)
516 );
517 assert_eq!(l.state(), ListenerHandshakeState::Rejected);
518 }
519
520 #[test]
521 fn successful_conclusion_reaches_connected() {
522 let mut l = ListenerHandshake::new(1, 0xC0FF_EE00, HandshakeConfig::default());
523 l.feed(&caller_induction(2)).unwrap();
524 let outputs = l.feed(&caller_conclusion(2, 1, 0xC0FF_EE00, 120)).unwrap();
525 assert_eq!(outputs.len(), 2);
526 assert!(matches!(outputs[0], HandshakeOutput::Send(_)));
527 assert!(matches!(outputs[1], HandshakeOutput::Connected(_)));
528 assert_eq!(l.state(), ListenerHandshakeState::Connected);
529 assert!(l.negotiated().is_some());
530 }
531
532 #[test]
533 fn feed_before_induction_seen_still_requires_induction_first() {
534 let mut l = ListenerHandshake::new(1, 1, HandshakeConfig::default());
535 let outputs = l.feed(&caller_induction(2));
536 assert!(outputs.is_ok());
537 }
538
539 #[test]
543 fn cif_fixed_len_is_the_documented_48_bytes() {
544 assert_eq!(HANDSHAKE_CIF_FIXED_LEN, 48);
545 }
546}