use axum::http::{Extensions, HeaderMap, HeaderValue, StatusCode, Version};
use axum::middleware::{self, Next};
use axum::Router;
use tower_http::compression::predicate::{And, DefaultPredicate};
use tower_http::compression::{CompressionLayer, Predicate};
use tower_http::cors::{AllowOrigin, CorsLayer};
mod etag;
mod jsonrpc_stdio;
pub use etag::compute_strong_etag;
pub use jsonrpc_stdio::{
read_jsonrpc_stdio_frame, write_jsonrpc_stdio_message, JsonRpcStdioFrame,
JsonRpcStdioFrameStyle,
};
pub const COMPRESSION_OPT_OUT_HEADER: &str = "x-compress";
pub const COMPRESSION_OPT_OUT_VALUE: &str = "never";
pub const COMPRESSION_MIN_SIZE_BYTES: u16 = 512;
#[derive(Clone, Debug, Default)]
pub struct CorsConfig {
pub allow_origins: Vec<String>,
pub allow_any_origin: bool,
pub allow_methods: Vec<String>,
pub allow_headers: Vec<String>,
pub expose_headers: Vec<String>,
pub allow_credentials: bool,
pub max_age_seconds: Option<u32>,
}
impl CorsConfig {
pub fn allow_any() -> Self {
Self {
allow_any_origin: true,
..Self::default()
}
}
}
#[derive(Clone, Debug, Default)]
pub struct TransportConfig {
pub compression: bool,
pub etag: bool,
pub cors: Option<CorsConfig>,
}
impl TransportConfig {
pub fn default_enabled() -> Self {
Self {
compression: true,
etag: true,
cors: None,
}
}
}
pub fn apply_transport_layers(mut router: Router, config: &TransportConfig) -> Router {
if config.etag {
router = etag::install_on(router);
}
if config.compression {
router = router.layer(compression_layer());
router = router.layer(middleware::from_fn(strip_compression_marker));
}
if let Some(cors) = &config.cors {
router = router.layer(build_cors_layer(cors));
}
router
}
fn compression_layer() -> CompressionLayer<And<HeaderOptOutPredicate, DefaultPredicate>> {
CompressionLayer::new()
.gzip(true)
.br(true)
.zstd(true)
.compress_when(HeaderOptOutPredicate.and(DefaultPredicate::new()))
}
#[derive(Clone, Copy, Debug, Default)]
pub struct HeaderOptOutPredicate;
impl Predicate for HeaderOptOutPredicate {
fn should_compress<B>(&self, response: &axum::http::Response<B>) -> bool {
opt_out_predicate(
response.status(),
response.version(),
response.headers(),
response.extensions(),
)
}
}
fn opt_out_predicate(
_status: StatusCode,
_version: Version,
headers: &HeaderMap,
_extensions: &Extensions,
) -> bool {
!headers
.get_all(COMPRESSION_OPT_OUT_HEADER)
.iter()
.any(|value| {
value
.to_str()
.map(|s| s.eq_ignore_ascii_case(COMPRESSION_OPT_OUT_VALUE))
.unwrap_or(false)
})
}
async fn strip_compression_marker(
req: axum::extract::Request,
next: Next,
) -> axum::response::Response {
let mut response = next.run(req).await;
response.headers_mut().remove(COMPRESSION_OPT_OUT_HEADER);
response
}
fn build_cors_layer(config: &CorsConfig) -> CorsLayer {
let mut layer = CorsLayer::new();
let wildcard_origin =
config.allow_any_origin || config.allow_origins.iter().any(|origin| origin == "*");
if wildcard_origin {
layer = layer.allow_origin(AllowOrigin::any());
} else if !config.allow_origins.is_empty() {
let origins: Vec<HeaderValue> = config
.allow_origins
.iter()
.filter_map(|origin| HeaderValue::from_str(origin).ok())
.collect();
layer = layer.allow_origin(AllowOrigin::list(origins));
}
let methods: Vec<axum::http::Method> = if config.allow_methods.is_empty() {
["GET", "POST", "PUT", "PATCH", "DELETE", "OPTIONS"]
.into_iter()
.filter_map(|m| m.parse().ok())
.collect()
} else {
config
.allow_methods
.iter()
.filter_map(|m| m.parse().ok())
.collect()
};
layer = layer.allow_methods(methods);
let header_names: Vec<axum::http::HeaderName> = if config.allow_headers.is_empty() {
["authorization", "content-type", "x-request-id"]
.into_iter()
.filter_map(|h| h.parse().ok())
.collect()
} else {
config
.allow_headers
.iter()
.filter_map(|h| h.parse().ok())
.collect()
};
layer = layer.allow_headers(header_names);
if !config.expose_headers.is_empty() {
let exposed: Vec<axum::http::HeaderName> = config
.expose_headers
.iter()
.filter_map(|h| h.parse().ok())
.collect();
layer = layer.expose_headers(exposed);
}
if config.allow_credentials && !wildcard_origin {
layer = layer.allow_credentials(true);
}
let max_age = config.max_age_seconds.unwrap_or(3600);
layer.max_age(std::time::Duration::from_secs(max_age as u64))
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::{to_bytes, Body};
use axum::http::{header, Method, Request, StatusCode};
use axum::response::IntoResponse;
use axum::routing::get;
use tower::ServiceExt;
fn sample_router() -> Router {
Router::new().route(
"/json",
get(|| async {
axum::Json(serde_json::json!({
"data": "x".repeat(2048),
}))
}),
)
}
#[tokio::test]
async fn compression_layer_gzips_when_accepted() {
let app = apply_transport_layers(sample_router(), &TransportConfig::default_enabled());
let response = app
.oneshot(
Request::builder()
.uri("/json")
.header(header::ACCEPT_ENCODING, "gzip")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
response
.headers()
.get(header::CONTENT_ENCODING)
.map(HeaderValue::to_str)
.transpose()
.unwrap(),
Some("gzip"),
);
}
#[tokio::test]
async fn compression_layer_skipped_without_accept_encoding() {
let app = apply_transport_layers(sample_router(), &TransportConfig::default_enabled());
let response = app
.oneshot(Request::builder().uri("/json").body(Body::empty()).unwrap())
.await
.unwrap();
assert!(!response.headers().contains_key(header::CONTENT_ENCODING));
}
#[tokio::test]
async fn cors_preflight_returns_allow_headers() {
let config = TransportConfig {
compression: false,
etag: false,
cors: Some(CorsConfig {
allow_origins: vec!["https://app.example.com".into()],
..Default::default()
}),
};
let app = apply_transport_layers(sample_router(), &config);
let response = app
.oneshot(
Request::builder()
.method(Method::OPTIONS)
.uri("/json")
.header(header::ORIGIN, "https://app.example.com")
.header("access-control-request-method", "GET")
.header("access-control-request-headers", "authorization")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(response.status().is_success());
assert_eq!(
response
.headers()
.get("access-control-allow-origin")
.unwrap(),
"https://app.example.com"
);
assert!(response
.headers()
.get("access-control-allow-methods")
.unwrap()
.to_str()
.unwrap()
.contains("GET"));
}
#[tokio::test]
async fn cors_disabled_emits_no_headers() {
let app = apply_transport_layers(sample_router(), &TransportConfig::default_enabled());
let response = app
.oneshot(
Request::builder()
.method(Method::OPTIONS)
.uri("/json")
.header(header::ORIGIN, "https://app.example.com")
.header("access-control-request-method", "GET")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(!response
.headers()
.contains_key("access-control-allow-origin"));
}
#[tokio::test]
async fn cors_list_wildcard_with_credentials_does_not_panic_or_set_credentials() {
let config = TransportConfig {
compression: false,
etag: false,
cors: Some(CorsConfig {
allow_origins: vec!["*".into()],
allow_credentials: true,
..Default::default()
}),
};
let app = apply_transport_layers(sample_router(), &config);
let response = app
.oneshot(
Request::builder()
.method(Method::OPTIONS)
.uri("/json")
.header(header::ORIGIN, "https://app.example.com")
.header("access-control-request-method", "GET")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(response.status().is_success());
assert_eq!(
response
.headers()
.get("access-control-allow-origin")
.unwrap(),
"*"
);
assert!(!response
.headers()
.contains_key("access-control-allow-credentials"));
}
#[tokio::test]
async fn handler_x_compress_never_skips_compression_and_strips_marker() {
let app = apply_transport_layers(
Router::new().route(
"/raw",
get(|| async {
let mut response = axum::Json(serde_json::json!({
"data": "x".repeat(2048),
}))
.into_response();
response.headers_mut().insert(
COMPRESSION_OPT_OUT_HEADER,
HeaderValue::from_static(COMPRESSION_OPT_OUT_VALUE),
);
response
}),
),
&TransportConfig::default_enabled(),
);
let response = app
.oneshot(
Request::builder()
.uri("/raw")
.header(header::ACCEPT_ENCODING, "gzip")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
assert!(
!response.headers().contains_key(header::CONTENT_ENCODING),
"x-compress: never must skip compression",
);
assert!(
!response.headers().contains_key(COMPRESSION_OPT_OUT_HEADER),
"marker header must be stripped before flushing to client",
);
}
#[tokio::test]
async fn handler_x_compress_other_value_still_compresses() {
let app = apply_transport_layers(
Router::new().route(
"/maybe",
get(|| async {
let mut response = axum::Json(serde_json::json!({
"data": "x".repeat(2048),
}))
.into_response();
response
.headers_mut()
.insert(COMPRESSION_OPT_OUT_HEADER, HeaderValue::from_static("auto"));
response
}),
),
&TransportConfig::default_enabled(),
);
let response = app
.oneshot(
Request::builder()
.uri("/maybe")
.header(header::ACCEPT_ENCODING, "gzip")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(
response
.headers()
.get(header::CONTENT_ENCODING)
.map(HeaderValue::to_str)
.transpose()
.unwrap(),
Some("gzip"),
);
assert!(!response.headers().contains_key(COMPRESSION_OPT_OUT_HEADER));
}
#[tokio::test]
async fn pipeline_yields_uncompressed_body_when_no_accept_encoding() {
let app = apply_transport_layers(sample_router(), &TransportConfig::default_enabled());
let response = app
.oneshot(Request::builder().uri("/json").body(Body::empty()).unwrap())
.await
.unwrap();
let bytes = to_bytes(response.into_body(), usize::MAX).await.unwrap();
let parsed: serde_json::Value = serde_json::from_slice(&bytes).unwrap();
assert!(parsed["data"].as_str().unwrap().starts_with("xxxx"));
}
}