From ff4c63f00d77423a083e2c264523a13f6a52e0f5 Mon Sep 17 00:00:00 2001 From: Samuel Bloch Date: Tue, 29 Sep 2026 08:41:55 -0400 Subject: [PATCH] chore: Resolves Issue #92, updates to use transitive dependency explicit and removes the deprecation warning for the authlib.jose library.Signed --- src/auth0_api_python/api_client.py | 19 +++++++++++-------- src/auth0_api_python/config.py | 3 +++ src/auth0_api_python/token_utils.py | 17 +++++++++++++---- tests/test_api_client.py | 16 ++++++++++++++++ 4 files changed, 43 insertions(+), 12 deletions(-) diff --git a/src/auth0_api_python/api_client.py b/src/auth0_api_python/api_client.py index a87ef14..d41d262 100644 --- a/src/auth0_api_python/api_client.py +++ b/src/auth0_api_python/api_client.py @@ -4,7 +4,7 @@ from typing import Any, Optional, Union import httpx -from authlib.jose import JsonWebKey, JsonWebToken +from joserfc import jwk, jwt from .cache import InMemoryCache from .config import ApiClientOptions @@ -60,6 +60,12 @@ def __init__(self, options: ApiClientOptions): if not options.audience: raise MissingRequiredArgumentError("audience") + if not isinstance(options.jwt_algorithms, list) or not options.jwt_algorithms or not all( + isinstance(algorithm, str) and algorithm for algorithm in options.jwt_algorithms + ): + raise ConfigurationError("jwt_algorithms must be a non-empty list of algorithm names") + self._jwt_algorithms = options.jwt_algorithms + # Validate domains parameter if provided if options.domains is not None: if isinstance(options.domains, list): @@ -113,10 +119,7 @@ def __init__(self, options: ApiClientOptions): self._cache_ttl = options.cache_ttl_seconds - self._jwt = JsonWebToken(["RS256"]) - self._dpop_algorithms = ["ES256"] - self._dpop_jwt = JsonWebToken(self._dpop_algorithms) def is_dpop_required(self) -> bool: """Check if DPoP authentication is required.""" @@ -524,12 +527,12 @@ async def verify_access_token( raise VerifyAccessTokenError("No matching key found in JWKS") # Import public key and verify signature - public_key = JsonWebKey.import_key(matching_key_dict) + public_key = jwk.import_key(matching_key_dict) if isinstance(access_token, str) and access_token.startswith("b'"): access_token = access_token[2:-1] try: - claims = self._jwt.decode(access_token, public_key) + claims = jwt.decode(access_token, public_key, algorithms=self._jwt_algorithms).claims except Exception as e: raise VerifyAccessTokenError(f"Signature verification failed: {str(e)}") from e @@ -606,9 +609,9 @@ async def verify_dpop_proof( if jwk_dict.get("crv") != "P-256": raise InvalidDpopProofError("Only P-256 curve is supported") - public_key = JsonWebKey.import_key(jwk_dict) + public_key = jwk.import_key(jwk_dict) try: - claims = self._dpop_jwt.decode(proof, public_key) + claims = jwt.decode(proof, public_key, algorithms=self._dpop_algorithms).claims except Exception as e: raise InvalidDpopProofError(f"JWT signature verification failed: {e}") diff --git a/src/auth0_api_python/config.py b/src/auth0_api_python/config.py index 1929d55..e7ad471 100644 --- a/src/auth0_api_python/config.py +++ b/src/auth0_api_python/config.py @@ -19,6 +19,7 @@ class ApiClientOptions: Can be a static list of domain strings or a callable that returns allowed domains dynamically. Optional if domain is provided. audience: The expected 'aud' claim in the token. + jwt_algorithms: Allowed access-token signing algorithms (default: ["RS256"]). custom_fetch: Optional callable that can replace the default HTTP fetch logic. cache_adapter: Custom cache implementation. If not provided, uses default InMemoryCache. cache_ttl_seconds: Time-to-live for cache entries in seconds (default: 600 = 10 minutes). @@ -49,10 +50,12 @@ def __init__( client_id: Optional[str] = None, client_secret: Optional[str] = None, timeout: float = 10.0, + jwt_algorithms: Optional[list[str]] = None, ): self.domain = domain self.domains = domains self.audience = audience + self.jwt_algorithms = ["RS256"] if jwt_algorithms is None else jwt_algorithms self.custom_fetch = custom_fetch self.cache_adapter = cache_adapter self.cache_ttl_seconds = cache_ttl_seconds diff --git a/src/auth0_api_python/token_utils.py b/src/auth0_api_python/token_utils.py index c234681..663f588 100644 --- a/src/auth0_api_python/token_utils.py +++ b/src/auth0_api_python/token_utils.py @@ -2,7 +2,9 @@ import uuid from typing import Any, Optional, Union -from authlib.jose import JsonWebKey, jwt +from joserfc import jwk, jwt +from joserfc.jws import JWSRegistry +from joserfc.registry import HeaderParameter from .utils import calculate_jwk_thumbprint, normalize_url_for_htu, sha256_base64url @@ -80,7 +82,7 @@ async def generate_token( token_claims["aud"] = audience - key = JsonWebKey.import_key(PRIVATE_JWK) + key = jwk.import_key(PRIVATE_JWK) header = {"alg": "RS256", "kid": PRIVATE_JWK["kid"]} token = jwt.encode(header, token_claims, key) @@ -166,8 +168,15 @@ async def generate_dpop_proof( if header_overrides: header.update(header_overrides) - key = JsonWebKey.import_key(PRIVATE_EC_JWK) - token = jwt.encode(header, proof_claims, key) + key = jwk.import_key(PRIVATE_EC_JWK) + registry = JWSRegistry( + header_registry={ + name: HeaderParameter("test override", lambda _: None) + for name in ("typ", "jwk") + if name in (header_overrides or {}) + } + ) + token = jwt.encode(header, proof_claims, key, registry=registry) # Ensure we return a string, not bytes return token.decode('utf-8') if isinstance(token, bytes) else token diff --git a/tests/test_api_client.py b/tests/test_api_client.py index 9931c39..d60ec31 100644 --- a/tests/test_api_client.py +++ b/tests/test_api_client.py @@ -65,6 +65,22 @@ async def test_init_missing_args(): _ = ApiClient(ApiClientOptions(domain="example.us.auth0.com", audience="")) +def test_custom_jwt_algorithms(): + client = ApiClient(ApiClientOptions( + domain="example.us.auth0.com", + audience="my-audience", + jwt_algorithms=["RS384", "RS512"], + )) + + assert client._jwt_algorithms == ["RS384", "RS512"] + with pytest.raises(ConfigurationError): + ApiClient(ApiClientOptions( + domain="example.us.auth0.com", + audience="my-audience", + jwt_algorithms=[], + )) + + @pytest.mark.asyncio async def test_verify_access_token_successfully(httpx_mock: HTTPXMock): """