Skip to main content

datagram_socket/
mmsg.rs

1// Copyright (C) 2025, Cloudflare, Inc.
2// All rights reserved.
3//
4// Redistribution and use in source and binary forms, with or without
5// modification, are permitted provided that the following conditions are
6// met:
7//
8//     * Redistributions of source code must retain the above copyright notice,
9//       this list of conditions and the following disclaimer.
10//
11//     * Redistributions in binary form must reproduce the above copyright
12//       notice, this list of conditions and the following disclaimer in the
13//       documentation and/or other materials provided with the distribution.
14//
15// THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS
16// IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO,
17// THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR
18// PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR
19// CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL,
20// EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO,
21// PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR
22// PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF
23// LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING
24// NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS
25// SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
26
27use std::io::IoSlice;
28use std::io::{
29    self,
30};
31use std::os::fd::AsRawFd;
32use std::os::fd::BorrowedFd;
33
34use smallvec::SmallVec;
35use tokio::io::ReadBuf;
36
37/// Maximum number of messages passed to a single `sendmmsg(2)` or
38/// `recvmmsg(2)` call.
39pub const MAX_MMSG: usize = 16;
40
41pub fn recvmmsg(fd: BorrowedFd, bufs: &mut [ReadBuf<'_>]) -> io::Result<usize> {
42    let mut msgvec: SmallVec<[libc::mmsghdr; MAX_MMSG]> = SmallVec::new();
43    let mut slices: SmallVec<[IoSlice; MAX_MMSG]> = SmallVec::new();
44
45    let mut ret = 0;
46
47    for bufs in bufs.chunks_mut(MAX_MMSG) {
48        msgvec.clear();
49        slices.clear();
50
51        for buf in bufs.iter_mut() {
52            // Safety: will not read the maybe uninitialized bytes.
53            let b = unsafe {
54                &mut *(buf.unfilled_mut() as *mut [std::mem::MaybeUninit<u8>]
55                    as *mut [u8])
56            };
57
58            slices.push(IoSlice::new(b));
59
60            msgvec.push(libc::mmsghdr {
61                msg_hdr: libc::msghdr {
62                    msg_name: std::ptr::null_mut(),
63                    msg_namelen: 0,
64                    msg_iov: slices.last_mut().unwrap() as *mut _ as *mut _,
65                    msg_iovlen: 1,
66                    msg_control: std::ptr::null_mut(),
67                    msg_controllen: 0,
68                    msg_flags: 0,
69                },
70                msg_len: buf.capacity().try_into().unwrap(),
71            });
72        }
73
74        // SAFETY: `slices` and `msgvec` are `SmallVec`s with inline capacity
75        // `MAX_MMSG`, and each chunk has at most `MAX_MMSG` elements, so the
76        // pushes above cannot trigger a reallocation that would invalidate
77        // the pointers into `slices` taken by `msg_iov` above. Neither
78        // vector is modified before the syscall returns, and `fd` remains
79        // valid for the duration of the call.
80        let result = unsafe {
81            libc::recvmmsg(
82                fd.as_raw_fd(),
83                msgvec.as_mut_ptr(),
84                msgvec.len() as _,
85                0,
86                std::ptr::null_mut(),
87            )
88        };
89
90        if result == -1 {
91            break;
92        }
93
94        for i in 0..result as usize {
95            let filled = msgvec[i].msg_len as usize;
96            unsafe { bufs[i].assume_init(filled) };
97            bufs[i].advance(filled);
98            ret += 1;
99        }
100
101        if (result as usize) < MAX_MMSG {
102            break;
103        }
104    }
105
106    if ret == 0 {
107        return Err(io::Error::last_os_error());
108    }
109
110    Ok(ret)
111}
112
113/// Sends multiple datagrams with a single system call where possible.
114///
115/// The returned value is the number of datagrams sent.
116pub fn sendmmsg(fd: BorrowedFd, bufs: &[ReadBuf<'_>]) -> io::Result<usize> {
117    sendmmsg_impl(fd, bufs, None)
118}
119
120/// Sends multiple datagrams, appending the same suffix to every datagram.
121///
122/// This is useful for situations where the datagrams carry UDP payloads, which
123/// can create an ambiguous situation: an empty UDP payload is valid in a UDP IP
124/// packet, but reading 0 bytes from a datagram socket is also a signal of a
125/// closed socket. So the suffix can add a marker to every payload to tell those
126/// cases apart, and this allows to send that efficiently.
127///
128/// The returned value is the number of datagrams sent. The suffix is framing
129/// supplied on the caller's behalf and does not affect that count.
130pub fn sendmmsg_with_suffix(
131    fd: BorrowedFd, bufs: &[ReadBuf<'_>], suffix: &[u8],
132) -> io::Result<usize> {
133    sendmmsg_impl(fd, bufs, Some(suffix))
134}
135
136fn sendmmsg_impl(
137    fd: BorrowedFd, bufs: &[ReadBuf<'_>], suffix: Option<&[u8]>,
138) -> io::Result<usize> {
139    if bufs.is_empty() {
140        return Ok(0);
141    }
142
143    let mut msgvec: SmallVec<[libc::mmsghdr; MAX_MMSG]> = SmallVec::new();
144    let mut iovecs: SmallVec<[libc::iovec; 2 * MAX_MMSG]> = SmallVec::new();
145
146    let mut ret = 0;
147
148    for bufs in bufs.chunks(MAX_MMSG) {
149        msgvec.clear();
150        iovecs.clear();
151
152        for buf in bufs {
153            iovecs.push(iovec(buf.filled()));
154
155            if let Some(suffix) = suffix {
156                iovecs.push(iovec(suffix));
157            }
158        }
159
160        // Populate all iovecs before taking pointers into the vector. This
161        // ensures none of the pointers can be invalidated by a reallocation.
162        let iovecs_per_message = if suffix.is_some() { 2 } else { 1 };
163        for message_iovecs in iovecs.chunks_exact_mut(iovecs_per_message) {
164            msgvec.push(libc::mmsghdr {
165                msg_hdr: libc::msghdr {
166                    msg_name: std::ptr::null_mut(),
167                    msg_namelen: 0,
168                    msg_iov: message_iovecs.as_mut_ptr(),
169                    msg_iovlen: iovecs_per_message as _,
170                    msg_control: std::ptr::null_mut(),
171                    msg_controllen: 0,
172                    msg_flags: 0,
173                },
174                // Output field populated by the kernel.
175                msg_len: 0,
176            });
177        }
178
179        // SAFETY: `iovecs` was fully populated before the pointers in
180        // `msgvec` were created, and neither vector is modified before the
181        // syscall returns. Each header points to one or two live iovecs, and
182        // `fd` remains valid for the duration of the call.
183        let result = unsafe {
184            libc::sendmmsg(
185                fd.as_raw_fd(),
186                msgvec.as_mut_ptr(),
187                msgvec.len() as _,
188                0,
189            )
190        };
191
192        if result == -1 {
193            let err = io::Error::last_os_error();
194
195            if ret == 0 {
196                return Err(err);
197            }
198
199            break;
200        }
201
202        ret += result as usize;
203
204        if (result as usize) < bufs.len() {
205            break;
206        }
207    }
208
209    Ok(ret)
210}
211
212fn iovec(buf: &[u8]) -> libc::iovec {
213    libc::iovec {
214        // `sendmmsg(2)` does not mutate the memory described by an iovec, but
215        // the C API represents the pointer as mutable.
216        iov_base: buf.as_ptr().cast_mut().cast(),
217        iov_len: buf.len(),
218    }
219}
220
221#[macro_export]
222macro_rules! poll_recvmmsg {
223    ($self: expr, $cx: ident, $bufs: ident) => {
224        loop {
225            match $self.poll_recv_ready($cx)? {
226                Poll::Ready(()) => {
227                    match $self.try_io(tokio::io::Interest::READABLE, || {
228                        $crate::mmsg::recvmmsg($self.as_fd(), $bufs)
229                    }) {
230                        Err(err) if err.kind() == io::ErrorKind::WouldBlock => {}  // Have to poll for recv ready
231                        res => break Poll::Ready(res),
232                    }
233                }
234                Poll::Pending => break Poll::Pending,
235            }
236        }
237    };
238}
239
240#[macro_export]
241macro_rules! poll_sendmmsg {
242    ($self: expr, $cx: ident, $bufs: ident) => {
243        loop {
244            match $self.poll_send_ready($cx)? {
245                Poll::Ready(()) => {
246                    match $self.try_io(tokio::io::Interest::WRITABLE, || {
247                        $crate::mmsg::sendmmsg($self.as_fd(), $bufs)
248                    }) {
249                        Err(err) if err.kind() == io::ErrorKind::WouldBlock => {} // Have to poll for send ready
250                        res => break Poll::Ready(res),
251                    }
252                }
253                Poll::Pending => break Poll::Pending,
254            }
255        }
256    };
257}
258
259#[cfg(test)]
260mod tests {
261    use std::io;
262    use std::os::fd::AsFd;
263
264    use tokio::io::ReadBuf;
265    use tokio::net::UnixDatagram;
266
267    use super::sendmmsg;
268    use super::sendmmsg_with_suffix;
269    use super::MAX_MMSG;
270    use crate::DatagramSocketRecvExt;
271    use crate::DatagramSocketSendExt;
272
273    #[tokio::test]
274    async fn recvmmsg() -> io::Result<()> {
275        let (s, mut r) = UnixDatagram::pair()?;
276        let mut bufs = [[0u8; 128]; 128];
277
278        for i in 0..5 {
279            s.send(&[i; 128]).await?;
280        }
281
282        let mut rbufs: Vec<_> =
283            bufs.iter_mut().map(|s| ReadBuf::new(&mut s[..])).collect();
284        assert_eq!(r.recv_many(&mut rbufs).await?, 5);
285
286        for (i, buf) in rbufs[0..5].iter().enumerate() {
287            assert_eq!(buf.filled(), &[i as u8; 128]);
288        }
289
290        for i in 0..92 {
291            s.send(&[i; 128]).await?;
292        }
293
294        let mut rbufs: Vec<_> =
295            bufs.iter_mut().map(|s| ReadBuf::new(&mut s[..])).collect();
296        assert_eq!(r.recv_many(&mut rbufs).await?, 92);
297
298        for (i, buf) in rbufs[0..92].iter().enumerate() {
299            assert_eq!(buf.filled(), &[i as u8; 128]);
300        }
301
302        Ok(())
303    }
304
305    #[tokio::test]
306    async fn send_many() -> io::Result<()> {
307        let (s, r) = UnixDatagram::pair()?;
308        let mut bufs: [_; 128] = std::array::from_fn(|i| [i as u8; 128]);
309
310        let wbufs: Vec<_> = bufs
311            .iter_mut()
312            .map(|s| {
313                let mut b = ReadBuf::new(&mut s[..]);
314                b.set_filled(128);
315                b
316            })
317            .collect();
318
319        assert_eq!(s.send_many(&wbufs[..5]).await?, 5);
320
321        let mut rbuf = [0u8; 128];
322
323        for i in 0..5 {
324            assert_eq!(r.recv(&mut rbuf).await?, 128);
325            assert_eq!(rbuf, [i as u8; 128]);
326        }
327
328        Ok(())
329    }
330
331    #[tokio::test]
332    async fn sendmmsg_with_suffix_appends_to_every_datagram() -> io::Result<()> {
333        let (s, r) = UnixDatagram::pair()?;
334        let suffix = b"-suffix";
335        let mut payloads: Vec<Vec<u8>> =
336            (0..MAX_MMSG + 4).map(|i| vec![i as u8; i]).collect();
337        let bufs: Vec<_> = payloads
338            .iter_mut()
339            .map(|payload| {
340                let len = payload.len();
341                let mut buf = ReadBuf::new(payload);
342                buf.set_filled(len);
343                buf
344            })
345            .collect();
346
347        assert_eq!(
348            sendmmsg_with_suffix(s.as_fd(), &bufs, suffix)?,
349            payloads.len()
350        );
351
352        let mut received = vec![0; MAX_MMSG + suffix.len() + 4];
353        for expected in &payloads {
354            let received_len = r.recv(&mut received).await?;
355            assert_eq!(
356                &received[..received_len],
357                [expected.as_slice(), suffix].concat()
358            );
359        }
360
361        Ok(())
362    }
363
364    #[test]
365    fn empty_send_batches_are_noops() -> io::Result<()> {
366        let (s, _r) = std::os::unix::net::UnixDatagram::pair()?;
367
368        assert_eq!(sendmmsg(s.as_fd(), &[])?, 0);
369        assert_eq!(sendmmsg_with_suffix(s.as_fd(), &[], b"suffix")?, 0);
370
371        Ok(())
372    }
373
374    #[test]
375    fn sendmmsg_with_suffix_reports_an_error_without_progress() -> io::Result<()>
376    {
377        let (s, r) = std::os::unix::net::UnixDatagram::pair()?;
378        drop(r);
379
380        let mut payload = *b"payload";
381        let payload_len = payload.len();
382        let mut buf = ReadBuf::new(&mut payload);
383        buf.set_filled(payload_len);
384
385        assert!(sendmmsg_with_suffix(s.as_fd(), &[buf], b"suffix").is_err());
386
387        Ok(())
388    }
389}