mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
22b653a893
commit
407a9f2962
4 changed files with 103 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue