mod common;
use std::sync::{Arc, Mutex};
use http::{header, HeaderMap, HeaderValue, Method};
use reqkey::{
middleware::{
client_ip_from_headers, extract_credential, AuthorizationOutcome, DenialCode, FailureMode,
KeyLocation, KeyScheme, Middleware, MiddlewareConfig, Mode, RequestContext,
ResponseContext,
},
Client, Error, Operation,
};
use serde_json::{json, Value};
use wiremock::{
matchers::{method, path},
Mock, MockServer, ResponseTemplate,
};
#[tokio::test]
async fn both_mode_validates_attaches_headers_and_records_redacted_metadata() {
let server = MockServer::start().await;
common::mount_allowed(&server).await;
common::mount_ingest(&server).await;
let config = MiddlewareConfig::builder("api_payments")
.key_location(KeyLocation::header("Authorization"))
.key_scheme(KeyScheme::Bearer)
.credit_cost_resolver(|request| Ok(u64::from(request.method == Method::POST) + 1))
.capture_query_params(true)
.capture_request_headers(true)
.capture_response_headers(true)
.capture_response_body(true)
.capture_client_ip(true)
.consumer_name_resolver(|_| Some("Acme".into()))
.build()
.expect("config");
let middleware = Middleware::new(common::client(&server), config);
let mut headers = HeaderMap::new();
headers.insert(
header::AUTHORIZATION,
HeaderValue::from_static("Bearer consumer_test"),
);
headers.insert(header::USER_AGENT, HeaderValue::from_static("test-agent"));
headers.insert("x-safe", HeaderValue::from_static("visible"));
let request = RequestContext::new(Method::POST, "/payments")
.with_query("page=2")
.with_headers(headers)
.with_client_ip("203.0.113.4");
let authorized = match middleware.authorize(request).await {
AuthorizationOutcome::Authorized(value) => value,
other => panic!("expected authorization, got {other:?}"),
};
assert_eq!(authorized.request_id(), Some("request_123"));
assert_eq!(
authorized
.response_headers()
.get("x-reqkey-credits-remaining")
.unwrap(),
"9"
);
let mut response_headers = HeaderMap::new();
response_headers.insert(
header::CONTENT_TYPE,
HeaderValue::from_static("application/json"),
);
response_headers.insert("x-response", HeaderValue::from_static("captured"));
middleware
.record(
authorized,
ResponseContext::new(201)
.with_headers(response_headers)
.with_body(r#"{"created":true}"#)
.with_latency_ms(12),
)
.await;
let requests = server.received_requests().await.expect("requests");
assert_eq!(requests.len(), 2);
let verify: Value = serde_json::from_slice(&requests[0].body).unwrap();
assert_eq!(verify["credits"], 2);
assert_eq!(verify["resource"], "/payments");
let ingest: Value = serde_json::from_slice(&requests[1].body).unwrap();
assert_eq!(ingest["requestId"], "request_123");
assert_eq!(ingest["path"], "/payments?page=2");
assert_eq!(ingest["queryParams"]["page"], "2");
assert_eq!(ingest["requestHeaders"]["x-safe"], "visible");
assert!(ingest["requestHeaders"].get("authorization").is_none());
assert_eq!(ingest["responseHeaders"]["x-response"], "captured");
assert_eq!(ingest["responseBody"], r#"{"created":true}"#);
assert_eq!(ingest["clientIp"], "203.0.113.4");
assert_eq!(ingest["consumerName"], "Acme");
assert_eq!(ingest["apiKey"], "consumer_test");
}
#[tokio::test]
async fn missing_and_rate_limited_requests_are_denied_and_ingested() {
let server = MockServer::start().await;
common::mount_ingest(&server).await;
let middleware = Middleware::new(
common::client(&server),
MiddlewareConfig::builder("api_payments")
.capture_response_body(true)
.build()
.unwrap(),
);
let denied = match middleware
.authorize(RequestContext::new(Method::GET, "/payments"))
.await
{
AuthorizationOutcome::Denied(value) => value,
other => panic!("expected denial, got {other:?}"),
};
assert_eq!(denied.status_code, 401);
assert_eq!(denied.error, DenialCode::MissingApiKey);
let requests = server.received_requests().await.unwrap();
let ingest: Value = serde_json::from_slice(&requests[0].body).unwrap();
assert_eq!(ingest["statusCode"], 401);
assert!(ingest["responseBody"]
.as_str()
.unwrap()
.contains("missing_api_key"));
let server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/key/validate"))
.respond_with(ResponseTemplate::new(429).set_body_json(json!({
"valid": false,
"retryAfter": 2.9,
"requestId": "request_limited"
})))
.mount(&server)
.await;
common::mount_ingest(&server).await;
let middleware = Middleware::new(common::client(&server), config());
let denied = match middleware.authorize(keyed_request()).await {
AuthorizationOutcome::Denied(value) => value,
other => panic!("expected denial, got {other:?}"),
};
assert_eq!(denied.status_code, 429);
assert_eq!(denied.retry_after, Some(2));
}
#[tokio::test]
async fn query_key_is_extracted_and_never_duplicated_into_captured_path() {
let server = MockServer::start().await;
common::mount_allowed(&server).await;
common::mount_ingest(&server).await;
let config = MiddlewareConfig::builder("api_payments")
.key_location(KeyLocation::query("api_key"))
.capture_query_params(true)
.build()
.unwrap();
let middleware = Middleware::new(common::client(&server), config);
let request =
RequestContext::new(Method::GET, "/payments").with_query("api_key=consumer_test&page=2");
let authorized = match middleware.authorize(request).await {
AuthorizationOutcome::Authorized(value) => value,
other => panic!("expected authorization, got {other:?}"),
};
middleware
.record(authorized, ResponseContext::new(200))
.await;
let requests = server.received_requests().await.unwrap();
let ingest: Value = serde_json::from_slice(&requests[1].body).unwrap();
assert_eq!(ingest["path"], "/payments?page=2");
assert!(ingest["queryParams"].get("api_key").is_none());
assert_eq!(ingest["apiKey"], "consumer_test");
}
#[tokio::test]
async fn excluded_paths_bypass_and_failure_modes_match_python() {
let server = MockServer::start().await;
let middleware = Middleware::new(
common::client(&server),
MiddlewareConfig::builder("api_payments")
.exclude_path("/health")
.exclude_path("/docs/*")
.build()
.unwrap(),
);
assert!(matches!(
middleware
.authorize(RequestContext::new(Method::GET, "/docs/openapi.json"))
.await,
AuthorizationOutcome::Bypass
));
let errors = Arc::new(Mutex::new(Vec::new()));
let captured = Arc::clone(&errors);
let unavailable = Client::builder()
.project_key("project_test")
.base_url("http://127.0.0.1:9")
.build()
.unwrap();
let closed = Middleware::new(unavailable.clone(), config());
assert!(matches!(
closed.authorize(keyed_request()).await,
AuthorizationOutcome::Denied(ref value)
if value.error == DenialCode::ReqKeyUnavailable
));
let open_config = MiddlewareConfig::builder("api_payments")
.failure_mode(FailureMode::Open)
.on_error(move |event| captured.lock().unwrap().push(event.clone()))
.build()
.unwrap();
let open = Middleware::new(unavailable, open_config);
let authorized = match open.authorize(keyed_request()).await {
AuthorizationOutcome::Authorized(value) => value,
other => panic!("expected fail-open, got {other:?}"),
};
assert!(authorized.failure.is_some());
assert_eq!(errors.lock().unwrap()[0].operation, Operation::Validate);
}
#[tokio::test]
async fn validate_and_ingest_modes_are_independent() {
let server = MockServer::start().await;
common::mount_allowed(&server).await;
let validate = Middleware::new(
common::client(&server),
MiddlewareConfig::builder("api_payments")
.mode(Mode::Validate)
.build()
.unwrap(),
);
let authorized = match validate.authorize(keyed_request()).await {
AuthorizationOutcome::Authorized(value) => value,
other => panic!("expected authorization, got {other:?}"),
};
validate.record(authorized, ResponseContext::new(200)).await;
assert_eq!(server.received_requests().await.unwrap().len(), 1);
let server = MockServer::start().await;
common::mount_ingest(&server).await;
let ingest = Middleware::new(
common::client(&server),
MiddlewareConfig::builder("api_payments")
.mode(Mode::Ingest)
.build()
.unwrap(),
);
let authorized = match ingest
.authorize(RequestContext::new(Method::GET, "/public"))
.await
{
AuthorizationOutcome::Authorized(value) => value,
other => panic!("expected ingest-only authorization, got {other:?}"),
};
assert!(authorized.decision.is_none());
ingest.record(authorized, ResponseContext::new(204)).await;
let requests = server.received_requests().await.unwrap();
assert_eq!(requests.len(), 1);
assert_eq!(requests[0].url.path(), "/ingest");
}
#[test]
fn extraction_configuration_and_proxy_helpers_are_strict() {
assert_eq!(
extract_credential(Some("Bearer consumer_test"), KeyScheme::Bearer).as_deref(),
Some("consumer_test")
);
assert_eq!(
extract_credential(Some("Basic consumer_test"), KeyScheme::Bearer),
None
);
assert!(matches!(
MiddlewareConfig::builder("api")
.key_location(KeyLocation::cookie("api_key"))
.key_scheme(KeyScheme::Bearer)
.build(),
Err(Error::Configuration(_))
));
let mut headers = HeaderMap::new();
headers.insert(
"x-real-ip",
HeaderValue::from_static("invalid, 198.51.100.7"),
);
assert_eq!(
client_ip_from_headers(&headers).as_deref(),
Some("198.51.100.7")
);
let mut forwarded = HeaderMap::new();
forwarded.insert(
"forwarded",
HeaderValue::from_static("for=\"[2001:db8::1]:443\";proto=https"),
);
assert_eq!(
client_ip_from_headers(&forwarded).as_deref(),
Some("2001:db8::1")
);
}
fn config() -> reqkey::middleware::MiddlewareConfig {
MiddlewareConfig::builder("api_payments")
.build()
.expect("test config")
}
fn keyed_request() -> RequestContext {
let mut headers = HeaderMap::new();
headers.insert("x-api-key", HeaderValue::from_static("consumer_test"));
RequestContext::new(Method::GET, "/payments").with_headers(headers)
}