1use crate::{Consumer, Delivery};
10use commonware_utils::futures::{AbortablePool, Aborter};
11use futures::future::Aborted;
12use std::collections::{hash_map::Entry as HashMapEntry, HashMap};
13
14#[derive(Clone, Debug, Eq, PartialEq)]
16pub struct Completion<K, S, Context = ()> {
17 pub context: Context,
19
20 pub delivery: Delivery<K, S>,
22
23 pub valid: bool,
25}
26
27struct Response<Context, V> {
29 context: Context,
30 value: V,
31 accepted: bool,
32}
33
34struct ActiveDelivery {
36 generation: u64,
37 _aborter: Aborter,
38}
39
40struct PooledCompletion<K, S, Context> {
42 generation: u64,
43 completion: Completion<K, S, Context>,
44}
45
46struct Entry<Context, V, State> {
48 delivery: Option<ActiveDelivery>,
49 response: Option<Response<Context, V>>,
50 state: Option<State>,
51}
52
53impl<Context, V, State> Entry<Context, V, State> {
54 const fn new(state: State) -> Self {
55 Self {
56 delivery: None,
57 response: None,
58 state: Some(state),
59 }
60 }
61}
62
63pub struct Tracker<Con, Context = (), State = ()>
71where
72 Con: Consumer,
73 Con::Value: Clone + Send + 'static,
74 Context: Clone + Send + 'static,
75{
76 entries: HashMap<Con::Key, Entry<Context, Con::Value, State>>,
77 deliveries: AbortablePool<PooledCompletion<Con::Key, Con::Subscriber, Context>>,
78 next_generation: u64,
79 consumer: Con,
80}
81
82impl<Con, Context, State> Tracker<Con, Context, State>
83where
84 Con: Consumer,
85 Con::Value: Clone + Send + 'static,
86 Context: Clone + Send + 'static,
87{
88 pub fn new(consumer: Con) -> Self {
90 Self {
91 entries: HashMap::new(),
92 deliveries: AbortablePool::default(),
93 next_generation: 0,
94 consumer,
95 }
96 }
97
98 pub fn contains(&self, key: &Con::Key) -> bool {
100 self.entries.contains_key(key)
101 }
102
103 pub(crate) fn insert_with_state(&mut self, key: Con::Key, state: State) -> bool {
108 match self.entries.entry(key) {
109 HashMapEntry::Vacant(entry) => {
110 entry.insert(Entry::new(state));
111 true
112 }
113 HashMapEntry::Occupied(_) => false,
114 }
115 }
116
117 pub fn remove(&mut self, key: &Con::Key) -> bool {
122 self.entries.remove(key).is_some()
123 }
124
125 pub(crate) fn remove_with_state(&mut self, key: &Con::Key) -> Option<Option<State>> {
130 self.entries.remove(key).map(|entry| entry.state)
131 }
132
133 pub(crate) fn take_state(&mut self, key: &Con::Key) -> Option<State> {
137 self.entries
138 .get_mut(key)
139 .and_then(|entry| entry.state.take())
140 }
141
142 pub fn retain<F: FnMut(&Con::Key) -> bool>(&mut self, mut predicate: F) -> usize {
147 self.entries.extract_if(|key, _| !predicate(key)).count()
148 }
149
150 pub fn drain(&mut self) -> usize {
154 let count = self.entries.len();
155 self.entries.clear();
156 count
157 }
158
159 pub fn deliver(
165 &mut self,
166 delivery: Delivery<Con::Key, Con::Subscriber>,
167 context: Context,
168 value: Con::Value,
169 ) {
170 let key = delivery.key.clone();
171 let entry = self.entries.get_mut(&key).expect("delivery entry");
172 entry.response = Some(Response {
173 context: context.clone(),
174 value: value.clone(),
175 accepted: false,
176 });
177 self.push_delivery(delivery, context, value);
178 }
179
180 pub fn redeliver(&mut self, delivery: Delivery<Con::Key, Con::Subscriber>) {
186 let key = delivery.key.clone();
187 let (context, value) = {
188 let entry = self.entries.get(&key).expect("delivery entry");
189 let response = entry.response.as_ref().expect("response");
190 assert!(response.accepted, "accepted response");
191 (response.context.clone(), response.value.clone())
192 };
193 self.push_delivery(delivery, context, value);
194 }
195
196 pub fn response_accepted(&self, key: &Con::Key) -> bool {
198 self.entries
199 .get(key)
200 .and_then(|entry| entry.response.as_ref())
201 .is_some_and(|response| response.accepted)
202 }
203
204 pub fn accept_response(&mut self, key: &Con::Key) {
208 let entry = self.entries.get_mut(key).expect("delivery entry");
209 let response = entry.response.as_mut().expect("response");
210 response.accepted = true;
211 }
212
213 pub fn discard_response(&mut self, key: &Con::Key) {
218 if let Some(entry) = self.entries.get_mut(key) {
219 entry.response = None;
220 }
221 }
222
223 pub async fn next_completion(
230 &mut self,
231 ) -> Result<Completion<Con::Key, Con::Subscriber, Context>, Aborted> {
232 let completed = self.deliveries.next_completed().await?;
233 let Some(entry) = self.entries.get_mut(&completed.completion.delivery.key) else {
234 return Err(Aborted);
235 };
236 if entry
237 .delivery
238 .as_ref()
239 .is_none_or(|delivery| delivery.generation != completed.generation)
240 {
241 return Err(Aborted);
242 }
243 entry.delivery = None;
244 Ok(completed.completion)
245 }
246
247 fn push_delivery(
249 &mut self,
250 delivery: Delivery<Con::Key, Con::Subscriber>,
251 context: Context,
252 value: Con::Value,
253 ) {
254 let generation = self.next_generation;
255 self.next_generation = self
256 .next_generation
257 .checked_add(1)
258 .expect("delivery generation overflow");
259 let key = delivery.key.clone();
260 let completed = delivery.clone();
261 let mut consumer = self.consumer.clone();
262 let receiver = consumer.deliver(delivery, value);
263 let aborter = self.deliveries.push(async move {
264 PooledCompletion {
265 generation,
266 completion: Completion {
267 context,
268 delivery: completed,
269 valid: receiver.await.unwrap_or(false),
270 },
271 }
272 });
273 let entry = self.entries.get_mut(&key).expect("delivery entry");
274 assert!(entry
275 .delivery
276 .replace(ActiveDelivery {
277 generation,
278 _aborter: aborter,
279 })
280 .is_none());
281 }
282}
283
284impl<Con, Context> Tracker<Con, Context>
285where
286 Con: Consumer,
287 Con::Value: Clone + Send + 'static,
288 Context: Clone + Send + 'static,
289{
290 pub fn insert(&mut self, key: Con::Key) -> bool {
295 self.insert_with_state(key, ())
296 }
297}
298
299#[cfg(test)]
300mod tests {
301 use super::*;
302 use crate::p2p::mocks::{Consumer as MockConsumer, Key as MockKey};
303 use bytes::Bytes;
304 use commonware_runtime::{deterministic::Runner, Runner as _};
305 use commonware_utils::{
306 channel::{fallible::FallibleExt, mpsc, oneshot},
307 non_empty_vec,
308 };
309
310 type TestTracker = Tracker<MockConsumer<MockKey, Bytes>, u8>;
311
312 fn delivery(key: MockKey) -> Delivery<MockKey, ()> {
313 Delivery {
314 key,
315 subscribers: non_empty_vec![((), tracing::Span::none())],
316 }
317 }
318
319 #[derive(Clone)]
320 struct PendingConsumer {
321 sender: mpsc::UnboundedSender<oneshot::Sender<bool>>,
322 }
323
324 impl PendingConsumer {
325 fn new() -> (Self, mpsc::UnboundedReceiver<oneshot::Sender<bool>>) {
326 let (sender, receiver) = mpsc::unbounded_channel();
327 (Self { sender }, receiver)
328 }
329 }
330
331 impl Consumer for PendingConsumer {
332 type Key = MockKey;
333 type Value = Bytes;
334 type Subscriber = ();
335
336 fn deliver(
337 &mut self,
338 _delivery: Delivery<Self::Key, Self::Subscriber>,
339 _value: Self::Value,
340 ) -> oneshot::Receiver<bool> {
341 let (sender, receiver) = oneshot::channel();
342 self.sender.send_lossy(sender);
343 receiver
344 }
345 }
346
347 #[test]
348 fn test_insert_contains_remove_round_trip() {
349 let runner = Runner::default();
350 runner.start(|_| async move {
351 let mut tracker = TestTracker::new(MockConsumer::dummy());
352
353 assert!(!tracker.contains(&MockKey(1)));
354 assert!(tracker.insert(MockKey(1)));
355 assert!(tracker.contains(&MockKey(1)));
356
357 assert!(!tracker.insert(MockKey(1)));
358 assert!(tracker.remove(&MockKey(1)));
359 assert!(!tracker.contains(&MockKey(1)));
360 assert!(!tracker.remove(&MockKey(1)));
361 });
362 }
363
364 #[test]
365 fn test_deliver_completes_with_context_and_consumer_result() {
366 let runner = Runner::default();
367 runner.start(|_| async move {
368 let (consumer, mut events) = MockConsumer::<MockKey, Bytes>::new();
369 let mut tracker = TestTracker::new(consumer);
370 let key = MockKey(7);
371 let value = Bytes::from("data");
372
373 tracker.insert(key.clone());
374 tracker.deliver(delivery(key.clone()), 9, value.clone());
375
376 let completed = tracker
377 .next_completion()
378 .await
379 .expect("delivery should complete");
380 assert_eq!(completed.context, 9);
381 assert_eq!(completed.delivery.key, key);
382 assert!(completed.valid);
383
384 let (delivered_key, delivered_value) = events.recv().await.unwrap();
385 assert_eq!(delivered_key, key);
386 assert_eq!(delivered_value, value);
387 });
388 }
389
390 #[test]
391 fn test_remove_aborts_in_flight_delivery() {
392 let runner = Runner::default();
393 runner.start(|_| async move {
394 let (consumer, _events) = MockConsumer::<MockKey, Bytes>::new();
395 let mut tracker = TestTracker::new(consumer);
396 let key = MockKey(1);
397
398 tracker.insert(key.clone());
399 tracker.deliver(delivery(key.clone()), 2, Bytes::from("v"));
400 assert!(tracker.remove(&key));
401
402 assert!(matches!(tracker.next_completion().await, Err(Aborted)));
403 });
404 }
405
406 #[test]
407 fn test_stale_same_key_completion_does_not_clear_new_delivery() {
408 let runner = Runner::default();
409 runner.start(|_| async move {
410 let (consumer, mut senders) = PendingConsumer::new();
411 let mut tracker = Tracker::<PendingConsumer, u8>::new(consumer);
412 let key = MockKey(1);
413
414 tracker.insert(key.clone());
415 tracker.deliver(delivery(key.clone()), 1, Bytes::from("old"));
416 let old_sender = senders.recv().await.unwrap();
417 old_sender.send(true).unwrap();
418 let stale = tracker.deliveries.next_completed().await.unwrap();
419
420 assert!(tracker.remove(&key));
421 tracker.insert(key.clone());
422 tracker.deliver(delivery(key.clone()), 2, Bytes::from("new"));
423 let new_sender = senders.recv().await.unwrap();
424
425 let _stale_aborter = tracker.deliveries.push(async move { stale });
426 assert!(matches!(tracker.next_completion().await, Err(Aborted)));
427
428 new_sender.send(true).unwrap();
429 let completed = tracker
430 .next_completion()
431 .await
432 .expect("new delivery should complete");
433 assert_eq!(completed.context, 2);
434 assert_eq!(completed.delivery.key, key);
435 assert!(completed.valid);
436 });
437 }
438
439 #[test]
440 fn test_redeliver_reuses_accepted_response_for_new_subscribers() {
441 let runner = Runner::default();
442 runner.start(|_| async move {
443 let (consumer, mut events) = MockConsumer::<MockKey, Bytes>::new();
444 let mut tracker = TestTracker::new(consumer);
445 let key = MockKey(5);
446 let value = Bytes::from("first");
447
448 tracker.insert(key.clone());
449 tracker.deliver(delivery(key.clone()), 3, value.clone());
450
451 let completed = tracker
452 .next_completion()
453 .await
454 .expect("first delivery should complete");
455 assert!(completed.valid);
456 tracker.accept_response(&key);
457 assert!(tracker.response_accepted(&key));
458
459 tracker.redeliver(delivery(key.clone()));
460 let redelivered = tracker
461 .next_completion()
462 .await
463 .expect("redelivery should complete");
464 assert_eq!(redelivered.context, 3);
465 assert_eq!(redelivered.delivery.key, key);
466 assert!(redelivered.valid);
467
468 let first = events.recv().await.unwrap();
469 let second = events.recv().await.unwrap();
470 assert_eq!(first, (key.clone(), value.clone()));
471 assert_eq!(second, (key, value));
472 });
473 }
474
475 #[test]
476 #[should_panic(expected = "accepted response")]
477 fn test_redeliver_requires_accepted_response() {
478 let runner = Runner::default();
479 runner.start(|_| async move {
480 let (consumer, _events) = MockConsumer::<MockKey, Bytes>::new();
481 let mut tracker = TestTracker::new(consumer);
482 let key = MockKey(7);
483
484 tracker.insert(key.clone());
485 tracker.deliver(delivery(key.clone()), 3, Bytes::from("first"));
486 let completed = tracker
487 .next_completion()
488 .await
489 .expect("first delivery should complete");
490 assert!(completed.valid);
491
492 tracker.redeliver(delivery(key));
493 });
494 }
495
496 #[test]
497 fn test_rejected_response_can_be_discarded_and_replaced() {
498 let runner = Runner::default();
499 runner.start(|_| async move {
500 let (mut consumer, _events) = MockConsumer::<MockKey, Bytes>::new();
501 let key = MockKey(8);
502 consumer.add_expected(key.clone(), Bytes::from("good"));
503 let mut tracker = TestTracker::new(consumer);
504
505 tracker.insert(key.clone());
506 tracker.deliver(delivery(key.clone()), 1, Bytes::from("bad"));
507 let rejected = tracker
508 .next_completion()
509 .await
510 .expect("rejected delivery should complete");
511 assert!(!rejected.valid);
512
513 tracker.discard_response(&key);
514 assert!(!tracker.response_accepted(&key));
515 tracker.deliver(delivery(key.clone()), 2, Bytes::from("good"));
516
517 let accepted = tracker
518 .next_completion()
519 .await
520 .expect("accepted delivery should complete");
521 assert_eq!(accepted.context, 2);
522 assert!(accepted.valid);
523 });
524 }
525}