use alloc::string::String;
use alloc::vec::Vec;
use crate::error::{Error, Result};
#[cfg(feature = "crypto")]
use crate::packet::KeyMaterial;
use crate::packet::handshake::{HS_EXT_FLAG_CONFIG, HS_EXT_FLAG_HSREQ};
#[cfg(feature = "crypto")]
use crate::packet::{Cipher, KmAuth, KmKeyFlag, StreamEncapsulation};
use crate::packet::{
ControlPacket, EncryptionField, GroupMembershipExtension, HandshakeExtensionMessageFlags,
HandshakePacket, HandshakeType, HsExtMessage,
};
pub const HANDSHAKE_VERSION_4: u32 = 4;
pub const HANDSHAKE_VERSION_5: u32 = 5;
pub const SRT_MAGIC_CODE: u16 = 0x4A17;
pub const INDUCTION_LEGACY_SOCKET_TYPE: u16 = 2;
pub const REJECTION_CODE_BASE: u32 = 1000;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[non_exhaustive]
pub enum RejectionReason {
Unknown,
System,
Peer,
Resource,
Rogue,
Backlog,
Ipe,
Close,
Version,
RdvCookie,
BadSecret,
Unsecure,
MessageApi,
Congestion,
Filter,
Group,
Reserved(u32),
}
impl RejectionReason {
pub fn from_bits(v: u32) -> Self {
match v {
1000 => RejectionReason::Unknown,
1001 => RejectionReason::System,
1002 => RejectionReason::Peer,
1003 => RejectionReason::Resource,
1004 => RejectionReason::Rogue,
1005 => RejectionReason::Backlog,
1006 => RejectionReason::Ipe,
1007 => RejectionReason::Close,
1008 => RejectionReason::Version,
1009 => RejectionReason::RdvCookie,
1010 => RejectionReason::BadSecret,
1011 => RejectionReason::Unsecure,
1012 => RejectionReason::MessageApi,
1013 => RejectionReason::Congestion,
1014 => RejectionReason::Filter,
1015 => RejectionReason::Group,
other => RejectionReason::Reserved(other),
}
}
pub fn to_bits(self) -> u32 {
match self {
RejectionReason::Unknown => 1000,
RejectionReason::System => 1001,
RejectionReason::Peer => 1002,
RejectionReason::Resource => 1003,
RejectionReason::Rogue => 1004,
RejectionReason::Backlog => 1005,
RejectionReason::Ipe => 1006,
RejectionReason::Close => 1007,
RejectionReason::Version => 1008,
RejectionReason::RdvCookie => 1009,
RejectionReason::BadSecret => 1010,
RejectionReason::Unsecure => 1011,
RejectionReason::MessageApi => 1012,
RejectionReason::Congestion => 1013,
RejectionReason::Filter => 1014,
RejectionReason::Group => 1015,
RejectionReason::Reserved(v) => v,
}
}
pub fn name(&self) -> &'static str {
match self {
RejectionReason::Unknown => "REJ_UNKNOWN",
RejectionReason::System => "REJ_SYSTEM",
RejectionReason::Peer => "REJ_PEER",
RejectionReason::Resource => "REJ_RESOURCE",
RejectionReason::Rogue => "REJ_ROGUE",
RejectionReason::Backlog => "REJ_BACKLOG",
RejectionReason::Ipe => "REJ_IPE",
RejectionReason::Close => "REJ_CLOSE",
RejectionReason::Version => "REJ_VERSION",
RejectionReason::RdvCookie => "REJ_RDVCOOKIE",
RejectionReason::BadSecret => "REJ_BADSECRET",
RejectionReason::Unsecure => "REJ_UNSECURE",
RejectionReason::MessageApi => "REJ_MESSAGEAPI",
RejectionReason::Congestion => "REJ_CONGESTION",
RejectionReason::Filter => "REJ_FILTER",
RejectionReason::Group => "REJ_GROUP",
RejectionReason::Reserved(_) => "reserved",
}
}
pub fn from_handshake_type(ht: HandshakeType) -> Option<Self> {
match ht {
HandshakeType::Reserved(v) if v >= REJECTION_CODE_BASE => {
Some(RejectionReason::from_bits(v))
}
_ => None,
}
}
pub fn to_handshake_type(self) -> HandshakeType {
HandshakeType::from_bits(self.to_bits())
}
}
broadcast_common::impl_spec_display!(RejectionReason, Reserved);
#[cfg(feature = "crypto")]
#[cfg_attr(docsrs, doc(cfg(feature = "crypto")))]
#[derive(Debug, Clone, PartialEq)]
pub struct CryptoConfig {
pub passphrase: Vec<u8>,
pub salt: [u8; crate::crypto::SALT_LEN],
pub sek: Vec<u8>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct HandshakeConfig {
pub latency_ms: u16,
pub mtu: u32,
pub max_flow_window_size: u32,
pub initial_seq_number: u32,
pub srt_version: u32,
pub flags: HandshakeExtensionMessageFlags,
pub encryption_field: EncryptionField,
pub stream_id: Option<String>,
pub group: Option<GroupMembershipExtension>,
pub local_ip: [u32; 4],
pub retransmit_after_ticks: u32,
pub max_retries: u32,
#[cfg(feature = "crypto")]
#[cfg_attr(docsrs, doc(cfg(feature = "crypto")))]
pub crypto: Option<CryptoConfig>,
}
impl Default for HandshakeConfig {
fn default() -> Self {
HandshakeConfig {
latency_ms: 120,
mtu: 1500,
max_flow_window_size: 8192,
initial_seq_number: 0,
srt_version: 0x0105_0000,
flags: HandshakeExtensionMessageFlags(
crate::packet::handshake::HS_MSG_FLAG_TSBPDSND
| crate::packet::handshake::HS_MSG_FLAG_TSBPDRCV
| crate::packet::handshake::HS_MSG_FLAG_CRYPT
| crate::packet::handshake::HS_MSG_FLAG_TLPKTDROP
| crate::packet::handshake::HS_MSG_FLAG_PERIODICNAK
| crate::packet::handshake::HS_MSG_FLAG_REXMITFLG,
),
encryption_field: EncryptionField::NoEncryption,
stream_id: None,
group: None,
local_ip: [0, 0, 0, 0],
retransmit_after_ticks: 3,
max_retries: 5,
#[cfg(feature = "crypto")]
crypto: None,
}
}
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
pub struct NegotiatedParams {
pub version: u32,
pub flags: HandshakeExtensionMessageFlags,
pub latency_ms: u16,
pub own_socket_id: u32,
pub peer_socket_id: u32,
pub stream_id: Option<String>,
pub group: Option<GroupMembershipExtension>,
#[cfg(feature = "crypto")]
#[cfg_attr(docsrs, doc(cfg(feature = "crypto")))]
pub sek: Option<Vec<u8>>,
#[cfg(feature = "crypto")]
#[cfg_attr(docsrs, doc(cfg(feature = "crypto")))]
pub salt: Option<[u8; crate::crypto::SALT_LEN]>,
}
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize))]
#[non_exhaustive]
pub enum HandshakeOutput {
Send(Vec<u8>),
Connected(NegotiatedParams),
Rejected(RejectionReason),
TimedOut,
}
pub(crate) fn build_bytes(hp: HandshakePacket<'_>) -> Result<Vec<u8>> {
let pkt = ControlPacket::Handshake(hp);
let mut buf = alloc::vec![0u8; pkt.serialized_len()];
pkt.serialize_into(&mut buf)?;
Ok(buf)
}
pub(crate) fn build_conclusion_extensions(
hs_ext_type: crate::packet::ExtensionType,
hs_msg: &HsExtMessage,
stream_id: Option<&str>,
group: Option<GroupMembershipExtension>,
) -> Result<(Vec<u8>, u16)> {
use crate::packet::handshake::{build_extension_block, encode_stream_id};
let mut ext_bytes = Vec::new();
ext_bytes.extend(build_extension_block(hs_ext_type, &hs_msg.to_bytes())?);
let mut ext_flags = HS_EXT_FLAG_HSREQ;
if let Some(sid) = stream_id {
let sid_bytes = encode_stream_id(sid);
ext_bytes.extend(build_extension_block(
crate::packet::ExtensionType::Sid,
&sid_bytes,
)?);
ext_flags |= HS_EXT_FLAG_CONFIG;
}
if let Some(g) = group {
ext_bytes.extend(build_extension_block(
crate::packet::ExtensionType::Group,
&g.to_bytes(),
)?);
ext_flags |= HS_EXT_FLAG_CONFIG;
}
Ok((ext_bytes, ext_flags))
}
#[derive(Debug, Default)]
pub(crate) struct ParsedPeerExtensions<'a> {
pub hs_msg: Option<HsExtMessage>,
pub stream_id: Option<String>,
pub group: Option<GroupMembershipExtension>,
#[cfg(feature = "crypto")]
pub km: Option<KeyMaterial<'a>>,
#[cfg(not(feature = "crypto"))]
_phantom: core::marker::PhantomData<&'a ()>,
}
pub(crate) fn parse_peer_extensions<'a>(
hp: &HandshakePacket<'a>,
) -> Result<ParsedPeerExtensions<'a>> {
use crate::packet::ExtensionType;
let mut out = ParsedPeerExtensions::default();
for block in hp.extensions.iter() {
let block = block.map_err(|_| Error::InvalidField {
what: "handshake extensions",
reason: "malformed extension block",
})?;
match block.ext_type {
ExtensionType::HsReq | ExtensionType::HsRsp => {
let msg = block.as_hs_ext_message().map_err(|_| Error::InvalidField {
what: "handshake extension message",
reason: "malformed HSREQ/HSRSP contents",
})?;
out.hs_msg = Some(msg);
}
ExtensionType::Sid => {
let sid = block.as_stream_id().map_err(|_| Error::InvalidField {
what: "stream ID extension",
reason: "malformed or non-UTF-8 contents",
})?;
out.stream_id = Some(sid);
}
ExtensionType::Group => {
let g = block
.as_group_membership()
.map_err(|_| Error::InvalidField {
what: "group membership extension",
reason: "malformed contents",
})?;
out.group = Some(g);
}
#[cfg(feature = "crypto")]
ExtensionType::KmReq | ExtensionType::KmRsp => {
let km = block.as_key_material().map_err(|_| Error::InvalidField {
what: "key material extension",
reason: "malformed contents",
})?;
out.km = Some(km);
}
_ => {}
}
}
Ok(out)
}
#[cfg(feature = "crypto")]
pub(crate) fn build_key_material_extension(crypto: &CryptoConfig) -> Result<Vec<u8>> {
let kek = crate::crypto::derive_kek(&crypto.passphrase, &crypto.salt, crypto.sek.len())?;
let (icv, wrapped) = crate::crypto::wrap_sek(&kek, &crypto.sek)?;
let km = KeyMaterial {
kk: KmKeyFlag::Even,
keki: 0,
cipher: Cipher::AesCtr,
auth: KmAuth::None,
se: StreamEncapsulation::Unspecified,
salt: &crypto.salt,
icv,
x_sek: &wrapped,
o_sek: None,
};
let mut buf = alloc::vec![0u8; km.serialized_len()];
km.serialize_into(&mut buf)?;
crate::packet::handshake::build_extension_block(crate::packet::ExtensionType::KmReq, &buf)
}
#[cfg(feature = "crypto")]
pub(crate) type RecoveredSek = (Vec<u8>, [u8; crate::crypto::SALT_LEN]);
#[cfg(feature = "crypto")]
pub(crate) fn recover_sek(passphrase: &[u8], km: &KeyMaterial<'_>) -> Result<RecoveredSek> {
if km.salt.len() != crate::crypto::SALT_LEN {
return Err(Error::InvalidField {
what: "Key Material Salt",
reason: "must be 16 bytes (the only Salt length this crate's codec accepts)",
});
}
let mut salt = [0u8; crate::crypto::SALT_LEN];
salt.copy_from_slice(km.salt);
let kek = crate::crypto::derive_kek(passphrase, &salt, km.x_sek.len())?;
let sek = crate::crypto::unwrap_sek(&kek, &km.icv, km.x_sek)?;
Ok((sek, salt))
}
#[cfg(feature = "crypto")]
pub(crate) fn echo_key_material_as_response(km: &KeyMaterial<'_>) -> Result<Vec<u8>> {
let mut buf = alloc::vec![0u8; km.serialized_len()];
km.serialize_into(&mut buf)?;
crate::packet::handshake::build_extension_block(crate::packet::ExtensionType::KmRsp, &buf)
}
#[cfg(feature = "crypto")]
pub(crate) fn verify_km_echo(crypto: &CryptoConfig, echoed: &KeyMaterial<'_>) -> bool {
let kek = match crate::crypto::derive_kek(&crypto.passphrase, &crypto.salt, crypto.sek.len()) {
Ok(k) => k,
Err(_) => return false,
};
let (icv, wrapped) = match crate::crypto::wrap_sek(&kek, &crypto.sek) {
Ok(v) => v,
Err(_) => return false,
};
echoed.salt == crypto.salt.as_slice() && echoed.icv == icv && echoed.x_sek == wrapped.as_slice()
}
pub(crate) fn negotiate_latency_ms(local_latency_ms: u16, peer_msg: &HsExtMessage) -> u16 {
local_latency_ms
.max(peer_msg.receiver_tsbpd_delay_ms)
.max(peer_msg.sender_tsbpd_delay_ms)
}
pub fn derive_cookie(peer_key: u64, time_bucket: u32, secret: u64) -> u32 {
let mut x = peer_key ^ (u64::from(time_bucket).wrapping_mul(0x9E37_79B9_7F4A_7C15)) ^ secret;
x ^= x >> 30;
x = x.wrapping_mul(0xBF58_476D_1CE4_E5B9);
x ^= x >> 27;
x = x.wrapping_mul(0x94D0_49BB_1331_11EB);
x ^= x >> 31;
(x as u32) | 1 }
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejection_reason_round_trips_table_7() {
let all = [
RejectionReason::Unknown,
RejectionReason::System,
RejectionReason::Peer,
RejectionReason::Resource,
RejectionReason::Rogue,
RejectionReason::Backlog,
RejectionReason::Ipe,
RejectionReason::Close,
RejectionReason::Version,
RejectionReason::RdvCookie,
RejectionReason::BadSecret,
RejectionReason::Unsecure,
RejectionReason::MessageApi,
RejectionReason::Congestion,
RejectionReason::Filter,
RejectionReason::Group,
];
for (i, r) in all.iter().enumerate() {
assert_eq!(r.to_bits(), 1000 + i as u32);
assert_eq!(RejectionReason::from_bits(r.to_bits()), *r);
assert_eq!(
RejectionReason::from_handshake_type(r.to_handshake_type()),
Some(*r)
);
}
}
#[test]
fn non_rejection_handshake_types_are_not_a_rejection() {
for ht in [
HandshakeType::Induction,
HandshakeType::Conclusion,
HandshakeType::Wavehand,
HandshakeType::Agreement,
HandshakeType::Done,
] {
assert_eq!(RejectionReason::from_handshake_type(ht), None);
}
}
#[test]
fn derive_cookie_is_deterministic_and_nonzero() {
let a = derive_cookie(0x1234_5678_9ABC_DEF0, 12345, 0xDEAD_BEEF);
let b = derive_cookie(0x1234_5678_9ABC_DEF0, 12345, 0xDEAD_BEEF);
assert_eq!(a, b);
assert_ne!(a, 0);
let c = derive_cookie(0x1234_5678_9ABC_DEF1, 12345, 0xDEAD_BEEF);
assert_ne!(a, c);
}
}