use super::shared::{H1_ONLY_HEADERS, ValidatedRequest};
use crate::{
BufWriter, Buffer, Conn, Headers, KnownHeaderName, Method, ProtocolSession, Status, TypeSet,
Version,
after_send::AfterSend,
body::BodyFraming,
h3::{Frame, FrameStream, H3Connection, H3Error, H3ErrorCode},
headers::{
date::current_date_header,
qpack::{FieldSection, PseudoHeaders},
},
received_body::ReceivedBodyState,
};
use futures_lite::{AsyncRead, AsyncWrite, AsyncWriteExt};
use std::{io, sync::Arc, time::Instant};
pub(crate) enum H3FirstFrame {
Request {
validated: ValidatedRequest,
start_time: Instant,
},
WebTransport { session_id: u64 },
}
impl<Transport> Conn<Transport>
where
Transport: AsyncRead + AsyncWrite + Unpin + Send + Sync + 'static,
{
pub(crate) async fn process_first_frame_h3(
h3_connection: &H3Connection,
transport: &mut Transport,
buffer: &mut Buffer,
stream_id: u64,
) -> Result<H3FirstFrame, H3Error> {
log::trace!("H3 bidi stream {stream_id}: started");
let start_time = Instant::now();
log::trace!("H3 bidi stream {stream_id}: waiting for first frame");
let field_section = {
let mut frame_stream = FrameStream::new(transport, buffer);
let mut frame = frame_stream
.next()
.await
.map_err(|e| match e {
H3Error::Protocol(_) => H3ErrorCode::FrameUnexpected.into(),
io @ H3Error::Io(_) => io,
})?
.ok_or(H3ErrorCode::RequestIncomplete)?;
match frame.frame() {
Frame::Headers(_) => {
log::trace!("H3 bidi stream {stream_id}: decoding HEADERS frame");
let encoded = frame.buffer_payload().await?;
let result = h3_connection
.decode_field_section(encoded, stream_id)
.await
.inspect_err(|e| {
log::debug!("H3 bidi stream {stream_id}: HEADERS decode error: {e:?}");
})?;
log::trace!("H3 bidi stream {stream_id}: HEADERS decoded:\n{result}");
result
}
Frame::WebTransport(session_id) => {
let session_id = *session_id;
return Ok(H3FirstFrame::WebTransport { session_id });
}
other => {
log::trace!("H3 bidi stream {stream_id}: unexpected first frame {other:?}");
return Err(H3ErrorCode::FrameUnexpected.into());
}
}
};
log::trace!("received:\n{field_section}");
let validated = ValidatedRequest::new(field_section).ok_or(H3ErrorCode::MessageError)?;
Ok(H3FirstFrame::Request {
validated,
start_time,
})
}
pub(crate) async fn send_h3(mut self) -> io::Result<Self> {
self.finalize_response_headers_h3();
let mut output_buffer = Vec::with_capacity(self.context.config.response_buffer_len);
self.encode_headers_h3(&mut output_buffer)?;
let upgrading = self.should_upgrade();
let max_buf = self.context.config.response_buffer_max_len;
let mut bufwriter = BufWriter::new_with_buffer(output_buffer, &mut self.transport, max_buf);
if self.method != Method::Head
&& !matches!(self.status, Some(Status::NotModified | Status::NoContent))
&& let Some(body) = self.response_body.take()
{
let trailers = body
.write_into(&mut bufwriter, BodyFraming::H3Data, &self.context.config)
.await?;
if !upgrading && let Some(trailers) = trailers {
let Some((h3, stream_id)) = self.protocol_session.as_h3() else {
return Err(io::ErrorKind::NotConnected.into());
};
log::trace!("sending trailers: {trailers}");
h3.encode_field_section_framed(
&FieldSection::new(PseudoHeaders::default(), &trailers),
bufwriter.buffer_mut(),
stream_id,
)?;
}
}
bufwriter.flush().await?;
self.after_send.call(true.into());
Ok(self)
}
fn encode_headers_h3(&mut self, buffer: &mut Vec<u8>) -> io::Result<()> {
let pseudo_headers = PseudoHeaders::default().with_status(self.response_status());
let field_section = FieldSection::new(pseudo_headers, &self.response_headers);
log::trace!("sending:\n{field_section}");
let Some((h3, stream_id)) = self.protocol_session.as_h3() else {
return Err(io::ErrorKind::NotConnected.into());
};
h3.encode_field_section_framed(&field_section, buffer, stream_id)
}
pub(crate) fn build_h3(
h3_connection: Arc<H3Connection>,
transport: Transport,
buffer: Buffer,
validated: ValidatedRequest,
start_time: Instant,
stream_id: u64,
) -> Self {
let ValidatedRequest {
method,
path,
authority,
scheme,
protocol,
request_headers,
} = validated;
let response_headers = h3_connection
.context()
.shared_state()
.get::<Headers>()
.cloned()
.unwrap_or_default();
let request_body_state = ReceivedBodyState::new_h3();
Conn {
context: h3_connection.context(),
transport,
request_headers,
method,
version: Version::Http3,
path,
buffer,
response_headers,
status: None,
state: TypeSet::new(),
response_body: None,
request_body_state,
secure: true,
after_send: AfterSend::default(),
start_time,
peer_ip: None,
authority,
scheme,
protocol,
protocol_session: ProtocolSession::Http3 {
connection: h3_connection,
stream_id,
},
request_trailers: None,
upgrade: false,
}
}
pub(super) fn finalize_response_headers_h3(&mut self) {
self.response_headers
.try_insert_with(KnownHeaderName::Date, current_date_header);
if !self.should_upgrade()
&& !matches!(self.status, Some(Status::NotModified | Status::NoContent))
&& let Some(len) = self.body_len()
{
self.response_headers
.try_insert(KnownHeaderName::ContentLength, len);
}
self.response_headers.remove_all(H1_ONLY_HEADERS);
}
}