Skip to main content

miraland_streamer/
sendmmsg.rs

1//! The `sendmmsg` module provides sendmmsg() API implementation
2
3#[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    /// IO Error during send: first error, num failed packets
26    #[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                // skip over the failing packet
110                1_usize
111            }
112            n => {
113                // if we fail to send all packets we advance to the failing
114                // packet and retry in order to capture the error code
115                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(_, /*num_failed*/ 1)));
287        let res = multi_target_send(&sender, &packets[0], &dest_refs);
288        assert_matches!(res, Err(SendPktsError::IoError(_, /*num_failed*/ 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        // test intermediate failures for batch_send
299        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        // test leading and trailing failures for batch_send
315        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        // test consecutive intermediate failures for batch_send
331        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        // test intermediate failures for multi_target_send
347        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        // test leading and trailing failures for multi_target_send
363        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}