Coverage for /home/airflow/.local/lib/python3.12/site-packages/airflow/api_fastapi/auth/tokens.py: 47%
305 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 14:22 +0000
1# Licensed to the Apache Software Foundation (ASF) under one
2# or more contributor license agreements. See the NOTICE file
3# distributed with this work for additional information
4# regarding copyright ownership. The ASF licenses this file
5# to you under the Apache License, Version 2.0 (the
6# "License"); you may not use this file except in compliance
7# with the License. You may obtain a copy of the License at
8#
9# http://www.apache.org/licenses/LICENSE-2.0
10#
11# Unless required by applicable law or agreed to in writing,
12# software distributed under the License is distributed on an
13# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14# KIND, either express or implied. See the License for the
15# specific language governing permissions and limitations
16# under the License.
17from __future__ import annotations
19import json
20import os
21import time
22import uuid
23from base64 import urlsafe_b64encode
24from collections.abc import Callable, Sequence
25from datetime import datetime
26from typing import TYPE_CHECKING, Any, Literal, overload
28import attrs
29import httpx
30import jwt
31import structlog
32from asgiref.sync import async_to_sync
33from cryptography.hazmat.backends import default_backend
34from cryptography.hazmat.primitives import hashes
35from cryptography.hazmat.primitives.serialization import load_pem_private_key
37from airflow._shared.timezones import timezone
38from airflow.models.revoked_token import RevokedToken
40if TYPE_CHECKING: 40 ↛ 41line 40 didn't jump to line 41 because the condition on line 40 was never true
41 from jwt.algorithms import AllowedKeys, AllowedPrivateKeys
43log = structlog.get_logger(logger_name=__name__)
45__all__ = [
46 "InvalidClaimError",
47 "JWKS",
48 "JWTGenerator",
49 "JWTValidator",
50 "generate_private_key",
51 "get_sig_validation_args",
52 "get_signing_args",
53 "get_signing_key",
54 "key_to_pem",
55 "key_to_jwk_dict",
56]
59class InvalidClaimError(ValueError):
60 """Raised when a claim in the JWT is invalid."""
62 def __init__(self, claim: str):
63 super().__init__(f"Invalid claim: {claim}")
66def key_to_jwk_dict(key: AllowedKeys, kid: str | None = None):
67 """Convert a public or private key into a valid JWKS dict."""
68 from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey, Ed25519PublicKey
69 from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey, RSAPublicKey
70 from jwt.algorithms import OKPAlgorithm, RSAAlgorithm
72 if isinstance(key, (RSAPrivateKey, Ed25519PrivateKey)):
73 key = key.public_key()
75 if isinstance(key, RSAPublicKey):
76 jwk_dict = RSAAlgorithm(RSAAlgorithm.SHA256).to_jwk(key, as_dict=True)
78 elif isinstance(key, Ed25519PublicKey):
79 jwk_dict = OKPAlgorithm().to_jwk(key, as_dict=True)
80 else:
81 raise ValueError(f"Unknown key object {type(key)}")
83 if not kid:
84 kid = thumbprint(jwk_dict)
86 jwk_dict["kid"] = kid
88 return jwk_dict
91def _guess_best_algorithm(key: AllowedPrivateKeys):
92 from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
93 from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey
95 if isinstance(key, RSAPrivateKey):
96 return "RS256"
97 if isinstance(key, Ed25519PrivateKey):
98 return "EdDSA"
99 raise ValueError(f"Unknown key object {type(key)}")
102@attrs.define(repr=False)
103class JWKS:
104 """A class to fetch and sync a set of JSON Web Keys."""
106 url: str
107 fetched_at: float = 0
108 last_fetch_attempt_at: float = 0
110 client: httpx.AsyncClient = attrs.field(factory=httpx.AsyncClient)
112 _jwks: jwt.PyJWKSet | None = None
113 refresh_jwks: bool = True
114 refresh_interval_secs: int = 3600
115 refresh_retry_interval_secs: int = 10
117 def __repr__(self) -> str:
118 return f"JWKS(url={self.url}, fetched_at={self.fetched_at})"
120 @classmethod
121 def from_private_key(cls, *keys: AllowedPrivateKeys | tuple[AllowedPrivateKeys, str]):
122 obj = cls(url=os.devnull)
123 keyset = [
124 # Each `key` is either the key directly or `(key, "my-kid")`
125 key_to_jwk_dict(*key) if isinstance(key, tuple) else key_to_jwk_dict(key)
126 for key in keys
127 ]
128 obj._jwks = jwt.PyJWKSet(keyset)
129 return obj
131 async def fetch_jwks(self) -> None:
132 if not self._should_fetch_jwks():
133 return
134 if self.url.startswith("http"):
135 data = await self._fetch_remote_jwks()
136 else:
137 data = self._fetch_local_jwks()
139 if not data:
140 return
142 self._jwks = jwt.PyJWKSet.from_dict(data)
143 log.debug("Fetched JWKS", url=self.url, keys=len(self._jwks.keys))
145 async def _fetch_remote_jwks(self) -> dict[str, Any] | None:
146 try:
147 log.debug(
148 "Fetching JWKS",
149 url=self.url,
150 last_fetched_secs_ago=int(time.monotonic() - self.fetched_at) if self.fetched_at else None,
151 )
152 if TYPE_CHECKING:
153 assert self.url
154 self.last_fetch_attempt_at = int(time.monotonic())
155 response = await self.client.get(self.url)
156 response.raise_for_status()
157 self.fetched_at = int(time.monotonic())
158 await response.aread()
159 await response.aclose()
160 return response.json()
161 except Exception:
162 log.exception("Failed to fetch remote JWKS", url=self.url)
163 return None
165 def _fetch_local_jwks(self) -> dict[str, Any] | None:
166 try:
167 with open(self.url) as jwks_file:
168 content = json.load(jwks_file)
169 self.fetched_at = int(time.monotonic())
170 return content
171 except Exception:
172 log.exception("Failed to read local JWKS", url=self.url)
173 return None
175 def _should_fetch_jwks(self) -> bool:
176 """
177 Check if we need to fetch the JWKS based on the last fetch time and the refresh interval.
179 If the JWKS URL is local, we only fetch it once. For remote JWKS URLs we fetch it based
180 on the refresh interval if refreshing has been enabled with a minimum interval between
181 attempts. The fetcher functions set the fetched_at timestamp to the current monotonic time
182 when the JWKS is fetched.
183 """
184 if not self.url.startswith("http"):
185 # Fetch local JWKS only if not already loaded
186 # This could be improved in future by looking at mtime of file.
187 return not self._jwks
188 # For remote fetches we check if the JWKS is not loaded (fetched_at = 0) or if the last fetch was more than
189 # refresh_interval_secs ago and the last fetch attempt was more than refresh_retry_interval_secs ago
190 now = time.monotonic()
191 return self.refresh_jwks and (
192 not self._jwks
193 or (
194 self.fetched_at == 0
195 or (
196 now - self.fetched_at > self.refresh_interval_secs
197 and now - self.last_fetch_attempt_at > self.refresh_retry_interval_secs
198 )
199 )
200 )
202 async def get_key(self, kid: str) -> jwt.PyJWK:
203 """Fetch the JWKS and find the matching key for the token."""
204 await self.fetch_jwks()
206 if self._jwks:
207 return self._jwks[kid]
209 # It didn't load!
210 raise KeyError(f"Key ID {kid} not found in keyset")
212 def status(self):
213 # https://svcs.hynek.me/en/stable/core-concepts.html#health-checks
214 if not self._should_fetch_jwks():
215 # Up-to-date, we are healthy
216 return
218 if self.fetched_at == 0:
219 raise RuntimeError("JWKS never fetched")
221 last_successful_fetch = time.monotonic() - self.fetched_at
222 if last_successful_fetch > 3 * self.refresh_interval_secs:
223 raise RuntimeError(f"JWKS last fetched {last_successful_fetch}s ago")
226def _conf_factory(section, key, **kwargs):
227 def factory() -> str:
228 from airflow.configuration import conf
230 return conf.get(section, key, **kwargs, suppress_warnings=True)
232 return factory
235@overload
236def _conf_list_factory(section, key, first_only: Literal[True], **kwargs) -> Callable[[], str]: ... 236 ↛ exitline 236 didn't return from function '_conf_list_factory' because
239@overload
240def _conf_list_factory( 240 ↛ exitline 240 didn't return from function '_conf_list_factory' because
241 section, key, first_only: Literal[False] = False, **kwargs
242) -> Callable[[], list[str]]: ...
245def _conf_list_factory(section, key, first_only: bool = False, **kwargs):
246 def factory() -> list[str] | str | None:
247 from airflow.configuration import conf
249 val = conf.getlist(section, key, **kwargs, suppress_warnings=True)
251 if first_only:
252 return val[0] if val else None
253 return val or []
255 return factory
258def _to_list(val: str | list[str]) -> list[str]:
259 if isinstance(val, str): 259 ↛ 260line 259 didn't jump to line 260 because the condition on line 259 was never true
260 val = [val]
261 return val
264@attrs.define(kw_only=True)
265class JWTValidator:
266 """
267 Validate the claims and validitory of a JWT.
269 This will either validate the JWT is signed with the symmetric key if ``secret_key`` is passed, or else
270 that it is signed by one of the public keys in the keyset in ``jwks`` attribute.
271 """
273 jwks: JWKS | None = None
274 secret_key: str | None = attrs.field(repr=False, default=None, converter=lambda v: None if v == "" else v)
275 issuer: str | list[str] | None = attrs.field(
276 factory=_conf_list_factory("api_auth", "jwt_issuer", fallback=None),
277 # Ensure we have None, instead of an empty list, else pyjwt will fail to validate it
278 converter=lambda v: None if v == [] else v,
279 )
280 # By default, we just validate these
281 required_claims: frozenset[str] = frozenset({"exp", "iat", "nbf"})
282 audience: str | Sequence[str]
283 algorithm: list[str] = attrs.field(
284 factory=_conf_list_factory("api_auth", "jwt_algorithm", fallback="GUESS"), converter=_to_list
285 )
287 leeway: float = attrs.field(factory=_conf_factory("api_auth", "jwt_leeway"), converter=int)
289 def __attrs_post_init__(self):
290 if not (self.jwks is None) ^ (self.secret_key is None): 290 ↛ 291line 290 didn't jump to line 291 because the condition on line 290 was never true
291 raise ValueError("Exactly one of private_key and secret_key must be specified")
293 if self.algorithm == ["GUESS"]: 293 ↛ exitline 293 didn't return from function '__attrs_post_init__' because the condition on line 293 was always true
294 if not self.jwks: 294 ↛ exitline 294 didn't return from function '__attrs_post_init__' because the condition on line 294 was always true
295 self.algorithm = ["HS512"]
297 def _get_kid_from_header(self, unvalidated: str) -> str:
298 header = jwt.get_unverified_header(unvalidated)
299 if "kid" not in header:
300 raise jwt.InvalidTokenError("Missing 'kid' in token header")
301 return header["kid"]
303 async def _get_validation_key(self, unvalidated: str) -> str | jwt.PyJWK:
304 if self.secret_key: 304 ↛ 307line 304 didn't jump to line 307 because the condition on line 304 was always true
305 return self.secret_key
307 if TYPE_CHECKING:
308 assert self.jwks is not None
310 kid = self._get_kid_from_header(unvalidated)
311 return await self.jwks.get_key(kid)
313 def validated_claims(
314 self, unvalidated: str, required_claims: dict[str, Any] | None = None
315 ) -> dict[str, Any]:
316 return async_to_sync(self.avalidated_claims)(unvalidated, required_claims)
318 async def avalidated_claims(
319 self, unvalidated: str, required_claims: dict[str, Any] | None = None
320 ) -> dict[str, Any]:
321 """Decode the JWT token, returning the validated claims or raising an exception."""
322 try:
323 key = await self._get_validation_key(unvalidated)
324 except KeyError:
325 raise jwt.InvalidTokenError("Kid did not match any validation keys")
326 algorithms = self.algorithm
327 validation_key: str | jwt.PyJWK | Any = key
328 if algorithms == ["GUESS"] and isinstance(key, jwt.PyJWK): 328 ↛ 329line 328 didn't jump to line 329 because the condition on line 328 was never true
329 if not key.algorithm_name:
330 raise jwt.InvalidTokenError("Missing algorithm in JWK")
331 algorithms = [key.algorithm_name]
332 validation_key = key.key
334 claims = jwt.decode(
335 unvalidated,
336 validation_key,
337 audience=self.audience,
338 issuer=self.issuer,
339 options={"require": list(self.required_claims)},
340 algorithms=algorithms,
341 leeway=self.leeway,
342 )
344 # Validate additional claims if provided
345 if required_claims: 345 ↛ 346line 345 didn't jump to line 346 because the condition on line 345 was never true
346 for claim, expected_value in required_claims.items():
347 if expected_value["essential"] and (
348 claim not in claims or claims[claim] != expected_value["value"]
349 ):
350 raise InvalidClaimError(claim)
352 return claims
354 def revoke_token(self, token: str) -> None:
355 """Validate the token, extract jti and exp, and revoke it in the database."""
356 try:
357 claims = self.validated_claims(token)
358 if (jti := claims.get("jti")) and (exp := claims.get("exp")):
359 RevokedToken.revoke(jti, datetime.fromtimestamp(exp, tz=timezone.utc))
360 except (jwt.InvalidTokenError, Exception):
361 log.warning("Failed to revoke token", exc_info=True)
363 def status(self):
364 if self.jwks:
365 self.jwks.status()
368def _pem_to_key(pem_data: str | bytes | AllowedPrivateKeys) -> AllowedPrivateKeys:
369 if isinstance(pem_data, str): 369 ↛ 370line 369 didn't jump to line 370 because the condition on line 369 was never true
370 pem_data = pem_data.encode()
371 elif not isinstance(pem_data, bytes): 371 ↛ 375line 371 didn't jump to line 375 because the condition on line 371 was always true
372 # Assume it's already a key object
373 return pem_data
375 return load_pem_private_key(pem_data, password=None) # type: ignore[return-value]
378def _load_key_from_configured_file() -> AllowedPrivateKeys | None:
379 from airflow.configuration import conf
381 path = conf.get("api_auth", "jwt_private_key_path", fallback=None)
382 if not path: 382 ↛ 385line 382 didn't jump to line 385 because the condition on line 382 was always true
383 return None
385 with open(path, mode="rb") as fh:
386 return _pem_to_key(fh.read())
389def _generate_kid(gen) -> str:
390 # Always check config first — both symmetric and asymmetric keys can have a configured kid
391 if kid := _conf_factory("api_auth", "jwt_kid", fallback=None)(): 391 ↛ 392line 391 didn't jump to line 392 because the condition on line 391 was never true
392 return kid
394 if not gen._private_key: 394 ↛ 398line 394 didn't jump to line 398 because the condition on line 394 was always true
395 return "not-used"
397 # Generate it from the thumbprint of the private key
398 info = key_to_jwk_dict(gen._private_key)
399 return info["kid"]
402@attrs.define(repr=False, kw_only=True)
403class JWTGenerator:
404 """Generate JWT tokens."""
406 _private_key: AllowedPrivateKeys | None = attrs.field(
407 repr=False, alias="private_key", converter=_pem_to_key, factory=_load_key_from_configured_file
408 )
409 """
410 Private key to sign generated tokens.
412 Should be either a private key object from the cryptography module, or a PEM-encoded byte string
413 """
414 _secret_key: str | None = attrs.field(
415 repr=False,
416 alias="secret_key",
417 default=None,
418 converter=lambda v: None if v == "" else v,
419 )
420 """A pre-shared secret key to sign tokens with symmetric encryption"""
422 kid: str = attrs.field(default=attrs.Factory(_generate_kid, takes_self=True))
423 valid_for: float
424 audience: str
425 issuer: str | list[str] | None = attrs.field(
426 factory=_conf_list_factory("api_auth", "jwt_issuer", first_only=True, fallback=None)
427 )
428 algorithm: str = attrs.field(
429 factory=_conf_list_factory("api_auth", "jwt_algorithm", first_only=True, fallback="GUESS")
430 )
432 def __attrs_post_init__(self):
433 if not (self._private_key is None) ^ (self._secret_key is None): 433 ↛ 434line 433 didn't jump to line 434 because the condition on line 433 was never true
434 raise ValueError("Exactly one of private_key and secret_key must be specified")
436 if self.algorithm == "GUESS": 436 ↛ exitline 436 didn't return from function '__attrs_post_init__' because the condition on line 436 was always true
437 if self._private_key: 437 ↛ 438line 437 didn't jump to line 438 because the condition on line 437 was never true
438 self.algorithm = _guess_best_algorithm(self._private_key)
439 else:
440 self.algorithm = "HS512"
442 @property
443 def signing_arg(self) -> AllowedPrivateKeys | str:
444 if callable(self._private_key): 444 ↛ 445line 444 didn't jump to line 445 because the condition on line 444 was never true
445 return self._private_key()
446 if self._private_key: 446 ↛ 447line 446 didn't jump to line 447 because the condition on line 446 was never true
447 return self._private_key
448 if TYPE_CHECKING: 448 ↛ 450line 448 didn't jump to line 450 because the condition on line 448 was never true
449 # Already handled at in post_init
450 assert self._secret_key
451 return self._secret_key
453 def generate(
454 self,
455 extras: dict[str, Any] | None = None,
456 headers: dict[str, Any] | None = None,
457 valid_for: float | None = None,
458 ) -> str:
459 """Generate a signed JWT for the subject."""
460 now = int(datetime.now(tz=timezone.utc).timestamp())
461 effective_valid_for = valid_for if valid_for is not None else self.valid_for
462 claims = {
463 "jti": uuid.uuid4().hex,
464 "iss": self.issuer,
465 "aud": self.audience,
466 "nbf": now,
467 "exp": int(now + effective_valid_for),
468 "iat": now,
469 }
471 # Remove iss and aud claims if they are falsy (None, [], "", etc.)
472 # Per RFC 7519, these are optional claims and should be omitted entirely
473 # rather than set to empty/invalid values: https://datatracker.ietf.org/doc/html/rfc7519#section-4.1.1
474 if not claims["iss"]: 474 ↛ 476line 474 didn't jump to line 476 because the condition on line 474 was always true
475 del claims["iss"]
476 if not claims["aud"]: 476 ↛ 477line 476 didn't jump to line 477 because the condition on line 476 was never true
477 del claims["aud"]
479 if extras is not None: 479 ↛ 481line 479 didn't jump to line 481 because the condition on line 479 was always true
480 claims = extras | claims
481 headers = {"alg": self.algorithm, **(headers or {})}
482 headers["kid"] = self.kid
483 return jwt.encode(claims, self.signing_arg, algorithm=self.algorithm, headers=headers)
486def generate_private_key(key_type: str = "RSA", key_size: int = 2048):
487 """
488 Generate a valid private key for testing.
490 Args:
491 key_type (str): Type of key to generate. Can be "RSA" or "Ed25516". Defaults to "RSA".
492 key_size (int): Size of the key in bits. Only applicable for RSA keys. Defaults to 2048.
494 Returns:
495 tuple: A tuple containing the private key in PEM format and the corresponding public key in PEM format.
496 """
497 from cryptography.hazmat.primitives.asymmetric import ed25519, rsa
499 if key_type == "RSA":
500 # Generate an RSA private key
502 return rsa.generate_private_key(public_exponent=65537, key_size=key_size, backend=default_backend())
503 if key_type == "Ed25519":
504 return ed25519.Ed25519PrivateKey.generate()
505 raise ValueError(f"unsupported key type: {key_type}")
508def key_to_pem(key: AllowedPrivateKeys) -> bytes:
509 from cryptography.hazmat.primitives import serialization
511 # Serialize the private key in PEM format
512 return key.private_bytes(
513 encoding=serialization.Encoding.PEM,
514 format=serialization.PrivateFormat.PKCS8,
515 encryption_algorithm=serialization.NoEncryption(),
516 )
519def thumbprint(jwk: dict[str, Any], hashalg=hashes.SHA256()) -> str:
520 """
521 Return the key thumbprint as specified by RFC 7638.
523 :param hashalg: A hash function (defaults to SHA256)
525 :return: A base64url encoded digest of the key
526 """
527 digest = hashes.Hash(hashalg, backend=default_backend())
528 jsonstr = json.dumps(jwk, separators=(",", ":"), sort_keys=True)
529 digest.update(jsonstr.encode("utf8"))
530 return base64url_encode(digest.finalize())
533def base64url_encode(payload):
534 if not isinstance(payload, bytes):
535 payload = payload.encode("utf-8")
536 encode = urlsafe_b64encode(payload)
537 return encode.decode("utf-8").rstrip("=")
540def get_signing_key(section: str, key: str, make_secret_key_if_needed: bool = True) -> str:
541 """
542 Get a signing shared key from the config.
544 If the config option is empty this will generate a random one and warn about the lack of it.
545 """
546 from airflow.configuration import conf
548 sentinel = object()
549 secret_key = conf.get(section, key, fallback=sentinel)
551 if not secret_key or secret_key is sentinel: 551 ↛ 552line 551 didn't jump to line 552 because the condition on line 551 was never true
552 if make_secret_key_if_needed:
553 log.warning(
554 "`%s/%s` was empty, using a generated one for now. Please set this in your config",
555 section,
556 key,
557 )
558 secret_key = base64url_encode(os.urandom(16))
559 # Set it back so any other callers get the same value for the duration of this process
560 conf.set(section, key, secret_key)
561 else:
562 raise ValueError(f"The value {section}/{key} must be set!")
564 # Mypy can't grock the `if not secret_key`
565 return secret_key
568def get_signing_args(make_secret_key_if_needed: bool = True) -> dict[str, Any]:
569 """
570 Return the args to splat into JWTGenerator for private or secret key.
572 Will use ``get_signing_key`` to generate a key if nothing else suitable is found.
573 """
574 # Try private key first
575 priv = _load_key_from_configured_file()
577 if priv is not None: 577 ↛ 578line 577 didn't jump to line 578 because the condition on line 577 was never true
578 return {"private_key": priv}
580 # Don't call this unless we have to as it might issue a warning
581 return {"secret_key": get_signing_key("api_auth", "jwt_secret", make_secret_key_if_needed)}
584def get_sig_validation_args(make_secret_key_if_needed: bool = True) -> dict[str, Any]:
585 from airflow.configuration import conf
587 sentinel = object()
589 # Try JWKS url first
590 url = conf.get("api_auth", "trusted_jwks_url", fallback=sentinel)
592 if url and url is not sentinel: 592 ↛ 593line 592 didn't jump to line 593 because the condition on line 592 was never true
593 jwks = JWKS(url=url)
594 return {"jwks": jwks}
596 key = _load_key_from_configured_file()
598 if key is not None: 598 ↛ 599line 598 didn't jump to line 599 because the condition on line 598 was never true
599 jwks = JWKS.from_private_key(key)
600 return {
601 "jwks": jwks,
602 "algorithm": conf.get("api_auth", "jwt_algorithm", fallback=None) or _guess_best_algorithm(key),
603 }
605 return {"secret_key": get_signing_key("api_auth", "jwt_secret", make_secret_key_if_needed)}