Skip to main content

commonware_p2p/utils/
limited.rs

1//! Rate-limited [`UnlimitedSender`] wrapper.
2
3use crate::{Recipients, UnlimitedSender};
4use commonware_actor::{Feedback, Unreliable};
5use commonware_cryptography::PublicKey;
6use commonware_runtime::{Clock, IoBufs, KeyedRateLimiter, Quota};
7use commonware_utils::{channel::ring, sync::Mutex};
8use futures::{FutureExt, StreamExt};
9use std::{cmp, fmt, sync::Arc, time::SystemTime};
10
11/// Provides peer snapshots for resolving [`Recipients::All`].
12pub trait Connected: Clone + Send + Sync + 'static {
13    type PublicKey: PublicKey;
14
15    /// Return the current peer snapshot.
16    fn peers(&self) -> Vec<Self::PublicKey> {
17        Vec::new()
18    }
19
20    /// Subscribe to peer updates.
21    ///
22    /// The receiver yields the current set of known peers whenever it changes.
23    /// New subscriptions should publish the current set promptly so callers do
24    /// not have to wait for the next membership change.
25    fn subscribe(&self) -> ring::Receiver<Vec<Self::PublicKey>>;
26}
27
28/// A wrapper around a [`UnlimitedSender`] that provides rate limiting with retry-time feedback.
29pub struct LimitedSender<E, S, P>
30where
31    E: Clock,
32    S: UnlimitedSender,
33    P: Connected<PublicKey = S::PublicKey>,
34{
35    sender: S,
36    state: Arc<Mutex<State<S::PublicKey, E>>>,
37    peers: P,
38}
39
40struct State<P: PublicKey, E: Clock> {
41    // Per-peer rate limiter shared by all clones
42    rate_limit: KeyedRateLimiter<P, E>,
43    // Latest peer updates from the source used for Recipients::All
44    peer_subscription: ring::Receiver<Vec<P>>,
45    // Snapshot used until the subscription yields a newer peer list
46    known_peers: Vec<P>,
47}
48
49impl<E, S, P> Clone for LimitedSender<E, S, P>
50where
51    E: Clock,
52    S: UnlimitedSender,
53    P: Connected<PublicKey = S::PublicKey>,
54{
55    fn clone(&self) -> Self {
56        Self {
57            sender: self.sender.clone(),
58            state: self.state.clone(),
59            peers: self.peers.clone(),
60        }
61    }
62}
63
64impl<E, S, P> fmt::Debug for LimitedSender<E, S, P>
65where
66    E: Clock,
67    S: UnlimitedSender,
68    P: Connected<PublicKey = S::PublicKey>,
69{
70    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
71        let known_peers = self.state.lock().known_peers.len();
72        f.debug_struct("LimitedSender")
73            .field("known_peers", &known_peers)
74            .finish_non_exhaustive()
75    }
76}
77
78impl<E, S, P> LimitedSender<E, S, P>
79where
80    E: Clock,
81    S: UnlimitedSender,
82    P: Connected<PublicKey = S::PublicKey>,
83{
84    /// Create a new [`LimitedSender`] with the given sender, [`Quota`], and peer source.
85    pub fn new(sender: S, quota: Quota, clock: E, peers: P) -> Self {
86        let state = Arc::new(Mutex::new(State {
87            rate_limit: KeyedRateLimiter::hashmap_with_clock(quota, clock),
88            peer_subscription: peers.subscribe(),
89            known_peers: peers.peers(),
90        }));
91        Self {
92            sender,
93            state,
94            peers,
95        }
96    }
97
98    /// Check that a given set of [`Recipients`] are within the rate limit.
99    ///
100    /// Returns a [`CheckedSender`] with only the recipients that are not
101    /// currently rate-limited. If _all_ recipients are rate-limited, returns
102    /// the earliest instant at which all recipients will be available.
103    pub fn check(
104        &mut self,
105        recipients: Recipients<S::PublicKey>,
106    ) -> Result<CheckedSender<'_, S>, SystemTime> {
107        let mut state = self.state.lock();
108        if matches!(&recipients, Recipients::All) {
109            if let Some(peers) = state.peer_subscription.next().now_or_never().flatten() {
110                state.known_peers = peers;
111                state.rate_limit.retain_recent();
112            }
113        }
114
115        let recipients = match recipients {
116            Recipients::One(peer) => match state.rate_limit.check_key(&peer) {
117                Ok(()) => Recipients::One(peer),
118                Err(not_until) => return Err(not_until.earliest_possible()),
119            },
120            Recipients::Some(peers) => {
121                let (allowed, max_retry) = filter_rate_limited(peers.iter(), &state.rate_limit);
122                if allowed.is_empty() {
123                    match max_retry {
124                        Some(retry) => return Err(retry),
125                        None => Recipients::Some(Vec::new()),
126                    }
127                } else {
128                    Recipients::Some(allowed)
129                }
130            }
131            Recipients::All => {
132                let (allowed, max_retry) =
133                    filter_rate_limited(state.known_peers.iter(), &state.rate_limit);
134                if allowed.is_empty() {
135                    match max_retry {
136                        Some(retry) => return Err(retry),
137                        None => Recipients::Some(Vec::new()),
138                    }
139                } else {
140                    Recipients::Some(allowed)
141                }
142            }
143        };
144        drop(state);
145
146        Ok(CheckedSender {
147            recipients,
148            sender: &mut self.sender,
149        })
150    }
151}
152
153/// Filters peers by rate limit, returning those that pass and the latest retry
154/// time among those that don't.
155pub(crate) fn filter_rate_limited<'a, K, C>(
156    peers: impl Iterator<Item = &'a K>,
157    rate_limit: &KeyedRateLimiter<K, C>,
158) -> (Vec<K>, Option<SystemTime>)
159where
160    K: PublicKey,
161    C: Clock,
162{
163    peers.fold(
164        (Vec::new(), None),
165        |(mut allowed, max_retry), p| match rate_limit.check_key(p) {
166            Ok(()) => {
167                allowed.push(p.clone());
168                (allowed, max_retry)
169            }
170            Err(not_until) => {
171                let earliest = not_until.earliest_possible();
172                let new_max = max_retry.map_or(earliest, |current| cmp::max(current, earliest));
173                (allowed, Some(new_max))
174            }
175        },
176    )
177}
178
179/// An exclusive reference to an [`UnlimitedSender`] with a pre-checked list of
180/// recipients that are not currently rate-limited.
181///
182/// A [`CheckedSender`] can only be acquired via [`LimitedSender::check`].
183#[derive(Debug)]
184pub struct CheckedSender<'a, S: UnlimitedSender> {
185    sender: &'a mut S,
186    recipients: Recipients<S::PublicKey>,
187}
188
189impl<'a, S: UnlimitedSender> CheckedSender<'a, S> {
190    /// Extracts the inner [`UnlimitedSender`] reference.
191    ///
192    /// # Warning
193    ///
194    /// Rate limiting has already been applied to the original recipients. Any
195    /// messages sent via the extracted sender will bypass the rate limiter.
196    #[commonware_macros::stability(ALPHA)]
197    pub(crate) fn into_inner(self) -> &'a mut S {
198        self.sender
199    }
200}
201
202impl<'a, S: UnlimitedSender> crate::CheckedSender for CheckedSender<'a, S> {
203    type PublicKey = S::PublicKey;
204
205    fn recipients(&self) -> Vec<Self::PublicKey> {
206        match &self.recipients {
207            Recipients::All => Vec::new(),
208            Recipients::Some(peers) => peers.clone(),
209            Recipients::One(peer) => vec![peer.clone()],
210        }
211    }
212
213    fn send(self, message: impl Into<IoBufs> + Send, priority: bool) -> Unreliable<Feedback> {
214        self.sender.send(self.recipients, message, priority)
215    }
216}
217
218#[cfg(test)]
219mod tests {
220    use super::*;
221    use crate::CheckedSender as _;
222    use commonware_cryptography::{ed25519, Signer as _};
223    use commonware_runtime::{deterministic::Runner, IoBuf, Quota, Runner as _};
224    use commonware_utils::{channel::ring, NZUsize, NZU32};
225    use futures::SinkExt;
226
227    type PublicKey = ed25519::PublicKey;
228    type SentMessage = (Recipients<PublicKey>, IoBuf, bool);
229
230    #[derive(Debug, Clone)]
231    struct MockSender {
232        sent: Arc<Mutex<Vec<SentMessage>>>,
233    }
234
235    impl MockSender {
236        fn new() -> Self {
237            Self {
238                sent: Arc::new(Mutex::new(Vec::new())),
239            }
240        }
241
242        fn sent_messages(&self) -> Vec<SentMessage> {
243            self.sent.lock().clone()
244        }
245    }
246
247    fn assert_sent_to(sender: &MockSender, index: usize, expected: &[PublicKey]) {
248        let messages = sender.sent_messages();
249        let Recipients::Some(sent) = &messages[index].0 else {
250            panic!("expected Recipients::Some");
251        };
252        assert_eq!(sent, expected);
253    }
254
255    impl UnlimitedSender for MockSender {
256        type PublicKey = PublicKey;
257
258        fn send(
259            &mut self,
260            recipients: Recipients<Self::PublicKey>,
261            message: impl Into<IoBufs> + Send,
262            priority: bool,
263        ) -> Unreliable<Feedback> {
264            let message = message.into().coalesce();
265            self.sent.lock().push((recipients, message, priority));
266            Unreliable::new(Feedback::Ok)
267        }
268    }
269
270    #[derive(Clone)]
271    struct MockPeers {
272        peers: Vec<PublicKey>,
273    }
274
275    #[derive(Clone)]
276    struct UpdatingPeers {
277        peers: Vec<PublicKey>,
278        receiver: Arc<Mutex<Option<ring::Receiver<Vec<PublicKey>>>>>,
279    }
280
281    impl MockPeers {
282        fn new() -> Self {
283            Self { peers: Vec::new() }
284        }
285
286        fn with_peers(peers: Vec<PublicKey>) -> Self {
287            Self { peers }
288        }
289    }
290
291    impl Connected for MockPeers {
292        type PublicKey = PublicKey;
293
294        fn peers(&self) -> Vec<Self::PublicKey> {
295            self.peers.clone()
296        }
297
298        fn subscribe(&self) -> ring::Receiver<Vec<Self::PublicKey>> {
299            let (_sender, receiver) = ring::channel(NZUsize!(16));
300            receiver
301        }
302    }
303
304    impl Connected for UpdatingPeers {
305        type PublicKey = PublicKey;
306
307        fn peers(&self) -> Vec<Self::PublicKey> {
308            self.peers.clone()
309        }
310
311        fn subscribe(&self) -> ring::Receiver<Vec<Self::PublicKey>> {
312            self.receiver
313                .lock()
314                .take()
315                .expect("subscription should only be created once")
316        }
317    }
318
319    fn key(seed: u64) -> PublicKey {
320        ed25519::PrivateKey::from_seed(seed).public_key()
321    }
322
323    fn quota_per_second(n: u32) -> Quota {
324        Quota::per_second(NZU32!(n))
325    }
326
327    #[test]
328    fn check_one_not_rate_limited() {
329        Runner::default().start(|context| async move {
330            let sender = MockSender::new();
331            let peers = MockPeers::new();
332            let mut limited = LimitedSender::new(sender, quota_per_second(10), context, peers);
333
334            let checked = limited.check(Recipients::One(key(1))).unwrap();
335            assert_eq!(
336                checked.send(IoBuf::from(b"hello"), false),
337                Unreliable::new(Feedback::Ok)
338            );
339        });
340    }
341
342    #[test]
343    fn check_one_rate_limited() {
344        Runner::default().start(|context| async move {
345            let sender = MockSender::new();
346            let peers = MockPeers::new();
347            let mut limited = LimitedSender::new(sender, quota_per_second(1), context, peers);
348
349            let peer = key(1);
350
351            // First check should succeed and consume the quota
352            let checked = limited.check(Recipients::One(peer.clone())).unwrap();
353            checked.send(IoBuf::from(b"first"), false);
354
355            // Second check should fail (rate limited)
356            let result = limited.check(Recipients::One(peer));
357            assert!(result.is_err());
358        });
359    }
360
361    #[test]
362    fn check_some_all_not_rate_limited() {
363        Runner::default().start(|context| async move {
364            let sender = MockSender::new();
365            let peers = MockPeers::new();
366            let mut limited =
367                LimitedSender::new(sender.clone(), quota_per_second(1), context, peers);
368
369            let peers_list = vec![key(1), key(2), key(3)];
370            let checked = limited.check(Recipients::Some(peers_list)).unwrap();
371            assert_eq!(
372                checked.send(IoBuf::from(b"hello"), false),
373                Unreliable::new(Feedback::Ok)
374            );
375            assert_sent_to(&sender, 0, &[key(1), key(2), key(3)]);
376        });
377    }
378
379    #[test]
380    fn check_some_filters_rate_limited_peers() {
381        Runner::default().start(|context| async move {
382            let sender = MockSender::new();
383            let peers = MockPeers::new();
384            let mut limited =
385                LimitedSender::new(sender.clone(), quota_per_second(1), context, peers);
386
387            let peer1 = key(1);
388            let peer2 = key(2);
389            let peer3 = key(3);
390
391            // Rate limit peer1 by sending to it first
392            let checked = limited.check(Recipients::One(peer1.clone())).unwrap();
393            checked.send(IoBuf::from(b"limit"), false);
394
395            // Now check with all three peers - peer1 should be filtered out
396            let expected = vec![peer2.clone(), peer3.clone()];
397            let checked = limited
398                .check(Recipients::Some(vec![peer1, peer2, peer3]))
399                .unwrap();
400            checked.send(IoBuf::from(b"filtered"), false);
401            assert_sent_to(&sender, 1, &expected);
402        });
403    }
404
405    #[test]
406    fn check_some_all_rate_limited_returns_error() {
407        Runner::default().start(|context| async move {
408            let sender = MockSender::new();
409            let peers = MockPeers::new();
410            let mut limited = LimitedSender::new(sender, quota_per_second(1), context, peers);
411
412            let peer1 = key(1);
413            let peer2 = key(2);
414
415            // Rate limit both peers
416            limited
417                .check(Recipients::One(peer1.clone()))
418                .unwrap()
419                .send(IoBuf::from(b"limit1"), false);
420
421            limited
422                .check(Recipients::One(peer2.clone()))
423                .unwrap()
424                .send(IoBuf::from(b"limit2"), false);
425
426            // Now both are rate limited - should return error with retry time
427            assert!(limited.check(Recipients::Some(vec![peer1, peer2])).is_err());
428        });
429    }
430
431    #[test]
432    fn check_some_empty_returns_as_is() {
433        Runner::default().start(|context| async move {
434            let sender = MockSender::new();
435            let peers = MockPeers::new();
436            let mut limited = LimitedSender::new(sender, quota_per_second(10), context, peers);
437
438            // Empty recipients should pass through
439            limited.check(Recipients::Some(Vec::new())).unwrap();
440        });
441    }
442
443    #[test]
444    fn check_all_uses_known_peers() {
445        Runner::default().start(|context| async move {
446            let sender = MockSender::new();
447            let peers = MockPeers::new();
448            let mut limited =
449                LimitedSender::new(sender.clone(), quota_per_second(10), context, peers);
450
451            // No known peers yet
452            let checked = limited.check(Recipients::All).unwrap();
453            assert!(crate::CheckedSender::recipients(&checked).is_empty());
454            checked.send(IoBuf::from(b"empty"), false);
455
456            // Verify that the sender received the message with empty Recipients::Some.
457            assert_sent_to(&sender, 0, &[]);
458        });
459    }
460
461    #[test]
462    fn check_all_filters_rate_limited_known_peers() {
463        Runner::default().start(|context| async move {
464            let sender = MockSender::new();
465            let peer1 = key(1);
466            let peer2 = key(2);
467            let peers = MockPeers::with_peers(vec![peer1.clone(), peer2.clone()]);
468            let mut limited =
469                LimitedSender::new(sender.clone(), quota_per_second(1), context, peers);
470
471            // Rate limit peer1
472            limited
473                .check(Recipients::One(peer1))
474                .unwrap()
475                .send(IoBuf::from(b"limit"), false);
476
477            // Check All should filter out peer1
478            let checked = limited.check(Recipients::All).unwrap();
479            checked.send(IoBuf::from(b"filtered"), false);
480            assert_sent_to(&sender, 1, &[peer2]);
481        });
482    }
483
484    #[test]
485    fn check_all_returns_error_when_all_known_peers_rate_limited() {
486        Runner::default().start(|context| async move {
487            let sender = MockSender::new();
488            let peer1 = key(1);
489            let peer2 = key(2);
490            let peers = MockPeers::with_peers(vec![peer1.clone(), peer2.clone()]);
491            let mut limited = LimitedSender::new(sender, quota_per_second(1), context, peers);
492
493            // Rate limit both peers
494            limited
495                .check(Recipients::One(peer1))
496                .unwrap()
497                .send(IoBuf::from(b"limit1"), false);
498
499            limited
500                .check(Recipients::One(peer2))
501                .unwrap()
502                .send(IoBuf::from(b"limit2"), false);
503
504            // Check All should fail since all known peers are rate limited
505            assert!(limited.check(Recipients::All).is_err());
506        });
507    }
508
509    #[test]
510    fn clone_shares_peer_updates() {
511        Runner::default().start(|context| async move {
512            let sender = MockSender::new();
513            let initial = key(1);
514            let updated = key(2);
515            let (updates, receiver) = ring::channel(NZUsize!(1));
516            let peers = UpdatingPeers {
517                peers: vec![initial],
518                receiver: Arc::new(Mutex::new(Some(receiver))),
519            };
520            let mut limited1 = LimitedSender::new(sender, quota_per_second(10), context, peers);
521
522            let mut limited2 = limited1.clone();
523            let mut updates = updates;
524            updates.send(vec![updated.clone()]).await.unwrap();
525
526            let checked = limited2.check(Recipients::All).unwrap();
527            assert_eq!(crate::CheckedSender::recipients(&checked), vec![updated]);
528
529            let checked = limited1.check(Recipients::All).unwrap();
530            assert_eq!(crate::CheckedSender::recipients(&checked), vec![key(2)]);
531        });
532    }
533
534    #[test]
535    fn checked_sender_sends_with_priority() {
536        Runner::default().start(|context| async move {
537            let sender = MockSender::new();
538            let peers = MockPeers::new();
539            let mut limited =
540                LimitedSender::new(sender.clone(), quota_per_second(10), context, peers);
541
542            let peer = key(1);
543            limited
544                .check(Recipients::One(peer))
545                .unwrap()
546                .send(IoBuf::from(b"priority"), true);
547
548            let messages = sender.sent_messages();
549            assert_eq!(messages.len(), 1);
550            assert!(messages[0].2); // priority flag
551        });
552    }
553
554    #[test]
555    fn rate_limit_shared_across_clones() {
556        Runner::default().start(|context| async move {
557            let sender = MockSender::new();
558            let peers = MockPeers::new();
559            let mut limited1 = LimitedSender::new(sender, quota_per_second(1), context, peers);
560            let mut limited2 = limited1.clone();
561
562            let peer = key(1);
563
564            // Rate limit peer via first instance
565            limited1
566                .check(Recipients::One(peer.clone()))
567                .unwrap()
568                .send(IoBuf::from(b"limit"), false);
569
570            // Second instance should see the rate limit
571            assert!(limited2.check(Recipients::One(peer)).is_err());
572        });
573    }
574}