use alloc::vec;
use alloc::vec::Vec;
use crate::error::{Error, Result};
use crate::handshake_sm::{
self, HANDSHAKE_VERSION_5, HandshakeConfig, HandshakeOutput, NegotiatedParams, RejectionReason,
SRT_MAGIC_CODE,
};
use crate::packet::{
ControlPacket, EncryptionField, ExtensionType, HandshakeExtensionFlags, HandshakeExtensions,
HandshakePacket, HandshakeType, HsExtMessage,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[non_exhaustive]
pub enum ListenerHandshakeState {
Idle,
AwaitingConclusion,
Connected,
Rejected,
TimedOut,
}
impl ListenerHandshakeState {
pub fn name(&self) -> &'static str {
match self {
ListenerHandshakeState::Idle => "Idle",
ListenerHandshakeState::AwaitingConclusion => "AwaitingConclusion",
ListenerHandshakeState::Connected => "Connected",
ListenerHandshakeState::Rejected => "Rejected",
ListenerHandshakeState::TimedOut => "TimedOut",
}
}
}
impl core::fmt::Display for ListenerHandshakeState {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str(self.name())
}
}
#[derive(Debug)]
pub struct ListenerHandshake {
own_socket_id: u32,
syn_cookie: u32,
config: HandshakeConfig,
state: ListenerHandshakeState,
peer_socket_id: u32,
last_sent: Option<Vec<u8>>,
ticks_since_send: u32,
retries: u32,
negotiated: Option<NegotiatedParams>,
}
impl ListenerHandshake {
pub fn new(own_socket_id: u32, syn_cookie: u32, config: HandshakeConfig) -> Self {
ListenerHandshake {
own_socket_id,
syn_cookie,
config,
state: ListenerHandshakeState::Idle,
peer_socket_id: 0,
last_sent: None,
ticks_since_send: 0,
retries: 0,
negotiated: None,
}
}
pub fn state(&self) -> ListenerHandshakeState {
self.state
}
pub fn negotiated(&self) -> Option<&NegotiatedParams> {
self.negotiated.as_ref()
}
pub fn feed(&mut self, packet: &ControlPacket<'_>) -> Result<Vec<HandshakeOutput>> {
let hp = match packet {
ControlPacket::Handshake(hp) => hp,
other => {
return Err(Error::UnexpectedControlPacket {
actual: other.control_type().name(),
});
}
};
match self.state {
ListenerHandshakeState::Idle => self.on_induction(hp),
ListenerHandshakeState::AwaitingConclusion => self.on_conclusion(hp),
_ => Err(Error::HandshakeOutOfSequence {
state: self.state.name(),
reason: "not awaiting an induction or conclusion",
}),
}
}
pub fn feed_bytes(&mut self, bytes: &[u8]) -> Result<Vec<HandshakeOutput>> {
let packet = ControlPacket::parse(bytes)?;
self.feed(&packet)
}
pub fn tick(&mut self) -> Vec<HandshakeOutput> {
if self.state != ListenerHandshakeState::AwaitingConclusion {
return Vec::new();
}
self.ticks_since_send += 1;
if self.ticks_since_send < self.config.retransmit_after_ticks {
return Vec::new();
}
self.ticks_since_send = 0;
self.retries += 1;
if self.retries > self.config.max_retries {
self.state = ListenerHandshakeState::TimedOut;
return vec![HandshakeOutput::TimedOut];
}
match self.last_sent.clone() {
Some(bytes) => vec![HandshakeOutput::Send(bytes)],
None => Vec::new(),
}
}
fn on_induction(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
if hp.handshake_type != HandshakeType::Induction {
return Err(Error::HandshakeOutOfSequence {
state: self.state.name(),
reason: "expected an INDUCTION handshake",
});
}
self.peer_socket_id = hp.srt_socket_id;
let hp_out = HandshakePacket {
timestamp: 0,
dest_socket_id: self.peer_socket_id,
version: HANDSHAKE_VERSION_5,
encryption_field: self.config.encryption_field,
extension_field: HandshakeExtensionFlags(SRT_MAGIC_CODE),
initial_seq_number: self.config.initial_seq_number,
mtu: self.config.mtu,
max_flow_window_size: self.config.max_flow_window_size,
handshake_type: HandshakeType::Induction,
srt_socket_id: self.own_socket_id,
syn_cookie: self.syn_cookie,
peer_ip: self.config.local_ip,
extensions: HandshakeExtensions(&[]),
};
let bytes = handshake_sm::build_bytes(hp_out)?;
self.last_sent = Some(bytes.clone());
self.ticks_since_send = 0;
self.state = ListenerHandshakeState::AwaitingConclusion;
Ok(vec![HandshakeOutput::Send(bytes)])
}
fn on_conclusion(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
if hp.handshake_type != HandshakeType::Conclusion {
return self.reject(RejectionReason::Rogue, hp);
}
if hp.version != HANDSHAKE_VERSION_5 {
return self.reject(RejectionReason::Version, hp);
}
if hp.syn_cookie != self.syn_cookie {
return self.reject(RejectionReason::Rogue, hp);
}
let parsed = match handshake_sm::parse_peer_extensions(hp) {
Ok(p) => p,
Err(_) => return self.reject(RejectionReason::Rogue, hp),
};
let peer_msg = match parsed.hs_msg {
Some(m) => m,
None => return self.reject(RejectionReason::Rogue, hp),
};
self.peer_socket_id = hp.srt_socket_id;
#[cfg(feature = "crypto")]
let crypto_cfg = self.config.crypto.clone();
#[cfg(feature = "crypto")]
type CryptoConclusionResult = (
Option<handshake_sm::RecoveredSek>,
Option<alloc::vec::Vec<u8>>,
);
#[cfg(feature = "crypto")]
let (crypto_result, km_echo_ext): CryptoConclusionResult = match (&crypto_cfg, &parsed.km) {
(Some(crypto), Some(km_req)) => {
match handshake_sm::recover_sek(&crypto.passphrase, km_req) {
Ok((sek, salt)) => {
let echo = match handshake_sm::echo_key_material_as_response(km_req) {
Ok(e) => e,
Err(_) => return self.reject(RejectionReason::Rogue, hp),
};
(Some((sek, salt)), Some(echo))
}
Err(_) => return self.reject(RejectionReason::BadSecret, hp),
}
}
(Some(_), None) | (None, Some(_)) => {
return self.reject(RejectionReason::Unsecure, hp);
}
(None, None) => (None, None),
};
let negotiated = NegotiatedParams {
version: HANDSHAKE_VERSION_5,
flags: crate::packet::HandshakeExtensionMessageFlags(
self.config.flags.0 & peer_msg.srt_flags.0,
),
latency_ms: handshake_sm::negotiate_latency_ms(self.config.latency_ms, &peer_msg),
own_socket_id: self.own_socket_id,
peer_socket_id: self.peer_socket_id,
stream_id: parsed.stream_id,
group: parsed.group,
#[cfg(feature = "crypto")]
sek: crypto_result.as_ref().map(|(sek, _)| sek.clone()),
#[cfg(feature = "crypto")]
salt: crypto_result.as_ref().map(|(_, salt)| *salt),
};
let hs_msg = HsExtMessage {
srt_version: self.config.srt_version,
srt_flags: self.config.flags,
receiver_tsbpd_delay_ms: self.config.latency_ms,
sender_tsbpd_delay_ms: self.config.latency_ms,
};
let (ext_bytes, ext_flags) = handshake_sm::build_conclusion_extensions(
ExtensionType::HsRsp,
&hs_msg,
None,
self.config.group,
)?;
#[cfg(feature = "crypto")]
let (ext_bytes, ext_flags) = {
let mut ext_bytes = ext_bytes;
let mut ext_flags = ext_flags;
if let Some(echo) = km_echo_ext {
ext_bytes.extend(echo);
ext_flags |= crate::packet::handshake::HS_EXT_FLAG_KMREQ;
}
(ext_bytes, ext_flags)
};
let hp_out = HandshakePacket {
timestamp: 0,
dest_socket_id: self.peer_socket_id,
version: HANDSHAKE_VERSION_5,
encryption_field: self.config.encryption_field,
extension_field: HandshakeExtensionFlags(ext_flags),
initial_seq_number: self.config.initial_seq_number,
mtu: self.config.mtu,
max_flow_window_size: self.config.max_flow_window_size,
handshake_type: HandshakeType::Conclusion,
srt_socket_id: self.own_socket_id,
syn_cookie: 0, peer_ip: self.config.local_ip,
extensions: HandshakeExtensions(&ext_bytes),
};
let bytes = handshake_sm::build_bytes(hp_out)?;
self.last_sent = Some(bytes.clone());
self.negotiated = Some(negotiated.clone());
self.state = ListenerHandshakeState::Connected;
Ok(vec![
HandshakeOutput::Send(bytes),
HandshakeOutput::Connected(negotiated),
])
}
fn reject(
&mut self,
reason: RejectionReason,
hp: &HandshakePacket<'_>,
) -> Result<Vec<HandshakeOutput>> {
self.state = ListenerHandshakeState::Rejected;
let hp_out = HandshakePacket {
timestamp: 0,
dest_socket_id: hp.srt_socket_id,
version: HANDSHAKE_VERSION_5,
encryption_field: EncryptionField::NoEncryption,
extension_field: HandshakeExtensionFlags(0),
initial_seq_number: 0,
mtu: self.config.mtu,
max_flow_window_size: self.config.max_flow_window_size,
handshake_type: reason.to_handshake_type(),
srt_socket_id: self.own_socket_id,
syn_cookie: 0,
peer_ip: self.config.local_ip,
extensions: HandshakeExtensions(&[]),
};
let bytes = handshake_sm::build_bytes(hp_out)?;
self.last_sent = Some(bytes.clone());
Ok(vec![
HandshakeOutput::Send(bytes),
HandshakeOutput::Rejected(reason),
])
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::packet::handshake::{HANDSHAKE_CIF_FIXED_LEN, HS_EXT_FLAG_HSREQ};
fn caller_induction(caller_id: u32) -> ControlPacket<'static> {
ControlPacket::Handshake(HandshakePacket {
timestamp: 0,
dest_socket_id: 0,
version: crate::handshake_sm::HANDSHAKE_VERSION_4,
encryption_field: EncryptionField::NoEncryption,
extension_field: HandshakeExtensionFlags(2),
initial_seq_number: 0,
mtu: 1500,
max_flow_window_size: 8192,
handshake_type: HandshakeType::Induction,
srt_socket_id: caller_id,
syn_cookie: 0,
peer_ip: [0; 4],
extensions: HandshakeExtensions(&[]),
})
}
#[test]
fn induction_response_wire_values_match_draft_4_3_1_1() {
let mut l = ListenerHandshake::new(0x9999, 0xC0FF_EE00, HandshakeConfig::default());
let outputs = l.feed(&caller_induction(0x1234)).unwrap();
assert_eq!(outputs.len(), 1);
let bytes = match &outputs[0] {
HandshakeOutput::Send(b) => b.clone(),
other => panic!("expected Send, got {other:?}"),
};
let pkt = ControlPacket::parse(&bytes).unwrap();
match pkt {
ControlPacket::Handshake(hp) => {
assert_eq!(hp.version, HANDSHAKE_VERSION_5);
assert_eq!(hp.extension_field.0, SRT_MAGIC_CODE);
assert_eq!(hp.handshake_type, HandshakeType::Induction);
assert_eq!(hp.srt_socket_id, 0x9999);
assert_eq!(hp.syn_cookie, 0xC0FF_EE00);
assert_eq!(hp.dest_socket_id, 0x1234);
}
_ => panic!("expected handshake"),
}
assert_eq!(l.state(), ListenerHandshakeState::AwaitingConclusion);
}
fn caller_conclusion(
caller_id: u32,
listener_id: u32,
cookie: u32,
latency_ms: u16,
) -> ControlPacket<'static> {
let hs_msg = HsExtMessage {
srt_version: 0x0105_0000,
srt_flags: crate::packet::HandshakeExtensionMessageFlags(0x6F),
receiver_tsbpd_delay_ms: latency_ms,
sender_tsbpd_delay_ms: latency_ms,
};
let ext = crate::packet::handshake::build_extension_block(
ExtensionType::HsReq,
&hs_msg.to_bytes(),
)
.unwrap();
let ext: &'static [u8] = Vec::leak(ext);
ControlPacket::Handshake(HandshakePacket {
timestamp: 0,
dest_socket_id: listener_id,
version: HANDSHAKE_VERSION_5,
encryption_field: EncryptionField::NoEncryption,
extension_field: HandshakeExtensionFlags(HS_EXT_FLAG_HSREQ),
initial_seq_number: 0,
mtu: 1500,
max_flow_window_size: 8192,
handshake_type: HandshakeType::Conclusion,
srt_socket_id: caller_id,
syn_cookie: cookie,
peer_ip: [0; 4],
extensions: HandshakeExtensions(ext),
})
}
#[test]
fn conclusion_bad_cookie_is_rejected() {
let mut l = ListenerHandshake::new(1, 0xC0FF_EE00, HandshakeConfig::default());
l.feed(&caller_induction(2)).unwrap();
let outputs = l.feed(&caller_conclusion(2, 1, 0xBAD_C00C, 120)).unwrap();
assert_eq!(outputs.len(), 2);
assert_eq!(
outputs[1],
HandshakeOutput::Rejected(RejectionReason::Rogue)
);
let bytes = match &outputs[0] {
HandshakeOutput::Send(b) => b,
other => panic!("expected Send, got {other:?}"),
};
let pkt = ControlPacket::parse(bytes).unwrap();
if let ControlPacket::Handshake(hp) = pkt {
assert_eq!(
hp.handshake_type,
RejectionReason::Rogue.to_handshake_type()
);
} else {
panic!("expected handshake");
}
assert_eq!(l.state(), ListenerHandshakeState::Rejected);
}
#[test]
fn conclusion_version_mismatch_is_rejected() {
let mut l = ListenerHandshake::new(1, 0xC0FF_EE00, HandshakeConfig::default());
l.feed(&caller_induction(2)).unwrap();
let mut bad = caller_conclusion(2, 1, 0xC0FF_EE00, 120);
if let ControlPacket::Handshake(hp) = &mut bad {
hp.version = 4;
}
let outputs = l.feed(&bad).unwrap();
assert_eq!(
outputs[1],
HandshakeOutput::Rejected(RejectionReason::Version)
);
assert_eq!(l.state(), ListenerHandshakeState::Rejected);
}
#[test]
fn conclusion_malformed_extension_is_rejected_not_panicking() {
let mut l = ListenerHandshake::new(1, 0xC0FF_EE00, HandshakeConfig::default());
l.feed(&caller_induction(2)).unwrap();
let bad_ext: &'static [u8] = &[0x00, 0x01, 0xFF, 0xFF];
let bad = ControlPacket::Handshake(HandshakePacket {
timestamp: 0,
dest_socket_id: 1,
version: HANDSHAKE_VERSION_5,
encryption_field: EncryptionField::NoEncryption,
extension_field: HandshakeExtensionFlags(HS_EXT_FLAG_HSREQ),
initial_seq_number: 0,
mtu: 1500,
max_flow_window_size: 8192,
handshake_type: HandshakeType::Conclusion,
srt_socket_id: 2,
syn_cookie: 0xC0FF_EE00,
peer_ip: [0; 4],
extensions: HandshakeExtensions(bad_ext),
});
let outputs = l.feed(&bad).unwrap();
assert_eq!(
outputs[1],
HandshakeOutput::Rejected(RejectionReason::Rogue)
);
assert_eq!(l.state(), ListenerHandshakeState::Rejected);
}
#[test]
fn successful_conclusion_reaches_connected() {
let mut l = ListenerHandshake::new(1, 0xC0FF_EE00, HandshakeConfig::default());
l.feed(&caller_induction(2)).unwrap();
let outputs = l.feed(&caller_conclusion(2, 1, 0xC0FF_EE00, 120)).unwrap();
assert_eq!(outputs.len(), 2);
assert!(matches!(outputs[0], HandshakeOutput::Send(_)));
assert!(matches!(outputs[1], HandshakeOutput::Connected(_)));
assert_eq!(l.state(), ListenerHandshakeState::Connected);
assert!(l.negotiated().is_some());
}
#[test]
fn feed_before_induction_seen_still_requires_induction_first() {
let mut l = ListenerHandshake::new(1, 1, HandshakeConfig::default());
let outputs = l.feed(&caller_induction(2));
assert!(outputs.is_ok());
}
#[test]
fn cif_fixed_len_is_the_documented_48_bytes() {
assert_eq!(HANDSHAKE_CIF_FIXED_LEN, 48);
}
}