Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
153 changes: 153 additions & 0 deletions services/api/src/csrf.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<CsrfConfig> {
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);
}
}
199 changes: 199 additions & 0 deletions services/api/src/versioning.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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("</api/v1>"));
}

#[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");
}
}
Loading
Loading