diff --git a/litellm/proxy/auth/reject_invalid_tokens.py b/litellm/proxy/auth/reject_invalid_tokens.py index 06f2cf8ee27..fb90bfc2e8c 100644 --- a/litellm/proxy/auth/reject_invalid_tokens.py +++ b/litellm/proxy/auth/reject_invalid_tokens.py @@ -64,6 +64,23 @@ class InvalidVirtualKeyCache: def _cache_key(cls, hashed_token: str) -> str: return "{}{}".format(cls._prefix, hashed_token) + @classmethod + async def delete_invalid_token_cache( + cls, + *, + hashed_token: str, + user_api_key_cache: Any, + ) -> None: + """Clear a stale negative-cache entry after the token is created/restored.""" + try: + await user_api_key_cache.async_delete_cache( + key=cls._cache_key(hashed_token) + ) + except Exception as e: + verbose_proxy_logger.debug( + "InvalidVirtualKeyCache.delete_invalid_token_cache: %s", e + ) + @classmethod async def allows_db_lookup( cls, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index a424f1558fb..1536efe4a3d 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -48,6 +48,7 @@ from litellm.proxy.auth.auth_checks import ( get_team_object, ) from litellm.proxy.auth.auth_utils import abbreviate_api_key +from litellm.proxy.auth.reject_invalid_tokens import InvalidVirtualKeyCache from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_keys from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time @@ -3031,7 +3032,7 @@ async def generate_key_helper_fn( # noqa: PLR0915 access_group_ids: Optional[list] = None, budget_limits: Optional[list] = None, # multiple concurrent budget windows ): - from litellm.proxy.proxy_server import premium_user, prisma_client + from litellm.proxy.proxy_server import premium_user, prisma_client, user_api_key_cache if prisma_client is None: raise Exception( @@ -3244,6 +3245,11 @@ async def generate_key_helper_fn( # noqa: PLR0915 ) key_data["token_id"] = getattr(create_key_response, "token", None) + if key_data["token_id"] is not None: + await InvalidVirtualKeyCache.delete_invalid_token_cache( + hashed_token=key_data["token_id"], + user_api_key_cache=user_api_key_cache, + ) key_data["litellm_budget_table"] = getattr( create_key_response, "litellm_budget_table", None ) diff --git a/tests/proxy_unit_tests/test_key_generate_invalid_token_cache.py b/tests/proxy_unit_tests/test_key_generate_invalid_token_cache.py new file mode 100644 index 00000000000..38388bf03e6 --- /dev/null +++ b/tests/proxy_unit_tests/test_key_generate_invalid_token_cache.py @@ -0,0 +1,46 @@ +import os +import sys +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +import litellm.proxy.proxy_server +from litellm.proxy._types import hash_token +from litellm.proxy.auth.reject_invalid_tokens import InvalidVirtualKeyCache +from litellm.proxy.management_endpoints.key_management_endpoints import ( + generate_key_helper_fn, +) + + +@pytest.mark.asyncio +async def test_generate_key_helper_clears_stale_invalid_token_cache(monkeypatch): + raw_key = "sk-user-supplied-key-that-was-previously-invalid" + hashed_key = hash_token(token=raw_key) + mock_prisma_client = MagicMock() + mock_prisma_client.insert_data = AsyncMock( + return_value=SimpleNamespace( + token=hashed_key, + litellm_budget_table=None, + created_at=None, + updated_at=None, + ) + ) + mock_cache = MagicMock() + mock_cache.async_delete_cache = AsyncMock() + + monkeypatch.setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client) + monkeypatch.setattr(litellm.proxy.proxy_server, "user_api_key_cache", mock_cache) + monkeypatch.setattr(litellm.proxy.proxy_server, "premium_user", False) + + await generate_key_helper_fn( + request_type="key", + key=raw_key, + table_name="key", + ) + + mock_cache.async_delete_cache.assert_awaited_once_with( + key=InvalidVirtualKeyCache._cache_key(hashed_key) + ) diff --git a/tests/proxy_unit_tests/test_reject_invalid_tokens.py b/tests/proxy_unit_tests/test_reject_invalid_tokens.py index 28979c11ed7..2b7db9229af 100644 --- a/tests/proxy_unit_tests/test_reject_invalid_tokens.py +++ b/tests/proxy_unit_tests/test_reject_invalid_tokens.py @@ -15,6 +15,7 @@ from litellm.proxy.auth.reject_invalid_tokens import InvalidVirtualKeyCache def _sk_key() -> str: return "sk-test-invalid-virtual-key" + def _general_settings_positive_ttl() -> dict: """Force negative-cache path on (avoid relying only on default constant).""" return {"invalid_virtual_key_cache_ttl": 3600} @@ -63,7 +64,7 @@ async def test_check_invalid_token_empty_cache_db_miss_records_negative_entry(): async def test_check_invalid_token_negative_cache_hit_short_circuits_even_if_db_has_row(): """ Hash already negative-cached → reject immediately without calling Prisma, - even if a row exists in DB (stale negative cache after key creation is possible). + even if a row exists in DB. """ api_key = _sk_key() neg_key = _negative_cache_key_for(api_key) @@ -88,6 +89,23 @@ async def test_check_invalid_token_negative_cache_hit_short_circuits_even_if_db_ user_api_key_cache.async_set_cache.assert_not_called() +@pytest.mark.asyncio +async def test_delete_invalid_token_cache_removes_negative_cache_key(): + 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_delete_cache = AsyncMock() + + await InvalidVirtualKeyCache.delete_invalid_token_cache( + hashed_token=hashed, + user_api_key_cache=user_api_key_cache, + ) + + user_api_key_cache.async_delete_cache.assert_awaited_once_with(key=neg_key) + + @pytest.mark.asyncio async def test_check_invalid_token_cache_miss_db_hit_allows_auth_flow(): """Negative cache empty, Prisma finds a verification row → preflight passes (False)."""