Enhance token validation and caching mechanisms

- Added `_deprecated_token_exists` method to `InvalidVirtualKeyCache` for checking the validity of deprecated tokens during the key-rotation grace period.
- Updated `check_invalid_token` to utilize the new method, ensuring that negative caching does not occur for valid deprecated tokens.
- Modified `_lookup_deprecated_key` to cache the `revoke_at` timestamp alongside the active token ID.
- Implemented unit tests for the new functionality, ensuring correct behavior in various scenarios.
This commit is contained in:
harish-berri 2026-04-29 17:42:11 +00:00
parent 407a9f2962
commit 63cad23c66
6 changed files with 228 additions and 3 deletions

View file

@ -16,6 +16,7 @@ Cache keys: ``invalid_vk:{hashed_token}``.
from __future__ import annotations
from datetime import datetime, timezone
from typing import Any, Dict, Optional, Union
from fastapi import HTTPException, status
@ -132,6 +133,37 @@ class InvalidVirtualKeyCache:
except Exception as e:
verbose_proxy_logger.debug("InvalidVirtualKeyCache.record_miss: %s", e)
@classmethod
async def _deprecated_token_exists(
cls,
*,
hashed_token: str,
prisma_client: Any,
) -> Optional[bool]:
"""
Return whether ``hashed_token`` is still valid through key-rotation grace period.
``None`` means the probe failed; callers should avoid negative-caching so the
regular auth path can attempt the full lookup.
"""
try:
deprecated_token_row = (
await prisma_client.db.litellm_deprecatedverificationtoken.find_first(
where={
"token": hashed_token,
"revoke_at": {"gt": datetime.now(timezone.utc)},
}
)
)
except Exception as e:
verbose_proxy_logger.debug(
"InvalidVirtualKeyCache._deprecated_token_exists: deprecated token probe failed, continuing to combined_view: %s",
e,
)
return None
return deprecated_token_row is not None
@classmethod
async def check_invalid_token(
cls,
@ -224,7 +256,17 @@ class InvalidVirtualKeyCache:
token_probe_failed = True
token_row = None
if not token_probe_failed and token_row is None:
if token_probe_failed:
return False
if token_row is None:
deprecated_token_exists = await cls._deprecated_token_exists(
hashed_token=hashed_token,
prisma_client=prisma_client,
)
if deprecated_token_exists is not False:
return False
await cls.record_miss(
hashed_token=hashed_token,
user_api_key_cache=user_api_key_cache,

View file

@ -3889,6 +3889,15 @@ async def _execute_virtual_key_regeneration(
updated_token_dict["key"] = new_token
updated_token_dict["token_id"] = updated_token_dict.pop("token")
await InvalidVirtualKeyCache.delete_invalid_token_cache(
hashed_token=new_token_hash,
user_api_key_cache=user_api_key_cache,
)
await InvalidVirtualKeyCache.delete_invalid_token_cache(
hashed_token=hashed_api_key,
user_api_key_cache=user_api_key_cache,
)
if hashed_api_key or key:
await _delete_cache_key_object(
hashed_token=_hash_token_if_needed(key),

View file

@ -2412,7 +2412,6 @@ async def _lookup_deprecated_key(
# Check cache first
cached = _deprecated_key_cache.get(hashed_token)
cached = _deprecated_key_cache.get(hashed_token)
if cached is not None:
active_token_id, cache_expires_at_ts, revoke_at_ts = cached
if now_ts < cache_expires_at_ts and now_ts < revoke_at_ts:
@ -2426,12 +2425,13 @@ async def _lookup_deprecated_key(
"token": hashed_token,
"revoke_at": {"gt": now},
},
select={"active_token_id": True},
select={"active_token_id": True, "revoke_at": True},
)
if deprecated_row and deprecated_row.active_token_id:
_deprecated_key_cache[hashed_token] = (
deprecated_row.active_token_id,
now_ts + _DEPRECATED_KEY_CACHE_TTL_SECONDS,
deprecated_row.revoke_at.timestamp(),
)
return deprecated_row.active_token_id
# Only cache positive results; negative lookups are fast on indexed columns

View file

@ -0,0 +1,35 @@
from datetime import datetime, timedelta, timezone
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm.proxy.utils import _deprecated_key_cache, _lookup_deprecated_key
@pytest.mark.asyncio
async def test_lookup_deprecated_key_handles_cached_entry():
"""
The first lookup should populate the deprecated-key cache; the second lookup
should read that cached entry without raising or hitting the database again.
"""
hashed_token = "old-token-hash"
active_token_id = "active-token-hash"
revoke_at = datetime.now(timezone.utc) + timedelta(minutes=5)
deprecated_row = MagicMock()
deprecated_row.active_token_id = active_token_id
deprecated_row.revoke_at = revoke_at
db = MagicMock()
db.litellm_deprecatedverificationtoken.find_first = AsyncMock(
return_value=deprecated_row
)
_deprecated_key_cache.clear()
try:
assert await _lookup_deprecated_key(db=db, hashed_token=hashed_token) == active_token_id
assert await _lookup_deprecated_key(db=db, hashed_token=hashed_token) == active_token_id
finally:
_deprecated_key_cache.clear()
db.litellm_deprecatedverificationtoken.find_first.assert_awaited_once()

View file

@ -40,8 +40,12 @@ async def test_check_invalid_token_empty_cache_db_miss_records_negative_entry():
user_api_key_cache.async_set_cache = AsyncMock()
find_first = AsyncMock(return_value=None)
deprecated_find_first = AsyncMock(return_value=None)
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_first = find_first
prisma_client.db.litellm_deprecatedverificationtoken.find_first = (
deprecated_find_first
)
result = await InvalidVirtualKeyCache.check_invalid_token(
api_key=api_key,
@ -52,6 +56,7 @@ async def test_check_invalid_token_empty_cache_db_miss_records_negative_entry():
assert result is True
find_first.assert_awaited_once_with(where={"token": hashed})
deprecated_find_first.assert_awaited_once()
user_api_key_cache.async_get_cache.assert_awaited_once_with(key=neg_key)
user_api_key_cache.async_set_cache.assert_awaited_once_with(
key=neg_key,
@ -74,8 +79,12 @@ async def test_check_invalid_token_negative_cache_hit_short_circuits_even_if_db_
user_api_key_cache.async_set_cache = AsyncMock()
find_first = AsyncMock(return_value=MagicMock(token=hash_token(token=api_key)))
deprecated_find_first = AsyncMock(return_value=None)
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_first = find_first
prisma_client.db.litellm_deprecatedverificationtoken.find_first = (
deprecated_find_first
)
result = await InvalidVirtualKeyCache.check_invalid_token(
api_key=api_key,
@ -86,6 +95,7 @@ async def test_check_invalid_token_negative_cache_hit_short_circuits_even_if_db_
assert result is True
find_first.assert_not_called()
deprecated_find_first.assert_not_called()
user_api_key_cache.async_set_cache.assert_not_called()
@ -118,8 +128,12 @@ async def test_check_invalid_token_cache_miss_db_hit_allows_auth_flow():
user_api_key_cache.async_set_cache = AsyncMock()
find_first = AsyncMock(return_value=MagicMock(token=hashed))
deprecated_find_first = AsyncMock(return_value=None)
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_first = find_first
prisma_client.db.litellm_deprecatedverificationtoken.find_first = (
deprecated_find_first
)
result = await InvalidVirtualKeyCache.check_invalid_token(
api_key=api_key,
@ -130,5 +144,42 @@ async def test_check_invalid_token_cache_miss_db_hit_allows_auth_flow():
assert result is False
find_first.assert_awaited_once_with(where={"token": hashed})
deprecated_find_first.assert_not_called()
user_api_key_cache.async_get_cache.assert_awaited_once_with(key=neg_key)
user_api_key_cache.async_set_cache.assert_not_called()
@pytest.mark.asyncio
async def test_check_invalid_token_allows_deprecated_key_in_grace_period():
"""
A rotated key can be absent from the active token table but still valid via
LiteLLM_DeprecatedVerificationToken. Do not negative-cache that hash.
"""
api_key = _sk_key()
hashed = hash_token(token=api_key)
neg_key = _negative_cache_key_for(api_key)
user_api_key_cache = MagicMock()
user_api_key_cache.async_get_cache = AsyncMock(return_value=None)
user_api_key_cache.async_set_cache = AsyncMock()
find_first = AsyncMock(return_value=None)
deprecated_find_first = AsyncMock(return_value=MagicMock(token=hashed))
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_first = find_first
prisma_client.db.litellm_deprecatedverificationtoken.find_first = (
deprecated_find_first
)
result = await InvalidVirtualKeyCache.check_invalid_token(
api_key=api_key,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
general_settings=_general_settings_positive_ttl(),
)
assert result is False
find_first.assert_awaited_once_with(where={"token": hashed})
deprecated_find_first.assert_awaited_once()
user_api_key_cache.async_get_cache.assert_awaited_once_with(key=neg_key)
user_api_key_cache.async_set_cache.assert_not_called()

View file

@ -9726,3 +9726,91 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha
call_kwargs = mock_delete_cache.call_args.kwargs
# The token hash should be passed as-is, NOT double-hashed
assert call_kwargs["hashed_token"] == token_hash
@pytest.mark.asyncio
async def test_execute_virtual_key_regeneration_clears_invalid_token_cache_for_new_key():
"""
Custom regenerated keys may have been attempted before regeneration and
negative-cached. Regeneration should clear that stale invalid-token cache entry.
"""
from litellm.proxy._types import RegenerateKeyRequest, hash_token
from litellm.proxy.auth.reject_invalid_tokens import InvalidVirtualKeyCache
from litellm.proxy.management_endpoints.key_management_endpoints import (
_execute_virtual_key_regeneration,
)
old_token_hash = "old-token-hash"
custom_new_key = "sk-custom-new-key-previously-invalid"
new_token_hash = hash_token(token=custom_new_key)
existing_key = LiteLLM_VerificationToken(
token=old_token_hash,
user_id="user-1",
models=["gpt-4"],
team_id=None,
max_budget=None,
tags=None,
)
class DictLikeResult:
def __init__(self, data):
self._data = data
def __iter__(self):
return iter(self._data.items())
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock(
return_value=DictLikeResult(
{"token": new_token_hash, "key_name": "sk-...alid", "user_id": "user-1"}
)
)
mock_prisma_client.jsonify_object = MagicMock(side_effect=lambda data: data)
mock_user_api_key_cache = MagicMock()
mock_user_api_key_cache.async_delete_cache = AsyncMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
with (
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook",
new_callable=AsyncMock,
),
patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data",
new_callable=AsyncMock,
return_value={},
),
):
await _execute_virtual_key_regeneration(
prisma_client=mock_prisma_client,
key_in_db=existing_key,
hashed_api_key=old_token_hash,
key=old_token_hash,
data=RegenerateKeyRequest(new_key=custom_new_key),
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=MagicMock(),
)
mock_user_api_key_cache.async_delete_cache.assert_any_await(
key=InvalidVirtualKeyCache._cache_key(new_token_hash)
)
mock_user_api_key_cache.async_delete_cache.assert_any_await(
key=InvalidVirtualKeyCache._cache_key(old_token_hash)
)