use axum::{body::Body, http::StatusCode, response::Response};
use governor::middleware::NoOpMiddleware;
use std::sync::Arc;
use tower_governor::governor::GovernorConfigBuilder;
pub use tower_governor::key_extractor::SmartIpKeyExtractor;
pub use tower_governor::GovernorLayer;
use crate::config::RateLimitConfig;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RateLimitType {
Validate,
Heartbeat,
Bind,
}
pub fn create_rate_limiter(
config: &RateLimitConfig,
limit_type: RateLimitType,
) -> GovernorLayer<SmartIpKeyExtractor, NoOpMiddleware> {
let rpm = match limit_type {
RateLimitType::Validate => config.validate_rpm,
RateLimitType::Heartbeat => config.heartbeat_rpm,
RateLimitType::Bind => config.bind_rpm,
};
let interval_ms = if rpm > 0 { 60_000 / rpm } else { 60_000 };
let governor_config = GovernorConfigBuilder::default()
.per_millisecond(interval_ms.into())
.burst_size(config.burst_size)
.key_extractor(SmartIpKeyExtractor)
.finish()
.expect("failed to build governor config");
GovernorLayer {
config: Arc::new(governor_config),
}
}
pub fn rate_limit_error_response(retry_after_secs: u64) -> Response<Body> {
let retry_after = retry_after_secs.max(1);
let body = serde_json::json!({
"error": "Too many requests",
"message": format!("Rate limit exceeded. Please retry after {} seconds.", retry_after),
"retry_after_seconds": retry_after
});
Response::builder()
.status(StatusCode::TOO_MANY_REQUESTS)
.header("Content-Type", "application/json")
.header("Retry-After", retry_after.to_string())
.body(Body::from(serde_json::to_string(&body).unwrap()))
.unwrap()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rate_limit_config_defaults() {
let config = RateLimitConfig::default();
assert!(config.enabled);
assert_eq!(config.validate_rpm, 100);
assert_eq!(config.heartbeat_rpm, 60);
assert_eq!(config.bind_rpm, 10);
assert_eq!(config.burst_size, 5);
}
#[test]
fn rate_limit_error_response_format() {
let response = rate_limit_error_response(30);
assert_eq!(response.status(), StatusCode::TOO_MANY_REQUESTS);
assert_eq!(
response
.headers()
.get("Retry-After")
.unwrap()
.to_str()
.unwrap(),
"30"
);
}
#[test]
fn rate_limit_error_minimum_retry_after() {
let response = rate_limit_error_response(0);
assert_eq!(
response
.headers()
.get("Retry-After")
.unwrap()
.to_str()
.unwrap(),
"1"
);
}
#[test]
fn create_rate_limiter_validate() {
let config = RateLimitConfig::default();
let _layer = create_rate_limiter(&config, RateLimitType::Validate);
}
#[test]
fn create_rate_limiter_heartbeat() {
let config = RateLimitConfig::default();
let _layer = create_rate_limiter(&config, RateLimitType::Heartbeat);
}
#[test]
fn create_rate_limiter_bind() {
let config = RateLimitConfig::default();
let _layer = create_rate_limiter(&config, RateLimitType::Bind);
}
}