Coverage for api/utils/throttle.py: 76%

91 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 06:14 +0000

1import abc 

2 

3from rest_framework.throttling import SimpleRateThrottle as BaseSimpleRateThrottle 

4 

5import structlog 

6from redis.exceptions import ConnectionError 

7 

8 

9logger = structlog.get_logger(__name__) 

10 

11 

12class SimpleRateThrottle(BaseSimpleRateThrottle, metaclass=abc.ABCMeta): 

13 """ 

14 Extends the ``SimpleRateThrottle`` class to provide additional functionality such as 

15 rate-limit headers in the response. 

16 """ 

17 

18 def allow_request(self, request, view): 

19 try: 

20 is_allowed = super().allow_request(request, view) 

21 except ConnectionError: 

22 logger.warning("Redis connect failed, allowing request.") 

23 is_allowed = True 

24 view.headers |= self.headers() 

25 return is_allowed 

26 

27 def headers(self): 

28 """ 

29 Get `X-RateLimit-` headers for this particular throttle. Each pair of headers 

30 contains the limit and the number of requests left in the limit. Since multiple 

31 rate limits can apply concurrently, the suffix identifies each pair uniquely. 

32 """ 

33 prefix = "X-RateLimit" 

34 suffix = self.scope or self.__class__.__name__.lower() 

35 if hasattr(self, "history"): 

36 return { 

37 f"{prefix}-Limit-{suffix}": self.rate, 

38 f"{prefix}-Available-{suffix}": self.num_requests - len(self.history), 

39 } 

40 else: 

41 return {} 

42 

43 def has_valid_token(self, request): 

44 if not request.auth: 

45 return False 

46 

47 application = getattr(request.auth, "application", None) 

48 if application is None: 48 ↛ 49line 48 didn't jump to line 49 because the condition on line 48 was never true

49 return False 

50 

51 return application.client_id and application.verified 

52 

53 def get_cache_key(self, request, view): 

54 return self.cache_format % { 

55 "scope": self.scope, 

56 "ident": self.get_ident(request), 

57 } 

58 

59 

60class AbstractAnonRateThrottle(SimpleRateThrottle, metaclass=abc.ABCMeta): 

61 """ 

62 Limits the rate of API calls that may be made by a anonymous users. 

63 

64 The IP address of the request will be used as the unique cache key. 

65 """ 

66 

67 def get_cache_key(self, request, view): 

68 # Do not apply this throttle to requests with valid tokens 

69 if self.has_valid_token(request): 

70 return None 

71 

72 if request.headers.get("referrer") == "openverse.org": 72 ↛ 74line 72 didn't jump to line 74 because the condition on line 72 was never true

73 # Use `ov_referrer` throttles instead 

74 return None 

75 

76 return super().get_cache_key(request, view) 

77 

78 

79class AbstractOpenverseReferrerRateThrottle(SimpleRateThrottle, metaclass=abc.ABCMeta): 

80 """Use a different limit for requests that appear to come from Openverse.org.""" 

81 

82 def get_cache_key(self, request, view): 

83 # Do not apply this throttle to requests with valid tokens 

84 if self.has_valid_token(request): 

85 return None 

86 

87 if request.headers.get("referrer") != "openverse.org": 

88 # Use regular anon throttles instead 

89 return None 

90 

91 return super().get_cache_key(request, view) 

92 

93 

94class BurstRateThrottle(AbstractAnonRateThrottle): 

95 scope = "anon_burst" 

96 

97 

98class SustainedRateThrottle(AbstractAnonRateThrottle): 

99 scope = "anon_sustained" 

100 

101 

102class HealthcheckAnonRateThrottle(AbstractAnonRateThrottle): 

103 scope = "anon_healthcheck" 

104 

105 

106class AnonThumbnailRateThrottle(AbstractAnonRateThrottle): 

107 scope = "anon_thumbnail" 

108 

109 

110class OpenverseReferrerBurstRateThrottle(AbstractOpenverseReferrerRateThrottle): 

111 scope = "ov_referrer_burst" 

112 

113 

114class OpenverseReferrerSustainedRateThrottle(AbstractOpenverseReferrerRateThrottle): 

115 scope = "ov_referrer_sustained" 

116 

117 

118class OpenverseReferrerAnonThumbnailRateThrottle(AbstractOpenverseReferrerRateThrottle): 

119 scope = "ov_referrer_thumbnail" 

120 

121 

122class TenPerDay(AbstractAnonRateThrottle): 

123 rate = "1000000/second" 

124 

125 

126class OnePerSecond(AbstractAnonRateThrottle): 

127 rate = "1000000/second" 

128 

129 

130class AbstractOAuth2IdRateThrottle(SimpleRateThrottle, metaclass=abc.ABCMeta): 

131 """ 

132 Ties a particular throttling scope from ``settings.py`` to a rate limit model. 

133 

134 See ``ThrottledApplication.rate_limit_model`` for an explanation of that concept. 

135 """ 

136 

137 scope: str 

138 """The name of the scope. Used to retrieve the rate limit from settings.""" 

139 applies_to_rate_limit_model: set[str] 

140 """ 

141 The set of ``ThrottledApplication.rate_limit_model`` to which the scope applies. 

142 

143 Use a ``set`` specifically to make checks O(1). All default throttles run on 

144 almost every single request and must be performant. 

145 """ 

146 

147 def get_cache_key(self, request, view): 

148 # Find the client ID associated with the access token. 

149 if not self.has_valid_token(request): 

150 return None 

151 

152 # `self.has_valid_token` call earlier ensures accessing `application` will not fail 

153 application = request.auth.application 

154 

155 if application.rate_limit_model not in self.applies_to_rate_limit_model: 

156 return None 

157 

158 return self.cache_format % {"scope": self.scope, "ident": application.client_id} 

159 

160 

161class OAuth2IdThumbnailRateThrottle(AbstractOAuth2IdRateThrottle): 

162 applies_to_rate_limit_model = {"standard", "enhanced"} 

163 scope = "oauth2_client_credentials_thumbnail" 

164 

165 

166class OAuth2IdSustainedRateThrottle(AbstractOAuth2IdRateThrottle): 

167 applies_to_rate_limit_model = {"standard"} 

168 scope = "oauth2_client_credentials_sustained" 

169 

170 

171class OAuth2IdBurstRateThrottle(AbstractOAuth2IdRateThrottle): 

172 applies_to_rate_limit_model = {"standard"} 

173 scope = "oauth2_client_credentials_burst" 

174 

175 

176class EnhancedOAuth2IdSustainedRateThrottle(AbstractOAuth2IdRateThrottle): 

177 applies_to_rate_limit_model = {"enhanced"} 

178 scope = "enhanced_oauth2_client_credentials_sustained" 

179 

180 

181class EnhancedOAuth2IdBurstRateThrottle(AbstractOAuth2IdRateThrottle): 

182 applies_to_rate_limit_model = {"enhanced"} 

183 scope = "enhanced_oauth2_client_credentials_burst" 

184 

185 

186class ExemptOAuth2IdRateThrottle(AbstractOAuth2IdRateThrottle): 

187 applies_to_rate_limit_model = {"exempt"} 

188 scope = "exempt_oauth2_client_credentials"