1pub(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
74const PACKET_RX_YIELD_AFTER: usize = 30;
78const 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 src_addr: SocketAddr,
112 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
121pub enum ConnectionMapCommand {
124 MapCid {
125 existing_cid: ConnectionId<'static>,
126 new_cid: ConnectionId<'static>,
127 },
128 UnmapCid(ConnectionId<'static>),
129}
130
131pub 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 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 #[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 reusable_cmsg_space: nix::cmsg_space!(
219 u32, nix::sys::time::TimeSpec, u16, sockaddr_in, sockaddr_in6, u32 ),
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 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 return Ok(());
323 };
324 let Ok(send_permit) = self.accept_sink.try_reserve() else {
325 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 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 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 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 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 return self.poll_recv_from(cx);
460 };
461
462 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 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 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 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 },
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 continue;
607 };
608
609 let _ = mark_bytes.insert(arr);
610 }
611 },
612 _ => {
613 },
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 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 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 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
735fn 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
746fn 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#[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 if self.shutdown_tx.is_some() && self.accept_sink.is_closed() {
792 self.shutdown_tx = None;
793 }
794
795 if self.shutdown_tx.is_none() &&
797 self.shutdown_rx.poll_recv(cx).is_ready()
798 {
799 return Poll::Ready(Ok(()));
800 }
801
802 if let Err(error) = self.incoming_packet_handler.update(cx) {
804 let sender = self.accept_sink.clone();
807 spawn_with_killswitch(async move {
808 let _ = sender.send(Err(error)).await;
809 });
810 }
811
812 let _ = self.poll_conn_map_commands(cx);
816
817 for _ in 0..PACKET_RX_YIELD_AFTER {
820 ready!(self.poll_process_packet(cx));
821 }
822 }
823 }
824}
825
826fn 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
843pub 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
862pub struct NewConnection {
865 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 handshake_start_time: Instant,
874}
875
876#[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 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(¶ms, 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 std::thread::spawn(move || test_connect(host_port, None));
1005
1006 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 time::advance(Duration::new(30, 0)).await;
1020 time::resume();
1021
1022 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 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 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 => (), 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 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(¶ms, 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 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 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 std::thread::sleep(Duration::from_secs(1));
1224
1225 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 let ipr = ipr.join().unwrap();
1238 assert!(ipr.conn_map_cmd_rx.is_empty());
1239 }
1240}