use socket2::{Domain, Protocol, Socket, TcpKeepalive, Type};
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;
pub fn udp(addr: SocketAddr) -> std::io::Result<UdpSocket> {
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);
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(any(feature = "noq", feature = "quinn", feature = "quiche"))]
pub(crate) 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) -> std::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("[::]: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("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("127.0.0.1:0".parse().unwrap()).unwrap();
assert!(socket.local_addr().unwrap().is_ipv4());
}
#[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");
}
}