1#![cfg_attr(coverage_nightly, feature(coverage_attribute))]
2
3use nonzero_ext::nonzero;
4use pin_project::pin_project;
5use std::io;
6use std::io::Error;
7use std::marker::PhantomData;
8use std::num::NonZeroU64;
9use std::pin::Pin;
10use std::sync::{Arc, Mutex};
11use std::task::{Context, Poll};
12use std::time::SystemTime;
13use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
14
15pub trait SpeedStat: 'static + Send + Sync {
16 fn get_write_speed(&self) -> u64;
17 fn get_write_sum_size(&self) -> u64;
18 fn get_read_speed(&self) -> u64;
19 fn get_read_sum_size(&self) -> u64;
20}
21
22pub trait SpeedTracker: SpeedStat {
23 fn add_write_data_size(&self, size: u64);
24 fn add_read_data_size(&self, size: u64);
25}
26
27pub trait TimePicker: 'static + Sync + Send {
28 fn now() -> u128;
29}
30
31pub struct SystemTimePicker;
32
33impl TimePicker for SystemTimePicker {
34 fn now() -> u128 {
35 SystemTime::now()
36 .duration_since(SystemTime::UNIX_EPOCH)
37 .unwrap()
38 .as_millis()
39 }
40}
41
42struct DataItem {
43 size: u64,
44 time: u64,
45}
46
47pub(crate) struct SpeedState<T: TimePicker> {
48 sum_size: u64,
49 last_time: u128,
50 speed_duration: NonZeroU64,
51 data_items: Vec<DataItem>,
52 _time_picker: PhantomData<T>,
53}
54
55impl<T: TimePicker> SpeedState<T> {
56 fn new(speed_duration: NonZeroU64) -> SpeedState<T> {
57 SpeedState {
58 sum_size: 0,
59 last_time: T::now(),
60 speed_duration,
61 data_items: vec![],
62 _time_picker: Default::default(),
63 }
64 }
65
66 pub fn add_data(&mut self, size: u64) {
67 self.sum_size += size;
68 let now = T::now().max(self.last_time);
69 self.clear_invalid_item(now);
70
71 if now / 1000 == self.last_time / 1000 {
72 if self.data_items.len() == 0 {
73 self.data_items.push(DataItem {
74 size,
75 time: (now / 1000) as u64,
76 });
77 } else {
78 let last_item = self.data_items.last_mut().unwrap();
79 if last_item.time == (now / 1000) as u64 {
80 last_item.size += size;
81 } else {
82 self.data_items.push(DataItem {
83 size,
84 time: (now / 1000) as u64,
85 });
86 }
87 }
88 } else {
89 let duration = now - self.last_time;
90 let mut pos = 0;
91 let mut offset = 1000 - self.last_time % 1000;
92 let mut sec = (self.last_time / 1000) as u64;
93 while pos < duration {
94 let mut weight = offset;
95 if pos + offset > duration {
96 weight = duration - pos;
97 }
98 let data_size = (size as u128 * weight / duration) as u64;
99 if self.data_items.len() == 0 {
100 self.data_items.push(DataItem {
101 size: data_size,
102 time: sec,
103 });
104 } else {
105 let last_item = self.data_items.last_mut().unwrap();
106 if last_item.time == sec {
107 last_item.size += data_size;
108 } else {
109 self.data_items.push(DataItem {
110 size: data_size,
111 time: sec,
112 })
113 }
114 }
115 pos += offset;
116 offset = 1000;
117 sec += 1;
118 }
119 }
120 self.last_time = now;
121 }
122
123 pub fn clear_invalid_item(&mut self, now: u128) {
124 let now = (now / 1000) as u64;
125 self.data_items
126 .retain(|item| now.saturating_sub(item.time) <= self.speed_duration.get());
127 }
128
129 pub fn get_speed(&self) -> u64 {
130 let now = (T::now().max(self.last_time) / 1000) as u64;
131 let mut sum_size = 0;
132 for item in self.data_items.iter() {
133 if item.time < now && now.saturating_sub(item.time) <= self.speed_duration.get() {
134 sum_size += item.size;
135 }
136 }
137
138 sum_size / self.speed_duration
139 }
140
141 pub fn get_sum_size(&self) -> u64 {
142 self.sum_size
143 }
144}
145
146pub struct SfoSpeedStat<T: TimePicker = SystemTimePicker> {
147 upload_state: Mutex<SpeedState<T>>,
148 download_state: Mutex<SpeedState<T>>,
149}
150
151impl SfoSpeedStat {
152 pub fn new() -> SfoSpeedStat {
153 Self {
154 upload_state: Mutex::new(SpeedState::new(nonzero!(5u64))),
155 download_state: Mutex::new(SpeedState::new(nonzero!(5u64))),
156 }
157 }
158
159 pub fn new_with_duration(duration: u64) -> SfoSpeedStat {
167 SfoSpeedStat {
168 upload_state: Mutex::new(SpeedState::new(NonZeroU64::new(duration).unwrap())),
169 download_state: Mutex::new(SpeedState::new(NonZeroU64::new(duration).unwrap())),
170 }
171 }
172}
173
174impl<T: TimePicker> SfoSpeedStat<T> {
175 pub(crate) fn new_with_time_picker() -> SfoSpeedStat<T> {
176 SfoSpeedStat {
177 upload_state: Mutex::new(SpeedState::new(nonzero!(5u64))),
178 download_state: Mutex::new(SpeedState::new(nonzero!(5u64))),
179 }
180 }
181}
182
183impl<T: TimePicker> SpeedTracker for SfoSpeedStat<T> {
184 fn add_write_data_size(&self, size: u64) {
185 self.upload_state.lock().unwrap().add_data(size);
186 }
187
188 fn add_read_data_size(&self, size: u64) {
189 self.download_state.lock().unwrap().add_data(size);
190 }
191}
192
193impl<T: TimePicker> SpeedStat for SfoSpeedStat<T> {
194 fn get_write_speed(&self) -> u64 {
195 self.upload_state.lock().unwrap().get_speed()
196 }
197
198 fn get_write_sum_size(&self) -> u64 {
199 self.upload_state.lock().unwrap().get_sum_size()
200 }
201
202 fn get_read_speed(&self) -> u64 {
203 self.download_state.lock().unwrap().get_speed()
204 }
205
206 fn get_read_sum_size(&self) -> u64 {
207 self.download_state.lock().unwrap().get_sum_size()
208 }
209}
210
211#[pin_project]
212pub struct StatStream<T: AsyncRead + AsyncWrite + Send + 'static> {
213 #[pin]
214 stream: T,
215 stat: Arc<dyn SpeedTracker>,
216}
217
218impl<T: AsyncRead + AsyncWrite + Send + 'static> StatStream<T> {
219 pub fn new(stream: T) -> StatStream<T> {
220 StatStream {
221 stream,
222 stat: Arc::new(SfoSpeedStat::new()),
223 }
224 }
225
226 pub fn new_with_tracker(stream: T, tracker: Arc<dyn SpeedTracker>) -> StatStream<T> {
227 StatStream {
228 stream,
229 stat: tracker,
230 }
231 }
232}
233
234impl<T: AsyncRead + AsyncWrite + Send + 'static> StatStream<T> {
235 pub(crate) fn new_test<S: TimePicker>(stream: T) -> StatStream<T> {
236 StatStream {
237 stream,
238 stat: Arc::new(SfoSpeedStat::<S>::new_with_time_picker()),
239 }
240 }
241
242 pub fn get_speed_stat(&self) -> Arc<dyn SpeedStat> {
243 self.stat.clone()
244 }
245
246 pub fn raw_stream(&mut self) -> &mut T {
247 &mut self.stream
248 }
249}
250
251impl<T: AsyncRead + AsyncWrite + Unpin + Send + 'static> AsyncRead for StatStream<T> {
252 fn poll_read(
253 self: Pin<&mut Self>,
254 cx: &mut Context<'_>,
255 buf: &mut ReadBuf<'_>,
256 ) -> Poll<io::Result<()>> {
257 let this = self.project();
258 let filled_len = buf.filled().len();
259 match this.stream.poll_read(cx, buf) {
260 Poll::Ready(res) => {
261 if res.is_ok() {
262 this.stat
263 .add_read_data_size((buf.filled().len() - filled_len) as u64);
264 }
265 Poll::Ready(res)
266 }
267 Poll::Pending => Poll::Pending,
268 }
269 }
270}
271
272impl<T: AsyncRead + AsyncWrite + Unpin + Send + 'static> AsyncWrite for StatStream<T> {
273 fn poll_write(
274 self: Pin<&mut Self>,
275 cx: &mut Context<'_>,
276 buf: &[u8],
277 ) -> Poll<Result<usize, io::Error>> {
278 let this = self.project();
279 match this.stream.poll_write(cx, buf) {
280 Poll::Ready(Ok(size)) => {
281 this.stat.add_write_data_size(size as u64);
282 Poll::Ready(Ok(size))
283 }
284 Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
285 Poll::Pending => Poll::Pending,
286 }
287 }
288
289 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
290 self.project().stream.poll_flush(cx)
291 }
292
293 fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
294 self.project().stream.poll_shutdown(cx)
295 }
296}
297
298#[pin_project]
299pub struct StatRead<T: AsyncRead + Send + 'static> {
300 #[pin]
301 reader: T,
302 stat: Arc<dyn SpeedTracker>,
303}
304
305impl<T: AsyncRead + Send + 'static> StatRead<T> {
306 pub fn new(reader: T) -> StatRead<T> {
307 StatRead {
308 reader,
309 stat: Arc::new(SfoSpeedStat::new()),
310 }
311 }
312
313 pub fn new_with_tracker(reader: T, tracker: Arc<dyn SpeedTracker>) -> StatRead<T> {
314 StatRead {
315 reader,
316 stat: tracker,
317 }
318 }
319}
320
321impl<T: AsyncRead + Send + 'static> StatRead<T> {
322 pub(crate) fn new_test<S: TimePicker>(reader: T) -> StatRead<T> {
323 StatRead {
324 reader,
325 stat: Arc::new(SfoSpeedStat::<S>::new_with_time_picker()),
326 }
327 }
328
329 pub fn get_speed_stat(&self) -> Arc<dyn SpeedStat> {
330 self.stat.clone()
331 }
332
333 pub fn raw_reader(&mut self) -> &mut T {
334 &mut self.reader
335 }
336}
337
338impl<T: AsyncRead + Unpin + Send + 'static> AsyncRead for StatRead<T> {
339 fn poll_read(
340 self: Pin<&mut Self>,
341 cx: &mut Context<'_>,
342 buf: &mut ReadBuf<'_>,
343 ) -> Poll<io::Result<()>> {
344 let this = self.project();
345 let filled_len = buf.filled().len();
346 match this.reader.poll_read(cx, buf) {
347 Poll::Ready(res) => {
348 if res.is_ok() {
349 this.stat
350 .add_read_data_size((buf.filled().len() - filled_len) as u64);
351 }
352 Poll::Ready(res)
353 }
354 Poll::Pending => Poll::Pending,
355 }
356 }
357}
358
359#[pin_project]
360pub struct StatWrite<T: AsyncWrite + Send + 'static> {
361 #[pin]
362 writer: T,
363 stat: Arc<dyn SpeedTracker>,
364}
365
366impl<T: AsyncWrite + Send + 'static> StatWrite<T> {
367 pub fn new(writer: T) -> StatWrite<T> {
368 StatWrite {
369 writer,
370 stat: Arc::new(SfoSpeedStat::new()),
371 }
372 }
373
374 pub fn new_with_tracker(writer: T, tracker: Arc<dyn SpeedTracker>) -> StatWrite<T> {
375 StatWrite {
376 writer,
377 stat: tracker,
378 }
379 }
380}
381
382impl<T: AsyncWrite + Send + 'static> StatWrite<T> {
383 pub(crate) fn new_test<S: TimePicker>(writer: T) -> StatWrite<T> {
384 StatWrite {
385 writer,
386 stat: Arc::new(SfoSpeedStat::<S>::new_with_time_picker()),
387 }
388 }
389
390 pub fn get_speed_stat(&self) -> Arc<dyn SpeedStat> {
391 self.stat.clone()
392 }
393
394 pub fn raw_writer(&mut self) -> &mut T {
395 &mut self.writer
396 }
397}
398
399impl<T: AsyncWrite + Unpin + Send + 'static> AsyncWrite for StatWrite<T> {
400 fn poll_write(
401 self: Pin<&mut Self>,
402 cx: &mut Context<'_>,
403 buf: &[u8],
404 ) -> Poll<Result<usize, io::Error>> {
405 let this = self.project();
406 match this.writer.poll_write(cx, buf) {
407 Poll::Ready(Ok(size)) => {
408 this.stat.add_write_data_size(size as u64);
409 Poll::Ready(Ok(size))
410 }
411 Poll::Ready(Err(err)) => Poll::Ready(Err(err)),
412 Poll::Pending => Poll::Pending,
413 }
414 }
415
416 fn poll_flush(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), io::Error>> {
417 self.project().writer.poll_flush(cx)
418 }
419
420 fn poll_shutdown(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Error>> {
421 self.project().writer.poll_shutdown(cx)
422 }
423}
424#[cfg(test)]
425#[cfg_attr(coverage_nightly, coverage(off))]
426mod tests {
427 use super::*;
428 use std::sync::atomic::{AtomicU64, Ordering};
429 use std::time::Duration;
430 use tokio::io::{AsyncReadExt, AsyncWriteExt};
431
432 struct PartialStream {
433 max_write_size: usize,
434 read_size: usize,
435 }
436
437 impl AsyncRead for PartialStream {
438 fn poll_read(
439 self: Pin<&mut Self>,
440 _cx: &mut Context<'_>,
441 buf: &mut ReadBuf<'_>,
442 ) -> Poll<io::Result<()>> {
443 let size = self.read_size.min(buf.remaining());
444 buf.put_slice(&[0; 3][..size]);
445 Poll::Ready(Ok(()))
446 }
447 }
448
449 impl AsyncWrite for PartialStream {
450 fn poll_write(
451 self: Pin<&mut Self>,
452 _cx: &mut Context<'_>,
453 buf: &[u8],
454 ) -> Poll<Result<usize, io::Error>> {
455 Poll::Ready(Ok(self.max_write_size.min(buf.len())))
456 }
457
458 fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
459 Poll::Ready(Ok(()))
460 }
461
462 fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<io::Result<()>> {
463 Poll::Ready(Ok(()))
464 }
465 }
466
467 #[tokio::test]
468 async fn test_stat_stream_partial_write_counts_returned_size() {
469 let tracker = Arc::new(SfoSpeedStat::new());
470 let stream = PartialStream {
471 max_write_size: 3,
472 read_size: 0,
473 };
474 let mut stream = StatStream::new_with_tracker(stream, tracker.clone());
475
476 let written = stream.write(&[0; 10]).await.unwrap();
477
478 assert_eq!(written, 3);
479 assert_eq!(tracker.get_write_sum_size(), 3);
480 }
481
482 #[tokio::test]
483 async fn test_stat_stream_partial_read_counts_newly_filled_size() {
484 let tracker = Arc::new(SfoSpeedStat::new());
485 let stream = PartialStream {
486 max_write_size: 0,
487 read_size: 3,
488 };
489 let mut stream = StatStream::new_with_tracker(stream, tracker.clone());
490 let mut storage = [0; 10];
491 let mut buf = ReadBuf::new(&mut storage);
492 buf.put_slice(&[0; 4]);
493
494 let result =
495 futures::future::poll_fn(|cx| Pin::new(&mut stream).poll_read(cx, &mut buf)).await;
496
497 assert!(result.is_ok());
498 assert_eq!(buf.filled().len(), 7);
499 assert_eq!(tracker.get_read_sum_size(), 3);
500 }
501
502 #[tokio::test]
503 async fn test_stat_stream_counts_only_transferred_bytes() {
504 let tracker = Arc::new(SfoSpeedStat::new());
505 let stream = PartialStream {
506 max_write_size: 3,
507 read_size: 3,
508 };
509 let mut stream = StatStream::new_with_tracker(stream, tracker.clone());
510
511 stream.write_all(&[0; 10]).await.unwrap();
512 assert_eq!(tracker.get_write_sum_size(), 10);
513
514 let mut storage = [0; 10];
515 let mut buf = ReadBuf::new(&mut storage);
516 buf.put_slice(&[0; 4]);
517 let result =
518 futures::future::poll_fn(|cx| Pin::new(&mut stream).poll_read(cx, &mut buf)).await;
519
520 assert!(result.is_ok());
521 assert_eq!(buf.filled().len(), 7);
522 assert_eq!(tracker.get_read_sum_size(), 3);
523 }
524
525 #[tokio::test]
526 async fn test_split_stat_io_counts_only_transferred_bytes() {
527 let write_tracker = Arc::new(SfoSpeedStat::new());
528 let writer = PartialStream {
529 max_write_size: 3,
530 read_size: 0,
531 };
532 let mut writer = StatWrite::new_with_tracker(writer, write_tracker.clone());
533
534 writer.write_all(&[0; 10]).await.unwrap();
535 assert_eq!(write_tracker.get_write_sum_size(), 10);
536
537 let zero_tracker = Arc::new(SfoSpeedStat::new());
538 let zero_writer = PartialStream {
539 max_write_size: 0,
540 read_size: 0,
541 };
542 let mut zero_writer = StatWrite::new_with_tracker(zero_writer, zero_tracker.clone());
543 let err = zero_writer.write_all(&[0]).await.unwrap_err();
544
545 assert_eq!(err.kind(), io::ErrorKind::WriteZero);
546 assert_eq!(zero_tracker.get_write_sum_size(), 0);
547
548 let read_tracker = Arc::new(SfoSpeedStat::new());
549 let reader = PartialStream {
550 max_write_size: 0,
551 read_size: 3,
552 };
553 let mut reader = StatRead::new_with_tracker(reader, read_tracker.clone());
554 let mut storage = [0; 10];
555 let mut buf = ReadBuf::new(&mut storage);
556 buf.put_slice(&[0; 4]);
557 let result =
558 futures::future::poll_fn(|cx| Pin::new(&mut reader).poll_read(cx, &mut buf)).await;
559
560 assert!(result.is_ok());
561 assert_eq!(buf.filled().len(), 7);
562 assert_eq!(read_tracker.get_read_sum_size(), 3);
563 }
564
565 #[test]
566 fn test_speed_state_new() {
567 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
569
570 struct MockTimePicker;
571
572 impl TimePicker for MockTimePicker {
573 fn now() -> u128 {
574 MOCK_TIME.load(Ordering::Relaxed) as u128
575 }
576 }
577
578 fn set_mock_time(time_ms: u64) {
580 MOCK_TIME.store(time_ms, Ordering::Relaxed);
581 }
582
583 set_mock_time(1000);
584 let state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
585
586 assert_eq!(state.sum_size, 0);
587 assert_eq!(state.last_time, 1000);
588 assert_eq!(state.data_items.len(), 0);
589 }
590
591 #[test]
592 fn test_add_data_same_second() {
593 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
595
596 struct MockTimePicker;
597
598 impl TimePicker for MockTimePicker {
599 fn now() -> u128 {
600 MOCK_TIME.load(Ordering::Relaxed) as u128
601 }
602 }
603
604 fn set_mock_time(time_ms: u64) {
606 MOCK_TIME.store(time_ms, Ordering::Relaxed);
607 }
608
609 set_mock_time(1500);
610 let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
611
612 state.add_data(100);
613 assert_eq!(state.sum_size, 100);
614 assert_eq!(state.data_items.len(), 1);
615 assert_eq!(state.data_items[0].size, 100);
616 assert_eq!(state.data_items[0].time, 1); state.add_data(200);
619 assert_eq!(state.sum_size, 300);
620 assert_eq!(state.data_items.len(), 1);
621 assert_eq!(state.data_items[0].size, 300); }
623
624 #[test]
625 fn test_add_data_different_seconds() {
626 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
628
629 struct MockTimePicker;
630
631 impl TimePicker for MockTimePicker {
632 fn now() -> u128 {
633 MOCK_TIME.load(Ordering::Relaxed) as u128
634 }
635 }
636
637 fn set_mock_time(time_ms: u64) {
639 MOCK_TIME.store(time_ms, Ordering::Relaxed);
640 }
641
642 fn advance_mock_time(delta_ms: u64) {
644 MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
645 }
646
647 set_mock_time(1000);
648 let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
649
650 state.add_data(100);
651 advance_mock_time(2000); state.add_data(200);
653
654 assert_eq!(state.sum_size, 300);
655 assert_eq!(state.data_items.len(), 2);
656 assert_eq!(state.data_items[0].size, 200);
657 assert_eq!(state.data_items[0].time, 1);
658 assert_eq!(state.data_items[1].size, 100);
659 assert_eq!(state.data_items[1].time, 2);
660
661 assert_eq!(state.get_speed(), 60);
662 advance_mock_time(500);
663 assert_eq!(state.get_speed(), 60);
664 state.add_data(300);
666 assert_eq!(state.sum_size, 600);
667 assert_eq!(state.get_speed(), 60);
668 advance_mock_time(500);
669 assert_eq!(state.get_speed(), 120);
670 }
671
672 #[test]
673 fn test_add_data_cross_seconds_distribution() {
674 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
676
677 struct MockTimePicker;
678
679 impl TimePicker for MockTimePicker {
680 fn now() -> u128 {
681 MOCK_TIME.load(Ordering::Relaxed) as u128
682 }
683 }
684
685 fn set_mock_time(time_ms: u64) {
687 MOCK_TIME.store(time_ms, Ordering::Relaxed);
688 }
689
690 fn advance_mock_time(delta_ms: u64) {
692 MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
693 }
694
695 set_mock_time(1500); let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
698 advance_mock_time(1500);
699
700 state.add_data(300);
704 advance_mock_time(1500);
705
706 assert_eq!(state.data_items.len(), 2);
710 assert_eq!(state.data_items[0].time, 1);
711 assert_eq!(state.data_items[0].size, 100);
712 assert_eq!(state.data_items[1].time, 2);
713 assert_eq!(state.data_items[1].size, 200);
714
715 set_mock_time(1500); let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
718 advance_mock_time(2000);
719
720 state.add_data(400);
724 advance_mock_time(1500);
725
726 assert_eq!(state.data_items.len(), 3);
730 assert_eq!(state.data_items[0].time, 1);
731 assert_eq!(state.data_items[0].size, 100);
732 assert_eq!(state.data_items[1].time, 2);
733 assert_eq!(state.data_items[1].size, 200);
734 assert_eq!(state.data_items[2].time, 3);
735 assert_eq!(state.data_items[2].size, 100);
736 }
737
738 #[test]
739 fn test_clear_invalid_item() {
740 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
742
743 struct MockTimePicker;
744
745 impl TimePicker for MockTimePicker {
746 fn now() -> u128 {
747 MOCK_TIME.load(Ordering::Relaxed) as u128
748 }
749 }
750
751 fn set_mock_time(time_ms: u64) {
753 MOCK_TIME.store(time_ms, Ordering::Relaxed);
754 }
755
756 let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
757
758 state.data_items.push(DataItem { size: 100, time: 5 }); state.data_items.push(DataItem { size: 200, time: 7 }); state.data_items.push(DataItem { size: 300, time: 8 }); set_mock_time(11000); state.clear_invalid_item(MockTimePicker::now());
765
766 assert_eq!(state.data_items.len(), 2);
767 assert_eq!(state.data_items[0].time, 7);
768 assert_eq!(state.data_items[1].time, 8);
769 }
770
771 #[test]
772 fn test_get_speed() {
773 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
775
776 struct MockTimePicker;
777
778 impl TimePicker for MockTimePicker {
779 fn now() -> u128 {
780 MOCK_TIME.load(Ordering::Relaxed) as u128
781 }
782 }
783
784 fn set_mock_time(time_ms: u64) {
786 MOCK_TIME.store(time_ms, Ordering::Relaxed);
787 }
788
789 fn advance_mock_time(delta_ms: u64) {
791 MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
792 }
793
794 set_mock_time(10000); let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
796
797 state.add_data(100);
798 advance_mock_time(1000);
799 state.add_data(200);
800 advance_mock_time(1000);
801 state.add_data(300);
802 advance_mock_time(1000);
803 state.add_data(400);
804 advance_mock_time(1000);
805 state.add_data(500);
806 advance_mock_time(1000);
807 state.add_data(600);
808 advance_mock_time(1000);
809 state.add_data(700);
810
811 let speed = state.get_speed();
812 assert_eq!(speed, 500);
814 }
815
816 #[test]
817 fn test_get_sum_size() {
818 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
820
821 struct MockTimePicker;
822
823 impl TimePicker for MockTimePicker {
824 fn now() -> u128 {
825 MOCK_TIME.load(Ordering::Relaxed) as u128
826 }
827 }
828
829 let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(10u64));
830 state.add_data(100);
831 state.add_data(200);
832 assert_eq!(state.get_sum_size(), 300);
833 }
834
835 #[test]
836 fn test_add_data_handles_clock_rollback() {
837 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
838
839 struct MockTimePicker;
840
841 impl TimePicker for MockTimePicker {
842 fn now() -> u128 {
843 MOCK_TIME.load(Ordering::Relaxed) as u128
844 }
845 }
846
847 fn set_mock_time(time_ms: u64) {
848 MOCK_TIME.store(time_ms, Ordering::Relaxed);
849 }
850
851 set_mock_time(10_000);
852 let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
853 state.add_data(100);
854
855 set_mock_time(9_000);
856 state.add_data(50);
857
858 assert_eq!(state.get_sum_size(), 150);
859 assert_eq!(state.last_time, 10_000);
860 assert_eq!(state.data_items.len(), 1);
861 assert_eq!(state.data_items[0].time, 10);
862 assert_eq!(state.data_items[0].size, 150);
863 }
864
865 #[test]
866 fn test_speed_queries_handle_future_items_without_overflow() {
867 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
868
869 struct MockTimePicker;
870
871 impl TimePicker for MockTimePicker {
872 fn now() -> u128 {
873 MOCK_TIME.load(Ordering::Relaxed) as u128
874 }
875 }
876
877 fn set_mock_time(time_ms: u64) {
878 MOCK_TIME.store(time_ms, Ordering::Relaxed);
879 }
880
881 let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
882 state.data_items.push(DataItem {
883 size: 100,
884 time: 12,
885 });
886
887 set_mock_time(11_000);
888 state.clear_invalid_item(MockTimePicker::now());
889 assert_eq!(state.data_items.len(), 1);
890 assert_eq!(state.get_speed(), 0);
891 }
892
893 #[test]
894 fn test_speed_stat_impl() {
895 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
897
898 struct MockTimePicker;
899
900 impl TimePicker for MockTimePicker {
901 fn now() -> u128 {
902 MOCK_TIME.load(Ordering::Relaxed) as u128
903 }
904 }
905
906 fn set_mock_time(time_ms: u64) {
908 MOCK_TIME.store(time_ms, Ordering::Relaxed);
909 }
910
911 let stat: SfoSpeedStat<MockTimePicker> = SfoSpeedStat::new_with_time_picker();
912
913 stat.add_write_data_size(100);
914 stat.add_read_data_size(200);
915
916 assert_eq!(stat.get_write_speed(), 0);
918 assert_eq!(stat.get_read_speed(), 0);
919
920 set_mock_time(5000);
922 stat.add_write_data_size(500);
923 stat.add_read_data_size(1000);
924 set_mock_time(6000);
925 stat.add_write_data_size(0);
926 stat.add_read_data_size(0);
927
928 assert!(stat.get_write_speed() > 0);
930 assert!(stat.get_read_speed() > 0);
931 }
932
933 #[tokio::test]
935 async fn test_stat_stream_creation() {
936 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
937
938 struct MockTimePicker;
939
940 impl TimePicker for MockTimePicker {
941 fn now() -> u128 {
942 MOCK_TIME.load(Ordering::Relaxed) as u128
943 }
944 }
945
946 fn advance_mock_time(delta_ms: u64) {
948 MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
949 }
950
951 struct MockStream {
953 future: Option<Pin<Box<dyn Future<Output = ()> + Send>>>,
954 }
955
956 impl AsyncRead for MockStream {
957 fn poll_read(
958 mut self: Pin<&mut Self>,
959 _cx: &mut Context<'_>,
960 _buf: &mut ReadBuf<'_>,
961 ) -> Poll<io::Result<()>> {
962 if self.future.is_none() {
963 self.future = Some(Box::pin(tokio::time::sleep(Duration::from_millis(10))));
964 }
965 match Pin::new(self.future.as_mut().unwrap()).poll(_cx) {
966 Poll::Ready(_) => {
967 self.future = None;
968 _buf.set_filled(10);
969 Poll::Ready(Ok(()))
970 }
971 Poll::Pending => Poll::Pending,
972 }
973 }
974 }
975
976 impl AsyncWrite for MockStream {
977 fn poll_write(
978 mut self: Pin<&mut Self>,
979 _cx: &mut Context<'_>,
980 buf: &[u8],
981 ) -> Poll<Result<usize, io::Error>> {
982 if self.future.is_none() {
983 self.future = Some(Box::pin(tokio::time::sleep(Duration::from_millis(10))));
984 }
985 match Pin::new(self.future.as_mut().unwrap()).poll(_cx) {
986 Poll::Ready(_) => {
987 self.future = None;
988 Poll::Ready(Ok(buf.len()))
989 }
990 Poll::Pending => Poll::Pending,
991 }
992 }
993
994 fn poll_flush(
995 self: Pin<&mut Self>,
996 _cx: &mut Context<'_>,
997 ) -> Poll<Result<(), io::Error>> {
998 Poll::Ready(Ok(()))
999 }
1000
1001 fn poll_shutdown(
1002 self: Pin<&mut Self>,
1003 _cx: &mut Context<'_>,
1004 ) -> Poll<Result<(), Error>> {
1005 Poll::Ready(Ok(()))
1006 }
1007 }
1008
1009 impl Unpin for MockStream {}
1010
1011 let stream = MockStream { future: None };
1012 let mut stat_stream = StatStream::new_test::<MockTimePicker>(stream);
1013 let speed_stat = stat_stream.get_speed_stat();
1014 let mut upload_size = 0;
1015 let mut download_size = 0;
1016 let mut buf = vec![0u8; 4096];
1017 advance_mock_time(500);
1018 for i in 0..100 {
1019 let size = stat_stream.write(&buf).await.unwrap();
1020 stat_stream.flush().await.unwrap();
1021 upload_size += size;
1022 let size = stat_stream.read(&mut buf).await.unwrap();
1023 download_size += size;
1024 advance_mock_time(1000);
1025 if i < 5 {
1026 assert_eq!(speed_stat.get_write_speed(), (upload_size / 5) as u64);
1027 assert_eq!(speed_stat.get_read_speed(), (download_size / 5) as u64);
1028 } else {
1029 assert_eq!(speed_stat.get_write_sum_size(), upload_size as u64);
1030 assert_eq!(speed_stat.get_read_sum_size(), download_size as u64);
1031 assert_eq!(speed_stat.get_write_speed(), (4096 * 5 - 2048) / 5);
1032 assert_eq!(speed_stat.get_read_speed(), (10 * 5 - 5) / 5);
1033 }
1034 }
1035 stat_stream.shutdown().await.unwrap();
1036 }
1037
1038 #[tokio::test]
1039 async fn test_stat_read_creation() {
1040 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
1041
1042 struct MockTimePicker;
1043
1044 impl TimePicker for MockTimePicker {
1045 fn now() -> u128 {
1046 MOCK_TIME.load(Ordering::Relaxed) as u128
1047 }
1048 }
1049
1050 fn advance_mock_time(delta_ms: u64) {
1052 MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
1053 }
1054
1055 struct MockReader {
1057 future: Option<Pin<Box<dyn Future<Output = ()> + Send>>>,
1058 }
1059
1060 impl AsyncRead for MockReader {
1061 fn poll_read(
1062 mut self: Pin<&mut Self>,
1063 _cx: &mut Context<'_>,
1064 _buf: &mut ReadBuf<'_>,
1065 ) -> Poll<io::Result<()>> {
1066 if self.future.is_none() {
1067 self.future = Some(Box::pin(tokio::time::sleep(Duration::from_millis(10))));
1068 }
1069 match Pin::new(self.future.as_mut().unwrap()).poll(_cx) {
1070 Poll::Ready(_) => {
1071 self.future = None;
1072 _buf.set_filled(10);
1073 Poll::Ready(Ok(()))
1074 }
1075 Poll::Pending => Poll::Pending,
1076 }
1077 }
1078 }
1079
1080 impl Unpin for MockReader {}
1081
1082 let reader = MockReader { future: None };
1083 let mut stat_reader = StatRead::new_test::<MockTimePicker>(reader);
1084 let speed_stat = stat_reader.get_speed_stat();
1085 let mut download_size = 0;
1086 let mut buf = vec![0u8; 4096];
1087 advance_mock_time(500);
1088 for i in 0..100 {
1089 let size = stat_reader.read(&mut buf).await.unwrap();
1090 download_size += size;
1091 advance_mock_time(1000);
1092 if i < 5 {
1093 assert_eq!(speed_stat.get_read_speed(), (download_size / 5) as u64);
1094 } else {
1095 assert_eq!(speed_stat.get_read_sum_size(), download_size as u64);
1096 assert_eq!(speed_stat.get_read_speed(), (10 * 5 - 5) / 5);
1097 }
1098 }
1099 }
1100
1101 #[tokio::test]
1102 async fn test_stat_write_creation() {
1103 static MOCK_TIME: AtomicU64 = AtomicU64::new(0);
1104
1105 struct MockTimePicker;
1106
1107 impl TimePicker for MockTimePicker {
1108 fn now() -> u128 {
1109 MOCK_TIME.load(Ordering::Relaxed) as u128
1110 }
1111 }
1112
1113 fn advance_mock_time(delta_ms: u64) {
1115 MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
1116 }
1117
1118 struct MockWriter {
1120 future: Option<Pin<Box<dyn Future<Output = ()> + Send>>>,
1121 }
1122
1123 impl AsyncWrite for MockWriter {
1124 fn poll_write(
1125 mut self: Pin<&mut Self>,
1126 _cx: &mut Context<'_>,
1127 buf: &[u8],
1128 ) -> Poll<Result<usize, io::Error>> {
1129 if self.future.is_none() {
1130 self.future = Some(Box::pin(tokio::time::sleep(Duration::from_millis(10))));
1131 }
1132 match Pin::new(self.future.as_mut().unwrap()).poll(_cx) {
1133 Poll::Ready(_) => {
1134 self.future = None;
1135 Poll::Ready(Ok(buf.len()))
1136 }
1137 Poll::Pending => Poll::Pending,
1138 }
1139 }
1140
1141 fn poll_flush(
1142 self: Pin<&mut Self>,
1143 _cx: &mut Context<'_>,
1144 ) -> Poll<Result<(), io::Error>> {
1145 Poll::Ready(Ok(()))
1146 }
1147
1148 fn poll_shutdown(
1149 self: Pin<&mut Self>,
1150 _cx: &mut Context<'_>,
1151 ) -> Poll<Result<(), Error>> {
1152 Poll::Ready(Ok(()))
1153 }
1154 }
1155
1156 impl Unpin for MockWriter {}
1157
1158 let writer = MockWriter { future: None };
1159 let mut stat_writer = StatWrite::new_test::<MockTimePicker>(writer);
1160 let speed_stat = stat_writer.get_speed_stat();
1161 let mut upload_size = 0;
1162 let buf = vec![0u8; 4096];
1163 advance_mock_time(500);
1164 for i in 0..100 {
1165 let size = stat_writer.write(&buf).await.unwrap();
1166 stat_writer.flush().await.unwrap();
1167 upload_size += size;
1168 advance_mock_time(1000);
1169 if i < 5 {
1170 assert_eq!(speed_stat.get_write_speed(), (upload_size / 5) as u64);
1171 } else {
1172 assert_eq!(speed_stat.get_write_sum_size(), upload_size as u64);
1173 assert_eq!(speed_stat.get_write_speed(), (4096 * 5 - 2048) / 5);
1174 }
1175 }
1176 stat_writer.shutdown().await.unwrap();
1177 }
1178}