use std::{
collections::VecDeque,
future::Future,
pin::Pin,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
task::{Context, Poll},
};
use bytes::Bytes;
use http::{HeaderMap, Method, Response, StatusCode, Uri};
use http_body::{Body, Frame};
use pin_project_lite::pin_project;
use super::hpack::Header;
use crate::early_hints::EarlyHintsReceiver;
#[derive(Debug)]
pub(crate) enum BodyMsg {
Data(Bytes),
Trailers(HeaderMap),
EndStream,
}
#[derive(Debug)]
pub(crate) enum StreamMsg {
Headers {
parts: http::response::Parts,
end_stream: bool,
},
Informational {
parts: http::response::Parts,
},
Data {
data: Bytes,
end_stream: bool,
},
Trailers {
trailers: HeaderMap,
},
Reset {
error_code: u32,
},
Closed,
}
pub(crate) struct StreamEntry {
pub(crate) body_tx: kanal::AsyncSender<BodyMsg>,
pub(crate) reset_tx: kanal::AsyncSender<u32>,
pub(crate) msg_rx: kanal::AsyncReceiver<StreamMsg>,
pub(crate) msg_tx: Option<kanal::AsyncSender<StreamMsg>>,
pub(crate) body_rx: Option<kanal::AsyncReceiver<BodyMsg>>,
pub(crate) reset_rx: Option<kanal::AsyncReceiver<u32>>,
pub(crate) wake_tx: Option<kanal::AsyncSender<()>>,
pub(crate) field_block: Vec<u8>,
pub(crate) pending_end_stream: bool,
pub(crate) request_started: bool,
pub(crate) remote_ended: bool,
pub(crate) local_ended: bool,
pub(crate) content_length: Option<u64>,
pub(crate) data_sum: u64,
pub(crate) trailers_seen: bool,
pub(crate) task_done: bool,
pub(crate) send_window: i64,
pub(crate) pending_data: VecDeque<(Bytes, bool)>,
pub(crate) header_list_size: usize,
}
impl StreamEntry {
#[inline]
pub(crate) fn new(
body_tx: kanal::AsyncSender<BodyMsg>,
reset_tx: kanal::AsyncSender<u32>,
msg_rx: kanal::AsyncReceiver<StreamMsg>,
) -> Self {
StreamEntry {
body_tx,
reset_tx,
msg_rx,
msg_tx: None,
body_rx: None,
reset_rx: None,
wake_tx: None,
field_block: Vec::new(),
pending_end_stream: false,
request_started: false,
remote_ended: false,
local_ended: false,
content_length: None,
data_sum: 0,
trailers_seen: false,
task_done: false,
send_window: 65535,
pending_data: VecDeque::new(),
header_list_size: 0,
}
}
#[inline]
pub(crate) fn extend_block(&mut self, block: &[u8]) {
self.field_block.extend_from_slice(block);
}
#[inline]
pub(crate) fn take_block(&mut self) -> Vec<u8> {
std::mem::take(&mut self.field_block)
}
#[inline]
pub(crate) async fn send_body(&mut self, msg: BodyMsg) -> bool {
self.body_tx.send(msg).await.is_ok()
}
#[inline]
pub(crate) fn send_reset(&self, code: u32) {
let _ = self.reset_tx.try_send(code);
}
}
pub(crate) struct H2Body {
inner: kanal::AsyncReceiver<BodyMsg>,
inner_fut: Option<Pin<Box<kanal::ReceiveFuture<'static, BodyMsg>>>>,
send_continue_body: Option<Arc<AtomicBool>>,
ended: bool,
}
impl H2Body {
#[inline]
pub(crate) fn new(
rx: kanal::AsyncReceiver<BodyMsg>,
send_continue_body: Option<Arc<AtomicBool>>,
) -> Self {
H2Body {
inner: rx,
inner_fut: None,
send_continue_body,
ended: false,
}
}
}
impl Body for H2Body {
type Data = Bytes;
type Error = std::io::Error;
#[inline]
fn poll_frame(
self: Pin<&mut Self>,
cx: &mut Context<'_>,
) -> Poll<Option<Result<Frame<Self::Data>, Self::Error>>> {
let this = self.get_mut();
if this.ended {
return Poll::Ready(None);
}
loop {
if let Some(inner_fut) = &mut this.inner_fut {
match Pin::new(inner_fut).poll(cx) {
Poll::Ready(Ok(BodyMsg::Data(data))) => {
this.inner_fut.take();
return Poll::Ready(Some(Ok(Frame::data(data))));
}
Poll::Ready(Ok(BodyMsg::Trailers(trailers))) => {
this.inner_fut.take();
return Poll::Ready(Some(Ok(Frame::trailers(trailers))));
}
Poll::Ready(Ok(BodyMsg::EndStream)) => {
this.ended = true;
this.inner_fut.take();
return Poll::Ready(None);
}
Poll::Ready(Err(_)) => {
this.ended = true;
this.inner_fut.take();
return Poll::Ready(None);
}
Poll::Pending => {
if let Some(scb) = this.send_continue_body.as_ref() {
scb.store(true, Ordering::Relaxed);
}
return Poll::Pending;
}
}
}
let fut = this.inner.recv();
let fut = unsafe {
std::mem::transmute::<
kanal::ReceiveFuture<'_, BodyMsg>,
kanal::ReceiveFuture<'static, BodyMsg>,
>(fut)
};
this.inner_fut = Some(Box::pin(fut));
}
}
}
pub(crate) struct ParsedRequest {
pub(crate) method: Method,
pub(crate) uri: Uri,
pub(crate) headers: HeaderMap,
pub(crate) content_length: Option<u64>,
pub(crate) expect_continue: bool,
pub(crate) is_connect: bool,
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) struct MalformedRequest;
const CONNECTION_SPECIFIC: &[&[u8]] = &[
b"connection",
b"keep-alive",
b"proxy-connection",
b"transfer-encoding",
b"upgrade",
];
#[inline]
pub(crate) fn is_connection_specific(name: &[u8]) -> bool {
CONNECTION_SPECIFIC
.iter()
.any(|forbidden| **forbidden == *name)
}
#[inline]
pub(crate) fn parse_request(headers: &[Header]) -> Result<ParsedRequest, MalformedRequest> {
let mut method: Option<&[u8]> = None;
let mut scheme: Option<&[u8]> = None;
let mut authority: Option<&[u8]> = None;
let mut path: Option<&[u8]> = None;
let mut protocol: Option<&[u8]> = None;
let mut regular = HeaderMap::new();
let mut content_length: Option<u64> = None;
let mut content_length_conflict = false;
let mut pseudo_phase = true;
for header in headers {
let name = header.name();
let value = header.value();
if name.first() == Some(&b':') {
if !pseudo_phase {
return Err(MalformedRequest);
}
match name {
b":method" => {
if method.is_some() {
return Err(MalformedRequest);
}
method = Some(value);
}
b":scheme" => {
if scheme.is_some() {
return Err(MalformedRequest);
}
scheme = Some(value);
}
b":authority" => {
if authority.is_some() {
return Err(MalformedRequest);
}
authority = Some(value);
}
b":path" => {
if path.is_some() {
return Err(MalformedRequest);
}
path = Some(value);
}
b":protocol" => {
if protocol.is_some() {
return Err(MalformedRequest);
}
protocol = Some(value);
}
_ => return Err(MalformedRequest),
}
} else {
pseudo_phase = false;
if is_connection_specific(name) {
return Err(MalformedRequest);
}
if name == b"te" && !te_is_trailers(value) {
return Err(MalformedRequest);
}
if name == b"content-length" {
let value = parse_content_length(value)?;
if let Some(previous) = content_length {
if previous != value {
content_length_conflict = true;
}
} else {
content_length = Some(value);
}
}
if name.iter().any(|byte| byte.is_ascii_uppercase()) {
return Err(MalformedRequest);
}
let name = http::header::HeaderName::from_bytes(name).map_err(|_| MalformedRequest)?;
let value =
http::header::HeaderValue::from_bytes(value).map_err(|_| MalformedRequest)?;
regular.append(name, value);
}
}
let is_connect = method == Some(&b"CONNECT"[..]);
let Some(method) = method else {
return Err(MalformedRequest);
};
if is_connect {
let Some(authority) = authority else {
return Err(MalformedRequest);
};
if scheme.is_some() || path.is_some() {
return Err(MalformedRequest);
}
if content_length_conflict {
return Err(MalformedRequest);
}
let uri = Uri::try_from(authority).map_err(|_| MalformedRequest)?;
return Ok(ParsedRequest {
method: Method::from_bytes(method).map_err(|_| MalformedRequest)?,
uri,
headers: regular,
content_length,
expect_continue: false,
is_connect: true,
});
}
if protocol.is_some() {
return Err(MalformedRequest);
}
let Some(scheme) = scheme else {
return Err(MalformedRequest);
};
let Some(path) = path else {
return Err(MalformedRequest);
};
if path.is_empty() {
return Err(MalformedRequest);
}
let uri = {
let scheme = std::str::from_utf8(scheme).map_err(|_| MalformedRequest)?;
let mut builder = Uri::builder();
builder = builder.scheme(scheme);
if let Some(authority) = authority.as_ref() {
let authority = std::str::from_utf8(authority).map_err(|_| MalformedRequest)?;
builder = builder.authority(authority);
}
let path = std::str::from_utf8(path).map_err(|_| MalformedRequest)?;
match builder.path_and_query(path).build() {
Ok(uri) => uri,
Err(_) => Uri::try_from(path).map_err(|_| MalformedRequest)?,
}
};
let expect_continue = regular
.get(http::header::EXPECT)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.eq_ignore_ascii_case("100-continue"));
if content_length_conflict {
return Err(MalformedRequest);
}
Ok(ParsedRequest {
method: Method::from_bytes(method).map_err(|_| MalformedRequest)?,
uri,
headers: regular,
content_length,
expect_continue,
is_connect: false,
})
}
#[inline]
pub(crate) fn te_is_trailers(value: &[u8]) -> bool {
std::str::from_utf8(value)
.ok()
.is_some_and(|value| value.split(',').all(|part| part.trim() == "trailers"))
}
#[inline]
pub(crate) fn parse_content_length(value: &[u8]) -> Result<u64, MalformedRequest> {
let start = match value
.iter()
.position(|byte| *byte != b' ' && *byte != b'\t')
{
None => return Err(MalformedRequest),
Some(start) => start,
};
let end = value
.iter()
.rposition(|byte| *byte != b' ' && *byte != b'\t')
.unwrap_or(start);
let value = &value[start..=end];
if value.is_empty() || value.iter().any(|byte| !byte.is_ascii_digit()) {
return Err(MalformedRequest);
}
let mut result: u64 = 0;
for &byte in value {
result = result
.checked_mul(10)
.and_then(|n| n.checked_add((byte - b'0') as u64))
.ok_or(MalformedRequest)?;
}
Ok(result)
}
#[inline]
pub(crate) fn parse_trailers(headers: &[Header]) -> Result<HeaderMap, MalformedRequest> {
let mut trailers = HeaderMap::new();
for header in headers {
let name = header.name();
if name.first() == Some(&b':') {
return Err(MalformedRequest);
}
if name.iter().any(|byte| byte.is_ascii_uppercase()) {
return Err(MalformedRequest);
}
let name = http::header::HeaderName::from_bytes(name).map_err(|_| MalformedRequest)?;
let value =
http::header::HeaderValue::from_bytes(header.value()).map_err(|_| MalformedRequest)?;
trailers.append(name, value);
}
Ok(trailers)
}
pub(crate) use service::StreamDriver;
mod service;
#[cfg(test)]
mod tests;