use std::collections::HashMap;
use std::sync::RwLock;
use std::sync::atomic::{AtomicU64, Ordering};
use super::bucket::TokenBucket;
use super::limiter::LoginRateLimitOutcome;
const MAX_TRACKED_BUCKETS: usize = 100_000;
pub(super) struct LoginLimiter {
buckets: RwLock<HashMap<String, TokenBucket>>,
ip_cap: AtomicU64,
user_cap: AtomicU64,
}
impl LoginLimiter {
pub(super) fn new() -> Self {
Self {
buckets: RwLock::new(HashMap::new()),
ip_cap: AtomicU64::new(30),
user_cap: AtomicU64::new(10),
}
}
pub(super) fn set_capacities(&self, ip_cap: u64, user_cap: u64) {
self.ip_cap.store(ip_cap, Ordering::Relaxed);
self.user_cap.store(user_cap, Ordering::Relaxed);
}
pub(super) fn capacities(&self) -> (u64, u64) {
(
self.ip_cap.load(Ordering::Relaxed),
self.user_cap.load(Ordering::Relaxed),
)
}
pub(super) fn check(&self, peer_addr: &str, username: &str) -> LoginRateLimitOutcome {
let ip_cap = self.ip_cap.load(Ordering::Relaxed);
if ip_cap == 0 {
return LoginRateLimitOutcome::Allowed;
}
let user_cap = self.user_cap.load(Ordering::Relaxed);
if let Some(retry_after_secs) = self.peek_empty(&format!("login_fail_ip:{peer_addr}")) {
return LoginRateLimitOutcome::IpExceeded { retry_after_secs };
}
if user_cap > 0
&& !username.is_empty()
&& let Some(retry_after_secs) = self.peek_empty(&format!("login_fail_user:{username}"))
{
return LoginRateLimitOutcome::UserExceeded { retry_after_secs };
}
let dos_cap = Self::dos_ceiling(ip_cap);
let dos_rate = (dos_cap as f64) / 60.0;
if let Some(retry_after_secs) =
self.consume(&format!("login_dos_ip:{peer_addr}"), dos_cap, dos_rate)
{
return LoginRateLimitOutcome::IpExceeded { retry_after_secs };
}
LoginRateLimitOutcome::Allowed
}
pub(super) fn record_failure(&self, peer_addr: &str, username: &str) {
let ip_cap = self.ip_cap.load(Ordering::Relaxed);
if ip_cap == 0 {
return;
}
let ip_rate = (ip_cap as f64) / 60.0;
let _ = self.consume(&format!("login_fail_ip:{peer_addr}"), ip_cap, ip_rate);
let user_cap = self.user_cap.load(Ordering::Relaxed);
if user_cap > 0 && !username.is_empty() {
let user_rate = (user_cap as f64) / 60.0;
let _ = self.consume(&format!("login_fail_user:{username}"), user_cap, user_rate);
}
}
pub(super) fn is_rate_limited(&self, peer_addr: &str, username: &str) -> bool {
let ip_cap = self.ip_cap.load(Ordering::Relaxed);
if ip_cap == 0 {
return false;
}
if self
.peek_empty(&format!("login_fail_ip:{peer_addr}"))
.is_some()
{
return true;
}
let user_cap = self.user_cap.load(Ordering::Relaxed);
if user_cap > 0
&& !username.is_empty()
&& self
.peek_empty(&format!("login_fail_user:{username}"))
.is_some()
{
return true;
}
self.peek_empty(&format!("login_dos_ip:{peer_addr}"))
.is_some()
}
pub(super) fn active_buckets(&self) -> usize {
self.buckets.read().unwrap_or_else(|p| p.into_inner()).len()
}
fn dos_ceiling(ip_cap: u64) -> u64 {
ip_cap.saturating_mul(4).max(120)
}
fn peek_empty(&self, key: &str) -> Option<u64> {
let buckets = self.buckets.read().unwrap_or_else(|p| p.into_inner());
let bucket = buckets.get(key)?;
if bucket.available() < 1 {
Some((bucket.retry_after_ms() / 1000).max(1))
} else {
None
}
}
fn consume(&self, key: &str, capacity: u64, rate_per_sec: f64) -> Option<u64> {
{
let buckets = self.buckets.read().unwrap_or_else(|p| p.into_inner());
if let Some(bucket) = buckets.get(key) {
return if bucket.try_acquire(1) {
None
} else {
Some((bucket.retry_after_ms() / 1000).max(1))
};
}
}
let mut buckets = self.buckets.write().unwrap_or_else(|p| p.into_inner());
if buckets.len() > MAX_TRACKED_BUCKETS {
buckets.retain(|_, b| b.available() < b.capacity());
}
let bucket = buckets
.entry(key.to_string())
.or_insert_with(|| TokenBucket::new(capacity, rate_per_sec));
if bucket.try_acquire(1) {
None
} else {
Some((bucket.retry_after_ms() / 1000).max(1))
}
}
}
impl Default for LoginLimiter {
fn default() -> Self {
Self::new()
}
}