1use 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#[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 segments: VecDeque<Segment>,
114 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 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 &s2n_quic_core::time::clock::Cached::new(clock),
190 subscriber
191 ))?;
192
193 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 &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 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 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 segment_alloc.free(segment.buffer);
328
329 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 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 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 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 let segment_count = segments.len();
421 drop(segments);
422 for segment in self.segments.drain(..segment_count) {
423 segment_alloc.free(segment.buffer);
425 }
426
427 ready!(result)?;
428 }
429
430 Ok(()).into()
431 }
432}