#[cfg(feature = "ws")]
pub use crate::ws;
#[cfg(feature = "ws")]
pub use crate::ws::WebSocketUpgrade;
use crate::http::error::Error;
use crate::http::response::Body;
use bytes::Bytes;
#[cfg(feature = "cookies")]
use cookie::{Cookie, CookieJar};
use hyper::header::HeaderMap;
use hyper::{Method, StatusCode, Uri};
use serde::de::DeserializeOwned;
use std::convert::Infallible;
use std::future::Future;
pub trait FromRequestParts<S>: Sized + Send {
type Rejection: crate::http::response::IntoResponse + Send + 'static;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
state: &S,
) -> Result<Self, Self::Rejection>;
}
pub trait FromRequest<S: Sync>: Sized + Send {
type Rejection: crate::http::response::IntoResponse + Send + 'static;
fn from_request(
req: hyper::Request<Body>,
state: &S,
) -> impl Future<Output = Result<Self, Self::Rejection>> + Send;
}
pub(crate) const DEFAULT_MAX_BODY_SIZE: usize = 2 * 1024 * 1024;
#[derive(Debug, Clone, Copy)]
pub(crate) struct MaxBodySize(pub usize);
pub(crate) fn max_body_size(extensions: &hyper::http::Extensions) -> usize {
extensions
.get::<MaxBodySize>()
.map_or(DEFAULT_MAX_BODY_SIZE, |m| m.0)
}
#[derive(Debug, Clone, Copy)]
pub struct DefaultBodyLimit {
limit: Option<usize>,
}
impl DefaultBodyLimit {
#[must_use]
pub const fn max(limit: usize) -> Self {
Self { limit: Some(limit) }
}
#[must_use]
pub const fn disable() -> Self {
Self { limit: None }
}
pub fn into_middleware<S>(
self,
) -> impl Fn(hyper::Request<Body>, crate::routing::middleware::Next<S>) -> BoxedResponseFuture
+ Clone
+ Send
+ Sync
+ 'static
where
S: Send + Sync + 'static,
{
let limit = self.limit.unwrap_or(usize::MAX);
move |mut req: hyper::Request<Body>, next: crate::routing::middleware::Next<S>| {
let _ = req.extensions_mut().insert(MaxBodySize(limit));
Box::pin(next.run(req)) as BoxedResponseFuture
}
}
}
type BoxedResponseFuture = std::pin::Pin<Box<dyn Future<Output = hyper::Response<Body>> + Send>>;
pub trait FromRef<S> {
fn from_ref(state: &S) -> Self;
}
impl<T: Clone> FromRef<T> for T {
fn from_ref(state: &T) -> Self {
state.clone()
}
}
#[derive(Debug, Clone, Copy)]
pub struct State<T>(pub T);
impl<S, T> FromRequestParts<S> for State<T>
where
T: FromRef<S> + Send + Sync + 'static,
{
type Rejection = Infallible;
fn from_request_parts(
_parts: &mut hyper::http::request::Parts,
state: &S,
) -> Result<Self, Self::Rejection> {
Ok(Self(T::from_ref(state)))
}
}
impl<S, T> FromRequest<S> for State<T>
where
S: Sync,
T: FromRef<S> + Send + Sync + 'static,
{
type Rejection = Infallible;
async fn from_request(req: hyper::Request<Body>, state: &S) -> Result<Self, Self::Rejection> {
let (mut parts, _) = req.into_parts();
Self::from_request_parts(&mut parts, state)
}
}
#[cfg(any(feature = "query", feature = "form"))]
#[derive(Debug, Clone)]
struct QueryIter<'de> {
input: &'de str,
}
struct CoercingCowDeserializer<'de> {
val: std::borrow::Cow<'de, str>,
}
impl<'de> serde::de::Deserializer<'de> for CoercingCowDeserializer<'de> {
type Error = serde::de::value::Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
match self.val {
std::borrow::Cow::Borrowed(s) => visitor.visit_borrowed_str(s),
std::borrow::Cow::Owned(s) => visitor.visit_string(s),
}
}
fn deserialize_str<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
match self.val {
std::borrow::Cow::Borrowed(s) => visitor.visit_borrowed_str(s),
std::borrow::Cow::Owned(s) => visitor.visit_string(s),
}
}
fn deserialize_string<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
self.deserialize_str(visitor)
}
fn deserialize_u8<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<u8>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_u8(n)
}
fn deserialize_u16<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<u16>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_u16(n)
}
fn deserialize_u32<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<u32>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_u32(n)
}
fn deserialize_u64<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<u64>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_u64(n)
}
fn deserialize_i8<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<i8>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_i8(n)
}
fn deserialize_i16<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<i16>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_i16(n)
}
fn deserialize_i32<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<i32>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_i32(n)
}
fn deserialize_i64<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<i64>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_i64(n)
}
fn deserialize_f32<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<f32>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_f32(n)
}
fn deserialize_f64<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let n = self
.val
.parse::<f64>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?;
visitor.visit_f64(n)
}
fn deserialize_bool<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let b = match self.val.as_ref() {
"true" | "1" => true,
"false" | "0" => false,
_ => self
.val
.parse::<bool>()
.map_err(|e| serde::de::Error::custom(e.to_string()))?,
};
visitor.visit_bool(b)
}
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_some(self)
}
fn deserialize_enum<V>(
self,
_name: &'static str,
_variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
use serde::de::IntoDeserializer;
visitor.visit_enum(self.val.into_deserializer())
}
serde::forward_to_deserialize_any! {
char bytes byte_buf unit unit_struct newtype_struct
seq tuple tuple_struct map struct identifier ignored_any
}
}
impl<'de> serde::de::IntoDeserializer<'de, serde::de::value::Error>
for CoercingCowDeserializer<'de>
{
type Deserializer = Self;
fn into_deserializer(self) -> Self {
self
}
}
#[cfg(any(feature = "query", feature = "form"))]
impl<'de> Iterator for QueryIter<'de> {
type Item = (std::borrow::Cow<'de, str>, CoercingCowDeserializer<'de>);
fn next(&mut self) -> Option<Self::Item> {
if self.input.is_empty() {
return None;
}
let bytes = self.input.as_bytes();
let len = bytes.len();
let end = bytes.iter().position(|&b| b == b'&').unwrap_or(len);
let pair_str = &self.input[..end];
if end < len {
self.input = &self.input[end + 1..];
} else {
self.input = "";
}
if pair_str.is_empty() {
return self.next();
}
let pair_bytes = pair_str.as_bytes();
let (key_raw, val_raw) = pair_bytes
.iter()
.position(|&b| b == b'=')
.map_or((pair_str, ""), |eq_idx| {
(&pair_str[..eq_idx], &pair_str[eq_idx + 1..])
});
let key = decode_query_param(key_raw);
let val = decode_query_param(val_raw);
Some((key, CoercingCowDeserializer { val }))
}
}
#[cfg(any(feature = "query", feature = "form"))]
fn decode_query_param(s: &str) -> std::borrow::Cow<'_, str> {
let bytes = s.as_bytes();
if !bytes.iter().any(|&b| b == b'%' || b == b'+') {
return std::borrow::Cow::Borrowed(s);
}
let mut decoded = Vec::with_capacity(bytes.len());
let mut i = 0;
while i < bytes.len() {
match bytes[i] {
b'%' if i + 2 < bytes.len() => {
if let Ok(hex) = std::str::from_utf8(&bytes[i + 1..i + 3])
&& let Ok(val) = u8::from_str_radix(hex, 16)
{
decoded.push(val);
i += 3;
continue;
}
decoded.push(b'%');
i += 1;
}
b'+' => {
decoded.push(b' ');
i += 1;
}
b => {
decoded.push(b);
i += 1;
}
}
}
String::from_utf8(decoded)
.map_or_else(|_| std::borrow::Cow::Borrowed(s), std::borrow::Cow::Owned)
}
#[derive(Debug, Clone)]
pub struct Path<T>(pub T);
#[derive(Debug, Clone)]
pub struct PathParams(pub Vec<(std::sync::Arc<str>, String)>);
struct PathDeserializer<'de> {
params: &'de [(std::sync::Arc<str>, String)],
}
impl<'de> PathDeserializer<'de> {
fn single_value(&self) -> Result<&'de str, serde::de::value::Error> {
match self.params {
[(_, v)] => Ok(v.as_str()),
_ => Err(serde::de::Error::custom(format!(
"wrong number of path parameters: expected 1, got {}",
self.params.len()
))),
}
}
const fn value_deserializer(val: &'de str) -> CoercingCowDeserializer<'de> {
CoercingCowDeserializer {
val: std::borrow::Cow::Borrowed(val),
}
}
fn map_deserializer(
&self,
) -> serde::de::value::MapDeserializer<
'de,
impl Iterator<Item = (std::borrow::Cow<'de, str>, CoercingCowDeserializer<'de>)>,
serde::de::value::Error,
> {
serde::de::value::MapDeserializer::new(self.params.iter().map(|(k, v)| {
(
std::borrow::Cow::Borrowed(k.as_ref()),
CoercingCowDeserializer {
val: std::borrow::Cow::Borrowed(v.as_str()),
},
)
}))
}
}
struct PathParamsSeqAccess<'de> {
iter: std::slice::Iter<'de, (std::sync::Arc<str>, String)>,
}
impl<'de> serde::de::SeqAccess<'de> for PathParamsSeqAccess<'de> {
type Error = serde::de::value::Error;
fn next_element_seed<T>(&mut self, seed: T) -> Result<Option<T::Value>, Self::Error>
where
T: serde::de::DeserializeSeed<'de>,
{
match self.iter.next() {
Some((_, v)) => seed
.deserialize(CoercingCowDeserializer {
val: std::borrow::Cow::Borrowed(v.as_str()),
})
.map(Some),
None => Ok(None),
}
}
fn size_hint(&self) -> Option<usize> {
Some(self.iter.len())
}
}
macro_rules! path_deserialize_scalar {
($($method:ident),* $(,)?) => {
$(
fn $method<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
let val = self.single_value()?;
PathDeserializer::value_deserializer(val).$method(visitor)
}
)*
};
}
impl<'de> serde::de::Deserializer<'de> for PathDeserializer<'de> {
type Error = serde::de::value::Error;
fn deserialize_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_map(self.map_deserializer())
}
path_deserialize_scalar!(
deserialize_bool,
deserialize_u8,
deserialize_u16,
deserialize_u32,
deserialize_u64,
deserialize_i8,
deserialize_i16,
deserialize_i32,
deserialize_i64,
deserialize_f32,
deserialize_f64,
deserialize_char,
deserialize_str,
deserialize_string,
deserialize_bytes,
deserialize_byte_buf,
);
fn deserialize_option<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_some(self)
}
fn deserialize_unit<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_unit()
}
fn deserialize_unit_struct<V>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_unit()
}
fn deserialize_newtype_struct<V>(
self,
_name: &'static str,
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_newtype_struct(self)
}
fn deserialize_seq<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_seq(PathParamsSeqAccess {
iter: self.params.iter(),
})
}
fn deserialize_tuple<V>(self, len: usize, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
if self.params.len() != len {
return Err(serde::de::Error::custom(format!(
"wrong number of path parameters: expected {len}, got {}",
self.params.len()
)));
}
self.deserialize_seq(visitor)
}
fn deserialize_tuple_struct<V>(
self,
_name: &'static str,
len: usize,
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
self.deserialize_tuple(len, visitor)
}
fn deserialize_map<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_map(self.map_deserializer())
}
fn deserialize_struct<V>(
self,
_name: &'static str,
_fields: &'static [&'static str],
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
self.deserialize_map(visitor)
}
fn deserialize_enum<V>(
self,
_name: &'static str,
_variants: &'static [&'static str],
visitor: V,
) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
use serde::de::IntoDeserializer;
let val = self.single_value()?;
visitor.visit_enum(val.into_deserializer())
}
fn deserialize_identifier<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
self.deserialize_str(visitor)
}
fn deserialize_ignored_any<V>(self, visitor: V) -> Result<V::Value, Self::Error>
where
V: serde::de::Visitor<'de>,
{
visitor.visit_unit()
}
}
impl<S, T> FromRequestParts<S> for Path<T>
where
T: DeserializeOwned + Send + Sync + 'static,
{
type Rejection = Error;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
let params = parts
.extensions
.get::<PathParams>()
.map_or(&[][..], |p| p.0.as_slice());
T::deserialize(PathDeserializer { params })
.map(Path)
.map_err(|e: serde::de::value::Error| Error::Rejection {
status: StatusCode::BAD_REQUEST,
message: format!("Failed to deserialize path parameters: {e}"),
})
}
}
impl<S, T> FromRequest<S> for Path<T>
where
S: Sync,
T: DeserializeOwned + Send + Sync + 'static,
{
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, state: &S) -> Result<Self, Self::Rejection> {
let (mut parts, _) = req.into_parts();
Self::from_request_parts(&mut parts, state)
}
}
#[cfg(feature = "query")]
#[derive(Debug, Clone)]
pub struct Query<T>(pub T);
#[cfg(feature = "query")]
impl<S, T> FromRequestParts<S> for Query<T>
where
T: DeserializeOwned + Send + Sync + 'static,
{
type Rejection = Error;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
let query_str = parts.uri.query().unwrap_or("");
let iter = QueryIter { input: query_str };
let map_de = serde::de::value::MapDeserializer::new(iter);
T::deserialize(map_de)
.map(Query)
.map_err(|e: serde::de::value::Error| Error::Rejection {
status: StatusCode::BAD_REQUEST,
message: format!("Failed to deserialize query parameters: {e}"),
})
}
}
#[cfg(feature = "query")]
impl<S, T> FromRequest<S> for Query<T>
where
S: Sync,
T: DeserializeOwned + Send + Sync + 'static,
{
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, state: &S) -> Result<Self, Self::Rejection> {
let (mut parts, _) = req.into_parts();
Self::from_request_parts(&mut parts, state)
}
}
#[derive(Debug, Clone)]
pub struct RawQuery(pub Option<String>);
impl<S> FromRequestParts<S> for RawQuery {
type Rejection = std::convert::Infallible;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
Ok(Self(parts.uri.query().map(str::to_string)))
}
}
#[cfg(feature = "json")]
fn is_json_content_type(content_type: &str) -> bool {
let essence = content_type.split(';').next().unwrap_or("").trim();
let Some((ty, subtype)) = essence.split_once('/') else {
return false;
};
if !ty.eq_ignore_ascii_case("application") {
return false;
}
subtype.eq_ignore_ascii_case("json") || subtype.to_ascii_lowercase().ends_with("+json")
}
#[cfg(feature = "json")]
#[derive(Debug, Clone)]
pub struct Json<T>(pub T);
#[cfg(feature = "json")]
impl<S, T> FromRequest<S> for Json<T>
where
S: Sync,
T: DeserializeOwned + Send + Sync + 'static,
{
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, _state: &S) -> Result<Self, Self::Rejection> {
let ct = req
.headers()
.get(hyper::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if !is_json_content_type(ct) {
return Err(Error::Rejection {
status: StatusCode::UNSUPPORTED_MEDIA_TYPE,
message: format!("Expected Content-Type: application/json, got: '{ct}'"),
});
}
let limit = max_body_size(req.extensions());
let body = req.into_body().collect_bytes(limit).await?;
serde_json::from_slice::<T>(&body).map(Json).map_err(|e| {
let status = match e.classify() {
serde_json::error::Category::Syntax | serde_json::error::Category::Eof => {
StatusCode::BAD_REQUEST
}
serde_json::error::Category::Data | serde_json::error::Category::Io => {
StatusCode::UNPROCESSABLE_ENTITY
}
};
Error::Rejection {
status,
message: format!("Failed to deserialize JSON payload: {e}"),
}
})
}
}
impl<S> FromRequestParts<S> for HeaderMap {
type Rejection = Infallible;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
Ok(parts.headers.clone())
}
}
impl<S> FromRequestParts<S> for Method {
type Rejection = Infallible;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
Ok(parts.method.clone())
}
}
impl<S> FromRequestParts<S> for Uri {
type Rejection = Infallible;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
Ok(parts.uri.clone())
}
}
impl<S: Sync> FromRequest<S> for Bytes {
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, _state: &S) -> Result<Self, Self::Rejection> {
let limit = max_body_size(req.extensions());
req.into_body().collect_bytes(limit).await
}
}
impl<S: Sync> FromRequest<S> for String {
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, _state: &S) -> Result<Self, Self::Rejection> {
let limit = max_body_size(req.extensions());
let body = req.into_body().collect_bytes(limit).await?;
Self::from_utf8(body.to_vec()).map_err(|e| Error::Rejection {
status: StatusCode::BAD_REQUEST,
message: format!("Request body is not valid UTF-8: {e}"),
})
}
}
#[cfg(feature = "form")]
#[derive(Debug, Clone)]
pub struct Form<T>(pub T);
#[cfg(feature = "form")]
impl<S, T> FromRequestParts<S> for Form<T>
where
T: DeserializeOwned + Send + Sync + 'static,
{
type Rejection = Error;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
let query_str = parts.uri.query().unwrap_or("");
let iter = QueryIter { input: query_str };
let map_de = serde::de::value::MapDeserializer::new(iter);
T::deserialize(map_de)
.map(Form)
.map_err(|e: serde::de::value::Error| Error::Rejection {
status: StatusCode::UNPROCESSABLE_ENTITY,
message: format!("Failed to deserialize form payload: {e}"),
})
}
}
#[cfg(feature = "form")]
impl<S, T> FromRequest<S> for Form<T>
where
S: Sync,
T: DeserializeOwned + Send + Sync + 'static,
{
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, state: &S) -> Result<Self, Self::Rejection> {
if req.method() == hyper::Method::GET || req.method() == hyper::Method::HEAD {
let (mut parts, _body) = req.into_parts();
return Self::from_request_parts(&mut parts, state);
}
let ct = req
.headers()
.get(hyper::header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let essence = ct.split(';').next().unwrap_or("").trim();
if !essence.eq_ignore_ascii_case("application/x-www-form-urlencoded") {
return Err(Error::Rejection {
status: StatusCode::UNSUPPORTED_MEDIA_TYPE,
message: format!(
"Expected Content-Type: application/x-www-form-urlencoded, got: '{ct}'"
),
});
}
let limit = max_body_size(req.extensions());
let body = req.into_body().collect_bytes(limit).await?;
let body_str = std::str::from_utf8(&body).map_err(|_| Error::Rejection {
status: StatusCode::BAD_REQUEST,
message: "Form body is not valid UTF-8".to_string(),
})?;
let iter = QueryIter { input: body_str };
let map_de = serde::de::value::MapDeserializer::new(iter);
T::deserialize(map_de)
.map(Form)
.map_err(|e: serde::de::value::Error| Error::Rejection {
status: StatusCode::UNPROCESSABLE_ENTITY,
message: format!("Failed to deserialize form payload: {e}"),
})
}
}
#[derive(Debug, Clone, Copy)]
pub struct Extension<T>(pub T);
impl<S, T> FromRequestParts<S> for Extension<T>
where
T: Clone + Send + Sync + 'static,
{
type Rejection = Error;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<T>()
.cloned()
.map(Extension)
.ok_or_else(|| Error::Rejection {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!("Missing extension: {}", std::any::type_name::<T>()),
})
}
}
impl<S, T> FromRequest<S> for Extension<T>
where
S: Sync,
T: Clone + Send + Sync + 'static,
{
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, state: &S) -> Result<Self, Self::Rejection> {
let (mut parts, _) = req.into_parts();
Self::from_request_parts(&mut parts, state)
}
}
#[cfg(feature = "cookies")]
#[derive(Debug, Clone)]
pub struct Cookies {
pub jar: CookieJar,
}
#[cfg(feature = "cookies")]
impl Cookies {
#[must_use]
pub fn new() -> Self {
Self {
jar: CookieJar::new(),
}
}
#[must_use]
pub fn get(&self, name: &str) -> Option<&Cookie<'static>> {
self.jar.get(name)
}
#[allow(clippy::should_implement_trait)]
#[must_use]
pub fn add(mut self, cookie: Cookie<'static>) -> Self {
self.jar.add(cookie);
self
}
#[must_use]
pub fn remove(mut self, cookie: Cookie<'static>) -> Self {
self.jar.remove(cookie);
self
}
}
#[cfg(feature = "cookies")]
impl Default for Cookies {
fn default() -> Self {
Self::new()
}
}
#[cfg(feature = "cookies")]
impl<S> FromRequestParts<S> for Cookies {
type Rejection = Infallible;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
let mut jar = CookieJar::new();
if let Some(cookie_header) = parts.headers.get(hyper::header::COOKIE)
&& let Ok(cookie_str) = cookie_header.to_str()
{
for c in Cookie::split_parse_encoded(cookie_str).flatten() {
jar.add_original(c.into_owned());
}
}
Ok(Self { jar })
}
}
impl<S: Sync> FromRequest<S> for hyper::Request<Bytes> {
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, _state: &S) -> Result<Self, Self::Rejection> {
let limit = max_body_size(req.extensions());
let (parts, body) = req.into_parts();
let bytes = body.collect_bytes(limit).await?;
Ok(Self::from_parts(parts, bytes))
}
}
#[derive(Debug)]
pub struct BodyStream(pub Body);
impl<S: Sync> FromRequest<S> for BodyStream {
type Rejection = Infallible;
async fn from_request(req: hyper::Request<Body>, _state: &S) -> Result<Self, Self::Rejection> {
Ok(Self(req.into_body()))
}
}
impl<S: Sync> FromRequest<S> for hyper::Request<Body> {
type Rejection = Infallible;
async fn from_request(req: hyper::Request<Body>, _state: &S) -> Result<Self, Self::Rejection> {
Ok(req)
}
}
#[derive(Debug, Clone)]
pub struct Host(pub String);
impl<S> FromRequestParts<S> for Host {
type Rejection = Error;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
if let Some(host) = parts
.headers
.get(hyper::header::HOST)
.and_then(|h| h.to_str().ok())
{
Ok(Self(host.to_string()))
} else if let Some(host) = parts.uri.host() {
Ok(Self(host.to_string()))
} else {
Err(Error::Rejection {
status: StatusCode::BAD_REQUEST,
message: "Missing Host header or authority in URI".to_string(),
})
}
}
}
#[cfg(feature = "original-uri")]
#[derive(Debug, Clone)]
pub struct OriginalUri(pub Uri);
#[cfg(feature = "original-uri")]
impl<S> FromRequestParts<S> for OriginalUri {
type Rejection = Infallible;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
let uri = parts
.extensions
.get::<Self>()
.map_or_else(|| parts.uri.clone(), |ou| ou.0.clone());
Ok(Self(uri))
}
}
#[cfg(feature = "matched-path")]
#[derive(Debug, Clone)]
pub struct MatchedPath(pub(crate) std::sync::Arc<str>);
#[cfg(feature = "matched-path")]
impl MatchedPath {
#[must_use]
pub fn as_str(&self) -> &str {
&self.0
}
}
#[cfg(feature = "matched-path")]
impl<S> FromRequestParts<S> for MatchedPath {
type Rejection = Error;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<Self>()
.cloned()
.ok_or_else(|| Error::Rejection {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: "No matched path found in request extensions".to_string(),
})
}
}
#[derive(Debug, Clone, Copy)]
pub struct ConnectInfo<T>(pub T);
impl<S, T> FromRequestParts<S> for ConnectInfo<T>
where
T: Clone + Send + Sync + 'static,
{
type Rejection = Error;
fn from_request_parts(
parts: &mut hyper::http::request::Parts,
_state: &S,
) -> Result<Self, Self::Rejection> {
parts
.extensions
.get::<Self>()
.cloned()
.ok_or_else(|| Error::Rejection {
status: StatusCode::INTERNAL_SERVER_ERROR,
message: format!(
"Missing ConnectInfo<{}> extension",
std::any::type_name::<T>()
),
})
}
}
impl<S, T> FromRequest<S> for ConnectInfo<T>
where
S: Sync,
T: Clone + Send + Sync + 'static,
{
type Rejection = Error;
async fn from_request(req: hyper::Request<Body>, state: &S) -> Result<Self, Self::Rejection> {
let (mut parts, _) = req.into_parts();
Self::from_request_parts(&mut parts, state)
}
}
macro_rules! impl_from_request_via_parts {
($ty:ty) => {
impl<S: Sync> FromRequest<S> for $ty {
type Rejection = <Self as FromRequestParts<S>>::Rejection;
async fn from_request(
req: hyper::Request<Body>,
state: &S,
) -> Result<Self, Self::Rejection> {
let (mut parts, _) = req.into_parts();
<Self as FromRequestParts<S>>::from_request_parts(&mut parts, state)
}
}
};
}
impl_from_request_via_parts!(RawQuery);
impl_from_request_via_parts!(HeaderMap);
impl_from_request_via_parts!(Method);
impl_from_request_via_parts!(Uri);
#[cfg(feature = "cookies")]
impl_from_request_via_parts!(Cookies);
impl_from_request_via_parts!(Host);
#[cfg(feature = "original-uri")]
impl_from_request_via_parts!(OriginalUri);
#[cfg(feature = "matched-path")]
impl_from_request_via_parts!(MatchedPath);
#[cfg(test)]
mod tests {
#![allow(clippy::unwrap_used)]
use super::*;
use hyper::http::Request;
use serde::Deserialize;
#[derive(Deserialize, Debug)]
#[allow(clippy::struct_excessive_bools)]
struct BigCoerce {
a: u16,
b: u64,
c: i8,
d: i16,
e: i32,
f: i64,
g: f32,
h: f64,
i: bool,
j: bool,
k: bool,
l: bool,
}
#[derive(Deserialize)]
#[allow(dead_code)]
struct BoolTest {
val: bool,
}
#[derive(Deserialize, PartialEq, Debug)]
enum Color {
Red,
Blue,
}
#[derive(Deserialize)]
struct EnumTest {
val: Color,
}
#[cfg(any(feature = "query", feature = "form"))]
#[test]
fn test_coercing_cow_deserializer() {
let query_str = "a=12&b=34&c=5&d=6&e=7&f=8&g=1.2&h=3.4&i=true&j=1&k=false&l=0";
let iter = QueryIter { input: query_str };
let map_de = serde::de::value::MapDeserializer::new(iter);
let data = BigCoerce::deserialize(map_de).unwrap();
assert_eq!(data.a, 12);
assert_eq!(data.b, 34);
assert_eq!(data.c, 5);
assert_eq!(data.d, 6);
assert_eq!(data.e, 7);
assert_eq!(data.f, 8);
assert!((data.g - 1.2).abs() < 0.001);
assert!((data.h - 3.4).abs() < 0.001);
assert!(data.i);
assert!(data.j);
assert!(!data.k);
assert!(!data.l);
let iter = QueryIter {
input: "val=not_bool",
};
let map_de = serde::de::value::MapDeserializer::new(iter);
assert!(BoolTest::deserialize(map_de).is_err());
let iter = QueryIter { input: "val=Red" };
let map_de = serde::de::value::MapDeserializer::new(iter);
let et = EnumTest::deserialize(map_de).unwrap();
assert_eq!(et.val, Color::Red);
}
#[cfg(any(feature = "query", feature = "form"))]
#[test]
fn test_query_iter_edge_cases() {
let query_str = "&&foo&&bar=baz";
let mut iter = QueryIter { input: query_str };
let first = iter.next().unwrap();
assert_eq!(first.0, "foo");
assert_eq!(first.1.val, "");
let second = iter.next().unwrap();
assert_eq!(second.0, "bar");
assert_eq!(second.1.val, "baz");
let query_str2 = "foo=bar%xy&baz=%";
let mut iter2 = QueryIter { input: query_str2 };
let first2 = iter2.next().unwrap();
assert_eq!(first2.0, "foo");
assert_eq!(first2.1.val, "bar%xy");
let second2 = iter2.next().unwrap();
assert_eq!(second2.0, "baz");
assert_eq!(second2.1.val, "%");
}
#[tokio::test]
async fn test_extractors_direct() {
let req = Request::builder()
.method("POST")
.uri("/path?q=1")
.header("x-test", "hello")
.body(Body::full(Bytes::from("body_bytes")))
.unwrap();
let (mut parts, body) = req.into_parts();
let headers = HeaderMap::from_request_parts(&mut parts, &()).unwrap();
assert_eq!(headers.get("x-test").unwrap(), "hello");
let method = Method::from_request_parts(&mut parts, &()).unwrap();
assert_eq!(method, "POST");
let uri = Uri::from_request_parts(&mut parts, &()).unwrap();
assert_eq!(uri.path(), "/path");
let req_bytes = Request::from_parts(parts.clone(), Body::full(Bytes::from("body_bytes")));
let bytes = Bytes::from_request(req_bytes, &()).await.unwrap();
assert_eq!(bytes.as_ref(), b"body_bytes");
let req_full = Request::from_parts(parts, body);
let extracted_req = <Request<Bytes>>::from_request(req_full, &()).await.unwrap();
assert_eq!(extracted_req.uri().path(), "/path");
}
#[cfg(feature = "cookies")]
#[test]
fn test_cookies_remove() {
use cookie::Cookie;
let cookies = Cookies::new().add(Cookie::new("foo", "bar"));
assert_eq!(cookies.get("foo").unwrap().value(), "bar");
let cookies = cookies.remove(Cookie::new("foo", ""));
assert!(cookies.get("foo").is_none());
}
#[test]
fn test_host_missing() {
let mut parts = Request::builder().uri("/").body(()).unwrap().into_parts().0;
let res = Host::from_request_parts(&mut parts, &());
assert!(res.is_err());
}
#[test]
fn test_connect_info_missing() {
let mut parts = Request::builder().uri("/").body(()).unwrap().into_parts().0;
let res = ConnectInfo::<std::net::SocketAddr>::from_request_parts(&mut parts, &());
assert!(res.is_err());
}
#[cfg(feature = "form")]
#[tokio::test]
async fn test_form_errors() {
#[derive(Deserialize, Debug)]
#[allow(dead_code)]
struct FormPayload {
foo: String,
}
let req = Request::builder()
.method("POST")
.header(hyper::header::CONTENT_TYPE, "text/plain")
.body(Body::full(Bytes::from("foo=bar")))
.unwrap();
let res = Form::<FormPayload>::from_request(req, &()).await;
assert!(res.is_err());
let utf8_req = Request::builder()
.method("POST")
.header(
hyper::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.body(Body::full(Bytes::from(vec![0xff, 0xff])))
.unwrap();
let utf8_result = Form::<FormPayload>::from_request(utf8_req, &()).await;
assert!(utf8_result.is_err());
let payload_req = Request::builder()
.method("POST")
.header(
hyper::header::CONTENT_TYPE,
"application/x-www-form-urlencoded",
)
.body(Body::full(Bytes::from("not_valid")))
.unwrap();
let payload_result = Form::<FormPayload>::from_request(payload_req, &()).await;
assert!(payload_result.is_err());
}
#[cfg(feature = "form")]
#[test]
fn test_form_from_request_parts_deserialize_error() {
#[derive(Deserialize, Debug)]
#[allow(dead_code)]
struct FormPayload {
foo: String,
}
let mut parts = Request::builder()
.uri("/search?bar=baz")
.body(())
.unwrap()
.into_parts()
.0;
let result = Form::<FormPayload>::from_request_parts(&mut parts, &());
assert!(result.is_err());
}
#[cfg(feature = "form")]
#[tokio::test]
async fn test_form_get_request_reads_from_query_string() {
#[derive(Deserialize, Debug, PartialEq)]
struct FormPayload {
foo: String,
}
let req = Request::builder()
.method("GET")
.uri("/search?foo=bar")
.body(Body::empty())
.unwrap();
let Form(payload) = Form::<FormPayload>::from_request(req, &()).await.unwrap();
assert_eq!(
payload,
FormPayload {
foo: "bar".to_string(),
}
);
}
#[cfg(feature = "query")]
#[test]
fn test_query_deserialize_error() {
#[derive(Deserialize, Debug)]
#[allow(dead_code)]
struct QueryPayload {
foo: u32,
}
let mut parts = Request::builder()
.uri("/?foo=not_a_number")
.body(())
.unwrap()
.into_parts()
.0;
let result = Query::<QueryPayload>::from_request_parts(&mut parts, &());
assert!(result.is_err());
}
#[test]
fn test_raw_query_present_and_absent() {
let mut with_query = Request::builder()
.uri("/path?a=1&b=2")
.body(())
.unwrap()
.into_parts()
.0;
let RawQuery(q) = RawQuery::from_request_parts(&mut with_query, &()).unwrap();
assert_eq!(q.as_deref(), Some("a=1&b=2"));
let mut without_query = Request::builder()
.uri("/path")
.body(())
.unwrap()
.into_parts()
.0;
let RawQuery(q2) = RawQuery::from_request_parts(&mut without_query, &()).unwrap();
assert!(q2.is_none());
}
#[cfg(feature = "cookies")]
#[test]
fn test_cookies_default() {
let cookies = Cookies::default();
assert!(cookies.get("anything").is_none());
}
#[tokio::test]
async fn test_body_stream_from_request() {
let req = Request::builder()
.body(Body::full(Bytes::from("stream me")))
.unwrap();
let BodyStream(body) = BodyStream::from_request(req, &()).await.unwrap();
let collected = body.collect_bytes(1024).await.unwrap();
assert_eq!(collected.as_ref(), b"stream me");
}
#[cfg(feature = "json")]
#[test]
fn test_is_json_content_type_without_a_slash_is_rejected() {
assert!(!is_json_content_type("not-a-media-type"));
}
fn make_path_parts(params: Vec<(&str, &str)>) -> hyper::http::request::Parts {
let mut parts = Request::builder().body(()).unwrap().into_parts().0;
let path_params = PathParams(
params
.into_iter()
.map(|(k, v)| (std::sync::Arc::from(k), v.to_string()))
.collect(),
);
parts.extensions.insert(path_params);
parts
}
#[test]
fn test_path_tuple_success_and_length_mismatch() {
let mut ok_parts = make_path_parts(vec![("id", "42"), ("name", "hello")]);
let Path((id, name)) =
Path::<(u32, String)>::from_request_parts(&mut ok_parts, &()).unwrap();
assert_eq!(id, 42);
assert_eq!(name, "hello");
let mut too_many = make_path_parts(vec![("a", "1"), ("b", "2"), ("c", "3")]);
assert!(Path::<(u32, String)>::from_request_parts(&mut too_many, &()).is_err());
let mut too_few = make_path_parts(vec![("a", "1")]);
assert!(Path::<(u32, String)>::from_request_parts(&mut too_few, &()).is_err());
}
#[test]
fn test_path_vec_seq_target() {
let mut parts = make_path_parts(vec![("a", "x"), ("b", "y"), ("c", "z")]);
let Path(values) = Path::<Vec<String>>::from_request_parts(&mut parts, &()).unwrap();
assert_eq!(
values,
vec!["x".to_string(), "y".to_string(), "z".to_string()]
);
}
#[test]
fn test_path_scalar_wrong_param_count() {
let mut zero = make_path_parts(vec![]);
assert!(Path::<u32>::from_request_parts(&mut zero, &()).is_err());
let mut two = make_path_parts(vec![("a", "1"), ("b", "2")]);
assert!(Path::<u32>::from_request_parts(&mut two, &()).is_err());
let mut one = make_path_parts(vec![("id", "7")]);
let Path(v) = Path::<u32>::from_request_parts(&mut one, &()).unwrap();
assert_eq!(v, 7);
}
#[test]
fn test_path_option_top_level_target() {
let mut parts = make_path_parts(vec![("id", "9")]);
let Path(v) = Path::<Option<u32>>::from_request_parts(&mut parts, &()).unwrap();
assert_eq!(v, Some(9));
}
#[test]
fn test_path_enum_target() {
let mut parts = make_path_parts(vec![("color", "Red")]);
let Path(c) = Path::<Color>::from_request_parts(&mut parts, &()).unwrap();
assert_eq!(c, Color::Red);
}
#[test]
fn test_path_unit_and_unit_struct_targets() {
#[derive(Deserialize, PartialEq, Debug)]
struct UnitStruct;
let mut parts = make_path_parts(vec![("a", "1"), ("b", "2")]);
let Path(unit_val) = Path::<()>::from_request_parts(&mut parts, &()).unwrap();
assert_eq!(unit_val, ());
let mut empty_parts = make_path_parts(vec![]);
let Path(u) = Path::<UnitStruct>::from_request_parts(&mut empty_parts, &()).unwrap();
assert_eq!(u, UnitStruct);
}
#[test]
fn test_path_newtype_struct_target() {
#[derive(Deserialize, PartialEq, Debug)]
struct Wrapper(u32);
let mut parts = make_path_parts(vec![("id", "77")]);
let Path(Wrapper(v)) = Path::<Wrapper>::from_request_parts(&mut parts, &()).unwrap();
assert_eq!(v, 77);
}
#[test]
fn test_path_ignored_any_top_level_target() {
let mut parts = make_path_parts(vec![("a", "1"), ("b", "2")]);
let result = Path::<serde::de::IgnoredAny>::from_request_parts(&mut parts, &());
assert!(result.is_ok());
}
struct IdentifierVisitor;
impl serde::de::Visitor<'_> for IdentifierVisitor {
type Value = String;
fn expecting(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("a string identifier")
}
fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
Ok(v.to_string())
}
}
#[test]
fn test_path_deserializer_identifier_direct() {
let params: Vec<(std::sync::Arc<str>, String)> =
vec![(std::sync::Arc::from("k"), "myvalue".to_string())];
let de = PathDeserializer { params: ¶ms };
let result =
serde::de::Deserializer::deserialize_identifier(de, IdentifierVisitor).unwrap();
assert_eq!(result, "myvalue");
}
}