Skip to main content

tokio_fast_udp/
socket.rs

1// FastUdpSocket and FastUdpSocketBuilder for Unix (AsyncFd + nix syscalls)
2// and Windows (tokio::net::UdpSocket).
3
4// ---------------------------------------------------------------------------
5// Unix implementation
6// ---------------------------------------------------------------------------
7#[cfg(unix)]
8mod unix_socket {
9    use std::{
10        io,
11        net::SocketAddr,
12        os::fd::{AsRawFd, OwnedFd},
13    };
14
15    // -- platform alias --------------------------------------------------------
16    #[cfg(target_os = "linux")]
17    use linux_inner as plat;
18    use tokio::io::{Interest, unix::AsyncFd};
19    #[cfg(all(unix, not(target_os = "linux")))]
20    use unix_inner as plat;
21
22    use crate::{
23        capability::Capabilities,
24        ecn::Ecn,
25        item::{ReceiveItem, SendItem},
26    };
27
28    // -- FastUdpSocket ---------------------------------------------------------
29    pub struct FastUdpSocket {
30        fd: AsyncFd<OwnedFd>,
31        caps: Capabilities,
32    }
33
34    impl FastUdpSocket {
35        pub fn build(addr: SocketAddr) -> FastUdpSocketBuilder {
36            FastUdpSocketBuilder {
37                addr,
38                disable_gso: false,
39                disable_gro: false,
40                disable_sendmmsg: false,
41                disable_ecn: false,
42                max_batch: 64,
43                enable_reuse_port: false,
44            }
45        }
46    }
47
48    impl FastUdpSocket {
49        pub async fn send(&self, item: SendItem<'_>) -> io::Result<()> {
50            let items = [item];
51            self.send_many(&items).await?;
52            Ok(())
53        }
54
55        pub async fn send_many(&self, items: &[SendItem<'_>]) -> io::Result<usize> {
56            if items.is_empty() {
57                return Ok(0);
58            }
59            let caps = &self.caps;
60            let limit = items.len().min(caps.max_batch);
61            let items = &items[..limit];
62            let mut gso_einval = false;
63            let mut chunk_offset = 0usize;
64            self.fd
65                .async_io(Interest::WRITABLE, |inner| {
66                    plat::send_batch(inner.as_raw_fd(), items, caps, &mut gso_einval, &mut chunk_offset)
67                })
68                .await
69        }
70
71        pub async fn receive(&self, item: &mut ReceiveItem<'_>) -> io::Result<()> {
72            self.receive_many(std::slice::from_mut(item)).await?;
73            Ok(())
74        }
75
76        pub async fn receive_many(&self, items: &mut [ReceiveItem<'_>]) -> io::Result<usize> {
77            if items.is_empty() {
78                return Ok(0);
79            }
80            let caps = &self.caps;
81            let limit = items.len().min(caps.max_batch);
82            let items = &mut items[..limit];
83            self.fd
84                .async_io(Interest::READABLE, |inner| plat::recv_batch(inner.as_raw_fd(), items, caps))
85                .await
86        }
87
88        pub fn capabilities(&self) -> &Capabilities {
89            &self.caps
90        }
91
92        pub fn local_addr(&self) -> io::Result<SocketAddr> {
93            plat::local_addr(self.fd.get_ref())
94        }
95    }
96
97    // -- FastUdpSocketBuilder --------------------------------------------------
98    pub struct FastUdpSocketBuilder {
99        addr: SocketAddr,
100        disable_gso: bool,
101        disable_gro: bool,
102        disable_sendmmsg: bool,
103        disable_ecn: bool,
104        max_batch: usize,
105        enable_reuse_port: bool,
106    }
107
108    impl FastUdpSocketBuilder {
109        pub fn gso(mut self, enable: bool) -> Self {
110            self.disable_gso = !enable;
111            self
112        }
113
114        pub fn gro(mut self, enable: bool) -> Self {
115            self.disable_gro = !enable;
116            self
117        }
118
119        pub fn sendmmsg(mut self, enable: bool) -> Self {
120            self.disable_sendmmsg = !enable;
121            self
122        }
123
124        pub fn ecn(mut self, enable: bool) -> Self {
125            self.disable_ecn = !enable;
126            self
127        }
128
129        pub fn reuse_port(mut self, enable: bool) -> Self {
130            self.enable_reuse_port = enable;
131            self
132        }
133
134        pub fn max_batch_size(mut self, n: usize) -> Self {
135            self.max_batch = n;
136            self
137        }
138
139        pub fn bind(self) -> io::Result<FastUdpSocket> {
140            let is_ipv6 = self.addr.is_ipv6();
141
142            let fd: OwnedFd = if self.enable_reuse_port {
143                create_socket_with_reuse_port(self.addr)?
144            } else {
145                let std_socket = std::net::UdpSocket::bind(self.addr)?;
146                std_socket.set_nonblocking(true)?;
147
148                use std::os::fd::{FromRawFd, IntoRawFd};
149                let raw_fd = std_socket.into_raw_fd();
150                unsafe { OwnedFd::from_raw_fd(raw_fd) }
151            };
152
153            let config = BuildConfig {
154                disable_gso: self.disable_gso,
155                disable_gro: self.disable_gro,
156                disable_sendmmsg: self.disable_sendmmsg,
157                disable_ecn: self.disable_ecn,
158                max_batch: self.max_batch,
159            };
160
161            let caps = plat::setup_and_probe(&fd, &config, is_ipv6)?;
162            let async_fd = AsyncFd::new(fd)?;
163
164            Ok(FastUdpSocket {
165                fd: async_fd,
166                caps,
167            })
168        }
169    }
170
171    // -- shared nix helpers (used by the platform inner modules via super::) --
172
173    use nix::sys::socket::{ControlMessageOwned, SockaddrStorage};
174
175    fn nix_err(e: nix::errno::Errno) -> io::Error {
176        e.into()
177    }
178
179    fn create_socket_with_reuse_port(addr: SocketAddr) -> io::Result<OwnedFd> {
180        use nix::sys::socket::{AddressFamily, SockFlag, SockType, bind, setsockopt, socket, sockopt::ReusePort};
181
182        let is_ipv6 = addr.is_ipv6();
183        let domain = if is_ipv6 {
184            AddressFamily::Inet6
185        } else {
186            AddressFamily::Inet
187        };
188
189        let fd = socket(
190            domain,
191            SockType::Datagram,
192            SockFlag::SOCK_NONBLOCK | SockFlag::SOCK_CLOEXEC,
193            None,
194        )
195        .map_err(nix_err)?;
196
197        setsockopt(&fd, ReusePort, &true).map_err(nix_err)?;
198
199        let sock_addr = SockaddrStorage::from(addr);
200        bind(fd.as_raw_fd(), &sock_addr).map_err(nix_err)?;
201
202        Ok(fd)
203    }
204
205    fn ecn_to_tos(ecn: Ecn) -> u8 {
206        ecn.to_tos_bits()
207    }
208
209    fn parse_ecn(cmsgs: &[ControlMessageOwned]) -> Option<Ecn> {
210        for cmsg in cmsgs {
211            match cmsg {
212                ControlMessageOwned::Ipv4Tos(tos) => return Some(Ecn::from_tos_bits(*tos)),
213                ControlMessageOwned::Ipv6TClass(tc) => return Some(Ecn::from_tos_bits(*tc as u8)),
214                _ => {}
215            }
216        }
217        None
218    }
219
220    fn storage_to_addr(addr: &SockaddrStorage) -> SocketAddr {
221        if let Some(v4) = addr.as_sockaddr_in() {
222            return (*v4).into();
223        }
224        if let Some(v6) = addr.as_sockaddr_in6() {
225            return (*v6).into();
226        }
227        SocketAddr::new(std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED), 0)
228    }
229
230    pub(crate) struct BuildConfig {
231        pub disable_gso: bool,
232        pub disable_gro: bool,
233        pub disable_sendmmsg: bool,
234        pub disable_ecn: bool,
235        pub max_batch: usize,
236    }
237
238    // -- Linux fast path ----------------------------------------------------
239    #[cfg(target_os = "linux")]
240    mod linux_inner {
241        use std::{
242            io,
243            net::SocketAddr,
244            os::fd::{AsRawFd, OwnedFd},
245        };
246
247        use nix::sys::socket::{ControlMessage, MsgFlags, SockaddrStorage, recvmsg, sendmsg};
248
249        use super::{BuildConfig, ecn_to_tos, nix_err, parse_ecn, storage_to_addr};
250        use crate::{
251            capability::Capabilities,
252            ecn::Ecn,
253            item::{ReceiveItem, SendItem},
254        };
255
256        fn send_one(fd: std::os::fd::RawFd, item: &SendItem, ecn_enabled: bool, gso_enabled: bool) -> io::Result<()> {
257            let iov = [std::io::IoSlice::new(item.data)];
258            let addr = SockaddrStorage::from(item.destination);
259
260            let has_gso = gso_enabled && item.segment_size.is_some();
261            let has_ecn = ecn_enabled && item.ecn.is_some();
262            let is_v4 = item.destination.is_ipv4();
263
264            match (has_gso, has_ecn, is_v4) {
265                (true, true, true) => {
266                    let s = item.segment_size.as_ref().unwrap();
267                    let tos = &ecn_to_tos(item.ecn.unwrap());
268                    sendmsg(
269                        fd,
270                        &iov,
271                        &[ControlMessage::UdpGsoSegments(s), ControlMessage::Ipv4Tos(tos)],
272                        MsgFlags::empty(),
273                        Some(&addr),
274                    )
275                }
276                (true, true, false) => {
277                    let s = item.segment_size.as_ref().unwrap();
278                    let tc = &(ecn_to_tos(item.ecn.unwrap()) as i32);
279                    sendmsg(
280                        fd,
281                        &iov,
282                        &[ControlMessage::UdpGsoSegments(s), ControlMessage::Ipv6TClass(tc)],
283                        MsgFlags::empty(),
284                        Some(&addr),
285                    )
286                }
287                (true, false, _) => {
288                    let s = item.segment_size.as_ref().unwrap();
289                    sendmsg(fd, &iov, &[ControlMessage::UdpGsoSegments(s)], MsgFlags::empty(), Some(&addr))
290                }
291                (false, true, true) => {
292                    let tos = &ecn_to_tos(item.ecn.unwrap());
293                    sendmsg(fd, &iov, &[ControlMessage::Ipv4Tos(tos)], MsgFlags::empty(), Some(&addr))
294                }
295                (false, true, false) => {
296                    let tc = &(ecn_to_tos(item.ecn.unwrap()) as i32);
297                    sendmsg(fd, &iov, &[ControlMessage::Ipv6TClass(tc)], MsgFlags::empty(), Some(&addr))
298                }
299                (false, false, _) => sendmsg::<SockaddrStorage>(fd, &iov, &[], MsgFlags::empty(), Some(&addr)),
300            }
301            .map_err(nix_err)?;
302            Ok(())
303        }
304
305        const SENDMMSG_MAX: usize = 64;
306        const SENDMMSG_CMSG: usize = 32;
307
308        fn fill_sockaddr(addr: &SocketAddr, ss: &mut libc::sockaddr_storage) -> libc::socklen_t {
309            unsafe {
310                std::ptr::write_bytes(ss as *mut _ as *mut u8, 0, std::mem::size_of::<libc::sockaddr_storage>());
311            }
312            match addr {
313                SocketAddr::V4(v4) => {
314                    let sin = ss as *mut _ as *mut libc::sockaddr_in;
315                    unsafe {
316                        (*sin).sin_family = libc::AF_INET as u16;
317                        (*sin).sin_port = v4.port().to_be();
318                        (*sin).sin_addr.s_addr = u32::from_be_bytes(v4.ip().octets()).to_be();
319                    }
320                    std::mem::size_of::<libc::sockaddr_in>() as libc::socklen_t
321                }
322                SocketAddr::V6(v6) => {
323                    let sin6 = ss as *mut _ as *mut libc::sockaddr_in6;
324                    unsafe {
325                        (*sin6).sin6_family = libc::AF_INET6 as u16;
326                        (*sin6).sin6_port = v6.port().to_be();
327                        (*sin6).sin6_addr.s6_addr = v6.ip().octets();
328                        (*sin6).sin6_flowinfo = v6.flowinfo().to_be();
329                        (*sin6).sin6_scope_id = v6.scope_id();
330                    }
331                    std::mem::size_of::<libc::sockaddr_in6>() as libc::socklen_t
332                }
333            }
334        }
335
336        fn encode_ecn_cmsg(buf: &mut [u8], ecn: Option<Ecn>, is_v6: bool) -> usize {
337            let ecn = match ecn {
338                Some(e) => e,
339                None => return 0,
340            };
341            let mut mhdr: libc::msghdr = unsafe { std::mem::zeroed() };
342            mhdr.msg_control = buf.as_mut_ptr() as *mut _;
343            mhdr.msg_controllen = buf.len() as _;
344            let cmsg = unsafe { libc::CMSG_FIRSTHDR(&mhdr) };
345            if is_v6 {
346                unsafe {
347                    (*cmsg).cmsg_level = libc::IPPROTO_IPV6;
348                    (*cmsg).cmsg_type = libc::IPV6_TCLASS;
349                    (*cmsg).cmsg_len = libc::CMSG_LEN(std::mem::size_of::<i32>() as u32) as _;
350                    *(libc::CMSG_DATA(cmsg) as *mut i32) = ecn.to_tos_bits() as i32;
351                }
352                unsafe { libc::CMSG_SPACE(std::mem::size_of::<i32>() as u32) as usize }
353            } else {
354                unsafe {
355                    (*cmsg).cmsg_level = libc::IPPROTO_IP;
356                    (*cmsg).cmsg_type = libc::IP_TOS;
357                    (*cmsg).cmsg_len = libc::CMSG_LEN(std::mem::size_of::<u8>() as u32) as _;
358                    *libc::CMSG_DATA(cmsg) = ecn.to_tos_bits();
359                }
360                unsafe { libc::CMSG_SPACE(std::mem::size_of::<u8>() as u32) as usize }
361            }
362        }
363
364        fn sendmmsg_batch(fd: std::os::fd::RawFd, items: &[SendItem], caps: &Capabilities) -> io::Result<usize> {
365            let mut total_sent = 0;
366
367            while total_sent < items.len() {
368                let chunk_len = (items.len() - total_sent).min(SENDMMSG_MAX);
369                let chunk = &items[total_sent..total_sent + chunk_len];
370
371                let mut addrs: [libc::sockaddr_storage; SENDMMSG_MAX] = unsafe { std::mem::zeroed() };
372                let mut iovs: [libc::iovec; SENDMMSG_MAX] = unsafe { std::mem::zeroed() };
373                let mut cmsg_bufs: [[u8; SENDMMSG_CMSG]; SENDMMSG_MAX] = [[0u8; SENDMMSG_CMSG]; SENDMMSG_MAX];
374                let mut msgs: [libc::mmsghdr; SENDMMSG_MAX] = unsafe { std::mem::zeroed() };
375
376                for i in 0..chunk_len {
377                    let item = &chunk[i];
378                    let addr_len = fill_sockaddr(&item.destination, &mut addrs[i]);
379                    let cmsg_len = if caps.ecn {
380                        encode_ecn_cmsg(&mut cmsg_bufs[i], item.ecn, item.destination.is_ipv6())
381                    } else {
382                        0
383                    };
384
385                    iovs[i].iov_base = item.data.as_ptr() as *mut _;
386                    iovs[i].iov_len = item.data.len();
387
388                    let mhdr = &mut msgs[i].msg_hdr;
389                    mhdr.msg_name = &mut addrs[i] as *mut _ as *mut _;
390                    mhdr.msg_namelen = addr_len;
391                    mhdr.msg_iov = &mut iovs[i] as *mut _ as *mut _;
392                    mhdr.msg_iovlen = 1;
393                    mhdr.msg_control = cmsg_bufs[i].as_mut_ptr() as *mut _;
394                    mhdr.msg_controllen = cmsg_len as _;
395                    mhdr.msg_flags = 0;
396                }
397
398                let ret = unsafe { libc::sendmmsg(fd, msgs.as_mut_ptr(), chunk_len as u32, 0) };
399                if ret < 0 {
400                    let e = io::Error::last_os_error();
401                    if total_sent > 0 && e.kind() == io::ErrorKind::WouldBlock {
402                        return Ok(total_sent);
403                    }
404                    return Err(e);
405                }
406                let sent = ret as usize;
407                total_sent += sent;
408                if sent < chunk_len {
409                    break;
410                }
411            }
412
413            Ok(total_sent)
414        }
415
416        pub fn send_batch(
417            fd: std::os::fd::RawFd,
418            items: &[SendItem<'_>],
419            caps: &Capabilities,
420            gso_einval: &mut bool,
421            chunk_offset: &mut usize,
422        ) -> io::Result<usize> {
423            if items.is_empty() {
424                return Ok(0);
425            }
426
427            // Single-item GSO fast path — only attempted once
428            if items.len() == 1 && items[0].segment_size.is_some() && caps.gso && !*gso_einval {
429                match send_one(fd, &items[0], caps.ecn, true) {
430                    Ok(()) => {
431                        *chunk_offset = 0;
432                        return Ok(1);
433                    }
434                    Err(e) if e.raw_os_error() == Some(libc::EINVAL) => {
435                        *gso_einval = true;
436                        // Fall through — GSO rejected by this interface,
437                        // manually chunk below
438                    }
439                    Err(e) if e.kind() == io::ErrorKind::WouldBlock => return Err(e),
440                    Err(e) => return Err(e),
441                }
442            }
443
444            // Multi-item non-segmented batch via sendmmsg
445            if caps.sendmmsg && items.len() > 1 && !items.iter().any(|i| i.segment_size.is_some()) {
446                return sendmmsg_batch(fd, items, caps);
447            }
448
449            let mut count = 0usize;
450            for item in items {
451                if let Some(seg) = item.segment_size {
452                    // Manual chunking — GSO unavailable or rejected
453                    let seg = seg as usize;
454                    let data_slice = &item.data[*chunk_offset..];
455                    for chunk in data_slice.chunks(seg) {
456                        let mut chunk_item = SendItem::new(item.destination, chunk);
457                        if let Some(ecn) = item.ecn {
458                            chunk_item = chunk_item.ecn(ecn);
459                        }
460                        match send_one(fd, &chunk_item, caps.ecn, false) {
461                            Ok(()) => *chunk_offset += chunk.len(),
462                            Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
463                                if count == 0 && *chunk_offset == 0 {
464                                    return Err(e);
465                                }
466                                return Ok(count);
467                            }
468                            Err(e) => return Err(e),
469                        }
470                    }
471                    *chunk_offset = 0;
472                    count += 1;
473                } else {
474                    match send_one(fd, item, caps.ecn, false) {
475                        Ok(()) => count += 1,
476                        Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
477                            if count == 0 {
478                                return Err(e);
479                            }
480                            return Ok(count);
481                        }
482                        Err(e) => return Err(e),
483                    }
484                }
485            }
486            Ok(count)
487        }
488
489        const RECV_CMSG_BUF_SIZE: usize = 128;
490
491        fn recv_one(fd: std::os::fd::RawFd, item: &mut ReceiveItem<'_>, ecn_enabled: bool) -> io::Result<()> {
492            let mut iov = [std::io::IoSliceMut::new(item.buf)];
493            let mut cmsg_buf = [0u8; RECV_CMSG_BUF_SIZE];
494
495            let result =
496                recvmsg::<SockaddrStorage>(fd, &mut iov, Some(&mut cmsg_buf), MsgFlags::empty()).map_err(nix_err)?;
497
498            item.len = result.bytes;
499            item.source = result
500                .address
501                .as_ref()
502                .map(storage_to_addr)
503                .unwrap_or_else(|| SocketAddr::new(std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED), 0));
504
505            if ecn_enabled && let Ok(cmsgs) = result.cmsgs() {
506                let collected: Vec<_> = cmsgs.collect();
507                item.ecn = parse_ecn(&collected);
508            }
509
510            Ok(())
511        }
512
513        fn recv_gro(fd: std::os::fd::RawFd, items: &mut [ReceiveItem<'_>], ecn_enabled: bool) -> io::Result<usize> {
514            let mut gro_buf = [0u8; 65536];
515            let mut iov = [std::io::IoSliceMut::new(&mut gro_buf)];
516            let mut cmsg_buf = [0u8; RECV_CMSG_BUF_SIZE];
517
518            let result =
519                recvmsg::<SockaddrStorage>(fd, &mut iov, Some(&mut cmsg_buf), MsgFlags::empty()).map_err(nix_err)?;
520
521            let total = result.bytes;
522            let source = result
523                .address
524                .as_ref()
525                .map(storage_to_addr)
526                .unwrap_or_else(|| SocketAddr::new(std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED), 0));
527
528            // Single-pass cmsg parsing — always inspect GRO segment info
529            // (GRO is independent of ECN). No heap allocation.
530            let mut ecn_val: Option<Ecn> = None;
531            let mut seg_size: Option<u32> = None;
532            if let Ok(cmsgs) = result.cmsgs() {
533                for cmsg in cmsgs {
534                    match cmsg {
535                        nix::sys::socket::ControlMessageOwned::UdpGroSegments(seg) if seg_size.is_none() => {
536                            seg_size = Some(seg as u32);
537                        }
538                        nix::sys::socket::ControlMessageOwned::Ipv4Tos(tos) if ecn_enabled && ecn_val.is_none() => {
539                            ecn_val = Some(Ecn::from_tos_bits(tos));
540                        }
541                        nix::sys::socket::ControlMessageOwned::Ipv6TClass(tc) if ecn_enabled && ecn_val.is_none() => {
542                            ecn_val = Some(Ecn::from_tos_bits(tc as u8));
543                        }
544                        _ => {}
545                    }
546                }
547            }
548
549            if seg_size.is_none() {
550                if items.is_empty() {
551                    return Ok(0);
552                }
553                let copy_len = total.min(items[0].buf.len());
554                items[0].buf[..copy_len].copy_from_slice(&gro_buf[..copy_len]);
555                items[0].len = copy_len;
556                items[0].source = source;
557                items[0].ecn = ecn_val;
558                return Ok(1);
559            }
560
561            let seg = seg_size.unwrap() as usize;
562            let mut offset = 0;
563            let mut count = 0;
564            for item in items.iter_mut() {
565                if offset >= total {
566                    break;
567                }
568                let seg_end = (offset + seg).min(total);
569                let copy_len = (seg_end - offset).min(item.buf.len());
570                item.buf[..copy_len].copy_from_slice(&gro_buf[offset..offset + copy_len]);
571                item.len = copy_len;
572                item.source = source;
573                item.ecn = ecn_val;
574                count += 1;
575                offset = seg_end;
576            }
577            Ok(count)
578        }
579
580        pub fn recv_batch(
581            fd: std::os::fd::RawFd,
582            items: &mut [ReceiveItem<'_>],
583            caps: &Capabilities,
584        ) -> io::Result<usize> {
585            if items.is_empty() {
586                return Ok(0);
587            }
588
589            if caps.gro {
590                return recv_gro(fd, items, caps.ecn);
591            }
592
593            let mut count = 0;
594            for item in items.iter_mut() {
595                match recv_one(fd, item, caps.ecn) {
596                    Ok(()) => count += 1,
597                    Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
598                        if count == 0 {
599                            return Err(e);
600                        }
601                        return Ok(count);
602                    }
603                    Err(e) => return Err(e),
604                }
605            }
606            Ok(count)
607        }
608
609        fn probe_gso(fd: &OwnedFd) -> bool {
610            use nix::sys::socket::{getsockopt, sockopt::UdpGsoSegment};
611            getsockopt(fd, UdpGsoSegment).is_ok()
612        }
613
614        fn probe_gro(fd: &OwnedFd) -> bool {
615            use nix::sys::socket::{setsockopt, sockopt::UdpGroSegment};
616            setsockopt(fd, UdpGroSegment, &true).is_ok()
617        }
618
619        fn enable_ecn_recv(fd: &OwnedFd, is_ipv6: bool) -> bool {
620            use nix::sys::socket::{
621                setsockopt,
622                sockopt::{IpRecvTos, Ipv6RecvTClass},
623            };
624            if is_ipv6 {
625                setsockopt(fd, Ipv6RecvTClass, &true).is_ok()
626            } else {
627                setsockopt(fd, IpRecvTos, &true).is_ok()
628            }
629        }
630
631        pub fn setup_and_probe(fd: &OwnedFd, config: &BuildConfig, is_ipv6: bool) -> io::Result<Capabilities> {
632            use nix::sys::socket::{setsockopt, sockopt::ReuseAddr};
633
634            setsockopt(fd, ReuseAddr, &true).map_err(nix_err)?;
635
636            let gso = if config.disable_gso { false } else { probe_gso(fd) };
637            let gro = if config.disable_gro { false } else { probe_gro(fd) };
638            let sendmmsg = !config.disable_sendmmsg;
639            let ecn = if config.disable_ecn {
640                false
641            } else {
642                enable_ecn_recv(fd, is_ipv6)
643            };
644
645            Ok(Capabilities {
646                gso,
647                gro,
648                sendmmsg,
649                ecn,
650                max_batch: config.max_batch,
651            })
652        }
653
654        pub fn local_addr(fd: &OwnedFd) -> io::Result<SocketAddr> {
655            use nix::sys::socket::getsockname;
656            let addr: SockaddrStorage = getsockname(fd.as_raw_fd()).map_err(nix_err)?;
657            Ok(storage_to_addr(&addr))
658        }
659    }
660
661    // -- Generic Unix fallback -------------------------------------------------
662    #[cfg(all(unix, not(target_os = "linux")))]
663    mod unix_inner {
664        use std::{
665            io,
666            net::SocketAddr,
667            os::fd::{AsRawFd, OwnedFd},
668        };
669
670        use nix::sys::socket::{ControlMessage, MsgFlags, SockaddrStorage, recvmsg, sendmsg};
671
672        use super::{BuildConfig, ecn_to_tos, nix_err, parse_ecn, storage_to_addr};
673        use crate::{
674            capability::Capabilities,
675            item::{ReceiveItem, SendItem},
676        };
677
678        fn send_one(fd: std::os::fd::RawFd, item: &SendItem, ecn_enabled: bool) -> io::Result<()> {
679            let iov = [std::io::IoSlice::new(item.data)];
680            let addr = SockaddrStorage::from(item.destination);
681
682            let has_ecn = ecn_enabled && item.ecn.is_some();
683            let is_v4 = item.destination.is_ipv4();
684
685            match (has_ecn, is_v4) {
686                (true, true) => {
687                    let tos = &ecn_to_tos(item.ecn.unwrap());
688                    sendmsg(fd, &iov, &[ControlMessage::Ipv4Tos(tos)], MsgFlags::empty(), Some(&addr))
689                }
690                (true, false) => {
691                    let tc = &(ecn_to_tos(item.ecn.unwrap()) as i32);
692                    sendmsg(fd, &iov, &[ControlMessage::Ipv6TClass(tc)], MsgFlags::empty(), Some(&addr))
693                }
694                (false, _) => sendmsg::<SockaddrStorage>(fd, &iov, &[], MsgFlags::empty(), Some(&addr)),
695            }
696            .map_err(nix_err)?;
697            Ok(())
698        }
699
700        pub fn send_batch(fd: std::os::fd::RawFd, items: &[SendItem<'_>], caps: &Capabilities) -> io::Result<usize> {
701            if items.is_empty() {
702                return Ok(0);
703            }
704
705            let mut count = 0;
706            for item in items {
707                match send_one(fd, item, caps.ecn) {
708                    Ok(()) => count += 1,
709                    Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
710                        if count == 0 {
711                            return Err(e);
712                        }
713                        return Ok(count);
714                    }
715                    Err(e) => return Err(e),
716                }
717            }
718            Ok(count)
719        }
720
721        const RECV_CMSG_BUF_SIZE: usize = 64;
722
723        fn recv_one(fd: std::os::fd::RawFd, item: &mut ReceiveItem<'_>, ecn_enabled: bool) -> io::Result<()> {
724            let mut iov = [std::io::IoSliceMut::new(item.buf)];
725            let mut cmsg_buf = [0u8; RECV_CMSG_BUF_SIZE];
726
727            let result =
728                recvmsg::<SockaddrStorage>(fd, &mut iov, Some(&mut cmsg_buf), MsgFlags::empty()).map_err(nix_err)?;
729
730            item.len = result.bytes;
731            item.source = result
732                .address
733                .as_ref()
734                .map(storage_to_addr)
735                .unwrap_or_else(|| SocketAddr::new(std::net::IpAddr::V4(std::net::Ipv4Addr::UNSPECIFIED), 0));
736
737            if ecn_enabled {
738                if let Ok(cmsgs) = result.cmsgs() {
739                    let collected: Vec<_> = cmsgs.collect();
740                    item.ecn = parse_ecn(&collected);
741                }
742            }
743
744            Ok(())
745        }
746
747        pub fn recv_batch(
748            fd: std::os::fd::RawFd,
749            items: &mut [ReceiveItem<'_>],
750            caps: &Capabilities,
751        ) -> io::Result<usize> {
752            if items.is_empty() {
753                return Ok(0);
754            }
755
756            let mut count = 0;
757            for item in items.iter_mut() {
758                match recv_one(fd, item, caps.ecn) {
759                    Ok(()) => count += 1,
760                    Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
761                        if count == 0 {
762                            return Err(e);
763                        }
764                        return Ok(count);
765                    }
766                    Err(e) => return Err(e),
767                }
768            }
769            Ok(count)
770        }
771
772        pub fn setup_and_probe(fd: &OwnedFd, config: &BuildConfig, is_ipv6: bool) -> io::Result<Capabilities> {
773            use nix::sys::socket::{setsockopt, sockopt::ReuseAddr};
774
775            setsockopt(fd, ReuseAddr, &true).map_err(nix_err)?;
776
777            let ecn = if config.disable_ecn {
778                false
779            } else {
780                use nix::sys::socket::sockopt::{IpRecvTos, Ipv6RecvTClass};
781                if is_ipv6 {
782                    setsockopt(fd, Ipv6RecvTClass, &true).is_ok()
783                } else {
784                    setsockopt(fd, IpRecvTos, &true).is_ok()
785                }
786            };
787
788            Ok(Capabilities {
789                gso: false,
790                gro: false,
791                sendmmsg: false,
792                ecn,
793                max_batch: config.max_batch,
794            })
795        }
796
797        pub fn local_addr(fd: &OwnedFd) -> io::Result<SocketAddr> {
798            use nix::sys::socket::getsockname;
799            let addr: SockaddrStorage = getsockname(fd.as_raw_fd()).map_err(nix_err)?;
800            Ok(storage_to_addr(&addr))
801        }
802    }
803}
804
805// ---------------------------------------------------------------------------
806// Windows implementation
807// ---------------------------------------------------------------------------
808#[cfg(windows)]
809mod windows_socket_inner {
810    use std::{io, net::SocketAddr};
811
812    use tokio::net::UdpSocket;
813
814    use crate::{
815        capability::Capabilities,
816        ecn::Ecn,
817        item::{ReceiveItem, SendItem},
818    };
819
820    pub struct FastUdpSocket {
821        socket: UdpSocket,
822        caps: Capabilities,
823    }
824
825    impl FastUdpSocket {
826        pub async fn send(&self, item: SendItem<'_>) -> io::Result<()> {
827            self.socket.send_to(item.data, item.destination).await?;
828            Ok(())
829        }
830
831        pub async fn send_many(&self, items: &[SendItem<'_>]) -> io::Result<usize> {
832            if items.is_empty() {
833                return Ok(0);
834            }
835            let mut count = 0;
836            for item in items {
837                match self.socket.send_to(item.data, item.destination).await {
838                    Ok(()) => count += 1,
839                    Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
840                        if count == 0 {
841                            return Err(e);
842                        }
843                        return Ok(count);
844                    }
845                    Err(e) => return Err(e),
846                }
847            }
848            Ok(count)
849        }
850
851        pub async fn receive(&self, item: &mut ReceiveItem<'_>) -> io::Result<()> {
852            let (n, addr) = self.socket.recv_from(item.buf).await?;
853            item.len = n;
854            item.source = addr;
855            Ok(())
856        }
857
858        pub async fn receive_many(&self, items: &mut [ReceiveItem<'_>]) -> io::Result<usize> {
859            if items.is_empty() {
860                return Ok(0);
861            }
862            let mut count = 0;
863            for item in items.iter_mut() {
864                match self.socket.recv_from(item.buf).await {
865                    Ok((n, addr)) => {
866                        item.len = n;
867                        item.source = addr;
868                        count += 1;
869                    }
870                    Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
871                        if count == 0 {
872                            return Err(e);
873                        }
874                        return Ok(count);
875                    }
876                    Err(e) => return Err(e),
877                }
878            }
879            Ok(count)
880        }
881
882        pub fn capabilities(&self) -> &Capabilities {
883            &self.caps
884        }
885
886        pub fn local_addr(&self) -> io::Result<SocketAddr> {
887            self.socket.local_addr()
888        }
889    }
890
891    impl FastUdpSocket {
892        pub fn build(addr: SocketAddr) -> FastUdpSocketBuilder {
893            FastUdpSocketBuilder::new(addr)
894        }
895    }
896
897    pub struct FastUdpSocketBuilder {
898        addr: SocketAddr,
899        max_batch: usize,
900    }
901
902    impl FastUdpSocketBuilder {
903        pub fn new(addr: SocketAddr) -> Self {
904            FastUdpSocketBuilder {
905                addr,
906                max_batch: 64,
907            }
908        }
909
910        pub fn disable_gso(self) -> Self {
911            self
912        }
913
914        pub fn disable_gro(self) -> Self {
915            self
916        }
917
918        pub fn disable_sendmmsg(self) -> Self {
919            self
920        }
921
922        pub fn disable_ecn(self) -> Self {
923            self
924        }
925
926        pub fn max_batch_size(mut self, n: usize) -> Self {
927            self.max_batch = n;
928            self
929        }
930
931        pub fn enable_reuse_port(self) -> Self {
932            self
933        }
934
935        pub fn bind(self) -> io::Result<FastUdpSocket> {
936            let std_socket = std::net::UdpSocket::bind(self.addr)?;
937            let socket = UdpSocket::from_std(std_socket);
938
939            Ok(FastUdpSocket {
940                socket,
941                caps: Capabilities::none(),
942            })
943        }
944    }
945}
946
947// -- re-exports ------------------------------------------------------------
948#[cfg(unix)]
949pub use unix_socket::{FastUdpSocket, FastUdpSocketBuilder};
950#[cfg(windows)]
951pub use windows_socket_inner::{FastUdpSocket, FastUdpSocketBuilder};