use std::time::Duration;
use axum::error_handling::HandleErrorLayer;
use axum::extract::DefaultBodyLimit;
use axum::http::StatusCode;
use axum::Router;
use tower::limit::GlobalConcurrencyLimitLayer;
use tower::load_shed::LoadShedLayer;
use tower::{BoxError, ServiceBuilder};
use tower_http::timeout::TimeoutLayer;
pub const DEFAULT_DRAIN_TIMEOUT: Duration = Duration::from_secs(25);
pub const DEFAULT_REQUEST_TIMEOUT: Duration = Duration::from_secs(20);
pub const DEFAULT_MAX_CONCURRENT_REQUESTS: usize = 1024;
pub const DEFAULT_MAX_CONNECTIONS: usize = 2048;
#[derive(Debug, Clone)]
pub struct ServeHygieneConfig {
pub drain_timeout: Duration,
pub request_timeout: Option<Duration>,
pub max_concurrent_requests: Option<usize>,
pub max_connections: Option<usize>,
pub max_body_bytes: Option<usize>,
}
impl Default for ServeHygieneConfig {
fn default() -> Self {
Self {
drain_timeout: DEFAULT_DRAIN_TIMEOUT,
request_timeout: Some(DEFAULT_REQUEST_TIMEOUT),
max_concurrent_requests: Some(DEFAULT_MAX_CONCURRENT_REQUESTS),
max_connections: Some(DEFAULT_MAX_CONNECTIONS),
max_body_bytes: None,
}
}
}
async fn shed_to_status(error: BoxError) -> StatusCode {
if error.is::<tower::load_shed::error::Overloaded>() {
StatusCode::SERVICE_UNAVAILABLE
} else {
StatusCode::INTERNAL_SERVER_ERROR
}
}
pub fn apply_server_hygiene(mut router: Router, config: &ServeHygieneConfig) -> Router {
if let Some(limit) = config.max_body_bytes {
router = router.layer(DefaultBodyLimit::max(limit));
}
if let Some(timeout) = config.request_timeout {
router = router.layer(TimeoutLayer::with_status_code(
StatusCode::REQUEST_TIMEOUT,
timeout,
));
}
if let Some(max) = config.max_concurrent_requests {
router = router.layer(
ServiceBuilder::new()
.layer(HandleErrorLayer::new(shed_to_status))
.layer(LoadShedLayer::new())
.layer(GlobalConcurrencyLimitLayer::new(max)),
);
}
router
}