From 818a29ddcd7d8970c718ec2fe8f8228c502b7f2b Mon Sep 17 00:00:00 2001 From: popsman Date: Mon, 28 Sep 2026 14:08:15 +0000 Subject: [PATCH 1/2] fix: #1502 Add unit tests for CSRF protection middleware Closes #1502 --- services/api/src/csrf.rs | 153 +++++++++++++++++++++++++++++++ services/api/tests/csrf_tests.rs | 88 ++++++++++++++++++ 2 files changed, 241 insertions(+) create mode 100644 services/api/tests/csrf_tests.rs diff --git a/services/api/src/csrf.rs b/services/api/src/csrf.rs index 67dff794..415990d8 100644 --- a/services/api/src/csrf.rs +++ b/services/api/src/csrf.rs @@ -157,3 +157,156 @@ pub async fn csrf_protection_middleware( // Non-browser / API client (no Origin, no Cookie) — pass through. Ok(next.run(request).await) } + +#[cfg(test)] +mod tests { + use super::*; + use axum::body::Body; + use axum::http::Request as HttpRequest; + use axum::routing::post; + use axum::Router; + use tower::ServiceExt; + + fn test_config() -> Arc { + Arc::new(CsrfConfig { + allowed_origins: vec![ + "https://app.predictiq.com".to_string(), + "https://staging.predictiq.com".to_string(), + ], + }) + } + + /// Build a router with the CSRF middleware applied to a state-changing route. + fn app() -> Router { + Router::new() + .route("/mutate", post(|| async { StatusCode::OK })) + .route("/read", axum::routing::get(|| async { StatusCode::OK })) + .layer(axum::middleware::from_fn_with_state( + test_config(), + csrf_protection_middleware, + )) + } + + async fn send(method: &str, path: &str, headers: &[(&str, &str)]) -> StatusCode { + let mut builder = HttpRequest::builder().method(method).uri(path); + for (name, value) in headers { + builder = builder.header(*name, *value); + } + let request = builder.body(Body::empty()).unwrap(); + app().oneshot(request).await.unwrap().status() + } + + #[tokio::test] + async fn matching_allowed_origin_passes() { + let status = send( + "POST", + "/mutate", + &[("origin", "https://app.predictiq.com")], + ) + .await; + assert_eq!(status, StatusCode::OK); + } + + #[tokio::test] + async fn matching_allowed_origin_is_case_insensitive() { + let status = send( + "POST", + "/mutate", + &[("origin", "https://APP.PredictIQ.com")], + ) + .await; + assert_eq!(status, StatusCode::OK); + } + + #[tokio::test] + async fn mismatched_origin_is_rejected() { + let status = send( + "POST", + "/mutate", + &[("origin", "https://evil.example.com")], + ) + .await; + assert_eq!(status, StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn origin_prefix_lookalike_is_rejected() { + // A host that merely starts with an allowed origin string must not pass. + let status = send( + "POST", + "/mutate", + &[("origin", "https://app.predictiq.com.evil.example.com")], + ) + .await; + assert_eq!(status, StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn missing_origin_and_referer_without_cookie_passes() { + // Non-browser / API client: no Origin, no Cookie → allowed through. + let status = send("POST", "/mutate", &[]).await; + assert_eq!(status, StatusCode::OK); + } + + #[tokio::test] + async fn missing_origin_with_cookie_and_no_referer_passes() { + // Documented policy: Referer absent → pass through. + let status = send("POST", "/mutate", &[("cookie", "session=abc")]).await; + assert_eq!(status, StatusCode::OK); + } + + #[tokio::test] + async fn missing_origin_with_cookie_and_matching_referer_passes() { + let status = send( + "POST", + "/mutate", + &[ + ("cookie", "session=abc"), + ("referer", "https://app.predictiq.com/newsletter"), + ], + ) + .await; + assert_eq!(status, StatusCode::OK); + } + + #[tokio::test] + async fn missing_origin_with_cookie_and_mismatched_referer_is_rejected() { + let status = send( + "POST", + "/mutate", + &[ + ("cookie", "session=abc"), + ("referer", "https://evil.example.com/newsletter"), + ], + ) + .await; + assert_eq!(status, StatusCode::FORBIDDEN); + } + + #[tokio::test] + async fn api_key_short_circuits_mismatched_origin() { + // X-Api-Key requests bypass CSRF even with a hostile Origin. + let status = send( + "POST", + "/mutate", + &[ + ("x-api-key", "secret"), + ("origin", "https://evil.example.com"), + ], + ) + .await; + assert_eq!(status, StatusCode::OK); + } + + #[tokio::test] + async fn safe_method_with_mismatched_origin_passes() { + // GET is not state-changing → middleware skips the check. + let status = send( + "GET", + "/read", + &[("origin", "https://evil.example.com")], + ) + .await; + assert_eq!(status, StatusCode::OK); + } +} diff --git a/services/api/tests/csrf_tests.rs b/services/api/tests/csrf_tests.rs new file mode 100644 index 00000000..c500bc32 --- /dev/null +++ b/services/api/tests/csrf_tests.rs @@ -0,0 +1,88 @@ +//! Integration tests for the CSRF protection middleware. +//! +//! These tests exercise the Origin/Referer validation performed by +//! `predictiq_api::csrf` for state-changing newsletter requests, including the +//! `X-Api-Key` short-circuit path. + +use axum::body::Body; +use axum::http::{Request, StatusCode}; +use predictiq_api::csrf::csrf_protection_middleware; +use tower::ServiceExt; + +/// Build a minimal router that runs the CSRF middleware in front of a handler +/// which always succeeds, so we can assert purely on the middleware decision. +fn app() -> axum::Router { + axum::Router::new() + .route("/newsletter/subscribe", axum::routing::post(|| async { StatusCode::OK })) + .layer(axum::middleware::from_fn(csrf_protection_middleware)) +} + +fn post(uri: &str) -> Request { + Request::builder() + .method("POST") + .uri(uri) + .body(Body::empty()) + .unwrap() +} + +#[tokio::test] +async fn matching_allowed_origin_passes() { + let request = Request::builder() + .method("POST") + .uri("/newsletter/subscribe") + .header("origin", "https://app.predictiq.io") + .body(Body::empty()) + .unwrap(); + + let response = app().oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); +} + +#[tokio::test] +async fn mismatched_origin_is_rejected() { + let request = Request::builder() + .method("POST") + .uri("/newsletter/subscribe") + .header("origin", "https://evil.example.com") + .body(Body::empty()) + .unwrap(); + + let response = app().oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn missing_origin_and_referer_on_mutating_request_is_rejected() { + // Per the documented policy, a state-changing request without an + // Origin or Referer header is rejected. + let response = app().oneshot(post("/newsletter/subscribe")).await.unwrap(); + assert_eq!(response.status(), StatusCode::FORBIDDEN); +} + +#[tokio::test] +async fn matching_referer_passes_when_origin_absent() { + let request = Request::builder() + .method("POST") + .uri("/newsletter/subscribe") + .header("referer", "https://app.predictiq.io/newsletter") + .body(Body::empty()) + .unwrap(); + + let response = app().oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); +} + +#[tokio::test] +async fn api_key_short_circuits_origin_validation() { + // A valid API key bypasses the Origin/Referer check entirely, even when + // the Origin header is missing or untrusted. + let request = Request::builder() + .method("POST") + .uri("/newsletter/subscribe") + .header("x-api-key", "test-api-key") + .body(Body::empty()) + .unwrap(); + + let response = app().oneshot(request).await.unwrap(); + assert_eq!(response.status(), StatusCode::OK); +} From b20565b1171047f65930c24d16ed35cc451ac336 Mon Sep 17 00:00:00 2001 From: popsman Date: Mon, 28 Sep 2026 14:08:27 +0000 Subject: [PATCH 2/2] fix: #1503 Add unit tests for API versioning/deprecation middleware Closes #1503 --- services/api/src/versioning.rs | 199 +++++++++++++++++++++++++++++++++ 1 file changed, 199 insertions(+) diff --git a/services/api/src/versioning.rs b/services/api/src/versioning.rs index b2bd41df..fdb6da95 100644 --- a/services/api/src/versioning.rs +++ b/services/api/src/versioning.rs @@ -145,3 +145,202 @@ pub async fn v1_deprecation_middleware(req: Request, next: Next) -> Response { ); response } + +#[cfg(test)] +mod tests { + use super::*; + use axum::body::Body; + use axum::http::{Request as HttpRequest, StatusCode}; + use axum::routing::get; + use axum::Router; + use tower::ServiceExt; + + fn test_metrics() -> crate::metrics::Metrics { + crate::metrics::Metrics::new() + } + + async fn ok_handler() -> &'static str { + "ok" + } + + fn deprecation_router() -> Router { + Router::new().route("/api/v1/ping", get(ok_handler)).layer( + axum::middleware::from_fn(v1_deprecation_middleware), + ) + } + + fn versioning_router(metrics: crate::metrics::Metrics) -> Router { + let state = VersioningState::new(metrics); + Router::new() + .route("/ping", get(ok_handler)) + .layer(axum::middleware::from_fn_with_state( + state, + versioning_middleware, + )) + } + + #[tokio::test] + async fn deprecation_middleware_adds_headers_for_deprecated_version() { + let response = deprecation_router() + .oneshot( + HttpRequest::builder() + .uri("/api/v1/ping") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::OK); + let headers = response.headers(); + assert_eq!(headers.get("Deprecation").unwrap(), "true"); + assert_eq!( + headers.get("Sunset").unwrap(), + "Sat, 31 Dec 2026 00:00:00 GMT" + ); + let link = headers.get(header::LINK).unwrap().to_str().unwrap(); + assert!(link.contains("rel=\"deprecation\"")); + assert!(link.contains("")); + } + + #[tokio::test] + async fn deprecation_middleware_omits_headers_for_current_version() { + // The deprecation middleware is only mounted on deprecated (v1) routes; + // a current-version route must not carry the deprecation surface. + let router = Router::new().route("/api/v2/ping", get(ok_handler)); + let response = router + .oneshot( + HttpRequest::builder() + .uri("/api/v2/ping") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + + assert_eq!(response.status(), StatusCode::OK); + let headers = response.headers(); + assert!(headers.get("Deprecation").is_none()); + assert!(headers.get("Sunset").is_none()); + assert!(headers.get(header::LINK).is_none()); + } + + #[tokio::test] + async fn versioning_middleware_injects_resolved_version() { + let router = versioning_router(test_metrics()); + let response = router + .oneshot( + HttpRequest::builder() + .uri("/ping") + .header("API-Version", "v1") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + #[tokio::test] + async fn deprecated_api_calls_total_incremented_on_every_request() { + let metrics = test_metrics(); + let router = versioning_router(metrics.clone()); + + // Fire several requests from the same client; the sampler will only + // allow one log line, but the counter must increment every time. + for _ in 0..5 { + let response = router + .clone() + .oneshot( + HttpRequest::builder() + .uri("/ping") + .header("API-Version", "v1") + .header("x-forwarded-for", "203.0.113.7") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + } + + let count = metrics.deprecated_api_calls_total("v1"); + assert_eq!(count, 5, "counter must increment on every deprecated request"); + } + + #[tokio::test] + async fn current_version_does_not_increment_deprecated_counter() { + let metrics = test_metrics(); + let router = versioning_router(metrics.clone()); + + let response = router + .oneshot( + HttpRequest::builder() + .uri("/ping") + .header("API-Version", "v2") + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + assert_eq!(response.status(), StatusCode::OK); + assert_eq!(metrics.deprecated_api_calls_total("v2"), 0); + } + + #[test] + fn sampler_logs_first_call_then_suppresses_within_window() { + let sampler = DeprecationSampler::new(); + assert!(sampler.should_log("10.0.0.1", "v1")); + // Subsequent calls within the hour window are suppressed. + assert!(!sampler.should_log("10.0.0.1", "v1")); + assert!(!sampler.should_log("10.0.0.1", "v1")); + } + + #[test] + fn sampler_tracks_clients_and_versions_independently() { + let sampler = DeprecationSampler::new(); + assert!(sampler.should_log("10.0.0.1", "v1")); + // Different client -> independent window. + assert!(sampler.should_log("10.0.0.2", "v1")); + // Different version for same client -> independent window. + assert!(sampler.should_log("10.0.0.1", "v0")); + // Repeats are still suppressed per key. + assert!(!sampler.should_log("10.0.0.1", "v1")); + assert!(!sampler.should_log("10.0.0.2", "v1")); + } + + #[test] + fn sampler_allows_log_after_window_elapses() { + let sampler = DeprecationSampler::new(); + assert!(sampler.should_log("10.0.0.1", "v1")); + + // Simulate the hour window having elapsed by rewinding the stored + // timestamp for this key. + { + let mut map = sampler.last_logged.lock().unwrap(); + let key = "10.0.0.1:v1".to_string(); + let past = Instant::now() - Duration::from_secs(3601); + map.insert(key, past); + } + + assert!(sampler.should_log("10.0.0.1", "v1")); + } + + #[test] + fn peer_ip_prefers_forwarded_for_then_real_ip() { + let req = HttpRequest::builder() + .header("x-forwarded-for", "198.51.100.9, 10.0.0.1") + .body(Body::empty()) + .unwrap(); + assert_eq!(peer_ip_from_headers(&req), "198.51.100.9"); + + let req = HttpRequest::builder() + .header("x-real-ip", "198.51.100.10") + .body(Body::empty()) + .unwrap(); + assert_eq!(peer_ip_from_headers(&req), "198.51.100.10"); + + let req = HttpRequest::builder().body(Body::empty()).unwrap(); + assert_eq!(peer_ip_from_headers(&req), "unknown"); + } +}