1use async_trait::async_trait;
2use std::fmt;
3use std::sync::Arc;
4use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
5use std::time::{Duration, Instant};
6
7use tracing::{debug, warn};
8
9use crate::constants;
10use crate::error::Result;
11use crate::filesystem::disk_writer::DiskWriter;
12
13#[derive(Clone, Debug, Default)]
14pub struct RateLimiterConfig {
15 pub max_download_bytes_per_sec: Option<u64>,
16 pub max_upload_bytes_per_sec: Option<u64>,
17 pub download_burst_bytes: Option<u64>,
18 pub upload_burst_bytes: Option<u64>,
19}
20
21impl RateLimiterConfig {
22 pub fn new(download_limit: Option<u64>, upload_limit: Option<u64>) -> Self {
23 Self {
24 max_download_bytes_per_sec: download_limit,
25 max_upload_bytes_per_sec: upload_limit,
26 download_burst_bytes: None,
27 upload_burst_bytes: None,
28 }
29 }
30
31 pub fn with_burst(mut self, download_burst: Option<u64>, upload_burst: Option<u64>) -> Self {
32 self.download_burst_bytes = download_burst;
33 self.upload_burst_bytes = upload_burst;
34 self
35 }
36
37 pub fn is_limited(&self) -> bool {
38 self.max_download_bytes_per_sec.is_some() || self.max_upload_bytes_per_sec.is_some()
39 }
40
41 pub fn download_rate(&self) -> Option<u64> {
42 self.max_download_bytes_per_sec
43 }
44
45 pub fn upload_rate(&self) -> Option<u64> {
46 self.max_upload_bytes_per_sec
47 }
48
49 pub fn download_burst(&self) -> Option<u64> {
50 self.download_burst_bytes
51 }
52
53 pub fn upload_burst(&self) -> Option<u64> {
54 self.upload_burst_bytes
55 }
56}
57
58const NS_PER_SEC: u64 = 1_000_000_000;
60
61const MIN_SLEEP: Duration = Duration::from_micros(1);
65
66pub struct TokenBucket {
80 tokens_milli: AtomicU64,
83 capacity_milli: u64,
85 rate_milli_per_sec: AtomicU64,
89 last_refill_elapsed_ns: AtomicU64,
92 unlimited: AtomicBool,
95 anchor: Instant,
98}
99
100impl fmt::Debug for TokenBucket {
101 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
102 f.debug_struct("TokenBucket")
103 .field("tokens_milli", &self.tokens_milli.load(Ordering::Relaxed))
104 .field("capacity_milli", &self.capacity_milli)
105 .field(
106 "rate_milli_per_sec",
107 &self.rate_milli_per_sec.load(Ordering::Relaxed),
108 )
109 .field("unlimited", &self.unlimited.load(Ordering::Relaxed))
110 .finish()
111 }
112}
113
114impl TokenBucket {
115 pub fn new(rate_bytes_per_sec: u64, burst_bytes: Option<u64>) -> Self {
120 let burst = burst_bytes.unwrap_or(constants::DEFAULT_BURST_BYTES as u64);
121 let anchor = Instant::now();
122 Self {
123 tokens_milli: AtomicU64::new(burst.saturating_mul(1000)),
124 capacity_milli: burst.saturating_mul(1000),
125 rate_milli_per_sec: AtomicU64::new(rate_bytes_per_sec.saturating_mul(1000)),
126 last_refill_elapsed_ns: AtomicU64::new(0),
127 unlimited: AtomicBool::new(false),
128 anchor,
129 }
130 }
131
132 pub fn unlimited() -> Self {
135 let anchor = Instant::now();
136 let huge = u64::MAX / 4;
138 Self {
139 tokens_milli: AtomicU64::new(huge),
140 capacity_milli: huge,
141 rate_milli_per_sec: AtomicU64::new(huge),
142 last_refill_elapsed_ns: AtomicU64::new(0),
143 unlimited: AtomicBool::new(true),
144 anchor,
145 }
146 }
147
148 pub fn is_unlimited(&self) -> bool {
150 self.unlimited.load(Ordering::Relaxed)
151 }
152
153 pub fn rate(&self) -> f64 {
156 if self.unlimited.load(Ordering::Relaxed) {
157 f64::MAX
158 } else {
159 self.rate_milli_per_sec.load(Ordering::Relaxed) as f64 / 1000.0
160 }
161 }
162
163 pub fn available_tokens(&self) -> f64 {
167 if self.unlimited.load(Ordering::Relaxed) {
168 return f64::MAX;
169 }
170 self.refill();
171 self.tokens_milli.load(Ordering::Relaxed) as f64 / 1000.0
172 }
173
174 #[inline]
176 fn now_ns(&self) -> u64 {
177 Instant::now()
180 .saturating_duration_since(self.anchor)
181 .as_nanos() as u64
182 }
183
184 fn refill(&self) {
193 if self.unlimited.load(Ordering::Relaxed) {
194 return;
195 }
196 let now = self.now_ns();
197 let last = self.last_refill_elapsed_ns.load(Ordering::Relaxed);
198 if now <= last {
199 return;
201 }
202 let elapsed_ns = now - last;
203 let added_milli = ((elapsed_ns as u128)
205 * (self.rate_milli_per_sec.load(Ordering::Relaxed) as u128)
206 / NS_PER_SEC as u128) as u64;
207 if added_milli == 0 {
208 return;
211 }
212 match self.last_refill_elapsed_ns.compare_exchange(
215 last,
216 now,
217 Ordering::Relaxed,
218 Ordering::Relaxed,
219 ) {
220 Ok(_) => {
221 loop {
223 let current = self.tokens_milli.load(Ordering::Relaxed);
224 let new = current.saturating_add(added_milli).min(self.capacity_milli);
225 match self.tokens_milli.compare_exchange_weak(
226 current,
227 new,
228 Ordering::Relaxed,
229 Ordering::Relaxed,
230 ) {
231 Ok(_) => break,
232 Err(_) => continue, }
234 }
235 }
236 Err(_) => {
237 }
239 }
240 }
241
242 pub async fn acquire(&self, bytes: u64) {
250 if self.unlimited.load(Ordering::Relaxed) {
251 return;
252 }
253 let needed_milli = bytes.saturating_mul(1000);
255
256 loop {
257 self.refill();
258 let current = self.tokens_milli.load(Ordering::Relaxed);
259 if current >= needed_milli {
260 match self.tokens_milli.compare_exchange_weak(
262 current,
263 current - needed_milli,
264 Ordering::Relaxed,
265 Ordering::Relaxed,
266 ) {
267 Ok(_) => return,
268 Err(_) => continue, }
270 }
271
272 let deficit_milli = needed_milli - current;
274 let rate_milli = self.rate_milli_per_sec.load(Ordering::Relaxed);
275 if rate_milli == 0 {
276 warn!("TokenBucket::acquire with rate=0; treating as unlimited");
279 return;
280 }
281 let wait_ns =
284 ((deficit_milli as u128) * NS_PER_SEC as u128 / rate_milli as u128) as u64;
285 let wait = Duration::from_nanos(wait_ns);
286
287 if wait < MIN_SLEEP {
288 std::hint::spin_loop();
290 continue;
291 }
292
293 debug!(
294 bytes = bytes,
295 deficit_milli = deficit_milli,
296 wait_ns = wait_ns,
297 "throttling: sleeping for token refill"
298 );
299 tokio::time::sleep(wait).await;
300
301 self.refill();
305 loop {
306 let cur = self.tokens_milli.load(Ordering::Relaxed);
307 let new = cur.saturating_sub(needed_milli);
308 match self.tokens_milli.compare_exchange_weak(
309 cur,
310 new,
311 Ordering::Relaxed,
312 Ordering::Relaxed,
313 ) {
314 Ok(_) => return,
315 Err(_) => continue,
316 }
317 }
318 }
319 }
320
321 pub fn try_acquire(&self, bytes: u64) -> bool {
324 if self.unlimited.load(Ordering::Relaxed) {
325 return true;
326 }
327 self.refill();
328 let needed_milli = bytes.saturating_mul(1000);
329 loop {
330 let current = self.tokens_milli.load(Ordering::Relaxed);
331 if current < needed_milli {
332 return false;
333 }
334 match self.tokens_milli.compare_exchange_weak(
335 current,
336 current - needed_milli,
337 Ordering::Relaxed,
338 Ordering::Relaxed,
339 ) {
340 Ok(_) => return true,
341 Err(_) => continue,
342 }
343 }
344 }
345
346 pub fn set_rate(&self, rate_bytes_per_sec: u64) {
349 self.rate_milli_per_sec
350 .store(rate_bytes_per_sec.saturating_mul(1000), Ordering::Relaxed);
351 }
352
353 pub fn set_unlimited(&self, unlimited: bool) {
356 self.unlimited.store(unlimited, Ordering::Relaxed);
357 }
358}
359
360struct RateLimiterInner {
366 download: TokenBucket,
367 upload: TokenBucket,
368 download_limited: AtomicBool,
369 upload_limited: AtomicBool,
370}
371
372#[derive(Clone)]
379pub struct RateLimiter {
380 inner: Arc<RateLimiterInner>,
381}
382
383impl RateLimiter {
384 pub fn new(config: &RateLimiterConfig) -> Self {
385 let dl_rate = config.download_rate();
386 let ul_rate = config.upload_rate();
387 let dl_burst = config.download_burst();
388 let ul_burst = config.upload_burst();
389
390 let download = match dl_rate {
391 Some(rate) if rate > 0 => TokenBucket::new(rate, dl_burst),
392 _ => TokenBucket::unlimited(),
393 };
394 let upload = match ul_rate {
395 Some(rate) if rate > 0 => TokenBucket::new(rate, ul_burst),
396 _ => TokenBucket::unlimited(),
397 };
398
399 Self {
400 inner: Arc::new(RateLimiterInner {
401 download,
402 upload,
403 download_limited: AtomicBool::new(dl_rate.is_some_and(|r| r > 0)),
404 upload_limited: AtomicBool::new(ul_rate.is_some_and(|r| r > 0)),
405 }),
406 }
407 }
408
409 pub fn unlimited() -> Self {
410 Self::new(&RateLimiterConfig::default())
411 }
412
413 pub async fn acquire_download(&self, bytes: u64) {
414 self.inner.download.acquire(bytes).await;
415 }
416
417 pub async fn acquire_upload(&self, bytes: u64) {
418 self.inner.upload.acquire(bytes).await;
419 }
420
421 #[allow(clippy::unused_async)]
424 pub async fn try_acquire_download(&self, bytes: u64) -> bool {
425 self.inner.download.try_acquire(bytes)
426 }
427
428 #[allow(clippy::unused_async)]
431 pub async fn try_acquire_upload(&self, bytes: u64) -> bool {
432 self.inner.upload.try_acquire(bytes)
433 }
434
435 pub fn is_download_limited(&self) -> bool {
436 self.inner.download_limited.load(Ordering::Relaxed)
437 }
438
439 pub fn is_upload_limited(&self) -> bool {
440 self.inner.upload_limited.load(Ordering::Relaxed)
441 }
442
443 pub async fn config(&self) -> RateLimiterConfig {
444 RateLimiterConfig::new(
445 if self.inner.download.is_unlimited() {
446 None
447 } else {
448 Some(self.inner.download.rate() as u64)
449 },
450 if self.inner.upload.is_unlimited() {
451 None
452 } else {
453 Some(self.inner.upload.rate() as u64)
454 },
455 )
456 }
457
458 pub fn set_download_rate(&self, rate: Option<u64>) {
462 match rate {
463 Some(r) if r > 0 => {
464 self.inner.download.set_unlimited(false);
465 self.inner.download.set_rate(r);
466 self.inner.download_limited.store(true, Ordering::Relaxed);
467 }
468 _ => {
469 self.inner.download.set_unlimited(true);
470 self.inner.download_limited.store(false, Ordering::Relaxed);
471 }
472 }
473 }
474
475 pub fn set_upload_rate(&self, rate: Option<u64>) {
478 match rate {
479 Some(r) if r > 0 => {
480 self.inner.upload.set_unlimited(false);
481 self.inner.upload.set_rate(r);
482 self.inner.upload_limited.store(true, Ordering::Relaxed);
483 }
484 _ => {
485 self.inner.upload.set_unlimited(true);
486 self.inner.upload_limited.store(false, Ordering::Relaxed);
487 }
488 }
489 }
490}
491
492pub struct ThrottledWriter<W> {
499 inner: W,
500 limiter: RateLimiter,
501 chunk_size: usize,
502}
503
504impl<W> ThrottledWriter<W>
505where
506 W: DiskWriter + Send,
507{
508 pub fn new(inner: W, limiter: RateLimiter) -> Self {
509 Self {
510 inner,
511 limiter,
512 chunk_size: constants::RATE_LIMITER_CHUNK_SIZE,
513 }
514 }
515
516 pub fn with_chunk_size(mut self, size: usize) -> Self {
517 self.chunk_size = size.max(constants::RATE_LIMITER_MIN_CHUNK_SIZE);
518 self
519 }
520
521 pub fn into_inner(self) -> W {
522 self.inner
523 }
524
525 pub fn limiter(&self) -> &RateLimiter {
526 &self.limiter
527 }
528}
529
530#[async_trait]
531impl<W> DiskWriter for ThrottledWriter<W>
532where
533 W: DiskWriter + Send,
534{
535 async fn write(&mut self, data: &[u8]) -> Result<()> {
536 if !self.limiter.is_download_limited() {
537 return self.inner.write(data).await;
538 }
539
540 if data.len() <= self.chunk_size {
556 self.limiter.acquire_download(data.len() as u64).await;
557 return self.inner.write(data).await;
558 }
559
560 let mut offset = 0usize;
561 while offset < data.len() {
562 let end = (offset + self.chunk_size).min(data.len());
563 let chunk = &data[offset..end];
564 self.limiter.acquire_download(chunk.len() as u64).await;
565 self.inner.write(chunk).await?;
566 offset = end;
567 }
568 Ok(())
569 }
570
571 async fn finalize(&mut self) -> Result<Vec<u8>> {
572 self.inner.finalize().await
573 }
574}
575
576#[cfg(test)]
577mod tests {
578 use super::*;
579 use std::sync::Arc;
580
581 #[tokio::test]
582 async fn test_token_bucket_unlimited() {
583 let tb = TokenBucket::unlimited();
584 assert!(tb.is_unlimited());
585 tb.acquire(1024 * 1024 * 1024).await;
586 assert!(tb.available_tokens() > 0.0);
587 }
588
589 #[tokio::test]
590 async fn test_token_bucket_basic_acquire() {
591 let tb = TokenBucket::new(10000, Some(5000));
592 assert!(!tb.is_unlimited());
593
594 let start = Instant::now();
595 tb.acquire(5000).await;
596 let elapsed = start.elapsed();
597 assert!(
598 elapsed < Duration::from_millis(100),
599 "burst should be instant: {:?}",
600 elapsed
601 );
602
603 tb.acquire(6000).await;
604 let total_elapsed = start.elapsed();
605 let expected_min = Duration::from_millis(100);
606 assert!(
607 total_elapsed >= expected_min.saturating_sub(Duration::from_millis(200)),
608 "should have waited for refill: got {:?} expected >= {:?}",
609 total_elapsed,
610 expected_min
611 );
612 }
613
614 #[tokio::test]
615 async fn test_token_bucket_try_acquire() {
616 let tb = TokenBucket::new(1000, Some(2000));
617
618 assert!(tb.try_acquire(1000));
619 assert!(tb.try_acquire(1000));
620 assert!(!tb.try_acquire(1));
621 }
622
623 #[test]
624 fn test_token_bucket_available_tokens() {
625 let tb = TokenBucket::new(1000, Some(5000));
626 let initial = tb.available_tokens();
627 assert!(
628 (initial - 5000.0).abs() < 0.01,
629 "initial tokens should be ~5000, got {}",
630 initial
631 );
632
633 tb.try_acquire(2000);
634 let after = tb.available_tokens();
635 assert!(
636 (after - 3000.0).abs() < 0.01,
637 "after acquiring 2000, should have ~3000, got {}",
638 after
639 );
640 }
641
642 #[test]
643 fn test_rate_limiter_config_default() {
644 let cfg = RateLimiterConfig::default();
645 assert!(!cfg.is_limited());
646 assert!(cfg.download_rate().is_none());
647 assert!(cfg.upload_rate().is_none());
648 }
649
650 #[test]
651 fn test_rate_limiter_config_new() {
652 let cfg = RateLimiterConfig::new(Some(1024), Some(512));
653 assert!(cfg.is_limited());
654 assert_eq!(cfg.download_rate(), Some(1024));
655 assert_eq!(cfg.upload_rate(), Some(512));
656 }
657
658 #[test]
659 fn test_rate_limiter_config_download_only() {
660 let cfg = RateLimiterConfig::new(Some(2048), None);
661 assert!(cfg.is_limited());
662 assert_eq!(cfg.download_rate(), Some(2048));
663 assert!(cfg.upload_rate().is_none());
664 }
665
666 #[tokio::test]
667 async fn test_rate_limiter_unlimited() {
668 let rl = RateLimiter::unlimited();
669 assert!(!rl.is_download_limited());
670 assert!(!rl.is_upload_limited());
671 rl.acquire_download(999999).await;
672 rl.acquire_upload(999999).await;
673 }
674
675 #[tokio::test]
676 async fn test_rate_limiter_with_limits() {
677 let cfg = RateLimiterConfig::new(Some(5000), Some(1000)).with_burst(Some(1000), Some(500));
678 let rl = RateLimiter::new(&cfg);
679 assert!(rl.is_download_limited());
680 assert!(rl.is_upload_limited());
681
682 let start = Instant::now();
683 rl.acquire_download(6000).await;
684 let elapsed = start.elapsed();
685 assert!(
686 elapsed >= Duration::from_millis(800),
687 "should throttle: got {:?}",
688 elapsed
689 );
690 }
691
692 #[tokio::test]
693 async fn test_throttled_writer_no_limit_passthrough() {
694 use crate::filesystem::disk_writer::ByteArrayDiskWriter;
695
696 let raw = ByteArrayDiskWriter::new();
697 let rl = RateLimiter::unlimited();
698 let mut tw = ThrottledWriter::new(raw, rl);
699
700 tw.write(b"hello world").await.unwrap();
701 tw.write(b" foo bar baz").await.unwrap();
702 let result = tw.finalize().await.unwrap();
703
704 assert_eq!(result, b"hello world foo bar baz");
705 }
706
707 #[tokio::test]
708 async fn test_throttled_writer_with_limit() {
709 use crate::filesystem::disk_writer::ByteArrayDiskWriter;
710
711 let raw = ByteArrayDiskWriter::new();
712 let cfg = RateLimiterConfig::new(Some(100_000), None).with_burst(Some(1000), None);
713 let rl = RateLimiter::new(&cfg);
714 let mut tw = ThrottledWriter::new(raw, rl);
715
716 let data = vec![0xABu8; 50_000];
717 let start = Instant::now();
718 tw.write(&data).await.unwrap();
719 let elapsed = start.elapsed();
720
721 let result = tw.finalize().await.unwrap();
722 assert_eq!(result.len(), 50_000);
723 assert!(
724 elapsed >= Duration::from_millis(400),
725 "50KB at 100KB/s with 1KB burst should take >= 400ms, got {:?}",
726 elapsed
727 );
728 }
729
730 #[tokio::test]
731 async fn test_throttled_writer_chunk_size() {
732 use crate::filesystem::disk_writer::ByteArrayDiskWriter;
733
734 let raw = ByteArrayDiskWriter::new();
735 let cfg = RateLimiterConfig::new(Some(1_000_000), None);
736 let rl = RateLimiter::new(&cfg);
737 let mut tw = ThrottledWriter::new(raw, rl).with_chunk_size(1024);
738
739 let large_data = vec![0x42u8; 10_000];
740 tw.write(&large_data).await.unwrap();
741 let result = tw.finalize().await.unwrap();
742 assert_eq!(result.len(), 10_000);
743 }
744
745 #[tokio::test]
746 async fn test_rate_limiter_zero_rate_means_unlimited() {
747 let cfg = RateLimiterConfig::new(Some(0), Some(0));
748 let rl = RateLimiter::new(&cfg);
749 assert!(!rl.is_download_limited());
750 assert!(!rl.is_upload_limited());
751 }
752
753 #[tokio::test]
765 async fn test_token_bucket_concurrent_no_deadlock() {
766 let bucket = Arc::new(TokenBucket::new(10_000_000, Some(10_000_000)));
769
770 let mut handles = Vec::with_capacity(4);
771 for task_id in 0..4u8 {
772 let b = bucket.clone();
773 handles.push(tokio::spawn(async move {
774 for _ in 0..1000 {
775 b.acquire(1000).await;
776 }
777 task_id }));
779 }
780
781 for (i, h) in handles.into_iter().enumerate() {
783 let id = h.await.expect("task should complete without panic");
784 assert_eq!(id as usize, i, "task ordering preserved");
785 }
786
787 let remaining = bucket.available_tokens();
790 assert!(
791 remaining > 5_000_000.0,
792 "should have ~6MB left after consuming 4MB, got {}",
793 remaining
794 );
795 }
796
797 #[tokio::test]
806 async fn test_throttled_writer_batches_token_acquisition() {
807 use crate::filesystem::disk_writer::ByteArrayDiskWriter;
808
809 let raw = ByteArrayDiskWriter::new();
812 let cfg = RateLimiterConfig::new(Some(10_000), None).with_burst(Some(0), None);
813 let rl = RateLimiter::new(&cfg);
814 let mut tw = ThrottledWriter::new(raw, rl).with_chunk_size(100);
816
817 let data = vec![0x77u8; 5_000];
818 let start = Instant::now();
819 tw.write(&data).await.unwrap();
820 let elapsed = start.elapsed();
821 let result = tw.finalize().await.unwrap();
822
823 assert_eq!(result.len(), 5_000, "data integrity preserved");
824 assert!(
825 result.iter().all(|&b| b == 0x77),
826 "all bytes should be 0x77"
827 );
828
829 assert!(
832 elapsed >= Duration::from_millis(450),
833 "batch acquire should wait ~500ms, got {:?}",
834 elapsed
835 );
836
837 assert!(
842 elapsed < Duration::from_secs(2),
843 "batch acquire should complete well under 2s, got {:?}",
844 elapsed
845 );
846 }
847
848 #[tokio::test]
852 async fn test_rate_limiter_clone_shares_state() {
853 let cfg = RateLimiterConfig::new(Some(10000), None).with_burst(Some(5000), None);
854 let rl1 = RateLimiter::new(&cfg);
855 let rl2 = rl1.clone();
856
857 assert!(rl1.is_download_limited());
858 assert!(rl2.is_download_limited());
859
860 rl1.acquire_download(3000).await;
862 rl2.acquire_download(3000).await;
863
864 let config = rl1.config().await;
865 assert!(config.download_rate().is_some());
866 }
867
868 #[tokio::test]
871 async fn test_rate_limiter_high_rate_low_latency() {
872 let cfg = RateLimiterConfig::new(Some(100_000_000), None).with_burst(Some(1_000_000), None);
873 let rl = RateLimiter::new(&cfg);
874
875 let start = Instant::now();
876 rl.acquire_download(100_000).await;
877 let elapsed = start.elapsed();
878 assert!(
879 elapsed < Duration::from_millis(50),
880 "100MB/s rate with 1MB burst should be near-instant for 100KB: got {:?}",
881 elapsed
882 );
883 }
884
885 #[tokio::test]
897 async fn test_token_bucket_set_rate() {
898 let tb = TokenBucket::new(1_000_000, Some(1000)); assert!(tb.try_acquire(1000));
901
902 tb.set_rate(10_000_000);
904 let rate = tb.rate();
905 assert!(
906 (rate - 10_000_000.0).abs() < 1.0,
907 "rate should be ~10 MB/s, got {}",
908 rate
909 );
910 }
911
912 #[tokio::test]
914 async fn test_token_bucket_set_rate_to_zero() {
915 let tb = TokenBucket::new(1_000_000, Some(1000));
916 tb.set_rate(0);
917 let rate = tb.rate();
918 assert!(
919 (rate - 0.0).abs() < 0.01,
920 "rate should be 0 after set_rate(0), got {}",
921 rate
922 );
923 }
924
925 #[tokio::test]
928 async fn test_token_bucket_set_unlimited() {
929 let tb = TokenBucket::new(1_000, None); assert!(!tb.is_unlimited());
931 tb.set_unlimited(true);
932 assert!(tb.is_unlimited());
933
934 let start = Instant::now();
936 tb.acquire(1_000_000_000).await;
937 let elapsed = start.elapsed();
938 assert!(
939 elapsed < Duration::from_millis(50),
940 "unlimited acquire should be instant, got {:?}",
941 elapsed
942 );
943 }
944
945 #[tokio::test]
948 async fn test_rate_limiter_set_download_rate() {
949 let rl = RateLimiter::new(&RateLimiterConfig::new(Some(1_000_000), None)); assert!(rl.is_download_limited());
951
952 rl.set_download_rate(Some(5_000_000));
954 assert!(rl.is_download_limited());
955 let config = rl.config().await;
956 assert_eq!(config.download_rate(), Some(5_000_000));
957
958 rl.set_download_rate(None);
960 assert!(!rl.is_download_limited());
961 let config = rl.config().await;
962 assert!(
963 config.download_rate().is_none(),
964 "download_rate should be None after set_download_rate(None), got {:?}",
965 config.download_rate()
966 );
967 }
968
969 #[tokio::test]
972 async fn test_rate_limiter_set_upload_rate() {
973 let rl = RateLimiter::new(&RateLimiterConfig::new(None, Some(500_000))); assert!(rl.is_upload_limited());
975
976 rl.set_upload_rate(Some(2_000_000)); assert!(rl.is_upload_limited());
978 let config = rl.config().await;
979 assert_eq!(config.upload_rate(), Some(2_000_000));
980
981 rl.set_upload_rate(Some(0));
983 assert!(!rl.is_upload_limited());
984 let config = rl.config().await;
985 assert!(config.upload_rate().is_none());
986 }
987
988 #[tokio::test]
993 async fn test_rate_limiter_clone_shares_inner_state() {
994 let rl = RateLimiter::new(&RateLimiterConfig::new(Some(1_000_000), None));
995 let rl_clone = rl.clone();
996
997 rl.set_download_rate(Some(5_000_000));
999
1000 let config = rl_clone.config().await;
1002 assert_eq!(
1003 config.download_rate(),
1004 Some(5_000_000),
1005 "clone should see updated rate via shared Arc<inner>"
1006 );
1007 assert!(
1008 rl_clone.is_download_limited(),
1009 "clone should see updated limited flag via shared Arc<inner>"
1010 );
1011
1012 rl_clone.set_download_rate(None);
1014 assert!(
1015 !rl.is_download_limited(),
1016 "original should see unlimited flag set by clone"
1017 );
1018 }
1019}