1#[cfg(unix)]
8mod unix_socket {
9 use std::{
10 io,
11 net::SocketAddr,
12 os::fd::{AsRawFd, OwnedFd},
13 };
14
15 #[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 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 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 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 #[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 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 }
439 Err(e) if e.kind() == io::ErrorKind::WouldBlock => return Err(e),
440 Err(e) => return Err(e),
441 }
442 }
443
444 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 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 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 #[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#[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#[cfg(unix)]
949pub use unix_socket::{FastUdpSocket, FastUdpSocketBuilder};
950#[cfg(windows)]
951pub use windows_socket_inner::{FastUdpSocket, FastUdpSocketBuilder};