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