use crate::WasiHttpHooks;
use http::header::Entry;
use http::{HeaderMap, HeaderName, HeaderValue};
use std::fmt;
use std::ops::Deref;
use std::sync::Arc;
use wasmtime::Result;
#[derive(Debug, Clone)]
pub struct FieldMap {
map: Arc<HeaderMap>,
limit: Limit,
size: usize,
}
#[derive(Debug, Clone)]
enum Limit {
Mutable(usize),
Immutable,
}
impl Default for FieldMap {
fn default() -> Self {
Self {
map: Arc::new(HeaderMap::new()),
size: 0,
limit: Limit::Immutable,
}
}
}
impl FieldMap {
pub fn new_immutable(hooks: &mut dyn WasiHttpHooks, mut map: HeaderMap) -> Self {
let forbidden_keys = Vec::from_iter(map.keys().filter_map(|name| {
if hooks.is_forbidden_header(name) {
Some(name.clone())
} else {
None
}
}));
for name in forbidden_keys {
map.remove(&name);
}
let size = Self::content_size(&map);
Self {
map: Arc::new(map),
size,
limit: Limit::Immutable,
}
}
pub fn new_mutable(limit: usize) -> Self {
Self {
map: Arc::new(HeaderMap::new()),
size: 0,
limit: Limit::Mutable(limit),
}
}
pub(crate) fn content_size(map: &HeaderMap) -> usize {
let mut sum = 0;
for key in map.keys() {
sum += header_name_size(key);
}
for value in map.values() {
sum += header_value_size(value);
}
sum
}
pub fn set(
&mut self,
hooks: &mut dyn WasiHttpHooks,
key: String,
values: Vec<Vec<u8>>,
) -> Result<(), FieldMapError> {
let key = key.parse()?;
if hooks.is_forbidden_header(&key) {
return Err(FieldMapError::Forbidden);
}
let (map, limit, size) = self.mutable()?;
let key_size = header_name_size(&key);
let values = values
.into_iter()
.map(|v| parse_header_value(&key, v))
.collect::<Result<Vec<_>, _>>()?;
let values_size = values.iter().map(header_value_size).sum::<usize>();
let mut values = values.into_iter();
let mut entry = match map.try_entry(key)? {
Entry::Vacant(e) => match values.next() {
Some(v) => {
update_size(size, limit, *size + values_size + key_size)?;
e.try_insert_entry(v)?
}
None => return Ok(()),
},
Entry::Occupied(mut e) => {
let prev_values_size = e.iter().map(header_value_size).sum::<usize>();
let _prev = match values.next() {
Some(v) => {
update_size(size, limit, *size - prev_values_size + values_size)?;
e.insert(v);
}
None => {
update_size(size, limit, *size - prev_values_size - key_size)?;
e.remove();
return Ok(());
}
};
e
}
};
for value in values {
entry.append(value);
}
Ok(())
}
pub fn remove_all(
&mut self,
hooks: &mut dyn WasiHttpHooks,
key: String,
) -> Result<Vec<HeaderValue>, FieldMapError> {
let key = key.parse()?;
if hooks.is_forbidden_header(&key) {
return Err(FieldMapError::Forbidden);
}
let (map, _limit, size) = self.mutable()?;
match map.try_entry(key)? {
Entry::Vacant { .. } => Ok(Vec::new()),
Entry::Occupied(e) => {
let (name, value_drain) = e.remove_entry_mult();
let mut removed = header_name_size(&name);
let values = value_drain.collect::<Vec<_>>();
for v in values.iter() {
removed += header_value_size(v);
}
*size -= removed;
Ok(values)
}
}
}
fn mutable(&mut self) -> Result<(&mut HeaderMap, usize, &mut usize), FieldMapError> {
match self.limit {
Limit::Immutable => Err(FieldMapError::Immutable),
Limit::Mutable(limit) => Ok((Arc::make_mut(&mut self.map), limit, &mut self.size)),
}
}
pub fn append(
&mut self,
hooks: &mut dyn WasiHttpHooks,
key: String,
value: Vec<u8>,
) -> Result<bool, FieldMapError> {
let key = key.parse()?;
if hooks.is_forbidden_header(&key) {
return Err(FieldMapError::Forbidden);
}
let value = parse_header_value(&key, value)?;
self.append_raw(key, value)
}
pub fn append_raw(
&mut self,
key: HeaderName,
value: HeaderValue,
) -> Result<bool, FieldMapError> {
let (map, limit, size) = self.mutable()?;
let key_size = header_name_size(&key);
let val_size = header_value_size(&value);
let new_size = if !map.contains_key(&key) {
*size + key_size + val_size
} else {
*size + val_size
};
update_size(size, limit, new_size)?;
let already_present = map.try_append(key, value)?;
self.size = new_size;
Ok(already_present)
}
pub fn set_mutable(&mut self, limit: usize) {
self.limit = Limit::Mutable(limit);
}
pub fn set_immutable(&mut self) {
self.limit = Limit::Immutable;
}
}
fn header_name_size(name: &HeaderName) -> usize {
name.as_str().len() + size_of::<HeaderName>()
}
fn header_value_size(value: &HeaderValue) -> usize {
value.len() + size_of::<HeaderValue>()
}
fn update_size(size: &mut usize, limit: usize, new: usize) -> Result<(), FieldMapError> {
if new > limit {
Err(FieldMapError::TotalSizeTooBig)
} else {
*size = new;
Ok(())
}
}
impl Deref for FieldMap {
type Target = HeaderMap;
fn deref(&self) -> &HeaderMap {
&self.map
}
}
impl From<FieldMap> for HeaderMap {
fn from(map: FieldMap) -> Self {
Arc::unwrap_or_clone(map.map)
}
}
#[derive(Debug, PartialEq, Eq, Clone, Copy)]
pub enum FieldMapError {
Immutable,
TooManyFields,
TotalSizeTooBig,
InvalidHeaderName,
InvalidHeaderValue,
Forbidden,
}
impl fmt::Display for FieldMapError {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
let s = match self {
FieldMapError::Immutable => "cannot mutate an immutable field map",
FieldMapError::TooManyFields => "too many fields in the field map",
FieldMapError::TotalSizeTooBig => "total size of fields exceeds limit",
FieldMapError::InvalidHeaderName => "invalid header name",
FieldMapError::InvalidHeaderValue => "invalid header value",
FieldMapError::Forbidden => "forbidden header name",
};
f.write_str(s)
}
}
impl std::error::Error for FieldMapError {}
impl From<http::header::MaxSizeReached> for FieldMapError {
fn from(_: http::header::MaxSizeReached) -> Self {
Self::TooManyFields
}
}
impl From<http::header::InvalidHeaderName> for FieldMapError {
fn from(_: http::header::InvalidHeaderName) -> Self {
Self::InvalidHeaderName
}
}
impl From<http::header::InvalidHeaderValue> for FieldMapError {
fn from(_: http::header::InvalidHeaderValue) -> Self {
Self::InvalidHeaderValue
}
}
fn parse_header_value(
name: &http::HeaderName,
value: Vec<u8>,
) -> Result<http::HeaderValue, FieldMapError> {
if name == http::header::CONTENT_LENGTH {
let s = str::from_utf8(value.as_ref()).or(Err(FieldMapError::InvalidHeaderValue))?;
if s.is_empty() || !s.bytes().all(|b| b.is_ascii_digit()) {
return Err(FieldMapError::InvalidHeaderValue);
}
let v: u64 = s.parse().or(Err(FieldMapError::InvalidHeaderValue))?;
Ok(v.into())
} else {
let value = value.try_into()?;
Ok(value)
}
}
#[cfg(test)]
mod tests {
use super::{FieldMap, FieldMapError, parse_header_value};
use crate::default_hooks;
use http::header::{CONTENT_LENGTH, CONTENT_TYPE};
#[test]
fn test_immutable() {
let mut map = FieldMap::default();
assert_eq!(
map.set(default_hooks(), "foo".to_owned(), vec![b"bar".to_vec()]),
Err(FieldMapError::Immutable)
);
assert_eq!(
map.append(default_hooks(), "foo".to_owned(), b"bar".to_vec()),
Err(FieldMapError::Immutable)
);
assert_eq!(
map.remove_all(default_hooks(), "foo".to_owned()),
Err(FieldMapError::Immutable)
);
}
#[test]
fn test_limits() {
let mut map = FieldMap::new_mutable(100);
loop {
match map.append(default_hooks(), "foo".to_owned(), b"bar".to_vec()) {
Ok(_) => {}
Err(FieldMapError::TotalSizeTooBig) => break,
Err(e) => panic!("unexpected error: {e}"),
}
}
map = FieldMap::new_mutable(100);
for i in 0.. {
match map.set(
default_hooks(),
"foo".to_owned(),
(0..i).map(|j| format!("bar{j}").into_bytes()).collect(),
) {
Ok(_) => {}
Err(FieldMapError::TotalSizeTooBig) => break,
Err(e) => panic!("unexpected error: {e}"),
}
}
map = FieldMap::new_mutable(100);
for i in 0.. {
match map.set(default_hooks(), format!("foo{i}"), vec![b"bar".to_vec()]) {
Ok(_) => {}
Err(FieldMapError::TotalSizeTooBig) => break,
Err(e) => panic!("unexpected error: {e}"),
}
}
}
#[test]
fn test_size() -> Result<(), FieldMapError> {
let mut map = FieldMap::new_mutable(2000);
let name = "foo".to_owned();
let hooks = default_hooks();
map.append(hooks, name.clone(), b"bar".to_vec())?;
assert!(map.size > 0);
map.remove_all(hooks, name.clone())?;
assert_eq!(map.size, 0);
map.set(hooks, name.clone(), vec![b"bar".to_vec()])?;
assert!(map.size > 0);
map.remove_all(hooks, name.clone())?;
assert_eq!(map.size, 0);
map.set(hooks, name.clone(), vec![])?;
assert_eq!(map.size, 0);
map.set(hooks, name.clone(), vec![b"bar".to_vec()])?;
assert!(map.size > 0);
map.set(hooks, name.clone(), vec![])?;
assert_eq!(map.size, 0);
map.set(hooks, name.clone(), vec![b"bar".to_vec()])?;
assert!(map.size > 0);
map.set(hooks, name.clone(), vec![b"bar".to_vec(), b"baz".to_vec()])?;
assert!(map.size > 0);
map.remove_all(hooks, name.clone())?;
assert_eq!(map.size, 0);
Ok(())
}
#[test]
fn content_length_rejects_non_digits() {
assert!(parse_header_value(&CONTENT_LENGTH, b"0".to_vec()).is_ok());
assert!(parse_header_value(&CONTENT_LENGTH, b"1234".to_vec()).is_ok());
assert!(parse_header_value(&CONTENT_LENGTH, b"+5".to_vec()).is_err());
assert!(parse_header_value(&CONTENT_LENGTH, b"-5".to_vec()).is_err());
assert!(parse_header_value(&CONTENT_LENGTH, b" 5".to_vec()).is_err());
assert!(parse_header_value(&CONTENT_LENGTH, b"".to_vec()).is_err());
assert!(parse_header_value(&CONTENT_TYPE, b"text/plain".to_vec()).is_ok());
}
}