From 407a9f2962ded502aa7250b95a43962462223ac5 Mon Sep 17 00:00:00 2001 From: harish-berri Date: Wed, 29 Apr 2026 06:11:52 +0000 Subject: [PATCH] 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. --- litellm/proxy/auth/reject_invalid_tokens.py | 30 +++++++-- litellm/proxy/auth/user_api_key_auth.py | 6 +- .../proxy/hooks/proxy_track_cost_callback.py | 14 ++++- .../hooks/test_proxy_track_cost_callback.py | 61 ++++++++++++++++++- 4 files changed, 103 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/auth/reject_invalid_tokens.py b/litellm/proxy/auth/reject_invalid_tokens.py index fb90bfc2e8c..310783973af 100644 --- a/litellm/proxy/auth/reject_invalid_tokens.py +++ b/litellm/proxy/auth/reject_invalid_tokens.py @@ -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 diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 357230affe2..874cdb93b1c 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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), diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index c9946f4e26f..d3ac3f36423 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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, diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 65e7f744c85..6966e2e91d0 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -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,