Skip to main content

mail4agent_server/
live.rs

1//! Wake-only long-poll plumbing for `GET /client/v3/sync`, plus the
2//! per-caller token bucket for `POST /client/v3/keys/claim`.
3//!
4//! `/sync` is pull, not push. A client polls with `since`, the server blocks
5//! up to `timeout`, and on wake it rebuilds the response from the database.
6//! The registry carries no payloads. A burst of writes against one key
7//! coalesces into one rebuilt response.
8//!
9//! Register, then read the database, then wait. [`LiveRegistry::register`]
10//! snapshots the generation before the read, so a wake that lands during
11//! the read is still observed by [`Registration::wait`]. Wake keys are
12//! `user:{user_id}`. Call [`LiveRegistry::wake_many`] only after the
13//! connection mutex is released.
14//!
15//! `Notify::notify_waiters` has no memory. [`Registration::wait`] enables
16//! its `Notified` future and then re-checks the generation, so a wake that
17//! fired before `wait` is not lost.
18
19use std::collections::HashMap;
20use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
21use std::sync::{Arc, Mutex};
22use std::time::{Duration, Instant};
23
24use tokio::sync::Notify;
25
26const CLAIM_RATE_CAPACITY: f64 = 100.0;
27const CLAIM_RATE_REFILL_PER_SEC: f64 = CLAIM_RATE_CAPACITY / 60.0;
28
29/// Per-caller token bucket for `POST /client/v3/keys/claim`. One token per
30/// requested device. A batch that would exceed the cap is refused whole.
31pub struct ClaimRateLimiter {
32    buckets: Mutex<HashMap<i64, (f64, Instant)>>,
33}
34
35impl Default for ClaimRateLimiter {
36    fn default() -> Self {
37        Self::new()
38    }
39}
40
41impl ClaimRateLimiter {
42    pub fn new() -> Self {
43        Self { buckets: Mutex::new(HashMap::new()) }
44    }
45
46    pub fn try_consume(&self, user_id: i64, cost: f64, now: Instant) -> bool {
47        let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
48        let (tokens, last_refill) = buckets.entry(user_id).or_insert((CLAIM_RATE_CAPACITY, now));
49        let elapsed = now.saturating_duration_since(*last_refill).as_secs_f64();
50        *tokens = (*tokens + elapsed * CLAIM_RATE_REFILL_PER_SEC).min(CLAIM_RATE_CAPACITY);
51        *last_refill = now;
52        if *tokens >= cost {
53            *tokens -= cost;
54            true
55        } else {
56            false
57        }
58    }
59}
60
61struct KeyEntry {
62    notify: Notify,
63    generation: AtomicU64,
64    interest: AtomicUsize,
65}
66
67impl KeyEntry {
68    fn new() -> Self {
69        Self { notify: Notify::new(), generation: AtomicU64::new(0), interest: AtomicUsize::new(0) }
70    }
71}
72
73#[derive(Debug, Clone, Copy, PartialEq, Eq)]
74pub enum WaitOutcome {
75    Woken,
76    TimedOut,
77}
78
79/// Registry of wake-only long-poll keys.
80pub struct LiveRegistry {
81    keys: Mutex<HashMap<String, Arc<KeyEntry>>>,
82}
83
84impl Default for LiveRegistry {
85    fn default() -> Self {
86        Self::new()
87    }
88}
89
90impl LiveRegistry {
91    pub fn new() -> Self {
92        Self { keys: Mutex::new(HashMap::new()) }
93    }
94
95    /// Register interest in `key` before reading the database.
96    pub fn register(&self, key: &str) -> Registration<'_> {
97        let entry = {
98            let mut keys = self.keys.lock().unwrap_or_else(|e| e.into_inner());
99            let entry = Arc::clone(keys.entry(key.to_string()).or_insert_with(|| Arc::new(KeyEntry::new())));
100            entry.interest.fetch_add(1, Ordering::AcqRel);
101            entry
102        };
103        let seen_generation = entry.generation.load(Ordering::Acquire);
104        Registration { registry: self, key: key.to_string(), entry, seen_generation }
105    }
106
107    pub fn wake(&self, key: &str) {
108        let keys = self.keys.lock().unwrap_or_else(|e| e.into_inner());
109        if let Some(entry) = keys.get(key) {
110            entry.generation.fetch_add(1, Ordering::AcqRel);
111            entry.notify.notify_waiters();
112        }
113    }
114
115    pub fn wake_many(&self, keys_to_wake: impl IntoIterator<Item = String>) {
116        for key in keys_to_wake {
117            self.wake(&key);
118        }
119    }
120
121    pub fn key_count(&self) -> usize {
122        self.keys.lock().unwrap_or_else(|e| e.into_inner()).len()
123    }
124}
125
126pub struct Registration<'a> {
127    registry: &'a LiveRegistry,
128    key: String,
129    entry: Arc<KeyEntry>,
130    seen_generation: u64,
131}
132
133impl Registration<'_> {
134    pub async fn wait(&self, timeout: Duration) -> WaitOutcome {
135        let notified = self.entry.notify.notified();
136        tokio::pin!(notified);
137        notified.as_mut().enable();
138
139        if self.entry.generation.load(Ordering::Acquire) != self.seen_generation {
140            return WaitOutcome::Woken;
141        }
142        match tokio::time::timeout(timeout, notified).await {
143            Ok(()) => WaitOutcome::Woken,
144            Err(_) => WaitOutcome::TimedOut,
145        }
146    }
147}
148
149impl Drop for Registration<'_> {
150    fn drop(&mut self) {
151        if self.entry.interest.fetch_sub(1, Ordering::AcqRel) == 1 {
152            let mut keys = self.registry.keys.lock().unwrap_or_else(|e| e.into_inner());
153            if let Some(current) = keys.get(&self.key) {
154                if Arc::ptr_eq(current, &self.entry) && current.interest.load(Ordering::Acquire) == 0 {
155                    keys.remove(&self.key);
156                }
157            }
158        }
159    }
160}
161
162#[cfg(test)]
163mod tests {
164    use super::*;
165
166    #[tokio::test]
167    async fn waking_one_user_key_does_not_wake_a_different_users_waiter() {
168        let registry = LiveRegistry::new();
169        let w1 = registry.register("user:1");
170        let w2 = registry.register("user:2");
171        registry.wake("user:1");
172        assert_eq!(w1.wait(Duration::from_secs(5)).await, WaitOutcome::Woken);
173        assert_eq!(w2.wait(Duration::from_millis(20)).await, WaitOutcome::TimedOut);
174    }
175
176    #[tokio::test]
177    async fn two_concurrent_waiters_on_the_same_key_both_wake() {
178        let registry = LiveRegistry::new();
179        let w1 = registry.register("user:1");
180        let w2 = registry.register("user:1");
181        let (o1, o2, ()) = tokio::join!(w1.wait(Duration::from_secs(5)), w2.wait(Duration::from_secs(5)), async {
182            for _ in 0..4 {
183                tokio::task::yield_now().await;
184            }
185            registry.wake("user:1");
186        });
187        assert_eq!(o1, WaitOutcome::Woken);
188        assert_eq!(o2, WaitOutcome::Woken);
189    }
190
191    #[tokio::test]
192    async fn wake_between_register_and_wait_is_not_lost() {
193        let registry = LiveRegistry::new();
194        let w = registry.register("user:1");
195        registry.wake("user:1");
196        let outcome = tokio::time::timeout(Duration::from_millis(50), w.wait(Duration::from_secs(5)))
197            .await
198            .expect("wake must not be lost");
199        assert_eq!(outcome, WaitOutcome::Woken);
200    }
201
202    #[tokio::test]
203    async fn a_previous_wake_is_not_replayed_onto_the_next_registration() {
204        let registry = LiveRegistry::new();
205        let w1 = registry.register("user:1");
206        registry.wake("user:1");
207        assert_eq!(w1.wait(Duration::from_secs(5)).await, WaitOutcome::Woken);
208        drop(w1);
209        let w2 = registry.register("user:1");
210        registry.wake("user:1");
211        assert_eq!(w2.wait(Duration::from_secs(5)).await, WaitOutcome::Woken);
212    }
213
214    #[tokio::test]
215    async fn wait_times_out_without_a_wake() {
216        let registry = LiveRegistry::new();
217        let w = registry.register("user:1");
218        assert_eq!(w.wait(Duration::from_millis(20)).await, WaitOutcome::TimedOut);
219    }
220
221    #[tokio::test]
222    async fn idle_keys_are_pruned() {
223        let registry = LiveRegistry::new();
224        assert_eq!(registry.key_count(), 0);
225        let w = registry.register("user:1");
226        assert_eq!(registry.key_count(), 1);
227        assert_eq!(w.wait(Duration::from_millis(20)).await, WaitOutcome::TimedOut);
228        drop(w);
229        assert_eq!(registry.key_count(), 0);
230        registry.wake("user:1");
231    }
232
233    #[test]
234    fn entry_pruned_after_last_guard_drops() {
235        let registry = LiveRegistry::new();
236        let w1 = registry.register("user:1");
237        let w2 = registry.register("user:1");
238        assert_eq!(registry.key_count(), 1);
239        drop(w1);
240        assert_eq!(registry.key_count(), 1);
241        drop(w2);
242        assert_eq!(registry.key_count(), 0);
243    }
244}