Coverage for /pythoncovmergedfiles/medio/medio/usr/local/lib/python3.11/site-packages/jwt/jwks_client.py: 26%

Shortcuts on this page

r m x   toggle line displays

j k   next/prev highlighted chunk

0   (zero) top of page

1   (one) first highlighted chunk

117 statements  

1from __future__ import annotations 

2 

3import http.client 

4import json 

5import math 

6import threading 

7import time 

8import urllib.request 

9from functools import lru_cache 

10from ssl import SSLContext 

11from typing import Any 

12from urllib.error import HTTPError, URLError 

13from urllib.parse import urlparse 

14 

15from .api_jwk import PyJWK, PyJWKSet 

16from .api_jwt import decode_complete as decode_token 

17from .exceptions import PyJWKClientConnectionError, PyJWKClientError 

18from .jwk_set_cache import JWKSetCache 

19 

20 

21class _NoRedirectHandler(urllib.request.HTTPRedirectHandler): 

22 def redirect_request( 

23 self, 

24 req: urllib.request.Request, 

25 fp: Any, 

26 code: int, 

27 msg: str, 

28 headers: Any, 

29 newurl: str, 

30 ) -> urllib.request.Request | None: 

31 return None 

32 

33 

34class PyJWKClient: 

35 def __init__( 

36 self, 

37 uri: str, 

38 cache_keys: bool = False, 

39 max_cached_keys: int = 16, 

40 cache_jwk_set: bool = True, 

41 lifespan: float = 300, 

42 headers: dict[str, Any] | None = None, 

43 timeout: float = 30, 

44 ssl_context: SSLContext | None = None, 

45 cooldown_duration: float = 30, 

46 ): 

47 """A client for retrieving signing keys from a JWKS endpoint. 

48 

49 ``PyJWKClient`` uses a two-tier caching system to avoid unnecessary 

50 network requests: 

51 

52 **Tier 1 — JWK Set cache** (enabled by default): 

53 Caches the entire JSON Web Key Set response from the endpoint. 

54 Controlled by: 

55 

56 - ``cache_jwk_set``: Set to ``True`` (the default) to enable this 

57 cache. When enabled, the JWK Set is fetched from the network only 

58 when the cache is empty or expired. 

59 - ``lifespan``: Time in seconds before the cached JWK Set expires. 

60 Defaults to ``300`` (5 minutes). Must be greater than 0. 

61 

62 Unknown key IDs can trigger one forced refresh after the cooldown 

63 period configured by ``cooldown_duration``. Every successful fetch 

64 starts this cooldown, including the initial fetch and cache-expiry 

65 fetches. A newly rotated key may therefore wait for the cooldown 

66 period before it is fetched; set ``cooldown_duration`` to ``0`` to 

67 disable this behavior. The cooldown is bypassed when 

68 ``cache_jwk_set`` is ``False``. 

69 

70 **Tier 2 — Signing key cache** (disabled by default): 

71 Caches individual signing keys (looked up by ``kid``) using an LRU 

72 cache with **no time-based expiration**. Keys are evicted only when 

73 the cache reaches its maximum size. Controlled by: 

74 

75 - ``cache_keys``: Set to ``True`` to enable this cache. 

76 Defaults to ``False``. 

77 - ``max_cached_keys``: Maximum number of signing keys to keep in 

78 the LRU cache. Defaults to ``16``. 

79 

80 :param uri: The URL of the JWKS endpoint. 

81 :type uri: str 

82 :param cache_keys: Enable the per-key LRU cache (Tier 2). 

83 :type cache_keys: bool 

84 :param max_cached_keys: Max entries in the signing key LRU cache. 

85 :type max_cached_keys: int 

86 :param cache_jwk_set: Enable the JWK Set response cache (Tier 1). 

87 :type cache_jwk_set: bool 

88 :param lifespan: TTL in seconds for the JWK Set cache. 

89 :type lifespan: float 

90 :param headers: Optional HTTP headers to include in requests. 

91 :type headers: dict or None 

92 :param timeout: HTTP request timeout in seconds. 

93 :type timeout: float 

94 :param ssl_context: Optional SSL context for the request. 

95 :type ssl_context: ssl.SSLContext or None 

96 :param cooldown_duration: Minimum time in seconds between forced 

97 refreshes after an unknown key ID. Defaults to ``30``. 

98 :type cooldown_duration: float 

99 """ 

100 if headers is None: 

101 headers = {} 

102 # urllib's default OpenerDirector also handles file://, ftp://, and 

103 # data: URIs. Reject anything that isn't http(s) eagerly so a caller 

104 # passing an attacker-influenced URL (e.g. taken from a `jku` token 

105 # header) can't read local files or reach other unintended schemes. 

106 scheme = urlparse(uri).scheme.lower() 

107 if scheme not in ("http", "https"): 

108 raise PyJWKClientError( 

109 f"Invalid JWKS URI scheme {scheme!r}: only 'http' and 'https' " 

110 f"are supported." 

111 ) 

112 self.uri = uri 

113 self.jwk_set_cache: JWKSetCache | None = None 

114 self.headers = headers 

115 self.timeout = timeout 

116 self.ssl_context = ssl_context 

117 if cooldown_duration < 0: 

118 raise PyJWKClientError( 

119 "Cooldown duration must be greater than or equal to 0, " 

120 f'the input is "{cooldown_duration}"' 

121 ) 

122 if not math.isfinite(cooldown_duration): 

123 raise PyJWKClientError( 

124 f'Cooldown duration must be finite, the input is "{cooldown_duration}"' 

125 ) 

126 self.cooldown_duration = cooldown_duration 

127 self._last_successful_fetch: float | None = None 

128 self._client_lock = threading.RLock() 

129 

130 if cache_jwk_set: 

131 # Init jwt set cache with default or given lifespan. 

132 # Default lifespan is 300 seconds (5 minutes). 

133 if lifespan <= 0: 

134 raise PyJWKClientError( 

135 f'Lifespan must be greater than 0, the input is "{lifespan}"' 

136 ) 

137 self.jwk_set_cache = JWKSetCache(lifespan) 

138 else: 

139 self.jwk_set_cache = None 

140 

141 if cache_keys: 

142 # Cache signing keys 

143 get_signing_key = lru_cache(maxsize=max_cached_keys)(self.get_signing_key) 

144 # Ignore mypy (https://github.com/python/mypy/issues/2427) 

145 self.get_signing_key = get_signing_key # type: ignore[method-assign] 

146 

147 def fetch_data(self) -> Any: 

148 """Fetch the JWK Set from the JWKS endpoint. 

149 

150 Makes an HTTP request to the configured ``uri`` and returns the 

151 parsed JSON response. If the JWK Set cache is enabled, the 

152 response is stored in the cache. 

153 

154 :returns: The parsed JWK Set as a dictionary. 

155 :raises PyJWKClientConnectionError: If the HTTP request fails. 

156 :raises PyJWKClientError: If the endpoint does not return a JSON 

157 object. 

158 :raises PyJWKSetError: If the JWK Set cache is enabled and the 

159 response contains no usable keys. 

160 """ 

161 try: 

162 r = urllib.request.Request(url=self.uri, headers=self.headers) 

163 handlers: list[Any] = [_NoRedirectHandler()] 

164 if self.ssl_context is not None: 

165 handlers.append(urllib.request.HTTPSHandler(context=self.ssl_context)) 

166 opener = urllib.request.build_opener(*handlers) 

167 with opener.open(r, timeout=self.timeout) as response: 

168 jwk_set = json.load(response) 

169 except (URLError, TimeoutError, http.client.HTTPException) as e: 

170 if isinstance(e, HTTPError): 

171 e.close() 

172 raise PyJWKClientConnectionError( 

173 f'Fail to fetch data from the url, err: "{e}"' 

174 ) from e 

175 

176 # Validate the payload before it reaches the cache, so an endpoint 

177 # returning something that isn't a JSON object is reported as such 

178 # rather than as a malformed JWK Set. 

179 jwk_set = self._as_jwk_set_payload(jwk_set) 

180 

181 # Only update the cache on a successful fetch. Writing in a 

182 # `finally` block with `jwk_set=None` on error clears any 

183 # previously-cached JWKS, turning a transient outage into a cache 

184 # wipe that breaks legitimate auth. 

185 if self.jwk_set_cache is not None: 

186 self.jwk_set_cache.put(jwk_set) 

187 self._last_successful_fetch = time.monotonic() 

188 return jwk_set 

189 

190 def get_jwk_set(self, refresh: bool = False) -> PyJWKSet: 

191 """Return the JWK Set, using the cache when available. 

192 

193 :param refresh: Force a fresh fetch from the endpoint, bypassing 

194 the cache. 

195 :type refresh: bool 

196 :returns: The JWK Set. 

197 :rtype: PyJWKSet 

198 :raises PyJWKClientError: If the endpoint does not return a JSON 

199 object. 

200 """ 

201 data = None 

202 fetched = False 

203 if self.jwk_set_cache is not None and not refresh: 

204 data = self.jwk_set_cache.get() 

205 

206 if data is None: 

207 data = self.fetch_data() 

208 fetched = True 

209 

210 # A cache hit is already parsed, so serve it as-is. Only a fresh 

211 # fetch reaches the check below, which still matters because 

212 # `fetch_data()` may be overridden by a subclass. 

213 if isinstance(data, PyJWKSet): 

214 jwk_set = data 

215 else: 

216 jwk_set = PyJWKSet.from_dict(self._as_jwk_set_payload(data)) 

217 

218 # `fetch_data()` caches the payload it received from the endpoint, 

219 # but a subclass may filter or transform it before returning. Cache 

220 # the value actually being returned, so later cache hits serve the 

221 # same key set as this call rather than the pre-transform one. 

222 if fetched and self.jwk_set_cache is not None: 

223 self.jwk_set_cache.put(jwk_set) 

224 

225 return jwk_set 

226 

227 @staticmethod 

228 def _as_jwk_set_payload(data: Any) -> dict[str, Any]: 

229 if not isinstance(data, dict): 

230 raise PyJWKClientError("The JWKS endpoint did not return a JSON object") 

231 

232 return data 

233 

234 def get_signing_keys(self, refresh: bool = False) -> list[PyJWK]: 

235 """Return all signing keys from the JWK Set. 

236 

237 Filters the JWK Set to keys whose ``use`` is ``"sig"`` (or 

238 unspecified) and that have a ``kid``. 

239 

240 :param refresh: Force a fresh fetch from the endpoint, bypassing 

241 the cache. 

242 :type refresh: bool 

243 :returns: A list of signing keys. 

244 :rtype: list[PyJWK] 

245 :raises PyJWKClientError: If no signing keys are found. 

246 """ 

247 jwk_set = self.get_jwk_set(refresh) 

248 return self._get_signing_keys_from_jwk_set(jwk_set) 

249 

250 @staticmethod 

251 def _get_signing_keys_from_jwk_set(jwk_set: PyJWKSet) -> list[PyJWK]: 

252 signing_keys = [ 

253 jwk_set_key 

254 for jwk_set_key in jwk_set.keys 

255 if jwk_set_key.public_key_use in ["sig", None] and jwk_set_key.key_id 

256 ] 

257 

258 if not signing_keys: 

259 raise PyJWKClientError("The JWKS endpoint did not contain any signing keys") 

260 

261 return signing_keys 

262 

263 def get_signing_key(self, kid: str) -> PyJWK: 

264 """Return the signing key matching the given ``kid``. 

265 

266 If no match is found in the current JWK Set, the set is refreshed 

267 from the endpoint and the lookup is retried once when the refresh 

268 cooldown permits it. 

269 

270 :param kid: The key ID to look up. 

271 :type kid: str 

272 :returns: The matching signing key. 

273 :rtype: PyJWK 

274 :raises PyJWKClientError: If no matching key is found after 

275 refreshing. 

276 """ 

277 with self._client_lock: 

278 signing_keys = self.get_signing_keys() 

279 signing_key = self.match_kid(signing_keys, kid) 

280 

281 if not signing_key: 

282 cooling_down = ( 

283 self.jwk_set_cache is not None 

284 and self._last_successful_fetch is not None 

285 and time.monotonic() - self._last_successful_fetch 

286 < self.cooldown_duration 

287 ) 

288 if not cooling_down: 

289 signing_keys = self.get_signing_keys(refresh=True) 

290 self._last_successful_fetch = time.monotonic() 

291 signing_key = self.match_kid(signing_keys, kid) 

292 

293 if not signing_key: 

294 raise PyJWKClientError( 

295 f'Unable to find a signing key that matches: "{kid}"' 

296 ) 

297 

298 return signing_key 

299 

300 def get_signing_key_from_jwt(self, token: str | bytes) -> PyJWK: 

301 """Return the signing key for a JWT by reading its ``kid`` header. 

302 

303 Extracts the ``kid`` from the token's unverified header and 

304 delegates to :meth:`get_signing_key`. 

305 

306 :param token: The encoded JWT. 

307 :type token: str or bytes 

308 :returns: The matching signing key. 

309 :rtype: PyJWK 

310 """ 

311 unverified = decode_token(token, options={"verify_signature": False}) 

312 header = unverified["header"] 

313 return self.get_signing_key(header.get("kid")) 

314 

315 @staticmethod 

316 def match_kid(signing_keys: list[PyJWK], kid: str) -> PyJWK | None: 

317 """Find a key in *signing_keys* that matches *kid*. 

318 

319 :param signing_keys: The list of keys to search. 

320 :type signing_keys: list[PyJWK] 

321 :param kid: The key ID to match. 

322 :type kid: str 

323 :returns: The matching key, or ``None`` if not found. 

324 :rtype: PyJWK or None 

325 """ 

326 signing_key = None 

327 

328 for key in signing_keys: 

329 if key.key_id == kid: 

330 signing_key = key 

331 break 

332 

333 return signing_key