1use crate::codec::UserError;
2use crate::frame::{Reason, StreamId};
3use crate::{client, server};
4
5use crate::frame::DEFAULT_INITIAL_WINDOW_SIZE;
6use crate::proto::*;
7
8use bytes::Bytes;
9use futures_core::Stream;
10use std::io;
11use std::marker::PhantomData;
12use std::pin::Pin;
13use std::task::{Context, Poll};
14use std::time::Duration;
15use tokio::io::AsyncRead;
16
17#[derive(Debug)]
19pub(crate) struct Connection<T, P, B: Buf = Bytes>
20where
21 P: Peer,
22{
23 codec: Codec<T, Prioritized<B>>,
25
26 inner: ConnectionInner<P, B>,
27}
28
29#[derive(Debug)]
32struct ConnectionInner<P, B: Buf = Bytes>
33where
34 P: Peer,
35{
36 state: State,
38
39 error: Option<frame::GoAway>,
44
45 go_away: GoAway,
47
48 ping_pong: PingPong,
50
51 settings: Settings,
53
54 streams: Streams<B, P>,
56
57 span: tracing::Span,
59
60 _phantom: PhantomData<P>,
62}
63
64struct DynConnection<'a, B: Buf = Bytes> {
65 state: &'a mut State,
66
67 go_away: &'a mut GoAway,
68
69 streams: DynStreams<'a, B>,
70
71 error: &'a mut Option<frame::GoAway>,
72
73 ping_pong: &'a mut PingPong,
74}
75
76#[derive(Debug, Clone)]
77pub(crate) struct Config {
78 pub next_stream_id: StreamId,
79 pub initial_max_send_streams: usize,
80 pub max_send_buffer_size: usize,
81 pub reset_stream_duration: Duration,
82 pub reset_stream_max: usize,
83 pub remote_reset_stream_max: usize,
84 pub local_error_reset_streams_max: Option<usize>,
85 pub settings: frame::Settings,
86 pub data_frame_budget: usize,
87}
88
89#[derive(Clone, Copy, Debug)]
90pub(crate) enum DataFrameBudget {
91 Auto,
92 Configured(usize),
93}
94
95impl DataFrameBudget {
96 pub(crate) fn resolve(self, connection_window: Option<WindowSize>) -> usize {
97 match self {
98 Self::Configured(budget) => budget,
99 Self::Auto => {
100 let window = connection_window.unwrap_or(DEFAULT_INITIAL_WINDOW_SIZE);
101 let budget = window as usize / 2;
102
103 budget.max(DEFAULT_DATA_FRAME_BUDGET)
104 }
105 }
106 }
107}
108
109#[derive(Debug)]
110enum State {
111 Open,
113
114 Closing(Reason, Initiator),
116
117 Closed(Reason, Initiator),
119}
120
121impl<T, P, B> Connection<T, P, B>
122where
123 T: AsyncRead + AsyncWrite + Unpin,
124 P: Peer,
125 B: Buf,
126{
127 pub fn new(codec: Codec<T, Prioritized<B>>, config: Config) -> Connection<T, P, B> {
128 fn streams_config(config: &Config) -> streams::Config {
129 streams::Config {
130 initial_max_send_streams: config.initial_max_send_streams,
131 local_max_buffer_size: config.max_send_buffer_size,
132 local_next_stream_id: config.next_stream_id,
133 local_push_enabled: config.settings.is_push_enabled().unwrap_or(true),
134 extended_connect_protocol_enabled: config
135 .settings
136 .is_extended_connect_protocol_enabled()
137 .unwrap_or(false),
138 local_reset_duration: config.reset_stream_duration,
139 local_reset_max: config.reset_stream_max,
140 remote_reset_max: config.remote_reset_stream_max,
141 remote_init_window_sz: DEFAULT_INITIAL_WINDOW_SIZE,
142 remote_max_initiated: config
143 .settings
144 .max_concurrent_streams()
145 .map(|max| max as usize),
146 local_max_error_reset_streams: config.local_error_reset_streams_max,
147 data_frame_budget: config.data_frame_budget,
148 }
149 }
150 let streams = Streams::new(streams_config(&config));
151 let span = tracing::debug_span!(parent: None, "Connection", peer = %P::NAME);
152 span.follows_from(tracing::Span::current());
153 Connection {
154 codec,
155 inner: ConnectionInner {
156 state: State::Open,
157 error: None,
158 go_away: GoAway::new(),
159 ping_pong: PingPong::new(),
160 settings: Settings::new(config.settings),
161 streams,
162 span,
163 _phantom: PhantomData,
164 },
165 }
166 }
167
168 pub(crate) fn set_target_window_size(&mut self, size: WindowSize) {
170 let _res = self.inner.streams.set_target_connection_window_size(size);
171 debug_assert!(_res.is_ok());
173 }
174
175 pub(crate) fn set_initial_window_size(&mut self, size: WindowSize) -> Result<(), UserError> {
177 let mut settings = frame::Settings::default();
178 settings.set_initial_window_size(Some(size));
179 self.inner.settings.send_settings(settings)
180 }
181
182 pub(crate) fn set_enable_connect_protocol(&mut self) -> Result<(), UserError> {
184 let mut settings = frame::Settings::default();
185 settings.set_enable_connect_protocol(Some(1));
186 self.inner.settings.send_settings(settings)
187 }
188
189 pub(crate) fn max_send_streams(&self) -> usize {
192 self.inner.streams.max_send_streams()
193 }
194
195 pub(crate) fn max_recv_streams(&self) -> usize {
198 self.inner.streams.max_recv_streams()
199 }
200
201 #[cfg(feature = "unstable")]
202 pub fn num_wired_streams(&self) -> usize {
203 self.inner.streams.num_wired_streams()
204 }
205
206 fn poll_ready(&mut self, cx: &mut Context) -> Poll<Result<(), Error>> {
211 let _e = self.inner.span.enter();
212 let span = tracing::trace_span!("poll_ready");
213 let _e = span.enter();
214 ready!(self.inner.ping_pong.send_pending_pong(cx, &mut self.codec))?;
216 ready!(self.inner.ping_pong.send_pending_ping(cx, &mut self.codec))?;
217 ready!(self
218 .inner
219 .settings
220 .poll_send(cx, &mut self.codec, &mut self.inner.streams))?;
221 ready!(self.inner.streams.send_pending_refusal(cx, &mut self.codec))?;
222
223 Poll::Ready(Ok(()))
224 }
225
226 fn poll_go_away(&mut self, cx: &mut Context) -> Poll<Option<io::Result<Reason>>> {
231 self.inner.go_away.send_pending_go_away(cx, &mut self.codec)
232 }
233
234 pub fn go_away_from_user(&mut self, e: Reason) {
235 self.inner.as_dyn().go_away_from_user(e)
236 }
237
238 fn take_error(&mut self, ours: Reason, initiator: Initiator) -> Result<(), Error> {
239 let (debug_data, theirs) = self
240 .inner
241 .error
242 .take()
243 .as_ref()
244 .map_or((Bytes::new(), Reason::NO_ERROR), |frame| {
245 (frame.debug_data().clone(), frame.reason())
246 });
247
248 match (ours, theirs) {
249 (Reason::NO_ERROR, Reason::NO_ERROR) => Ok(()),
250 (ours, Reason::NO_ERROR) => Err(Error::GoAway(Bytes::new(), ours, initiator)),
251 (_, theirs) => Err(Error::remote_go_away(debug_data, theirs)),
256 }
257 }
258
259 pub fn maybe_close_connection_if_no_streams(&mut self) {
262 if !self.inner.streams.has_streams_or_other_references() {
265 self.inner.as_dyn().go_away_now(Reason::NO_ERROR);
266 }
267 }
268
269 pub fn has_streams(&self) -> bool {
271 self.inner.streams.has_streams()
272 }
273
274 pub fn has_streams_or_other_references(&self) -> bool {
276 self.inner.streams.has_streams_or_other_references()
279 }
280
281 pub(crate) fn take_user_pings(&mut self) -> Option<UserPings> {
282 self.inner.ping_pong.take_user_pings()
283 }
284
285 pub fn poll(&mut self, cx: &mut Context) -> Poll<Result<(), Error>> {
287 let span = self.inner.span.clone();
292 let _e = span.enter();
293 let span = tracing::trace_span!("poll");
294 let _e = span.enter();
295
296 loop {
297 tracing::trace!(connection.state = ?self.inner.state);
298 match self.inner.state {
300 State::Open => {
302 let result = match self.poll2(cx) {
303 Poll::Ready(result) => result,
304 Poll::Pending => {
306 ready!(self.inner.streams.poll_complete(cx, &mut self.codec))?;
310
311 if (self.inner.error.is_some()
312 || self.inner.go_away.should_close_on_idle())
313 && !self.inner.streams.has_streams()
314 {
315 self.inner.as_dyn().go_away_now(Reason::NO_ERROR);
316 continue;
317 }
318
319 return Poll::Pending;
320 }
321 };
322
323 self.inner.as_dyn().handle_poll2_result(result)?
324 }
325 State::Closing(reason, initiator) => {
326 tracing::trace!("connection closing after flush");
327 ready!(self.codec.shutdown(cx))?;
329
330 self.inner.state = State::Closed(reason, initiator);
332 }
333 State::Closed(reason, initiator) => {
334 return Poll::Ready(self.take_error(reason, initiator));
335 }
336 }
337 }
338 }
339
340 fn poll2(&mut self, cx: &mut Context) -> Poll<Result<(), Error>> {
341 self.clear_expired_reset_streams();
345
346 loop {
347 if let Some(reason) = ready!(self.poll_go_away(cx)?) {
353 if self.inner.go_away.should_close_now() {
354 if self.inner.go_away.is_user_initiated() {
355 return Poll::Ready(Ok(()));
358 } else {
359 return Poll::Ready(Err(Error::library_go_away(reason)));
360 }
361 }
362 debug_assert_eq!(
364 reason,
365 Reason::NO_ERROR,
366 "graceful GOAWAY should be NO_ERROR"
367 );
368 }
369 ready!(self.poll_ready(cx))?;
370
371 match self
372 .inner
373 .as_dyn()
374 .recv_frame(ready!(Pin::new(&mut self.codec).poll_next(cx)?))?
375 {
376 ReceivedFrame::Settings(frame) => {
377 self.inner.settings.recv_settings(
378 frame,
379 &mut self.codec,
380 &mut self.inner.streams,
381 )?;
382 }
383 ReceivedFrame::Continue => (),
384 ReceivedFrame::Done => {
385 return Poll::Ready(Ok(()));
386 }
387 }
388 }
389 }
390
391 fn clear_expired_reset_streams(&mut self) {
392 self.inner.streams.clear_expired_reset_streams();
393 }
394}
395
396impl<P, B> ConnectionInner<P, B>
397where
398 P: Peer,
399 B: Buf,
400{
401 fn as_dyn(&mut self) -> DynConnection<'_, B> {
402 let ConnectionInner {
403 state,
404 go_away,
405 streams,
406 error,
407 ping_pong,
408 ..
409 } = self;
410 let streams = streams.as_dyn();
411 DynConnection {
412 state,
413 go_away,
414 streams,
415 error,
416 ping_pong,
417 }
418 }
419}
420
421impl<B> DynConnection<'_, B>
422where
423 B: Buf,
424{
425 fn go_away(&mut self, id: StreamId, e: Reason) {
426 let frame = frame::GoAway::new(id, e);
427 self.streams.send_go_away(id);
428 self.go_away.go_away(frame);
429 }
430
431 fn go_away_now(&mut self, e: Reason) {
432 let last_processed_id = self.streams.last_processed_id();
433 let frame = frame::GoAway::new(last_processed_id, e);
434 self.go_away.go_away_now(frame);
435 }
436
437 fn go_away_now_data(&mut self, e: Reason, data: Bytes) {
438 let last_processed_id = self.streams.last_processed_id();
439 let frame = frame::GoAway::with_debug_data(last_processed_id, e, data);
440 self.go_away.go_away_now(frame);
441 }
442
443 fn go_away_from_user(&mut self, e: Reason) {
444 let last_processed_id = self.streams.last_processed_id();
445 let frame = frame::GoAway::new(last_processed_id, e);
446 self.go_away.go_away_from_user(frame);
447
448 self.streams.handle_error(Error::user_go_away(e));
450 }
451
452 fn handle_poll2_result(&mut self, result: Result<(), Error>) -> Result<(), Error> {
453 match result {
454 Ok(()) => {
456 *self.state = State::Closing(Reason::NO_ERROR, Initiator::Library);
457 Ok(())
458 }
459 Err(Error::GoAway(debug_data, reason, initiator)) => {
463 self.handle_go_away(reason, debug_data, initiator);
464 Ok(())
465 }
466 Err(Error::Reset(id, reason, initiator)) => {
471 if initiator == Initiator::Remote {
472 tracing::trace!(?id, ?reason, ?initiator, "stream reset");
473 return Ok(());
474 }
475
476 debug_assert_eq!(initiator, Initiator::Library);
477 tracing::trace!(?id, ?reason, ?initiator, "stream error");
478 match self.streams.send_reset(id, reason) {
479 Ok(()) => (),
480 Err(crate::proto::error::GoAway { debug_data, reason }) => {
481 self.handle_go_away(reason, debug_data, Initiator::Library);
482 }
483 }
484 Ok(())
485 }
486 Err(Error::Io(kind, inner)) => {
491 tracing::debug!(error = ?kind, "Connection::poll; IO error");
492 let e = Error::Io(kind, inner);
493
494 self.streams.handle_error(e.clone());
496
497 if self.streams.is_buffer_empty()
504 && matches!(kind, io::ErrorKind::UnexpectedEof)
505 && (self.streams.is_server()
506 || self.error.as_ref().map(|f| f.reason() == Reason::NO_ERROR)
507 == Some(true))
508 {
509 *self.state = State::Closed(Reason::NO_ERROR, Initiator::Library);
510 return Ok(());
511 }
512
513 Err(e)
515 }
516 }
517 }
518
519 fn handle_go_away(&mut self, reason: Reason, debug_data: Bytes, initiator: Initiator) {
520 let e = Error::GoAway(debug_data.clone(), reason, initiator);
521 tracing::debug!(error = ?e, "Connection::poll; connection error");
522
523 if self
526 .go_away
527 .going_away()
528 .map_or(false, |frame| frame.reason() == reason)
529 {
530 tracing::trace!(" -> already going away");
531 *self.state = State::Closing(reason, initiator);
532 return;
533 }
534
535 self.streams.handle_error(e);
537 self.go_away_now_data(reason, debug_data);
538 }
539
540 fn recv_frame(&mut self, frame: Option<Frame>) -> Result<ReceivedFrame, Error> {
541 use crate::frame::Frame::*;
542 match frame {
543 Some(Headers(frame)) => {
544 tracing::trace!(?frame, "recv HEADERS");
545 self.streams.recv_headers(frame)?;
546 }
547 Some(Data(frame)) => {
548 tracing::trace!(?frame, "recv DATA");
549 self.streams.recv_data(frame)?;
550 }
551 Some(Reset(frame)) => {
552 tracing::trace!(?frame, "recv RST_STREAM");
553 self.streams.recv_reset(frame)?;
554 }
555 Some(PushPromise(frame)) => {
556 tracing::trace!(?frame, "recv PUSH_PROMISE");
557 self.streams.recv_push_promise(frame)?;
558 }
559 Some(Settings(frame)) => {
560 tracing::trace!(?frame, "recv SETTINGS");
561 return Ok(ReceivedFrame::Settings(frame));
562 }
563 Some(GoAway(frame)) => {
564 tracing::trace!(?frame, "recv GOAWAY");
565 self.streams.recv_go_away(&frame)?;
570 *self.error = Some(frame);
571 }
572 Some(Ping(frame)) => {
573 tracing::trace!(?frame, "recv PING");
574 let status = self.ping_pong.recv_ping(frame);
575 if status.is_shutdown() {
576 assert!(
577 self.go_away.is_going_away(),
578 "received unexpected shutdown ping"
579 );
580
581 let last_processed_id = self.streams.last_processed_id();
582 self.go_away(last_processed_id, Reason::NO_ERROR);
583 }
584 }
585 Some(WindowUpdate(frame)) => {
586 tracing::trace!(?frame, "recv WINDOW_UPDATE");
587 self.streams.recv_window_update(frame)?;
588 }
589 Some(Priority(frame)) => {
590 tracing::trace!(?frame, "recv PRIORITY");
591 }
593 None => {
594 tracing::trace!("codec closed");
595 self.streams.recv_eof(false).expect("mutex poisoned");
596 return Ok(ReceivedFrame::Done);
597 }
598 }
599 Ok(ReceivedFrame::Continue)
600 }
601}
602
603enum ReceivedFrame {
604 Settings(frame::Settings),
605 Continue,
606 Done,
607}
608
609impl<T, B> Connection<T, client::Peer, B>
610where
611 T: AsyncRead + AsyncWrite,
612 B: Buf,
613{
614 pub(crate) fn streams(&self) -> &Streams<B, client::Peer> {
615 &self.inner.streams
616 }
617}
618
619impl<T, B> Connection<T, server::Peer, B>
620where
621 T: AsyncRead + AsyncWrite + Unpin,
622 B: Buf,
623{
624 pub fn next_incoming(&mut self) -> Option<StreamRef<B>> {
625 self.inner.streams.next_incoming()
626 }
627
628 pub fn go_away_gracefully(&mut self) {
630 if self.inner.go_away.is_going_away() {
631 return;
633 }
634
635 self.inner.as_dyn().go_away(StreamId::MAX, Reason::NO_ERROR);
647
648 self.inner.ping_pong.ping_shutdown();
651 }
652}
653
654impl<T, P, B> Drop for Connection<T, P, B>
655where
656 P: Peer,
657 B: Buf,
658{
659 fn drop(&mut self) {
660 let _ = self.inner.streams.recv_eof(true);
662 }
663}
664
665#[cfg(test)]
666mod tests {
667 use super::*;
668
669 #[test]
670 fn auto_data_frame_budget_scales_with_connection_window() {
671 assert_eq!(
672 DataFrameBudget::Auto.resolve(None),
673 DEFAULT_INITIAL_WINDOW_SIZE as usize / 2
674 );
675 assert_eq!(
676 DataFrameBudget::Auto.resolve(Some(DEFAULT_INITIAL_WINDOW_SIZE)),
677 DEFAULT_INITIAL_WINDOW_SIZE as usize / 2
678 );
679 assert_eq!(DataFrameBudget::Auto.resolve(Some(1024 * 1024)), 512 * 1024);
680 }
681
682 #[test]
683 fn auto_data_frame_budget_has_minimum() {
684 assert_eq!(
685 DataFrameBudget::Auto.resolve(Some(1)),
686 DEFAULT_DATA_FRAME_BUDGET
687 );
688 assert_eq!(
689 DataFrameBudget::Auto.resolve(Some(MAX_WINDOW_SIZE)),
690 MAX_WINDOW_SIZE as usize / 2
691 );
692 }
693
694 #[test]
695 fn configured_data_frame_budget_is_unchanged() {
696 assert_eq!(
697 DataFrameBudget::Configured(123).resolve(Some(MAX_WINDOW_SIZE)),
698 123
699 );
700 }
701}