feat: ж/д слой + cant_swim в скоринг + профили вне модели
Определение специфических рекомендаций (матрица профилей §8):
1. Ж/д слой (закрыт мёртвый railway ×2.5 у РАС):
- /api/v1/water/{case_id} отдаёт railway=rail как LineString
(без service/industrial/military веток), кэш общий v2;
- railway_warning «перекрыть/проверить немедленно» по профилям;
- SearchMap: Polyline слой ж/д (тёмно-красный), счётчики 💧/🚂.
2. cant_swim → профиль не_умеет_плавать (water ×3.0, без изменения
радиуса, critical_warning «обследовать водоёмы НЕМЕДЛЕННО»):
- раньше чекбокс влиял только на текст, в скоринге был пробел;
- derive в analyze._derive_profiles — работает и для closed_cases.
3. unmodeled_profiles: ДЦП/слабое зрение/слух — честная пометка
«вне поведенческой модели» с пояснением (vector_tasks B12:
профили без аналога не выдавать за учтённые); блок на фронте
в карточке здоровья.
Площадь воды: сферический эксцесс, проверен на квадрате 53° (744017 м²
vs 743272 точного). Тесты: 202 passed (новый test_cant_swim_profile).
This commit is contained in:
@@ -0,0 +1,78 @@
|
||||
from .api_jwk import PyJWK, PyJWKSet
|
||||
from .api_jws import (
|
||||
PyJWS,
|
||||
get_algorithm_by_name,
|
||||
get_unverified_header,
|
||||
register_algorithm,
|
||||
unregister_algorithm,
|
||||
)
|
||||
from .api_jwt import PyJWT, decode, decode_complete, encode
|
||||
from .exceptions import (
|
||||
DecodeError,
|
||||
ExpiredSignatureError,
|
||||
ImmatureSignatureError,
|
||||
InvalidAlgorithmError,
|
||||
InvalidAudienceError,
|
||||
InvalidIssuedAtError,
|
||||
InvalidIssuerError,
|
||||
InvalidKeyError,
|
||||
InvalidSignatureError,
|
||||
InvalidTokenError,
|
||||
MissingRequiredClaimError,
|
||||
PyJWKClientConnectionError,
|
||||
PyJWKClientError,
|
||||
PyJWKError,
|
||||
PyJWKSetError,
|
||||
PyJWTError,
|
||||
)
|
||||
from .jwks_client import PyJWKClient
|
||||
from .warnings import InsecureKeyLengthWarning
|
||||
|
||||
__version__ = "2.13.0"
|
||||
|
||||
__title__ = "PyJWT"
|
||||
__description__ = "JSON Web Token implementation in Python"
|
||||
__url__ = "https://pyjwt.readthedocs.io"
|
||||
__uri__ = __url__
|
||||
__doc__ = f"{__description__} <{__uri__}>"
|
||||
|
||||
__author__ = "José Padilla"
|
||||
__email__ = "hello@jpadilla.com"
|
||||
|
||||
__license__ = "MIT"
|
||||
__copyright__ = "Copyright 2015-2026 José Padilla"
|
||||
|
||||
|
||||
__all__ = [
|
||||
"PyJWS",
|
||||
"PyJWT",
|
||||
"PyJWKClient",
|
||||
"PyJWK",
|
||||
"PyJWKSet",
|
||||
"decode",
|
||||
"decode_complete",
|
||||
"encode",
|
||||
"get_unverified_header",
|
||||
"register_algorithm",
|
||||
"unregister_algorithm",
|
||||
"get_algorithm_by_name",
|
||||
# Warnings
|
||||
"InsecureKeyLengthWarning",
|
||||
# Exceptions
|
||||
"DecodeError",
|
||||
"ExpiredSignatureError",
|
||||
"ImmatureSignatureError",
|
||||
"InvalidAlgorithmError",
|
||||
"InvalidAudienceError",
|
||||
"InvalidIssuedAtError",
|
||||
"InvalidIssuerError",
|
||||
"InvalidKeyError",
|
||||
"InvalidSignatureError",
|
||||
"InvalidTokenError",
|
||||
"MissingRequiredClaimError",
|
||||
"PyJWKClientConnectionError",
|
||||
"PyJWKClientError",
|
||||
"PyJWKError",
|
||||
"PyJWKSetError",
|
||||
"PyJWTError",
|
||||
]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,188 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import time
|
||||
from collections.abc import Iterator
|
||||
from typing import Any
|
||||
|
||||
from .algorithms import get_default_algorithms, has_crypto, requires_cryptography
|
||||
from .exceptions import (
|
||||
InvalidKeyError,
|
||||
MissingCryptographyError,
|
||||
PyJWKError,
|
||||
PyJWKSetError,
|
||||
PyJWTError,
|
||||
)
|
||||
from .types import JWKDict
|
||||
|
||||
|
||||
class PyJWK:
|
||||
def __init__(self, jwk_data: JWKDict, algorithm: str | None = None) -> None:
|
||||
"""A class that represents a `JSON Web Key <https://www.rfc-editor.org/rfc/rfc7517>`_.
|
||||
|
||||
:param jwk_data: The decoded JWK data.
|
||||
:type jwk_data: dict[str, typing.Any]
|
||||
:param algorithm: The key algorithm. If not specified, the key's ``alg`` will be used.
|
||||
:type algorithm: str or None
|
||||
:raises InvalidKeyError: If the key type (``kty``) is not found or unsupported, or if the curve (``crv``) is not found or unsupported.
|
||||
:raises MissingCryptographyError: If the algorithm requires ``cryptography`` to be installed and it is not available.
|
||||
:raises PyJWKError: If unable to find an algorithm for the key.
|
||||
"""
|
||||
self._jwk_data = jwk_data
|
||||
|
||||
kty = self._jwk_data.get("kty", None)
|
||||
if not kty:
|
||||
raise InvalidKeyError(f"kty is not found: {self._jwk_data}")
|
||||
|
||||
if not algorithm and isinstance(self._jwk_data, dict):
|
||||
algorithm = self._jwk_data.get("alg", None)
|
||||
|
||||
if not algorithm:
|
||||
# Determine alg with kty (and crv).
|
||||
crv = self._jwk_data.get("crv", None)
|
||||
if kty == "EC":
|
||||
if crv == "P-256" or not crv:
|
||||
algorithm = "ES256"
|
||||
elif crv == "P-384":
|
||||
algorithm = "ES384"
|
||||
elif crv == "P-521":
|
||||
algorithm = "ES512"
|
||||
elif crv == "secp256k1":
|
||||
algorithm = "ES256K"
|
||||
else:
|
||||
raise InvalidKeyError(f"Unsupported crv: {crv}")
|
||||
elif kty == "RSA":
|
||||
algorithm = "RS256"
|
||||
elif kty == "oct":
|
||||
algorithm = "HS256"
|
||||
elif kty == "OKP":
|
||||
if not crv:
|
||||
raise InvalidKeyError(f"crv is not found: {self._jwk_data}")
|
||||
if crv == "Ed25519":
|
||||
algorithm = "EdDSA"
|
||||
else:
|
||||
raise InvalidKeyError(f"Unsupported crv: {crv}")
|
||||
else:
|
||||
raise InvalidKeyError(f"Unsupported kty: {kty}")
|
||||
|
||||
if not has_crypto and algorithm in requires_cryptography:
|
||||
raise MissingCryptographyError(
|
||||
f"{algorithm} requires 'cryptography' to be installed."
|
||||
)
|
||||
|
||||
self.algorithm_name = algorithm
|
||||
|
||||
try:
|
||||
self.Algorithm = get_default_algorithms()[algorithm]
|
||||
except KeyError:
|
||||
raise PyJWKError(
|
||||
f"Unable to find an algorithm for key: {self._jwk_data}",
|
||||
) from None
|
||||
|
||||
self.key = self.Algorithm.from_jwk(self._jwk_data)
|
||||
|
||||
@staticmethod
|
||||
def from_dict(obj: JWKDict, algorithm: str | None = None) -> PyJWK:
|
||||
"""Creates a :class:`PyJWK` object from a JSON-like dictionary.
|
||||
|
||||
:param obj: The JWK data, as a dictionary
|
||||
:type obj: dict[str, typing.Any]
|
||||
:param algorithm: The key algorithm. If not specified, the key's ``alg`` will be used.
|
||||
:type algorithm: str or None
|
||||
:rtype: PyJWK
|
||||
"""
|
||||
return PyJWK(obj, algorithm)
|
||||
|
||||
@staticmethod
|
||||
def from_json(data: str, algorithm: None = None) -> PyJWK:
|
||||
"""Create a :class:`PyJWK` object from a JSON string.
|
||||
Implicitly calls :meth:`PyJWK.from_dict()`.
|
||||
|
||||
:param str data: The JWK data, as a JSON string.
|
||||
:param algorithm: The key algorithm. If not specific, the key's ``alg`` will be used.
|
||||
:type algorithm: str or None
|
||||
|
||||
:rtype: PyJWK
|
||||
"""
|
||||
obj = json.loads(data)
|
||||
return PyJWK.from_dict(obj, algorithm)
|
||||
|
||||
@property
|
||||
def key_type(self) -> str | None:
|
||||
"""The `kty` property from the JWK.
|
||||
|
||||
:rtype: str or None
|
||||
"""
|
||||
return self._jwk_data.get("kty", None)
|
||||
|
||||
@property
|
||||
def key_id(self) -> str | None:
|
||||
"""The `kid` property from the JWK.
|
||||
|
||||
:rtype: str or None
|
||||
"""
|
||||
return self._jwk_data.get("kid", None)
|
||||
|
||||
@property
|
||||
def public_key_use(self) -> str | None:
|
||||
"""The `use` property from the JWK.
|
||||
|
||||
:rtype: str or None
|
||||
"""
|
||||
return self._jwk_data.get("use", None)
|
||||
|
||||
|
||||
class PyJWKSet:
|
||||
def __init__(self, keys: list[JWKDict]) -> None:
|
||||
self.keys: list[PyJWK] = []
|
||||
|
||||
if not keys:
|
||||
raise PyJWKSetError("The JWK Set did not contain any keys")
|
||||
|
||||
if not isinstance(keys, list):
|
||||
raise PyJWKSetError("Invalid JWK Set value")
|
||||
|
||||
for key in keys:
|
||||
try:
|
||||
self.keys.append(PyJWK(key))
|
||||
except PyJWTError as error:
|
||||
if isinstance(error, MissingCryptographyError):
|
||||
raise error
|
||||
# skip unusable keys
|
||||
continue
|
||||
|
||||
if len(self.keys) == 0:
|
||||
raise PyJWKSetError(
|
||||
"The JWK Set did not contain any usable keys. Perhaps 'cryptography' is not installed?"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def from_dict(obj: dict[str, Any]) -> PyJWKSet:
|
||||
keys = obj.get("keys", [])
|
||||
return PyJWKSet(keys)
|
||||
|
||||
@staticmethod
|
||||
def from_json(data: str) -> PyJWKSet:
|
||||
obj = json.loads(data)
|
||||
return PyJWKSet.from_dict(obj)
|
||||
|
||||
def __getitem__(self, kid: str) -> PyJWK:
|
||||
for key in self.keys:
|
||||
if key.key_id == kid:
|
||||
return key
|
||||
raise KeyError(f"keyset has no key for kid: {kid}")
|
||||
|
||||
def __iter__(self) -> Iterator[PyJWK]:
|
||||
return iter(self.keys)
|
||||
|
||||
|
||||
class PyJWTSetWithTimestamp:
|
||||
def __init__(self, jwk_set: PyJWKSet):
|
||||
self.jwk_set = jwk_set
|
||||
self.timestamp = time.monotonic()
|
||||
|
||||
def get_jwk_set(self) -> PyJWKSet:
|
||||
return self.jwk_set
|
||||
|
||||
def get_timestamp(self) -> float:
|
||||
return self.timestamp
|
||||
@@ -0,0 +1,456 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import binascii
|
||||
import json
|
||||
import warnings
|
||||
from collections.abc import Sequence
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from .algorithms import (
|
||||
Algorithm,
|
||||
get_default_algorithms,
|
||||
has_crypto,
|
||||
requires_cryptography,
|
||||
)
|
||||
from .api_jwk import PyJWK
|
||||
from .exceptions import (
|
||||
DecodeError,
|
||||
InvalidAlgorithmError,
|
||||
InvalidKeyError,
|
||||
InvalidSignatureError,
|
||||
InvalidTokenError,
|
||||
)
|
||||
from .utils import base64url_decode, base64url_encode
|
||||
from .warnings import InsecureKeyLengthWarning, RemovedInPyjwt3Warning
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .algorithms import AllowedPrivateKeys, AllowedPublicKeys
|
||||
from .types import SigOptions
|
||||
|
||||
_ALGORITHM_UNSET = object()
|
||||
|
||||
|
||||
class PyJWS:
|
||||
header_typ = "JWT"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
algorithms: Sequence[str] | None = None,
|
||||
options: SigOptions | None = None,
|
||||
) -> None:
|
||||
self._algorithms = get_default_algorithms()
|
||||
self._valid_algs = (
|
||||
set(algorithms) if algorithms is not None else set(self._algorithms)
|
||||
)
|
||||
|
||||
# Remove algorithms that aren't on the whitelist
|
||||
for key in list(self._algorithms.keys()):
|
||||
if key not in self._valid_algs:
|
||||
del self._algorithms[key]
|
||||
|
||||
self.options: SigOptions = self._get_default_options()
|
||||
if options is not None:
|
||||
self.options = {**self.options, **options}
|
||||
|
||||
@staticmethod
|
||||
def _get_default_options() -> SigOptions:
|
||||
return {"verify_signature": True, "enforce_minimum_key_length": False}
|
||||
|
||||
def register_algorithm(self, alg_id: str, alg_obj: Algorithm) -> None:
|
||||
"""
|
||||
Registers a new Algorithm for use when creating and verifying tokens.
|
||||
|
||||
:param str alg_id: the ID of the Algorithm
|
||||
:param alg_obj: the Algorithm object
|
||||
:type alg_obj: Algorithm
|
||||
"""
|
||||
if alg_id in self._algorithms:
|
||||
raise ValueError("Algorithm already has a handler.")
|
||||
|
||||
if not isinstance(alg_obj, Algorithm):
|
||||
raise TypeError("Object is not of type `Algorithm`")
|
||||
|
||||
self._algorithms[alg_id] = alg_obj
|
||||
self._valid_algs.add(alg_id)
|
||||
|
||||
def unregister_algorithm(self, alg_id: str) -> None:
|
||||
"""
|
||||
Unregisters an Algorithm for use when creating and verifying tokens
|
||||
:param str alg_id: the ID of the Algorithm
|
||||
:raises KeyError: if algorithm is not registered.
|
||||
"""
|
||||
if alg_id not in self._algorithms:
|
||||
raise KeyError(
|
||||
"The specified algorithm could not be removed"
|
||||
" because it is not registered."
|
||||
)
|
||||
|
||||
del self._algorithms[alg_id]
|
||||
self._valid_algs.remove(alg_id)
|
||||
|
||||
def get_algorithms(self) -> list[str]:
|
||||
"""
|
||||
Returns a list of supported values for the `alg` parameter.
|
||||
|
||||
:rtype: list[str]
|
||||
"""
|
||||
return list(self._valid_algs)
|
||||
|
||||
def get_algorithm_by_name(self, alg_name: str) -> Algorithm:
|
||||
"""
|
||||
For a given string name, return the matching Algorithm object.
|
||||
|
||||
Example usage:
|
||||
>>> jws_obj = PyJWS()
|
||||
>>> jws_obj.get_algorithm_by_name("RS256")
|
||||
|
||||
:param alg_name: The name of the algorithm to retrieve
|
||||
:type alg_name: str
|
||||
:rtype: Algorithm
|
||||
"""
|
||||
try:
|
||||
return self._algorithms[alg_name]
|
||||
except KeyError as e:
|
||||
if not has_crypto and alg_name in requires_cryptography:
|
||||
raise NotImplementedError(
|
||||
f"Algorithm '{alg_name}' could not be found. Do you have cryptography installed?"
|
||||
) from e
|
||||
raise NotImplementedError("Algorithm not supported") from e
|
||||
|
||||
def encode(
|
||||
self,
|
||||
payload: bytes,
|
||||
key: AllowedPrivateKeys | PyJWK | str | bytes,
|
||||
algorithm: str | None = _ALGORITHM_UNSET, # type: ignore[assignment]
|
||||
headers: dict[str, Any] | None = None,
|
||||
json_encoder: type[json.JSONEncoder] | None = None,
|
||||
is_payload_detached: bool = False,
|
||||
sort_headers: bool = True,
|
||||
) -> str:
|
||||
segments: list[bytes] = []
|
||||
|
||||
# declare a new var to narrow the type for type checkers
|
||||
if algorithm is _ALGORITHM_UNSET:
|
||||
if isinstance(key, PyJWK):
|
||||
algorithm_ = key.algorithm_name
|
||||
else:
|
||||
algorithm_ = "HS256"
|
||||
elif algorithm is None:
|
||||
if isinstance(key, PyJWK):
|
||||
algorithm_ = key.algorithm_name
|
||||
else:
|
||||
algorithm_ = "none"
|
||||
else:
|
||||
algorithm_ = algorithm
|
||||
|
||||
# Prefer headers values if present to function parameters.
|
||||
if headers:
|
||||
headers_alg = headers.get("alg")
|
||||
if headers_alg:
|
||||
algorithm_ = headers["alg"]
|
||||
|
||||
headers_b64 = headers.get("b64")
|
||||
if headers_b64 is False:
|
||||
is_payload_detached = True
|
||||
|
||||
# Header
|
||||
header: dict[str, Any] = {"typ": self.header_typ, "alg": algorithm_}
|
||||
|
||||
if headers:
|
||||
self._validate_headers(headers, encoding=True)
|
||||
header.update(headers)
|
||||
|
||||
if not header["typ"]:
|
||||
del header["typ"]
|
||||
|
||||
if is_payload_detached:
|
||||
header["b64"] = False
|
||||
# RFC 7797 §3: producers MUST list "b64" in "crit" whenever
|
||||
# "b64" appears in the protected header, so b64-unaware
|
||||
# verifiers don't silently treat an unencoded payload as
|
||||
# base64-encoded.
|
||||
existing_crit = header.get("crit", [])
|
||||
if not isinstance(existing_crit, list):
|
||||
raise InvalidTokenError("Invalid 'crit' header: must be a list")
|
||||
if "b64" not in existing_crit:
|
||||
header["crit"] = [*existing_crit, "b64"]
|
||||
elif "b64" in header:
|
||||
# True is the standard value for b64, so no need for it
|
||||
del header["b64"]
|
||||
|
||||
json_header = json.dumps(
|
||||
header, separators=(",", ":"), cls=json_encoder, sort_keys=sort_headers
|
||||
).encode()
|
||||
|
||||
segments.append(base64url_encode(json_header))
|
||||
|
||||
if is_payload_detached:
|
||||
msg_payload = payload
|
||||
else:
|
||||
msg_payload = base64url_encode(payload)
|
||||
segments.append(msg_payload)
|
||||
|
||||
# Segments
|
||||
signing_input = b".".join(segments)
|
||||
|
||||
alg_obj = self.get_algorithm_by_name(algorithm_)
|
||||
if isinstance(key, PyJWK):
|
||||
key = key.key
|
||||
key = alg_obj.prepare_key(key)
|
||||
|
||||
key_length_msg = alg_obj.check_key_length(key)
|
||||
if key_length_msg:
|
||||
if self.options.get("enforce_minimum_key_length", False):
|
||||
raise InvalidKeyError(key_length_msg)
|
||||
else:
|
||||
warnings.warn(key_length_msg, InsecureKeyLengthWarning, stacklevel=2)
|
||||
|
||||
signature = alg_obj.sign(signing_input, key)
|
||||
|
||||
segments.append(base64url_encode(signature))
|
||||
|
||||
# Don't put the payload content inside the encoded token when detached
|
||||
if is_payload_detached:
|
||||
segments[1] = b""
|
||||
encoded_string = b".".join(segments)
|
||||
|
||||
return encoded_string.decode("utf-8")
|
||||
|
||||
def decode_complete(
|
||||
self,
|
||||
jwt: str | bytes,
|
||||
key: AllowedPublicKeys | PyJWK | str | bytes = "",
|
||||
algorithms: Sequence[str] | None = None,
|
||||
options: SigOptions | None = None,
|
||||
detached_payload: bytes | None = None,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> dict[str, Any]:
|
||||
if kwargs:
|
||||
warnings.warn(
|
||||
"passing additional kwargs to decode_complete() is deprecated "
|
||||
"and will be removed in pyjwt version 3. "
|
||||
f"Unsupported kwargs: {tuple(kwargs.keys())}",
|
||||
RemovedInPyjwt3Warning,
|
||||
stacklevel=2,
|
||||
)
|
||||
merged_options: SigOptions
|
||||
if options is None:
|
||||
merged_options = self.options
|
||||
else:
|
||||
merged_options = {**self.options, **options}
|
||||
|
||||
verify_signature = merged_options["verify_signature"]
|
||||
|
||||
if verify_signature and not algorithms and not isinstance(key, PyJWK):
|
||||
raise DecodeError(
|
||||
'It is required that you pass in a value for the "algorithms" argument when calling decode().'
|
||||
)
|
||||
|
||||
payload, signing_input, header, signature = self._load(jwt)
|
||||
|
||||
self._validate_headers(header)
|
||||
|
||||
if header.get("b64", True) is False:
|
||||
# RFC 7797 §3: when "b64" is present in the protected header,
|
||||
# it MUST also appear in "crit". A token that sets b64=false
|
||||
# without declaring it critical is malformed.
|
||||
crit = header.get("crit") or []
|
||||
if not isinstance(crit, list) or "b64" not in crit:
|
||||
raise InvalidTokenError(
|
||||
"The 'b64' header parameter requires 'b64' to be listed in 'crit'."
|
||||
)
|
||||
if detached_payload is None:
|
||||
raise DecodeError(
|
||||
'It is required that you pass in a value for the "detached_payload" argument to decode a message having the b64 header set to false.'
|
||||
)
|
||||
payload = detached_payload
|
||||
signing_input = b".".join([signing_input.rsplit(b".", 1)[0], payload])
|
||||
|
||||
if verify_signature:
|
||||
self._verify_signature(
|
||||
signing_input,
|
||||
header,
|
||||
signature,
|
||||
key,
|
||||
algorithms,
|
||||
options=merged_options,
|
||||
)
|
||||
|
||||
return {
|
||||
"payload": payload,
|
||||
"header": header,
|
||||
"signature": signature,
|
||||
}
|
||||
|
||||
def decode(
|
||||
self,
|
||||
jwt: str | bytes,
|
||||
key: AllowedPublicKeys | PyJWK | str | bytes = "",
|
||||
algorithms: Sequence[str] | None = None,
|
||||
options: SigOptions | None = None,
|
||||
detached_payload: bytes | None = None,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> Any:
|
||||
if kwargs:
|
||||
warnings.warn(
|
||||
"passing additional kwargs to decode() is deprecated "
|
||||
"and will be removed in pyjwt version 3. "
|
||||
f"Unsupported kwargs: {tuple(kwargs.keys())}",
|
||||
RemovedInPyjwt3Warning,
|
||||
stacklevel=2,
|
||||
)
|
||||
decoded = self.decode_complete(
|
||||
jwt, key, algorithms, options, detached_payload=detached_payload
|
||||
)
|
||||
return decoded["payload"]
|
||||
|
||||
def get_unverified_header(self, jwt: str | bytes) -> dict[str, Any]:
|
||||
"""Returns back the JWT header parameters as a `dict`
|
||||
|
||||
Note: The signature is not verified so the header parameters
|
||||
should not be fully trusted until signature verification is complete
|
||||
"""
|
||||
headers = self._load(jwt)[2]
|
||||
self._validate_headers(headers)
|
||||
|
||||
return headers
|
||||
|
||||
def _load(self, jwt: str | bytes) -> tuple[bytes, bytes, dict[str, Any], bytes]:
|
||||
if isinstance(jwt, str):
|
||||
jwt = jwt.encode("utf-8")
|
||||
|
||||
if not isinstance(jwt, bytes):
|
||||
raise DecodeError(f"Invalid token type. Token must be a {bytes}")
|
||||
|
||||
try:
|
||||
signing_input, crypto_segment = jwt.rsplit(b".", 1)
|
||||
header_segment, payload_segment = signing_input.split(b".", 1)
|
||||
except ValueError as err:
|
||||
raise DecodeError("Not enough segments") from err
|
||||
|
||||
try:
|
||||
header_data = base64url_decode(header_segment)
|
||||
except (TypeError, binascii.Error) as err:
|
||||
raise DecodeError("Invalid header padding") from err
|
||||
|
||||
try:
|
||||
header: dict[str, Any] = json.loads(header_data)
|
||||
except ValueError as e:
|
||||
raise DecodeError(f"Invalid header string: {e}") from e
|
||||
|
||||
if not isinstance(header, dict):
|
||||
raise DecodeError("Invalid header string: must be a json object")
|
||||
|
||||
if header.get("b64", True) is False:
|
||||
# Detached payload form (RFC 7515 Appendix F): the compact-form
|
||||
# payload segment must be empty; the caller supplies the actual
|
||||
# payload via the `detached_payload` argument in decode_complete.
|
||||
# Skipping the base64 decode here removes an unauthenticated work
|
||||
# amplifier — otherwise an attacker can inflate the unused
|
||||
# segment to force CPU + memory cost before the signature is
|
||||
# even checked.
|
||||
if payload_segment:
|
||||
raise DecodeError("Payload segment must be empty when 'b64' is false.")
|
||||
payload = b""
|
||||
else:
|
||||
try:
|
||||
payload = base64url_decode(payload_segment)
|
||||
except (TypeError, binascii.Error) as err:
|
||||
raise DecodeError("Invalid payload padding") from err
|
||||
|
||||
try:
|
||||
signature = base64url_decode(crypto_segment)
|
||||
except (TypeError, binascii.Error) as err:
|
||||
raise DecodeError("Invalid crypto padding") from err
|
||||
|
||||
return (payload, signing_input, header, signature)
|
||||
|
||||
def _verify_signature(
|
||||
self,
|
||||
signing_input: bytes,
|
||||
header: dict[str, Any],
|
||||
signature: bytes,
|
||||
key: AllowedPublicKeys | PyJWK | str | bytes = "",
|
||||
algorithms: Sequence[str] | None = None,
|
||||
options: SigOptions | None = None,
|
||||
) -> None:
|
||||
effective_options = options if options is not None else self.options
|
||||
|
||||
if algorithms is None and isinstance(key, PyJWK):
|
||||
algorithms = [key.algorithm_name]
|
||||
try:
|
||||
alg = header["alg"]
|
||||
except KeyError:
|
||||
raise InvalidAlgorithmError("Algorithm not specified") from None
|
||||
|
||||
if not alg or (algorithms is not None and alg not in algorithms):
|
||||
raise InvalidAlgorithmError("The specified alg value is not allowed")
|
||||
|
||||
if isinstance(key, PyJWK):
|
||||
# The PyJWK has a fixed algorithm bound at construction time.
|
||||
# Verification must use that algorithm, not whatever the token
|
||||
# header advertises, otherwise the caller's allow-list check
|
||||
# above degenerates into a string compare with no behavioural
|
||||
# effect on which algorithm actually verifies the signature.
|
||||
if alg != key.algorithm_name:
|
||||
raise InvalidAlgorithmError(
|
||||
f"Token algorithm {alg!r} does not match the key's "
|
||||
f"algorithm {key.algorithm_name!r}"
|
||||
)
|
||||
alg_obj = key.Algorithm
|
||||
prepared_key = key.key
|
||||
else:
|
||||
try:
|
||||
alg_obj = self.get_algorithm_by_name(alg)
|
||||
except NotImplementedError as e:
|
||||
raise InvalidAlgorithmError("Algorithm not supported") from e
|
||||
prepared_key = alg_obj.prepare_key(key)
|
||||
|
||||
key_length_msg = alg_obj.check_key_length(prepared_key)
|
||||
if key_length_msg:
|
||||
if effective_options.get("enforce_minimum_key_length", False):
|
||||
raise InvalidKeyError(key_length_msg)
|
||||
else:
|
||||
warnings.warn(key_length_msg, InsecureKeyLengthWarning, stacklevel=4)
|
||||
|
||||
if not alg_obj.verify(signing_input, prepared_key, signature):
|
||||
raise InvalidSignatureError("Signature verification failed")
|
||||
|
||||
# Extensions that PyJWT actually understands and supports
|
||||
_supported_crit: set[str] = {"b64"}
|
||||
|
||||
def _validate_headers(
|
||||
self, headers: dict[str, Any], *, encoding: bool = False
|
||||
) -> None:
|
||||
if "kid" in headers:
|
||||
self._validate_kid(headers["kid"])
|
||||
if not encoding and "crit" in headers:
|
||||
self._validate_crit(headers)
|
||||
|
||||
def _validate_kid(self, kid: Any) -> None:
|
||||
if not isinstance(kid, str):
|
||||
raise InvalidTokenError("Key ID header parameter must be a string")
|
||||
|
||||
def _validate_crit(self, headers: dict[str, Any]) -> None:
|
||||
crit = headers["crit"]
|
||||
if not isinstance(crit, list) or len(crit) == 0:
|
||||
raise InvalidTokenError("Invalid 'crit' header: must be a non-empty list")
|
||||
for ext in crit:
|
||||
if not isinstance(ext, str):
|
||||
raise InvalidTokenError("Invalid 'crit' header: values must be strings")
|
||||
if ext not in self._supported_crit:
|
||||
raise InvalidTokenError(f"Unsupported critical extension: {ext}")
|
||||
if ext not in headers:
|
||||
raise InvalidTokenError(
|
||||
f"Critical extension '{ext}' is missing from headers"
|
||||
)
|
||||
|
||||
|
||||
_jws_global_obj = PyJWS()
|
||||
encode = _jws_global_obj.encode
|
||||
decode_complete = _jws_global_obj.decode_complete
|
||||
decode = _jws_global_obj.decode
|
||||
register_algorithm = _jws_global_obj.register_algorithm
|
||||
unregister_algorithm = _jws_global_obj.unregister_algorithm
|
||||
get_algorithm_by_name = _jws_global_obj.get_algorithm_by_name
|
||||
get_unverified_header = _jws_global_obj.get_unverified_header
|
||||
@@ -0,0 +1,593 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
import warnings
|
||||
from calendar import timegm
|
||||
from collections.abc import Container, Iterable, Sequence
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Any, Union, cast
|
||||
|
||||
from .api_jws import PyJWS, _ALGORITHM_UNSET, _jws_global_obj
|
||||
from .exceptions import (
|
||||
DecodeError,
|
||||
ExpiredSignatureError,
|
||||
ImmatureSignatureError,
|
||||
InvalidAudienceError,
|
||||
InvalidIssuedAtError,
|
||||
InvalidIssuerError,
|
||||
InvalidJTIError,
|
||||
InvalidSubjectError,
|
||||
MissingRequiredClaimError,
|
||||
)
|
||||
from .warnings import RemovedInPyjwt3Warning
|
||||
|
||||
if TYPE_CHECKING or bool(os.getenv("SPHINX_BUILD", "")):
|
||||
import sys
|
||||
|
||||
if sys.version_info >= (3, 10):
|
||||
from typing import TypeAlias
|
||||
else:
|
||||
# Python 3.9 and lower
|
||||
from typing_extensions import TypeAlias
|
||||
|
||||
from .algorithms import AllowedPrivateKeys, AllowedPublicKeys
|
||||
from .api_jwk import PyJWK
|
||||
from .types import FullOptions, Options, SigOptions
|
||||
|
||||
AllowedPrivateKeyTypes: TypeAlias = Union[AllowedPrivateKeys, PyJWK, str, bytes]
|
||||
AllowedPublicKeyTypes: TypeAlias = Union[AllowedPublicKeys, PyJWK, str, bytes]
|
||||
|
||||
|
||||
class PyJWT:
|
||||
def __init__(self, options: Options | None = None) -> None:
|
||||
self.options: FullOptions
|
||||
self.options = self._get_default_options()
|
||||
if options is not None:
|
||||
self.options = self._merge_options(options)
|
||||
|
||||
self._jws = PyJWS(options=self._get_sig_options())
|
||||
|
||||
@staticmethod
|
||||
def _get_default_options() -> FullOptions:
|
||||
return {
|
||||
"verify_signature": True,
|
||||
"verify_exp": True,
|
||||
"verify_nbf": True,
|
||||
"verify_iat": True,
|
||||
"verify_aud": True,
|
||||
"verify_iss": True,
|
||||
"verify_sub": True,
|
||||
"verify_jti": True,
|
||||
"require": [],
|
||||
"strict_aud": False,
|
||||
"enforce_minimum_key_length": False,
|
||||
}
|
||||
|
||||
def _get_sig_options(self) -> SigOptions:
|
||||
return {
|
||||
"verify_signature": self.options["verify_signature"],
|
||||
"enforce_minimum_key_length": self.options.get(
|
||||
"enforce_minimum_key_length", False
|
||||
),
|
||||
}
|
||||
|
||||
def _merge_options(self, options: Options | None = None) -> FullOptions:
|
||||
if options is None:
|
||||
return self.options
|
||||
|
||||
# (defensive) set defaults for verify_x to False if verify_signature is False
|
||||
if not options.get("verify_signature", True):
|
||||
options["verify_exp"] = options.get("verify_exp", False)
|
||||
options["verify_nbf"] = options.get("verify_nbf", False)
|
||||
options["verify_iat"] = options.get("verify_iat", False)
|
||||
options["verify_aud"] = options.get("verify_aud", False)
|
||||
options["verify_iss"] = options.get("verify_iss", False)
|
||||
options["verify_sub"] = options.get("verify_sub", False)
|
||||
options["verify_jti"] = options.get("verify_jti", False)
|
||||
return {**self.options, **options}
|
||||
|
||||
def encode(
|
||||
self,
|
||||
payload: dict[str, Any],
|
||||
key: AllowedPrivateKeyTypes,
|
||||
algorithm: str | None = _ALGORITHM_UNSET, # type: ignore[assignment]
|
||||
headers: dict[str, Any] | None = None,
|
||||
json_encoder: type[json.JSONEncoder] | None = None,
|
||||
sort_headers: bool = True,
|
||||
) -> str:
|
||||
"""Encode the ``payload`` as JSON Web Token.
|
||||
|
||||
:param payload: JWT claims, e.g. ``dict(iss=..., aud=..., sub=...)``
|
||||
:type payload: dict[str, typing.Any]
|
||||
:param key: a key suitable for the chosen algorithm:
|
||||
|
||||
* for **asymmetric algorithms**: PEM-formatted private key, a multiline string
|
||||
* for **symmetric algorithms**: plain string, sufficiently long for security
|
||||
|
||||
:type key: str or bytes or PyJWK or :py:class:`jwt.algorithms.AllowedPrivateKeys`
|
||||
:param algorithm: algorithm to sign the token with, e.g. ``"ES256"``.
|
||||
If ``headers`` includes ``alg``, it will be preferred to this parameter.
|
||||
If ``key`` is a :class:`PyJWK` object, by default the key algorithm will be used.
|
||||
:type algorithm: str or None
|
||||
:param headers: additional JWT header fields, e.g. ``dict(kid="my-key-id")``.
|
||||
:type headers: dict[str, typing.Any] or None
|
||||
:param json_encoder: custom JSON encoder for ``payload`` and ``headers``
|
||||
:type json_encoder: json.JSONEncoder or None
|
||||
|
||||
:rtype: str
|
||||
:returns: a JSON Web Token
|
||||
|
||||
:raises TypeError: if ``payload`` is not a ``dict``
|
||||
"""
|
||||
# Check that we get a dict
|
||||
if not isinstance(payload, dict):
|
||||
raise TypeError(
|
||||
"Expecting a dict object, as JWT only supports "
|
||||
"JSON objects as payloads."
|
||||
)
|
||||
|
||||
# Payload
|
||||
payload = payload.copy()
|
||||
for time_claim in ["exp", "iat", "nbf"]:
|
||||
# Convert datetime to a intDate value in known time-format claims
|
||||
if isinstance(payload.get(time_claim), datetime):
|
||||
payload[time_claim] = timegm(payload[time_claim].utctimetuple())
|
||||
|
||||
# Issue #1039, iss being set to non-string
|
||||
if "iss" in payload and not isinstance(payload["iss"], str):
|
||||
raise TypeError("Issuer (iss) must be a string.")
|
||||
|
||||
json_payload = self._encode_payload(
|
||||
payload,
|
||||
headers=headers,
|
||||
json_encoder=json_encoder,
|
||||
)
|
||||
|
||||
return self._jws.encode(
|
||||
json_payload,
|
||||
key,
|
||||
algorithm,
|
||||
headers,
|
||||
json_encoder,
|
||||
sort_headers=sort_headers,
|
||||
)
|
||||
|
||||
def _encode_payload(
|
||||
self,
|
||||
payload: dict[str, Any],
|
||||
headers: dict[str, Any] | None = None,
|
||||
json_encoder: type[json.JSONEncoder] | None = None,
|
||||
) -> bytes:
|
||||
"""
|
||||
Encode a given payload to the bytes to be signed.
|
||||
|
||||
This method is intended to be overridden by subclasses that need to
|
||||
encode the payload in a different way, e.g. compress the payload.
|
||||
"""
|
||||
return json.dumps(
|
||||
payload,
|
||||
separators=(",", ":"),
|
||||
cls=json_encoder,
|
||||
).encode("utf-8")
|
||||
|
||||
def decode_complete(
|
||||
self,
|
||||
jwt: str | bytes,
|
||||
key: AllowedPublicKeyTypes = "",
|
||||
algorithms: Sequence[str] | None = None,
|
||||
options: Options | None = None,
|
||||
# deprecated arg, remove in pyjwt3
|
||||
verify: bool | None = None,
|
||||
# could be used as passthrough to api_jws, consider removal in pyjwt3
|
||||
detached_payload: bytes | None = None,
|
||||
# passthrough arguments to _validate_claims
|
||||
# consider putting in options
|
||||
audience: str | Iterable[str] | None = None,
|
||||
issuer: str | Container[str] | None = None,
|
||||
subject: str | None = None,
|
||||
leeway: float | timedelta = 0,
|
||||
# kwargs
|
||||
**kwargs: Any,
|
||||
) -> dict[str, Any]:
|
||||
"""Identical to ``jwt.decode`` except for return value which is a dictionary containing the token header (JOSE Header),
|
||||
the token payload (JWT Payload), and token signature (JWT Signature) on the keys "header", "payload",
|
||||
and "signature" respectively.
|
||||
|
||||
:param jwt: the token to be decoded
|
||||
:type jwt: str or bytes
|
||||
:param key: the key suitable for the allowed algorithm
|
||||
:type key: str or bytes or PyJWK or :py:class:`jwt.algorithms.AllowedPublicKeys`
|
||||
|
||||
:param algorithms: allowed algorithms, e.g. ``["ES256"]``
|
||||
|
||||
.. warning::
|
||||
|
||||
Do **not** compute the ``algorithms`` parameter based on
|
||||
the ``alg`` from the token itself, or on any other data
|
||||
that an attacker may be able to influence, as that might
|
||||
expose you to various vulnerabilities (see `RFC 8725 §2.1
|
||||
<https://www.rfc-editor.org/rfc/rfc8725.html#section-2.1>`_). Instead,
|
||||
either hard-code a fixed value for ``algorithms``, or
|
||||
configure it in the same place you configure the
|
||||
``key``. Make sure not to mix symmetric and asymmetric
|
||||
algorithms that interpret the ``key`` in different ways
|
||||
(e.g. HS\\* and RS\\*).
|
||||
:type algorithms: typing.Sequence[str] or None
|
||||
|
||||
:param jwt.types.Options options: extended decoding and validation options
|
||||
Refer to :py:class:`jwt.types.Options` for more information.
|
||||
|
||||
:param audience: optional, the value for ``verify_aud`` check
|
||||
:type audience: str or typing.Iterable[str] or None
|
||||
:param issuer: optional, the value for ``verify_iss`` check
|
||||
:type issuer: str or typing.Container[str] or None
|
||||
:param leeway: a time margin in seconds for the expiration check
|
||||
:type leeway: float or datetime.timedelta
|
||||
:rtype: dict[str, typing.Any]
|
||||
:returns: Decoded JWT with the JOSE Header on the key ``header``, the JWS
|
||||
Payload on the key ``payload``, and the JWS Signature on the key ``signature``.
|
||||
"""
|
||||
if kwargs:
|
||||
warnings.warn(
|
||||
"passing additional kwargs to decode_complete() is deprecated "
|
||||
"and will be removed in pyjwt version 3. "
|
||||
f"Unsupported kwargs: {tuple(kwargs.keys())}",
|
||||
RemovedInPyjwt3Warning,
|
||||
stacklevel=2,
|
||||
)
|
||||
|
||||
if options is None:
|
||||
verify_signature = True
|
||||
else:
|
||||
verify_signature = options.get("verify_signature", True)
|
||||
|
||||
# If the user has set the legacy `verify` argument, and it doesn't match
|
||||
# what the relevant `options` entry for the argument is, inform the user
|
||||
# that they're likely making a mistake.
|
||||
if verify is not None and verify != verify_signature:
|
||||
warnings.warn(
|
||||
"The `verify` argument to `decode` does nothing in PyJWT 2.0 and newer. "
|
||||
"The equivalent is setting `verify_signature` to False in the `options` dictionary. "
|
||||
"This invocation has a mismatch between the kwarg and the option entry.",
|
||||
category=DeprecationWarning,
|
||||
stacklevel=2,
|
||||
)
|
||||
|
||||
merged_options = self._merge_options(options)
|
||||
|
||||
sig_options: SigOptions = {
|
||||
"verify_signature": verify_signature,
|
||||
"enforce_minimum_key_length": merged_options.get(
|
||||
"enforce_minimum_key_length", False
|
||||
),
|
||||
}
|
||||
decoded = self._jws.decode_complete(
|
||||
jwt,
|
||||
key=key,
|
||||
algorithms=algorithms,
|
||||
options=sig_options,
|
||||
detached_payload=detached_payload,
|
||||
)
|
||||
|
||||
payload = self._decode_payload(decoded)
|
||||
|
||||
self._validate_claims(
|
||||
payload,
|
||||
merged_options,
|
||||
audience=audience,
|
||||
issuer=issuer,
|
||||
leeway=leeway,
|
||||
subject=subject,
|
||||
)
|
||||
|
||||
decoded["payload"] = payload
|
||||
return decoded
|
||||
|
||||
def _decode_payload(self, decoded: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
Decode the payload from a JWS dictionary (payload, signature, header).
|
||||
|
||||
This method is intended to be overridden by subclasses that need to
|
||||
decode the payload in a different way, e.g. decompress compressed
|
||||
payloads.
|
||||
"""
|
||||
try:
|
||||
payload: dict[str, Any] = json.loads(decoded["payload"])
|
||||
except ValueError as e:
|
||||
raise DecodeError(f"Invalid payload string: {e}") from e
|
||||
if not isinstance(payload, dict):
|
||||
raise DecodeError("Invalid payload string: must be a json object")
|
||||
return payload
|
||||
|
||||
def decode(
|
||||
self,
|
||||
jwt: str | bytes,
|
||||
key: AllowedPublicKeys | PyJWK | str | bytes = "",
|
||||
algorithms: Sequence[str] | None = None,
|
||||
options: Options | None = None,
|
||||
# deprecated arg, remove in pyjwt3
|
||||
verify: bool | None = None,
|
||||
# could be used as passthrough to api_jws, consider removal in pyjwt3
|
||||
detached_payload: bytes | None = None,
|
||||
# passthrough arguments to _validate_claims
|
||||
# consider putting in options
|
||||
audience: str | Iterable[str] | None = None,
|
||||
subject: str | None = None,
|
||||
issuer: str | Container[str] | None = None,
|
||||
leeway: float | timedelta = 0,
|
||||
# kwargs
|
||||
**kwargs: Any,
|
||||
) -> dict[str, Any]:
|
||||
"""Verify the ``jwt`` token signature and return the token claims.
|
||||
|
||||
:param jwt: the token to be decoded
|
||||
:type jwt: str or bytes
|
||||
:param key: the key suitable for the allowed algorithm
|
||||
:type key: str or bytes or PyJWK or :py:class:`jwt.algorithms.AllowedPublicKeys`
|
||||
|
||||
:param algorithms: allowed algorithms, e.g. ``["ES256"]``
|
||||
If ``key`` is a :class:`PyJWK` object, allowed algorithms will default to the key algorithm.
|
||||
|
||||
.. warning::
|
||||
|
||||
Do **not** compute the ``algorithms`` parameter based on
|
||||
the ``alg`` from the token itself, or on any other data
|
||||
that an attacker may be able to influence, as that might
|
||||
expose you to various vulnerabilities (see `RFC 8725 §2.1
|
||||
<https://www.rfc-editor.org/rfc/rfc8725.html#section-2.1>`_). Instead,
|
||||
either hard-code a fixed value for ``algorithms``, or
|
||||
configure it in the same place you configure the
|
||||
``key``. Make sure not to mix symmetric and asymmetric
|
||||
algorithms that interpret the ``key`` in different ways
|
||||
(e.g. HS\\* and RS\\*).
|
||||
:type algorithms: typing.Sequence[str] or None
|
||||
|
||||
:param jwt.types.Options options: extended decoding and validation options
|
||||
Refer to :py:class:`jwt.types.Options` for more information.
|
||||
|
||||
:param audience: optional, the value for ``verify_aud`` check
|
||||
:type audience: str or typing.Iterable[str] or None
|
||||
:param subject: optional, the value for ``verify_sub`` check
|
||||
:type subject: str or None
|
||||
:param issuer: optional, the value for ``verify_iss`` check
|
||||
:type issuer: str or typing.Container[str] or None
|
||||
:param leeway: a time margin in seconds for the expiration check
|
||||
:type leeway: float or datetime.timedelta
|
||||
:rtype: dict[str, typing.Any]
|
||||
:returns: the JWT claims
|
||||
"""
|
||||
if kwargs:
|
||||
warnings.warn(
|
||||
"passing additional kwargs to decode() is deprecated "
|
||||
"and will be removed in pyjwt version 3. "
|
||||
f"Unsupported kwargs: {tuple(kwargs.keys())}",
|
||||
RemovedInPyjwt3Warning,
|
||||
stacklevel=2,
|
||||
)
|
||||
decoded = self.decode_complete(
|
||||
jwt,
|
||||
key,
|
||||
algorithms,
|
||||
options,
|
||||
verify=verify,
|
||||
detached_payload=detached_payload,
|
||||
audience=audience,
|
||||
subject=subject,
|
||||
issuer=issuer,
|
||||
leeway=leeway,
|
||||
)
|
||||
return cast(dict[str, Any], decoded["payload"])
|
||||
|
||||
def _validate_claims(
|
||||
self,
|
||||
payload: dict[str, Any],
|
||||
options: FullOptions,
|
||||
audience: Iterable[str] | str | None = None,
|
||||
issuer: Container[str] | str | None = None,
|
||||
subject: str | None = None,
|
||||
leeway: float | timedelta = 0,
|
||||
) -> None:
|
||||
if isinstance(leeway, timedelta):
|
||||
leeway = leeway.total_seconds()
|
||||
|
||||
if audience is not None and not isinstance(audience, (str, Iterable)):
|
||||
raise TypeError("audience must be a string, iterable or None")
|
||||
|
||||
self._validate_required_claims(payload, options["require"])
|
||||
|
||||
now = datetime.now(tz=timezone.utc).timestamp()
|
||||
|
||||
if "iat" in payload and options["verify_iat"]:
|
||||
self._validate_iat(payload, now, leeway)
|
||||
|
||||
if "nbf" in payload and options["verify_nbf"]:
|
||||
self._validate_nbf(payload, now, leeway)
|
||||
|
||||
if "exp" in payload and options["verify_exp"]:
|
||||
self._validate_exp(payload, now, leeway)
|
||||
|
||||
if options["verify_iss"]:
|
||||
self._validate_iss(payload, issuer)
|
||||
|
||||
if options["verify_aud"]:
|
||||
self._validate_aud(
|
||||
payload, audience, strict=options.get("strict_aud", False)
|
||||
)
|
||||
|
||||
if options["verify_sub"]:
|
||||
self._validate_sub(payload, subject)
|
||||
|
||||
if options["verify_jti"]:
|
||||
self._validate_jti(payload)
|
||||
|
||||
def _validate_required_claims(
|
||||
self,
|
||||
payload: dict[str, Any],
|
||||
claims: Iterable[str],
|
||||
) -> None:
|
||||
for claim in claims:
|
||||
if payload.get(claim) is None:
|
||||
raise MissingRequiredClaimError(claim)
|
||||
|
||||
def _validate_sub(
|
||||
self, payload: dict[str, Any], subject: str | None = None
|
||||
) -> None:
|
||||
"""
|
||||
Checks whether "sub" if in the payload is valid or not.
|
||||
This is an Optional claim
|
||||
|
||||
:param payload(dict): The payload which needs to be validated
|
||||
:param subject(str): The subject of the token
|
||||
"""
|
||||
|
||||
if "sub" not in payload:
|
||||
return
|
||||
|
||||
if not isinstance(payload["sub"], str):
|
||||
raise InvalidSubjectError("Subject must be a string")
|
||||
|
||||
if subject is not None:
|
||||
if payload.get("sub") != subject:
|
||||
raise InvalidSubjectError("Invalid subject")
|
||||
|
||||
def _validate_jti(self, payload: dict[str, Any]) -> None:
|
||||
"""
|
||||
Checks whether "jti" if in the payload is valid or not
|
||||
This is an Optional claim
|
||||
|
||||
:param payload(dict): The payload which needs to be validated
|
||||
"""
|
||||
|
||||
if "jti" not in payload:
|
||||
return
|
||||
|
||||
if not isinstance(payload.get("jti"), str):
|
||||
raise InvalidJTIError("JWT ID must be a string")
|
||||
|
||||
def _validate_iat(
|
||||
self,
|
||||
payload: dict[str, Any],
|
||||
now: float,
|
||||
leeway: float,
|
||||
) -> None:
|
||||
try:
|
||||
iat = int(payload["iat"])
|
||||
except ValueError:
|
||||
raise InvalidIssuedAtError(
|
||||
"Issued At claim (iat) must be an integer."
|
||||
) from None
|
||||
if iat > (now + leeway):
|
||||
raise ImmatureSignatureError("The token is not yet valid (iat)")
|
||||
|
||||
def _validate_nbf(
|
||||
self,
|
||||
payload: dict[str, Any],
|
||||
now: float,
|
||||
leeway: float,
|
||||
) -> None:
|
||||
try:
|
||||
nbf = int(payload["nbf"])
|
||||
except ValueError:
|
||||
raise DecodeError("Not Before claim (nbf) must be an integer.") from None
|
||||
|
||||
if nbf > (now + leeway):
|
||||
raise ImmatureSignatureError("The token is not yet valid (nbf)")
|
||||
|
||||
def _validate_exp(
|
||||
self,
|
||||
payload: dict[str, Any],
|
||||
now: float,
|
||||
leeway: float,
|
||||
) -> None:
|
||||
try:
|
||||
exp = int(payload["exp"])
|
||||
except ValueError:
|
||||
raise DecodeError(
|
||||
"Expiration Time claim (exp) must be an integer."
|
||||
) from None
|
||||
|
||||
if exp <= (now - leeway):
|
||||
raise ExpiredSignatureError("Signature has expired")
|
||||
|
||||
def _validate_aud(
|
||||
self,
|
||||
payload: dict[str, Any],
|
||||
audience: str | Iterable[str] | None,
|
||||
*,
|
||||
strict: bool = False,
|
||||
) -> None:
|
||||
if audience is None:
|
||||
if "aud" not in payload or not payload["aud"]:
|
||||
return
|
||||
# Application did not specify an audience, but
|
||||
# the token has the 'aud' claim
|
||||
raise InvalidAudienceError("Invalid audience")
|
||||
|
||||
if "aud" not in payload or not payload["aud"]:
|
||||
# Application specified an audience, but it could not be
|
||||
# verified since the token does not contain a claim.
|
||||
raise MissingRequiredClaimError("aud")
|
||||
|
||||
audience_claims = payload["aud"]
|
||||
|
||||
# In strict mode, we forbid list matching: the supplied audience
|
||||
# must be a string, and it must exactly match the audience claim.
|
||||
if strict:
|
||||
# Only a single audience is allowed in strict mode.
|
||||
if not isinstance(audience, str):
|
||||
raise InvalidAudienceError("Invalid audience (strict)")
|
||||
|
||||
# Only a single audience claim is allowed in strict mode.
|
||||
if not isinstance(audience_claims, str):
|
||||
raise InvalidAudienceError("Invalid claim format in token (strict)")
|
||||
|
||||
if audience != audience_claims:
|
||||
raise InvalidAudienceError("Audience doesn't match (strict)")
|
||||
|
||||
return
|
||||
|
||||
if isinstance(audience_claims, str):
|
||||
audience_claims = [audience_claims]
|
||||
if not isinstance(audience_claims, list):
|
||||
raise InvalidAudienceError("Invalid claim format in token")
|
||||
if any(not isinstance(c, str) for c in audience_claims):
|
||||
raise InvalidAudienceError("Invalid claim format in token")
|
||||
|
||||
if isinstance(audience, str):
|
||||
audience = [audience]
|
||||
|
||||
if all(aud not in audience_claims for aud in audience):
|
||||
raise InvalidAudienceError("Audience doesn't match")
|
||||
|
||||
def _validate_iss(
|
||||
self, payload: dict[str, Any], issuer: Container[str] | str | None
|
||||
) -> None:
|
||||
if issuer is None:
|
||||
return
|
||||
|
||||
if "iss" not in payload:
|
||||
raise MissingRequiredClaimError("iss")
|
||||
|
||||
iss = payload["iss"]
|
||||
if not isinstance(iss, str):
|
||||
raise InvalidIssuerError("Payload Issuer (iss) must be a string")
|
||||
|
||||
if isinstance(issuer, str):
|
||||
if iss != issuer:
|
||||
raise InvalidIssuerError("Invalid issuer")
|
||||
else:
|
||||
try:
|
||||
if iss not in issuer:
|
||||
raise InvalidIssuerError("Invalid issuer")
|
||||
except TypeError:
|
||||
raise InvalidIssuerError(
|
||||
'Issuer param must be "str" or "Container[str]"'
|
||||
) from None
|
||||
|
||||
|
||||
_jwt_global_obj = PyJWT()
|
||||
_jwt_global_obj._jws = _jws_global_obj
|
||||
encode = _jwt_global_obj.encode
|
||||
decode_complete = _jwt_global_obj.decode_complete
|
||||
decode = _jwt_global_obj.decode
|
||||
@@ -0,0 +1,113 @@
|
||||
class PyJWTError(Exception):
|
||||
"""
|
||||
Base class for all exceptions
|
||||
"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class InvalidTokenError(PyJWTError):
|
||||
"""Base exception when ``decode()`` fails on a token"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class DecodeError(InvalidTokenError):
|
||||
"""Raised when a token cannot be decoded because it failed validation"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class InvalidSignatureError(DecodeError):
|
||||
"""Raised when a token's signature doesn't match the one provided as part of
|
||||
the token."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class ExpiredSignatureError(InvalidTokenError):
|
||||
"""Raised when a token's ``exp`` claim indicates that it has expired"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class InvalidAudienceError(InvalidTokenError):
|
||||
"""Raised when a token's ``aud`` claim does not match one of the expected
|
||||
audience values"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class InvalidIssuerError(InvalidTokenError):
|
||||
"""Raised when a token's ``iss`` claim does not match the expected issuer"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class InvalidIssuedAtError(InvalidTokenError):
|
||||
"""Raised when a token's ``iat`` claim is non-numeric"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class ImmatureSignatureError(InvalidTokenError):
|
||||
"""Raised when a token's ``nbf`` or ``iat`` claims represent a time in the future"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class InvalidKeyError(PyJWTError):
|
||||
"""Raised when the specified key is not in the proper format"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class InvalidAlgorithmError(InvalidTokenError):
|
||||
"""Raised when the specified algorithm is not recognized by PyJWT"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class MissingRequiredClaimError(InvalidTokenError):
|
||||
"""Raised when a claim that is required to be present is not contained
|
||||
in the claimset"""
|
||||
|
||||
def __init__(self, claim: str) -> None:
|
||||
self.claim = claim
|
||||
|
||||
def __str__(self) -> str:
|
||||
return f'Token is missing the "{self.claim}" claim'
|
||||
|
||||
|
||||
class PyJWKError(PyJWTError):
|
||||
pass
|
||||
|
||||
|
||||
class MissingCryptographyError(PyJWKError):
|
||||
"""Raised if the algorithm requires ``cryptography`` to be installed and it is not available."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class PyJWKSetError(PyJWTError):
|
||||
pass
|
||||
|
||||
|
||||
class PyJWKClientError(PyJWTError):
|
||||
pass
|
||||
|
||||
|
||||
class PyJWKClientConnectionError(PyJWKClientError):
|
||||
pass
|
||||
|
||||
|
||||
class InvalidSubjectError(InvalidTokenError):
|
||||
"""Raised when a token's ``sub`` claim is not a string or doesn't match the expected ``subject``"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class InvalidJTIError(InvalidTokenError):
|
||||
"""Raised when a token's ``jti`` claim is not a string"""
|
||||
|
||||
pass
|
||||
@@ -0,0 +1,66 @@
|
||||
import json
|
||||
import platform
|
||||
import sys
|
||||
|
||||
from . import __version__ as pyjwt_version
|
||||
|
||||
try:
|
||||
import cryptography
|
||||
|
||||
cryptography_version = cryptography.__version__
|
||||
except ModuleNotFoundError:
|
||||
cryptography_version = ""
|
||||
|
||||
|
||||
def info() -> dict[str, dict[str, str]]:
|
||||
"""
|
||||
Generate information for a bug report.
|
||||
Based on the requests package help utility module.
|
||||
"""
|
||||
try:
|
||||
platform_info = {
|
||||
"system": platform.system(),
|
||||
"release": platform.release(),
|
||||
}
|
||||
except OSError:
|
||||
platform_info = {"system": "Unknown", "release": "Unknown"}
|
||||
|
||||
implementation = platform.python_implementation()
|
||||
|
||||
if implementation == "CPython":
|
||||
implementation_version = platform.python_version()
|
||||
elif implementation == "PyPy":
|
||||
pypy_version_info = sys.pypy_version_info # type: ignore[attr-defined]
|
||||
implementation_version = (
|
||||
f"{pypy_version_info.major}."
|
||||
f"{pypy_version_info.minor}."
|
||||
f"{pypy_version_info.micro}"
|
||||
)
|
||||
if pypy_version_info.releaselevel != "final":
|
||||
implementation_version = "".join(
|
||||
[
|
||||
implementation_version,
|
||||
pypy_version_info.releaselevel,
|
||||
]
|
||||
)
|
||||
else:
|
||||
implementation_version = "Unknown"
|
||||
|
||||
return {
|
||||
"platform": platform_info,
|
||||
"implementation": {
|
||||
"name": implementation,
|
||||
"version": implementation_version,
|
||||
},
|
||||
"cryptography": {"version": cryptography_version},
|
||||
"pyjwt": {"version": pyjwt_version},
|
||||
}
|
||||
|
||||
|
||||
def main() -> None:
|
||||
"""Pretty-print the bug information as JSON."""
|
||||
print(json.dumps(info(), sort_keys=True, indent=2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,31 @@
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from .api_jwk import PyJWKSet, PyJWTSetWithTimestamp
|
||||
|
||||
|
||||
class JWKSetCache:
|
||||
def __init__(self, lifespan: float) -> None:
|
||||
self.jwk_set_with_timestamp: Optional[PyJWTSetWithTimestamp] = None
|
||||
self.lifespan = lifespan
|
||||
|
||||
def put(self, jwk_set: PyJWKSet) -> None:
|
||||
if jwk_set is not None:
|
||||
self.jwk_set_with_timestamp = PyJWTSetWithTimestamp(jwk_set)
|
||||
else:
|
||||
# clear cache
|
||||
self.jwk_set_with_timestamp = None
|
||||
|
||||
def get(self) -> Optional[PyJWKSet]:
|
||||
if self.jwk_set_with_timestamp is None or self.is_expired():
|
||||
return None
|
||||
|
||||
return self.jwk_set_with_timestamp.get_jwk_set()
|
||||
|
||||
def is_expired(self) -> bool:
|
||||
return (
|
||||
self.jwk_set_with_timestamp is not None
|
||||
and self.lifespan > -1
|
||||
and time.monotonic()
|
||||
> self.jwk_set_with_timestamp.get_timestamp() + self.lifespan
|
||||
)
|
||||
@@ -0,0 +1,246 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import urllib.request
|
||||
from functools import lru_cache
|
||||
from ssl import SSLContext
|
||||
from typing import Any
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from .api_jwk import PyJWK, PyJWKSet
|
||||
from .api_jwt import decode_complete as decode_token
|
||||
from .exceptions import PyJWKClientConnectionError, PyJWKClientError
|
||||
from .jwk_set_cache import JWKSetCache
|
||||
|
||||
|
||||
class PyJWKClient:
|
||||
def __init__(
|
||||
self,
|
||||
uri: str,
|
||||
cache_keys: bool = False,
|
||||
max_cached_keys: int = 16,
|
||||
cache_jwk_set: bool = True,
|
||||
lifespan: float = 300,
|
||||
headers: dict[str, Any] | None = None,
|
||||
timeout: float = 30,
|
||||
ssl_context: SSLContext | None = None,
|
||||
):
|
||||
"""A client for retrieving signing keys from a JWKS endpoint.
|
||||
|
||||
``PyJWKClient`` uses a two-tier caching system to avoid unnecessary
|
||||
network requests:
|
||||
|
||||
**Tier 1 — JWK Set cache** (enabled by default):
|
||||
Caches the entire JSON Web Key Set response from the endpoint.
|
||||
Controlled by:
|
||||
|
||||
- ``cache_jwk_set``: Set to ``True`` (the default) to enable this
|
||||
cache. When enabled, the JWK Set is fetched from the network only
|
||||
when the cache is empty or expired.
|
||||
- ``lifespan``: Time in seconds before the cached JWK Set expires.
|
||||
Defaults to ``300`` (5 minutes). Must be greater than 0.
|
||||
|
||||
**Tier 2 — Signing key cache** (disabled by default):
|
||||
Caches individual signing keys (looked up by ``kid``) using an LRU
|
||||
cache with **no time-based expiration**. Keys are evicted only when
|
||||
the cache reaches its maximum size. Controlled by:
|
||||
|
||||
- ``cache_keys``: Set to ``True`` to enable this cache.
|
||||
Defaults to ``False``.
|
||||
- ``max_cached_keys``: Maximum number of signing keys to keep in
|
||||
the LRU cache. Defaults to ``16``.
|
||||
|
||||
:param uri: The URL of the JWKS endpoint.
|
||||
:type uri: str
|
||||
:param cache_keys: Enable the per-key LRU cache (Tier 2).
|
||||
:type cache_keys: bool
|
||||
:param max_cached_keys: Max entries in the signing key LRU cache.
|
||||
:type max_cached_keys: int
|
||||
:param cache_jwk_set: Enable the JWK Set response cache (Tier 1).
|
||||
:type cache_jwk_set: bool
|
||||
:param lifespan: TTL in seconds for the JWK Set cache.
|
||||
:type lifespan: float
|
||||
:param headers: Optional HTTP headers to include in requests.
|
||||
:type headers: dict or None
|
||||
:param timeout: HTTP request timeout in seconds.
|
||||
:type timeout: float
|
||||
:param ssl_context: Optional SSL context for the request.
|
||||
:type ssl_context: ssl.SSLContext or None
|
||||
"""
|
||||
if headers is None:
|
||||
headers = {}
|
||||
# urllib's default OpenerDirector also handles file://, ftp://, and
|
||||
# data: URIs. Reject anything that isn't http(s) eagerly so a caller
|
||||
# passing an attacker-influenced URL (e.g. taken from a `jku` token
|
||||
# header) can't read local files or reach other unintended schemes.
|
||||
scheme = urlparse(uri).scheme.lower()
|
||||
if scheme not in ("http", "https"):
|
||||
raise PyJWKClientError(
|
||||
f"Invalid JWKS URI scheme {scheme!r}: only 'http' and 'https' "
|
||||
f"are supported."
|
||||
)
|
||||
self.uri = uri
|
||||
self.jwk_set_cache: JWKSetCache | None = None
|
||||
self.headers = headers
|
||||
self.timeout = timeout
|
||||
self.ssl_context = ssl_context
|
||||
|
||||
if cache_jwk_set:
|
||||
# Init jwt set cache with default or given lifespan.
|
||||
# Default lifespan is 300 seconds (5 minutes).
|
||||
if lifespan <= 0:
|
||||
raise PyJWKClientError(
|
||||
f'Lifespan must be greater than 0, the input is "{lifespan}"'
|
||||
)
|
||||
self.jwk_set_cache = JWKSetCache(lifespan)
|
||||
else:
|
||||
self.jwk_set_cache = None
|
||||
|
||||
if cache_keys:
|
||||
# Cache signing keys
|
||||
get_signing_key = lru_cache(maxsize=max_cached_keys)(self.get_signing_key)
|
||||
# Ignore mypy (https://github.com/python/mypy/issues/2427)
|
||||
self.get_signing_key = get_signing_key # type: ignore[method-assign]
|
||||
|
||||
def fetch_data(self) -> Any:
|
||||
"""Fetch the JWK Set from the JWKS endpoint.
|
||||
|
||||
Makes an HTTP request to the configured ``uri`` and returns the
|
||||
parsed JSON response. If the JWK Set cache is enabled, the
|
||||
response is stored in the cache.
|
||||
|
||||
:returns: The parsed JWK Set as a dictionary.
|
||||
:raises PyJWKClientConnectionError: If the HTTP request fails.
|
||||
"""
|
||||
try:
|
||||
r = urllib.request.Request(url=self.uri, headers=self.headers)
|
||||
with urllib.request.urlopen(
|
||||
r, timeout=self.timeout, context=self.ssl_context
|
||||
) as response:
|
||||
jwk_set = json.load(response)
|
||||
except (URLError, TimeoutError) as e:
|
||||
if isinstance(e, HTTPError):
|
||||
e.close()
|
||||
raise PyJWKClientConnectionError(
|
||||
f'Fail to fetch data from the url, err: "{e}"'
|
||||
) from e
|
||||
|
||||
# Only update the cache on a successful fetch. Writing in a
|
||||
# `finally` block with `jwk_set=None` on error clears any
|
||||
# previously-cached JWKS, turning a transient outage into a cache
|
||||
# wipe that breaks legitimate auth.
|
||||
if self.jwk_set_cache is not None:
|
||||
self.jwk_set_cache.put(jwk_set)
|
||||
return jwk_set
|
||||
|
||||
def get_jwk_set(self, refresh: bool = False) -> PyJWKSet:
|
||||
"""Return the JWK Set, using the cache when available.
|
||||
|
||||
:param refresh: Force a fresh fetch from the endpoint, bypassing
|
||||
the cache.
|
||||
:type refresh: bool
|
||||
:returns: The JWK Set.
|
||||
:rtype: PyJWKSet
|
||||
:raises PyJWKClientError: If the endpoint does not return a JSON
|
||||
object.
|
||||
"""
|
||||
data = None
|
||||
if self.jwk_set_cache is not None and not refresh:
|
||||
data = self.jwk_set_cache.get()
|
||||
|
||||
if data is None:
|
||||
data = self.fetch_data()
|
||||
|
||||
if not isinstance(data, dict):
|
||||
raise PyJWKClientError("The JWKS endpoint did not return a JSON object")
|
||||
|
||||
return PyJWKSet.from_dict(data)
|
||||
|
||||
def get_signing_keys(self, refresh: bool = False) -> list[PyJWK]:
|
||||
"""Return all signing keys from the JWK Set.
|
||||
|
||||
Filters the JWK Set to keys whose ``use`` is ``"sig"`` (or
|
||||
unspecified) and that have a ``kid``.
|
||||
|
||||
:param refresh: Force a fresh fetch from the endpoint, bypassing
|
||||
the cache.
|
||||
:type refresh: bool
|
||||
:returns: A list of signing keys.
|
||||
:rtype: list[PyJWK]
|
||||
:raises PyJWKClientError: If no signing keys are found.
|
||||
"""
|
||||
jwk_set = self.get_jwk_set(refresh)
|
||||
signing_keys = [
|
||||
jwk_set_key
|
||||
for jwk_set_key in jwk_set.keys
|
||||
if jwk_set_key.public_key_use in ["sig", None] and jwk_set_key.key_id
|
||||
]
|
||||
|
||||
if not signing_keys:
|
||||
raise PyJWKClientError("The JWKS endpoint did not contain any signing keys")
|
||||
|
||||
return signing_keys
|
||||
|
||||
def get_signing_key(self, kid: str) -> PyJWK:
|
||||
"""Return the signing key matching the given ``kid``.
|
||||
|
||||
If no match is found in the current JWK Set, the set is
|
||||
refreshed from the endpoint and the lookup is retried once.
|
||||
|
||||
:param kid: The key ID to look up.
|
||||
:type kid: str
|
||||
:returns: The matching signing key.
|
||||
:rtype: PyJWK
|
||||
:raises PyJWKClientError: If no matching key is found after
|
||||
refreshing.
|
||||
"""
|
||||
signing_keys = self.get_signing_keys()
|
||||
signing_key = self.match_kid(signing_keys, kid)
|
||||
|
||||
if not signing_key:
|
||||
# If no matching signing key from the jwk set, refresh the jwk set and try again.
|
||||
signing_keys = self.get_signing_keys(refresh=True)
|
||||
signing_key = self.match_kid(signing_keys, kid)
|
||||
|
||||
if not signing_key:
|
||||
raise PyJWKClientError(
|
||||
f'Unable to find a signing key that matches: "{kid}"'
|
||||
)
|
||||
|
||||
return signing_key
|
||||
|
||||
def get_signing_key_from_jwt(self, token: str | bytes) -> PyJWK:
|
||||
"""Return the signing key for a JWT by reading its ``kid`` header.
|
||||
|
||||
Extracts the ``kid`` from the token's unverified header and
|
||||
delegates to :meth:`get_signing_key`.
|
||||
|
||||
:param token: The encoded JWT.
|
||||
:type token: str or bytes
|
||||
:returns: The matching signing key.
|
||||
:rtype: PyJWK
|
||||
"""
|
||||
unverified = decode_token(token, options={"verify_signature": False})
|
||||
header = unverified["header"]
|
||||
return self.get_signing_key(header.get("kid"))
|
||||
|
||||
@staticmethod
|
||||
def match_kid(signing_keys: list[PyJWK], kid: str) -> PyJWK | None:
|
||||
"""Find a key in *signing_keys* that matches *kid*.
|
||||
|
||||
:param signing_keys: The list of keys to search.
|
||||
:type signing_keys: list[PyJWK]
|
||||
:param kid: The key ID to match.
|
||||
:type kid: str
|
||||
:returns: The matching key, or ``None`` if not found.
|
||||
:rtype: PyJWK or None
|
||||
"""
|
||||
signing_key = None
|
||||
|
||||
for key in signing_keys:
|
||||
if key.key_id == kid:
|
||||
signing_key = key
|
||||
break
|
||||
|
||||
return signing_key
|
||||
@@ -0,0 +1,69 @@
|
||||
from typing import Any, Callable, TypedDict
|
||||
|
||||
JWKDict = dict[str, Any]
|
||||
|
||||
HashlibHash = Callable[..., Any]
|
||||
|
||||
|
||||
class SigOptions(TypedDict, total=False):
|
||||
"""Options for PyJWS class (TypedDict). Note that this is a smaller set of options than
|
||||
for :py:func:`jwt.decode()`."""
|
||||
|
||||
verify_signature: bool
|
||||
"""verify the JWT cryptographic signature"""
|
||||
enforce_minimum_key_length: bool
|
||||
"""Default: ``False``. Raise :py:class:`jwt.exceptions.InvalidKeyError` instead of warning when keys are below minimum recommended length."""
|
||||
|
||||
|
||||
class Options(TypedDict, total=False):
|
||||
"""Options for :py:func:`jwt.decode()` and :py:func:`jwt.decode_complete()` (TypedDict).
|
||||
|
||||
.. warning::
|
||||
|
||||
Some claims, such as ``exp``, ``iat``, ``jti``, ``nbf``, and ``sub``,
|
||||
will only be verified if present. Please refer to the documentation below
|
||||
for which ones, and make sure to include them in the ``require`` param
|
||||
if you want to make sure that they are always present (and therefore always verified
|
||||
if ``verify_{claim} = True`` for that claim).
|
||||
"""
|
||||
|
||||
verify_signature: bool
|
||||
"""Default: ``True``. Verify the JWT cryptographic signature."""
|
||||
require: list[str]
|
||||
"""Default: ``[]``. List of claims that must be present.
|
||||
Example: ``require=["exp", "iat", "nbf"]``.
|
||||
**Only verifies that the claims exists**. Does not verify that the claims are valid."""
|
||||
strict_aud: bool
|
||||
"""Default: ``False``. (requires ``verify_aud=True``) Check that the ``aud`` claim is a single value (not a list), and matches ``audience`` exactly."""
|
||||
verify_aud: bool
|
||||
"""Default: ``verify_signature``. Check that ``aud`` (audience) claim matches ``audience``."""
|
||||
verify_exp: bool
|
||||
"""Default: ``verify_signature``. Check that ``exp`` (expiration) claim value is in the future (if present in payload). """
|
||||
verify_iat: bool
|
||||
"""Default: ``verify_signature``. Check that ``iat`` (issued at) claim value is an integer (if present in payload). """
|
||||
verify_iss: bool
|
||||
"""Default: ``verify_signature``. Check that ``iss`` (issuer) claim matches ``issuer``. """
|
||||
verify_jti: bool
|
||||
"""Default: ``verify_signature``. Check that ``jti`` (JWT ID) claim is a string (if present in payload). """
|
||||
verify_nbf: bool
|
||||
"""Default: ``verify_signature``. Check that ``nbf`` (not before) claim value is in the past (if present in payload). """
|
||||
verify_sub: bool
|
||||
"""Default: ``verify_signature``. Check that ``sub`` (subject) claim is a string and matches ``subject`` (if present in payload). """
|
||||
enforce_minimum_key_length: bool
|
||||
"""Default: ``False``. Raise :py:class:`jwt.exceptions.InvalidKeyError` instead of warning when keys are below minimum recommended length."""
|
||||
|
||||
|
||||
# The only difference between Options and FullOptions is that FullOptions
|
||||
# required _every_ value to be there; Options doesn't require any
|
||||
class FullOptions(TypedDict):
|
||||
verify_signature: bool
|
||||
require: list[str]
|
||||
strict_aud: bool
|
||||
verify_aud: bool
|
||||
verify_exp: bool
|
||||
verify_iat: bool
|
||||
verify_iss: bool
|
||||
verify_jti: bool
|
||||
verify_nbf: bool
|
||||
verify_sub: bool
|
||||
enforce_minimum_key_length: bool
|
||||
@@ -0,0 +1,142 @@
|
||||
import base64
|
||||
import binascii
|
||||
import re
|
||||
from typing import Optional, Union
|
||||
|
||||
try:
|
||||
from cryptography.hazmat.primitives.asymmetric.ec import EllipticCurve
|
||||
from cryptography.hazmat.primitives.asymmetric.utils import (
|
||||
decode_dss_signature,
|
||||
encode_dss_signature,
|
||||
)
|
||||
except ModuleNotFoundError:
|
||||
pass
|
||||
|
||||
|
||||
def force_bytes(value: Union[bytes, str]) -> bytes:
|
||||
if isinstance(value, str):
|
||||
return value.encode("utf-8")
|
||||
elif isinstance(value, bytes):
|
||||
return value
|
||||
else:
|
||||
raise TypeError("Expected a string value")
|
||||
|
||||
|
||||
def base64url_decode(input: Union[bytes, str]) -> bytes:
|
||||
input_bytes = force_bytes(input)
|
||||
|
||||
rem = len(input_bytes) % 4
|
||||
|
||||
if rem > 0:
|
||||
input_bytes += b"=" * (4 - rem)
|
||||
|
||||
return base64.urlsafe_b64decode(input_bytes)
|
||||
|
||||
|
||||
def base64url_encode(input: bytes) -> bytes:
|
||||
return base64.urlsafe_b64encode(input).replace(b"=", b"")
|
||||
|
||||
|
||||
def to_base64url_uint(val: int, *, bit_length: Optional[int] = None) -> bytes:
|
||||
if val < 0:
|
||||
raise ValueError("Must be a positive integer")
|
||||
|
||||
int_bytes = bytes_from_int(val, bit_length=bit_length)
|
||||
|
||||
if len(int_bytes) == 0:
|
||||
int_bytes = b"\x00"
|
||||
|
||||
return base64url_encode(int_bytes)
|
||||
|
||||
|
||||
def from_base64url_uint(val: Union[bytes, str]) -> int:
|
||||
data = base64url_decode(force_bytes(val))
|
||||
return int.from_bytes(data, byteorder="big")
|
||||
|
||||
|
||||
def number_to_bytes(num: int, num_bytes: int) -> bytes:
|
||||
padded_hex = "%0*x" % (2 * num_bytes, num)
|
||||
return binascii.a2b_hex(padded_hex.encode("ascii"))
|
||||
|
||||
|
||||
def bytes_to_number(string: bytes) -> int:
|
||||
return int(binascii.b2a_hex(string), 16)
|
||||
|
||||
|
||||
def bytes_from_int(val: int, *, bit_length: Optional[int] = None) -> bytes:
|
||||
if bit_length is None:
|
||||
bit_length = val.bit_length()
|
||||
byte_length = (bit_length + 7) // 8
|
||||
|
||||
return val.to_bytes(byte_length, "big", signed=False)
|
||||
|
||||
|
||||
def der_to_raw_signature(der_sig: bytes, curve: "EllipticCurve") -> bytes:
|
||||
num_bits = curve.key_size
|
||||
num_bytes = (num_bits + 7) // 8
|
||||
|
||||
r, s = decode_dss_signature(der_sig)
|
||||
|
||||
return number_to_bytes(r, num_bytes) + number_to_bytes(s, num_bytes)
|
||||
|
||||
|
||||
def raw_to_der_signature(raw_sig: bytes, curve: "EllipticCurve") -> bytes:
|
||||
num_bits = curve.key_size
|
||||
num_bytes = (num_bits + 7) // 8
|
||||
|
||||
if len(raw_sig) != 2 * num_bytes:
|
||||
raise ValueError("Invalid signature")
|
||||
|
||||
r = bytes_to_number(raw_sig[:num_bytes])
|
||||
s = bytes_to_number(raw_sig[num_bytes:])
|
||||
|
||||
return bytes(encode_dss_signature(r, s))
|
||||
|
||||
|
||||
# Based on https://github.com/hynek/pem/blob/7ad94db26b0bc21d10953f5dbad3acfdfacf57aa/src/pem/_core.py#L224-L252
|
||||
_PEMS = {
|
||||
b"CERTIFICATE",
|
||||
b"TRUSTED CERTIFICATE",
|
||||
b"PRIVATE KEY",
|
||||
b"PUBLIC KEY",
|
||||
b"ENCRYPTED PRIVATE KEY",
|
||||
b"OPENSSH PRIVATE KEY",
|
||||
b"DSA PRIVATE KEY",
|
||||
b"RSA PRIVATE KEY",
|
||||
b"RSA PUBLIC KEY",
|
||||
b"EC PRIVATE KEY",
|
||||
b"DH PARAMETERS",
|
||||
b"NEW CERTIFICATE REQUEST",
|
||||
b"CERTIFICATE REQUEST",
|
||||
b"SSH2 PUBLIC KEY",
|
||||
b"SSH2 ENCRYPTED PRIVATE KEY",
|
||||
b"X509 CRL",
|
||||
}
|
||||
|
||||
_PEM_RE = re.compile(
|
||||
b"----[- ]BEGIN ("
|
||||
+ b"|".join(_PEMS)
|
||||
+ b""")[- ]----\r?
|
||||
.+?\r?
|
||||
----[- ]END \\1[- ]----\r?\n?""",
|
||||
re.DOTALL,
|
||||
)
|
||||
|
||||
|
||||
def is_pem_format(key: bytes) -> bool:
|
||||
return bool(_PEM_RE.search(key))
|
||||
|
||||
|
||||
# Based on https://github.com/pyca/cryptography/blob/bcb70852d577b3f490f015378c75cba74986297b/src/cryptography/hazmat/primitives/serialization/ssh.py#L40-L46
|
||||
_SSH_KEY_FORMATS = (
|
||||
b"ssh-ed25519",
|
||||
b"ssh-rsa",
|
||||
b"ssh-dss",
|
||||
b"ecdsa-sha2-nistp256",
|
||||
b"ecdsa-sha2-nistp384",
|
||||
b"ecdsa-sha2-nistp521",
|
||||
)
|
||||
|
||||
|
||||
def is_ssh_key(key: bytes) -> bool:
|
||||
return key.startswith(_SSH_KEY_FORMATS)
|
||||
@@ -0,0 +1,11 @@
|
||||
class RemovedInPyjwt3Warning(DeprecationWarning):
|
||||
"""Warning for features that will be removed in PyJWT 3."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class InsecureKeyLengthWarning(UserWarning):
|
||||
"""Warning emitted when a cryptographic key is shorter than the minimum
|
||||
recommended length. See :ref:`key-length-validation` for details."""
|
||||
|
||||
pass
|
||||
Reference in New Issue
Block a user