1use std::fmt;
2use std::io;
3
4use crate::codec::UserError;
5use crate::frame::{self, Reason, StreamId};
6use crate::proto::{self, Error, Initiator, PollReset};
7
8use self::Inner::*;
9use self::Peer::*;
10
11#[derive(Clone)]
52pub struct State {
53 inner: Inner,
54}
55
56#[derive(Debug, Clone)]
57enum Inner {
58 Idle,
59 ReservedLocal,
61 ReservedRemote,
62 Open { local: Peer, remote: Peer },
63 HalfClosedLocal(Peer), HalfClosedRemote(Peer),
65 Closed(Cause),
66}
67
68#[derive(Debug, Copy, Clone, Default)]
69enum Peer {
70 #[default]
71 AwaitingHeaders,
72 Streaming,
73}
74
75#[derive(Debug, Clone)]
76enum Cause {
77 EndStream,
78 Error(Error),
79 ErrorAfterEndStream(Error),
81
82 ScheduledLibraryReset(Reason),
90}
91
92impl State {
93 pub fn send_open(&mut self, eos: bool) -> Result<(), UserError> {
95 let local = Streaming;
96
97 self.inner = match self.inner {
98 Idle => {
99 if eos {
100 HalfClosedLocal(AwaitingHeaders)
101 } else {
102 Open {
103 local,
104 remote: AwaitingHeaders,
105 }
106 }
107 }
108 Open {
109 local: AwaitingHeaders,
110 remote,
111 } => {
112 if eos {
113 HalfClosedLocal(remote)
114 } else {
115 Open { local, remote }
116 }
117 }
118 HalfClosedRemote(AwaitingHeaders) | ReservedLocal => {
119 if eos {
120 Closed(Cause::EndStream)
121 } else {
122 HalfClosedRemote(local)
123 }
124 }
125 _ => {
126 return Err(UserError::UnexpectedFrameType);
128 }
129 };
130
131 Ok(())
132 }
133
134 pub fn recv_open(&mut self, frame: &frame::Headers) -> Result<bool, Error> {
138 let mut initial = false;
139 let eos = frame.is_end_stream();
140
141 self.inner = match self.inner {
142 Idle => {
143 initial = true;
144
145 if eos {
146 HalfClosedRemote(AwaitingHeaders)
147 } else {
148 Open {
149 local: AwaitingHeaders,
150 remote: if frame.is_informational() {
151 tracing::trace!("skipping 1xx response headers");
152 AwaitingHeaders
153 } else {
154 Streaming
155 },
156 }
157 }
158 }
159 ReservedRemote => {
160 initial = true;
161
162 if eos {
163 Closed(Cause::EndStream)
164 } else if frame.is_informational() {
165 tracing::trace!("skipping 1xx response headers");
166 ReservedRemote
167 } else {
168 HalfClosedLocal(Streaming)
169 }
170 }
171 Open {
172 local,
173 remote: AwaitingHeaders,
174 } => {
175 if eos {
176 HalfClosedRemote(local)
177 } else {
178 Open {
179 local,
180 remote: if frame.is_informational() {
181 tracing::trace!("skipping 1xx response headers");
182 AwaitingHeaders
183 } else {
184 Streaming
185 },
186 }
187 }
188 }
189 HalfClosedLocal(AwaitingHeaders) => {
190 if eos {
191 Closed(Cause::EndStream)
192 } else if frame.is_informational() {
193 tracing::trace!("skipping 1xx response headers");
194 HalfClosedLocal(AwaitingHeaders)
195 } else {
196 HalfClosedLocal(Streaming)
197 }
198 }
199 ref state => {
200 proto_err!(conn: "recv_open: in unexpected state {:?}", state);
202 return Err(Error::library_go_away(Reason::PROTOCOL_ERROR));
203 }
204 };
205
206 Ok(initial)
207 }
208
209 pub fn reserve_remote(&mut self) -> Result<(), Error> {
211 match self.inner {
212 Idle => {
213 self.inner = ReservedRemote;
214 Ok(())
215 }
216 ref state => {
217 proto_err!(conn: "reserve_remote: in unexpected state {:?}", state);
218 Err(Error::library_go_away(Reason::PROTOCOL_ERROR))
219 }
220 }
221 }
222
223 pub fn reserve_local(&mut self) -> Result<(), UserError> {
225 match self.inner {
226 Idle => {
227 self.inner = ReservedLocal;
228 Ok(())
229 }
230 _ => Err(UserError::UnexpectedFrameType),
231 }
232 }
233
234 pub fn recv_close(&mut self) -> Result<(), Error> {
236 match self.inner {
237 Open { local, .. } => {
238 tracing::trace!("recv_close: Open => HalfClosedRemote({:?})", local);
240 self.inner = HalfClosedRemote(local);
241 Ok(())
242 }
243 HalfClosedLocal(..) => {
244 tracing::trace!("recv_close: HalfClosedLocal => Closed");
245 self.inner = Closed(Cause::EndStream);
246 Ok(())
247 }
248 ref state => {
249 proto_err!(conn: "recv_close: in unexpected state {:?}", state);
250 Err(Error::library_go_away(Reason::PROTOCOL_ERROR))
251 }
252 }
253 }
254
255 pub fn recv_reset(&mut self, frame: frame::Reset, queued: bool) {
261 let recv_end_stream = self.is_recv_end_stream();
262 match self.inner {
263 Closed(..) if !queued => {}
266 ref state => {
281 tracing::trace!(
282 "recv_reset; frame={:?}; state={:?}; queued={:?}",
283 frame,
284 state,
285 queued
286 );
287 let error = Error::remote_reset(frame.stream_id(), frame.reason());
288 self.inner = Closed(if recv_end_stream {
290 Cause::ErrorAfterEndStream(error)
291 } else {
292 Cause::Error(error)
293 });
294 }
295 }
296 }
297
298 pub fn handle_error(&mut self, err: &proto::Error) {
300 match self.inner {
301 Closed(..) => {}
302 _ => {
303 tracing::trace!("handle_error; err={:?}", err);
304 self.inner = Closed(Cause::Error(err.clone()));
305 }
306 }
307 }
308
309 pub fn recv_eof(&mut self) {
310 match self.inner {
311 Closed(..) => {}
312 ref state => {
313 tracing::trace!("recv_eof; state={:?}", state);
314 self.inner = Closed(Cause::Error(
315 io::Error::new(
316 io::ErrorKind::BrokenPipe,
317 "stream closed because of a broken pipe",
318 )
319 .into(),
320 ));
321 }
322 }
323 }
324
325 pub fn send_close(&mut self) {
327 match self.inner {
328 Open { remote, .. } => {
329 tracing::trace!("send_close: Open => HalfClosedLocal({:?})", remote);
331 self.inner = HalfClosedLocal(remote);
332 }
333 HalfClosedRemote(..) => {
334 tracing::trace!("send_close: HalfClosedRemote => Closed");
335 self.inner = Closed(Cause::EndStream);
336 }
337 ref state => panic!("send_close: unexpected state {:?}", state),
338 }
339 }
340
341 pub fn set_reset(&mut self, stream_id: StreamId, reason: Reason, initiator: Initiator) {
343 self.inner = Closed(Cause::Error(Error::Reset(stream_id, reason, initiator)));
344 }
345
346 pub fn set_scheduled_reset(&mut self, reason: Reason) {
348 debug_assert!(!self.is_closed());
349 self.inner = Closed(Cause::ScheduledLibraryReset(reason));
350 }
351
352 pub fn get_scheduled_reset(&self) -> Option<Reason> {
353 match self.inner {
354 Closed(Cause::ScheduledLibraryReset(reason)) => Some(reason),
355 _ => None,
356 }
357 }
358
359 pub fn is_scheduled_reset(&self) -> bool {
360 matches!(self.inner, Closed(Cause::ScheduledLibraryReset(..)))
361 }
362
363 pub fn is_local_error(&self) -> bool {
364 match self.inner {
365 Closed(Cause::Error(ref e) | Cause::ErrorAfterEndStream(ref e)) => e.is_local(),
366 Closed(Cause::ScheduledLibraryReset(..)) => true,
367 _ => false,
368 }
369 }
370
371 pub fn is_remote_reset(&self) -> bool {
372 matches!(
373 self.inner,
374 Closed(Cause::Error(Error::Reset(_, _, Initiator::Remote)))
375 | Closed(Cause::ErrorAfterEndStream(Error::Reset(
376 _,
377 _,
378 Initiator::Remote
379 )))
380 )
381 }
382
383 pub fn is_reset(&self) -> bool {
385 match self.inner {
386 Closed(Cause::EndStream) => false,
387 Closed(_) => true,
388 _ => false,
389 }
390 }
391
392 pub fn is_send_streaming(&self) -> bool {
393 matches!(
394 self.inner,
395 Open {
396 local: Streaming,
397 ..
398 } | HalfClosedRemote(Streaming)
399 )
400 }
401
402 pub fn is_recv_headers(&self) -> bool {
404 matches!(
405 self.inner,
406 Idle | Open {
407 remote: AwaitingHeaders,
408 ..
409 } | HalfClosedLocal(AwaitingHeaders)
410 | ReservedRemote
411 )
412 }
413
414 pub fn is_recv_streaming(&self) -> bool {
415 matches!(
416 self.inner,
417 Open {
418 remote: Streaming,
419 ..
420 } | HalfClosedLocal(Streaming)
421 )
422 }
423
424 pub fn is_recv_end_stream(&self) -> bool {
425 matches!(
427 self.inner,
428 Closed(Cause::EndStream | Cause::ErrorAfterEndStream(_)) | HalfClosedRemote(..)
429 )
430 }
431
432 pub fn is_closed(&self) -> bool {
433 matches!(self.inner, Closed(_))
434 }
435
436 pub fn is_send_closed(&self) -> bool {
437 matches!(
438 self.inner,
439 Closed(..) | HalfClosedLocal(..) | ReservedRemote
440 )
441 }
442
443 pub fn is_idle(&self) -> bool {
444 matches!(self.inner, Idle)
445 }
446
447 pub fn ensure_recv_open(&self) -> Result<bool, proto::Error> {
448 match self.inner {
450 Closed(Cause::Error(ref e)) => Err(e.clone()),
451 Closed(Cause::ScheduledLibraryReset(reason)) => {
452 Err(proto::Error::library_go_away(reason))
453 }
454 Closed(Cause::EndStream | Cause::ErrorAfterEndStream(_))
455 | HalfClosedRemote(..)
456 | ReservedLocal => Ok(false),
457 _ => Ok(true),
458 }
459 }
460
461 pub(super) fn ensure_reason(&self, mode: PollReset) -> Result<Option<Reason>, crate::Error> {
463 match self.inner {
464 Closed(Cause::Error(Error::Reset(_, reason, _)))
465 | Closed(Cause::ErrorAfterEndStream(Error::Reset(_, reason, _)))
466 | Closed(Cause::Error(Error::GoAway(_, reason, _)))
467 | Closed(Cause::ErrorAfterEndStream(Error::GoAway(_, reason, _)))
468 | Closed(Cause::ScheduledLibraryReset(reason)) => Ok(Some(reason)),
469 Closed(Cause::Error(ref e) | Cause::ErrorAfterEndStream(ref e)) => {
470 Err(e.clone().into())
471 }
472 Open {
473 local: Streaming, ..
474 }
475 | HalfClosedRemote(Streaming) => match mode {
476 PollReset::AwaitingHeaders => Err(UserError::PollResetAfterSendResponse.into()),
477 PollReset::Streaming => Ok(None),
478 },
479 _ => Ok(None),
480 }
481 }
482}
483
484impl Default for State {
485 fn default() -> State {
486 State { inner: Inner::Idle }
487 }
488}
489
490impl fmt::Debug for State {
492 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
493 self.inner.fmt(f)
494 }
495}
496
497#[cfg(test)]
498mod tests {
499 use super::*;
500 use http::HeaderMap;
501
502 #[test]
503 fn recv_reset_preserves_received_end_stream() {
504 let stream_id = StreamId::from(1);
505 let mut state = State::default();
506 state.send_open(false).unwrap();
507
508 let mut headers = frame::Headers::new(stream_id, Default::default(), HeaderMap::new());
509 headers.set_end_stream();
510 state.recv_open(&headers).unwrap();
511 assert!(state.is_recv_end_stream());
512
513 state.recv_reset(frame::Reset::new(stream_id, Reason::NO_ERROR), true);
514
515 assert!(state.is_recv_end_stream());
516 assert_eq!(state.ensure_recv_open().unwrap(), false);
517 assert_eq!(
518 state.ensure_reason(PollReset::Streaming).unwrap(),
519 Some(Reason::NO_ERROR)
520 );
521 }
522}