use axum::Router;
use std::sync::Arc;
use std::time::Duration;
use tower_http::compression::CompressionLayer;
use tower_http::timeout::TimeoutLayer;
use tower_http::trace::TraceLayer as TowerTraceLayer;
use super::builder::MiddlewareBuilder;
#[derive(Clone)]
pub enum TowerLayer {
Compression(CompressionLayer),
Timeout(TimeoutLayer),
TowerTrace(Arc<dyn Fn(Router) -> Router + Send + Sync>),
}
impl TowerLayer {
pub fn apply(self, router: Router) -> Router {
match self {
TowerLayer::Compression(layer) => router.layer(layer),
TowerLayer::Timeout(layer) => router.layer(layer),
TowerLayer::TowerTrace(apply_fn) => apply_fn(router),
}
}
pub fn kind_name(&self) -> &'static str {
match self {
TowerLayer::Compression(_) => "compression",
TowerLayer::Timeout(_) => "timeout",
TowerLayer::TowerTrace(_) => "tower_trace",
}
}
pub fn is_compression(&self) -> bool {
matches!(self, TowerLayer::Compression(_))
}
pub fn is_timeout(&self) -> bool {
matches!(self, TowerLayer::Timeout(_))
}
pub fn is_tower_trace(&self) -> bool {
matches!(self, TowerLayer::TowerTrace(_))
}
}
impl std::fmt::Debug for TowerLayer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
TowerLayer::Compression(_) => f.debug_tuple("Compression").finish(),
TowerLayer::Timeout(_) => f.debug_tuple("Timeout").finish(),
TowerLayer::TowerTrace(_) => f.debug_tuple("TowerTrace").finish(),
}
}
}
pub fn compression_layer() -> CompressionLayer {
CompressionLayer::new()
}
pub fn timeout_layer(duration: Duration) -> TimeoutLayer {
TimeoutLayer::with_status_code(axum::http::StatusCode::GATEWAY_TIMEOUT, duration)
}
pub fn tower_trace_layer() -> Arc<dyn Fn(Router) -> Router + Send + Sync> {
Arc::new(|router: Router| router.layer(TowerTraceLayer::new_for_http()))
}
#[derive(Clone)]
pub struct TowerCompat {
builder: MiddlewareBuilder,
tower_layers: Vec<TowerLayer>,
}
impl TowerCompat {
pub fn from_builder(builder: MiddlewareBuilder) -> Self {
Self {
builder,
tower_layers: Vec::new(),
}
}
pub fn php_global() -> Self {
Self::from_builder(MiddlewareBuilder::php_global_builder())
}
pub fn default_builder() -> Self {
Self::from_builder(MiddlewareBuilder::default_builder())
}
pub fn with_tower_layer(mut self, layer: TowerLayer) -> Self {
self.tower_layers.push(layer);
self
}
pub fn with_compression(self) -> Self {
self.with_tower_layer(TowerLayer::Compression(compression_layer()))
}
pub fn with_timeout(self, duration: Duration) -> Self {
self.with_tower_layer(TowerLayer::Timeout(timeout_layer(duration)))
}
pub fn with_tower_trace(self) -> Self {
self.with_tower_layer(TowerLayer::TowerTrace(tower_trace_layer()))
}
pub fn builder(&self) -> &MiddlewareBuilder {
&self.builder
}
pub fn tower_layers(&self) -> &[TowerLayer] {
&self.tower_layers
}
pub fn tower_layer_count(&self) -> usize {
self.tower_layers.len()
}
pub fn has_tower_layer(&self, kind_name: &str) -> bool {
self.tower_layers.iter().any(|l| l.kind_name() == kind_name)
}
pub fn apply(self, router: Router) -> Router {
let mut router = self.builder.apply(router);
for layer in self.tower_layers.into_iter().rev() {
router = layer.apply(router);
}
router
}
}
impl std::fmt::Debug for TowerCompat {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TowerCompat")
.field("builder", &self.builder)
.field(
"tower_layers",
&self
.tower_layers
.iter()
.map(|l| l.kind_name())
.collect::<Vec<_>>(),
)
.finish()
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use axum::http::{Request, StatusCode};
use http_body_util::BodyExt;
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_request_with_header(method: &str, uri: &str, key: &str, value: &str) -> Request<Body> {
Request::builder()
.method(method)
.uri(uri)
.header(key, value)
.body(Body::empty())
.unwrap()
}
fn make_router() -> Router {
Router::new()
.route("/", axum::routing::get(|| async { "hello world" }))
.route(
"/large",
axum::routing::get(|| async {
"a".repeat(1000)
}),
)
}
#[test]
fn test_tower_layer_compression_variant() {
let layer = TowerLayer::Compression(compression_layer());
assert!(layer.is_compression());
assert!(!layer.is_timeout());
assert!(!layer.is_tower_trace());
assert_eq!(layer.kind_name(), "compression");
}
#[test]
fn test_tower_layer_timeout_variant() {
let layer = TowerLayer::Timeout(timeout_layer(Duration::from_secs(30)));
assert!(!layer.is_compression());
assert!(layer.is_timeout());
assert!(!layer.is_tower_trace());
assert_eq!(layer.kind_name(), "timeout");
}
#[test]
fn test_tower_layer_tower_trace_variant() {
let layer = TowerLayer::TowerTrace(tower_trace_layer());
assert!(!layer.is_compression());
assert!(!layer.is_timeout());
assert!(layer.is_tower_trace());
assert_eq!(layer.kind_name(), "tower_trace");
}
#[test]
fn test_tower_layer_clone() {
let layer = TowerLayer::Compression(compression_layer());
let cloned = layer.clone();
assert_eq!(layer.kind_name(), cloned.kind_name());
}
#[test]
fn test_tower_layer_debug_format() {
let layer = TowerLayer::Compression(compression_layer());
let debug_str = format!("{:?}", layer);
assert!(debug_str.contains("Compression"));
}
#[test]
fn test_compression_layer_default() {
let _layer = compression_layer();
}
#[test]
fn test_timeout_layer_with_duration() {
let _layer = timeout_layer(Duration::from_secs(30));
}
#[test]
fn test_timeout_layer_zero_duration() {
let _layer = timeout_layer(Duration::from_secs(0));
}
#[test]
fn test_tower_trace_layer_for_http() {
let _layer = tower_trace_layer();
}
#[test]
fn test_tower_compat_from_builder() {
let builder = MiddlewareBuilder::new();
let compat = TowerCompat::from_builder(builder);
assert_eq!(compat.tower_layer_count(), 0);
assert!(compat.tower_layers().is_empty());
}
#[test]
fn test_tower_compat_php_global() {
let compat = TowerCompat::php_global();
assert_eq!(compat.builder().chain().len(), 2);
assert_eq!(compat.tower_layer_count(), 0);
}
#[test]
fn test_tower_compat_default_builder() {
let compat = TowerCompat::default_builder();
assert_eq!(compat.builder().chain().len(), 5);
assert_eq!(compat.tower_layer_count(), 0);
}
#[test]
fn test_tower_compat_clone() {
let compat = TowerCompat::php_global().with_compression();
let cloned = compat.clone();
assert_eq!(compat.tower_layer_count(), cloned.tower_layer_count());
}
#[test]
fn test_tower_compat_debug_format() {
let compat = TowerCompat::php_global()
.with_compression()
.with_timeout(Duration::from_secs(30));
let debug_str = format!("{:?}", compat);
assert!(debug_str.contains("TowerCompat"));
assert!(debug_str.contains("compression"));
assert!(debug_str.contains("timeout"));
}
#[test]
fn test_tower_compat_with_tower_layer() {
let compat = TowerCompat::php_global()
.with_tower_layer(TowerLayer::Compression(compression_layer()));
assert_eq!(compat.tower_layer_count(), 1);
assert!(compat.has_tower_layer("compression"));
}
#[test]
fn test_tower_compat_with_compression() {
let compat = TowerCompat::php_global().with_compression();
assert_eq!(compat.tower_layer_count(), 1);
assert!(compat.has_tower_layer("compression"));
assert!(compat.tower_layers()[0].is_compression());
}
#[test]
fn test_tower_compat_with_timeout() {
let compat = TowerCompat::php_global().with_timeout(Duration::from_secs(30));
assert_eq!(compat.tower_layer_count(), 1);
assert!(compat.has_tower_layer("timeout"));
assert!(compat.tower_layers()[0].is_timeout());
}
#[test]
fn test_tower_compat_with_tower_trace() {
let compat = TowerCompat::php_global().with_tower_trace();
assert_eq!(compat.tower_layer_count(), 1);
assert!(compat.has_tower_layer("tower_trace"));
assert!(compat.tower_layers()[0].is_tower_trace());
}
#[test]
fn test_tower_compat_chained_builders() {
let compat = TowerCompat::php_global()
.with_compression()
.with_timeout(Duration::from_secs(30))
.with_tower_trace();
assert_eq!(compat.tower_layer_count(), 3);
assert!(compat.tower_layers()[0].is_compression());
assert!(compat.tower_layers()[1].is_timeout());
assert!(compat.tower_layers()[2].is_tower_trace());
}
#[test]
fn test_tower_compat_has_tower_layer_negative() {
let compat = TowerCompat::php_global().with_compression();
assert!(!compat.has_tower_layer("timeout"));
assert!(!compat.has_tower_layer("tower_trace"));
}
#[tokio::test]
async fn test_tower_compat_apply_preserves_routes() {
let app = make_router();
let app = TowerCompat::php_global().apply(app);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = read_body(resp).await;
assert_eq!(body, "hello world");
}
#[tokio::test]
async fn test_tower_compat_apply_with_compression() {
let app = make_router();
let app = TowerCompat::php_global().with_compression().apply(app);
let resp = app
.oneshot(make_request_with_header(
"GET",
"/large",
"accept-encoding",
"gzip",
))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let content_encoding = resp.headers().get("content-encoding");
assert!(
content_encoding.is_some(),
"Response should have Content-Encoding: gzip header"
);
}
#[tokio::test]
async fn test_tower_compat_apply_with_timeout_passes_within_timeout() {
let app = make_router();
let app = TowerCompat::php_global()
.with_timeout(Duration::from_secs(30))
.apply(app);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_tower_compat_apply_with_tower_trace() {
let app = make_router();
let app = TowerCompat::php_global().with_tower_trace().apply(app);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let body = read_body(resp).await;
assert_eq!(body, "hello world");
}
#[tokio::test]
async fn test_tower_compat_apply_with_all_layers() {
let app = make_router();
let app = TowerCompat::php_global()
.with_compression()
.with_timeout(Duration::from_secs(30))
.with_tower_trace()
.apply(app);
let resp = app
.oneshot(make_request_with_header(
"GET",
"/large",
"accept-encoding",
"gzip",
))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_tower_compat_apply_execution_order() {
let app = make_router();
let app = TowerCompat::php_global().with_compression().apply(app);
let req = Request::builder()
.method("GET")
.uri("/large")
.header("origin", "https://example.com")
.header("accept-encoding", "gzip")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let allow_origin = resp.headers().get("access-control-allow-origin");
assert!(allow_origin.is_some(), "CORS should be applied");
let content_encoding = resp.headers().get("content-encoding");
assert!(content_encoding.is_some(), "Compression should be applied");
}
#[tokio::test]
async fn test_independent_compression_layer() {
let app = make_router().layer(compression_layer());
let resp = app
.oneshot(make_request_with_header(
"GET",
"/large",
"accept-encoding",
"gzip",
))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let content_encoding = resp.headers().get("content-encoding");
assert!(content_encoding.is_some());
}
#[tokio::test]
async fn test_independent_timeout_layer_passes() {
let app = make_router().layer(timeout_layer(Duration::from_secs(30)));
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_independent_tower_trace_layer() {
let app = (tower_trace_layer())(make_router());
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_tower_layer_apply_compression() {
let layer = TowerLayer::Compression(compression_layer());
let app = make_router();
let app = layer.apply(app);
let resp = app
.oneshot(make_request_with_header(
"GET",
"/large",
"accept-encoding",
"gzip",
))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_tower_layer_apply_timeout() {
let layer = TowerLayer::Timeout(timeout_layer(Duration::from_secs(30)));
let app = make_router();
let app = layer.apply(app);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn test_tower_layer_apply_tower_trace() {
let layer = TowerLayer::TowerTrace(tower_trace_layer());
let app = make_router();
let app = layer.apply(app);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[test]
fn test_r5_php_no_compression_middleware() {
let compat = TowerCompat::php_global().with_compression();
assert!(compat.has_tower_layer("compression"));
}
#[test]
fn test_r5_php_no_timeout_middleware() {
let compat = TowerCompat::php_global().with_timeout(Duration::from_secs(30));
assert!(compat.has_tower_layer("timeout"));
}
#[test]
fn test_r5_php_no_tower_trace_layer() {
let compat = TowerCompat::php_global().with_tower_trace();
assert!(compat.has_tower_layer("tower_trace"));
}
#[test]
fn test_r5_cors_already_in_middleware_builder() {
let compat = TowerCompat::php_global();
assert!(compat.builder().cors().is_some());
assert!(!compat.has_tower_layer("cors"));
}
#[test]
fn test_r5_sz_rust_trace_vs_tower_http_trace() {
let tower_layer = tower_trace_layer();
let _ = tower_layer;
}
#[test]
fn test_r5_execution_order_aligns_php_nginx() {
let compat = TowerCompat::php_global()
.with_compression()
.with_timeout(Duration::from_secs(30));
assert_eq!(compat.tower_layers()[0].kind_name(), "compression");
assert_eq!(compat.tower_layers()[1].kind_name(), "timeout");
}
#[test]
fn test_r5_compression_default_mime_types_align_nginx() {
let _layer = compression_layer();
}
#[tokio::test]
async fn test_r5_tower_compat_preserves_middleware_builder_behavior() {
let app = make_router();
let app = TowerCompat::php_global().apply(app);
let resp = app
.oneshot(make_request_with_header(
"GET",
"/",
"origin",
"https://example.com",
))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let allow_origin = resp.headers().get("access-control-allow-origin");
assert!(allow_origin.is_some(), "CORS should be applied");
}
#[tokio::test]
async fn integration_tower_compat_full_stack() {
let app = make_router();
let app = TowerCompat::php_global()
.with_compression()
.with_timeout(Duration::from_secs(30))
.with_tower_trace()
.apply(app);
let req = Request::builder()
.method("GET")
.uri("/large")
.header("origin", "https://example.com")
.header("accept-encoding", "gzip")
.body(Body::empty())
.unwrap();
let resp = app.oneshot(req).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert!(resp.headers().get("access-control-allow-origin").is_some());
assert!(resp.headers().get("content-encoding").is_some());
}
#[tokio::test]
async fn integration_tower_compat_options_preflight() {
let app = make_router();
let app = TowerCompat::php_global().with_compression().apply(app);
let resp = app
.oneshot(make_request_with_header(
"OPTIONS",
"/",
"origin",
"https://example.com",
))
.await
.unwrap();
assert!(resp.status().is_success() || resp.status() == StatusCode::NO_CONTENT);
}
#[tokio::test]
async fn integration_tower_compat_no_tower_layers() {
let app = make_router();
let app = TowerCompat::php_global().apply(app);
let resp = app.oneshot(make_request("GET", "/")).await.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test]
async fn integration_tower_compat_default_builder_with_compression() {
let app = make_router();
let app = TowerCompat::default_builder().with_compression().apply(app);
let resp = app
.oneshot(make_request_with_header(
"GET",
"/large",
"accept-encoding",
"gzip",
))
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
assert!(resp.headers().get("content-encoding").is_some());
}
}