use bitflags::bitflags;
use http::StatusCode;
use http::header::{ALLOW, CONNECTION};
use std::borrow::Cow;
use std::fmt::Display;
use std::sync::Arc;
use super::{Error, ErrorSourceRef, Errors};
use crate::middleware::{BoxFuture, Middleware};
use crate::response::{Finalize, Response, ResponseBuilder};
use crate::util::sealed;
use crate::{Next, Request};
pub trait Policy: sealed::Sealed {
fn apply(&self, error: &Error, flags: PolicyFlags) -> PolicyFlags;
}
pub struct And<T, U>(T, U);
pub struct CanonicalReasonPhrase;
pub struct ConnectionClose;
pub struct Json;
pub struct PolicyBuilder<T> {
policy: T,
}
pub struct PolicyFlags {
mask: PolicyMask,
}
pub struct Inspect<F> {
op: F,
}
pub struct Rescue<T> {
policy: Arc<T>,
}
struct Render<'a> {
flags: &'a PolicyMask,
error: &'a Error,
}
bitflags! {
#[derive(Clone, Copy, Eq, PartialEq)]
struct PolicyMask: u8 {
const CANONICAL_REASON_PHRASE = 1 << 3;
const CONNECTION_CLOSE = 1 << 4;
const JSON_RESPONSE = 1 << 5;
}
}
sealed!(
And<T, U>,
CanonicalReasonPhrase,
ConnectionClose,
Inspect<F>,
Json
);
pub fn canonical_reason_phrase() -> PolicyBuilder<CanonicalReasonPhrase> {
PolicyBuilder {
policy: CanonicalReasonPhrase,
}
}
pub fn connection_close() -> PolicyBuilder<ConnectionClose> {
PolicyBuilder {
policy: ConnectionClose,
}
}
pub fn inspect<F>(op: F) -> PolicyBuilder<Inspect<F>>
where
F: Fn(&dyn Display) + Copy + Send + Sync + 'static,
{
PolicyBuilder {
policy: Inspect { op },
}
}
pub fn json() -> PolicyBuilder<Json> {
PolicyBuilder { policy: Json }
}
impl<T> PolicyBuilder<T> {
pub fn and<U>(self, other: PolicyBuilder<U>) -> PolicyBuilder<And<T, U>> {
PolicyBuilder {
policy: And(self.policy, other.policy),
}
}
pub fn build(self) -> Rescue<T> {
Rescue {
policy: Arc::new(self.policy),
}
}
}
impl<T: Policy, U: Policy> Policy for And<T, U> {
fn apply(&self, error: &Error, flags: PolicyFlags) -> PolicyFlags {
let flags = self.0.apply(error, flags);
self.1.apply(error, flags)
}
}
impl Policy for CanonicalReasonPhrase {
fn apply(&self, _: &Error, flags: PolicyFlags) -> PolicyFlags {
flags.set(PolicyMask::CANONICAL_REASON_PHRASE)
}
}
impl Policy for ConnectionClose {
fn apply(&self, _: &Error, flags: PolicyFlags) -> PolicyFlags {
flags.set(PolicyMask::CONNECTION_CLOSE)
}
}
impl<F> Policy for Inspect<F>
where
F: Fn(&dyn Display) + Send + Sync + 'static,
{
fn apply(&self, error: &Error, flags: PolicyFlags) -> PolicyFlags {
(self.op)(error);
flags
}
}
impl Policy for Json {
fn apply(&self, _: &Error, flags: PolicyFlags) -> PolicyFlags {
flags.set(PolicyMask::JSON_RESPONSE)
}
}
impl PolicyFlags {
fn new() -> Self {
Self {
mask: PolicyMask::empty(),
}
}
fn set(self, flag: PolicyMask) -> Self {
Self {
mask: self.mask | flag,
}
}
}
impl<T, App> Middleware<App> for Rescue<T>
where
T: Policy + Send + Sync + 'static,
{
fn call(&self, request: Request<App>, next: Next<App>) -> BoxFuture {
let policy = Arc::clone(&self.policy);
let future = next.call(request);
Box::pin(async move {
future.await.or_else(|error| {
let flags = policy.apply(&error, PolicyFlags::new());
Render::new(&flags.mask, &error)
.finalize(Response::build())
.or_else(|residual| {
log!(warn(rescue = 0), "a residual error occurred in rescue");
log!(warn(rescue = 1), "{}", &residual);
Ok(error.into())
})
})
})
}
}
impl<'a> Render<'a> {
fn new(flags: &'a PolicyMask, error: &'a Error) -> Self {
Self { flags, error }
}
}
impl Finalize for Render<'_> {
fn finalize(self, mut builder: ResponseBuilder) -> Result<Response, Error> {
builder = builder.status(self.error.status);
if let ErrorSourceRef::AllowMethod(error) = self.error.as_source()
&& self.error.status == StatusCode::METHOD_NOT_ALLOWED
&& let Some(allow) = error.allows()
{
builder = builder.header(ALLOW, allow);
}
if self.flags.contains(PolicyMask::CONNECTION_CLOSE) {
builder = builder.header(CONNECTION, "close");
}
if self.flags.contains(PolicyMask::JSON_RESPONSE) {
if self.flags.contains(PolicyMask::CANONICAL_REASON_PHRASE)
&& let Some(reason_phrase) = self.error.status.canonical_reason()
{
let mut errors = Errors::new(self.error.status);
errors.push(Cow::Borrowed(reason_phrase));
builder.json(&errors)
} else {
let errors = self.error.repr_json();
builder.json(&errors)
}
} else if self.flags.contains(PolicyMask::CANONICAL_REASON_PHRASE)
&& let Some(reason_phrase) = self.error.status.canonical_reason()
{
builder.text(reason_phrase)
} else {
builder.text(self.error.to_string())
}
}
}