use super::{
JsonResponse, Request, RequestError, StreamRequest,
parts::{AxumRequestPreparationExt, PreparedRequestParts},
};
use sword_core::State;
use axum::{
body::to_bytes,
extract::{FromRef, FromRequest as AxumFromRequest, Request as AxumReq},
http::request::Parts,
response::IntoResponse,
};
pub trait FromRequest: Sized {
type Rejection: IntoResponse;
fn from_request(
req: AxumReq,
state: &State,
) -> impl Future<Output = Result<Self, Self::Rejection>> + Send;
}
pub trait FromRequestParts: Sized {
type Rejection: IntoResponse;
fn from_request_parts(
parts: &mut Parts,
state: &State,
) -> impl Future<Output = Result<Self, Self::Rejection>> + Send;
}
impl FromRequest for Request {
type Rejection = JsonResponse;
async fn from_request(req: AxumReq, _: &State) -> Result<Self, Self::Rejection> {
let PreparedRequestParts {
params,
parts,
body,
body_limit,
} = req.prepare().await?;
let body_bytes = to_bytes(body, body_limit).await.map_err(|err| {
let inner = err.into_inner();
RequestError::from_body_read_error(inner.as_ref())
})?;
Ok(Self {
params,
body_bytes,
method: parts.method,
headers: parts.headers,
uri: parts.uri,
extensions: parts.extensions,
next: None,
})
}
}
impl FromRequest for StreamRequest {
type Rejection = JsonResponse;
async fn from_request(req: AxumReq, _: &State) -> Result<Self, Self::Rejection> {
let PreparedRequestParts {
params,
parts,
body,
body_limit,
} = req.prepare().await?;
Ok(Self {
params,
body,
method: parts.method,
headers: parts.headers,
uri: parts.uri,
extensions: parts.extensions,
next: None,
body_limit,
})
}
}
impl TryFrom<Request> for AxumReq {
type Error = RequestError;
fn try_from(req: Request) -> Result<Self, Self::Error> {
use axum::body::Body;
let mut builder = AxumReq::builder().method(req.method).uri(req.uri);
for (key, value) in &req.headers {
builder = builder.header(key, value);
}
let body = Body::from(req.body_bytes);
let mut request = builder.body(body).map_err(|_| {
RequestError::parse_error(
"Failed to build axum request",
"Error building request".to_string(),
)
})?;
*request.extensions_mut() = req.extensions;
Ok(request)
}
}
impl TryFrom<StreamRequest> for AxumReq {
type Error = RequestError;
fn try_from(req: StreamRequest) -> Result<Self, Self::Error> {
let mut builder = AxumReq::builder().method(req.method).uri(req.uri);
for (key, value) in &req.headers {
builder = builder.header(key, value);
}
let mut request = builder.body(req.body).map_err(|_| {
RequestError::parse_error(
"Failed to build axum request",
"Error building request".to_string(),
)
})?;
*request.extensions_mut() = req.extensions;
Ok(request)
}
}
impl<S> AxumFromRequest<S> for Request
where
S: Send + Sync + 'static,
State: FromRef<S>,
{
type Rejection = JsonResponse;
async fn from_request(req: AxumReq, state: &S) -> Result<Self, Self::Rejection> {
let state = State::from_ref(state);
<Self as FromRequest>::from_request(req, &state).await
}
}
impl<S> AxumFromRequest<S> for StreamRequest
where
S: Send + Sync + 'static,
State: FromRef<S>,
{
type Rejection = JsonResponse;
async fn from_request(req: AxumReq, state: &S) -> Result<Self, Self::Rejection> {
let state = State::from_ref(state);
<Self as FromRequest>::from_request(req, &state).await
}
}