use super::*;
use crate::spec::MaskSpec;
use crate::validate::validate_mask;
use axum::extract::FromRequestParts;
use axum::http::{header, request::Parts, HeaderValue, StatusCode};
use axum::response::{IntoResponse, Response};
use serde::Deserialize;
use std::marker::PhantomData;
#[derive(Debug, Clone)]
pub struct MaskRejection {
pub status: StatusCode,
pub code: &'static str,
pub message: String,
}
impl MaskRejection {
fn invalid_argument(msg: impl Into<String>) -> Self {
Self {
status: StatusCode::BAD_REQUEST,
code: "INVALID_ARGUMENT",
message: msg.into(),
}
}
fn missing_mask() -> Self {
Self::invalid_argument("missing required field mask (use ?fields=... or 'x-fields' header)")
}
}
impl IntoResponse for MaskRejection {
fn into_response(self) -> Response {
let mut res = (self.status, format!("{}: {}", self.code, self.message)).into_response();
res.headers_mut().insert(
header::CONTENT_TYPE,
HeaderValue::from_static("text/plain; charset=utf-8"),
);
res
}
}
#[derive(Deserialize)]
struct QueryFields {
#[serde(rename = "fields")]
fields: Option<String>,
}
fn extract_mask_from_parts(parts: &Parts) -> Option<String> {
if let Some(qs) = parts.uri.query() {
if let Ok(q) = serde_urlencoded::from_str::<QueryFields>(qs) {
if let Some(v) = q.fields {
let v = v.trim();
if !v.is_empty() {
return Some(v.to_string());
}
}
}
}
if let Some(val) = parts.headers.get("x-fields") {
if let Ok(s) = val.to_str() {
let s = s.trim();
if !s.is_empty() {
return Some(s.to_string());
}
}
}
None
}
pub struct MaskRequired<T>(pub FieldMask, pub PhantomData<T>);
pub struct MaskOptional<T>(pub FieldMask, pub PhantomData<T>);
impl<S, T> FromRequestParts<S> for MaskRequired<T>
where
S: Send + Sync,
T: MaskSpec,
{
type Rejection = MaskRejection;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
let raw = extract_mask_from_parts(parts).ok_or_else(MaskRejection::missing_mask)?;
let mask =
FieldMask::parse(&raw).map_err(|e| MaskRejection::invalid_argument(e.to_string()))?;
validate_mask(&mask, T::mask_spec())
.map_err(|e| MaskRejection::invalid_argument(e.to_string()))?;
Ok(MaskRequired(mask, PhantomData))
}
}
impl<S, T> FromRequestParts<S> for MaskOptional<T>
where
S: Send + Sync,
T: MaskSpec,
{
type Rejection = MaskRejection;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
let mask = match extract_mask_from_parts(parts) {
None => FieldMask::all(),
Some(raw) => FieldMask::parse(&raw)
.map_err(|e| MaskRejection::invalid_argument(e.to_string()))?,
};
validate_mask(&mask, T::mask_spec())
.map_err(|e| MaskRejection::invalid_argument(e.to_string()))?;
Ok(MaskOptional(mask, PhantomData))
}
}