use std::collections::VecDeque;
use std::ops::Range;
use bytes::{Bytes, BytesMut};
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
use crate::helpers::compression::Compression;
use crate::helpers::scan;
use crate::helpers::text::Text;
use crate::models::{Body, ConnectionID, HeaderCase, Headers, Limits, Message, Method, Role, Version};
use crate::tls::Security;
use crate::protocol::base::Connection;
use crate::protocol::common::{self, Buffer, Error};
use crate::helpers::sync;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct H1Limits {
pub max_message_size: u64,
pub max_message_body_size: u64,
pub max_decompressed_body_size: u64,
pub max_startline_size: u32,
pub max_headers_size: u64,
pub max_header_count: u16,
pub max_chunk_header_size: u32,
pub inline_body_size: u64,
pub max_concurrent_streams: u32,
pub read_chunk_size: u64,
pub idle_capacity: u64,
pub read_timeout: f64,
pub write_timeout: f64,
pub receive_timeout: f64,
pub send_timeout: f64,
}
impl Default for H1Limits {
fn default() -> Self {
Limits::default().into()
}
}
impl From<Limits> for H1Limits {
fn from(limits: Limits) -> Self {
Self {
max_message_size: limits.max_message_size,
max_message_body_size: limits.max_message_body_size,
max_decompressed_body_size: limits.max_decompressed_body_size,
max_startline_size: limits.max_startline_size,
max_headers_size: limits.max_headers_size,
max_header_count: limits.max_header_count,
max_chunk_header_size: limits.max_chunk_header_size,
inline_body_size: limits.inline_body_size,
max_concurrent_streams: limits.max_concurrent_streams,
read_chunk_size: limits.read_chunk_size,
idle_capacity: limits.idle_capacity,
read_timeout: limits.read_timeout,
write_timeout: limits.write_timeout,
receive_timeout: limits.receive_timeout,
send_timeout: limits.send_timeout,
}
}
}
pub struct Line;
impl Line {
pub async fn end<T>(buffer: &mut Buffer, transport: &mut T, max: usize, timeout: f64) -> Result<usize, Error>
where
T: AsyncRead + Unpin,
{
let mut searched = 0;
loop {
if let Some(offset) = scan::find(&buffer.as_slice()[searched..], b'\n') {
let end = searched + offset;
if end == 0 || buffer.as_slice()[end - 1] != b'\r' {
return Err(Error::Protocol("line is not terminated by CRLF".into()));
}
if end - 1 > max {
return Err(Error::Limit(format!("line exceeds {max} octets")));
}
return Ok(end - 1);
}
searched = buffer.len();
if searched > max {
return Err(Error::Limit(format!("line exceeds {max} octets")));
}
if !buffer.fill(transport, timeout).await? {
return Err(Error::Closed);
}
}
}
}
pub struct Persistence;
impl Persistence {
pub fn keep_alive(headers: Option<&Headers>, version: Version) -> bool {
let mut close = false;
let mut keep = false;
if let Some(headers) = headers {
for value in headers.get_all("connection") {
for token in value.split(',') {
let token = token.trim();
if token.eq_ignore_ascii_case("close") {
close = true;
} else if token.eq_ignore_ascii_case("keep-alive") {
keep = true;
}
}
}
}
if close {
return false;
}
match version {
Version::V1_0 => keep,
Version::V1_1 | Version::V2_0 | Version::V3_0 => true,
}
}
}
pub struct Expectation;
impl Expectation {
pub const CONTINUE: &'static str = "100-continue";
pub const STATUS: u16 = 100;
pub fn requested(headers: Option<&Headers>, version: Version) -> bool {
if version == Version::V1_0 {
return false;
}
headers.is_some_and(|headers| {
headers
.get_all("expect")
.flat_map(|value| value.split(','))
.any(|token| token.trim().eq_ignore_ascii_case(Self::CONTINUE))
})
}
}
pub struct StartLine;
impl StartLine {
pub fn write(message: &Message, out: &mut BytesMut) -> Result<(), Error> {
if message.version.major() != 1 {
return Err(Error::Version(format!("{} has no start line", message.version)));
}
let version = message.version.as_str();
if let Some(method) = message.method {
let method = method.as_str();
let target = message.target.as_deref().unwrap_or("/");
if !crate::models::URL::is_target(target) {
return Err(Error::Protocol(format!("request target {target:?} is malformed")));
}
let start = out.len();
out.resize(start + method.len() + target.len() + version.len() + 2, 0);
let line = &mut out[start..];
scan::copy(line, method.as_bytes());
line[method.len()] = b' ';
scan::copy(&mut line[method.len() + 1..], target.as_bytes());
line[method.len() + target.len() + 1] = b' ';
scan::copy(&mut line[method.len() + target.len() + 2..], version.as_bytes());
return Ok(());
}
if let Some(status_code) = message.status_code {
let reason = crate::responses::Status::reason(status_code);
let digits = [
b'0' + (status_code / 100 % 10) as u8,
b'0' + (status_code / 10 % 10) as u8,
b'0' + (status_code % 10) as u8,
];
if (100..1000).contains(&status_code) {
let start = out.len();
out.resize(start + version.len() + reason.len() + 5, 0);
let line = &mut out[start..];
scan::copy(line, version.as_bytes());
line[version.len()] = b' ';
line[version.len() + 1..version.len() + 4].copy_from_slice(&digits);
line[version.len() + 4] = b' ';
scan::copy(&mut line[version.len() + 5..], reason.as_bytes());
return Ok(());
}
out.extend_from_slice(version.as_bytes());
out.extend_from_slice(b" ");
Number::write_decimal(status_code as u64, out);
out.extend_from_slice(b" ");
out.extend_from_slice(reason.as_bytes());
return Ok(());
}
Err(Error::Protocol("message is neither a request nor a response".into()))
}
pub fn encode(message: &Message) -> Result<String, Error> {
let mut out = BytesMut::new();
Self::write(message, &mut out)?;
Ok(String::from_utf8(out.to_vec()).unwrap_or_default())
}
#[inline]
pub fn parse(line: &str) -> Result<Message, Error> {
Self::parse_bytes(line.as_bytes())
}
pub fn parse_bytes(line: &[u8]) -> Result<Message, Error> {
if line.starts_with(b"HTTP/") {
let Some(first) = scan::find(line, b' ') else {
return Err(Error::Protocol("status line has no status code".into()));
};
let (version, rest) = Self::split(line, first);
let (status_code, reason) = match scan::find(rest, b' ') {
Some(second) => Self::split(rest, second),
None => return Err(Error::Protocol("status line has no reason phrase".into())),
};
if status_code.len() != 3 || !status_code.iter().all(u8::is_ascii_digit) {
return Err(Error::Protocol(format!("status code {:?} is not three digits", String::from_utf8_lossy(status_code))));
}
if !Octets::is_reason_bytes(reason) {
return Err(Error::Protocol("reason phrase contains a control character other than a tab".into()));
}
let status_code = u16::from(status_code[0] - b'0') * 100 + u16::from(status_code[1] - b'0') * 10 + u16::from(status_code[2] - b'0');
return Ok(Message::response(status_code, Self::version_bytes(version)?));
}
let Some(first) = scan::find(line, b' ') else {
return Err(Error::Protocol("request line has no target".into()));
};
let (method, rest) = Self::split(line, first);
let (target, version) = match scan::find(rest, b' ') {
Some(second) => Self::split(rest, second),
None => return Err(Error::Protocol("request line has no version".into())),
};
let Some(method) = std::str::from_utf8(method).ok().and_then(|method| method.parse::<Method>().ok()) else {
return Err(Error::Protocol(format!("method {:?} is not recognised", String::from_utf8_lossy(method))));
};
let target = match std::str::from_utf8(target) {
Ok(text) if Octets::is_target(text) => text,
_ => return Err(Error::Protocol(format!("request target {:?} is malformed", String::from_utf8_lossy(target)))),
};
Ok(Message::request(method, target, Self::version_bytes(version)?))
}
#[inline]
pub fn error_status(line: &str) -> u16 {
Self::error_status_bytes(line.as_bytes())
}
pub fn error_status_bytes(line: &[u8]) -> u16 {
let Some(first) = scan::find(line, b' ') else {
return 400;
};
let (method, rest) = Self::split(line, first);
let (target, version) = match scan::find(rest, b' ') {
Some(second) => Self::split(rest, second),
None => return 400,
};
if !matches!(std::str::from_utf8(method), Ok(method) if method.parse::<Method>().is_ok()) {
return 501;
}
if !matches!(std::str::from_utf8(target), Ok(target) if Octets::is_target(target)) {
return 400;
}
if Self::version_bytes(version).is_err() {
return 505;
}
400
}
#[inline]
pub fn version(text: &str) -> Result<Version, Error> {
Self::version_bytes(text.as_bytes())
}
pub fn version_bytes(text: &[u8]) -> Result<Version, Error> {
match text {
b"HTTP/1.0" => Ok(Version::V1_0),
b"HTTP/1.1" => Ok(Version::V1_1),
_ => Err(Error::Version(format!("{:?} is not an HTTP/1.x version", String::from_utf8_lossy(text)))),
}
}
pub fn split(line: &[u8], at: usize) -> (&[u8], &[u8]) {
(&line[..at], &line[at + 1..])
}
}
pub struct Octets;
impl Octets {
pub const TOKEN: u8 = 1 << 0;
pub const FIELD: u8 = 1 << 1;
pub const TARGET: u8 = 1 << 2;
pub const TABLE: &'static [u8; 256] = &{
let mut octets = [0u8; 256];
let mut value = 0usize;
while value < 256 {
let byte = value as u8;
let token = byte.is_ascii_alphanumeric()
|| matches!(
byte,
b'!' | b'#' | b'$' | b'%' | b'&' | b'\'' | b'*' | b'+' | b'-' | b'.' | b'^' | b'_' | b'`' | b'|' | b'~'
);
let field = byte == b'\t' || (byte >= 0x20 && byte != 0x7f);
let target = byte > 0x20 && byte != 0x7f;
octets[value] = (token as u8) | (field as u8) << 1 | (target as u8) << 2;
value += 1;
}
octets
};
pub fn is_control(byte: u8) -> bool {
byte < 0x20 || byte == 0x7f
}
pub fn is_token(text: &str) -> bool {
Self::is_token_bytes(text.as_bytes())
}
#[inline]
pub fn is_target(text: &str) -> bool {
Self::is_target_bytes(text.as_bytes())
}
#[inline]
pub fn is_reason(text: &str) -> bool {
Self::is_reason_bytes(text.as_bytes())
}
#[inline]
pub fn is_token_bytes(text: &[u8]) -> bool {
!text.is_empty() && scan::all_in_class(text, Self::TABLE, Self::TOKEN)
}
#[inline]
pub fn is_target_bytes(text: &[u8]) -> bool {
!text.is_empty() && scan::all_visible(text)
}
#[inline]
pub fn is_reason_bytes(text: &[u8]) -> bool {
scan::all_in_class(text, Self::TABLE, Self::FIELD)
}
}
pub struct Number;
impl Number {
pub fn decimal(mut value: u64, digits: &mut [u8; 20]) -> usize {
let mut index = digits.len();
loop {
index -= 1;
digits[index] = b'0' + (value % 10) as u8;
value /= 10;
if value == 0 {
return index;
}
}
}
pub fn write_decimal(value: u64, out: &mut BytesMut) {
let mut digits = [0u8; 20];
let index = Self::decimal(value, &mut digits);
out.extend_from_slice(&digits[index..]);
}
pub fn hexadecimal(mut value: u64, digits: &mut [u8; 16]) -> usize {
let mut index = digits.len();
loop {
index -= 1;
digits[index] = b"0123456789abcdef"[(value & 0xf) as usize];
value >>= 4;
if value == 0 {
return index;
}
}
}
pub fn write_hexadecimal(value: u64, out: &mut Vec<u8>) {
let mut digits = [0u8; 16];
let index = Self::hexadecimal(value, &mut digits);
out.extend_from_slice(&digits[index..]);
}
}
pub struct FieldSpans {
pub name: Range<usize>,
pub value: Range<usize>,
pub ascii: bool,
}
pub struct Field;
impl Field {
pub fn write(name: &str, value: &str, case: HeaderCase, out: &mut BytesMut) -> Result<(), Error> {
if !Octets::is_token(name) {
return Err(Error::Protocol(format!("field name {name:?} is not a token")));
}
if !scan::is_field_value(value.as_bytes()) {
return Err(Error::Protocol(format!("field value of {name:?} contains a control character")));
}
let start = out.len();
let line = name.len() + value.len() + 4;
out.resize(start + line, 0);
let (head, tail) = out[start..].split_at_mut(name.len());
scan::copy(head, name.as_bytes());
case.apply_in_place(head);
tail[0] = b':';
tail[1] = b' ';
scan::copy(&mut tail[2..], value.as_bytes());
tail[value.len() + 2] = b'\r';
tail[value.len() + 3] = b'\n';
Ok(())
}
pub fn encode(name: &str, value: &str, case: HeaderCase) -> Result<String, Error> {
let mut out = BytesMut::new();
Self::write(name, value, case, &mut out)?;
Ok(String::from_utf8(out.to_vec()).unwrap_or_default())
}
pub fn write_all(headers: &Headers, case: HeaderCase, out: &mut BytesMut) -> Result<(), Error> {
out.reserve(Self::size(headers) as usize);
for (name, value) in headers.iter() {
Self::write(name, value, case, out)?;
}
Ok(())
}
pub fn encode_all(headers: &Headers, case: HeaderCase) -> Result<String, Error> {
let mut out = BytesMut::new();
Self::write_all(headers, case, &mut out)?;
Ok(String::from_utf8(out.to_vec()).unwrap_or_default())
}
pub fn write_content_length(length: u64, case: HeaderCase, out: &mut BytesMut) {
let mut digits = [0u8; 20];
let index = Number::decimal(length, &mut digits);
let name = match case {
HeaderCase::Title => "Content-Length",
HeaderCase::Lower => "content-length",
};
let start = out.len();
let written = digits.len() - index;
out.resize(start + name.len() + written + 4, 0);
let line = &mut out[start..];
scan::copy(line, name.as_bytes());
line[name.len()] = b':';
line[name.len() + 1] = b' ';
scan::copy(&mut line[name.len() + 2..], &digits[index..]);
line[name.len() + written + 2] = b'\r';
line[name.len() + written + 3] = b'\n';
}
pub fn size(headers: &Headers) -> u64 {
headers.iter().map(|(name, value)| (name.len() + value.len() + 4) as u64).sum()
}
#[inline]
pub fn name_end(line: &[u8]) -> Result<usize, Error> {
let Some(at) = scan::find(line, b':') else {
return Err(Error::Protocol(format!("field line {:?} has no colon", String::from_utf8_lossy(line))));
};
if at == 0 {
return Err(Error::Protocol("field line has an empty name".into()));
}
if !scan::all_in_class(&line[..at], Octets::TABLE, Octets::TOKEN) {
return Err(Error::Protocol(format!("field name {:?} is not a token", String::from_utf8_lossy(&line[..at]))));
}
Ok(at)
}
#[inline]
pub fn spans(line: &[u8]) -> Result<FieldSpans, Error> {
let colon = Self::name_end(line)?;
let rest = &line[colon + 1..];
let start = rest.iter().position(|byte| !matches!(byte, b' ' | b'\t')).unwrap_or(rest.len());
let end = rest.iter().rposition(|byte| !matches!(byte, b' ' | b'\t')).map_or(start, |index| index + 1);
let class = scan::classify_field_value(&rest[start..end]);
if class & scan::VALUE_CONTROL != 0 {
return Err(Error::Protocol(format!(
"field value of {:?} contains a control character",
String::from_utf8_lossy(&line[..colon])
)));
}
let value = colon + 1;
Ok(FieldSpans { name: 0..colon, value: value + start..value + end, ascii: class & scan::VALUE_OBS_TEXT == 0 })
}
pub fn parse(line: &str) -> Result<(String, String), Error> {
let spans = Self::spans(line.as_bytes())?;
let name = line.get(spans.name).unwrap_or_default().to_ascii_lowercase();
let value = line.get(spans.value).unwrap_or_default().to_owned();
Ok((name, value))
}
pub fn parse_bytes(line: &[u8]) -> Result<(Text, Text), Error> {
let spans = Self::spans(line)?;
let name = unsafe { Text::from_verified_ascii_lowercase(&line[spans.name]) };
let value = match spans.ascii {
true => unsafe { Text::from_verified_ascii(&line[spans.value]) },
false => Text::from_utf8_lossy(&line[spans.value]),
};
Ok((name, value))
}
pub fn parse_lines(lines: impl IntoIterator<Item = String>) -> Result<Headers, Error> {
let mut headers = Headers::new();
for line in lines {
if line.starts_with([' ', '\t']) {
return Err(Error::Protocol("field line is folded onto a continuation line".into()));
}
let (name, value) = Self::parse_bytes(line.as_bytes())?;
headers.append_lowercase(name, value);
}
Ok(headers)
}
pub fn parse_block(block: &[u8], max_count: usize) -> Result<Headers, Error> {
let mut headers = Headers::with_capacity((block.len() / 32).min(max_count) + 1);
let mut rest = block;
while !rest.is_empty() {
let Some(end) = scan::find(rest, b'\n') else {
return Err(Error::Protocol("line is not terminated by CRLF".into()));
};
if end == 0 || rest[end - 1] != b'\r' {
return Err(Error::Protocol("line is not terminated by CRLF".into()));
}
let line = &rest[..end - 1];
rest = &rest[end + 1..];
if headers.len() >= max_count {
return Err(Error::Limit(format!("more than {max_count} header fields")));
}
if matches!(line.first(), Some(b' ' | b'\t')) {
return Err(Error::Protocol("field line is folded onto a continuation line".into()));
}
let (name, value) = Self::parse_bytes(line)?;
headers.append_lowercase(name, value);
}
Ok(headers)
}
pub fn block_end(data: &[u8], searched: &mut usize) -> Option<(usize, usize)> {
if data.len() >= 2 && data[0] == b'\r' && data[1] == b'\n' {
return Some((0, 2));
}
let mut at = *searched;
while let Some(offset) = scan::find(&data[at..], b'\n') {
let line_end = at + offset;
if data.len() < line_end + 3 {
break;
}
if data[line_end + 1] == b'\r' && data[line_end + 2] == b'\n' {
return Some((line_end + 1, line_end + 3));
}
at = line_end + 1;
}
*searched = at;
None
}
}
pub struct Chunk;
impl Chunk {
pub fn write(data: &[u8], out: &mut BytesMut) {
let mut digits = [0u8; 16];
let index = Number::hexadecimal(data.len() as u64, &mut digits);
out.reserve(digits.len() - index + data.len() + 4);
out.extend_from_slice(&digits[index..]);
out.extend_from_slice(b"\r\n");
out.extend_from_slice(data);
out.extend_from_slice(b"\r\n");
}
pub fn encode(data: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(data.len() + 20);
Number::write_hexadecimal(data.len() as u64, &mut out);
out.extend_from_slice(b"\r\n");
out.extend_from_slice(data);
out.extend_from_slice(b"\r\n");
out
}
pub fn size(digits: &[u8]) -> Result<usize, Error> {
let malformed = || Error::Protocol(format!("chunk size {:?} is not hexadecimal", String::from_utf8_lossy(digits)));
if digits.is_empty() {
return Err(malformed());
}
let mut size = 0usize;
for digit in digits {
let value = match digit {
b'0'..=b'9' => digit - b'0',
b'a'..=b'f' => digit - b'a' + 10,
b'A'..=b'F' => digit - b'A' + 10,
_ => return Err(malformed()),
};
size = size
.checked_mul(16)
.and_then(|size| size.checked_add(value as usize))
.ok_or_else(|| Error::Protocol("chunk size is too large to address".into()))?;
}
Ok(size)
}
pub fn parse_size(data: &[u8]) -> Result<Option<(usize, usize)>, Error> {
let Some(end) = scan::find(data, b'\n') else {
return Ok(None);
};
if end == 0 || data[end - 1] != b'\r' {
return Err(Error::Protocol("chunk size line is not terminated by CRLF".into()));
}
let line = &data[..end - 1];
let digits = match scan::find(line, b';') {
Some(at) => &line[..at],
None => line,
};
Ok(Some((end + 1, Self::size(digits)?)))
}
pub fn decode(data: &[u8]) -> Result<(usize, Range<usize>), Error> {
let Some((start, size)) = Self::parse_size(data)? else {
return Ok((0, 0..0));
};
if size == 0 {
return Ok((start, 0..0));
}
let overflow = || Error::Protocol("chunk size is too large to address".into());
let end = start.checked_add(size).ok_or_else(overflow)?;
let terminator = end.checked_add(2).ok_or_else(overflow)?;
if data.len() < terminator {
return Ok((0, 0..0));
}
if &data[end..terminator] != b"\r\n" {
return Err(Error::Protocol("chunk data is not terminated by CRLF".into()));
}
Ok((terminator, start..end))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BodyLength {
None,
Chunked,
Fixed(u64),
Close,
}
impl BodyLength {
pub const CHUNKED: &'static str = "chunked";
pub fn of(message: &Message, method: Option<Method>) -> Result<Self, Error> {
let headers = message.headers.as_ref();
if let Some(status_code) = message.status_code {
if message.bodyless(method) {
return Ok(Self::None);
}
if method == Some(Method::CONNECT) && (200..300).contains(&status_code) {
return Ok(Self::None);
}
}
let encoded = headers.is_some_and(|headers| headers.contains("transfer-encoding"));
let measured = headers.is_some_and(|headers| headers.contains("content-length"));
if encoded && measured {
return Err(Error::Protocol("Transfer-Encoding and Content-Length are both present".into()));
}
if encoded {
let codings = headers.into_iter().flat_map(|headers| headers.get_all("transfer-encoding")).flat_map(|value| value.split(','));
let applied = codings.filter(|coding| coding.trim().eq_ignore_ascii_case(Self::CHUNKED)).count();
if applied > 1 {
return Err(Error::Protocol("the chunked transfer coding is applied more than once".into()));
}
let last = headers.and_then(|headers| headers.get_all("transfer-encoding").last()).unwrap_or_default();
let last = last.rsplit(',').next().unwrap_or_default().trim();
if !last.eq_ignore_ascii_case(Self::CHUNKED) {
return if message.is_request() || applied > 0 {
Err(Error::Protocol("Transfer-Encoding does not end with chunked".into()))
} else {
Ok(Self::Close)
};
}
return Ok(Self::Chunked);
}
if measured {
let mut values = headers
.into_iter()
.flat_map(|headers| headers.get_all("content-length"))
.flat_map(|value| value.split(','))
.map(str::trim);
let first = values.next().unwrap_or_default();
let length = Self::content_length(first)?;
for value in values {
if Self::content_length(value)? != length {
return Err(Error::Protocol("Content-Length values disagree".into()));
}
}
return Ok(Self::Fixed(length));
}
if message.is_request() { Ok(Self::None) } else { Ok(Self::Close) }
}
pub fn content_length(value: &str) -> Result<u64, Error> {
if value.is_empty() || !value.bytes().all(|byte| byte.is_ascii_digit()) {
return Err(Error::Protocol(format!("Content-Length {value:?} is not a number")));
}
value.parse().map_err(|_| Error::Protocol(format!("Content-Length {value:?} does not fit")))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Exchange {
pub method: Method,
pub accepted: Option<Compression>,
}
impl Exchange {
pub fn new(method: Method, accepted: Option<Compression>) -> Self {
Self { method, accepted }
}
pub fn of(message: &Message) -> Option<Self> {
Some(Self::new(message.method?, message.accepted()))
}
}
pub struct H1Connection<T> {
transport: T,
role: Role,
version: Version,
id: ConnectionID,
client: Option<std::net::SocketAddr>,
limits: H1Limits,
buffer: Buffer,
scratch: BytesMut,
pending: VecDeque<Exchange>,
closing: bool,
request_finalizer: crate::finalizer::RequestFinalizer,
response_finalizer: crate::finalizer::ResponseFinalizer,
security: Security,
}
impl<T> H1Connection<T>
where
T: AsyncRead + AsyncWrite + Unpin,
{
pub fn new(transport: T, role: Role, id: ConnectionID, limits: impl Into<H1Limits>) -> Self {
let limits = limits.into();
Self::resume(transport, role, id, limits, Buffer::new())
}
pub fn resume(transport: T, role: Role, id: ConnectionID, limits: impl Into<H1Limits>, buffer: Buffer) -> Self {
let limits = limits.into();
let mut buffer = buffer;
buffer.set_chunk_size(limits.read_chunk_size as usize);
Self {
transport,
role,
version: Version::V1_1,
id,
client: None,
limits,
buffer,
scratch: BytesMut::new(),
pending: VecDeque::new(),
closing: false,
request_finalizer: crate::finalizer::RequestFinalizer::default(),
response_finalizer: crate::finalizer::ResponseFinalizer::new(None),
security: Security::default(),
}
}
pub fn with_version(mut self, version: Version) -> Self {
if version.major() == 1 {
self.version = version;
}
self
}
pub fn limits(&self) -> H1Limits {
self.limits
}
pub fn with_request_finalizer(mut self, finalizer: crate::finalizer::RequestFinalizer) -> Self {
self.request_finalizer = finalizer;
self
}
pub fn with_response_finalizer(mut self, finalizer: crate::finalizer::ResponseFinalizer) -> Self {
self.response_finalizer = finalizer;
self
}
pub fn with_security(mut self, security: Security) -> Self {
self.security = security;
self
}
pub fn with_client(mut self, client: Option<std::net::SocketAddr>) -> Self {
self.client = client;
self
}
pub fn buffer_capacity(&self) -> usize {
self.buffer.capacity()
}
pub fn scratch_capacity(&self) -> usize {
self.scratch.capacity()
}
pub fn upgrade(self) -> (T, Buffer) {
(self.transport, self.buffer)
}
pub async fn write(&mut self, data: &[u8]) -> Result<(), Error> {
sync::Timeout::within(self.limits.write_timeout, self.transport.write_all(data)).await??;
Ok(())
}
pub async fn write_flushed(&mut self, data: &[u8]) -> Result<(), Error> {
let transport = &mut self.transport;
sync::Timeout::within(self.limits.write_timeout, async move {
transport.write_all(data).await?;
transport.flush().await
})
.await??;
Ok(())
}
pub async fn flush(&mut self) -> Result<(), Error> {
sync::Timeout::within(self.limits.write_timeout, self.transport.flush()).await??;
Ok(())
}
pub async fn reject(&mut self, status_code: u16) -> Result<(), Error> {
self.closing = true;
self.send_message(Message::response(status_code, self.version)).await
}
pub fn pipeline_depth(&self) -> usize {
(self.limits.max_concurrent_streams as usize).max(1)
}
pub fn accepted(&self, message: &Message) -> Option<Compression> {
message.is_response().then(|| self.pending.front()?.accepted).flatten()
}
pub fn answering(&self, message: &Message) -> Option<Method> {
message.is_response().then(|| self.pending.front().map(|exchange| exchange.method)).flatten()
}
pub async fn send_message(&mut self, message: Message) -> Result<(), Error> {
if message.method.is_some() && self.pending.len() >= self.pipeline_depth() {
let reason = format!("more than {} requests are awaiting a response", self.pipeline_depth());
return Err(Error::Limit(reason));
}
let mut message = message;
self.request_finalizer.finalize(self.role, &mut message);
self.response_finalizer.finalize(self.role, self.security.secure, &mut message);
message.materialize().await?;
message.compress(self.accepted(&message))?;
let body = match message.body.take().map(Body::into_inline) {
Some(Ok(data)) => Some(data),
Some(Err(path)) => Some(Bytes::from(Box::pin(tokio::fs::read(path)).await?)),
None => None,
};
let case = HeaderCase::from_version(message.version);
let headers = message.headers.as_ref();
let chunked = headers.is_some_and(|headers| {
headers
.get_all("transfer-encoding")
.any(|value| value.rsplit(',').next().unwrap_or_default().trim().eq_ignore_ascii_case("chunked"))
});
let bodyless = matches!(message.status_code, Some(100..=199 | 204 | 304));
let framed = !message.bodyless(self.answering(&message));
let has_length = headers.is_some_and(|headers| headers.contains("content-length"));
let inline = body.as_ref().is_some_and(|body| body.len() <= self.limits.inline_body_size as usize);
let estimate = 64
+ headers.map_or(0, |headers| headers.len() * 40)
+ if inline { body.as_ref().map_or(0, Bytes::len) + 16 } else { 0 };
let mut out = std::mem::take(&mut self.scratch);
out.clear();
out.reserve(estimate);
let closing = self.closing;
let head = (|| -> Result<(), Error> {
StartLine::write(&message, &mut out)?;
out.extend_from_slice(b"\r\n");
if let Some(headers) = headers {
Field::write_all(headers, case, &mut out)?;
}
if !chunked && !has_length && !bodyless {
match &body {
Some(body) => Field::write_content_length(body.len() as u64, case, &mut out),
None if message.is_response() => Field::write_content_length(0, case, &mut out),
None => {}
}
}
if closing && !headers.is_some_and(|headers| headers.contains("connection")) {
Field::write("connection", "close", case, &mut out)?;
}
out.extend_from_slice(b"\r\n");
Ok(())
})();
if let Err(error) = head {
self.scratch = out;
return Err(error);
}
let trailing = match (chunked && framed, body.filter(|_| framed)) {
(true, body) => {
if let Some(body) = body.filter(|body| !body.is_empty()) {
Chunk::write(&body, &mut out);
}
out.extend_from_slice(b"0\r\n");
if let Some(trailers) = &message.trailers {
Field::write_all(trailers, case, &mut out)?;
}
out.extend_from_slice(b"\r\n");
None
}
(false, Some(body)) if inline => {
out.extend_from_slice(&body);
None
}
(false, body) => body,
};
let method = message.method;
let informational = message.is_informational();
drop(message);
if let Some(body) = trailing {
self.write(&out).await?;
self.write_flushed(&body).await?;
} else {
self.write_flushed(&out).await?;
}
out.clear();
common::Buffer::reclaim_bytes(&mut out, self.limits.idle_capacity as usize);
self.scratch = out;
match method {
Some(method) => self.pending.push_back(Exchange::new(method, None)),
None if self.role.is_server() && !informational => drop(self.pending.pop_front()),
None => {}
}
Ok(())
}
pub async fn receive_message(&mut self) -> Result<Message, Error> {
let max = self.limits.max_startline_size as usize;
let mut length = Line::end(&mut self.buffer, &mut self.transport, max, self.limits.read_timeout).await?;
if length == 0 && self.role.is_server() {
self.buffer.consume(2);
length = Line::end(&mut self.buffer, &mut self.transport, max, self.limits.read_timeout).await?;
}
let mut message = match StartLine::parse_bytes(&self.buffer.as_slice()[..length]) {
Ok(message) => {
self.buffer.consume(length + 2);
message
}
Err(error) => {
let status = self
.role
.is_server()
.then(|| StartLine::error_status_bytes(&self.buffer.as_slice()[..length]));
self.buffer.consume(length + 2);
if let Some(status) = status {
let _ = Box::pin(self.reject(status)).await;
}
return Err(error);
}
};
let (headers, block) = self.read_header_block().await?;
let head = length as u64 + 4 + block as u64;
message.headers = Some(headers);
message.connection_id = Some(self.id.clone());
message.client = self.client;
self.security.apply(&mut message);
self.closing = self.closing || !Persistence::keep_alive(message.headers.as_ref(), message.version);
let limit = self.limits.max_message_size;
let Some(budget) = limit.checked_sub(head) else {
return Err(Error::Limit(format!("message head of {head} octets exceeds {limit}")));
};
if let Some(exchange) = self.role.is_server().then(|| Exchange::of(&message)).flatten() {
if self.pending.len() >= self.pipeline_depth() {
let reason = format!("more than {} requests are awaiting an answer", self.pipeline_depth());
return Err(Error::Limit(reason));
}
self.pending.push_back(exchange);
}
let method = if message.is_response() { self.pending.pop_front().map(|exchange| exchange.method) } else { None };
let length = BodyLength::of(&message, method)?;
if self.role.is_server() && length != BodyLength::None && Expectation::requested(message.headers.as_ref(), message.version) {
let interim = Message::response(Expectation::STATUS, self.version);
Box::pin(self.send_message(interim)).await?;
}
message.body = match length {
BodyLength::None => None,
_ => self.receive_body(length, budget).await?.map(Body::Data),
};
if length == BodyLength::Chunked {
message.trailers = Some(self.read_header_block().await?.0);
}
message.decompress(self.limits.max_decompressed_body_size)?;
self.buffer.reclaim(self.limits.idle_capacity as usize);
Ok(message)
}
pub async fn receive_headers(&mut self) -> Result<Headers, Error> {
Ok(self.read_header_block().await?.0)
}
pub async fn read_header_block(&mut self) -> Result<(Headers, usize), Error> {
let max = self.limits.max_headers_size;
let mut searched = 0usize;
let (fields, consumed) = loop {
if let Some(found) = Field::block_end(self.buffer.as_slice(), &mut searched) {
break found;
}
if self.buffer.len() as u64 > max {
return Err(Error::Limit(format!("header block exceeds {max} octets")));
}
if !self.buffer.fill(&mut self.transport, self.limits.read_timeout).await? {
return Err(Error::Closed);
}
};
if fields as u64 > max {
return Err(Error::Limit(format!("header block exceeds {max} octets")));
}
let headers = Field::parse_block(&self.buffer.as_slice()[..fields], self.limits.max_header_count as usize)?;
self.buffer.consume(consumed);
Ok((headers, consumed))
}
pub async fn receive_body(&mut self, length: BodyLength, budget: u64) -> Result<Option<Bytes>, Error> {
Ok(self.read_body(length, budget).await?.filter(|body| !body.is_empty()))
}
pub async fn read_body(&mut self, length: BodyLength, budget: u64) -> Result<Option<Bytes>, Error> {
let limit = self.limits.max_message_body_size.min(budget);
match length {
BodyLength::None => Ok(None),
BodyLength::Fixed(size) => {
if size > limit {
return Err(Error::Limit(format!("body of {size} octets exceeds {limit}")));
}
let size = usize::try_from(size).map_err(|_| Error::Limit(format!("body of {size} octets exceeds what this platform can address")))?;
self.buffer.require(&mut self.transport, size, self.limits.read_timeout).await?;
Ok(Some(self.buffer.take(size).freeze()))
}
BodyLength::Chunked => {
let mut body = BytesMut::new();
loop {
match Chunk::parse_size(self.buffer.as_slice())? {
Some((_, size)) if (body.len() as u64).saturating_add(size as u64) > limit => {
return Err(Error::Limit(format!("chunked body exceeds {limit} octets")));
}
None if self.buffer.len() > self.limits.max_chunk_header_size as usize => {
return Err(Error::Limit(format!("chunk size line exceeds {} octets", self.limits.max_chunk_header_size)));
}
_ => {}
}
let (consumed, chunk) = Chunk::decode(self.buffer.as_slice())?;
if consumed == 0 {
if !self.buffer.fill(&mut self.transport, self.limits.read_timeout).await? {
return Err(Error::Closed);
}
continue;
}
if chunk.is_empty() {
self.buffer.consume(consumed);
return Ok(Some(body.freeze()));
}
body.extend_from_slice(&self.buffer.as_slice()[chunk]);
self.buffer.consume(consumed);
}
}
BodyLength::Close => {
while self.buffer.fill(&mut self.transport, self.limits.read_timeout).await? {
if self.buffer.len() as u64 > limit {
return Err(Error::Limit(format!("body exceeds {limit} octets")));
}
}
let size = self.buffer.len();
Ok(Some(self.buffer.take(size).freeze()))
}
}
}
}
impl<T> Connection for H1Connection<T>
where
T: AsyncRead + AsyncWrite + Unpin,
{
fn version(&self) -> Version {
self.version
}
fn role(&self) -> Role {
self.role
}
fn id(&self) -> ConnectionID {
self.id.clone()
}
fn reusable(&self) -> bool {
!self.closing
}
fn security(&self) -> Security {
self.security
}
fn client(&self) -> Option<std::net::SocketAddr> {
self.client
}
async fn send(&mut self, message: Message) -> Result<(), Error> {
let timeout = self.limits.send_timeout;
let sending = std::pin::pin!(self.send_message(message));
sync::Timeout::within(timeout, sending).await?
}
async fn receive(&mut self) -> Result<Message, Error> {
let timeout = self.limits.receive_timeout;
let receiving = std::pin::pin!(self.receive_message());
sync::Timeout::within(timeout, receiving).await?
}
async fn close(&mut self) {
let _ = self.transport.flush().await;
let _ = self.transport.shutdown().await;
}
}