diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 3f2196a2e15..0f855369f84 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -4733,6 +4733,13 @@ async def _execute_virtual_key_regeneration( user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, ) + await invalidate_cached_object_permissions( + object_permission_ids=( + key_in_db.object_permission_id, + non_default_values.get("object_permission_id"), + ), + user_api_key_cache=user_api_key_cache, + ) response: Final = GenerateKeyResponse.model_validate(updated_token_dict) asyncio.create_task( diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index a0e9ac00aa9..4781dcf3440 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -182,14 +182,22 @@ async def invalidate_cached_object_permissions( the entity's cache entry re-attaches the old grants until the management-object TTL expires. Pass both the outgoing and incoming ids, since a change can also mint a new row. Non-string ids, which untyped update payloads can carry, are ignored. + + The local eviction only covers this worker and the shared backend, so each key is also broadcast: + other workers hold their own in-memory copy and would keep applying revoked grants until its TTL. """ - for object_permission_id in dict.fromkeys(pid for pid in object_permission_ids if isinstance(pid, str)): + from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import publish_auth_cache_invalidation + + cache_keys: Final = tuple( + object_permission_cache_key(object_permission_id) + for object_permission_id in dict.fromkeys(pid for pid in object_permission_ids if isinstance(pid, str)) + ) + for cache_key in cache_keys: try: - await user_api_key_cache.async_delete_cache(key=object_permission_cache_key(object_permission_id)) + await user_api_key_cache.async_delete_cache(key=cache_key) except Exception as e: # noqa: BLE001 # a cache we cannot clear still expires; never fail the write - verbose_proxy_logger.warning( - "Failed to invalidate cached object permission %r: %s", object_permission_id, e - ) + verbose_proxy_logger.warning("Failed to invalidate cached object permission %r: %s", cache_key, e) + await publish_auth_cache_invalidation(cache_key=cache_key) async def _set_object_permission( diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index dfab243f8cd..cabb91bdb4a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -15686,3 +15686,112 @@ async def test_key_update_invalidates_cached_object_permission(monkeypatch): ) assert reread is not None assert reread.mcp_tool_permissions == {"server-1": grants["new"]} + + +@pytest.mark.asyncio +async def test_key_regeneration_invalidates_cached_object_permission(monkeypatch): + """Regression: regenerating a key with new permissions must not keep serving the old grants.""" + from litellm.proxy._types import LiteLLM_ObjectPermissionBase, RegenerateKeyRequest + from litellm.proxy.auth.auth_checks import get_object_permission + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _execute_virtual_key_regeneration, + ) + + permission_id = "objperm-regenerate" + grants = {"served": ["tool_a"]} + + def _row(**kwargs): + row = MagicMock() + row.dict.return_value = { + "object_permission_id": permission_id, + "mcp_tool_permissions": {"server-1": grants["served"]}, + } + return row + + mock_prisma_client = _make_regenerate_mock_prisma() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(side_effect=_row) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=MagicMock(object_permission_id=permission_id) + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + existing_key = _make_regenerate_existing_key() + existing_key.object_permission_id = permission_id + user_api_key_cache = UserApiKeyCache() + assert ( + await get_object_permission( + object_permission_id=permission_id, + prisma_client=mock_prisma_client, + user_api_key_cache=user_api_key_cache, + ) + ).mcp_tool_permissions == {"server-1": ["tool_a"]} + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_new_token", + new_callable=AsyncMock, + return_value="sk-newtoken1234ab12", + ), + 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" + ), + ): + await _execute_virtual_key_regeneration( + prisma_client=mock_prisma_client, + key_in_db=existing_key, + hashed_api_key="abc123", + key="abc123", + data=RegenerateKeyRequest( + object_permission=LiteLLM_ObjectPermissionBase( + mcp_tool_permissions={"server-1": ["tool_a", "tool_b"]} + ) + ), + user_api_key_dict=_make_regenerate_user_api_key_dict(), + litellm_changed_by=None, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=AsyncMock(), + ) + + grants["served"] = ["tool_a", "tool_b"] + reread = await get_object_permission( + object_permission_id=permission_id, + prisma_client=mock_prisma_client, + user_api_key_cache=user_api_key_cache, + ) + assert reread is not None + assert reread.mcp_tool_permissions == {"server-1": ["tool_a", "tool_b"]} + + +@pytest.mark.asyncio +async def test_invalidate_cached_object_permissions_broadcasts_to_other_workers(): + """Other workers hold their own in-memory copy, so eviction has to be broadcast, not just local.""" + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.management_helpers.object_permission_utils import ( + invalidate_cached_object_permissions, + ) + + user_api_key_cache = UserApiKeyCache() + user_api_key_cache.async_delete_cache = AsyncMock(side_effect=Exception("redis down")) + + with patch( + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation", + new_callable=AsyncMock, + ) as mock_publish: + await invalidate_cached_object_permissions( + object_permission_ids=("objperm-old", "objperm-old", None, 42, "objperm-new"), + user_api_key_cache=user_api_key_cache, + ) + + assert [call.kwargs["cache_key"] for call in mock_publish.await_args_list] == [ + "object_permission_id:objperm-old", + "object_permission_id:objperm-new", + ]