From 63cad23c66f5761358630db5a80dd23abc8caf4c Mon Sep 17 00:00:00 2001 From: harish-berri Date: Wed, 29 Apr 2026 17:42:11 +0000 Subject: [PATCH] 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. --- litellm/proxy/auth/reject_invalid_tokens.py | 44 +++++++++- .../key_management_endpoints.py | 9 ++ litellm/proxy/utils.py | 4 +- .../test_deprecated_key_lookup.py | 35 ++++++++ .../test_reject_invalid_tokens.py | 51 +++++++++++ .../test_key_management_endpoints.py | 88 +++++++++++++++++++ 6 files changed, 228 insertions(+), 3 deletions(-) create mode 100644 tests/proxy_unit_tests/test_deprecated_key_lookup.py diff --git a/litellm/proxy/auth/reject_invalid_tokens.py b/litellm/proxy/auth/reject_invalid_tokens.py index 310783973af..ad98b690dec 100644 --- a/litellm/proxy/auth/reject_invalid_tokens.py +++ b/litellm/proxy/auth/reject_invalid_tokens.py @@ -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, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 1536efe4a3d..588ed81fe3b 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -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), diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 712853a33c4..686d1bd5f92 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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 diff --git a/tests/proxy_unit_tests/test_deprecated_key_lookup.py b/tests/proxy_unit_tests/test_deprecated_key_lookup.py new file mode 100644 index 00000000000..58a5208586e --- /dev/null +++ b/tests/proxy_unit_tests/test_deprecated_key_lookup.py @@ -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() diff --git a/tests/proxy_unit_tests/test_reject_invalid_tokens.py b/tests/proxy_unit_tests/test_reject_invalid_tokens.py index 2b7db9229af..6de3e515448 100644 --- a/tests/proxy_unit_tests/test_reject_invalid_tokens.py +++ b/tests/proxy_unit_tests/test_reject_invalid_tokens.py @@ -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() diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index f256ffd8661..59241a8f34e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -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) + )