mod fixtures;
use fixtures::*;
use rstest::rstest;
use serial_test::serial;
use std::sync::Arc;
use reinhardt_http::Middleware;
use reinhardt_http::Request;
use reinhardt_middleware::csrf::CsrfMiddleware;
use reinhardt_middleware::locale::LocaleMiddleware;
#[cfg(feature = "rate-limit")]
use reinhardt_middleware::rate_limit::{RateLimitConfig, RateLimitMiddleware, RateLimitStrategy};
use reinhardt_middleware::circuit_breaker::{
CircuitBreakerConfig, CircuitBreakerMiddleware, CircuitState,
};
use bytes::Bytes;
use hyper::Method;
use std::time::Duration;
#[rstest]
#[tokio::test]
#[serial(csrf)]
async fn test_csrf_post_without_token_returns_error() {
let middleware = Arc::new(CsrfMiddleware::new());
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request = create_test_request("POST", "/form/submit");
let result = middleware.process(request, handler.clone()).await;
assert!(result.is_err(), "CSRF validation should fail without token");
assert_eq!(
handler.count(),
0,
"Handler should not be called for CSRF error"
);
}
#[rstest]
#[tokio::test]
#[serial(csrf)]
async fn test_csrf_post_with_invalid_token_returns_error() {
let middleware = Arc::new(CsrfMiddleware::new());
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request = Request::builder()
.method(Method::POST)
.uri("/form/submit")
.header("X-CSRF-Token", "invalid_token_12345")
.header("Cookie", "csrftoken=different_token_67890")
.header("Referer", "https://example.com/")
.body(Bytes::new())
.build()
.unwrap();
let result = middleware.process(request, handler.clone()).await;
assert!(
result.is_err(),
"CSRF validation should fail with invalid token"
);
assert_eq!(
handler.count(),
0,
"Handler should not be called for CSRF error"
);
}
#[rstest]
#[tokio::test]
#[serial(csrf)]
async fn test_csrf_post_with_mismatched_token_returns_error() {
let middleware = Arc::new(CsrfMiddleware::new());
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request = Request::builder()
.method(Method::POST)
.uri("/api/data")
.header("X-CSRF-Token", "token_a")
.header("Cookie", "csrftoken=token_b")
.header("Referer", "https://example.com/")
.body(Bytes::new())
.build()
.unwrap();
let result = middleware.process(request, handler.clone()).await;
assert!(
result.is_err(),
"CSRF validation should fail with mismatched tokens"
);
}
#[cfg(feature = "rate-limit")]
#[rstest]
#[tokio::test]
#[serial(rate_limit)]
async fn test_rate_limit_exceeded_returns_429() {
let config =
RateLimitConfig::new(RateLimitStrategy::PerIp, 1.0, 0.001).with_cost_per_request(1.0);
let middleware = Arc::new(RateLimitMiddleware::new(config));
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request1 = create_test_request("GET", "/api/data");
let response1 = middleware.process(request1, handler.clone()).await.unwrap();
assert_status(&response1, 200);
let request2 = create_test_request("GET", "/api/data");
let response2 = middleware.process(request2, handler.clone()).await.unwrap();
assert_status(&response2, 429);
}
#[cfg(feature = "rate-limit")]
#[rstest]
#[tokio::test]
#[serial(rate_limit)]
async fn test_rate_limit_adds_retry_after_header() {
let config =
RateLimitConfig::new(RateLimitStrategy::PerIp, 1.0, 1.0).with_cost_per_request(1.0);
let middleware = Arc::new(RateLimitMiddleware::new(config));
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request1 = create_test_request("GET", "/api/data");
let _ = middleware.process(request1, handler.clone()).await.unwrap();
let request2 = create_test_request("GET", "/api/data");
let response2 = middleware.process(request2, handler.clone()).await.unwrap();
assert_status(&response2, 429);
assert!(
response2.headers.get("retry-after").is_some(),
"Response should have Retry-After header"
);
}
#[cfg(feature = "rate-limit")]
#[rstest]
#[tokio::test]
#[serial(rate_limit)]
async fn test_rate_limit_per_user_isolation() {
let config =
RateLimitConfig::new(RateLimitStrategy::PerUser, 1.0, 0.001).with_cost_per_request(1.0);
let middleware = Arc::new(RateLimitMiddleware::new(config));
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request1 = Request::builder()
.method(Method::GET)
.uri("/api/data")
.body(Bytes::new())
.build()
.unwrap();
request1.extensions.insert("user_1".to_string());
let _ = middleware.process(request1, handler.clone()).await.unwrap();
let request2 = Request::builder()
.method(Method::GET)
.uri("/api/data")
.body(Bytes::new())
.build()
.unwrap();
request2.extensions.insert("user_1".to_string());
let response2 = middleware.process(request2, handler.clone()).await.unwrap();
assert_status(&response2, 429);
let request3 = Request::builder()
.method(Method::GET)
.uri("/api/data")
.body(Bytes::new())
.build()
.unwrap();
request3.extensions.insert("user_2".to_string());
let response3 = middleware.process(request3, handler.clone()).await.unwrap();
assert_status(&response3, 200);
}
#[rstest]
#[tokio::test]
#[serial(circuit_breaker)]
async fn test_circuit_breaker_open_state_returns_503() {
let config = CircuitBreakerConfig::new(0.5, 2, Duration::from_secs(60))
.with_half_open_success_threshold(2);
let middleware = Arc::new(CircuitBreakerMiddleware::new(config));
let failure_handler = Arc::new(ConfigurableTestHandler::always_failure());
for _ in 0..3 {
let request = create_test_request("GET", "/api/data");
let _ = middleware.process(request, failure_handler.clone()).await;
}
assert_eq!(
middleware.state(),
CircuitState::Open,
"Circuit should be open after failures"
);
let success_handler = Arc::new(ConfigurableTestHandler::always_success());
let request = create_test_request("GET", "/api/data");
let response = middleware
.process(request, success_handler.clone())
.await
.unwrap();
assert_status(&response, 503);
assert_eq!(
success_handler.count(),
0,
"Handler should not be called when circuit is open"
);
}
#[rstest]
#[tokio::test]
#[serial(circuit_breaker)]
async fn test_circuit_breaker_counts_failures() {
let config = CircuitBreakerConfig::new(0.5, 5, Duration::from_secs(60));
let middleware = Arc::new(CircuitBreakerMiddleware::new(config));
let failure_handler = Arc::new(ConfigurableTestHandler::always_failure());
for _ in 0..2 {
let request = create_test_request("GET", "/api/data");
let _ = middleware.process(request, failure_handler.clone()).await;
}
assert_eq!(
middleware.state(),
CircuitState::Closed,
"Circuit should still be closed (min_requests not met)"
);
for _ in 0..3 {
let request = create_test_request("GET", "/api/data");
let _ = middleware.process(request, failure_handler.clone()).await;
}
assert_eq!(
middleware.state(),
CircuitState::Open,
"Circuit should be open after exceeding threshold"
);
}
#[rstest]
#[tokio::test]
#[serial(locale)]
async fn test_locale_invalid_accept_language_fallback() {
let middleware = Arc::new(LocaleMiddleware::new());
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request = Request::builder()
.method(Method::GET)
.uri("/page")
.header("Accept-Language", "invalid-locale-format")
.body(Bytes::new())
.build()
.unwrap();
let response = middleware.process(request, handler.clone()).await.unwrap();
assert_status(&response, 200);
}
#[rstest]
#[tokio::test]
#[serial(locale)]
async fn test_locale_empty_accept_language_uses_default() {
let middleware = Arc::new(LocaleMiddleware::new());
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request = create_test_request("GET", "/page");
let response = middleware.process(request, handler.clone()).await.unwrap();
assert_status(&response, 200);
}
#[rstest]
#[tokio::test]
#[serial(locale)]
async fn test_locale_unsupported_locale_fallback() {
let middleware = Arc::new(LocaleMiddleware::new());
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request = Request::builder()
.method(Method::GET)
.uri("/page")
.header("Accept-Language", "x-klingon")
.body(Bytes::new())
.build()
.unwrap();
let response = middleware.process(request, handler.clone()).await.unwrap();
assert_status(&response, 200);
}
#[rstest]
#[tokio::test]
#[serial(error_propagation)]
async fn test_middleware_propagates_handler_error_status() {
let middleware = Arc::new(LocaleMiddleware::new());
let handler = Arc::new(ConfigurableTestHandler::always_failure());
let request = create_test_request("GET", "/api/data");
let response = middleware.process(request, handler).await.unwrap();
assert_status(&response, 500);
}
#[rstest]
#[tokio::test]
#[serial(error_chain)]
async fn test_middleware_chain_error_handling() {
use reinhardt_middleware::csp::CspMiddleware;
let inner_middleware = Arc::new(CspMiddleware::new());
let outer_middleware = Arc::new(LocaleMiddleware::new());
let failure_handler = Arc::new(ConfigurableTestHandler::always_failure());
struct MiddlewareHandler {
middleware: Arc<dyn Middleware>,
handler: Arc<dyn reinhardt_http::Handler>,
}
#[async_trait::async_trait]
impl reinhardt_http::Handler for MiddlewareHandler {
async fn handle(
&self,
request: Request,
) -> reinhardt_core::exception::Result<reinhardt_http::Response> {
self.middleware.process(request, self.handler.clone()).await
}
}
let composite_handler = Arc::new(MiddlewareHandler {
middleware: inner_middleware,
handler: failure_handler,
});
let request = create_test_request("GET", "/api/data");
let response = outer_middleware
.process(request, composite_handler)
.await
.unwrap();
assert_status(&response, 500);
}
#[rstest]
#[tokio::test]
#[serial(cache)]
async fn test_cache_does_not_cache_error_responses(cache_middleware: Arc<CacheMiddleware>) {
let failure_handler = Arc::new(ConfigurableTestHandler::always_failure());
let request1 = create_test_request("GET", "/api/error");
let response1 = cache_middleware
.process(request1, failure_handler.clone())
.await
.unwrap();
assert_status(&response1, 500);
let success_handler = Arc::new(ConfigurableTestHandler::always_success());
let request2 = create_test_request("GET", "/api/error");
let response2 = cache_middleware
.process(request2, success_handler.clone())
.await
.unwrap();
assert_status(&response2, 200);
assert_eq!(
success_handler.count(),
1,
"Handler should be called (error not cached)"
);
}
#[rstest]
#[tokio::test]
#[serial(cache)]
async fn test_cache_does_not_cache_post_requests(cache_middleware: Arc<CacheMiddleware>) {
let handler = Arc::new(ConfigurableTestHandler::always_success());
let request1 = Request::builder()
.method(Method::POST)
.uri("/api/data")
.body(Bytes::new())
.build()
.unwrap();
let _ = cache_middleware
.process(request1, handler.clone())
.await
.unwrap();
let request2 = Request::builder()
.method(Method::POST)
.uri("/api/data")
.body(Bytes::new())
.build()
.unwrap();
let _ = cache_middleware
.process(request2, handler.clone())
.await
.unwrap();
assert_eq!(
handler.count(),
2,
"POST requests should not be cached - handler should be called twice"
);
}
use reinhardt_middleware::cache::CacheMiddleware;
struct ErrorReturningHandler {
error_variant: ErrorVariant,
}
enum ErrorVariant {
NotFound,
Unauthorized,
BadRequest,
}
#[async_trait::async_trait]
impl reinhardt_http::Handler for ErrorReturningHandler {
async fn handle(
&self,
_request: Request,
) -> reinhardt_core::exception::Result<reinhardt_http::Response> {
match self.error_variant {
ErrorVariant::NotFound => Err(reinhardt_core::exception::Error::NotFound(
"not found".into(),
)),
ErrorVariant::Unauthorized => Err(reinhardt_core::exception::Error::Authentication(
"unauthorized".into(),
)),
ErrorVariant::BadRequest => {
Err(reinhardt_core::exception::Error::Http("bad request".into()))
}
}
}
}
#[rstest]
#[tokio::test]
#[serial(security_error_chain)]
async fn test_security_headers_present_on_error_response() {
use reinhardt_http::middleware::MiddlewareChain;
use reinhardt_middleware::security_middleware::SecurityMiddleware;
let handler = Arc::new(ErrorReturningHandler {
error_variant: ErrorVariant::Unauthorized,
});
let mut chain = MiddlewareChain::new(handler);
chain.add_middleware(Arc::new(SecurityMiddleware::new()));
let request = create_test_request("GET", "/api/protected");
let response = reinhardt_http::Handler::handle(&chain, request)
.await
.unwrap();
assert_eq!(response.status, hyper::StatusCode::UNAUTHORIZED);
assert!(
response.headers.contains_key("X-Content-Type-Options"),
"Error response should have X-Content-Type-Options header"
);
}
#[rstest]
#[tokio::test]
#[serial(xframe_error_chain)]
async fn test_xframe_header_present_on_error_response() {
use reinhardt_http::middleware::MiddlewareChain;
use reinhardt_middleware::xframe::{XFrameOptions, XFrameOptionsMiddleware};
let handler = Arc::new(ErrorReturningHandler {
error_variant: ErrorVariant::NotFound,
});
let mut chain = MiddlewareChain::new(handler);
chain.add_middleware(Arc::new(XFrameOptionsMiddleware::new(XFrameOptions::Deny)));
let request = create_test_request("GET", "/api/missing");
let response = reinhardt_http::Handler::handle(&chain, request)
.await
.unwrap();
assert_eq!(response.status, hyper::StatusCode::NOT_FOUND);
assert_eq!(
response
.headers
.get("X-Frame-Options")
.map(|v| v.to_str().unwrap()),
Some("DENY"),
"Error response should have X-Frame-Options: DENY header"
);
}
#[rstest]
#[tokio::test]
#[serial(multi_middleware_error_chain)]
async fn test_multiple_middleware_headers_on_error_response() {
use reinhardt_http::middleware::MiddlewareChain;
use reinhardt_middleware::security_middleware::SecurityMiddleware;
use reinhardt_middleware::xframe::{XFrameOptions, XFrameOptionsMiddleware};
let handler = Arc::new(ErrorReturningHandler {
error_variant: ErrorVariant::BadRequest,
});
let mut chain = MiddlewareChain::new(handler);
chain.add_middleware(Arc::new(SecurityMiddleware::new()));
chain.add_middleware(Arc::new(XFrameOptionsMiddleware::new(XFrameOptions::Deny)));
let request = create_test_request("GET", "/api/invalid");
let response = reinhardt_http::Handler::handle(&chain, request)
.await
.unwrap();
assert_eq!(response.status, hyper::StatusCode::BAD_REQUEST);
assert!(
response.headers.contains_key("X-Content-Type-Options"),
"Error response should have X-Content-Type-Options from SecurityMiddleware"
);
assert_eq!(
response
.headers
.get("X-Frame-Options")
.map(|v| v.to_str().unwrap()),
Some("DENY"),
"Error response should have X-Frame-Options from XFrameOptionsMiddleware"
);
}