1use crate::data::{Error, PieceInfo};
2use crate::pwp::{Bitfield, BlockInfo};
3use std::collections::BTreeMap;
4use std::rc::Rc;
5
6#[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 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 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 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 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 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 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 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 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 pub fn accounted_bytes(&self) -> usize {
189 self.total_bytes
190 }
191
192 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 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 a.remove_block_internal(10, 5);
383
384 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 a.blocks_start_end.insert(0, 10);
398 a.total_bytes = 10;
399
400 a.remove_block_internal(5, 5);
402
403 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 a.blocks_start_end.insert(0, 10);
416 a.total_bytes = 10;
417
418 a.remove_block_internal(0, 5);
420
421 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 a.blocks_start_end.insert(0, 20);
434 a.total_bytes = 20;
435
436 a.remove_block_internal(5, 10);
438
439 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 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 a.remove_block_internal(8, 20);
460
461 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 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 a.remove_block_internal(4, 27);
482
483 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}