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::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
73const PACKET_RX_YIELD_AFTER: usize = 30;
77const 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 src_addr: SocketAddr,
111 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
120pub enum ConnectionMapCommand {
123 MapCid {
124 existing_cid: ConnectionId<'static>,
125 new_cid: ConnectionId<'static>,
126 },
127 UnmapCid(ConnectionId<'static>),
128}
129
130pub 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 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 #[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 reusable_cmsg_space: nix::cmsg_space!(
218 u32, nix::sys::time::TimeSpec, u16, sockaddr_in, sockaddr_in6, u32 ),
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 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 return Ok(());
321 };
322 let Ok(send_permit) = self.accept_sink.try_reserve() else {
323 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 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 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 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 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 return self.poll_recv_from(cx);
457 };
458
459 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 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 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 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 },
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 continue;
604 };
605
606 let _ = mark_bytes.insert(arr);
607 }
608 },
609 _ => {
610 },
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 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 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 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
732fn 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
743fn 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#[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 if self.shutdown_tx.is_some() && self.accept_sink.is_closed() {
789 self.shutdown_tx = None;
790 }
791
792 if self.shutdown_tx.is_none() &&
794 self.shutdown_rx.poll_recv(cx).is_ready()
795 {
796 return Poll::Ready(Ok(()));
797 }
798
799 if let Err(error) = self.incoming_packet_handler.update(cx) {
801 let sender = self.accept_sink.clone();
804 spawn_with_killswitch(async move {
805 let _ = sender.send(Err(error)).await;
806 });
807 }
808
809 let _ = self.poll_conn_map_commands(cx);
813
814 for _ in 0..PACKET_RX_YIELD_AFTER {
817 ready!(self.poll_process_packet(cx));
818 }
819 }
820 }
821}
822
823fn 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
840pub 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
859pub struct NewConnection {
862 conn: Box<QuicheConnection>,
864 pending_cid: Option<ConnectionId<'static>>,
865 initial_pkt: Option<Incoming>,
866 cid_generator: Option<SharedConnectionIdGenerator>,
867 handshake_start_time: Instant,
870}
871
872#[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 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(¶ms, 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 std::thread::spawn(move || test_connect(host_port));
991
992 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 time::advance(Duration::new(30, 0)).await;
1002 time::resume();
1003
1004 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 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(¶ms, 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 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 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 std::thread::sleep(Duration::from_secs(1));
1094
1095 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 let ipr = ipr.join().unwrap();
1108 assert!(ipr.conn_map_cmd_rx.is_empty());
1109 }
1110}