use std::hash::Hash;
use std::net::SocketAddr;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::time::Duration;
use dashmap::DashMap;
use freenet_stdlib::prelude::ContractInstanceId;
use tokio::time::Instant;
use crate::util::time_source::TimeSource;
pub(crate) const EMIT_REFILL_INTERVAL: Duration = Duration::from_secs(60);
pub(crate) const EMIT_BURST: f64 = 2.0;
pub(crate) const RESPOND_REFILL_INTERVAL: Duration = Duration::from_secs(30);
pub(crate) const RESPOND_BURST: f64 = 1.0;
pub(crate) const RESPOND_GLOBAL_REFILL_INTERVAL: Duration = Duration::from_secs(5);
pub(crate) const RESPOND_GLOBAL_BURST: f64 = 2.0;
pub(crate) const MAX_TRACKED_KEYS: usize = 16_384;
const CLEANUP_AGE: Duration = Duration::from_secs(5 * 60);
struct Bucket {
tokens: f64,
last_refill: Instant,
}
impl Bucket {
fn refill(&mut self, now: Instant, capacity: f64, interval: Duration) {
let elapsed = now.saturating_duration_since(self.last_refill);
let gained = elapsed.as_secs_f64() / interval.as_secs_f64();
self.tokens = (self.tokens + gained).min(capacity);
self.last_refill = now;
}
}
pub(crate) struct TokenBucketLimiter<K: Eq + Hash + Clone> {
buckets: DashMap<K, Bucket>,
size: AtomicUsize,
max_tracked: usize,
capacity: f64,
refill_interval: Duration,
time_source: Arc<dyn TimeSource + Send + Sync>,
allowed_total: AtomicU64,
suppressed_total: AtomicU64,
}
impl<K: Eq + Hash + Clone> TokenBucketLimiter<K> {
pub fn new(
time_source: Arc<dyn TimeSource + Send + Sync>,
capacity: f64,
refill_interval: Duration,
max_tracked: usize,
) -> Self {
Self {
buckets: DashMap::new(),
size: AtomicUsize::new(0),
max_tracked,
capacity,
refill_interval,
time_source,
allowed_total: AtomicU64::new(0),
suppressed_total: AtomicU64::new(0),
}
}
pub fn check_and_record(&self, key: K) -> bool {
let now = self.time_source.now();
use dashmap::mapref::entry::Entry;
match self.buckets.entry(key) {
Entry::Occupied(mut occ) => {
let bucket = occ.get_mut();
bucket.refill(now, self.capacity, self.refill_interval);
if bucket.tokens >= 1.0 {
bucket.tokens -= 1.0;
self.allowed_total.fetch_add(1, Ordering::Relaxed);
true
} else {
self.suppressed_total.fetch_add(1, Ordering::Relaxed);
false
}
}
Entry::Vacant(vac) => {
let prev = self.size.fetch_add(1, Ordering::Relaxed);
if prev >= self.max_tracked {
self.size.fetch_sub(1, Ordering::Relaxed);
self.suppressed_total.fetch_add(1, Ordering::Relaxed);
return false;
}
vac.insert(Bucket {
tokens: self.capacity - 1.0,
last_refill: now,
});
self.allowed_total.fetch_add(1, Ordering::Relaxed);
true
}
}
}
pub fn cleanup(&self) {
let now = self.time_source.now();
let capacity = self.capacity;
let interval = self.refill_interval;
let mut removed = 0usize;
self.buckets.retain(|_, bucket| {
let elapsed = now.saturating_duration_since(bucket.last_refill);
let gained = elapsed.as_secs_f64() / interval.as_secs_f64();
let tokens_now = (bucket.tokens + gained).min(capacity);
let recovered = tokens_now >= capacity;
let idle = elapsed > CLEANUP_AGE;
if recovered && idle {
removed += 1;
false
} else {
true
}
});
if removed > 0 {
self.size.fetch_sub(removed, Ordering::Relaxed);
}
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn allowed_total(&self) -> u64 {
self.allowed_total.load(Ordering::Relaxed)
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn suppressed_total(&self) -> u64 {
self.suppressed_total.load(Ordering::Relaxed)
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn len(&self) -> usize {
self.buckets.len()
}
}
pub(crate) type ResyncEmitLimiter = TokenBucketLimiter<ContractInstanceId>;
pub(crate) type ResyncResponseLimiter = TokenBucketLimiter<(SocketAddr, ContractInstanceId)>;
pub(crate) type ResyncResponseGlobalLimiter = TokenBucketLimiter<ContractInstanceId>;
pub(crate) fn new_emit_limiter(
time_source: Arc<dyn TimeSource + Send + Sync>,
) -> ResyncEmitLimiter {
TokenBucketLimiter::new(
time_source,
EMIT_BURST,
EMIT_REFILL_INTERVAL,
MAX_TRACKED_KEYS,
)
}
pub(crate) fn new_response_limiter(
time_source: Arc<dyn TimeSource + Send + Sync>,
) -> ResyncResponseLimiter {
TokenBucketLimiter::new(
time_source,
RESPOND_BURST,
RESPOND_REFILL_INTERVAL,
MAX_TRACKED_KEYS,
)
}
pub(crate) fn new_response_global_limiter(
time_source: Arc<dyn TimeSource + Send + Sync>,
) -> ResyncResponseGlobalLimiter {
TokenBucketLimiter::new(
time_source,
RESPOND_GLOBAL_BURST,
RESPOND_GLOBAL_REFILL_INTERVAL,
MAX_TRACKED_KEYS,
)
}
pub(crate) const OUTSTANDING_RESYNC_TTL: Duration = Duration::from_secs(60);
pub(crate) struct OutstandingResyncRequests {
entries: DashMap<(ContractInstanceId, SocketAddr), Instant>,
size: AtomicUsize,
max_tracked: usize,
ttl: Duration,
time_source: Arc<dyn TimeSource + Send + Sync>,
}
impl OutstandingResyncRequests {
pub fn new(time_source: Arc<dyn TimeSource + Send + Sync>) -> Self {
Self {
entries: DashMap::new(),
size: AtomicUsize::new(0),
max_tracked: MAX_TRACKED_KEYS,
ttl: OUTSTANDING_RESYNC_TTL,
time_source,
}
}
pub fn record(&self, contract: ContractInstanceId, target: SocketAddr) {
let now = self.time_source.now();
use dashmap::mapref::entry::Entry;
match self.entries.entry((contract, target)) {
Entry::Occupied(mut occ) => {
*occ.get_mut() = now;
}
Entry::Vacant(vac) => {
let prev = self.size.fetch_add(1, Ordering::Relaxed);
if prev >= self.max_tracked {
self.size.fetch_sub(1, Ordering::Relaxed);
return;
}
vac.insert(now);
}
}
}
pub fn consume(&self, contract: ContractInstanceId, source: SocketAddr) -> bool {
let now = self.time_source.now();
match self.entries.remove(&(contract, source)) {
Some((_, emitted)) => {
self.size.fetch_sub(1, Ordering::Relaxed);
now.saturating_duration_since(emitted) <= self.ttl
}
None => false,
}
}
pub fn cleanup(&self) {
let now = self.time_source.now();
let ttl = self.ttl;
let mut removed = 0usize;
self.entries.retain(|_, emitted| {
if now.saturating_duration_since(*emitted) > ttl {
removed += 1;
false
} else {
true
}
});
if removed > 0 {
self.size.fetch_sub(removed, Ordering::Relaxed);
}
}
#[cfg_attr(not(test), allow(dead_code))]
pub fn len(&self) -> usize {
self.entries.len()
}
}
pub(crate) fn new_outstanding_resync_requests(
time_source: Arc<dyn TimeSource + Send + Sync>,
) -> OutstandingResyncRequests {
OutstandingResyncRequests::new(time_source)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::util::time_source::SharedMockTimeSource;
fn mk_contract(byte: u8) -> ContractInstanceId {
ContractInstanceId::new([byte; 32])
}
fn mk_peer(byte: u8) -> SocketAddr {
SocketAddr::from(([10, 0, 0, byte], 30000 + byte as u16))
}
#[test]
fn outstanding_resync_solicited_response_consumes_and_replay_is_rejected() {
let ts = SharedMockTimeSource::new();
let m = new_outstanding_resync_requests(Arc::new(ts.clone()));
let c = mk_contract(1);
let peer = mk_peer(1);
assert!(
!m.consume(c, peer),
"a response with no outstanding request must be rejected"
);
m.record(c, peer);
assert_eq!(m.len(), 1);
assert!(
m.consume(c, peer),
"the solicited response must be accepted"
);
assert_eq!(m.len(), 0, "consume must remove the entry");
assert!(
!m.consume(c, peer),
"a replayed/duplicate response must be rejected after the first consume"
);
}
#[test]
fn outstanding_resync_is_scoped_per_contract_and_peer() {
let ts = SharedMockTimeSource::new();
let m = new_outstanding_resync_requests(Arc::new(ts.clone()));
let c = mk_contract(1);
m.record(c, mk_peer(1));
assert!(!m.consume(c, mk_peer(2)));
assert!(!m.consume(mk_contract(2), mk_peer(1)));
assert!(m.consume(c, mk_peer(1)));
}
#[test]
fn outstanding_resync_ttl_expiry_rejects_and_cleanup_reaps() {
let ts = SharedMockTimeSource::new();
let m = new_outstanding_resync_requests(Arc::new(ts.clone()));
let c = mk_contract(1);
let peer = mk_peer(1);
m.record(c, peer);
ts.advance_time(OUTSTANDING_RESYNC_TTL - Duration::from_millis(1));
assert!(
m.consume(c, peer),
"a response within the TTL must be accepted"
);
m.record(c, peer);
ts.advance_time(OUTSTANDING_RESYNC_TTL + Duration::from_millis(1));
assert!(
!m.consume(c, peer),
"a response arriving past the TTL must be rejected as unsolicited"
);
assert_eq!(m.len(), 0, "a consumed (even if stale) entry is removed");
m.record(mk_contract(9), mk_peer(9));
assert_eq!(m.len(), 1);
ts.advance_time(OUTSTANDING_RESYNC_TTL + Duration::from_secs(1));
m.cleanup();
assert_eq!(m.len(), 0, "cleanup must reap entries older than the TTL");
}
#[test]
fn outstanding_resync_is_strictly_capped() {
let ts = SharedMockTimeSource::new();
let m = new_outstanding_resync_requests(Arc::new(ts.clone()));
for i in 0..MAX_TRACKED_KEYS {
let c = ContractInstanceId::new([(i % 256) as u8; 32]);
let peer = SocketAddr::from(([10, (i >> 16) as u8, (i >> 8) as u8, i as u8], 40000));
m.record(c, peer);
}
assert_eq!(m.len(), MAX_TRACKED_KEYS, "map fills to exactly the cap");
let overflow_c = ContractInstanceId::new([0xAB; 32]);
let overflow_peer = SocketAddr::from(([172, 16, 0, 1], 41000));
m.record(overflow_c, overflow_peer);
assert_eq!(m.len(), MAX_TRACKED_KEYS, "over-cap insert is dropped");
assert!(!m.consume(overflow_c, overflow_peer));
}
#[test]
fn emit_allows_burst_then_throttles() {
let ts = SharedMockTimeSource::new();
let l = new_emit_limiter(Arc::new(ts.clone()));
let c = mk_contract(1);
assert!(l.check_and_record(c));
assert!(l.check_and_record(c));
assert!(!l.check_and_record(c));
assert_eq!(l.allowed_total(), 2);
assert_eq!(l.suppressed_total(), 1);
}
#[test]
fn emit_refills_one_per_interval() {
let ts = SharedMockTimeSource::new();
let l = new_emit_limiter(Arc::new(ts.clone()));
let c = mk_contract(1);
assert!(l.check_and_record(c));
assert!(l.check_and_record(c));
assert!(!l.check_and_record(c));
ts.advance_time(EMIT_REFILL_INTERVAL + Duration::from_secs(1));
assert!(l.check_and_record(c));
assert!(!l.check_and_record(c), "only one token refilled");
}
#[test]
fn emit_flood_is_bounded_to_ceil_window_over_interval() {
let ts = SharedMockTimeSource::new();
let l = new_emit_limiter(Arc::new(ts.clone()));
let c = mk_contract(1);
for _ in 0..300 {
l.check_and_record(c);
ts.advance_time(Duration::from_secs(1));
}
let allowed = l.allowed_total();
assert!(
(6..=8).contains(&allowed),
"flood over 5 min must admit ~7 ResyncRequests, got {allowed}"
);
assert_eq!(allowed + l.suppressed_total(), 300);
}
#[test]
fn emit_different_contracts_independent() {
let ts = SharedMockTimeSource::new();
let l = new_emit_limiter(Arc::new(ts.clone()));
assert!(l.check_and_record(mk_contract(1)));
assert!(l.check_and_record(mk_contract(2)));
assert!(l.check_and_record(mk_contract(1)));
assert!(l.check_and_record(mk_contract(2)));
assert!(!l.check_and_record(mk_contract(1)));
assert!(!l.check_and_record(mk_contract(2)));
}
#[test]
fn responder_burst_one_then_throttles() {
let ts = SharedMockTimeSource::new();
let l = new_response_limiter(Arc::new(ts.clone()));
let key = (mk_peer(1), mk_contract(1));
assert!(l.check_and_record(key));
assert!(!l.check_and_record(key), "responder burst is 1");
ts.advance_time(RESPOND_REFILL_INTERVAL + Duration::from_secs(1));
assert!(l.check_and_record(key));
}
#[test]
fn responder_flood_from_one_peer_is_bounded() {
let ts = SharedMockTimeSource::new();
let l = new_response_limiter(Arc::new(ts.clone()));
let key = (mk_peer(1), mk_contract(1));
for _ in 0..300 {
l.check_and_record(key);
ts.advance_time(Duration::from_secs(1));
}
let allowed = l.allowed_total();
assert!(
(10..=12).contains(&allowed),
"responder flood over 5 min must send ~11 responses, got {allowed}"
);
}
#[test]
fn global_responder_cap_bounds_across_many_distinct_peers() {
let ts = SharedMockTimeSource::new();
let l = new_response_global_limiter(Arc::new(ts.clone()));
let contract = mk_contract(1);
for _ in 0..300 {
for p in 0..45u8 {
let _ = p; l.check_and_record(contract);
}
ts.advance_time(Duration::from_secs(1));
}
let allowed = l.allowed_total();
assert!(
(60..=64).contains(&allowed),
"global cap must bound one contract to ~62 responses over 5 min \
regardless of requester count, got {allowed}"
);
assert!(
allowed < 100,
"global cap keeps a single forked contract well under the ~9,733/day \
production storm, got {allowed} in 5 min"
);
}
#[test]
fn responder_different_peers_independent() {
let ts = SharedMockTimeSource::new();
let l = new_response_limiter(Arc::new(ts.clone()));
let c = mk_contract(1);
assert!(l.check_and_record((mk_peer(1), c)));
assert!(l.check_and_record((mk_peer(2), c)));
assert!(!l.check_and_record((mk_peer(1), c)));
assert!(!l.check_and_record((mk_peer(2), c)));
}
#[test]
fn tracked_keys_capped() {
let ts = SharedMockTimeSource::new();
let l = TokenBucketLimiter::new(Arc::new(ts.clone()), 1.0, RESPOND_REFILL_INTERVAL, 4);
for i in 0..4u8 {
assert!(l.check_and_record((mk_peer(i), mk_contract(i))));
}
assert_eq!(l.len(), 4);
assert!(!l.check_and_record((mk_peer(99), mk_contract(99))));
assert_eq!(l.len(), 4);
}
#[test]
fn cleanup_removes_recovered_idle_keys() {
let ts = SharedMockTimeSource::new();
let l = new_emit_limiter(Arc::new(ts.clone()));
l.check_and_record(mk_contract(1));
l.check_and_record(mk_contract(2));
assert_eq!(l.len(), 2);
ts.advance_time(CLEANUP_AGE + EMIT_REFILL_INTERVAL * 3);
l.cleanup();
assert_eq!(l.len(), 0, "recovered idle keys must be swept");
assert_eq!(l.size.load(Ordering::Relaxed), 0);
}
#[test]
fn cleanup_preserves_recently_used_keys() {
let ts = SharedMockTimeSource::new();
let l = new_emit_limiter(Arc::new(ts.clone()));
l.check_and_record(mk_contract(1));
ts.advance_time(Duration::from_secs(10));
l.cleanup();
assert_eq!(l.len(), 1, "recently-used key must be preserved");
}
}