use http::header::{self as h, HeaderMap, HeaderName};
use std::fmt::Debug;
use super::bytes::{Contains, Trim, contains, trim};
use super::error::{InvalidHeader, OnError};
use super::on::{self, On};
use super::{Any, OkOr, Predicate, any, media, ok_or, or};
use crate::Request;
pub type Accept<T> = Contains<Trim<media::AllOr<T>>>;
pub struct Header<T> {
predicate: On<OkOr<T, HeaderName>, on::Header>,
}
pub struct Opt<T> {
predicate: on::Opt<OkOr<T, HeaderName>, on::Header>,
}
pub fn accept<T>(predicate: T) -> Header<Accept<T>> {
header(
h::ACCEPT,
contains(trim(or((media::all(), predicate))), b','),
)
}
pub fn content_type<T>(predicate: T) -> Header<T> {
header(h::CONTENT_TYPE, predicate)
}
pub fn content_length() -> Header<Any> {
exists(h::CONTENT_LENGTH)
}
pub fn exists<K>(key: K) -> Header<Any>
where
K: TryInto<HeaderName>,
K::Error: Debug,
{
header(key, any())
}
pub fn header<K, V>(key: K, value: V) -> Header<V>
where
K: TryInto<http::HeaderName>,
K::Error: Debug,
{
let key = key.try_into().expect("invalid header name");
let value = ok_or(value, key.clone());
Header {
predicate: on::header(value, key),
}
}
impl<T> Header<T> {
pub fn opt(self) -> Opt<T> {
Opt {
predicate: self.predicate.opt(),
}
}
}
impl<T> Predicate<HeaderMap> for Header<T>
where
for<'a> T: Predicate<[u8]> + 'a,
{
type Error<'a> = InvalidHeader<'a>;
fn cmp<'a>(&'a self, headers: &HeaderMap) -> Result<(), Self::Error<'a>> {
self.predicate.cmp(headers).map_err(InvalidHeader::new)
}
}
impl<T, App> Predicate<Request<App>> for Header<T>
where
for<'a> T: Predicate<[u8]> + 'a,
{
type Error<'a> = InvalidHeader<'a>;
fn cmp<'a>(&'a self, request: &Request<App>) -> Result<(), Self::Error<'a>> {
self.cmp(request.headers())
}
}
impl<T> Predicate<HeaderMap> for Opt<T>
where
for<'a> T: Predicate<[u8]> + 'a,
{
type Error<'a> = InvalidHeader<'a>;
fn cmp<'a>(&'a self, headers: &HeaderMap) -> Result<(), Self::Error<'a>> {
self.predicate
.cmp(headers)
.map_err(|error| InvalidHeader::new(OnError::Predicate(error)))
}
}
impl<T, App> Predicate<Request<App>> for Opt<T>
where
for<'a> T: Predicate<[u8]> + 'a,
{
type Error<'a> = InvalidHeader<'a>;
fn cmp<'a>(&'a self, request: &Request<App>) -> Result<(), Self::Error<'a>> {
self.cmp(request.headers())
}
}