Skip to main content

sfo_io/
stat_stream.rs

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    /// Creates a new SfoSpeedStat instance with the specified duration
160    ///
161    /// # Parameters
162    /// * `duration` - The duration for statistics, in seconds
163    ///
164    /// # Returns
165    /// Returns a new SfoSpeedStat instance containing initialized upload and download states
166    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        // Mock TimePicker for testing
568        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        // Helper function to set mock time
579        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        // Mock TimePicker for testing
594        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        // Helper function to set mock time
605        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); // 1000 / 1000 = 1
617
618        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); // 合并到同一秒
622    }
623
624    #[test]
625    fn test_add_data_different_seconds() {
626        // Mock TimePicker for testing
627        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        // Helper function to set mock time
638        fn set_mock_time(time_ms: u64) {
639            MOCK_TIME.store(time_ms, Ordering::Relaxed);
640        }
641
642        // Helper function to advance mock time
643        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); // 时间前进到3000ms
652        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        //
665        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        // Mock TimePicker for testing
675        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        // Helper function to set mock time
686        fn set_mock_time(time_ms: u64) {
687            MOCK_TIME.store(time_ms, Ordering::Relaxed);
688        }
689
690        // Helper function to advance mock time
691        fn advance_mock_time(delta_ms: u64) {
692            MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
693        }
694
695        // 测试跨秒时数据如何分配到不同的秒中
696        set_mock_time(1500); // 1.5秒
697        let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
698        advance_mock_time(1500);
699
700        // 从1500ms到3000ms,增加1500ms,跨越2个完整的秒(2s和3s)
701        // 1.5s到2s有500ms,2s到3s有1000ms
702        // 总共1500ms,添加300字节数据
703        state.add_data(300);
704        advance_mock_time(1500);
705
706        // 应该创建两个数据项: 一个在第2秒,一个在第3秒
707        // 第2秒应该有 300 * 500/1500 = 100 字节
708        // 第3秒应该有 300 * 1000/1500 = 200 字节
709        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        // 测试跨秒时数据如何分配到不同的秒中
716        set_mock_time(1500); // 1.5秒
717        let mut state: SpeedState<MockTimePicker> = SpeedState::new(nonzero!(5u64));
718        advance_mock_time(2000);
719
720        // 从1500ms到3000ms,增加1500ms,跨越2个完整的秒(2s和3s)
721        // 1.5s到2s有500ms,2s到3s有1000ms
722        // 总共1500ms,添加300字节数据
723        state.add_data(400);
724        advance_mock_time(1500);
725
726        // 应该创建两个数据项: 一个在第2秒,一个在第3秒
727        // 第2秒应该有 300 * 500/1500 = 100 字节
728        // 第3秒应该有 300 * 1000/1500 = 200 字节
729        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        // Mock TimePicker for testing
741        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        // Helper function to set mock time
752        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        // 添加几个不同时间的数据项
759        state.data_items.push(DataItem { size: 100, time: 5 }); // 5秒时的数据,应该被清除
760        state.data_items.push(DataItem { size: 200, time: 7 }); // 7秒时的数据,应该保留
761        state.data_items.push(DataItem { size: 300, time: 8 }); // 8秒时的数据,应该保留
762
763        set_mock_time(11000); // 当前时间10秒
764        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        // Mock TimePicker for testing
774        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        // Helper function to set mock time
785        fn set_mock_time(time_ms: u64) {
786            MOCK_TIME.store(time_ms, Ordering::Relaxed);
787        }
788
789        // Helper function to advance mock time
790        fn advance_mock_time(delta_ms: u64) {
791            MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
792        }
793
794        set_mock_time(10000); // 10秒
795        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        // 应该计算 300+400+500+600+700 = 2500 字节在4秒内 => 2500/5 = 500 bytes/sec
813        assert_eq!(speed, 500);
814    }
815
816    #[test]
817    fn test_get_sum_size() {
818        // Mock TimePicker for testing
819        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        // Mock TimePicker for testing
896        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        // Helper function to set mock time
907        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        // 由于没有时间流逝,速度为0
917        assert_eq!(stat.get_write_speed(), 0);
918        assert_eq!(stat.get_read_speed(), 0);
919
920        // 模拟时间流逝后再次检查
921        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        // 现在应该有速度了
929        assert!(stat.get_write_speed() > 0);
930        assert!(stat.get_read_speed() > 0);
931    }
932
933    // 注意: StatStream的测试需要tokio运行时,这里省略了复杂的AsyncRead/AsyncWrite mock
934    #[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        // Helper function to advance mock time
947        fn advance_mock_time(delta_ms: u64) {
948            MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
949        }
950
951        // 创建一个简单的mock stream用于测试
952        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        // Helper function to advance mock time
1051        fn advance_mock_time(delta_ms: u64) {
1052            MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
1053        }
1054
1055        // 创建一个简单的mock reader用于测试
1056        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        // Helper function to advance mock time
1114        fn advance_mock_time(delta_ms: u64) {
1115            MOCK_TIME.fetch_add(delta_ms, Ordering::Relaxed);
1116        }
1117
1118        // 创建一个简单的mock writer用于测试
1119        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}