reinhardt_middleware/
redirect_fallback.rs1use 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#[non_exhaustive]
15#[derive(Debug, Clone, Serialize, Deserialize)]
16pub struct RedirectResponseConfig {
17 pub fallback_url: String,
19 pub path_patterns: Option<Vec<String>>,
21 pub redirect_status: Option<u16>,
23}
24
25impl RedirectResponseConfig {
26 pub fn new(fallback_url: String) -> Self {
37 Self {
38 fallback_url,
39 path_patterns: None,
40 redirect_status: None,
41 }
42 }
43
44 pub fn with_patterns(mut self, patterns: Vec<String>) -> Self {
55 self.path_patterns = Some(patterns);
56 self
57 }
58
59 pub fn with_status(mut self, status: u16) -> Self {
70 self.redirect_status = Some(status);
71 self
72 }
73}
74
75pub struct RedirectFallbackMiddleware {
118 config: RedirectResponseConfig,
119 compiled_patterns: Option<Vec<Regex>>,
120}
121
122impl RedirectFallbackMiddleware {
123 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 fn matches_pattern(&self, path: &str) -> bool {
147 match &self.compiled_patterns {
148 None => true, Some(patterns) => patterns.iter().any(|re| re.is_match(path)),
150 }
151 }
152
153 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 fn should_redirect(&self, path: &str) -> bool {
163 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 let response = match handler.handle(request).await {
176 Ok(resp) => resp,
177 Err(e) => Response::from(e),
178 };
179
180 if response.status != StatusCode::NOT_FOUND {
182 return Ok(response);
183 }
184
185 if !self.matches_pattern(&path) || !self.should_redirect(&path) {
187 return Ok(response);
188 }
189
190 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 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 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 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 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 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 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 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 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 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}