use super::{rejection::*, take_body, Extension, FromRequest, RequestParts}; use crate::BoxError; use async_trait::async_trait; use bytes::Bytes; use futures_util::stream::Stream; use http::{Extensions, HeaderMap, Method, Request, Uri, Version}; use std::{ convert::Infallible, pin::Pin, task::{Context, Poll}, }; #[async_trait] impl FromRequest for Request where B: Send, { type Rejection = RequestAlreadyExtracted; async fn from_request(req: &mut RequestParts) -> Result { let req = std::mem::replace( req, RequestParts { method: req.method.clone(), version: req.version, uri: req.uri.clone(), headers: None, extensions: None, body: None, }, ); let err = match req.try_into_request() { Ok(req) => return Ok(req), Err(err) => err, }; match err.downcast::() { Ok(err) => return Err(err), Err(err) => unreachable!( "Unexpected error type from `try_into_request`: `{:?}`. This is a bug in axum, please file an issue", err, ), } } } #[async_trait] impl FromRequest for Body where B: Send, { type Rejection = BodyAlreadyExtracted; async fn from_request(req: &mut RequestParts) -> Result { let body = take_body(req)?; Ok(Self(body)) } } #[async_trait] impl FromRequest for Method where B: Send, { type Rejection = Infallible; async fn from_request(req: &mut RequestParts) -> Result { Ok(req.method().clone()) } } #[async_trait] impl FromRequest for Uri where B: Send, { type Rejection = Infallible; async fn from_request(req: &mut RequestParts) -> Result { Ok(req.uri().clone()) } } /// Extractor that gets the original request URI regardless of nesting. /// /// This is necessary since [`Uri`](http::Uri), when used as an extractor, will /// have the prefix stripped if used in a nested service. /// /// # Example /// /// ``` /// use axum::{ /// handler::get, /// Router, /// extract::OriginalUri, /// http::Uri /// }; /// /// let api_routes = Router::new() /// .route( /// "/users", /// get(|uri: Uri, OriginalUri(original_uri): OriginalUri| async { /// // `uri` is `/users` /// // `original_uri` is `/api/users` /// }), /// ); /// /// let app = Router::new().nest("/api", api_routes); /// # async { /// # axum::Server::bind(&"".parse().unwrap()).serve(app.into_make_service()).await.unwrap(); /// # }; /// ``` #[derive(Debug, Clone)] pub struct OriginalUri(pub Uri); #[async_trait] impl FromRequest for OriginalUri where B: Send, { type Rejection = Infallible; async fn from_request(req: &mut RequestParts) -> Result { let uri = Extension::::from_request(req) .await .unwrap_or_else(|_| Extension(OriginalUri(req.uri().clone()))) .0; Ok(uri) } } #[async_trait] impl FromRequest for Version where B: Send, { type Rejection = Infallible; async fn from_request(req: &mut RequestParts) -> Result { Ok(req.version()) } } #[async_trait] impl FromRequest for HeaderMap where B: Send, { type Rejection = HeadersAlreadyExtracted; async fn from_request(req: &mut RequestParts) -> Result { req.take_headers().ok_or(HeadersAlreadyExtracted) } } #[async_trait] impl FromRequest for Extensions where B: Send, { type Rejection = ExtensionsAlreadyExtracted; async fn from_request(req: &mut RequestParts) -> Result { req.take_extensions().ok_or(ExtensionsAlreadyExtracted) } } /// Extractor that extracts the request body as a [`Stream`]. /// /// # Example /// /// ```rust,no_run /// use axum::{ /// extract::BodyStream, /// handler::get, /// Router, /// }; /// use futures::StreamExt; /// /// async fn handler(mut stream: BodyStream) { /// while let Some(chunk) = stream.next().await { /// // ... /// } /// } /// /// let app = Router::new().route("/users", get(handler)); /// # async { /// # axum::Server::bind(&"".parse().unwrap()).serve(app.into_make_service()).await.unwrap(); /// # }; /// ``` /// /// [`Stream`]: https://docs.rs/futures/latest/futures/stream/trait.Stream.html #[derive(Debug)] pub struct BodyStream(B); impl Stream for BodyStream where B: http_body::Body + Unpin, { type Item = Result; fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { Pin::new(&mut self.0).poll_data(cx) } } #[async_trait] impl FromRequest for BodyStream where B: http_body::Body + Unpin + Send, { type Rejection = BodyAlreadyExtracted; async fn from_request(req: &mut RequestParts) -> Result { let body = take_body(req)?; let stream = BodyStream(body); Ok(stream) } } /// Extractor that extracts the request body. /// /// # Example /// /// ```rust,no_run /// use axum::{ /// extract::Body, /// handler::get, /// Router, /// }; /// use futures::StreamExt; /// /// async fn handler(Body(body): Body) { /// // ... /// } /// /// let app = Router::new().route("/users", get(handler)); /// # async { /// # axum::Server::bind(&"".parse().unwrap()).serve(app.into_make_service()).await.unwrap(); /// # }; /// ``` #[derive(Debug, Default, Clone)] pub struct Body(pub B); #[async_trait] impl FromRequest for Bytes where B: http_body::Body + Send, B::Data: Send, B::Error: Into, { type Rejection = BytesRejection; async fn from_request(req: &mut RequestParts) -> Result { let body = take_body(req)?; let bytes = hyper::body::to_bytes(body) .await .map_err(FailedToBufferBody::from_err)?; Ok(bytes) } } #[async_trait] impl FromRequest for String where B: http_body::Body + Send, B::Data: Send, B::Error: Into, { type Rejection = StringRejection; async fn from_request(req: &mut RequestParts) -> Result { let body = take_body(req)?; let bytes = hyper::body::to_bytes(body) .await .map_err(FailedToBufferBody::from_err)? .to_vec(); let string = String::from_utf8(bytes).map_err(InvalidUtf8::from_err)?; Ok(string) } } #[cfg(test)] mod tests { use super::*; use crate::{body::Body, handler::post, tests::*, Router}; use http::StatusCode; #[tokio::test] async fn multiple_request_extractors() { async fn handler(_: Request, _: Request) {} let app = Router::new().route("/", post(handler)); let addr = run_in_background(app).await; let client = reqwest::Client::new(); let res = client .post(format!("http://{}", addr)) .body("hi there") .send() .await .unwrap(); assert_eq!(res.status(), StatusCode::INTERNAL_SERVER_ERROR); assert_eq!( res.text().await.unwrap(), "Cannot have two request body extractors for a single handler" ); } }