|
4 | 4 | from typing import Any, Optional, Union |
5 | 5 |
|
6 | 6 | import httpx |
7 | | -from authlib.jose import JsonWebKey, JsonWebToken |
| 7 | +from joserfc import jwk, jwt |
8 | 8 |
|
9 | 9 | from .cache import InMemoryCache |
10 | 10 | from .config import ApiClientOptions |
@@ -60,6 +60,12 @@ def __init__(self, options: ApiClientOptions): |
60 | 60 | if not options.audience: |
61 | 61 | raise MissingRequiredArgumentError("audience") |
62 | 62 |
|
| 63 | + if not isinstance(options.jwt_algorithms, list) or not options.jwt_algorithms or not all( |
| 64 | + isinstance(algorithm, str) and algorithm for algorithm in options.jwt_algorithms |
| 65 | + ): |
| 66 | + raise ConfigurationError("jwt_algorithms must be a non-empty list of algorithm names") |
| 67 | + self._jwt_algorithms = options.jwt_algorithms |
| 68 | + |
63 | 69 | # Validate domains parameter if provided |
64 | 70 | if options.domains is not None: |
65 | 71 | if isinstance(options.domains, list): |
@@ -113,10 +119,7 @@ def __init__(self, options: ApiClientOptions): |
113 | 119 |
|
114 | 120 | self._cache_ttl = options.cache_ttl_seconds |
115 | 121 |
|
116 | | - self._jwt = JsonWebToken(["RS256"]) |
117 | | - |
118 | 122 | self._dpop_algorithms = ["ES256"] |
119 | | - self._dpop_jwt = JsonWebToken(self._dpop_algorithms) |
120 | 123 |
|
121 | 124 | def is_dpop_required(self) -> bool: |
122 | 125 | """Check if DPoP authentication is required.""" |
@@ -524,12 +527,12 @@ async def verify_access_token( |
524 | 527 | raise VerifyAccessTokenError("No matching key found in JWKS") |
525 | 528 |
|
526 | 529 | # Import public key and verify signature |
527 | | - public_key = JsonWebKey.import_key(matching_key_dict) |
| 530 | + public_key = jwk.import_key(matching_key_dict) |
528 | 531 |
|
529 | 532 | if isinstance(access_token, str) and access_token.startswith("b'"): |
530 | 533 | access_token = access_token[2:-1] |
531 | 534 | try: |
532 | | - claims = self._jwt.decode(access_token, public_key) |
| 535 | + claims = jwt.decode(access_token, public_key, algorithms=self._jwt_algorithms).claims |
533 | 536 | except Exception as e: |
534 | 537 | raise VerifyAccessTokenError(f"Signature verification failed: {str(e)}") from e |
535 | 538 |
|
@@ -606,9 +609,9 @@ async def verify_dpop_proof( |
606 | 609 | if jwk_dict.get("crv") != "P-256": |
607 | 610 | raise InvalidDpopProofError("Only P-256 curve is supported") |
608 | 611 |
|
609 | | - public_key = JsonWebKey.import_key(jwk_dict) |
| 612 | + public_key = jwk.import_key(jwk_dict) |
610 | 613 | try: |
611 | | - claims = self._dpop_jwt.decode(proof, public_key) |
| 614 | + claims = jwt.decode(proof, public_key, algorithms=self._dpop_algorithms).claims |
612 | 615 | except Exception as e: |
613 | 616 | raise InvalidDpopProofError(f"JWT signature verification failed: {e}") |
614 | 617 |
|
|
0 commit comments