Make extractors easier to write (#36)

Previously extractors worked directly on `Request<B>` which meant you
had to do weird tricks like `mem::take(req.headers_mut())` to get owned
parts of the request.

This changes that instead to use a new `RequestParts` type that have
methods to "take" each part of the request. Without having to do weird
tricks.

Also removed the need to have `B: Default` for body extractors.
This commit is contained in:
David Pedersen
2021-07-22 13:23:50 +02:00
committed by GitHub
parent e544fe1c39
commit f32d325e55
12 changed files with 441 additions and 193 deletions
+34 -4
View File
@@ -7,7 +7,8 @@
//! ```
use axum::{
extract::{ContentLengthLimit, Extension, UrlParams},
async_trait,
extract::{extractor_middleware, ContentLengthLimit, Extension, RequestParts, UrlParams},
prelude::*,
response::IntoResponse,
routing::BoxRoute,
@@ -24,8 +25,7 @@ use std::{
};
use tower::{BoxError, ServiceBuilder};
use tower_http::{
add_extension::AddExtensionLayer, auth::RequireAuthorizationLayer,
compression::CompressionLayer, trace::TraceLayer,
add_extension::AddExtensionLayer, compression::CompressionLayer, trace::TraceLayer,
};
#[tokio::main]
@@ -118,10 +118,40 @@ fn admin_routes() -> BoxRoute<hyper::Body> {
route("/keys", delete(delete_all_keys))
.route("/key/:key", delete(remove_key))
// Require beare auth for all admin routes
.layer(RequireAuthorizationLayer::bearer("secret-token"))
.layer(extractor_middleware::<RequireAuth>())
.boxed()
}
/// An extractor that performs authorization.
// TODO: when https://github.com/hyperium/http-body/pull/46 is merged we can use
// `tower_http::auth::RequireAuthorization` instead
struct RequireAuth;
#[async_trait]
impl<B> extract::FromRequest<B> for RequireAuth
where
B: Send,
{
type Rejection = StatusCode;
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
let auth_header = req
.headers()
.and_then(|headers| headers.get(http::header::AUTHORIZATION))
.and_then(|value| value.to_str().ok());
if let Some(value) = auth_header {
if let Some(token) = value.strip_prefix("Bearer ") {
if token == "secret-token" {
return Ok(Self);
}
}
}
Err(StatusCode::UNAUTHORIZED)
}
}
fn handle_error(error: BoxError) -> impl IntoResponse {
if error.is::<tower::timeout::error::Elapsed>() {
return (StatusCode::REQUEST_TIMEOUT, Cow::from("request timed out"));
+6 -2
View File
@@ -1,5 +1,9 @@
use axum::response::IntoResponse;
use axum::{async_trait, extract::FromRequest, prelude::*};
use axum::{
async_trait,
extract::{FromRequest, RequestParts},
prelude::*,
};
use http::Response;
use http::StatusCode;
use std::net::SocketAddr;
@@ -36,7 +40,7 @@ where
{
type Rejection = Response<Body>;
async fn from_request(req: &mut Request<B>) -> Result<Self, Self::Rejection> {
async fn from_request(req: &mut RequestParts<B>) -> Result<Self, Self::Rejection> {
let params = extract::UrlParamsMap::from_request(req)
.await
.map_err(IntoResponse::into_response)?;