Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 11 additions & 8 deletions src/auth0_api_python/api_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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}")

Expand Down
3 changes: 3 additions & 0 deletions src/auth0_api_python/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down Expand Up @@ -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
Expand Down
17 changes: 13 additions & 4 deletions src/auth0_api_python/token_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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

Expand Down
16 changes: 16 additions & 0 deletions tests/test_api_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
"""
Expand Down