use std::convert::Infallible;
use std::time::Duration;
use axum::http::StatusCode;
use axum::response::Response;
use axum::error_handling::HandleErrorLayer;
use axum::Router;
use tower::{BoxError, ServiceBuilder};
pub use tower_resilience_bulkhead::BulkheadLayer;
pub use tower_resilience_circuitbreaker::{
CircuitBreakerConfigBuilder, CircuitBreakerLayer, CircuitState,
};
use tower_resilience_circuitbreaker::classifier::FnClassifier;
pub type HttpFailureClassifier = FnClassifier<HttpClassifierFn>;
pub type HttpClassifierFn = fn(&Result<Response, Infallible>) -> bool;
fn is_server_error(result: &Result<Response, Infallible>) -> bool {
result
.as_ref()
.is_ok_and(|response| response.status().is_server_error())
}
#[derive(Debug, Clone)]
pub struct ResilienceConfig {
pub circuit_breaker_enabled: bool,
pub circuit_breaker_threshold: f64,
pub circuit_breaker_min_requests: u64,
pub circuit_breaker_wait_duration: Duration,
pub bulkhead_enabled: bool,
pub bulkhead_max_concurrent: usize,
pub bulkhead_max_wait: Duration,
}
impl Default for ResilienceConfig {
fn default() -> Self {
Self {
circuit_breaker_enabled: true,
circuit_breaker_threshold: 0.5, circuit_breaker_min_requests: 10,
circuit_breaker_wait_duration: Duration::from_secs(30),
bulkhead_enabled: true,
bulkhead_max_concurrent: 100,
bulkhead_max_wait: Duration::from_secs(5),
}
}
}
impl From<&crate::config::ResilienceConfig> for ResilienceConfig {
fn from(config: &crate::config::ResilienceConfig) -> Self {
Self {
circuit_breaker_enabled: config.circuit_breaker_enabled,
circuit_breaker_threshold: config.circuit_breaker_threshold,
circuit_breaker_min_requests: config.circuit_breaker_min_requests,
circuit_breaker_wait_duration: config.circuit_breaker_wait_duration(),
bulkhead_enabled: config.bulkhead_enabled,
bulkhead_max_concurrent: config.bulkhead_max_concurrent,
bulkhead_max_wait: config.bulkhead_max_wait(),
}
}
}
impl ResilienceConfig {
pub fn new() -> Self {
Self::default()
}
pub fn with_circuit_breaker(mut self, enabled: bool) -> Self {
self.circuit_breaker_enabled = enabled;
self
}
pub fn with_circuit_breaker_threshold(mut self, threshold: f64) -> Self {
self.circuit_breaker_threshold = threshold.clamp(0.0, 1.0);
self
}
pub fn with_bulkhead(mut self, enabled: bool) -> Self {
self.bulkhead_enabled = enabled;
self
}
pub fn with_bulkhead_max_concurrent(mut self, max: usize) -> Self {
self.bulkhead_max_concurrent = max;
self
}
pub fn circuit_breaker_layer(&self) -> Option<CircuitBreakerLayer> {
if !self.circuit_breaker_enabled {
return None;
}
Some(
self.circuit_breaker_builder()
.on_state_transition(Self::log_state_transition)
.build_with_handle()
.0,
)
}
pub fn http_circuit_breaker_layer(
&self,
) -> Option<CircuitBreakerLayer<HttpFailureClassifier>> {
if !self.circuit_breaker_enabled {
return None;
}
Some(
self.circuit_breaker_builder()
.on_state_transition(Self::log_state_transition)
.failure_classifier(is_server_error as HttpClassifierFn)
.build_with_handle()
.0,
)
}
fn circuit_breaker_builder(&self) -> CircuitBreakerConfigBuilder {
CircuitBreakerLayer::builder()
.name("acton-circuit-breaker")
.failure_rate_threshold(self.circuit_breaker_threshold)
.sliding_window_size(self.circuit_breaker_min_requests as usize)
.wait_duration_in_open(self.circuit_breaker_wait_duration)
}
fn log_state_transition(from: CircuitState, to: CircuitState) {
tracing::warn!(
from = ?from,
to = ?to,
"Circuit breaker state transition"
);
}
pub fn bulkhead_layer(&self) -> Option<BulkheadLayer> {
if !self.bulkhead_enabled {
return None;
}
Some(
BulkheadLayer::builder()
.name("acton-bulkhead")
.max_concurrent_calls(self.bulkhead_max_concurrent)
.max_wait_duration(self.bulkhead_max_wait)
.on_call_permitted(|concurrent| {
tracing::debug!(
concurrent_requests = concurrent,
"Request permitted through bulkhead"
);
})
.on_call_rejected(|max| {
tracing::warn!(
max_concurrent = max,
"Request rejected by bulkhead - max concurrent limit reached"
);
})
.build_with_handle()
.0,
)
}
}
async fn handle_resilience_error<E: Into<BoxError>>(error: E) -> (StatusCode, &'static str) {
let error: BoxError = error.into();
tracing::warn!(error = %error, "Request rejected by resilience middleware");
(
StatusCode::SERVICE_UNAVAILABLE,
"Service temporarily unavailable",
)
}
pub fn apply_resilience(app: Router, config: &ResilienceConfig) -> Router {
let app = match config.bulkhead_layer() {
Some(bulkhead) => app.layer(
ServiceBuilder::new()
.layer(HandleErrorLayer::new(handle_resilience_error))
.layer(bulkhead),
),
None => app,
};
match config.http_circuit_breaker_layer() {
Some(circuit_breaker) => app.layer(
ServiceBuilder::new()
.layer(HandleErrorLayer::new(handle_resilience_error))
.layer(circuit_breaker),
),
None => app,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_config() {
let config = ResilienceConfig::default();
assert!(config.circuit_breaker_enabled);
assert!(config.bulkhead_enabled);
assert_eq!(config.circuit_breaker_threshold, 0.5);
assert_eq!(config.bulkhead_max_concurrent, 100);
}
#[test]
fn test_builder_pattern() {
let config = ResilienceConfig::new()
.with_circuit_breaker(false)
.with_bulkhead_max_concurrent(50);
assert!(!config.circuit_breaker_enabled);
assert_eq!(config.bulkhead_max_concurrent, 50);
}
#[test]
fn test_threshold_clamping() {
let config = ResilienceConfig::new().with_circuit_breaker_threshold(1.5);
assert_eq!(config.circuit_breaker_threshold, 1.0);
let config = ResilienceConfig::new().with_circuit_breaker_threshold(-0.5);
assert_eq!(config.circuit_breaker_threshold, 0.0);
}
#[test]
fn test_circuit_breaker_layer_creation() {
let config = ResilienceConfig::new().with_circuit_breaker(true);
assert!(config.circuit_breaker_layer().is_some());
assert!(config.http_circuit_breaker_layer().is_some());
let config = ResilienceConfig::new().with_circuit_breaker(false);
assert!(config.circuit_breaker_layer().is_none());
assert!(config.http_circuit_breaker_layer().is_none());
}
#[test]
fn test_bulkhead_layer_creation() {
let config = ResilienceConfig::new().with_bulkhead(true);
assert!(config.bulkhead_layer().is_some());
let config = ResilienceConfig::new().with_bulkhead(false);
assert!(config.bulkhead_layer().is_none());
}
#[test]
fn classifier_counts_5xx_as_failure_not_2xx() {
let ok = Ok(Response::new(axum::body::Body::empty()));
assert!(!is_server_error(&ok));
let mut server_error = Response::new(axum::body::Body::empty());
*server_error.status_mut() = StatusCode::INTERNAL_SERVER_ERROR;
assert!(is_server_error(&Ok(server_error)));
let mut bad_request = Response::new(axum::body::Body::empty());
*bad_request.status_mut() = StatusCode::BAD_REQUEST;
assert!(!is_server_error(&Ok(bad_request)));
}
}