tor_basic_utils/token_bucket.rs
1//! A token bucket implementation.
2
3use std::fmt::Debug;
4use web_time_compat::{Duration, Instant};
5
6/// A token bucket.
7///
8/// Calculations are performed at microsecond resolution.
9/// You likely want to call [`refill()`](Self::refill) each time you want to access or perform an
10/// operation on the token bucket.
11///
12/// This is partially inspired by tor's `token_bucket_ctr_t`,
13/// but the implementation is quite a bit different.
14/// We use larger values here (for example `u64`),
15/// and we aim to avoid drift when refills occur at times that aren't exactly in period with the
16/// refill rate.
17///
18/// It's possible that we could relax these requirements to reduce memory usage and computation
19/// complexity, but that optimization should probably only be made if/when needed since it would
20/// make the code more difficult to reason about, and possibly more complex.
21#[derive(Debug)]
22pub struct TokenBucket<I> {
23 /// The refill rate in tokens/second.
24 rate: u64,
25 /// The max amount of tokens in the bucket.
26 /// Commonly referred to as the "burst".
27 bucket_max: u64,
28 /// Current amount of tokens in the bucket.
29 // It's possible that in the future we may want a token bucket to allow negative values. For
30 // example we might want to send a few extra bytes over the allowed limit if it would mean that
31 // we send a complete TLS record.
32 bucket: u64,
33 /// Time that the most recent token was added to the bucket.
34 ///
35 /// While this can be thought of as the last time the bucket was partially refilled, it more
36 /// specifically is the time that the most recent token was added. For example if the bucket
37 /// refills one token every 100 ms, and the bucket is refilled at time 510 ms, the bucket would
38 /// gain 5 tokens and the stored time would be 500 ms.
39 added_tokens_at: I,
40}
41
42impl<I: TokenBucketInstant> TokenBucket<I> {
43 /// A new [`TokenBucket`] with a given `rate` in tokens/second and a `max` token limit.
44 ///
45 /// The bucket will initially be full.
46 /// The value `max` is commonly referred to as the "burst".
47 pub fn new(config: &TokenBucketConfig, now: I) -> Self {
48 Self {
49 rate: config.rate,
50 bucket_max: config.bucket_max,
51 bucket: config.bucket_max,
52 added_tokens_at: now,
53 }
54 }
55
56 /// Are there no tokens in the bucket?
57 pub fn is_empty(&self) -> bool {
58 self.bucket == 0
59 }
60
61 /// The maximum number of tokens that this bucket can hold.
62 pub fn max(&self) -> u64 {
63 self.bucket_max
64 }
65
66 /// Remove `count` tokens from the bucket.
67 pub fn drain(&mut self, count: u64) -> Result<BecameEmpty, InsufficientTokensError> {
68 Ok(self.claim(count)?.commit())
69 }
70
71 /// Drain all tokens in the bucket.
72 ///
73 /// Return how many were removed.
74 pub fn drain_all(&mut self) -> u64 {
75 let count = self.bucket;
76 self.bucket = 0;
77 count
78 }
79
80 /// Claim a number of tokens.
81 ///
82 /// The claim will be held by the returned [`ClaimedTokens`], and committed when dropped.
83 ///
84 /// **Note:** You probably want to call [`refill()`](Self::refill) before this.
85 // Since the `ClaimedTokens` holds a `&mut` to this `TokenBucket`, we don't need to worry about
86 // other calls accessing the `TokenBucket` before the `ClaimedTokens` are committed.
87 pub fn claim(&mut self, count: u64) -> Result<ClaimedTokens<I>, InsufficientTokensError> {
88 if count > self.bucket {
89 return Err(InsufficientTokensError {
90 available: self.bucket,
91 });
92 }
93
94 Ok(ClaimedTokens::new(self, count))
95 }
96
97 /// Adjust the refill rate and max tokens of the bucket.
98 ///
99 /// The token bucket is refilled up to `now` before changing the rate.
100 ///
101 /// If the new max is smaller than the existing number of tokens,
102 /// the number of tokens will be reduced to the new max.
103 ///
104 /// A rate and/or max of 0 is allowed.
105 pub fn adjust(&mut self, now: I, config: &TokenBucketConfig) {
106 // make sure that the bucket gets the tokens it is owed before we change the rate
107 self.refill(now);
108
109 // If the old rate was small (or 0), the `refill()` might not have updated
110 // `added_tokens_at`.
111 //
112 // For example if the bucket has a rate of 0 and was last refilled 10 seconds ago, it will
113 // not have gained any tokens in the last 10 seconds. If we were to only update the rate to
114 // 100 tokens/second now, the bucket would immediately become eligible to refill 1000
115 // tokens. We only want the rate change to become effective now, not in the past, so we
116 // ensure this by resetting `added_tokens_at`.
117 self.added_tokens_at = std::cmp::max(self.added_tokens_at, now);
118
119 self.rate = config.rate;
120 self.bucket_max = config.bucket_max;
121 self.bucket = std::cmp::min(self.bucket, self.bucket_max);
122 }
123
124 /// An estimated time at which the bucket will have `tokens` available.
125 ///
126 /// It is not guaranteed that `tokens` will be available at the returned time.
127 ///
128 /// If there are already enough tokens available, a time in the past may be returned.
129 ///
130 /// A value of `None` implies "never",
131 /// for example if the refill rate is 0,
132 /// the bucket max is too small,
133 /// or the time is too large to be represented as an `I`.
134 pub fn tokens_available_at(&self, tokens: u64) -> Result<I, NeverEnoughTokensError> {
135 let tokens_needed = tokens.saturating_sub(self.bucket);
136
137 // check if we currently have enough tokens before considering refilling
138 if tokens_needed == 0 {
139 return Ok(self.added_tokens_at);
140 }
141
142 // if the rate is 0, we'll never get more tokens
143 if self.rate == 0 {
144 return Err(NeverEnoughTokensError::ZeroRate);
145 }
146
147 // if more tokens are wanted than the capacity of the bucket, we'll never get enough
148 if tokens > self.bucket_max {
149 return Err(NeverEnoughTokensError::ExceedsMaxTokens);
150 }
151
152 // this may underestimate the time if either argument is very large
153 let time_needed = Self::tokens_to_duration(tokens_needed, self.rate)
154 .ok_or(NeverEnoughTokensError::ZeroRate)?;
155
156 // Always return at least 1 microsecond since:
157 // 1. We don't want to return `Duration::ZERO` if the tokens aren't ready,
158 // which may occur if the rate is very large (<1 ns/token).
159 // 2. Clocks generally don't operate at <1 us resolution.
160 let time_needed = std::cmp::max(time_needed, Duration::from_micros(1));
161
162 self.added_tokens_at
163 .checked_add(time_needed)
164 .ok_or(NeverEnoughTokensError::InstantNotRepresentable)
165 }
166
167 /// Refill the bucket.
168 pub fn refill(&mut self, now: I) -> BecameNonEmpty {
169 // time since we last added tokens
170 let elapsed = now.saturating_duration_since(self.added_tokens_at);
171
172 // If we exceeded the threshold, update the timestamp and return.
173 // This is taken from tor, which has the comment below:
174 //
175 // > Skip over updates that include an overflow or a very large jump. This can happen for
176 // > platform specific reasons, such as the old ~48 day windows timer.
177 //
178 // It's unclear if this type of OS bug is still common enough that this check is useful,
179 // but it shouldn't hurt.
180 if elapsed > I::IGNORE_THRESHOLD {
181 tracing::debug!(
182 "Time jump of {elapsed:?} is larger than {:?}; not refilling token bucket",
183 I::IGNORE_THRESHOLD,
184 );
185 self.added_tokens_at = now;
186 return BecameNonEmpty::No;
187 }
188
189 let old_bucket = self.bucket;
190
191 // Compute how much we should increment the bucket by.
192 // This may be underestimated in some cases.
193 let bucket_inc = Self::duration_to_tokens(elapsed, self.rate);
194
195 self.bucket = std::cmp::min(self.bucket_max, self.bucket.saturating_add(bucket_inc));
196
197 // Compute how much we should increment the `last_added_tokens` time by. This avoids
198 // drifting if the `bucket_inc` was underestimated, and avoids rounding errors which could
199 // cause the token bucket to effectively use a lower rate. For example if the rate was
200 // "1 token / sec" and the elapsed time was "1.2 sec", we only want to refill 1 token and
201 // increment the time by 1 second.
202 //
203 // While the docs for `tokens_to_duration` say that a smaller than expected duration may be
204 // returned, we have a test `test_duration_token_round_trip` which ensures that
205 // `tokens_to_duration` returns the expected value when used with the result from
206 // `duration_to_tokens`.
207 let added_tokens_at_inc =
208 Self::tokens_to_duration(bucket_inc, self.rate).unwrap_or(Duration::ZERO);
209
210 self.added_tokens_at = self
211 .added_tokens_at
212 .checked_add(added_tokens_at_inc)
213 .expect("overflowed time");
214 debug_assert!(self.added_tokens_at <= now);
215
216 if old_bucket == 0 && self.bucket != 0 {
217 BecameNonEmpty::Yes
218 } else {
219 BecameNonEmpty::No
220 }
221 }
222
223 /// How long would it take to refill `tokens` at `rate`?
224 ///
225 /// The result is rounded up to the nearest microsecond.
226 /// If the number of `tokens` is large,
227 /// the result may be much lower than the expected duration due to saturating 64-bit arithmetic.
228 ///
229 /// `None` will be returned if the `rate` is 0.
230 fn tokens_to_duration(tokens: u64, rate: u64) -> Option<Duration> {
231 // Perform the calculation in microseconds rather than nanoseconds since timers typically
232 // have microsecond granularity, and it lowers the chance that the calculation overflows the
233 // `u64::MAX` limit compared to nanoseconds. In the case that the calculation saturates, the
234 // returned duration will be shorter than the real value.
235 //
236 // For example with `tokens = u64::MAX` and `rate = u64::MAX` we'd expect a result of 1
237 // second, but:
238 // u64::MAX.saturating_mul(1000 * 1000).div_ceil(u64::MAX) = 1 microsecond
239 //
240 // The `div_ceil` ensures we always round up to the nearest microsecond.
241 //
242 // dimensional analysis:
243 // (tokens) * (microseconds / second) / (tokens / second) = microseconds
244 if rate == 0 {
245 return None;
246 }
247 let micros = tokens.saturating_mul(1000 * 1000).div_ceil(rate);
248 Some(Duration::from_micros(micros))
249 }
250
251 /// How many tokens would be refilled within `time` at `rate`?
252 ///
253 /// The `time` is truncated to microsecond granularity.
254 /// If the `time` or `rate` is large,
255 /// the result may be much lower than the expected number of tokens due to saturating 64-bit
256 /// arithmetic.
257 fn duration_to_tokens(time: Duration, rate: u64) -> u64 {
258 let micros = u64::try_from(time.as_micros()).unwrap_or(u64::MAX);
259 // dimensional analysis:
260 // (tokens / second) * (microseconds) / (microseconds / second) = tokens
261 rate.saturating_mul(micros) / (1000 * 1000)
262 }
263}
264
265/// The refill rate and token max for a [`TokenBucket`].
266#[derive(Clone, Debug)]
267#[allow(clippy::exhaustive_structs)] // constructed directly by callers configuring the bucket
268pub struct TokenBucketConfig {
269 /// The refill rate in tokens/second.
270 pub rate: u64,
271 /// The max amount of tokens in the bucket.
272 /// Commonly referred to as the "burst".
273 pub bucket_max: u64,
274}
275
276/// A handle to a number of claimed tokens.
277///
278/// Dropping this handle will commit the claim.
279#[derive(Debug)]
280pub struct ClaimedTokens<'a, I> {
281 /// The bucket that the claim is for.
282 bucket: &'a mut TokenBucket<I>,
283 /// How many tokens to remove from the bucket.
284 count: u64,
285}
286
287impl<'a, I> ClaimedTokens<'a, I> {
288 /// Create a new [`ClaimedTokens`] that will remove `count` tokens from the token `bucket` when
289 /// dropped.
290 fn new(bucket: &'a mut TokenBucket<I>, count: u64) -> Self {
291 Self { bucket, count }
292 }
293
294 /// Commit the claimed tokens.
295 ///
296 /// This is equivalent to just dropping the [`ClaimedTokens`], but also returns whether the
297 /// token bucket became empty or not.
298 pub fn commit(mut self) -> BecameEmpty {
299 self.commit_impl()
300 }
301
302 /// Reduce the claim to a fewer number of tokens than the original claim.
303 ///
304 /// If `count` is larger than the original claim, an error will be returned containing the
305 /// current number of claimed tokens.
306 pub fn reduce(&mut self, count: u64) -> Result<(), InsufficientTokensError> {
307 if count > self.count {
308 return Err(InsufficientTokensError {
309 available: self.count,
310 });
311 }
312
313 self.count = count;
314 Ok(())
315 }
316
317 /// Discard the claim.
318 ///
319 /// This does not remove any tokens from the token bucket.
320 pub fn discard(mut self) {
321 self.count = 0;
322 }
323
324 /// The commit implementation.
325 ///
326 /// After calling [`commit_impl()`](Self::commit_impl),
327 /// the [`ClaimedTokens`] should no longer be used and should be dropped immediately.
328 fn commit_impl(&mut self) -> BecameEmpty {
329 // when the `ClaimedTokens` was created by the `TokenBucket`, it should have ensured that
330 // there were enough tokens
331 self.bucket.bucket = self
332 .bucket
333 .bucket
334 .checked_sub(self.count)
335 .unwrap_or_else(|| {
336 panic!(
337 "claim commit failed: {}, {}",
338 self.count, self.bucket.bucket,
339 )
340 });
341
342 // when `self` is dropped some time after this function ends,
343 // we don't want to subtract again
344 self.count = 0;
345
346 if self.bucket.bucket > 0 {
347 BecameEmpty::No
348 } else {
349 BecameEmpty::Yes
350 }
351 }
352}
353
354impl<'a, I> std::ops::Drop for ClaimedTokens<'a, I> {
355 fn drop(&mut self) {
356 self.commit_impl();
357 }
358}
359
360/// An operation was attempted to reduce the number of tokens,
361/// but the token bucket did not have enough tokens.
362#[derive(Copy, Clone, Debug, PartialEq, Eq, thiserror::Error)]
363#[error("insufficient tokens for operation")]
364pub struct InsufficientTokensError {
365 /// The number of tokens that are available to drain/commit.
366 available: u64,
367}
368
369impl InsufficientTokensError {
370 /// Get the number of tokens that are available to drain/commit.
371 pub fn available_tokens(&self) -> u64 {
372 self.available
373 }
374}
375
376/// The token bucket will never have the requested number of tokens.
377#[derive(Copy, Clone, Debug, PartialEq, Eq, thiserror::Error)]
378#[allow(clippy::exhaustive_enums)] // callers exhaustively match on these variants
379#[error("there will never be enough tokens for this operation")]
380pub enum NeverEnoughTokensError {
381 /// The request exceeds the bucket's maximum number of tokens.
382 ExceedsMaxTokens,
383 /// The refill rate is 0.
384 ZeroRate,
385 /// The time is not representable.
386 ///
387 /// For example the if the rate is low and a large number of tokens were requested, it may be
388 /// too far in the future that it cannot be represented as a time value.
389 InstantNotRepresentable,
390}
391
392/// The token bucket transitioned from "empty" to "non-empty".
393#[derive(Copy, Clone, Debug, PartialEq, Eq)]
394#[allow(clippy::exhaustive_enums)] // a simple yes/no status that callers match on
395pub enum BecameNonEmpty {
396 /// Token bucket became non-empty.
397 Yes,
398 /// Token bucket remains empty.
399 No,
400}
401
402/// The token bucket transitioned from "non-empty" to "empty".
403#[derive(Copy, Clone, Debug, PartialEq, Eq)]
404#[allow(clippy::exhaustive_enums)] // a simple yes/no status that callers match on
405pub enum BecameEmpty {
406 /// Token bucket became empty.
407 Yes,
408 /// Token bucket remains non-empty.
409 No,
410}
411
412/// Any type implementing this must be represented as a measurement of a monotonically nondecreasing
413/// clock.
414pub trait TokenBucketInstant: Copy + Clone + Debug + PartialEq + Eq + PartialOrd + Ord {
415 /// An unrealistically large time jump.
416 ///
417 /// We assume that any time change larger than this indicates a broken monotonic clock,
418 /// and the bucket will not be refilled.
419 const IGNORE_THRESHOLD: Duration;
420
421 /// See [`Instant::checked_add`].
422 fn checked_add(&self, duration: Duration) -> Option<Self>;
423
424 /// See [`Instant::checked_duration_since`].
425 fn checked_duration_since(&self, earlier: Self) -> Option<Duration>;
426
427 /// See [`Instant::saturating_duration_since`].
428 fn saturating_duration_since(&self, earlier: Self) -> Duration {
429 self.checked_duration_since(earlier).unwrap_or_default()
430 }
431}
432
433impl TokenBucketInstant for Instant {
434 // This value is taken from tor (see `elapsed_ticks <= UINT32_MAX/4` in
435 // `src/lib/evloop/token_bucket.c`).
436 const IGNORE_THRESHOLD: Duration = Duration::from_secs((u32::MAX / 4) as u64);
437
438 #[inline]
439 fn checked_add(&self, duration: Duration) -> Option<Self> {
440 self.checked_add(duration)
441 }
442
443 #[inline]
444 fn checked_duration_since(&self, earlier: Self) -> Option<Duration> {
445 self.checked_duration_since(earlier)
446 }
447
448 #[inline]
449 fn saturating_duration_since(&self, earlier: Self) -> Duration {
450 self.saturating_duration_since(earlier)
451 }
452}
453
454#[cfg(test)]
455mod test {
456 #![allow(clippy::unwrap_used)]
457
458 use super::*;
459
460 use rand::RngExt;
461
462 #[derive(Copy, Clone, Debug, PartialEq, Eq, PartialOrd, Ord)]
463 struct MillisTimestamp(u64);
464
465 impl TokenBucketInstant for MillisTimestamp {
466 const IGNORE_THRESHOLD: Duration = Duration::from_millis(1_000_000_000);
467
468 fn checked_add(&self, duration: Duration) -> Option<Self> {
469 let duration = u64::try_from(duration.as_millis()).ok()?;
470 self.0.checked_add(duration).map(Self)
471 }
472
473 fn checked_duration_since(&self, earlier: Self) -> Option<Duration> {
474 Some(Duration::from_millis(self.0.checked_sub(earlier.0)?))
475 }
476 }
477
478 #[test]
479 fn adjust_now() {
480 let time = MillisTimestamp(100);
481
482 let config = TokenBucketConfig {
483 rate: 10,
484 bucket_max: 100,
485 };
486 let mut tb = TokenBucket::new(&config, time);
487 assert_eq!(tb.bucket, 100);
488 assert_eq!(tb.bucket_max, 100);
489 assert_eq!(tb.rate, 10);
490
491 tb.adjust(
492 time,
493 &TokenBucketConfig {
494 rate: 20,
495 bucket_max: 100,
496 },
497 );
498 assert_eq!(tb.bucket, 100);
499 assert_eq!(tb.bucket_max, 100);
500
501 tb.adjust(
502 time,
503 &TokenBucketConfig {
504 rate: 20,
505 bucket_max: 40,
506 },
507 );
508 assert_eq!(tb.bucket, 40);
509 assert_eq!(tb.bucket_max, 40);
510
511 tb.adjust(
512 time,
513 &TokenBucketConfig {
514 rate: 20,
515 bucket_max: 100,
516 },
517 );
518 assert_eq!(tb.bucket, 40);
519 assert_eq!(tb.bucket_max, 100);
520
521 tb.adjust(
522 time,
523 &TokenBucketConfig {
524 rate: 200,
525 bucket_max: 100,
526 },
527 );
528 assert_eq!(tb.bucket, 40);
529 assert_eq!(tb.bucket_max, 100);
530 assert_eq!(tb.rate, 200);
531 }
532
533 #[test]
534 fn adjust_future() {
535 let config = TokenBucketConfig {
536 rate: 10,
537 bucket_max: 100,
538 };
539 let mut tb = TokenBucket::new(&config, MillisTimestamp(100));
540 assert_eq!(tb.bucket, 100);
541 assert_eq!(tb.bucket_max, 100);
542 assert_eq!(tb.rate, 10);
543
544 // at 300 ms: increase rate and max; bucket was already full, so doesn't gain any tokens
545 tb.adjust(
546 MillisTimestamp(300),
547 &TokenBucketConfig {
548 rate: 20,
549 bucket_max: 200,
550 },
551 );
552 assert_eq!(tb.bucket, 100);
553 assert_eq!(tb.bucket_max, 200);
554
555 // at 500 ms: no changes; bucket is refilled during `adjust()`, so gains 4 tokens
556 tb.adjust(
557 MillisTimestamp(500),
558 &TokenBucketConfig {
559 rate: 20,
560 bucket_max: 200,
561 },
562 );
563 assert_eq!(tb.bucket, 104);
564 assert_eq!(tb.bucket_max, 200);
565
566 // at 700 ms: lower rate and max; bucket is lowered to new max, so loses 4 tokens
567 tb.adjust(
568 MillisTimestamp(700),
569 &TokenBucketConfig {
570 rate: 0,
571 bucket_max: 100,
572 },
573 );
574 assert_eq!(tb.bucket, 100);
575 assert_eq!(tb.bucket_max, 100);
576
577 // at 900 ms: raise rate and max; rate was previously 0 so doesn't gain any tokens
578 tb.adjust(
579 MillisTimestamp(900),
580 &TokenBucketConfig {
581 rate: 100,
582 bucket_max: 200,
583 },
584 );
585 assert_eq!(tb.bucket, 100);
586 assert_eq!(tb.bucket_max, 200);
587 }
588
589 #[test]
590 fn adjust_zero() {
591 let time = MillisTimestamp(100);
592
593 let config = TokenBucketConfig {
594 rate: 10,
595 bucket_max: 100,
596 };
597
598 let mut tb = TokenBucket::new(&config, time);
599 tb.adjust(
600 time,
601 &TokenBucketConfig {
602 rate: 0,
603 bucket_max: 200,
604 },
605 );
606 assert_eq!(tb.bucket, 100);
607 assert_eq!(tb.bucket_max, 200);
608 assert_eq!(tb.rate, 0);
609 // bucket should not increase
610 tb.refill(MillisTimestamp(10_000_000));
611 assert_eq!(tb.bucket, 100);
612
613 let mut tb = TokenBucket::new(&config, time);
614 tb.adjust(
615 time,
616 &TokenBucketConfig {
617 rate: 10,
618 bucket_max: 0,
619 },
620 );
621 assert_eq!(tb.bucket, 0);
622 assert_eq!(tb.bucket_max, 0);
623 assert_eq!(tb.rate, 10);
624 // bucket should stay empty
625 tb.refill(MillisTimestamp(10_000_000));
626 assert_eq!(tb.bucket, 0);
627
628 let mut tb = TokenBucket::new(&config, time);
629 tb.adjust(
630 time,
631 &TokenBucketConfig {
632 rate: 0,
633 bucket_max: 0,
634 },
635 );
636 assert_eq!(tb.bucket, 0);
637 assert_eq!(tb.bucket_max, 0);
638 assert_eq!(tb.rate, 0);
639 // bucket should stay empty
640 tb.refill(MillisTimestamp(10_000_000));
641 assert_eq!(tb.bucket, 0);
642 }
643
644 #[test]
645 fn is_empty() {
646 // increases 10 tokens/second (one every 100 ms)
647 let config = TokenBucketConfig {
648 rate: 10,
649 bucket_max: 100,
650 };
651 let mut tb = TokenBucket::new(&config, MillisTimestamp(100));
652 assert!(!tb.is_empty());
653
654 tb.drain(99).unwrap();
655 assert!(!tb.is_empty());
656
657 tb.drain(1).unwrap();
658 assert!(tb.is_empty());
659
660 tb.refill(MillisTimestamp(199));
661 assert!(tb.is_empty());
662
663 tb.refill(MillisTimestamp(200));
664 assert!(!tb.is_empty());
665 }
666
667 #[test]
668 fn correctness() {
669 // increases 10 tokens/second (one every 100 ms)
670 let config = TokenBucketConfig {
671 rate: 10,
672 bucket_max: 100,
673 };
674 let mut tb = TokenBucket::new(&config, MillisTimestamp(100));
675
676 tb.drain(50).unwrap();
677 assert_eq!(tb.bucket, 50);
678
679 tb.refill(MillisTimestamp(1100));
680 assert_eq!(tb.bucket, 60);
681
682 tb.drain(50).unwrap();
683 assert_eq!(tb.bucket, 10);
684
685 tb.refill(MillisTimestamp(2100));
686 assert_eq!(tb.bucket, 20);
687
688 tb.refill(MillisTimestamp(2101));
689 assert_eq!(tb.bucket, 20);
690 tb.refill(MillisTimestamp(2199));
691 assert_eq!(tb.bucket, 20);
692 tb.refill(MillisTimestamp(2200));
693 assert_eq!(tb.bucket, 21);
694 }
695
696 #[test]
697 fn rounding() {
698 // increases 10 tokens/second (one every 100 ms)
699 let config = TokenBucketConfig {
700 rate: 10,
701 bucket_max: 100,
702 };
703 let mut tb = TokenBucket::new(&config, MillisTimestamp(0));
704 tb.drain(100).unwrap();
705
706 // ensure that refilling at 150 ms does not change the `added_tokens_at` time to 150 ms,
707 // otherwise the next refill wouldn't occur until 250 ms instead of 200 ms
708 tb.refill(MillisTimestamp(99));
709 assert_eq!(tb.bucket, 0);
710 tb.refill(MillisTimestamp(150));
711 assert_eq!(tb.bucket, 1);
712 tb.refill(MillisTimestamp(199));
713 assert_eq!(tb.bucket, 1);
714 tb.refill(MillisTimestamp(200));
715 assert_eq!(tb.bucket, 2);
716 }
717
718 #[test]
719 fn tokens_available_at() {
720 // increases 10 tokens/second (one every 100 ms)
721 let config = TokenBucketConfig {
722 rate: 10,
723 bucket_max: 100,
724 };
725 let mut tb = TokenBucket::new(&config, MillisTimestamp(0));
726
727 // bucket is empty at 0 ms, next token at 100 ms
728 tb.drain(100).unwrap();
729
730 assert_eq!(tb.tokens_available_at(0), Ok(MillisTimestamp(0)));
731 assert_eq!(tb.tokens_available_at(1), Ok(MillisTimestamp(100)));
732 assert_eq!(tb.tokens_available_at(2), Ok(MillisTimestamp(200)));
733
734 // bucket is still empty at 40 ms, next token at 100 ms
735 tb.refill(MillisTimestamp(40));
736
737 assert_eq!(tb.tokens_available_at(0), Ok(MillisTimestamp(0)));
738 assert_eq!(tb.tokens_available_at(1), Ok(MillisTimestamp(100)));
739 assert_eq!(tb.tokens_available_at(2), Ok(MillisTimestamp(200)));
740
741 // bucket has 1 token at 100 ms, next token at 200 ms
742 tb.refill(MillisTimestamp(100));
743
744 assert_eq!(tb.tokens_available_at(0), Ok(MillisTimestamp(100)));
745 assert_eq!(tb.tokens_available_at(1), Ok(MillisTimestamp(100)));
746 assert_eq!(tb.tokens_available_at(2), Ok(MillisTimestamp(200)));
747
748 // bucket is empty at 100 ms, next token at 200 ms
749 tb.drain(1).unwrap();
750
751 assert_eq!(tb.tokens_available_at(0), Ok(MillisTimestamp(100)));
752 assert_eq!(tb.tokens_available_at(1), Ok(MillisTimestamp(200)));
753 assert_eq!(tb.tokens_available_at(2), Ok(MillisTimestamp(300)));
754
755 // bucket is empty at 140 ms, next token at 200 ms
756 tb.refill(MillisTimestamp(140));
757
758 assert_eq!(tb.tokens_available_at(0), Ok(MillisTimestamp(100)));
759 assert_eq!(tb.tokens_available_at(1), Ok(MillisTimestamp(200)));
760 assert_eq!(tb.tokens_available_at(2), Ok(MillisTimestamp(300)));
761
762 // bucket has 1 token at 210 ms, next token at 300 ms
763 tb.refill(MillisTimestamp(210));
764
765 assert_eq!(tb.tokens_available_at(0), Ok(MillisTimestamp(200)));
766 assert_eq!(tb.tokens_available_at(1), Ok(MillisTimestamp(200)));
767 assert_eq!(tb.tokens_available_at(2), Ok(MillisTimestamp(300)));
768
769 use NeverEnoughTokensError as NETE;
770
771 assert_eq!(tb.tokens_available_at(100), Ok(MillisTimestamp(10_100)));
772 assert_eq!(tb.tokens_available_at(101), Err(NETE::ExceedsMaxTokens));
773 assert_eq!(
774 tb.tokens_available_at(u64::MAX),
775 Err(NETE::ExceedsMaxTokens),
776 );
777
778 // set the refill rate to 0; note that adjusting the rate also resets `added_tokens_at`
779 tb.adjust(
780 MillisTimestamp(210),
781 &TokenBucketConfig {
782 rate: 0,
783 bucket_max: 100,
784 },
785 );
786
787 assert_eq!(tb.tokens_available_at(0), Ok(MillisTimestamp(210)));
788 assert_eq!(tb.tokens_available_at(1), Ok(MillisTimestamp(210)));
789 assert_eq!(tb.tokens_available_at(2), Err(NETE::ZeroRate));
790 }
791
792 #[test]
793 fn test_duration_token_round_trip() {
794 let tokens_to_duration = TokenBucket::<Instant>::tokens_to_duration;
795 let duration_to_tokens = TokenBucket::<Instant>::duration_to_tokens;
796
797 // start with some hand-picked cases
798 let mut duration_rate_pairs = vec![
799 (Duration::from_nanos(0), 1),
800 (Duration::from_nanos(1), 1),
801 (Duration::from_micros(2), 1),
802 (Duration::MAX, 1),
803 (Duration::from_nanos(0), 3),
804 (Duration::from_nanos(1), 3),
805 (Duration::from_micros(2), 3),
806 (Duration::MAX, 3),
807 (Duration::from_nanos(0), 1000),
808 (Duration::from_nanos(1), 1000),
809 (Duration::from_micros(2), 1000),
810 (Duration::MAX, 1000),
811 (Duration::from_nanos(0), u64::MAX),
812 (Duration::from_nanos(1), u64::MAX),
813 (Duration::from_micros(2), u64::MAX),
814 (Duration::MAX, u64::MAX),
815 ];
816
817 let mut rng = rand::rng();
818
819 // add some fuzzing
820 for _ in 0..10_000 {
821 let secs = rng.random();
822 let nanos = rng.random();
823 // Duration::new() may panic, so just skip if there's a panic rather than trying to
824 // write our own logic to avoid the panic in the first place
825 let Ok(random_duration) = std::panic::catch_unwind(|| Duration::new(secs, nanos))
826 else {
827 continue;
828 };
829 let random_rate = rng.random();
830 duration_rate_pairs.push((random_duration, random_rate));
831 }
832
833 // for various combinations of durations and rates, we ensure that after an initial
834 // `duration_to_tokens` calculation which may truncate, a round-trip between
835 // `tokens_to_duration` and `duration_to_tokens` isn't lossy
836 for (original_duration, rate) in duration_rate_pairs {
837 // this may give a smaller number of tokens than expected (see docs on
838 // `TokenBucket::duration_to_tokens`)
839 let tokens = duration_to_tokens(original_duration, rate);
840
841 // we want to ensure that converting these `tokens` to a duration and then back to
842 // tokens is not lossy, which implies that `tokens_to_duration` is returning the
843 // expected value and not a truncated value due to saturating arithmetic
844 let duration = tokens_to_duration(tokens, rate).unwrap();
845 assert_eq!(tokens, duration_to_tokens(duration, rate));
846 }
847 }
848}