1use std::{
26 collections::VecDeque,
27 io,
28 io::IoSliceMut,
29 net::SocketAddr,
30 time::{Duration, Instant},
31};
32
33use ana_gotatun::packet::{Packet, PacketBufPool};
34use quinn_udp::{RecvMeta, Transmit, UdpSockRef, UdpSocketState};
35use tokio::{io::Interest, net::UdpSocket};
36
37const MAX_BATCH_SIZE: usize = 64;
38
39const MAX_UDP_PAYLOAD_SIZE: usize = u16::MAX as usize - 20 - 8;
45
46const DISCARD_LOG_INTERVAL: Duration = Duration::from_secs(60);
51
52#[derive(Debug, Clone, Copy, PartialEq, Eq)]
58pub struct SocketNotWritable;
59
60impl std::fmt::Display for SocketNotWritable {
61 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
62 f.write_str("socket not writable, queued datagrams were left in place")
63 }
64}
65
66impl std::error::Error for SocketNotWritable {}
67
68pub enum RecvBatchError<E> {
70 Io(io::Error),
72 Handler(E),
74}
75
76#[derive(Debug)]
78pub enum QueuePacketError {
79 Full {
81 packet: Packet,
83 target: SocketAddr,
85 },
86 PacketTooLarge {
88 packet: Packet,
90 target: SocketAddr,
92 packet_len: usize,
94 max_packet_size: usize,
96 },
97}
98
99pub struct UdpBatchReceiver<const BATCH_SIZE: usize, const BUFFER_SIZE: usize = 4096> {
109 state: UdpSocketState,
110 recv_meta: [RecvMeta; BATCH_SIZE],
111 recv_slots: [Packet; BATCH_SIZE],
112}
113
114impl<const BATCH_SIZE: usize, const BUFFER_SIZE: usize> UdpBatchReceiver<BATCH_SIZE, BUFFER_SIZE> {
115 pub fn new(socket: &UdpSocket, pool: &PacketBufPool<BUFFER_SIZE>) -> io::Result<Self> {
121 assert!(
122 BATCH_SIZE > 0,
123 "UdpBatchReceiver BATCH_SIZE must be greater than zero"
124 );
125 assert!(
126 BATCH_SIZE <= MAX_BATCH_SIZE,
127 "UdpBatchReceiver BATCH_SIZE must not exceed MAX_BATCH_SIZE"
128 );
129 let state = UdpSocketState::new(UdpSockRef::from(socket))?;
130 let recv_slots = std::array::from_fn(|_| pool.get());
131 Ok(Self {
132 state,
133 recv_meta: std::array::from_fn(|_| RecvMeta::default()),
134 recv_slots,
135 })
136 }
137
138 pub async fn recv_batch<E, F>(
140 &mut self,
141 socket: &UdpSocket,
142 pool: &PacketBufPool<BUFFER_SIZE>,
143 mut handler: F,
144 ) -> Result<(), RecvBatchError<E>>
145 where
146 F: FnMut(Packet, SocketAddr) -> Result<(), E>,
147 {
148 let received = loop {
149 socket.readable().await.map_err(RecvBatchError::Io)?;
150 match socket.try_io(Interest::READABLE, || self.try_recv(socket)) {
151 Ok(count) => break count,
152 Err(err) if err.kind() == io::ErrorKind::WouldBlock => continue,
153 Err(err) => return Err(RecvBatchError::Io(err)),
154 }
155 };
156
157 for index in 0..received {
158 self.handle_received(index, pool, &mut handler)
159 .map_err(RecvBatchError::Handler)?;
160 }
161
162 Ok(())
163 }
164
165 fn handle_received<E, F>(
166 &mut self,
167 index: usize,
168 pool: &PacketBufPool<BUFFER_SIZE>,
169 handler: &mut F,
170 ) -> Result<(), E>
171 where
172 F: FnMut(Packet, SocketAddr) -> Result<(), E>,
173 {
174 let meta = self.recv_meta[index];
178 if meta.len == 0 {
179 return Ok(());
180 }
181 let stride = if meta.stride == 0 {
182 meta.len
183 } else {
184 meta.stride
185 };
186 if stride >= meta.len {
187 let mut packet = std::mem::replace(&mut self.recv_slots[index], pool.get());
190 packet.truncate(meta.len);
191 handler(packet, meta.addr)?;
192 return Ok(());
193 }
194
195 let packet = std::mem::replace(&mut self.recv_slots[index], pool.get());
199 for chunk in packet[..meta.len].chunks(stride) {
200 let mut segment = pool.get();
201 segment[..chunk.len()].copy_from_slice(chunk);
202 segment.truncate(chunk.len());
203 handler(segment, meta.addr)?;
204 }
205 Ok(())
206 }
207
208 fn try_recv(&mut self, socket: &UdpSocket) -> io::Result<usize> {
209 let mut bufs_uninit: [std::mem::MaybeUninit<IoSliceMut<'_>>; BATCH_SIZE] =
213 std::array::from_fn(|_| std::mem::MaybeUninit::uninit());
214 for (index, packet) in self.recv_slots.iter_mut().enumerate() {
215 bufs_uninit[index].write(IoSliceMut::new(packet.as_mut()));
216 }
217 let bufs = unsafe {
224 std::slice::from_raw_parts_mut(
225 bufs_uninit.as_mut_ptr() as *mut IoSliceMut<'_>,
226 BATCH_SIZE,
227 )
228 };
229 self.state
230 .recv(UdpSockRef::from(socket), bufs, &mut self.recv_meta)
231 }
232}
233
234pub struct UdpBatchSender<const BATCH_SIZE: usize, const MAX_PACKET_SIZE: usize = 4096> {
242 state: UdpSocketState,
243 queued_packets: VecDeque<(SocketAddr, Packet)>,
244 scratch: Vec<u8>,
245 discarded_datagrams: u64,
247 last_discard_log: Option<Instant>,
249 discards_since_log: u64,
251}
252
253impl<const BATCH_SIZE: usize, const MAX_PACKET_SIZE: usize>
254 UdpBatchSender<BATCH_SIZE, MAX_PACKET_SIZE>
255{
256 pub fn new(socket: &UdpSocket) -> io::Result<Self> {
261 assert!(
262 BATCH_SIZE > 0,
263 "UdpBatchSender BATCH_SIZE must be greater than zero"
264 );
265 assert!(
266 BATCH_SIZE <= MAX_BATCH_SIZE,
267 "UdpBatchSender BATCH_SIZE must not exceed MAX_BATCH_SIZE"
268 );
269 Ok(Self {
270 state: UdpSocketState::new(UdpSockRef::from(socket))?,
271 queued_packets: VecDeque::with_capacity(BATCH_SIZE),
272 scratch: Vec::with_capacity(MAX_PACKET_SIZE * BATCH_SIZE),
273 discarded_datagrams: 0,
274 last_discard_log: None,
275 discards_since_log: 0,
276 })
277 }
278
279 pub fn is_empty(&self) -> bool {
281 self.queued_packets.is_empty()
282 }
283
284 pub fn is_full(&self) -> bool {
286 self.queued_packets.len() == BATCH_SIZE
287 }
288
289 pub fn try_queue_packet(
297 &mut self,
298 packet: Packet,
299 target: SocketAddr,
300 ) -> Result<(), QueuePacketError> {
301 let packet_len = packet.len();
302 if packet.len() > MAX_PACKET_SIZE {
303 return Err(QueuePacketError::PacketTooLarge {
304 packet,
305 target,
306 packet_len,
307 max_packet_size: MAX_PACKET_SIZE,
308 });
309 }
310 if self.is_full() {
311 return Err(QueuePacketError::Full { packet, target });
312 }
313 self.queued_packets.push_back((target, packet));
314 Ok(())
315 }
316
317 pub fn try_flush_best_effort(&mut self, socket: &UdpSocket) -> Result<(), SocketNotWritable> {
331 let mut coalesce = true;
335
336 while !self.is_empty() {
337 let (target, segment_size, segments) = self.fill_scratch_from_front(coalesce);
338 let result = socket.try_io(Interest::WRITABLE, || {
339 let transmit = Transmit {
340 destination: target,
341 ecn: None,
342 contents: &self.scratch,
343 segment_size: (segments > 1).then_some(segment_size),
344 src_ip: None,
345 };
346 self.state.try_send(UdpSockRef::from(socket), &transmit)
347 });
348
349 match result {
350 Ok(()) => self.drop_prefix(segments),
351 Err(err) if err.kind() == io::ErrorKind::WouldBlock => {
352 return Err(SocketNotWritable);
353 }
354 Err(err) if segments > 1 => {
355 tracing::debug!(
356 ?target,
357 segments,
358 segment_size,
359 err = ?err,
360 "coalesced transmit failed, retrying datagrams individually"
361 );
362 coalesce = false;
363 }
364 Err(err) => {
365 self.drop_prefix(segments);
366 self.discarded_datagrams += 1;
367 self.log_discard(target, segment_size, &err);
368 }
369 }
370 }
371 Ok(())
372 }
373
374 pub async fn flush(&mut self, socket: &UdpSocket) -> io::Result<()> {
379 while !self.is_empty() {
380 socket.writable().await?;
381 if self.try_flush_best_effort(socket).is_ok() {
382 break;
383 }
384 }
385 Ok(())
386 }
387
388 pub fn take_discarded_datagrams(&mut self) -> u64 {
394 std::mem::take(&mut self.discarded_datagrams)
395 }
396
397 fn log_discard(&mut self, target: SocketAddr, datagram_len: usize, err: &io::Error) {
401 let now = Instant::now();
402 let due = self
403 .last_discard_log
404 .is_none_or(|last| now.duration_since(last) >= DISCARD_LOG_INTERVAL);
405
406 if !due {
407 self.discards_since_log += 1;
408 tracing::debug!(?target, datagram_len, err = ?err, "discarding refused datagram");
409 return;
410 }
411
412 tracing::warn!(
413 ?target,
414 datagram_len,
415 err = ?err,
416 suppressed_discards = self.discards_since_log,
417 "discarding refused datagram"
418 );
419 self.last_discard_log = Some(now);
420 self.discards_since_log = 0;
421 }
422
423 fn drop_prefix(&mut self, count: usize) {
424 self.queued_packets.drain(..count);
425 }
426
427 fn fill_scratch_from_front(&mut self, coalesce: bool) -> (SocketAddr, usize, usize) {
432 self.scratch.clear();
433 let (target, first_packet) = self
434 .queued_packets
435 .front()
436 .expect("filling the scratch buffer requires a non-empty queue");
437 let target = *target;
438 let segment_size = first_packet.len();
439 let mut segments = 0;
440 let max_segments = if coalesce && segment_size > 0 {
447 self.state
448 .max_gso_segments()
449 .min(BATCH_SIZE)
450 .min(MAX_UDP_PAYLOAD_SIZE / segment_size)
451 .max(1)
452 } else {
453 1
454 };
455
456 for (queued_target, packet) in self.queued_packets.iter().take(max_segments) {
460 if *queued_target != target || packet.len() != segment_size {
461 break;
462 }
463 self.scratch.extend_from_slice(&packet[..]);
464 segments += 1;
465 }
466
467 (target, segment_size, segments)
468 }
469}
470
471#[cfg(test)]
472mod tests {
473 use std::{net::SocketAddr, time::Duration};
474
475 use ana_gotatun::packet::PacketBufPool;
476 use tokio::net::UdpSocket;
477
478 use super::{MAX_BATCH_SIZE, MAX_UDP_PAYLOAD_SIZE, UdpBatchReceiver, UdpBatchSender};
479
480 const TEST_PACKET_SIZE: usize = 128;
481
482 fn packet_pool() -> PacketBufPool<TEST_PACKET_SIZE> {
483 PacketBufPool::new(MAX_BATCH_SIZE)
484 }
485
486 async fn bound_socket() -> UdpSocket {
487 UdpSocket::bind("127.0.0.1:0").await.unwrap()
488 }
489
490 fn packet_from_bytes(
491 pool: &PacketBufPool<TEST_PACKET_SIZE>,
492 bytes: &[u8],
493 ) -> ana_gotatun::packet::Packet {
494 let mut packet = pool.get();
495 packet[..bytes.len()].copy_from_slice(bytes);
496 packet.truncate(bytes.len());
497 packet
498 }
499
500 #[tokio::test]
501 async fn flushes_partially_full_sender_batch() {
502 let sender_socket = bound_socket().await;
503 let receiver_socket = bound_socket().await;
504 let pool = packet_pool();
505 let mut sender =
506 UdpBatchSender::<MAX_BATCH_SIZE, TEST_PACKET_SIZE>::new(&sender_socket).unwrap();
507
508 sender
509 .try_queue_packet(
510 packet_from_bytes(&pool, b"one"),
511 receiver_socket.local_addr().unwrap(),
512 )
513 .unwrap();
514 sender
515 .try_queue_packet(
516 packet_from_bytes(&pool, b"two"),
517 receiver_socket.local_addr().unwrap(),
518 )
519 .unwrap();
520
521 sender.flush(&sender_socket).await.unwrap();
522
523 let mut buf = [0u8; TEST_PACKET_SIZE];
524 let (n1, _) = receiver_socket.recv_from(&mut buf).await.unwrap();
525 let first = buf[..n1].to_vec();
526 let (n2, _) = receiver_socket.recv_from(&mut buf).await.unwrap();
527 let second = buf[..n2].to_vec();
528
529 assert!(sender.is_empty());
530 assert_eq!(vec![first, second], vec![b"one".to_vec(), b"two".to_vec()]);
531 }
532
533 #[tokio::test]
534 async fn flushes_sender_batch_with_mixed_targets() {
535 let sender_socket = bound_socket().await;
536 let first_target = bound_socket().await;
537 let second_target = bound_socket().await;
538 let pool = packet_pool();
539 let mut sender =
540 UdpBatchSender::<MAX_BATCH_SIZE, TEST_PACKET_SIZE>::new(&sender_socket).unwrap();
541
542 sender
543 .try_queue_packet(
544 packet_from_bytes(&pool, b"alpha"),
545 first_target.local_addr().unwrap(),
546 )
547 .unwrap();
548 sender
549 .try_queue_packet(
550 packet_from_bytes(&pool, b"beta"),
551 second_target.local_addr().unwrap(),
552 )
553 .unwrap();
554 sender
555 .try_queue_packet(
556 packet_from_bytes(&pool, b"gamma"),
557 first_target.local_addr().unwrap(),
558 )
559 .unwrap();
560
561 sender.flush(&sender_socket).await.unwrap();
562
563 let mut buf = [0u8; TEST_PACKET_SIZE];
564 let (n_first_a, _) = first_target.recv_from(&mut buf).await.unwrap();
565 let first_a = buf[..n_first_a].to_vec();
566 let (n_second, _) = second_target.recv_from(&mut buf).await.unwrap();
567 let second = buf[..n_second].to_vec();
568 let (n_first_b, _) = first_target.recv_from(&mut buf).await.unwrap();
569 let first_b = buf[..n_first_b].to_vec();
570
571 assert_eq!(first_a, b"alpha".to_vec());
572 assert_eq!(second, b"beta".to_vec());
573 assert_eq!(first_b, b"gamma".to_vec());
574 }
575
576 #[tokio::test]
579 async fn flush_discards_undeliverable_datagram_and_drains_the_rest() {
580 const OVERSIZED_PACKET_SIZE: usize = MAX_UDP_PAYLOAD_SIZE + 1;
581
582 let sender_socket = bound_socket().await;
583 let receiver_socket = bound_socket().await;
584 let target = receiver_socket.local_addr().unwrap();
585 let pool = PacketBufPool::<OVERSIZED_PACKET_SIZE>::new(2);
586 let mut sender =
587 UdpBatchSender::<MAX_BATCH_SIZE, OVERSIZED_PACKET_SIZE>::new(&sender_socket).unwrap();
588
589 let mut oversized = pool.get();
590 oversized.truncate(OVERSIZED_PACKET_SIZE);
591 sender.try_queue_packet(oversized, target).unwrap();
592
593 let mut deliverable = pool.get();
594 deliverable[..5].copy_from_slice(b"after");
595 deliverable.truncate(5);
596 sender.try_queue_packet(deliverable, target).unwrap();
597
598 sender.flush(&sender_socket).await.unwrap();
599
600 let mut buf = [0u8; 64];
601 let (len, _) = receiver_socket.recv_from(&mut buf).await.unwrap();
602
603 assert!(sender.is_empty(), "the queue must not retain the discard");
604 assert_eq!(&buf[..len], b"after");
605 assert_eq!(
606 sender.take_discarded_datagrams(),
607 1,
608 "the discard must be accountable, the log is only sampled"
609 );
610 assert_eq!(
611 sender.take_discarded_datagrams(),
612 0,
613 "taking the count must reset it"
614 );
615 }
616
617 #[tokio::test]
620 async fn coalescing_respects_the_maximum_udp_payload_size() {
621 const LARGE_PACKET_SIZE: usize = 2048;
622
623 let sender_socket = bound_socket().await;
624 let target = bound_socket().await.local_addr().unwrap();
625 let pool = PacketBufPool::<LARGE_PACKET_SIZE>::new(MAX_BATCH_SIZE);
626 let mut sender =
627 UdpBatchSender::<MAX_BATCH_SIZE, LARGE_PACKET_SIZE>::new(&sender_socket).unwrap();
628
629 for _ in 0..MAX_BATCH_SIZE {
630 let mut packet = pool.get();
631 packet.truncate(LARGE_PACKET_SIZE);
632 sender.try_queue_packet(packet, target).unwrap();
633 }
634
635 let (_, segment_size, segments) = sender.fill_scratch_from_front(true);
636
637 assert_eq!(segment_size, LARGE_PACKET_SIZE);
638 assert!(segments >= 1);
639 assert!(
640 segments * segment_size <= MAX_UDP_PAYLOAD_SIZE,
641 "coalesced {segments} segments of {segment_size} bytes exceed the maximum UDP payload"
642 );
643 }
644
645 #[tokio::test]
652 async fn flush_delivers_every_zero_length_datagram() {
653 const PACKETS: usize = 3;
654
655 let sender_socket = bound_socket().await;
656 let receiver_socket = bound_socket().await;
657 let target = receiver_socket.local_addr().unwrap();
658 let pool = packet_pool();
659 let mut sender =
660 UdpBatchSender::<MAX_BATCH_SIZE, TEST_PACKET_SIZE>::new(&sender_socket).unwrap();
661
662 for _ in 0..PACKETS {
663 sender
664 .try_queue_packet(packet_from_bytes(&pool, b""), target)
665 .unwrap();
666 }
667
668 let (_, segment_size, segments) = sender.fill_scratch_from_front(true);
669 assert_eq!(segment_size, 0);
670 assert_eq!(segments, 1, "zero-length datagrams must be sent one by one");
671
672 sender.flush(&sender_socket).await.unwrap();
673
674 assert!(sender.is_empty());
675 assert_eq!(sender.take_discarded_datagrams(), 0);
676
677 let mut buf = [0u8; TEST_PACKET_SIZE];
678 for index in 0..PACKETS {
679 let received =
680 tokio::time::timeout(Duration::from_secs(5), receiver_socket.recv_from(&mut buf))
681 .await
682 .unwrap_or_else(|_| panic!("zero-length datagram {index} never arrived"))
683 .unwrap();
684
685 assert_eq!(received.0, 0);
686 }
687 }
688
689 #[tokio::test]
692 async fn flush_delivers_more_datagrams_than_one_transmit_can_carry() {
693 const LARGE_PACKET_SIZE: usize = 2048;
694 const PACKETS: usize = MAX_BATCH_SIZE;
695
696 let sender_socket = bound_socket().await;
697 let receiver_socket = bound_socket().await;
698 let target = receiver_socket.local_addr().unwrap();
699 let pool = PacketBufPool::<LARGE_PACKET_SIZE>::new(PACKETS);
700 let mut sender =
701 UdpBatchSender::<MAX_BATCH_SIZE, LARGE_PACKET_SIZE>::new(&sender_socket).unwrap();
702
703 for index in 0..PACKETS {
704 let mut packet = pool.get();
705 packet.truncate(LARGE_PACKET_SIZE);
706 packet[0] = index as u8;
707 sender.try_queue_packet(packet, target).unwrap();
708 }
709
710 let receive = tokio::spawn(async move {
713 let mut buf = [0u8; LARGE_PACKET_SIZE];
714 let mut received = Vec::with_capacity(PACKETS);
715 for _ in 0..PACKETS {
716 let (len, _) = receiver_socket.recv_from(&mut buf).await.unwrap();
717 assert_eq!(len, LARGE_PACKET_SIZE);
718 received.push(buf[0]);
719 }
720 received
721 });
722
723 sender.flush(&sender_socket).await.unwrap();
724
725 let received = tokio::time::timeout(Duration::from_secs(10), receive)
726 .await
727 .expect("all datagrams should arrive")
728 .unwrap();
729
730 assert!(sender.is_empty());
731 assert_eq!(
732 received,
733 (0..PACKETS).map(|index| index as u8).collect::<Vec<_>>(),
734 "datagrams arrived out of order or were lost"
735 );
736 }
737
738 #[tokio::test]
739 async fn receive_with_stride_smaller_than_length_splits_segments() {
740 let socket = bound_socket().await;
741 let pool = packet_pool();
742 let mut receiver =
743 UdpBatchReceiver::<MAX_BATCH_SIZE, TEST_PACKET_SIZE>::new(&socket, &pool).unwrap();
744 let source = "127.0.0.1:30000".parse::<SocketAddr>().unwrap();
745
746 receiver.recv_meta[0].addr = source;
747 receiver.recv_meta[0].len = 10;
748 receiver.recv_meta[0].stride = 4;
749 receiver.recv_slots[0][..10].copy_from_slice(b"abcdefghij");
750
751 let mut seen = Vec::new();
752 receiver
753 .handle_received(0, &pool, &mut |packet, addr| {
754 seen.push((packet[..].to_vec(), addr));
755 Ok::<(), ()>(())
756 })
757 .unwrap();
758
759 assert_eq!(
760 seen,
761 vec![
762 (b"abcd".to_vec(), source),
763 (b"efgh".to_vec(), source),
764 (b"ij".to_vec(), source),
765 ]
766 );
767 }
768
769 #[tokio::test]
770 async fn receive_with_stride_at_least_length_uses_single_packet() {
771 let socket = bound_socket().await;
772 let pool = packet_pool();
773 let mut receiver =
774 UdpBatchReceiver::<MAX_BATCH_SIZE, TEST_PACKET_SIZE>::new(&socket, &pool).unwrap();
775 let source = "127.0.0.1:30001".parse::<SocketAddr>().unwrap();
776
777 receiver.recv_meta[0].addr = source;
778 receiver.recv_meta[0].len = 5;
779 receiver.recv_meta[0].stride = 5;
780 receiver.recv_slots[0][..5].copy_from_slice(b"hello");
781
782 let mut seen = Vec::new();
783 receiver
784 .handle_received(0, &pool, &mut |packet, addr| {
785 seen.push((packet[..].to_vec(), addr));
786 Ok::<(), ()>(())
787 })
788 .unwrap();
789
790 assert_eq!(seen, vec![(b"hello".to_vec(), source)]);
791 }
792
793 #[test]
794 fn refuses_to_grow_beyond_batch_capacity() {
795 let runtime = tokio::runtime::Runtime::new().unwrap();
796 runtime.block_on(async {
797 let socket = bound_socket().await;
798 let pool = packet_pool();
799 let mut sender =
800 UdpBatchSender::<MAX_BATCH_SIZE, TEST_PACKET_SIZE>::new(&socket).unwrap();
801
802 for _ in 0..MAX_BATCH_SIZE {
803 sender
804 .try_queue_packet(packet_from_bytes(&pool, b"x"), socket.local_addr().unwrap())
805 .unwrap();
806 }
807
808 assert!(
809 sender
810 .try_queue_packet(
811 packet_from_bytes(&pool, b"overflow"),
812 socket.local_addr().unwrap()
813 )
814 .is_err()
815 );
816 });
817 }
818}