use socket2::{Domain, Protocol, Socket, TcpKeepalive, Type};
use std::io;
use std::net::{SocketAddr, TcpListener, UdpSocket};
use std::sync::Once;
use std::time::Duration;
const KEEPALIVE_IDLE: Duration = Duration::from_secs(30);
const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(10);
const UDP_BUFFER: usize = 8 * 1024 * 1024;
#[cfg(any(target_os = "linux", target_os = "android"))]
const RECV_SYSCTL: Option<&str> = Some("net.core.rmem_max");
#[cfg(any(
target_vendor = "apple",
target_os = "freebsd",
target_os = "netbsd",
target_os = "openbsd"
))]
const RECV_SYSCTL: Option<&str> = Some("kern.ipc.maxsockbuf");
#[cfg(not(any(
target_os = "linux",
target_os = "android",
target_vendor = "apple",
target_os = "freebsd",
target_os = "netbsd",
target_os = "openbsd"
)))]
const RECV_SYSCTL: Option<&str> = None;
#[cfg(any(target_os = "linux", target_os = "android"))]
const SEND_SYSCTL: Option<&str> = Some("net.core.wmem_max");
#[cfg(not(any(target_os = "linux", target_os = "android")))]
const SEND_SYSCTL: Option<&str> = RECV_SYSCTL;
#[derive(Clone, Copy, Debug)]
pub struct Udp {
addr: SocketAddr,
reuse_port: bool,
}
impl Udp {
pub fn new(addr: SocketAddr) -> Self {
Self {
addr,
reuse_port: false,
}
}
pub fn with_reuse_port(mut self, enabled: bool) -> Self {
self.reuse_port = enabled;
self
}
}
pub fn udp(options: Udp) -> io::Result<UdpSocket> {
let addr = options.addr;
let domain = if addr.is_ipv4() { Domain::IPV4 } else { Domain::IPV6 };
let socket = Socket::new(domain, Type::DGRAM, Some(Protocol::UDP))?;
make_dual_stack(&socket, addr);
grow_buffers(&socket);
if options.reuse_port {
set_reuse_port(&socket)?;
}
socket.bind(&addr.into())?;
Ok(socket.into())
}
fn grow_buffers(socket: &Socket) {
for direction in [Direction::Recv, Direction::Send] {
direction.grow(socket);
}
}
#[derive(Clone, Copy)]
enum Direction {
Recv,
Send,
}
impl Direction {
fn grow(self, socket: &Socket) {
if self.size(socket).is_ok_and(sufficient) {
return;
}
match self.set_size(socket, UDP_BUFFER).and_then(|()| self.size(socket)) {
Ok(reported) if sufficient(reported) => {}
Ok(reported) => self.warn_short(granted(reported)),
Err(err) => self.warn_failed(&err),
}
}
fn size(self, socket: &Socket) -> std::io::Result<usize> {
match self {
Self::Recv => socket.recv_buffer_size(),
Self::Send => socket.send_buffer_size(),
}
}
fn set_size(self, socket: &Socket, size: usize) -> std::io::Result<()> {
match self {
Self::Recv => socket.set_recv_buffer_size(size),
Self::Send => socket.set_send_buffer_size(size),
}
}
fn name(self) -> &'static str {
match self {
Self::Recv => "receive",
Self::Send => "send",
}
}
fn sysctl(self) -> Option<&'static str> {
match self {
Self::Recv => RECV_SYSCTL,
Self::Send => SEND_SYSCTL,
}
}
fn warned(self) -> &'static Once {
static RECV: Once = Once::new();
static SEND: Once = Once::new();
match self {
Self::Recv => &RECV,
Self::Send => &SEND,
}
}
fn warn_short(self, granted: usize) {
self.warned().call_once(|| self.emit_short(granted));
}
fn emit_short(self, granted: usize) {
let name = self.name();
match self.sysctl() {
Some(sysctl) => tracing::warn!(
wanted = UDP_BUFFER,
granted,
"UDP {name} buffer is smaller than requested; raise `{sysctl}` or expect packet loss under load"
),
None => tracing::warn!(
wanted = UDP_BUFFER,
granted,
"UDP {name} buffer is smaller than requested; expect packet loss under load"
),
}
}
fn warn_failed(self, err: &std::io::Error) {
let name = self.name();
self.warned()
.call_once(|| tracing::warn!(%err, "failed to set the UDP {name} buffer size"));
}
}
fn sufficient(reported: usize) -> bool {
granted(reported) >= UDP_BUFFER
}
fn granted(reported: usize) -> usize {
if cfg!(any(target_os = "linux", target_os = "android")) {
reported / 2
} else {
reported
}
}
#[cfg(target_os = "linux")]
fn set_reuse_port(socket: &Socket) -> io::Result<()> {
socket.set_reuse_port(true)
}
#[cfg(not(target_os = "linux"))]
fn set_reuse_port(_socket: &Socket) -> io::Result<()> {
Err(io::Error::new(
io::ErrorKind::Unsupported,
"SO_REUSEPORT load balancing is Linux-only",
))
}
pub fn udp_is_dual_stack(socket: &UdpSocket) -> bool {
match socket.local_addr() {
Ok(addr) if addr.is_ipv6() => socket2::SockRef::from(socket).only_v6().is_ok_and(|only| !only),
_ => false,
}
}
pub fn tcp(addr: SocketAddr) -> io::Result<TcpListener> {
let domain = if addr.is_ipv4() { Domain::IPV4 } else { Domain::IPV6 };
let socket = Socket::new(domain, Type::STREAM, Some(Protocol::TCP))?;
make_dual_stack(&socket, addr);
#[cfg(not(windows))]
socket.set_reuse_address(true)?;
let keepalive = TcpKeepalive::new()
.with_time(KEEPALIVE_IDLE)
.with_interval(KEEPALIVE_INTERVAL);
if let Err(err) = socket.set_tcp_keepalive(&keepalive) {
tracing::warn!(%err, "failed to enable TCP keepalive; dead peers may linger");
}
socket.bind(&addr.into())?;
socket.listen(1024)?;
let listener: TcpListener = socket.into();
listener.set_nonblocking(true)?;
Ok(listener)
}
fn make_dual_stack(socket: &Socket, addr: SocketAddr) {
if addr.is_ipv6()
&& let Err(err) = socket.set_only_v6(false)
{
tracing::warn!(%err, "failed to enable dual-stack IPv6 socket; IPv4 clients may be unreachable");
}
}
#[cfg(test)]
mod tests {
use super::*;
fn skip_if_no_ipv6(err: &std::io::Error) -> bool {
const NO_IPV6_ERRNOS: &[i32] = &[97, 99, 93, 10047, 10049, 10043];
let no_ipv6 = matches!(
err.kind(),
std::io::ErrorKind::AddrNotAvailable | std::io::ErrorKind::Unsupported
) || err.raw_os_error().is_some_and(|code| NO_IPV6_ERRNOS.contains(&code));
if no_ipv6 {
eprintln!("skipping: host has no IPv6 support ({err})");
}
no_ipv6
}
#[test]
fn udp_ipv6_is_dual_stack() {
let socket = match udp(Udp::new("[::]:0".parse().unwrap())) {
Ok(socket) => socket,
Err(err) if skip_if_no_ipv6(&err) => return,
Err(err) => panic!("failed to bind IPv6 UDP socket: {err}"),
};
let socket = Socket::from(socket);
assert!(!socket.only_v6().unwrap(), "IPv6 socket should be dual-stack");
}
#[test]
fn udp_buffers_grow() {
fn check(direction: Direction) {
let plain = Socket::from(std::net::UdpSocket::bind("127.0.0.1:0").unwrap());
let before = direction.size(&plain).unwrap();
let tuned = Socket::from(udp(Udp::new("127.0.0.1:0".parse().unwrap())).unwrap());
let after = direction.size(&tuned).unwrap();
if sufficient(before) {
assert_eq!(after, before, "{} buffer should be left alone", direction.name());
} else {
assert!(after > before, "{} buffer should grow past {before}", direction.name());
}
}
check(Direction::Recv);
check(Direction::Send);
}
#[test]
fn sufficient_accounts_for_the_doubled_report() {
assert!(sufficient(UDP_BUFFER * 2));
assert!(!sufficient(512 * 1024));
}
#[tracing_test::traced_test]
#[test]
fn a_clamped_buffer_warns_and_names_the_sysctl() {
Direction::Recv.emit_short(512 * 1024);
assert!(logs_contain("UDP receive buffer is smaller than requested"));
if let Some(sysctl) = Direction::Recv.sysctl() {
assert!(logs_contain(sysctl));
}
}
#[test]
fn udp_ipv4_still_binds() {
let socket = udp(Udp::new("127.0.0.1:0".parse().unwrap())).unwrap();
assert!(socket.local_addr().unwrap().is_ipv4());
}
#[test]
#[cfg(target_os = "linux")]
fn udp_reuse_port_shares_a_port() {
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let first = udp(Udp::new(addr).with_reuse_port(true)).unwrap();
let bound = first.local_addr().unwrap();
let second = udp(Udp::new(bound).with_reuse_port(true)).unwrap();
assert_eq!(second.local_addr().unwrap(), bound);
}
#[test]
#[cfg(target_os = "linux")]
fn udp_without_reuse_port_keeps_the_port() {
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let first = udp(Udp::new(addr)).unwrap();
let bound = first.local_addr().unwrap();
assert!(udp(Udp::new(bound).with_reuse_port(true)).is_err());
}
#[test]
#[cfg(target_os = "linux")]
fn udp_reuse_port_spreads_datagrams() {
const MEMBERS: usize = 4;
const SENDERS: usize = 64;
let mut group = Vec::with_capacity(MEMBERS);
let mut addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
for _ in 0..MEMBERS {
let socket = udp(Udp::new(addr).with_reuse_port(true)).unwrap();
addr = socket.local_addr().unwrap();
socket.set_nonblocking(true).unwrap();
group.push(socket);
}
for _ in 0..SENDERS {
let sender = udp(Udp::new("127.0.0.1:0".parse().unwrap())).unwrap();
sender.send_to(b"quic", addr).unwrap();
}
let mut fed = 0;
let mut total = 0;
for socket in &group {
let mut received = 0;
let mut buf = [0u8; 8];
while socket.recv_from(&mut buf).is_ok() {
received += 1;
}
total += received;
fed += usize::from(received > 0);
}
assert_eq!(total, SENDERS, "every datagram reached exactly one member");
assert!(fed > 1, "only {fed} of {MEMBERS} members were fed");
}
#[test]
#[cfg(not(target_os = "linux"))]
fn udp_reuse_port_is_linux_only() {
let addr: SocketAddr = "127.0.0.1:0".parse().unwrap();
let err = udp(Udp::new(addr).with_reuse_port(true)).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::Unsupported);
}
#[test]
fn tcp_ipv6_is_dual_stack() {
let listener = match tcp("[::]:0".parse().unwrap()) {
Ok(listener) => listener,
Err(err) if skip_if_no_ipv6(&err) => return,
Err(err) => panic!("failed to bind IPv6 TCP listener: {err}"),
};
let socket = Socket::from(listener);
assert!(!socket.only_v6().unwrap(), "IPv6 listener should be dual-stack");
}
}