diff --git a/examples/key_value_store.rs b/examples/key_value_store.rs
index 02b5658c..bbd2a0d4 100644
--- a/examples/key_value_store.rs
+++ b/examples/key_value_store.rs
@@ -1,7 +1,7 @@
#![allow(warnings)]
use bytes::Bytes;
-use http::{Request, StatusCode};
+use http::{Request, Response, StatusCode};
use hyper::Server;
use serde::Deserialize;
use std::{
@@ -14,7 +14,12 @@ use tower::{make::Shared, ServiceBuilder};
use tower_http::{
add_extension::AddExtensionLayer, compression::CompressionLayer, trace::TraceLayer,
};
-use tower_web::{body::Body, extract, response, Error};
+use tower_web::{
+ body::Body,
+ extract,
+ response::{self, IntoResponse},
+ Error,
+};
#[tokio::main]
async fn main() {
@@ -54,16 +59,16 @@ async fn get(
_req: Request
,
params: extract::UrlParams<(String,)>,
state: extract::Extension,
-) -> Result {
+) -> Result {
let state = state.into_inner();
let db = &state.lock().unwrap().db;
- let (key,) = params.into_inner();
+ let key = params.into_inner();
if let Some(value) = db.get(&key) {
Ok(value.clone())
} else {
- Err(Error::Status(StatusCode::NOT_FOUND))
+ Err(NotFound)
}
}
@@ -72,14 +77,23 @@ async fn set(
params: extract::UrlParams<(String,)>,
value: extract::BytesMaxLength<{ 1024 * 5_000 }>, // ~5mb
state: extract::Extension,
-) -> response::Empty {
+) {
let state = state.into_inner();
let db = &mut state.lock().unwrap().db;
- let (key,) = params.into_inner();
+ let key = params.into_inner();
let value = value.into_inner();
db.insert(key.to_string(), value);
-
- response::Empty
+}
+
+struct NotFound;
+
+impl IntoResponse for NotFound {
+ fn into_response(self) -> Response {
+ Response::builder()
+ .status(StatusCode::NOT_FOUND)
+ .body(Body::empty())
+ .unwrap()
+ }
}
diff --git a/src/extract.rs b/src/extract.rs
index 3a9f7d53..9f0aa718 100644
--- a/src/extract.rs
+++ b/src/extract.rs
@@ -1,36 +1,85 @@
-use crate::{body::Body, Error};
+use crate::{
+ body::Body,
+ response::{BoxIntoResponse, IntoResponse},
+ Error,
+};
use async_trait::async_trait;
use bytes::Bytes;
-use http::{header, Request, StatusCode};
+use http::{header, Request};
use serde::de::DeserializeOwned;
-use std::{collections::HashMap, str::FromStr};
+use std::{collections::HashMap, convert::Infallible, str::FromStr};
#[async_trait]
-pub trait FromRequest: Sized {
- async fn from_request(req: &mut Request) -> Result;
-}
+pub trait FromRequest: Sized {
+ type Rejection: IntoResponse;
-fn take_body(req: &mut Request) -> Body {
- struct BodyAlreadyTaken;
-
- if req.extensions_mut().insert(BodyAlreadyTaken).is_some() {
- panic!("Cannot have two request body on extractors")
- } else {
- let body = std::mem::take(req.body_mut());
- body
- }
+ async fn from_request(req: &mut Request) -> Result;
}
#[async_trait]
-impl FromRequest for Option
+impl FromRequest for Option
where
- T: FromRequest,
+ T: FromRequest,
{
- async fn from_request(req: &mut Request) -> Result