Refactor InvalidVirtualKeyCache and enhance token validation

- Introduced `check_invalid_hashed_token` method to `InvalidVirtualKeyCache` for validating hashed tokens, improving security and efficiency.
- Updated `user_api_key_auth.py` to utilize the new method for checking invalid tokens.
- Modified `_ProxyDBLogger` to skip key lookups for negative cached invalid keys, optimizing performance.
- Added unit tests to ensure correct behavior of the new validation logic and caching mechanism.
This commit is contained in:
harish-berri 2026-04-29 06:11:52 +00:00
parent 22b653a893
commit 407a9f2962
4 changed files with 103 additions and 8 deletions

View file

@ -25,6 +25,8 @@ from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm.proxy._types import hash_token
INVALID_VIRTUAL_KEY_CACHE_PREFIX = "invalid_vk:"
class InvalidVirtualKeyCache:
"""Settings + negative cache for virtual keys that are not in the DB (or not yet)."""
@ -149,8 +151,6 @@ class InvalidVirtualKeyCache:
Malformed keys raise ``HTTPException`` (401) with masking details instead of returning bool.
"""
ttl_seconds = cls.configured_ttl_seconds(general_settings)
if isinstance(api_key, str):
_masked_key = (
"{}****{}".format(api_key[:4], api_key[-4:])
@ -177,7 +177,29 @@ class InvalidVirtualKeyCache:
detail="LiteLLM Virtual Key expected.",
)
hashed_token = hash_token(token=api_key)
return await cls.check_invalid_hashed_token(
hashed_token=hash_token(token=api_key),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
general_settings=general_settings,
)
@classmethod
async def check_invalid_hashed_token(
cls,
*,
hashed_token: str,
prisma_client: Any,
user_api_key_cache: Any,
general_settings: Any,
) -> bool:
"""
Hashed-token preflight for callers that no longer have the raw ``sk-`` key.
Returns ``True`` if the request should be rejected as an invalid virtual key.
Returns ``False`` if preflight passed; caller may load the full key object.
"""
ttl_seconds = cls.configured_ttl_seconds(general_settings)
if ttl_seconds is None:
return False
@ -196,7 +218,7 @@ class InvalidVirtualKeyCache:
)
except Exception as e:
verbose_proxy_logger.debug(
"InvalidVirtualKeyCache.check_invalid_token: verification token probe failed, continuing to combined_view: %s",
"InvalidVirtualKeyCache.check_invalid_hashed_token: verification token probe failed, continuing to combined_view: %s",
e,
)
token_probe_failed = True

View file

@ -1166,12 +1166,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
if valid_token is None:
if await InvalidVirtualKeyCache.check_invalid_token(
is_invalid_token = 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,
):
)
if is_invalid_token:
raise ProxyException(
message="Authentication Error at InvalidVirtualKeyCache, Invalid proxy server token passed. Token (hash) = {}. Unable to find token in cache or `LiteLLM_VerificationTokenTable`".format(
hash_token(token=api_key),

View file

@ -17,6 +17,7 @@ from litellm.proxy.auth.auth_checks import (
get_team_object,
log_db_metrics,
)
from litellm.proxy.auth.reject_invalid_tokens import InvalidVirtualKeyCache
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy.utils import ProxyUpdateSpend
@ -303,6 +304,7 @@ class _ProxyDBLogger(CustomLogger):
return metadata
from litellm.proxy.proxy_server import (
general_settings,
prisma_client,
proxy_logging_obj,
user_api_key_cache,
@ -310,8 +312,18 @@ class _ProxyDBLogger(CustomLogger):
# Step 1: If key fields are missing, look up the full key object
if metadata.get("user_api_key_alias") is None:
is_invalid_token = await InvalidVirtualKeyCache.check_invalid_hashed_token(
hashed_token=api_key_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
general_settings=general_settings,
)
if is_invalid_token:
return metadata
try:
key_obj = await get_key_object(
key_obj = await cast(Any, get_key_object)(
hashed_token=api_key_hash,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,

View file

@ -12,7 +12,8 @@ sys.path.insert(
from datetime import datetime
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy._types import UserAPIKeyAuth, hash_token
from litellm.proxy.auth.reject_invalid_tokens import InvalidVirtualKeyCache
from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger
from litellm.types.utils import StandardLoggingPayload
@ -278,6 +279,12 @@ async def test_enrich_failure_metadata_with_full_key_lookup():
new_callable=AsyncMock,
return_value=mock_key_obj,
),
patch.object(
InvalidVirtualKeyCache,
"check_invalid_hashed_token",
new_callable=AsyncMock,
return_value=False,
),
patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
new_callable=AsyncMock,
@ -348,6 +355,52 @@ async def test_enrich_failure_metadata_skips_when_no_api_key():
mock_get_key.assert_not_called()
@pytest.mark.asyncio
async def test_enrich_failure_metadata_skips_key_lookup_for_negative_cached_invalid_key():
"""
Invalid auth failures are already negative-cached. Failure metadata enrichment
should not re-query the key combined view for those same invalid keys.
"""
raw_api_key = "sk-invalid-key"
hashed_api_key = hash_token(raw_api_key)
mock_cache = MagicMock()
mock_cache.async_get_cache = AsyncMock(return_value="")
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch(
"litellm.proxy.proxy_server.general_settings",
{"invalid_virtual_key_cache_ttl": 3600},
),
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()),
patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_key_object",
new_callable=AsyncMock,
) as mock_get_key,
patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
new_callable=AsyncMock,
) as mock_get_team,
):
metadata = {
"user_api_key": hashed_api_key,
"user_api_key_alias": None,
"user_api_key_user_id": None,
"user_api_key_team_id": None,
"user_api_key_team_alias": None,
}
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
assert result == metadata
mock_cache.async_get_cache.assert_awaited_once_with(
key=InvalidVirtualKeyCache._cache_key(hashed_api_key)
)
mock_get_key.assert_not_called()
mock_get_team.assert_not_called()
@pytest.mark.asyncio
async def test_async_post_call_failure_hook_enriches_auth_error_metadata():
"""
@ -389,6 +442,12 @@ async def test_async_post_call_failure_hook_enriches_auth_error_metadata():
new_callable=AsyncMock,
return_value=mock_key_obj,
),
patch.object(
InvalidVirtualKeyCache,
"check_invalid_hashed_token",
new_callable=AsyncMock,
return_value=False,
),
patch(
"litellm.proxy.hooks.proxy_track_cost_callback.get_team_object",
new_callable=AsyncMock,