Skip to main content

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}