1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
40#[cfg_attr(feature = "serde", derive(serde::Serialize))]
41#[non_exhaustive]
42pub enum CallerHandshakeState {
43 Idle,
45 AwaitingInductionResponse,
47 AwaitingConclusionResponse,
49 Connected,
51 Rejected,
53 TimedOut,
55}
56
57impl CallerHandshakeState {
58 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#[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 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 pub fn state(&self) -> CallerHandshakeState {
109 self.state
110 }
111
112 pub fn negotiated(&self) -> Option<&NegotiatedParams> {
114 self.negotiated.as_ref()
115 }
116
117 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, 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 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 pub fn feed_bytes(&mut self, bytes: &[u8]) -> Result<Vec<HandshakeOutput>> {
180 let packet = ControlPacket::parse(bytes)?;
181 self.feed(&packet)
182 }
183
184 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 self.state = CallerHandshakeState::Rejected;
226 return Ok(vec![HandshakeOutput::Rejected(RejectionReason::Version)]);
227 }
228 if hp.extension_field.0 != SRT_MAGIC_CODE {
229 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 #[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 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 #[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()); let out = c.tick(); assert_eq!(out, vec![HandshakeOutput::Send(first.clone())]);
526
527 assert_eq!(c.tick(), Vec::new());
528 let out = c.tick(); assert_eq!(out, vec![HandshakeOutput::TimedOut]);
530 assert_eq!(c.state(), CallerHandshakeState::TimedOut);
531 }
532}