mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Enhance token validation and caching mechanisms
- Added `_deprecated_token_exists` method to `InvalidVirtualKeyCache` for checking the validity of deprecated tokens during the key-rotation grace period. - Updated `check_invalid_token` to utilize the new method, ensuring that negative caching does not occur for valid deprecated tokens. - Modified `_lookup_deprecated_key` to cache the `revoke_at` timestamp alongside the active token ID. - Implemented unit tests for the new functionality, ensuring correct behavior in various scenarios.
This commit is contained in:
parent
407a9f2962
commit
63cad23c66
6 changed files with 228 additions and 3 deletions
|
|
@ -16,6 +16,7 @@ Cache keys: ``invalid_vk:{hashed_token}``.
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
from fastapi import HTTPException, status
|
||||
|
|
@ -132,6 +133,37 @@ class InvalidVirtualKeyCache:
|
|||
except Exception as e:
|
||||
verbose_proxy_logger.debug("InvalidVirtualKeyCache.record_miss: %s", e)
|
||||
|
||||
@classmethod
|
||||
async def _deprecated_token_exists(
|
||||
cls,
|
||||
*,
|
||||
hashed_token: str,
|
||||
prisma_client: Any,
|
||||
) -> Optional[bool]:
|
||||
"""
|
||||
Return whether ``hashed_token`` is still valid through key-rotation grace period.
|
||||
|
||||
``None`` means the probe failed; callers should avoid negative-caching so the
|
||||
regular auth path can attempt the full lookup.
|
||||
"""
|
||||
try:
|
||||
deprecated_token_row = (
|
||||
await prisma_client.db.litellm_deprecatedverificationtoken.find_first(
|
||||
where={
|
||||
"token": hashed_token,
|
||||
"revoke_at": {"gt": datetime.now(timezone.utc)},
|
||||
}
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"InvalidVirtualKeyCache._deprecated_token_exists: deprecated token probe failed, continuing to combined_view: %s",
|
||||
e,
|
||||
)
|
||||
return None
|
||||
|
||||
return deprecated_token_row is not None
|
||||
|
||||
@classmethod
|
||||
async def check_invalid_token(
|
||||
cls,
|
||||
|
|
@ -224,7 +256,17 @@ class InvalidVirtualKeyCache:
|
|||
token_probe_failed = True
|
||||
token_row = None
|
||||
|
||||
if not token_probe_failed and token_row is None:
|
||||
if token_probe_failed:
|
||||
return False
|
||||
|
||||
if token_row is None:
|
||||
deprecated_token_exists = await cls._deprecated_token_exists(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
if deprecated_token_exists is not False:
|
||||
return False
|
||||
|
||||
await cls.record_miss(
|
||||
hashed_token=hashed_token,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
|
|
|
|||
|
|
@ -3889,6 +3889,15 @@ async def _execute_virtual_key_regeneration(
|
|||
updated_token_dict["key"] = new_token
|
||||
updated_token_dict["token_id"] = updated_token_dict.pop("token")
|
||||
|
||||
await InvalidVirtualKeyCache.delete_invalid_token_cache(
|
||||
hashed_token=new_token_hash,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
await InvalidVirtualKeyCache.delete_invalid_token_cache(
|
||||
hashed_token=hashed_api_key,
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
)
|
||||
|
||||
if hashed_api_key or key:
|
||||
await _delete_cache_key_object(
|
||||
hashed_token=_hash_token_if_needed(key),
|
||||
|
|
|
|||
|
|
@ -2412,7 +2412,6 @@ async def _lookup_deprecated_key(
|
|||
|
||||
# Check cache first
|
||||
cached = _deprecated_key_cache.get(hashed_token)
|
||||
cached = _deprecated_key_cache.get(hashed_token)
|
||||
if cached is not None:
|
||||
active_token_id, cache_expires_at_ts, revoke_at_ts = cached
|
||||
if now_ts < cache_expires_at_ts and now_ts < revoke_at_ts:
|
||||
|
|
@ -2426,12 +2425,13 @@ async def _lookup_deprecated_key(
|
|||
"token": hashed_token,
|
||||
"revoke_at": {"gt": now},
|
||||
},
|
||||
select={"active_token_id": True},
|
||||
select={"active_token_id": True, "revoke_at": True},
|
||||
)
|
||||
if deprecated_row and deprecated_row.active_token_id:
|
||||
_deprecated_key_cache[hashed_token] = (
|
||||
deprecated_row.active_token_id,
|
||||
now_ts + _DEPRECATED_KEY_CACHE_TTL_SECONDS,
|
||||
deprecated_row.revoke_at.timestamp(),
|
||||
)
|
||||
return deprecated_row.active_token_id
|
||||
# Only cache positive results; negative lookups are fast on indexed columns
|
||||
|
|
|
|||
35
tests/proxy_unit_tests/test_deprecated_key_lookup.py
Normal file
35
tests/proxy_unit_tests/test_deprecated_key_lookup.py
Normal file
|
|
@ -0,0 +1,35 @@
|
|||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.utils import _deprecated_key_cache, _lookup_deprecated_key
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_lookup_deprecated_key_handles_cached_entry():
|
||||
"""
|
||||
The first lookup should populate the deprecated-key cache; the second lookup
|
||||
should read that cached entry without raising or hitting the database again.
|
||||
"""
|
||||
hashed_token = "old-token-hash"
|
||||
active_token_id = "active-token-hash"
|
||||
revoke_at = datetime.now(timezone.utc) + timedelta(minutes=5)
|
||||
|
||||
deprecated_row = MagicMock()
|
||||
deprecated_row.active_token_id = active_token_id
|
||||
deprecated_row.revoke_at = revoke_at
|
||||
|
||||
db = MagicMock()
|
||||
db.litellm_deprecatedverificationtoken.find_first = AsyncMock(
|
||||
return_value=deprecated_row
|
||||
)
|
||||
|
||||
_deprecated_key_cache.clear()
|
||||
try:
|
||||
assert await _lookup_deprecated_key(db=db, hashed_token=hashed_token) == active_token_id
|
||||
assert await _lookup_deprecated_key(db=db, hashed_token=hashed_token) == active_token_id
|
||||
finally:
|
||||
_deprecated_key_cache.clear()
|
||||
|
||||
db.litellm_deprecatedverificationtoken.find_first.assert_awaited_once()
|
||||
|
|
@ -40,8 +40,12 @@ async def test_check_invalid_token_empty_cache_db_miss_records_negative_entry():
|
|||
user_api_key_cache.async_set_cache = AsyncMock()
|
||||
|
||||
find_first = AsyncMock(return_value=None)
|
||||
deprecated_find_first = AsyncMock(return_value=None)
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_verificationtoken.find_first = find_first
|
||||
prisma_client.db.litellm_deprecatedverificationtoken.find_first = (
|
||||
deprecated_find_first
|
||||
)
|
||||
|
||||
result = await InvalidVirtualKeyCache.check_invalid_token(
|
||||
api_key=api_key,
|
||||
|
|
@ -52,6 +56,7 @@ async def test_check_invalid_token_empty_cache_db_miss_records_negative_entry():
|
|||
|
||||
assert result is True
|
||||
find_first.assert_awaited_once_with(where={"token": hashed})
|
||||
deprecated_find_first.assert_awaited_once()
|
||||
user_api_key_cache.async_get_cache.assert_awaited_once_with(key=neg_key)
|
||||
user_api_key_cache.async_set_cache.assert_awaited_once_with(
|
||||
key=neg_key,
|
||||
|
|
@ -74,8 +79,12 @@ async def test_check_invalid_token_negative_cache_hit_short_circuits_even_if_db_
|
|||
user_api_key_cache.async_set_cache = AsyncMock()
|
||||
|
||||
find_first = AsyncMock(return_value=MagicMock(token=hash_token(token=api_key)))
|
||||
deprecated_find_first = AsyncMock(return_value=None)
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_verificationtoken.find_first = find_first
|
||||
prisma_client.db.litellm_deprecatedverificationtoken.find_first = (
|
||||
deprecated_find_first
|
||||
)
|
||||
|
||||
result = await InvalidVirtualKeyCache.check_invalid_token(
|
||||
api_key=api_key,
|
||||
|
|
@ -86,6 +95,7 @@ async def test_check_invalid_token_negative_cache_hit_short_circuits_even_if_db_
|
|||
|
||||
assert result is True
|
||||
find_first.assert_not_called()
|
||||
deprecated_find_first.assert_not_called()
|
||||
user_api_key_cache.async_set_cache.assert_not_called()
|
||||
|
||||
|
||||
|
|
@ -118,8 +128,12 @@ async def test_check_invalid_token_cache_miss_db_hit_allows_auth_flow():
|
|||
user_api_key_cache.async_set_cache = AsyncMock()
|
||||
|
||||
find_first = AsyncMock(return_value=MagicMock(token=hashed))
|
||||
deprecated_find_first = AsyncMock(return_value=None)
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_verificationtoken.find_first = find_first
|
||||
prisma_client.db.litellm_deprecatedverificationtoken.find_first = (
|
||||
deprecated_find_first
|
||||
)
|
||||
|
||||
result = await InvalidVirtualKeyCache.check_invalid_token(
|
||||
api_key=api_key,
|
||||
|
|
@ -130,5 +144,42 @@ async def test_check_invalid_token_cache_miss_db_hit_allows_auth_flow():
|
|||
|
||||
assert result is False
|
||||
find_first.assert_awaited_once_with(where={"token": hashed})
|
||||
deprecated_find_first.assert_not_called()
|
||||
user_api_key_cache.async_get_cache.assert_awaited_once_with(key=neg_key)
|
||||
user_api_key_cache.async_set_cache.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_invalid_token_allows_deprecated_key_in_grace_period():
|
||||
"""
|
||||
A rotated key can be absent from the active token table but still valid via
|
||||
LiteLLM_DeprecatedVerificationToken. Do not negative-cache that hash.
|
||||
"""
|
||||
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_get_cache = AsyncMock(return_value=None)
|
||||
user_api_key_cache.async_set_cache = AsyncMock()
|
||||
|
||||
find_first = AsyncMock(return_value=None)
|
||||
deprecated_find_first = AsyncMock(return_value=MagicMock(token=hashed))
|
||||
prisma_client = MagicMock()
|
||||
prisma_client.db.litellm_verificationtoken.find_first = find_first
|
||||
prisma_client.db.litellm_deprecatedverificationtoken.find_first = (
|
||||
deprecated_find_first
|
||||
)
|
||||
|
||||
result = 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_positive_ttl(),
|
||||
)
|
||||
|
||||
assert result is False
|
||||
find_first.assert_awaited_once_with(where={"token": hashed})
|
||||
deprecated_find_first.assert_awaited_once()
|
||||
user_api_key_cache.async_get_cache.assert_awaited_once_with(key=neg_key)
|
||||
user_api_key_cache.async_set_cache.assert_not_called()
|
||||
|
|
|
|||
|
|
@ -9726,3 +9726,91 @@ async def test_execute_virtual_key_regeneration_cache_invalidation_with_token_ha
|
|||
call_kwargs = mock_delete_cache.call_args.kwargs
|
||||
# The token hash should be passed as-is, NOT double-hashed
|
||||
assert call_kwargs["hashed_token"] == token_hash
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_virtual_key_regeneration_clears_invalid_token_cache_for_new_key():
|
||||
"""
|
||||
Custom regenerated keys may have been attempted before regeneration and
|
||||
negative-cached. Regeneration should clear that stale invalid-token cache entry.
|
||||
"""
|
||||
from litellm.proxy._types import RegenerateKeyRequest, hash_token
|
||||
from litellm.proxy.auth.reject_invalid_tokens import InvalidVirtualKeyCache
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
_execute_virtual_key_regeneration,
|
||||
)
|
||||
|
||||
old_token_hash = "old-token-hash"
|
||||
custom_new_key = "sk-custom-new-key-previously-invalid"
|
||||
new_token_hash = hash_token(token=custom_new_key)
|
||||
|
||||
existing_key = LiteLLM_VerificationToken(
|
||||
token=old_token_hash,
|
||||
user_id="user-1",
|
||||
models=["gpt-4"],
|
||||
team_id=None,
|
||||
max_budget=None,
|
||||
tags=None,
|
||||
)
|
||||
|
||||
class DictLikeResult:
|
||||
def __init__(self, data):
|
||||
self._data = data
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self._data.items())
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock(
|
||||
return_value=DictLikeResult(
|
||||
{"token": new_token_hash, "key_name": "sk-...alid", "user_id": "user-1"}
|
||||
)
|
||||
)
|
||||
mock_prisma_client.jsonify_object = MagicMock(side_effect=lambda data: data)
|
||||
|
||||
mock_user_api_key_cache = MagicMock()
|
||||
mock_user_api_key_cache.async_delete_cache = AsyncMock()
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
api_key="sk-admin",
|
||||
user_id="admin-user",
|
||||
)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._insert_deprecated_key",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_rotated_hook",
|
||||
new_callable=AsyncMock,
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data",
|
||||
new_callable=AsyncMock,
|
||||
return_value={},
|
||||
),
|
||||
):
|
||||
await _execute_virtual_key_regeneration(
|
||||
prisma_client=mock_prisma_client,
|
||||
key_in_db=existing_key,
|
||||
hashed_api_key=old_token_hash,
|
||||
key=old_token_hash,
|
||||
data=RegenerateKeyRequest(new_key=custom_new_key),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
litellm_changed_by=None,
|
||||
user_api_key_cache=mock_user_api_key_cache,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
mock_user_api_key_cache.async_delete_cache.assert_any_await(
|
||||
key=InvalidVirtualKeyCache._cache_key(new_token_hash)
|
||||
)
|
||||
mock_user_api_key_cache.async_delete_cache.assert_any_await(
|
||||
key=InvalidVirtualKeyCache._cache_key(old_token_hash)
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue