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
« 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
6from fastapi import APIRouter, Depends, HTTPException, Query
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
23router: Final = APIRouter()
25_TOKEN_HASH_PATTERN: Final = re.compile(r"[0-9a-f]{64}")
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
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)
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)
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
71class _JWTKeyMappingRecord(Protocol):
72 """A ``LiteLLM_JWTKeyMapping`` row, viewed through the columns these endpoints read."""
74 @property
75 def id(self) -> str: ... 75 ↛ exitline 75 didn't return from function 'id' because
77 @property
78 def jwt_issuer(self) -> str: ... 78 ↛ exitline 78 didn't return from function 'jwt_issuer' because
80 @property
81 def jwt_claim_name(self) -> str: ... 81 ↛ exitline 81 didn't return from function 'jwt_claim_name' because
83 @property
84 def jwt_claim_value(self) -> str: ... 84 ↛ exitline 84 didn't return from function 'jwt_claim_value' because
86 @property
87 def description(self) -> str | None: ... 87 ↛ exitline 87 didn't return from function 'description' because
89 @property
90 def is_active(self) -> bool: ... 90 ↛ exitline 90 didn't return from function 'is_active' because
92 @property
93 def created_at(self) -> datetime: ... 93 ↛ exitline 93 didn't return from function 'created_at' because
95 @property
96 def updated_at(self) -> datetime: ... 96 ↛ exitline 96 didn't return from function 'updated_at' because
98 @property
99 def created_by(self) -> str | None: ... 99 ↛ exitline 99 didn't return from function 'created_by' because
101 @property
102 def updated_by(self) -> str | None: ... 102 ↛ exitline 102 didn't return from function 'updated_by' because
105class _JWTKeyMappingTable(Protocol):
106 """The Prisma table actions these endpoints issue against the JWT key mapping table."""
108 async def create(self, *, data: Mapping[str, object]) -> _JWTKeyMappingRecord: ... 108 ↛ exitline 108 didn't return from function 'create' because
110 async def find_unique(self, *, where: Mapping[str, object]) -> _JWTKeyMappingRecord | None: ... 110 ↛ exitline 110 didn't return from function 'find_unique' because
112 async def update(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> _JWTKeyMappingRecord: ... 112 ↛ exitline 112 didn't return from function 'update' because
114 async def delete(self, *, where: Mapping[str, object]) -> _JWTKeyMappingRecord | None: ... 114 ↛ exitline 114 didn't return from function 'delete' because
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
118 async def count(self) -> int: ... 118 ↛ exitline 118 didn't return from function 'count' because
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
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 )
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
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")
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")
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
172 new_mapping: Final = await _mapping_table(prisma_client).create(data=create_data)
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)
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.")
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
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")
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")
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
224 try:
225 # Get old mapping for cache invalidation
226 old_mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": data.id})
228 if old_mapping is None:
229 raise HTTPException(status_code=404, detail="Mapping not found")
231 updated_mapping: Final = await _mapping_table(prisma_client).update(where={"id": data.id}, data=update_data)
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")
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)
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.")
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
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")
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")
279 try:
280 # Get old mapping for cache invalidation
281 old_mapping: Final = await _mapping_table(prisma_client).find_unique(where={"id": data.id})
283 if old_mapping is None:
284 raise HTTPException(status_code=404, detail="Mapping not found")
286 await _mapping_table(prisma_client).delete(where={"id": data.id})
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.")
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
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")
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")
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.")
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
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")
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")
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.")