Skip to main content

tokio_quiche/quic/io/
worker.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
27use std::net::SocketAddr;
28use std::ops::ControlFlow;
29use std::sync::Arc;
30use std::task::Poll;
31use std::time::Duration;
32use std::time::Instant;
33#[cfg(feature = "perf-quic-listener-metrics")]
34use std::time::SystemTime;
35
36use super::connection_stage::Close;
37use super::connection_stage::ConnectionStage;
38use super::connection_stage::ConnectionStageContext;
39use super::connection_stage::Handshake;
40use super::connection_stage::RunningApplication;
41use super::gso::*;
42use super::utilization_estimator::BandwidthReporter;
43
44use crate::metrics::labels;
45use crate::metrics::Metrics;
46use crate::quic::connection::ApplicationOverQuic;
47use crate::quic::connection::HandshakeError;
48use crate::quic::connection::Incoming;
49use crate::quic::connection::QuicConnectionStats;
50use crate::quic::connection::SharedConnectionIdGenerator;
51use crate::quic::hooks::ConnectionHook;
52use crate::quic::router::ConnectionMapCommand;
53use crate::quic::QuicheConnection;
54use crate::QuicResult;
55
56use boring::ssl::SslRef;
57use datagram_socket::DatagramSocketSend;
58use datagram_socket::DatagramSocketSendExt;
59use datagram_socket::MaybeConnectedSocket;
60use datagram_socket::QuicAuditStats;
61use foundations::telemetry::log;
62use quiche::ConnectionId;
63use quiche::Error as QuicheError;
64use quiche::SendInfo;
65use tokio::select;
66use tokio::sync::mpsc;
67use tokio::time;
68
69// Number of incoming packets to be buffered in the incoming channel.
70pub(crate) const INCOMING_QUEUE_SIZE: usize = 2048;
71
72// Check if there are any incoming packets while sending data every this number
73// of sent packets
74pub(crate) const CHECK_INCOMING_QUEUE_RATIO: usize = INCOMING_QUEUE_SIZE / 16;
75
76const RELEASE_TIMER_THRESHOLD: Duration = Duration::from_micros(250);
77
78/// Stop queuing GSO packets, if packet size is below this threshold.
79const GSO_THRESHOLD: usize = 1_000;
80
81/// Size of each full egress buffer borrowed for a send burst.
82///
83/// Matches the maximum quiche buffer size so GSO batching is unaffected while
84/// a connection is actively sending. Unlike a persistent per-connection
85/// buffer, this memory is returned to the per-worker [`SEND_BUF_POOL`] before
86/// the worker sleeps. The free list retains it for reuse (rather than truly
87/// freeing it) but no idle connection owns an egress buffer.
88const SEND_BUFFER_SIZE: usize = crate::buf_factory::BufFactory::MAX_BUF_SIZE;
89
90/// Size of a temporary egress buffer when pooling is disabled.
91///
92/// The cold handshake and connection-close paths generate a single datagram at
93/// a time. When pooling is enabled, they borrow a full-size buffer from
94/// [`SEND_BUF_POOL`]; otherwise a one-MTU buffer is enough.
95const TRANSIENT_SEND_BUFFER_SIZE: usize = 1500;
96
97/// Allocates a zero-initialized egress buffer on the heap.
98///
99/// This is the cold path that fills [`SEND_BUF_POOL`] on a miss; steady-state
100/// bursts borrow a recycled buffer via [`PooledSendBuf::acquire`] and never
101/// hit this. The buffer is boxed (never a stack array) so that holding it
102/// across the `.await` in [`IoWorker::flush_buffer_to_socket`] keeps the
103/// worker's futures small and avoids uncontrolled stack growth.
104fn alloc_send_buffer() -> Box<[u8]> {
105    vec![0u8; SEND_BUFFER_SIZE].into_boxed_slice()
106}
107
108thread_local! {
109    /// Per-runtime-worker free-list of egress scratch buffers.
110    ///
111    /// A buffer is borrowed for a single send burst and returned on drop, so
112    /// its pages stay resident across bursts (no per-burst page fault or
113    /// kernel zero-fill) while no idle connection retains a buffer: the pool
114    /// holds at most [`SEND_BUF_POOL_CAP`] buffers per worker thread,
115    /// independent of the connection count.
116    static SEND_BUF_POOL: std::cell::RefCell<Vec<Box<[u8]>>> =
117        const { std::cell::RefCell::new(Vec::new()) };
118}
119
120/// Upper bound on egress buffers parked per worker thread. The natural
121/// high-water mark is the number of connection tasks simultaneously suspended
122/// at a flush `.await` on one runtime thread; returns beyond the cap are freed
123/// so a burst spike cannot pin unbounded memory to a thread.
124///
125/// This is a fixed per-worker-thread reservation, independent of the
126/// connection count: up to `SEND_BUF_POOL_CAP * SEND_BUFFER_SIZE`
127/// (16 * 64 KiB = 1 MiB) per runtime worker thread. The pool is not shrunk
128/// once grown, so after a burst it stays at its high-water mark for the
129/// process lifetime.
130const SEND_BUF_POOL_CAP: usize = 16;
131
132/// Egress scratch buffer borrowed from the per-thread [`SEND_BUF_POOL`] and
133/// returned to it on drop.
134///
135/// Behaves like the `Box<[u8]>` it replaces via `Deref`/`DerefMut`, so idle
136/// connections still retain no egress buffer, but the backing pages are
137/// recycled instead of re-faulted (and re-zeroed by the kernel) every burst.
138struct PooledSendBuf(Box<[u8]>);
139
140impl PooledSendBuf {
141    fn acquire() -> Self {
142        let buf = SEND_BUF_POOL
143            .with(|pool| pool.borrow_mut().pop())
144            .unwrap_or_else(|| {
145                // Pool miss (cold path): allocate a fresh buffer and record it.
146                // Re-using a parked buffer (the hot path) does not touch any
147                // counter.
148                crate::metrics::quic::send_buffer_pool_allocated().inc();
149                alloc_send_buffer()
150            });
151        // No re-zeroing on reuse: quiche writes only the bytes it emits and
152        // the flush path transmits solely `send_buf[..bytes_written]`, so any
153        // stale bytes left by a previous burst are never sent.
154        Self(buf)
155    }
156}
157
158impl Drop for PooledSendBuf {
159    fn drop(&mut self) {
160        // Move the buffer out, leaving an empty (non-allocating) boxed slice
161        // behind, so it can be returned to the pool. Storing a plain
162        // `Box<[u8]>` rather than an `Option` keeps `Deref`/`DerefMut`
163        // panic-free.
164        let buf = std::mem::take(&mut self.0);
165        // Returns to *this* thread's pool. With tokio work-stealing a task may
166        // migrate across the flush `.await`, so a buffer can be acquired on one
167        // worker and returned on another; this is benign, and the per-thread
168        // cap keeps the total bounded by `cap * workers`.
169        SEND_BUF_POOL.with(|pool| {
170            let mut pool = pool.borrow_mut();
171            if pool.len() < SEND_BUF_POOL_CAP {
172                pool.push(buf);
173            } else {
174                // Pool already at capacity (cold path, only under burst
175                // spikes): drop the buffer instead of parking it, and record
176                // the discard. Returning to a non-full pool (the hot path)
177                // does not touch any counter.
178                crate::metrics::quic::send_buffer_pool_discarded().inc();
179            }
180        });
181    }
182}
183
184impl std::ops::Deref for PooledSendBuf {
185    type Target = [u8];
186
187    fn deref(&self) -> &[u8] {
188        &self.0
189    }
190}
191
192impl std::ops::DerefMut for PooledSendBuf {
193    fn deref_mut(&mut self) -> &mut [u8] {
194        &mut self.0
195    }
196}
197
198/// Egress buffer used briefly outside the main write loop.
199///
200/// This preserves the runtime pooling switch for handshake and close paths:
201/// pooling borrows a full-size recycled buffer, while disabling it allocates a
202/// small one-off buffer.
203enum TransientSendBuf {
204    Pooled(PooledSendBuf),
205    Unpooled(Box<[u8]>),
206}
207
208impl TransientSendBuf {
209    fn acquire(pool_send_buffer: bool) -> Self {
210        if pool_send_buffer {
211            Self::Pooled(PooledSendBuf::acquire())
212        } else {
213            Self::Unpooled(
214                vec![0u8; TRANSIENT_SEND_BUFFER_SIZE].into_boxed_slice(),
215            )
216        }
217    }
218}
219
220impl AsRef<[u8]> for TransientSendBuf {
221    fn as_ref(&self) -> &[u8] {
222        match self {
223            Self::Pooled(buf) => &buf[..],
224            Self::Unpooled(buf) => &buf[..],
225        }
226    }
227}
228
229impl AsMut<[u8]> for TransientSendBuf {
230    fn as_mut(&mut self) -> &mut [u8] {
231        match self {
232            Self::Pooled(buf) => &mut buf[..],
233            Self::Unpooled(buf) => &mut buf[..],
234        }
235    }
236}
237
238pub struct WriterConfig {
239    pub pending_cid: Option<ConnectionId<'static>>,
240    pub peer_addr: SocketAddr,
241    pub local_addr: SocketAddr,
242    pub with_gso: bool,
243    pub pacing_offload: bool,
244    pub with_pktinfo: bool,
245    /// Whether the worker borrows its egress buffer from a per-worker-thread
246    /// pool for each send burst. When `false`, the worker keeps a persistent
247    /// per-connection buffer for its lifetime instead.
248    pub pool_send_buffer: bool,
249}
250
251#[derive(Default)]
252pub(crate) struct WriteState {
253    conn_established: bool,
254    bytes_written: usize,
255    segment_size: usize,
256    num_pkts: usize,
257    tx_time: Option<Instant>,
258    has_pending_data: bool,
259    // If pacer schedules packets too far into the future, we want to pause
260    // sending, until the future arrives
261    next_release_time: Option<Instant>,
262    // The selected source and destination addresses for the current write
263    // cycle.
264    selected_path: Option<(SocketAddr, SocketAddr)>,
265    // Iterator over the network paths that haven't been flushed yet.
266    pending_paths: quiche::SocketAddrIter,
267}
268
269pub(crate) struct IoWorkerParams<Tx, M> {
270    pub(crate) socket: MaybeConnectedSocket<Tx>,
271    pub(crate) shutdown_tx: mpsc::Sender<()>,
272    pub(crate) cfg: WriterConfig,
273    pub(crate) audit_log_stats: Arc<QuicAuditStats>,
274    pub(crate) write_state: WriteState,
275    pub(crate) conn_map_cmd_tx: mpsc::UnboundedSender<ConnectionMapCommand>,
276    pub(crate) cid_generator: Option<SharedConnectionIdGenerator>,
277    #[cfg(feature = "perf-quic-listener-metrics")]
278    pub(crate) init_rx_time: Option<SystemTime>,
279    pub(crate) metrics: M,
280}
281
282fn notify_path_events(
283    connection_hook: Option<&(dyn ConnectionHook + Send + Sync + 'static)>,
284    qconn: &mut QuicheConnection,
285) {
286    while let Some(path_event) = qconn.path_event_next() {
287        if let Some(hook) = connection_hook {
288            hook.on_path_event(qconn, &path_event);
289        }
290    }
291}
292
293pub(crate) struct IoWorker<Tx, M, S> {
294    socket: MaybeConnectedSocket<Tx>,
295    /// A field that signals to the listener task that the connection has gone
296    /// away (nothing is sent here, listener task just detects the sender
297    /// has dropped)
298    shutdown_tx: mpsc::Sender<()>,
299    cfg: WriterConfig,
300    audit_log_stats: Arc<QuicAuditStats>,
301    write_state: WriteState,
302    conn_map_cmd_tx: mpsc::UnboundedSender<ConnectionMapCommand>,
303    cid_generator: Option<SharedConnectionIdGenerator>,
304    #[cfg(feature = "perf-quic-listener-metrics")]
305    init_rx_time: Option<SystemTime>,
306    metrics: M,
307    conn_stage: S,
308    bw_estimator: BandwidthReporter,
309}
310
311impl<Tx, M, S> IoWorker<Tx, M, S>
312where
313    Tx: DatagramSocketSend + Send,
314    M: Metrics,
315    S: ConnectionStage,
316{
317    pub(crate) fn new(params: IoWorkerParams<Tx, M>, conn_stage: S) -> Self {
318        let bw_estimator =
319            BandwidthReporter::new(params.metrics.utilized_bandwidth());
320
321        log::trace!("Creating IoWorker with stage: {conn_stage:?}");
322
323        Self {
324            socket: params.socket,
325            shutdown_tx: params.shutdown_tx,
326            cfg: params.cfg,
327            audit_log_stats: params.audit_log_stats,
328            write_state: params.write_state,
329            conn_map_cmd_tx: params.conn_map_cmd_tx,
330            cid_generator: params.cid_generator,
331            #[cfg(feature = "perf-quic-listener-metrics")]
332            init_rx_time: params.init_rx_time,
333            metrics: params.metrics,
334            conn_stage,
335            bw_estimator,
336        }
337    }
338
339    fn fill_available_scids(&self, qconn: &mut QuicheConnection) {
340        if qconn.scids_left() == 0 {
341            return;
342        }
343        let Some(cid_generator) = self.cid_generator.as_deref() else {
344            return;
345        };
346
347        let current_cid = qconn.source_id().into_owned();
348        for _ in 0..qconn.scids_left() {
349            // We don't emit stateless resets, so any unguessable value is fine
350            let reset_token = random_u128();
351            let new_cid = cid_generator.new_connection_id();
352
353            if self
354                .conn_map_cmd_tx
355                .send(ConnectionMapCommand::MapCid {
356                    existing_cid: current_cid.clone(),
357                    new_cid: new_cid.clone(),
358                })
359                .is_err()
360            {
361                // Can't do anything if the connection map is gone
362                return;
363            }
364
365            if qconn.new_scid(&new_cid, reset_token, false).is_err() {
366                // This only fails if we have reached the CID limit already
367                return;
368            }
369        }
370    }
371
372    fn unmap_cid(&self, cid: ConnectionId<'static>) {
373        // If the connection map is gone, the ID is already "unmapped"
374        let _ = self
375            .conn_map_cmd_tx
376            .send(ConnectionMapCommand::UnmapCid(cid));
377    }
378
379    fn refresh_connection_ids(&self, qconn: &mut QuicheConnection) {
380        // Top up the connection's active CIDs
381        self.fill_available_scids(qconn);
382
383        // Remove retired CIDs from the ingress router
384        while let Some(retired_cid) = qconn.retired_scid_next() {
385            self.unmap_cid(retired_cid);
386        }
387    }
388
389    async fn work_loop<A: ApplicationOverQuic>(
390        &mut self, qconn: &mut QuicheConnection,
391        ctx: &mut ConnectionStageContext<A>,
392    ) -> QuicResult<()> {
393        const DEFAULT_SLEEP: Duration = Duration::from_secs(60);
394        let mut current_deadline: Option<Instant> = None;
395        let sleep = time::sleep(DEFAULT_SLEEP);
396        tokio::pin!(sleep);
397
398        // Without pooling, the IO worker owns one egress buffer for the entire
399        // connection, including idle periods. With pooling, this remains `None`
400        // and each send burst borrows a buffer from the per-worker pool.
401        let mut persistent_send_buf: Option<Box<[u8]>> =
402            (!self.cfg.pool_send_buffer).then(alloc_send_buffer);
403
404        loop {
405            let now = Instant::now();
406
407            self.write_state.has_pending_data = true;
408
409            // Transient egress buffer for this wakeup's send burst when pooling
410            // is enabled. Borrowed from the per-worker pool on demand (see
411            // below) and returned after the burst, before the worker sleeps in
412            // the `select!` further down, so idle connections still hold no
413            // egress buffer. Stays `None` when a persistent buffer is used.
414            let mut pooled_send_buf: Option<PooledSendBuf> = None;
415
416            while self.write_state.has_pending_data {
417                let mut packets_sent = 0;
418
419                // Drain received packets periodically because they contain ACKs
420                // and the bounded receive queue stalls new packets when full.
421                let mut did_recv = false;
422                while let Some(pkt) = ctx
423                    .in_pkt
424                    .take()
425                    .or_else(|| ctx.incoming_pkt_receiver.try_recv().ok())
426                {
427                    self.process_incoming(qconn, pkt)?;
428                    did_recv = true;
429                }
430
431                // Deliver transport-generated events before the application can
432                // consume them.
433                notify_path_events(ctx.connection_hook.as_deref(), qconn);
434
435                self.conn_stage.on_read(did_recv, qconn, ctx)?;
436
437                // Deliver events generated by application reads or writes.
438                notify_path_events(ctx.connection_hook.as_deref(), qconn);
439                self.refresh_connection_ids(qconn);
440
441                let can_release = match self.write_state.next_release_time {
442                    None => true,
443                    Some(next_release) =>
444                        next_release
445                            .checked_duration_since(now)
446                            .unwrap_or_default() <
447                            RELEASE_TIMER_THRESHOLD,
448                };
449
450                self.write_state.has_pending_data &= can_release;
451
452                while self.write_state.has_pending_data &&
453                    packets_sent < CHECK_INCOMING_QUEUE_RATIO
454                {
455                    // Use the persistent per-connection buffer when pooling is
456                    // disabled. Otherwise borrow a buffer from the per-worker
457                    // pool on the first gather of this send burst and reuse it
458                    // for the remainder of the burst; it is returned at the end
459                    // of the enclosing block (below), before the worker sleeps,
460                    // so idle connections hold no egress buffer.
461                    let send_buf: &mut [u8] =
462                        if let Some(buf) = persistent_send_buf.as_deref_mut() {
463                            buf
464                        } else {
465                            &mut pooled_send_buf
466                                .get_or_insert_with(PooledSendBuf::acquire)[..]
467                        };
468
469                    self.gather_data_from_quiche_conn(qconn, send_buf, false)?;
470
471                    // Break if the connection is closed
472                    if qconn.is_closed() {
473                        return Ok(());
474                    }
475
476                    let mut flush_operation_token =
477                        TrackMidHandshakeFlush::new(self.metrics.clone());
478
479                    self.flush_buffer_to_socket(&send_buf[..]).await;
480
481                    flush_operation_token.mark_complete();
482
483                    packets_sent += self.write_state.num_pkts;
484
485                    if let ControlFlow::Break(reason) =
486                        self.conn_stage.on_flush(qconn, ctx)
487                    {
488                        return reason;
489                    }
490                }
491            }
492
493            // Return the borrowed egress buffer to the per-worker pool before
494            // sleeping so it is not held while the connection is idle. The
495            // persistent buffer (when pooling is disabled) is intentionally
496            // kept across sleeps for the connection's lifetime.
497            drop(pooled_send_buf);
498
499            self.bw_estimator.update(qconn, now);
500
501            self.audit_log_stats
502                .set_max_bandwidth(self.bw_estimator.max_bandwidth);
503            self.audit_log_stats.set_max_loss_pct(
504                (self.bw_estimator.max_loss_pct * 100_f32).round() as u8,
505            );
506
507            let new_deadline = min_of_some(
508                qconn.timeout_instant(),
509                self.write_state.next_release_time,
510            );
511            let new_deadline =
512                min_of_some(new_deadline, self.conn_stage.wait_deadline());
513
514            if new_deadline != current_deadline {
515                current_deadline = new_deadline;
516
517                sleep
518                    .as_mut()
519                    .reset(new_deadline.unwrap_or(now + DEFAULT_SLEEP).into());
520            }
521
522            let incoming_recv = &mut ctx.incoming_pkt_receiver;
523            let application = &mut ctx.application;
524
525            select! {
526                biased;
527                () = &mut sleep => {
528                    // It's very important that we keep the timeout arm at the top of this loop so
529                    // that we poll it every time we need to. Since this is a biased `select!`, if
530                    // we put this behind another arm, we could theoretically starve the sleep arm
531                    // and hang connections.
532                    //
533                    // See https://docs.rs/tokio/latest/tokio/macro.select.html#fairness for more
534                    qconn.on_timeout();
535
536                    self.write_state.next_release_time = None;
537                    current_deadline = None;
538                    sleep.as_mut().reset((now + DEFAULT_SLEEP).into());
539                }
540                Some(pkt) = incoming_recv.recv() => ctx.in_pkt = Some(pkt),
541                directive = self.wait_for_data_or_handshake(qconn, application) => {
542                    match directive? {
543                        WaitForDataOrHandshakeDirective::Flush(send_buf) => {
544                            // The handshake data was gathered into this
545                            // on-demand buffer; flush it here (outside the
546                            // select! so the flush cannot be cancelled), then
547                            // return it to the pool or drop it.
548                            self.flush_buffer_to_socket(send_buf.as_ref()).await;
549                        }
550                        WaitForDataOrHandshakeDirective::Noop => {}
551                    }
552                },
553            };
554
555            if let ControlFlow::Break(reason) = self.conn_stage.post_wait(qconn) {
556                return reason;
557            }
558        }
559    }
560
561    #[cfg(feature = "perf-quic-listener-metrics")]
562    fn measure_complete_handshake_time(&mut self) {
563        if let Some(init_rx_time) = self.init_rx_time.take() {
564            if let Ok(delta) = init_rx_time.elapsed() {
565                self.metrics
566                    .handshake_time_seconds(
567                        labels::QuicHandshakeStage::HandshakeResponse,
568                    )
569                    .observe(delta.as_nanos() as u64);
570            }
571        }
572    }
573
574    /// Gathers one or more packets from quiche into `send_buf`.
575    ///
576    /// A single-packet gather leaves a full buffer available for the next
577    /// packet instead of generating a short packet in the remaining tail.
578    fn gather_data_from_quiche_conn(
579        &mut self, qconn: &mut QuicheConnection, send_buf: &mut [u8],
580        single_packet: bool,
581    ) -> QuicResult<usize> {
582        let mut segment_size = None;
583        let mut send_info = None;
584
585        self.write_state.num_pkts = 0;
586        self.write_state.bytes_written = 0;
587
588        self.write_state.selected_path = None;
589
590        let now = Instant::now();
591
592        let send_buf = {
593            let trunc = UDP_MAX_GSO_PACKET_SIZE.min(send_buf.len());
594            &mut send_buf[..trunc]
595        };
596
597        #[cfg(feature = "gcongestion")]
598        let gcongestion_enabled = true;
599
600        #[cfg(not(feature = "gcongestion"))]
601        let gcongestion_enabled = qconn.gcongestion_enabled().unwrap_or(false);
602
603        let initial_release_decision = if gcongestion_enabled {
604            let initial_release_decision = qconn
605                .get_next_release_time()
606                .filter(|_| self.pacing_enabled(qconn));
607
608            if let Some(future_release_time) =
609                initial_release_decision.as_ref().and_then(|v| v.time(now))
610            {
611                let max_into_fut = qconn.max_release_into_future();
612
613                if future_release_time.duration_since(now) >= max_into_fut {
614                    self.write_state.next_release_time =
615                        Some(now + max_into_fut.mul_f32(0.8));
616                    self.write_state.has_pending_data = false;
617                    return Ok(0);
618                }
619            }
620
621            initial_release_decision
622        } else {
623            None
624        };
625
626        let buffer_write_outcome = loop {
627            let outcome = self.write_packet_to_buffer(
628                qconn,
629                send_buf,
630                &mut send_info,
631                segment_size,
632            );
633
634            let packet_size = match outcome {
635                Ok(0) => break Ok(0),
636
637                Ok(bytes_written) => bytes_written,
638
639                Err(e) => break Err(e),
640            };
641
642            // Flush after one packet when GSO is disabled or the caller needs
643            // each packet to start with a full buffer.
644            if single_packet || !self.cfg.with_gso {
645                break outcome;
646            }
647
648            #[cfg(not(feature = "gcongestion"))]
649            let max_send_size = if !gcongestion_enabled {
650                // Only call qconn.send_quantum when !gcongestion_enabled.
651                tune_max_send_size(
652                    segment_size,
653                    qconn.send_quantum(),
654                    send_buf.len(),
655                )
656            } else {
657                usize::MAX
658            };
659
660            #[cfg(feature = "gcongestion")]
661            let max_send_size = usize::MAX;
662
663            // If segment_size is known, update the maximum of
664            // GSO sender buffer size to the multiple of
665            // segment_size.
666            let buffer_is_full = self.write_state.num_pkts ==
667                UDP_MAX_SEGMENT_COUNT ||
668                self.write_state.bytes_written >= max_send_size;
669
670            if buffer_is_full {
671                break outcome;
672            }
673
674            // Flush to network when the newly generated packet size is
675            // different from previously written packet, as GSO needs packets
676            // to have the same size, except for the last one in the buffer.
677            // The last packet may be smaller than the previous size.
678            match segment_size {
679                Some(size)
680                    if packet_size != size || packet_size < GSO_THRESHOLD =>
681                    break outcome,
682                None => segment_size = Some(packet_size),
683                _ => (),
684            }
685
686            if gcongestion_enabled {
687                // Start a new batch if the next packet has a different release
688                // time or cannot be part of a burst.
689                if let Some(initial_release_decision) = initial_release_decision {
690                    match qconn.get_next_release_time() {
691                        Some(release)
692                            if release.can_burst() ||
693                                release.time_eq(
694                                    &initial_release_decision,
695                                    now,
696                                ) => {},
697                        _ => break outcome,
698                    }
699                }
700            }
701        };
702
703        let tx_time = if gcongestion_enabled {
704            initial_release_decision
705                .filter(|_| self.pacing_enabled(qconn))
706                // Return the time from the release decision if release_decision.time > now, else None.
707                .and_then(|v| v.time(now))
708        } else {
709            send_info
710                .filter(|_| self.pacing_enabled(qconn))
711                .map(|v| v.at)
712        };
713
714        self.write_state.conn_established = qconn.is_established();
715        self.write_state.tx_time = tx_time;
716        self.write_state.segment_size =
717            segment_size.unwrap_or(self.write_state.bytes_written);
718
719        if !gcongestion_enabled {
720            if let Some(time) = tx_time {
721                const DEFAULT_MAX_INTO_FUTURE: Duration =
722                    Duration::from_millis(1);
723                if time
724                    .checked_duration_since(now)
725                    .map(|d| d > DEFAULT_MAX_INTO_FUTURE)
726                    .unwrap_or(false)
727                {
728                    self.write_state.next_release_time =
729                        Some(now + DEFAULT_MAX_INTO_FUTURE.mul_f32(0.8));
730                    self.write_state.has_pending_data = false;
731                    return Ok(0);
732                }
733            }
734        }
735
736        buffer_write_outcome
737    }
738
739    /// Selects a network path, if none already selected.
740    ///
741    /// This will return the first path available in the write state's
742    /// `pending_paths` iterator. If that is empty a new iterator will be
743    /// created by querying quiche itself.
744    ///
745    /// Note that the connection's statically configured local address will be
746    /// used to query quiche for available paths, so this can't handle multiple
747    /// local addresses currently.
748    fn select_path(
749        &mut self, qconn: &QuicheConnection,
750    ) -> Option<(SocketAddr, SocketAddr)> {
751        if self.write_state.selected_path.is_some() {
752            return self.write_state.selected_path;
753        }
754
755        let from = self.cfg.local_addr;
756
757        // Initialize paths iterator.
758        if self.write_state.pending_paths.len() == 0 {
759            self.write_state.pending_paths = qconn.paths_iter(from);
760        }
761
762        let to = self.write_state.pending_paths.next()?;
763
764        Some((from, to))
765    }
766
767    #[cfg(not(feature = "gcongestion"))]
768    fn pacing_enabled(&self, qconn: &QuicheConnection) -> bool {
769        self.cfg.pacing_offload && qconn.pacing_enabled()
770    }
771
772    #[cfg(feature = "gcongestion")]
773    fn pacing_enabled(&self, _qconn: &QuicheConnection) -> bool {
774        self.cfg.pacing_offload
775    }
776
777    fn write_packet_to_buffer(
778        &mut self, qconn: &mut QuicheConnection, send_buf: &mut [u8],
779        send_info: &mut Option<SendInfo>, segment_size: Option<usize>,
780    ) -> QuicResult<usize> {
781        let mut send_buf = &mut send_buf[self.write_state.bytes_written..];
782        if send_buf.len() > segment_size.unwrap_or(usize::MAX) {
783            // Never let the buffer be longer than segment size, for GSO to
784            // function properly.
785            send_buf = &mut send_buf[..segment_size.unwrap_or(usize::MAX)];
786        }
787
788        // On the first call to `select_path()` a path will be chosen based on
789        // the local address the connection initially landed on. Once a path is
790        // selected following calls to `select_path()` will return it, until it
791        // is reset at the start of the next write cycle.
792        //
793        // The path is then passed to `send_on_path()` which will only generate
794        // packets meant for that path, this way a single GSO buffer will only
795        // contain packets that belong to the same network path, which is
796        // required because the from/to addresses for each `sendmsg()` call
797        // apply to the whole GSO buffer.
798        let (from, to) = self.select_path(qconn).unzip();
799
800        match qconn.send_on_path(send_buf, from, to) {
801            Ok((packet_size, info)) => {
802                let _ = send_info.get_or_insert(info);
803
804                self.write_state.bytes_written += packet_size;
805                self.write_state.num_pkts += 1;
806
807                let from = send_info.as_ref().map(|info| info.from);
808                let to = send_info.as_ref().map(|info| info.to);
809
810                self.write_state.selected_path = from.zip(to);
811
812                self.write_state.has_pending_data = true;
813
814                Ok(packet_size)
815            },
816
817            Err(QuicheError::Done) => {
818                // Flush the current buffer to network. If no other path needs
819                // to be flushed to the network also yield the work loop task.
820                //
821                // Otherwise the write loop will start again and the next path
822                // will be selected.
823                let has_pending_paths = self.write_state.pending_paths.len() > 0;
824
825                // Keep writing if there are paths left to try.
826                self.write_state.has_pending_data = has_pending_paths;
827
828                Ok(0)
829            },
830
831            Err(e) => {
832                let error_code = if let Some(local_error) = qconn.local_error() {
833                    local_error.error_code
834                } else {
835                    let internal_error_code =
836                        quiche::WireErrorCode::InternalError as u64;
837                    let _ = qconn.close(false, internal_error_code, &[]);
838
839                    internal_error_code
840                };
841
842                self.audit_log_stats
843                    .set_sent_conn_close_transport_error_code(error_code as i64);
844
845                Err(Box::new(e))
846            },
847        }
848    }
849
850    async fn flush_buffer_to_socket(&mut self, send_buf: &[u8]) {
851        if self.write_state.bytes_written > 0 {
852            let current_send_buf = &send_buf[..self.write_state.bytes_written];
853
854            let (from, to) = self.write_state.selected_path.unzip();
855
856            let to = to.unwrap_or(self.cfg.peer_addr);
857            let from = from.filter(|_| self.cfg.with_pktinfo);
858
859            let send_res = if let (Some(udp_socket), true) =
860                (self.socket.as_udp_socket(), self.cfg.with_gso)
861            {
862                // Only UDP supports GSO.
863                send_to(
864                    udp_socket,
865                    to,
866                    from,
867                    current_send_buf,
868                    self.write_state.segment_size,
869                    self.write_state.tx_time,
870                    self.metrics
871                        .write_errors(labels::QuicWriteError::WouldBlock),
872                    self.metrics.send_to_wouldblock_duration_s(),
873                )
874                .await
875            } else {
876                self.socket.send_to(current_send_buf, to).await
877            };
878
879            #[cfg(feature = "perf-quic-listener-metrics")]
880            self.measure_complete_handshake_time();
881
882            match send_res {
883                Ok(n) =>
884                    if n < self.write_state.bytes_written {
885                        self.metrics
886                            .write_errors(labels::QuicWriteError::Partial)
887                            .inc();
888                    },
889
890                Err(_) => {
891                    self.metrics.write_errors(labels::QuicWriteError::Err).inc();
892                },
893            }
894        }
895    }
896
897    /// Process the incoming packet
898    fn process_incoming(
899        &mut self, qconn: &mut QuicheConnection, mut pkt: Incoming,
900    ) -> QuicResult<()> {
901        let recv_info = quiche::RecvInfo {
902            from: pkt.peer_addr,
903            to: pkt.local_addr,
904        };
905
906        if let Some(gro) = pkt.gro {
907            for dgram in pkt.buf.chunks_mut(gro as usize) {
908                qconn.recv(dgram, recv_info)?;
909            }
910        } else {
911            qconn.recv(&mut pkt.buf, recv_info)?;
912        }
913
914        Ok(())
915    }
916
917    // Process application data after establishment. Before then, a BoringSSL
918    // wakeup might require quiche to send handshake packets.
919    //
920    // TODO(erittenhouse): Decouple `wait_for_data` from the application.
921    // `wait_for_quiche` depends on IOW methods, preventing a default
922    // `ConnectionStage` implementation.
923    //
924    // # Cancel safety
925    //
926    // This future is polled as an arm of the `select!` in `Self::work_loop`, so
927    // it must be cancel-safe. Another arm may complete first and drop it at any
928    // `.await`. `ApplicationOverQuic::wait_for_data` is also cancel-safe, and
929    // the handshake branch retains only its local `send_buf` across `.await`.
930    // Cancellation returns or frees that buffer. The next poll gathers its
931    // bytes again. Preserve this property when modifying the function.
932    async fn wait_for_data_or_handshake<A: ApplicationOverQuic>(
933        &mut self, qconn: &mut QuicheConnection, quic_application: &mut A,
934    ) -> QuicResult<WaitForDataOrHandshakeDirective> {
935        if quic_application.should_act() {
936            // Poll the application to make progress.
937            //
938            // Once the connection has been established (i.e. the handshake is
939            // complete), we only poll the application.
940            //
941            // The exception is 0-RTT in TLS 1.3, where the full handshake is
942            // still in progress but we have 0-RTT keys to process early data.
943            // This means TLS callbacks might only be polled on the next timeout
944            // or when a packet is received from the peer.
945            quic_application.wait_for_data(qconn).await?;
946            Ok(WaitForDataOrHandshakeDirective::Noop)
947        } else {
948            // Poll quiche to make progress on handshake callbacks, gathering
949            // any handshake packets into an on-demand buffer that the caller
950            // flushes. `wait_for_quiche()` returns it only after generating a
951            // packet, so pending handshake waits do not retain a buffer.
952            let send_buf = self.wait_for_quiche(qconn).await?;
953            Ok(WaitForDataOrHandshakeDirective::Flush(send_buf))
954        }
955    }
956
957    /// Check if Quiche has any packets to send
958    ///
959    /// If yes: fills buffer and updates self.write_state.bytes_written
960    /// If no: Poll::Pending
961    ///
962    /// # Example
963    ///
964    /// This function can be used, for example, to drive an asynchronous TLS
965    /// handshake. Each call to `gather_data_from_quiche_conn` attempts to
966    /// progress the handshake via a call to `quiche::Connection.send()` -
967    /// once one of the `gather_data_from_quiche_conn()` calls writes to the
968    /// send buffer, we signal to the caller which has to take care of flushing
969    ///
970    /// # Cancel safety
971    ///
972    /// This future is awaited (indirectly) as an arm of the `select!` in
973    /// [`Self::work_loop`], so it MUST be cancel safe. The `poll_fn` below
974    /// holds no state across polls other than what lives in `self.write_state`,
975    /// so dropping the future between polls loses nothing: the next call simply
976    /// re-gathers. Take care to preserve this property when modifying it.
977    async fn wait_for_quiche(
978        &mut self, qconn: &mut QuicheConnection,
979    ) -> QuicResult<TransientSendBuf> {
980        let send_buf = std::future::poll_fn(|_| {
981            // Allocate inside this closure so a pending poll immediately
982            // returns its buffer to the pool instead of retaining it across
983            // the select! wait.
984            let mut send_buf =
985                TransientSendBuf::acquire(self.cfg.pool_send_buffer);
986
987            match self.gather_data_from_quiche_conn(
988                qconn,
989                send_buf.as_mut(),
990                true,
991            ) {
992                Ok(bytes_written) => {
993                    // Do not call `gather()` twice without an intervening
994                    // `flush()`. Consecutive calls may overwrite data or delay
995                    // handshake completion.
996                    if bytes_written == 0 && self.write_state.bytes_written == 0 {
997                        Poll::Pending
998                    } else {
999                        Poll::Ready(Ok(send_buf))
1000                    }
1001                },
1002                _ => Poll::Ready(Err(quiche::Error::TlsFail)),
1003            }
1004        })
1005        .await?;
1006        Ok(send_buf)
1007    }
1008}
1009
1010/// Whether caller of [`wait_for_data_or_handshake`] is required to
1011/// call [`flush_buffer_to_socket`].
1012///
1013/// `Flush` carries the on-demand buffer the handshake data was gathered into so
1014/// the caller can flush it and then return it to the pool or drop it.
1015#[must_use]
1016enum WaitForDataOrHandshakeDirective {
1017    Noop,
1018    Flush(TransientSendBuf),
1019}
1020
1021pub struct Running<Tx, M, A> {
1022    pub(crate) params: IoWorkerParams<Tx, M>,
1023    pub(crate) context: ConnectionStageContext<A>,
1024    /// See [`QuicConnectionParams::quiche_conn`].
1025    pub(crate) qconn: Box<QuicheConnection>,
1026}
1027
1028impl<Tx, M, A> Running<Tx, M, A> {
1029    pub fn ssl(&mut self) -> &mut SslRef {
1030        // Deref to pick `Connection::as_mut` over `Box::as_mut`.
1031        (*self.qconn).as_mut()
1032    }
1033}
1034
1035pub(crate) struct Closing<Tx, M, A> {
1036    pub(crate) params: IoWorkerParams<Tx, M>,
1037    pub(crate) context: ConnectionStageContext<A>,
1038    pub(crate) work_loop_result: QuicResult<()>,
1039    /// See [`QuicConnectionParams::quiche_conn`].
1040    pub(crate) qconn: Box<QuicheConnection>,
1041}
1042
1043pub enum RunningOrClosing<Tx, M, A> {
1044    Running(Running<Tx, M, A>),
1045    Closing(Closing<Tx, M, A>),
1046}
1047
1048impl<Tx, M> IoWorker<Tx, M, Handshake>
1049where
1050    Tx: DatagramSocketSend + Send,
1051    M: Metrics,
1052{
1053    pub(crate) async fn run<A>(
1054        mut self, mut qconn: Box<QuicheConnection>,
1055        mut ctx: ConnectionStageContext<A>,
1056    ) -> RunningOrClosing<Tx, M, A>
1057    where
1058        A: ApplicationOverQuic,
1059    {
1060        // The `ex_data` waker must remain stable for this task. Moving a future
1061        // with an async callback to another task leaves a stale waker that
1062        // wakes the wrong task.
1063        std::future::poll_fn(|cx| {
1064            // Deref to pick `Connection::as_mut` over `Box::as_mut`.
1065            let ssl = (*qconn).as_mut();
1066            ssl.set_task_waker(Some(cx.waker().clone()));
1067
1068            Poll::Ready(())
1069        })
1070        .await;
1071
1072        #[cfg(target_os = "linux")]
1073        if let Some(incoming) = ctx.in_pkt.as_mut() {
1074            self.audit_log_stats
1075                .set_initial_so_mark_data(incoming.so_mark_data.take());
1076        }
1077
1078        let mut work_loop_result = self.work_loop(&mut qconn, &mut ctx).await;
1079        notify_path_events(ctx.connection_hook.as_deref(), &mut qconn);
1080        if work_loop_result.is_ok() && qconn.is_closed() {
1081            work_loop_result = Err(HandshakeError::ConnectionClosed.into());
1082        }
1083
1084        if let Err(err) = &work_loop_result {
1085            self.metrics.failed_handshakes(err.into()).inc();
1086
1087            return RunningOrClosing::Closing(Closing {
1088                params: self.into(),
1089                context: ctx,
1090                work_loop_result,
1091                qconn,
1092            });
1093        };
1094
1095        let on_conn_established_result =
1096            self.on_conn_established(&mut qconn, &mut ctx.application);
1097        notify_path_events(ctx.connection_hook.as_deref(), &mut qconn);
1098
1099        match on_conn_established_result {
1100            Ok(()) => RunningOrClosing::Running(Running {
1101                params: self.into(),
1102                context: ctx,
1103                qconn,
1104            }),
1105            Err(e) => {
1106                foundations::telemetry::log::warn!(
1107                    "Handshake stage on_connection_established failed"; "error"=>%e
1108                );
1109
1110                RunningOrClosing::Closing(Closing {
1111                    params: self.into(),
1112                    context: ctx,
1113                    work_loop_result,
1114                    qconn,
1115                })
1116            },
1117        }
1118    }
1119
1120    fn on_conn_established<App: ApplicationOverQuic>(
1121        &mut self, qconn: &mut QuicheConnection, driver: &mut App,
1122    ) -> QuicResult<()> {
1123        // Only calculate the QUIC handshake duration and call the driver's
1124        // on_conn_established hook if this is the first time
1125        // is_established == true.
1126        if self.audit_log_stats.transport_handshake_duration_us() == -1 {
1127            self.conn_stage.handshake_info.set_elapsed();
1128            let handshake_info = &self.conn_stage.handshake_info;
1129
1130            self.audit_log_stats
1131                .set_transport_handshake_duration(handshake_info.elapsed());
1132
1133            driver.on_conn_established(qconn, handshake_info)?;
1134        }
1135
1136        if let Some(cid) = self.cfg.pending_cid.take() {
1137            self.unmap_cid(cid);
1138        }
1139
1140        Ok(())
1141    }
1142}
1143
1144impl<Tx, M, S> From<IoWorker<Tx, M, S>> for IoWorkerParams<Tx, M> {
1145    fn from(value: IoWorker<Tx, M, S>) -> Self {
1146        Self {
1147            socket: value.socket,
1148            shutdown_tx: value.shutdown_tx,
1149            cfg: value.cfg,
1150            audit_log_stats: value.audit_log_stats,
1151            write_state: value.write_state,
1152            conn_map_cmd_tx: value.conn_map_cmd_tx,
1153            cid_generator: value.cid_generator,
1154            #[cfg(feature = "perf-quic-listener-metrics")]
1155            init_rx_time: value.init_rx_time,
1156            metrics: value.metrics,
1157        }
1158    }
1159}
1160
1161impl<Tx, M> IoWorker<Tx, M, RunningApplication>
1162where
1163    Tx: DatagramSocketSend + Send,
1164    M: Metrics,
1165{
1166    pub(crate) async fn run<A: ApplicationOverQuic>(
1167        mut self, mut qconn: Box<QuicheConnection>,
1168        mut ctx: ConnectionStageContext<A>,
1169    ) -> Closing<Tx, M, A> {
1170        // Perform a single call to process_reads()/process_writes(),
1171        // unconditionally, to ensure that any application data (e.g.
1172        // STREAM frames or datagrams) processed by the Handshake
1173        // stage are properly passed to the application.
1174        let on_read_result = self.conn_stage.on_read(true, &mut qconn, &mut ctx);
1175        notify_path_events(ctx.connection_hook.as_deref(), &mut qconn);
1176        if let Err(e) = on_read_result {
1177            return Closing {
1178                params: self.into(),
1179                context: ctx,
1180                work_loop_result: Err(e),
1181                qconn,
1182            };
1183        };
1184
1185        let work_loop_result = self.work_loop(&mut qconn, &mut ctx).await;
1186        notify_path_events(ctx.connection_hook.as_deref(), &mut qconn);
1187
1188        Closing {
1189            params: self.into(),
1190            context: ctx,
1191            work_loop_result,
1192            qconn,
1193        }
1194    }
1195}
1196
1197impl<Tx, M> IoWorker<Tx, M, Close>
1198where
1199    Tx: DatagramSocketSend + Send,
1200    M: Metrics,
1201{
1202    pub(crate) async fn close<A: ApplicationOverQuic>(
1203        mut self, qconn: &mut QuicheConnection,
1204        ctx: &mut ConnectionStageContext<A>,
1205    ) {
1206        if self.conn_stage.work_loop_result.is_ok() &&
1207            self.bw_estimator.max_bandwidth > 0
1208        {
1209            let metrics = &self.metrics;
1210
1211            metrics
1212                .max_bandwidth_mbps()
1213                .observe(self.bw_estimator.max_bandwidth as f64 * 1e-6);
1214
1215            metrics
1216                .max_loss_pct()
1217                .observe(self.bw_estimator.max_loss_pct as f64 * 100.);
1218        }
1219
1220        if ctx.application.should_act() {
1221            ctx.application.on_conn_close(
1222                qconn,
1223                &self.metrics,
1224                &self.conn_stage.work_loop_result,
1225            );
1226            notify_path_events(ctx.connection_hook.as_deref(), qconn);
1227        }
1228
1229        // TODO: this assumes that the tidy_up operation can be completed in one
1230        // send (ignoring flow/congestion control constraints). We should
1231        // guarantee that it gets sent by doublechecking the
1232        // gathered/flushed byte totals and retry if they don't match.
1233        //
1234        // This runs once per connection at close and sends a single
1235        // CONNECTION_CLOSE datagram, so acquire a buffer only for this send.
1236        let mut send_buf = TransientSendBuf::acquire(self.cfg.pool_send_buffer);
1237        let _ =
1238            self.gather_data_from_quiche_conn(qconn, send_buf.as_mut(), false);
1239        self.flush_buffer_to_socket(send_buf.as_ref()).await;
1240
1241        *ctx.stats.lock().unwrap() = QuicConnectionStats::from_conn(qconn);
1242
1243        if let Some(err) = qconn.peer_error() {
1244            if err.is_app {
1245                self.audit_log_stats
1246                    .set_recvd_conn_close_application_error_code(
1247                        err.error_code as _,
1248                    );
1249            } else {
1250                self.audit_log_stats
1251                    .set_recvd_conn_close_transport_error_code(
1252                        err.error_code as _,
1253                    );
1254            }
1255        }
1256
1257        if let Some(err) = qconn.local_error() {
1258            if err.is_app {
1259                self.audit_log_stats
1260                    .set_sent_conn_close_application_error_code(
1261                        err.error_code as _,
1262                    );
1263            } else {
1264                self.audit_log_stats
1265                    .set_sent_conn_close_transport_error_code(
1266                        err.error_code as _,
1267                    );
1268            }
1269        }
1270
1271        self.close_connection(qconn);
1272
1273        if let Err(work_loop_error) = self.conn_stage.work_loop_result {
1274            self.audit_log_stats
1275                .set_connection_close_reason(work_loop_error);
1276        }
1277    }
1278
1279    fn close_connection(&mut self, qconn: &mut QuicheConnection) {
1280        if let Some(cid) = self.cfg.pending_cid.take() {
1281            self.unmap_cid(cid);
1282        }
1283        while let Some(retired_cid) = qconn.retired_scid_next() {
1284            self.unmap_cid(retired_cid);
1285        }
1286        for cid in qconn.source_ids().cloned() {
1287            self.unmap_cid(cid.into_owned());
1288        }
1289
1290        self.metrics.connections_in_memory().dec();
1291    }
1292}
1293
1294/// Returns the minimum of `v1` and `v2`, ignoring `None`s.
1295fn min_of_some<T: Ord>(v1: Option<T>, v2: Option<T>) -> Option<T> {
1296    match (v1, v2) {
1297        (Some(a), Some(b)) => Some(a.min(b)),
1298        (Some(v), _) | (_, Some(v)) => Some(v),
1299        (None, None) => None,
1300    }
1301}
1302
1303/// A Token which increment the skipped_mid_handshake_flush_count metric on
1304/// `Drop` unless it is marked complete.
1305struct TrackMidHandshakeFlush<M: Metrics> {
1306    complete: bool,
1307    metrics: M,
1308}
1309
1310impl<M: Metrics> TrackMidHandshakeFlush<M> {
1311    fn new(metrics: M) -> Self {
1312        Self {
1313            complete: false,
1314            metrics,
1315        }
1316    }
1317
1318    fn mark_complete(&mut self) {
1319        self.complete = true;
1320    }
1321}
1322
1323impl<M: Metrics> Drop for TrackMidHandshakeFlush<M> {
1324    fn drop(&mut self) {
1325        if !self.complete {
1326            self.metrics.skipped_mid_handshake_flush_count().inc();
1327        }
1328    }
1329}
1330
1331fn random_u128() -> u128 {
1332    let mut buf = [0; 16];
1333    boring::rand::rand_bytes(&mut buf).expect("boring's RAND_bytes never fails");
1334    u128::from_ne_bytes(buf)
1335}
1336
1337#[cfg(test)]
1338mod pooled_send_buf_tests {
1339    use super::*;
1340
1341    // Each pool test runs on a freshly spawned thread so the thread-local
1342    // `SEND_BUF_POOL` starts empty (const-initialized) and cannot interfere
1343    // with other tests sharing the harness's worker threads.
1344
1345    #[test]
1346    fn caps_retained_buffers() {
1347        std::thread::spawn(|| {
1348            // Acquire more than the cap at once (all misses, so all fresh
1349            // allocations), then drop them. Only `SEND_BUF_POOL_CAP` may be
1350            // parked; the remainder are freed.
1351            let bufs: Vec<PooledSendBuf> = (0..SEND_BUF_POOL_CAP + 4)
1352                .map(|_| PooledSendBuf::acquire())
1353                .collect();
1354            drop(bufs);
1355
1356            let retained = SEND_BUF_POOL.with(|pool| pool.borrow().len());
1357            assert_eq!(retained, SEND_BUF_POOL_CAP);
1358        })
1359        .join()
1360        .unwrap();
1361    }
1362
1363    #[test]
1364    fn reuses_a_returned_buffer() {
1365        std::thread::spawn(|| {
1366            let first_ptr = {
1367                let buf = PooledSendBuf::acquire();
1368                assert_eq!(buf.len(), SEND_BUFFER_SIZE);
1369                buf.as_ptr()
1370            }; // returned to the pool here
1371
1372            let reused = PooledSendBuf::acquire();
1373            assert_eq!(
1374                reused.as_ptr(),
1375                first_ptr,
1376                "acquire should hand back the pooled allocation"
1377            );
1378        })
1379        .join()
1380        .unwrap();
1381    }
1382
1383    #[test]
1384    fn returns_to_the_dropping_thread() {
1385        // Acquire on one thread, drop on another: the buffer lands in the
1386        // dropping thread's pool (work-stealing migration across `.await` is
1387        // benign).
1388        let buf = std::thread::spawn(PooledSendBuf::acquire).join().unwrap();
1389
1390        std::thread::spawn(move || {
1391            assert_eq!(SEND_BUF_POOL.with(|pool| pool.borrow().len()), 0);
1392            drop(buf);
1393            assert_eq!(SEND_BUF_POOL.with(|pool| pool.borrow().len()), 1);
1394        })
1395        .join()
1396        .unwrap();
1397    }
1398}