1#![allow(clippy::useless_conversion)]
2
3use crate::udp::UdpConfig;
4use socket2::{Domain, Protocol, Socket, Type as SockType};
5use std::io;
6use std::net::IpAddr;
7use std::net::{Ipv4Addr, Ipv6Addr, SocketAddr, UdpSocket as StdUdpSocket};
8
9#[derive(Debug)]
11pub struct UdpSocket {
12 socket: Socket,
13}
14
15#[derive(Clone, Debug, Eq, PartialEq)]
17pub struct UdpRecvMeta {
18 pub bytes_read: usize,
20 pub source_addr: SocketAddr,
22 pub destination_addr: Option<IpAddr>,
24 pub interface_index: Option<u32>,
26}
27
28#[derive(Clone, Debug, Default, Eq, PartialEq)]
30pub struct UdpSendMeta {
31 pub source_addr: Option<IpAddr>,
33 pub interface_index: Option<u32>,
35}
36
37impl UdpSocket {
38 pub fn from_config(config: &UdpConfig) -> io::Result<Self> {
40 config.validate()?;
41
42 let socket = Socket::new(
43 config.socket_family.to_domain(),
44 config.socket_type.to_sock_type(),
45 Some(Protocol::UDP),
46 )?;
47
48 socket.set_nonblocking(false)?;
49
50 if let Some(flag) = config.reuseaddr {
52 socket.set_reuse_address(flag)?;
53 }
54 #[cfg(any(
55 target_os = "android",
56 target_os = "dragonfly",
57 target_os = "freebsd",
58 target_os = "fuchsia",
59 target_os = "ios",
60 target_os = "linux",
61 target_os = "macos",
62 target_os = "netbsd",
63 target_os = "openbsd",
64 target_os = "tvos",
65 target_os = "visionos",
66 target_os = "watchos"
67 ))]
68 if let Some(flag) = config.reuseport {
69 socket.set_reuse_port(flag)?;
70 }
71 if let Some(flag) = config.broadcast {
72 socket.set_broadcast(flag)?;
73 }
74 if let Some(ttl) = config.ttl {
75 socket.set_ttl_v4(ttl)?;
76 }
77 if let Some(hoplimit) = config.hoplimit {
78 socket.set_unicast_hops_v6(hoplimit)?;
79 }
80 if let Some(timeout) = config.read_timeout {
81 socket.set_read_timeout(Some(timeout))?;
82 }
83 if let Some(timeout) = config.write_timeout {
84 socket.set_write_timeout(Some(timeout))?;
85 }
86 if let Some(size) = config.recv_buffer_size {
87 socket.set_recv_buffer_size(size)?;
88 }
89 if let Some(size) = config.send_buffer_size {
90 socket.set_send_buffer_size(size)?;
91 }
92 if let Some(tos) = config.tos {
93 socket.set_tos_v4(tos)?;
94 }
95 crate::apply_tclass_v6(&socket, config.tclass_v6)?;
96 if let Some(only_v6) = config.only_v6 {
97 socket.set_only_v6(only_v6)?;
98 }
99 if let Some(on) = config.recv_pktinfo {
100 crate::udp::set_recv_pktinfo(&socket, config.socket_family, on)?;
101 }
102
103 #[cfg(any(target_os = "linux", target_os = "android", target_os = "fuchsia"))]
105 if let Some(iface) = &config.bind_device {
106 socket.bind_device(Some(iface.as_bytes()))?;
107 }
108
109 if let Some(addr) = config.bind_addr {
111 socket.bind(&addr.into())?;
112 }
113
114 Ok(Self { socket })
115 }
116
117 pub fn new(domain: Domain, sock_type: SockType) -> io::Result<Self> {
119 let socket = Socket::new(domain, sock_type, Some(Protocol::UDP))?;
120 socket.set_nonblocking(false)?;
121 Ok(Self { socket })
122 }
123
124 pub fn v4_dgram() -> io::Result<Self> {
126 Self::new(Domain::IPV4, SockType::DGRAM)
127 }
128
129 pub fn v6_dgram() -> io::Result<Self> {
131 Self::new(Domain::IPV6, SockType::DGRAM)
132 }
133
134 pub fn raw_v4() -> io::Result<Self> {
136 Self::new(Domain::IPV4, SockType::RAW)
137 }
138
139 pub fn raw_v6() -> io::Result<Self> {
141 Self::new(Domain::IPV6, SockType::RAW)
142 }
143
144 pub fn send_to(&self, buf: &[u8], target: SocketAddr) -> io::Result<usize> {
146 self.socket.send_to(buf, &target.into())
147 }
148
149 #[cfg(unix)]
154 pub fn send_msg(
155 &self,
156 buf: &[u8],
157 target: SocketAddr,
158 meta: Option<&UdpSendMeta>,
159 ) -> io::Result<usize> {
160 use nix::sys::socket::{ControlMessage, MsgFlags, SockaddrIn, SockaddrIn6, sendmsg};
161 use std::io::IoSlice;
162 use std::os::fd::AsRawFd;
163
164 let iov = [IoSlice::new(buf)];
165 let raw_fd = self.socket.as_raw_fd();
166 let packet_info_meta =
167 meta.filter(|meta| meta.source_addr.is_some() || meta.interface_index.is_some());
168
169 match target {
170 SocketAddr::V4(addr) => {
171 let sockaddr = SockaddrIn::from(addr);
172 #[cfg(any(
173 target_os = "android",
174 target_os = "linux",
175 target_os = "netbsd",
176 target_vendor = "apple"
177 ))]
178 {
179 if let Some(meta) = packet_info_meta {
180 if meta.source_addr.is_some_and(|src| !src.is_ipv4()) {
181 return Err(io::Error::new(
182 io::ErrorKind::InvalidInput,
183 "source_addr family does not match target",
184 ));
185 }
186 let mut pktinfo: libc::in_pktinfo = unsafe { std::mem::zeroed() };
189 if let Some(src) = meta.source_addr.and_then(|ip| match ip {
190 IpAddr::V4(v4) => Some(v4),
191 IpAddr::V6(_) => None,
192 }) {
193 #[cfg(target_os = "netbsd")]
194 {
195 pktinfo.ipi_addr.s_addr = u32::from_ne_bytes(src.octets());
196 }
197 #[cfg(not(target_os = "netbsd"))]
198 {
199 pktinfo.ipi_spec_dst.s_addr = u32::from_ne_bytes(src.octets());
200 }
201 }
202 if let Some(ifindex) = meta.interface_index {
203 pktinfo.ipi_ifindex = ifindex.try_into().map_err(|_| {
204 io::Error::new(
205 io::ErrorKind::InvalidInput,
206 "interface_index is out of range for this platform",
207 )
208 })?;
209 }
210 let cmsgs = [ControlMessage::Ipv4PacketInfo(&pktinfo)];
211 return sendmsg(raw_fd, &iov, &cmsgs, MsgFlags::empty(), Some(&sockaddr))
212 .map_err(|e| io::Error::from_raw_os_error(e as i32));
213 }
214 }
215 if packet_info_meta.is_some() {
216 return Err(io::Error::new(
217 io::ErrorKind::Unsupported,
218 "send_msg packet-info metadata is not supported on this platform",
219 ));
220 }
221 sendmsg(raw_fd, &iov, &[], MsgFlags::empty(), Some(&sockaddr))
222 .map_err(|e| io::Error::from_raw_os_error(e as i32))
223 }
224 SocketAddr::V6(addr) => {
225 let sockaddr = SockaddrIn6::from(addr);
226 #[cfg(any(
227 target_os = "android",
228 target_os = "freebsd",
229 target_os = "linux",
230 target_os = "netbsd",
231 target_vendor = "apple"
232 ))]
233 {
234 if let Some(meta) = packet_info_meta {
235 if meta.source_addr.is_some_and(|src| !src.is_ipv6()) {
236 return Err(io::Error::new(
237 io::ErrorKind::InvalidInput,
238 "source_addr family does not match target",
239 ));
240 }
241 let mut pktinfo: libc::in6_pktinfo = unsafe { std::mem::zeroed() };
244 if let Some(src) = meta.source_addr.and_then(|ip| match ip {
245 IpAddr::V4(_) => None,
246 IpAddr::V6(v6) => Some(v6),
247 }) {
248 pktinfo.ipi6_addr.s6_addr = src.octets();
249 }
250 if let Some(ifindex) = meta.interface_index {
251 pktinfo.ipi6_ifindex = ifindex.try_into().map_err(|_| {
252 io::Error::new(
253 io::ErrorKind::InvalidInput,
254 "interface_index is out of range for this platform",
255 )
256 })?;
257 }
258 let cmsgs = [ControlMessage::Ipv6PacketInfo(&pktinfo)];
259 return sendmsg(raw_fd, &iov, &cmsgs, MsgFlags::empty(), Some(&sockaddr))
260 .map_err(|e| io::Error::from_raw_os_error(e as i32));
261 }
262 }
263 if packet_info_meta.is_some() {
264 return Err(io::Error::new(
265 io::ErrorKind::Unsupported,
266 "send_msg packet-info metadata is not supported on this platform",
267 ));
268 }
269 sendmsg(raw_fd, &iov, &[], MsgFlags::empty(), Some(&sockaddr))
270 .map_err(|e| io::Error::from_raw_os_error(e as i32))
271 }
272 }
273 }
274
275 #[cfg(not(unix))]
277 pub fn send_msg(
278 &self,
279 _buf: &[u8],
280 _target: SocketAddr,
281 _meta: Option<&UdpSendMeta>,
282 ) -> io::Result<usize> {
283 Err(io::Error::new(
284 io::ErrorKind::Unsupported,
285 "send_msg is only supported on Unix",
286 ))
287 }
288
289 pub fn recv_from(&self, buf: &mut [u8]) -> io::Result<(usize, SocketAddr)> {
291 let buf_maybe = unsafe {
294 std::slice::from_raw_parts_mut(
295 buf.as_mut_ptr() as *mut std::mem::MaybeUninit<u8>,
296 buf.len(),
297 )
298 };
299
300 let (n, addr) = self.socket.recv_from(buf_maybe)?;
301 let addr = addr
302 .as_socket()
303 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid address format"))?;
304
305 Ok((n, addr))
306 }
307
308 #[cfg(unix)]
314 pub fn recv_msg(&self, buf: &mut [u8]) -> io::Result<UdpRecvMeta> {
315 use nix::sys::socket::{ControlMessageOwned, MsgFlags, SockaddrStorage, recvmsg};
316 use std::io::IoSliceMut;
317 use std::os::fd::AsRawFd;
318
319 let mut iov = [IoSliceMut::new(buf)];
320 #[cfg(any(
321 target_os = "android",
322 target_os = "fuchsia",
323 target_os = "linux",
324 target_vendor = "apple",
325 target_os = "netbsd"
326 ))]
327 let mut cmsgspace = nix::cmsg_space!(libc::in_pktinfo, libc::in6_pktinfo);
328 #[cfg(all(
329 not(any(
330 target_os = "android",
331 target_os = "fuchsia",
332 target_os = "linux",
333 target_vendor = "apple",
334 target_os = "netbsd"
335 )),
336 any(target_os = "freebsd", target_os = "openbsd")
337 ))]
338 let mut cmsgspace = nix::cmsg_space!(libc::in6_pktinfo);
339 #[cfg(all(
340 not(any(
341 target_os = "android",
342 target_os = "fuchsia",
343 target_os = "linux",
344 target_vendor = "apple",
345 target_os = "netbsd"
346 )),
347 not(any(target_os = "freebsd", target_os = "openbsd"))
348 ))]
349 let mut cmsgspace = nix::cmsg_space!(libc::c_int);
350 let msg = recvmsg::<SockaddrStorage>(
351 self.socket.as_raw_fd(),
352 &mut iov,
353 Some(&mut cmsgspace),
354 MsgFlags::empty(),
355 )
356 .map_err(|e| io::Error::from_raw_os_error(e as i32))?;
357
358 let source_addr = msg
359 .address
360 .and_then(|addr: SockaddrStorage| {
361 if let Some(v4) = addr.as_sockaddr_in() {
362 return Some(SocketAddr::from(*v4));
363 }
364 if let Some(v6) = addr.as_sockaddr_in6() {
365 return Some(SocketAddr::from(*v6));
366 }
367 None
368 })
369 .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "invalid source address"))?;
370
371 let mut destination_addr = None;
372 let mut interface_index = None;
373
374 if let Ok(cmsgs) = msg.cmsgs() {
375 for cmsg in cmsgs {
376 match cmsg {
377 #[cfg(any(
378 target_os = "android",
379 target_os = "fuchsia",
380 target_os = "linux",
381 target_vendor = "apple",
382 target_os = "netbsd"
383 ))]
384 ControlMessageOwned::Ipv4PacketInfo(info) => {
385 destination_addr = Some(IpAddr::V4(std::net::Ipv4Addr::from(
386 info.ipi_addr.s_addr.to_ne_bytes(),
387 )));
388 interface_index = Some(info.ipi_ifindex.try_into().map_err(|_| {
389 io::Error::new(
390 io::ErrorKind::InvalidData,
391 "received invalid interface index",
392 )
393 })?);
394 }
395 #[cfg(any(
396 target_os = "android",
397 target_os = "freebsd",
398 target_os = "linux",
399 target_os = "macos",
400 target_os = "ios",
401 target_os = "tvos",
402 target_os = "visionos",
403 target_os = "watchos",
404 target_os = "netbsd",
405 target_os = "openbsd"
406 ))]
407 ControlMessageOwned::Ipv6PacketInfo(info) => {
408 destination_addr =
409 Some(IpAddr::V6(std::net::Ipv6Addr::from(info.ipi6_addr.s6_addr)));
410 interface_index = Some(info.ipi6_ifindex.try_into().map_err(|_| {
411 io::Error::new(
412 io::ErrorKind::InvalidData,
413 "received invalid interface index",
414 )
415 })?);
416 }
417 _ => {}
418 }
419 }
420 }
421
422 Ok(UdpRecvMeta {
423 bytes_read: msg.bytes,
424 source_addr,
425 destination_addr,
426 interface_index,
427 })
428 }
429
430 #[cfg(not(unix))]
432 pub fn recv_msg(&self, _buf: &mut [u8]) -> io::Result<UdpRecvMeta> {
433 Err(io::Error::new(
434 io::ErrorKind::Unsupported,
435 "recv_msg is only supported on Unix",
436 ))
437 }
438
439 pub fn set_ttl(&self, ttl: u32) -> io::Result<()> {
440 self.socket.set_ttl_v4(ttl)
441 }
442
443 pub fn ttl(&self) -> io::Result<u32> {
444 self.socket.ttl_v4()
445 }
446
447 pub fn set_hoplimit(&self, hops: u32) -> io::Result<()> {
448 self.socket.set_unicast_hops_v6(hops)
449 }
450
451 pub fn hoplimit(&self) -> io::Result<u32> {
452 self.socket.unicast_hops_v6()
453 }
454
455 pub fn set_reuseaddr(&self, on: bool) -> io::Result<()> {
456 self.socket.set_reuse_address(on)
457 }
458
459 pub fn reuseaddr(&self) -> io::Result<bool> {
460 self.socket.reuse_address()
461 }
462
463 #[cfg(any(
464 target_os = "android",
465 target_os = "dragonfly",
466 target_os = "freebsd",
467 target_os = "fuchsia",
468 target_os = "ios",
469 target_os = "linux",
470 target_os = "macos",
471 target_os = "netbsd",
472 target_os = "openbsd",
473 target_os = "tvos",
474 target_os = "visionos",
475 target_os = "watchos"
476 ))]
477 pub fn set_reuseport(&self, on: bool) -> io::Result<()> {
478 self.socket.set_reuse_port(on)
479 }
480
481 #[cfg(any(
482 target_os = "android",
483 target_os = "dragonfly",
484 target_os = "freebsd",
485 target_os = "fuchsia",
486 target_os = "ios",
487 target_os = "linux",
488 target_os = "macos",
489 target_os = "netbsd",
490 target_os = "openbsd",
491 target_os = "tvos",
492 target_os = "visionos",
493 target_os = "watchos"
494 ))]
495 pub fn reuseport(&self) -> io::Result<bool> {
496 self.socket.reuse_port()
497 }
498
499 pub fn set_broadcast(&self, on: bool) -> io::Result<()> {
500 self.socket.set_broadcast(on)
501 }
502
503 pub fn broadcast(&self) -> io::Result<bool> {
504 self.socket.broadcast()
505 }
506
507 pub fn join_multicast_v4(&self, group: &Ipv4Addr, interface: &Ipv4Addr) -> io::Result<()> {
509 self.socket.join_multicast_v4(group, interface)
510 }
511
512 pub fn leave_multicast_v4(&self, group: &Ipv4Addr, interface: &Ipv4Addr) -> io::Result<()> {
514 self.socket.leave_multicast_v4(group, interface)
515 }
516
517 pub fn join_multicast_v6(&self, group: &Ipv6Addr, interface: u32) -> io::Result<()> {
519 self.socket.join_multicast_v6(group, interface)
520 }
521
522 pub fn leave_multicast_v6(&self, group: &Ipv6Addr, interface: u32) -> io::Result<()> {
524 self.socket.leave_multicast_v6(group, interface)
525 }
526
527 pub fn set_recv_buffer_size(&self, size: usize) -> io::Result<()> {
528 self.socket.set_recv_buffer_size(size)
529 }
530
531 pub fn recv_buffer_size(&self) -> io::Result<usize> {
532 self.socket.recv_buffer_size()
533 }
534
535 pub fn set_send_buffer_size(&self, size: usize) -> io::Result<()> {
536 self.socket.set_send_buffer_size(size)
537 }
538
539 pub fn send_buffer_size(&self) -> io::Result<usize> {
540 self.socket.send_buffer_size()
541 }
542
543 pub fn set_tos(&self, tos: u32) -> io::Result<()> {
544 self.socket.set_tos_v4(tos)
545 }
546
547 pub fn tos(&self) -> io::Result<u32> {
548 self.socket.tos_v4()
549 }
550
551 #[cfg(any(
552 target_os = "android",
553 target_os = "dragonfly",
554 target_os = "freebsd",
555 target_os = "fuchsia",
556 target_os = "linux",
557 target_os = "macos",
558 target_os = "netbsd",
559 target_os = "openbsd"
560 ))]
561 pub fn set_tclass_v6(&self, tclass: u32) -> io::Result<()> {
562 self.socket.set_tclass_v6(tclass)
563 }
564
565 #[cfg(any(
566 target_os = "android",
567 target_os = "dragonfly",
568 target_os = "freebsd",
569 target_os = "fuchsia",
570 target_os = "linux",
571 target_os = "macos",
572 target_os = "netbsd",
573 target_os = "openbsd"
574 ))]
575 pub fn tclass_v6(&self) -> io::Result<u32> {
576 self.socket.tclass_v6()
577 }
578
579 pub fn set_only_v6(&self, only_v6: bool) -> io::Result<()> {
580 self.socket.set_only_v6(only_v6)
581 }
582
583 pub fn only_v6(&self) -> io::Result<bool> {
584 self.socket.only_v6()
585 }
586
587 pub fn set_keepalive(&self, on: bool) -> io::Result<()> {
588 self.socket.set_keepalive(on)
589 }
590
591 pub fn keepalive(&self) -> io::Result<bool> {
592 self.socket.keepalive()
593 }
594
595 pub fn set_recv_pktinfo_v4(&self, on: bool) -> io::Result<()> {
597 crate::udp::set_recv_pktinfo_v4(&self.socket, on)
598 }
599
600 pub fn set_recv_pktinfo_v6(&self, on: bool) -> io::Result<()> {
602 crate::udp::set_recv_pktinfo_v6(&self.socket, on)
603 }
604
605 pub fn recv_pktinfo_v4(&self) -> io::Result<bool> {
607 crate::udp::recv_pktinfo_v4(&self.socket)
608 }
609
610 pub fn recv_pktinfo_v6(&self) -> io::Result<bool> {
612 crate::udp::recv_pktinfo_v6(&self.socket)
613 }
614
615 pub fn local_addr(&self) -> io::Result<SocketAddr> {
617 self.socket
618 .local_addr()?
619 .as_socket()
620 .ok_or_else(|| io::Error::other("failed to retrieve local address"))
621 }
622
623 pub fn to_std(self) -> io::Result<StdUdpSocket> {
625 Ok(self.socket.into())
626 }
627
628 pub fn from_socket(socket: Socket) -> Self {
630 Self { socket }
631 }
632
633 pub fn socket(&self) -> &Socket {
635 &self.socket
636 }
637
638 pub fn into_socket(self) -> Socket {
640 self.socket
641 }
642
643 #[cfg(unix)]
644 pub fn as_raw_fd(&self) -> std::os::unix::io::RawFd {
645 use std::os::fd::AsRawFd;
646 self.socket.as_raw_fd()
647 }
648
649 #[cfg(windows)]
650 pub fn as_raw_socket(&self) -> std::os::windows::io::RawSocket {
651 use std::os::windows::io::AsRawSocket;
652 self.socket.as_raw_socket()
653 }
654}
655
656#[cfg(test)]
657mod tests {
658 use super::*;
659
660 #[test]
661 fn create_v4_socket() {
662 let sock = UdpSocket::v4_dgram().expect("create socket");
663 sock.socket
664 .bind(&"0.0.0.0:0".parse::<SocketAddr>().unwrap().into())
665 .expect("bind");
666 let addr = sock.local_addr().expect("addr");
667 assert!(addr.is_ipv4());
668 }
669
670 #[test]
671 fn v4_socket_options_and_family_mismatch() {
672 let sock = UdpSocket::v4_dgram().expect("create socket");
673 sock.socket
674 .bind(&"127.0.0.1:0".parse::<SocketAddr>().unwrap().into())
675 .expect("bind");
676
677 sock.set_ttl(37).expect("set ttl");
678 assert_eq!(sock.ttl().expect("ttl"), 37);
679 sock.set_broadcast(true).expect("set broadcast");
680 assert!(sock.broadcast().expect("broadcast"));
681
682 let mismatch = sock.send_to(&[], "[::1]:9".parse().unwrap());
683 assert!(mismatch.is_err(), "IPv6 target must fail on IPv4 socket");
684 }
685
686 #[test]
687 fn v4_multicast_membership_round_trip() {
688 let sock = UdpSocket::v4_dgram().expect("create socket");
689 sock.socket
690 .bind(&"0.0.0.0:0".parse::<SocketAddr>().unwrap().into())
691 .expect("bind");
692 let group = Ipv4Addr::new(239, 255, 0, 1);
693
694 sock.join_multicast_v4(&group, &Ipv4Addr::UNSPECIFIED)
695 .expect("join multicast");
696 sock.leave_multicast_v4(&group, &Ipv4Addr::UNSPECIFIED)
697 .expect("leave multicast");
698 }
699
700 #[test]
701 fn v6_hop_limit_round_trip() {
702 let sock = UdpSocket::v6_dgram().expect("create socket");
703 sock.socket
704 .bind(&"[::1]:0".parse::<SocketAddr>().unwrap().into())
705 .expect("bind");
706
707 sock.set_hoplimit(23).expect("set hop limit");
708 assert_eq!(sock.hoplimit().expect("hop limit"), 23);
709 }
710}