Skip to main content

s2n_quic_dc/stream/send/
queue.rs

1// Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use crate::{
5    event::{self, ConnectionPublisher},
6    msg::{addr, segment},
7    stream::{
8        send::{
9            application::{self, transmission},
10            buffer,
11            state::Transmission,
12        },
13        shared,
14        socket::Socket,
15    },
16};
17use bytes::buf::UninitSlice;
18use core::task::{Context, Poll};
19use s2n_quic_core::{
20    assume, buffer::reader, ensure, inet::ExplicitCongestionNotification, ready, time::Clock,
21};
22use s2n_quic_platform::features::Gso;
23use std::{collections::VecDeque, io};
24
25/// An enqueued segment waiting to be transmitted on the socket
26#[derive(Debug)]
27pub struct Segment {
28    ecn: ExplicitCongestionNotification,
29    buffer: buffer::Segment,
30    offset: u16,
31}
32
33impl Segment {
34    #[inline]
35    fn as_slice(&self) -> &[u8] {
36        &self.buffer[self.offset as usize..]
37    }
38}
39
40pub struct Message<'a> {
41    batch: &'a mut Option<Vec<Transmission>>,
42    queue: &'a mut Queue,
43    max_segments: usize,
44    segment_alloc: &'a buffer::Allocator,
45}
46
47impl application::state::Message for Message<'_> {
48    #[inline]
49    fn max_segments(&self) -> usize {
50        self.max_segments
51    }
52
53    #[inline]
54    fn push<P: FnOnce(&mut UninitSlice) -> transmission::Event<()>>(
55        &mut self,
56        buffer_len: usize,
57        p: P,
58    ) -> Option<usize> {
59        let (mut buffer, buf_source) = self.segment_alloc.alloc(buffer_len);
60
61        let transmission = {
62            let buffer = buffer.make_mut();
63
64            debug_assert!(buffer.capacity() >= buffer_len);
65
66            let slice = UninitSlice::uninit(buffer.spare_capacity_mut());
67
68            let transmission = p(slice);
69
70            unsafe {
71                let packet_len = transmission.info.packet_len;
72                assume!(buffer.capacity() >= packet_len as usize);
73                buffer.set_len(packet_len as usize);
74            }
75
76            transmission
77        };
78
79        let transmission::Event {
80            packet_number,
81            info,
82            has_more_app_data,
83        } = transmission;
84
85        let ecn = info.ecn;
86
87        if let Some(batch) = self.batch.as_mut() {
88            let info = info.map(|_| buffer.clone());
89
90            batch.push(transmission::Event {
91                packet_number,
92                info,
93                has_more_app_data,
94            });
95        }
96
97        self.queue.segments.push_back(Segment {
98            ecn,
99            buffer,
100            offset: 0,
101        });
102
103        match buf_source {
104            buffer::Source::Pool => None,
105            buffer::Source::Fresh => Some(buffer_len),
106        }
107    }
108}
109
110#[derive(Debug, Default)]
111pub struct Queue {
112    /// Holds any segments that haven't been flushed to the socket
113    segments: VecDeque<Segment>,
114    /// How many bytes we've accepted from the caller of `poll_write`, but actually returned
115    /// `Poll::Pending` for. This many bytes will be skipped the next time `poll_write` is called.
116    ///
117    /// This functionality ensures that we don't return to the application until we've flushed all
118    /// outstanding packets to the underlying socket. Experience has shown applications rely on
119    /// TCP's behavior, which never really requires `flush` or `shutdown` to progress the stream.
120    accepted_len: usize,
121}
122
123impl Queue {
124    #[inline]
125    pub fn is_empty(&self) -> bool {
126        self.segments.is_empty()
127    }
128
129    #[inline]
130    pub fn accepted_len(&self) -> usize {
131        self.accepted_len
132    }
133
134    #[inline]
135    pub fn push_buffer<B, F, E>(
136        &mut self,
137        buf: &mut B,
138        batch: &mut Option<Vec<Transmission>>,
139        max_segments: usize,
140        segment_alloc: &buffer::Allocator,
141        push: F,
142    ) -> Result<(), E>
143    where
144        B: reader::Storage,
145        F: FnOnce(&mut Message, &mut reader::storage::Tracked<B>) -> Result<(), E>,
146    {
147        let mut message = Message {
148            batch,
149            queue: self,
150            max_segments,
151            segment_alloc,
152        };
153
154        let mut buf = buf.track_read();
155
156        push(&mut message, &mut buf)?;
157
158        // record how many bytes we encrypted/buffered so we only return Ready once everything has
159        // been flushed
160        self.accepted_len += buf.consumed_len();
161
162        Ok(())
163    }
164
165    #[inline]
166    pub fn poll_flush<S, C, Sub>(
167        &mut self,
168        cx: &mut Context,
169        limit: usize,
170        socket: &S,
171        addr: &addr::Addr,
172        segment_alloc: &buffer::Allocator,
173        gso: &Gso,
174        clock: &C,
175        subscriber: &shared::Subscriber<Sub>,
176    ) -> Poll<Result<usize, io::Error>>
177    where
178        S: ?Sized + Socket,
179        C: ?Sized + Clock,
180        Sub: event::Subscriber,
181    {
182        ready!(self.poll_flush_segments(
183            cx,
184            socket,
185            addr,
186            segment_alloc,
187            gso,
188            // cache the timestamps to avoid fetching too many
189            &s2n_quic_core::time::clock::Cached::new(clock),
190            subscriber
191        ))?;
192
193        // Consume accepted credits
194        let accepted = limit.min(self.accepted_len);
195        self.accepted_len -= accepted;
196        Poll::Ready(Ok(accepted))
197    }
198
199    #[inline]
200    fn poll_flush_segments<S, C, Sub>(
201        &mut self,
202        cx: &mut Context,
203        socket: &S,
204        addr: &addr::Addr,
205        segment_alloc: &buffer::Allocator,
206        gso: &Gso,
207        clock: &C,
208        subscriber: &shared::Subscriber<Sub>,
209    ) -> Poll<Result<(), io::Error>>
210    where
211        S: ?Sized + Socket,
212        C: ?Sized + Clock,
213        Sub: event::Subscriber,
214    {
215        ensure!(!self.segments.is_empty(), Poll::Ready(Ok(())));
216
217        let default_addr = addr::Addr::new(Default::default());
218
219        let addr = if socket.features().is_connected() {
220            // no need to load the socket addr if the stream is already connected
221            &default_addr
222        } else {
223            addr
224        };
225
226        if socket.features().is_stream() {
227            self.poll_flush_segments_stream(cx, socket, addr, segment_alloc, clock, subscriber)
228        } else {
229            self.poll_flush_segments_datagram(
230                cx,
231                socket,
232                addr,
233                segment_alloc,
234                gso,
235                clock,
236                subscriber,
237            )
238        }
239    }
240
241    #[inline]
242    fn poll_flush_segments_stream<S, C, Sub>(
243        &mut self,
244        cx: &mut Context,
245        socket: &S,
246        addr: &addr::Addr,
247        segment_alloc: &buffer::Allocator,
248        clock: &C,
249        subscriber: &shared::Subscriber<Sub>,
250    ) -> Poll<Result<(), io::Error>>
251    where
252        S: ?Sized + Socket,
253        C: ?Sized + Clock,
254        Sub: event::Subscriber,
255    {
256        while !self.segments.is_empty() {
257            let mut provided_len = 0;
258            let segments = segment::Batch::new(
259                self.segments.iter().map(|v| {
260                    let slice = v.as_slice();
261                    provided_len += slice.len();
262                    (v.ecn, v.as_slice())
263                }),
264                &socket.features(),
265            );
266
267            let ecn = segments.ecn();
268
269            let result = socket.poll_send(cx, addr, ecn, &segments);
270
271            let now = clock.get_time();
272
273            drop(segments);
274
275            match result {
276                Poll::Ready(Ok(written_len)) => {
277                    subscriber.publisher(now).on_stream_write_socket_flushed(
278                        event::builder::StreamWriteSocketFlushed {
279                            provided_len,
280                            committed_len: written_len,
281                        },
282                    );
283
284                    self.consume_segments(written_len, segment_alloc);
285
286                    // keep trying to drain the buffer
287                    continue;
288                }
289                Poll::Ready(Err(err)) => {
290                    subscriber.publisher(now).on_stream_write_socket_errored(
291                        event::builder::StreamWriteSocketErrored {
292                            provided_len,
293                            errno: err.raw_os_error(),
294                        },
295                    );
296
297                    // the socket encountered an error so clear everything out since we're shutting
298                    // down
299                    self.segments.clear();
300                    self.accepted_len = 0;
301                    return Err(err).into();
302                }
303                Poll::Pending => {
304                    subscriber.publisher(now).on_stream_write_socket_blocked(
305                        event::builder::StreamWriteSocketBlocked { provided_len },
306                    );
307
308                    return Poll::Pending;
309                }
310            }
311        }
312
313        Ok(()).into()
314    }
315
316    #[inline]
317    fn consume_segments(&mut self, consumed: usize, segment_alloc: &buffer::Allocator) {
318        ensure!(consumed > 0);
319
320        let mut remaining = consumed;
321
322        while let Some(mut segment) = self.segments.pop_front() {
323            if let Some(r) = remaining.checked_sub(segment.as_slice().len()) {
324                remaining = r;
325
326                // try to reuse the buffer for future allocations
327                segment_alloc.free(segment.buffer);
328
329                // if we don't have any remaining bytes to pop then we're done
330                ensure!(remaining > 0, break);
331
332                continue;
333            }
334
335            segment.offset += core::mem::take(&mut remaining) as u16;
336
337            debug_assert!(!segment.as_slice().is_empty());
338
339            self.segments.push_front(segment);
340            break;
341        }
342
343        debug_assert_eq!(
344            remaining, 0,
345            "consumed ({consumed}) with too many bytes remaining ({remaining})"
346        );
347    }
348
349    #[inline]
350    fn poll_flush_segments_datagram<S, C, Sub>(
351        &mut self,
352        cx: &mut Context,
353        socket: &S,
354        addr: &addr::Addr,
355        segment_alloc: &buffer::Allocator,
356        gso: &Gso,
357        clock: &C,
358        subscriber: &shared::Subscriber<Sub>,
359    ) -> Poll<Result<(), io::Error>>
360    where
361        S: ?Sized + Socket,
362        C: ?Sized + Clock,
363        Sub: event::Subscriber,
364    {
365        let mut max_segments = gso.max_segments();
366
367        while !self.segments.is_empty() {
368            let mut provided_len = 0;
369
370            // construct all of the segments we're going to send in this batch
371            let segments = segment::Batch::new(
372                self.segments
373                    .iter()
374                    .map(|v| {
375                        let slice = v.as_slice();
376                        provided_len += slice.len();
377                        (v.ecn, slice)
378                    })
379                    .take(max_segments),
380                &socket.features(),
381            );
382
383            let ecn = segments.ecn();
384
385            let result = socket.poll_send(cx, addr, ecn, &segments);
386
387            let now = clock.get_time();
388
389            match &result {
390                Poll::Ready(Ok(_len)) => {
391                    subscriber.publisher(now).on_stream_write_socket_flushed(
392                        event::builder::StreamWriteSocketFlushed {
393                            provided_len,
394                            // if the syscall went through, then we wrote the whole thing
395                            committed_len: provided_len,
396                        },
397                    );
398                }
399                Poll::Ready(Err(error)) => {
400                    subscriber.publisher(now).on_stream_write_socket_errored(
401                        event::builder::StreamWriteSocketErrored {
402                            provided_len,
403                            errno: error.raw_os_error(),
404                        },
405                    );
406
407                    if gso.handle_socket_error(error).is_some() {
408                        // update the max_segments value if it was changed due to the error
409                        max_segments = 1;
410                    }
411                }
412                Poll::Pending => {
413                    subscriber.publisher(now).on_stream_write_socket_blocked(
414                        event::builder::StreamWriteSocketBlocked { provided_len },
415                    );
416                }
417            };
418
419            // consume the segments that we transmitted
420            let segment_count = segments.len();
421            drop(segments);
422            for segment in self.segments.drain(..segment_count) {
423                // try to reuse the buffer for future allocations
424                segment_alloc.free(segment.buffer);
425            }
426
427            ready!(result)?;
428        }
429
430        Ok(()).into()
431    }
432}