Skip to main content

reinhardt_middleware/
redirect_fallback.rs

1//! Redirect fallback middleware
2//!
3//! Provides automatic redirection for 404 errors to a fallback URL.
4//! Useful for handling missing pages gracefully.
5
6use async_trait::async_trait;
7use hyper::StatusCode;
8use regex::Regex;
9use reinhardt_http::{Handler, Middleware, Request, Response, Result};
10use serde::{Deserialize, Serialize};
11use std::sync::Arc;
12
13/// Configuration for redirect fallback behavior
14#[non_exhaustive]
15#[derive(Debug, Clone, Serialize, Deserialize)]
16pub struct RedirectResponseConfig {
17	/// The fallback URL to redirect to on 404 errors
18	pub fallback_url: String,
19	/// Optional path patterns to match (if None, matches all 404s)
20	pub path_patterns: Option<Vec<String>>,
21	/// Status code to use for redirect (default: 302 Found)
22	pub redirect_status: Option<u16>,
23}
24
25impl RedirectResponseConfig {
26	/// Create a new configuration with a fallback URL
27	///
28	/// # Examples
29	///
30	/// ```
31	/// use reinhardt_middleware::RedirectResponseConfig;
32	///
33	/// let config = RedirectResponseConfig::new("/404".to_string());
34	/// assert_eq!(config.fallback_url, "/404");
35	/// ```
36	pub fn new(fallback_url: String) -> Self {
37		Self {
38			fallback_url,
39			path_patterns: None,
40			redirect_status: None,
41		}
42	}
43
44	/// Add path patterns to match
45	///
46	/// # Examples
47	///
48	/// ```
49	/// use reinhardt_middleware::RedirectResponseConfig;
50	///
51	/// let config = RedirectResponseConfig::new("/404".to_string())
52	///     .with_patterns(vec!["/api/.*".to_string()]);
53	/// ```
54	pub fn with_patterns(mut self, patterns: Vec<String>) -> Self {
55		self.path_patterns = Some(patterns);
56		self
57	}
58
59	/// Set custom redirect status code
60	///
61	/// # Examples
62	///
63	/// ```
64	/// use reinhardt_middleware::RedirectResponseConfig;
65	///
66	/// let config = RedirectResponseConfig::new("/404".to_string())
67	///     .with_status(301);
68	/// ```
69	pub fn with_status(mut self, status: u16) -> Self {
70		self.redirect_status = Some(status);
71		self
72	}
73}
74
75/// Middleware that redirects 404 errors to a fallback URL
76///
77/// # Examples
78///
79/// ```
80/// use std::sync::Arc;
81/// use reinhardt_middleware::{RedirectFallbackMiddleware, RedirectResponseConfig};
82/// use reinhardt_http::{Handler, Middleware, Request, Response};
83/// use hyper::{StatusCode, Method, Version, HeaderMap};
84/// use bytes::Bytes;
85///
86/// struct NotFoundHandler;
87///
88/// #[async_trait::async_trait]
89/// impl Handler for NotFoundHandler {
90///     async fn handle(&self, _request: Request) -> reinhardt_core::exception::Result<Response> {
91///         Ok(Response::new(StatusCode::NOT_FOUND))
92///     }
93/// }
94///
95/// # tokio_test::block_on(async {
96/// let config = RedirectResponseConfig::new("/404".to_string());
97/// let middleware = RedirectFallbackMiddleware::new(config);
98/// let handler = Arc::new(NotFoundHandler);
99///
100/// let request = Request::builder()
101///     .method(Method::GET)
102///     .uri("/missing")
103///     .version(Version::HTTP_11)
104///     .headers(HeaderMap::new())
105///     .body(Bytes::new())
106///     .build()
107///     .unwrap();
108///
109/// let response = middleware.process(request, handler).await.unwrap();
110/// assert_eq!(response.status, StatusCode::FOUND);
111/// assert_eq!(
112///     response.headers.get(hyper::header::LOCATION).unwrap(),
113///     "/404"
114/// );
115/// # });
116/// ```
117pub struct RedirectFallbackMiddleware {
118	config: RedirectResponseConfig,
119	compiled_patterns: Option<Vec<Regex>>,
120}
121
122impl RedirectFallbackMiddleware {
123	/// Create a new RedirectFallbackMiddleware with the given configuration
124	///
125	/// # Examples
126	///
127	/// ```
128	/// use reinhardt_middleware::{RedirectFallbackMiddleware, RedirectResponseConfig};
129	///
130	/// let config = RedirectResponseConfig::new("/404".to_string());
131	/// let middleware = RedirectFallbackMiddleware::new(config);
132	/// ```
133	pub fn new(config: RedirectResponseConfig) -> Self {
134		let compiled_patterns = config
135			.path_patterns
136			.as_ref()
137			.map(|patterns| patterns.iter().filter_map(|p| Regex::new(p).ok()).collect());
138
139		Self {
140			config,
141			compiled_patterns,
142		}
143	}
144
145	/// Check if the path matches any configured patterns
146	fn matches_pattern(&self, path: &str) -> bool {
147		match &self.compiled_patterns {
148			None => true, // No patterns means match all
149			Some(patterns) => patterns.iter().any(|re| re.is_match(path)),
150		}
151	}
152
153	/// Get the redirect status code to use
154	fn redirect_status(&self) -> StatusCode {
155		self.config
156			.redirect_status
157			.and_then(|code| StatusCode::from_u16(code).ok())
158			.unwrap_or(StatusCode::FOUND)
159	}
160
161	/// Check if we should redirect to avoid loops
162	fn should_redirect(&self, path: &str) -> bool {
163		// Prevent redirect loop: don't redirect if already at fallback URL
164		path != self.config.fallback_url
165	}
166}
167
168#[async_trait]
169impl Middleware for RedirectFallbackMiddleware {
170	async fn process(&self, request: Request, handler: Arc<dyn Handler>) -> Result<Response> {
171		let path = request.uri.path().to_string();
172
173		// Convert errors to responses so post-processing always runs,
174		// even when invoked outside MiddlewareChain. (#3244)
175		let response = match handler.handle(request).await {
176			Ok(resp) => resp,
177			Err(e) => Response::from(e),
178		};
179
180		// Only redirect on 404 errors
181		if response.status != StatusCode::NOT_FOUND {
182			return Ok(response);
183		}
184
185		// Check if we should redirect (pattern match and loop prevention)
186		if !self.matches_pattern(&path) || !self.should_redirect(&path) {
187			return Ok(response);
188		}
189
190		// Create redirect response
191		let mut redirect_response = Response::new(self.redirect_status());
192		redirect_response.headers.insert(
193			hyper::header::LOCATION,
194			self.config
195				.fallback_url
196				.parse()
197				.unwrap_or_else(|_| hyper::header::HeaderValue::from_static("/")),
198		);
199
200		Ok(redirect_response)
201	}
202}
203
204#[cfg(test)]
205mod tests {
206	use super::*;
207	use bytes::Bytes;
208	use hyper::{HeaderMap, Method, StatusCode, Version};
209
210	struct NotFoundHandler;
211
212	#[async_trait]
213	impl Handler for NotFoundHandler {
214		async fn handle(&self, _request: Request) -> Result<Response> {
215			Ok(Response::new(StatusCode::NOT_FOUND))
216		}
217	}
218
219	struct OkHandler;
220
221	#[async_trait]
222	impl Handler for OkHandler {
223		async fn handle(&self, _request: Request) -> Result<Response> {
224			Ok(Response::new(StatusCode::OK).with_body(Bytes::from("OK")))
225		}
226	}
227
228	#[tokio::test]
229	async fn test_redirect_on_404() {
230		let config = RedirectResponseConfig::new("/404".to_string());
231		let middleware = RedirectFallbackMiddleware::new(config);
232		let handler = Arc::new(NotFoundHandler);
233
234		let request = Request::builder()
235			.method(Method::GET)
236			.uri("/missing")
237			.version(Version::HTTP_11)
238			.headers(HeaderMap::new())
239			.body(Bytes::new())
240			.build()
241			.unwrap();
242
243		let response = middleware.process(request, handler).await.unwrap();
244
245		assert_eq!(response.status, StatusCode::FOUND);
246		assert_eq!(
247			response.headers.get(hyper::header::LOCATION).unwrap(),
248			"/404"
249		);
250	}
251
252	#[tokio::test]
253	async fn test_no_redirect_on_200() {
254		let config = RedirectResponseConfig::new("/404".to_string());
255		let middleware = RedirectFallbackMiddleware::new(config);
256		let handler = Arc::new(OkHandler);
257
258		let request = Request::builder()
259			.method(Method::GET)
260			.uri("/existing")
261			.version(Version::HTTP_11)
262			.headers(HeaderMap::new())
263			.body(Bytes::new())
264			.build()
265			.unwrap();
266
267		let response = middleware.process(request, handler).await.unwrap();
268
269		assert_eq!(response.status, StatusCode::OK);
270		assert!(!response.headers.contains_key(hyper::header::LOCATION));
271	}
272
273	#[tokio::test]
274	async fn test_pattern_matching_redirect() {
275		let config = RedirectResponseConfig::new("/404".to_string())
276			.with_patterns(vec!["/api/.*".to_string()]);
277		let middleware = RedirectFallbackMiddleware::new(config);
278		let handler = Arc::new(NotFoundHandler);
279
280		// Should redirect for /api/* paths
281		let request = Request::builder()
282			.method(Method::GET)
283			.uri("/api/missing")
284			.version(Version::HTTP_11)
285			.headers(HeaderMap::new())
286			.body(Bytes::new())
287			.build()
288			.unwrap();
289
290		let response = middleware.process(request, handler).await.unwrap();
291
292		assert_eq!(response.status, StatusCode::FOUND);
293		assert_eq!(
294			response.headers.get(hyper::header::LOCATION).unwrap(),
295			"/404"
296		);
297	}
298
299	#[tokio::test]
300	async fn test_pattern_no_match_no_redirect() {
301		let config = RedirectResponseConfig::new("/404".to_string())
302			.with_patterns(vec!["/api/.*".to_string()]);
303		let middleware = RedirectFallbackMiddleware::new(config);
304		let handler = Arc::new(NotFoundHandler);
305
306		// Should NOT redirect for non-/api/* paths
307		let request = Request::builder()
308			.method(Method::GET)
309			.uri("/other/missing")
310			.version(Version::HTTP_11)
311			.headers(HeaderMap::new())
312			.body(Bytes::new())
313			.build()
314			.unwrap();
315
316		let response = middleware.process(request, handler).await.unwrap();
317
318		assert_eq!(response.status, StatusCode::NOT_FOUND);
319		assert!(!response.headers.contains_key(hyper::header::LOCATION));
320	}
321
322	#[tokio::test]
323	async fn test_custom_redirect_status() {
324		let config = RedirectResponseConfig::new("/404".to_string()).with_status(301);
325		let middleware = RedirectFallbackMiddleware::new(config);
326		let handler = Arc::new(NotFoundHandler);
327
328		let request = Request::builder()
329			.method(Method::GET)
330			.uri("/missing")
331			.version(Version::HTTP_11)
332			.headers(HeaderMap::new())
333			.body(Bytes::new())
334			.build()
335			.unwrap();
336
337		let response = middleware.process(request, handler).await.unwrap();
338
339		assert_eq!(response.status, StatusCode::MOVED_PERMANENTLY);
340		assert_eq!(
341			response.headers.get(hyper::header::LOCATION).unwrap(),
342			"/404"
343		);
344	}
345
346	#[tokio::test]
347	async fn test_prevent_redirect_loop() {
348		let config = RedirectResponseConfig::new("/404".to_string());
349		let middleware = RedirectFallbackMiddleware::new(config);
350		let handler = Arc::new(NotFoundHandler);
351
352		// Request to the fallback URL itself should not redirect
353		let request = Request::builder()
354			.method(Method::GET)
355			.uri("/404")
356			.version(Version::HTTP_11)
357			.headers(HeaderMap::new())
358			.body(Bytes::new())
359			.build()
360			.unwrap();
361
362		let response = middleware.process(request, handler).await.unwrap();
363
364		assert_eq!(response.status, StatusCode::NOT_FOUND);
365		assert!(!response.headers.contains_key(hyper::header::LOCATION));
366	}
367
368	#[tokio::test]
369	async fn test_multiple_pattern_matching() {
370		let config = RedirectResponseConfig::new("/error".to_string())
371			.with_patterns(vec!["/api/.*".to_string(), "/v1/.*".to_string()]);
372		let middleware = RedirectFallbackMiddleware::new(config);
373		let handler = Arc::new(NotFoundHandler);
374
375		// Test first pattern
376		let request1 = Request::builder()
377			.method(Method::GET)
378			.uri("/api/test")
379			.version(Version::HTTP_11)
380			.headers(HeaderMap::new())
381			.body(Bytes::new())
382			.build()
383			.unwrap();
384
385		let response1 = middleware.process(request1, handler.clone()).await.unwrap();
386		assert_eq!(response1.status, StatusCode::FOUND);
387
388		// Test second pattern
389		let request2 = Request::builder()
390			.method(Method::GET)
391			.uri("/v1/test")
392			.version(Version::HTTP_11)
393			.headers(HeaderMap::new())
394			.body(Bytes::new())
395			.build()
396			.unwrap();
397
398		let response2 = middleware.process(request2, handler).await.unwrap();
399		assert_eq!(response2.status, StatusCode::FOUND);
400	}
401
402	#[tokio::test]
403	async fn test_different_http_methods() {
404		let config = RedirectResponseConfig::new("/404".to_string());
405		let middleware = RedirectFallbackMiddleware::new(config);
406		let handler = Arc::new(NotFoundHandler);
407
408		// Test POST
409		let request = Request::builder()
410			.method(Method::POST)
411			.uri("/missing")
412			.version(Version::HTTP_11)
413			.headers(HeaderMap::new())
414			.body(Bytes::new())
415			.build()
416			.unwrap();
417
418		let response = middleware.process(request, handler).await.unwrap();
419
420		assert_eq!(response.status, StatusCode::FOUND);
421		assert_eq!(
422			response.headers.get(hyper::header::LOCATION).unwrap(),
423			"/404"
424		);
425	}
426
427	#[tokio::test]
428	async fn test_no_patterns_matches_all() {
429		let config = RedirectResponseConfig::new("/fallback".to_string());
430		let middleware = RedirectFallbackMiddleware::new(config);
431		let handler = Arc::new(NotFoundHandler);
432
433		// Any path should redirect when no patterns are specified
434		let paths = vec!["/api/test", "/admin/test", "/any/path/here"];
435
436		for path in paths {
437			let request = Request::builder()
438				.method(Method::GET)
439				.uri(path)
440				.version(Version::HTTP_11)
441				.headers(HeaderMap::new())
442				.body(Bytes::new())
443				.build()
444				.unwrap();
445
446			let response = middleware.process(request, handler.clone()).await.unwrap();
447
448			assert_eq!(response.status, StatusCode::FOUND);
449			assert_eq!(
450				response.headers.get(hyper::header::LOCATION).unwrap(),
451				"/fallback"
452			);
453		}
454	}
455
456	#[tokio::test]
457	async fn test_complex_pattern_matching() {
458		let config = RedirectResponseConfig::new("/404".to_string())
459			.with_patterns(vec!["/api/v[0-9]+/.*".to_string()]);
460		let middleware = RedirectFallbackMiddleware::new(config);
461		let handler = Arc::new(NotFoundHandler);
462
463		// Should match /api/v1/, /api/v2/, etc.
464		let request1 = Request::builder()
465			.method(Method::GET)
466			.uri("/api/v1/users")
467			.version(Version::HTTP_11)
468			.headers(HeaderMap::new())
469			.body(Bytes::new())
470			.build()
471			.unwrap();
472
473		let response1 = middleware.process(request1, handler.clone()).await.unwrap();
474		assert_eq!(response1.status, StatusCode::FOUND);
475
476		// Should NOT match /api/version/
477		let request2 = Request::builder()
478			.method(Method::GET)
479			.uri("/api/version/users")
480			.version(Version::HTTP_11)
481			.headers(HeaderMap::new())
482			.body(Bytes::new())
483			.build()
484			.unwrap();
485
486		let response2 = middleware.process(request2, handler).await.unwrap();
487		assert_eq!(response2.status, StatusCode::NOT_FOUND);
488	}
489}