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

1""" 

2Auth0 Token Handler 

3 

4Handles Auth0 JWT token validation and role extraction using RS256 asymmetric verification. 

5Fetches Auth0's public keys (JWKS) and validates tokens against them. 

6 

7Usage: 

8 from server.authentication.auth0_handler import auth0_handler 

9 

10 payload = await auth0_handler.verify_token(token) 

11 user_info = auth0_handler.extract_user_info(payload) 

12""" 

13 

14import asyncio 

15import logging 

16import time 

17 

18import httpx 

19from jose import jwt, JWTError 

20from typing import Dict, Optional 

21from fastapi import HTTPException, status 

22from .auth0_config import get_auth0_settings 

23 

24logger = logging.getLogger(__name__) 

25 

26 

27class Auth0TokenHandler: 

28 """ 

29 Handles Auth0 JWT token validation using RS256 asymmetric verification. 

30 

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. 

33 

34 Attributes: 

35 settings: Auth0 configuration settings 

36 _jwks_cache: Cached JWKS to minimize API calls to Auth0 

37 """ 

38 

39 JWKS_CACHE_TTL = 21600 # 6 hours in seconds 

40 

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() 

47 

48 async def get_jwks(self) -> Dict: 

49 """ 

50 Fetch Auth0 JSON Web Key Set (JWKS) for token verification. 

51 

52 The JWKS contains Auth0's public keys used to verify JWT signatures. 

53 Results are cached to minimize API calls to Auth0. 

54 

55 Returns: 

56 Dict: JWKS containing public keys 

57 

58 Raises: 

59 HTTPException: If JWKS fetch fails 

60 

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 

82 

83 jwks_url = f"https://{self.settings.AUTH0_DOMAIN}/.well-known/jwks.json" 

84 

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 ) 

98 

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 

103 

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). 

107 

108 The token header contains a 'kid' field identifying which public key 

109 from the JWKS should be used to verify the signature. 

110 

111 Args: 

112 token: JWT token string 

113 jwks: JSON Web Key Set from Auth0 

114 

115 Returns: 

116 str: PEM-formatted public key for verification 

117 

118 Raises: 

119 HTTPException: If kid is missing or signing key not found 

120 

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") 

132 

133 if not kid: 

134 raise HTTPException( 

135 status_code=status.HTTP_401_UNAUTHORIZED, 

136 detail="Token missing key ID (kid)" 

137 ) 

138 

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 

147 

148 # Decode base64url-encoded modulus and exponent 

149 n_bytes = key["n"] 

150 e_bytes = key["e"] 

151 

152 # Add padding if necessary for base64 decoding 

153 n_bytes += "=" * (4 - len(n_bytes) % 4) 

154 e_bytes += "=" * (4 - len(e_bytes) % 4) 

155 

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 ) 

164 

165 # Construct RSA public key 

166 public_numbers = rsa.RSAPublicNumbers(e, n) 

167 public_key = public_numbers.public_key(default_backend()) 

168 

169 # Convert to PEM format 

170 pem = public_key.public_bytes( 

171 encoding=serialization.Encoding.PEM, 

172 format=serialization.PublicFormat.SubjectPublicKeyInfo 

173 ) 

174 

175 return pem.decode("utf-8") 

176 

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 ) 

182 

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 ) 

191 

192 async def verify_token(self, token: str) -> Dict: 

193 """ 

194 Verify and decode Auth0 JWT token. 

195 

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 

203 

204 Args: 

205 token: JWT token string from Authorization header 

206 

207 Returns: 

208 Dict: Decoded token payload containing user info and claims 

209 

210 Raises: 

211 HTTPException (401): If token is invalid, expired, or verification fails 

212 

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() 

227 

228 # Get signing key for this token 

229 signing_key = self.get_signing_key(token, jwks) 

230 

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 ) 

239 

240 return payload 

241 

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 ) 

256 

257 def extract_user_info(self, payload: Dict) -> Dict: 

258 """ 

259 Extract user information from verified token payload. 

260 

261 Parses the JWT payload to extract user identity, email, roles, 

262 and metadata into a standardized format. 

263 

264 Args: 

265 payload: Decoded JWT token payload from verify_token() 

266 

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) 

275 

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, []) 

295 

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 } 

309 

310 

311# Singleton instance for use across the application 

312auth0_handler = Auth0TokenHandler()