use std::collections::{BTreeMap, HashMap, VecDeque};
use std::io;
use std::net::{SocketAddr, ToSocketAddrs, UdpSocket};
use std::sync::Arc;
use std::time::{Duration, Instant};
use crate::control_table::ControlTable;
use crate::fusion::{FusionPolicy, ImmediateUpConservativeDown, SensorSnapshot};
use crate::interleave::Interleaver;
use crate::control_frame::{
decode_control, encode_control, is_control, AckFrame, ControlPacket, LinkFrame, LossAcctFrame,
LossFrame, NakFrame, PathFrame, PmtuFrame, RingFrame, TimingFrame,
};
use crate::link_sensor::{platform_sensor, LinkClass, LinkSensor};
use crate::net_events::NetEventObserver;
use crate::path_model_sensor::PathModel;
use crate::path_sensor::PathSensor;
use crate::rtt_shape_sensor::RttShape;
use crate::reliable_udp::{
datagram_epoch, is_outer_datagram, Decoder, Encoder, Feedback, NAK_NONE,
};
const CONTROL_RECV_BUF: usize = 256;
fn feedback_from_control(cp: &ControlPacket) -> Feedback {
let (nak_block, nak_mask) = cp.nak.map(|n| (n.block, n.mask)).unwrap_or((NAK_NONE, 0));
let loss = cp.loss.unwrap_or_default();
Feedback {
ack_through: cp.ack.map(|a| a.ack_through).unwrap_or(0),
nak_block,
nak_mask,
loss_x255: loss.loss_x255,
burstiness_x255: loss.burstiness_x255,
owd_trend_class: loss.owd_trend_class,
loss_class: loss.loss_class,
}
}
pub type SensOMaticRsSender = ReliableUdpSender;
pub type SensOMaticRsReceiver = ReliableUdpReceiver;
pub type SensOMaticSender = ReliableUdpSender;
pub type SensOMaticReceiver = ReliableUdpReceiver;
const SESSION_CHALLENGE_TIMEOUT: Duration = Duration::from_millis(500);
const PEER_SILENCE_TIMEOUT: Duration = Duration::from_secs(2);
const UNSPECIFIED_PEER: SocketAddr = SocketAddr::new(
std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED),
0,
);
const RECV_BUF: usize = 2048;
#[cfg(target_os = "linux")]
type MmsgLen = libc::c_uint;
#[cfg(target_os = "freebsd")]
type MmsgLen = usize;
const NAK_COOLDOWN: Duration = Duration::from_millis(12);
const SOCK_BUF_BYTES: usize = 8 << 20;
fn size_socket_buffers(sock: &UdpSocket) {
let s = socket2::SockRef::from(sock);
s.set_recv_buffer_size(SOCK_BUF_BYTES).ok();
s.set_send_buffer_size(SOCK_BUF_BYTES).ok();
}
const ACK_INTERVAL: Duration = Duration::from_millis(1);
const MAX_NAKS_PER_CYCLE: usize = 64;
const HEARTBEAT_INTERVAL: Duration = Duration::from_millis(20);
const BW_PROBE_INTERVAL: Duration = Duration::from_secs(2);
const BW_PROBE_PAIRS: u8 = 8;
const BW_PROBE_TRAIN: u8 = 12;
const BW_PROBE_BYTES: usize = 1400;
#[cfg_attr(not(target_os = "linux"), allow(dead_code))]
const TRACE_INTERVAL: Duration = Duration::from_secs(3);
const MAX_TRACE_HOPS: u8 = 8;
const FORECAST_TICK: Duration = Duration::from_millis(50);
const FORECAST_HEADROOM: f64 = 2.0;
const LINK_SAMPLE_INTERVAL: Duration = Duration::from_millis(200);
const MIN_PACED_WINDOW: u32 = 4;
const PACE_TARGET_MS: f32 = 10.0;
const MIN_PACE_INTERVAL_US: u64 = 1000;
const DEAD_RTT_MULTIPLE: u64 = 8;
const DEAD_FLOOR_US: u64 = 250_000;
const MIN_RECOVERY_BYTES_PER_S: u64 = 1_000_000;
const RECOVERY_BUCKET_BYTES: f64 = 8192.0;
const RECOVERY_GRACE_RTTS: u64 = 4;
const WIFI_SHAPE_CONFIDENCE: f32 = 0.15;
pub struct ReliableUdpSender {
sock: crate::dgram::DgramSock,
enc: Encoder,
interleaver: Interleaver,
control: Arc<ControlTable>,
fusion: Box<dyn FusionPolicy + Send>,
link_sensor: Box<dyn LinkSensor + Send>,
link_stress: f32,
link_class: LinkClass,
link_quality: u8,
class_shift: f32,
link_phy_kbps: u32,
link_mcs_norm: f32,
congestion_fraction: f32,
ctrl_out: u32,
ctrl_recv: u32,
peer_seq: u32,
rev_loss: f32,
last_fwd_loss: f32,
path_sensor: PathSensor,
net_events: NetEventObserver,
peer_pmtu: u16,
peer_pmtu_shift: f32,
net_event_shift_peak: f32,
path_model: PathModel,
rtt_shape: RttShape,
block_send_us: VecDeque<(u32, u64)>,
flow_window_max: u32,
pacing_enabled: bool,
paced_window: f32,
last_pace_us: u64,
last_feedback_at: Instant,
link_dead: bool,
last_probe_at: Instant,
proactive_recovery: bool,
dead_episodes: u64,
probes_sent: u64,
recovered_blocks: u64,
recovery_dgrams: VecDeque<Vec<u8>>,
recovery_tokens: f64,
last_recovery_us: u64,
recovery_grace_until_us: u64,
recovery_target: u32,
recovery_started_us: u64,
recovery_interval_us: u64,
start: Instant,
last_hb: Instant,
last_link_sample: Instant,
last_bw_probe: Instant,
bw_probe_round: u8,
avail_bw_kbps: u64,
wbest_capacity_kbps: u64,
#[cfg_attr(not(target_os = "linux"), allow(dead_code))]
trace_peer: SocketAddr,
#[cfg_attr(not(target_os = "linux"), allow(dead_code))]
last_trace: Instant,
#[cfg_attr(not(target_os = "linux"), allow(dead_code))]
trace_round: u8,
#[cfg_attr(not(target_os = "linux"), allow(dead_code))]
trace_send_us: Vec<u64>,
trace_hops: Vec<crate::trace_sensor::TraceHop>,
asym: crate::trace_sensor::PathAsymmetry,
ce_rate: f32,
forecast_bps: u64,
leo_period_s: f32,
leo_conf: f32,
leo_secs_to_spike: f32,
}
#[cfg(target_os = "windows")]
type LpfnWsaSendMsg = unsafe extern "system" fn(
usize,
*const windows_sys::Win32::Networking::WinSock::WSAMSG,
u32,
*mut u32,
*mut core::ffi::c_void,
*const core::ffi::c_void,
) -> i32;
#[cfg(target_os = "windows")]
static WSASENDMSG_PTR: std::sync::OnceLock<Option<LpfnWsaSendMsg>> =
std::sync::OnceLock::new();
#[cfg(target_os = "windows")]
fn load_wsasendmsg(sock: usize) -> Option<LpfnWsaSendMsg> {
*WSASENDMSG_PTR.get_or_init(|| {
use windows_sys::Win32::Networking::WinSock::WSAIoctl;
const SIO_GET_EXTENSION_FUNCTION_POINTER: u32 = 0xC800_0006;
let guid = windows_sys::core::GUID {
data1: 0xa441_e712,
data2: 0x754f,
data3: 0x43ca,
data4: [0x84, 0xa7, 0x0d, 0xee, 0x44, 0xcf, 0x60, 0x6d],
};
let mut func: usize = 0;
let mut bytes: u32 = 0;
let rc = unsafe {
WSAIoctl(
sock,
SIO_GET_EXTENSION_FUNCTION_POINTER,
&guid as *const _ as *const core::ffi::c_void,
size_of::<windows_sys::core::GUID>() as u32,
&mut func as *mut usize as *mut core::ffi::c_void,
size_of::<usize>() as u32,
&mut bytes,
std::ptr::null_mut(),
None,
)
};
if rc != 0 || func == 0 {
None
} else {
let p = func as *const core::ffi::c_void;
Some(unsafe { std::mem::transmute::<*const core::ffi::c_void, LpfnWsaSendMsg>(p) })
}
})
}
#[cfg(target_os = "windows")]
type LpfnWsaRecvMsg = unsafe extern "system" fn(
usize,
*mut windows_sys::Win32::Networking::WinSock::WSAMSG,
*mut u32,
*mut core::ffi::c_void,
*const core::ffi::c_void,
) -> i32;
#[cfg(target_os = "windows")]
static WSARECVMSG_PTR: std::sync::OnceLock<Option<LpfnWsaRecvMsg>> =
std::sync::OnceLock::new();
#[cfg(target_os = "windows")]
fn load_wsarecvmsg(sock: usize) -> Option<LpfnWsaRecvMsg> {
*WSARECVMSG_PTR.get_or_init(|| {
use windows_sys::Win32::Networking::WinSock::WSAIoctl;
const SIO_GET_EXTENSION_FUNCTION_POINTER: u32 = 0xC800_0006;
let guid = windows_sys::core::GUID {
data1: 0xf689_d7c8,
data2: 0x6f1f,
data3: 0x436b,
data4: [0x8a, 0x53, 0xe5, 0x4f, 0xe3, 0x51, 0xc3, 0x22],
};
let mut func: usize = 0;
let mut bytes: u32 = 0;
let rc = unsafe {
WSAIoctl(
sock,
SIO_GET_EXTENSION_FUNCTION_POINTER,
&guid as *const _ as *const core::ffi::c_void,
size_of::<windows_sys::core::GUID>() as u32,
&mut func as *mut usize as *mut core::ffi::c_void,
size_of::<usize>() as u32,
&mut bytes,
std::ptr::null_mut(),
None,
)
};
if rc != 0 || func == 0 {
None
} else {
let p = func as *const core::ffi::c_void;
Some(unsafe { std::mem::transmute::<*const core::ffi::c_void, LpfnWsaRecvMsg>(p) })
}
})
}
#[cfg(target_os = "windows")]
fn uso_enabled() -> bool {
static EN: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*EN.get_or_init(|| std::env::var("SUBETHA_USO").map(|v| v != "0").unwrap_or(true))
}
#[cfg(target_os = "windows")]
static USO_OFFLOAD: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
#[cfg(target_os = "windows")]
static USO_FALLBACK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
pub fn uso_stats() -> (u64, u64) {
#[cfg(target_os = "windows")]
{
use std::sync::atomic::Ordering::Relaxed;
(USO_OFFLOAD.load(Relaxed), USO_FALLBACK.load(Relaxed))
}
#[cfg(not(target_os = "windows"))]
{
(0, 0)
}
}
impl ReliableUdpSender {
pub fn bind(
local: impl ToSocketAddrs,
peer: SocketAddr,
k: usize,
r: usize,
max_item: usize,
) -> io::Result<Self> {
Self::bind_with_control(local, peer, k, r, max_item, Arc::new(ControlTable::new()))
}
pub fn bind_with_control(
local: impl ToSocketAddrs,
peer: SocketAddr,
k: usize,
r: usize,
max_item: usize,
control: Arc<ControlTable>,
) -> io::Result<Self> {
if !(1..=crate::reliable_udp::MAX_SHARDS).contains(&k) {
return Err(io::Error::new(
io::ErrorKind::InvalidInput,
format!(
"RS data-shard count k={k} out of range: need 1 <= k <= {}",
crate::reliable_udp::MAX_SHARDS
),
));
}
let sock = UdpSocket::bind(local)?;
sock.connect(peer)?;
sock.set_nonblocking(true)?;
size_socket_buffers(&sock);
#[cfg(target_os = "linux")]
{
use std::os::fd::AsRawFd;
crate::trace_sensor::enable_icmp_errors(sock.as_raw_fd());
}
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
{
enable_ttl_ecn(&sock);
set_ect(&sock);
}
let sock = crate::dgram::DgramSock::from_udp(sock);
let depth = control.interleave_depth() as usize;
let now = Instant::now();
let enc = Encoder::new(k, r, max_item);
let flow_window_max = enc.flow_window();
Ok(Self {
sock,
enc,
interleaver: Interleaver::new(depth),
control,
fusion: Box::new(ImmediateUpConservativeDown::new(8)),
link_sensor: platform_sensor(None),
link_stress: 0.0,
link_class: LinkClass::Unknown,
link_quality: 0,
class_shift: 0.0,
link_phy_kbps: 0,
link_mcs_norm: 0.0,
congestion_fraction: 0.0,
ctrl_out: 0,
ctrl_recv: 0,
peer_seq: 0,
rev_loss: 0.0,
last_fwd_loss: 0.0,
path_sensor: PathSensor::new(),
net_events: NetEventObserver::start(None),
peer_pmtu: 0,
peer_pmtu_shift: 0.0,
net_event_shift_peak: 0.0,
path_model: PathModel::new(k * max_item),
rtt_shape: RttShape::new(),
block_send_us: VecDeque::new(),
flow_window_max,
pacing_enabled: true,
paced_window: flow_window_max as f32,
last_pace_us: 0,
last_feedback_at: now,
link_dead: false,
last_probe_at: now,
proactive_recovery: true,
dead_episodes: 0,
probes_sent: 0,
recovered_blocks: 0,
recovery_dgrams: VecDeque::new(),
recovery_tokens: 0.0,
last_recovery_us: 0,
recovery_grace_until_us: 0,
recovery_target: 0,
recovery_started_us: 0,
recovery_interval_us: 0,
start: now,
last_hb: now.checked_sub(HEARTBEAT_INTERVAL).unwrap_or(now),
last_link_sample: now.checked_sub(LINK_SAMPLE_INTERVAL).unwrap_or(now),
last_bw_probe: now,
bw_probe_round: 0,
avail_bw_kbps: 0,
wbest_capacity_kbps: 0,
trace_peer: peer,
last_trace: now,
trace_round: 0,
trace_send_us: vec![0u64; MAX_TRACE_HOPS as usize + 1],
trace_hops: Vec::new(),
asym: crate::trace_sensor::PathAsymmetry::new(),
ce_rate: 0.0,
forecast_bps: 0,
leo_period_s: 0.0,
leo_conf: 0.0,
leo_secs_to_spike: 0.0,
})
}
pub fn link_stress(&self) -> f32 {
self.link_stress
}
pub fn path_observation(&self) -> Option<(u8, u8, u8)> {
self.path_sensor.last()
}
pub fn net_event_count(&self) -> u64 {
self.net_events.event_count()
}
pub fn local_pmtu(&self) -> u16 {
self.net_events.pmtu().unwrap_or(0)
}
pub fn peer_pmtu(&self) -> u16 {
self.peer_pmtu
}
pub fn net_event_shift(&self) -> f32 {
self.net_events.path_shift().max(self.peer_pmtu_shift)
}
pub fn net_event_shift_peak(&self) -> f32 {
self.net_event_shift_peak
}
pub fn inject_path_event(&self) {
self.net_events.inject_event();
}
pub fn inject_pmtu(&self, mtu: u16) {
self.net_events.inject_pmtu(mtu);
}
pub fn congestion_fraction(&self) -> f32 {
self.congestion_fraction
}
pub fn rev_loss(&self) -> f32 {
self.rev_loss
}
pub fn link_backend(&self) -> &'static str {
self.link_sensor.backend()
}
pub fn btlbw_bps(&self) -> u64 {
self.path_model.btlbw_bps()
}
pub fn rtprop_us(&self) -> u64 {
self.path_model.rtprop_us()
}
pub fn bdp_blocks(&self) -> u64 {
self.path_model.bdp_blocks()
}
pub fn backhaul_hops(&self) -> u8 {
self.path_model.backhaul_hops(
self.link_phy_kbps as u64 * 1000,
self.link_mcs_norm,
self.congestion_fraction,
)
}
pub fn first_hop_mbps(&self) -> f32 {
self.link_phy_kbps as f32 / 1000.0
}
fn inferred_link_class(&self) -> LinkClass {
if self.link_class == LinkClass::Unknown
&& self.rtt_shape.wifi_confidence() > WIFI_SHAPE_CONFIDENCE
{
LinkClass::Wifi
} else {
self.link_class
}
}
pub fn rtt_bimodality(&self) -> f32 {
self.rtt_shape.bimodality().map(|b| b as f32).unwrap_or(-1.0)
}
pub fn rtt_wifi_confidence(&self) -> f32 {
self.rtt_shape.wifi_confidence()
}
pub fn queue_delay_ms(&self) -> f32 {
self.path_model.queue_delay_us() as f32 / 1000.0
}
pub fn rtt_mean_ms(&self) -> f32 {
self.path_model.rtt_mean_us() as f32 / 1000.0
}
pub fn flow_window(&self) -> u32 {
self.enc.flow_window()
}
pub fn set_pacing(&mut self, enabled: bool) {
self.pacing_enabled = enabled;
if !enabled {
self.enc.set_flow_window(self.flow_window_max);
}
}
pub fn set_proactive_recovery(&mut self, enabled: bool) {
self.proactive_recovery = enabled;
}
pub fn link_dead(&self) -> bool {
self.link_dead
}
pub fn liveness_stats(&self) -> (u64, u64, u64) {
(self.dead_episodes, self.probes_sent, self.recovered_blocks)
}
pub fn recovery_interval_ms(&self) -> f32 {
self.recovery_interval_us as f32 / 1000.0
}
pub fn control(&self) -> &Arc<ControlTable> {
&self.control
}
pub fn coding_counts(&self) -> (u64, u64) {
self.enc.coding_counts()
}
pub fn with_sensor(mut self, sensor: Box<dyn LinkSensor + Send>) -> Self {
self.link_sensor = sensor;
self
}
pub fn with_fusion(mut self, policy: Box<dyn FusionPolicy + Send>) -> Self {
self.fusion = policy;
self
}
pub fn enable_tower(&mut self, d: usize, r_outer: usize) {
self.enc.enable_tower(d, r_outer);
}
pub fn set_sock(&mut self, sock: crate::dgram::DgramSock) {
self.sock = sock;
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.sock.local_addr()
}
pub fn send_item(&mut self, item: &[u8]) -> io::Result<()> {
self.sync_interleave()?;
let block = self.enc.push(item);
if !block.is_empty() {
let pkts = self.interleaver.add_block(block);
self.send_batch(&pkts)?;
let sealed = self.enc.next_block_id().wrapping_sub(1);
self.block_send_us
.push_back((sealed, self.start.elapsed().as_micros() as u64));
self.maybe_sample_link();
self.maybe_send_heartbeat()?;
self.maybe_send_bw_probe()?;
self.maybe_send_trace()?;
self.drain_feedback()?;
}
Ok(())
}
pub fn flush(&mut self) -> io::Result<()> {
let block = self.enc.flush();
if !block.is_empty() {
let pkts = self.interleaver.add_block(block);
self.send_batch(&pkts)?;
}
let tail = self.interleaver.flush();
self.send_batch(&tail)?;
Ok(())
}
fn sync_interleave(&mut self) -> io::Result<()> {
let want = self.control.interleave_depth() as usize;
if want != self.interleaver.depth() {
let pkts = self.interleaver.set_depth(want);
self.send_batch(&pkts)?;
}
Ok(())
}
pub fn flow_blocked(&self) -> bool {
self.enc.flow_blocked()
}
pub fn pending_len(&self) -> usize {
self.enc.pending_len()
}
pub fn pump_feedback(&mut self) -> io::Result<()> {
self.maybe_sample_link();
self.maybe_send_heartbeat()?;
self.drain_feedback()
}
pub fn fb_loss(&self) -> f64 {
self.last_fwd_loss as f64
}
pub fn drain_until_acked(&mut self, timeout: Duration) -> io::Result<bool> {
let start = Instant::now();
while self.enc.pending_len() > 0 {
if start.elapsed() > timeout {
return Ok(false);
}
self.maybe_send_heartbeat()?;
self.drain_feedback()?;
std::thread::sleep(Duration::from_micros(200));
}
Ok(true)
}
fn drain_feedback(&mut self) -> io::Result<()> {
let mut buf = [0u8; CONTROL_RECV_BUF];
loop {
let res = match self.sock.as_udp() {
Some(u) => recv_with_ttl(u, &mut buf),
None => self.sock.recv(&mut buf).map(|n| (n, None)),
};
match res {
Ok((n, ttl)) => {
if let Some(t) = ttl {
self.asym
.observe_reverse(crate::path_sensor::hop_count_from_ttl(t));
}
if let Some(cp) = decode_control(&buf[..n]) {
if cp.session_announce.is_some_and(|e| e != self.enc.epoch()) {
continue;
}
}
if let Some(cp) = decode_control(&buf[..n]) {
self.ctrl_recv = self.ctrl_recv.wrapping_add(1);
if let Some(sc) = cp.session_challenge {
let mut ans = ControlPacket::new();
ans.session_response = Some(sc);
let wire = encode_control(&ans);
self.sock.send(&wire).ok();
}
self.last_feedback_at = Instant::now();
let was_dead = self.link_dead;
self.link_dead = false;
if was_dead {
self.recovery_target = self.enc.next_block_id();
self.recovery_started_us = self.start.elapsed().as_micros() as u64;
}
if let Some(la) = cp.loss_acct {
if la.seq > self.peer_seq {
self.peer_seq = la.seq;
}
let missed = self.peer_seq.saturating_sub(self.ctrl_recv);
self.rev_loss =
(missed as f32 / self.peer_seq.max(1) as f32).clamp(0.0, 1.0);
}
if let Some(p) = cp.path {
self.path_sensor.observe(p.ttl, p.ecn, p.hop_count);
self.asym.observe_forward(p.hop_count);
if p.ect_count > 0 {
self.ce_rate =
(p.ce_count as f32 / p.ect_count as f32).clamp(0.0, 1.0);
}
}
if let Some(pm) = cp.pmtu {
if self.peer_pmtu != 0 && pm.pmtu != 0 && pm.pmtu < self.peer_pmtu {
self.peer_pmtu_shift = 1.0;
}
if pm.pmtu != 0 {
self.peer_pmtu = pm.pmtu;
}
}
if let Some(ab) = cp.avail_bw {
self.avail_bw_kbps = ab.avail_kbps;
self.wbest_capacity_kbps = ab.capacity_kbps;
}
if let Some(fc) = cp.forecast {
self.forecast_bps = fc.forecast_kbps * 1000 / 8;
}
if let Some(pe) = cp.periodicity {
self.leo_period_s = pe.period_ds as f32 / 10.0;
self.leo_secs_to_spike = pe.secs_to_spike_ds as f32 / 10.0;
self.leo_conf = pe.confidence_x255 as f32 / 255.0;
}
let fb = feedback_from_control(&cp);
let rtx = self.enc.on_feedback(&fb);
self.send_batch(&rtx)?;
if was_dead && self.proactive_recovery {
let gap = self.enc.retransmit_all_data();
if !gap.is_empty() {
self.recovered_blocks += self.enc.pending_len() as u64;
self.recovery_dgrams.extend(gap);
self.last_recovery_us = self.start.elapsed().as_micros() as u64;
self.recovery_tokens = 0.0;
}
}
if self.recovery_target != 0 && fb.ack_through >= self.recovery_target {
self.recovery_interval_us = (self.start.elapsed().as_micros() as u64)
.saturating_sub(self.recovery_started_us);
self.recovery_target = 0;
}
let now_us = self.start.elapsed().as_micros() as u64;
let mut rtt_us = 0u64;
let mut newest_send = 0u64;
while let Some(&(id, sent)) = self.block_send_us.front() {
if id < fb.ack_through {
newest_send = sent;
rtt_us = now_us.saturating_sub(sent);
self.block_send_us.pop_front();
} else {
break;
}
}
self.path_model
.on_ack(fb.ack_through as u64, now_us, rtt_us, newest_send);
if rtt_us > 0 {
self.rtt_shape.observe(rtt_us as f64);
}
self.apply_fusion(&fb);
}
}
Err(e)
if e.kind() == io::ErrorKind::WouldBlock
|| e.kind() == io::ErrorKind::TimedOut
|| e.kind() == io::ErrorKind::ConnectionReset
|| e.kind() == io::ErrorKind::ConnectionRefused
|| e.kind() == io::ErrorKind::HostUnreachable
|| e.kind() == io::ErrorKind::NetworkUnreachable =>
{
break;
}
Err(e) => return Err(e),
}
}
self.drain_recovery()?;
self.check_liveness()?;
Ok(())
}
fn drain_recovery(&mut self) -> io::Result<()> {
if self.recovery_dgrams.is_empty() {
return Ok(());
}
let now_us = self.start.elapsed().as_micros() as u64;
let elapsed = now_us.saturating_sub(self.last_recovery_us);
self.last_recovery_us = now_us;
let rate_bytes = (self.path_model.btlbw_bps() / 8).max(MIN_RECOVERY_BYTES_PER_S) as f64;
self.recovery_tokens += rate_bytes * elapsed as f64 / 1_000_000.0;
if self.recovery_tokens > RECOVERY_BUCKET_BYTES {
self.recovery_tokens = RECOVERY_BUCKET_BYTES;
}
while let Some(front) = self.recovery_dgrams.front() {
let size = front.len() as f64;
if self.recovery_tokens < size {
break;
}
self.recovery_tokens -= size;
let dgram = self.recovery_dgrams.pop_front().expect("front exists");
self.send(&dgram)?;
}
let grace = RECOVERY_GRACE_RTTS * self.path_model.rtt_now_us().max(MIN_PACE_INTERVAL_US);
self.recovery_grace_until_us = now_us + grace;
Ok(())
}
fn check_liveness(&mut self) -> io::Result<()> {
let silence_us = self.last_feedback_at.elapsed().as_micros() as u64;
let dead_timeout =
(DEAD_RTT_MULTIPLE * self.path_model.rtt_now_us()).max(DEAD_FLOOR_US);
if silence_us <= dead_timeout {
return Ok(());
}
if !self.link_dead {
self.link_dead = true;
self.dead_episodes += 1;
}
if self.last_probe_at.elapsed().as_micros() as u64 >= dead_timeout
&& let Some(oldest) = self.enc.oldest_pending()
{
let probe = self.enc.probe_block(oldest);
if !probe.is_empty() {
self.send_batch(&probe)?;
self.probes_sent += 1;
}
self.last_probe_at = Instant::now();
}
Ok(())
}
fn maybe_send_heartbeat(&mut self) -> io::Result<()> {
if self.last_hb.elapsed() >= HEARTBEAT_INTERVAL {
let mut cp = ControlPacket::new();
cp.session_announce = Some(self.enc.epoch());
cp.timing = Some(TimingFrame {
send_ts: self.start.elapsed().as_micros() as u64,
echo_ts: 0,
});
cp.ring = Some(RingFrame {
fill_pct: self.enc.in_flight().min(255) as u8,
ring_kind: 0,
producers: 1,
consumers: 1,
trend: 1,
flags: 0,
});
self.ctrl_out = self.ctrl_out.wrapping_add(1);
cp.loss_acct = Some(LossAcctFrame {
seq: self.ctrl_out,
last_recv_seq: self.ctrl_recv,
});
cp.link = Some(LinkFrame {
class: self.inferred_link_class().as_u8(),
quality: self.link_quality,
});
if let Some(pm) = self.net_events.pmtu() {
cp.pmtu = Some(PmtuFrame { pmtu: pm });
}
let buf = encode_control(&cp);
self.send(&buf)?;
self.last_hb = Instant::now();
}
Ok(())
}
fn maybe_send_bw_probe(&mut self) -> io::Result<()> {
if self.last_bw_probe.elapsed() < BW_PROBE_INTERVAL {
return Ok(());
}
let round = self.bw_probe_round;
self.bw_probe_round = self.bw_probe_round.wrapping_add(1);
let total = 2 * BW_PROBE_PAIRS + BW_PROBE_TRAIN;
for idx in 0..total {
let mut cp = ControlPacket::new();
cp.bw_probe.push(crate::control_frame::BwProbeFrame {
probe_id: round,
idx,
send_ts: self.start.elapsed().as_micros() as u64,
});
let mut buf = encode_control(&cp);
crate::control_frame::pad_control_to(&mut buf, BW_PROBE_BYTES);
self.send(&buf)?;
}
self.last_bw_probe = Instant::now();
Ok(())
}
pub fn avail_bw_bps(&self) -> (u64, u64) {
(self.avail_bw_kbps * 1000, self.wbest_capacity_kbps * 1000)
}
#[cfg(target_os = "linux")]
fn maybe_send_trace(&mut self) -> io::Result<()> {
use std::os::fd::AsRawFd;
let Some(fd) = self.sock.as_udp().map(|u| u.as_raw_fd()) else {
return Ok(());
};
if self.last_trace.elapsed() >= TRACE_INTERVAL {
self.trace_round = self.trace_round.wrapping_add(1);
let now = self.start.elapsed().as_micros() as u64;
for ttl in 1..=MAX_TRACE_HOPS {
let mut cp = ControlPacket::new();
cp.trace.push(crate::control_frame::TraceFrame {
hop_ttl: ttl,
probe_id: self.trace_round,
});
let buf = encode_control(&cp);
crate::trace_sensor::send_at_ttl(fd, self.trace_peer, &buf, ttl)
.ok();
self.trace_send_us[ttl as usize] = now;
}
self.last_trace = Instant::now();
}
let now = self.start.elapsed().as_micros() as u64;
for (router, payload) in crate::trace_sensor::drain_icmp_errors(fd) {
if let Some(cp) = decode_control(&payload)
&& let Some(tf) = cp.trace.first()
{
let ttl = tf.hop_ttl;
let sent = self.trace_send_us.get(ttl as usize).copied().unwrap_or(0);
let rtt_us = now.saturating_sub(sent);
if !self.trace_hops.iter().any(|h| h.ttl == ttl) {
self.trace_hops.push(crate::trace_sensor::TraceHop {
ttl,
addr: router,
rtt_us,
});
self.trace_hops.sort_by_key(|h| h.ttl);
}
}
}
Ok(())
}
#[cfg(not(target_os = "linux"))]
fn maybe_send_trace(&mut self) -> io::Result<()> {
Ok(())
}
pub fn trace_hops(&self) -> &[crate::trace_sensor::TraceHop] {
&self.trace_hops
}
pub fn path_asymmetry(&self) -> (Option<u8>, Option<u8>, Option<u8>) {
(self.asym.forward(), self.asym.reverse(), self.asym.asymmetry())
}
pub fn ce_rate(&self) -> f32 {
self.ce_rate
}
pub fn forecast_bps(&self) -> u64 {
self.forecast_bps * 8
}
fn leo_prearm_shift(&self) -> f32 {
const LEO_PRE_ARM_WINDOW_S: f32 = 2.0;
if self.leo_conf >= 0.4
&& self.leo_period_s > 0.0
&& self.leo_secs_to_spike <= LEO_PRE_ARM_WINDOW_S
{
self.leo_conf
} else {
0.0
}
}
pub fn leo_cadence(&self) -> (f32, f32, f32) {
(self.leo_period_s, self.leo_conf, self.leo_secs_to_spike)
}
fn maybe_sample_link(&mut self) {
if self.last_link_sample.elapsed() >= LINK_SAMPLE_INTERVAL {
let snap = self.link_sensor.sample();
self.link_stress = snap.link_stress();
if self.link_class != LinkClass::Unknown && self.link_class != snap.class {
self.class_shift = 1.0;
}
self.link_class = snap.class;
self.link_quality = snap
.signal_quality
.unwrap_or(((1.0 - self.link_stress) * 100.0) as u8);
self.link_phy_kbps = snap.phy_rate_kbps.unwrap_or(0);
self.link_mcs_norm = snap.mcs_norm.unwrap_or(0.0);
self.last_link_sample = Instant::now();
}
self.class_shift *= 0.9;
self.peer_pmtu_shift *= 0.9;
}
fn apply_fusion(&mut self, fb: &crate::reliable_udp::Feedback) {
if fb.loss_class != 0 {
let contribution = match fb.loss_class {
2 => 1.0,
3 => 0.5,
_ => 0.0,
};
self.congestion_fraction += (contribution - self.congestion_fraction) * 0.125;
}
let event_shift = self.net_events.path_shift().max(self.peer_pmtu_shift);
self.net_event_shift_peak = self.net_event_shift_peak.max(event_shift);
let snap = SensorSnapshot {
loss: fb.loss_x255 as f32 / 255.0,
burstiness: fb.burstiness_x255 as f32 / 255.0,
owd_trend: match fb.owd_trend_class {
2 => 0.1,
0 => -0.1,
_ => 0.0,
},
link_stress: self.link_stress,
path_shift: self
.path_sensor
.path_shift()
.max(self.class_shift)
.max(event_shift)
.max(self.leo_prearm_shift()),
ecn_ce: self.ce_rate.max(self.path_sensor.ecn_ce()),
congestion_fraction: self.congestion_fraction,
rev_loss: self.rev_loss,
queue_delay_ms: self.path_model.queue_delay_us() as f32 / 1000.0,
backhaul_hops: self.backhaul_hops(),
};
self.last_fwd_loss = snap.loss;
let d = self.fusion.decide(&snap);
self.control.set_level(d.level);
self.control.set_parity_r(d.parity_r);
self.control.set_interleave_depth(d.interleave_depth);
self.enc.set_parity_covering(d.parity_r as usize, snap.loss);
self.pace_flow_window(snap.queue_delay_ms);
}
fn pace_flow_window(&mut self, queue_delay_ms: f32) {
if !self.pacing_enabled {
return;
}
let now = self.start.elapsed().as_micros() as u64;
if !self.recovery_dgrams.is_empty() || now < self.recovery_grace_until_us {
return;
}
let interval = self.path_model.rtt_now_us().max(MIN_PACE_INTERVAL_US);
if now < self.last_pace_us + interval {
return;
}
self.last_pace_us = now;
let off_target = (PACE_TARGET_MS - queue_delay_ms) / PACE_TARGET_MS;
let w = self.paced_window;
let step = off_target.clamp(-0.5 * w, w);
self.paced_window = (w + step).clamp(MIN_PACED_WINDOW as f32, self.flow_window_max as f32);
let mut target = self.paced_window.round() as u32;
let btlbw = self.path_model.btlbw_bps();
if self.forecast_bps > 0 && btlbw > 0 {
let ratio =
(self.forecast_bps as f64 * FORECAST_HEADROOM / btlbw as f64).clamp(0.1, 1.0);
let cap = ((self.flow_window_max as f64 * ratio).ceil() as u32).max(MIN_PACED_WINDOW);
target = target.min(cap);
}
if target != self.enc.flow_window() {
self.enc.set_flow_window(target);
}
}
fn send(&self, pkt: &[u8]) -> io::Result<()> {
let mut spins = 0u32;
loop {
match self.sock.send(pkt) {
Ok(_) => return Ok(()),
Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
spins += 1;
if spins > 20_000 {
return Ok(());
}
std::thread::sleep(Duration::from_micros(50));
}
Err(e)
if matches!(
e.kind(),
io::ErrorKind::ConnectionReset
| io::ErrorKind::ConnectionRefused
| io::ErrorKind::HostUnreachable
| io::ErrorKind::NetworkUnreachable
) =>
{
return Ok(());
}
Err(e) => return Err(e),
}
}
}
fn send_batch(&self, pkts: &[Vec<u8>]) -> io::Result<()> {
if pkts.is_empty() {
return Ok(());
}
#[cfg(target_os = "linux")]
{
self.send_gso(pkts)
}
#[cfg(target_os = "freebsd")]
{
self.send_mmsg(pkts)
}
#[cfg(target_os = "windows")]
{
if uso_enabled() {
self.send_uso(pkts)
} else {
for pkt in pkts {
self.send(pkt)?;
}
Ok(())
}
}
#[cfg(not(any(
target_os = "linux",
target_os = "freebsd",
target_os = "windows"
)))]
{
for pkt in pkts {
self.send(pkt)?;
}
Ok(())
}
}
#[cfg(target_os = "linux")]
fn send_gso(&self, pkts: &[Vec<u8>]) -> io::Result<()> {
use std::os::fd::AsRawFd;
let fd = match self.sock.as_udp() {
Some(u) => u.as_raw_fd(),
None => {
for p in pkts {
self.sock.send(p)?;
}
return Ok(());
}
};
let mut buf: Vec<u8> = Vec::with_capacity(64 * 1500);
let mut i = 0usize;
while i < pkts.len() {
let seg = pkts[i].len();
buf.clear();
let mut j = i;
while j < pkts.len()
&& pkts[j].len() == seg
&& (j - i) < 64
&& buf.len() + seg <= 61440
{
buf.extend_from_slice(&pkts[j]);
j += 1;
}
if j - i <= 1 || seg == 0 || seg > u16::MAX as usize {
self.send(&pkts[i])?;
i += 1;
continue;
}
if !self.send_gso_buf(fd, &buf, seg as u16)? {
return self.send_mmsg(&pkts[i..]);
}
i = j;
}
Ok(())
}
#[cfg(target_os = "linux")]
fn send_gso_buf(&self, fd: libc::c_int, buf: &[u8], seg_size: u16) -> io::Result<bool> {
const UDP_SEGMENT: libc::c_int = 103;
let mut iov = libc::iovec {
iov_base: buf.as_ptr() as *mut libc::c_void,
iov_len: buf.len(),
};
let mut cmsg_space = [0u64; 8]; let mut msg: libc::msghdr = unsafe { std::mem::zeroed() };
msg.msg_iov = &mut iov;
msg.msg_iovlen = 1;
msg.msg_control = cmsg_space.as_mut_ptr() as *mut libc::c_void;
msg.msg_controllen = unsafe { libc::CMSG_SPACE(size_of::<u16>() as u32) } as _;
unsafe {
let cmsg = libc::CMSG_FIRSTHDR(&msg);
(*cmsg).cmsg_level = libc::SOL_UDP;
(*cmsg).cmsg_type = UDP_SEGMENT;
(*cmsg).cmsg_len = libc::CMSG_LEN(size_of::<u16>() as u32) as _;
std::ptr::write_unaligned(libc::CMSG_DATA(cmsg) as *mut u16, seg_size);
}
let mut spins = 0u32;
loop {
let n = unsafe { libc::sendmsg(fd, &msg, 0) };
if n >= 0 {
return Ok(true);
}
let err = io::Error::last_os_error();
match err.raw_os_error() {
Some(libc::ENOPROTOOPT)
| Some(libc::EOPNOTSUPP)
| Some(libc::EINVAL)
| Some(libc::EIO) => {
return Ok(false);
}
_ => match err.kind() {
io::ErrorKind::WouldBlock => {
spins += 1;
if spins > 20_000 {
return Ok(true);
}
std::thread::sleep(Duration::from_micros(50));
}
io::ErrorKind::ConnectionReset
| io::ErrorKind::ConnectionRefused
| io::ErrorKind::HostUnreachable
| io::ErrorKind::NetworkUnreachable => {
return Ok(true);
}
_ => return Err(err),
},
}
}
}
#[cfg(target_os = "windows")]
fn send_uso(&self, pkts: &[Vec<u8>]) -> io::Result<()> {
use std::os::windows::io::AsRawSocket;
if self.sock.as_udp().is_none() {
for p in pkts {
self.sock.send(p)?;
}
return Ok(());
}
let sock = self.sock.as_udp().expect("Udp checked above").as_raw_socket() as usize;
let mut buf: Vec<u8> = Vec::with_capacity(64 * 1500);
let mut i = 0usize;
while i < pkts.len() {
let seg = pkts[i].len();
buf.clear();
let mut j = i;
while j < pkts.len()
&& pkts[j].len() == seg
&& (j - i) < 64
&& buf.len() + seg <= 61440
{
buf.extend_from_slice(&pkts[j]);
j += 1;
}
if j - i <= 1 || seg == 0 || seg > u16::MAX as usize {
self.send(&pkts[i])?;
i += 1;
continue;
}
if !self.send_uso_buf(sock, &buf, seg as u32)? {
for pkt in &pkts[i..] {
self.send(pkt)?;
}
return Ok(());
}
i = j;
}
Ok(())
}
#[cfg(target_os = "windows")]
fn send_uso_buf(&self, sock: usize, buf: &[u8], seg_size: u32) -> io::Result<bool> {
use windows_sys::Win32::Networking::WinSock::{WSAGetLastError, WSABUF, WSAMSG};
let Some(wsasendmsg) = load_wsasendmsg(sock) else {
return Ok(false);
};
const IPPROTO_UDP: i32 = 17;
const UDP_SEND_MSG_SIZE: i32 = 2;
const SOCKET_ERROR: i32 = -1;
const WSAEINVAL: i32 = 10022;
const WSAEWOULDBLOCK: i32 = 10035;
const WSAEMSGSIZE: i32 = 10040;
const WSAENOPROTOOPT: i32 = 10042;
const WSAECONNRESET: i32 = 10054;
const WSAECONNREFUSED: i32 = 10061;
let mut data = WSABUF {
len: buf.len() as u32,
buf: buf.as_ptr() as *mut u8,
};
let mut ctrl = [0u64; 4];
let cp = ctrl.as_mut_ptr() as *mut u8;
unsafe {
std::ptr::write_unaligned(cp as *mut usize, 20usize);
std::ptr::write_unaligned(cp.add(8) as *mut i32, IPPROTO_UDP);
std::ptr::write_unaligned(cp.add(12) as *mut i32, UDP_SEND_MSG_SIZE);
std::ptr::write_unaligned(cp.add(16) as *mut u32, seg_size);
}
let msg = WSAMSG {
name: std::ptr::null_mut(),
namelen: 0,
lpBuffers: &mut data,
dwBufferCount: 1,
Control: WSABUF { len: 24, buf: cp },
dwFlags: 0,
};
let mut sent = 0u32;
let mut spins = 0u32;
loop {
let rc = unsafe {
wsasendmsg(
sock,
&msg,
0,
&mut sent,
std::ptr::null_mut(),
std::ptr::null(),
)
};
if rc != SOCKET_ERROR {
USO_OFFLOAD.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(true);
}
let err = unsafe { WSAGetLastError() };
match err {
WSAEINVAL | WSAENOPROTOOPT | WSAEMSGSIZE => {
USO_FALLBACK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
return Ok(false);
}
WSAEWOULDBLOCK => {
spins += 1;
if spins > 20_000 {
return Ok(true);
}
std::thread::sleep(Duration::from_micros(50));
}
WSAECONNRESET | WSAECONNREFUSED => return Ok(true),
_ => return Err(io::Error::from_raw_os_error(err)),
}
}
}
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
fn send_mmsg(&self, pkts: &[Vec<u8>]) -> io::Result<()> {
use std::os::fd::AsRawFd;
let fd = match self.sock.as_udp() {
Some(u) => u.as_raw_fd(),
None => {
for p in pkts {
self.sock.send(p)?;
}
return Ok(());
}
};
let mut iovecs: Vec<libc::iovec> = pkts
.iter()
.map(|p| libc::iovec {
iov_base: p.as_ptr() as *mut libc::c_void,
iov_len: p.len(),
})
.collect();
let mut msgs: Vec<libc::mmsghdr> = Vec::with_capacity(pkts.len());
for i in 0..pkts.len() {
let mut hdr: libc::mmsghdr = unsafe { std::mem::zeroed() };
hdr.msg_hdr.msg_iov = iovecs.as_mut_ptr().wrapping_add(i);
hdr.msg_hdr.msg_iovlen = 1 as _;
msgs.push(hdr);
}
let mut sent = 0usize;
let mut spins = 0u32;
while sent < msgs.len() {
let count = (msgs.len() - sent) as MmsgLen;
let n = unsafe { libc::sendmmsg(fd, msgs.as_mut_ptr().add(sent), count, 0) };
if n > 0 {
sent += n as usize;
spins = 0;
continue;
}
let err = io::Error::last_os_error();
match err.kind() {
io::ErrorKind::WouldBlock => {
spins += 1;
if spins > 20_000 {
return Ok(());
}
std::thread::sleep(Duration::from_micros(50));
}
io::ErrorKind::ConnectionReset | io::ErrorKind::ConnectionRefused => {
return Ok(());
}
_ => return Err(err),
}
}
Ok(())
}
}
struct RsSession {
sock: std::sync::Arc<crate::dgram::DgramSock>,
dec: Decoder,
last_data_at: Instant,
peer: Option<SocketAddr>,
recv_count: u64,
nak_history: BTreeMap<u32, Instant>,
last_feedback: Instant,
ctrl_out: u32,
ctrl_recv: u32,
peer_acked: u32,
ctrl_out_at_last_hb: u32,
peer_acked_at_last_hb: u32,
fb_loss_est: f32,
wbest: crate::wbest_sensor::WBestEstimator,
wbest_round: Option<u8>,
wbest_avail_kbps: u64,
wbest_capacity_kbps: u64,
peer_link_class: u8,
peer_link_quality: u8,
ack_interval: Duration,
fb_drop_pct: u32,
fb_drop_rng: u64,
fb_delay: Duration,
fb_pending: VecDeque<(Instant, Vec<u8>)>,
nak_batch: usize,
max_hold: Duration,
head_block: u32,
head_since: Instant,
start: Instant,
debug_drop_pct: u32,
drop_rng: u64,
ge_loss_p: u32,
ge_loss_r: u32,
ge_bad: bool,
drop_block_mod: u32,
burst_at: u64,
burst_len: u64,
connected: bool,
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
rbufs: Vec<Vec<u8>>,
#[cfg(target_os = "linux")]
gro_on: bool,
#[cfg(target_os = "linux")]
gro_buf: Vec<u8>,
last_ttl: u8,
last_tos: u8,
ce_count: u64,
ect_count: u64,
forecast: crate::forecast_sensor::ArrivalForecast,
fc_bytes: u64,
fc_last: Instant,
periodicity: crate::periodicity_sensor::PeriodicitySensor,
peer_pmtu: u16,
local_pmtu: u16,
}
pub struct ReliableUdpReceiver {
sock: std::sync::Arc<crate::dgram::DgramSock>,
sessions: HashMap<u32, RsSession>,
order: Vec<u32>,
pending_admissions: HashMap<u32, (SocketAddr, u64, Instant)>,
session_ceiling: Option<usize>,
session_refusals: u64,
session_nonce_seq: u64,
session_changed: bool,
session_admissions: u64,
session_admission_failures: u64,
start: Instant,
net_events: NetEventObserver,
multi_peer: bool,
net_event_shift_peak: f32,
cfg: RsSessionConfig,
}
#[derive(Clone, Copy)]
struct RsSessionConfig {
max_hold: Duration,
fb_delay: Duration,
nak_batch: usize,
debug_drop_pct: u32,
drop_rng: u64,
ge_loss_p: u32,
ge_loss_r: u32,
drop_block_mod: u32,
burst_at: u64,
burst_len: u64,
fb_drop_pct: u32,
fb_drop_rng: u64,
}
#[cfg(target_os = "linux")]
fn enable_gro(sock: &UdpSocket) -> bool {
use std::os::fd::AsRawFd;
const UDP_GRO: libc::c_int = 104;
let on: libc::c_int = 1;
let rc = unsafe {
libc::setsockopt(
sock.as_raw_fd(),
libc::SOL_UDP,
UDP_GRO,
&on as *const libc::c_int as *const libc::c_void,
size_of::<libc::c_int>() as libc::socklen_t,
)
};
rc == 0
}
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
fn enable_ttl_ecn(sock: &UdpSocket) {
use std::os::fd::AsRawFd;
let fd = sock.as_raw_fd();
let on: libc::c_int = 1;
let set = |opt: libc::c_int| unsafe {
libc::setsockopt(
fd,
libc::IPPROTO_IP,
opt,
&on as *const libc::c_int as *const libc::c_void,
size_of::<libc::c_int>() as libc::socklen_t,
);
};
set(libc::IP_RECVTTL);
set(libc::IP_RECVTOS);
}
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
fn set_ect(sock: &UdpSocket) {
use std::os::fd::AsRawFd;
let tos: libc::c_int = 0b10;
unsafe {
libc::setsockopt(
sock.as_raw_fd(),
libc::IPPROTO_IP,
libc::IP_TOS,
&tos as *const libc::c_int as *const libc::c_void,
size_of::<libc::c_int>() as libc::socklen_t,
);
}
}
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
fn cmsg_scalar_u8(cmsg: *const libc::cmsghdr) -> u8 {
unsafe {
let hdr_len = libc::CMSG_LEN(0) as usize;
let total: usize = (*cmsg).cmsg_len as _;
let payload = total.saturating_sub(hdr_len);
if payload >= size_of::<libc::c_int>() {
let mut v: libc::c_int = 0;
std::ptr::copy_nonoverlapping(
libc::CMSG_DATA(cmsg),
&mut v as *mut libc::c_int as *mut u8,
size_of::<libc::c_int>(),
);
v as u8
} else if payload >= 1 {
let mut b: u8 = 0;
std::ptr::copy_nonoverlapping(libc::CMSG_DATA(cmsg), &mut b, 1);
b
} else {
0
}
}
}
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
fn recv_with_ttl(sock: &UdpSocket, buf: &mut [u8]) -> io::Result<(usize, Option<u8>)> {
use std::mem::zeroed;
use std::os::fd::AsRawFd;
unsafe {
let mut iov = libc::iovec {
iov_base: buf.as_mut_ptr() as *mut libc::c_void,
iov_len: buf.len(),
};
let mut cbuf = [0u8; 64];
let mut msg: libc::msghdr = zeroed();
msg.msg_iov = &mut iov;
msg.msg_iovlen = 1;
msg.msg_control = cbuf.as_mut_ptr() as *mut libc::c_void;
msg.msg_controllen = cbuf.len() as _;
let n = libc::recvmsg(sock.as_raw_fd(), &mut msg, 0);
if n < 0 {
return Err(io::Error::last_os_error());
}
let mut ttl = None;
let mut cmsg = libc::CMSG_FIRSTHDR(&msg);
while !cmsg.is_null() {
if (*cmsg).cmsg_level == libc::IPPROTO_IP
&& ((*cmsg).cmsg_type == libc::IP_TTL || (*cmsg).cmsg_type == libc::IP_RECVTTL)
{
ttl = Some(cmsg_scalar_u8(cmsg));
}
cmsg = libc::CMSG_NXTHDR(&msg, cmsg);
}
Ok((n as usize, ttl))
}
}
#[cfg(not(any(target_os = "linux", target_os = "freebsd")))]
fn recv_with_ttl(sock: &UdpSocket, buf: &mut [u8]) -> io::Result<(usize, Option<u8>)> {
sock.recv(buf).map(|n| (n, None))
}
#[cfg(target_os = "windows")]
fn enable_ttl_ecn_win(sock: &UdpSocket) {
use std::os::windows::io::AsRawSocket;
use windows_sys::Win32::Networking::WinSock::{
setsockopt, IPPROTO_IP, IP_ECN, IP_HOPLIMIT, IP_RECVTOS,
};
let s = sock.as_raw_socket() as usize;
let on: i32 = 1;
let set = |opt: i32| unsafe {
setsockopt(
s,
IPPROTO_IP,
opt,
&on as *const i32 as *const u8,
size_of::<i32>() as i32,
);
};
set(IP_HOPLIMIT);
set(IP_RECVTOS);
set(IP_ECN);
}
#[cfg(target_os = "linux")]
fn gro_wanted() -> bool {
static EN: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*EN.get_or_init(|| std::env::var("SUBETHA_GRO").map(|v| v != "0").unwrap_or(true))
}
#[cfg(target_os = "linux")]
static GRO_RECVMSG: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
#[cfg(target_os = "linux")]
static GRO_SEGMENTS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
pub fn gro_stats() -> (u64, u64) {
#[cfg(target_os = "linux")]
{
use std::sync::atomic::Ordering::Relaxed;
(GRO_RECVMSG.load(Relaxed), GRO_SEGMENTS.load(Relaxed))
}
#[cfg(not(target_os = "linux"))]
{
(0, 0)
}
}
impl RsSession {
fn new(sock: std::sync::Arc<crate::dgram::DgramSock>, cfg: RsSessionConfig) -> Self {
Self {
sock,
dec: Decoder::new(),
last_data_at: Instant::now(),
peer: None,
recv_count: 0,
nak_history: BTreeMap::new(),
last_feedback: Instant::now(),
ctrl_out: 0,
ctrl_recv: 0,
peer_acked: 0,
ctrl_out_at_last_hb: 0,
peer_acked_at_last_hb: 0,
fb_loss_est: 0.0,
wbest: crate::wbest_sensor::WBestEstimator::new(BW_PROBE_BYTES),
wbest_round: None,
wbest_avail_kbps: 0,
wbest_capacity_kbps: 0,
peer_link_class: 0,
peer_link_quality: 0,
ack_interval: ACK_INTERVAL,
fb_drop_pct: cfg.fb_drop_pct,
fb_drop_rng: cfg.fb_drop_rng,
fb_delay: cfg.fb_delay,
fb_pending: VecDeque::new(),
nak_batch: cfg.nak_batch,
max_hold: cfg.max_hold,
head_block: 0,
head_since: Instant::now(),
start: Instant::now(),
debug_drop_pct: cfg.debug_drop_pct,
drop_rng: cfg.drop_rng,
ge_loss_p: cfg.ge_loss_p,
ge_loss_r: cfg.ge_loss_r,
ge_bad: false,
drop_block_mod: cfg.drop_block_mod,
burst_at: cfg.burst_at,
burst_len: cfg.burst_len,
connected: false,
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
rbufs: Vec::new(),
#[cfg(target_os = "linux")]
gro_on: false,
#[cfg(target_os = "linux")]
gro_buf: Vec::new(),
last_ttl: 0,
last_tos: 0,
ce_count: 0,
ect_count: 0,
forecast: crate::forecast_sensor::ArrivalForecast::new(),
fc_bytes: 0,
fc_last: Instant::now(),
periodicity: crate::periodicity_sensor::PeriodicitySensor::new(),
peer_pmtu: 0,
local_pmtu: 0,
}
}
pub fn recv_count(&self) -> u64 {
self.recv_count
}
pub fn peak_loss_x255(&self) -> u8 {
self.dec.peak_loss_x255()
}
pub fn false_recovery_count(&self) -> u64 {
self.dec.false_recovery_count()
}
pub fn set_ge_burst(&mut self, on: bool) {
self.dec.set_ge_burst(on);
}
pub fn mean_burst_len(&self) -> f32 {
self.dec.mean_burst_len()
}
pub fn owd_skew(&self) -> f64 {
self.dec.owd_skew()
}
pub fn owd_trend_debiased(&self) -> f64 {
self.dec.owd_trend_debiased()
}
pub fn ack_interval(&self) -> Duration {
self.ack_interval
}
fn update_feedback_cadence(&mut self) {
let d_out = self.ctrl_out.saturating_sub(self.ctrl_out_at_last_hb);
let d_peer = self.peer_acked.saturating_sub(self.peer_acked_at_last_hb);
self.ctrl_out_at_last_hb = self.ctrl_out;
self.peer_acked_at_last_hb = self.peer_acked;
if d_out >= 10 {
let fb_loss = d_out.saturating_sub(d_peer) as f32 / d_out as f32;
self.fb_loss_est = fb_loss;
self.ack_interval = if fb_loss > 0.2 {
ACK_INTERVAL / 4
} else {
ACK_INTERVAL
};
}
}
pub fn feedback_loss_est(&self) -> f32 {
self.fb_loss_est
}
pub fn peer_link(&self) -> (u8, u8) {
(self.peer_link_class, self.peer_link_quality)
}
pub fn accecn_counts(&self) -> (u64, u64) {
(self.ce_count, self.ect_count)
}
pub fn forecast_bps(&self) -> u64 {
(self.forecast.forecast_bps() * 8.0) as u64
}
pub fn leo_cadence(&self) -> Option<(f64, f64, f64)> {
let (period, conf) = self.periodicity.detected_period()?;
Some((period, conf, self.periodicity.secs_to_next_spike().unwrap_or(0.0)))
}
pub fn peer_pmtu(&self) -> u16 {
self.peer_pmtu
}
pub fn head_status(&self) -> Option<(u32, u32, usize, bool)> {
self.dec.head_status()
}
pub fn set_debug_loss(&mut self, pct: u32) {
self.debug_drop_pct = pct.min(100);
}
fn drop_whole_block(&self, buf: &[u8]) -> bool {
if self.drop_block_mod == 0 || buf.len() < 5 || is_outer_datagram(buf) {
return false;
}
let bid = u32::from_le_bytes([buf[1], buf[2], buf[3], buf[4]]);
bid.is_multiple_of(self.drop_block_mod)
}
#[inline]
fn in_burst(&self) -> bool {
self.burst_len != 0
&& self.recv_count >= self.burst_at
&& self.recv_count < self.burst_at + self.burst_len
}
#[inline]
fn roll_drop(&mut self) -> bool {
if self.ge_loss_r > 0 {
let drop = self.ge_bad;
self.drop_rng = self
.drop_rng
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let roll = ((self.drop_rng >> 33) as u32) % 10000;
if self.ge_bad {
if roll < self.ge_loss_r {
self.ge_bad = false;
}
} else if roll < self.ge_loss_p {
self.ge_bad = true;
}
return drop;
}
if self.debug_drop_pct == 0 {
return false;
}
self.drop_rng = self
.drop_rng
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((self.drop_rng >> 33) as u32) % 100 < self.debug_drop_pct
}
fn recv_demux_drain(&mut self, out: &mut Vec<Vec<u8>>) -> io::Result<bool> {
let mut buf = [0u8; RECV_BUF];
let mut idle = true;
loop {
match self.sock.recv_from(&mut buf) {
Ok((n, src)) => {
self.peer = Some(src);
self.process_datagram(&buf[..n], out);
idle = false;
}
Err(e) if e.kind() == io::ErrorKind::WouldBlock => break,
Err(e) => return Err(e),
}
}
Ok(idle)
}
fn process_datagram(&mut self, buf: &[u8], out: &mut Vec<Vec<u8>>) {
self.recv_count += 1;
self.last_data_at = Instant::now();
if self.roll_drop() {
return;
}
if is_control(buf) {
if let Some(cp) = decode_control(buf) {
self.ctrl_recv = self.ctrl_recv.wrapping_add(1);
if let Some(la) = cp.loss_acct
&& la.last_recv_seq > self.peer_acked
{
self.peer_acked = la.last_recv_seq;
}
if let Some(lk) = cp.link {
self.peer_link_class = lk.class;
self.peer_link_quality = lk.quality;
}
if let Some(pm) = cp.pmtu
&& pm.pmtu != 0
{
self.peer_pmtu = pm.pmtu;
}
if let Some(announced) = cp.session_announce
&& self.dec.session_epoch().is_some_and(|e| e != announced)
{
self.dec.note_unknown_epoch(announced);
}
if let Some(t) = cp.timing {
let recv_ts = self.start.elapsed().as_micros() as u64;
self.dec.on_heartbeat(t.send_ts, recv_ts);
let owd = recv_ts as f64 - t.send_ts as f64;
self.periodicity.observe(owd, recv_ts);
}
if !cp.bw_probe.is_empty() {
let arrival_us = self.start.elapsed().as_nanos() as f64 / 1000.0;
self.ingest_bw_probe(&cp.bw_probe, arrival_us);
}
self.update_feedback_cadence();
}
} else if self.drop_whole_block(buf) {
} else if self.in_burst() {
} else {
let ecn = self.last_tos & 0b11;
if ecn != 0 {
self.ect_count += 1;
if ecn == crate::path_sensor::ECN_CE {
self.ce_count += 1;
}
}
self.fc_bytes += buf.len() as u64;
let recv_us = self.start.elapsed().as_micros() as u64;
out.extend(self.dec.on_packet_at(buf, recv_us));
}
}
fn maybe_observe_forecast(&mut self) {
let dt = self.fc_last.elapsed();
if dt >= FORECAST_TICK {
self.forecast.observe(self.fc_bytes, dt.as_secs_f64());
self.fc_bytes = 0;
self.fc_last = Instant::now();
}
}
fn ingest_bw_probe(&mut self, probes: &[crate::control_frame::BwProbeFrame], arrival_us: f64) {
let pair_probes = 2 * BW_PROBE_PAIRS;
for f in probes {
if self.wbest_round != Some(f.probe_id) {
self.wbest.reset();
self.wbest_round = Some(f.probe_id);
}
if f.idx < pair_probes {
self.wbest.on_pair_probe(f.idx % 2, arrival_us);
} else {
self.wbest.on_train_probe(arrival_us);
}
}
if let Some(c) = self.wbest.effective_capacity_bps() {
self.wbest_capacity_kbps = (c / 1000.0) as u64;
}
if let Some(a) = self.wbest.available_bps() {
self.wbest_avail_kbps = (a / 1000.0) as u64;
}
}
pub fn wbest_bps(&self) -> (u64, u64) {
(self.wbest_avail_kbps * 1000, self.wbest_capacity_kbps * 1000)
}
fn recv_into(&mut self, out: &mut Vec<Vec<u8>>) -> io::Result<bool> {
#[cfg(target_os = "linux")]
if self.connected {
if self.gro_on {
return self.recv_gro(out);
}
return self.recv_batch(out);
}
#[cfg(target_os = "freebsd")]
if self.connected {
return self.recv_batch(out);
}
#[cfg(target_os = "windows")]
if self.connected {
return self.recv_wsamsg(out);
}
let mut buf = [0u8; RECV_BUF];
match self.sock.recv_from(&mut buf) {
Ok((n, src)) => {
self.peer = Some(src);
#[cfg(target_os = "linux")]
{
self.connected = self.sock.connect(src).is_ok();
self.gro_on = self.connected
&& gro_wanted()
&& self.sock.as_udp().map(enable_gro).unwrap_or(false);
}
#[cfg(target_os = "freebsd")]
{
self.connected = self.sock.connect(src).is_ok();
}
#[cfg(target_os = "windows")]
{
self.connected = self.sock.connect(src).is_ok();
}
self.process_datagram(&buf[..n], out);
Ok(false)
}
Err(e)
if matches!(
e.kind(),
io::ErrorKind::WouldBlock
| io::ErrorKind::TimedOut
| io::ErrorKind::ConnectionReset
| io::ErrorKind::ConnectionRefused
) =>
{
Ok(true)
}
Err(e) => Err(e),
}
}
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
fn observe_ttl_tos(&mut self, msg: &libc::msghdr) {
unsafe {
let mut cmsg = libc::CMSG_FIRSTHDR(msg as *const libc::msghdr);
while !cmsg.is_null() {
let level = (*cmsg).cmsg_level;
let cty = (*cmsg).cmsg_type;
if level == libc::IPPROTO_IP
&& (cty == libc::IP_TTL || cty == libc::IP_RECVTTL)
{
self.last_ttl = cmsg_scalar_u8(cmsg);
} else if level == libc::IPPROTO_IP
&& (cty == libc::IP_TOS || cty == libc::IP_RECVTOS)
{
self.last_tos = cmsg_scalar_u8(cmsg);
}
cmsg = libc::CMSG_NXTHDR(msg as *const libc::msghdr, cmsg);
}
}
}
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
fn recv_batch(&mut self, out: &mut Vec<Vec<u8>>) -> io::Result<bool> {
use std::os::fd::AsRawFd;
const RECV_BATCH: usize = 32;
const CMSG_WORDS: usize = 8;
if self.rbufs.len() < RECV_BATCH {
self.rbufs.resize_with(RECV_BATCH, || vec![0u8; RECV_BUF]);
}
if self.sock.as_udp().is_none() {
return self.recv_demux_drain(out);
}
let fd = self.sock.as_udp().expect("Udp checked above").as_raw_fd();
let mut iovecs: Vec<libc::iovec> = self
.rbufs
.iter_mut()
.take(RECV_BATCH)
.map(|b| libc::iovec {
iov_base: b.as_mut_ptr() as *mut libc::c_void,
iov_len: RECV_BUF,
})
.collect();
let mut ctrl: Vec<[u64; CMSG_WORDS]> = vec![[0u64; CMSG_WORDS]; RECV_BATCH];
let mut msgs: Vec<libc::mmsghdr> = Vec::with_capacity(RECV_BATCH);
for (i, slot) in ctrl.iter_mut().enumerate() {
let mut hdr: libc::mmsghdr = unsafe { std::mem::zeroed() };
hdr.msg_hdr.msg_iov = iovecs.as_mut_ptr().wrapping_add(i);
hdr.msg_hdr.msg_iovlen = 1 as _;
hdr.msg_hdr.msg_control = slot.as_mut_ptr() as *mut libc::c_void;
hdr.msg_hdr.msg_controllen = (CMSG_WORDS * size_of::<u64>()) as _;
msgs.push(hdr);
}
let mut ts = libc::timespec {
tv_sec: 0,
tv_nsec: 4_000_000,
};
let n = unsafe {
libc::recvmmsg(
fd,
msgs.as_mut_ptr(),
RECV_BATCH as MmsgLen,
libc::MSG_WAITFORONE,
&mut ts as *mut libc::timespec,
)
};
if n == 0 {
return Ok(true);
}
if n < 0 {
let e = io::Error::last_os_error();
return match e.kind() {
io::ErrorKind::WouldBlock
| io::ErrorKind::TimedOut
| io::ErrorKind::ConnectionReset
| io::ErrorKind::ConnectionRefused => Ok(true),
_ => Err(e),
};
}
for (i, msg) in msgs.iter().take(n as usize).enumerate() {
let len = msg.msg_len as usize;
if len == 0 || len > RECV_BUF {
continue;
}
self.observe_ttl_tos(&msg.msg_hdr);
let mut tmp = [0u8; RECV_BUF];
tmp[..len].copy_from_slice(&self.rbufs[i][..len]);
self.process_datagram(&tmp[..len], out);
}
Ok(false)
}
#[cfg(target_os = "linux")]
fn recv_gro(&mut self, out: &mut Vec<Vec<u8>>) -> io::Result<bool> {
use std::os::fd::AsRawFd;
use std::sync::atomic::Ordering::Relaxed;
const UDP_GRO: libc::c_int = 104;
const GRO_BUF: usize = 65536;
if self.gro_buf.len() < GRO_BUF {
self.gro_buf.resize(GRO_BUF, 0);
}
if self.sock.as_udp().is_none() {
return self.recv_demux_drain(out);
}
let fd = self.sock.as_udp().expect("Udp checked above").as_raw_fd();
let mut got_any = false;
let mut first = true;
loop {
let mut iov = libc::iovec {
iov_base: self.gro_buf.as_mut_ptr() as *mut libc::c_void,
iov_len: GRO_BUF,
};
let mut cmsg_space = [0u64; 16];
let mut msg: libc::msghdr = unsafe { std::mem::zeroed() };
msg.msg_iov = &mut iov;
msg.msg_iovlen = 1;
msg.msg_control = cmsg_space.as_mut_ptr() as *mut libc::c_void;
msg.msg_controllen = (cmsg_space.len() * size_of::<u64>()) as _;
let flags = if first { 0 } else { libc::MSG_DONTWAIT };
let n = unsafe { libc::recvmsg(fd, &mut msg, flags) };
if n < 0 {
let e = io::Error::last_os_error();
return match e.kind() {
io::ErrorKind::WouldBlock
| io::ErrorKind::TimedOut
| io::ErrorKind::ConnectionReset
| io::ErrorKind::ConnectionRefused => Ok(!got_any),
_ => Err(e),
};
}
let n = n as usize;
let mut seg = n;
unsafe {
let mut cmsg = libc::CMSG_FIRSTHDR(&msg);
while !cmsg.is_null() {
let level = (*cmsg).cmsg_level;
let cty = (*cmsg).cmsg_type;
if level == libc::SOL_UDP && cty == UDP_GRO {
let mut s: libc::c_int = 0;
std::ptr::copy_nonoverlapping(
libc::CMSG_DATA(cmsg),
&mut s as *mut libc::c_int as *mut u8,
size_of::<libc::c_int>(),
);
if s > 0 {
seg = s as usize;
}
} else if level == libc::IPPROTO_IP && cty == libc::IP_TTL {
let mut t: libc::c_int = 0;
std::ptr::copy_nonoverlapping(
libc::CMSG_DATA(cmsg),
&mut t as *mut libc::c_int as *mut u8,
size_of::<libc::c_int>(),
);
self.last_ttl = t as u8;
} else if level == libc::IPPROTO_IP && cty == libc::IP_TOS {
let mut tos: u8 = 0;
std::ptr::copy_nonoverlapping(libc::CMSG_DATA(cmsg), &mut tos, 1);
self.last_tos = tos;
}
cmsg = libc::CMSG_NXTHDR(&msg, cmsg);
}
}
if seg == 0 {
seg = n;
}
let mut off = 0usize;
let mut segs = 0u64;
while off < n {
let end = (off + seg).min(n);
let len = end - off;
if len > 0 && len <= RECV_BUF {
let mut tmp = [0u8; RECV_BUF];
tmp[..len].copy_from_slice(&self.gro_buf[off..end]);
self.process_datagram(&tmp[..len], out);
segs += 1;
}
off = end;
}
GRO_RECVMSG.fetch_add(1, Relaxed);
GRO_SEGMENTS.fetch_add(segs, Relaxed);
got_any = true;
first = false;
if out.len() > 4096 {
return Ok(false);
}
}
}
#[cfg(target_os = "windows")]
fn recv_wsamsg(&mut self, out: &mut Vec<Vec<u8>>) -> io::Result<bool> {
use std::os::windows::io::AsRawSocket;
use windows_sys::Win32::Networking::WinSock::{WSAGetLastError, WSABUF, WSAMSG};
if self.sock.as_udp().is_none() {
return self.recv_demux_drain(out);
}
let sock = self.sock.as_udp().expect("Udp checked above").as_raw_socket() as usize;
let Some(wsarecvmsg) = load_wsarecvmsg(sock) else {
let mut buf = [0u8; RECV_BUF];
return match self.sock.recv(&mut buf) {
Ok(n) => {
self.process_datagram(&buf[..n], out);
Ok(false)
}
Err(e)
if matches!(
e.kind(),
io::ErrorKind::WouldBlock
| io::ErrorKind::TimedOut
| io::ErrorKind::ConnectionReset
| io::ErrorKind::ConnectionRefused
) =>
{
Ok(true)
}
Err(e) => Err(e),
};
};
const SOCKET_ERROR: i32 = -1;
const WSAEMSGSIZE: i32 = 10040;
const WSAEWOULDBLOCK: i32 = 10035;
const WSAETIMEDOUT: i32 = 10060;
const WSAECONNRESET: i32 = 10054;
const WSAECONNREFUSED: i32 = 10061;
let mut buf = [0u8; RECV_BUF];
let mut data = WSABUF {
len: RECV_BUF as u32,
buf: buf.as_mut_ptr(),
};
let mut ctrl = [0u64; 16];
let mut msg = WSAMSG {
name: std::ptr::null_mut(),
namelen: 0,
lpBuffers: &mut data,
dwBufferCount: 1,
Control: WSABUF {
len: (ctrl.len() * size_of::<u64>()) as u32,
buf: ctrl.as_mut_ptr() as *mut u8,
},
dwFlags: 0,
};
let mut recvd = 0u32;
let rc = unsafe {
wsarecvmsg(
sock,
&mut msg,
&mut recvd,
std::ptr::null_mut(),
std::ptr::null(),
)
};
if rc == SOCKET_ERROR {
let err = unsafe { WSAGetLastError() };
return match err {
WSAEWOULDBLOCK | WSAETIMEDOUT | WSAECONNRESET | WSAECONNREFUSED
| WSAEMSGSIZE => Ok(true),
_ => Err(io::Error::from_raw_os_error(err)),
};
}
let n = recvd as usize;
if n == 0 || n > RECV_BUF {
return Ok(true);
}
let ctrl_len = (msg.Control.len as usize).min(ctrl.len() * size_of::<u64>());
let cbytes = unsafe { std::slice::from_raw_parts(ctrl.as_ptr() as *const u8, ctrl_len) };
self.observe_wsa_cmsgs(cbytes);
self.process_datagram(&buf[..n], out);
Ok(false)
}
#[cfg(target_os = "windows")]
fn observe_wsa_cmsgs(&mut self, control: &[u8]) {
use windows_sys::Win32::Networking::WinSock::{IPPROTO_IP, IP_ECN, IP_TOS, IP_TTL};
const HDR: usize = 16;
let lvl_ip = IPPROTO_IP;
let mut off = 0usize;
while off + HDR <= control.len() {
let cmsg_len =
unsafe { std::ptr::read_unaligned(control.as_ptr().add(off) as *const usize) };
if cmsg_len < HDR || off + cmsg_len > control.len() {
break;
}
let level =
unsafe { std::ptr::read_unaligned(control.as_ptr().add(off + 8) as *const i32) };
let cty =
unsafe { std::ptr::read_unaligned(control.as_ptr().add(off + 12) as *const i32) };
if level == lvl_ip && cmsg_len - HDR >= size_of::<i32>() {
let val = unsafe {
std::ptr::read_unaligned(control.as_ptr().add(off + HDR) as *const i32)
};
if cty == IP_TTL {
self.last_ttl = val as u8;
} else if cty == IP_TOS || cty == IP_ECN {
self.last_tos = val as u8;
}
}
off += (cmsg_len + 7) & !7;
}
}
fn service(&mut self, timed_out: bool) -> io::Result<Vec<Vec<u8>>> {
let mut out = Vec::new();
self.flush_delayed_feedback();
if let Some(peer) = self.peer {
let base = self.dec.feedback(timed_out);
let now = Instant::now();
if timed_out || self.last_feedback.elapsed() >= self.ack_interval {
let mut ack = base;
ack.nak_block = NAK_NONE;
ack.nak_mask = 0;
self.queue_feedback(peer, &ack);
self.last_feedback = now;
}
let catching_up = self.dec.next_needed() > self.dec.highest_seen();
let gaps = self.dec.missing_blocks(self.nak_batch, timed_out || catching_up);
for (block, mask) in gaps {
if mask == 0 {
continue;
}
let fresh = self
.nak_history
.get(&block)
.is_none_or(|t| now.duration_since(*t) >= NAK_COOLDOWN);
if fresh {
let mut nfb = base;
nfb.nak_block = block;
nfb.nak_mask = mask;
self.queue_feedback(peer, &nfb);
self.nak_history.insert(block, now);
}
}
let nd = self.dec.next_needed();
self.nak_history = self.nak_history.split_off(&nd);
}
let head = self.dec.next_needed();
if head != self.head_block {
self.head_block = head;
self.head_since = Instant::now();
} else if self.head_since.elapsed() > self.max_hold && self.dec.window_len() > 0 {
out.extend(self.dec.skip_head());
self.head_block = self.dec.next_needed();
self.head_since = Instant::now();
}
Ok(out)
}
fn control_bytes(&self, fb: &Feedback) -> Vec<u8> {
let mut cp = ControlPacket::new();
cp.session_announce = self.dec.session_epoch();
cp.ack = Some(AckFrame {
ack_through: fb.ack_through,
});
if fb.nak_block != NAK_NONE {
cp.nak = Some(NakFrame {
block: fb.nak_block,
mask: fb.nak_mask,
});
}
cp.loss = Some(LossFrame {
loss_x255: fb.loss_x255,
burstiness_x255: fb.burstiness_x255,
owd_trend_class: fb.owd_trend_class,
loss_class: fb.loss_class,
});
if self.last_ttl != 0 {
cp.path = Some(PathFrame {
ttl: self.last_ttl,
ecn: self.last_tos & 0b11,
hop_count: crate::path_sensor::hop_count_from_ttl(self.last_ttl),
ce_count: self.ce_count,
ect_count: self.ect_count,
});
}
if self.local_pmtu != 0 {
cp.pmtu = Some(PmtuFrame { pmtu: self.local_pmtu });
}
cp.loss_acct = Some(LossAcctFrame {
seq: self.ctrl_out,
last_recv_seq: self.ctrl_recv,
});
if self.wbest_capacity_kbps != 0 {
cp.avail_bw = Some(crate::control_frame::AvailBwFrame {
avail_kbps: self.wbest_avail_kbps,
capacity_kbps: self.wbest_capacity_kbps,
});
}
let fc_kbps = (self.forecast.forecast_bps() * 8.0 / 1000.0) as u64;
if fc_kbps != 0 {
cp.forecast = Some(crate::control_frame::ForecastFrame {
forecast_kbps: fc_kbps,
});
}
if let Some((period_s, conf)) = self.periodicity.detected_period() {
let to_spike = self.periodicity.secs_to_next_spike().unwrap_or(0.0);
cp.periodicity = Some(crate::control_frame::PeriodicityFrame {
period_ds: (period_s * 10.0).round() as u64,
secs_to_spike_ds: (to_spike * 10.0).round() as u64,
confidence_x255: (conf.clamp(0.0, 1.0) * 255.0) as u8,
});
}
encode_control(&cp)
}
fn queue_feedback(&mut self, peer: SocketAddr, fb: &Feedback) {
self.maybe_observe_forecast();
self.ctrl_out = self.ctrl_out.wrapping_add(1);
let fbuf = self.control_bytes(fb);
if self.fb_drop_pct > 0 {
self.fb_drop_rng = self
.fb_drop_rng
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
if ((self.fb_drop_rng >> 33) as u32) % 100 < self.fb_drop_pct {
return;
}
}
if self.fb_delay.is_zero() {
self.send_feedback_bytes(&fbuf, peer);
} else {
self.fb_pending
.push_back((Instant::now() + self.fb_delay, fbuf));
}
}
fn send_feedback_bytes(&self, bytes: &[u8], peer: SocketAddr) {
if self.connected {
if self.sock.send(bytes).is_ok() {
return;
}
}
self.sock.send_to(bytes, peer).ok();
}
fn flush_delayed_feedback(&mut self) {
if self.fb_pending.is_empty() {
return;
}
let now = Instant::now();
let Some(peer) = self.peer else { return };
while let Some((due, _)) = self.fb_pending.front() {
if *due > now {
break;
}
let (_, bytes) = self.fb_pending.pop_front().unwrap();
self.send_feedback_bytes(&bytes, peer);
}
}
pub fn nudge_feedback(&mut self) -> io::Result<()> {
if let Some(peer) = self.peer {
self.ctrl_out = self.ctrl_out.wrapping_add(1);
let fb = self.dec.feedback(true);
let fbuf = self.control_bytes(&fb);
self.send_feedback_bytes(&fbuf, peer);
}
Ok(())
}
}
impl ReliableUdpReceiver {
pub fn bind(local: impl ToSocketAddrs) -> io::Result<Self> {
let sock = UdpSocket::bind(local)?;
sock.set_read_timeout(Some(Duration::from_millis(4)))?;
size_socket_buffers(&sock);
#[cfg(any(target_os = "linux", target_os = "freebsd"))]
enable_ttl_ecn(&sock);
#[cfg(target_os = "windows")]
enable_ttl_ecn_win(&sock);
let sock = crate::dgram::DgramSock::from_udp(sock);
Ok(Self {
sock: std::sync::Arc::new(sock),
sessions: HashMap::new(),
order: Vec::new(),
pending_admissions: HashMap::new(),
session_ceiling: None,
session_refusals: 0,
session_nonce_seq: 0,
session_changed: false,
session_admissions: 0,
session_admission_failures: 0,
start: Instant::now(),
net_events: NetEventObserver::start(None),
multi_peer: false,
net_event_shift_peak: 0.0,
cfg: RsSessionConfig {
max_hold: Duration::from_secs(60),
fb_delay: Duration::ZERO,
nak_batch: MAX_NAKS_PER_CYCLE,
debug_drop_pct: 0,
drop_rng: 0x9E3779B97F4A7C15,
ge_loss_p: 0,
ge_loss_r: 0,
drop_block_mod: 0,
burst_at: 0,
burst_len: 0,
fb_drop_pct: 0,
fb_drop_rng: 0x243F6A8885A308D3,
},
})
}
fn open_session(&mut self, epoch: u32) -> Option<&mut RsSession> {
if !self.sessions.contains_key(&epoch) {
if self.session_ceiling.is_some_and(|max| self.sessions.len() >= max) {
self.session_refusals += 1;
return None;
}
if !self.sessions.is_empty() {
self.dissolve_peer_association();
}
let s = RsSession::new(std::sync::Arc::clone(&self.sock), self.cfg);
self.sessions.insert(epoch, s);
self.order.push(epoch);
}
self.sessions.get_mut(&epoch)
}
fn dissolve_peer_association(&mut self) {
if self.sock.connect(UNSPECIFIED_PEER).is_ok() {
for s in self.sessions.values_mut() {
s.connected = false;
#[cfg(target_os = "linux")]
{
s.gro_on = false;
}
}
}
}
fn solo_connected(&self) -> bool {
match self.order.first().and_then(|e| self.sessions.get(e)) {
Some(s) => s.connected || s.peer.is_none(),
None => true,
}
}
fn release_silent_peer(&mut self) {
if self.multi_peer || self.sessions.len() != 1 {
return;
}
let stale = self
.order
.first()
.and_then(|e| self.sessions.get(e))
.is_some_and(|s| s.connected && s.last_data_at.elapsed() > PEER_SILENCE_TIMEOUT);
if stale {
self.dissolve_peer_association();
}
}
pub fn poll_from(&mut self) -> io::Result<Vec<(u32, Vec<u8>)>> {
let mut tagged: Vec<(u32, Vec<u8>)> = Vec::new();
let shift = self.net_events.path_shift();
if shift > self.net_event_shift_peak {
self.net_event_shift_peak = shift;
}
let pmtu = self.net_events.pmtu().unwrap_or(0);
self.release_silent_peer();
let timed_out = if !self.multi_peer && self.sessions.len() == 1 && self.solo_connected() {
let mut items = Vec::new();
let epoch = self.order[0];
let t = {
let s = self.sessions.get_mut(&epoch).expect("len == 1");
s.local_pmtu = pmtu;
s.recv_into(&mut items)?
};
tagged.extend(items.into_iter().map(|i| (epoch, i)));
t
} else {
self.drain_unconnected(&mut tagged, pmtu)?
};
self.expire_stale_admissions();
self.send_admission_challenges()?;
let ids = self.order.clone();
for epoch in ids {
if let Some(mut s) = self.sessions.remove(&epoch) {
s.local_pmtu = pmtu;
let r = s.service(timed_out);
self.sessions.insert(epoch, s);
tagged.extend(r?.into_iter().map(|i| (epoch, i)));
}
}
Ok(tagged)
}
pub fn poll(&mut self) -> io::Result<Vec<Vec<u8>>> {
Ok(self.poll_from()?.into_iter().map(|(_, item)| item).collect())
}
fn drain_unconnected(
&mut self,
tagged: &mut Vec<(u32, Vec<u8>)>,
pmtu: u16,
) -> io::Result<bool> {
let mut buf = [0u8; RECV_BUF];
match self.sock.recv_from(&mut buf) {
Ok((n, src)) => {
self.route_datagram(&buf[..n], src, tagged, pmtu);
Ok(false)
}
Err(e)
if matches!(
e.kind(),
io::ErrorKind::WouldBlock
| io::ErrorKind::TimedOut
| io::ErrorKind::ConnectionReset
) =>
{
Ok(true)
}
Err(e) => Err(e),
}
}
fn route_datagram(
&mut self,
buf: &[u8],
src: SocketAddr,
tagged: &mut Vec<(u32, Vec<u8>)>,
pmtu: u16,
) {
if self.try_admit(buf, src) {
return;
}
let epoch = match datagram_epoch(buf) {
Some(e) => e,
None => {
let by_addr = self
.order
.iter()
.find(|e| self.sessions.get(e).and_then(|s| s.peer) == Some(src))
.copied();
match by_addr.or_else(|| self.order.last().copied()) {
Some(e) => e,
None => return,
}
}
};
if !self.sessions.contains_key(&epoch) && !self.sessions.is_empty() {
self.begin_admission(epoch, src);
return;
}
let mut items = Vec::new();
if let Some(s) = self.open_session(epoch) {
s.local_pmtu = pmtu;
if s.peer != Some(src) {
s.peer = Some(src);
}
s.process_datagram(buf, &mut items);
}
tagged.extend(items.into_iter().map(|i| (epoch, i)));
}
fn begin_admission(&mut self, epoch: u32, addr: SocketAddr) {
if self.session_ceiling.is_some_and(|max| self.pending_admissions.len() >= max) {
self.session_refusals += 1;
return;
}
if let Some((a, _, _)) = self.pending_admissions.get(&epoch)
&& *a == addr
{
return;
}
self.session_nonce_seq = self.session_nonce_seq.wrapping_add(1);
let entropy = self.start.elapsed().as_nanos() as u64;
let mut x = entropy ^ self.session_nonce_seq.rotate_left(32) ^ u64::from(epoch);
x = (x ^ (x >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
x = (x ^ (x >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
let nonce = (x ^ (x >> 31)) & crate::control_frame::NONCE_MASK;
self.pending_admissions.insert(epoch, (addr, nonce, Instant::now()));
}
fn send_admission_challenges(&mut self) -> io::Result<()> {
let pending: Vec<(u32, SocketAddr, u64)> = self
.pending_admissions
.iter()
.map(|(e, (a, n, _))| (*e, *a, *n))
.collect();
for (epoch, addr, nonce) in pending {
let mut cp = ControlPacket::new();
cp.session_challenge = Some(crate::control_frame::SessionFrame { epoch, nonce });
let wire = encode_control(&cp);
self.sock.send_to(&wire, addr)?;
}
Ok(())
}
fn expire_stale_admissions(&mut self) {
let before = self.pending_admissions.len();
self.pending_admissions
.retain(|_, (_, _, sent)| sent.elapsed() <= SESSION_CHALLENGE_TIMEOUT);
self.session_admission_failures += (before - self.pending_admissions.len()) as u64;
}
fn try_admit(&mut self, buf: &[u8], src: SocketAddr) -> bool {
if !is_control(buf) {
return false;
}
let Some(cp) = decode_control(buf) else { return false };
let Some(sr) = cp.session_response else { return false };
let Some((addr, nonce, _)) = self.pending_admissions.get(&sr.epoch).copied() else {
return false;
};
if addr != src || nonce != sr.nonce {
return false;
}
self.pending_admissions.remove(&sr.epoch);
if let Some(s) = self.open_session(sr.epoch) {
s.peer = Some(src);
s.dec.adopt_epoch(sr.epoch);
s.nak_history.clear();
s.last_data_at = Instant::now();
self.session_admissions += 1;
self.session_changed = true;
return true;
}
false
}
pub fn with_multi_peer(mut self) -> Self {
self.multi_peer = true;
self
}
pub fn set_sock(&mut self, sock: crate::dgram::DgramSock) {
if sock.backend() == crate::dgram::DgramBackend::Demux {
self.multi_peer = true;
}
let sock = std::sync::Arc::new(sock);
self.sock = std::sync::Arc::clone(&sock);
for s in self.sessions.values_mut() {
s.sock = std::sync::Arc::clone(&sock);
s.connected = false;
}
}
pub fn session_epoch(&self) -> Option<u32> {
self.order.last().copied()
}
pub fn nudge_feedback(&mut self) -> io::Result<()> {
let ids = self.order.clone();
for epoch in ids {
if let Some(mut s) = self.sessions.remove(&epoch) {
let r = s.nudge_feedback();
self.sessions.insert(epoch, s);
r?;
}
}
Ok(())
}
pub fn with_debug_loss(mut self, pct: u32, seed: u64) -> Self {
self.cfg.debug_drop_pct = pct.min(100);
self.cfg.drop_rng = seed | 1;
self
}
pub fn with_gilbert_loss(mut self, p_per_10k: u32, r_per_10k: u32, seed: u64) -> Self {
self.cfg.ge_loss_p = p_per_10k;
self.cfg.ge_loss_r = r_per_10k.max(1);
self.cfg.drop_rng = seed | 1;
self
}
pub fn with_max_hold(mut self, hold: Duration) -> Self {
self.cfg.max_hold = hold;
self
}
pub fn with_feedback_delay(mut self, delay: Duration) -> Self {
self.cfg.fb_delay = delay;
self
}
pub fn with_nak_batch(mut self, batch: usize) -> Self {
self.cfg.nak_batch = batch.max(1);
self
}
pub fn with_feedback_drop(mut self, pct: u32) -> Self {
self.cfg.fb_drop_pct = pct.min(100);
self
}
pub fn with_block_drop_mod(mut self, m: u32) -> Self {
self.cfg.drop_block_mod = m;
self
}
pub fn with_burst_loss(mut self, at: u64, len: u64) -> Self {
self.cfg.burst_at = at;
self.cfg.burst_len = len;
self
}
pub fn local_addr(&self) -> io::Result<SocketAddr> {
self.sock.local_addr()
}
pub fn net_event_count(&self) -> u64 {
self.net_events.event_count()
}
pub fn local_pmtu(&self) -> u16 {
self.net_events.pmtu().unwrap_or(0)
}
pub fn net_event_shift(&self) -> f32 {
self.net_events.path_shift()
}
pub fn net_event_shift_peak(&self) -> f32 {
self.net_event_shift_peak
}
pub fn inject_path_event(&self) {
self.net_events.inject_event();
}
pub fn inject_pmtu(&self, mtu: u16) {
self.net_events.inject_pmtu(mtu);
}
pub fn recv_count(&self) -> u64 {
self.sessions.values().map(|s| s.recv_count()).sum()
}
pub fn peer_pmtu(&self) -> u16 {
self.order.last().and_then(|e| self.sessions.get(e)).map(|s| s.peer_pmtu()).unwrap_or(0)
}
pub fn peer_pmtu_of(&self, epoch: u32) -> Option<u16> {
self.sessions.get(&epoch).map(|s| s.peer_pmtu())
}
pub fn live_sessions(&self) -> Vec<u32> {
self.order.clone()
}
pub fn with_session_ceiling(mut self, max: usize) -> Self {
self.session_ceiling = Some(max.max(1));
self
}
pub fn session_refusals(&self) -> u64 {
self.session_refusals
}
fn newest(&self) -> Option<&RsSession> {
self.order.last().and_then(|e| self.sessions.get(e))
}
pub fn peak_loss_x255(&self) -> u8 {
self.sessions.values().map(|s| s.peak_loss_x255()).max().unwrap_or(0)
}
pub fn false_recovery_count(&self) -> u64 {
self.sessions.values().map(|s| s.false_recovery_count()).sum()
}
pub fn set_ge_burst(&mut self, on: bool) {
for s in self.sessions.values_mut() {
s.set_ge_burst(on);
}
}
pub fn set_debug_loss(&mut self, pct: u32) {
self.cfg.debug_drop_pct = pct.min(100);
for s in self.sessions.values_mut() {
s.set_debug_loss(pct);
}
}
pub fn mean_burst_len(&self) -> f32 {
self.newest().map(|s| s.mean_burst_len()).unwrap_or(0.0)
}
pub fn owd_skew(&self) -> f64 {
self.newest().map(|s| s.owd_skew()).unwrap_or(0.0)
}
pub fn owd_trend_debiased(&self) -> f64 {
self.newest().map(|s| s.owd_trend_debiased()).unwrap_or(0.0)
}
pub fn ack_interval(&self) -> Duration {
self.newest().map(|s| s.ack_interval()).unwrap_or(ACK_INTERVAL)
}
pub fn feedback_loss_est(&self) -> f32 {
self.newest().map(|s| s.feedback_loss_est()).unwrap_or(0.0)
}
pub fn peer_link(&self) -> (u8, u8) {
self.newest().map(|s| s.peer_link()).unwrap_or((0, 0))
}
pub fn accecn_counts(&self) -> (u64, u64) {
self.newest().map(|s| s.accecn_counts()).unwrap_or((0, 0))
}
pub fn forecast_bps(&self) -> u64 {
self.newest().map(|s| s.forecast_bps()).unwrap_or(0)
}
pub fn leo_cadence(&self) -> Option<(f64, f64, f64)> {
self.newest().and_then(|s| s.leo_cadence())
}
pub fn wbest_bps(&self) -> (u64, u64) {
self.newest().map(|s| s.wbest_bps()).unwrap_or((0, 0))
}
pub fn head_status(&self) -> Option<(u32, u32, usize, bool)> {
self.newest().and_then(|s| s.head_status())
}
pub fn take_session_changed(&mut self) -> bool {
std::mem::replace(&mut self.session_changed, false)
}
pub fn session_adoption_counts(&self) -> (u64, u64) {
(self.session_admissions, self.session_admission_failures)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicBool, Ordering as AtomicOrdering};
use std::sync::mpsc;
use std::sync::Arc;
#[test]
fn bind_rejects_oversized_k() {
let peer: SocketAddr = "127.0.0.1:9".parse().unwrap();
assert!(
ReliableUdpSender::bind("127.0.0.1:0", peer, 33, 1, 64).is_err(),
"k=33 > MAX_SHARDS must be rejected"
);
assert!(
ReliableUdpSender::bind("127.0.0.1:0", peer, 0, 1, 64).is_err(),
"k=0 must be rejected"
);
assert!(
ReliableUdpSender::bind("127.0.0.1:0", peer, 16, 8, 64).is_ok(),
"k=16 r=8 (k+r=24) must be accepted"
);
}
fn loopback_round_trip(n: u64, k: usize, r: usize, loss_pct: u32, seed: u64) {
let (addr_tx, addr_rx) = mpsc::channel();
let (done_tx, done_rx) = mpsc::channel();
let rx = std::thread::spawn(move || {
let mut recv = ReliableUdpReceiver::bind("127.0.0.1:0")
.unwrap()
.with_debug_loss(loss_pct, seed);
addr_tx.send(recv.local_addr().unwrap()).unwrap();
let mut got: Vec<u64> = Vec::new();
let start = Instant::now();
while (got.len() as u64) < n {
if start.elapsed() > Duration::from_secs(20) {
break;
}
for item in recv.poll().unwrap() {
got.push(u64::from_le_bytes(item.try_into().unwrap()));
}
}
for _ in 0..10 {
recv.nudge_feedback().ok();
std::thread::sleep(Duration::from_millis(2));
}
done_tx.send(()).ok();
got
});
let recv_addr = addr_rx.recv().unwrap();
let tx = std::thread::spawn(move || {
let mut send =
ReliableUdpSender::bind("127.0.0.1:0", recv_addr, k, r, 8).unwrap();
for i in 0..n {
while send.flow_blocked() {
send.drain_until_acked(Duration::from_millis(50)).ok();
}
send.send_item(&i.to_le_bytes()).unwrap();
}
send.flush().unwrap();
send.drain_until_acked(Duration::from_secs(15)).unwrap();
done_rx.recv_timeout(Duration::from_secs(20)).ok();
});
let got = rx.join().unwrap();
tx.join().unwrap();
let expected: Vec<u64> = (0..n).collect();
assert_eq!(got, expected, "loopback exact in-order delivery");
}
#[test]
fn loopback_clean() {
loopback_round_trip(500, 8, 2, 0, 1);
}
#[test]
fn loopback_lossy_fec() {
loopback_round_trip(500, 8, 3, 12, 7);
}
#[test]
fn loopback_heavy_arq() {
loopback_round_trip(300, 8, 2, 30, 1234);
}
#[test]
fn two_concurrent_rs_senders_both_deliver() {
const PER: u64 = 200;
const SENDERS: u64 = 2;
let mut recv = ReliableUdpReceiver::bind("127.0.0.1:0").unwrap().with_multi_peer();
let addr = recv.local_addr().unwrap();
let gate = Arc::new(std::sync::Barrier::new(SENDERS as usize));
let done = Arc::new(AtomicBool::new(false));
let mut txs = Vec::new();
let mut epochs = Vec::new();
for s in 0..SENDERS {
let stop = Arc::clone(&done);
let gate = Arc::clone(&gate);
let (etx, erx) = std::sync::mpsc::channel();
txs.push(std::thread::spawn(move || {
let mut send = ReliableUdpSender::bind("127.0.0.1:0", addr, 4, 2, 8).unwrap();
etx.send(send.enc.epoch()).ok();
gate.wait();
for i in 0..PER {
send.send_item(&((s << 56) | i).to_le_bytes()).unwrap();
}
send.flush().unwrap();
while !stop.load(AtomicOrdering::Relaxed) {
send.drain_until_acked(Duration::from_millis(50)).ok();
}
}));
epochs.push(erx.recv_timeout(Duration::from_secs(5)).unwrap());
}
assert_ne!(epochs[0], epochs[1], "independent senders must draw distinct epochs");
let mut got: Vec<u64> = Vec::new();
let start = Instant::now();
while (got.len() as u64) < PER * SENDERS && start.elapsed() < Duration::from_secs(25) {
for item in recv.poll().unwrap() {
got.push(u64::from_le_bytes(item.try_into().unwrap()));
}
}
done.store(true, AtomicOrdering::Relaxed);
for t in txs {
t.join().ok();
}
let (adopted, unanswered) = recv.session_adoption_counts();
for s in 0..SENDERS {
let mine: Vec<u64> =
got.iter().filter(|v| (*v >> 56) == s).map(|v| v & 0x00FF_FFFF_FFFF_FFFF).collect();
assert_eq!(
mine,
(0..PER).collect::<Vec<_>>(),
"sender {s} (epoch {}) must deliver every item in order alongside the other sender",
epochs[s as usize],
);
}
assert!(
adopted <= 1,
"receiver thrashed between the two peers: {adopted} adoptions, {unanswered} \
unanswered, for {SENDERS} concurrent senders",
);
}
#[test]
fn restarted_sender_is_adopted_after_the_challenge() {
const N: u64 = 40;
let mut recv = ReliableUdpReceiver::bind("127.0.0.1:0").unwrap();
let addr = recv.local_addr().unwrap();
let mut first = ReliableUdpSender::bind("127.0.0.1:0", addr, 4, 2, 8).unwrap();
for i in 0..N {
first.send_item(&i.to_le_bytes()).unwrap();
}
first.flush().unwrap();
let mut seen = 0u64;
let start = Instant::now();
while seen < N && start.elapsed() < Duration::from_secs(10) {
seen += recv.poll().unwrap().len() as u64;
first.pump_feedback().ok();
}
assert_eq!(seen, N, "first session did not deliver");
let epoch_a = recv.session_epoch();
assert!(epoch_a.is_some(), "no session epoch learned");
drop(first);
let mut second = ReliableUdpSender::bind("127.0.0.1:0", addr, 4, 2, 8).unwrap();
assert_ne!(
second.enc.epoch(),
recv.session_epoch().unwrap(),
"the replacement drew the same epoch as its predecessor"
);
for i in 0..N {
second.send_item(&(1000 + i).to_le_bytes()).unwrap();
}
second.flush().unwrap();
let done = Arc::new(AtomicBool::new(false));
let stop = Arc::clone(&done);
let tx = std::thread::spawn(move || {
while !stop.load(AtomicOrdering::Relaxed) {
second.drain_until_acked(Duration::from_millis(50)).ok();
}
});
let mut got = Vec::new();
let start = Instant::now();
while (got.len() as u64) < N && start.elapsed() < Duration::from_secs(25) {
got.extend(recv.poll().unwrap());
}
done.store(true, AtomicOrdering::Relaxed);
tx.join().ok();
let (adopted, unanswered) = recv.session_adoption_counts();
assert_eq!(
got.len() as u64,
N,
"restarted sender delivered {}/{N} (adopted {adopted}, unanswered {unanswered})",
got.len(),
);
assert_eq!(adopted, 1, "expected exactly one adoption");
assert!(recv.take_session_changed(), "session_changed never raised");
}
#[test]
fn path_event_registers_and_pmtu_round_trips() {
let n = 400u64;
let (addr_tx, addr_rx) = mpsc::channel();
let (rres_tx, rres_rx) = mpsc::channel();
let rx = std::thread::spawn(move || {
let mut recv = ReliableUdpReceiver::bind("127.0.0.1:0").unwrap();
recv.inject_pmtu(1400);
recv.inject_path_event();
addr_tx.send(recv.local_addr().unwrap()).unwrap();
let mut got = 0u64;
let start = Instant::now();
while got < n {
if start.elapsed() > Duration::from_secs(20) {
break;
}
for _item in recv.poll().unwrap() {
got += 1;
}
}
for _ in 0..60 {
recv.nudge_feedback().ok();
std::thread::sleep(Duration::from_millis(2));
}
rres_tx
.send((recv.net_event_count(), recv.peer_pmtu(), recv.local_pmtu()))
.unwrap();
got
});
let recv_addr = addr_rx.recv().unwrap();
let (sres_tx, sres_rx) = mpsc::channel();
let tx = std::thread::spawn(move || {
let mut send = ReliableUdpSender::bind("127.0.0.1:0", recv_addr, 8, 2, 8).unwrap();
send.inject_pmtu(1280);
for i in 0..n {
while send.flow_blocked() {
send.drain_until_acked(Duration::from_millis(50)).ok();
}
send.send_item(&i.to_le_bytes()).unwrap();
}
send.flush().unwrap();
send.drain_until_acked(Duration::from_secs(15)).unwrap();
for _ in 0..60 {
send.pump_feedback().ok();
std::thread::sleep(Duration::from_millis(2));
}
sres_tx
.send((send.net_event_count(), send.peer_pmtu(), send.local_pmtu()))
.unwrap();
});
let got = rx.join().unwrap();
tx.join().unwrap();
assert_eq!(got, n, "all items delivered");
let (recv_events, recv_peer_pmtu, recv_local_pmtu) = rres_rx.recv().unwrap();
let (send_events, send_peer_pmtu, send_local_pmtu) = sres_rx.recv().unwrap();
assert!(recv_events >= 1, "receiver path event registered");
assert!(send_events >= 1, "sender path event registered");
assert_eq!(
send_peer_pmtu, recv_local_pmtu,
"sender learned the receiver's MTU via the feedback Pmtu frame"
);
assert_eq!(
recv_peer_pmtu, send_local_pmtu,
"receiver learned the sender's MTU via the heartbeat Pmtu frame"
);
}
}