use std::collections::HashMap;
use std::future::Future;
use std::num::NonZeroU32;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use cdk_http_client::Transport;
use serde::de::DeserializeOwned;
use serde::Serialize;
use tokio::sync::{watch, OnceCell};
use url::Url;
use web_time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use crate::database::{self, WalletDatabase, KVSTORE_NAMESPACE_KEY_MAX_LEN};
use crate::{AuthToken, HttpError, RawResponse};
const KV_NAMESPACE: &str = "rate_limiter";
type BudgetDb = Arc<dyn WalletDatabase<database::Error> + Send + Sync>;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RateLimitConfig {
pub capacity: NonZeroU32,
pub refill_per_minute: NonZeroU32,
}
impl RateLimitConfig {
pub fn new(capacity: NonZeroU32, refill_per_minute: NonZeroU32) -> Self {
Self {
capacity,
refill_per_minute,
}
}
pub fn try_new(capacity: u32, refill_per_minute: u32) -> Option<Self> {
Some(Self::new(
NonZeroU32::new(capacity)?,
NonZeroU32::new(refill_per_minute)?,
))
}
}
impl Default for RateLimitConfig {
fn default() -> Self {
Self {
capacity: NonZeroU32::new(20).expect("20 is non-zero"),
refill_per_minute: NonZeroU32::new(20).expect("20 is non-zero"),
}
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
trait BudgetStore: std::fmt::Debug + Send + Sync {
async fn load(&self) -> Option<Vec<u8>>;
async fn store(&self, value: &[u8]) -> bool;
}
#[derive(Debug)]
struct KvBudgetStore {
db: BudgetDb,
key: String,
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl BudgetStore for KvBudgetStore {
async fn load(&self) -> Option<Vec<u8>> {
match self.db.kv_read(KV_NAMESPACE, "", &self.key).await {
Ok(value) => value,
Err(err) => {
tracing::warn!("rate limiter failed to load persisted budget: {err}");
None
}
}
}
async fn store(&self, value: &[u8]) -> bool {
match self.db.kv_write(KV_NAMESPACE, "", &self.key, value).await {
Ok(()) => true,
Err(err) => {
tracing::warn!("rate limiter failed to persist budget: {err}");
false
}
}
}
}
#[derive(Debug)]
struct BucketState {
arrival_time: Instant,
emission_interval: Duration,
tolerance: Duration,
enabled: bool,
}
#[derive(Debug)]
struct Writer {
desired: watch::Sender<u64>,
progress: watch::Receiver<u64>,
}
#[derive(Debug)]
struct TokenBucketInner {
state: Mutex<BucketState>,
persistence: Option<Arc<dyn BudgetStore>>,
started: OnceCell<Option<Writer>>,
}
fn params_from_config(config: RateLimitConfig) -> (Duration, Duration) {
let emission_interval = Duration::from_secs(60) / config.refill_per_minute.get();
let tolerance = emission_interval * (config.capacity.get() - 1);
(emission_interval, tolerance)
}
#[derive(Debug, Clone)]
pub struct TokenBucket {
inner: Arc<TokenBucketInner>,
}
impl TokenBucket {
pub fn new(config: RateLimitConfig) -> Self {
Self::build(config, None)
}
fn persisted(config: RateLimitConfig, key: &str, db: Option<BudgetDb>) -> Self {
let persistence = db.map(|db| {
Arc::new(KvBudgetStore {
db,
key: key.to_string(),
}) as Arc<dyn BudgetStore>
});
Self::build(config, persistence)
}
#[cfg(test)]
fn with_store(config: RateLimitConfig, store: Arc<dyn BudgetStore>) -> Self {
Self::build(config, Some(store))
}
fn build(config: RateLimitConfig, persistence: Option<Arc<dyn BudgetStore>>) -> Self {
let (emission_interval, tolerance) = params_from_config(config);
Self {
inner: Arc::new(TokenBucketInner {
state: Mutex::new(BucketState {
arrival_time: Instant::now(),
emission_interval,
tolerance,
enabled: true,
}),
persistence,
started: OnceCell::new(),
}),
}
}
pub fn set_config(&self, config: RateLimitConfig) {
let (emission_interval, tolerance) = params_from_config(config);
let mut state = lock(&self.inner.state);
state.emission_interval = emission_interval;
state.tolerance = tolerance;
state.enabled = true;
}
pub fn set_enabled(&self, enabled: bool) {
lock(&self.inner.state).enabled = enabled;
}
pub async fn acquire<F, T>(&self, action: F) -> T
where
F: Future<Output = T>,
{
let wait = self.reserve_slot().await;
if !wait.is_zero() {
sleep(wait).await;
}
action.await
}
async fn reserve_slot(&self) -> Duration {
self.ensure_started().await;
match self.advance() {
Some((wait, millis)) => {
self.publish(millis);
wait
}
None => Duration::ZERO,
}
}
pub async fn try_acquire(&self) -> bool {
self.ensure_started().await;
let reserved = {
let mut state = lock(&self.inner.state);
if !state.enabled {
return true;
}
let base = state.arrival_time.max(Instant::now());
let ahead = base.saturating_duration_since(Instant::now());
if ahead > state.tolerance {
None
} else {
state.arrival_time = base + state.emission_interval;
Some(tat_to_unix_millis(state.arrival_time))
}
};
match reserved {
Some(millis) => {
self.publish(millis);
true
}
None => false,
}
}
fn is_fully_recovered(&self) -> bool {
lock(&self.inner.state).arrival_time <= Instant::now()
}
fn handle_count(&self) -> usize {
Arc::strong_count(&self.inner)
}
fn advance(&self) -> Option<(Duration, u64)> {
let mut state = lock(&self.inner.state);
if !state.enabled {
return None;
}
let now = Instant::now();
let base = state.arrival_time.max(now);
let ahead = base.saturating_duration_since(now);
let wait = ahead.saturating_sub(state.tolerance);
state.arrival_time = base + state.emission_interval;
Some((wait, tat_to_unix_millis(state.arrival_time)))
}
async fn ensure_started(&self) {
self.inner
.started
.get_or_init(|| async {
let store = self.inner.persistence.clone()?;
let loaded = match with_timeout(LOAD_TIMEOUT, store.load()).await {
Some(Some(bytes)) => <[u8; 8]>::try_from(bytes.as_slice())
.ok()
.map(u64::from_be_bytes),
_ => None,
};
let seed = if let Some(stored) = loaded {
let tat_wall = UNIX_EPOCH + Duration::from_millis(stored);
let mut state = lock(&self.inner.state);
let debt = tat_wall
.duration_since(SystemTime::now())
.unwrap_or_default()
.min(state.tolerance);
state.arrival_time = Instant::now() + debt;
tat_to_unix_millis(state.arrival_time)
} else {
0
};
let (desired_tx, desired_rx) = watch::channel(seed);
let (progress_tx, progress_rx) = watch::channel(seed);
crate::task::spawn(run_writer(store, desired_rx, progress_tx));
Some(Writer {
desired: desired_tx,
progress: progress_rx,
})
})
.await;
}
fn publish(&self, millis: u64) {
if let Some(Some(writer)) = self.inner.started.get() {
writer.desired.send_if_modified(|current| {
if *current < millis {
*current = millis;
true
} else {
false
}
});
}
}
pub async fn flush(&self) {
let Some(Some(writer)) = self.inner.started.get() else {
return;
};
let target = {
let state = lock(&self.inner.state);
tat_to_unix_millis(state.arrival_time)
};
writer.desired.send_if_modified(|current| {
if *current < target {
*current = target;
true
} else {
false
}
});
let mut progress = writer.progress.clone();
let wait = async {
while *progress.borrow_and_update() < target {
if progress.changed().await.is_err() {
break;
}
}
};
let _ = with_timeout(FLUSH_TIMEOUT, wait).await;
}
}
#[derive(Debug, Clone, Copy)]
struct ManagerSettings {
config: RateLimitConfig,
enabled: bool,
}
#[derive(Clone)]
pub struct RateLimiterManager {
settings: Arc<Mutex<ManagerSettings>>,
db: Option<BudgetDb>,
buckets: Arc<Mutex<HashMap<String, TokenBucket>>>,
}
impl std::fmt::Debug for RateLimiterManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RateLimiterManager")
.field("settings", &*lock(&self.settings))
.finish_non_exhaustive()
}
}
impl RateLimiterManager {
pub fn new(config: RateLimitConfig, db: Option<BudgetDb>) -> Self {
Self {
settings: Arc::new(Mutex::new(ManagerSettings {
config,
enabled: true,
})),
db,
buckets: Arc::new(Mutex::new(HashMap::new())),
}
}
pub fn bucket_for(&self, url: &Url) -> TokenBucket {
let origin = origin_key(url);
let key = origin.clone().unwrap_or_else(|| url.to_string());
let mut buckets = lock(&self.buckets);
if let Some(bucket) = buckets.get(&key) {
return bucket.clone();
}
buckets.retain(|_, bucket| bucket.handle_count() > 1 || !bucket.is_fully_recovered());
let db = match origin {
Some(_) => self.db.clone(),
None => None,
};
let settings = *lock(&self.settings);
let bucket = TokenBucket::persisted(settings.config, &key, db);
bucket.set_enabled(settings.enabled);
buckets.insert(key, bucket.clone());
bucket
}
pub fn set_config(&self, config: RateLimitConfig) {
{
let mut settings = lock(&self.settings);
settings.config = config;
settings.enabled = true;
}
for bucket in lock(&self.buckets).values() {
bucket.set_config(config);
}
}
pub fn set_enabled(&self, enabled: bool) {
lock(&self.settings).enabled = enabled;
for bucket in lock(&self.buckets).values() {
bucket.set_enabled(enabled);
}
}
pub fn is_enabled(&self) -> bool {
lock(&self.settings).enabled
}
pub async fn flush(&self) {
let buckets: Vec<TokenBucket> = lock(&self.buckets).values().cloned().collect();
futures::future::join_all(buckets.iter().map(TokenBucket::flush)).await;
}
pub fn origin_count(&self) -> usize {
lock(&self.buckets).len()
}
}
async fn run_writer(
store: Arc<dyn BudgetStore>,
mut desired: watch::Receiver<u64>,
progress: watch::Sender<u64>,
) {
while desired.changed().await.is_ok() {
let millis = *desired.borrow_and_update();
store.store(&millis.to_be_bytes()).await;
let _ = progress.send(millis);
}
let millis = *desired.borrow();
store.store(&millis.to_be_bytes()).await;
let _ = progress.send(millis);
}
const FLUSH_TIMEOUT: Duration = Duration::from_secs(5);
const LOAD_TIMEOUT: Duration = Duration::from_millis(200);
async fn with_timeout<F: Future>(duration: Duration, fut: F) -> Option<F::Output> {
use futures::future::{select, Either};
let fut = std::pin::pin!(fut);
let timeout = std::pin::pin!(sleep(duration));
match select(fut, timeout).await {
Either::Left((output, _)) => Some(output),
Either::Right(((), _)) => None,
}
}
fn lock<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
mutex
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
fn tat_to_unix_millis(arrival_time: Instant) -> u64 {
let ahead = arrival_time
.checked_duration_since(Instant::now())
.unwrap_or_default();
(SystemTime::now() + ahead)
.duration_since(UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0)
}
pub fn origin_key(url: &Url) -> Option<String> {
let host = url.host_str()?;
let authority = match url.port() {
Some(port) => format!("{host}:{port}"),
None => host.to_string(),
};
let sanitized: String = authority
.chars()
.map(|c| {
if c.is_ascii_alphanumeric() || c == '_' || c == '-' {
c
} else {
'_'
}
})
.take(KVSTORE_NAMESPACE_KEY_MAX_LEN)
.collect();
if sanitized.is_empty() {
None
} else {
Some(sanitized)
}
}
#[cfg(not(target_arch = "wasm32"))]
async fn sleep(duration: Duration) {
tokio::time::sleep(duration).await;
}
#[cfg(target_arch = "wasm32")]
async fn sleep(duration: Duration) {
gloo_timers::future::TimeoutFuture::new(duration.as_millis() as u32).await;
}
#[derive(Debug, Clone)]
pub struct RateLimitedTransport<T> {
inner: T,
limiter: RateLimiterManager,
}
impl<T> RateLimitedTransport<T> {
pub fn with_manager(inner: T, limiter: RateLimiterManager) -> Self {
Self { inner, limiter }
}
}
impl<T: Default> Default for RateLimitedTransport<T> {
fn default() -> Self {
Self::with_manager(
T::default(),
RateLimiterManager::new(RateLimitConfig::default(), None),
)
}
}
#[cfg_attr(target_arch = "wasm32", async_trait(?Send))]
#[cfg_attr(not(target_arch = "wasm32"), async_trait)]
impl<T: Transport> Transport for RateLimitedTransport<T> {
async fn ws_connect(
&self,
url: &str,
headers: &[(&str, &str)],
) -> Result<
(
cdk_http_client::ws::WsSender,
cdk_http_client::ws::WsReceiver,
),
cdk_http_client::ws::WsError,
> {
self.inner.ws_connect(url, headers).await
}
fn with_proxy(
&mut self,
proxy: Url,
host_matcher: Option<&str>,
accept_invalid_certs: bool,
) -> Result<(), HttpError> {
self.inner
.with_proxy(proxy, host_matcher, accept_invalid_certs)
}
#[cfg(all(feature = "bip353", not(target_arch = "wasm32")))]
async fn resolve_dns_txt(&self, domain: &str) -> Result<Vec<String>, HttpError> {
self.inner.resolve_dns_txt(domain).await
}
async fn http_get<R>(&self, url: Url, auth: Option<AuthToken>) -> Result<R, HttpError>
where
R: DeserializeOwned,
{
let bucket = self.limiter.bucket_for(&url);
bucket.acquire(self.inner.http_get(url, auth)).await
}
async fn http_get_raw(
&self,
url: Url,
auth: Option<AuthToken>,
) -> Result<RawResponse, HttpError> {
let bucket = self.limiter.bucket_for(&url);
bucket.acquire(self.inner.http_get_raw(url, auth)).await
}
async fn http_post<P, R>(
&self,
url: Url,
auth_token: Option<AuthToken>,
payload: &P,
) -> Result<R, HttpError>
where
P: Serialize + Send + Sync,
R: DeserializeOwned,
{
let bucket = self.limiter.bucket_for(&url);
bucket
.acquire(self.inner.http_post(url, auth_token, payload))
.await
}
async fn http_post_form_raw<P>(
&self,
url: Url,
auth_token: Option<AuthToken>,
payload: &P,
) -> Result<RawResponse, HttpError>
where
P: Serialize + Send + Sync,
{
let bucket = self.limiter.bucket_for(&url);
bucket
.acquire(self.inner.http_post_form_raw(url, auth_token, payload))
.await
}
}
#[cfg(all(test, not(target_arch = "wasm32")))]
mod tests {
use std::time::Instant as StdInstant;
use super::*;
fn config(capacity: u32, refill_per_minute: u32) -> RateLimitConfig {
RateLimitConfig::new(
NonZeroU32::new(capacity).unwrap_or(NonZeroU32::MIN),
NonZeroU32::new(refill_per_minute).unwrap_or(NonZeroU32::MIN),
)
}
#[test]
fn default_config_values() {
let cfg = RateLimitConfig::default();
assert_eq!(cfg.capacity.get(), 20);
assert_eq!(cfg.refill_per_minute.get(), 20);
assert!(cfg.capacity.get() + cfg.refill_per_minute.get() < 60);
}
#[test]
fn try_new_rejects_zero() {
assert!(RateLimitConfig::try_new(0, 45).is_none());
assert!(RateLimitConfig::try_new(10, 0).is_none());
let cfg = RateLimitConfig::try_new(10, 45).expect("non-zero");
assert_eq!(cfg.capacity.get(), 10);
assert_eq!(cfg.refill_per_minute.get(), 45);
}
#[derive(Debug, Default)]
struct StubStore {
fail: bool,
}
#[async_trait]
impl BudgetStore for StubStore {
async fn load(&self) -> Option<Vec<u8>> {
None
}
async fn store(&self, _value: &[u8]) -> bool {
!self.fail
}
}
#[derive(Debug)]
struct GatedStore {
loads: Arc<std::sync::atomic::AtomicUsize>,
writes: Arc<Mutex<Vec<u64>>>,
gate: Arc<tokio::sync::Semaphore>,
}
#[async_trait]
impl BudgetStore for GatedStore {
async fn load(&self) -> Option<Vec<u8>> {
self.loads.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
None
}
async fn store(&self, value: &[u8]) -> bool {
let _permit = self.gate.acquire().await.expect("gate not closed");
if let Ok(raw) = <[u8; 8]>::try_from(value) {
lock(&self.writes).push(u64::from_be_bytes(raw));
}
true
}
}
#[derive(Debug)]
struct PreloadedStore {
preload: u64,
writes: Arc<Mutex<Vec<u64>>>,
}
#[async_trait]
impl BudgetStore for PreloadedStore {
async fn load(&self) -> Option<Vec<u8>> {
Some(self.preload.to_be_bytes().to_vec())
}
async fn store(&self, value: &[u8]) -> bool {
if let Ok(raw) = <[u8; 8]>::try_from(value) {
lock(&self.writes).push(u64::from_be_bytes(raw));
}
true
}
}
#[derive(Debug)]
struct DropSignalStore {
on_drop: Mutex<Option<tokio::sync::oneshot::Sender<()>>>,
}
#[async_trait]
impl BudgetStore for DropSignalStore {
async fn load(&self) -> Option<Vec<u8>> {
None
}
async fn store(&self, _value: &[u8]) -> bool {
true
}
}
impl Drop for DropSignalStore {
fn drop(&mut self) {
if let Some(tx) = lock(&self.on_drop).take() {
let _ = tx.send(());
}
}
}
#[tokio::test]
async fn writer_task_terminates_when_last_handle_drops() {
let (tx, rx) = tokio::sync::oneshot::channel();
let store: Arc<dyn BudgetStore> = Arc::new(DropSignalStore {
on_drop: Mutex::new(Some(tx)),
});
let bucket = TokenBucket::with_store(config(5, 300), store);
bucket.acquire(async {}).await;
drop(bucket);
tokio::time::timeout(Duration::from_secs(2), rx)
.await
.expect("writer task should terminate and drop its store")
.expect("drop signal sender dropped without sending");
}
#[tokio::test]
async fn far_future_persisted_budget_heals_to_clamped_value() {
let now_millis = SystemTime::now()
.duration_since(UNIX_EPOCH)
.expect("after epoch")
.as_millis() as u64;
let far_future = now_millis + 3_600_000;
let writes = Arc::new(Mutex::new(Vec::new()));
let store = Arc::new(PreloadedStore {
preload: far_future,
writes: writes.clone(),
});
let bucket = TokenBucket::with_store(config(5, 300), store);
bucket.acquire(async {}).await;
bucket.flush().await;
let writes = lock(&writes).clone();
assert!(!writes.is_empty(), "the clamped budget must be persisted");
let max_written = writes.iter().copied().max().expect("non-empty");
assert!(
max_written < far_future,
"persisted budget {max_written} should heal below the far-future seed {far_future}",
);
assert!(max_written <= now_millis + 60_000);
}
#[tokio::test]
async fn writer_coalesces_to_the_latest_value() {
let store = Arc::new(GatedStore {
loads: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
writes: Arc::new(Mutex::new(Vec::new())),
gate: Arc::new(tokio::sync::Semaphore::new(0)),
});
let bucket = TokenBucket::with_store(config(1000, 60000), store.clone());
for _ in 0..50 {
bucket.acquire(async {}).await;
}
store.gate.add_permits(10);
bucket.flush().await;
let writes = lock(&store.writes).clone();
assert!(!writes.is_empty(), "the latest value must be persisted");
assert!(
writes.len() <= 2,
"writes should coalesce, got {}",
writes.len()
);
assert!(writes.windows(2).all(|w| w[0] <= w[1]));
assert_eq!(store.loads.load(std::sync::atomic::Ordering::SeqCst), 1);
}
#[tokio::test]
async fn failing_store_never_blocks_acquire() {
let store = Arc::new(StubStore { fail: true });
let bucket = TokenBucket::with_store(config(5, 300), store);
let start = StdInstant::now();
for _ in 0..5 {
bucket.acquire(async {}).await;
}
assert!(start.elapsed() < Duration::from_millis(100));
}
#[derive(Debug)]
struct HangingStore;
#[async_trait]
impl BudgetStore for HangingStore {
async fn load(&self) -> Option<Vec<u8>> {
None
}
async fn store(&self, _value: &[u8]) -> bool {
std::future::pending().await
}
}
#[tokio::test(start_paused = true)]
async fn flush_is_bounded_when_the_store_hangs() {
let bucket = TokenBucket::with_store(config(5, 300), Arc::new(HangingStore));
bucket.acquire(async {}).await;
assert!(
with_timeout(FLUSH_TIMEOUT * 2, bucket.flush())
.await
.is_some(),
"flush must give up rather than wait on a hung store",
);
}
#[tokio::test]
async fn flush_on_an_unused_bucket_persists_nothing() {
let writes = Arc::new(Mutex::new(Vec::new()));
let store = Arc::new(PreloadedStore {
preload: 0,
writes: writes.clone(),
});
let bucket = TokenBucket::with_store(config(5, 300), store);
bucket.flush().await;
assert!(lock(&writes).is_empty());
}
#[tokio::test]
async fn flush_returns_on_failing_store() {
let store = Arc::new(StubStore { fail: true });
let bucket = TokenBucket::with_store(config(5, 300), store);
bucket.acquire(async {}).await;
tokio::time::timeout(Duration::from_secs(2), bucket.flush())
.await
.expect("flush must not hang on a failing store");
}
fn parse(url: &str) -> Url {
Url::parse(url).expect("valid url")
}
#[test]
fn kv_key_is_sanitized() {
let key = origin_key(&parse("https://mint.example.com:3338")).expect("host present");
assert!(!key.contains('.'));
assert!(!key.contains(':'));
assert!(key
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-'));
}
#[test]
fn kv_key_ignores_default_port() {
let implicit = origin_key(&parse("https://mint.example.com"));
let explicit = origin_key(&parse("https://mint.example.com:443"));
assert_eq!(implicit, explicit);
assert_eq!(implicit.as_deref(), Some("mint_example_com"));
let other = origin_key(&parse("https://mint.example.com:8443"));
assert_ne!(implicit, other);
}
#[test]
fn kv_key_ignores_scheme_and_path() {
let plain = origin_key(&parse("http://mint.example.com/v1/keys"));
let secure = origin_key(&parse("https://mint.example.com/other/mint"));
assert_eq!(plain, secure);
}
#[tokio::test]
async fn manager_paces_each_origin_independently() {
let manager = RateLimiterManager::new(config(1, 300), None);
let mint = manager.bucket_for(&parse("https://mint.example.com/v1/mint"));
let same_host = manager.bucket_for(&parse("https://mint.example.com/v1/melt"));
let lnurl = manager.bucket_for(&parse("https://pay.example.org/.well-known/lnurlp/alice"));
assert!(mint.try_acquire().await);
assert!(!same_host.try_acquire().await, "same host shares a budget");
assert!(lnurl.try_acquire().await, "another host is independent");
}
#[tokio::test]
async fn manager_toggles_reach_existing_and_future_buckets() {
let manager = RateLimiterManager::new(config(1, 60), None);
let existing = manager.bucket_for(&parse("https://mint.example.com"));
assert!(existing.try_acquire().await);
assert!(!existing.try_acquire().await);
manager.set_enabled(false);
assert!(existing.try_acquire().await, "existing bucket is disabled");
let later = manager.bucket_for(&parse("https://pay.example.org"));
for _ in 0..5 {
assert!(later.try_acquire().await, "new bucket inherits disabled");
}
manager.set_config(config(1, 60));
assert!(!existing.try_acquire().await, "existing bucket paces again");
let newest = manager.bucket_for(&parse("https://third.example.net"));
assert!(newest.try_acquire().await);
assert!(!newest.try_acquire().await, "new bucket inherits config");
}
#[tokio::test]
async fn manager_reports_whether_pacing_is_on() {
let manager = RateLimiterManager::new(config(1, 60), None);
assert!(manager.is_enabled(), "a fresh manager paces");
manager.set_enabled(false);
assert!(!manager.is_enabled());
manager.set_config(config(2, 120));
assert!(manager.is_enabled(), "set_config re-enables pacing");
}
#[tokio::test]
async fn manager_evicts_only_recovered_unheld_origins() {
let manager = RateLimiterManager::new(config(1, 60), None);
let held = manager.bucket_for(&parse("https://held.example.com"));
manager.bucket_for(&parse("https://stale.example.com"));
assert_eq!(manager.origin_count(), 2);
manager.bucket_for(&parse("https://fresh.example.com"));
assert_eq!(manager.origin_count(), 2, "stale origin was evicted");
assert!(held.try_acquire().await);
assert!(!held.try_acquire().await);
assert!(
!manager
.bucket_for(&parse("https://held.example.com"))
.try_acquire()
.await
);
}
#[tokio::test]
async fn fresh_bucket_starts_full() {
let bucket = TokenBucket::new(config(5, 300));
let start = StdInstant::now();
for _ in 0..5 {
bucket.acquire(async {}).await;
}
assert!(start.elapsed() < Duration::from_millis(100));
}
#[tokio::test]
async fn acquiring_past_capacity_blocks() {
let bucket = TokenBucket::new(config(3, 300));
for _ in 0..3 {
bucket.acquire(async {}).await;
}
let start = StdInstant::now();
bucket.acquire(async {}).await;
assert!(start.elapsed() >= Duration::from_millis(150));
}
#[tokio::test]
async fn try_acquire_respects_burst() {
let bucket = TokenBucket::new(config(2, 600));
assert!(bucket.try_acquire().await);
assert!(bucket.try_acquire().await);
assert!(!bucket.try_acquire().await);
assert!(!bucket.try_acquire().await);
}
#[tokio::test]
async fn disabled_bucket_admits_immediately() {
let bucket = TokenBucket::new(config(1, 60));
bucket.set_enabled(false);
let start = StdInstant::now();
for _ in 0..10 {
bucket.acquire(async {}).await;
}
assert!(start.elapsed() < Duration::from_millis(100));
assert!(bucket.try_acquire().await);
assert!(bucket.try_acquire().await);
}
#[tokio::test]
async fn re_enabling_resumes_pacing() {
let bucket = TokenBucket::new(config(1, 300));
bucket.set_enabled(false);
for _ in 0..5 {
bucket.acquire(async {}).await;
}
bucket.set_config(config(1, 300));
bucket.acquire(async {}).await;
let start = StdInstant::now();
bucket.acquire(async {}).await;
assert!(start.elapsed() >= Duration::from_millis(150));
}
#[tokio::test]
async fn set_config_changes_rate() {
let bucket = TokenBucket::new(config(1, 60));
bucket.set_config(config(5, 600));
let start = StdInstant::now();
for _ in 0..5 {
bucket.acquire(async {}).await;
}
assert!(start.elapsed() < Duration::from_millis(100));
}
#[tokio::test]
async fn set_config_re_enables_a_disabled_bucket() {
let bucket = TokenBucket::new(config(2, 600));
bucket.set_enabled(false);
bucket.set_config(config(2, 600));
assert!(bucket.try_acquire().await);
assert!(bucket.try_acquire().await);
assert!(!bucket.try_acquire().await);
}
#[tokio::test]
async fn clones_share_one_budget() {
let bucket = TokenBucket::new(config(2, 600));
let clone = bucket.clone();
assert!(bucket.try_acquire().await);
assert!(clone.try_acquire().await);
assert!(!bucket.try_acquire().await);
assert!(!clone.try_acquire().await);
}
#[tokio::test]
async fn concurrent_acquires_all_complete() {
let bucket = TokenBucket::new(config(4, 6000));
let mut handles = Vec::new();
for _ in 0..8 {
let bucket = bucket.clone();
handles.push(tokio::spawn(async move { bucket.acquire(async {}).await }));
}
for handle in handles {
handle.await.unwrap();
}
}
#[derive(Debug, Clone, Default)]
struct CountingTransport {
http_calls: Arc<std::sync::atomic::AtomicUsize>,
proxied: Arc<std::sync::atomic::AtomicBool>,
}
impl CountingTransport {
fn http_calls(&self) -> usize {
self.http_calls.load(std::sync::atomic::Ordering::SeqCst)
}
fn bump(&self) {
self.http_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}
}
#[async_trait]
impl Transport for CountingTransport {
fn with_proxy(
&mut self,
_proxy: Url,
_host_matcher: Option<&str>,
_accept_invalid_certs: bool,
) -> Result<(), HttpError> {
self.proxied
.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(())
}
async fn http_get<R>(&self, _url: Url, _auth: Option<AuthToken>) -> Result<R, HttpError>
where
R: DeserializeOwned,
{
self.bump();
Err(HttpError::Other("mock".to_string()))
}
async fn http_get_raw(
&self,
_url: Url,
_auth: Option<AuthToken>,
) -> Result<RawResponse, HttpError> {
self.bump();
Err(HttpError::Other("mock".to_string()))
}
async fn http_post<P, R>(
&self,
_url: Url,
_auth: Option<AuthToken>,
_payload: &P,
) -> Result<R, HttpError>
where
P: Serialize + Send + Sync,
R: DeserializeOwned,
{
self.bump();
Err(HttpError::Other("mock".to_string()))
}
async fn http_post_form_raw<P>(
&self,
_url: Url,
_auth: Option<AuthToken>,
_payload: &P,
) -> Result<RawResponse, HttpError>
where
P: Serialize + Send + Sync,
{
self.bump();
Err(HttpError::Other("mock".to_string()))
}
#[cfg(all(feature = "bip353", not(target_arch = "wasm32")))]
async fn resolve_dns_txt(&self, _domain: &str) -> Result<Vec<String>, HttpError> {
Ok(Vec::new())
}
}
fn url() -> Url {
Url::parse("http://localhost/").expect("valid url")
}
#[tokio::test]
async fn transport_paces_http_and_delegates() {
let inner = CountingTransport::default();
let counter = inner.clone();
let transport = RateLimitedTransport::with_manager(
inner,
RateLimiterManager::new(config(2, 300), None),
);
let start = StdInstant::now();
let _ = transport.http_get_raw(url(), None).await;
let _ = transport.http_get_raw(url(), None).await;
assert!(
start.elapsed() < Duration::from_millis(100),
"burst should not pace"
);
let start = StdInstant::now();
let _ = transport.http_get_raw(url(), None).await;
assert!(
start.elapsed() >= Duration::from_millis(150),
"third call should be paced"
);
assert_eq!(counter.http_calls(), 3);
}
#[tokio::test]
async fn transport_paces_by_destination_host() {
let transport = RateLimitedTransport::with_manager(
CountingTransport::default(),
RateLimiterManager::new(config(1, 300), None),
);
let mint = parse("https://mint.example.com/v1/keys");
let lnurl = parse("https://pay.example.org/.well-known/lnurlp/alice");
let _ = transport.http_get_raw(mint.clone(), None).await;
let start = StdInstant::now();
let _ = transport.http_get_raw(lnurl, None).await;
assert!(
start.elapsed() < Duration::from_millis(100),
"another host must not be paced by the mint's budget"
);
let start = StdInstant::now();
let _ = transport.http_get_raw(mint, None).await;
assert!(
start.elapsed() >= Duration::from_millis(150),
"the mint's own budget is still enforced"
);
}
#[tokio::test]
async fn transport_passes_proxy_through_unthrottled() {
let inner = CountingTransport::default();
let flag = inner.proxied.clone();
let counter = inner.clone();
let mut transport = RateLimitedTransport::with_manager(
inner,
RateLimiterManager::new(config(1, 300), None),
);
transport.with_proxy(url(), None, false).expect("proxy set");
assert!(flag.load(std::sync::atomic::Ordering::SeqCst));
assert_eq!(counter.http_calls(), 0);
}
#[tokio::test]
async fn transports_sharing_a_manager_share_the_budget() {
let manager = RateLimiterManager::new(config(1, 300), None);
let a = RateLimitedTransport::with_manager(CountingTransport::default(), manager.clone());
let b = RateLimitedTransport::with_manager(CountingTransport::default(), manager.clone());
let _ = a.http_get_raw(url(), None).await;
let start = StdInstant::now();
let _ = b.http_get_raw(url(), None).await;
assert!(
start.elapsed() >= Duration::from_millis(150),
"shared bucket should force B to pace"
);
}
}