1use 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
11pub trait Connected: Clone + Send + Sync + 'static {
13 type PublicKey: PublicKey;
14
15 fn peers(&self) -> Vec<Self::PublicKey> {
17 Vec::new()
18 }
19
20 fn subscribe(&self) -> ring::Receiver<Vec<Self::PublicKey>>;
26}
27
28pub 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 rate_limit: KeyedRateLimiter<P, E>,
43 peer_subscription: ring::Receiver<Vec<P>>,
45 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 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 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
153pub(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#[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 #[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 let checked = limited.check(Recipients::One(peer.clone())).unwrap();
353 checked.send(IoBuf::from(b"first"), false);
354
355 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 let checked = limited.check(Recipients::One(peer1.clone())).unwrap();
393 checked.send(IoBuf::from(b"limit"), false);
394
395 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 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 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 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 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 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 limited
473 .check(Recipients::One(peer1))
474 .unwrap()
475 .send(IoBuf::from(b"limit"), false);
476
477 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 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 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); });
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 limited1
566 .check(Recipients::One(peer.clone()))
567 .unwrap()
568 .send(IoBuf::from(b"limit"), false);
569
570 assert!(limited2.check(Recipients::One(peer)).is_err());
572 });
573 }
574}