use ::axum::{
Json,
extract::{FromRequest, FromRequestParts, Request},
http::{HeaderValue, StatusCode, header::CONTENT_TYPE, request::Parts},
response::{IntoResponse, Response},
};
use ::garde::Validate;
use ::serde::{Serialize, Serializer, ser::SerializeStruct};
use std::{
fmt::Display,
marker::PhantomData,
ops::{Deref, DerefMut},
};
#[derive(Debug, Clone, Copy)]
pub struct Validated<T>(pub T);
pub trait ValidatePayload {
type Inner: Validate<Context = ()>;
fn payload(&self) -> &Self::Inner;
}
pub trait ValidationContext<C> {
fn validation_context(&self) -> &C;
}
impl<C> ValidationContext<C> for C {
fn validation_context(&self) -> &C {
self
}
}
pub trait ValidatePayloadWith<C> {
type Inner: Validate<Context = C>;
fn payload(&self) -> &Self::Inner;
}
impl<U, C> ValidatePayloadWith<C> for Json<U>
where
U: Validate<Context = C>,
{
type Inner = U;
fn payload(&self) -> &U {
&self.0
}
}
impl<U, C> ValidatePayloadWith<C> for ::axum::Form<U>
where
U: Validate<Context = C>,
{
type Inner = U;
fn payload(&self) -> &U {
&self.0
}
}
impl<U, C> ValidatePayloadWith<C> for ::axum::extract::Query<U>
where
U: Validate<Context = C>,
{
type Inner = U;
fn payload(&self) -> &U {
&self.0
}
}
impl<U, C> ValidatePayloadWith<C> for ::axum::extract::Path<U>
where
U: Validate<Context = C>,
{
type Inner = U;
fn payload(&self) -> &U {
&self.0
}
}
impl<U, C> ValidatePayloadWith<C> for crate::multipart::TypedMultipart<U>
where
U: Validate<Context = C>,
{
type Inner = U;
fn payload(&self) -> &U {
&self.0
}
}
#[derive(Debug, Clone, Copy)]
pub struct ValidatedWith<C, T>(pub T, PhantomData<fn() -> C>);
impl<C, T> ValidatedWith<C, T> {
#[must_use]
pub const fn new(value: T) -> Self {
Self(value, PhantomData)
}
#[must_use]
pub fn into_inner(self) -> T {
self.0
}
#[must_use]
pub const fn get(&self) -> &T {
&self.0
}
#[must_use]
pub const fn get_mut(&mut self) -> &mut T {
&mut self.0
}
}
impl<C, T> Deref for ValidatedWith<C, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl<C, T> DerefMut for ValidatedWith<C, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.0
}
}
impl<T> ValidatePayload for T
where
T: ValidatePayloadWith<()>,
{
type Inner = <T as ValidatePayloadWith<()>>::Inner;
fn payload(&self) -> &Self::Inner {
<Self as ValidatePayloadWith<()>>::payload(self)
}
}
impl<S, T> FromRequest<S> for Validated<T>
where
S: Send + Sync,
T: FromRequest<S> + ValidatePayload + Send,
{
type Rejection = Response;
async fn from_request(req: Request, state: &S) -> Result<Self, Self::Rejection> {
let extracted = T::from_request(req, state)
.await
.map_err(IntoResponse::into_response)?;
match extracted.payload().validate() {
Ok(()) => Ok(Self(extracted)),
Err(report) => Err(build_validation_response(&report)),
}
}
}
impl<S, T> FromRequestParts<S> for Validated<T>
where
S: Send + Sync,
T: FromRequestParts<S> + ValidatePayload + Send,
{
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
let extracted = T::from_request_parts(parts, state)
.await
.map_err(IntoResponse::into_response)?;
match extracted.payload().validate() {
Ok(()) => Ok(Self(extracted)),
Err(report) => Err(build_validation_response(&report)),
}
}
}
impl<S, C, T> FromRequest<S> for ValidatedWith<C, T>
where
S: Send + Sync + ValidationContext<C>,
C: Send + Sync + 'static,
T: FromRequest<S> + ValidatePayloadWith<C> + Send,
{
type Rejection = Response;
async fn from_request(req: Request, state: &S) -> Result<Self, Self::Rejection> {
let extracted = T::from_request(req, state)
.await
.map_err(IntoResponse::into_response)?;
match extracted
.payload()
.validate_with(state.validation_context())
{
Ok(()) => Ok(Self::new(extracted)),
Err(report) => Err(build_validation_response(&report)),
}
}
}
impl<S, C, T> FromRequestParts<S> for ValidatedWith<C, T>
where
S: Send + Sync + ValidationContext<C>,
C: Send + Sync + 'static,
T: FromRequestParts<S> + ValidatePayloadWith<C> + Send,
{
type Rejection = Response;
async fn from_request_parts(parts: &mut Parts, state: &S) -> Result<Self, Self::Rejection> {
let extracted = T::from_request_parts(parts, state)
.await
.map_err(IntoResponse::into_response)?;
match extracted
.payload()
.validate_with(state.validation_context())
{
Ok(()) => Ok(Self::new(extracted)),
Err(report) => Err(build_validation_response(&report)),
}
}
}
const FALLBACK_VALIDATION_ENVELOPE: &[u8] =
br#"{"errors":[{"message":"request validation failed","path":""}]}"#;
fn build_validation_response(report: &::garde::Report) -> Response {
struct DisplayValue<T>(T);
impl<T> Serialize for DisplayValue<T>
where
T: Display,
{
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.collect_str(&self.0)
}
}
struct ValidationEnvelope<'a> {
report: &'a ::garde::Report,
}
impl Serialize for ValidationEnvelope<'_> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut envelope = serializer.serialize_struct("ValidationEnvelope", 1)?;
envelope.serialize_field(
"errors",
&ValidationErrors {
report: self.report,
},
)?;
envelope.end()
}
}
struct ValidationErrors<'a> {
report: &'a ::garde::Report,
}
impl Serialize for ValidationErrors<'_> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.collect_seq(
self.report
.iter()
.map(|(path, err)| ValidationError { path, err }),
)
}
}
struct ValidationError<'a> {
path: &'a ::garde::Path,
err: &'a ::garde::Error,
}
impl Serialize for ValidationError<'_> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut error = serializer.serialize_struct("ValidationError", 2)?;
error.serialize_field("message", &DisplayValue(self.err.message()))?;
error.serialize_field("path", &DisplayValue(self.path))?;
error.end()
}
}
let body = ::serde_json::to_vec(&ValidationEnvelope { report })
.unwrap_or_else(|_| FALLBACK_VALIDATION_ENVELOPE.to_vec());
(
StatusCode::UNPROCESSABLE_ENTITY,
[(CONTENT_TYPE, HeaderValue::from_static("application/json"))],
body,
)
.into_response()
}
#[cfg(test)]
mod tests {
use super::FALLBACK_VALIDATION_ENVELOPE;
#[test]
fn fallback_envelope_matches_the_snapshot_locked_shape() {
let parsed: serde_json::Value = serde_json::from_slice(FALLBACK_VALIDATION_ENVELOPE)
.expect("the fallback envelope must be valid JSON");
let errors = parsed["errors"]
.as_array()
.expect("the fallback envelope must carry an `errors` array");
assert_eq!(errors.len(), 1);
assert_eq!(errors[0]["message"], "request validation failed");
assert_eq!(errors[0]["path"], "");
let text = std::str::from_utf8(FALLBACK_VALIDATION_ENVELOPE).expect("UTF-8");
assert!(
text.find(r#""message""#) < text.find(r#""path""#),
"fallback field order drifted: {text}"
);
}
}