reinhardt_http/
exception.rs1use async_trait::async_trait;
13use std::sync::Arc;
14
15use crate::{Error, Handler, Request, Response, Result};
16
17#[async_trait]
58pub trait ExceptionHandler: Send + Sync + 'static {
59 async fn handle_exception(&self, request: &Request, error: Error) -> Response;
67}
68
69#[doc(hidden)]
71#[derive(Debug, Clone, Copy)]
72pub struct ExceptionHandlerInvoked;
73
74pub struct ExceptionHandlingHandler {
135 inner: Arc<dyn Handler>,
136 exception_handler: Arc<dyn ExceptionHandler>,
137}
138
139impl ExceptionHandlingHandler {
140 pub fn new(inner: Arc<dyn Handler>, exception_handler: Arc<dyn ExceptionHandler>) -> Self {
142 Self {
143 inner,
144 exception_handler,
145 }
146 }
147}
148
149#[async_trait]
150impl Handler for ExceptionHandlingHandler {
151 async fn handle(&self, mut request: Request) -> Result<Response> {
152 request.install_exception_handler(Arc::clone(&self.exception_handler));
159 let mut context = request.clone_for_di();
160 match self.inner.handle(request).await {
161 Ok(response) => Ok(response),
162 Err(error) => {
163 context.sync_path_params_from_shared_state();
164 context.extensions.insert(ExceptionHandlerInvoked);
165 Ok(self
166 .exception_handler
167 .handle_exception(&context, error)
168 .await)
169 }
170 }
171 }
172}
173
174#[cfg(test)]
175mod tests {
176 use super::*;
177 use bytes::Bytes;
178 use hyper::{HeaderMap, Method, StatusCode, Version};
179 use rstest::rstest;
180 use std::sync::Mutex;
181
182 #[derive(Debug, PartialEq, Eq)]
184 struct MarkerContext(&'static str);
185
186 #[derive(Debug, PartialEq, Eq)]
188 struct ObservedRequest {
189 method: Method,
190 path: String,
191 request_id: Option<String>,
192 item_id: Option<String>,
193 di_marker: Option<&'static str>,
194 }
195
196 struct ObservingHandler {
198 observed: Arc<Mutex<Vec<ObservedRequest>>>,
199 }
200
201 #[async_trait]
202 impl ExceptionHandler for ObservingHandler {
203 async fn handle_exception(&self, request: &Request, error: Error) -> Response {
204 let di_marker = request
205 .get_di_context::<MarkerContext>()
206 .map(|marker| marker.0);
207 self.observed.lock().unwrap().push(ObservedRequest {
208 method: request.method.clone(),
209 path: request.uri.path().to_string(),
210 request_id: request.get_header("x-request-id"),
211 item_id: request.path_params.get("id").cloned(),
212 di_marker,
213 });
214
215 Response::new(StatusCode::IM_A_TEAPOT).with_body(error.to_string())
217 }
218 }
219
220 struct FailingHandler {
224 factory: fn() -> Error,
225 }
226
227 #[async_trait]
228 impl Handler for FailingHandler {
229 async fn handle(&self, _request: Request) -> Result<Response> {
230 Err((self.factory)())
231 }
232 }
233
234 struct OkHandler;
236
237 #[async_trait]
238 impl Handler for OkHandler {
239 async fn handle(&self, _request: Request) -> Result<Response> {
240 Ok(Response::ok().with_body("ok"))
241 }
242 }
243
244 struct CountingHandler {
246 calls: Arc<Mutex<usize>>,
247 }
248
249 #[async_trait]
250 impl Handler for CountingHandler {
251 async fn handle(&self, _request: Request) -> Result<Response> {
252 *self.calls.lock().unwrap() += 1;
253 Err(Error::Internal("boom".to_string()))
254 }
255 }
256
257 fn build_request(method: Method, uri: &str) -> Request {
258 Request::builder()
259 .method(method)
260 .uri(uri)
261 .version(Version::HTTP_11)
262 .headers(HeaderMap::new())
263 .body(Bytes::new())
264 .build()
265 .unwrap()
266 }
267
268 #[rstest]
269 #[tokio::test]
270 async fn test_inner_error_is_converted_by_installed_handler() {
271 let observed = Arc::new(Mutex::new(Vec::new()));
273 let handler = ExceptionHandlingHandler::new(
274 Arc::new(FailingHandler {
275 factory: || Error::NotFound("no route".to_string()),
276 }),
277 Arc::new(ObservingHandler {
278 observed: Arc::clone(&observed),
279 }),
280 );
281 let request = build_request(Method::GET, "/missing");
282
283 let response = handler.handle(request).await.unwrap();
285
286 assert_eq!(response.status, StatusCode::IM_A_TEAPOT);
288 assert_eq!(observed.lock().unwrap().len(), 1);
289 }
290
291 #[rstest]
292 #[tokio::test]
293 async fn test_handler_receives_original_request_context() {
294 let observed = Arc::new(Mutex::new(Vec::new()));
296 let handler = ExceptionHandlingHandler::new(
297 Arc::new(FailingHandler {
298 factory: || Error::Internal("boom".to_string()),
299 }),
300 Arc::new(ObservingHandler {
301 observed: Arc::clone(&observed),
302 }),
303 );
304
305 let mut headers = HeaderMap::new();
306 headers.insert("x-request-id", "req-42".parse().unwrap());
307 let mut request = Request::builder()
308 .method(Method::POST)
309 .uri("/api/items/7")
310 .version(Version::HTTP_11)
311 .headers(headers)
312 .body(Bytes::from_static(b"payload"))
313 .build()
314 .unwrap();
315 request.path_params.insert("id", "7");
316
317 handler.handle(request).await.unwrap();
319
320 let observed = observed.lock().unwrap();
322 assert_eq!(observed.len(), 1);
323 assert_eq!(observed[0].method, Method::POST);
324 assert_eq!(observed[0].path, "/api/items/7");
325 assert_eq!(observed[0].request_id, Some("req-42".to_string()));
326 assert_eq!(observed[0].item_id, Some("7".to_string()));
327 }
328
329 #[rstest]
330 #[tokio::test]
331 async fn test_handler_shares_request_extensions() {
332 let observed = Arc::new(Mutex::new(Vec::new()));
334 let handler = ExceptionHandlingHandler::new(
335 Arc::new(FailingHandler {
336 factory: || Error::Internal("boom".to_string()),
337 }),
338 Arc::new(ObservingHandler {
339 observed: Arc::clone(&observed),
340 }),
341 );
342 let mut request = build_request(Method::GET, "/");
343 request.set_di_context(MarkerContext("di-visible"));
344
345 handler.handle(request).await.unwrap();
347
348 let observed = observed.lock().unwrap();
350 assert_eq!(observed.len(), 1);
351 assert_eq!(observed[0].di_marker, Some("di-visible"));
352 }
353
354 #[rstest]
355 #[tokio::test]
356 async fn test_handler_is_not_invoked_for_successful_response() {
357 let observed = Arc::new(Mutex::new(Vec::new()));
359 let handler = ExceptionHandlingHandler::new(
360 Arc::new(OkHandler),
361 Arc::new(ObservingHandler {
362 observed: Arc::clone(&observed),
363 }),
364 );
365
366 let response = handler
368 .handle(build_request(Method::GET, "/"))
369 .await
370 .unwrap();
371
372 assert_eq!(response.status, StatusCode::OK);
374 assert_eq!(String::from_utf8(response.body.to_vec()).unwrap(), "ok");
375 assert!(observed.lock().unwrap().is_empty());
376 }
377
378 #[rstest]
379 #[tokio::test]
380 async fn test_inner_handler_runs_exactly_once() {
381 let calls = Arc::new(Mutex::new(0_usize));
383 let observed = Arc::new(Mutex::new(Vec::new()));
384 let handler = ExceptionHandlingHandler::new(
385 Arc::new(CountingHandler {
386 calls: Arc::clone(&calls),
387 }),
388 Arc::new(ObservingHandler {
389 observed: Arc::clone(&observed),
390 }),
391 );
392
393 handler
395 .handle(build_request(Method::GET, "/"))
396 .await
397 .unwrap();
398
399 assert_eq!(*calls.lock().unwrap(), 1);
401 assert_eq!(observed.lock().unwrap().len(), 1);
402 }
403}