use std::sync::Arc;
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use dashmap::mapref::entry::Entry;
use super::window;
fn epoch_key(key: &str, epoch: u64) -> String {
format!("sp:rl:{key}:{epoch}")
}
fn unix_now() -> Duration {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or(Duration::ZERO)
}
const KEY_TTL: Duration = Duration::from_secs(600);
const MAX_KEYS: usize = 500_000;
#[derive(Debug, Clone)]
struct KeyState {
epoch: u64,
window_secs: u64,
local_count: u64,
carryover: u64,
carryover_epoch: u64,
remote_estimate: u64,
estimate_at: Instant,
last_seen: Instant,
seen_since_tick: bool,
}
impl KeyState {
fn new(epoch: u64, window_secs: u64) -> Self {
let now = Instant::now();
Self {
epoch,
window_secs,
local_count: 0,
carryover: 0,
carryover_epoch: 0,
remote_estimate: 0,
estimate_at: now,
last_seen: now,
seen_since_tick: true,
}
}
fn roll_to(&mut self, epoch: u64) {
if self.epoch != epoch {
if self.local_count > 0 {
if self.carryover > 0 && self.carryover_epoch != self.epoch {
self.carryover = self.local_count;
} else {
self.carryover += self.local_count;
}
self.carryover_epoch = self.epoch;
}
self.epoch = epoch;
self.local_count = 0;
}
}
}
pub struct GlobalCounters {
client: redis::Client,
conn: tokio::sync::OnceCell<redis::aio::ConnectionManager>,
interval: Duration,
states: dashmap::DashMap<String, KeyState>,
}
impl GlobalCounters {
pub fn build(url: &str, interval: Duration) -> Result<Arc<Self>, String> {
let client = redis::Client::open(url)
.map_err(|e| format!("invalid shared-store URL for rate limiting: {e}"))?;
Ok(Arc::new(Self {
client,
conn: tokio::sync::OnceCell::new(),
interval: interval.max(Duration::from_millis(50)),
states: dashmap::DashMap::new(),
}))
}
pub fn fleet_remaining(&self, key: &str, budget: u64, window: Duration) -> u64 {
let Some(state) = self.state_for(key, window) else {
return budget;
};
let stale_after = (self.interval * 4).max(Duration::from_secs(2));
let remote = if state.estimate_at.elapsed() < stale_after {
state.remote_estimate
} else {
0
};
let used = remote + state.local_count + state.carryover;
budget.saturating_sub(used)
}
pub fn record(&self, key: &str, window: Duration) {
if let Some(mut state) = self.state_for(key, window) {
state.local_count += 1;
}
}
fn state_for(
&self,
key: &str,
window: Duration,
) -> Option<dashmap::mapref::one::RefMut<'_, String, KeyState>> {
let now = unix_now();
let window_secs = window.as_secs().max(1);
let ep = window::epoch(now, window);
let over_cap = self.states.len() >= MAX_KEYS;
let mut state = match self.states.entry(key.to_string()) {
Entry::Occupied(o) => o.into_ref(),
Entry::Vacant(_) if over_cap => return None,
Entry::Vacant(v) => v.insert(KeyState::new(ep, window_secs)),
};
if state.window_secs != window_secs {
*state = KeyState::new(ep, window_secs);
}
state.roll_to(ep);
state.last_seen = Instant::now();
state.seen_since_tick = true;
Some(state)
}
#[must_use]
pub fn spawn(self: &Arc<Self>) -> bool {
if tokio::runtime::Handle::try_current().is_err() {
tracing::warn!("rate-limit reconciler not started: no Tokio runtime in this context");
return false;
}
let this = self.clone();
tokio::spawn(async move {
let mut ticker = tokio::time::interval(this.interval);
loop {
ticker.tick().await;
this.reconcile().await;
}
});
true
}
async fn connection(&self) -> redis::RedisResult<redis::aio::ConnectionManager> {
self.conn
.get_or_try_init(|| redis::aio::ConnectionManager::new(self.client.clone()))
.await
.cloned()
}
async fn reconcile(&self) {
self.reconcile_at(unix_now()).await;
}
async fn reconcile_at(&self, now: Duration) {
self.evict_stale();
let plans = self.claim_plans(now);
if plans.is_empty() {
return;
}
let mut conn = match self.connection().await {
Ok(c) => c,
Err(e) => {
tracing::warn!("rate-limit shared store unavailable, staying local: {e}");
return;
}
};
if let Err(e) = self.push_deltas(&mut conn, &plans).await {
tracing::warn!("rate-limit delta push failed, dropping this interval: {e}");
return;
}
let reads = match self.read_epochs(&mut conn, &plans).await {
Ok(r) => r,
Err(e) => {
tracing::warn!("rate-limit aggregate read failed: {e}");
return;
}
};
self.apply_estimates(now, &plans, &reads);
}
fn claim_plans(&self, now: Duration) -> Vec<PushPlan> {
let keys: Vec<String> = self.states.iter().map(|e| e.key().clone()).collect();
let mut plans: Vec<PushPlan> = Vec::new();
for key in keys {
if let Some(mut s) = self.states.get_mut(&key) {
let active = s.local_count > 0 || s.carryover > 0 || s.seen_since_tick;
s.seen_since_tick = false;
if !active {
continue;
}
let window = Duration::from_secs(s.window_secs);
let read_epoch = window::epoch(now, window);
let claim = s.local_count;
s.local_count = 0;
plans.push(PushPlan {
key: key.clone(),
push_epoch: s.epoch,
read_epoch,
window,
delta: claim,
is_carryover: false,
});
if s.carryover > 0 {
let carry = s.carryover;
s.carryover = 0;
plans.push(PushPlan {
key: key.clone(),
push_epoch: s.carryover_epoch,
read_epoch,
window,
delta: carry,
is_carryover: true,
});
}
}
}
plans
}
fn apply_estimates(&self, now: Duration, plans: &[PushPlan], reads: &[(u64, u64)]) {
for (p, (cur, prev)) in plans.iter().zip(reads) {
if p.is_carryover {
continue;
}
if let Some(mut s) = self.states.get_mut(&p.key) {
if s.window_secs != p.window.as_secs() || s.epoch > p.read_epoch {
continue;
}
let elapsed = window::elapsed_in_window(now, p.window);
let est = window::sliding_estimate(*cur, *prev, elapsed, p.window);
s.remote_estimate = est.ceil().max(0.0) as u64;
s.estimate_at = Instant::now();
}
}
}
async fn push_deltas(
&self,
conn: &mut redis::aio::ConnectionManager,
plans: &[PushPlan],
) -> redis::RedisResult<()> {
let mut pipe = redis::pipe();
pipe.atomic();
let mut any = false;
for p in plans.iter().filter(|p| p.delta > 0) {
any = true;
let k = epoch_key(&p.key, p.push_epoch);
let ttl_ms = (p.window.as_millis() as u64).saturating_mul(2).max(1);
pipe.cmd("INCRBY").arg(&k).arg(p.delta).ignore();
pipe.cmd("PEXPIRE").arg(&k).arg(ttl_ms).ignore();
}
if !any {
return Ok(());
}
pipe.query_async(conn).await
}
async fn read_epochs(
&self,
conn: &mut redis::aio::ConnectionManager,
plans: &[PushPlan],
) -> redis::RedisResult<Vec<(u64, u64)>> {
let mut keys: Vec<String> = Vec::with_capacity(plans.len() * 2);
for p in plans {
keys.push(epoch_key(&p.key, p.read_epoch));
keys.push(epoch_key(&p.key, p.read_epoch.saturating_sub(1)));
}
let vals: Vec<Option<i64>> = redis::cmd("MGET").arg(&keys).query_async(conn).await?;
Ok(plans
.iter()
.enumerate()
.map(|(i, _)| {
let cur = vals.get(i * 2).copied().flatten().unwrap_or(0).max(0) as u64;
let prev = vals.get(i * 2 + 1).copied().flatten().unwrap_or(0).max(0) as u64;
(cur, prev)
})
.collect())
}
fn evict_stale(&self) {
let now = Instant::now();
self.states.retain(|_, s| {
if s.local_count > 0 || s.carryover > 0 {
return true;
}
let threshold = KEY_TTL.max(Duration::from_secs(s.window_secs.saturating_mul(2)));
now.duration_since(s.last_seen) < threshold
});
}
}
struct PushPlan {
key: String,
push_epoch: u64,
read_epoch: u64,
window: Duration,
delta: u64,
is_carryover: bool,
}
#[cfg(test)]
mod tests {
use super::*;
const W: Duration = Duration::from_secs(60);
fn counters() -> Arc<GlobalCounters> {
GlobalCounters::build("redis://127.0.0.1/", Duration::from_millis(500)).unwrap()
}
#[test]
fn budget_admits_until_local_count_reaches_it() {
let g = counters();
for _ in 0..3 {
assert!(g.fleet_remaining("k", 3, W) > 0);
g.record("k", W);
}
assert_eq!(g.fleet_remaining("k", 3, W), 0);
}
#[test]
fn budget_accounts_for_remote_estimate() {
let g = counters();
let now = unix_now();
let ep = window::epoch(now, W);
let mut state = KeyState::new(ep, W.as_secs());
state.remote_estimate = 2;
g.states.insert("k".to_string(), state);
assert!(g.fleet_remaining("k", 3, W) > 0);
g.record("k", W);
assert_eq!(g.fleet_remaining("k", 3, W), 0);
}
#[test]
fn idle_zero_delta_keys_are_not_refreshed_each_tick() {
use std::collections::HashSet;
let g = counters();
let now = unix_now();
g.record("delta", W);
let _ = g.fleet_remaining("gated", 100, W);
g.record("idle", W);
let _ = g.claim_plans(now);
g.record("delta", W);
let _ = g.fleet_remaining("gated", 100, W);
let plans = g.claim_plans(now);
let keys: HashSet<&str> = plans.iter().map(|p| p.key.as_str()).collect();
assert!(
keys.contains("delta"),
"a key with a pending delta must be pushed"
);
assert!(
keys.contains("gated"),
"an actively-gated key must be refreshed"
);
assert!(
!keys.contains("idle"),
"an idle zero-delta key must not be refreshed every tick"
);
}
#[test]
fn roll_preserves_unclaimed_deltas_as_carryover() {
let mut s = KeyState::new(10, 60);
s.local_count = 5;
s.roll_to(11);
assert_eq!(s.carryover, 5);
assert_eq!(s.carryover_epoch, 10);
assert_eq!(s.local_count, 0);
assert_eq!(s.epoch, 11);
}
#[test]
fn consecutive_rolls_do_not_collapse_epochs() {
let mut s = KeyState::new(10, 60);
s.local_count = 3;
s.roll_to(11); s.local_count = 4;
s.roll_to(12); assert_eq!(
s.carryover, 4,
"keeps the latest window, not 3 + 4 collapsed"
);
assert_eq!(s.carryover_epoch, 11);
}
#[test]
fn stale_estimate_not_applied_after_window_reset() {
let g = counters();
let key = "k";
g.record(key, Duration::from_secs(120));
g.states.get_mut(key).unwrap().remote_estimate = 5;
let now = unix_now();
let plan = PushPlan {
key: key.to_string(),
push_epoch: window::epoch(now, W),
read_epoch: window::epoch(now, W),
window: W,
delta: 0,
is_carryover: false,
};
g.apply_estimates(now, &[plan], &[(100, 0)]);
assert_eq!(
g.states.get(key).unwrap().remote_estimate,
5,
"an old-window estimate must not overwrite the reset state"
);
}
#[test]
fn independent_keys_have_independent_budgets() {
let g = counters();
assert!(g.fleet_remaining("a", 1, W) > 0);
g.record("a", W);
assert_eq!(g.fleet_remaining("a", 1, W), 0);
assert!(g.fleet_remaining("b", 1, W) > 0);
}
#[tokio::test]
async fn reconciles_across_instances() {
let Ok(url) = std::env::var("SHIELD_REDIS_TEST_URL") else {
eprintln!("SKIP reconciles_across_instances: SHIELD_REDIS_TEST_URL not set");
return;
};
let nonce = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos();
let key = format!("it:{}:{nonce}", std::process::id());
let a = GlobalCounters::build(&url, Duration::from_millis(200)).unwrap();
let b = GlobalCounters::build(&url, Duration::from_millis(200)).unwrap();
a.record(&key, W);
a.record(&key, W);
a.reconcile().await;
assert!(
b.fleet_remaining(&key, 3, W) > 0,
"B admits its first request for the key"
);
b.record(&key, W);
b.reconcile().await;
assert_eq!(
b.fleet_remaining(&key, 3, W),
0,
"B must reject once the combined budget is reached"
);
}
#[tokio::test]
async fn carryover_is_pushed_to_its_epoch() {
let Ok(url) = std::env::var("SHIELD_REDIS_TEST_URL") else {
eprintln!("SKIP carryover_is_pushed_to_its_epoch: SHIELD_REDIS_TEST_URL not set");
return;
};
let win = Duration::from_secs(3600);
let nonce = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos();
let key = format!("it3:{}:{nonce}", std::process::id());
let g = GlobalCounters::build(&url, Duration::from_millis(200)).unwrap();
let cur = window::epoch(unix_now(), win);
let mut st = KeyState::new(cur, win.as_secs());
st.carryover = 3;
st.carryover_epoch = cur - 1;
g.states.insert(key.clone(), st);
g.reconcile().await;
let client = redis::Client::open(url).unwrap();
let mut conn = client.get_multiplexed_async_connection().await.unwrap();
let prev: i64 = redis::cmd("GET")
.arg(epoch_key(&key, cur - 1))
.query_async(&mut conn)
.await
.unwrap_or(0);
assert_eq!(prev, 3, "carryover must be published to its own epoch");
assert_eq!(
g.states.get(&key).unwrap().carryover,
0,
"carryover must be cleared once published"
);
}
#[tokio::test]
async fn estimate_refreshes_without_local_deltas() {
let Ok(url) = std::env::var("SHIELD_REDIS_TEST_URL") else {
eprintln!(
"SKIP estimate_refreshes_without_local_deltas: SHIELD_REDIS_TEST_URL not set"
);
return;
};
let win = Duration::from_secs(3600);
let nonce = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos();
let key = format!("it4:{}:{nonce}", std::process::id());
let g = GlobalCounters::build(&url, Duration::from_millis(200)).unwrap();
let cur = window::epoch(unix_now(), win);
let mut st = KeyState::new(cur, win.as_secs());
st.remote_estimate = 99;
g.states.insert(key.clone(), st);
g.reconcile().await;
assert_eq!(
g.states.get(&key).unwrap().remote_estimate,
0,
"estimate must be refreshed even with no local deltas"
);
}
#[tokio::test]
async fn repeated_reconcile_does_not_double_push() {
let Ok(url) = std::env::var("SHIELD_REDIS_TEST_URL") else {
eprintln!(
"SKIP repeated_reconcile_does_not_double_push: SHIELD_REDIS_TEST_URL not set"
);
return;
};
let win = Duration::from_secs(3600);
let nonce = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos();
let key = format!("it2:{}:{nonce}", std::process::id());
let a = GlobalCounters::build(&url, Duration::from_millis(200)).unwrap();
a.record(&key, win);
a.record(&key, win);
a.reconcile().await;
a.reconcile().await;
a.reconcile().await;
let epoch = window::epoch(unix_now(), win);
let redis_key = epoch_key(&key, epoch);
let client = redis::Client::open(url).unwrap();
let mut conn = client.get_multiplexed_async_connection().await.unwrap();
let count: i64 = redis::cmd("GET")
.arg(&redis_key)
.query_async(&mut conn)
.await
.unwrap_or(0);
assert_eq!(
count, 2,
"shared counter must reflect the 2 admits exactly once"
);
}
}