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
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
1from __future__ import annotations
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
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
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
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.
49 ``PyJWKClient`` uses a two-tier caching system to avoid unnecessary
50 network requests:
52 **Tier 1 — JWK Set cache** (enabled by default):
53 Caches the entire JSON Web Key Set response from the endpoint.
54 Controlled by:
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.
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``.
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:
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``.
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()
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
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]
147 def fetch_data(self) -> Any:
148 """Fetch the JWK Set from the JWKS endpoint.
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.
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
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)
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
190 def get_jwk_set(self, refresh: bool = False) -> PyJWKSet:
191 """Return the JWK Set, using the cache when available.
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()
206 if data is None:
207 data = self.fetch_data()
208 fetched = True
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))
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)
225 return jwk_set
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")
232 return data
234 def get_signing_keys(self, refresh: bool = False) -> list[PyJWK]:
235 """Return all signing keys from the JWK Set.
237 Filters the JWK Set to keys whose ``use`` is ``"sig"`` (or
238 unspecified) and that have a ``kid``.
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)
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 ]
258 if not signing_keys:
259 raise PyJWKClientError("The JWKS endpoint did not contain any signing keys")
261 return signing_keys
263 def get_signing_key(self, kid: str) -> PyJWK:
264 """Return the signing key matching the given ``kid``.
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.
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)
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)
293 if not signing_key:
294 raise PyJWKClientError(
295 f'Unable to find a signing key that matches: "{kid}"'
296 )
298 return signing_key
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.
303 Extracts the ``kid`` from the token's unverified header and
304 delegates to :meth:`get_signing_key`.
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"))
315 @staticmethod
316 def match_kid(signing_keys: list[PyJWK], kid: str) -> PyJWK | None:
317 """Find a key in *signing_keys* that matches *kid*.
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
328 for key in signing_keys:
329 if key.key_id == kid:
330 signing_key = key
331 break
333 return signing_key