negative flow flagged by greptime where generated api key could be spoofed / used by the user in which case, the invalid key needs to be deleted if there is a collision

This commit is contained in:
harish-berri 2026-04-29 01:05:43 +00:00
parent 1a885a755d
commit a34469c62f
4 changed files with 89 additions and 2 deletions

View file

@ -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,

View file

@ -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
)

View file

@ -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)
)

View file

@ -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)."""