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#[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 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 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 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 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 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 false
107 }
108 }
109
110 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 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 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(¤t_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(¤t_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}