mirror of
https://github.com/tokio-rs/axum.git
synced 2026-09-06 00:00:17 +02:00
Change HeaderMap extractor to clone the headers (#698)
* Change `HeaderMap` extractor to clone the headers * fix docs * changelog * inline variable * also add changelog item to axum * don't list types from axum in axum-core's changelog * document that `HeaderMap::from_request` clones the headers * fix typo * a few more typos
This commit is contained in:
+16
-1
@@ -7,7 +7,22 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
|
|
||||||
# Unreleased
|
# Unreleased
|
||||||
|
|
||||||
- None.
|
- **breaking:** Using `HeaderMap` as an extractor will no longer remove the headers and thus
|
||||||
|
they'll still be accessible to other extractors, such as `axum::extract::Json`. Instead
|
||||||
|
`HeaderMap` will clone the headers. You should prefer to use `TypedHeader` to extract only the
|
||||||
|
headers you need ([#698])
|
||||||
|
|
||||||
|
This includes these breaking changes:
|
||||||
|
- `RequestParts::take_headers` has been removed.
|
||||||
|
- `RequestParts::headers` returns `&HeaderMap`.
|
||||||
|
- `RequestParts::headers_mut` returns `&mut HeaderMap`.
|
||||||
|
- `HeadersAlreadyExtracted` has been removed.
|
||||||
|
- The `HeadersAlreadyExtracted` variant has been removed from these rejections:
|
||||||
|
- `RequestAlreadyExtracted`
|
||||||
|
- `RequestPartsAlreadyExtracted`
|
||||||
|
- `<HeaderMap as FromRequest<_>>::Error` has been changed to `std::convert::Infallible`.
|
||||||
|
|
||||||
|
[#698]: https://github.com/tokio-rs/axum/pull/698
|
||||||
|
|
||||||
# 0.1.1 (06. December, 2021)
|
# 0.1.1 (06. December, 2021)
|
||||||
|
|
||||||
|
|||||||
@@ -77,7 +77,7 @@ pub struct RequestParts<B> {
|
|||||||
method: Method,
|
method: Method,
|
||||||
uri: Uri,
|
uri: Uri,
|
||||||
version: Version,
|
version: Version,
|
||||||
headers: Option<HeaderMap>,
|
headers: HeaderMap,
|
||||||
extensions: Option<Extensions>,
|
extensions: Option<Extensions>,
|
||||||
body: Option<B>,
|
body: Option<B>,
|
||||||
}
|
}
|
||||||
@@ -107,7 +107,7 @@ impl<B> RequestParts<B> {
|
|||||||
method,
|
method,
|
||||||
uri,
|
uri,
|
||||||
version,
|
version,
|
||||||
headers: Some(headers),
|
headers,
|
||||||
extensions: Some(extensions),
|
extensions: Some(extensions),
|
||||||
body: Some(body),
|
body: Some(body),
|
||||||
}
|
}
|
||||||
@@ -117,14 +117,11 @@ impl<B> RequestParts<B> {
|
|||||||
///
|
///
|
||||||
/// Fails if
|
/// Fails if
|
||||||
///
|
///
|
||||||
/// - The full [`HeaderMap`] has been extracted, that is [`take_headers`]
|
|
||||||
/// have been called.
|
|
||||||
/// - The full [`Extensions`] has been extracted, that is
|
/// - The full [`Extensions`] has been extracted, that is
|
||||||
/// [`take_extensions`] have been called.
|
/// [`take_extensions`] have been called.
|
||||||
/// - The request body has been extracted, that is [`take_body`] have been
|
/// - The request body has been extracted, that is [`take_body`] have been
|
||||||
/// called.
|
/// called.
|
||||||
///
|
///
|
||||||
/// [`take_headers`]: RequestParts::take_headers
|
|
||||||
/// [`take_extensions`]: RequestParts::take_extensions
|
/// [`take_extensions`]: RequestParts::take_extensions
|
||||||
/// [`take_body`]: RequestParts::take_body
|
/// [`take_body`]: RequestParts::take_body
|
||||||
pub fn try_into_request(self) -> Result<Request<B>, RequestAlreadyExtracted> {
|
pub fn try_into_request(self) -> Result<Request<B>, RequestAlreadyExtracted> {
|
||||||
@@ -132,7 +129,7 @@ impl<B> RequestParts<B> {
|
|||||||
method,
|
method,
|
||||||
uri,
|
uri,
|
||||||
version,
|
version,
|
||||||
mut headers,
|
headers,
|
||||||
mut extensions,
|
mut extensions,
|
||||||
mut body,
|
mut body,
|
||||||
} = self;
|
} = self;
|
||||||
@@ -148,14 +145,7 @@ impl<B> RequestParts<B> {
|
|||||||
*req.method_mut() = method;
|
*req.method_mut() = method;
|
||||||
*req.uri_mut() = uri;
|
*req.uri_mut() = uri;
|
||||||
*req.version_mut() = version;
|
*req.version_mut() = version;
|
||||||
|
*req.headers_mut() = headers;
|
||||||
if let Some(headers) = headers.take() {
|
|
||||||
*req.headers_mut() = headers;
|
|
||||||
} else {
|
|
||||||
return Err(RequestAlreadyExtracted::HeadersAlreadyExtracted(
|
|
||||||
HeadersAlreadyExtracted,
|
|
||||||
));
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(extensions) = extensions.take() {
|
if let Some(extensions) = extensions.take() {
|
||||||
*req.extensions_mut() = extensions;
|
*req.extensions_mut() = extensions;
|
||||||
@@ -199,22 +189,13 @@ impl<B> RequestParts<B> {
|
|||||||
}
|
}
|
||||||
|
|
||||||
/// Gets a reference to the request headers.
|
/// Gets a reference to the request headers.
|
||||||
///
|
pub fn headers(&self) -> &HeaderMap {
|
||||||
/// Returns `None` if the headers has been taken by another extractor.
|
&self.headers
|
||||||
pub fn headers(&self) -> Option<&HeaderMap> {
|
|
||||||
self.headers.as_ref()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Gets a mutable reference to the request headers.
|
/// Gets a mutable reference to the request headers.
|
||||||
///
|
pub fn headers_mut(&mut self) -> &mut HeaderMap {
|
||||||
/// Returns `None` if the headers has been taken by another extractor.
|
&mut self.headers
|
||||||
pub fn headers_mut(&mut self) -> Option<&mut HeaderMap> {
|
|
||||||
self.headers.as_mut()
|
|
||||||
}
|
|
||||||
|
|
||||||
/// Takes the headers out of the request, leaving a `None` in its place.
|
|
||||||
pub fn take_headers(&mut self) -> Option<HeaderMap> {
|
|
||||||
self.headers.take()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
/// Gets a reference to the request extensions.
|
/// Gets a reference to the request extensions.
|
||||||
|
|||||||
@@ -8,13 +8,6 @@ define_rejection! {
|
|||||||
pub struct BodyAlreadyExtracted;
|
pub struct BodyAlreadyExtracted;
|
||||||
}
|
}
|
||||||
|
|
||||||
define_rejection! {
|
|
||||||
#[status = INTERNAL_SERVER_ERROR]
|
|
||||||
#[body = "Headers taken by other extractor"]
|
|
||||||
/// Rejection used if the headers has been taken by another extractor.
|
|
||||||
pub struct HeadersAlreadyExtracted;
|
|
||||||
}
|
|
||||||
|
|
||||||
define_rejection! {
|
define_rejection! {
|
||||||
#[status = INTERNAL_SERVER_ERROR]
|
#[status = INTERNAL_SERVER_ERROR]
|
||||||
#[body = "Extensions taken by other extractor"]
|
#[body = "Extensions taken by other extractor"]
|
||||||
@@ -47,7 +40,6 @@ composite_rejection! {
|
|||||||
/// [`Request<_>`]: http::Request
|
/// [`Request<_>`]: http::Request
|
||||||
pub enum RequestAlreadyExtracted {
|
pub enum RequestAlreadyExtracted {
|
||||||
BodyAlreadyExtracted,
|
BodyAlreadyExtracted,
|
||||||
HeadersAlreadyExtracted,
|
|
||||||
ExtensionsAlreadyExtracted,
|
ExtensionsAlreadyExtracted,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -79,7 +71,6 @@ composite_rejection! {
|
|||||||
///
|
///
|
||||||
/// Contains one variant for each way the [`http::request::Parts`] extractor can fail.
|
/// Contains one variant for each way the [`http::request::Parts`] extractor can fail.
|
||||||
pub enum RequestPartsAlreadyExtracted {
|
pub enum RequestPartsAlreadyExtracted {
|
||||||
HeadersAlreadyExtracted,
|
|
||||||
ExtensionsAlreadyExtracted,
|
ExtensionsAlreadyExtracted,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ where
|
|||||||
method: req.method.clone(),
|
method: req.method.clone(),
|
||||||
version: req.version,
|
version: req.version,
|
||||||
uri: req.uri.clone(),
|
uri: req.uri.clone(),
|
||||||
headers: None,
|
headers: HeaderMap::new(),
|
||||||
extensions: None,
|
extensions: None,
|
||||||
body: None,
|
body: None,
|
||||||
},
|
},
|
||||||
@@ -65,15 +65,20 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/// Clone the headers from the request.
|
||||||
|
///
|
||||||
|
/// Prefer using [`TypedHeader`] to extract only the headers you need.
|
||||||
|
///
|
||||||
|
/// [`TypedHeader`]: https://docs.rs/axum/latest/axum/extract/struct.TypedHeader.html
|
||||||
#[async_trait]
|
#[async_trait]
|
||||||
impl<B> FromRequest<B> for HeaderMap
|
impl<B> FromRequest<B> for HeaderMap
|
||||||
where
|
where
|
||||||
B: Send,
|
B: Send,
|
||||||
{
|
{
|
||||||
type Rejection = HeadersAlreadyExtracted;
|
type Rejection = Infallible;
|
||||||
|
|
||||||
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
||||||
req.take_headers().ok_or(HeadersAlreadyExtracted)
|
Ok(req.headers().clone())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -143,7 +148,10 @@ where
|
|||||||
let method = unwrap_infallible(Method::from_request(req).await);
|
let method = unwrap_infallible(Method::from_request(req).await);
|
||||||
let uri = unwrap_infallible(Uri::from_request(req).await);
|
let uri = unwrap_infallible(Uri::from_request(req).await);
|
||||||
let version = unwrap_infallible(Version::from_request(req).await);
|
let version = unwrap_infallible(Version::from_request(req).await);
|
||||||
let headers = HeaderMap::from_request(req).await?;
|
let headers = match HeaderMap::from_request(req).await {
|
||||||
|
Ok(headers) => headers,
|
||||||
|
Err(err) => match err {},
|
||||||
|
};
|
||||||
let extensions = Extensions::from_request(req).await?;
|
let extensions = Extensions::from_request(req).await?;
|
||||||
|
|
||||||
let mut temp_request = Request::new(());
|
let mut temp_request = Request::new(());
|
||||||
|
|||||||
@@ -13,9 +13,28 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
|
|||||||
overwriting old values.
|
overwriting old values.
|
||||||
- **breaking:** Require `Output = ()` on `WebSocketStream::on_upgrade` ([#644])
|
- **breaking:** Require `Output = ()` on `WebSocketStream::on_upgrade` ([#644])
|
||||||
- **breaking:** Make `TypedHeaderRejectionReason` `#[non_exhaustive]` ([#665])
|
- **breaking:** Make `TypedHeaderRejectionReason` `#[non_exhaustive]` ([#665])
|
||||||
|
- **breaking:** Using `HeaderMap` as an extractor will no longer remove the headers and thus
|
||||||
|
they'll still be accessible to other extractors, such as `axum::extract::Json`. Instead
|
||||||
|
`HeaderMap` will clone the headers. You should prefer to use `TypedHeader` to extract only the
|
||||||
|
headers you need ([#698])
|
||||||
|
|
||||||
|
This includes these breaking changes:
|
||||||
|
- `RequestParts::take_headers` has been removed.
|
||||||
|
- `RequestParts::headers` returns `&HeaderMap`.
|
||||||
|
- `RequestParts::headers_mut` returns `&mut HeaderMap`.
|
||||||
|
- `HeadersAlreadyExtracted` has been removed.
|
||||||
|
- The `HeadersAlreadyExtracted` removed variant has been removed from these rejections:
|
||||||
|
- `RequestAlreadyExtracted`
|
||||||
|
- `RequestPartsAlreadyExtracted`
|
||||||
|
- `JsonRejection`
|
||||||
|
- `FormRejection`
|
||||||
|
- `ContentLengthLimitRejection`
|
||||||
|
- `WebSocketUpgradeRejection`
|
||||||
|
- `<HeaderMap as FromRequest<_>>::Error` has been changed to `std::convert::Infallible`.
|
||||||
|
|
||||||
[#644]: https://github.com/tokio-rs/axum/pull/644
|
[#644]: https://github.com/tokio-rs/axum/pull/644
|
||||||
[#665]: https://github.com/tokio-rs/axum/pull/665
|
[#665]: https://github.com/tokio-rs/axum/pull/665
|
||||||
|
[#698]: https://github.com/tokio-rs/axum/pull/698
|
||||||
|
|
||||||
# 0.4.4 (13. January, 2021)
|
# 0.4.4 (13. January, 2021)
|
||||||
|
|
||||||
|
|||||||
@@ -320,10 +320,6 @@ async fn handler(result: Result<Json<Value>, JsonRejection>) -> impl IntoRespons
|
|||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
StatusCode::INTERNAL_SERVER_ERROR,
|
||||||
"Failed to buffer request body".to_string(),
|
"Failed to buffer request body".to_string(),
|
||||||
)),
|
)),
|
||||||
JsonRejection::HeadersAlreadyExtracted(_) => Err((
|
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
|
||||||
"Headers already extracted".to_string(),
|
|
||||||
)),
|
|
||||||
// we must provide a catch-all case since `JsonRejection` is marked
|
// we must provide a catch-all case since `JsonRejection` is marked
|
||||||
// `#[non_exhaustive]`
|
// `#[non_exhaustive]`
|
||||||
_ => Err((
|
_ => Err((
|
||||||
@@ -377,9 +373,7 @@ where
|
|||||||
type Rejection = (StatusCode, &'static str);
|
type Rejection = (StatusCode, &'static str);
|
||||||
|
|
||||||
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
||||||
let user_agent = req.headers().and_then(|headers| headers.get(USER_AGENT));
|
if let Some(user_agent) = req.headers().get(USER_AGENT) {
|
||||||
|
|
||||||
if let Some(user_agent) = user_agent {
|
|
||||||
Ok(ExtractUserAgent(user_agent.clone()))
|
Ok(ExtractUserAgent(user_agent.clone()))
|
||||||
} else {
|
} else {
|
||||||
Err((StatusCode::BAD_REQUEST, "`User-Agent` header is missing"))
|
Err((StatusCode::BAD_REQUEST, "`User-Agent` header is missing"))
|
||||||
|
|||||||
@@ -39,14 +39,7 @@ where
|
|||||||
type Rejection = ContentLengthLimitRejection<T::Rejection>;
|
type Rejection = ContentLengthLimitRejection<T::Rejection>;
|
||||||
|
|
||||||
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
||||||
let content_length = req
|
let content_length = req.headers().get(http::header::CONTENT_LENGTH);
|
||||||
.headers()
|
|
||||||
.ok_or_else(|| {
|
|
||||||
ContentLengthLimitRejection::HeadersAlreadyExtracted(
|
|
||||||
HeadersAlreadyExtracted::default(),
|
|
||||||
)
|
|
||||||
})?
|
|
||||||
.get(http::header::CONTENT_LENGTH);
|
|
||||||
|
|
||||||
let content_length =
|
let content_length =
|
||||||
content_length.and_then(|value| value.to_str().ok()?.parse::<u64>().ok());
|
content_length.and_then(|value| value.to_str().ok()?.parse::<u64>().ok());
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ use tower_service::Service;
|
|||||||
/// async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
/// async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
||||||
/// let auth_header = req
|
/// let auth_header = req
|
||||||
/// .headers()
|
/// .headers()
|
||||||
/// .and_then(|headers| headers.get(http::header::AUTHORIZATION))
|
/// .get(http::header::AUTHORIZATION)
|
||||||
/// .and_then(|value| value.to_str().ok());
|
/// .and_then(|value| value.to_str().ok());
|
||||||
///
|
///
|
||||||
/// match auth_header {
|
/// match auth_header {
|
||||||
@@ -291,7 +291,6 @@ mod tests {
|
|||||||
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
||||||
if let Some(auth) = req
|
if let Some(auth) = req
|
||||||
.headers()
|
.headers()
|
||||||
.expect("headers already extracted")
|
|
||||||
.get("authorization")
|
.get("authorization")
|
||||||
.and_then(|v| v.to_str().ok())
|
.and_then(|v| v.to_str().ok())
|
||||||
{
|
{
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ where
|
|||||||
.map_err(FailedToDeserializeQueryString::new::<T, _>)?;
|
.map_err(FailedToDeserializeQueryString::new::<T, _>)?;
|
||||||
Ok(Form(value))
|
Ok(Form(value))
|
||||||
} else {
|
} else {
|
||||||
if !has_content_type(req, &mime::APPLICATION_WWW_FORM_URLENCODED)? {
|
if !has_content_type(req, &mime::APPLICATION_WWW_FORM_URLENCODED) {
|
||||||
return Err(InvalidFormContentType.into());
|
return Err(InvalidFormContentType.into());
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -78,24 +78,20 @@ pub use self::typed_header::TypedHeader;
|
|||||||
pub(crate) fn has_content_type<B>(
|
pub(crate) fn has_content_type<B>(
|
||||||
req: &RequestParts<B>,
|
req: &RequestParts<B>,
|
||||||
expected_content_type: &mime::Mime,
|
expected_content_type: &mime::Mime,
|
||||||
) -> Result<bool, HeadersAlreadyExtracted> {
|
) -> bool {
|
||||||
let content_type = if let Some(content_type) = req
|
let content_type = if let Some(content_type) = req.headers().get(header::CONTENT_TYPE) {
|
||||||
.headers()
|
|
||||||
.ok_or_else(HeadersAlreadyExtracted::default)?
|
|
||||||
.get(header::CONTENT_TYPE)
|
|
||||||
{
|
|
||||||
content_type
|
content_type
|
||||||
} else {
|
} else {
|
||||||
return Ok(false);
|
return false;
|
||||||
};
|
};
|
||||||
|
|
||||||
let content_type = if let Ok(content_type) = content_type.to_str() {
|
let content_type = if let Ok(content_type) = content_type.to_str() {
|
||||||
content_type
|
content_type
|
||||||
} else {
|
} else {
|
||||||
return Ok(false);
|
return false;
|
||||||
};
|
};
|
||||||
|
|
||||||
Ok(content_type.starts_with(expected_content_type.as_ref()))
|
content_type.starts_with(expected_content_type.as_ref())
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(crate) fn take_body<B>(req: &mut RequestParts<B>) -> Result<B, BodyAlreadyExtracted> {
|
pub(crate) fn take_body<B>(req: &mut RequestParts<B>) -> Result<B, BodyAlreadyExtracted> {
|
||||||
|
|||||||
@@ -58,7 +58,7 @@ where
|
|||||||
|
|
||||||
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
||||||
let stream = BodyStream::from_request(req).await?;
|
let stream = BodyStream::from_request(req).await?;
|
||||||
let headers = req.headers().ok_or_else(HeadersAlreadyExtracted::default)?;
|
let headers = req.headers();
|
||||||
let boundary = parse_boundary(headers).ok_or(InvalidBoundary)?;
|
let boundary = parse_boundary(headers).ok_or(InvalidBoundary)?;
|
||||||
let multipart = multer::Multipart::new(stream, boundary);
|
let multipart = multer::Multipart::new(stream, boundary);
|
||||||
Ok(Self { inner: multipart })
|
Ok(Self { inner: multipart })
|
||||||
@@ -179,7 +179,6 @@ composite_rejection! {
|
|||||||
pub enum MultipartRejection {
|
pub enum MultipartRejection {
|
||||||
BodyAlreadyExtracted,
|
BodyAlreadyExtracted,
|
||||||
InvalidBoundary,
|
InvalidBoundary,
|
||||||
HeadersAlreadyExtracted,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -124,7 +124,6 @@ composite_rejection! {
|
|||||||
InvalidFormContentType,
|
InvalidFormContentType,
|
||||||
FailedToDeserializeQueryString,
|
FailedToDeserializeQueryString,
|
||||||
BytesRejection,
|
BytesRejection,
|
||||||
HeadersAlreadyExtracted,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -139,7 +138,6 @@ composite_rejection! {
|
|||||||
InvalidJsonBody,
|
InvalidJsonBody,
|
||||||
MissingJsonContentType,
|
MissingJsonContentType,
|
||||||
BytesRejection,
|
BytesRejection,
|
||||||
HeadersAlreadyExtracted,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -195,8 +193,6 @@ pub enum ContentLengthLimitRejection<T> {
|
|||||||
#[allow(missing_docs)]
|
#[allow(missing_docs)]
|
||||||
LengthRequired(LengthRequired),
|
LengthRequired(LengthRequired),
|
||||||
#[allow(missing_docs)]
|
#[allow(missing_docs)]
|
||||||
HeadersAlreadyExtracted(HeadersAlreadyExtracted),
|
|
||||||
#[allow(missing_docs)]
|
|
||||||
Inner(T),
|
Inner(T),
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -208,7 +204,6 @@ where
|
|||||||
match self {
|
match self {
|
||||||
Self::PayloadTooLarge(inner) => inner.into_response(),
|
Self::PayloadTooLarge(inner) => inner.into_response(),
|
||||||
Self::LengthRequired(inner) => inner.into_response(),
|
Self::LengthRequired(inner) => inner.into_response(),
|
||||||
Self::HeadersAlreadyExtracted(inner) => inner.into_response(),
|
|
||||||
Self::Inner(inner) => inner.into_response(),
|
Self::Inner(inner) => inner.into_response(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -222,7 +217,6 @@ where
|
|||||||
match self {
|
match self {
|
||||||
Self::PayloadTooLarge(inner) => inner.fmt(f),
|
Self::PayloadTooLarge(inner) => inner.fmt(f),
|
||||||
Self::LengthRequired(inner) => inner.fmt(f),
|
Self::LengthRequired(inner) => inner.fmt(f),
|
||||||
Self::HeadersAlreadyExtracted(inner) => inner.fmt(f),
|
|
||||||
Self::Inner(inner) => inner.fmt(f),
|
Self::Inner(inner) => inner.fmt(f),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -236,7 +230,6 @@ where
|
|||||||
match self {
|
match self {
|
||||||
Self::PayloadTooLarge(inner) => Some(inner),
|
Self::PayloadTooLarge(inner) => Some(inner),
|
||||||
Self::LengthRequired(inner) => Some(inner),
|
Self::LengthRequired(inner) => Some(inner),
|
||||||
Self::HeadersAlreadyExtracted(inner) => Some(inner),
|
|
||||||
Self::Inner(inner) => Some(inner),
|
Self::Inner(inner) => Some(inner),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -44,16 +44,7 @@ where
|
|||||||
type Rejection = TypedHeaderRejection;
|
type Rejection = TypedHeaderRejection;
|
||||||
|
|
||||||
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
||||||
let headers = if let Some(headers) = req.headers() {
|
match req.headers().typed_try_get::<T>() {
|
||||||
headers
|
|
||||||
} else {
|
|
||||||
return Err(TypedHeaderRejection {
|
|
||||||
name: T::name(),
|
|
||||||
reason: TypedHeaderRejectionReason::Missing,
|
|
||||||
});
|
|
||||||
};
|
|
||||||
|
|
||||||
match headers.typed_try_get::<T>() {
|
|
||||||
Ok(Some(value)) => Ok(Self(value)),
|
Ok(Some(value)) => Ok(Self(value)),
|
||||||
Ok(None) => Err(TypedHeaderRejection {
|
Ok(None) => Err(TypedHeaderRejection {
|
||||||
name: T::name(),
|
name: T::name(),
|
||||||
|
|||||||
+19
-43
@@ -249,27 +249,24 @@ where
|
|||||||
return Err(MethodNotGet.into());
|
return Err(MethodNotGet.into());
|
||||||
}
|
}
|
||||||
|
|
||||||
if !header_contains(req, header::CONNECTION, "upgrade")? {
|
if !header_contains(req, header::CONNECTION, "upgrade") {
|
||||||
return Err(InvalidConnectionHeader.into());
|
return Err(InvalidConnectionHeader.into());
|
||||||
}
|
}
|
||||||
|
|
||||||
if !header_eq(req, header::UPGRADE, "websocket")? {
|
if !header_eq(req, header::UPGRADE, "websocket") {
|
||||||
return Err(InvalidUpgradeHeader.into());
|
return Err(InvalidUpgradeHeader.into());
|
||||||
}
|
}
|
||||||
|
|
||||||
if !header_eq(req, header::SEC_WEBSOCKET_VERSION, "13")? {
|
if !header_eq(req, header::SEC_WEBSOCKET_VERSION, "13") {
|
||||||
return Err(InvalidWebSocketVersionHeader.into());
|
return Err(InvalidWebSocketVersionHeader.into());
|
||||||
}
|
}
|
||||||
|
|
||||||
let sec_websocket_key = if let Some(key) = req
|
let sec_websocket_key =
|
||||||
.headers_mut()
|
if let Some(key) = req.headers_mut().remove(header::SEC_WEBSOCKET_KEY) {
|
||||||
.ok_or_else(HeadersAlreadyExtracted::default)?
|
key
|
||||||
.remove(header::SEC_WEBSOCKET_KEY)
|
} else {
|
||||||
{
|
return Err(WebSocketKeyHeaderMissing.into());
|
||||||
key
|
};
|
||||||
} else {
|
|
||||||
return Err(WebSocketKeyHeaderMissing.into());
|
|
||||||
};
|
|
||||||
|
|
||||||
let on_upgrade = req
|
let on_upgrade = req
|
||||||
.extensions_mut()
|
.extensions_mut()
|
||||||
@@ -277,11 +274,7 @@ where
|
|||||||
.remove::<OnUpgrade>()
|
.remove::<OnUpgrade>()
|
||||||
.unwrap();
|
.unwrap();
|
||||||
|
|
||||||
let sec_websocket_protocol = req
|
let sec_websocket_protocol = req.headers().get(header::SEC_WEBSOCKET_PROTOCOL).cloned();
|
||||||
.headers()
|
|
||||||
.ok_or_else(HeadersAlreadyExtracted::default)?
|
|
||||||
.get(header::SEC_WEBSOCKET_PROTOCOL)
|
|
||||||
.cloned();
|
|
||||||
|
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
config: Default::default(),
|
config: Default::default(),
|
||||||
@@ -293,41 +286,25 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn header_eq<B>(
|
fn header_eq<B>(req: &RequestParts<B>, key: HeaderName, value: &'static str) -> bool {
|
||||||
req: &RequestParts<B>,
|
if let Some(header) = req.headers().get(&key) {
|
||||||
key: HeaderName,
|
header.as_bytes().eq_ignore_ascii_case(value.as_bytes())
|
||||||
value: &'static str,
|
|
||||||
) -> Result<bool, HeadersAlreadyExtracted> {
|
|
||||||
if let Some(header) = req
|
|
||||||
.headers()
|
|
||||||
.ok_or_else(HeadersAlreadyExtracted::default)?
|
|
||||||
.get(&key)
|
|
||||||
{
|
|
||||||
Ok(header.as_bytes().eq_ignore_ascii_case(value.as_bytes()))
|
|
||||||
} else {
|
} else {
|
||||||
Ok(false)
|
false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn header_contains<B>(
|
fn header_contains<B>(req: &RequestParts<B>, key: HeaderName, value: &'static str) -> bool {
|
||||||
req: &RequestParts<B>,
|
let header = if let Some(header) = req.headers().get(&key) {
|
||||||
key: HeaderName,
|
|
||||||
value: &'static str,
|
|
||||||
) -> Result<bool, HeadersAlreadyExtracted> {
|
|
||||||
let header = if let Some(header) = req
|
|
||||||
.headers()
|
|
||||||
.ok_or_else(HeadersAlreadyExtracted::default)?
|
|
||||||
.get(&key)
|
|
||||||
{
|
|
||||||
header
|
header
|
||||||
} else {
|
} else {
|
||||||
return Ok(false);
|
return false;
|
||||||
};
|
};
|
||||||
|
|
||||||
if let Ok(header) = std::str::from_utf8(header.as_bytes()) {
|
if let Ok(header) = std::str::from_utf8(header.as_bytes()) {
|
||||||
Ok(header.to_ascii_lowercase().contains(value))
|
header.to_ascii_lowercase().contains(value)
|
||||||
} else {
|
} else {
|
||||||
Ok(false)
|
false
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -585,7 +562,6 @@ pub mod rejection {
|
|||||||
InvalidUpgradeHeader,
|
InvalidUpgradeHeader,
|
||||||
InvalidWebSocketVersionHeader,
|
InvalidWebSocketVersionHeader,
|
||||||
WebSocketKeyHeaderMissing,
|
WebSocketKeyHeaderMissing,
|
||||||
HeadersAlreadyExtracted,
|
|
||||||
ExtensionsAlreadyExtracted,
|
ExtensionsAlreadyExtracted,
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+7
-11
@@ -96,7 +96,7 @@ where
|
|||||||
type Rejection = JsonRejection;
|
type Rejection = JsonRejection;
|
||||||
|
|
||||||
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
|
||||||
if json_content_type(req)? {
|
if json_content_type(req) {
|
||||||
let bytes = Bytes::from_request(req).await?;
|
let bytes = Bytes::from_request(req).await?;
|
||||||
|
|
||||||
let value = serde_json::from_slice(&bytes).map_err(InvalidJsonBody::from_err)?;
|
let value = serde_json::from_slice(&bytes).map_err(InvalidJsonBody::from_err)?;
|
||||||
@@ -108,33 +108,29 @@ where
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
fn json_content_type<B>(req: &RequestParts<B>) -> Result<bool, HeadersAlreadyExtracted> {
|
fn json_content_type<B>(req: &RequestParts<B>) -> bool {
|
||||||
let content_type = if let Some(content_type) = req
|
let content_type = if let Some(content_type) = req.headers().get(header::CONTENT_TYPE) {
|
||||||
.headers()
|
|
||||||
.ok_or_else(HeadersAlreadyExtracted::default)?
|
|
||||||
.get(header::CONTENT_TYPE)
|
|
||||||
{
|
|
||||||
content_type
|
content_type
|
||||||
} else {
|
} else {
|
||||||
return Ok(false);
|
return false;
|
||||||
};
|
};
|
||||||
|
|
||||||
let content_type = if let Ok(content_type) = content_type.to_str() {
|
let content_type = if let Ok(content_type) = content_type.to_str() {
|
||||||
content_type
|
content_type
|
||||||
} else {
|
} else {
|
||||||
return Ok(false);
|
return false;
|
||||||
};
|
};
|
||||||
|
|
||||||
let mime = if let Ok(mime) = content_type.parse::<mime::Mime>() {
|
let mime = if let Ok(mime) = content_type.parse::<mime::Mime>() {
|
||||||
mime
|
mime
|
||||||
} else {
|
} else {
|
||||||
return Ok(false);
|
return false;
|
||||||
};
|
};
|
||||||
|
|
||||||
let is_json_content_type = mime.type_() == "application"
|
let is_json_content_type = mime.type_() == "application"
|
||||||
&& (mime.subtype() == "json" || mime.suffix().map_or(false, |name| name == "json"));
|
&& (mime.subtype() == "json" || mime.suffix().map_or(false, |name| name == "json"));
|
||||||
|
|
||||||
Ok(is_json_content_type)
|
is_json_content_type
|
||||||
}
|
}
|
||||||
|
|
||||||
impl<T> Deref for Json<T> {
|
impl<T> Deref for Json<T> {
|
||||||
|
|||||||
@@ -73,9 +73,6 @@ where
|
|||||||
JsonRejection::MissingJsonContentType(err) => {
|
JsonRejection::MissingJsonContentType(err) => {
|
||||||
(StatusCode::BAD_REQUEST, err.to_string().into())
|
(StatusCode::BAD_REQUEST, err.to_string().into())
|
||||||
}
|
}
|
||||||
JsonRejection::HeadersAlreadyExtracted(err) => {
|
|
||||||
(StatusCode::INTERNAL_SERVER_ERROR, err.to_string().into())
|
|
||||||
}
|
|
||||||
err => (
|
err => (
|
||||||
StatusCode::INTERNAL_SERVER_ERROR,
|
StatusCode::INTERNAL_SERVER_ERROR,
|
||||||
format!("Unknown internal error: {}", err).into(),
|
format!("Unknown internal error: {}", err).into(),
|
||||||
|
|||||||
Reference in New Issue
Block a user