use std::{fmt::Display, str::FromStr};
use axum::extract::OptionalFromRequestParts;
use crate::headers::parser::HeaderParser;
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum AcceptHeaderItem {
Wildcard,
PartialWildcard(String),
MediaType(String),
}
impl Display for AcceptHeaderItem {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Wildcard => write!(f, "*/*"),
Self::PartialWildcard(media_type) => write!(f, "{media_type}/*"),
Self::MediaType(media_type) => write!(f, "{media_type}"),
}
}
}
impl FromStr for AcceptHeaderItem {
type Err = String;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let mut parser = HeaderParser::new(s);
if let Some(item) = parser.parse_accept_header_item() {
if !parser.is_at_end() {
return Err(format!(
"Found additional data after media type or wildcard. Expexted text to end after '{item}'"
));
}
return Ok(item);
} else {
return Err("Unable to parse accept header item, expected some/type[+foo], some/* or */* as imput.".to_string());
}
}
}
impl<'de> serde::de::Deserialize<'de> for AcceptHeaderItem {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
deserializer.deserialize_str(AcceptHeaderItemVisitor)
}
}
struct AcceptHeaderItemVisitor;
impl<'de> serde::de::Visitor<'de> for AcceptHeaderItemVisitor {
type Value = AcceptHeaderItem;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
write!(
f,
"The media_type part of an HTTP Accept header, i.e. some/type[+foo], some/* or */*"
)
}
fn visit_str<E>(self, v: &str) -> Result<Self::Value, E>
where
E: serde::de::Error,
{
AcceptHeaderItem::from_str(v).map_err(|e| E::custom(e.to_string()))
}
}
impl serde::Serialize for AcceptHeaderItem {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
serializer.serialize_str(self.to_string().as_str())
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct AcceptHeader {
pub items: Vec<(AcceptHeaderItem, f32)>,
}
impl AcceptHeader {
pub fn new(header_text: &str) -> Self {
HeaderParser::new(header_text).parse_accept()
}
}
impl FromStr for AcceptHeader {
type Err = ();
fn from_str(s: &str) -> Result<Self, Self::Err> {
Ok(HeaderParser::new(s).parse_accept())
}
}
impl<S> OptionalFromRequestParts<S> for AcceptHeader
where
S: Sync,
{
type Rejection = ();
async fn from_request_parts(
parts: &mut axum::http::request::Parts,
_state: &S,
) -> Result<Option<Self>, Self::Rejection> {
Ok(parts
.headers
.get("accept")
.and_then(|v| Some(AcceptHeader::new(v.to_str().ok()?))))
}
}