use std::hash::{BuildHasher, Hash, Hasher};
use std::sync::Arc;
use alloc::collections::VecDeque;
use core::time::Duration;
use tokio::net::UdpSocket;
use tokio::sync::mpsc;
use tokio::time::Instant;
use crate::arq::{Receiver as ArqReceiver, Sender as ArqSender};
use crate::caller::{CallerHandshake, CallerHandshakeState};
use crate::error::{Error, Result};
use crate::handshake_sm::{HandshakeConfig, HandshakeOutput, derive_cookie};
use crate::listener::{ListenerHandshake, ListenerHandshakeState};
use crate::livecc::{LiveCC, MaxBwConfig};
use crate::packet::misc::KeepAlivePacket;
use crate::packet::{ControlPacket, DataPacket, SrtPacket};
use crate::tsbpd::TsbpdScheduler;
const MAX_DATAGRAM: usize = 1500;
const TICK_INTERVAL_MS: u64 = 2;
const DEFAULT_DRIFT_US: u64 = 0;
const DEFAULT_TLPKT_DROP_ENABLED: bool = false;
const DEFAULT_MAX_BW: MaxBwConfig = MaxBwConfig::Set(125_000_000);
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Debug)]
struct OutboundPacket {
bytes: Vec<u8>,
is_data: bool,
}
impl OutboundPacket {
fn data(bytes: Vec<u8>) -> Self {
OutboundPacket {
bytes,
is_data: true,
}
}
fn control(bytes: Vec<u8>) -> Self {
OutboundPacket {
bytes,
is_data: false,
}
}
}
#[derive(Debug)]
pub struct SrtSocket {
peer_addr: std::net::SocketAddr,
to_driver: mpsc::UnboundedSender<Vec<u8>>,
from_driver: mpsc::UnboundedReceiver<Vec<u8>>,
driver: Option<tokio::task::JoinHandle<()>>,
}
impl SrtSocket {
pub async fn connect<A: tokio::net::ToSocketAddrs>(
remote_addr: A,
config: HandshakeConfig,
) -> Result<Self> {
let local = "0.0.0.0:0".parse::<std::net::SocketAddr>().unwrap();
Self::connect_from(local, remote_addr, config).await
}
pub async fn connect_from<A: tokio::net::ToSocketAddrs>(
local_addr: std::net::SocketAddr,
remote_addr: A,
config: HandshakeConfig,
) -> Result<Self> {
let socket = UdpSocket::bind(local_addr)
.await
.map_err(|e| io_err("bind", e))?;
let peer = resolve_one(remote_addr).await?;
let socket = Arc::new(socket);
let own_socket_id = config.initial_seq_number;
let mut hs = CallerHandshake::new(own_socket_id, config.clone());
let induction = hs.start().map_err(|_| Error::InvalidField {
what: "caller start",
reason: "start failed",
})?;
socket
.send_to(&induction, peer)
.await
.map_err(|e| io_err("send induction", e))?;
let mut buf = [0u8; MAX_DATAGRAM];
loop {
match hs.state() {
CallerHandshakeState::Connected => break,
CallerHandshakeState::Rejected | CallerHandshakeState::TimedOut => {
return Err(Error::InvalidField {
what: "hs state",
reason: "rejected or timed out",
});
}
_ => {}
}
let n = tokio::time::timeout(HANDSHAKE_TIMEOUT, socket.recv_from(&mut buf)).await;
match n {
Ok(Ok((len, _src))) => {
let bytes = &buf[..len];
let outcomes = hs.feed_bytes(bytes).map_err(|_| Error::InvalidField {
what: "hs feed",
reason: "feed failed",
})?;
for outcome in outcomes {
match outcome {
HandshakeOutput::Send(bytes) => {
socket
.send_to(&bytes, peer)
.await
.map_err(|e| io_err("send hs", e))?;
}
HandshakeOutput::Connected(params) => {
let peer_isn = require_peer_isn(bytes)?;
let epoch = Instant::now();
let tsbpd_delay_ms = u64::from(config.latency_ms);
let tsbpd_time_base = 0u64;
let conn = SrtSocket::spawn(
socket,
peer,
config.initial_seq_number,
peer_isn,
params.peer_socket_id,
tsbpd_time_base,
tsbpd_delay_ms,
epoch,
);
return Ok(conn);
}
HandshakeOutput::Rejected(_) => {
return Err(Error::InvalidField {
what: "hs rejected",
reason: "peer rejected",
});
}
HandshakeOutput::TimedOut => {
return Err(Error::InvalidField {
what: "hs timeout",
reason: "caller timed out",
});
}
}
}
}
Ok(Err(e)) => return Err(io_err("recv hs", e)),
Err(_) => {
for outcome in hs.tick() {
match outcome {
HandshakeOutput::Send(bytes) => {
socket
.send_to(&bytes, peer)
.await
.map_err(|e| io_err("retransmit", e))?;
}
HandshakeOutput::TimedOut => {
return Err(Error::InvalidField {
what: "hs timeout",
reason: "retransmit exhausted",
});
}
_ => {}
}
}
}
}
}
Err(Error::InvalidField {
what: "handshake",
reason: "unreachable",
})
}
#[allow(clippy::too_many_arguments)]
fn spawn(
udp: Arc<UdpSocket>,
peer_addr: std::net::SocketAddr,
our_initial_seq: u32,
peer_initial_seq: u32,
peer_socket_id: u32,
tsbpd_time_base: u64,
tsbpd_delay_ms: u64,
epoch: Instant,
) -> Self {
let (to_driver, app_out) = mpsc::unbounded_channel::<Vec<u8>>();
let (deliver, from_driver) = mpsc::unbounded_channel::<Vec<u8>>();
let driver = Driver {
udp,
peer_addr,
peer_socket_id,
sender: ArqSender::new(peer_socket_id),
receiver: ArqReceiver::new(peer_socket_id, peer_initial_seq),
tsbpd: TsbpdScheduler::new(
peer_initial_seq,
tsbpd_time_base,
tsbpd_delay_ms,
DEFAULT_DRIFT_US,
DEFAULT_TLPKT_DROP_ENABLED,
None,
),
livecc: LiveCC::new(DEFAULT_MAX_BW),
next_message_number: 1,
next_send_seq: our_initial_seq,
epoch,
staged: std::collections::BTreeMap::new(),
outbound: VecDeque::new(),
deliver,
peer_shutdown: false,
};
let handle = tokio::spawn(driver.run(app_out));
SrtSocket {
peer_addr,
to_driver,
from_driver,
driver: Some(handle),
}
}
pub async fn send(&mut self, payload: &[u8]) -> Result<()> {
self.to_driver
.send(payload.to_vec())
.map_err(|_| Error::Io {
kind: std::io::ErrorKind::BrokenPipe,
context: "send",
})
}
pub fn peer_addr(&self) -> std::net::SocketAddr {
self.peer_addr
}
pub async fn recv(&mut self) -> Result<Option<Vec<u8>>> {
Ok(self.from_driver.recv().await)
}
}
impl Drop for SrtSocket {
fn drop(&mut self) {
if let Some(handle) = self.driver.take() {
handle.abort();
}
}
}
struct Driver {
udp: Arc<UdpSocket>,
peer_addr: std::net::SocketAddr,
peer_socket_id: u32,
sender: ArqSender,
receiver: ArqReceiver,
tsbpd: TsbpdScheduler,
livecc: LiveCC,
next_message_number: u32,
next_send_seq: u32,
epoch: Instant,
staged: std::collections::BTreeMap<u32, Vec<u8>>,
outbound: VecDeque<OutboundPacket>,
deliver: mpsc::UnboundedSender<Vec<u8>>,
peer_shutdown: bool,
}
impl Driver {
async fn run(mut self, mut app_out: mpsc::UnboundedReceiver<Vec<u8>>) {
let udp = Arc::clone(&self.udp);
let mut buf = [0u8; MAX_DATAGRAM];
let mut ticker = tokio::time::interval(Duration::from_millis(TICK_INTERVAL_MS));
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
let mut app_open = true;
let mut shutting_down = false;
loop {
tokio::select! {
maybe = app_out.recv(), if app_open => {
match maybe {
Some(payload) => self.send_one(&payload),
None => {
app_open = false;
shutting_down = true;
}
}
}
r = udp.recv_from(&mut buf) => {
match r {
Ok((len, src)) if src == self.peer_addr => {
let _ = self.ingress(&buf[..len]);
}
Ok(_) => {} Err(_) => break, }
}
_ = ticker.tick() => {
self.tick_engines();
}
}
if self.flush_outbound().await.is_err() {
break;
}
if self.peer_shutdown || shutting_down {
break;
}
}
}
fn send_one(&mut self, payload: &[u8]) {
let now = self.elapsed();
self.livecc.on_data_packet(payload.len() as u64);
let bytes = self
.sender
.on_data(self.next_send_seq, self.next_message_number, payload, now);
self.next_send_seq = self.next_send_seq.wrapping_add(1);
self.next_message_number = self.next_message_number.wrapping_add(1);
self.outbound.push_back(OutboundPacket::data(bytes));
}
fn ingress(&mut self, bytes: &[u8]) -> Result<()> {
let now = self.elapsed();
let packet = SrtPacket::parse(bytes)?;
match packet {
SrtPacket::Data(d) => {
let outcome = self.receiver.feed_data(d.seq_number, now);
if let Some(nak_bytes) = outcome.nak {
self.outbound.push_back(OutboundPacket::control(nak_bytes));
}
self.staged
.entry(d.seq_number)
.or_insert_with(|| d.data.to_vec());
let tsbpd_out = self.tsbpd.feed_data(d.seq_number, d.timestamp, now);
for &seq in &tsbpd_out.delivered {
if let Some(payload) = self.staged.remove(&seq) {
let _ = self.deliver.send(payload);
}
}
}
SrtPacket::Control(ref c) => match c {
ControlPacket::Ack(ack) => {
if let Some(ackack_bytes) = self.sender.on_ack(ack, now) {
self.outbound
.push_back(OutboundPacket::control(ackack_bytes));
}
}
ControlPacket::Nak(nak) => {
self.sender.on_nak(nak);
}
ControlPacket::AckAck(ackack) => {
self.receiver.on_ackack(ackack, now);
}
ControlPacket::KeepAlive(_) => {
let pkt = ControlPacket::KeepAlive(KeepAlivePacket {
timestamp: self.elapsed_us(),
dest_socket_id: self.peer_socket_id,
});
let mut buf = vec![0u8; pkt.serialized_len()];
let _ = pkt.serialize_into(&mut buf);
self.outbound.push_back(OutboundPacket::control(buf));
}
ControlPacket::Shutdown(_) => {
self.peer_shutdown = true;
}
_ => {}
},
}
Ok(())
}
fn tick_engines(&mut self) {
let now = self.elapsed();
for bytes in self.sender.tick(now) {
if let Ok(dp) = DataPacket::parse(&bytes) {
self.livecc.on_data_packet(dp.data.len() as u64);
}
self.outbound.push_back(OutboundPacket::data(bytes));
}
for bytes in self.receiver.tick(now) {
self.outbound.push_back(OutboundPacket::control(bytes));
}
let tsbpd_out = self.tsbpd.tick(now);
for &seq in &tsbpd_out.delivered {
if let Some(payload) = self.staged.remove(&seq) {
let _ = self.deliver.send(payload);
}
}
}
fn elapsed(&self) -> Duration {
Instant::now().duration_since(self.epoch)
}
fn elapsed_us(&self) -> u32 {
self.elapsed().as_micros().min(u128::from(u32::MAX)) as u32
}
async fn flush_outbound(&mut self) -> Result<()> {
while let Some(item) = self.outbound.pop_front() {
if item.is_data {
let period = self.livecc.on_ack_received();
if period > Duration::ZERO {
tokio::time::sleep(period).await;
}
}
self.udp
.send_to(&item.bytes, self.peer_addr)
.await
.map_err(|e| io_err("send", e))?;
}
Ok(())
}
}
#[derive(Debug)]
pub struct SrtListener {
udp: Arc<UdpSocket>,
config: HandshakeConfig,
next_socket_id: u32,
cookie_secret: u64,
pending: std::collections::HashMap<std::net::SocketAddr, PendingListener>,
outbound_queue: std::collections::HashMap<std::net::SocketAddr, VecDeque<Vec<u8>>>,
}
#[derive(Debug)]
struct PendingListener {
handshake: ListenerHandshake,
params: Option<HandshakeOutput>,
peer_initial_seq: u32,
}
impl SrtListener {
pub async fn bind<A: tokio::net::ToSocketAddrs>(
addr: A,
config: HandshakeConfig,
) -> Result<Self> {
let socket = UdpSocket::bind(addr).await.map_err(|e| io_err("bind", e))?;
Ok(SrtListener {
udp: Arc::new(socket),
config,
next_socket_id: 1,
cookie_secret: random_u64(),
pending: std::collections::HashMap::new(),
outbound_queue: std::collections::HashMap::new(),
})
}
pub fn local_addr(&self) -> Result<std::net::SocketAddr> {
self.udp.local_addr().map_err(|e| io_err("local_addr", e))
}
pub async fn accept(&mut self) -> Result<SrtSocket> {
let mut buf = [0u8; MAX_DATAGRAM];
loop {
if let Some(conn) = self.drain_completed() {
return conn;
}
let n = tokio::time::timeout(Duration::from_millis(100), self.udp.recv_from(&mut buf))
.await;
match n {
Ok(Ok((len, src))) => {
let _ = self.handle_datagram(src, &buf[..len]);
self.flush_for_peer(src).await?;
}
Ok(Err(e)) => return Err(io_err("recv_from", e)),
Err(_) => {
self.tick_pending();
self.flush_all().await?;
}
}
}
}
fn handle_datagram(&mut self, src: std::net::SocketAddr, bytes: &[u8]) -> Result<()> {
let packet = SrtPacket::parse(bytes).map_err(|_| Error::InvalidField {
what: "parse",
reason: "non-SRT datagram",
})?;
let ctrl = match packet {
SrtPacket::Control(c) => c,
_ => return Ok(()),
};
let is_new = !self.pending.contains_key(&src);
if is_new {
let peer_isn = match &ctrl {
ControlPacket::Handshake(hp) => hp.initial_seq_number,
_ => return Ok(()),
};
let own_socket_id = self.next_socket_id;
self.next_socket_id = self.next_socket_id.wrapping_add(1);
let peer_key = addr_to_u64(&src);
let time_bucket = unix_time_bucket();
let syn_cookie = derive_cookie(peer_key, time_bucket, self.cookie_secret);
let hs = ListenerHandshake::new(own_socket_id, syn_cookie, self.config.clone());
self.pending.insert(
src,
PendingListener {
handshake: hs,
params: None,
peer_initial_seq: peer_isn,
},
);
}
let entry = self.pending.get_mut(&src).ok_or(Error::InvalidField {
what: "pending",
reason: "no pending entry",
})?;
let outcomes = entry
.handshake
.feed(&ctrl)
.map_err(|_| Error::InvalidField {
what: "listener feed",
reason: "feed failed",
})?;
for outcome in outcomes {
match outcome {
HandshakeOutput::Send(bytes) => {
self.outbound_queue.entry(src).or_default().push_back(bytes);
}
HandshakeOutput::Connected(_) => {
entry.params = Some(HandshakeOutput::Connected(
entry.handshake.negotiated().unwrap().clone(),
));
}
HandshakeOutput::Rejected(_) => {
self.pending.remove(&src);
return Err(Error::InvalidField {
what: "hs rejected",
reason: "peer rejected",
});
}
HandshakeOutput::TimedOut => {
self.pending.remove(&src);
return Err(Error::InvalidField {
what: "hs timeout",
reason: "listener",
});
}
}
}
Ok(())
}
fn tick_pending(&mut self) {
let mut to_remove = Vec::new();
for (addr, entry) in self.pending.iter_mut() {
for outcome in entry.handshake.tick() {
match outcome {
HandshakeOutput::Send(bytes) => {
self.outbound_queue
.entry(*addr)
.or_default()
.push_back(bytes);
}
HandshakeOutput::TimedOut => {
to_remove.push(*addr);
}
_ => {}
}
}
}
for addr in to_remove {
self.pending.remove(&addr);
}
}
fn drain_completed(&mut self) -> Option<Result<SrtSocket>> {
let addr = self
.pending
.iter()
.find(|(_, p)| {
p.params.is_some()
&& matches!(p.handshake.state(), ListenerHandshakeState::Connected)
})
.map(|(addr, _)| *addr)?;
let entry = self.pending.remove(&addr)?;
let peer_initial_seq = entry.peer_initial_seq;
let peer_socket_id = entry
.handshake
.negotiated()
.expect("filtered to Connected state")
.peer_socket_id;
let our_initial_seq = self.config.initial_seq_number;
let tsbpd_delay_ms = u64::from(self.config.latency_ms);
let tsbpd_time_base = 0;
let epoch = Instant::now();
let conn = SrtSocket::spawn(
Arc::clone(&self.udp),
addr,
our_initial_seq,
peer_initial_seq,
peer_socket_id,
tsbpd_time_base,
tsbpd_delay_ms,
epoch,
);
Some(Ok(conn))
}
async fn flush_for_peer(&mut self, addr: std::net::SocketAddr) -> Result<()> {
if let Some(queue) = self.outbound_queue.get_mut(&addr) {
while let Some(bytes) = queue.pop_front() {
self.udp
.send_to(&bytes, addr)
.await
.map_err(|e| io_err("send_to", e))?;
}
}
Ok(())
}
async fn flush_all(&mut self) -> Result<()> {
let addrs: Vec<std::net::SocketAddr> = self.outbound_queue.keys().copied().collect();
for addr in addrs {
self.flush_for_peer(addr).await?;
}
Ok(())
}
}
fn io_err(context: &'static str, e: std::io::Error) -> Error {
Error::Io {
kind: e.kind(),
context,
}
}
fn addr_to_u64(addr: &std::net::SocketAddr) -> u64 {
let mut hasher = std::collections::hash_map::DefaultHasher::new();
addr.hash(&mut hasher);
hasher.finish()
}
fn unix_time_bucket() -> u32 {
let secs = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
(secs / 60) as u32
}
fn random_u64() -> u64 {
std::collections::hash_map::RandomState::new()
.build_hasher()
.finish()
}
async fn resolve_one<A: tokio::net::ToSocketAddrs>(addr: A) -> Result<std::net::SocketAddr> {
let mut addrs = tokio::net::lookup_host(addr)
.await
.map_err(|e| io_err("resolve", e))?;
addrs.next().ok_or(Error::InvalidField {
what: "resolve",
reason: "no addrs",
})
}
fn require_peer_isn(bytes: &[u8]) -> Result<u32> {
match SrtPacket::parse(bytes) {
Ok(SrtPacket::Control(ControlPacket::Handshake(hp))) => Ok(hp.initial_seq_number),
_ => Err(Error::InvalidField {
what: "peer isn",
reason: "handshake reached Connected but its final packet did not re-parse as a \
Handshake control packet; refusing to seed ARQ/TSBPD with a fabricated ISN",
}),
}
}
#[cfg(test)]
mod isn_tests {
use super::*;
use crate::packet::{
EncryptionField, HandshakeExtensionFlags, HandshakeExtensions, HandshakePacket,
HandshakeType,
};
fn handshake_bytes(initial_seq_number: u32) -> Vec<u8> {
let hp = HandshakePacket {
timestamp: 0,
dest_socket_id: 0,
version: 5,
encryption_field: EncryptionField::NoEncryption,
extension_field: HandshakeExtensionFlags(0),
initial_seq_number,
mtu: 1500,
max_flow_window_size: 8192,
handshake_type: HandshakeType::Conclusion,
srt_socket_id: 42,
syn_cookie: 0,
peer_ip: [0; 4],
extensions: HandshakeExtensions(&[]),
};
crate::handshake_sm::build_bytes(hp).expect("build handshake bytes")
}
#[test]
fn require_peer_isn_extracts_nonzero_isn() {
let bytes = handshake_bytes(0xABCD_1234);
assert_eq!(require_peer_isn(&bytes).unwrap(), 0xABCD_1234);
}
#[test]
fn require_peer_isn_distinguishes_genuine_zero_from_parse_failure() {
let bytes = handshake_bytes(0);
assert_eq!(require_peer_isn(&bytes).unwrap(), 0);
let err = require_peer_isn(&[0u8; 4]).unwrap_err();
assert!(matches!(
err,
Error::InvalidField {
what: "peer isn",
..
}
));
}
#[test]
fn require_peer_isn_rejects_non_handshake_control_packet() {
let ka = ControlPacket::KeepAlive(KeepAlivePacket {
timestamp: 0,
dest_socket_id: 7,
});
let mut buf = alloc::vec![0u8; ka.serialized_len()];
ka.serialize_into(&mut buf).expect("serialize keepalive");
let err = require_peer_isn(&buf).unwrap_err();
assert!(matches!(
err,
Error::InvalidField {
what: "peer isn",
..
}
));
}
}