use axum::{
extract::{FromRequest, Request},
http::StatusCode,
response::{IntoResponse, Response},
Json,
};
use serde::{de::DeserializeOwned, Serialize};
use crate::error::json_rejection_to_api_error;
use crate::ApiError;
#[derive(Debug, Clone)]
pub struct ApiJson<T>(pub T);
impl<T, S> FromRequest<S> for ApiJson<T>
where
T: DeserializeOwned,
S: Send + Sync,
{
type Rejection = (StatusCode, Json<ApiError>);
async fn from_request(req: Request, state: &S) -> Result<Self, Self::Rejection> {
let Json(value) = Json::<T>::from_request(req, state)
.await
.map_err(json_rejection_to_api_error)?;
Ok(ApiJson(value))
}
}
impl<T: Serialize> IntoResponse for ApiJson<T> {
fn into_response(self) -> Response {
Json(self.0).into_response()
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::{body::Body, http::header::CONTENT_TYPE, http::Request};
use serde::Deserialize;
#[derive(Debug, Deserialize)]
struct Input {
name: String,
}
async fn extract(body: &str, with_content_type: bool) -> Result<Input, (StatusCode, ApiError)> {
let mut builder = Request::builder().method("POST").uri("/");
if with_content_type {
builder = builder.header(CONTENT_TYPE, "application/json");
}
let req = builder.body(Body::from(body.to_owned())).unwrap();
ApiJson::<Input>::from_request(req, &())
.await
.map(|ApiJson(v)| v)
.map_err(|(status, Json(err))| (status, err))
}
#[tokio::test]
async fn valid_body_extracts() {
let input = extract(r#"{"name":"abc"}"#, true).await.unwrap();
assert_eq!(input.name, "abc");
}
#[tokio::test]
async fn malformed_json_is_invalid_json() {
let (status, err) = extract("{not json", true).await.unwrap_err();
assert_eq!(status, StatusCode::BAD_REQUEST);
assert_eq!(err.code, "INVALID_JSON");
}
#[tokio::test]
async fn wrong_shape_is_invalid_body() {
let (status, err) = extract(r#"{"name":123}"#, true).await.unwrap_err();
assert_eq!(status, StatusCode::UNPROCESSABLE_ENTITY);
assert_eq!(err.code, "INVALID_BODY");
}
#[tokio::test]
async fn missing_content_type_is_unsupported_media_type() {
let (status, err) = extract(r#"{"name":"abc"}"#, false).await.unwrap_err();
assert_eq!(status, StatusCode::UNSUPPORTED_MEDIA_TYPE);
assert_eq!(err.code, "UNSUPPORTED_MEDIA_TYPE");
}
#[tokio::test]
async fn serializes_as_response() {
#[derive(Serialize)]
struct Out {
id: u32,
}
let res = ApiJson(Out { id: 7 }).into_response();
assert_eq!(res.status(), StatusCode::OK);
let bytes = axum::body::to_bytes(res.into_body(), usize::MAX)
.await
.unwrap();
let v: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert_eq!(v["id"], 7);
}
}