1use socket2::{Domain, Protocol, Socket, TcpKeepalive, Type};
13use std::net::{SocketAddr, TcpListener, UdpSocket};
14use std::sync::Once;
15use std::time::Duration;
16
17const KEEPALIVE_IDLE: Duration = Duration::from_secs(30);
25const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(10);
26
27const UDP_BUFFER: usize = 8 * 1024 * 1024;
42
43#[cfg(any(target_os = "linux", target_os = "android"))]
46const RECV_SYSCTL: Option<&str> = Some("net.core.rmem_max");
47#[cfg(any(
48 target_vendor = "apple",
49 target_os = "freebsd",
50 target_os = "netbsd",
51 target_os = "openbsd"
52))]
53const RECV_SYSCTL: Option<&str> = Some("kern.ipc.maxsockbuf");
54#[cfg(not(any(
55 target_os = "linux",
56 target_os = "android",
57 target_vendor = "apple",
58 target_os = "freebsd",
59 target_os = "netbsd",
60 target_os = "openbsd"
61)))]
62const RECV_SYSCTL: Option<&str> = None;
63
64#[cfg(any(target_os = "linux", target_os = "android"))]
67const SEND_SYSCTL: Option<&str> = Some("net.core.wmem_max");
68#[cfg(not(any(target_os = "linux", target_os = "android")))]
69const SEND_SYSCTL: Option<&str> = RECV_SYSCTL;
70
71pub fn udp(addr: SocketAddr) -> std::io::Result<UdpSocket> {
76 let domain = if addr.is_ipv4() { Domain::IPV4 } else { Domain::IPV6 };
77 let socket = Socket::new(domain, Type::DGRAM, Some(Protocol::UDP))?;
78 make_dual_stack(&socket, addr);
79 grow_buffers(&socket);
80 socket.bind(&addr.into())?;
81 Ok(socket.into())
82}
83
84fn grow_buffers(socket: &Socket) {
86 for direction in [Direction::Recv, Direction::Send] {
87 direction.grow(socket);
88 }
89}
90
91#[derive(Clone, Copy)]
94enum Direction {
95 Recv,
96 Send,
97}
98
99impl Direction {
100 fn grow(self, socket: &Socket) {
108 if self.size(socket).is_ok_and(sufficient) {
110 return;
111 }
112
113 match self.set_size(socket, UDP_BUFFER).and_then(|()| self.size(socket)) {
114 Ok(reported) if sufficient(reported) => {}
115 Ok(reported) => self.warn_short(granted(reported)),
116 Err(err) => self.warn_failed(&err),
117 }
118 }
119
120 fn size(self, socket: &Socket) -> std::io::Result<usize> {
121 match self {
122 Self::Recv => socket.recv_buffer_size(),
123 Self::Send => socket.send_buffer_size(),
124 }
125 }
126
127 fn set_size(self, socket: &Socket, size: usize) -> std::io::Result<()> {
128 match self {
129 Self::Recv => socket.set_recv_buffer_size(size),
130 Self::Send => socket.set_send_buffer_size(size),
131 }
132 }
133
134 fn name(self) -> &'static str {
135 match self {
136 Self::Recv => "receive",
137 Self::Send => "send",
138 }
139 }
140
141 fn sysctl(self) -> Option<&'static str> {
142 match self {
143 Self::Recv => RECV_SYSCTL,
144 Self::Send => SEND_SYSCTL,
145 }
146 }
147
148 fn warned(self) -> &'static Once {
151 static RECV: Once = Once::new();
152 static SEND: Once = Once::new();
153
154 match self {
155 Self::Recv => &RECV,
156 Self::Send => &SEND,
157 }
158 }
159
160 fn warn_short(self, granted: usize) {
162 self.warned().call_once(|| self.emit_short(granted));
163 }
164
165 fn emit_short(self, granted: usize) {
168 let name = self.name();
169 match self.sysctl() {
170 Some(sysctl) => tracing::warn!(
171 wanted = UDP_BUFFER,
172 granted,
173 "UDP {name} buffer is smaller than requested; raise `{sysctl}` or expect packet loss under load"
174 ),
175 None => tracing::warn!(
176 wanted = UDP_BUFFER,
177 granted,
178 "UDP {name} buffer is smaller than requested; expect packet loss under load"
179 ),
180 }
181 }
182
183 fn warn_failed(self, err: &std::io::Error) {
185 let name = self.name();
186 self.warned()
187 .call_once(|| tracing::warn!(%err, "failed to set the UDP {name} buffer size"));
188 }
189}
190
191fn sufficient(reported: usize) -> bool {
193 granted(reported) >= UDP_BUFFER
194}
195
196fn granted(reported: usize) -> usize {
202 if cfg!(any(target_os = "linux", target_os = "android")) {
203 reported / 2
204 } else {
205 reported
206 }
207}
208
209#[cfg(any(feature = "noq", feature = "quinn", feature = "quiche"))]
217pub(crate) fn udp_is_dual_stack(socket: &UdpSocket) -> bool {
218 match socket.local_addr() {
219 Ok(addr) if addr.is_ipv6() => socket2::SockRef::from(socket).only_v6().is_ok_and(|only| !only),
220 _ => false,
221 }
222}
223
224pub fn tcp(addr: SocketAddr) -> std::io::Result<TcpListener> {
229 let domain = if addr.is_ipv4() { Domain::IPV4 } else { Domain::IPV6 };
230 let socket = Socket::new(domain, Type::STREAM, Some(Protocol::TCP))?;
231 make_dual_stack(&socket, addr);
232 #[cfg(not(windows))]
235 socket.set_reuse_address(true)?;
236 let keepalive = TcpKeepalive::new()
242 .with_time(KEEPALIVE_IDLE)
243 .with_interval(KEEPALIVE_INTERVAL);
244 if let Err(err) = socket.set_tcp_keepalive(&keepalive) {
245 tracing::warn!(%err, "failed to enable TCP keepalive; dead peers may linger");
246 }
247 socket.bind(&addr.into())?;
248 socket.listen(1024)?;
249 let listener: TcpListener = socket.into();
250 listener.set_nonblocking(true)?;
251 Ok(listener)
252}
253
254fn make_dual_stack(socket: &Socket, addr: SocketAddr) {
258 if addr.is_ipv6()
259 && let Err(err) = socket.set_only_v6(false)
260 {
261 tracing::warn!(%err, "failed to enable dual-stack IPv6 socket; IPv4 clients may be unreachable");
262 }
263}
264
265#[cfg(test)]
266mod tests {
267 use super::*;
268
269 fn skip_if_no_ipv6(err: &std::io::Error) -> bool {
275 const NO_IPV6_ERRNOS: &[i32] = &[97, 99, 93, 10047, 10049, 10043];
278 let no_ipv6 = matches!(
279 err.kind(),
280 std::io::ErrorKind::AddrNotAvailable | std::io::ErrorKind::Unsupported
281 ) || err.raw_os_error().is_some_and(|code| NO_IPV6_ERRNOS.contains(&code));
282 if no_ipv6 {
283 eprintln!("skipping: host has no IPv6 support ({err})");
284 }
285 no_ipv6
286 }
287
288 #[test]
289 fn udp_ipv6_is_dual_stack() {
290 let socket = match udp("[::]:0".parse().unwrap()) {
293 Ok(socket) => socket,
294 Err(err) if skip_if_no_ipv6(&err) => return,
295 Err(err) => panic!("failed to bind IPv6 UDP socket: {err}"),
296 };
297 let socket = Socket::from(socket);
298 assert!(!socket.only_v6().unwrap(), "IPv6 socket should be dual-stack");
299 }
300
301 #[test]
302 fn udp_buffers_grow() {
303 fn check(direction: Direction) {
304 let plain = Socket::from(std::net::UdpSocket::bind("127.0.0.1:0").unwrap());
305 let before = direction.size(&plain).unwrap();
306
307 let tuned = Socket::from(udp("127.0.0.1:0".parse().unwrap()).unwrap());
308 let after = direction.size(&tuned).unwrap();
309
310 if sufficient(before) {
314 assert_eq!(after, before, "{} buffer should be left alone", direction.name());
315 } else {
316 assert!(after > before, "{} buffer should grow past {before}", direction.name());
317 }
318 }
319
320 check(Direction::Recv);
321 check(Direction::Send);
322 }
323
324 #[test]
325 fn sufficient_accounts_for_the_doubled_report() {
326 assert!(sufficient(UDP_BUFFER * 2));
328 assert!(!sufficient(512 * 1024));
329 }
330
331 #[tracing_test::traced_test]
332 #[test]
333 fn a_clamped_buffer_warns_and_names_the_sysctl() {
334 Direction::Recv.emit_short(512 * 1024);
335
336 assert!(logs_contain("UDP receive buffer is smaller than requested"));
337 if let Some(sysctl) = Direction::Recv.sysctl() {
338 assert!(logs_contain(sysctl));
339 }
340 }
341
342 #[test]
343 fn udp_ipv4_still_binds() {
344 let socket = udp("127.0.0.1:0".parse().unwrap()).unwrap();
345 assert!(socket.local_addr().unwrap().is_ipv4());
346 }
347
348 #[test]
349 fn tcp_ipv6_is_dual_stack() {
350 let listener = match tcp("[::]:0".parse().unwrap()) {
351 Ok(listener) => listener,
352 Err(err) if skip_if_no_ipv6(&err) => return,
353 Err(err) => panic!("failed to bind IPv6 TCP listener: {err}"),
354 };
355 let socket = Socket::from(listener);
356 assert!(!socket.only_v6().unwrap(), "IPv6 listener should be dual-stack");
357 }
358}