Coverage for server / authentication / auth0_handler.py: 95%
79 statements
« prev ^ index » next coverage.py v7.13.4, created at 2026-10-04 09:33 +0000
« prev ^ index » next coverage.py v7.13.4, created at 2026-10-04 09:33 +0000
1"""
2Auth0 Token Handler
4Handles Auth0 JWT token validation and role extraction using RS256 asymmetric verification.
5Fetches Auth0's public keys (JWKS) and validates tokens against them.
7Usage:
8 from server.authentication.auth0_handler import auth0_handler
10 payload = await auth0_handler.verify_token(token)
11 user_info = auth0_handler.extract_user_info(payload)
12"""
14import asyncio
15import logging
16import time
18import httpx
19from jose import jwt, JWTError
20from typing import Dict, Optional
21from fastapi import HTTPException, status
22from .auth0_config import get_auth0_settings
24logger = logging.getLogger(__name__)
27class Auth0TokenHandler:
28 """
29 Handles Auth0 JWT token validation using RS256 asymmetric verification.
31 This class fetches Auth0's JSON Web Key Set (JWKS) containing public keys
32 and uses them to verify JWT tokens signed by Auth0's private key.
34 Attributes:
35 settings: Auth0 configuration settings
36 _jwks_cache: Cached JWKS to minimize API calls to Auth0
37 """
39 JWKS_CACHE_TTL = 21600 # 6 hours in seconds
41 def __init__(self):
42 """Initialize Auth0 token handler with settings."""
43 self.settings = get_auth0_settings()
44 self._jwks_cache: Optional[Dict] = None
45 self._jwks_fetched_at: Optional[float] = None
46 self._jwks_lock = asyncio.Lock()
48 async def get_jwks(self) -> Dict:
49 """
50 Fetch Auth0 JSON Web Key Set (JWKS) for token verification.
52 The JWKS contains Auth0's public keys used to verify JWT signatures.
53 Results are cached to minimize API calls to Auth0.
55 Returns:
56 Dict: JWKS containing public keys
58 Raises:
59 HTTPException: If JWKS fetch fails
61 Example JWKS:
62 {
63 "keys": [
64 {
65 "kty": "RSA",
66 "kid": "abc123",
67 "use": "sig",
68 "n": "base64_modulus",
69 "e": "base64_exponent"
70 }
71 ]
72 }
73 """
74 async with self._jwks_lock:
75 # Double-check cache inside lock to prevent thundering herd
76 if (
77 self._jwks_cache
78 and self._jwks_fetched_at
79 and (time.time() - self._jwks_fetched_at) < self.JWKS_CACHE_TTL
80 ):
81 return self._jwks_cache
83 jwks_url = f"https://{self.settings.AUTH0_DOMAIN}/.well-known/jwks.json"
85 try:
86 async with httpx.AsyncClient() as client:
87 response = await client.get(jwks_url, timeout=10.0)
88 response.raise_for_status()
89 self._jwks_cache = response.json()
90 self._jwks_fetched_at = time.time()
91 return self._jwks_cache
92 except httpx.HTTPError as e:
93 logger.error(f"Failed to fetch Auth0 JWKS: {e}")
94 raise HTTPException(
95 status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
96 detail="Authentication service temporarily unavailable"
97 )
99 def invalidate_jwks_cache(self) -> None:
100 """Clear the JWKS cache, forcing a refetch on next verification."""
101 self._jwks_cache = None
102 self._jwks_fetched_at = None
104 def get_signing_key(self, token: str, jwks: Dict) -> str:
105 """
106 Extract the signing key from JWKS based on token's kid (key ID).
108 The token header contains a 'kid' field identifying which public key
109 from the JWKS should be used to verify the signature.
111 Args:
112 token: JWT token string
113 jwks: JSON Web Key Set from Auth0
115 Returns:
116 str: PEM-formatted public key for verification
118 Raises:
119 HTTPException: If kid is missing or signing key not found
121 Example:
122 >>> key = handler.get_signing_key(token, jwks)
123 >>> print(key)
124 -----BEGIN PUBLIC KEY-----
125 MIIBIjANBgkqhki...
126 -----END PUBLIC KEY-----
127 """
128 try:
129 # Get the kid from token header
130 unverified_header = jwt.get_unverified_header(token)
131 kid = unverified_header.get("kid")
133 if not kid:
134 raise HTTPException(
135 status_code=status.HTTP_401_UNAUTHORIZED,
136 detail="Token missing key ID (kid)"
137 )
139 # Find matching key in JWKS
140 for key in jwks.get("keys", []):
141 if key.get("kid") == kid:
142 # Construct PEM key from modulus and exponent
143 from cryptography.hazmat.primitives.asymmetric import rsa
144 from cryptography.hazmat.primitives import serialization
145 from cryptography.hazmat.backends import default_backend
146 import base64
148 # Decode base64url-encoded modulus and exponent
149 n_bytes = key["n"]
150 e_bytes = key["e"]
152 # Add padding if necessary for base64 decoding
153 n_bytes += "=" * (4 - len(n_bytes) % 4)
154 e_bytes += "=" * (4 - len(e_bytes) % 4)
156 n = int.from_bytes(
157 base64.urlsafe_b64decode(n_bytes),
158 byteorder="big"
159 )
160 e = int.from_bytes(
161 base64.urlsafe_b64decode(e_bytes),
162 byteorder="big"
163 )
165 # Construct RSA public key
166 public_numbers = rsa.RSAPublicNumbers(e, n)
167 public_key = public_numbers.public_key(default_backend())
169 # Convert to PEM format
170 pem = public_key.public_bytes(
171 encoding=serialization.Encoding.PEM,
172 format=serialization.PublicFormat.SubjectPublicKeyInfo
173 )
175 return pem.decode("utf-8")
177 # No matching key found
178 raise HTTPException(
179 status_code=status.HTTP_401_UNAUTHORIZED,
180 detail="Unable to find appropriate signing key in JWKS"
181 )
183 except Exception as e:
184 if isinstance(e, HTTPException):
185 raise
186 logger.error(f"Error extracting signing key: {e}")
187 raise HTTPException(
188 status_code=status.HTTP_401_UNAUTHORIZED,
189 detail="Invalid authentication credentials"
190 )
192 async def verify_token(self, token: str) -> Dict:
193 """
194 Verify and decode Auth0 JWT token.
196 This method performs comprehensive token validation:
197 1. Fetches Auth0's public keys (JWKS)
198 2. Extracts the appropriate signing key
199 3. Verifies token signature using RS256
200 4. Validates audience (API identifier)
201 5. Validates issuer (Auth0 domain)
202 6. Checks token expiration
204 Args:
205 token: JWT token string from Authorization header
207 Returns:
208 Dict: Decoded token payload containing user info and claims
210 Raises:
211 HTTPException (401): If token is invalid, expired, or verification fails
213 Example payload:
214 {
215 "iss": "https://tenant.auth0.com/",
216 "sub": "auth0|123456",
217 "aud": "https://api.eruditiontx.com",
218 "exp": 1234567890,
219 "iat": 1234567890,
220 "email": "user@example.com",
221 "https://eruditiontx.com/roles": ["teacher"]
222 }
223 """
224 try:
225 # Get JWKS from Auth0
226 jwks = await self.get_jwks()
228 # Get signing key for this token
229 signing_key = self.get_signing_key(token, jwks)
231 # Verify and decode token
232 payload = jwt.decode(
233 token,
234 signing_key,
235 algorithms=self.settings.AUTH0_ALGORITHMS,
236 audience=self.settings.AUTH0_API_IDENTIFIER,
237 issuer=self.settings.AUTH0_ISSUER,
238 )
240 return payload
242 except JWTError as e:
243 logger.error(f"Auth0 token validation failed: {e}")
244 raise HTTPException(
245 status_code=status.HTTP_401_UNAUTHORIZED,
246 detail="Invalid authentication credentials"
247 )
248 except HTTPException:
249 # Re-raise HTTPExceptions from get_jwks or get_signing_key
250 raise
251 except Exception as e:
252 raise HTTPException(
253 status_code=status.HTTP_401_UNAUTHORIZED,
254 detail="Could not validate credentials"
255 )
257 def extract_user_info(self, payload: Dict) -> Dict:
258 """
259 Extract user information from verified token payload.
261 Parses the JWT payload to extract user identity, email, roles,
262 and metadata into a standardized format.
264 Args:
265 payload: Decoded JWT token payload from verify_token()
267 Returns:
268 Dict: Standardized user information containing:
269 - sub: Auth0 user ID (e.g., "auth0|123456")
270 - email: User email address
271 - email_verified: Email verification status
272 - name: User's full name
273 - roles: List of role strings from custom claim
274 - auth0_metadata: Token metadata (issuer, audience, exp, iat)
276 Example:
277 >>> user_info = handler.extract_user_info(payload)
278 >>> print(user_info)
279 {
280 "sub": "auth0|123456",
281 "email": "teacher@example.com",
282 "email_verified": True,
283 "name": "John Doe",
284 "roles": ["teacher"],
285 "auth0_metadata": {
286 "iss": "https://tenant.auth0.com/",
287 "aud": "https://api.eruditiontx.com",
288 "exp": 1234567890,
289 "iat": 1234567890
290 }
291 }
292 """
293 # Extract roles from custom claim namespace
294 roles = payload.get(self.settings.AUTH0_ROLES_CLAIM, [])
296 return {
297 "sub": payload.get("sub"), # Auth0 user ID
298 "email": payload.get("email"),
299 "email_verified": payload.get("email_verified", False),
300 "name": payload.get("name"),
301 "roles": roles,
302 "auth0_metadata": {
303 "iss": payload.get("iss"),
304 "aud": payload.get("aud"),
305 "exp": payload.get("exp"),
306 "iat": payload.get("iat"),
307 }
308 }
311# Singleton instance for use across the application
312auth0_handler = Auth0TokenHandler()