mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(jwt): evict mapping cache after DB write in /jwt/key/mapping update and delete
Evicting before the mutation commits left a race: a concurrent JWT request could re-cache the old mapping between the eviction and the commit, keeping a deleted or renamed claim authorized until the cache TTL expired. Flagged by review on PR #39808.
This commit is contained in:
parent
52b746e8ea
commit
49cac6fc5b
2 changed files with 93 additions and 6 deletions
|
|
@ -170,18 +170,20 @@ async def update_jwt_key_mapping(
|
|||
if old_mapping is None:
|
||||
raise HTTPException(status_code=404, detail="Mapping not found")
|
||||
|
||||
old_cache_key: Final = jwt_key_mapping_cache_key(old_mapping.jwt_claim_name, old_mapping.jwt_claim_value)
|
||||
await evict_and_broadcast(cache_keys=(old_cache_key,), user_api_key_cache=user_api_key_cache)
|
||||
|
||||
updated_mapping: Final = await _mapping_table(prisma_client).update(where={"id": data.id}, data=update_data)
|
||||
|
||||
if updated_mapping is None:
|
||||
raise HTTPException(status_code=404, detail="Mapping not found")
|
||||
|
||||
# Evict only after the write commits: a concurrent request between an
|
||||
# early eviction and the commit would re-cache the old mapping and keep
|
||||
# it authorized until TTL.
|
||||
old_cache_key: Final = jwt_key_mapping_cache_key(old_mapping.jwt_claim_name, old_mapping.jwt_claim_value)
|
||||
new_cache_key: Final = jwt_key_mapping_cache_key(
|
||||
updated_mapping.jwt_claim_name, updated_mapping.jwt_claim_value
|
||||
)
|
||||
await evict_and_broadcast(cache_keys=(new_cache_key,), user_api_key_cache=user_api_key_cache)
|
||||
cache_keys: Final = (old_cache_key,) if old_cache_key == new_cache_key else (old_cache_key, new_cache_key)
|
||||
await evict_and_broadcast(cache_keys=cache_keys, user_api_key_cache=user_api_key_cache)
|
||||
|
||||
return _to_response(updated_mapping)
|
||||
except HTTPException:
|
||||
|
|
@ -221,10 +223,12 @@ async def delete_jwt_key_mapping(
|
|||
if old_mapping is None:
|
||||
raise HTTPException(status_code=404, detail="Mapping not found")
|
||||
|
||||
await _mapping_table(prisma_client).delete(where={"id": data.id})
|
||||
|
||||
# Evict only after the row is gone, else a concurrent request can
|
||||
# re-cache the deleted mapping and keep it authorized until TTL.
|
||||
cache_key: Final = jwt_key_mapping_cache_key(old_mapping.jwt_claim_name, old_mapping.jwt_claim_value)
|
||||
await evict_and_broadcast(cache_keys=(cache_key,), user_api_key_cache=user_api_key_cache)
|
||||
|
||||
await _mapping_table(prisma_client).delete(where={"id": data.id})
|
||||
return {"status": "success"}
|
||||
except HTTPException:
|
||||
raise
|
||||
|
|
|
|||
|
|
@ -1333,3 +1333,86 @@ def test_jwt_client_id_field_does_not_raise_on_duplicate():
|
|||
virtual_key_claim_field="new_field",
|
||||
)
|
||||
assert auth.virtual_key_claim_field == "new_field"
|
||||
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# Tests: cache eviction must happen AFTER the DB write commits
|
||||
# ──────────────────────────────────────────────
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_evicts_cache_after_row_is_gone():
|
||||
"""A JWT request racing the delete must not keep the removed mapping authorized.
|
||||
|
||||
The DB delete simulates a concurrent request re-caching the mapping mid-write.
|
||||
If the endpoint evicts before the delete commits, that repopulated entry
|
||||
survives until TTL and the deleted mapping stays usable.
|
||||
"""
|
||||
from litellm.proxy._types import DeleteJWTKeyMappingRequest
|
||||
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
|
||||
|
||||
cache_key = jwt_key_mapping_cache_key("email", "user@example.com")
|
||||
user_api_key_cache = DualCache()
|
||||
await user_api_key_cache.async_set_cache(key=cache_key, value="hashed_token")
|
||||
|
||||
mock_prisma = _mock_prisma()
|
||||
mock_prisma.db.litellm_jwtkeymapping.find_unique.return_value = _mock_mapping()
|
||||
|
||||
async def concurrent_reader_repopulates(**kwargs):
|
||||
await user_api_key_cache.async_set_cache(key=cache_key, value="hashed_token")
|
||||
return _mock_mapping()
|
||||
|
||||
mock_prisma.db.litellm_jwtkeymapping.delete.side_effect = concurrent_reader_repopulates
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
):
|
||||
result = await delete_jwt_key_mapping(
|
||||
data=DeleteJWTKeyMappingRequest(id="mapping-1"),
|
||||
user_api_key_dict=_make_admin_auth(),
|
||||
)
|
||||
|
||||
assert result == {"status": "success"}
|
||||
assert await user_api_key_cache.async_get_cache(cache_key) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_evicts_old_and_new_cache_keys_after_write():
|
||||
"""Renaming a mapping's claim must leave neither claim serving stale cache.
|
||||
|
||||
The DB update simulates a concurrent request re-caching the OLD mapping
|
||||
mid-write. Both the old claim's entry (would restore the pre-rename token)
|
||||
and the new claim's __NO_MAPPING__ sentinel (would 403 the renamed claim)
|
||||
must be gone once the endpoint returns.
|
||||
"""
|
||||
from litellm.proxy._types import UpdateJWTKeyMappingRequest
|
||||
from litellm.proxy.auth.auth_checks import jwt_key_mapping_cache_key
|
||||
|
||||
old_cache_key = jwt_key_mapping_cache_key("email", "user@example.com")
|
||||
new_cache_key = jwt_key_mapping_cache_key("email", "renamed@example.com")
|
||||
user_api_key_cache = DualCache()
|
||||
await user_api_key_cache.async_set_cache(key=old_cache_key, value="hashed_token")
|
||||
await user_api_key_cache.async_set_cache(key=new_cache_key, value="__NO_MAPPING__")
|
||||
|
||||
mock_prisma = _mock_prisma()
|
||||
mock_prisma.db.litellm_jwtkeymapping.find_unique.return_value = _mock_mapping()
|
||||
|
||||
async def concurrent_reader_repopulates(**kwargs):
|
||||
await user_api_key_cache.async_set_cache(key=old_cache_key, value="hashed_token")
|
||||
return _mock_mapping(claim_value="renamed@example.com")
|
||||
|
||||
mock_prisma.db.litellm_jwtkeymapping.update.side_effect = concurrent_reader_repopulates
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), # test-quality-ok: proxy_server module global is the endpoint's only injection point
|
||||
):
|
||||
result = await update_jwt_key_mapping(
|
||||
data=UpdateJWTKeyMappingRequest(id="mapping-1", jwt_claim_value="renamed@example.com"),
|
||||
user_api_key_dict=_make_admin_auth(),
|
||||
)
|
||||
|
||||
assert result.jwt_claim_value == "renamed@example.com"
|
||||
assert await user_api_key_cache.async_get_cache(old_cache_key) is None
|
||||
assert await user_api_key_cache.async_get_cache(new_cache_key) is None
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue