Skip to main content

mtorrent_core/data/
piece_tracker.rs

1use crate::pwp;
2use derive_more::Debug;
3use std::borrow::Borrow;
4use std::collections::{BTreeMap, HashMap, HashSet};
5use std::net::SocketAddr;
6
7#[derive(PartialEq, Eq, PartialOrd, Ord, Hash, Clone, Copy, Debug)]
8struct PieceIndex(usize);
9
10impl Borrow<usize> for PieceIndex {
11    fn borrow(&self) -> &usize {
12        &self.0
13    }
14}
15
16fn available_pieces(bitfield: &pwp::Bitfield) -> impl Iterator<Item = usize> + Clone + '_ {
17    bitfield
18        .iter()
19        .enumerate()
20        .filter_map(|(index, bit)| (bit == true).then_some(index))
21}
22
23/// Keeps track of missing pieces (as opposed to blocks) and their owners.
24#[derive(Debug)]
25pub struct PieceTracker {
26    piece_index_to_owners: HashMap<PieceIndex, HashSet<SocketAddr>>,
27    owners_to_piece_indices: HashMap<SocketAddr, pwp::Bitfield>,
28
29    owner_count_to_piece_indices: BTreeMap<usize, HashSet<PieceIndex>>,
30    piece_count_to_owners: BTreeMap<usize, HashSet<SocketAddr>>,
31
32    #[debug(skip)]
33    bitfield_factory: Box<dyn Fn() -> pwp::Bitfield>,
34}
35
36impl PieceTracker {
37    pub fn new(piece_count: usize) -> Self {
38        let indices = (0..piece_count).map(PieceIndex).collect::<HashSet<PieceIndex>>();
39        Self {
40            piece_index_to_owners: indices.iter().map(|index| (*index, HashSet::new())).collect(),
41            owners_to_piece_indices: HashMap::new(),
42            owner_count_to_piece_indices: BTreeMap::from([(0usize, indices)]),
43            piece_count_to_owners: BTreeMap::new(),
44            bitfield_factory: Box::new(move || pwp::Bitfield::repeat(false, piece_count)),
45        }
46    }
47
48    /// Get an iterator over the not-yet-downloaded pieces, ordered by the number of
49    /// peers that own each piece, such that pieces with fewest owners are yielded first.
50    pub fn missing_pieces_rarest_first(&self) -> impl Iterator<Item = usize> + '_ {
51        self.owner_count_to_piece_indices
52            .iter()
53            .skip_while(|(count, _indices)| **count == 0usize)
54            .flat_map(|(_count, indices)| indices.iter().map(|i| i.0))
55    }
56
57    #[cfg(test)]
58    pub fn get_poorest_peers(&self) -> impl Iterator<Item = &SocketAddr> + Clone {
59        self.piece_count_to_owners.values().flat_map(HashSet::iter)
60    }
61
62    /// Get addresses of all peers that own a particular piece.
63    pub fn get_piece_owners(
64        &self,
65        piece_index: usize,
66    ) -> impl Iterator<Item = &SocketAddr> + Clone {
67        self.piece_index_to_owners.get(&piece_index).into_iter().flat_map(HashSet::iter)
68    }
69
70    /// Get piece indices of all pieces owned by a particular peer.
71    pub fn get_peer_pieces(&self, peer: &SocketAddr) -> impl Iterator<Item = usize> + Clone + '_ {
72        self.owners_to_piece_indices.get(peer).into_iter().flat_map(available_pieces)
73    }
74
75    /// Check if `peer` ownes `piece_index`.
76    pub fn has_peer_piece(&self, peer: &SocketAddr, piece_index: usize) -> bool {
77        self.owners_to_piece_indices
78            .get(peer)
79            .and_then(|pieces| pieces.get(piece_index))
80            .is_some_and(|piece_present| piece_present == true)
81    }
82
83    /// Record that `piece_owner` owns `piece_index`.
84    pub fn add_single_record(&mut self, piece_owner: &SocketAddr, piece_index: usize) -> bool {
85        let piece_index = PieceIndex(piece_index);
86
87        if let Some(piece_owners) = self.piece_index_to_owners.get_mut(&piece_index) {
88            let peer_pieces = self
89                .owners_to_piece_indices
90                .entry(*piece_owner)
91                .or_insert_with(&self.bitfield_factory);
92
93            let updated_peer_pieces = !peer_pieces.replace(piece_index.0, true);
94            let updated_piece_owners = piece_owners.insert(*piece_owner);
95            assert_eq!(updated_piece_owners, updated_peer_pieces, "Inconsistent internal state");
96
97            if updated_peer_pieces {
98                self.change_owner_count_for_piece(piece_index, |prev_count| prev_count + 1);
99                self.change_piece_count_for_owner(piece_owner, |prev_count| prev_count + 1);
100                true
101            } else {
102                false
103            }
104        } else {
105            // forgotten (i.e. already downloaded) or invalid piece
106            false
107        }
108    }
109
110    /// Record that `peer` owns pieces represented by the `bitfield`. This won't invalidate any
111    /// previous records for the same peer, i.e. it will never remove pieces.
112    pub fn add_bitfield_record(&mut self, peer: &SocketAddr, bitfield: &pwp::Bitfield) {
113        for piece_index in available_pieces(bitfield) {
114            self.add_single_record(peer, piece_index);
115        }
116    }
117
118    /// Erase all records pertaining to the specified peer.
119    pub fn forget_peer(&mut self, peer: &SocketAddr) {
120        if let Some(pieces) = self.owners_to_piece_indices.remove(peer) {
121            for piece_index in available_pieces(&pieces).map(PieceIndex) {
122                let owners = self
123                    .piece_index_to_owners
124                    .get_mut(&piece_index)
125                    .expect("Invalid internal state");
126                owners.remove(peer);
127                self.change_owner_count_for_piece(piece_index, |prev_count| {
128                    prev_count.saturating_sub(1)
129                });
130            }
131            let removed =
132                self.piece_count_to_owners.iter_mut().find_map(|(piece_count, owners)| {
133                    let owner_count = owners.len();
134                    owners.remove(peer).then_some((piece_count, owner_count - 1))
135                });
136            if let Some((&piece_count, 0)) = removed {
137                self.piece_count_to_owners.remove(&piece_count);
138            }
139        }
140    }
141
142    /// Erase all records pertaining to the specified piece.
143    pub fn forget_piece(&mut self, piece_index: usize) {
144        if let Some(owners) = self.piece_index_to_owners.remove(&piece_index) {
145            for owner in owners {
146                let pieces =
147                    self.owners_to_piece_indices.get_mut(&owner).expect("Invalid internal state");
148                pieces.set(piece_index, false);
149                self.change_piece_count_for_owner(&owner, |prev_count| {
150                    prev_count.saturating_sub(1)
151                });
152            }
153            let removed =
154                self.owner_count_to_piece_indices.iter_mut().find_map(|(owner_count, pieces)| {
155                    let indices_count = pieces.len();
156                    pieces.remove(&piece_index).then_some((owner_count, indices_count - 1))
157                });
158            if let Some((&owner_count, 0)) = removed {
159                self.owner_count_to_piece_indices.remove(&owner_count);
160            }
161        }
162    }
163
164    fn change_owner_count_for_piece<F>(&mut self, piece_index: PieceIndex, op: F)
165    where
166        F: FnOnce(usize) -> usize,
167    {
168        if let Some((current_owner_count, indices)) = self
169            .owner_count_to_piece_indices
170            .iter_mut()
171            .find_map(|(count, indices)| indices.remove(&piece_index).then_some((*count, indices)))
172        {
173            if indices.is_empty() {
174                self.owner_count_to_piece_indices.remove(&current_owner_count);
175            }
176            let new_owner_count = op(current_owner_count);
177            self.owner_count_to_piece_indices
178                .entry(new_owner_count)
179                .and_modify(|indices| {
180                    indices.insert(piece_index);
181                })
182                .or_insert_with(|| HashSet::from([piece_index]));
183        }
184    }
185
186    fn change_piece_count_for_owner<F>(&mut self, peer: &SocketAddr, op: F)
187    where
188        F: FnOnce(usize) -> usize,
189    {
190        let current_piece_count = if let Some((current_piece_count, owners)) = self
191            .piece_count_to_owners
192            .iter_mut()
193            .find_map(|(count, owners)| owners.remove(peer).then_some((*count, owners)))
194        {
195            if owners.is_empty() {
196                self.piece_count_to_owners.remove(&current_piece_count);
197            }
198            current_piece_count
199        } else {
200            0
201        };
202        let new_piece_count = op(current_piece_count);
203        if new_piece_count > 0 {
204            self.piece_count_to_owners
205                .entry(new_piece_count)
206                .and_modify(|owners| {
207                    owners.insert(*peer);
208                })
209                .or_insert_with(|| HashSet::from([*peer]));
210        }
211    }
212}
213
214#[cfg(test)]
215mod tests {
216    use super::*;
217    use bitvec::prelude::*;
218    use std::net::{Ipv4Addr, SocketAddrV4};
219
220    fn ip(port: u16) -> SocketAddr {
221        SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, port))
222    }
223
224    #[test]
225    fn test_add_records_and_get_owners() {
226        let mut pa = PieceTracker::new(4);
227        assert!(pa.get_piece_owners(4).next().is_none());
228        assert_eq!(0, pa.get_piece_owners(0).count());
229        assert_eq!(0, pa.get_piece_owners(1).count());
230        assert_eq!(0, pa.get_piece_owners(2).count());
231        assert_eq!(0, pa.get_piece_owners(3).count());
232
233        let added = pa.add_single_record(&ip(6000), 3);
234        assert!(added);
235        let added = pa.add_single_record(&ip(6000), 3);
236        assert!(!added);
237        assert_eq!(0, pa.get_piece_owners(0).count());
238        assert_eq!(0, pa.get_piece_owners(1).count());
239        assert_eq!(0, pa.get_piece_owners(2).count());
240        assert_eq!(HashSet::from([&ip(6000)]), pa.get_piece_owners(3).collect());
241
242        let added = pa.add_single_record(&ip(6666), 4);
243        assert!(!added);
244
245        let added = pa.add_single_record(&ip(6000), 2);
246        assert!(added);
247        assert_eq!(0, pa.get_piece_owners(0).count());
248        assert_eq!(0, pa.get_piece_owners(1).count());
249        assert_eq!(HashSet::from([&ip(6000)]), pa.get_piece_owners(2).collect());
250        assert_eq!(HashSet::from([&ip(6000)]), pa.get_piece_owners(3).collect());
251
252        pa.add_bitfield_record(&ip(6001), &BitVec::from_bitslice(bits![u8, Msb0; 1, 0, 0, 1]));
253        assert_eq!(HashSet::from([&ip(6001)]), pa.get_piece_owners(0).collect());
254        assert_eq!(0, pa.get_piece_owners(1).count());
255        assert_eq!(HashSet::from([&ip(6000)]), pa.get_piece_owners(2).collect());
256        assert_eq!(HashSet::from([&ip(6000), &ip(6001)]), pa.get_piece_owners(3).collect());
257
258        pa.add_bitfield_record(&ip(6002), &BitVec::repeat(true, 8));
259        assert_eq!(HashSet::from([&ip(6001), &ip(6002)]), pa.get_piece_owners(0).collect());
260        assert_eq!(HashSet::from([&ip(6002)]), pa.get_piece_owners(1).collect());
261        assert_eq!(HashSet::from([&ip(6000), &ip(6002)]), pa.get_piece_owners(2).collect());
262        assert_eq!(
263            HashSet::from([&ip(6000), &ip(6001), &ip(6002)]),
264            pa.get_piece_owners(3).collect()
265        );
266    }
267
268    #[test]
269    fn test_add_records_and_get_rarest_and_poorest() {
270        let mut pa = PieceTracker::new(4);
271        assert!(pa.missing_pieces_rarest_first().next().is_none());
272
273        pa.add_bitfield_record(&ip(6000), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 1, 0]));
274        pa.add_bitfield_record(&ip(6001), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 0, 0]));
275        pa.add_bitfield_record(&ip(6002), &BitVec::from_bitslice(bits![u8, Msb0; 1, 0, 0, 0]));
276        {
277            let mut rarest = pa.missing_pieces_rarest_first();
278            assert_eq!(2, rarest.next().unwrap());
279            assert_eq!(1, rarest.next().unwrap());
280            assert_eq!(0, rarest.next().unwrap());
281            assert!(rarest.next().is_none());
282        }
283        {
284            let mut poorest = pa.get_poorest_peers();
285            assert_eq!(&ip(6002), poorest.next().unwrap());
286            assert_eq!(&ip(6001), poorest.next().unwrap());
287            assert_eq!(&ip(6000), poorest.next().unwrap());
288            assert!(poorest.next().is_none());
289        }
290
291        pa.add_bitfield_record(&ip(6003), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 1, 1]));
292        {
293            let mut rarest = pa.missing_pieces_rarest_first();
294            assert_eq!(3, rarest.next().unwrap());
295            assert_eq!(2, rarest.next().unwrap());
296            assert_eq!(1, rarest.next().unwrap());
297            assert_eq!(0, rarest.next().unwrap());
298            assert!(rarest.next().is_none());
299        }
300        {
301            let mut poorest = pa.get_poorest_peers();
302            assert_eq!(&ip(6002), poorest.next().unwrap());
303            assert_eq!(&ip(6001), poorest.next().unwrap());
304            assert_eq!(&ip(6000), poorest.next().unwrap());
305            assert_eq!(&ip(6003), poorest.next().unwrap());
306            assert!(poorest.next().is_none());
307        }
308
309        pa.add_single_record(&ip(6002), 1);
310        {
311            let mut rarest = pa.missing_pieces_rarest_first();
312            assert_eq!(3, rarest.next().unwrap());
313            assert_eq!(2, rarest.next().unwrap());
314            assert_eq!(HashSet::from([0, 1]), rarest.collect());
315        }
316        {
317            let mut richest = pa.get_poorest_peers().collect::<Vec<_>>().into_iter().rev();
318            assert_eq!(&ip(6003), richest.next().unwrap());
319            assert_eq!(&ip(6000), richest.next().unwrap());
320            assert_eq!(HashSet::from([&ip(6001), &ip(6002)]), richest.collect());
321        }
322    }
323
324    #[test]
325    fn test_add_records_and_forget_piece() {
326        let mut pa = PieceTracker::new(4);
327        pa.add_bitfield_record(&ip(6000), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 1, 1]));
328        pa.add_bitfield_record(&ip(6001), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 1, 0]));
329        pa.add_bitfield_record(&ip(6002), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 0, 0]));
330        pa.add_single_record(&ip(6003), 0);
331
332        pa.forget_piece(0);
333        assert!(pa.get_piece_owners(0).next().is_none());
334        assert_eq!(HashSet::new(), pa.get_peer_pieces(&ip(6003)).collect());
335        assert_eq!(HashSet::from([1]), pa.get_peer_pieces(&ip(6002)).collect());
336        assert_eq!(HashSet::from([1, 2]), pa.get_peer_pieces(&ip(6001)).collect());
337        assert_eq!(HashSet::from([1, 2, 3]), pa.get_peer_pieces(&ip(6000)).collect());
338        {
339            let mut rarest = pa.missing_pieces_rarest_first();
340            assert_eq!(3, rarest.next().unwrap());
341            assert_eq!(2, rarest.next().unwrap());
342            assert_eq!(1, rarest.next().unwrap());
343            assert!(rarest.next().is_none());
344
345            let mut poorest = pa.get_poorest_peers();
346            assert_eq!(&ip(6002), poorest.next().unwrap());
347            assert_eq!(&ip(6001), poorest.next().unwrap());
348            assert_eq!(&ip(6000), poorest.next().unwrap());
349            assert!(poorest.next().is_none());
350        }
351
352        pa.forget_piece(3);
353        assert!(pa.get_piece_owners(3).next().is_none());
354        assert_eq!(HashSet::new(), pa.get_peer_pieces(&ip(6003)).collect());
355        assert_eq!(HashSet::from([1]), pa.get_peer_pieces(&ip(6002)).collect());
356        assert_eq!(HashSet::from([1, 2]), pa.get_peer_pieces(&ip(6001)).collect());
357        assert_eq!(HashSet::from([1, 2]), pa.get_peer_pieces(&ip(6000)).collect());
358        {
359            let mut rarest = pa.missing_pieces_rarest_first();
360            assert_eq!(2, rarest.next().unwrap());
361            assert_eq!(1, rarest.next().unwrap());
362            assert!(rarest.next().is_none());
363
364            let mut poorest = pa.get_poorest_peers();
365            assert_eq!(&ip(6002), poorest.next().unwrap());
366            assert_eq!(HashSet::from([&ip(6001), &ip(6000)]), poorest.collect());
367        }
368
369        pa.forget_piece(1);
370        pa.forget_piece(2);
371        assert_eq!(HashSet::new(), pa.get_peer_pieces(&ip(6003)).collect());
372        assert_eq!(HashSet::new(), pa.get_peer_pieces(&ip(6002)).collect());
373        assert_eq!(HashSet::new(), pa.get_peer_pieces(&ip(6001)).collect());
374        assert_eq!(HashSet::new(), pa.get_peer_pieces(&ip(6000)).collect());
375        assert!(pa.missing_pieces_rarest_first().next().is_none());
376    }
377
378    #[test]
379    fn test_add_records_and_forget_peer() {
380        let mut pa = PieceTracker::new(4);
381        pa.add_bitfield_record(&ip(6000), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 1, 1]));
382        pa.add_bitfield_record(&ip(6001), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 1, 0]));
383        pa.add_bitfield_record(&ip(6002), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 0, 0]));
384        pa.add_single_record(&ip(6003), 0);
385
386        pa.forget_peer(&ip(6000));
387        assert!(pa.get_peer_pieces(&ip(6000)).next().is_none());
388        assert_eq!(HashSet::new(), pa.get_piece_owners(3).collect());
389        assert_eq!(HashSet::from([&ip(6001)]), pa.get_piece_owners(2).collect());
390        assert_eq!(HashSet::from([&ip(6001), &ip(6002)]), pa.get_piece_owners(1).collect());
391        assert_eq!(
392            HashSet::from([&ip(6001), &ip(6002), &ip(6003)]),
393            pa.get_piece_owners(0).collect()
394        );
395        {
396            let mut rarest = pa.missing_pieces_rarest_first();
397            assert_eq!(2, rarest.next().unwrap());
398            assert_eq!(1, rarest.next().unwrap());
399            assert_eq!(0, rarest.next().unwrap());
400            assert!(rarest.next().is_none());
401
402            let mut poorest = pa.get_poorest_peers();
403            assert_eq!(&ip(6003), poorest.next().unwrap());
404            assert_eq!(&ip(6002), poorest.next().unwrap());
405            assert_eq!(&ip(6001), poorest.next().unwrap());
406            assert!(poorest.next().is_none());
407        }
408
409        pa.forget_peer(&ip(6003));
410        assert!(pa.get_peer_pieces(&ip(6003)).next().is_none());
411        assert_eq!(HashSet::new(), pa.get_piece_owners(3).collect());
412        assert_eq!(HashSet::from([&ip(6001)]), pa.get_piece_owners(2).collect());
413        assert_eq!(HashSet::from([&ip(6001), &ip(6002)]), pa.get_piece_owners(1).collect());
414        assert_eq!(HashSet::from([&ip(6001), &ip(6002)]), pa.get_piece_owners(0).collect());
415        {
416            let mut rarest = pa.missing_pieces_rarest_first();
417            assert_eq!(2, rarest.next().unwrap());
418            assert_eq!(HashSet::from([1, 0]), rarest.collect());
419
420            let mut poorest = pa.get_poorest_peers();
421            assert_eq!(&ip(6002), poorest.next().unwrap());
422            assert_eq!(&ip(6001), poorest.next().unwrap());
423            assert!(poorest.next().is_none());
424        }
425    }
426
427    #[test]
428    fn test_dont_leak_empty_owner_count_entries() {
429        let mut pa = PieceTracker::new(4);
430        assert_eq!(1, pa.owner_count_to_piece_indices.len());
431
432        pa.add_single_record(&ip(6000), 0);
433        let mut keys = pa.owner_count_to_piece_indices.keys().cloned();
434        assert_eq!(0, keys.next().unwrap());
435        assert_eq!(1, keys.next().unwrap());
436        assert!(keys.next().is_none());
437
438        pa.add_single_record(&ip(6001), 0);
439        let mut keys = pa.owner_count_to_piece_indices.keys().cloned();
440        assert_eq!(0, keys.next().unwrap());
441        assert_eq!(2, keys.next().unwrap());
442        assert!(keys.next().is_none());
443
444        pa.forget_piece(0);
445        let mut keys = pa.owner_count_to_piece_indices.keys().cloned();
446        assert_eq!(0, keys.next().unwrap());
447        assert!(keys.next().is_none());
448    }
449
450    #[test]
451    fn test_dont_leak_empty_piece_count_entries() {
452        let mut pa = PieceTracker::new(4);
453        assert_eq!(0, pa.piece_count_to_owners.len());
454
455        pa.add_single_record(&ip(6000), 0);
456        let mut keys = pa.piece_count_to_owners.keys().cloned();
457        assert_eq!(1, keys.next().unwrap());
458        assert!(keys.next().is_none());
459
460        pa.add_single_record(&ip(6000), 1);
461        let mut keys = pa.piece_count_to_owners.keys().cloned();
462        assert_eq!(2, keys.next().unwrap());
463        assert!(keys.next().is_none());
464
465        pa.forget_peer(&ip(6000));
466        assert!(pa.piece_count_to_owners.is_empty());
467    }
468
469    #[test]
470    fn test_process_entire_bitfield_ignoring_forgotten_pieces() {
471        let mut pa = PieceTracker::new(4);
472        pa.forget_piece(0);
473
474        pa.add_bitfield_record(&ip(6000), &BitVec::from_bitslice(bits![u8, Msb0; 1, 1, 1, 1]));
475        assert!(pa.get_piece_owners(0).next().is_none());
476        assert_eq!(HashSet::from([&ip(6000)]), pa.get_piece_owners(1).collect());
477        assert_eq!(HashSet::from([&ip(6000)]), pa.get_piece_owners(2).collect());
478        assert_eq!(HashSet::from([&ip(6000)]), pa.get_piece_owners(3).collect());
479        assert_eq!(HashSet::from([1, 2, 3]), pa.missing_pieces_rarest_first().collect());
480        assert_eq!(HashSet::from([1, 2, 3]), pa.get_peer_pieces(&ip(6000)).collect());
481    }
482}