Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py: 82%

178 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1import re 

2from collections.abc import Mapping, Sequence 

3from datetime import datetime 

4from typing import Final, Protocol 

5 

6from fastapi import APIRouter, Depends, HTTPException, Query 

7 

8from litellm.proxy._types import ( 

9 CreateJWTKeyMappingRequest, 

10 DeleteJWTKeyMappingRequest, 

11 JWTKeyMappingResponse, 

12 LitellmUserRoles, 

13 UpdateJWTKeyMappingRequest, 

14 UserAPIKeyAuth, 

15 hash_token, 

16) 

17from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key 

18from litellm.proxy.auth.user_api_key_auth import user_api_key_auth 

19from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast 

20from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view 

21from litellm.repositories.table_repositories import JWTKeyMappingRepository 

22 

23router: Final = APIRouter() 

24 

25_TOKEN_HASH_PATTERN: Final = re.compile(r"[0-9a-f]{64}") 

26 

27 

28def _validated_token_hash(token: str) -> str: 

29 """Guards a plaintext key from being stored as a hash of a hash, which would never match.""" 

30 if _TOKEN_HASH_PATTERN.fullmatch(token) is None: 

31 raise HTTPException( 

32 status_code=400, 

33 detail=( 

34 "`token` must be the SHA-256 hash of a virtual key " 

35 "(64 lowercase hex characters). Pass the plaintext as `key` instead." 

36 ), 

37 ) 

38 return token 

39 

40 

41_EXACTLY_ONE_IDENTIFIER: Final = ( 

42 "Provide exactly one of `key` (the plaintext virtual key) or `token` (its SHA-256 hash)." 

43) 

44_AT_MOST_ONE_IDENTIFIER: Final = ( 

45 "Provide at most one of `key` (the plaintext virtual key) or `token` (its SHA-256 hash)." 

46) 

47 

48 

49def _token_hash_for_create(data: CreateJWTKeyMappingRequest) -> str: 

50 """Resolve the token hash to store, from either the plaintext key or its hash.""" 

51 if data.key is not None and data.token is not None: 

52 raise HTTPException(status_code=400, detail=_EXACTLY_ONE_IDENTIFIER) 

53 if data.token is not None: 

54 return _validated_token_hash(data.token) 

55 if data.key is not None: 

56 return hash_token(data.key) 

57 raise HTTPException(status_code=400, detail=_EXACTLY_ONE_IDENTIFIER) 

58 

59 

60def _token_hash_for_update(data: UpdateJWTKeyMappingRequest) -> str | None: 

61 """Resolve the token hash to store, or None to leave the mapped key alone.""" 

62 if data.key is not None and data.token is not None: 

63 raise HTTPException(status_code=400, detail=_AT_MOST_ONE_IDENTIFIER) 

64 if data.token is not None: 

65 return _validated_token_hash(data.token) 

66 if data.key is not None: 

67 return hash_token(data.key) 

68 return None 

69 

70 

71class _JWTKeyMappingRecord(Protocol): 

72 """A ``LiteLLM_JWTKeyMapping`` row, viewed through the columns these endpoints read.""" 

73 

74 @property 

75 def id(self) -> str: ... 75 ↛ exitline 75 didn't return from function 'id' because

76 

77 @property 

78 def jwt_issuer(self) -> str: ... 78 ↛ exitline 78 didn't return from function 'jwt_issuer' because

79 

80 @property 

81 def jwt_claim_name(self) -> str: ... 81 ↛ exitline 81 didn't return from function 'jwt_claim_name' because

82 

83 @property 

84 def jwt_claim_value(self) -> str: ... 84 ↛ exitline 84 didn't return from function 'jwt_claim_value' because

85 

86 @property 

87 def description(self) -> str | None: ... 87 ↛ exitline 87 didn't return from function 'description' because

88 

89 @property 

90 def is_active(self) -> bool: ... 90 ↛ exitline 90 didn't return from function 'is_active' because

91 

92 @property 

93 def created_at(self) -> datetime: ... 93 ↛ exitline 93 didn't return from function 'created_at' because

94 

95 @property 

96 def updated_at(self) -> datetime: ... 96 ↛ exitline 96 didn't return from function 'updated_at' because

97 

98 @property 

99 def created_by(self) -> str | None: ... 99 ↛ exitline 99 didn't return from function 'created_by' because

100 

101 @property 

102 def updated_by(self) -> str | None: ... 102 ↛ exitline 102 didn't return from function 'updated_by' because

103 

104 

105class _JWTKeyMappingTable(Protocol): 

106 """The Prisma table actions these endpoints issue against the JWT key mapping table.""" 

107 

108 async def create(self, *, data: Mapping[str, object]) -> _JWTKeyMappingRecord: ... 108 ↛ exitline 108 didn't return from function 'create' because

109 

110 async def find_unique(self, *, where: Mapping[str, object]) -> _JWTKeyMappingRecord | None: ... 110 ↛ exitline 110 didn't return from function 'find_unique' because

111 

112 async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> _JWTKeyMappingRecord: ... 112 ↛ exitline 112 didn't return from function 'update' because

113 

114 async def delete(self, *, where: Mapping[str, object]) -> _JWTKeyMappingRecord | None: ... 114 ↛ exitline 114 didn't return from function 'delete' because

115 

116 async def find_many(self, *, skip: int, take: int, order: Mapping[str, str]) -> Sequence[_JWTKeyMappingRecord]: ... 116 ↛ exitline 116 didn't return from function 'find_many' because

117 

118 async def count(self) -> int: ... 118 ↛ exitline 118 didn't return from function 'count' because

119 

120 

121def _mapping_table(prisma_client: object) -> _JWTKeyMappingTable: 

122 """View the JWT key mapping repository's untyped Prisma table through the actions used here.""" 

123 return JWTKeyMappingRepository(prisma_client).table 

124 

125 

126def _to_response(mapping: _JWTKeyMappingRecord) -> JWTKeyMappingResponse: 

127 """Convert a Prisma mapping object to a safe response (no hashed token).""" 

128 return JWTKeyMappingResponse( 

129 id=mapping.id, 

130 jwt_issuer=mapping.jwt_issuer or None, 

131 jwt_claim_name=mapping.jwt_claim_name, 

132 jwt_claim_value=mapping.jwt_claim_value, 

133 description=mapping.description, 

134 is_active=mapping.is_active, 

135 created_at=mapping.created_at, 

136 updated_at=mapping.updated_at, 

137 created_by=mapping.created_by, 

138 updated_by=mapping.updated_by, 

139 ) 

140 

141 

142@router.post( 

143 "/jwt/key/mapping/new", 

144 tags=["JWT Key Mapping"], 

145 response_model=JWTKeyMappingResponse, 

146) 

147async def create_jwt_key_mapping( 

148 data: CreateJWTKeyMappingRequest, 

149 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

150): 

151 from litellm.proxy.proxy_server import prisma_client, user_api_key_cache 

152 

153 if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: 153 ↛ 154line 153 didn't jump to line 154 because the condition on line 153 was never true

154 raise HTTPException(status_code=403, detail="Only proxy admins can create JWT key mappings") 

155 

156 if prisma_client is None: 156 ↛ 157line 156 didn't jump to line 157 because the condition on line 156 was never true

157 raise HTTPException(status_code=500, detail="Database not connected") 

158 

159 try: 

160 hashed_key: Final = _token_hash_for_create(data) 

161 create_data: Final = { 

162 "jwt_issuer": data.jwt_issuer or "", 

163 "jwt_claim_name": data.jwt_claim_name, 

164 "jwt_claim_value": data.jwt_claim_value, 

165 "token": hashed_key, 

166 "created_by": user_api_key_dict.user_id, 

167 "updated_by": user_api_key_dict.user_id, 

168 } 

169 if data.description is not None: 

170 create_data["description"] = data.description 

171 

172 new_mapping: Final = await _mapping_table(prisma_client).create(data=create_data) 

173 

174 cache_key: Final = jwt_key_mapping_cache_key(data.jwt_claim_name, data.jwt_claim_value, data.jwt_issuer) 

175 await evict_and_broadcast(cache_keys=(cache_key,), user_api_key_cache=user_api_key_cache) 

176 

177 return _to_response(new_mapping) 

178 except HTTPException: 

179 raise 

180 except Exception as e: 

181 error_str: Final = str(e).lower() 

182 if "unique" in error_str or "p2002" in error_str: 

183 raise HTTPException( 

184 status_code=409, 

185 detail=( 

186 f"A mapping for claim '{data.jwt_claim_name}' = '{data.jwt_claim_value}' " 

187 f"already exists for issuer '{data.jwt_issuer}'." 

188 ), 

189 ) 

190 if "foreign" in error_str or "p2003" in error_str: 190 ↛ 195line 190 didn't jump to line 195 because the condition on line 190 was always true

191 raise HTTPException( 

192 status_code=400, 

193 detail="The provided key does not match an existing virtual key.", 

194 ) 

195 raise HTTPException(status_code=500, detail="Failed to create JWT key mapping.") 

196 

197 

198@router.post( 

199 "/jwt/key/mapping/update", 

200 tags=["JWT Key Mapping"], 

201 response_model=JWTKeyMappingResponse, 

202) 

203async def update_jwt_key_mapping( 

204 data: UpdateJWTKeyMappingRequest, 

205 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

206): 

207 from litellm.proxy.proxy_server import prisma_client, user_api_key_cache 

208 

209 if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: 209 ↛ 210line 209 didn't jump to line 210 because the condition on line 209 was never true

210 raise HTTPException(status_code=403, detail="Only proxy admins can update JWT key mappings") 

211 

212 if prisma_client is None: 212 ↛ 213line 212 didn't jump to line 213 because the condition on line 212 was never true

213 raise HTTPException(status_code=500, detail="Database not connected") 

214 

215 update_data: Final = data.model_dump(exclude_unset=True, exclude={"id", "key", "token"}) 

216 token_hash: Final = _token_hash_for_update(data) 

217 if token_hash is not None: 

218 update_data["token"] = token_hash 

219 if "jwt_issuer" in update_data: 

220 # DB column is NOT NULL (see schema.prisma); "" is the global/unscoped sentinel. 

221 update_data["jwt_issuer"] = update_data["jwt_issuer"] or "" 

222 update_data["updated_by"] = user_api_key_dict.user_id 

223 

224 try: 

225 # Get old mapping for cache invalidation 

226 old_mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": data.id}) 

227 

228 if old_mapping is None: 

229 raise HTTPException(status_code=404, detail="Mapping not found") 

230 

231 updated_mapping: Final = await _mapping_table(prisma_client).update(where={"id": data.id}, data=update_data) 

232 

233 if updated_mapping is None: 233 ↛ 234line 233 didn't jump to line 234 because the condition on line 233 was never true

234 raise HTTPException(status_code=404, detail="Mapping not found") 

235 

236 # Evict only after the write commits: a concurrent request between an 

237 # early eviction and the commit would re-cache the old mapping and keep 

238 # it authorized until TTL. 

239 old_cache_key: Final = jwt_key_mapping_cache_key( 

240 old_mapping.jwt_claim_name, old_mapping.jwt_claim_value, old_mapping.jwt_issuer 

241 ) 

242 new_cache_key: Final = jwt_key_mapping_cache_key( 

243 updated_mapping.jwt_claim_name, updated_mapping.jwt_claim_value, updated_mapping.jwt_issuer 

244 ) 

245 cache_keys: Final = (old_cache_key,) if old_cache_key == new_cache_key else (old_cache_key, new_cache_key) 

246 await evict_and_broadcast(cache_keys=cache_keys, user_api_key_cache=user_api_key_cache) 

247 

248 return _to_response(updated_mapping) 

249 except HTTPException: 

250 raise 

251 except Exception as e: 

252 error_str: Final = str(e).lower() 

253 if "unique" in error_str or "p2002" in error_str: 253 ↛ 254line 253 didn't jump to line 254 because the condition on line 253 was never true

254 raise HTTPException( 

255 status_code=409, 

256 detail="A mapping with those claim values already exists.", 

257 ) 

258 if "foreign" in error_str or "p2003" in error_str: 

259 raise HTTPException( 

260 status_code=400, 

261 detail="The provided key does not match an existing virtual key.", 

262 ) 

263 raise HTTPException(status_code=500, detail="Failed to update JWT key mapping.") 

264 

265 

266@router.post("/jwt/key/mapping/delete", tags=["JWT Key Mapping"]) 

267async def delete_jwt_key_mapping( 

268 data: DeleteJWTKeyMappingRequest, 

269 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

270): 

271 from litellm.proxy.proxy_server import prisma_client, user_api_key_cache 

272 

273 if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: 273 ↛ 274line 273 didn't jump to line 274 because the condition on line 273 was never true

274 raise HTTPException(status_code=403, detail="Only proxy admins can delete JWT key mappings") 

275 

276 if prisma_client is None: 276 ↛ 277line 276 didn't jump to line 277 because the condition on line 276 was never true

277 raise HTTPException(status_code=500, detail="Database not connected") 

278 

279 try: 

280 # Get old mapping for cache invalidation 

281 old_mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": data.id}) 

282 

283 if old_mapping is None: 

284 raise HTTPException(status_code=404, detail="Mapping not found") 

285 

286 await _mapping_table(prisma_client).delete(where={"id": data.id}) 

287 

288 # Evict only after the row is gone, else a concurrent request can 

289 # re-cache the deleted mapping and keep it authorized until TTL. 

290 cache_key: Final = jwt_key_mapping_cache_key( 

291 old_mapping.jwt_claim_name, old_mapping.jwt_claim_value, old_mapping.jwt_issuer 

292 ) 

293 await evict_and_broadcast(cache_keys=(cache_key,), user_api_key_cache=user_api_key_cache) 

294 return {"status": "success"} 

295 except HTTPException: 

296 raise 

297 except Exception: 

298 raise HTTPException(status_code=500, detail="Failed to delete JWT key mapping.") 

299 

300 

301@router.get( 

302 "/jwt/key/mapping/list", 

303 tags=["JWT Key Mapping"], 

304) 

305async def list_jwt_key_mappings( 

306 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

307 page: int = Query(1, description="Page number", ge=1), 

308 size: int = Query(50, description="Page size", ge=1, le=100), 

309): 

310 from litellm.proxy.proxy_server import prisma_client 

311 

312 # Admin Viewer follows the read-parity rule. 

313 if not _user_has_admin_view(user_api_key_dict): 313 ↛ 314line 313 didn't jump to line 314 because the condition on line 313 was never true

314 raise HTTPException(status_code=403, detail="Only proxy admins can list JWT key mappings") 

315 

316 if prisma_client is None: 316 ↛ 317line 316 didn't jump to line 317 because the condition on line 316 was never true

317 raise HTTPException(status_code=500, detail="Database not connected") 

318 

319 try: 

320 skip: Final = (page - 1) * size 

321 mappings: Final = await _mapping_table(prisma_client).find_many( 

322 skip=skip, 

323 take=size, 

324 order={"created_at": "desc"}, 

325 ) 

326 total_count: Final = await _mapping_table(prisma_client).count() 

327 return { 

328 "mappings": [_to_response(m) for m in mappings], 

329 "total_count": total_count, 

330 "current_page": page, 

331 "total_pages": -(-total_count // size), # ceiling division 

332 } 

333 except HTTPException: 

334 raise 

335 except Exception: 

336 raise HTTPException(status_code=500, detail="Failed to list JWT key mappings.") 

337 

338 

339@router.get( 

340 "/jwt/key/mapping/info", 

341 tags=["JWT Key Mapping"], 

342 response_model=JWTKeyMappingResponse, 

343) 

344async def info_jwt_key_mapping( 

345 id: str, 

346 user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), 

347): 

348 from litellm.proxy.proxy_server import prisma_client 

349 

350 # Admin Viewer follows the read-parity rule. 

351 if not _user_has_admin_view(user_api_key_dict): 351 ↛ 352line 351 didn't jump to line 352 because the condition on line 351 was never true

352 raise HTTPException(status_code=403, detail="Only proxy admins can get JWT key mapping info") 

353 

354 if prisma_client is None: 354 ↛ 355line 354 didn't jump to line 355 because the condition on line 354 was never true

355 raise HTTPException(status_code=500, detail="Database not connected") 

356 

357 try: 

358 mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": id}) 

359 if mapping is None: 

360 raise HTTPException(status_code=404, detail="Mapping not found") 

361 return _to_response(mapping) 

362 except HTTPException: 

363 raise 

364 except Exception: 

365 raise HTTPException(status_code=500, detail="Failed to get JWT key mapping info.")