#![allow(clippy::new_without_default)]
use bytes::BufMut;
use http::header::{AsHeaderName, HeaderName, HeaderValue};
use http::request::Builder as ReqBuilder;
use http::request::Parts as ReqParts;
use http::response::Builder as RespBuilder;
use http::response::Parts as RespParts;
use http::uri::Uri;
use pingora_error::{ErrorType::*, OrErr, Result};
use std::borrow::Cow;
use std::ops::Deref;
pub use http::method::Method;
pub use http::status::StatusCode;
pub use http::version::Version;
pub use http::HeaderMap as HMap;
pub mod authority;
use authority::{raw_target_authority, RawTargetAuthority};
mod case_header_name;
use case_header_name::CaseHeaderName;
pub use case_header_name::IntoCaseHeaderName;
pub mod prelude {
pub use crate::RequestHeader;
pub use crate::ResponseHeader;
}
type CaseMap = HMap<CaseHeaderName>;
pub enum HeaderNameVariant<'a> {
Case(&'a CaseHeaderName),
Titled(&'a str),
}
#[derive(Debug)]
pub struct RequestHeader {
base: ReqParts,
header_name_map: Option<CaseMap>,
raw_target: RawTarget,
send_end_stream: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
enum RawTarget {
FromUri,
Verbatim(Box<[u8]>),
Lossy(Box<[u8]>),
}
impl RawTarget {
fn bytes(&self) -> Option<&[u8]> {
match self {
Self::FromUri => None,
Self::Verbatim(target) | Self::Lossy(target) => Some(target.as_ref()),
}
}
fn is_utf8(&self) -> bool {
match self {
Self::FromUri | Self::Verbatim(_) => true,
Self::Lossy(_) => false,
}
}
}
impl AsRef<ReqParts> for RequestHeader {
fn as_ref(&self) -> &ReqParts {
&self.base
}
}
impl Deref for RequestHeader {
type Target = ReqParts;
fn deref(&self) -> &Self::Target {
&self.base
}
}
impl RequestHeader {
fn new_no_case(size_hint: Option<usize>) -> Self {
let mut base = ReqBuilder::new().body(()).unwrap().into_parts().0;
base.headers.reserve(http_header_map_upper_bound(size_hint));
RequestHeader {
base,
header_name_map: None,
raw_target: RawTarget::FromUri,
send_end_stream: true,
}
}
pub fn build(
method: impl TryInto<Method>,
path: &[u8],
size_hint: Option<usize>,
) -> Result<Self> {
let mut req = Self::build_no_case(method, path, size_hint)?;
req.header_name_map = Some(CaseMap::with_capacity(http_header_map_upper_bound(
size_hint,
)));
Ok(req)
}
pub fn build_no_case(
method: impl TryInto<Method>,
path: &[u8],
size_hint: Option<usize>,
) -> Result<Self> {
let mut req = Self::new_no_case(size_hint);
req.base.method = method
.try_into()
.explain_err(InvalidHTTPHeader, |_| "invalid method")?;
req.set_raw_path(path)?;
Ok(req)
}
pub fn append_header(
&mut self,
name: impl IntoCaseHeaderName,
value: impl TryInto<HeaderValue>,
) -> Result<bool> {
let header_value = value
.try_into()
.explain_err(InvalidHTTPHeader, |_| "invalid value while append")?;
append_header_value(
self.header_name_map.as_mut(),
&mut self.base.headers,
name,
header_value,
)
}
pub fn insert_header(
&mut self,
name: impl IntoCaseHeaderName,
value: impl TryInto<HeaderValue>,
) -> Result<()> {
let header_value = value
.try_into()
.explain_err(InvalidHTTPHeader, |_| "invalid value while insert")?;
insert_header_value(
self.header_name_map.as_mut(),
&mut self.base.headers,
name,
header_value,
)
}
pub fn remove_header<'a, N: ?Sized>(&mut self, name: &'a N) -> Option<HeaderValue>
where
&'a N: 'a + AsHeaderName,
{
remove_header(self.header_name_map.as_mut(), &mut self.base.headers, name)
}
pub fn header_to_h1_wire(&self, buf: &mut impl BufMut) {
header_to_h1_wire(self.header_name_map.as_ref(), &self.base.headers, buf)
}
pub fn case_header_iter(&self) -> impl Iterator<Item = (&CaseHeaderName, &HeaderValue)> + '_ {
case_header_iter(self.header_name_map.as_ref(), &self.base.headers)
}
pub fn has_case(&self) -> bool {
self.header_name_map.is_some()
}
pub fn map<F: FnMut(HeaderNameVariant, &HeaderValue) -> Result<()>>(
&self,
mut f: F,
) -> Result<()> {
let key_map = self.header_name_map.as_ref();
let value_map = &self.base.headers;
if let Some(key_map) = key_map {
let iter = key_map.iter().zip(value_map.iter());
for ((header, case_header), (header2, val)) in iter {
if header != header2 {
panic!("header iter mismatch {}, {}", header, header2)
}
f(HeaderNameVariant::Case(case_header), val)?;
}
} else {
for (header, value) in value_map {
let titled_header =
case_header_name::titled_header_name_str(header).unwrap_or(header.as_str());
f(HeaderNameVariant::Titled(titled_header), value)?;
}
}
Ok(())
}
pub fn extensions_mut(&mut self) -> &mut http::Extensions {
&mut self.base.extensions
}
pub fn set_method(&mut self, method: Method) {
self.base.method = method;
}
pub fn set_uri(&mut self, uri: http::Uri) {
self.base.uri = uri;
self.raw_target = RawTarget::FromUri;
}
pub fn set_raw_path(&mut self, path: &[u8]) -> Result<()> {
let parsed = parse_request_target(path)?;
self.base.uri = parsed.uri;
self.raw_target = parsed.raw_target;
Ok(())
}
pub fn set_send_end_stream(&mut self, send_end_stream: bool) {
self.send_end_stream = send_end_stream;
}
pub fn send_end_stream(&self) -> Option<bool> {
if self.base.version != Version::HTTP_2 {
return None;
}
Some(self.send_end_stream)
}
pub fn raw_path(&self) -> &[u8] {
self.raw_target.bytes().unwrap_or_else(|| {
self.base
.uri
.path_and_query()
.map(|path| path.as_str().as_bytes())
.or_else(|| {
self.base
.uri
.authority()
.map(|authority| authority.as_str().as_bytes())
})
.unwrap_or_default()
})
}
pub fn raw_path_is_utf8(&self) -> bool {
self.raw_target.is_utf8()
}
pub fn uri_file_extension(&self) -> Option<&str> {
let (_, ext) = self
.uri
.path_and_query()
.and_then(|pq| pq.path().rsplit_once('.'))?;
Some(ext)
}
pub fn set_version(&mut self, version: Version) {
self.base.version = version;
}
pub fn as_owned_parts(&self) -> ReqParts {
clone_req_parts(&self.base)
}
}
impl Clone for RequestHeader {
fn clone(&self) -> Self {
Self {
base: self.as_owned_parts(),
header_name_map: self.header_name_map.clone(),
raw_target: self.raw_target.clone(),
send_end_stream: self.send_end_stream,
}
}
}
impl From<ReqParts> for RequestHeader {
fn from(parts: ReqParts) -> RequestHeader {
Self {
base: parts,
header_name_map: None,
raw_target: RawTarget::FromUri,
send_end_stream: true,
}
}
}
impl From<RequestHeader> for ReqParts {
fn from(resp: RequestHeader) -> ReqParts {
resp.base
}
}
#[derive(Debug)]
pub struct ResponseHeader {
base: RespParts,
header_name_map: Option<CaseMap>,
reason_phrase: Option<String>,
}
impl AsRef<RespParts> for ResponseHeader {
fn as_ref(&self) -> &RespParts {
&self.base
}
}
impl Deref for ResponseHeader {
type Target = RespParts;
fn deref(&self) -> &Self::Target {
&self.base
}
}
impl Clone for ResponseHeader {
fn clone(&self) -> Self {
Self {
base: self.as_owned_parts(),
header_name_map: self.header_name_map.clone(),
reason_phrase: self.reason_phrase.clone(),
}
}
}
impl From<RespParts> for ResponseHeader {
fn from(parts: RespParts) -> ResponseHeader {
Self {
base: parts,
header_name_map: None,
reason_phrase: None,
}
}
}
impl From<ResponseHeader> for RespParts {
fn from(resp: ResponseHeader) -> RespParts {
resp.base
}
}
impl From<Box<ResponseHeader>> for Box<RespParts> {
fn from(resp: Box<ResponseHeader>) -> Box<RespParts> {
Box::new(resp.base)
}
}
impl ResponseHeader {
fn new(size_hint: Option<usize>) -> Self {
let mut resp_header = Self::new_no_case(size_hint);
resp_header.header_name_map = Some(CaseMap::with_capacity(http_header_map_upper_bound(
size_hint,
)));
resp_header
}
fn new_no_case(size_hint: Option<usize>) -> Self {
let mut base = RespBuilder::new().body(()).unwrap().into_parts().0;
base.headers.reserve(http_header_map_upper_bound(size_hint));
ResponseHeader {
base,
header_name_map: None,
reason_phrase: None,
}
}
pub fn build(code: impl TryInto<StatusCode>, size_hint: Option<usize>) -> Result<Self> {
let mut resp = Self::new(size_hint);
resp.base.status = code
.try_into()
.explain_err(InvalidHTTPHeader, |_| "invalid status")?;
Ok(resp)
}
pub fn build_no_case(code: impl TryInto<StatusCode>, size_hint: Option<usize>) -> Result<Self> {
let mut resp = Self::new_no_case(size_hint);
resp.base.status = code
.try_into()
.explain_err(InvalidHTTPHeader, |_| "invalid status")?;
Ok(resp)
}
pub fn append_header(
&mut self,
name: impl IntoCaseHeaderName,
value: impl TryInto<HeaderValue>,
) -> Result<bool> {
let header_value = value
.try_into()
.explain_err(InvalidHTTPHeader, |_| "invalid value while append")?;
append_header_value(
self.header_name_map.as_mut(),
&mut self.base.headers,
name,
header_value,
)
}
pub fn insert_header(
&mut self,
name: impl IntoCaseHeaderName,
value: impl TryInto<HeaderValue>,
) -> Result<()> {
let header_value = value
.try_into()
.explain_err(InvalidHTTPHeader, |_| "invalid value while insert")?;
insert_header_value(
self.header_name_map.as_mut(),
&mut self.base.headers,
name,
header_value,
)
}
pub fn remove_header<'a, N: ?Sized>(&mut self, name: &'a N) -> Option<HeaderValue>
where
&'a N: 'a + AsHeaderName,
{
remove_header(self.header_name_map.as_mut(), &mut self.base.headers, name)
}
pub fn header_to_h1_wire(&self, buf: &mut impl BufMut) {
header_to_h1_wire(self.header_name_map.as_ref(), &self.base.headers, buf)
}
pub fn case_header_iter(&self) -> impl Iterator<Item = (&CaseHeaderName, &HeaderValue)> + '_ {
case_header_iter(self.header_name_map.as_ref(), &self.base.headers)
}
pub fn has_case(&self) -> bool {
self.header_name_map.is_some()
}
pub fn map<F: FnMut(HeaderNameVariant, &HeaderValue) -> Result<()>>(
&self,
mut f: F,
) -> Result<()> {
let key_map = self.header_name_map.as_ref();
let value_map = &self.base.headers;
if let Some(key_map) = key_map {
let iter = key_map.iter().zip(value_map.iter());
for ((header, case_header), (header2, val)) in iter {
if header != header2 {
panic!("header iter mismatch {}, {}", header, header2)
}
f(HeaderNameVariant::Case(case_header), val)?;
}
} else {
for (header, value) in value_map {
let titled_header =
case_header_name::titled_header_name_str(header).unwrap_or(header.as_str());
f(HeaderNameVariant::Titled(titled_header), value)?;
}
}
Ok(())
}
pub fn extensions_mut(&mut self) -> &mut http::Extensions {
&mut self.base.extensions
}
pub fn set_status(&mut self, status: impl TryInto<StatusCode>) -> Result<()> {
self.base.status = status
.try_into()
.explain_err(InvalidHTTPHeader, |_| "invalid status")?;
Ok(())
}
pub fn set_version(&mut self, version: Version) {
self.base.version = version
}
pub fn set_reason_phrase(&mut self, reason_phrase: Option<&str>) -> Result<()> {
if reason_phrase == self.base.status.canonical_reason() {
self.reason_phrase = None;
return Ok(());
}
self.reason_phrase = reason_phrase.map(str::to_string);
Ok(())
}
pub fn get_reason_phrase(&self) -> Option<&str> {
self.reason_phrase
.as_deref()
.or_else(|| self.base.status.canonical_reason())
}
pub fn as_owned_parts(&self) -> RespParts {
clone_resp_parts(&self.base)
}
pub fn set_content_length(&mut self, len: usize) -> Result<()> {
self.insert_header(http::header::CONTENT_LENGTH, len)
}
}
fn path_and_query_uri(path_and_query: &str, target: &str) -> Result<Uri> {
Uri::builder()
.path_and_query(path_and_query)
.build()
.explain_err(InvalidHTTPHeader, |_| format!("invalid uri {target}"))
}
struct ParsedRequestTarget {
uri: Uri,
raw_target: RawTarget,
}
fn parse_request_target(target: &[u8]) -> Result<ParsedRequestTarget> {
let target = match target.iter().position(|&byte| byte == b'#') {
Some(fragment_start) => &target[..fragment_start],
None => target,
};
if target.is_empty() || matches!(target.first(), Some(b'/' | b'?')) || target == b"*" {
return Ok(match std::str::from_utf8(target) {
Ok(target) => ParsedRequestTarget {
uri: path_and_query_uri(target, target)?,
raw_target: RawTarget::FromUri,
},
Err(_) => {
let lossy = String::from_utf8_lossy(target);
ParsedRequestTarget {
uri: path_and_query_uri(&lossy, &lossy)?,
raw_target: RawTarget::Lossy(target.into()),
}
}
});
}
let lossy_target = String::from_utf8_lossy(target);
let uri = match raw_target_authority(target) {
RawTargetAuthority::Absolute { path_and_query, .. } => {
let path_and_query = String::from_utf8_lossy(path_and_query);
match path_and_query.as_ref() {
"" => Uri::default(),
pq if !pq.starts_with('/') => path_and_query_uri(&format!("/{pq}"), &lossy_target)?,
pq => path_and_query_uri(pq, &lossy_target)?,
}
}
RawTargetAuthority::None | RawTargetAuthority::AmbiguousAuthority => Uri::default(),
};
Ok(ParsedRequestTarget {
uri,
raw_target: match lossy_target {
Cow::Borrowed(_) => RawTarget::Verbatim(target.into()),
Cow::Owned(_) => RawTarget::Lossy(target.into()),
},
})
}
fn clone_req_parts(me: &ReqParts) -> ReqParts {
let mut parts = ReqBuilder::new()
.method(me.method.clone())
.uri(me.uri.clone())
.version(me.version)
.body(())
.unwrap()
.into_parts()
.0;
parts.headers = me.headers.clone();
parts.extensions = me.extensions.clone();
parts
}
fn clone_resp_parts(me: &RespParts) -> RespParts {
let mut parts = RespBuilder::new()
.status(me.status)
.version(me.version)
.body(())
.unwrap()
.into_parts()
.0;
parts.headers = me.headers.clone();
parts.extensions = me.extensions.clone();
parts
}
fn http_header_map_upper_bound(size_hint: Option<usize>) -> usize {
const PINGORA_MAX_HEADER_COUNT: usize = 4096;
const INIT_HEADER_SIZE: usize = 8;
std::cmp::min(
size_hint.unwrap_or(INIT_HEADER_SIZE),
PINGORA_MAX_HEADER_COUNT,
)
}
#[inline]
fn append_header_value<T>(
name_map: Option<&mut CaseMap>,
value_map: &mut HMap<T>,
name: impl IntoCaseHeaderName,
value: T,
) -> Result<bool> {
let case_header_name = name.into_case_header_name();
let header_name: HeaderName = case_header_name
.as_slice()
.try_into()
.or_err(InvalidHTTPHeader, "invalid header name")?;
if let Some(name_map) = name_map {
name_map
.try_append(header_name.clone(), case_header_name)
.or_err(InvalidHTTPHeader, "header name map size overflows MAX_SIZE")?;
}
value_map.try_append(header_name, value).or_err(
InvalidHTTPHeader,
"header value map size overflows MAX_SIZE",
)
}
#[inline]
fn insert_header_value<T>(
name_map: Option<&mut CaseMap>,
value_map: &mut HMap<T>,
name: impl IntoCaseHeaderName,
value: T,
) -> Result<()> {
let case_header_name = name.into_case_header_name();
let header_name: HeaderName = case_header_name
.as_slice()
.try_into()
.or_err(InvalidHTTPHeader, "invalid header name")?;
if let Some(name_map) = name_map {
name_map.insert(header_name.clone(), case_header_name);
}
value_map.insert(header_name, value);
Ok(())
}
#[inline]
fn remove_header<'a, T, N: ?Sized>(
name_map: Option<&mut CaseMap>,
value_map: &mut HMap<T>,
name: &'a N,
) -> Option<T>
where
&'a N: 'a + AsHeaderName,
{
let removed = value_map.remove(name);
if removed.is_some() {
if let Some(name_map) = name_map {
name_map.remove(name);
}
}
removed
}
pub fn header_value_from_raw(raw: impl Into<bytes::Bytes>) -> HeaderValue {
let normalized = normalize_field_value(raw.into());
unsafe { HeaderValue::from_maybe_shared_unchecked(normalized) }
}
pub fn header_value_from_slice(raw: &[u8]) -> HeaderValue {
header_value_from_raw(bytes::Bytes::copy_from_slice(raw))
}
fn normalize_field_value(raw: bytes::Bytes) -> bytes::Bytes {
if !raw.iter().any(|b| matches!(b, b'\r' | b'\n' | b'\0')) {
return raw;
}
let Some(first_nl) = raw.iter().position(|b| *b == b'\n') else {
let replaced: Vec<u8> = raw
.iter()
.map(|&b| if matches!(b, b'\r' | b'\0') { b' ' } else { b })
.collect();
return bytes::Bytes::from(replaced);
};
fn push_with_replacement(dst: &mut Vec<u8>, src: &[u8]) {
dst.extend(
src.iter()
.map(|&b| if matches!(b, b'\r' | b'\0') { b' ' } else { b }),
);
}
let head = raw[..first_nl].trim_ascii_end();
let mut unfolded = Vec::with_capacity(raw.len());
push_with_replacement(&mut unfolded, head);
for line in raw[first_nl + 1..].split(|b| *b == b'\n') {
let line = line.trim_ascii();
if line.is_empty() {
continue;
}
if !unfolded.is_empty() {
unfolded.push(b' ');
}
push_with_replacement(&mut unfolded, line);
}
bytes::Bytes::from(unfolded)
}
#[inline]
fn header_to_h1_wire(key_map: Option<&CaseMap>, value_map: &HMap, buf: &mut impl BufMut) {
const CRLF: &[u8; 2] = b"\r\n";
const HEADER_KV_DELIMITER: &[u8; 2] = b": ";
if let Some(key_map) = key_map {
case_header_iter(key_map.into(), value_map).for_each(|(case_header, val)| {
buf.put_slice(case_header.as_slice());
buf.put_slice(HEADER_KV_DELIMITER);
buf.put_slice(val.as_ref());
buf.put_slice(CRLF);
});
} else {
for (header, value) in value_map {
let titled_header =
case_header_name::titled_header_name_str(header).unwrap_or(header.as_str());
buf.put_slice(titled_header.as_bytes());
buf.put_slice(HEADER_KV_DELIMITER);
buf.put_slice(value.as_ref());
buf.put_slice(CRLF);
}
}
}
#[inline]
fn case_header_iter<'a>(
name_map: Option<&'a CaseMap>,
value_map: &'a HMap,
) -> impl Iterator<Item = (&'a CaseHeaderName, &'a HeaderValue)> + 'a {
name_map.into_iter().flat_map(|name_map| {
name_map
.iter()
.zip(value_map.iter())
.map(|((h1, name), (h2, value))| {
assert_eq!(h1, h2, "header iter mismatch {}, {}", h1, h2);
(name, value)
})
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn header_map_upper_bound() {
assert_eq!(8, http_header_map_upper_bound(None));
assert_eq!(16, http_header_map_upper_bound(Some(16)));
assert_eq!(4096, http_header_map_upper_bound(Some(7777)));
}
#[test]
fn test_single_header() {
let mut req = RequestHeader::build("GET", b"/", None).unwrap();
req.insert_header("foo", "bar").unwrap();
req.insert_header("FoO", "Bar").unwrap();
let mut buf: Vec<u8> = vec![];
req.header_to_h1_wire(&mut buf);
assert_eq!(buf, b"FoO: Bar\r\n");
req.case_header_iter().enumerate().for_each(|(i, (k, v))| {
let name = String::from_utf8_lossy(k.as_slice()).into_owned();
let value = String::from_utf8_lossy(v.as_ref()).into_owned();
match i + 1 {
1 => {
assert_eq!(name, "FoO");
assert_eq!(value, "Bar");
}
_ => panic!("too many headers"),
}
});
let mut resp = ResponseHeader::new(None);
resp.insert_header("foo", "bar").unwrap();
resp.insert_header("FoO", "Bar").unwrap();
let mut buf: Vec<u8> = vec![];
resp.header_to_h1_wire(&mut buf);
assert_eq!(buf, b"FoO: Bar\r\n");
resp.case_header_iter().enumerate().for_each(|(i, (k, v))| {
let name = String::from_utf8_lossy(k.as_slice()).into_owned();
let value = String::from_utf8_lossy(v.as_ref()).into_owned();
match i + 1 {
1 => {
assert_eq!(name, "FoO");
assert_eq!(value, "Bar");
}
_ => panic!("too many headers"),
}
});
}
#[test]
fn test_single_header_no_case() {
let mut req = RequestHeader::new_no_case(None);
req.insert_header("foo", "bar").unwrap();
req.insert_header("FoO", "Bar").unwrap();
let mut buf: Vec<u8> = vec![];
req.header_to_h1_wire(&mut buf);
assert_eq!(buf, b"foo: Bar\r\n");
assert!(req.case_header_iter().next().is_none());
let mut resp = ResponseHeader::new_no_case(None);
resp.insert_header("foo", "bar").unwrap();
resp.insert_header("FoO", "Bar").unwrap();
let mut buf: Vec<u8> = vec![];
resp.header_to_h1_wire(&mut buf);
assert_eq!(buf, b"foo: Bar\r\n");
assert!(resp.case_header_iter().next().is_none());
}
#[test]
fn test_multiple_header() {
let mut req = RequestHeader::build("GET", b"/", None).unwrap();
req.append_header("FoO", "Bar").unwrap();
req.append_header("fOO", "bar").unwrap();
req.append_header("BAZ", "baR").unwrap();
req.append_header(http::header::CONTENT_LENGTH, "0")
.unwrap();
req.append_header("a", "b").unwrap();
req.remove_header("a");
let mut buf: Vec<u8> = vec![];
req.header_to_h1_wire(&mut buf);
assert_eq!(
buf,
b"FoO: Bar\r\nfOO: bar\r\nBAZ: baR\r\nContent-Length: 0\r\n"
);
req.case_header_iter().enumerate().for_each(|(i, (k, v))| {
let name = String::from_utf8_lossy(k.as_slice()).into_owned();
let value = String::from_utf8_lossy(v.as_ref()).into_owned();
match i + 1 {
1 => {
assert_eq!(name, "FoO");
assert_eq!(value, "Bar");
}
2 => {
assert_eq!(name, "fOO");
assert_eq!(value, "bar");
}
3 => {
assert_eq!(name, "BAZ");
assert_eq!(value, "baR");
}
4 => {
assert_eq!(name, "Content-Length");
assert_eq!(value, "0");
}
_ => panic!("too many headers"),
}
});
let mut resp = ResponseHeader::new(None);
resp.append_header("FoO", "Bar").unwrap();
resp.append_header("fOO", "bar").unwrap();
resp.append_header("BAZ", "baR").unwrap();
resp.append_header(http::header::CONTENT_LENGTH, "0")
.unwrap();
resp.append_header("a", "b").unwrap();
resp.remove_header("a");
let mut buf: Vec<u8> = vec![];
resp.header_to_h1_wire(&mut buf);
assert_eq!(
buf,
b"FoO: Bar\r\nfOO: bar\r\nBAZ: baR\r\nContent-Length: 0\r\n"
);
resp.case_header_iter().enumerate().for_each(|(i, (k, v))| {
let name = String::from_utf8_lossy(k.as_slice()).into_owned();
let value = String::from_utf8_lossy(v.as_ref()).into_owned();
match i + 1 {
1 => {
assert_eq!(name, "FoO");
assert_eq!(value, "Bar");
}
2 => {
assert_eq!(name, "fOO");
assert_eq!(value, "bar");
}
3 => {
assert_eq!(name, "BAZ");
assert_eq!(value, "baR");
}
4 => {
assert_eq!(name, "Content-Length");
assert_eq!(value, "0");
}
_ => panic!("too many headers"),
}
});
}
#[test]
fn test_invalid_path() {
let raw_path = b"Hello\xF0\x90\x80World";
let req = RequestHeader::build("GET", &raw_path[..], None).unwrap();
assert_eq!("/", req.uri.path_and_query().unwrap());
assert_eq!(raw_path, req.raw_path());
assert!(!req.raw_path_is_utf8());
}
#[test]
fn test_override_invalid_path() {
let raw_path = b"Hello\xF0\x90\x80World";
let mut req = RequestHeader::build("GET", &raw_path[..], None).unwrap();
assert_eq!("/", req.uri.path_and_query().unwrap());
assert_eq!(raw_path, req.raw_path());
let new_path = "/HelloWorld";
req.set_uri(Uri::builder().path_and_query(new_path).build().unwrap());
assert_eq!(new_path, req.uri.path_and_query().unwrap());
assert_eq!(new_path.as_bytes(), req.raw_path());
assert!(req.raw_path_is_utf8());
}
#[test]
fn test_invalid_path_with_leading_slash_reaches_the_uri() {
for (raw_path, expected) in [
(&b"/Hello\xF0\x90\x80World"[..], "/Hello\u{FFFD}World"),
(b"/Hello\xF0\x90\x80World?q=1", "/Hello\u{FFFD}World?q=1"),
(b"/\xF0\x90\x80", "/\u{FFFD}"),
] {
let req = RequestHeader::build("GET", raw_path, None).unwrap();
let label = String::from_utf8_lossy(raw_path);
assert_eq!(expected, req.uri.path_and_query().unwrap(), "{label}");
assert_eq!(raw_path, req.raw_path(), "{label}");
assert!(!req.raw_path_is_utf8(), "{label}");
}
}
#[test]
fn test_absolute_form_http() {
let req = RequestHeader::build("GET", b"http://host/path?query=1", None).unwrap();
assert_eq!("/path?query=1", req.uri.path_and_query().unwrap().as_str());
assert_eq!("/path", req.uri.path());
assert_eq!(b"http://host/path?query=1", req.raw_path());
}
#[test]
fn test_absolute_form_https() {
let req = RequestHeader::build("GET", b"https://example.com/a/b/c?d=e", None).unwrap();
assert_eq!("/a/b/c?d=e", req.uri.path_and_query().unwrap().as_str());
assert_eq!("/a/b/c", req.uri.path());
}
#[test]
fn test_absolute_form_no_path() {
let req = RequestHeader::build("GET", b"http://host", None).unwrap();
assert_eq!("/", req.uri.path());
assert_eq!(Some("/"), req.uri.path_and_query().map(|pq| pq.as_str()));
assert_eq!(b"http://host", req.raw_path());
}
#[test]
fn test_absolute_form_root() {
let req = RequestHeader::build("GET", b"http://host/", None).unwrap();
assert_eq!("/", req.uri.path());
}
#[test]
fn test_absolute_form_no_path_with_query() {
let req = RequestHeader::build("GET", b"http://host?query=1", None).unwrap();
assert_eq!("/", req.uri.path());
assert_eq!(Some("query=1"), req.uri.query());
assert_eq!("/?query=1", req.uri.path_and_query().unwrap().as_str());
assert_eq!(b"http://host?query=1", req.raw_path());
}
#[test]
fn test_absolute_form_uri_has_no_authority() {
let req = RequestHeader::build("GET", b"http://host:8080/path?q=1", None).unwrap();
assert_eq!(None, req.uri.scheme_str());
assert_eq!(None, req.uri.authority());
assert_eq!("/path", req.uri.path());
assert_eq!(b"http://host:8080/path?q=1", req.raw_path());
}
#[test]
fn test_fragment_is_not_forwarded() {
for (target, raw, path) in [
(&b"http://host/p#frag"[..], &b"http://host/p"[..], "/p"),
(b"http://host#frag", b"http://host", "/"),
(b"http://host?q=1#frag", b"http://host?q=1", "/"),
(b"http://host#@evil.example/", b"http://host", "/"),
] {
let req = RequestHeader::build("GET", target, None).unwrap();
let target = String::from_utf8_lossy(target);
assert_eq!(raw, req.raw_path(), "{target}");
assert_eq!(path, req.uri.path(), "{target}");
}
let req = RequestHeader::build("GET", b"/p#frag", None).unwrap();
assert_eq!(b"/p", req.raw_path());
let req = RequestHeader::build("CONNECT", b"host:443#x", None).unwrap();
assert_eq!(b"host:443", req.raw_path());
let req = RequestHeader::build("GET", b"http://host/p\xff#frag", None).unwrap();
assert_eq!(b"http://host/p\xff", req.raw_path());
assert!(!req.raw_path_is_utf8());
let req = RequestHeader::build("GET", b"/a\xff#frag", None).unwrap();
assert_eq!(b"/a\xff", req.raw_path());
assert!(!req.raw_path_is_utf8());
}
#[test]
fn test_target_that_is_only_a_fragment_falls_back_to_root() {
for target in [&b""[..], b"#", b"#frag", b"#/admin"] {
let req = RequestHeader::build("GET", target, None).unwrap();
let label = String::from_utf8_lossy(target);
assert_eq!(b"/", req.raw_path(), "{label}");
assert_eq!(RawTarget::FromUri, req.raw_target, "{label}");
}
let req = RequestHeader::build("OPTIONS", b"*#frag", None).unwrap();
assert_eq!(b"*", req.raw_path());
assert_eq!(Some("*"), req.uri.path_and_query().map(|pq| pq.as_str()));
let req = RequestHeader::build("GET", b"?q=1#frag", None).unwrap();
assert_eq!(b"?q=1", req.raw_path());
assert_eq!(Some("q=1"), req.uri.query());
}
#[test]
fn test_non_utf8_target_still_yields_a_path() {
let req = RequestHeader::build("GET", b"http://host/p\xff", None).unwrap();
assert_eq!(b"http://host/p\xff", req.raw_path());
assert!(!req.raw_path_is_utf8());
assert_eq!(None, req.uri.authority());
assert_eq!("/p\u{FFFD}", req.uri.path());
let req = RequestHeader::build("CONNECT", b"ho\xffst:443", None).unwrap();
assert_eq!(b"ho\xffst:443", req.raw_path());
assert_eq!("/", req.uri.path());
assert!(!req.raw_path_is_utf8());
let req = RequestHeader::build("GET", b"/p-\xff", None).unwrap();
assert_eq!(b"/p-\xff", req.raw_path());
assert_eq!("/p-\u{FFFD}", req.uri.path());
}
#[test]
fn test_query_only_target_keeps_its_query() {
let req = RequestHeader::build("GET", b"?q=1", None).unwrap();
assert_eq!(b"?q=1", req.raw_path());
assert_eq!(Some("q=1"), req.uri.query());
assert_eq!("/", req.uri.path());
assert_eq!(RawTarget::FromUri, req.raw_target);
}
#[test]
fn test_set_uri_clears_non_origin_form_target() {
let mut req = RequestHeader::build("GET", b"http://host/abs?q=1", None).unwrap();
req.set_uri("/replaced".parse().unwrap());
assert_eq!(b"/replaced", req.raw_path());
assert_eq!(RawTarget::FromUri, req.raw_target);
let mut req = RequestHeader::build("CONNECT", b"example.com:443", None).unwrap();
req.set_uri("/replaced".parse().unwrap());
assert_eq!(b"/replaced", req.raw_path());
assert!(req.raw_path_is_utf8());
}
#[test]
fn test_unclassifiable_targets_are_anchored_to_the_root() {
for target in [
&b"foo:bar://evil.example/admin"[..],
b"myproto:x://evil.example/admin",
b"myproto:opaque",
] {
let req = RequestHeader::build("GET", target, None).unwrap();
let label = String::from_utf8_lossy(target);
assert_eq!(
RawTargetAuthority::None,
raw_target_authority(req.raw_path()),
"{label}"
);
assert_eq!("/", req.uri.path(), "{label}");
assert_eq!(target, req.raw_path(), "{label}");
}
let req = RequestHeader::build("GET", b"http:///path", None).unwrap();
assert_eq!("/", req.uri.path());
assert_eq!(b"http:///path", req.raw_path());
}
#[test]
fn test_relative_target_without_a_scheme_is_anchored_to_the_root() {
for target in [&b"foo/bar"[..], b"host/admin", b"foo", b"foo?q=1"] {
let req = RequestHeader::build("GET", target, None).unwrap();
let label = String::from_utf8_lossy(target);
assert_eq!(
RawTargetAuthority::None,
raw_target_authority(req.raw_path()),
"{label}"
);
assert_eq!("/", req.uri.path_and_query().unwrap(), "{label}");
assert_eq!(target, req.raw_path(), "{label}");
}
}
#[test]
fn test_origin_form_raw_path_is_byte_identical() {
for target in [
&b"/"[..],
b"/index.html",
b"/a/b/c?d=e&f=g",
b"/%2e%2e/x",
b"/a+b/c%20d",
b"*",
] {
let req = RequestHeader::build("GET", target, None).unwrap();
assert_eq!(
target,
req.raw_path(),
"{}",
String::from_utf8_lossy(target)
);
assert!(req.raw_path_is_utf8());
}
}
#[test]
fn test_non_origin_form_survives_clone_and_parts_round_trip() {
let req = RequestHeader::build("GET", b"http://host:8080/p?q=1", None).unwrap();
let cloned = req.clone();
assert_eq!(req.raw_path(), cloned.raw_path());
assert_eq!(req.uri.path(), cloned.uri.path());
assert_eq!(req.raw_path_is_utf8(), cloned.raw_path_is_utf8());
let from_parts = RequestHeader::from(req.as_owned_parts());
assert_eq!(b"/p?q=1", from_parts.raw_path());
assert!(from_parts.raw_path_is_utf8());
}
#[test]
fn test_connect_authority_form() {
let req = RequestHeader::build("CONNECT", b"example.com:443", None).unwrap();
assert_eq!(b"example.com:443", req.raw_path());
assert_eq!(None, req.uri.authority());
assert_eq!("/", req.uri.path());
}
#[test]
fn test_connect_authority_form_shapes_pass_through() {
for target in [
&b"[v7.x]:443"[..],
b"[vF.a:b~!$&'()*+,;=]:8443",
b"[::1]:443",
b"[2001:db8::1]:8443",
b"127.0.0.1:443",
b"sub.example.com:8080",
b"host-with-dash:1",
b"a_b:443",
] {
let req = RequestHeader::build("CONNECT", target, None).unwrap();
let label = String::from_utf8_lossy(target);
assert_eq!(target, req.raw_path(), "{label}");
assert_eq!("/", req.uri.path(), "{label}");
}
}
#[test]
fn test_set_raw_path_replaces_all_target_state() {
let mut req = RequestHeader::build("GET", b"/path-\xff", None).unwrap();
assert!(!req.raw_path_is_utf8());
assert!(matches!(req.raw_target, RawTarget::Lossy(_)));
req.set_raw_path(b"/plain").unwrap();
assert_eq!(b"/plain", req.raw_path());
assert_eq!("/plain", req.uri.path());
assert!(req.raw_path_is_utf8());
assert_eq!(RawTarget::FromUri, req.raw_target);
req.set_raw_path(b"http://host/abs?q=1").unwrap();
assert_eq!(b"http://host/abs?q=1", req.raw_path());
assert_eq!("/abs", req.uri.path());
assert!(req.raw_path_is_utf8());
}
#[test]
fn test_raw_target_variant_per_request_target_form() {
let req = RequestHeader::build("GET", b"http://host/path", None).unwrap();
assert_eq!(
RawTarget::Verbatim(b"http://host/path".to_vec().into()),
req.raw_target
);
assert!(req.raw_path_is_utf8());
let req = RequestHeader::build("CONNECT", b"example.com:443", None).unwrap();
assert_eq!(
RawTarget::Verbatim(b"example.com:443".to_vec().into()),
req.raw_target
);
assert!(req.raw_path_is_utf8());
let req = RequestHeader::build("GET", b"/path-\xff", None).unwrap();
assert_eq!(
RawTarget::Lossy(b"/path-\xff".to_vec().into()),
req.raw_target
);
assert!(!req.raw_path_is_utf8());
for target in [&b"/path"[..], b"*"] {
let req = RequestHeader::build("GET", target, None).unwrap();
let label = String::from_utf8_lossy(target);
assert_eq!(RawTarget::FromUri, req.raw_target, "{label}");
assert!(req.raw_path_is_utf8(), "{label}");
}
}
#[test]
fn test_set_raw_path_clears_stale_connect_fallback() {
let mut req = RequestHeader::build("CONNECT", b"example.com:443", None).unwrap();
assert_eq!(b"example.com:443", req.raw_path());
req.set_method(Method::GET);
req.set_raw_path(b"/ok").unwrap();
assert_eq!(b"/ok", req.raw_path());
assert_eq!("/ok", req.uri.path());
}
#[test]
fn test_absolute_form_set_raw_path_mutation() {
let mut req = RequestHeader::build("GET", b"/original", None).unwrap();
assert_eq!("/original", req.uri.path());
req.set_raw_path(b"http://host/mutated?q=1").unwrap();
assert_eq!("/mutated?q=1", req.uri.path_and_query().unwrap().as_str());
assert_eq!("/mutated", req.uri.path());
}
#[test]
fn test_absolute_form_with_port() {
let req = RequestHeader::build("GET", b"http://host:8080/path", None).unwrap();
assert_eq!("/path", req.uri.path());
}
#[test]
fn test_absolute_form_uppercase_scheme() {
let req = RequestHeader::build("GET", b"HTTP://HOST/path", None).unwrap();
assert_eq!("/path", req.uri.path());
}
#[test]
fn test_absolute_form_non_http_scheme() {
let req = RequestHeader::build("GET", b"ftp://host/path", None).unwrap();
assert_eq!("/path", req.uri.path());
}
#[test]
fn test_origin_form_unchanged() {
let req = RequestHeader::build("GET", b"/path?q=1", None).unwrap();
assert_eq!("/path?q=1", req.uri.path_and_query().unwrap().as_str());
}
#[test]
fn test_origin_form_with_scheme_in_query() {
let req = RequestHeader::build("GET", b"/redir?url=http://other", None).unwrap();
assert_eq!(
"/redir?url=http://other",
req.uri.path_and_query().unwrap().as_str()
);
}
#[test]
fn test_asterisk_form_unchanged() {
let req = RequestHeader::build("OPTIONS", b"*", None).unwrap();
assert_eq!("*", req.uri.path_and_query().unwrap().as_str());
}
#[test]
fn test_authority_form_raw_path() {
let mut req = RequestHeader::new_no_case(None);
req.set_method(Method::CONNECT);
req.set_uri(Uri::builder().authority("pingora.org:443").build().unwrap());
assert!(req.uri.path_and_query().is_none());
assert_eq!(b"pingora.org:443", req.raw_path());
assert!(req.raw_path_is_utf8());
}
#[test]
fn test_reason_phrase() {
let mut resp = ResponseHeader::new(None);
let reason = resp.get_reason_phrase().unwrap();
assert_eq!(reason, "OK");
resp.set_reason_phrase(Some("FooBar")).unwrap();
let reason = resp.get_reason_phrase().unwrap();
assert_eq!(reason, "FooBar");
resp.set_reason_phrase(Some("OK")).unwrap();
let reason = resp.get_reason_phrase().unwrap();
assert_eq!(reason, "OK");
resp.set_reason_phrase(None).unwrap();
let reason = resp.get_reason_phrase().unwrap();
assert_eq!(reason, "OK");
}
#[test]
fn set_test_send_end_stream() {
let mut req = RequestHeader::build("GET", b"/", None).unwrap();
req.set_send_end_stream(true);
assert!(req.send_end_stream().is_none());
let mut req = RequestHeader::build("GET", b"/", None).unwrap();
req.set_version(Version::HTTP_2);
assert!(req.send_end_stream().unwrap());
req.set_send_end_stream(false);
assert!(!req.send_end_stream().unwrap());
}
#[test]
fn set_test_set_content_length() {
let mut resp = ResponseHeader::new(None);
resp.set_content_length(10).unwrap();
assert_eq!(
b"10",
resp.headers
.get(http::header::CONTENT_LENGTH)
.map(|d| d.as_bytes())
.unwrap()
);
}
#[test]
fn normalize_field_value_no_fold_is_zero_copy() {
let input = bytes::Bytes::from_static(b"text/html; charset=utf-8");
let out = normalize_field_value(input.clone());
assert_eq!(out, input);
assert_eq!(out.as_ptr(), input.as_ptr());
}
#[test]
fn normalize_field_value_single_fold() {
let input = bytes::Bytes::from_static(b"obs\r\n fold");
assert_eq!(&normalize_field_value(input)[..], b"obs fold");
}
#[test]
fn normalize_field_value_multiple_folds_mixed_ws() {
let input = bytes::Bytes::from_static(b"obs\r\n fold\r\n\t line");
assert_eq!(&normalize_field_value(input)[..], b"obs fold line");
}
#[test]
fn normalize_field_value_collapses_long_indent() {
let input =
bytes::Bytes::from_static(b"default-src 'self';\r\n script-src 'self' blob:");
assert_eq!(
&normalize_field_value(input)[..],
b"default-src 'self'; script-src 'self' blob:"
);
}
#[test]
fn normalize_field_value_removes_all_cr_and_lf() {
let input = bytes::Bytes::from_static(b"a\r\n b\r\n c\r\n d");
let out = normalize_field_value(input);
assert!(!out.contains(&b'\r'));
assert!(!out.contains(&b'\n'));
assert_eq!(&out[..], b"a b c d");
}
#[test]
fn normalize_field_value_bare_lf() {
let input = bytes::Bytes::from_static(b"obs\n fold");
assert_eq!(&normalize_field_value(input)[..], b"obs fold");
}
#[test]
fn normalize_field_value_empty() {
let input = bytes::Bytes::new();
let out = normalize_field_value(input.clone());
assert_eq!(out, input);
}
#[test]
fn header_value_from_raw_round_trip() {
let hv = header_value_from_raw(bytes::Bytes::from_static(
b"default-src 'self';\r\n script-src 'self'",
));
assert_eq!(hv.as_bytes(), b"default-src 'self'; script-src 'self'");
}
#[test]
fn header_value_from_raw_passthrough() {
let hv = header_value_from_raw(bytes::Bytes::from_static(b"application/json"));
assert_eq!(hv.as_bytes(), b"application/json");
}
#[test]
fn normalize_field_value_leading_newline_no_spurious_space() {
let input = bytes::Bytes::from_static(b"\r\n fold");
assert_eq!(&normalize_field_value(input)[..], b"fold");
}
#[test]
fn normalize_field_value_trailing_newline_no_spurious_space() {
let input = bytes::Bytes::from_static(b"foo\r\n");
assert_eq!(&normalize_field_value(input)[..], b"foo");
let input = bytes::Bytes::from_static(b"a\r\n b\r\n");
assert_eq!(&normalize_field_value(input)[..], b"a b");
}
#[test]
fn normalize_field_value_only_newlines() {
assert_eq!(
&normalize_field_value(bytes::Bytes::from_static(b"\r\n"))[..],
b""
);
assert_eq!(
&normalize_field_value(bytes::Bytes::from_static(b"\n"))[..],
b""
);
assert_eq!(
&normalize_field_value(bytes::Bytes::from_static(b"\r\n\r\n"))[..],
b""
);
}
#[test]
fn normalize_field_value_replaces_bare_cr_with_sp() {
let cases: &[(&[u8], &[u8])] = &[
(b"foo\rbar", b"foo bar"),
(b"foo\r", b"foo "),
(b"\rfoo", b" foo"),
(b"\r", b" "),
(b"\r\r\r", b" "),
];
for (input, expected) in cases {
let out = normalize_field_value(bytes::Bytes::copy_from_slice(input));
assert_eq!(&out[..], *expected, "input = {input:?}");
assert!(!out.contains(&b'\r'));
assert!(!out.contains(&b'\n'));
assert!(!out.contains(&b'\0'));
}
}
#[test]
fn normalize_field_value_replaces_cr_mid_segment_with_sp() {
let input = bytes::Bytes::from_static(b"foo\rbar\r\n baz");
assert_eq!(&normalize_field_value(input)[..], b"foo bar baz");
}
#[test]
fn normalize_field_value_replaces_nul_with_sp() {
let cases: &[(&[u8], &[u8])] = &[
(b"foo\0bar", b"foo bar"),
(b"\0\0\0", b" "),
(b"foo\0", b"foo "),
(b"foo\r\0\nbar", b"foo bar"),
(b"foo\r\n\0bar", b"foo bar"),
];
for (input, expected) in cases {
let out = normalize_field_value(bytes::Bytes::copy_from_slice(input));
assert_eq!(&out[..], *expected, "input = {input:?}");
assert!(!out.contains(&b'\r'));
assert!(!out.contains(&b'\n'));
assert!(!out.contains(&b'\0'));
}
}
#[test]
fn header_value_from_raw_handles_invalid_bytes() {
let hv = header_value_from_raw(bytes::Bytes::from_static(b"foo\rbar\0baz"));
assert_eq!(hv.as_bytes(), b"foo bar baz");
}
}