From e257f061a41ee45b25398eba2cead559b0c255ac Mon Sep 17 00:00:00 2001 From: colin99d Date: Tue, 11 Apr 2023 11:32:33 -0400 Subject: [PATCH] Sending the request --- slowapi/extension.py | 4 +--- slowapi/wrappers.py | 12 +++++++++--- 2 files changed, 10 insertions(+), 6 deletions(-) diff --git a/slowapi/extension.py b/slowapi/extension.py index 811d577..d31abe3 100644 --- a/slowapi/extension.py +++ b/slowapi/extension.py @@ -482,7 +482,7 @@ class Limiter: limit_for_header = None for lim in limits: limit_scope = lim.scope or endpoint - if lim.is_exempt: + if lim.is_exempt(request): continue if lim.methods is not None and request.method.lower() not in lim.methods: continue @@ -699,11 +699,9 @@ class Limiter: else: self._route_limits.setdefault(name, []).extend(static_limits) - connection_type: Optional[str] = None sig = inspect.signature(func) for idx, parameter in enumerate(sig.parameters.values()): if parameter.name == "request" or parameter.name == "websocket": - connection_type = parameter.name break else: raise Exception( diff --git a/slowapi/wrappers.py b/slowapi/wrappers.py index a1741c5..54b0a31 100644 --- a/slowapi/wrappers.py +++ b/slowapi/wrappers.py @@ -2,6 +2,7 @@ import inspect from typing import Callable, Iterator, List, Optional, Union from limits import RateLimitItem, parse_many # type: ignore +from starlette.requests import Request class Limit(object): @@ -31,13 +32,18 @@ class Limit(object): self.cost = cost self.override_defaults = override_defaults - @property - def is_exempt(self) -> bool: + def is_exempt(self, request: Request) -> bool: """ Check if the limit is exempt. Return True to exempt the route from the limit. """ - return self.exempt_when() if self.exempt_when is not None else False + if self.exempt_when is None: + return False + params = inspect.signature(self.exempt_when).parameters + param_len = len(params) + if param_len == 1: + return self.exempt_when(request) + return self.exempt_when() @property def scope(self) -> str: