1use 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
37pub 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 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 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
113pub fn sendmmsg(fd: BorrowedFd, bufs: &[ReadBuf<'_>]) -> io::Result<usize> {
117 sendmmsg_impl(fd, bufs, None)
118}
119
120pub 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 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 msg_len: 0,
176 });
177 }
178
179 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 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 => {} 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 => {} 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}