use std::collections::HashMap;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::atomic::Ordering;
use log::{error, info};
use portable_atomic::AtomicU64;
use wacore::time::Instant;
use whatsapp_rust::RetryAdmission;
use whatsapp_rust::prelude::*;
struct TokenBucket {
tokens: f64,
last_refill: Instant,
}
impl TokenBucket {
fn new(initial: f64, now: Instant) -> Self {
Self {
tokens: initial,
last_refill: now,
}
}
fn try_take(&mut self, now: Instant, burst: f64, refill_per_sec: f64) -> bool {
let elapsed = now
.saturating_duration_since(self.last_refill)
.as_secs_f64();
self.tokens = (self.tokens + elapsed * refill_per_sec).min(burst);
self.last_refill = now;
if self.tokens >= 1.0 {
self.tokens -= 1.0;
true
} else {
false
}
}
}
struct RetryQuarantine {
buckets: Mutex<HashMap<(String, String), TokenBucket>>,
burst: f64,
refill_per_sec: f64,
max_pairs: usize,
quarantined_total: AtomicU64,
}
impl RetryQuarantine {
fn new(burst: u32, refill_per_day: u32, max_pairs: usize) -> Self {
Self {
buckets: Mutex::new(HashMap::new()),
burst: burst as f64,
refill_per_sec: refill_per_day as f64 / 86_400.0,
max_pairs,
quarantined_total: AtomicU64::new(0),
}
}
fn quarantined_total(&self) -> u64 {
self.quarantined_total.load(Ordering::Relaxed)
}
}
impl RetryAdmission for RetryQuarantine {
fn admit(&self, chat: &Jid, requester: &Jid, _retry_count: u8) -> bool {
if self.burst == 0.0 {
return true;
}
let now = Instant::now();
let key = (chat.user.to_string(), requester.user.to_string());
let mut buckets = self.buckets.lock().expect("quarantine mutex poisoned");
if !buckets.contains_key(&key) && buckets.len() >= self.max_pairs {
return true;
}
let bucket = buckets
.entry(key)
.or_insert_with(|| TokenBucket::new(self.burst, now));
let allowed = bucket.try_take(now, self.burst, self.refill_per_sec);
drop(buckets);
if !allowed {
self.quarantined_total.fetch_add(1, Ordering::Relaxed);
}
allowed
}
}
fn main() {
env_logger::Builder::from_env(env_logger::Env::default().default_filter_or("info")).init();
let rt = tokio::runtime::Builder::new_multi_thread()
.enable_all()
.build()
.expect("failed to build tokio runtime");
rt.block_on(async {
let store = match SqliteStore::new("whatsapp.db").await {
Ok(store) => store,
Err(e) => {
error!("failed to create SQLite backend: {e}");
return;
}
};
let bot = match Bot::builder()
.with_backend(store)
.on_qr_code(|code, _timeout| async move {
info!("scan to pair:\n{code}");
})
.on_connected(|_client| async {
info!("connected; inbound group retry receipts now pass through the quarantine");
})
.build()
.await
{
Ok(bot) => bot,
Err(e) => {
error!("failed to build bot: {e}");
return;
}
};
let quarantine = Arc::new(RetryQuarantine::new(2, 2, 32_768));
if !bot.client().set_retry_admission(quarantine.clone()) {
error!("retry admission policy was already set");
return;
}
bot.run().await;
info!(
"quarantined {} retry receipt(s) this session",
quarantine.quarantined_total()
);
});
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
#[test]
fn bounds_per_pair_and_isolates_pairs() {
let q = RetryQuarantine::new(2, 0, 100);
let g: Jid = "123-456@g.us".parse().unwrap();
let a: Jid = "111@lid".parse().unwrap();
let b: Jid = "222@lid".parse().unwrap();
assert!(q.admit(&g, &a, 1));
assert!(q.admit(&g, &a, 1));
assert!(!q.admit(&g, &a, 1), "third receipt quarantined");
assert!(q.admit(&g, &b, 1));
assert_eq!(q.quarantined_total(), 1);
}
#[test]
fn burst_zero_disables() {
let q = RetryQuarantine::new(0, 0, 100);
let g: Jid = "123-456@g.us".parse().unwrap();
let a: Jid = "111@lid".parse().unwrap();
for _ in 0..10 {
assert!(q.admit(&g, &a, 1));
}
assert_eq!(q.quarantined_total(), 0);
}
#[test]
fn fails_open_past_capacity_for_new_pairs() {
let q = RetryQuarantine::new(1, 0, 1);
let g: Jid = "123-456@g.us".parse().unwrap();
let a: Jid = "111@lid".parse().unwrap();
let b: Jid = "222@lid".parse().unwrap();
assert!(q.admit(&g, &a, 1));
assert!(!q.admit(&g, &a, 1), "tracked pair still bounded");
assert!(q.admit(&g, &b, 1), "new pair fails open past cap");
}
#[test]
fn refill_restores_admission_gradually() {
let t0 = Instant::now();
let (burst, refill_per_sec) = (2.0, 1.0); let mut b = TokenBucket::new(burst, t0);
assert!(b.try_take(t0, burst, refill_per_sec));
assert!(b.try_take(t0, burst, refill_per_sec));
assert!(!b.try_take(t0, burst, refill_per_sec), "burst exhausted");
let t1 = t0 + Duration::from_secs(1);
assert!(b.try_take(t1, burst, refill_per_sec), "one token refilled");
assert!(
!b.try_take(t1, burst, refill_per_sec),
"only one token, not a full reset"
);
}
#[test]
fn idle_does_not_accrue_beyond_burst() {
let t0 = Instant::now();
let (burst, refill_per_sec) = (2.0, 1.0);
let mut b = TokenBucket::new(burst, t0);
let much_later = t0 + Duration::from_secs(3600);
assert!(b.try_take(much_later, burst, refill_per_sec));
assert!(b.try_take(much_later, burst, refill_per_sec));
assert!(
!b.try_take(much_later, burst, refill_per_sec),
"tokens clamp at burst regardless of idle time"
);
}
}