Skip to main content

mtorrent_core/data/
block_accountant.rs

1use crate::data::{Error, PieceInfo};
2use crate::pwp::{Bitfield, BlockInfo};
3use std::collections::BTreeMap;
4use std::rc::Rc;
5
6/// Keeps track of downloaded data on per-block basis.
7#[derive(Debug)]
8pub struct BlockAccountant {
9    pieces: Rc<PieceInfo>,
10    blocks_start_end: BTreeMap<usize, usize>,
11    total_bytes: usize,
12}
13
14impl BlockAccountant {
15    pub fn new(pieces: Rc<PieceInfo>) -> Self {
16        BlockAccountant {
17            pieces,
18            blocks_start_end: BTreeMap::new(),
19            total_bytes: 0,
20        }
21    }
22
23    /// Try to add a received block to the internal records. Fails if the block has invalid
24    /// offset, index or length.
25    pub fn submit_block(&mut self, block_info: &BlockInfo) -> Result<usize, Error> {
26        let result = self.pieces.global_offset(
27            block_info.piece_index,
28            block_info.in_piece_offset,
29            block_info.block_length,
30        );
31        if let Ok(global_offset) = result {
32            self.submit_block_internal(global_offset, block_info.block_length);
33        }
34        result
35    }
36
37    fn submit_block_internal(&mut self, global_offset: usize, length: usize) {
38        let start = global_offset;
39        let mut end = global_offset + length;
40
41        while let Some(next_block) = self.blocks_start_end.range_mut(global_offset..).next() {
42            let (next_start, next_end) = { (*next_block.0, *next_block.1) };
43            if next_start > end {
44                break;
45            }
46            if next_end > end {
47                end = next_end;
48            }
49            self.blocks_start_end.remove(&next_start);
50            self.total_bytes -= next_end - next_start;
51        }
52
53        if let Some(prev_block) = self.blocks_start_end.range_mut(..global_offset).last() {
54            let (_prev_start, prev_end) = prev_block;
55            if *prev_end >= start {
56                if end > *prev_end {
57                    self.total_bytes += end - *prev_end;
58                    *prev_end = end;
59                }
60                return;
61            }
62        }
63
64        self.blocks_start_end.insert(start, end);
65        self.total_bytes += end - start;
66    }
67
68    /// Mark a piece as downloaded. Fails if the piece index is invalid.
69    pub fn submit_piece(&mut self, piece_index: usize) -> bool {
70        let piece_length = self.pieces.piece_len(piece_index);
71        if let Ok(offset) = self.pieces.global_offset(piece_index, 0, piece_length) {
72            self.submit_block_internal(offset, piece_length);
73            true
74        } else {
75            false
76        }
77    }
78
79    /// Update internal records from a bitfield. All pieces present in the bitfield
80    /// will be marked as downloaded. Fails if the bitfield has unexpected length.
81    pub fn submit_bitfield(&mut self, bitfield: &Bitfield) -> bool {
82        if bitfield.len() < self.pieces.piece_count() {
83            return false;
84        }
85        for (piece_index, is_piece_present) in bitfield.iter().enumerate() {
86            if *is_piece_present {
87                self.submit_piece(piece_index);
88            }
89        }
90        true
91    }
92
93    /// Remove a piece from the internal records, i.e. no longer consider it as downloaded.
94    pub fn remove_piece(&mut self, piece_index: usize) {
95        let piece_length = self.pieces.piece_len(piece_index);
96        if let Ok(global_offset) = self.pieces.global_offset(piece_index, 0, piece_length) {
97            self.remove_block_internal(global_offset, piece_length);
98        }
99    }
100
101    fn remove_block_internal(&mut self, global_offset: usize, length: usize) {
102        let start = global_offset;
103        let end = global_offset + length;
104
105        if let Some(prev_block) = self.blocks_start_end.range_mut(..global_offset).last() {
106            let (_prev_start, prev_end) = prev_block;
107            let prev_end_copy = *prev_end;
108            if *prev_end > start {
109                self.total_bytes -= *prev_end - start;
110                *prev_end = start;
111            }
112            if prev_end_copy > end {
113                self.blocks_start_end.insert(end, prev_end_copy);
114                self.total_bytes += prev_end_copy - end;
115            }
116        }
117
118        while let Some(next_block) = self.blocks_start_end.range_mut(global_offset..).next() {
119            let (next_start, next_end) = { (*next_block.0, *next_block.1) };
120            if next_start >= end {
121                break;
122            }
123            self.blocks_start_end.remove(&next_start);
124            self.total_bytes -= next_end - next_start;
125            if next_end > end {
126                self.blocks_start_end.insert(end, next_end);
127                self.total_bytes += next_end - end;
128            }
129        }
130    }
131
132    fn max_block_length_at(&self, global_offset: usize) -> Option<usize> {
133        if let Some((_start, end)) = self.blocks_start_end.range(..=global_offset).last() {
134            if *end > global_offset {
135                Some(*end - global_offset)
136            } else {
137                None
138            }
139        } else {
140            None
141        }
142    }
143
144    /// Check whether `length` bytes at `global_offset` have been downloaded.
145    pub fn has_exact_block_at(&self, global_offset: usize, length: usize) -> bool {
146        if let Some(block_length) = self.max_block_length_at(global_offset) {
147            block_length >= length
148        } else {
149            false
150        }
151    }
152
153    /// Check presense of an exact block among the downloaded data.
154    pub fn has_exact_block(&self, block_info: &BlockInfo) -> bool {
155        if let Ok(global_offset) = self.pieces.global_offset(
156            block_info.piece_index,
157            block_info.in_piece_offset,
158            block_info.block_length,
159        ) {
160            self.has_exact_block_at(global_offset, block_info.block_length)
161        } else {
162            false
163        }
164    }
165
166    /// Check whether the piece at `piece_index` has been downloaded.
167    pub fn has_piece(&self, piece_index: usize) -> bool {
168        let piece_len = self.pieces.piece_len(piece_index);
169        if let Ok(global_offset) = self.pieces.global_offset(piece_index, 0, piece_len) {
170            self.has_exact_block_at(global_offset, piece_len)
171        } else {
172            false
173        }
174    }
175
176    /// Represent the internal state as a bitfield. Partially downloaded pieces won't be included.
177    pub fn generate_bitfield(&self) -> Bitfield {
178        let mut bitfield = Bitfield::repeat(false, self.pieces.piece_count());
179        for (piece_index, mut is_piece_present) in bitfield.iter_mut().enumerate() {
180            if self.has_piece(piece_index) {
181                is_piece_present.set(true);
182            }
183        }
184        bitfield
185    }
186
187    /// The total number of downloaded bytes.
188    pub fn accounted_bytes(&self) -> usize {
189        self.total_bytes
190    }
191
192    /// The total number of missing bytes.
193    pub fn missing_bytes(&self) -> usize {
194        self.pieces.total_len() - self.total_bytes
195    }
196}
197
198#[cfg(test)]
199mod tests {
200    use super::*;
201    use std::iter;
202
203    fn piece_info() -> Rc<PieceInfo> {
204        Rc::new(PieceInfo::new(iter::repeat_n([0u8; 20], 86), 3, 256).unwrap())
205    }
206
207    #[test]
208    fn test_accountant_submit_one_block() {
209        let p = piece_info();
210        let mut a = BlockAccountant::new(p);
211        a.submit_block_internal(10, 10);
212
213        assert_eq!(1, a.blocks_start_end.len());
214        assert_eq!(Some(&20), a.blocks_start_end.get(&10));
215        assert_eq!(10, a.accounted_bytes());
216    }
217
218    #[test]
219    fn test_accountant_merge_into_preceding_block() {
220        let p = piece_info();
221        let mut a = BlockAccountant::new(p);
222        a.submit_block_internal(10, 10);
223        a.submit_block_internal(20, 10);
224
225        assert_eq!(1, a.blocks_start_end.len());
226        assert_eq!(Some(&30), a.blocks_start_end.get(&10));
227        assert_eq!(20, a.accounted_bytes());
228    }
229
230    #[test]
231    fn test_accountant_merge_overlapping_into_preceding_block() {
232        let p = piece_info();
233        let mut a = BlockAccountant::new(p);
234        a.submit_block_internal(10, 10);
235        a.submit_block_internal(15, 15);
236
237        assert_eq!(1, a.blocks_start_end.len());
238        assert_eq!(Some(&30), a.blocks_start_end.get(&10));
239        assert_eq!(20, a.accounted_bytes());
240    }
241
242    #[test]
243    fn test_accountant_merge_into_following_block() {
244        let p = piece_info();
245        let mut a = BlockAccountant::new(p);
246        a.submit_block_internal(10, 10);
247        a.submit_block_internal(0, 10);
248
249        assert_eq!(1, a.blocks_start_end.len());
250        assert_eq!(Some(&20), a.blocks_start_end.get(&0));
251        assert_eq!(20, a.accounted_bytes());
252    }
253
254    #[test]
255    fn test_accountant_merge_overlapping_into_following_block() {
256        let p = piece_info();
257        let mut a = BlockAccountant::new(p);
258        a.submit_block_internal(10, 10);
259        a.submit_block_internal(0, 15);
260
261        assert_eq!(1, a.blocks_start_end.len());
262        assert_eq!(Some(&20), a.blocks_start_end.get(&0));
263        assert_eq!(20, a.accounted_bytes());
264    }
265
266    #[test]
267    fn test_accountant_replace_overlapping_block() {
268        let p = piece_info();
269        let mut a = BlockAccountant::new(p);
270        a.submit_block_internal(10, 10);
271        a.submit_block_internal(5, 20);
272
273        assert_eq!(1, a.blocks_start_end.len());
274        assert_eq!(Some(&25), a.blocks_start_end.get(&5));
275        assert_eq!(20, a.accounted_bytes());
276    }
277
278    #[test]
279    fn test_accountant_ignore_overlapping_block() {
280        let p = piece_info();
281        let mut a = BlockAccountant::new(p);
282        a.submit_block_internal(5, 20);
283        a.submit_block_internal(10, 10);
284
285        assert_eq!(1, a.blocks_start_end.len());
286        assert_eq!(Some(&25), a.blocks_start_end.get(&5));
287        assert_eq!(20, a.accounted_bytes());
288    }
289
290    #[test]
291    fn test_accountant_merge_with_following_and_preceding_blocks() {
292        let p = piece_info();
293        let mut a = BlockAccountant::new(p);
294        a.submit_block_internal(10, 5);
295        a.submit_block_internal(0, 5);
296
297        assert_eq!(2, a.blocks_start_end.len());
298        assert_eq!(Some(&5), a.blocks_start_end.get(&0));
299        assert_eq!(Some(&15), a.blocks_start_end.get(&10));
300        assert_eq!(10, a.accounted_bytes());
301
302        a.submit_block_internal(5, 5);
303
304        assert_eq!(1, a.blocks_start_end.len());
305        assert_eq!(Some(&15), a.blocks_start_end.get(&0));
306        assert_eq!(15, a.accounted_bytes());
307    }
308
309    #[test]
310    fn test_accountant_merge_with_overlapping_following_and_preceding_blocks() {
311        let p = piece_info();
312        let mut a = BlockAccountant::new(p);
313        a.submit_block_internal(10, 5);
314        a.submit_block_internal(0, 5);
315
316        a.submit_block_internal(2, 10);
317
318        assert_eq!(1, a.blocks_start_end.len());
319        assert_eq!(Some(&15), a.blocks_start_end.get(&0));
320        assert_eq!(15, a.accounted_bytes());
321    }
322
323    #[test]
324    fn test_accountant_block_length_with_one_block() {
325        let p = piece_info();
326        let mut a = BlockAccountant::new(p);
327        a.submit_block_internal(10, 10);
328
329        assert_eq!(None, a.max_block_length_at(9));
330        assert_eq!(Some(10), a.max_block_length_at(10));
331        assert_eq!(Some(9), a.max_block_length_at(11));
332        assert_eq!(Some(1), a.max_block_length_at(19));
333        assert_eq!(None, a.max_block_length_at(20));
334    }
335
336    #[test]
337    fn test_accountant_block_length_with_two_blocks() {
338        let p = piece_info();
339        let mut a = BlockAccountant::new(p);
340        a.submit_block_internal(10, 10);
341        a.submit_block_internal(30, 10);
342
343        assert_eq!(Some(1), a.max_block_length_at(19));
344        for pos in 20..30 {
345            assert_eq!(None, a.max_block_length_at(pos), "pos={pos}");
346        }
347        assert_eq!(Some(10), a.max_block_length_at(30));
348        assert_eq!(Some(9), a.max_block_length_at(31));
349        assert_eq!(Some(1), a.max_block_length_at(39));
350        assert_eq!(None, a.max_block_length_at(40));
351    }
352
353    #[test]
354    fn test_accountant_has_exact_block_with_one_block() {
355        let p = piece_info();
356        let mut a = BlockAccountant::new(p);
357        a.submit_block_internal(10, 10);
358
359        for len in 0..=10 {
360            assert!(!a.has_exact_block_at(9, len), "len={len}");
361            assert!(a.has_exact_block_at(10, len), "len={len}");
362        }
363        assert!(a.has_exact_block_at(11, 9));
364        assert!(!a.has_exact_block_at(11, 10));
365
366        assert!(a.has_exact_block_at(19, 1));
367        assert!(!a.has_exact_block_at(19, 2));
368    }
369
370    #[test]
371    fn test_accountant_remove_exact_block() {
372        let p = piece_info();
373        let mut a = BlockAccountant::new(p);
374
375        // given
376        a.blocks_start_end.insert(0, 5);
377        a.blocks_start_end.insert(10, 15);
378        a.blocks_start_end.insert(20, 25);
379        a.total_bytes = 15;
380
381        // when
382        a.remove_block_internal(10, 5);
383
384        // then
385        assert_eq!(2, a.blocks_start_end.len());
386        assert_eq!(Some(&5), a.blocks_start_end.get(&0));
387        assert_eq!(Some(&25), a.blocks_start_end.get(&20));
388        assert_eq!(10, a.total_bytes);
389    }
390
391    #[test]
392    fn test_accountant_shrink_block_from_tail_end() {
393        let p = piece_info();
394        let mut a = BlockAccountant::new(p);
395
396        // given
397        a.blocks_start_end.insert(0, 10);
398        a.total_bytes = 10;
399
400        // when
401        a.remove_block_internal(5, 5);
402
403        // then
404        assert_eq!(1, a.blocks_start_end.len());
405        assert_eq!(Some(&5), a.blocks_start_end.get(&0));
406        assert_eq!(5, a.total_bytes);
407    }
408
409    #[test]
410    fn test_accountant_shrink_block_from_head_end() {
411        let p = piece_info();
412        let mut a = BlockAccountant::new(p);
413
414        // given
415        a.blocks_start_end.insert(0, 10);
416        a.total_bytes = 10;
417
418        // when
419        a.remove_block_internal(0, 5);
420
421        // then
422        assert_eq!(1, a.blocks_start_end.len());
423        assert_eq!(Some(&10), a.blocks_start_end.get(&5));
424        assert_eq!(5, a.total_bytes);
425    }
426
427    #[test]
428    fn test_accountant_split_block_into_two() {
429        let p = piece_info();
430        let mut a = BlockAccountant::new(p);
431
432        // given
433        a.blocks_start_end.insert(0, 20);
434        a.total_bytes = 20;
435
436        // when
437        a.remove_block_internal(5, 10);
438
439        // then
440        assert_eq!(2, a.blocks_start_end.len());
441        assert_eq!(Some(&5), a.blocks_start_end.get(&0));
442        assert_eq!(Some(&20), a.blocks_start_end.get(&15));
443        assert_eq!(10, a.total_bytes);
444    }
445
446    #[test]
447    fn test_accountant_remove_multiple_nonadjacent_blocks() {
448        let p = piece_info();
449        let mut a = BlockAccountant::new(p);
450
451        // given
452        a.blocks_start_end.insert(0, 5);
453        a.blocks_start_end.insert(10, 15);
454        a.blocks_start_end.insert(20, 25);
455        a.blocks_start_end.insert(30, 35);
456        a.total_bytes = 20;
457
458        // when
459        a.remove_block_internal(8, 20);
460
461        // then
462        assert_eq!(2, a.blocks_start_end.len());
463        assert_eq!(Some(&5), a.blocks_start_end.get(&0));
464        assert_eq!(Some(&35), a.blocks_start_end.get(&30));
465        assert_eq!(10, a.total_bytes);
466    }
467
468    #[test]
469    fn test_accountant_remove_multiple_nonadjacent_blocks_and_shrink() {
470        let p = piece_info();
471        let mut a = BlockAccountant::new(p);
472
473        // given
474        a.blocks_start_end.insert(0, 5);
475        a.blocks_start_end.insert(10, 15);
476        a.blocks_start_end.insert(20, 25);
477        a.blocks_start_end.insert(30, 35);
478        a.total_bytes = 20;
479
480        // when
481        a.remove_block_internal(4, 27);
482
483        // then
484        assert_eq!(2, a.blocks_start_end.len());
485        assert_eq!(Some(&4), a.blocks_start_end.get(&0));
486        assert_eq!(Some(&35), a.blocks_start_end.get(&31));
487        assert_eq!(8, a.total_bytes);
488    }
489}