diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml index 70fb9cc..0c933c9 100644 --- a/.github/workflows/lint.yml +++ b/.github/workflows/lint.yml @@ -62,5 +62,10 @@ jobs: with: rustflags: "" + # Large-window UDP tests request 16 MiB; avoid silently capping their + # socket buffers below the in-flight window. UDP has no receiver flow control. + - name: Configure socket buffer limits + run: sudo sysctl -w net.core.rmem_max=16777216 net.core.wmem_max=16777216 + - name: Run tests run: cargo test --workspace --all-features --locked diff --git a/Cargo.lock b/Cargo.lock index 5f967e9..82d789a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -664,6 +664,7 @@ dependencies = [ "flux-timing", "flux-utils", "httparse", + "io-uring", "libc", "mio", "serde", diff --git a/crates/flux-network/Cargo.toml b/crates/flux-network/Cargo.toml index d05f118..8eda13f 100644 --- a/crates/flux-network/Cargo.toml +++ b/crates/flux-network/Cargo.toml @@ -20,6 +20,9 @@ serde.workspace = true tracing.workspace = true wincode = { workspace = true, optional = true } +[target.'cfg(target_os = "linux")'.dependencies] +io-uring.workspace = true + [dev-dependencies] core_affinity.workspace = true criterion.workspace = true diff --git a/crates/flux-network/benches/udp_pipeline.rs b/crates/flux-network/benches/udp_pipeline.rs index 0638c9b..964f5ae 100644 --- a/crates/flux-network/benches/udp_pipeline.rs +++ b/crates/flux-network/benches/udp_pipeline.rs @@ -11,6 +11,11 @@ //! MiB message. Recovery should resend one datagram, not the message. //! - `bcast`: one sender broadcasting to 8 receivers on one listener socket. //! +//! `FLUX_BENCH_TRANSPORT=udp,uring` selects backends; `FLUX_BENCH_REVERSE=1` +//! reverses their order. `FLUX_BENCH_SCALE=16` increases burst/broadcast sample +//! counts. `FLUX_BENCH_SIZE=2m` filters paced/burst/broadcast sizes. CPU ns/B +//! sums sender and receiver thread CPU time; it excludes the loss relay. +//! //! Run with `cargo bench -p flux-network --bench udp_pipeline`. use std::{ @@ -53,8 +58,29 @@ fn udp_config() -> UdpConfig { UdpConfig { max_message_size: 4 * 1024 * 1024, ..UdpConfig::lan() } } -fn transports() -> [(&'static str, Transport); 2] { - [("tcp", Transport::default()), ("udp", Transport::Udp(udp_config()))] +fn transports() -> Vec<(&'static str, Transport)> { + let mut transports = vec![("tcp", Transport::default()), ("udp", Transport::Udp(udp_config()))]; + #[cfg(target_os = "linux")] + transports.push(( + "uring", + Transport::Udp(UdpConfig { + io: flux_network::udp::UdpIo::Uring(flux_network::udp::UringConfig::default()), + ..udp_config() + }), + )); + if let Ok(filter) = std::env::var("FLUX_BENCH_TRANSPORT") { + transports.retain(|(name, _)| filter.split(',').any(|selected| selected == *name)); + assert!(!transports.is_empty(), "unknown FLUX_BENCH_TRANSPORT"); + } + if std::env::var_os("FLUX_BENCH_REVERSE").is_some() { + transports.reverse(); + } + transports +} + +fn sizes() -> impl Iterator { + let filter = std::env::var("FLUX_BENCH_SIZE").ok(); + SIZES.into_iter().filter(move |(name, _)| filter.as_deref().is_none_or(|s| s == *name)) } fn connector(transport: Transport) -> NetworkDriver { @@ -65,13 +91,17 @@ fn connector(transport: Transport) -> NetworkDriver { fn burst_plan(size: usize) -> (usize, usize) { let count = (64 * 1024 * 1024 / size).clamp(16, 4096); let window = (UdpConfig::default().send_window / 2 / size.div_ceil(1171)).clamp(1, 256); - (count, window) + let scale: usize = std::env::var("FLUX_BENCH_SCALE") + .map_or(1, |s| s.parse().expect("invalid FLUX_BENCH_SCALE")); + assert!((1..=64).contains(&scale)); + (count * scale, window) } struct Stats { latencies_ns: Vec, elapsed: Duration, bytes: usize, + cpu: Duration, } impl Stats { @@ -81,11 +111,12 @@ impl Stats { let pct = |p: f64| l[((l.len() - 1) as f64 * p) as usize] as f64 / 1000.0; let mibps = self.bytes as f64 / self.elapsed.as_secs_f64() / (1024.0 * 1024.0); println!( - "{name:<28} n={:<6} p50={:>9.1}µs p99={:>9.1}µs max={:>9.1}µs {mibps:>9.0} MiB/s {extra}", + "{name:<28} n={:<6} p50={:>9.1}µs p99={:>9.1}µs max={:>9.1}µs {mibps:>9.0} MiB/s cpu={:.3}ns/B {extra}", l.len(), pct(0.5), pct(0.99), pct(1.0), + self.cpu.as_nanos() as f64 / self.bytes as f64, ); } } @@ -136,6 +167,17 @@ struct Scenario { window: usize, } +fn thread_cpu_time() -> Duration { + let mut time: libc::timespec = unsafe { std::mem::zeroed() }; + assert_eq!( + unsafe { + libc::clock_gettime(libc::CLOCK_THREAD_CPUTIME_ID, std::ptr::from_mut(&mut time)) + }, + 0 + ); + Duration::new(time.tv_sec as u64, time.tv_nsec as u32) +} + /// Sends `count` messages from the server while the receiver thread counts /// them and records one-way latency. fn run(sc: Scenario) -> Stats { @@ -153,6 +195,7 @@ fn run(sc: Scenario) -> Stats { thread::spawn(move || { pin(0); let mut lat = Vec::with_capacity(expected); + let cpu_start = thread_cpu_time(); while lat.len() < expected && !stop.load(Ordering::Relaxed) { for r in &mut receivers { r.poll_with(|e| { @@ -164,12 +207,13 @@ fn run(sc: Scenario) -> Stats { got.store(lat.len(), Ordering::Relaxed); } let _ = done_tx.send(()); - lat + (lat, thread_cpu_time() - cpu_start) }) }; pin(1); let start = Instant::now(); + let cpu_start = thread_cpu_time(); let mut sent = 0; let mut next_send = start; while done_rx.try_recv().is_err() { @@ -188,9 +232,10 @@ fn run(sc: Scenario) -> Stats { } } let elapsed = start.elapsed(); - let latencies_ns = rx_thread.join().unwrap(); - assert_eq!(latencies_ns.len(), expected, "receiver timed out"); - Stats { latencies_ns, elapsed, bytes: expected * size } + let sender_cpu = thread_cpu_time() - cpu_start; + let (latencies_ns, receiver_cpu) = rx_thread.join().unwrap(); + assert_eq!(latencies_ns.len(), expected, "receiver timed out after {sent} sends"); + Stats { latencies_ns, elapsed, bytes: expected * size, cpu: sender_cpu + receiver_cpu } } /// Raises the relay's socket buffers so it never drops a burst itself. @@ -276,7 +321,7 @@ impl Drop for Relay { fn main() { pin(usize::MAX); println!("== paced: one message per 100µs, one receiver =="); - for (size_name, size) in SIZES { + for (size_name, size) in sizes() { for (name, transport) in transports() { let addr = free_addr(); let s = run(Scenario { @@ -294,7 +339,7 @@ fn main() { } println!("\n== burst: bounded outstanding, one receiver =="); - for (size_name, size) in SIZES { + for (size_name, size) in sizes() { let (count, window) = burst_plan(size); for (name, transport) in transports() { let addr = free_addr(); @@ -313,12 +358,14 @@ fn main() { } println!("\n== loss1: one 2 MiB message, exactly one datagram dropped by a relay =="); + for (name, transport) in + transports().into_iter().filter(|(_, t)| matches!(t, Transport::Udp(_))) { let server_addr = free_addr(); // Drop the 900th of 1789 data datagrams. let relay = Relay::start(server_addr, 900); let s = run(Scenario { - transport: Transport::Udp(udp_config()), + transport, listen: server_addr, dial: relay.addr, clients: 1, @@ -330,11 +377,11 @@ fn main() { let data = relay.data.load(Ordering::Relaxed); let retx = relay.retransmits.load(Ordering::Relaxed); drop(relay); - s.row("loss1/udp/2m", &format!("datagrams={data} retransmitted={retx}")); + s.row(&format!("loss1/{name}/2m"), &format!("datagrams={data} retransmitted={retx}")); } println!("\n== bcast: one sender, {BCAST_PEERS} receivers on one listener =="); - for (size_name, size) in SIZES { + for (size_name, size) in sizes() { let (count, window) = burst_plan(size); let count = count / BCAST_PEERS; for (name, transport) in transports() { diff --git a/crates/flux-network/src/network_driver.rs b/crates/flux-network/src/network_driver.rs index 153d5bf..49f7a59 100644 --- a/crates/flux-network/src/network_driver.rs +++ b/crates/flux-network/src/network_driver.rs @@ -96,9 +96,12 @@ impl Inner { } } -/// Poll-driven message transport built on `mio`, over TCP or reliable UDP +/// Poll-driven message transport over TCP or reliable UDP /// (see [`Transport`]). The API and events are identical for both. /// +/// TCP uses `mio`; UDP selects syscall or Linux `io_uring` I/O through +/// [`UdpConfig`]. +/// /// Manages: /// - **Outbound (client) connections** created via [`connect`]. These are /// **auto-retried** on failure/disconnect: TCP on its configured reconnect @@ -245,7 +248,7 @@ impl NetworkDriver { /// /// This call: /// 1) attempts outbound reconnects if due - /// 2) polls `mio` with a zero timeout + /// 2) polls readiness or drains `io_uring` completions without waiting /// 3) for each event calls `handler` with the appropriate [`PollEvent`] /// 4) returns whether any IO events were processed /// diff --git a/crates/flux-network/src/udp/connector.rs b/crates/flux-network/src/udp/connector.rs index 9054f0f..f9a8e68 100644 --- a/crates/flux-network/src/udp/connector.rs +++ b/crates/flux-network/src/udp/connector.rs @@ -17,13 +17,13 @@ use flux::spine::{SpineProducerWithDCache, SpineProducers}; use flux_communication::Timer; use flux_timing::{Duration, Instant, Nanos}; use flux_utils::{DCache, DCachePtr, safe_panic}; -use mio::{Events, Interest, Poll, Registry, Token, event::Event, net::UdpSocket}; +use mio::{Events, Interest, Poll, Registry, Token}; use tracing::{debug, info, warn}; use super::{ UdpConfig, peer::{MsgStore, PushOutcome, RxPayload, SendOutcome, Staged, UdpPeer, send_batch}, - sys::{BATCH, RecvBatch, SendBatch}, + sys::{BATCH, RecvBatch, SendBatch, UdpSocket}, wire::{HEADER_SIZE, Header, Kind}, }; use crate::{ @@ -134,7 +134,11 @@ impl UdpManager { store: MsgStore::new(), batch: SendBatch::new(), staged: [Staged { peer: 0, seq: 0 }; BATCH], - recv: Some(RecvBatch::new(udp.max_datagram_size)), + recv: match udp.io { + super::UdpIo::Syscall => Some(RecvBatch::new(udp.max_datagram_size)), + #[cfg(target_os = "linux")] + super::UdpIo::Uring(_) => None, + }, pending_disconnects: Vec::new(), next_token: 0, tick_interval: udp.min_rto / 2_u32, @@ -165,6 +169,13 @@ impl UdpManager { if let Some(size) = self.config.socket_buf_size { set_socket_buf_size(&socket, size); } + #[cfg(target_os = "linux")] + if let super::UdpIo::Uring(config) = self.udp.io { + socket.ring = Some(std::cell::RefCell::new(super::sys::uring::Ring::new( + socket.as_raw_fd(), + config, + )?)); + } let token = self.next_token(); self.registry.register(&mut socket, token, Interest::READABLE)?; self.sockets.push(Endpoint { token, socket, listener, writable_armed: false }); @@ -333,7 +344,7 @@ impl UdpManager { /// peers. Arms WRITABLE if the kernel stopped accepting. fn flush_socket(&mut self, k: usize, now: Instant) { let entry = &mut self.sockets[k]; - let fd = entry.socket.as_raw_fd(); + let socket = &entry.socket; let mut n = 0; for i in 0..self.peers.len() { if self.peers[i].socket_token != entry.token { @@ -344,7 +355,14 @@ impl UdpManager { self.staged[n] = Staged { peer: i, seq }; n += 1; if n == BATCH { - if !Self::dispatch(&mut self.batch, &self.staged, &mut self.peers, fd, n, now) { + if !Self::dispatch( + &mut self.batch, + &self.staged, + &mut self.peers, + socket, + n, + now, + ) { arm_writable(&self.registry, entry); return; } @@ -352,7 +370,8 @@ impl UdpManager { } } } - if n != 0 && !Self::dispatch(&mut self.batch, &self.staged, &mut self.peers, fd, n, now) { + if n != 0 && !Self::dispatch(&mut self.batch, &self.staged, &mut self.peers, socket, n, now) + { arm_writable(&self.registry, entry); } } @@ -363,11 +382,11 @@ impl UdpManager { batch: &mut SendBatch, staged: &[Staged; BATCH], peers: &mut [UdpPeer], - fd: i32, + socket: &UdpSocket, n: usize, now: Instant, ) -> bool { - let accepted = send_batch(batch, fd, n); + let accepted = send_batch(batch, socket, n); for s in &staged[..accepted] { peers[s.peer].mark_sent(s.seq, now); } @@ -516,12 +535,18 @@ impl UdpManager { } /// Readable/writable event on the socket at `k`. - fn handle_event(&mut self, k: usize, event: &Event, dcache: Option<&DCache>, deliver: &mut F) - where + fn handle_event( + &mut self, + k: usize, + readable: bool, + writable: bool, + dcache: Option<&DCache>, + deliver: &mut F, + ) where F: for<'a> FnMut(PollEvent>), { let now = Instant::now(); - if event.is_readable() { + if readable { let fd = self.sockets[k].socket.as_raw_fd(); let mut recv = self.recv.take().expect("recv batch in use"); loop { @@ -542,22 +567,18 @@ impl UdpManager { self.on_datagram(k, &dgram, dcache, deliver); } } + self.flush_acks(k, Instant::now()); } self.recv = Some(recv); } - if event.is_writable() { + if writable { self.sockets[k].writable_armed = false; self.flush_socket(k, now); } + self.flush_acks(k, now); let entry = &mut self.sockets[k]; - let token = entry.token; - for peer in self.peers.iter_mut().filter(|p| p.socket_token == token) { - if peer.take_ack_due() && peer.send_ack(&entry.socket, now) == SendOutcome::WouldBlock { - arm_writable(&self.registry, entry); - } - } - if event.is_writable() && !entry.writable_armed { + if writable && !entry.writable_armed { if let Err(err) = self.registry.reregister(&mut entry.socket, entry.token, Interest::READABLE) { @@ -566,6 +587,16 @@ impl UdpManager { } } + fn flush_acks(&mut self, k: usize, now: Instant) { + let entry = &mut self.sockets[k]; + let token = entry.token; + for peer in self.peers.iter_mut().filter(|p| p.socket_token == token) { + if peer.take_ack_due() && peer.send_ack(&entry.socket, now) == SendOutcome::WouldBlock { + arm_writable(&self.registry, entry); + } + } + } + #[inline] fn drain_pending_disconnects(&mut self, deliver: &mut F) -> bool where @@ -589,6 +620,42 @@ impl UdpManager { self.next_tick = now + self.tick_interval; self.tick(now); } + #[cfg(target_os = "linux")] + if let super::UdpIo::Uring(config) = self.udp.io { + for k in 0..self.sockets.len() { + // Drain in bounded passes, emitting ACKs between passes so + // slow callbacks cannot hold back the sender's whole window. + for _ in 0..BATCH { + let work = self.sockets[k].socket.ring().poll(); + o |= work; + let mut received_any = false; + // ACK processing can reap more receives while sending. + for _ in 0..config.recv_entries { + let Some(received) = self.sockets[k].socket.ring().receive() else { break }; + received_any = true; + if let Some((datagrams, from)) = + received.datagrams(self.udp.max_datagram_size) + { + for bytes in datagrams { + let Some(header) = Header::decode(bytes) else { continue }; + let dgram = + Datagram { header, payload: &bytes[HEADER_SIZE..], from, now }; + self.on_datagram(k, &dgram, dcache, deliver); + } + } + self.sockets[k].socket.ring().recycle(received); + } + if !work && !received_any { + break; + } + self.handle_event(k, false, false, dcache, deliver); + } + let writable = self.sockets[k].writable_armed; + self.handle_event(k, false, writable, dcache, deliver); + self.sockets[k].socket.ring().submit(); + } + return o | self.drain_pending_disconnects(deliver); + } // Taken out so `handle_event` can borrow `self`; put back below. let mut events = std::mem::replace(&mut self.events, Events::with_capacity(0)); if let Err(e) = self.poll.poll(&mut events, Some(std::time::Duration::ZERO)) { @@ -602,7 +669,7 @@ impl UdpManager { debug!(token = ?event.token(), "ignoring stale udp readiness event"); continue; }; - self.handle_event(k, event, dcache, deliver); + self.handle_event(k, event.is_readable(), event.is_writable(), dcache, deliver); } self.events = events; o |= self.drain_pending_disconnects(deliver); diff --git a/crates/flux-network/src/udp/mod.rs b/crates/flux-network/src/udp/mod.rs index 7a675ac..3c28640 100644 --- a/crates/flux-network/src/udp/mod.rs +++ b/crates/flux-network/src/udp/mod.rs @@ -17,6 +17,38 @@ mod wire; pub(crate) use connector::UdpManager; +/// Socket I/O implementation. The wire protocol is identical for both backends. +#[derive(Clone, Copy, Debug, Default)] +pub enum UdpIo { + #[default] + Syscall, + /// Linux `io_uring` with bounded per-socket buffers. Requires synchronous + /// cancellation support (Linux 6.0+); socket creation fails if unavailable. + #[cfg(target_os = "linux")] + Uring(UringConfig), +} + +/// Per-socket operation limits. Buffers are allocated when the socket opens. +/// +/// Each entry reserves approximately 64 KiB; defaults use about 6 MiB/socket. +/// `recv_entries` must be a power of two, both counts must be nonzero, and +/// their sum must not exceed 4096. +#[cfg(target_os = "linux")] +#[derive(Clone, Copy, Debug)] +pub struct UringConfig { + /// Outstanding sends, including GSO groups and control packets. + pub send_entries: u16, + /// Buffers shared by the socket's multishot receive. + pub recv_entries: u16, +} + +#[cfg(target_os = "linux")] +impl Default for UringConfig { + fn default() -> Self { + Self { send_entries: 64, recv_entries: 32 } + } +} + /// Tuning for [`crate::Transport::Udp`]. /// /// The retransmit timeout is measured from acks (RFC 6298) and clamped to @@ -26,6 +58,9 @@ pub(crate) use connector::UdpManager; /// `max_datagram_size`. #[derive(Clone, Copy, Debug)] pub struct UdpConfig { + /// Local I/O backend; peers may use different backends. Queued `io_uring` + /// sends progress through `NetworkDriver::poll_with`, like other backlogs. + pub io: UdpIo, /// Datagram size including the 29-byte header. 1200 stays under the /// 1280-byte IPv6 minimum MTU. pub max_datagram_size: usize, @@ -52,6 +87,7 @@ pub struct UdpConfig { impl Default for UdpConfig { fn default() -> Self { Self { + io: UdpIo::Syscall, max_datagram_size: 1200, send_window: 16 * 1024, recv_window: 16 * 1024, @@ -93,6 +129,11 @@ impl UdpConfig { } pub(crate) fn validate(&self) { + #[cfg(target_os = "linux")] + if let UdpIo::Uring(config) = self.io { + assert!(config.send_entries > 0 && config.recv_entries.is_power_of_two()); + assert!(u32::from(config.send_entries) + u32::from(config.recv_entries) <= 4096); + } assert!( self.max_datagram_size > wire::HEADER_SIZE && self.max_datagram_size <= wire::MAX_DATAGRAM_SIZE, diff --git a/crates/flux-network/src/udp/peer.rs b/crates/flux-network/src/udp/peer.rs index f6fa452..abf7146 100644 --- a/crates/flux-network/src/udp/peer.rs +++ b/crates/flux-network/src/udp/peer.rs @@ -1,17 +1,17 @@ //! Per-peer reliability state. One [`UdpPeer`] per remote address; the //! connector owns the sockets and passes them in. -use std::{collections::VecDeque, io, net::SocketAddr, os::fd::AsRawFd}; +use std::{collections::VecDeque, io, net::SocketAddr}; use flux_communication::Timer; use flux_timing::{Duration, Instant, Nanos}; use flux_utils::{DCache, DCacheRef}; -use mio::{Token, net::UdpSocket}; +use mio::Token; use tracing::{debug, warn}; use super::{ UdpConfig, - sys::{BATCH, SendBatch, SockAddr}, + sys::{BATCH, SendBatch, SockAddr, UdpSocket}, wire::{HEADER_SIZE, Header, Kind, fragment_count, write_session}, }; @@ -43,11 +43,12 @@ pub(crate) enum PushOutcome { TooLarge, } -/// Sends a filled batch of `n` datagrams. Returns how many the kernel took; -/// errors other than `WouldBlock` count as sent so the RTO path retries them. +/// Sends a filled batch of `n` datagrams. Returns how many the socket backend +/// accepted; errors other than `WouldBlock` count as sent so the RTO path +/// retries them. #[inline] -pub(crate) fn send_batch(batch: &mut SendBatch, fd: i32, n: usize) -> usize { - match batch.send(fd) { +pub(crate) fn send_batch(batch: &mut SendBatch, socket: &UdpSocket, n: usize) -> usize { + match socket.send_batch(batch) { Ok(k) => k, Err(e) if e.kind() == io::ErrorKind::WouldBlock => 0, Err(e) => { @@ -357,7 +358,7 @@ impl TxWindow { range: std::ops::Range, mut pred: impl FnMut(&Self, u64) -> bool, store: &MsgStore, - fd: i32, + socket: &UdpSocket, to: &SockAddr, batch: &mut SendBatch, now: Instant, @@ -373,7 +374,7 @@ impl TxWindow { pushed[n] = seq; n += 1; if n == BATCH { - let accepted = send_batch(batch, fd, n); + let accepted = send_batch(batch, socket, n); for &seq in &pushed[..accepted] { self.mark_sent(seq, true, now); } @@ -385,7 +386,7 @@ impl TxWindow { } } if n != 0 { - let accepted = send_batch(batch, fd, n); + let accepted = send_batch(batch, socket, n); for &seq in &pushed[..accepted] { self.mark_sent(seq, true, now); } @@ -422,7 +423,7 @@ impl TxWindow { rto: Duration, reorder: Duration, store: &mut MsgStore, - fd: i32, + socket: &UdpSocket, to: &SockAddr, batch: &mut SendBatch, now: Instant, @@ -487,7 +488,7 @@ impl TxWindow { if let Some(end) = self.hole_scan_end() { let pred = |tx: &Self, seq: u64| tx.is_hole(seq, rto, reorder, now); - let (_, blocked) = self.send_where(self.base..end, pred, store, fd, to, batch, now); + let (_, blocked) = self.send_where(self.base..end, pred, store, socket, to, batch, now); acked.blocked = blocked; } acked @@ -529,7 +530,7 @@ impl TxWindow { max_rto: Duration, reorder: Duration, store: &MsgStore, - fd: i32, + socket: &UdpSocket, to: &SockAddr, batch: &mut SendBatch, now: Instant, @@ -556,7 +557,7 @@ impl TxWindow { } false }; - self.send_where(self.base..self.next_send, pred, store, fd, to, batch, now) + self.send_where(self.base..self.next_send, pred, store, socket, to, batch, now) } /// Restarts every retained message from its first fragment under @@ -875,9 +876,12 @@ impl UdpPeer { } .encode(&mut self.ctrl); let n = self.rx.write_bitmap(n_bits, &mut self.ctrl[HEADER_SIZE..]); - self.ack_due = false; - self.last_send = now; - send_datagram(socket, self.addr, &self.ctrl[..HEADER_SIZE + n]) + let outcome = send_datagram(socket, self.addr, &self.ctrl[..HEADER_SIZE + n]); + self.ack_due = outcome == SendOutcome::WouldBlock; + if outcome == SendOutcome::Done { + self.last_send = now; + } + outcome } /// Stages the message in store `slot`. @@ -1098,7 +1102,7 @@ impl UdpPeer { self.rto.current(), self.rto.reorder_window(), store, - socket.as_raw_fd(), + socket, &self.native_addr, batch, now, @@ -1133,7 +1137,7 @@ impl UdpPeer { self.config.max_rto, self.rto.reorder_window(), store, - socket.as_raw_fd(), + socket, &self.native_addr, batch, now, @@ -1155,6 +1159,43 @@ impl UdpPeer { mod tests { use super::*; + #[cfg(target_os = "linux")] + #[test] + fn blocked_ack_is_retried_after_send_capacity_returns() { + use std::os::fd::AsRawFd; + + use crate::udp::{UringConfig, sys::uring::Ring}; + + let remote = std::net::UdpSocket::bind("127.0.0.1:0").unwrap(); + remote.set_read_timeout(Some(std::time::Duration::from_secs(2))).unwrap(); + let addr = remote.local_addr().unwrap(); + let (mut socket, _) = sock(); + socket.ring = Some(std::cell::RefCell::new( + Ring::new(socket.as_raw_fd(), UringConfig { send_entries: 1, recv_entries: 1 }) + .unwrap(), + )); + let mut peer = UdpPeer::new(addr, Token(1), Token(0), 7, cfg(), None); + socket.send_to(b"occupy", addr).unwrap(); + assert_eq!(peer.send_ack(&socket, Instant::now()), SendOutcome::WouldBlock); + assert!(peer.take_ack_due(), "backpressure must retain the pending ACK"); + let deadline = std::time::Instant::now() + std::time::Duration::from_secs(2); + loop { + socket.ring.as_ref().unwrap().borrow_mut().poll(); + if peer.send_ack(&socket, Instant::now()) == SendOutcome::Done { + break; + } + assert!(std::time::Instant::now() < deadline); + } + assert!(!peer.take_ack_due()); + socket.ring.as_ref().unwrap().borrow_mut().submit(); + let mut bytes = [0; 1200]; + assert_eq!(remote.recv(&mut bytes).unwrap(), 6); + let n = remote.recv(&mut bytes).unwrap(); + let header = Header::decode(&bytes[..n]).unwrap(); + assert_eq!(header.kind, Kind::Ack); + assert_eq!(header.session, 7); + } + fn cfg() -> UdpConfig { UdpConfig { send_window: 64, @@ -1203,13 +1244,13 @@ mod tests { let (s, to) = sock(); let mut batch = SendBatch::new(); let now = Instant::now(); - let fd = s.as_raw_fd(); + let socket = &s; send_all(&mut tx, 0..8); assert_eq!(tx.next_send, 8); // Ack 0..3 cumulatively, 5 and 7 selectively: 3, 4, 6 are holes. let bits = 0b0000_1010_u64.to_le_bytes(); let zero = Duration::ZERO; - let acked = tx.on_ack(3, 8, &bits, zero, zero, &mut store, fd, &to, &mut batch, now); + let acked = tx.on_ack(3, 8, &bits, zero, zero, &mut store, socket, &to, &mut batch, now); assert!(!acked.blocked); assert!(acked.rtt.is_some()); assert_eq!(tx.base, 3); @@ -1217,7 +1258,7 @@ mod tests { assert_eq!(store.free.len(), 2); assert_eq!(tx.free(), 59, "slots of the unfinished message stay reserved"); assert!(tx.is_acked(5) && tx.is_acked(7) && !tx.is_acked(4)); - let acked = tx.on_ack(8, 0, &[], zero, zero, &mut store, fd, &to, &mut batch, now); + let acked = tx.on_ack(8, 0, &[], zero, zero, &mut store, socket, &to, &mut batch, now); assert!(acked.rtt.is_none(), "holes were retransmitted, no clean sample"); assert_eq!(tx.free(), 64); assert!(tx.messages.is_empty()); @@ -1231,7 +1272,7 @@ mod tests { let mut store = MsgStore::new(); let mut tx = TxWindow::new(config.send_window); let (s, to) = sock(); - let fd = s.as_raw_fd(); + let socket = &s; let mut batch = SendBatch::new(); let now = Instant::now(); let zero = Duration::ZERO; @@ -1239,7 +1280,7 @@ mod tests { assert!(push(&mut tx, &mut store, stride, &[1])); } send_all(&mut tx, 0..64); - tx.on_ack(64, 0, &[], zero, zero, &mut store, fd, &to, &mut batch, now); + tx.on_ack(64, 0, &[], zero, zero, &mut store, socket, &to, &mut batch, now); assert_eq!(tx.base, 64); for _ in 0..8 { assert!(push(&mut tx, &mut store, stride, &[2])); @@ -1247,7 +1288,7 @@ mod tests { send_all(&mut tx, 64..72); // Old ack: cumulative 1, selective bit for seq 2, which aliases seq 66. let bits = 0b1_u64.to_le_bytes(); - tx.on_ack(1, 8, &bits, zero, zero, &mut store, fd, &to, &mut batch, now); + tx.on_ack(1, 8, &bits, zero, zero, &mut store, socket, &to, &mut batch, now); assert_eq!(tx.base, 64); assert!(!tx.is_acked(66)); } @@ -1259,13 +1300,13 @@ mod tests { let mut store = MsgStore::new(); let mut tx = TxWindow::new(config.send_window); let (s, to) = sock(); - let fd = s.as_raw_fd(); + let socket = &s; let mut batch = SendBatch::new(); let now = Instant::now(); let zero = Duration::ZERO; assert!(push(&mut tx, &mut store, stride, &vec![0; stride * 4])); send_all(&mut tx, 0..4); - tx.on_ack(2, 0, &[], zero, zero, &mut store, fd, &to, &mut batch, now); + tx.on_ack(2, 0, &[], zero, zero, &mut store, socket, &to, &mut batch, now); assert_eq!(tx.base, 2); tx.rewind(0xabcd); assert_eq!((tx.base, tx.next_send), (0, 0), "restarts from the first fragment"); diff --git a/crates/flux-network/src/udp/sys.rs b/crates/flux-network/src/udp/sys.rs index 758db34..258ae17 100644 --- a/crates/flux-network/src/udp/sys.rs +++ b/crates/flux-network/src/udp/sys.rs @@ -1,4 +1,5 @@ -//! Batched datagram syscalls. Linux uses `sendmmsg`/`recvmmsg`; elsewhere the +//! Datagram I/O with an optional Linux completion backend. +//! Batched syscalls use `sendmmsg`/`recvmmsg` on Linux; elsewhere the //! same API loops over `sendmsg`/`recvmsg`. Both are non-blocking and send or //! receive whatever is available right now: there is no accumulation delay. @@ -9,6 +10,94 @@ use std::{ ptr, slice, }; +#[cfg(target_os = "linux")] +pub(crate) mod uring; + +/// Datagram socket and its optional completion backend. +pub(crate) struct UdpSocket { + #[cfg(target_os = "linux")] + // Retire kernel requests before closing the socket below. + pub(crate) ring: Option>, + socket: mio::net::UdpSocket, +} + +impl UdpSocket { + pub(crate) fn bind(addr: SocketAddr) -> io::Result { + Ok(Self { + socket: mio::net::UdpSocket::bind(addr)?, + #[cfg(target_os = "linux")] + ring: None, + }) + } + + #[cfg(test)] + pub(crate) fn local_addr(&self) -> io::Result { + self.socket.local_addr() + } + + /// Panics unless the socket was opened with [`crate::udp::UdpIo::Uring`]. + #[cfg(target_os = "linux")] + pub(crate) fn ring(&self) -> std::cell::RefMut<'_, uring::Ring> { + self.ring.as_ref().expect("io_uring socket").borrow_mut() + } + + pub(crate) fn send_to(&self, bytes: &[u8], addr: SocketAddr) -> io::Result { + #[cfg(target_os = "linux")] + if let Some(ring) = &self.ring { + return ring.borrow_mut().send(bytes, SockAddr::new(addr)); + } + self.socket.send_to(bytes, addr) + } + + pub(crate) fn send_batch(&self, batch: &mut SendBatch) -> io::Result { + #[cfg(target_os = "linux")] + if let Some(ring) = &self.ring { + return ring.borrow_mut().send_batch(batch); + } + batch.send(std::os::fd::AsRawFd::as_raw_fd(self)) + } +} + +impl std::os::fd::AsRawFd for UdpSocket { + fn as_raw_fd(&self) -> RawFd { + self.socket.as_raw_fd() + } +} + +impl mio::event::Source for UdpSocket { + fn register( + &mut self, + registry: &mio::Registry, + token: mio::Token, + interests: mio::Interest, + ) -> io::Result<()> { + #[cfg(target_os = "linux")] + if self.ring.is_some() { + return Ok(()); + } + self.socket.register(registry, token, interests) + } + fn reregister( + &mut self, + registry: &mio::Registry, + token: mio::Token, + interests: mio::Interest, + ) -> io::Result<()> { + #[cfg(target_os = "linux")] + if self.ring.is_some() { + return Ok(()); + } + self.socket.reregister(registry, token, interests) + } + fn deregister(&mut self, registry: &mio::Registry) -> io::Result<()> { + #[cfg(target_os = "linux")] + if self.ring.is_some() { + return Ok(()); + } + self.socket.deregister(registry) + } +} + /// Datagrams per syscall. pub(crate) const BATCH: usize = 32; @@ -368,36 +457,50 @@ impl RecvBatch { /// entries. pub(crate) fn datagrams(&self, i: usize) -> Option<(std::slice::Chunks<'_, u8>, SocketAddr)> { debug_assert!(i < self.len); - let hdr = &self.hdrs[i]; - if hdr.msg_hdr.msg_flags & (libc::MSG_TRUNC | libc::MSG_CTRUNC) != 0 { - return None; - } - let len = hdr.msg_len as usize; - #[cfg(not(target_os = "linux"))] - let segment_size = len; - #[cfg(target_os = "linux")] - let segment_size = if hdr.msg_hdr.msg_controllen != 0 { - let control = &self.controls[i]; - let control_len = - unsafe { libc::CMSG_LEN(mem::size_of::() as _) as usize }; - if hdr.msg_hdr.msg_controllen < control_len || - control.header.cmsg_len != control_len || - control.header.cmsg_level != libc::SOL_UDP || - control.header.cmsg_type != libc::UDP_GRO - { - return None; - } - usize::try_from(control.size).ok()? - } else { - len - }; - if segment_size == 0 || segment_size > self.datagram_size { + let start = i * self.stride; + decode_datagrams( + &self.hdrs[i], + &self.addrs[i], + #[cfg(target_os = "linux")] + &self.controls[i], + &self.bufs[start..start + self.stride], + self.datagram_size, + ) + } +} + +fn decode_datagrams<'a>( + hdr: &MMsgHdr, + addr: &libc::sockaddr_storage, + #[cfg(target_os = "linux")] control: &GroControl, + bytes: &'a [u8], + datagram_size: usize, +) -> Option<(std::slice::Chunks<'a, u8>, SocketAddr)> { + if hdr.msg_hdr.msg_flags & (libc::MSG_TRUNC | libc::MSG_CTRUNC) != 0 { + return None; + } + let len = hdr.msg_len as usize; + #[cfg(not(target_os = "linux"))] + let segment_size = len; + #[cfg(target_os = "linux")] + let segment_size = if hdr.msg_hdr.msg_controllen != 0 { + let control_len = unsafe { libc::CMSG_LEN(mem::size_of::() as _) as usize }; + if hdr.msg_hdr.msg_controllen < control_len || + control.header.cmsg_len != control_len || + control.header.cmsg_level != libc::SOL_UDP || + control.header.cmsg_type != libc::UDP_GRO + { return None; } - let from = SockAddr::decode(&self.addrs[i], hdr.msg_hdr.msg_namelen)?; - let start = i * self.stride; - Some((self.bufs[start..start + len].chunks(segment_size), from)) + usize::try_from(control.size).ok()? + } else { + len + }; + if segment_size == 0 || segment_size > datagram_size { + return None; } + let from = SockAddr::decode(addr, hdr.msg_hdr.msg_namelen)?; + Some((bytes[..len].chunks(segment_size), from)) } #[cfg(test)] diff --git a/crates/flux-network/src/udp/sys/uring.rs b/crates/flux-network/src/udp/sys/uring.rs new file mode 100644 index 0000000..c0c8134 --- /dev/null +++ b/crates/flux-network/src/udp/sys/uring.rs @@ -0,0 +1,683 @@ +//! Bounded completion I/O. Kernel requests own pooled buffers, never peer +//! state. + +use std::{collections::VecDeque, io, mem, os::fd::RawFd, ptr}; + +use io_uring::{IoUring, opcode, types}; +use tracing::{debug, warn}; + +use super::{ + GRO_CONTROL_SPACE, GroControl, MMsgHdr, SEGMENT_CONTROL_SPACE, SegmentControl, SendBatch, + SockAddr, decode_datagrams, iovec_mut, +}; +use crate::udp::UringConfig; + +const RX_BIT: u64 = 1 << 63; +const BUFFER_SIZE: usize = 65_535; + +const NAME_SIZE: usize = mem::size_of::(); +const RX_SIZE: usize = BUFFER_SIZE + 16 + NAME_SIZE + GRO_CONTROL_SPACE; + +struct Rx { + bytes: Box<[u8]>, + len: usize, +} + +impl Rx { + fn new() -> Self { + Self { bytes: vec![0; RX_SIZE].into_boxed_slice(), len: 0 } + } +} + +struct Provided { + base: std::ptr::NonNull, + layout: std::alloc::Layout, + tail: u16, + mask: u16, +} + +impl Provided { + fn new(entries: u16) -> Self { + let page = unsafe { libc::sysconf(libc::_SC_PAGESIZE) }; + assert!(page > 0); + let layout = std::alloc::Layout::from_size_align( + usize::from(entries) * mem::size_of::(), + page as usize, + ) + .unwrap(); + let base = std::ptr::NonNull::new(unsafe { std::alloc::alloc_zeroed(layout) }.cast()) + .unwrap_or_else(|| std::alloc::handle_alloc_error(layout)); + Self { base, layout, tail: 0, mask: entries - 1 } + } + + fn provide(&mut self, index: usize, bytes: &mut [u8]) { + // Only consumed entries are reused. Publish the buffer after its + // descriptor, and do not touch its bytes until the receive CQE. + unsafe { + let entry = &mut *self.base.as_ptr().add(usize::from(self.tail & self.mask)); + entry.set_addr(bytes.as_mut_ptr() as u64); + entry.set_len(bytes.len() as u32); + entry.set_bid(index as u16); + self.tail = self.tail.wrapping_add(1); + let tail = types::BufRingEntry::tail(self.base.as_ptr()) + .cast::(); + (*tail).store(self.tail, std::sync::atomic::Ordering::Release); + } + } +} + +impl Drop for Provided { + fn drop(&mut self) { + unsafe { + std::alloc::dealloc(self.base.as_ptr().cast(), self.layout); + } + } +} + +fn receive_header() -> libc::msghdr { + let mut header: libc::msghdr = unsafe { mem::zeroed() }; + header.msg_namelen = NAME_SIZE as _; + header.msg_controllen = GRO_CONTROL_SPACE; + header +} + +struct Tx { + header: libc::msghdr, + addr: SockAddr, + control: SegmentControl, + iov: libc::iovec, + bytes: Box<[u8]>, + len: usize, + segment: usize, + offset: usize, + fallback: bool, +} + +impl Tx { + fn new() -> Box { + Box::new(Self { + header: unsafe { mem::zeroed() }, + addr: SockAddr::new("0.0.0.0:0".parse().unwrap()), + control: unsafe { mem::zeroed() }, + iov: unsafe { mem::zeroed() }, + bytes: vec![0; BUFFER_SIZE].into_boxed_slice(), + len: 0, + segment: 0, + offset: 0, + fallback: false, + }) + } + + fn prepare(&mut self, fd: RawFd, index: usize) -> io_uring::squeue::Entry { + let end = if self.fallback { (self.offset + self.segment).min(self.len) } else { self.len }; + self.iov = iovec_mut(&mut self.bytes[self.offset..end]); + self.header.msg_name = ptr::from_mut(&mut self.addr.storage).cast(); + self.header.msg_namelen = self.addr.len; + self.header.msg_iov = ptr::from_mut(&mut self.iov); + self.header.msg_iovlen = 1; + self.header.msg_control = ptr::null_mut(); + self.header.msg_controllen = 0; + if self.segment != 0 && !self.fallback { + self.control.header = libc::cmsghdr { + cmsg_len: unsafe { libc::CMSG_LEN(mem::size_of::() as _) as usize }, + cmsg_level: libc::SOL_UDP, + cmsg_type: libc::UDP_SEGMENT, + }; + self.control.size = self.segment as u16; + self.header.msg_control = ptr::from_mut(&mut self.control).cast(); + self.header.msg_controllen = SEGMENT_CONTROL_SPACE; + } + opcode::SendMsg::new(types::Fd(fd), ptr::from_ref(&self.header)) + .flags(libc::MSG_DONTWAIT as _) + .build() + .user_data(index as u64) + } +} + +pub(crate) struct Received { + index: usize, + slot: Rx, +} + +impl Received { + pub(crate) fn datagrams( + &self, + max_size: usize, + ) -> Option<(std::slice::Chunks<'_, u8>, std::net::SocketAddr)> { + let rx = &self.slot; + let out = types::RecvMsgOut::parse(&rx.bytes[..rx.len], &receive_header()).ok()?; + if out.is_name_data_truncated() || + out.is_control_data_truncated() || + out.is_payload_truncated() + { + return None; + } + let mut addr: libc::sockaddr_storage = unsafe { mem::zeroed() }; + let mut control: GroControl = unsafe { mem::zeroed() }; + // The parser bounds these slices by the configured metadata capacities. + unsafe { + ptr::copy_nonoverlapping( + out.name_data().as_ptr(), + ptr::from_mut(&mut addr).cast(), + out.name_data().len(), + ); + ptr::copy_nonoverlapping( + out.control_data().as_ptr(), + ptr::from_mut(&mut control).cast(), + out.control_data().len(), + ); + } + let mut hdr: MMsgHdr = unsafe { mem::zeroed() }; + hdr.msg_len = out.payload_data().len() as u32; + hdr.msg_hdr.msg_namelen = out.incoming_name_len(); + hdr.msg_hdr.msg_controllen = out.control_data().len(); + // Return a slice of the owned buffer, independent of the parser's borrow. + let offset = out.payload_data().as_ptr() as usize - rx.bytes.as_ptr() as usize; + decode_datagrams( + &hdr, + &addr, + &control, + &rx.bytes[offset..offset + hdr.msg_len as usize], + max_size, + ) + } +} + +pub(crate) struct Ring { + io: IoUring, + fd: RawFd, + owner: std::thread::ThreadId, + rx: Vec>, + provided: Option, + // Boxed: the multishot SQE points at it and `Ring` itself moves. + receive_header: Box, + rx_active: bool, + available: usize, + // Vec indexing must not reborrow headers retained by other in-flight SQEs. + #[allow(clippy::vec_box)] + tx: Vec>, + free_tx: Vec, + ready: VecDeque, + completions: Vec<(u64, i32, u32)>, +} + +// SAFETY: requests point into stable heap allocations. Moving the owner does +// not move those allocations; RefCell on the socket excludes concurrent access. +#[allow(clippy::non_send_fields_in_send_ty)] +unsafe impl Send for Ring {} + +impl Ring { + pub(crate) fn new(fd: RawFd, config: UringConfig) -> io::Result { + let count = u32::from(config.send_entries) + u32::from(config.recv_entries); + let ring = IoUring::builder() + .setup_coop_taskrun() + .setup_taskrun_flag() + .build(count.next_power_of_two())?; + // Require the teardown primitive before posting any borrowed pointers. + cancel(&ring)?; + let provided = Provided::new(config.recv_entries); + // SAFETY: the aligned descriptor ring outlives its registration. + unsafe { + ring.submitter().register_buf_ring_with_flags( + provided.base.as_ptr() as u64, + config.recv_entries, + 0, + 0, + )?; + } + let mut this = Self { + io: ring, + fd, + owner: std::thread::current().id(), + rx: (0..config.recv_entries).map(|_| Some(Rx::new())).collect(), + provided: Some(provided), + receive_header: Box::new(receive_header()), + rx_active: false, + available: config.recv_entries.into(), + tx: (0..config.send_entries).map(|_| Tx::new()).collect(), + free_tx: (0..usize::from(config.send_entries)).rev().collect(), + ready: VecDeque::with_capacity(config.recv_entries.into()), + completions: Vec::with_capacity(count as usize + 1), + }; + for i in 0..this.rx.len() { + this.provided.as_mut().unwrap().provide(i, &mut this.rx[i].as_mut().unwrap().bytes); + } + this.arm_receive(); + this.io.submit()?; + Ok(this) + } + + fn push(&mut self, entry: &io_uring::squeue::Entry) { + // At most one SQE per slot. Ring capacity covers every TX and RX slot. + // SAFETY: slots remain allocated and unchanged until their CQE. + unsafe { + self.io.submission().push(entry).expect("UDP operation slots exceed SQ capacity"); + }; + } + + fn arm_receive(&mut self) { + if self.rx_active || self.available == 0 { + return; + } + let entry = + opcode::RecvMsgMulti::new(types::Fd(self.fd), ptr::from_ref(&*self.receive_header), 0) + .build() + .user_data(RX_BIT); + self.push(&entry); + self.rx_active = true; + } + + fn arm_send(&mut self, index: usize) { + let entry = self.tx[index].prepare(self.fd, index); + self.push(&entry); + } + + /// A successful enqueue owns the bytes, like acceptance into a socket + /// send buffer. Later errors are packet loss, recovered by the protocol. + pub(crate) fn send(&mut self, bytes: &[u8], to: SockAddr) -> io::Result { + if bytes.len() > BUFFER_SIZE { + return Err(io::Error::from_raw_os_error(libc::EMSGSIZE)); + } + let index = self.free_tx.pop().ok_or(io::ErrorKind::WouldBlock)?; + let tx = &mut self.tx[index]; + tx.bytes[..bytes.len()].copy_from_slice(bytes); + tx.len = bytes.len(); + tx.addr = to; + tx.segment = 0; + tx.offset = 0; + tx.fallback = false; + self.arm_send(index); + Ok(bytes.len()) + } + + pub(crate) fn send_batch(&mut self, batch: &mut SendBatch) -> io::Result { + if self.free_tx.is_empty() { + self.poll(); + } + let n = mem::take(&mut batch.len); + let mut accepted = 0; + while accepted < n { + let Some(index) = self.free_tx.pop() else { break }; + // A batch may straddle peers or message tails. Segment each + // compatible run instead of losing GSO for the entire batch. + let size = batch.iovs[accepted].iter().map(|iov| iov.iov_len).sum::(); + let mut end = accepted + 1; + if batch.gso && size != 0 { + while end < n && batch.addrs[end] == batch.addrs[accepted] { + let len = batch.iovs[end].iter().map(|iov| iov.iov_len).sum::(); + if len == 0 || len > size || (end - accepted) * size + len > BUFFER_SIZE { + break; + } + end += 1; + if len < size { + break; + } + } + } + let segment = if end > accepted + 1 { size } else { 0 }; + let tx = &mut self.tx[index]; + tx.len = 0; + tx.addr = batch.addrs[accepted]; + tx.segment = segment; + tx.offset = 0; + tx.fallback = false; + for i in accepted..end { + for iov in &batch.iovs[i] { + // SAFETY: SendBatch's slices remain live throughout this call. + let bytes = unsafe { + std::slice::from_raw_parts(iov.iov_base.cast::(), iov.iov_len) + }; + tx.bytes[tx.len..tx.len + bytes.len()].copy_from_slice(bytes); + tx.len += bytes.len(); + } + } + self.arm_send(index); + accepted = end; + } + if accepted == 0 && n != 0 { Err(io::ErrorKind::WouldBlock.into()) } else { Ok(accepted) } + } + + pub(crate) fn submit(&mut self) { + let owner = std::thread::current().id(); + if self.owner != owner { + // Receive task work belongs to the submitting thread. Retire those + // requests before a moved driver starts polling on another thread. + cancel(&self.io).expect("couldn't transfer UDP io_uring to this thread"); + self.owner = owner; + } + let needs_enter = { + let sq = self.io.submission(); + !sq.is_empty() || sq.taskrun() + }; + if needs_enter { + if let Err(err) = self.io.submit() { + debug!(?err, "UDP io_uring submit failed"); + } + } + } + + pub(crate) fn poll(&mut self) -> bool { + self.submit(); + self.completions.clear(); + self.completions + .extend(self.io.completion().map(|c| (c.user_data(), c.result(), c.flags()))); + for i in 0..self.completions.len() { + let (id, result, flags) = self.completions[i]; + if id & RX_BIT != 0 { + if !io_uring::cqueue::more(flags) { + self.rx_active = false; + } + if let Some(index) = io_uring::cqueue::buffer_select(flags) { + let index = usize::from(index); + self.available -= 1; + if result >= 0 { + self.rx[index].as_mut().unwrap().len = result as usize; + self.ready.push_back(index); + } else { + self.provided + .as_mut() + .unwrap() + .provide(index, &mut self.rx[index].as_mut().unwrap().bytes); + self.available += 1; + } + } + if result < 0 && result != -libc::ENOBUFS && result != -libc::ECANCELED { + debug!(result, "UDP io_uring receive failed"); + } + } else { + let index = id as usize; + let tx = &mut self.tx[index]; + if result < 0 && tx.segment != 0 && !tx.fallback { + // GSO errors reject the whole message. Retry its original + // datagrams without offload, retaining the same buffer. + tx.fallback = true; + self.arm_send(index); + } else if tx.fallback && result >= 0 && tx.offset + tx.iov.iov_len < tx.len { + tx.offset += tx.iov.iov_len; + self.arm_send(index); + } else { + if result < 0 { + debug!(result, "UDP io_uring send failed"); + } + self.free_tx.push(index); + } + } + } + self.arm_receive(); + !self.completions.is_empty() + } + + pub(crate) fn receive(&mut self) -> Option { + let index = self.ready.pop_front()?; + Some(Received { index, slot: self.rx[index].take().unwrap() }) + } + + pub(crate) fn recycle(&mut self, received: Received) { + let index = received.index; + self.rx[index] = Some(received.slot); + self.provided.as_mut().unwrap().provide(index, &mut self.rx[index].as_mut().unwrap().bytes); + self.available += 1; + self.arm_receive(); + } +} + +fn cancel(ring: &IoUring) -> io::Result<()> { + loop { + match ring.submitter().register_sync_cancel(None, types::CancelBuilder::any()) { + Ok(()) => return Ok(()), + Err(err) if err.raw_os_error() == Some(libc::ENOENT) => return Ok(()), + Err(err) if err.kind() == io::ErrorKind::Interrupted => {} + Err(err) => return Err(err), + } + } +} + +impl Drop for Ring { + fn drop(&mut self) { + // Unsubmitted SQEs are never submitted again. Cancel all kernel-owned + // operations synchronously before releasing their memory or socket. + match cancel(&self.io) { + Ok(()) => { + if let Err(err) = self.io.submitter().unregister_buf_ring(0) { + warn!(?err, "UDP buffer ring unregister failed; retaining receive buffers"); + mem::forget(mem::take(&mut self.rx)); + mem::forget(self.provided.take()); + } + } + Err(err) => { + // A failed cancellation cannot justify freeing kernel buffers. + warn!(?err, "UDP io_uring cancellation failed; retaining operation buffers"); + mem::forget(mem::take(&mut self.rx)); + mem::forget(mem::take(&mut self.tx)); + mem::forget(self.provided.take()); + // The multishot request also borrows this header. + mem::forget(mem::replace(&mut self.receive_header, Box::new(receive_header()))); + } + } + } +} + +#[cfg(test)] +mod tests { + use std::{ + net::UdpSocket, + os::fd::AsRawFd, + time::{Duration, Instant}, + }; + + use super::*; + + fn config() -> UringConfig { + UringConfig { send_entries: 2, recv_entries: 2 } + } + + fn receive(ring: &mut Ring) -> Received { + let deadline = Instant::now() + Duration::from_secs(2); + loop { + ring.poll(); + if let Some(rx) = ring.receive() { + return rx; + } + assert!(Instant::now() < deadline, "receive completion timed out"); + } + } + + #[test] + fn owned_sends_backpressure_and_receive_recycling() { + for bind in ["127.0.0.1:0", "[::1]:0"] { + let socket = UdpSocket::bind(bind).unwrap(); + let remote = UdpSocket::bind(bind).unwrap(); + remote.set_read_timeout(Some(Duration::from_secs(2))).unwrap(); + let mut ring = Ring::new(socket.as_raw_fd(), config()).unwrap(); + let to = SockAddr::new(remote.local_addr().unwrap()); + let mut payload = vec![1; 100]; + assert_eq!(ring.send(&payload, to).unwrap(), 100); + payload.fill(2); + assert_eq!(ring.send(&payload, to).unwrap(), 100); + payload.fill(3); + assert_eq!(ring.send(&payload, to).unwrap_err().kind(), io::ErrorKind::WouldBlock); + ring.poll(); + let mut buf = [0; 1200]; + let mut seen = Vec::new(); + for _ in 0..2 { + let n = remote.recv(&mut buf).unwrap(); + assert_eq!(n, 100); + assert!(buf[..n].iter().all(|b| *b == buf[0])); + seen.push(buf[0]); + } + seen.sort_unstable(); + assert_eq!(seen, [1, 2]); + for i in 0..16 { + remote.send_to(&[i; 10], socket.local_addr().unwrap()).unwrap(); + let rx = receive(&mut ring); + let (mut datagrams, from) = rx.datagrams(1200).unwrap(); + assert_eq!(from, remote.local_addr().unwrap()); + assert_eq!(datagrams.next().unwrap(), &[i; 10]); + assert!(datagrams.next().is_none()); + ring.recycle(rx); + } + // Drop with receives both submitted and queued for rearming. + drop(ring); + drop(socket); + } + } + + #[test] + fn gso_and_fallback_preserve_boundaries() { + for fallback in [false, true] { + let socket = UdpSocket::bind("127.0.0.1:0").unwrap(); + let remote = UdpSocket::bind("127.0.0.1:0").unwrap(); + remote.set_nonblocking(true).unwrap(); + if fallback { + let enabled: libc::c_int = 1; + assert_eq!( + unsafe { + libc::setsockopt( + socket.as_raw_fd(), + libc::SOL_SOCKET, + libc::SO_NO_CHECK, + ptr::from_ref(&enabled).cast(), + mem::size_of_val(&enabled) as _, + ) + }, + 0 + ); + } + let mut ring = Ring::new(socket.as_raw_fd(), config()).unwrap(); + let mut batch = SendBatch::new(); + batch.enable_gso(socket.as_raw_fd()).unwrap(); + let to = SockAddr::new(remote.local_addr().unwrap()); + let headers = [[1; 29], [2; 29], [3; 29], [4; 29]]; + let payload = [0x5a; 1171]; + for (i, header) in headers.iter().enumerate() { + batch.push(header, &payload[..if i == 3 { 17 } else { 1171 }], &to); + } + assert_eq!(ring.send_batch(&mut batch).unwrap(), 4); + let deadline = Instant::now() + Duration::from_secs(2); + let mut seen = [false; 4]; + let mut count = 0; + let mut buf = [0; 2048]; + while count < 4 { + ring.poll(); + if let Ok(n) = remote.recv(&mut buf) { + let i = usize::from(buf[0] - 1); + assert!(!seen[i]); + seen[i] = true; + assert_eq!(&buf[..29], &headers[i]); + assert_eq!(&buf[29..n], &payload[..if i == 3 { 17 } else { 1171 }]); + count += 1; + } + assert!(Instant::now() < deadline, "GSO/fallback timed out"); + } + } + } + + #[test] + fn exhausted_receive_pool_rearms_and_tail_wraps() { + let socket = UdpSocket::bind("127.0.0.1:0").unwrap(); + let remote = UdpSocket::bind("127.0.0.1:0").unwrap(); + let mut ring = Ring::new(socket.as_raw_fd(), config()).unwrap(); + for i in 0..8 { + remote.send_to(&[i], socket.local_addr().unwrap()).unwrap(); + } + let deadline = Instant::now() + Duration::from_secs(2); + while ring.rx_active || ring.available != 0 { + ring.poll(); + assert!(Instant::now() < deadline, "buffer exhaustion CQE missing"); + } + let mut seen = [false; 8]; + for _ in 0..8 { + let rx = receive(&mut ring); + let (mut packets, _) = rx.datagrams(1200).unwrap(); + let i = usize::from(packets.next().unwrap()[0]); + assert!(!seen[i]); + seen[i] = true; + ring.recycle(rx); + } + assert!(seen.into_iter().all(|v| v)); + for _ in 0..65_536 { + remote.send_to(b"wrap", socket.local_addr().unwrap()).unwrap(); + let rx = receive(&mut ring); + assert_eq!(rx.datagrams(1200).unwrap().0.next().unwrap(), b"wrap"); + ring.recycle(rx); + } + } + + #[test] + fn driver_moves_after_submitting_receives() { + let socket = UdpSocket::bind("127.0.0.1:0").unwrap(); + let remote = UdpSocket::bind("127.0.0.1:0").unwrap(); + let mut ring = Ring::new(socket.as_raw_fd(), config()).unwrap(); + // The original issuer exits while the socket and ring move onward. + let (mut ring, socket) = std::thread::spawn(move || { + ring.poll(); + (ring, socket) + }) + .join() + .unwrap(); + remote.send_to(b"moved", socket.local_addr().unwrap()).unwrap(); + let rx = receive(&mut ring); + assert_eq!(rx.datagrams(1200).unwrap().0.next().unwrap(), b"moved"); + ring.recycle(rx); + } + + #[test] + fn mixed_gso_runs_keep_their_destinations_and_lengths() { + let socket = UdpSocket::bind("127.0.0.1:0").unwrap(); + let peers = + [UdpSocket::bind("127.0.0.1:0").unwrap(), UdpSocket::bind("127.0.0.1:0").unwrap()]; + let to = peers.each_ref().map(|p| SockAddr::new(p.local_addr().unwrap())); + for p in &peers { + p.set_read_timeout(Some(Duration::from_secs(2))).unwrap(); + } + let mut ring = + Ring::new(socket.as_raw_fd(), UringConfig { send_entries: 8, recv_entries: 2 }) + .unwrap(); + let mut batch = SendBatch::new(); + batch.enable_gso(socket.as_raw_fd()).unwrap(); + let headers: [[u8; 29]; 8] = std::array::from_fn(|i| [i as u8; 29]); + let payload = [0x5a; 1171]; + let lengths = [1171, 17, 1171, 1171, 1171, 3, 1171, 1171]; + for i in 0..8 { + batch.push(&headers[i], &payload[..lengths[i]], &to[i / 4]); + } + assert_eq!(ring.send_batch(&mut batch).unwrap(), 8); + ring.poll(); + let mut seen = [false; 8]; + for (peer_index, peer) in peers.iter().enumerate() { + for _ in 0..4 { + let mut buf = [0; 2048]; + let len = peer.recv(&mut buf).unwrap(); + let i = usize::from(buf[0]); + assert_eq!(i / 4, peer_index); + assert!(!seen[i]); + seen[i] = true; + assert_eq!(&buf[..29], &headers[i]); + assert_eq!(&buf[29..len], &payload[..lengths[i]]); + } + } + } + + #[test] + fn gro_and_oversized_receives() { + let socket = UdpSocket::bind("127.0.0.1:0").unwrap(); + let remote = UdpSocket::bind("127.0.0.1:0").unwrap(); + super::super::RecvBatch::enable_gro(socket.as_raw_fd()).unwrap(); + let mut ring = Ring::new(socket.as_raw_fd(), config()).unwrap(); + let mut batch = SendBatch::new(); + batch.enable_gso(remote.as_raw_fd()).unwrap(); + let to = SockAddr::new(socket.local_addr().unwrap()); + for _ in 0..4 { + batch.push(&[1; 29], &[2; 1171], &to); + } + assert_eq!(batch.send(remote.as_raw_fd()).unwrap(), 4); + let rx = receive(&mut ring); + let (datagrams, _) = rx.datagrams(1200).unwrap(); + assert_eq!(datagrams.count(), 4); + ring.recycle(rx); + remote.send_to(&[0; 1201], socket.local_addr().unwrap()).unwrap(); + let rx = receive(&mut ring); + assert!(rx.datagrams(1200).is_none()); + ring.recycle(rx); + } +} diff --git a/crates/flux-network/tests/support/udp_connector.rs b/crates/flux-network/tests/support/udp_connector.rs new file mode 100644 index 0000000..b0fa649 --- /dev/null +++ b/crates/flux-network/tests/support/udp_connector.rs @@ -0,0 +1,668 @@ +use std::{ + net::{Ipv4Addr, SocketAddr, UdpSocket}, + sync::{ + Arc, + atomic::{AtomicBool, AtomicUsize, Ordering}, + }, + thread, + time::{Duration, Instant}, +}; + +use flux_network::{NetworkDriver, PollEvent, SendBehavior, Transport, UdpConfig}; +use mio::Token; + +fn free_addr() -> SocketAddr { + UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).unwrap().local_addr().unwrap() +} + +fn udp(config: UdpConfig) -> NetworkDriver { + let config = UdpConfig { io: IO, ..config }; + NetworkDriver::default().with_transport(Transport::Udp(config)) +} + +/// Handshake between a fresh listener and client, returning the accepted and +/// the outbound token. `dial` differs from `addr` when a relay sits between. +fn connect_via( + server: &mut NetworkDriver, + client: &mut NetworkDriver, + addr: SocketAddr, + dial: SocketAddr, +) -> (Token, Token) { + server.listen_at(addr).unwrap(); + let client_token = client.connect(dial).unwrap(); + let mut accepted = None; + let deadline = Instant::now() + Duration::from_secs(5); + while accepted.is_none() { + assert!(Instant::now() < deadline, "no accept"); + server.poll_with(|e| { + if let PollEvent::Accept { stream, .. } = e { + accepted = Some(stream); + } + }); + client.poll_with(|_| {}); + thread::sleep(Duration::from_micros(50)); + } + while client.currently_disconnected().count() != 0 { + assert!(Instant::now() < deadline, "handshake"); + server.poll_with(|_| {}); + client.poll_with(|_| {}); + thread::sleep(Duration::from_micros(50)); + } + // A hello retry may still be in flight if the ack took longer than one + // RTO. Let it land now, while the peer it belongs to still exists. + for _ in 0..5 { + server.poll_with(|_| {}); + client.poll_with(|_| {}); + } + (accepted.unwrap(), client_token) +} + +fn connect_pair( + server: &mut NetworkDriver, + client: &mut NetworkDriver, + addr: SocketAddr, +) -> (Token, Token) { + connect_via(server, client, addr, addr) +} + +fn checksum(bytes: &[u8]) -> u64 { + bytes + .iter() + .fold(0xcbf2_9ce4_8422_2325_u64, |h, b| (h ^ u64::from(*b)).wrapping_mul(0x100_0000_01b3)) +} + +/// Test message: 4-byte id, then `len` bytes derived from the id. +fn make_msg(id: u32, len: usize) -> Vec { + let mut v = Vec::with_capacity(4 + len); + v.extend_from_slice(&id.to_le_bytes()); + v.extend((0..len).map(|i| (id as usize).wrapping_mul(31).wrapping_add(i * 7) as u8)); + v +} + +fn msg_id(payload: &[u8]) -> u32 { + u32::from_le_bytes(payload[..4].try_into().unwrap()) +} + +#[test] +fn udp_roundtrip_before_handshake_completes() { + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()); + server.listen_at(addr).unwrap(); + let mut client = udp(UdpConfig::lan()); + let tok = client.connect(addr).unwrap(); + // Queued before the hello ack arrives; must go out once it does. + client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"ping")); + + let mut accepted = None; + let mut request_seen = false; + let mut reply_seen = false; + let deadline = Instant::now() + Duration::from_secs(5); + while !reply_seen { + assert!(Instant::now() < deadline, "roundtrip timed out"); + server.poll_with(|e| match e { + PollEvent::Accept { stream, .. } => accepted = Some(stream), + PollEvent::Message { token, payload, .. } => { + assert_eq!(Some(token), accepted); + assert_eq!(payload, b"ping"); + request_seen = true; + } + _ => {} + }); + if request_seen && !reply_seen { + server.write_or_enqueue_with(SendBehavior::Single(accepted.unwrap()), |b| { + b.extend_from_slice(b"pong"); + }); + request_seen = false; + } + client.poll_with(|e| { + if let PollEvent::Message { token, payload, .. } = e { + assert_eq!(token, tok); + assert_eq!(payload, b"pong"); + reply_seen = true; + } + }); + thread::sleep(Duration::from_micros(50)); + } +} + +#[test] +fn udp_broadcast_mixed_sizes_to_two_subscribers() { + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()); + server.listen_at(addr).unwrap(); + let mut a = udp(UdpConfig::lan()); + let mut b = udp(UdpConfig::lan()); + a.connect(addr).unwrap(); + b.connect(addr).unwrap(); + let mut accepted = 0; + let deadline = Instant::now() + Duration::from_secs(5); + while accepted < 2 { + assert!(Instant::now() < deadline, "accepts"); + server.poll_with(|e| { + if let PollEvent::Accept { .. } = e { + accepted += 1; + } + }); + a.poll_with(|_| {}); + b.poll_with(|_| {}); + } + + // 1 byte, one datagram, one stride exactly, and a 2 MiB message. + let sizes = [1usize, 100, 1171, 1172, 5000, 2 * 1024 * 1024]; + let msgs: Vec> = + sizes.iter().enumerate().map(|(i, s)| make_msg(i as u32, *s)).collect(); + for m in &msgs { + server.write_or_enqueue_with(SendBehavior::Broadcast, |buf| buf.extend_from_slice(m)); + } + + let expected: Vec = msgs.iter().map(|m| checksum(m)).collect(); + let mut got_a = Vec::new(); + let mut got_b = Vec::new(); + let deadline = Instant::now() + Duration::from_secs(10); + while got_a.len() < msgs.len() || got_b.len() < msgs.len() { + assert!(Instant::now() < deadline, "broadcast delivery"); + server.poll_with(|_| {}); + a.poll_with(|e| { + if let PollEvent::Message { payload, .. } = e { + got_a.push((msg_id(payload), checksum(payload))); + } + }); + b.poll_with(|e| { + if let PollEvent::Message { payload, .. } = e { + got_b.push((msg_id(payload), checksum(payload))); + } + }); + } + for got in [got_a, got_b] { + assert_eq!(got.len(), msgs.len()); + for (id, sum) in got { + assert_eq!(sum, expected[id as usize], "message {id} corrupted"); + } + } +} + +/// Forwards datagrams between one client and the server, dropping every +/// `drop_every`-th datagram in each direction. +struct LossyRelay { + addr: SocketAddr, + stop: Arc, + dropped: Arc, + handle: Option>, +} + +impl LossyRelay { + fn start(server: SocketAddr, drop_every: usize) -> Self { + let socket = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + socket.set_read_timeout(Some(Duration::from_millis(5))).unwrap(); + let addr = socket.local_addr().unwrap(); + let stop = Arc::new(AtomicBool::new(false)); + let dropped = Arc::new(AtomicUsize::new(0)); + let (stop_c, dropped_c) = (stop.clone(), dropped.clone()); + let handle = thread::spawn(move || { + let mut buf = vec![0u8; 65_536]; + let mut client: Option = None; + let mut count = 0usize; + while !stop_c.load(Ordering::Relaxed) { + let Ok((n, from)) = socket.recv_from(&mut buf) else { continue }; + let to = if from == server { + let Some(c) = client else { continue }; + c + } else { + client = Some(from); + server + }; + count += 1; + if count.is_multiple_of(drop_every) { + dropped_c.fetch_add(1, Ordering::Relaxed); + continue; + } + let _ = socket.send_to(&buf[..n], to); + } + }); + Self { addr, stop, dropped, handle: Some(handle) } + } +} + +impl Drop for LossyRelay { + fn drop(&mut self) { + self.stop.store(true, Ordering::Relaxed); + self.handle.take().unwrap().join().unwrap(); + } +} + +#[test] +fn udp_delivers_everything_exactly_once_under_loss() { + const N: u32 = 400; + let server_addr = free_addr(); + let relay = LossyRelay::start(server_addr, 7); + + let mut server = udp(UdpConfig::lan()); + let mut client = udp(UdpConfig::lan()); + let (accepted, _) = connect_via(&mut server, &mut client, server_addr, relay.addr); + + // Both directions at once: server pushes to the client, client replies. + let msgs: Vec> = (0..N).map(|i| make_msg(i, 1 + (i as usize * 613) % 4000)).collect(); + for m in &msgs { + server.write_or_enqueue_with(SendBehavior::Single(accepted), |b| b.extend_from_slice(m)); + } + + let mut seen = vec![false; N as usize]; + let mut received = 0; + let mut echoed_back = 0; + let deadline = Instant::now() + Duration::from_secs(20); + while received < N || echoed_back < N { + assert!(Instant::now() < deadline, "loss recovery: {received} rx, {echoed_back} echoed"); + let mut echo = Vec::new(); + client.poll_with(|e| { + if let PollEvent::Message { payload, .. } = e { + let id = msg_id(payload) as usize; + assert_eq!(checksum(payload), checksum(&msgs[id]), "message {id} corrupted"); + assert!(!seen[id], "message {id} delivered twice"); + seen[id] = true; + received += 1; + echo.push(id as u32); + } + }); + for id in echo { + client.write_or_enqueue_with(SendBehavior::Broadcast, |b| { + b.extend_from_slice(&id.to_le_bytes()); + }); + } + server.poll_with(|e| { + if let PollEvent::Message { payload, .. } = e { + assert_eq!(payload.len(), 4); + echoed_back += 1; + } + }); + thread::sleep(Duration::from_micros(20)); + } + assert!(relay.dropped.load(Ordering::Relaxed) > 0, "relay dropped nothing"); +} + +#[test] +fn udp_client_disconnect_is_a_new_session() { + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()); + let mut client = udp(UdpConfig::lan()); + let (first, tok) = connect_pair(&mut server, &mut client, addr); + + client.disconnect(tok); + assert_eq!(client.currently_disconnected().count(), 1); + client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"after")); + + let mut events = Vec::new(); + let mut payload_on = None; + let mut reconnected = false; + let deadline = Instant::now() + Duration::from_secs(5); + while payload_on.is_none() || !reconnected { + assert!(Instant::now() < deadline, "reconnect"); + server.poll_with(|e| match e { + PollEvent::Disconnect { token } => events.push(("disconnect", token)), + PollEvent::Accept { stream, .. } => events.push(("accept", stream)), + PollEvent::Message { token, payload, .. } => { + assert_eq!(payload, b"after"); + payload_on = Some(token); + } + PollEvent::Reconnect { .. } => unreachable!(), + }); + client.poll_with(|e| { + if let PollEvent::Reconnect { token } = e { + assert_eq!(token, tok); + reconnected = true; + } + }); + thread::sleep(Duration::from_micros(50)); + } + assert_eq!(events[0], ("disconnect", first)); + assert_eq!(events[1].0, "accept"); + assert_ne!(events[1].1, first); + assert_eq!(payload_on, Some(events[1].1), "backlog replayed on the new session"); + assert_eq!(client.currently_disconnected().count(), 0); +} + +#[test] +fn udp_drop_backlog_on_disconnect_discards_queued() { + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()); + let mut client = udp(UdpConfig::lan()).with_drop_outbound_backlog_on_disconnect(true); + let (_, tok) = connect_pair(&mut server, &mut client, addr); + + client.disconnect(tok); + client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"lost")); + let mut reconnected = false; + let mut got = 0; + let deadline = Instant::now() + Duration::from_millis(500); + while Instant::now() < deadline { + server.poll_with(|e| { + if let PollEvent::Message { .. } = e { + got += 1; + } + }); + client.poll_with(|e| { + if let PollEvent::Reconnect { .. } = e { + reconnected = true; + } + }); + } + assert!(reconnected); + assert_eq!(got, 0); +} + +#[test] +fn udp_server_restart_reconnects_client() { + let addr = free_addr(); + let mut client = udp(UdpConfig::lan()).with_user_timeout(2_000); + let tok; + { + let mut server = udp(UdpConfig::lan()); + let (_, t) = connect_pair(&mut server, &mut client, addr); + tok = t; + } + // Old server gone. A new one on the same port must be rejoined without + // waiting for the 2s peer timeout: its reset acks trigger renegotiation. + let mut server = udp(UdpConfig::lan()); + server.listen_at(addr).unwrap(); + let start = Instant::now(); + let mut disconnected = false; + let mut reconnected = false; + let mut accepted = false; + while !(disconnected && reconnected && accepted) { + assert!(start.elapsed() < Duration::from_secs(5), "server restart recovery"); + client.poll_with(|e| match e { + PollEvent::Disconnect { token } => { + assert_eq!(token, tok); + disconnected = true; + } + PollEvent::Reconnect { token } => { + assert_eq!(token, tok); + reconnected = true; + } + _ => {} + }); + server.poll_with(|e| { + if let PollEvent::Accept { .. } = e { + accepted = true; + } + }); + thread::sleep(Duration::from_micros(50)); + } + assert!(start.elapsed() < Duration::from_millis(1500), "took the slow timeout path"); +} + +#[test] +fn udp_peer_timeout_disconnects_silent_client() { + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()).with_user_timeout(300); + let mut client = udp(UdpConfig::lan()); + let (accepted, _) = connect_pair(&mut server, &mut client, addr); + drop(client); + + let mut disconnected = None; + let start = Instant::now(); + while disconnected.is_none() { + assert!(start.elapsed() < Duration::from_secs(5), "peer timeout"); + server.poll_with(|e| { + if let PollEvent::Disconnect { token } = e { + disconnected = Some(token); + } + }); + thread::sleep(Duration::from_millis(1)); + } + assert_eq!(disconnected, Some(accepted)); + assert!(start.elapsed() >= Duration::from_millis(250)); +} + +#[test] +fn udp_backlog_limit_disconnects_non_consuming_peer() { + let addr = free_addr(); + let mut server = + udp(UdpConfig::lan()).with_max_backlog(8, flux_timing::Duration::from_millis(50)); + let mut client = udp(UdpConfig::lan()); + let (accepted, _) = connect_pair(&mut server, &mut client, addr); + // Client stops polling: nothing gets acked. + + let mut disconnected = false; + let deadline = Instant::now() + Duration::from_secs(5); + while !disconnected { + assert!(Instant::now() < deadline, "backlog disconnect"); + server.write_or_enqueue_with(SendBehavior::Single(accepted), |b| { + b.extend_from_slice(&[0; 1000]); + }); + server.poll_with(|e| { + if let PollEvent::Disconnect { token } = e { + assert_eq!(token, accepted); + disconnected = true; + } + }); + thread::sleep(Duration::from_millis(1)); + } + let _ = &client; +} + +#[test] +fn udp_ignores_junk_datagrams() { + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()); + let mut client = udp(UdpConfig::lan()); + let (accepted, _) = connect_pair(&mut server, &mut client, addr); + let junk = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + junk.send_to(b"not flux", addr).unwrap(); + junk.send_to(&[0xFF; 1200], addr).unwrap(); + junk.send_to(&[0; 2000], addr).unwrap(); + server.write_or_enqueue_with(SendBehavior::Single(accepted), |b| b.extend_from_slice(b"ok")); + let mut got = false; + let deadline = Instant::now() + Duration::from_secs(5); + while !got { + assert!(Instant::now() < deadline, "junk tolerance"); + server.poll_with(|e| assert!(!matches!(e, PollEvent::Accept { .. }))); + client.poll_with(|e| { + if let PollEvent::Message { payload, .. } = e { + assert_eq!(payload, b"ok"); + got = true; + } + }); + } +} + +/// A peer dropped mid-broadcast must not take the shared payload with it. +#[test] +fn udp_broadcast_survives_dropping_a_peer_mid_way() { + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()).with_max_backlog(4, flux_timing::Duration::ZERO); + server.listen_at(addr).unwrap(); + let mut stalled = udp(UdpConfig::lan()); + let mut live = udp(UdpConfig::lan()); + stalled.connect(addr).unwrap(); + live.connect(addr).unwrap(); + let mut accepted = 0; + let deadline = Instant::now() + Duration::from_secs(5); + while accepted < 2 || live.currently_disconnected().count() != 0 { + assert!(Instant::now() < deadline, "accepts"); + server.poll_with(|e| accepted += usize::from(matches!(e, PollEvent::Accept { .. }))); + stalled.poll_with(|_| {}); + live.poll_with(|_| {}); + } + // `stalled` stops polling: its unacked count grows past the backlog limit + // while `live` keeps consuming. + let msgs: Vec> = (0..12).map(|i| make_msg(i, 3000)).collect(); + let mut got = Vec::new(); + let mut dropped = 0; + let deadline = Instant::now() + Duration::from_secs(5); + for m in &msgs { + server.write_or_enqueue_with(SendBehavior::Broadcast, |b| b.extend_from_slice(m)); + let until = Instant::now() + Duration::from_millis(20); + while Instant::now() < until { + server.poll_with(|e| dropped += usize::from(matches!(e, PollEvent::Disconnect { .. }))); + live.poll_with(|e| { + if let PollEvent::Message { payload, .. } = e { + got.push((msg_id(payload), checksum(payload))); + } + }); + } + } + while got.len() < msgs.len() { + assert!(Instant::now() < deadline, "live peer delivery: got {}", got.len()); + server.poll_with(|_| {}); + live.poll_with(|e| { + if let PollEvent::Message { payload, .. } = e { + got.push((msg_id(payload), checksum(payload))); + } + }); + } + assert_eq!(dropped, 1, "the stalled peer is dropped exactly once"); + for (id, sum) in got { + assert_eq!(sum, checksum(&msgs[id as usize]), "message {id} corrupted"); + } +} + +/// Messages in flight when the session is cut arrive whole under the new +/// session, however much of them the receiver had already acked. +#[test] +fn udp_reconnect_replays_messages_queued_before_disconnect() { + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()); + let mut client = udp(UdpConfig::lan()); + let (_, tok) = connect_pair(&mut server, &mut client, addr); + let big = make_msg(7, 500_000); + client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(&big)); + client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"small")); + // Let some fragments through and get acked, then cut the session. How + // much lands first depends on the receive buffer; either message may. + let mut got = Vec::new(); + for _ in 0..3 { + client.poll_with(|_| {}); + server.poll_with(|e| { + if let PollEvent::Message { payload, .. } = e { + got.push(payload.to_vec()); + } + }); + } + client.disconnect(tok); + + // Everything not cumulatively acked is replayed, so the small message may + // arrive twice: delivery across a reconnect is at-least-once. + let deadline = Instant::now() + Duration::from_secs(5); + while !got.iter().any(|m| m.len() == big.len()) { + assert!(Instant::now() < deadline, "replay"); + server.poll_with(|e| { + if let PollEvent::Message { payload, .. } = e { + got.push(payload.to_vec()); + } + }); + client.poll_with(|_| {}); + } + assert!(got.iter().any(|m| m == b"small")); + assert!(got.iter().any(|m| checksum(m) == checksum(&big)), "big message replayed intact"); +} + +/// A server that drops an accepted peer makes the client renegotiate. +#[test] +fn udp_server_disconnect_reconnects_client() { + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()).with_user_timeout(5_000); + let mut client = udp(UdpConfig::lan()).with_user_timeout(5_000); + let (accepted, tok) = connect_pair(&mut server, &mut client, addr); + server.disconnect(accepted); + // Client traffic hits a listener with no session for it and gets reset. + client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"x")); + + let start = Instant::now(); + let (mut disconnected, mut reconnected, mut reaccepted) = (false, false, false); + while !(disconnected && reconnected && reaccepted) { + assert!(start.elapsed() < Duration::from_secs(5), "server-side disconnect recovery"); + client.poll_with(|e| match e { + PollEvent::Disconnect { token } => { + assert_eq!(token, tok); + disconnected = true; + } + PollEvent::Reconnect { token } => { + assert_eq!(token, tok); + reconnected = true; + } + _ => {} + }); + server.poll_with(|e| reaccepted |= matches!(e, PollEvent::Accept { .. })); + thread::sleep(Duration::from_micros(50)); + } + assert!(start.elapsed() < Duration::from_secs(1), "did not wait for the peer timeout"); +} + +/// A message that cannot fit the send window drops the peer instead of +/// vanishing silently. +#[test] +fn udp_window_exhaustion_disconnects_instead_of_dropping() { + let addr = free_addr(); + let config = UdpConfig { send_window: 64, max_message_size: 64 * 1171, ..UdpConfig::lan() }; + let mut server = udp(config); + let mut client = udp(config); + let (accepted, _) = connect_pair(&mut server, &mut client, addr); + // Client never polls: nothing is acked, the window fills, the 65th + // single-fragment message cannot be queued. + let mut disconnected = None; + for _ in 0..70 { + server.write_or_enqueue_with(SendBehavior::Single(accepted), |b| b.extend_from_slice(b"m")); + server.poll_with(|e| { + if let PollEvent::Disconnect { token } = e { + disconnected = Some(token); + } + }); + } + assert_eq!(disconnected, Some(accepted)); + let _ = &client; +} + +#[test] +fn sustained_large_messages_keep_acknowledgements_current() { + const COUNT: usize = 64; + let addr = free_addr(); + let mut server = udp(UdpConfig::lan()).with_socket_buf_size(16 * 1024 * 1024); + let mut client = udp(UdpConfig::lan()).with_socket_buf_size(16 * 1024 * 1024); + let (accepted, _) = connect_pair(&mut server, &mut client, addr); + let got = Arc::new(AtomicUsize::new(0)); + let progress = got.clone(); + let deadline = Instant::now() + Duration::from_secs(10); + let receiver = thread::spawn(move || { + let mut seen = [false; COUNT]; + while progress.load(Ordering::Relaxed) < COUNT { + client.poll_with(|event| { + if let PollEvent::Message { payload, .. } = event { + assert_eq!(payload.len(), 2 * 1024 * 1024); + let id = msg_id(payload) as usize; + assert!(!seen[id]); + assert!(payload[4..].iter().all(|b| *b == 0x5a)); + seen[id] = true; + progress.fetch_add(1, Ordering::Relaxed); + } + }); + assert!( + Instant::now() < deadline, + "large-message receiver stalled at {}", + progress.load(Ordering::Relaxed) + ); + } + }); + let mut sent = 0; + let mut payload = vec![0x5a; 2 * 1024 * 1024]; + while got.load(Ordering::Relaxed) < COUNT { + if sent < COUNT && sent - got.load(Ordering::Relaxed) < 4 { + payload[..4].copy_from_slice(&(sent as u32).to_le_bytes()); + server.write_or_enqueue_with(SendBehavior::Single(accepted), |buf| { + buf.extend_from_slice(&payload); + }); + sent += 1; + } + server.poll_with(|event| { + assert!( + !matches!(event, PollEvent::Disconnect { .. }), + "sender disconnected at {sent} sent, {} received", + got.load(Ordering::Relaxed) + ); + }); + assert!( + Instant::now() < deadline, + "large-message sender stalled at {sent} sent, {} received", + got.load(Ordering::Relaxed) + ); + } + receiver.join().unwrap(); +} diff --git a/crates/flux-network/tests/tcp_dcache.rs b/crates/flux-network/tests/tcp_dcache.rs index a843e1c..d4f6758 100644 --- a/crates/flux-network/tests/tcp_dcache.rs +++ b/crates/flux-network/tests/tcp_dcache.rs @@ -98,6 +98,15 @@ fn dcache_multi_stream_udp() { dcache_multi_stream(Transport::Udp(UdpConfig::lan())); } +#[cfg(target_os = "linux")] +#[test] +fn dcache_multi_stream_udp_uring() { + dcache_multi_stream(Transport::Udp(UdpConfig { + io: flux_network::udp::UdpIo::Uring(flux_network::udp::UringConfig::default()), + ..UdpConfig::lan() + })); +} + /// Two streams into the same dcache-backed spine queue. /// Verifies dcache bytes match the queue message (same shmem region). #[allow(clippy::significant_drop_tightening)] diff --git a/crates/flux-network/tests/udp_connector.rs b/crates/flux-network/tests/udp_connector.rs index 59d0ff1..c4b30e7 100644 --- a/crates/flux-network/tests/udp_connector.rs +++ b/crates/flux-network/tests/udp_connector.rs @@ -1,611 +1,3 @@ -use std::{ - net::{Ipv4Addr, SocketAddr, UdpSocket}, - sync::{ - Arc, - atomic::{AtomicBool, AtomicUsize, Ordering}, - }, - thread, - time::{Duration, Instant}, -}; +const IO: flux_network::udp::UdpIo = flux_network::udp::UdpIo::Syscall; -use flux_network::{NetworkDriver, PollEvent, SendBehavior, Transport, UdpConfig}; -use mio::Token; - -fn free_addr() -> SocketAddr { - UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).unwrap().local_addr().unwrap() -} - -fn udp(config: UdpConfig) -> NetworkDriver { - NetworkDriver::default().with_transport(Transport::Udp(config)) -} - -/// Handshake between a fresh listener and client, returning the accepted and -/// the outbound token. `dial` differs from `addr` when a relay sits between. -fn connect_via( - server: &mut NetworkDriver, - client: &mut NetworkDriver, - addr: SocketAddr, - dial: SocketAddr, -) -> (Token, Token) { - server.listen_at(addr).unwrap(); - let client_token = client.connect(dial).unwrap(); - let mut accepted = None; - let deadline = Instant::now() + Duration::from_secs(5); - while accepted.is_none() { - assert!(Instant::now() < deadline, "no accept"); - server.poll_with(|e| { - if let PollEvent::Accept { stream, .. } = e { - accepted = Some(stream); - } - }); - client.poll_with(|_| {}); - thread::sleep(Duration::from_micros(50)); - } - while client.currently_disconnected().count() != 0 { - assert!(Instant::now() < deadline, "handshake"); - server.poll_with(|_| {}); - client.poll_with(|_| {}); - thread::sleep(Duration::from_micros(50)); - } - // A hello retry may still be in flight if the ack took longer than one - // RTO. Let it land now, while the peer it belongs to still exists. - for _ in 0..5 { - server.poll_with(|_| {}); - client.poll_with(|_| {}); - } - (accepted.unwrap(), client_token) -} - -fn connect_pair( - server: &mut NetworkDriver, - client: &mut NetworkDriver, - addr: SocketAddr, -) -> (Token, Token) { - connect_via(server, client, addr, addr) -} - -fn checksum(bytes: &[u8]) -> u64 { - bytes - .iter() - .fold(0xcbf2_9ce4_8422_2325_u64, |h, b| (h ^ u64::from(*b)).wrapping_mul(0x100_0000_01b3)) -} - -/// Test message: 4-byte id, then `len` bytes derived from the id. -fn make_msg(id: u32, len: usize) -> Vec { - let mut v = Vec::with_capacity(4 + len); - v.extend_from_slice(&id.to_le_bytes()); - v.extend((0..len).map(|i| (id as usize).wrapping_mul(31).wrapping_add(i * 7) as u8)); - v -} - -fn msg_id(payload: &[u8]) -> u32 { - u32::from_le_bytes(payload[..4].try_into().unwrap()) -} - -#[test] -fn udp_roundtrip_before_handshake_completes() { - let addr = free_addr(); - let mut server = udp(UdpConfig::lan()); - server.listen_at(addr).unwrap(); - let mut client = udp(UdpConfig::lan()); - let tok = client.connect(addr).unwrap(); - // Queued before the hello ack arrives; must go out once it does. - client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"ping")); - - let mut accepted = None; - let mut request_seen = false; - let mut reply_seen = false; - let deadline = Instant::now() + Duration::from_secs(5); - while !reply_seen { - assert!(Instant::now() < deadline, "roundtrip timed out"); - server.poll_with(|e| match e { - PollEvent::Accept { stream, .. } => accepted = Some(stream), - PollEvent::Message { token, payload, .. } => { - assert_eq!(Some(token), accepted); - assert_eq!(payload, b"ping"); - request_seen = true; - } - _ => {} - }); - if request_seen && !reply_seen { - server.write_or_enqueue_with(SendBehavior::Single(accepted.unwrap()), |b| { - b.extend_from_slice(b"pong"); - }); - request_seen = false; - } - client.poll_with(|e| { - if let PollEvent::Message { token, payload, .. } = e { - assert_eq!(token, tok); - assert_eq!(payload, b"pong"); - reply_seen = true; - } - }); - thread::sleep(Duration::from_micros(50)); - } -} - -#[test] -fn udp_broadcast_mixed_sizes_to_two_subscribers() { - let addr = free_addr(); - let mut server = udp(UdpConfig::lan()); - server.listen_at(addr).unwrap(); - let mut a = udp(UdpConfig::lan()); - let mut b = udp(UdpConfig::lan()); - a.connect(addr).unwrap(); - b.connect(addr).unwrap(); - let mut accepted = 0; - let deadline = Instant::now() + Duration::from_secs(5); - while accepted < 2 { - assert!(Instant::now() < deadline, "accepts"); - server.poll_with(|e| { - if let PollEvent::Accept { .. } = e { - accepted += 1; - } - }); - a.poll_with(|_| {}); - b.poll_with(|_| {}); - } - - // 1 byte, one datagram, one stride exactly, and a 2 MiB message. - let sizes = [1usize, 100, 1171, 1172, 5000, 2 * 1024 * 1024]; - let msgs: Vec> = - sizes.iter().enumerate().map(|(i, s)| make_msg(i as u32, *s)).collect(); - for m in &msgs { - server.write_or_enqueue_with(SendBehavior::Broadcast, |buf| buf.extend_from_slice(m)); - } - - let expected: Vec = msgs.iter().map(|m| checksum(m)).collect(); - let mut got_a = Vec::new(); - let mut got_b = Vec::new(); - let deadline = Instant::now() + Duration::from_secs(10); - while got_a.len() < msgs.len() || got_b.len() < msgs.len() { - assert!(Instant::now() < deadline, "broadcast delivery"); - server.poll_with(|_| {}); - a.poll_with(|e| { - if let PollEvent::Message { payload, .. } = e { - got_a.push((msg_id(payload), checksum(payload))); - } - }); - b.poll_with(|e| { - if let PollEvent::Message { payload, .. } = e { - got_b.push((msg_id(payload), checksum(payload))); - } - }); - } - for got in [got_a, got_b] { - assert_eq!(got.len(), msgs.len()); - for (id, sum) in got { - assert_eq!(sum, expected[id as usize], "message {id} corrupted"); - } - } -} - -/// Forwards datagrams between one client and the server, dropping every -/// `drop_every`-th datagram in each direction. -struct LossyRelay { - addr: SocketAddr, - stop: Arc, - dropped: Arc, - handle: Option>, -} - -impl LossyRelay { - fn start(server: SocketAddr, drop_every: usize) -> Self { - let socket = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); - socket.set_read_timeout(Some(Duration::from_millis(5))).unwrap(); - let addr = socket.local_addr().unwrap(); - let stop = Arc::new(AtomicBool::new(false)); - let dropped = Arc::new(AtomicUsize::new(0)); - let (stop_c, dropped_c) = (stop.clone(), dropped.clone()); - let handle = thread::spawn(move || { - let mut buf = vec![0u8; 65_536]; - let mut client: Option = None; - let mut count = 0usize; - while !stop_c.load(Ordering::Relaxed) { - let Ok((n, from)) = socket.recv_from(&mut buf) else { continue }; - let to = if from == server { - let Some(c) = client else { continue }; - c - } else { - client = Some(from); - server - }; - count += 1; - if count.is_multiple_of(drop_every) { - dropped_c.fetch_add(1, Ordering::Relaxed); - continue; - } - let _ = socket.send_to(&buf[..n], to); - } - }); - Self { addr, stop, dropped, handle: Some(handle) } - } -} - -impl Drop for LossyRelay { - fn drop(&mut self) { - self.stop.store(true, Ordering::Relaxed); - self.handle.take().unwrap().join().unwrap(); - } -} - -#[test] -fn udp_delivers_everything_exactly_once_under_loss() { - const N: u32 = 400; - let server_addr = free_addr(); - let relay = LossyRelay::start(server_addr, 7); - - let mut server = udp(UdpConfig::lan()); - let mut client = udp(UdpConfig::lan()); - let (accepted, _) = connect_via(&mut server, &mut client, server_addr, relay.addr); - - // Both directions at once: server pushes to the client, client replies. - let msgs: Vec> = (0..N).map(|i| make_msg(i, 1 + (i as usize * 613) % 4000)).collect(); - for m in &msgs { - server.write_or_enqueue_with(SendBehavior::Single(accepted), |b| b.extend_from_slice(m)); - } - - let mut seen = vec![false; N as usize]; - let mut received = 0; - let mut echoed_back = 0; - let deadline = Instant::now() + Duration::from_secs(20); - while received < N || echoed_back < N { - assert!(Instant::now() < deadline, "loss recovery: {received} rx, {echoed_back} echoed"); - let mut echo = Vec::new(); - client.poll_with(|e| { - if let PollEvent::Message { payload, .. } = e { - let id = msg_id(payload) as usize; - assert_eq!(checksum(payload), checksum(&msgs[id]), "message {id} corrupted"); - assert!(!seen[id], "message {id} delivered twice"); - seen[id] = true; - received += 1; - echo.push(id as u32); - } - }); - for id in echo { - client.write_or_enqueue_with(SendBehavior::Broadcast, |b| { - b.extend_from_slice(&id.to_le_bytes()); - }); - } - server.poll_with(|e| { - if let PollEvent::Message { payload, .. } = e { - assert_eq!(payload.len(), 4); - echoed_back += 1; - } - }); - thread::sleep(Duration::from_micros(20)); - } - assert!(relay.dropped.load(Ordering::Relaxed) > 0, "relay dropped nothing"); -} - -#[test] -fn udp_client_disconnect_is_a_new_session() { - let addr = free_addr(); - let mut server = udp(UdpConfig::lan()); - let mut client = udp(UdpConfig::lan()); - let (first, tok) = connect_pair(&mut server, &mut client, addr); - - client.disconnect(tok); - assert_eq!(client.currently_disconnected().count(), 1); - client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"after")); - - let mut events = Vec::new(); - let mut payload_on = None; - let mut reconnected = false; - let deadline = Instant::now() + Duration::from_secs(5); - while payload_on.is_none() || !reconnected { - assert!(Instant::now() < deadline, "reconnect"); - server.poll_with(|e| match e { - PollEvent::Disconnect { token } => events.push(("disconnect", token)), - PollEvent::Accept { stream, .. } => events.push(("accept", stream)), - PollEvent::Message { token, payload, .. } => { - assert_eq!(payload, b"after"); - payload_on = Some(token); - } - PollEvent::Reconnect { .. } => unreachable!(), - }); - client.poll_with(|e| { - if let PollEvent::Reconnect { token } = e { - assert_eq!(token, tok); - reconnected = true; - } - }); - thread::sleep(Duration::from_micros(50)); - } - assert_eq!(events[0], ("disconnect", first)); - assert_eq!(events[1].0, "accept"); - assert_ne!(events[1].1, first); - assert_eq!(payload_on, Some(events[1].1), "backlog replayed on the new session"); - assert_eq!(client.currently_disconnected().count(), 0); -} - -#[test] -fn udp_drop_backlog_on_disconnect_discards_queued() { - let addr = free_addr(); - let mut server = udp(UdpConfig::lan()); - let mut client = udp(UdpConfig::lan()).with_drop_outbound_backlog_on_disconnect(true); - let (_, tok) = connect_pair(&mut server, &mut client, addr); - - client.disconnect(tok); - client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"lost")); - let mut reconnected = false; - let mut got = 0; - let deadline = Instant::now() + Duration::from_millis(500); - while Instant::now() < deadline { - server.poll_with(|e| { - if let PollEvent::Message { .. } = e { - got += 1; - } - }); - client.poll_with(|e| { - if let PollEvent::Reconnect { .. } = e { - reconnected = true; - } - }); - } - assert!(reconnected); - assert_eq!(got, 0); -} - -#[test] -fn udp_server_restart_reconnects_client() { - let addr = free_addr(); - let mut client = udp(UdpConfig::lan()).with_user_timeout(2_000); - let tok; - { - let mut server = udp(UdpConfig::lan()); - let (_, t) = connect_pair(&mut server, &mut client, addr); - tok = t; - } - // Old server gone. A new one on the same port must be rejoined without - // waiting for the 2s peer timeout: its reset acks trigger renegotiation. - let mut server = udp(UdpConfig::lan()); - server.listen_at(addr).unwrap(); - let start = Instant::now(); - let mut disconnected = false; - let mut reconnected = false; - let mut accepted = false; - while !(disconnected && reconnected && accepted) { - assert!(start.elapsed() < Duration::from_secs(5), "server restart recovery"); - client.poll_with(|e| match e { - PollEvent::Disconnect { token } => { - assert_eq!(token, tok); - disconnected = true; - } - PollEvent::Reconnect { token } => { - assert_eq!(token, tok); - reconnected = true; - } - _ => {} - }); - server.poll_with(|e| { - if let PollEvent::Accept { .. } = e { - accepted = true; - } - }); - thread::sleep(Duration::from_micros(50)); - } - assert!(start.elapsed() < Duration::from_millis(1500), "took the slow timeout path"); -} - -#[test] -fn udp_peer_timeout_disconnects_silent_client() { - let addr = free_addr(); - let mut server = udp(UdpConfig::lan()).with_user_timeout(300); - let mut client = udp(UdpConfig::lan()); - let (accepted, _) = connect_pair(&mut server, &mut client, addr); - drop(client); - - let mut disconnected = None; - let start = Instant::now(); - while disconnected.is_none() { - assert!(start.elapsed() < Duration::from_secs(5), "peer timeout"); - server.poll_with(|e| { - if let PollEvent::Disconnect { token } = e { - disconnected = Some(token); - } - }); - thread::sleep(Duration::from_millis(1)); - } - assert_eq!(disconnected, Some(accepted)); - assert!(start.elapsed() >= Duration::from_millis(250)); -} - -#[test] -fn udp_backlog_limit_disconnects_non_consuming_peer() { - let addr = free_addr(); - let mut server = - udp(UdpConfig::lan()).with_max_backlog(8, flux_timing::Duration::from_millis(50)); - let mut client = udp(UdpConfig::lan()); - let (accepted, _) = connect_pair(&mut server, &mut client, addr); - // Client stops polling: nothing gets acked. - - let mut disconnected = false; - let deadline = Instant::now() + Duration::from_secs(5); - while !disconnected { - assert!(Instant::now() < deadline, "backlog disconnect"); - server.write_or_enqueue_with(SendBehavior::Single(accepted), |b| { - b.extend_from_slice(&[0; 1000]); - }); - server.poll_with(|e| { - if let PollEvent::Disconnect { token } = e { - assert_eq!(token, accepted); - disconnected = true; - } - }); - thread::sleep(Duration::from_millis(1)); - } - let _ = &client; -} - -#[test] -fn udp_ignores_junk_datagrams() { - let addr = free_addr(); - let mut server = udp(UdpConfig::lan()); - let mut client = udp(UdpConfig::lan()); - let (accepted, _) = connect_pair(&mut server, &mut client, addr); - let junk = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); - junk.send_to(b"not flux", addr).unwrap(); - junk.send_to(&[0xFF; 1200], addr).unwrap(); - junk.send_to(&[0; 2000], addr).unwrap(); - server.write_or_enqueue_with(SendBehavior::Single(accepted), |b| b.extend_from_slice(b"ok")); - let mut got = false; - let deadline = Instant::now() + Duration::from_secs(5); - while !got { - assert!(Instant::now() < deadline, "junk tolerance"); - server.poll_with(|e| assert!(!matches!(e, PollEvent::Accept { .. }))); - client.poll_with(|e| { - if let PollEvent::Message { payload, .. } = e { - assert_eq!(payload, b"ok"); - got = true; - } - }); - } -} - -/// A peer dropped mid-broadcast must not take the shared payload with it. -#[test] -fn udp_broadcast_survives_dropping_a_peer_mid_way() { - let addr = free_addr(); - let mut server = udp(UdpConfig::lan()).with_max_backlog(4, flux_timing::Duration::ZERO); - server.listen_at(addr).unwrap(); - let mut stalled = udp(UdpConfig::lan()); - let mut live = udp(UdpConfig::lan()); - stalled.connect(addr).unwrap(); - live.connect(addr).unwrap(); - let mut accepted = 0; - let deadline = Instant::now() + Duration::from_secs(5); - while accepted < 2 || live.currently_disconnected().count() != 0 { - assert!(Instant::now() < deadline, "accepts"); - server.poll_with(|e| accepted += usize::from(matches!(e, PollEvent::Accept { .. }))); - stalled.poll_with(|_| {}); - live.poll_with(|_| {}); - } - // `stalled` stops polling: its unacked count grows past the backlog limit - // while `live` keeps consuming. - let msgs: Vec> = (0..12).map(|i| make_msg(i, 3000)).collect(); - let mut got = Vec::new(); - let mut dropped = 0; - let deadline = Instant::now() + Duration::from_secs(5); - for m in &msgs { - server.write_or_enqueue_with(SendBehavior::Broadcast, |b| b.extend_from_slice(m)); - let until = Instant::now() + Duration::from_millis(20); - while Instant::now() < until { - server.poll_with(|e| dropped += usize::from(matches!(e, PollEvent::Disconnect { .. }))); - live.poll_with(|e| { - if let PollEvent::Message { payload, .. } = e { - got.push((msg_id(payload), checksum(payload))); - } - }); - } - } - while got.len() < msgs.len() { - assert!(Instant::now() < deadline, "live peer delivery: got {}", got.len()); - server.poll_with(|_| {}); - live.poll_with(|e| { - if let PollEvent::Message { payload, .. } = e { - got.push((msg_id(payload), checksum(payload))); - } - }); - } - assert_eq!(dropped, 1, "the stalled peer is dropped exactly once"); - for (id, sum) in got { - assert_eq!(sum, checksum(&msgs[id as usize]), "message {id} corrupted"); - } -} - -/// Messages in flight when the session is cut arrive whole under the new -/// session, however much of them the receiver had already acked. -#[test] -fn udp_reconnect_replays_messages_queued_before_disconnect() { - let addr = free_addr(); - let mut server = udp(UdpConfig::lan()); - let mut client = udp(UdpConfig::lan()); - let (_, tok) = connect_pair(&mut server, &mut client, addr); - let big = make_msg(7, 500_000); - client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(&big)); - client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"small")); - // Let some fragments through and get acked, then cut the session. How - // much lands first depends on the receive buffer; either message may. - let mut got = Vec::new(); - for _ in 0..3 { - client.poll_with(|_| {}); - server.poll_with(|e| { - if let PollEvent::Message { payload, .. } = e { - got.push(payload.to_vec()); - } - }); - } - client.disconnect(tok); - - // Everything not cumulatively acked is replayed, so the small message may - // arrive twice: delivery across a reconnect is at-least-once. - let deadline = Instant::now() + Duration::from_secs(5); - while !got.iter().any(|m| m.len() == big.len()) { - assert!(Instant::now() < deadline, "replay"); - server.poll_with(|e| { - if let PollEvent::Message { payload, .. } = e { - got.push(payload.to_vec()); - } - }); - client.poll_with(|_| {}); - } - assert!(got.iter().any(|m| m == b"small")); - assert!(got.iter().any(|m| checksum(m) == checksum(&big)), "big message replayed intact"); -} - -/// A server that drops an accepted peer makes the client renegotiate. -#[test] -fn udp_server_disconnect_reconnects_client() { - let addr = free_addr(); - let mut server = udp(UdpConfig::lan()).with_user_timeout(5_000); - let mut client = udp(UdpConfig::lan()).with_user_timeout(5_000); - let (accepted, tok) = connect_pair(&mut server, &mut client, addr); - server.disconnect(accepted); - // Client traffic hits a listener with no session for it and gets reset. - client.write_or_enqueue_with(SendBehavior::Single(tok), |b| b.extend_from_slice(b"x")); - - let start = Instant::now(); - let (mut disconnected, mut reconnected, mut reaccepted) = (false, false, false); - while !(disconnected && reconnected && reaccepted) { - assert!(start.elapsed() < Duration::from_secs(5), "server-side disconnect recovery"); - client.poll_with(|e| match e { - PollEvent::Disconnect { token } => { - assert_eq!(token, tok); - disconnected = true; - } - PollEvent::Reconnect { token } => { - assert_eq!(token, tok); - reconnected = true; - } - _ => {} - }); - server.poll_with(|e| reaccepted |= matches!(e, PollEvent::Accept { .. })); - thread::sleep(Duration::from_micros(50)); - } - assert!(start.elapsed() < Duration::from_secs(1), "did not wait for the peer timeout"); -} - -/// A message that cannot fit the send window drops the peer instead of -/// vanishing silently. -#[test] -fn udp_window_exhaustion_disconnects_instead_of_dropping() { - let addr = free_addr(); - let config = UdpConfig { send_window: 64, max_message_size: 64 * 1171, ..UdpConfig::lan() }; - let mut server = udp(config); - let mut client = udp(config); - let (accepted, _) = connect_pair(&mut server, &mut client, addr); - // Client never polls: nothing is acked, the window fills, the 65th - // single-fragment message cannot be queued. - let mut disconnected = None; - for _ in 0..70 { - server.write_or_enqueue_with(SendBehavior::Single(accepted), |b| b.extend_from_slice(b"m")); - server.poll_with(|e| { - if let PollEvent::Disconnect { token } = e { - disconnected = Some(token); - } - }); - } - assert_eq!(disconnected, Some(accepted)); - let _ = &client; -} +include!("support/udp_connector.rs"); diff --git a/crates/flux-network/tests/udp_uring.rs b/crates/flux-network/tests/udp_uring.rs new file mode 100644 index 0000000..994fc29 --- /dev/null +++ b/crates/flux-network/tests/udp_uring.rs @@ -0,0 +1,90 @@ +#![cfg(target_os = "linux")] + +const IO: flux_network::udp::UdpIo = + flux_network::udp::UdpIo::Uring(flux_network::udp::UringConfig { + send_entries: 64, + recv_entries: 32, + }); + +include!("support/udp_connector.rs"); + +#[test] +fn duplex_progress_with_minimum_ring_capacity() { + let config = UdpConfig { + io: flux_network::udp::UdpIo::Uring(flux_network::udp::UringConfig { + send_entries: 1, + recv_entries: 1, + }), + ..UdpConfig::lan() + }; + let mut server = NetworkDriver::default().with_transport(Transport::Udp(config)); + let mut client = NetworkDriver::default().with_transport(Transport::Udp(config)); + let (accepted, outbound) = connect_pair(&mut server, &mut client, free_addr()); + for id in 0..32 { + let payload = make_msg(id, 64 * 1024); + server.write_or_enqueue_with(SendBehavior::Single(accepted), |buf| { + buf.extend_from_slice(&payload); + }); + client.write_or_enqueue_with(SendBehavior::Single(outbound), |buf| { + buf.extend_from_slice(&payload); + }); + let (mut server_got, mut client_got) = (false, false); + let deadline = Instant::now() + Duration::from_secs(5); + while !server_got || !client_got { + for (driver, got) in [(&mut server, &mut server_got), (&mut client, &mut client_got)] { + driver.poll_with(|event| match event { + PollEvent::Message { payload: bytes, .. } => { + assert_eq!(bytes, payload); + assert!(!*got); + *got = true; + } + PollEvent::Disconnect { .. } => panic!("duplex peer disconnected"), + _ => {} + }); + } + assert!(Instant::now() < deadline, "duplex message {id} stalled"); + } + } +} + +#[test] +fn mixed_backends_interoperate() { + for server_uring in [false, true] { + let config = |uring| UdpConfig { + io: if uring { IO } else { flux_network::udp::UdpIo::Syscall }, + ..UdpConfig::lan() + }; + let mut server = + NetworkDriver::default().with_transport(Transport::Udp(config(server_uring))); + let mut client = + NetworkDriver::default().with_transport(Transport::Udp(config(!server_uring))); + let addr = free_addr(); + let (accepted, outbound) = connect_pair(&mut server, &mut client, addr); + let payload = make_msg(17, 256 * 1024); + server.write_or_enqueue_with(SendBehavior::Single(accepted), |buf| { + buf.extend_from_slice(&payload); + }); + client.write_or_enqueue_with(SendBehavior::Single(outbound), |buf| { + buf.extend_from_slice(&payload); + }); + let deadline = Instant::now() + Duration::from_secs(5); + let (mut server_got, mut client_got) = (false, false); + while !server_got || !client_got { + server.poll_with(|e| { + if let PollEvent::Message { payload: got, .. } = e { + assert_eq!(got, payload); + assert!(!server_got); + server_got = true; + } + }); + client.poll_with(|e| { + if let PollEvent::Message { payload: got, .. } = e { + assert_eq!(got, payload); + assert!(!client_got); + client_got = true; + } + }); + assert!(Instant::now() < deadline); + } + } +}