use std::{convert::TryFrom, sync::Arc};
use bytes::{Buf, Bytes};
use http::{Request, StatusCode};
use tokio::sync::mpsc::UnboundedSender;
#[cfg(feature = "tracing")]
use tracing::instrument;
use super::{connection::RequestEnd, stream::RequestStream};
use crate::{
connection::{self, RequestDecodeState},
error::{
Code, StreamError,
connection_error_creators::{CloseStream, HandleFrameStreamErrorOnRequestStream},
internal_error::InternalConnectionError,
},
frame::{FrameStream, FrameStreamError},
proto::{
frame::{Frame, PayloadLen},
headers::Header,
},
qpack,
quic::{self, SendStream, StreamId},
shared_state::{ConnectionState, SharedState},
};
pub struct RequestResolver<C, B>
where
C: quic::Connection<B>,
C::BidiStream: quic::SendStream<B>,
B: Buf,
{
#[doc(hidden)]
pub frame_stream: FrameStream<C::BidiStream, B>,
pub(super) request_end_send: UnboundedSender<StreamId>,
pub(super) send_grease_frame: bool,
pub(super) max_field_section_size: u64,
pub(super) shared: Arc<SharedState>,
pub(super) decode_state: RequestDecodeState,
}
impl<C, B> ConnectionState for RequestResolver<C, B>
where
C: quic::Connection<B>,
B: Buf,
{
fn shared_state(&self) -> &SharedState {
&self.shared
}
}
impl<C, B> CloseStream for RequestResolver<C, B>
where
C: quic::Connection<B>,
B: Buf,
{
}
impl<C, B> RequestResolver<C, B>
where
C: quic::Connection<B>,
B: Buf,
{
#[cfg_attr(feature = "tracing", instrument(skip_all, level = "trace"))]
#[allow(clippy::type_complexity)]
pub async fn resolve_request(
mut self,
) -> Result<(Request<()>, RequestStream<C::BidiStream, B>), StreamError> {
let frame = std::future::poll_fn(|cx| self.frame_stream.poll_next(cx)).await;
let req = self.accept_with_frame(frame)?;
req.resolve().await
}
#[cfg_attr(feature = "tracing", instrument(skip_all, level = "trace"))]
pub fn accept_with_frame(
mut self,
frame: Result<Option<Frame<PayloadLen>>, FrameStreamError>,
) -> Result<ResolvedRequest<C, B>, StreamError> {
let encoded = match frame {
Ok(Some(Frame::Headers(h))) => h,
Ok(None) => {
self.frame_stream.reset(Code::H3_REQUEST_INCOMPLETE.value());
return Err(StreamError::StreamError {
code: Code::H3_REQUEST_INCOMPLETE,
reason: "stream terminated without headers".to_string(),
});
}
Ok(Some(_)) => {
return Err(
self.handle_connection_error_on_stream(InternalConnectionError::new(
Code::H3_FRAME_UNEXPECTED,
"first request frame is not headers".to_string(),
)),
);
}
Err(e) => {
return Err(self.handle_frame_stream_error_on_request_stream(e));
}
};
let request_stream = RequestStream {
request_end: Arc::new(RequestEnd {
request_end: self.request_end_send.clone(),
stream_id: self.frame_stream.send_id(),
}),
inner: connection::RequestStream::with_decode_state(
self.frame_stream,
self.max_field_section_size,
self.shared.clone(),
self.send_grease_frame,
self.decode_state,
),
};
Ok(ResolvedRequest::new(
request_stream,
encoded,
self.max_field_section_size,
))
}
}
pub struct ResolvedRequest<C, B>
where
C: quic::Connection<B>,
B: Buf,
{
request_stream: RequestStream<C::BidiStream, B>,
encoded: Bytes,
max_field_section_size: u64,
}
impl<B, C> ResolvedRequest<C, B>
where
C: quic::Connection<B>,
B: Buf,
{
pub fn new(
request_stream: RequestStream<C::BidiStream, B>,
encoded: Bytes,
max_field_section_size: u64,
) -> Self {
Self {
request_stream,
encoded,
max_field_section_size,
}
}
#[cfg_attr(feature = "tracing", instrument(skip_all, level = "trace"))]
#[allow(clippy::type_complexity)]
pub async fn resolve(
mut self,
) -> Result<(Request<()>, RequestStream<C::BidiStream, B>), StreamError> {
let decoded = match std::future::poll_fn(|cx| {
self.request_stream
.inner
.poll_decode_field_section(cx, &mut self.encoded)
})
.await
{
Ok(decoded) => decoded,
Err(qpack::DecoderError::HeaderTooLong(cancel_size)) => {
self.request_stream
.send_response(
http::Response::builder()
.status(StatusCode::REQUEST_HEADER_FIELDS_TOO_LARGE)
.body(())
.expect("header too big response"),
)
.await?;
return Err(StreamError::HeaderTooBig {
actual_size: cancel_size,
max_size: self.max_field_section_size,
});
}
Err(error) => {
let code = if error.is_internal() {
Code::H3_INTERNAL_ERROR
} else {
Code::QPACK_DECOMPRESSION_FAILED
};
return Err(self.request_stream.handle_connection_error_on_stream(
InternalConnectionError::new(
code,
format!("failed to decode request headers: {error}"),
),
));
}
};
let fields = decoded.fields;
let result = match Header::try_from(fields) {
Ok(header) => match header.into_request_parts() {
Ok(parts) => Ok(parts),
Err(err) => Err(err),
},
Err(err) => Err(err),
};
let (method, uri, protocol, headers, pseudo_sensitivity) = match result {
Ok(parts) => parts,
Err(err) => {
let error_code = err.code();
self.request_stream.stop_stream(error_code);
self.request_stream.stop_sending(error_code);
return Err(StreamError::StreamError {
code: error_code,
reason: format!("rejected request headers: {err}"),
});
}
};
let mut req = http::Request::new(());
*req.method_mut() = method;
*req.uri_mut() = uri;
*req.headers_mut() = headers;
if !pseudo_sensitivity.is_empty() {
req.extensions_mut().insert(pseudo_sensitivity);
}
if let Some(protocol) = protocol {
req.extensions_mut().insert(protocol);
}
*req.version_mut() = http::Version::HTTP_3;
#[cfg(feature = "tracing")]
tracing::trace!("replying with: {:?}", req);
Ok((req, self.request_stream))
}
}