use crate::middleware::BoxFuture;
use crate::{Error, Middleware, Next, Request, Response, Result};
pub struct ErrorBoundary;
pub struct InspectErrorBoundary<F> {
inspect: F,
}
pub struct MapErrorBoundary<F> {
map: F,
}
pub struct OrElseErrorBoundary<F> {
or_else: F,
}
impl ErrorBoundary {
pub fn inspect<State, F>(inspect: F) -> InspectErrorBoundary<F>
where
F: Fn(&Error, &State) + Copy + Send + Sync + 'static,
State: Send + Sync + 'static,
{
InspectErrorBoundary { inspect }
}
pub fn map<State, F>(map: F) -> MapErrorBoundary<F>
where
F: Fn(Error, &State) -> Error + Copy + Send + Sync + 'static,
State: Send + Sync + 'static,
{
MapErrorBoundary { map }
}
pub fn or_else<State, F>(or_else: F) -> OrElseErrorBoundary<F>
where
F: Fn(Error, &State) -> Result<Response, Error> + Copy + Send + Sync + 'static,
State: Send + Sync + 'static,
{
OrElseErrorBoundary { or_else }
}
}
impl<State> Middleware<State> for ErrorBoundary
where
State: Send + Sync + 'static,
{
fn call(&self, request: Request<State>, next: Next<State>) -> BoxFuture<Result<Response>> {
Box::pin(async {
let result = next.call(request).await;
Ok(result.unwrap_or_else(|error| error.into_response()))
})
}
}
impl<State, F> Middleware<State> for InspectErrorBoundary<F>
where
F: Fn(&Error, &State) + Copy + Send + Sync + 'static,
State: Send + Sync + 'static,
{
fn call(&self, request: Request<State>, next: Next<State>) -> BoxFuture<Result<Response>> {
let inspect = self.inspect;
let state = request.state().clone();
Box::pin(async move {
let result = next.call(request).await;
result.inspect_err(|error| {
inspect(error, &state);
})
})
}
}
impl<State, F> Middleware<State> for MapErrorBoundary<F>
where
F: Fn(Error, &State) -> Error + Copy + Send + Sync + 'static,
State: Send + Sync + 'static,
{
fn call(&self, request: Request<State>, next: Next<State>) -> BoxFuture<Result<Response>> {
let map = self.map;
let state = request.state().clone();
Box::pin(async move {
let result = next.call(request).await;
result.or_else(|error| {
let error = map(error, &state);
Ok(error.into_response())
})
})
}
}
impl<State, F> Middleware<State> for OrElseErrorBoundary<F>
where
F: Fn(Error, &State) -> Result<Response, Error> + Copy + Send + Sync + 'static,
State: Send + Sync + 'static,
{
fn call(&self, request: Request<State>, next: Next<State>) -> BoxFuture<Result<Response>> {
let or_else = self.or_else;
let state = request.state().clone();
Box::pin(async move {
let result = next.call(request).await;
result.or_else(|error| {
or_else(error, &state)
})
})
}
}