Skip to main content

rama_http/layer/
error_handling.rs

1//! Middleware to turn [`Service`] errors into [`Response`]s.
2//!
3//! # Example
4//!
5//! ```
6//! use rama_core::{
7//!     service::service_fn,
8//!     Service, Layer,
9//!     telemetry::tracing,
10//! };
11//! use rama_http::{
12//!     service::client::HttpClientExt,
13//!     layer::{error_handling::ErrorHandlerLayer, timeout::TimeoutLayer},
14//!     service::web::WebService,
15//!     service::web::response::IntoResponse,
16//!     Body, Request, Response, StatusCode,
17//! };
18//! use std::time::Duration;
19//!
20//! # async fn some_expensive_io_operation() -> Result<(), std::io::Error> {
21//! #     Ok(())
22//! # }
23//!
24//! async fn handler(_req: Request) -> Result<Response, std::io::Error> {
25//!     some_expensive_io_operation().await?;
26//!     Ok(StatusCode::OK.into_response())
27//! }
28//!
29//! # #[tokio::main]
30//! # async fn main() {
31//!     let home_handler = (
32//!         ErrorHandlerLayer::new().error_mapper(|err| {
33//!             tracing::error!("Error: {err:?}");
34//!             StatusCode::INTERNAL_SERVER_ERROR.into_response()
35//!         }),
36//!         TimeoutLayer::new(Duration::from_secs(5)),
37//!         ).into_layer(service_fn(handler));
38//!
39//!     let service = WebService::default().with_get("/", home_handler);
40//!
41//!     _ = service.serve(Request::builder()
42//!         .method("GET")
43//!         .uri("/")
44//!         .body(Body::empty())
45//!         .unwrap()).await;
46//! # }
47//! ```
48
49use crate::service::web::response::{ErrorResponse, IntoResponse};
50use crate::{Request, Response};
51use rama_core::{Layer, Service};
52use rama_utils::macros::define_inner_service_accessors;
53use std::convert::Infallible;
54
55/// A [`Layer`] that wraps a [`Service`] and converts errors into [`Response`]s.
56#[derive(Debug, Clone)]
57pub struct ErrorHandlerLayer<F = ()> {
58    error_mapper: F,
59}
60
61impl Default for ErrorHandlerLayer {
62    fn default() -> Self {
63        Self::new()
64    }
65}
66
67impl ErrorHandlerLayer {
68    /// Create a new [`ErrorHandlerLayer`].
69    #[must_use]
70    pub const fn new() -> Self {
71        Self { error_mapper: () }
72    }
73
74    /// Set the error mapper function (not set by default).
75    ///
76    /// The error mapper function is called with the error,
77    /// and should return an [`IntoResponse`] implementation.
78    pub fn error_mapper<F>(self, error_mapper: F) -> ErrorHandlerLayer<F> {
79        ErrorHandlerLayer { error_mapper }
80    }
81}
82
83impl<S, F: Clone> Layer<S> for ErrorHandlerLayer<F> {
84    type Service = ErrorHandler<S, F>;
85
86    fn layer(&self, inner: S) -> Self::Service {
87        ErrorHandler::new(inner).error_mapper(self.error_mapper.clone())
88    }
89
90    fn into_layer(self, inner: S) -> Self::Service {
91        ErrorHandler::new(inner).error_mapper(self.error_mapper)
92    }
93}
94
95/// A [`Service`] adapter that handles errors by converting them into [`Response`]s.
96#[derive(Debug, Clone)]
97pub struct ErrorHandler<S, F = ()> {
98    inner: S,
99    error_mapper: F,
100}
101
102impl<S> ErrorHandler<S> {
103    /// Create a new [`ErrorHandler`] wrapping the given service.
104    pub const fn new(inner: S) -> Self {
105        Self {
106            inner,
107            error_mapper: (),
108        }
109    }
110
111    define_inner_service_accessors!();
112
113    /// Set the error mapper function (not set by default).
114    ///
115    /// The error mapper function is called with the error,
116    /// and should return an [`IntoResponse`] implementation.
117    pub fn error_mapper<F>(self, error_mapper: F) -> ErrorHandler<S, F> {
118        ErrorHandler {
119            inner: self.inner,
120            error_mapper,
121        }
122    }
123}
124
125impl<S, Body> Service<Request<Body>> for ErrorHandler<S, ()>
126where
127    S: Service<Request<Body>, Output: IntoResponse, Error: Into<ErrorResponse>>,
128    Body: Send + 'static,
129{
130    type Output = Response;
131    type Error = Infallible;
132
133    async fn serve(&self, req: Request<Body>) -> Result<Self::Output, Self::Error> {
134        match self.inner.serve(req).await {
135            Ok(response) => Ok(response.into_response()),
136            Err(error) => Ok(error.into().into_response()),
137        }
138    }
139}
140
141impl<S, F, R, Body> Service<Request<Body>> for ErrorHandler<S, F>
142where
143    S: Service<Request<Body>, Output: IntoResponse>,
144    F: Fn(S::Error) -> R + Clone + Send + Sync + 'static,
145    R: IntoResponse + 'static,
146    Body: Send + 'static,
147{
148    type Output = Response;
149    type Error = Infallible;
150
151    async fn serve(&self, req: Request<Body>) -> Result<Self::Output, Self::Error> {
152        match self.inner.serve(req).await {
153            Ok(response) => Ok(response.into_response()),
154            Err(error) => Ok((self.error_mapper)(error).into_response()),
155        }
156    }
157}