A token bucket for rate limiting. Args: query_per_second (float): The rate of the token bucket.
| 424 | |
| 425 | |
| 426 | class TokenBucket: |
| 427 | """A token bucket for rate limiting. |
| 428 | |
| 429 | Args: |
| 430 | query_per_second (float): The rate of the token bucket. |
| 431 | """ |
| 432 | |
| 433 | def __init__(self, rate, verbose=False): |
| 434 | self._rate = rate |
| 435 | self._tokens = threading.Semaphore(0) |
| 436 | self.started = False |
| 437 | self._request_queue = Queue() |
| 438 | self.logger = get_logger() |
| 439 | self.verbose = verbose |
| 440 | |
| 441 | def _add_tokens(self): |
| 442 | """Add tokens to the bucket.""" |
| 443 | while True: |
| 444 | if self._tokens._value < self._rate: |
| 445 | self._tokens.release() |
| 446 | sleep(1 / self._rate) |
| 447 | |
| 448 | def get_token(self): |
| 449 | """Get a token from the bucket.""" |
| 450 | if not self.started: |
| 451 | self.started = True |
| 452 | threading.Thread(target=self._add_tokens, daemon=True).start() |
| 453 | self._tokens.acquire() |
| 454 | if self.verbose: |
| 455 | cur_time = time.time() |
| 456 | while not self._request_queue.empty(): |
| 457 | if cur_time - self._request_queue.queue[0] > 60: |
| 458 | self._request_queue.get() |
| 459 | else: |
| 460 | break |
| 461 | self._request_queue.put(cur_time) |
| 462 | self.logger.info(f'Current RPM {self._request_queue.qsize()}.') |