1#[cfg(target_os = "linux")]
4#[allow(deprecated)]
5use nix::sys::socket::InetAddr;
6#[cfg(target_os = "linux")]
7use {
8 itertools::izip,
9 libc::{iovec, mmsghdr, sockaddr_in, sockaddr_in6, sockaddr_storage},
10 std::os::unix::io::AsRawFd,
11};
12use {
13 miraland_sdk::transport::TransportError,
14 std::{
15 borrow::Borrow,
16 io,
17 iter::repeat,
18 net::{SocketAddr, UdpSocket},
19 },
20 thiserror::Error,
21};
22
23#[derive(Debug, Error)]
24pub enum SendPktsError {
25 #[error("IO Error, some packets could not be sent")]
27 IoError(io::Error, usize),
28}
29
30impl From<SendPktsError> for TransportError {
31 fn from(err: SendPktsError) -> Self {
32 Self::Custom(format!("{err:?}"))
33 }
34}
35
36#[cfg(not(target_os = "linux"))]
37pub fn batch_send<S, T>(sock: &UdpSocket, packets: &[(T, S)]) -> Result<(), SendPktsError>
38where
39 S: Borrow<SocketAddr>,
40 T: AsRef<[u8]>,
41{
42 let mut num_failed = 0;
43 let mut erropt = None;
44 for (p, a) in packets {
45 if let Err(e) = sock.send_to(p.as_ref(), a.borrow()) {
46 num_failed += 1;
47 if erropt.is_none() {
48 erropt = Some(e);
49 }
50 }
51 }
52
53 if let Some(err) = erropt {
54 Err(SendPktsError::IoError(err, num_failed))
55 } else {
56 Ok(())
57 }
58}
59
60#[cfg(target_os = "linux")]
61fn mmsghdr_for_packet(
62 packet: &[u8],
63 dest: &SocketAddr,
64 iov: &mut iovec,
65 addr: &mut sockaddr_storage,
66 hdr: &mut mmsghdr,
67) {
68 const SIZE_OF_SOCKADDR_IN: usize = std::mem::size_of::<sockaddr_in>();
69 const SIZE_OF_SOCKADDR_IN6: usize = std::mem::size_of::<sockaddr_in6>();
70
71 *iov = iovec {
72 iov_base: packet.as_ptr() as *mut libc::c_void,
73 iov_len: packet.len(),
74 };
75 hdr.msg_hdr.msg_iov = iov;
76 hdr.msg_hdr.msg_iovlen = 1;
77 hdr.msg_hdr.msg_name = addr as *mut _ as *mut _;
78
79 #[allow(deprecated)]
80 match InetAddr::from_std(dest) {
81 InetAddr::V4(dest) => {
82 unsafe {
83 std::ptr::write(addr as *mut _ as *mut _, dest);
84 }
85 hdr.msg_hdr.msg_namelen = SIZE_OF_SOCKADDR_IN as u32;
86 }
87 InetAddr::V6(dest) => {
88 unsafe {
89 std::ptr::write(addr as *mut _ as *mut _, dest);
90 }
91 hdr.msg_hdr.msg_namelen = SIZE_OF_SOCKADDR_IN6 as u32;
92 }
93 };
94}
95
96#[cfg(target_os = "linux")]
97fn sendmmsg_retry(sock: &UdpSocket, hdrs: &mut [mmsghdr]) -> Result<(), SendPktsError> {
98 let sock_fd = sock.as_raw_fd();
99 let mut total_sent = 0;
100 let mut erropt = None;
101
102 let mut pkts = &mut *hdrs;
103 while !pkts.is_empty() {
104 let npkts = match unsafe { libc::sendmmsg(sock_fd, &mut pkts[0], pkts.len() as u32, 0) } {
105 -1 => {
106 if erropt.is_none() {
107 erropt = Some(io::Error::last_os_error());
108 }
109 1_usize
111 }
112 n => {
113 total_sent += n as usize;
116 n as usize
117 }
118 };
119 pkts = &mut pkts[npkts..];
120 }
121
122 if let Some(err) = erropt {
123 Err(SendPktsError::IoError(err, hdrs.len() - total_sent))
124 } else {
125 Ok(())
126 }
127}
128
129#[cfg(target_os = "linux")]
130pub fn batch_send<S, T>(sock: &UdpSocket, packets: &[(T, S)]) -> Result<(), SendPktsError>
131where
132 S: Borrow<SocketAddr>,
133 T: AsRef<[u8]>,
134{
135 let size = packets.len();
136 #[allow(clippy::uninit_assumed_init)]
137 let iovec = std::mem::MaybeUninit::<iovec>::uninit();
138 let mut iovs = vec![unsafe { iovec.assume_init() }; size];
139 let mut addrs = vec![unsafe { std::mem::zeroed() }; size];
140 let mut hdrs = vec![unsafe { std::mem::zeroed() }; size];
141 for ((pkt, dest), hdr, iov, addr) in izip!(packets, &mut hdrs, &mut iovs, &mut addrs) {
142 mmsghdr_for_packet(pkt.as_ref(), dest.borrow(), iov, addr, hdr);
143 }
144 sendmmsg_retry(sock, &mut hdrs)
145}
146
147pub fn multi_target_send<S, T>(
148 sock: &UdpSocket,
149 packet: T,
150 dests: &[S],
151) -> Result<(), SendPktsError>
152where
153 S: Borrow<SocketAddr>,
154 T: AsRef<[u8]>,
155{
156 let dests = dests.iter().map(Borrow::borrow);
157 let pkts: Vec<_> = repeat(&packet).zip(dests).collect();
158 batch_send(sock, &pkts)
159}
160
161#[cfg(test)]
162mod tests {
163 use {
164 crate::{
165 packet::Packet,
166 recvmmsg::recv_mmsg,
167 sendmmsg::{batch_send, multi_target_send, SendPktsError},
168 },
169 assert_matches::assert_matches,
170 miraland_sdk::packet::PACKET_DATA_SIZE,
171 std::{
172 io::ErrorKind,
173 net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, UdpSocket},
174 },
175 };
176
177 #[test]
178 pub fn test_send_mmsg_one_dest() {
179 let reader = UdpSocket::bind("127.0.0.1:0").expect("bind");
180 let addr = reader.local_addr().unwrap();
181 let sender = UdpSocket::bind("127.0.0.1:0").expect("bind");
182
183 let packets: Vec<_> = (0..32).map(|_| vec![0u8; PACKET_DATA_SIZE]).collect();
184 let packet_refs: Vec<_> = packets.iter().map(|p| (&p[..], &addr)).collect();
185
186 let sent = batch_send(&sender, &packet_refs[..]).ok();
187 assert_eq!(sent, Some(()));
188
189 let mut packets = vec![Packet::default(); 32];
190 let recv = recv_mmsg(&reader, &mut packets[..]).unwrap();
191 assert_eq!(32, recv);
192 }
193
194 #[test]
195 pub fn test_send_mmsg_multi_dest() {
196 let reader = UdpSocket::bind("127.0.0.1:0").expect("bind");
197 let addr = reader.local_addr().unwrap();
198
199 let reader2 = UdpSocket::bind("127.0.0.1:0").expect("bind");
200 let addr2 = reader2.local_addr().unwrap();
201
202 let sender = UdpSocket::bind("127.0.0.1:0").expect("bind");
203
204 let packets: Vec<_> = (0..32).map(|_| vec![0u8; PACKET_DATA_SIZE]).collect();
205 let packet_refs: Vec<_> = packets
206 .iter()
207 .enumerate()
208 .map(|(i, p)| {
209 if i < 16 {
210 (&p[..], &addr)
211 } else {
212 (&p[..], &addr2)
213 }
214 })
215 .collect();
216
217 let sent = batch_send(&sender, &packet_refs[..]).ok();
218 assert_eq!(sent, Some(()));
219
220 let mut packets = vec![Packet::default(); 32];
221 let recv = recv_mmsg(&reader, &mut packets[..]).unwrap();
222 assert_eq!(16, recv);
223
224 let mut packets = vec![Packet::default(); 32];
225 let recv = recv_mmsg(&reader2, &mut packets[..]).unwrap();
226 assert_eq!(16, recv);
227 }
228
229 #[test]
230 pub fn test_multicast_msg() {
231 let reader = UdpSocket::bind("127.0.0.1:0").expect("bind");
232 let addr = reader.local_addr().unwrap();
233
234 let reader2 = UdpSocket::bind("127.0.0.1:0").expect("bind");
235 let addr2 = reader2.local_addr().unwrap();
236
237 let reader3 = UdpSocket::bind("127.0.0.1:0").expect("bind");
238 let addr3 = reader3.local_addr().unwrap();
239
240 let reader4 = UdpSocket::bind("127.0.0.1:0").expect("bind");
241 let addr4 = reader4.local_addr().unwrap();
242
243 let sender = UdpSocket::bind("127.0.0.1:0").expect("bind");
244
245 let packet = Packet::default();
246
247 let sent = multi_target_send(
248 &sender,
249 packet.data(..).unwrap(),
250 &[&addr, &addr2, &addr3, &addr4],
251 )
252 .ok();
253 assert_eq!(sent, Some(()));
254
255 let mut packets = vec![Packet::default(); 32];
256 let recv = recv_mmsg(&reader, &mut packets[..]).unwrap();
257 assert_eq!(1, recv);
258
259 let mut packets = vec![Packet::default(); 32];
260 let recv = recv_mmsg(&reader2, &mut packets[..]).unwrap();
261 assert_eq!(1, recv);
262
263 let mut packets = vec![Packet::default(); 32];
264 let recv = recv_mmsg(&reader3, &mut packets[..]).unwrap();
265 assert_eq!(1, recv);
266
267 let mut packets = vec![Packet::default(); 32];
268 let recv = recv_mmsg(&reader4, &mut packets[..]).unwrap();
269 assert_eq!(1, recv);
270 }
271
272 #[test]
273 fn test_intermediate_failures_mismatched_bind() {
274 let packets: Vec<_> = (0..3).map(|_| vec![0u8; PACKET_DATA_SIZE]).collect();
275 let ip4 = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 8080);
276 let ip6 = SocketAddr::new(IpAddr::V6(Ipv6Addr::LOCALHOST), 8080);
277 let packet_refs: Vec<_> = vec![
278 (&packets[0][..], &ip4),
279 (&packets[1][..], &ip6),
280 (&packets[2][..], &ip4),
281 ];
282 let dest_refs: Vec<_> = vec![&ip4, &ip6, &ip4];
283
284 let sender = UdpSocket::bind("0.0.0.0:0").expect("bind");
285 let res = batch_send(&sender, &packet_refs[..]);
286 assert_matches!(res, Err(SendPktsError::IoError(_, 1)));
287 let res = multi_target_send(&sender, &packets[0], &dest_refs);
288 assert_matches!(res, Err(SendPktsError::IoError(_, 1)));
289 }
290
291 #[test]
292 fn test_intermediate_failures_unreachable_address() {
293 let packets: Vec<_> = (0..5).map(|_| vec![0u8; PACKET_DATA_SIZE]).collect();
294 let ipv4local = SocketAddr::new(IpAddr::V4(Ipv4Addr::LOCALHOST), 8080);
295 let ipv4broadcast = SocketAddr::new(IpAddr::V4(Ipv4Addr::BROADCAST), 8080);
296 let sender = UdpSocket::bind("0.0.0.0:0").expect("bind");
297
298 let packet_refs: Vec<_> = vec![
300 (&packets[0][..], &ipv4local),
301 (&packets[1][..], &ipv4broadcast),
302 (&packets[2][..], &ipv4local),
303 (&packets[3][..], &ipv4broadcast),
304 (&packets[4][..], &ipv4local),
305 ];
306 match batch_send(&sender, &packet_refs[..]) {
307 Ok(()) => panic!(),
308 Err(SendPktsError::IoError(ioerror, num_failed)) => {
309 assert_matches!(ioerror.kind(), ErrorKind::PermissionDenied);
310 assert_eq!(num_failed, 2);
311 }
312 }
313
314 let packet_refs: Vec<_> = vec![
316 (&packets[0][..], &ipv4broadcast),
317 (&packets[1][..], &ipv4local),
318 (&packets[2][..], &ipv4broadcast),
319 (&packets[3][..], &ipv4local),
320 (&packets[4][..], &ipv4broadcast),
321 ];
322 match batch_send(&sender, &packet_refs[..]) {
323 Ok(()) => panic!(),
324 Err(SendPktsError::IoError(ioerror, num_failed)) => {
325 assert_matches!(ioerror.kind(), ErrorKind::PermissionDenied);
326 assert_eq!(num_failed, 3);
327 }
328 }
329
330 let packet_refs: Vec<_> = vec![
332 (&packets[0][..], &ipv4local),
333 (&packets[1][..], &ipv4local),
334 (&packets[2][..], &ipv4broadcast),
335 (&packets[3][..], &ipv4broadcast),
336 (&packets[4][..], &ipv4local),
337 ];
338 match batch_send(&sender, &packet_refs[..]) {
339 Ok(()) => panic!(),
340 Err(SendPktsError::IoError(ioerror, num_failed)) => {
341 assert_matches!(ioerror.kind(), ErrorKind::PermissionDenied);
342 assert_eq!(num_failed, 2);
343 }
344 }
345
346 let dest_refs: Vec<_> = vec![
348 &ipv4local,
349 &ipv4broadcast,
350 &ipv4local,
351 &ipv4broadcast,
352 &ipv4local,
353 ];
354 match multi_target_send(&sender, &packets[0], &dest_refs) {
355 Ok(()) => panic!(),
356 Err(SendPktsError::IoError(ioerror, num_failed)) => {
357 assert_matches!(ioerror.kind(), ErrorKind::PermissionDenied);
358 assert_eq!(num_failed, 2);
359 }
360 }
361
362 let dest_refs: Vec<_> = vec![
364 &ipv4broadcast,
365 &ipv4local,
366 &ipv4broadcast,
367 &ipv4local,
368 &ipv4broadcast,
369 ];
370 match multi_target_send(&sender, &packets[0], &dest_refs) {
371 Ok(()) => panic!(),
372 Err(SendPktsError::IoError(ioerror, num_failed)) => {
373 assert_matches!(ioerror.kind(), ErrorKind::PermissionDenied);
374 assert_eq!(num_failed, 3);
375 }
376 }
377 }
378}