use axum::Router;
use tower_http::cors::CorsLayer;
use super::auth::{auth_middleware, AuthConfig};
use super::chain::MiddlewareChain;
use super::cors;
use super::log::{log_middleware_with_config, LogConfig};
use super::order::MiddlewareKind;
#[cfg(test)]
use super::order::{DEFAULT_ORDER, PHP_GLOBAL_ORDER};
use super::rate_limit::{rate_limit_middleware, RateLimitConfig};
use super::trace::{trace_middleware, TraceConfig};
#[derive(Debug, Clone)]
pub struct MiddlewareBuilder {
chain: MiddlewareChain,
cors: Option<CorsLayer>,
log: Option<LogConfig>,
auth: Option<AuthConfig>,
rate_limit: Option<RateLimitConfig>,
trace: Option<TraceConfig>,
}
impl MiddlewareBuilder {
pub fn new() -> Self {
Self {
chain: MiddlewareChain::new(),
cors: None,
log: None,
auth: None,
rate_limit: None,
trace: None,
}
}
pub fn default_builder() -> Self {
Self {
chain: MiddlewareChain::default_chain(),
cors: None,
log: None,
auth: None,
rate_limit: None,
trace: None,
}
}
pub fn php_global_builder() -> Self {
Self {
chain: MiddlewareChain::php_global(),
cors: Some(cors::cors_layer()),
log: None,
auth: None,
rate_limit: None,
trace: None,
}
}
pub fn with_chain(mut self, chain: MiddlewareChain) -> Self {
self.chain = chain;
self
}
pub fn with_cors(mut self, layer: CorsLayer) -> Self {
self.cors = Some(layer);
self
}
pub fn with_log(mut self, config: LogConfig) -> Self {
self.log = Some(config);
self
}
pub fn with_auth(mut self, config: AuthConfig) -> Self {
self.auth = Some(config);
self
}
pub fn with_rate_limit(mut self, config: RateLimitConfig) -> Self {
self.rate_limit = Some(config);
self
}
pub fn with_trace(mut self, config: TraceConfig) -> Self {
self.trace = Some(config);
self
}
pub fn remove_kind(&mut self, kind: MiddlewareKind) -> usize {
let removed = self.chain.remove_kind(kind);
if removed > 0 {
match kind {
MiddlewareKind::Trace => self.trace = None,
MiddlewareKind::Cors => self.cors = None,
MiddlewareKind::Log => self.log = None,
MiddlewareKind::RateLimit => self.rate_limit = None,
MiddlewareKind::Auth => self.auth = None,
}
}
removed
}
pub fn remove_from(&mut self, kind: MiddlewareKind) -> usize {
let removed_kinds: Vec<MiddlewareKind> = if let Some(pos) = self.chain.position(kind) {
self.chain.order()[pos..].to_vec()
} else {
return 0;
};
let removed = self.chain.remove_from(kind);
for k in removed_kinds {
match k {
MiddlewareKind::Trace => self.trace = None,
MiddlewareKind::Cors => self.cors = None,
MiddlewareKind::Log => self.log = None,
MiddlewareKind::RateLimit => self.rate_limit = None,
MiddlewareKind::Auth => self.auth = None,
}
}
removed
}
pub fn chain(&self) -> &MiddlewareChain {
&self.chain
}
pub fn cors(&self) -> Option<&CorsLayer> {
self.cors.as_ref()
}
pub fn log(&self) -> Option<&LogConfig> {
self.log.as_ref()
}
pub fn auth(&self) -> Option<&AuthConfig> {
self.auth.as_ref()
}
pub fn rate_limit(&self) -> Option<&RateLimitConfig> {
self.rate_limit.as_ref()
}
pub fn trace(&self) -> Option<&TraceConfig> {
self.trace.as_ref()
}
pub fn is_enabled(&self, kind: MiddlewareKind) -> bool {
if !self.chain.contains(kind) {
return false;
}
match kind {
MiddlewareKind::Trace => self.trace.is_some(),
MiddlewareKind::Cors => self.cors.is_some(),
MiddlewareKind::Log => self.log.is_some(),
MiddlewareKind::RateLimit => self.rate_limit.is_some(),
MiddlewareKind::Auth => self.auth.is_some(),
}
}
pub fn apply(self, mut router: Router) -> Router {
let mut cors = self.cors;
let mut log = self.log;
let mut auth = self.auth;
let mut rate_limit = self.rate_limit;
let mut trace = self.trace;
for kind in self.chain.service_builder_order() {
router = match kind {
MiddlewareKind::Trace => {
if let Some(cfg) = trace.take() {
router.layer(axum::middleware::from_fn_with_state(cfg, trace_middleware))
} else {
router
}
}
MiddlewareKind::Cors => {
if let Some(layer) = cors.take() {
router.layer(layer)
} else {
router
}
}
MiddlewareKind::Log => {
if let Some(cfg) = log.take() {
router.layer(axum::middleware::from_fn_with_state(
cfg,
log_middleware_with_config,
))
} else {
router
}
}
MiddlewareKind::RateLimit => {
if let Some(cfg) = rate_limit.take() {
router.layer(axum::middleware::from_fn_with_state(
cfg,
rate_limit_middleware,
))
} else {
router
}
}
MiddlewareKind::Auth => {
if let Some(cfg) = auth.take() {
router.layer(axum::middleware::from_fn_with_state(cfg, auth_middleware))
} else {
router
}
}
};
}
router
}
}
impl Default for MiddlewareBuilder {
fn default() -> Self {
Self::default_builder()
}
}
impl std::fmt::Display for MiddlewareBuilder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "MiddlewareBuilder(chain={}, ", self.chain)?;
write!(
f,
"cors={}, log={}, auth={}, rate_limit={}, trace={})",
self.cors.is_some(),
self.log.is_some(),
self.auth.is_some(),
self.rate_limit.is_some(),
self.trace.is_some()
)
}
}
pub fn default_builder() -> MiddlewareBuilder {
MiddlewareBuilder::default_builder()
}
pub fn php_global_builder() -> MiddlewareBuilder {
MiddlewareBuilder::php_global_builder()
}
pub fn with_default_cors() -> MiddlewareBuilder {
MiddlewareBuilder::php_global_builder()
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::http::Request;
use axum::http::StatusCode;
use http_body_util::BodyExt;
use std::sync::Arc;
use std::time::Duration;
use crate::orm::SlidingWindowRateLimiter;
use crate::orm::SzTracer;
use tower::ServiceExt;
async fn read_body(resp: axum::response::Response) -> String {
let bytes = resp.into_body().collect().await.unwrap().to_bytes();
String::from_utf8(bytes.to_vec()).unwrap()
}
fn make_request(method: &str, uri: &str) -> Request<Body> {
Request::builder()
.method(method)
.uri(uri)
.body(Body::empty())
.unwrap()
}
fn make_trace_config() -> TraceConfig {
let tracer: Arc<dyn Tracer + Send + Sync> = Arc::new(SzTracer::new("test-service"));
TraceConfig::new(tracer)
}
fn make_rate_limit_config() -> RateLimitConfig {
let limiter: Arc<dyn RateLimiter + Send + Sync> =
Arc::new(SlidingWindowRateLimiter::new(1000, Duration::from_secs(60)));
RateLimitConfig::new(limiter)
}
use crate::orm::RateLimiter;
use crate::orm::Tracer;
#[test]
fn test_new_creates_empty_builder() {
let builder = MiddlewareBuilder::new();
assert!(builder.chain().is_empty());
assert_eq!(builder.chain().len(), 0);
assert!(builder.cors().is_none());
assert!(builder.log().is_none());
assert!(builder.auth().is_none());
assert!(builder.rate_limit().is_none());
assert!(builder.trace().is_none());
}
#[test]
fn test_default_builder_uses_default_order() {
let builder = MiddlewareBuilder::default_builder();
assert_eq!(builder.chain().order(), DEFAULT_ORDER);
assert_eq!(builder.chain().len(), 5);
assert!(builder.cors().is_none());
assert!(builder.log().is_none());
assert!(builder.auth().is_none());
assert!(builder.rate_limit().is_none());
assert!(builder.trace().is_none());
}
#[test]
fn test_default_trait_uses_default_builder() {
let builder = MiddlewareBuilder::default();
assert_eq!(builder.chain().order(), DEFAULT_ORDER);
}
#[test]
fn test_php_global_builder_uses_php_global_order() {
let builder = MiddlewareBuilder::php_global_builder();
assert_eq!(builder.chain().order(), PHP_GLOBAL_ORDER);
assert_eq!(builder.chain().len(), 2);
assert!(builder.cors().is_some());
assert!(builder.log().is_none());
assert!(builder.auth().is_none());
assert!(builder.rate_limit().is_none());
assert!(builder.trace().is_none());
}
#[test]
fn test_with_chain_replaces_chain() {
let custom_chain = MiddlewareChain::new()
.push(MiddlewareKind::Cors)
.push(MiddlewareKind::Log);
let builder = MiddlewareBuilder::new().with_chain(custom_chain);
assert_eq!(
builder.chain().order(),
[MiddlewareKind::Cors, MiddlewareKind::Log]
);
}
#[test]
fn test_with_cors_sets_layer() {
let builder = MiddlewareBuilder::new().with_cors(cors::cors_layer());
assert!(builder.cors().is_some());
}
#[test]
fn test_with_log_sets_config() {
let config = LogConfig::default().with_exclude_paths(vec!["/health".to_string()]);
let builder = MiddlewareBuilder::new().with_log(config);
assert!(builder.log().is_some());
assert_eq!(
builder.log().unwrap().exclude_paths,
vec!["/health".to_string()]
);
}
#[test]
fn test_with_auth_sets_config() {
let config = AuthConfig::default().with_secret("custom-secret");
let builder = MiddlewareBuilder::new().with_auth(config);
assert!(builder.auth().is_some());
assert_eq!(builder.auth().unwrap().secret, "custom-secret");
}
#[test]
fn test_with_rate_limit_sets_config() {
let config = make_rate_limit_config();
let builder = MiddlewareBuilder::new().with_rate_limit(config);
assert!(builder.rate_limit().is_some());
}
#[test]
fn test_with_trace_sets_config() {
let config = make_trace_config();
let builder = MiddlewareBuilder::new().with_trace(config);
assert!(builder.trace().is_some());
}
#[test]
fn test_chained_with_xxx_builders() {
let builder = MiddlewareBuilder::default_builder()
.with_cors(cors::cors_layer())
.with_log(LogConfig::default())
.with_auth(AuthConfig::default())
.with_rate_limit(make_rate_limit_config())
.with_trace(make_trace_config());
assert!(builder.cors().is_some());
assert!(builder.log().is_some());
assert!(builder.auth().is_some());
assert!(builder.rate_limit().is_some());
assert!(builder.trace().is_some());
}
#[test]
fn test_remove_kind_removes_from_chain_and_config() {
let mut builder = MiddlewareBuilder::default_builder()
.with_auth(AuthConfig::default())
.with_log(LogConfig::default());
assert!(builder.auth().is_some());
let removed = builder.remove_kind(MiddlewareKind::Auth);
assert_eq!(removed, 1);
assert!(builder.auth().is_none());
assert!(!builder.chain().contains(MiddlewareKind::Auth));
}
#[test]
fn test_remove_kind_not_present_returns_zero() {
let mut builder = MiddlewareBuilder::php_global_builder();
let removed = builder.remove_kind(MiddlewareKind::Auth);
assert_eq!(removed, 0);
}
#[test]
fn test_remove_from_removes_kind_and_after() {
let mut builder = MiddlewareBuilder::default_builder()
.with_rate_limit(make_rate_limit_config())
.with_auth(AuthConfig::default());
let removed = builder.remove_from(MiddlewareKind::RateLimit);
assert_eq!(removed, 2);
assert!(builder.rate_limit().is_none());
assert!(builder.auth().is_none());
assert!(!builder.chain().contains(MiddlewareKind::RateLimit));
assert!(!builder.chain().contains(MiddlewareKind::Auth));
}
#[test]
fn test_is_enabled_true_when_chain_and_config_present() {
let builder = MiddlewareBuilder::default_builder().with_auth(AuthConfig::default());
assert!(builder.is_enabled(MiddlewareKind::Auth));
}
#[test]
fn test_is_enabled_false_when_config_missing() {
let builder = MiddlewareBuilder::default_builder();
assert!(!builder.is_enabled(MiddlewareKind::Auth));
}
#[test]
fn test_is_enabled_false_when_not_in_chain() {
let builder = MiddlewareBuilder::new().with_auth(AuthConfig::default());
assert!(!builder.is_enabled(MiddlewareKind::Auth));
}
#[test]
fn test_apply_empty_builder_returns_router_unchanged() {
let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
let builder = MiddlewareBuilder::new();
let app = builder.apply(router);
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = read_body(resp).await;
assert_eq!(body, "ok");
});
}
#[test]
fn test_apply_with_cors_only() {
let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
let builder = MiddlewareBuilder::new()
.with_chain(MiddlewareChain::new().push(MiddlewareKind::Cors))
.with_cors(cors::cors_layer());
let app = builder.apply(router);
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let req = Request::builder()
.method("GET")
.uri("/")
.header("origin", "https://example.com")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert!(resp.headers().contains_key("access-control-allow-origin"));
});
}
#[test]
fn test_apply_skips_middlewares_without_config() {
let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
let builder = MiddlewareBuilder::default_builder();
let app = builder.apply(router);
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
});
}
#[test]
fn test_apply_with_all_configs_does_not_panic() {
let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
let builder = MiddlewareBuilder::default_builder()
.with_cors(cors::cors_layer())
.with_log(LogConfig::default())
.with_auth(AuthConfig::default())
.with_rate_limit(make_rate_limit_config())
.with_trace(make_trace_config());
let app = builder.apply(router);
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
});
}
#[test]
fn test_apply_preserves_business_order() {
let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
let builder = MiddlewareBuilder::new()
.with_chain(
MiddlewareChain::new()
.push(MiddlewareKind::Cors)
.push(MiddlewareKind::Auth),
)
.with_cors(cors::cors_layer())
.with_auth(AuthConfig::default());
let app = builder.apply(router);
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
});
}
#[test]
fn test_default_builder_helper() {
let builder = default_builder();
assert_eq!(builder.chain().order(), DEFAULT_ORDER);
}
#[test]
fn test_php_global_builder_helper() {
let builder = php_global_builder();
assert_eq!(builder.chain().order(), PHP_GLOBAL_ORDER);
assert!(builder.cors().is_some());
}
#[test]
fn test_with_default_cors_helper() {
let builder = with_default_cors();
assert_eq!(builder.chain().order(), PHP_GLOBAL_ORDER);
assert!(builder.cors().is_some());
}
#[test]
fn test_display_empty_builder() {
let builder = MiddlewareBuilder::new();
let s = format!("{builder}");
assert!(s.contains("MiddlewareBuilder"));
assert!(s.contains("chain=MiddlewareChain[]"));
assert!(s.contains("cors=false"));
}
#[test]
fn test_display_full_builder() {
let builder = MiddlewareBuilder::default_builder()
.with_cors(cors::cors_layer())
.with_log(LogConfig::default())
.with_auth(AuthConfig::default())
.with_rate_limit(make_rate_limit_config())
.with_trace(make_trace_config());
let s = format!("{builder}");
assert!(s.contains("cors=true"));
assert!(s.contains("log=true"));
assert!(s.contains("auth=true"));
assert!(s.contains("rate_limit=true"));
assert!(s.contains("trace=true"));
}
#[test]
fn test_clone_preserves_state() {
let builder = MiddlewareBuilder::default_builder()
.with_cors(cors::cors_layer())
.with_log(LogConfig::default())
.with_auth(AuthConfig::default());
let cloned = builder.clone();
assert_eq!(builder.chain(), cloned.chain());
assert!(cloned.cors().is_some());
assert!(cloned.log().is_some());
assert!(cloned.auth().is_some());
}
#[test]
fn r5_1_php_global_order_matches_php_app_middleware() {
let builder = php_global_builder();
assert_eq!(
builder.chain().order(),
[MiddlewareKind::Trace, MiddlewareKind::Cors]
);
}
#[test]
fn r5_2_php_global_builder_has_default_cors() {
let builder = php_global_builder();
assert!(builder.cors().is_some());
}
#[test]
fn r5_3_php_global_builder_trace_config_none_by_default() {
let builder = php_global_builder();
assert!(builder.trace().is_none());
}
#[test]
fn r5_4_default_order_aligns_with_php_extension() {
let builder = default_builder();
assert_eq!(
builder.chain().order(),
[
MiddlewareKind::Trace,
MiddlewareKind::Cors,
MiddlewareKind::Log,
MiddlewareKind::RateLimit,
MiddlewareKind::Auth
]
);
assert!(builder.chain().order().starts_with(PHP_GLOBAL_ORDER));
}
#[test]
fn r5_5_php_public_api_skip_auth_via_remove_from() {
let mut builder = default_builder().with_auth(AuthConfig::default());
let removed = builder.remove_from(MiddlewareKind::RateLimit);
assert_eq!(removed, 2);
assert!(!builder.is_enabled(MiddlewareKind::Auth));
assert!(!builder.is_enabled(MiddlewareKind::RateLimit));
assert!(builder.chain().contains(MiddlewareKind::Trace));
assert!(builder.chain().contains(MiddlewareKind::Cors));
assert!(builder.chain().contains(MiddlewareKind::Log));
}
#[test]
fn r5_6_service_builder_order_reverses_for_router_layer() {
let builder = default_builder();
let sb_order = builder.chain().service_builder_order();
assert_eq!(
sb_order,
[
MiddlewareKind::Auth,
MiddlewareKind::RateLimit,
MiddlewareKind::Log,
MiddlewareKind::Cors,
MiddlewareKind::Trace,
]
);
}
#[test]
fn r5_7_php_global_builder_skip_middlewares_without_config() {
let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
let app = php_global_builder().apply(router);
let rt = tokio::runtime::Runtime::new().unwrap();
rt.block_on(async {
let req = Request::builder()
.method("GET")
.uri("/")
.header("origin", "https://example.com")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert!(resp.headers().contains_key("access-control-allow-origin"));
});
}
#[test]
fn r5_8_is_enabled_aligns_with_php_middleware_registration() {
let builder = MiddlewareBuilder::default_builder()
.with_cors(cors::cors_layer())
.with_auth(AuthConfig::default());
assert!(builder.is_enabled(MiddlewareKind::Cors));
assert!(builder.is_enabled(MiddlewareKind::Auth));
assert!(!builder.is_enabled(MiddlewareKind::Trace));
assert!(!builder.is_enabled(MiddlewareKind::Log));
assert!(!builder.is_enabled(MiddlewareKind::RateLimit));
}
#[tokio::test]
async fn integration_apply_returns_working_router() {
let router = Router::new().route("/health", axum::routing::get(|| async { "ok" }));
let app = MiddlewareBuilder::new()
.with_chain(MiddlewareChain::new().push(MiddlewareKind::Cors))
.with_cors(cors::cors_layer())
.apply(router);
let resp = app.oneshot(make_request("GET", "/health")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = read_body(resp).await;
assert_eq!(body, "ok");
}
#[tokio::test]
async fn integration_cors_preflight_passes_through() {
let router = Router::new().route(
"/api",
axum::routing::get(|| async { "ok" }).post(|| async { "created" }),
);
let app = MiddlewareBuilder::new()
.with_chain(MiddlewareChain::new().push(MiddlewareKind::Cors))
.with_cors(cors::cors_layer())
.apply(router);
let req = Request::builder()
.method("OPTIONS")
.uri("/api")
.header("origin", "https://example.com")
.header("access-control-request-method", "POST")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert!(resp.status().is_success());
assert!(resp.headers().contains_key("access-control-allow-methods"));
}
#[tokio::test]
async fn integration_auth_rejects_unauthenticated_request() {
let router = Router::new().route("/protected", axum::routing::get(|| async { "ok" }));
let app = MiddlewareBuilder::new()
.with_chain(MiddlewareChain::new().push(MiddlewareKind::Auth))
.with_auth(AuthConfig::default())
.apply(router);
let resp = app
.oneshot(make_request("GET", "/protected"))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
let body = read_body(resp).await;
assert!(body.contains("\"code\":-1"));
}
#[tokio::test]
async fn integration_log_does_not_block_request() {
let router = Router::new().route("/", axum::routing::get(|| async { "ok" }));
let app = MiddlewareBuilder::new()
.with_chain(MiddlewareChain::new().push(MiddlewareKind::Log))
.with_log(LogConfig::default())
.apply(router);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = read_body(resp).await;
assert_eq!(body, "ok");
}
}