Skip to main content

tokio_quiche/quic/router/
mod.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
27pub(crate) mod acceptor;
28pub(crate) mod connector;
29
30use super::connection::ConnectionMap;
31use super::connection::HandshakeInfo;
32use super::connection::Incoming;
33use super::connection::InitialQuicConnection;
34use super::connection::QuicConnectionParams;
35use super::io::worker::WriterConfig;
36use super::QuicheConnection;
37use crate::metrics::labels;
38use crate::metrics::quic_expensive_metrics_ip_reduce;
39use crate::metrics::Metrics;
40use crate::quic::connection::SharedConnectionIdGenerator;
41use crate::quic::DscpHandle;
42use crate::settings::Config;
43use datagram_socket::DatagramSocketRecv;
44use datagram_socket::DatagramSocketSend;
45use foundations::telemetry::log;
46use quiche::ConnectionId;
47use quiche::Header;
48use quiche::MAX_CONN_ID_LEN;
49use std::default::Default;
50use std::future::Future;
51use std::io;
52use std::net::SocketAddr;
53use std::pin::Pin;
54use std::sync::Arc;
55use std::task::ready;
56use std::task::Context;
57use std::task::Poll;
58use std::time::Instant;
59use std::time::SystemTime;
60use task_killswitch::spawn_with_killswitch;
61use tokio::sync::mpsc;
62
63#[cfg(target_os = "linux")]
64use foundations::telemetry::metrics::Counter;
65#[cfg(target_os = "linux")]
66use foundations::telemetry::metrics::TimeHistogram;
67#[cfg(target_os = "linux")]
68use libc::sockaddr_in;
69#[cfg(target_os = "linux")]
70use libc::sockaddr_in6;
71
72type ConnStream<Tx, M> = mpsc::Receiver<io::Result<InitialQuicConnection<Tx, M>>>;
73
74/// How many incoming packets (GRO batches) to process before checking the
75/// `ConnectionMapCommand` queue again. 30 means "check the command queue once
76/// every 30 packets".
77const PACKET_RX_YIELD_AFTER: usize = 30;
78/// `ConnectionMapCommand` processing batch size to amortize receive operations.
79const CONN_MAP_CMD_BATCH_SIZE: usize = 128;
80
81#[cfg(feature = "perf-quic-listener-metrics")]
82mod listener_stage_timer {
83    use foundations::telemetry::metrics::TimeHistogram;
84    use std::time::Instant;
85
86    pub(super) struct ListenerStageTimer {
87        start: Instant,
88        time_hist: TimeHistogram,
89    }
90
91    impl ListenerStageTimer {
92        pub(super) fn new(
93            start: Instant, time_hist: TimeHistogram,
94        ) -> ListenerStageTimer {
95            ListenerStageTimer { start, time_hist }
96        }
97    }
98
99    impl Drop for ListenerStageTimer {
100        fn drop(&mut self) {
101            self.time_hist
102                .observe((Instant::now() - self.start).as_nanos() as u64);
103        }
104    }
105}
106
107#[derive(Debug)]
108struct PollRecvData {
109    buf: Vec<u8>,
110    // The packet's source, e.g., the peer's address
111    src_addr: SocketAddr,
112    // The packet's original destination. If the original destination is
113    // different from the local listening address, this will be `None`.
114    dst_addr_override: Option<SocketAddr>,
115    rx_time: Option<SystemTime>,
116    gro: Option<i32>,
117    #[cfg(target_os = "linux")]
118    so_mark_data: Option<[u8; 4]>,
119}
120
121/// A message to the listener notifiying a mapping for a connection should be
122/// removed.
123pub enum ConnectionMapCommand {
124    MapCid {
125        existing_cid: ConnectionId<'static>,
126        new_cid: ConnectionId<'static>,
127    },
128    UnmapCid(ConnectionId<'static>),
129}
130
131/// An `InboundPacketRouter` maintains a map of quic connections and routes
132/// [`Incoming`] packets from the [recv half][rh] of a datagram socket to those
133/// connections or some quic initials handler. There is only 1
134/// `InboundPacketRouter` per socket.
135///
136/// [rh]: datagram_socket::DatagramSocketRecv
137///
138/// When a packet (or batch of packets) is received, the router will either
139/// route those packets to an established
140/// [`QuicConnection`](super::QuicConnection) or have a them handled by a
141/// `InitialPacketHandler` which either acts as a quic listener or
142/// quic connector, a server or client respectively.
143///
144/// If you only have a single connection, or if you need more control over the
145/// socket, use `QuicConnection` directly instead.
146pub struct InboundPacketRouter<Tx, Rx, M, I>
147where
148    Tx: DatagramSocketSend + Send + 'static,
149    M: Metrics,
150{
151    socket_tx: Arc<Tx>,
152    socket_rx: Rx,
153    local_addr: SocketAddr,
154    config: Config,
155    conns: ConnectionMap,
156    incoming_packet_handler: I,
157    shutdown_tx: Option<mpsc::Sender<()>>,
158    shutdown_rx: mpsc::Receiver<()>,
159    conn_map_cmd_tx: mpsc::UnboundedSender<ConnectionMapCommand>,
160    conn_map_cmd_rx: mpsc::UnboundedReceiver<ConnectionMapCommand>,
161    /// Reusable buffer to receive a batch of `ConnectionMapCommand`s in
162    /// `poll_conn_map_commands`. Always fully drained after use, so its length
163    /// should be 0 outside of `poll_conn_map_commands`.
164    conn_map_cmd_buf: Vec<ConnectionMapCommand>,
165    accept_sink: mpsc::Sender<io::Result<InitialQuicConnection<Tx, M>>>,
166    metrics: M,
167    #[cfg(target_os = "linux")]
168    udp_drop_count: u32,
169
170    #[cfg(target_os = "linux")]
171    reusable_cmsg_space: Vec<u8>,
172
173    #[cfg(target_os = "linux")]
174    buf: Vec<u8>,
175
176    // We keep the metrics in here, to avoid cloning them each packet
177    #[cfg(target_os = "linux")]
178    metrics_handshake_time_seconds: TimeHistogram,
179    #[cfg(target_os = "linux")]
180    metrics_udp_drop_count: Counter,
181}
182
183impl<Tx, Rx, M, I> InboundPacketRouter<Tx, Rx, M, I>
184where
185    Tx: DatagramSocketSend + Send + 'static,
186    Rx: DatagramSocketRecv,
187    M: Metrics,
188    I: InitialPacketHandler,
189{
190    pub(crate) fn new(
191        config: Config, socket_tx: Arc<Tx>, socket_rx: Rx,
192        local_addr: SocketAddr, incoming_packet_handler: I, metrics: M,
193    ) -> (Self, ConnStream<Tx, M>) {
194        let (shutdown_tx, shutdown_rx) = mpsc::channel(1);
195        let (accept_sink, accept_stream) = mpsc::channel(config.listen_backlog);
196        let (conn_map_cmd_tx, conn_map_cmd_rx) = mpsc::unbounded_channel();
197
198        (
199            InboundPacketRouter {
200                local_addr,
201                socket_tx,
202                socket_rx,
203                conns: ConnectionMap::default(),
204                incoming_packet_handler,
205                shutdown_tx: Some(shutdown_tx),
206                shutdown_rx,
207                conn_map_cmd_tx,
208                conn_map_cmd_rx,
209                conn_map_cmd_buf: Vec::with_capacity(4),
210                accept_sink,
211                #[cfg(target_os = "linux")]
212                udp_drop_count: 0,
213                #[cfg(target_os = "linux")]
214                // Specify CMSG space. Even if they're not all currently used, the cmsg buffer may
215                // have been configured by a previous version of Tokio-Quiche with the socket
216                // re-used on graceful restart. As such, this vector should _only grow_, and care
217                // should be taken when adding new cmsgs.
218                reusable_cmsg_space: nix::cmsg_space!(
219                    u32, // GRO
220                    nix::sys::time::TimeSpec, // timestamp
221                    u16, // drop count
222                    sockaddr_in, // IP_RECVORIGDSTADDR
223                    sockaddr_in6, // IPV6_RECVORIGDSTADDR
224                    u32 // SO_MARK
225                ),
226
227                config,
228
229                #[cfg(target_os = "linux")]
230                buf: Vec::new(),
231                #[cfg(target_os = "linux")]
232                metrics_handshake_time_seconds: metrics.handshake_time_seconds(labels::QuicHandshakeStage::QueueWaiting),
233                #[cfg(target_os = "linux")]
234                metrics_udp_drop_count: metrics.udp_drop_count(),
235
236                metrics,
237
238            },
239            accept_stream,
240        )
241    }
242
243    fn on_incoming(&mut self, mut incoming: Incoming) -> io::Result<()> {
244        #[cfg(feature = "perf-quic-listener-metrics")]
245        let start = std::time::Instant::now();
246
247        if let Some(dcid) = short_dcid(&incoming.buf) {
248            if let Some(ev_sender) = self.conns.get(&dcid) {
249                let _ = ev_sender.try_send(incoming);
250                return Ok(());
251            }
252        }
253
254        let hdr = Header::from_slice(&mut incoming.buf, MAX_CONN_ID_LEN)
255            .map_err(|e| match e {
256                quiche::Error::BufferTooShort | quiche::Error::InvalidPacket =>
257                    labels::QuicInvalidInitialPacketError::FailedToParse.into(),
258                e => io::Error::other(e),
259            })?;
260
261        if let Some(ev_sender) = self.conns.get(&hdr.dcid) {
262            let _ = ev_sender.try_send(incoming);
263            return Ok(());
264        }
265
266        #[cfg(feature = "perf-quic-listener-metrics")]
267        let _timer = listener_stage_timer::ListenerStageTimer::new(
268            start,
269            self.metrics.handshake_time_seconds(
270                labels::QuicHandshakeStage::HandshakeProtocol,
271            ),
272        );
273
274        if self.shutdown_tx.is_none() {
275            return Ok(());
276        }
277
278        let local_addr = incoming.local_addr;
279        let peer_addr = incoming.peer_addr;
280
281        #[cfg(feature = "perf-quic-listener-metrics")]
282        let init_rx_time = incoming.rx_time;
283
284        let new_connection = self.incoming_packet_handler.handle_initials(
285            incoming,
286            hdr,
287            self.config.as_mut(),
288        )?;
289
290        match new_connection {
291            Some(new_connection) => self.spawn_new_connection(
292                new_connection,
293                local_addr,
294                peer_addr,
295                #[cfg(feature = "perf-quic-listener-metrics")]
296                init_rx_time,
297            ),
298            None => Ok(()),
299        }
300    }
301
302    /// Creates a new [`QuicConnection`](super::QuicConnection) and spawns an
303    /// associated io worker.
304    fn spawn_new_connection(
305        &mut self, new_connection: NewConnection, local_addr: SocketAddr,
306        peer_addr: SocketAddr,
307        #[cfg(feature = "perf-quic-listener-metrics")] init_rx_time: Option<
308            SystemTime,
309        >,
310    ) -> io::Result<()> {
311        let NewConnection {
312            conn,
313            pending_cid,
314            cid_generator,
315            handshake_start_time,
316            initial_pkt,
317            enable_per_connection_dscp,
318        } = new_connection;
319
320        let Some(ref shutdown_tx) = self.shutdown_tx else {
321            // Do not create new connections while shutting down.
322            return Ok(());
323        };
324        let Ok(send_permit) = self.accept_sink.try_reserve() else {
325            // Drop the connection when the backlog is full. The client retries.
326            return Err(
327                labels::QuicInvalidInitialPacketError::AcceptQueueOverflow.into(),
328            );
329        };
330
331        let scid = conn.source_id().into_owned();
332        let writer_cfg = WriterConfig {
333            peer_addr,
334            local_addr,
335            pending_cid: pending_cid.clone(),
336            with_gso: self.config.has_gso,
337            pacing_offload: self.config.pacing_offload,
338            with_pktinfo: if self.local_addr.is_ipv4() {
339                self.config.has_ippktinfo
340            } else {
341                self.config.has_ipv6pktinfo
342            },
343            pool_send_buffer: self.config.pool_send_buffer,
344        };
345
346        let handshake_info = HandshakeInfo::new(
347            handshake_start_time,
348            self.config.handshake_timeout,
349        );
350
351        let conn = InitialQuicConnection::new(QuicConnectionParams {
352            writer_cfg,
353            initial_pkt,
354            dscp_handle: enable_per_connection_dscp.then(DscpHandle::new),
355            shutdown_tx: shutdown_tx.clone(),
356            conn_map_cmd_tx: self.conn_map_cmd_tx.clone(),
357            scid: scid.clone(),
358            cid_generator,
359            metrics: self.metrics.clone(),
360            connection_hook: self.config.connection_hook.clone(),
361            #[cfg(feature = "perf-quic-listener-metrics")]
362            init_rx_time,
363            handshake_info,
364            quiche_conn: conn,
365            socket: Arc::clone(&self.socket_tx),
366            local_addr,
367            peer_addr,
368        });
369
370        conn.audit_log_stats
371            .set_transport_handshake_start(instant_to_system(
372                handshake_start_time,
373            ));
374
375        self.conns.insert(&scid, &conn);
376
377        // Add the client-generated "pending" connection ID to the map as well.
378        // This is only required for QUIC servers, because clients can send
379        // Initial packets with arbitrary DCIDs to servers.
380        if let Some(pending_cid) = pending_cid {
381            self.conns.map_cid(&scid, &pending_cid);
382        }
383
384        self.metrics.accepted_initial_packet_count().inc();
385        if self.config.enable_expensive_packet_count_metrics {
386            if let Some(peer_ip) =
387                quic_expensive_metrics_ip_reduce(conn.peer_addr().ip())
388            {
389                self.metrics
390                    .expensive_accepted_initial_packet_count(peer_ip)
391                    .inc();
392            }
393        }
394
395        send_permit.send(Ok(conn));
396        Ok(())
397    }
398}
399
400impl<Tx, Rx, M, I> InboundPacketRouter<Tx, Rx, M, I>
401where
402    Tx: DatagramSocketSend + Send + Sync + 'static,
403    Rx: DatagramSocketRecv,
404    M: Metrics,
405    I: InitialPacketHandler,
406{
407    /// [`InboundPacketRouter::poll_recv_from`] should be used if the underlying
408    /// system or socket does not support rx_time nor GRO.
409    fn poll_recv_from(
410        &mut self, cx: &mut Context<'_>,
411    ) -> Poll<io::Result<PollRecvData>> {
412        let mut buf = Vec::with_capacity(datagram_socket::MAX_DATAGRAM_SIZE);
413        // We use ReadBuf's ability to write to uninitialized memory to avoid
414        // the cost of having to initialize the Vec.
415        let mut read_buf = tokio::io::ReadBuf::uninit(buf.spare_capacity_mut());
416        let addr = ready!(self.socket_rx.poll_recv_from(cx, &mut read_buf))?;
417        let n = read_buf.filled().len();
418        unsafe {
419            // Safety: ReadBuf has guaranteed that `n` initialized bytes have
420            // been written to the buffer, so we can set the vec's length
421            // accordingly
422            buf.set_len(n);
423        }
424        Poll::Ready(Ok(PollRecvData {
425            buf,
426            src_addr: addr,
427            rx_time: None,
428            gro: None,
429            dst_addr_override: None,
430            #[cfg(target_os = "linux")]
431            so_mark_data: None,
432        }))
433    }
434
435    fn poll_recv_and_rx_time(
436        &mut self, cx: &mut Context<'_>,
437    ) -> Poll<io::Result<PollRecvData>> {
438        #[cfg(not(target_os = "linux"))]
439        {
440            self.poll_recv_from(cx)
441        }
442
443        #[cfg(target_os = "linux")]
444        {
445            use libc::SOL_SOCKET;
446            use libc::SO_MARK;
447            use nix::errno::Errno;
448            use nix::sys::socket::*;
449            use std::net::SocketAddrV4;
450            use std::net::SocketAddrV6;
451            use std::os::fd::AsRawFd;
452            use tokio::io::Interest;
453
454            use crate::buf_factory::BufFactory;
455
456            let Some(udp_socket) = self.socket_rx.as_udp_socket() else {
457                // the given socket is not a UDP socket, fall back to the
458                // simple poll_recv_from.
459                return self.poll_recv_from(cx);
460            };
461
462            // Note, the resize will be a no-op after the first call since
463            // we never truncate the `self.buf`
464            self.buf.resize(BufFactory::MAX_BUF_SIZE, 0u8);
465            loop {
466                let iov_s = &mut [io::IoSliceMut::new(&mut self.buf)];
467                match udp_socket.try_io(Interest::READABLE, || {
468                    recvmsg::<SockaddrStorage>(
469                        udp_socket.as_raw_fd(),
470                        iov_s,
471                        Some(&mut self.reusable_cmsg_space),
472                        MsgFlags::empty(),
473                    )
474                    .map_err(|x| x.into())
475                }) {
476                    Ok(r) => {
477                        let filled_buf =
478                            r.iovs().next().map(Vec::from).unwrap_or_default();
479                        // Verify that the `recvmsg` slices total `r.bytes`.
480                        debug_assert_eq!(r.bytes, filled_buf.len());
481
482                        let address = match r.address {
483                            Some(inner) => inner,
484                            _ => return Poll::Ready(Err(Errno::EINVAL.into())),
485                        };
486
487                        let peer_addr = match address.family() {
488                            Some(AddressFamily::Inet) => SocketAddrV4::from(
489                                *address.as_sockaddr_in().unwrap(),
490                            )
491                            .into(),
492                            Some(AddressFamily::Inet6) => SocketAddrV6::from(
493                                *address.as_sockaddr_in6().unwrap(),
494                            )
495                            .into(),
496                            _ => {
497                                return Poll::Ready(Err(Errno::EINVAL.into()));
498                            },
499                        };
500
501                        let mut rx_time = None;
502                        let mut gro = None;
503                        let mut dst_addr_override = None;
504                        let mut mark_bytes: Option<[u8; 4]> = None;
505
506                        let Ok(cmsgs) = r.cmsgs() else {
507                            // Best-effort if we can't read cmsgs.
508                            return Poll::Ready(Ok(PollRecvData {
509                                buf: filled_buf,
510                                src_addr: peer_addr,
511                                dst_addr_override,
512                                rx_time,
513                                gro,
514                                so_mark_data: mark_bytes,
515                            }));
516                        };
517
518                        for cmsg in cmsgs {
519                            match cmsg {
520                                ControlMessageOwned::RxqOvfl(c) => {
521                                    if c != self.udp_drop_count {
522                                        self.metrics_udp_drop_count.inc_by(
523                                            (c - self.udp_drop_count) as u64,
524                                        );
525                                        self.udp_drop_count = c;
526                                    }
527                                },
528                                ControlMessageOwned::ScmTimestampns(val) => {
529                                    rx_time = SystemTime::UNIX_EPOCH
530                                        .checked_add(val.into());
531                                    if let Some(delta) =
532                                        rx_time.and_then(|rx_time| {
533                                            rx_time.elapsed().ok()
534                                        })
535                                    {
536                                        self.metrics_handshake_time_seconds
537                                            .observe(delta.as_nanos() as u64);
538                                    }
539                                },
540                                ControlMessageOwned::UdpGroSegments(val) =>
541                                    gro = Some(val),
542                                ControlMessageOwned::Ipv4OrigDstAddr(val) => {
543                                    let source_addr = std::net::Ipv4Addr::from(
544                                        u32::to_be(val.sin_addr.s_addr),
545                                    );
546                                    let source_port = u16::to_be(val.sin_port);
547
548                                    let parsed_addr =
549                                        SocketAddr::V4(SocketAddrV4::new(
550                                            source_addr,
551                                            source_port,
552                                        ));
553
554                                    dst_addr_override = resolve_dst_addr(
555                                        &self.local_addr,
556                                        &parsed_addr,
557                                    );
558                                },
559                                ControlMessageOwned::Ipv6OrigDstAddr(val) => {
560                                    // IPv6 is a byte array and needs no swap.
561                                    // IPv4 is parsed as a `u32` and does.
562                                    let source_addr = std::net::Ipv6Addr::from(
563                                        val.sin6_addr.s6_addr,
564                                    );
565                                    let source_port = u16::to_be(val.sin6_port);
566                                    let source_flowinfo =
567                                        u32::to_be(val.sin6_flowinfo);
568                                    let source_scope =
569                                        u32::to_be(val.sin6_scope_id);
570
571                                    let parsed_addr =
572                                        SocketAddr::V6(SocketAddrV6::new(
573                                            source_addr,
574                                            source_port,
575                                            source_flowinfo,
576                                            source_scope,
577                                        ));
578
579                                    dst_addr_override = resolve_dst_addr(
580                                        &self.local_addr,
581                                        &parsed_addr,
582                                    );
583                                },
584                                ControlMessageOwned::Ipv4PacketInfo(_) |
585                                ControlMessageOwned::Ipv6PacketInfo(_) => {
586                                    // We only want the destination address from
587                                    // IP_RECVORIGDSTADDR, but we'll get these
588                                    // messages because we set IP_PKTINFO on the
589                                    // socket.
590                                },
591                                ControlMessageOwned::Unknown(raw_cmsg) => {
592                                    let UnknownCmsg {
593                                        cmsg_header,
594                                        data_bytes,
595                                    } = raw_cmsg;
596
597                                    if cmsg_header.cmsg_level == SOL_SOCKET &&
598                                        cmsg_header.cmsg_type == SO_MARK
599                                    {
600                                        let Ok(arr) =
601                                            <[u8; 4]>::try_from(data_bytes)
602                                        else {
603                                            // SO_MARK is a `u32`. This should
604                                            // always succeed.
605                                            // https://elixir.bootlin.com/linux/v6.17/source/include/net/sock.h#L487
606                                            continue;
607                                        };
608
609                                        let _ = mark_bytes.insert(arr);
610                                    }
611                                },
612                                _ => {
613                                    // Unrecognized cmsg received, just ignore
614                                    // it.
615                                },
616                            };
617                        }
618
619                        return Poll::Ready(Ok(PollRecvData {
620                            buf: filled_buf,
621                            src_addr: peer_addr,
622                            dst_addr_override,
623                            rx_time,
624                            gro,
625                            so_mark_data: mark_bytes,
626                        }));
627                    },
628                    Err(e) if e.kind() == io::ErrorKind::WouldBlock => {
629                        // NOTE: we manually poll the socket here to register
630                        // interest in the socket to become
631                        // writable for the given `cx`. Under the hood, tokio's
632                        // implementation just checks for
633                        // EWOULDBLOCK and if socket is busy registers provided
634                        // waker to be invoked when the
635                        // socket is free and consequently drive the event loop.
636                        ready!(udp_socket.poll_recv_ready(cx))?
637                    },
638                    Err(e) => return Poll::Ready(Err(e)),
639                }
640            }
641        }
642    }
643
644    fn poll_process_packet(&mut self, cx: &mut Context) -> Poll<()> {
645        let pkt_data = match ready!(self.poll_recv_and_rx_time(cx)) {
646            Ok(v) => v,
647            Err(e) => {
648                log::error!("Incoming packet router encountered recvmsg error"; "error" => e);
649                return Poll::Ready(());
650            },
651        };
652
653        let PollRecvData {
654            buf,
655            src_addr: peer_addr,
656            dst_addr_override,
657            rx_time,
658            gro,
659            #[cfg(target_os = "linux")]
660            so_mark_data,
661        } = pkt_data;
662
663        let send_from = if let Some(dst_addr) = dst_addr_override {
664            log::trace!("overriding local address"; "actual_local" => dst_addr, "configured_local" => self.local_addr);
665            dst_addr
666        } else {
667            self.local_addr
668        };
669
670        let res = self.on_incoming(Incoming {
671            peer_addr,
672            local_addr: send_from,
673            buf,
674            rx_time,
675            gro,
676            #[cfg(target_os = "linux")]
677            so_mark_data,
678        });
679
680        // Only error handling below - if `on_incoming` was successful,
681        // we return here
682        let Err(e) = res else {
683            return Poll::Ready(());
684        };
685
686        let err_type = initial_packet_error_type(&e);
687        self.metrics
688            .rejected_initial_packet_count(err_type.clone())
689            .inc();
690
691        if self.config.enable_expensive_packet_count_metrics {
692            if let Some(peer_ip) =
693                quic_expensive_metrics_ip_reduce(peer_addr.ip())
694            {
695                self.metrics
696                    .expensive_rejected_initial_packet_count(
697                        err_type.clone(),
698                        peer_ip,
699                    )
700                    .inc();
701            }
702        }
703
704        if matches!(err_type, labels::QuicInvalidInitialPacketError::Unexpected) {
705            // don't block packet routing on errors
706            let _ = self.accept_sink.try_send(Err(e));
707        }
708
709        Poll::Ready(())
710    }
711
712    fn poll_conn_map_commands(&mut self, cx: &mut Context) -> Poll<()> {
713        let cmd_rx = &mut self.conn_map_cmd_rx;
714        let buf = &mut self.conn_map_cmd_buf;
715        debug_assert!(buf.is_empty());
716
717        while ready!(cmd_rx.poll_recv_many(cx, buf, CONN_MAP_CMD_BATCH_SIZE)) > 0
718        {
719            for cmd in buf.drain(..) {
720                match cmd {
721                    ConnectionMapCommand::MapCid {
722                        existing_cid,
723                        new_cid,
724                    } => self.conns.map_cid(&existing_cid, &new_cid),
725                    ConnectionMapCommand::UnmapCid(cid) =>
726                        self.conns.unmap_cid(&cid),
727                }
728            }
729        }
730
731        Poll::Ready(())
732    }
733}
734
735// Quickly extract the connection id of a short quic packet without allocating
736fn short_dcid(buf: &[u8]) -> Option<ConnectionId<'_>> {
737    let is_short_dcid = buf.first()? >> 7 == 0;
738
739    if is_short_dcid {
740        buf.get(1..1 + MAX_CONN_ID_LEN).map(ConnectionId::from_ref)
741    } else {
742        None
743    }
744}
745
746/// Converts an [`Instant`] to a [`SystemTime`], based on the current delta
747/// between both clocks.
748fn instant_to_system(ts: Instant) -> SystemTime {
749    let now = Instant::now();
750    let system_now = SystemTime::now();
751    if let Some(delta) = now.checked_duration_since(ts) {
752        return system_now - delta;
753    }
754
755    let delta = ts.checked_duration_since(now).expect("now < ts");
756    system_now + delta
757}
758
759/// Determine if we should store the destination address for a packet, based on
760/// an address parsed from a
761/// [`ControlMessageOwned`](nix::sys::socket::ControlMessageOwned).
762///
763/// This is to prevent overriding the destination address if the packet was
764/// originally addressed to `local`, as that would cause us to incorrectly
765/// address packets when sending.
766///
767/// Returns the parsed address if it should be stored.
768#[cfg(target_os = "linux")]
769fn resolve_dst_addr(
770    local: &SocketAddr, parsed: &SocketAddr,
771) -> Option<SocketAddr> {
772    if local != parsed {
773        return Some(*parsed);
774    }
775
776    None
777}
778
779impl<Tx, Rx, M, I> Future for InboundPacketRouter<Tx, Rx, M, I>
780where
781    Tx: DatagramSocketSend + Send + Sync + 'static,
782    Rx: DatagramSocketRecv + Unpin,
783    M: Metrics,
784    I: InitialPacketHandler + Unpin,
785{
786    type Output = io::Result<()>;
787
788    fn poll(mut self: Pin<&mut Self>, cx: &mut Context) -> Poll<io::Result<()>> {
789        loop {
790            // First, check whether the app stopped accepting connections.
791            if self.shutdown_tx.is_some() && self.accept_sink.is_closed() {
792                self.shutdown_tx = None;
793            }
794
795            // Second, check if all connections have shut down and we can exit.
796            if self.shutdown_tx.is_none() &&
797                self.shutdown_rx.poll_recv(cx).is_ready()
798            {
799                return Poll::Ready(Ok(()));
800            }
801
802            // Third, run the generic `InitialPacketHandler` update.
803            if let Err(error) = self.incoming_packet_handler.update(cx) {
804                // An error here is so rare that it's easier to spawn a separate
805                // task
806                let sender = self.accept_sink.clone();
807                spawn_with_killswitch(async move {
808                    let _ = sender.send(Err(error)).await;
809                });
810            }
811
812            // Fourth, update ConnectionMap before receiving packets so SCID
813            // destinations are current. A pending result means all available
814            // commands were processed and the next command will wake us.
815            let _ = self.poll_conn_map_commands(cx);
816
817            // Finally, process up to `PACKET_RX_YIELD_AFTER` packet batches. If
818            // no more packets are available, wait to be woken again.
819            for _ in 0..PACKET_RX_YIELD_AFTER {
820                ready!(self.poll_process_packet(cx));
821            }
822        }
823    }
824}
825
826/// Categorizes errors that are returned when handling packets which are not
827/// associated with an established connection. The purpose is to suppress
828/// logging of 'expected' errors (e.g. junk data sent to the UDP socket) to
829/// prevent DoS.
830fn initial_packet_error_type(
831    e: &io::Error,
832) -> labels::QuicInvalidInitialPacketError {
833    Some(e)
834        .filter(|e| e.kind() == io::ErrorKind::Other)
835        .and_then(io::Error::get_ref)
836        .and_then(|e| e.downcast_ref())
837        .map_or(
838            labels::QuicInvalidInitialPacketError::Unexpected,
839            Clone::clone,
840        )
841}
842
843/// An [`InitialPacketHandler`] handles unknown quic initials and processes
844/// them; generally accepting new connections (acting as a server), or
845/// establishing a connection to a server (acting as a client). An
846/// [`InboundPacketRouter`] holds an instance of this trait and routes
847/// [`Incoming`] packets to it when it receives initials.
848///
849/// The handler produces [`quiche::Connection`]s which are then turned into
850/// [`QuicConnection`](super::QuicConnection), IoWorker pair.
851pub trait InitialPacketHandler {
852    fn update(&mut self, _ctx: &mut Context<'_>) -> io::Result<()> {
853        Ok(())
854    }
855
856    fn handle_initials(
857        &mut self, incoming: Incoming, hdr: Header<'static>,
858        quiche_config: &mut quiche::Config,
859    ) -> io::Result<Option<NewConnection>>;
860}
861
862/// A [`NewConnection`] describes a new [`quiche::Connection`] that can be
863/// driven by an io worker.
864pub struct NewConnection {
865    /// See [`QuicConnectionParams::quiche_conn`].
866    conn: Box<QuicheConnection>,
867    pending_cid: Option<ConnectionId<'static>>,
868    initial_pkt: Option<Incoming>,
869    enable_per_connection_dscp: bool,
870    cid_generator: Option<SharedConnectionIdGenerator>,
871    /// When the handshake started. Should be called before [`quiche::accept`]
872    /// or [`quiche::connect`].
873    handshake_start_time: Instant,
874}
875
876// TODO: the router module is private so we can't move these to /tests
877// TODO: Rewrite tests to be Windows compatible
878#[cfg(all(test, unix))]
879mod tests {
880    use super::acceptor::ConnectionAcceptor;
881    use super::acceptor::ConnectionAcceptorConfig;
882    use super::*;
883
884    use crate::http3::settings::Http3Settings;
885    use crate::metrics::DefaultMetrics;
886    use crate::quic::connection::SimpleConnectionIdGenerator;
887    use crate::quic::Dscp;
888    use crate::settings::Config;
889    use crate::settings::Hooks;
890    use crate::settings::QuicSettings;
891    use crate::settings::TlsCertificatePaths;
892    use crate::socket::SocketCapabilities;
893    use crate::ConnectionIdGenerator as _;
894    use crate::ConnectionParams;
895    use crate::ServerH3Driver;
896
897    use datagram_socket::MAX_DATAGRAM_SIZE;
898    use futures::FutureExt as _;
899    use h3i::actions::h3::Action;
900    use h3i::actions::h3::WaitType;
901    use std::net::Ipv4Addr;
902    use std::sync::Arc;
903    use std::time::Duration;
904    use tokio::net::UdpSocket;
905    use tokio::time;
906
907    const TEST_CERT_FILE: &str = concat!(
908        env!("CARGO_MANIFEST_DIR"),
909        "/",
910        "../quiche/examples/cert.crt"
911    );
912    const TEST_KEY_FILE: &str = concat!(
913        env!("CARGO_MANIFEST_DIR"),
914        "/",
915        "../quiche/examples/cert.key"
916    );
917
918    fn test_connect(host_port: String, wait_before_close: Option<Duration>) {
919        let h3i_config = h3i::config::Config::new()
920            .with_host_port("test.com".to_string())
921            .with_idle_timeout(2000)
922            .with_connect_to(host_port)
923            .verify_peer(false)
924            .build()
925            .unwrap();
926
927        let conn_close = h3i::quiche::ConnectionError {
928            is_app: true,
929            error_code: h3i::quiche::WireErrorCode::NoError as _,
930            reason: Vec::new(),
931        };
932        let mut actions = Vec::new();
933        if let Some(duration) = wait_before_close {
934            actions.push(Action::Wait {
935                wait_type: WaitType::WaitDuration(duration),
936            });
937        }
938        actions.push(Action::ConnectionClose { error: conn_close });
939
940        let _ = h3i::client::sync_client::connect(h3i_config, actions, None);
941    }
942
943    #[tokio::test]
944    async fn test_timeout() {
945        // Configure a short idle timeout to speed up connection reclamation as
946        // quiche doesn't support time mocking
947        let quic_settings = QuicSettings {
948            max_idle_timeout: Some(Duration::from_millis(1)),
949            max_recv_udp_payload_size: MAX_DATAGRAM_SIZE,
950            max_send_udp_payload_size: MAX_DATAGRAM_SIZE,
951            ..Default::default()
952        };
953
954        let tls_cert_settings = TlsCertificatePaths {
955            cert: TEST_CERT_FILE,
956            private_key: TEST_KEY_FILE,
957            kind: crate::settings::CertificateKind::X509,
958        };
959
960        let mut params = ConnectionParams::new_server(
961            quic_settings,
962            tls_cert_settings,
963            Hooks::default(),
964        );
965        params.enable_per_connection_dscp = true;
966        let config = Config::new(&params, SocketCapabilities::default()).unwrap();
967
968        let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
969        let local_addr = socket.local_addr().unwrap();
970        let host_port = local_addr.to_string();
971        let socket_tx = Arc::new(socket);
972        let socket_rx = Arc::clone(&socket_tx);
973
974        let acceptor = ConnectionAcceptor::new(
975            ConnectionAcceptorConfig {
976                disable_client_ip_validation: config.disable_client_ip_validation,
977                enable_per_connection_dscp: params.enable_per_connection_dscp,
978                qlog_dir: config.qlog_dir.clone(),
979                qlog_compression: config.qlog_compression,
980                keylog_file: config
981                    .keylog_file
982                    .as_ref()
983                    .and_then(|f| f.try_clone().ok()),
984                #[cfg(target_os = "linux")]
985                with_pktinfo: false,
986            },
987            Arc::clone(&socket_tx),
988            Default::default(),
989            Arc::new(SimpleConnectionIdGenerator),
990            DefaultMetrics,
991        );
992
993        let (socket_driver, mut incoming) = InboundPacketRouter::new(
994            config,
995            socket_tx,
996            socket_rx,
997            local_addr,
998            acceptor,
999            DefaultMetrics,
1000        );
1001        tokio::spawn(socket_driver);
1002
1003        // Start a request and drop it after connection establishment
1004        std::thread::spawn(move || test_connect(host_port, None));
1005
1006        // Wait for a new connection
1007        time::pause();
1008
1009        let (h3_driver, _) = ServerH3Driver::new(Http3Settings::default());
1010        let conn = incoming.recv().await.unwrap().unwrap();
1011        let dscp_handle = conn.dscp_handle().unwrap().clone();
1012        assert_eq!(dscp_handle.get(), None);
1013        dscp_handle.set(Dscp::new(0));
1014        assert_eq!(conn.dscp_handle().unwrap().get(), Dscp::new(0));
1015        let drop_check = conn.incoming_ev_sender.clone();
1016        let _conn = conn.start(h3_driver);
1017
1018        // Poll incoming events until the connection is dropped.
1019        time::advance(Duration::new(30, 0)).await;
1020        time::resume();
1021
1022        // This is a smoke test. A failure leaves `notified()` unresolved and
1023        // hangs the test.
1024        drop_check.closed().await;
1025    }
1026
1027    #[cfg(all(target_os = "linux", not(feature = "fuzzing")))]
1028    #[tokio::test]
1029    async fn test_worker_dscp_across_handshake_ipv4() {
1030        check_worker_dscp_across_handshake("127.0.0.1:0").await;
1031    }
1032
1033    #[cfg(all(target_os = "linux", not(feature = "fuzzing")))]
1034    #[tokio::test]
1035    async fn test_worker_dscp_across_handshake_ipv6() {
1036        check_worker_dscp_across_handshake("[::1]:0").await;
1037    }
1038
1039    #[cfg(all(target_os = "linux", not(feature = "fuzzing")))]
1040    async fn check_worker_dscp_across_handshake(bind_addr: &str) {
1041        use crate::quic::io::gso::test::enable_tos;
1042        use crate::quic::io::gso::test::recv_tos;
1043        use futures::StreamExt;
1044        use nix::sys::socket::setsockopt;
1045        use nix::sys::socket::sockopt;
1046
1047        let socket = std::net::UdpSocket::bind(bind_addr).unwrap();
1048        let server_addr = socket.local_addr().unwrap();
1049        if server_addr.is_ipv4() {
1050            setsockopt(&socket, sockopt::Ipv4Tos, &0).unwrap();
1051        } else {
1052            setsockopt(&socket, sockopt::Ipv6TClass, &0).unwrap();
1053        }
1054
1055        let mut settings = QuicSettings::default();
1056        settings.disable_client_ip_validation = true;
1057        let tls_cert = TlsCertificatePaths {
1058            cert: TEST_CERT_FILE,
1059            private_key: TEST_KEY_FILE,
1060            kind: crate::settings::CertificateKind::X509,
1061        };
1062        let mut params =
1063            ConnectionParams::new_server(settings, tls_cert, Hooks::default());
1064        params.enable_per_connection_dscp = true;
1065        let mut incoming = crate::listen(vec![socket], params, DefaultMetrics)
1066            .unwrap()
1067            .remove(0);
1068
1069        let frontend = UdpSocket::bind(bind_addr).await.unwrap();
1070        let backend = UdpSocket::bind(bind_addr).await.unwrap();
1071        enable_tos(&backend);
1072        let host_port = frontend.local_addr().unwrap().to_string();
1073        let (marks_tx, mut marks_rx) = tokio::sync::mpsc::unbounded_channel();
1074        let relay = tokio::spawn(async move {
1075            let mut client_addr = None;
1076            let mut buf = [0u8; 2048];
1077            loop {
1078                tokio::select! {
1079                    received = frontend.recv_from(&mut buf) => {
1080                        let (len, addr) = received.unwrap();
1081                        client_addr = Some(addr);
1082                        backend.send_to(&buf[..len], server_addr).await.unwrap();
1083                    }
1084                    (packet, tos) = recv_tos(&backend) => {
1085                        marks_tx.send(tos >> 2).unwrap();
1086                        frontend.send_to(&packet, client_addr.unwrap()).await.unwrap();
1087                    }
1088                }
1089            }
1090        });
1091        let client = tokio::task::spawn_blocking(move || {
1092            test_connect(host_port, Some(Duration::from_secs(1)))
1093        });
1094
1095        let conn = time::timeout(Duration::from_secs(5), incoming.next())
1096            .await
1097            .expect("no initial packet")
1098            .expect("listener closed")
1099            .expect("initial packet rejected");
1100        let dscp = conn.dscp_handle().unwrap().clone();
1101        dscp.set(Dscp::new(34));
1102        let (driver, controller) = ServerH3Driver::new(Http3Settings::default());
1103        let (_, mut worker) =
1104            time::timeout(Duration::from_secs(5), conn.handshake(driver))
1105                .await
1106                .expect("handshake timed out")
1107                .expect("handshake failed");
1108
1109        // No stateless packets are sent when client IP validation is disabled.
1110        assert_eq!(
1111            time::timeout(Duration::from_secs(5), marks_rx.recv())
1112                .await
1113                .expect("no handshake packet"),
1114            Some(34)
1115        );
1116
1117        // Force a packet in the running worker to verify the handle survives
1118        // handshake() and resume(), not just the initial atomic clone.
1119        dscp.set(Dscp::new(28));
1120        worker.qconn.send_ack_eliciting().unwrap();
1121        InitialQuicConnection::resume(worker);
1122        time::timeout(Duration::from_secs(5), async {
1123            loop {
1124                match marks_rx.recv().await.expect("UDP relay closed") {
1125                    34 => (), // Handshake packets already in flight.
1126                    28 => break,
1127                    mark => panic!("unexpected post-handshake DSCP {mark}"),
1128                }
1129            }
1130        })
1131        .await
1132        .expect("no post-handshake packet with DSCP 28");
1133
1134        client.await.unwrap();
1135        drop(controller);
1136        relay.abort();
1137    }
1138
1139    struct NoopDatagramSender;
1140    impl DatagramSocketSend for NoopDatagramSender {
1141        fn poll_send(
1142            &self, _cx: &mut Context, buf: &[u8],
1143        ) -> Poll<io::Result<usize>> {
1144            Poll::Ready(Ok(buf.len()))
1145        }
1146
1147        fn poll_send_to(
1148            &self, _cx: &mut Context, buf: &[u8], _addr: SocketAddr,
1149        ) -> Poll<io::Result<usize>> {
1150            Poll::Ready(Ok(buf.len()))
1151        }
1152    }
1153
1154    struct AlwaysReadyReceiver;
1155    impl DatagramSocketRecv for AlwaysReadyReceiver {
1156        fn poll_recv(
1157            &mut self, _cx: &mut Context, buf: &mut tokio::io::ReadBuf,
1158        ) -> Poll<io::Result<()>> {
1159            // Short header packet:
1160            // 1 byte descriptor + 20 byte DCID + 1 byte packet number + payload
1161            const DUMMY_QUIC_PACKET: &[u8] =
1162                b"\x40THIS_20_BYTE_CONN_ID\x06payload_payload_payload";
1163            buf.put_slice(DUMMY_QUIC_PACKET);
1164            Poll::Ready(Ok(()))
1165        }
1166    }
1167
1168    struct NoopInitialHandler;
1169    impl InitialPacketHandler for NoopInitialHandler {
1170        fn handle_initials(
1171            &mut self, _incoming: Incoming, _hdr: Header<'static>,
1172            _quiche_config: &mut quiche::Config,
1173        ) -> io::Result<Option<NewConnection>> {
1174            Ok(None)
1175        }
1176    }
1177
1178    #[test]
1179    fn test_poll_packet_always_ready() {
1180        let tls_cert_settings = TlsCertificatePaths {
1181            cert: TEST_CERT_FILE,
1182            private_key: TEST_KEY_FILE,
1183            kind: crate::settings::CertificateKind::X509,
1184        };
1185        let params = ConnectionParams::new_server(
1186            QuicSettings::default(),
1187            tls_cert_settings,
1188            Hooks::default(),
1189        );
1190
1191        let config = Config::new(&params, SocketCapabilities::default()).unwrap();
1192        let local_addr = SocketAddr::new(Ipv4Addr::UNSPECIFIED.into(), 0);
1193
1194        let (mut ipr, accept_stream) = InboundPacketRouter::new(
1195            config,
1196            Arc::new(NoopDatagramSender),
1197            AlwaysReadyReceiver,
1198            local_addr,
1199            NoopInitialHandler,
1200            DefaultMetrics,
1201        );
1202        let conn_map_cmd_tx = ipr.conn_map_cmd_tx.clone();
1203
1204        // Keep polling the IPR in a busy loop until it resolves
1205        let (ipr_notifier, ipr_done) = std::sync::mpsc::sync_channel::<()>(0);
1206        let ipr = std::thread::spawn(move || {
1207            let mut cx = Context::from_waker(std::task::Waker::noop());
1208            while ipr.poll_unpin(&mut cx).is_pending() {
1209                std::thread::sleep(Duration::from_millis(10));
1210            }
1211            drop(ipr_notifier);
1212            ipr
1213        });
1214
1215        // Fill the `conn_map_cmd` channel with some messages to process
1216        for _ in 0..20 {
1217            let random_cid = SimpleConnectionIdGenerator.new_connection_id();
1218            conn_map_cmd_tx
1219                .send(ConnectionMapCommand::UnmapCid(random_cid))
1220                .unwrap();
1221        }
1222        // Give the IPR some time to process the ConnectionMapCommands
1223        std::thread::sleep(Duration::from_secs(1));
1224
1225        // Shut the IPR down by dropping the accept_stream receiver. We wait for
1226        // up to 10 seconds for IPR::poll to resolve. If it doesn't, it's not
1227        // checking the shutdown condition regularly.
1228        drop(accept_stream);
1229        let ipr_done_res = ipr_done.recv_timeout(Duration::from_secs(10));
1230        assert_eq!(
1231            ipr_done_res,
1232            Err(std::sync::mpsc::RecvTimeoutError::Disconnected)
1233        );
1234
1235        // Check that the ConnectionMapCommands we added above were actually
1236        // processed
1237        let ipr = ipr.join().unwrap();
1238        assert!(ipr.conn_map_cmd_rx.is_empty());
1239    }
1240}