//! Types and traits for extracting data from requests. //! //! A handler function must always take `Request` as its first argument //! but any arguments following are called "extractors". Any type that //! implements [`FromRequest`](FromRequest) can be used as an extractor. //! //! For example, [`Json`] is an extractor that consumes the request body and //! deserializes it as JSON into some target type: //! //! ```rust,no_run //! use tower_web::prelude::*; //! use serde::Deserialize; //! //! #[derive(Deserialize)] //! struct CreateUser { //! email: String, //! password: String, //! } //! //! async fn create_user(req: Request, payload: extract::Json) { //! let payload: CreateUser = payload.0; //! //! // ... //! } //! //! let app = route("/users", post(create_user)); //! # async { //! # hyper::Server::bind(&"".parse().unwrap()).serve(tower::make::Shared::new(app)).await; //! # }; //! ``` //! //! Technically extractors can also be used as "guards", for example to require //! that requests are authorized. However the recommended way to do that is //! using Tower middleware, such as [`tower_http::auth::RequireAuthorization`]. //! Extractors have to be applied to each handler, whereas middleware can be //! applied to a whole stack at once, which is typically what you want for //! authorization. //! //! # Defining custom extractors //! //! You can also define your own extractors by implementing [`FromRequest`]: //! //! ```rust,no_run //! use tower_web::{async_trait, extract::FromRequest, prelude::*}; //! use http::{StatusCode, header::{HeaderValue, USER_AGENT}}; //! //! struct ExtractUserAgent(HeaderValue); //! //! #[async_trait] //! impl FromRequest for ExtractUserAgent { //! type Rejection = (StatusCode, &'static str); //! //! async fn from_request(req: &mut Request) -> Result { //! if let Some(user_agent) = req.headers().get(USER_AGENT) { //! Ok(ExtractUserAgent(user_agent.clone())) //! } else { //! Err((StatusCode::BAD_REQUEST, "`User-Agent` header is missing")) //! } //! } //! } //! //! async fn handler(req: Request, user_agent: ExtractUserAgent) { //! let user_agent: HeaderValue = user_agent.0; //! //! // ... //! } //! //! let app = route("/foo", get(handler)); //! # async { //! # hyper::Server::bind(&"".parse().unwrap()).serve(tower::make::Shared::new(app)).await; //! # }; //! ``` //! //! # Multiple extractors //! //! Handlers can also contain multiple extractors: //! //! ```rust,no_run //! use tower_web::prelude::*; //! use std::collections::HashMap; //! //! async fn handler( //! req: Request, //! // Extract captured parameters from the URL //! params: extract::UrlParamsMap, //! // Parse query string into a `HashMap` //! query_params: extract::Query>, //! // Buffer the request body into a `Bytes` //! bytes: bytes::Bytes, //! ) { //! // ... //! } //! //! let app = route("/foo", get(handler)); //! # async { //! # hyper::Server::bind(&"".parse().unwrap()).serve(tower::make::Shared::new(app)).await; //! # }; //! ``` //! //! # Optional extractors //! //! Wrapping extractors in `Option` will make them optional: //! //! ```rust,no_run //! use tower_web::{extract::Json, prelude::*}; //! use serde_json::Value; //! //! async fn create_user(req: Request, payload: Option>) { //! if let Some(payload) = payload { //! // We got a valid JSON payload //! } else { //! // Payload wasn't valid JSON //! } //! } //! //! let app = route("/users", post(create_user)); //! # async { //! # hyper::Server::bind(&"".parse().unwrap()).serve(tower::make::Shared::new(app)).await; //! # }; //! ``` //! //! # Reducing boilerplate //! //! If you're feeling adventorous you can even deconstruct the extractors //! directly on the function signature: //! //! ```rust,no_run //! use tower_web::{extract::Json, prelude::*}; //! use serde_json::Value; //! //! async fn create_user(req: Request, Json(value): Json) { //! // `value` is of type `Value` //! } //! //! let app = route("/users", post(create_user)); //! # async { //! # hyper::Server::bind(&"".parse().unwrap()).serve(tower::make::Shared::new(app)).await; //! # }; //! ``` use crate::{body::Body, response::IntoResponse}; use async_trait::async_trait; use bytes::Bytes; use http::{header, Request, Response}; use rejection::{ BodyAlreadyTaken, FailedToBufferBody, InvalidJsonBody, InvalidUrlParam, InvalidUtf8, LengthRequired, MissingExtension, MissingJsonContentType, MissingRouteParams, PayloadTooLarge, QueryStringMissing, UrlParamsAlreadyTaken, }; use serde::de::DeserializeOwned; use std::{collections::HashMap, convert::Infallible, str::FromStr}; pub mod rejection; /// Types that can be created from requests. /// /// See the [module docs](crate::extract) for more details. #[async_trait] pub trait FromRequest: Sized { /// If the extractor fails it'll use this "rejection" type. A rejection is /// a kind of error that can be converted into a response. type Rejection: IntoResponse; /// Perform the extraction. async fn from_request(req: &mut Request) -> Result; } #[async_trait] impl FromRequest for Option where T: FromRequest, { type Rejection = Infallible; async fn from_request(req: &mut Request) -> Result, Self::Rejection> { Ok(T::from_request(req).await.ok()) } } /// Extractor that deserializes query strings into some type. /// /// `T` is expected to implement [`serde::Deserialize`]. /// /// # Example /// /// ```rust,no_run /// use tower_web::prelude::*; /// use serde::Deserialize; /// /// #[derive(Deserialize)] /// struct Pagination { /// page: usize, /// per_page: usize, /// } /// /// // This will parse query strings like `?page=2&per_page=30` into `Pagination` /// // structs. /// async fn list_things(req: Request, pagination: extract::Query) { /// let pagination: Pagination = pagination.0; /// /// // ... /// } /// let app = route("/list_things", get(list_things)); /// ``` /// /// If the query string cannot be parsed it will reject the request with a `404 /// Bad Request` response. #[derive(Debug, Clone, Copy, Default)] pub struct Query(pub T); #[async_trait] impl FromRequest for Query where T: DeserializeOwned, { type Rejection = QueryStringMissing; async fn from_request(req: &mut Request) -> Result { let query = req.uri().query().ok_or(QueryStringMissing)?; let value = serde_urlencoded::from_str(query).map_err(|_| QueryStringMissing)?; Ok(Query(value)) } } /// Extractor that deserializes request bodies into some type. /// /// `T` is expected to implement [`serde::Deserialize`]. /// /// # Example /// /// ```rust,no_run /// use tower_web::prelude::*; /// use serde::Deserialize; /// /// #[derive(Deserialize)] /// struct CreateUser { /// email: String, /// password: String, /// } /// /// async fn create_user(req: Request, payload: extract::Json) { /// let payload: CreateUser = payload.0; /// /// // ... /// } /// /// let app = route("/users", post(create_user)); /// ``` /// /// If the query string cannot be parsed it will reject the request with a `404 /// Bad Request` response. /// /// The request is required to have a `Content-Type: application/json` header. #[derive(Debug, Clone, Copy, Default)] pub struct Json(pub T); #[async_trait] impl FromRequest for Json where T: DeserializeOwned, { type Rejection = Response; async fn from_request(req: &mut Request) -> Result { use bytes::Buf; if has_content_type(req, "application/json") { let body = take_body(req).map_err(IntoResponse::into_response)?; let buf = hyper::body::aggregate(body) .await .map_err(InvalidJsonBody::from_err) .map_err(IntoResponse::into_response)?; let value = serde_json::from_reader(buf.reader()) .map_err(InvalidJsonBody::from_err) .map_err(IntoResponse::into_response)?; Ok(Json(value)) } else { Err(MissingJsonContentType.into_response()) } } } fn has_content_type(req: &Request, expected_content_type: &str) -> bool { let content_type = if let Some(content_type) = req.headers().get(header::CONTENT_TYPE) { content_type } else { return false; }; let content_type = if let Ok(content_type) = content_type.to_str() { content_type } else { return false; }; content_type.starts_with(expected_content_type) } /// Extractor that gets a value from request extensions. /// /// This is commonly used to share state across handlers. /// /// # Example /// /// ```rust,no_run /// use tower_web::{AddExtensionLayer, prelude::*}; /// use std::sync::Arc; /// /// // Some shared state used throughout our application /// struct State { /// // ... /// } /// /// async fn handler(req: Request, state: extract::Extension>) { /// // ... /// } /// /// let state = Arc::new(State { /* ... */ }); /// /// let app = route("/", get(handler)) /// // Add middleware that inserts the state into all incoming request's /// // extensions. /// .layer(AddExtensionLayer::new(state)); /// ``` /// /// If the extension is missing it will reject the request with a `500 Interal /// Server Error` response. #[derive(Debug, Clone, Copy)] pub struct Extension(pub T); #[async_trait] impl FromRequest for Extension where T: Clone + Send + Sync + 'static, { type Rejection = MissingExtension; async fn from_request(req: &mut Request) -> Result { let value = req .extensions() .get::() .ok_or(MissingExtension) .map(|x| x.clone())?; Ok(Extension(value)) } } #[async_trait] impl FromRequest for Bytes { type Rejection = Response; async fn from_request(req: &mut Request) -> Result { let body = take_body(req).map_err(IntoResponse::into_response)?; let bytes = hyper::body::to_bytes(body) .await .map_err(FailedToBufferBody::from_err) .map_err(IntoResponse::into_response)?; Ok(bytes) } } #[async_trait] impl FromRequest for String { type Rejection = Response; async fn from_request(req: &mut Request) -> Result { let body = take_body(req).map_err(IntoResponse::into_response)?; let bytes = hyper::body::to_bytes(body) .await .map_err(FailedToBufferBody::from_err) .map_err(IntoResponse::into_response)? .to_vec(); let string = String::from_utf8(bytes) .map_err(InvalidUtf8::from_err) .map_err(IntoResponse::into_response)?; Ok(string) } } #[async_trait] impl FromRequest for Body { type Rejection = BodyAlreadyTaken; async fn from_request(req: &mut Request) -> Result { take_body(req) } } /// Extractor that will buffer request bodies up to a certain size. /// /// # Example /// /// ```rust,no_run /// use tower_web::prelude::*; /// /// async fn handler(req: Request, body: extract::BytesMaxLength<1024>) { /// // ... /// } /// /// let app = route("/", post(handler)); /// ``` /// /// This requires the request to have a `Content-Length` header. #[derive(Debug, Clone)] pub struct BytesMaxLength(pub Bytes); #[async_trait] impl FromRequest for BytesMaxLength { type Rejection = Response; async fn from_request(req: &mut Request) -> Result { let content_length = req.headers().get(http::header::CONTENT_LENGTH).cloned(); let body = take_body(req).map_err(|reject| reject.into_response())?; let content_length = content_length.and_then(|value| value.to_str().ok()?.parse::().ok()); if let Some(length) = content_length { if length > N { return Err(PayloadTooLarge.into_response()); } } else { return Err(LengthRequired.into_response()); }; let bytes = hyper::body::to_bytes(body) .await .map_err(|e| FailedToBufferBody::from_err(e).into_response())?; Ok(BytesMaxLength(bytes)) } } /// Extractor that will get captures from the URL. /// /// # Example /// /// ```rust,no_run /// use tower_web::prelude::*; /// /// async fn users_show(req: Request, params: extract::UrlParamsMap) { /// let id: Option<&str> = params.get("id"); /// /// // ... /// } /// /// let app = route("/users/:id", get(users_show)); /// ``` /// /// Note that you can only have one URL params extractor per handler. If you /// have multiple it'll response with `500 Internal Server Error`. #[derive(Debug)] pub struct UrlParamsMap(HashMap); impl UrlParamsMap { /// Look up the value for a key. pub fn get(&self, key: &str) -> Option<&str> { self.0.get(key).map(|s| &**s) } /// Look up the value for a key and parse it into a value of type `T`. pub fn get_typed(&self, key: &str) -> Option> where T: FromStr, { self.get(key).map(str::parse) } } #[async_trait] impl FromRequest for UrlParamsMap { type Rejection = Response; async fn from_request(req: &mut Request) -> Result { if let Some(params) = req .extensions_mut() .get_mut::>() { if let Some(params) = params.take() { Ok(Self(params.0.into_iter().collect())) } else { Err(UrlParamsAlreadyTaken.into_response()) } } else { Err(MissingRouteParams.into_response()) } } } /// Extractor that will get captures from the URL and parse them. /// /// # Example /// /// ```rust,no_run /// use tower_web::{extract::UrlParams, prelude::*}; /// use uuid::Uuid; /// /// async fn users_teams_show( /// req: Request, /// UrlParams(params): UrlParams<(Uuid, Uuid)>, /// ) { /// let user_id: Uuid = params.0; /// let team_id: Uuid = params.1; /// /// // ... /// } /// /// let app = route("/users/:user_id/team/:team_id", get(users_teams_show)); /// ``` /// /// Note that you can only have one URL params extractor per handler. If you /// have multiple it'll response with `500 Internal Server Error`. #[derive(Debug)] pub struct UrlParams(pub T); macro_rules! impl_parse_url { () => {}; ( $head:ident, $($tail:ident),* $(,)? ) => { #[async_trait] impl<$head, $($tail,)*> FromRequest for UrlParams<($head, $($tail,)*)> where $head: FromStr + Send, $( $tail: FromStr + Send, )* { type Rejection = Response; #[allow(non_snake_case)] async fn from_request(req: &mut Request) -> Result { let params = if let Some(params) = req .extensions_mut() .get_mut::>() { if let Some(params) = params.take() { params.0 } else { return Err(UrlParamsAlreadyTaken.into_response()); } } else { return Err(MissingRouteParams.into_response()) }; if let [(_, $head), $((_, $tail),)*] = &*params { let $head = if let Ok(x) = $head.parse::<$head>() { x } else { return Err(InvalidUrlParam::new::<$head>().into_response()); }; $( let $tail = if let Ok(x) = $tail.parse::<$tail>() { x } else { return Err(InvalidUrlParam::new::<$tail>().into_response()); }; )* Ok(UrlParams(($head, $($tail,)*))) } else { return Err(MissingRouteParams.into_response()) } } } impl_parse_url!($($tail,)*); }; } impl_parse_url!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13, T14, T15, T16); fn take_body(req: &mut Request) -> Result { struct BodyAlreadyTakenExt; if req.extensions_mut().insert(BodyAlreadyTakenExt).is_some() { Err(BodyAlreadyTaken) } else { let body = std::mem::take(req.body_mut()); Ok(body) } } macro_rules! impl_from_request_tuple { () => {}; ( $head:ident, $($tail:ident),* $(,)? ) => { #[allow(non_snake_case)] #[async_trait] impl FromRequest for ($head, $($tail,)*) where R: IntoResponse, $head: FromRequest + Send, $( $tail: FromRequest + Send, )* { type Rejection = R; async fn from_request(req: &mut Request) -> Result { let $head = FromRequest::from_request(req).await?; $( let $tail = FromRequest::from_request(req).await?; )* Ok(($head, $($tail,)*)) } } impl_from_request_tuple!($($tail,)*); }; } impl_from_request_tuple!(T1, T2, T3, T4, T5, T6, T7, T8, T9, T10, T11, T12, T13, T14, T15, T16);