From 33899c0ef1ce5f9ee13b51acda6ea83fa7a38038 Mon Sep 17 00:00:00 2001 From: Marat Sarbasov Date: Sat, 15 Jan 2022 16:02:55 +0300 Subject: [PATCH] Enable dynamic limits dependant on key. --- slowapi/extension.py | 2 +- slowapi/wrappers.py | 23 ++++++++++++++++++----- 2 files changed, 19 insertions(+), 6 deletions(-) diff --git a/slowapi/extension.py b/slowapi/extension.py index e3e33af..242b2c0 100644 --- a/slowapi/extension.py +++ b/slowapi/extension.py @@ -503,7 +503,7 @@ class Limiter: if name in self._dynamic_route_limits: for lim in self._dynamic_route_limits[name]: try: - dynamic_limits.extend(list(lim)) + dynamic_limits.extend(list(lim.with_request(request))) except ValueError as e: self.logger.error( "failed to load ratelimit for view function %s (%s)", diff --git a/slowapi/wrappers.py b/slowapi/wrappers.py index fd9b61d..6e78b41 100644 --- a/slowapi/wrappers.py +++ b/slowapi/wrappers.py @@ -1,3 +1,4 @@ +import inspect from typing import Callable, Iterator, List, Optional, Union from limits import RateLimitItem, parse_many # type: ignore @@ -74,13 +75,21 @@ class LimitGroup(object): self.error_message = error_message self.exempt_when = exempt_when self.override_defaults = override_defaults + self.request = None def __iter__(self) -> Iterator[Limit]: - limit_items: List[RateLimitItem] = parse_many( - self.__limit_provider() - if callable(self.__limit_provider) - else self.__limit_provider - ) + if callable(self.__limit_provider): + if "key" in inspect.signature(self.__limit_provider).parameters.keys(): + assert ( + "request" in inspect.signature(self.key_function).parameters.keys() + ) + assert self.request + limit_raw = self.__limit_provider(self.key_function(self.request)) + else: + limit_raw = self.__limit_provider() + else: + limit_raw = self.__limit_provider + limit_items: List[RateLimitItem] = parse_many(limit_raw) for limit in limit_items: yield Limit( limit, @@ -92,3 +101,7 @@ class LimitGroup(object): self.exempt_when, self.override_defaults, ) + + def with_request(self, request): + self.request = request + return self