use alloc::vec;
use alloc::vec::Vec;
use crate::error::{Error, Result};
use crate::handshake_sm::{
self, HANDSHAKE_VERSION_5, HandshakeConfig, HandshakeOutput, NegotiatedParams, RejectionReason,
};
use crate::packet::{
ControlPacket, EncryptionField, ExtensionType, HandshakeExtensionFlags,
HandshakeExtensionMessageFlags, HandshakeExtensions, HandshakePacket, HandshakeType,
HsExtMessage,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[non_exhaustive]
pub enum RendezvousHandshakeState {
Idle,
Waving,
Attention,
Initiated,
Connected,
Rejected,
TimedOut,
}
impl RendezvousHandshakeState {
pub fn name(&self) -> &'static str {
match self {
RendezvousHandshakeState::Idle => "Idle",
RendezvousHandshakeState::Waving => "Waving",
RendezvousHandshakeState::Attention => "Attention",
RendezvousHandshakeState::Initiated => "Initiated",
RendezvousHandshakeState::Connected => "Connected",
RendezvousHandshakeState::Rejected => "Rejected",
RendezvousHandshakeState::TimedOut => "TimedOut",
}
}
}
impl core::fmt::Display for RendezvousHandshakeState {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str(self.name())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[non_exhaustive]
pub enum RendezvousRole {
Initiator,
Responder,
}
impl RendezvousRole {
pub fn name(&self) -> &'static str {
match self {
RendezvousRole::Initiator => "Initiator",
RendezvousRole::Responder => "Responder",
}
}
}
broadcast_common::impl_spec_display!(RendezvousRole);
enum PeerHsExt {
None,
HsReq(HsExtMessage),
HsRsp(HsExtMessage),
}
fn parse_peer_hs_ext(hp: &HandshakePacket<'_>) -> Result<PeerHsExt> {
let mut found = PeerHsExt::None;
for block in hp.extensions.iter() {
let block = block.map_err(|_| Error::InvalidField {
what: "rendezvous handshake extensions",
reason: "malformed extension block",
})?;
match block.ext_type {
ExtensionType::HsReq => {
let msg = block.as_hs_ext_message().map_err(|_| Error::InvalidField {
what: "HSREQ extension message",
reason: "malformed contents",
})?;
found = PeerHsExt::HsReq(msg);
}
ExtensionType::HsRsp => {
let msg = block.as_hs_ext_message().map_err(|_| Error::InvalidField {
what: "HSRSP extension message",
reason: "malformed contents",
})?;
found = PeerHsExt::HsRsp(msg);
}
_ => {}
}
}
Ok(found)
}
#[derive(Debug)]
pub struct RendezvousHandshake {
own_socket_id: u32,
own_cookie: u32,
config: HandshakeConfig,
state: RendezvousHandshakeState,
role: Option<RendezvousRole>,
peer_socket_id: u32,
peer_cookie: u32,
peer_hs_msg: Option<HsExtMessage>,
last_sent: Option<Vec<u8>>,
ticks_since_send: u32,
retries: u32,
negotiated: Option<NegotiatedParams>,
}
impl RendezvousHandshake {
pub fn new(own_socket_id: u32, own_cookie: u32, config: HandshakeConfig) -> Self {
RendezvousHandshake {
own_socket_id,
own_cookie,
config,
state: RendezvousHandshakeState::Idle,
role: None,
peer_socket_id: 0,
peer_cookie: 0,
peer_hs_msg: None,
last_sent: None,
ticks_since_send: 0,
retries: 0,
negotiated: None,
}
}
pub fn state(&self) -> RendezvousHandshakeState {
self.state
}
pub fn role(&self) -> Option<RendezvousRole> {
self.role
}
pub fn negotiated(&self) -> Option<&NegotiatedParams> {
self.negotiated.as_ref()
}
pub fn start(&mut self) -> Result<Vec<u8>> {
if self.state != RendezvousHandshakeState::Idle {
return Err(Error::HandshakeOutOfSequence {
state: self.state.name(),
reason: "start() called after the handshake already began",
});
}
let hp = HandshakePacket {
timestamp: 0,
dest_socket_id: 0, version: HANDSHAKE_VERSION_5,
encryption_field: self.config.encryption_field,
extension_field: HandshakeExtensionFlags(0),
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::Wavehand,
srt_socket_id: self.own_socket_id,
syn_cookie: self.own_cookie,
peer_ip: self.config.local_ip,
extensions: HandshakeExtensions(&[]),
};
let bytes = handshake_sm::build_bytes(hp)?;
self.last_sent = Some(bytes.clone());
self.ticks_since_send = 0;
self.state = RendezvousHandshakeState::Waving;
Ok(bytes)
}
pub fn feed(&mut self, packet: &ControlPacket<'_>) -> Result<Vec<HandshakeOutput>> {
let hp = match packet {
ControlPacket::Handshake(hp) => hp,
other => {
if self.state == RendezvousHandshakeState::Initiated
&& self.role == Some(RendezvousRole::Responder)
{
return self.enter_connected();
}
return Err(Error::UnexpectedControlPacket {
actual: other.control_type().name(),
});
}
};
if let Some(reason) = RejectionReason::from_handshake_type(hp.handshake_type) {
return Ok(self.reject(reason));
}
if hp.version != HANDSHAKE_VERSION_5 {
return Ok(self.reject(RejectionReason::Version));
}
match self.state {
RendezvousHandshakeState::Idle => Err(Error::HandshakeOutOfSequence {
state: self.state.name(),
reason: "feed() called before start()",
}),
RendezvousHandshakeState::Waving => self.on_waving(hp),
RendezvousHandshakeState::Attention => self.on_attention(hp),
RendezvousHandshakeState::Initiated => self.on_initiated(hp),
RendezvousHandshakeState::Connected => self.on_connected(hp),
RendezvousHandshakeState::Rejected | RendezvousHandshakeState::TimedOut => {
Err(Error::HandshakeOutOfSequence {
state: self.state.name(),
reason: "handshake already reached a terminal state",
})
}
}
}
pub fn feed_bytes(&mut self, bytes: &[u8]) -> Result<Vec<HandshakeOutput>> {
let packet = ControlPacket::parse(bytes)?;
self.feed(&packet)
}
pub fn on_recovery_trigger(&mut self) -> Result<Vec<HandshakeOutput>> {
if self.state == RendezvousHandshakeState::Initiated
&& self.role == Some(RendezvousRole::Responder)
{
self.enter_connected()
} else {
Ok(Vec::new())
}
}
pub fn tick(&mut self) -> Vec<HandshakeOutput> {
if matches!(
self.state,
RendezvousHandshakeState::Idle
| RendezvousHandshakeState::Connected
| RendezvousHandshakeState::Rejected
| RendezvousHandshakeState::TimedOut
) {
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 = RendezvousHandshakeState::TimedOut;
return vec![HandshakeOutput::TimedOut];
}
match self.last_sent.clone() {
Some(bytes) => vec![HandshakeOutput::Send(bytes)],
None => Vec::new(),
}
}
fn on_waving(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
if !matches!(
hp.handshake_type,
HandshakeType::Wavehand | HandshakeType::Conclusion
) {
return Ok(self.reject(RejectionReason::Rogue));
}
self.peer_socket_id = hp.srt_socket_id;
let role = match self.resolve_role(hp.syn_cookie) {
Ok(r) => r,
Err(_) => return Ok(self.reject(RejectionReason::RdvCookie)),
};
self.peer_cookie = hp.syn_cookie;
self.role = Some(role);
if hp.handshake_type == HandshakeType::Wavehand {
self.state = RendezvousHandshakeState::Attention;
match role {
RendezvousRole::Initiator => self.send_conclusion(Some(ExtensionType::HsReq)),
RendezvousRole::Responder => self.send_conclusion(None),
}
} else {
self.on_attention_conclusion(hp, role)
}
}
fn on_attention(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
let role = self
.role
.expect("role is always resolved before Attention is reached");
if hp.handshake_type == HandshakeType::Conclusion {
return self.on_attention_conclusion(hp, role);
}
if hp.handshake_type == HandshakeType::Wavehand {
return Ok(self.resend());
}
Ok(self.reject(RejectionReason::Rogue))
}
fn on_attention_conclusion(
&mut self,
hp: &HandshakePacket<'_>,
role: RendezvousRole,
) -> Result<Vec<HandshakeOutput>> {
self.peer_socket_id = hp.srt_socket_id;
let ext = match parse_peer_hs_ext(hp) {
Ok(e) => e,
Err(_) => return Ok(self.reject(RejectionReason::Rogue)),
};
match (role, ext) {
(RendezvousRole::Initiator, PeerHsExt::None) => {
self.state = RendezvousHandshakeState::Initiated;
self.send_conclusion(Some(ExtensionType::HsReq))
}
(RendezvousRole::Initiator, PeerHsExt::HsRsp(msg)) => {
self.peer_hs_msg = Some(msg);
self.enter_connected()
}
(RendezvousRole::Responder, PeerHsExt::None) => {
self.state = RendezvousHandshakeState::Attention;
self.send_conclusion(None)
}
(RendezvousRole::Responder, PeerHsExt::HsReq(msg)) => {
self.peer_hs_msg = Some(msg);
self.state = RendezvousHandshakeState::Initiated;
self.send_conclusion(Some(ExtensionType::HsRsp))
}
_ => Ok(self.reject(RejectionReason::Rogue)),
}
}
fn on_initiated(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
let role = self
.role
.expect("role is always resolved before Initiated is reached");
match role {
RendezvousRole::Initiator => {
if hp.handshake_type != HandshakeType::Conclusion {
return Ok(self.reject(RejectionReason::Rogue));
}
let ext = match parse_peer_hs_ext(hp) {
Ok(e) => e,
Err(_) => return Ok(self.reject(RejectionReason::Rogue)),
};
match ext {
PeerHsExt::None => self.send_conclusion(Some(ExtensionType::HsReq)),
PeerHsExt::HsRsp(msg) => {
self.peer_hs_msg = Some(msg);
self.enter_connected()
}
PeerHsExt::HsReq(_) => Ok(self.reject(RejectionReason::Rogue)),
}
}
RendezvousRole::Responder => {
if hp.handshake_type == HandshakeType::Agreement {
return self.enter_connected();
}
if hp.handshake_type != HandshakeType::Conclusion {
return Ok(self.reject(RejectionReason::Rogue));
}
let ext = match parse_peer_hs_ext(hp) {
Ok(e) => e,
Err(_) => return Ok(self.reject(RejectionReason::Rogue)),
};
match ext {
PeerHsExt::HsReq(msg) => {
self.peer_hs_msg = Some(msg);
self.send_conclusion(Some(ExtensionType::HsRsp))
}
_ => Ok(self.reject(RejectionReason::Rogue)),
}
}
}
}
fn on_connected(&mut self, hp: &HandshakePacket<'_>) -> Result<Vec<HandshakeOutput>> {
if hp.handshake_type == HandshakeType::Conclusion {
return self.send_agreement();
}
Ok(Vec::new())
}
fn resolve_role(&self, peer_cookie: u32) -> Result<RendezvousRole> {
if peer_cookie == self.own_cookie {
return Err(Error::InvalidField {
what: "rendezvous cookie",
reason: "identical to the peer's cookie (collision, L2119-2124)",
});
}
Ok(if self.own_cookie > peer_cookie {
RendezvousRole::Initiator
} else {
RendezvousRole::Responder
})
}
fn reject(&mut self, reason: RejectionReason) -> Vec<HandshakeOutput> {
self.state = RendezvousHandshakeState::Rejected;
vec![HandshakeOutput::Rejected(reason)]
}
fn resend(&mut self) -> Vec<HandshakeOutput> {
match &self.last_sent {
Some(bytes) => vec![HandshakeOutput::Send(bytes.clone())],
None => Vec::new(),
}
}
fn send_conclusion(&mut self, ext_type: Option<ExtensionType>) -> Result<Vec<HandshakeOutput>> {
let (ext_bytes, ext_flags): (Vec<u8>, u16) = match ext_type {
None => (Vec::new(), 0),
Some(t) => {
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,
};
handshake_sm::build_conclusion_extensions(t, &hs_msg, None, None)?
}
};
let hp = 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: self.own_cookie,
peer_ip: self.config.local_ip,
extensions: HandshakeExtensions(&ext_bytes),
};
let bytes = handshake_sm::build_bytes(hp)?;
self.last_sent = Some(bytes.clone());
self.ticks_since_send = 0;
self.retries = 0;
Ok(vec![HandshakeOutput::Send(bytes)])
}
fn send_agreement(&mut self) -> Result<Vec<HandshakeOutput>> {
let hp = HandshakePacket {
timestamp: 0,
dest_socket_id: self.peer_socket_id,
version: HANDSHAKE_VERSION_5,
encryption_field: EncryptionField::NoEncryption,
extension_field: HandshakeExtensionFlags(0),
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::Agreement,
srt_socket_id: self.own_socket_id,
syn_cookie: self.own_cookie,
peer_ip: self.config.local_ip,
extensions: HandshakeExtensions(&[]),
};
let bytes = handshake_sm::build_bytes(hp)?;
self.last_sent = Some(bytes.clone());
self.ticks_since_send = 0;
self.retries = 0;
Ok(vec![HandshakeOutput::Send(bytes)])
}
fn enter_connected(&mut self) -> Result<Vec<HandshakeOutput>> {
let negotiated = self.build_negotiated();
self.negotiated = Some(negotiated.clone());
self.state = RendezvousHandshakeState::Connected;
let mut out = self.send_agreement()?;
out.push(HandshakeOutput::Connected(negotiated));
Ok(out)
}
fn build_negotiated(&self) -> NegotiatedParams {
let peer_msg = self
.peer_hs_msg
.expect("peer_hs_msg is always captured before any transition reaches Connected");
NegotiatedParams {
version: HANDSHAKE_VERSION_5,
flags: 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: None,
group: None,
#[cfg(feature = "crypto")]
sek: None,
#[cfg(feature = "crypto")]
salt: None,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::packet::handshake::HS_EXT_FLAG_HSREQ;
#[test]
fn start_is_idempotent_guard() {
let mut r = RendezvousHandshake::new(1, 500, HandshakeConfig::default());
assert!(r.start().is_ok());
assert!(r.start().is_err());
}
#[test]
fn wavehand_wire_values_match_draft_4_3_2() {
let mut r = RendezvousHandshake::new(0xAAAA_BBBB, 0xC0FF_EE00, HandshakeConfig::default());
let bytes = r.start().unwrap();
let pkt = ControlPacket::parse(&bytes).unwrap();
match pkt {
ControlPacket::Handshake(hp) => {
assert_eq!(hp.version, HANDSHAKE_VERSION_5);
assert_eq!(hp.handshake_type, HandshakeType::Wavehand);
assert_eq!(hp.srt_socket_id, 0xAAAA_BBBB);
assert_eq!(hp.syn_cookie, 0xC0FF_EE00);
assert_eq!(hp.extension_field.0, 0);
}
_ => panic!("expected handshake"),
}
assert_eq!(r.state(), RendezvousHandshakeState::Waving);
}
fn wavehand(socket_id: u32, cookie: u32) -> ControlPacket<'static> {
ControlPacket::Handshake(HandshakePacket {
timestamp: 0,
dest_socket_id: 0,
version: HANDSHAKE_VERSION_5,
encryption_field: EncryptionField::NoEncryption,
extension_field: HandshakeExtensionFlags(0),
initial_seq_number: 0,
mtu: 1500,
max_flow_window_size: 8192,
handshake_type: HandshakeType::Wavehand,
srt_socket_id: socket_id,
syn_cookie: cookie,
peer_ip: [0; 4],
extensions: HandshakeExtensions(&[]),
})
}
#[test]
fn greater_cookie_wins_initiator() {
let mut a = RendezvousHandshake::new(1, 500, HandshakeConfig::default());
let mut b = RendezvousHandshake::new(2, 100, HandshakeConfig::default());
a.start().unwrap();
b.start().unwrap();
a.feed(&wavehand(2, 100)).unwrap();
b.feed(&wavehand(1, 500)).unwrap();
assert_eq!(a.role(), Some(RendezvousRole::Initiator));
assert_eq!(b.role(), Some(RendezvousRole::Responder));
assert_eq!(a.state(), RendezvousHandshakeState::Attention);
assert_eq!(b.state(), RendezvousHandshakeState::Attention);
}
#[test]
fn identical_cookies_are_rejected_as_a_collision() {
let mut a = RendezvousHandshake::new(1, 0x00C0_FFEE, HandshakeConfig::default());
a.start().unwrap();
let outputs = a.feed(&wavehand(2, 0x00C0_FFEE)).unwrap();
assert_eq!(
outputs,
vec![HandshakeOutput::Rejected(RejectionReason::RdvCookie)]
);
assert_eq!(a.state(), RendezvousHandshakeState::Rejected);
}
#[test]
fn initiator_attention_entry_sends_hsreq() {
let mut a = RendezvousHandshake::new(1, 500, HandshakeConfig::default());
a.start().unwrap();
let outputs = a.feed(&wavehand(2, 100)).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.handshake_type, HandshakeType::Conclusion);
assert_eq!(hp.extension_field.0 & HS_EXT_FLAG_HSREQ, HS_EXT_FLAG_HSREQ);
let blocks: Vec<_> = hp.extensions.iter().map(|b| b.unwrap()).collect();
assert_eq!(blocks.len(), 1);
assert_eq!(blocks[0].ext_type, ExtensionType::HsReq);
}
_ => panic!("expected handshake"),
}
}
#[test]
fn responder_attention_entry_sends_empty_conclusion() {
let mut b = RendezvousHandshake::new(2, 100, HandshakeConfig::default());
b.start().unwrap();
let outputs = b.feed(&wavehand(1, 500)).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.handshake_type, HandshakeType::Conclusion);
assert_eq!(hp.extension_field.0, 0);
assert_eq!(hp.extensions.iter().count(), 0);
}
_ => panic!("expected handshake"),
}
}
#[test]
fn malformed_extension_mid_flow_is_rejected_not_panicking() {
let mut a = RendezvousHandshake::new(1, 500, HandshakeConfig::default());
a.start().unwrap();
a.feed(&wavehand(2, 100)).unwrap();
assert_eq!(a.state(), RendezvousHandshakeState::Attention);
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: 100,
peer_ip: [0; 4],
extensions: HandshakeExtensions(bad_ext),
});
let outputs = a.feed(&bad).unwrap();
assert_eq!(
outputs,
vec![HandshakeOutput::Rejected(RejectionReason::Rogue)]
);
assert_eq!(a.state(), RendezvousHandshakeState::Rejected);
}
#[test]
fn feed_before_start_is_out_of_sequence() {
let mut r = RendezvousHandshake::new(1, 500, HandshakeConfig::default());
assert!(matches!(
r.feed(&wavehand(2, 100)),
Err(Error::HandshakeOutOfSequence { .. })
));
}
#[test]
fn feed_rejects_non_handshake_packets_outside_recovery_case() {
use crate::packet::misc::KeepAlivePacket;
let mut r = RendezvousHandshake::new(1, 500, HandshakeConfig::default());
r.start().unwrap();
let ka = ControlPacket::KeepAlive(KeepAlivePacket {
timestamp: 0,
dest_socket_id: 0,
});
assert!(matches!(
r.feed(&ka),
Err(Error::UnexpectedControlPacket { .. })
));
}
}