#![cfg(any(
feature = "axum",
feature = "actix-web",
feature = "rocket",
feature = "warp"
))]
mod common;
use reqkey::middleware::{Middleware, MiddlewareConfig};
use wiremock::MockServer;
async fn engine(server: &MockServer) -> Middleware {
common::mount_allowed(server).await;
common::mount_ingest(server).await;
Middleware::new(
common::client(server),
MiddlewareConfig::builder("api_payments")
.build()
.expect("config"),
)
}
#[cfg(feature = "axum")]
mod axum_tests {
use crate::common;
use axum::{
body::{to_bytes, Body},
extract::Extension,
http::{Request, StatusCode},
routing::get,
Json, Router,
};
use reqkey::{axum::ReqKeyLayer, VerificationResult};
use serde_json::{json, Value};
use tower::ServiceExt;
use wiremock::MockServer;
#[tokio::test]
async fn layer_gates_exposes_decision_and_records_response() {
let server = MockServer::start().await;
let app = Router::new()
.route(
"/payments",
get(
|Extension(decision): Extension<VerificationResult>| async move {
Json::<Value>(json!({"remaining": decision.credits_remaining}))
},
),
)
.layer(ReqKeyLayer::new(super::engine(&server).await));
let response = app
.oneshot(
Request::builder()
.uri("/payments")
.header("x-api-key", "consumer_test")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get("x-reqkey-request-id")
.unwrap()
.to_str()
.unwrap(),
"request_123"
);
let body = to_bytes(response.into_body(), 1_024).await.unwrap();
assert_eq!(
serde_json::from_slice::<Value>(&body).unwrap()["remaining"],
9
);
assert_eq!(server.received_requests().await.unwrap().len(), 2);
}
#[tokio::test]
async fn layer_short_circuits_a_missing_key() {
let server = MockServer::start().await;
common::mount_ingest(&server).await;
let middleware = reqkey::middleware::Middleware::new(
common::client(&server),
reqkey::middleware::MiddlewareConfig::builder("api_payments")
.build()
.unwrap(),
);
let app = Router::new()
.route("/payments", get(|| async { "should not run" }))
.layer(ReqKeyLayer::new(middleware));
let response = app
.oneshot(Request::get("/payments").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
}
}
#[cfg(feature = "actix-web")]
mod actix_tests {
use actix_web::{test, web, App, HttpMessage, HttpRequest, HttpResponse};
use reqkey::{actix_web::ReqKeyMiddleware, VerificationResult};
use wiremock::MockServer;
#[actix_web::test]
async fn transform_gates_exposes_decision_and_records_response() {
let server = MockServer::start().await;
let app = test::init_service(
App::new()
.wrap(ReqKeyMiddleware::new(super::engine(&server).await))
.route(
"/payments",
web::get().to(|request: HttpRequest| async move {
let remaining = request
.extensions()
.get::<VerificationResult>()
.and_then(|decision| decision.credits_remaining);
HttpResponse::Created().json(serde_json::json!({"remaining": remaining}))
}),
),
)
.await;
let request = test::TestRequest::get()
.uri("/payments")
.insert_header(("x-api-key", "consumer_test"))
.to_request();
let response = test::call_service(&app, request).await;
assert_eq!(response.status(), actix_web::http::StatusCode::CREATED);
assert_eq!(
response
.headers()
.get("x-reqkey-request-id")
.unwrap()
.to_str()
.unwrap(),
"request_123"
);
assert_eq!(server.received_requests().await.unwrap().len(), 2);
}
}
#[cfg(feature = "rocket")]
mod rocket_tests {
use reqkey::rocket::{default_catchers, ReqKeyFairing, ReqKeyGuard};
use rocket::{http::Status, local::asynchronous::Client as LocalClient};
use wiremock::MockServer;
#[rocket::get("/payments")]
#[allow(clippy::needless_pass_by_value)]
fn payments(guard: ReqKeyGuard) -> String {
guard
.decision()
.and_then(|decision| decision.credits_remaining)
.unwrap_or_default()
.to_string()
}
#[rocket::async_test]
async fn guard_and_fairing_gate_then_record() {
let server = MockServer::start().await;
let rocket = rocket::build()
.attach(ReqKeyFairing::new(super::engine(&server).await))
.register("/", default_catchers())
.mount("/", rocket::routes![payments]);
let client = LocalClient::tracked(rocket).await.unwrap();
let response = client
.get("/payments")
.header(rocket::http::Header::new("X-API-Key", "consumer_test"))
.dispatch()
.await;
assert_eq!(response.status(), Status::Ok);
assert_eq!(
response.headers().get_one("X-ReqKey-Request-ID"),
Some("request_123")
);
assert_eq!(response.into_string().await.as_deref(), Some("9"));
assert_eq!(server.received_requests().await.unwrap().len(), 2);
}
}
#[cfg(feature = "warp")]
mod warp_tests {
use std::convert::Infallible;
use reqkey::warp::{recover, ReqKeyFilter, WarpRequest};
use warp::Filter;
use wiremock::MockServer;
#[tokio::test]
async fn filter_gates_and_record_wraps_the_reply() {
let server = MockServer::start().await;
let reqkey = ReqKeyFilter::new(super::engine(&server).await);
let route = warp::path("payments")
.and(reqkey.filter())
.and_then(|guard: WarpRequest| async move {
Ok::<_, Infallible>(guard.record("created").await)
})
.recover(recover);
let response = warp::test::request()
.path("/payments")
.header("x-api-key", "consumer_test")
.reply(&route)
.await;
assert_eq!(response.status(), 200);
assert_eq!(response.headers()["x-reqkey-request-id"], "request_123");
assert_eq!(response.body(), "created");
assert_eq!(server.received_requests().await.unwrap().len(), 2);
}
}