noq_proto/connection/streams/
mod.rs1use std::{
2 collections::{BinaryHeap, hash_map},
3 io,
4};
5
6use bytes::Bytes;
7use thiserror::Error;
8use tracing::trace;
9
10use super::spaces::Retransmits;
11use crate::{
12 Dir, StreamId, VarInt,
13 connection::streams::state::{StreamRecv, get_or_insert_recv, get_or_insert_send},
14 frame,
15};
16
17mod recv;
18use recv::Recv;
19pub use recv::{Chunks, ReadError, ReadableError};
20
21mod send;
22pub(crate) use send::{ByteSlice, BytesArray, Written};
23use send::{BytesSource, Send, SendState};
24pub use send::{FinishError, WriteError};
25
26mod state;
27#[allow(unreachable_pub)] pub use state::StreamsState;
29
30pub struct Streams<'a> {
32 pub(super) state: &'a mut StreamsState,
33 pub(super) conn_state: &'a super::State,
34}
35
36#[allow(clippy::needless_lifetimes)] impl<'a> Streams<'a> {
38 #[cfg(fuzzing)]
39 pub fn new(state: &'a mut StreamsState, conn_state: &'a super::State) -> Self {
40 Self { state, conn_state }
41 }
42
43 pub fn open(&mut self, dir: Dir) -> Option<StreamId> {
47 if self.conn_state.is_closed() {
48 return None;
49 }
50
51 if self.state.next[dir as usize] >= self.state.max[dir as usize] {
52 self.state.streams_blocked[dir as usize] = true;
53 return None;
54 }
55
56 self.state.next[dir as usize] += 1;
57 let id = StreamId::new(self.state.side, dir, self.state.next[dir as usize] - 1);
58 self.state.insert_local(id);
59 self.state.send_streams += 1;
60 Some(id)
61 }
62
63 pub fn accept(&mut self, dir: Dir) -> Option<StreamId> {
68 if self.state.next_remote[dir as usize] == self.state.next_reported_remote[dir as usize] {
69 return None;
70 }
71
72 let x = self.state.next_reported_remote[dir as usize];
73 self.state.next_reported_remote[dir as usize] = x + 1;
74 if dir == Dir::Bi {
75 self.state.send_streams += 1;
76 }
77
78 Some(StreamId::new(!self.state.side, dir, x))
79 }
80
81 #[cfg(fuzzing)]
82 pub fn state(&mut self) -> &mut StreamsState {
83 self.state
84 }
85
86 pub fn send_streams(&self) -> usize {
88 self.state.send_streams
89 }
90
91 pub fn remote_open_streams(&self, dir: Dir) -> u64 {
98 self.state.next_remote[dir as usize]
100 - (self.state.max_remote[dir as usize]
101 - self.state.allocated_remote_count[dir as usize])
102 }
103}
104
105pub struct RecvStream<'a> {
107 pub(super) id: StreamId,
108 pub(super) state: &'a mut StreamsState,
109 pub(super) pending: &'a mut Retransmits,
110}
111
112impl RecvStream<'_> {
113 pub fn is_ordered(&self) -> Result<bool, ClosedStream> {
118 let Some(stream) = self.state.recv.get(&self.id) else {
119 return Err(ClosedStream { _private: () });
120 };
121 let Some(stream) = stream.as_ref().and_then(StreamRecv::as_open_recv) else {
122 return Ok(true);
123 };
124 if stream.stopped {
125 return Err(ClosedStream { _private: () });
126 }
127 Ok(stream.assembler.is_ordered())
128 }
129
130 pub fn read(&mut self, ordered: bool) -> Result<Chunks<'_>, ReadableError> {
147 Chunks::new(self.id, ordered, self.state, self.pending)
148 }
149
150 pub fn stop(&mut self, error_code: VarInt) -> Result<(), ClosedStream> {
155 let mut entry = match self.state.recv.entry(self.id) {
156 hash_map::Entry::Occupied(s) => s,
157 hash_map::Entry::Vacant(_) => return Err(ClosedStream { _private: () }),
158 };
159 let stream = get_or_insert_recv(self.state.stream_receive_window)(entry.get_mut());
160
161 let (read_credits, stop_sending) = stream.stop()?;
162 if stop_sending.should_transmit() {
163 self.pending.stop_sending.push(frame::StopSending {
164 id: self.id,
165 error_code,
166 });
167 }
168
169 if !stream.final_offset_unknown() {
173 let recv = entry.remove().expect("must have recv when stopping");
174 self.state.stream_recv_freed(self.id, recv);
175 }
176
177 if self.state.add_read_credits(read_credits).should_transmit() {
178 self.pending.max_data = true;
179 }
180
181 Ok(())
182 }
183
184 pub fn bytes_read(&self) -> Result<u64, ClosedStream> {
189 let recv = self
190 .state
191 .recv
192 .get(&self.id)
193 .and_then(|s| s.as_ref())
194 .and_then(|s| s.as_open_recv())
195 .ok_or(ClosedStream { _private: () })?;
196 Ok(recv.assembler.bytes_read())
197 }
198
199 pub fn received_reset(&mut self) -> Result<Option<VarInt>, ClosedStream> {
204 let hash_map::Entry::Occupied(entry) = self.state.recv.entry(self.id) else {
205 return Err(ClosedStream { _private: () });
206 };
207 let Some(s) = entry.get().as_ref().and_then(|s| s.as_open_recv()) else {
208 return Ok(None);
209 };
210 if s.stopped {
211 return Err(ClosedStream { _private: () });
212 }
213 let Some(code) = s.reset_code() else {
214 return Ok(None);
215 };
216
217 let (_, recv) = entry.remove_entry();
220 self.state
221 .stream_recv_freed(self.id, recv.expect("must have recv on reset"));
222 self.state.queue_max_stream_id(self.pending);
223
224 Ok(Some(code))
225 }
226}
227
228pub struct SendStream<'a> {
230 pub(super) id: StreamId,
231 pub(super) state: &'a mut StreamsState,
232 pub(super) pending: &'a mut Retransmits,
233 pub(super) conn_state: &'a super::State,
234}
235
236#[allow(clippy::needless_lifetimes)] impl<'a> SendStream<'a> {
238 #[cfg(fuzzing)]
239 pub fn new(
240 id: StreamId,
241 state: &'a mut StreamsState,
242 pending: &'a mut Retransmits,
243 conn_state: &'a super::State,
244 ) -> Self {
245 Self {
246 id,
247 state,
248 pending,
249 conn_state,
250 }
251 }
252
253 pub fn write(&mut self, data: &[u8]) -> Result<usize, WriteError> {
257 Ok(self.write_source(&mut ByteSlice::from_slice(data))?.bytes)
258 }
259
260 pub fn write_chunks(&mut self, data: &mut &mut [Bytes]) -> Result<usize, WriteError> {
268 let written = self.write_source(&mut BytesArray::from_chunks(data))?;
269 *data = &mut std::mem::take(data)[written.chunks..];
270 Ok(written.bytes)
271 }
272
273 fn write_source<'b, B: BytesSource<'b>>(
274 &mut self,
275 source: &'b mut B,
276 ) -> Result<Written, WriteError> {
277 if self.conn_state.is_closed() {
278 trace!(%self.id, "write blocked; connection draining");
279 return Err(WriteError::Blocked);
280 }
281
282 let limit = self.state.write_limit();
283
284 let max_send_data = self.state.max_send_data(self.id);
285
286 let stream = self
287 .state
288 .send
289 .get_mut(&self.id)
290 .map(get_or_insert_send(max_send_data))
291 .ok_or(WriteError::ClosedStream)?;
292
293 if limit == 0 {
294 trace!(
295 stream = %self.id, max_data = self.state.max_data, data_sent = self.state.data_sent,
296 "write blocked by connection-level flow control or send window"
297 );
298 if !stream.connection_blocked {
299 stream.connection_blocked = true;
300 self.state.connection_blocked.push(self.id);
301 }
302 return Err(WriteError::Blocked);
303 }
304
305 let was_pending = stream.is_pending();
306 let written = stream.write(source, limit)?;
307 self.state.data_sent += written.bytes as u64;
308 self.state.unacked_data += written.bytes as u64;
309 trace!(stream = %self.id, "wrote {} bytes", written.bytes);
310 if !was_pending {
311 self.state.pending.push_pending(self.id, stream.priority);
312 }
313 Ok(written)
314 }
315
316 pub fn stopped(&self) -> Result<Option<VarInt>, ClosedStream> {
318 match self.state.send.get(&self.id).as_ref() {
319 Some(Some(s)) => Ok(s.stop_reason),
320 Some(None) => Ok(None),
321 None => Err(ClosedStream { _private: () }),
322 }
323 }
324
325 pub fn finish(&mut self) -> Result<(), FinishError> {
331 let max_send_data = self.state.max_send_data(self.id);
332 let stream = self
333 .state
334 .send
335 .get_mut(&self.id)
336 .map(get_or_insert_send(max_send_data))
337 .ok_or(FinishError::ClosedStream)?;
338
339 let was_pending = stream.is_pending();
340 stream.finish()?;
341 if !was_pending {
342 self.state.pending.push_pending(self.id, stream.priority);
343 }
344
345 Ok(())
346 }
347
348 pub fn reset(&mut self, error_code: VarInt) -> Result<(), ClosedStream> {
353 let max_send_data = self.state.max_send_data(self.id);
354 let stream = self
355 .state
356 .send
357 .get_mut(&self.id)
358 .map(get_or_insert_send(max_send_data))
359 .ok_or(ClosedStream { _private: () })?;
360
361 if matches!(stream.state, SendState::ResetSent) {
362 return Err(ClosedStream { _private: () });
364 }
365
366 self.state.unacked_data -= stream.pending.unacked();
370 stream.reset();
371 self.pending.reset_stream.push((self.id, error_code));
372
373 Ok(())
375 }
376
377 pub fn set_priority(&mut self, priority: i32) -> Result<(), ClosedStream> {
382 let max_send_data = self.state.max_send_data(self.id);
383 let stream = self
384 .state
385 .send
386 .get_mut(&self.id)
387 .map(get_or_insert_send(max_send_data))
388 .ok_or(ClosedStream { _private: () })?;
389
390 stream.priority = priority;
391 Ok(())
392 }
393
394 pub fn priority(&self) -> Result<i32, ClosedStream> {
399 let stream = self
400 .state
401 .send
402 .get(&self.id)
403 .ok_or(ClosedStream { _private: () })?;
404
405 Ok(stream.as_ref().map(|s| s.priority).unwrap_or_default())
406 }
407}
408
409struct PendingStreamsQueue {
411 streams: BinaryHeap<PendingStream>,
412 next: Option<PendingStream>,
416 recency: u64,
420}
421
422impl PendingStreamsQueue {
423 fn new() -> Self {
424 Self {
425 streams: BinaryHeap::new(),
426 next: None,
427 recency: u64::MAX,
428 }
429 }
430
431 fn reinsert_pending(&mut self, id: StreamId, priority: i32) {
433 assert!(self.next.is_none());
434
435 self.next = Some(PendingStream {
436 priority,
437 recency: self.recency, id,
439 });
440 }
441
442 fn push_pending(&mut self, id: StreamId, priority: i32) {
445 self.recency -= 1;
455 self.streams.push(PendingStream {
456 priority,
457 recency: self.recency,
458 id,
459 });
460 }
461
462 fn pop(&mut self) -> Option<PendingStream> {
463 self.next.take().or_else(|| self.streams.pop())
464 }
465
466 fn clear(&mut self) {
467 self.next = None;
468 self.streams.clear();
469 }
470
471 fn iter(&self) -> impl Iterator<Item = &PendingStream> {
472 self.next.iter().chain(self.streams.iter())
473 }
474
475 #[cfg(test)]
476 fn len(&self) -> usize {
477 self.streams.len() + self.next.is_some() as usize
478 }
479}
480
481#[derive(Clone, PartialEq, Eq, PartialOrd, Ord)]
483struct PendingStream {
484 priority: i32,
488 recency: u64,
495 id: StreamId,
500}
501
502#[derive(Debug, PartialEq, Eq)]
504pub enum StreamEvent {
505 Opened {
507 dir: Dir,
509 },
510 Readable {
512 id: StreamId,
514 },
515 Writable {
519 id: StreamId,
521 },
522 Finished {
524 id: StreamId,
526 },
527 Stopped {
529 id: StreamId,
531 error_code: VarInt,
533 },
534 Available {
536 dir: Dir,
538 },
539}
540
541#[derive(Copy, Clone, Debug, Default, Eq, PartialEq)]
546#[must_use = "A frame might need to be enqueued"]
547pub struct ShouldTransmit(bool);
548
549impl ShouldTransmit {
550 pub fn should_transmit(self) -> bool {
552 self.0
553 }
554}
555
556#[derive(Debug, Default, Error, Clone, PartialEq, Eq)]
558#[error("closed stream")]
559pub struct ClosedStream {
560 _private: (),
561}
562
563impl From<ClosedStream> for io::Error {
564 fn from(x: ClosedStream) -> Self {
565 Self::new(io::ErrorKind::NotConnected, x)
566 }
567}
568
569#[derive(Debug, Copy, Clone, Eq, PartialEq)]
570enum StreamHalf {
571 Send,
572 Recv,
573}
574
575pub(super) trait BytesOrSlice<'a>: AsRef<[u8]> + 'a {
577 fn len(&self) -> usize {
578 self.as_ref().len()
579 }
580 fn is_empty(&self) -> bool {
581 self.as_ref().is_empty()
582 }
583 fn into_bytes(self) -> Bytes;
584}
585
586impl BytesOrSlice<'_> for Bytes {
587 fn into_bytes(self) -> Bytes {
588 self
589 }
590}
591
592impl<'a> BytesOrSlice<'a> for &'a [u8] {
593 fn into_bytes(self) -> Bytes {
594 Bytes::copy_from_slice(self)
595 }
596}