diff --git a/examples/webtransport_server.rs b/examples/webtransport_server.rs index 98f1ab05..87508b73 100644 --- a/examples/webtransport_server.rs +++ b/examples/webtransport_server.rs @@ -295,9 +295,16 @@ async fn handle_session_and_echo_all_inbound_messages( tokio::spawn( async move { log_result!(echo_stream(send, stream).await); }); } stream = session.accept_bi() => { - if let Some(server::AcceptedBi::BidiStream(_, stream)) = stream? { - let (send, recv) = quic::BidiStream::split(stream); - tokio::spawn( async move { log_result!(echo_stream(send, recv).await); }); + if let Some(resolver) = stream? { + tokio::spawn(async move { + log_result!(async move { + if let server::AcceptedBi::BidiStream(_, stream) = resolver.resolve().await? { + let (send, recv) = quic::BidiStream::split(stream); + echo_stream(send, recv).await?; + } + Ok::<_, anyhow::Error>(()) + }.await); + }); } } else => { diff --git a/h3-webtransport/src/server.rs b/h3-webtransport/src/server.rs index 8d472267..a2d711f7 100644 --- a/h3-webtransport/src/server.rs +++ b/h3-webtransport/src/server.rs @@ -17,7 +17,7 @@ use h3::{ frame::FrameStream, proto::frame::Frame, quic::{self, OpenStreams, WriteBuf}, - server::{Connection, RequestStream}, + server::{Connection, RequestResolver, RequestStream}, ConnectionState, SharedState, }; use h3::{ @@ -177,45 +177,22 @@ where } } - /// Accepts an incoming bidirectional stream or request - pub async fn accept_bi(&self) -> Result>, StreamError> { + /// Accepts an incoming bidirectional stream. + /// + /// The returned resolver reads the first frame to determine whether the stream contains a + /// WebTransport bidirectional stream or an HTTP/3 request. + pub async fn accept_bi(&self) -> Result>, ConnectionError> { let stream = poll_fn(|cx| { let mut conn = self.server_conn.lock().unwrap(); conn.poll_accept_request_stream(cx) }) - .await; + .await?; - let stream = match stream { - Ok(Some(s)) => FrameStream::new(BufRecvStream::new(s)), - Ok(None) => { - // FIXME: is proper HTTP GoAway shutdown required? - return Ok(None); - } - Err(err) => return Err(StreamError::ConnectionError(err)), - }; - - let mut resolver = { self.server_conn.lock().unwrap().create_resolver(stream) }; - // Read the first frame. - // - // This will determine if it is a webtransport bi-stream or a request stream - let frame = poll_fn(|cx| resolver.frame_stream.poll_next(cx)).await; - - match frame { - Ok(None) => Ok(None), - Ok(Some(Frame::WebTransportStream(session_id))) => { - // Take the stream out of the framed reader and split it in half like Paul Allen - let stream = resolver.frame_stream.into_inner(); - Ok(Some(AcceptedBi::BidiStream( - session_id, - BidiStream::new(stream), - ))) - } - // Make the underlying HTTP/3 connection handle the rest - frame => { - let (req, resp) = resolver.accept_with_frame(frame)?.resolve().await?; - Ok(Some(AcceptedBi::Request(req, resp))) - } - } + Ok(stream.map(|stream| { + let stream = FrameStream::new(BufRecvStream::new(stream)); + let resolver = self.server_conn.lock().unwrap().create_resolver(stream); + BiStreamResolver { resolver } + })) } /// Open a new bidirectional stream @@ -248,6 +225,39 @@ where } } +/// Resolves an accepted bidirectional stream as either WebTransport or HTTP/3. +pub struct BiStreamResolver +where + C: quic::Connection, + B: Buf, +{ + resolver: RequestResolver, +} + +impl BiStreamResolver +where + C: quic::Connection, + B: Buf, +{ + /// Reads the first frame and classifies the accepted stream. + pub async fn resolve(mut self) -> Result, StreamError> { + let frame = poll_fn(|cx| self.resolver.frame_stream.poll_next(cx)).await; + + match frame { + Ok(Some(Frame::WebTransportStream(session_id))) => { + // Take the stream out of the framed reader and split it in half like Paul Allen + let stream = self.resolver.frame_stream.into_inner(); + Ok(AcceptedBi::BidiStream(session_id, BidiStream::new(stream))) + } + // Make the underlying HTTP/3 connection handle the rest + frame => { + let (req, resp) = self.resolver.accept_with_frame(frame)?.resolve().await?; + Ok(AcceptedBi::Request(req, resp)) + } + } + } +} + /// Streams are opened, but the initial webtransport header has not been sent type PendingStreams = ( BidiStream<>::BidiStream, B>,