From 6593566cdbe2978869ea70beafdcaf26d75c9d11 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 15:55:09 -0700 Subject: [PATCH] fix(key_management): invalidate cached object permissions on key update (#36719) * fix(key_management): invalidate cached object permissions on key update Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(key_management): broadcast permission cache eviction and cover key regenerate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(key_management): evict object permission rows through evict_and_broadcast Co-Authored-By: bot_apk * test(integration): cover mcp tool permission widen, narrow and clear on both workers Co-Authored-By: bot_apk * fix(key_management): evict object permission rows before the key object Co-Authored-By: bot_apk --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: bot_apk --- .../key_management_endpoints.py | 24 +- .../object_permission_utils.py | 20 +- tests/integration/mcp/test_mcp_lifecycle.py | 80 +++++- .../test_key_management_endpoints.py | 269 ++++++++++++++++++ 4 files changed, 386 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 1845f11c0fa..854b80ba2c3 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -108,6 +108,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( _set_object_permission, attach_object_permission_to_dict, handle_update_object_permission_common, + invalidate_cached_object_permissions, validate_key_mcp_servers_against_team, validate_key_search_tools_against_team, validate_key_vector_stores_against_team, @@ -2839,7 +2840,14 @@ async def _process_single_key_update( await prisma_client.update_data(token=key_request.key, data=_data), ) - # Delete cache + # Permission row first: a key-object miss between the two evictions would re-cache stale grants + await invalidate_cached_object_permissions( + object_permission_ids=( + existing_key_row.object_permission_id, + non_default_values.get("object_permission_id"), + ), + user_api_key_cache=user_api_key_cache, + ) await _delete_cache_key_object( hashed_token=_hash_token_if_needed(key_request.key), user_api_key_cache=user_api_key_cache, @@ -3455,6 +3463,13 @@ async def update_key_fn( # Delete - key from cache, since it's been updated! # key updated - a new model could have been added to this key. it should not block requests after this is done + await invalidate_cached_object_permissions( + object_permission_ids=( + existing_key_row.object_permission_id, + non_default_values.get("object_permission_id"), + ), + user_api_key_cache=user_api_key_cache, + ) await _delete_cache_key_object( hashed_token=_hash_token_if_needed(key), user_api_key_cache=user_api_key_cache, @@ -5569,6 +5584,13 @@ async def _execute_virtual_key_regeneration( updated_token_dict["key"] = new_token updated_token_dict["token_id"] = updated_token_dict.pop("token") + 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, + ) if hashed_api_key or key: await _delete_cache_key_object( hashed_token=_hash_token_if_needed(key), diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index 437e6763502..4aaa77f8d45 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -4,7 +4,7 @@ organizations, teams, and keys. """ import json -from collections.abc import Mapping, Sequence +from collections.abc import Iterable, Mapping, Sequence from collections.abc import Set as AbstractSet from dataclasses import dataclass from types import MappingProxyType @@ -17,6 +17,8 @@ from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import ObjectPermissionDict, SpecialMCPServerName, SpecialMCPServerNames +from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key from litellm.proxy.utils import PrismaClient from litellm.repositories.object_permission_repository import ObjectPermissionRepository from litellm.repositories.table_repositories import MCPServerRepository @@ -181,6 +183,22 @@ async def handle_update_object_permission_common( return created_object_permission_row.object_permission_id +async def invalidate_cached_object_permissions( + object_permission_ids: Iterable[object], + user_api_key_cache: UserApiKeyCache, +) -> None: + """Drop permission rows an entitlement change makes stale. + + ``get_object_permission`` caches a row under its own id separate from the entity's cache entry, and an + upsert keeps that id, so pass both the outgoing and incoming ids since a change can also mint a new row. + """ + 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)) + ) + await evict_and_broadcast(cache_keys, user_api_key_cache) + + async def _set_object_permission( data_json: dict, prisma_client: PrismaClient | None, diff --git a/tests/integration/mcp/test_mcp_lifecycle.py b/tests/integration/mcp/test_mcp_lifecycle.py index 814e1e769d8..fa253f03520 100644 --- a/tests/integration/mcp/test_mcp_lifecycle.py +++ b/tests/integration/mcp/test_mcp_lifecycle.py @@ -449,9 +449,79 @@ def test_key_grant_added_by_key_update_is_visible_to_mcp_tool_listing_before_the seen: Final = eventually( lambda: _granted_view(gateway, key), lambda view: view.tools != (), seconds=15, return_last_on_timeout=True ) - if seen.tools == (): - pytest.skip( - "BUG: a server granted through POST /key/update is missing from /mcp tools/list until the 60s " - "key cache TTL expires; no invalidation is published" - ) assert set(seen.tools) == {f"{alias}-add", f"{alias}-multiply", f"{alias}-fail"}, seen.raw + + +def _update_tool_permissions( + gateway: Gateway, key: str, identity: str, permissions: dict[str, list[str]] | None +) -> None: + updated: Final = gateway.request( + "POST", + "/key/update", + {"key": key, "object_permission": {"mcp_servers": [identity], "mcp_tool_permissions": permissions}}, + ) + assert updated.status_code == 200, updated.text + + +def _listing_on_both( + gateway: Gateway, peer: Gateway, key: str, expected: set[str] +) -> None: + for worker in (gateway, peer): + listing: Final = eventually( + functools.partial(_granted_view, worker, key), + functools.partial(_matches_grants, expected), + seconds=15, + return_last_on_timeout=True, + ) + assert set(listing.tools) == expected, (worker.client.base_url, listing.raw) + + +def _multiply_outcome_on_both(gateway: Gateway, peer: Gateway, key: str, alias: str) -> tuple[Outcome, Outcome]: + return ( + McpCaller(gateway, key, "mcp").call(f"{alias}-multiply", {"a": 2, "b": 3}), + McpCaller(peer, key, "mcp").call(f"{alias}-multiply", {"a": 2, "b": 3}), + ) + + +def test_key_update_tool_permission_widen_narrow_and_clear_apply_on_both_workers( + gateway: Gateway, peer: Gateway +) -> None: + with mcp_peer() as upstream, gateway.scenario() as scenario: + alias: Final = "perm" + uuid.uuid4().hex[:8] + identity: Final = register_mcp(scenario, upstream, alias) + key: Final = scenario.key( + object_permission={"mcp_servers": [identity], "mcp_tool_permissions": {identity: ["add"]}} + ) + add_only: Final = {f"{alias}-add"} + all_tools: Final = {f"{alias}-add", f"{alias}-multiply", f"{alias}-fail"} + upstream.drain() + + _listing_on_both(gateway, peer, key, add_only) + denied: Final = _multiply_outcome_on_both(gateway, peer, key, alias) + assert all(call.error is not None and call.text != "6" for call in denied), [call.raw for call in denied] + assert tool_calls(upstream.drain()) == (), "a denied call reached the peer" + + _update_tool_permissions(gateway, key, identity, {identity: ["add", "multiply"]}) + _listing_on_both(gateway, peer, key, {f"{alias}-add", f"{alias}-multiply"}) + widened: Final = _multiply_outcome_on_both(gateway, peer, key, alias) + assert [call.text for call in widened] == ["6", "6"], [call.raw for call in widened] + + _update_tool_permissions(gateway, key, identity, {identity: ["add"]}) + _listing_on_both(gateway, peer, key, add_only) + upstream.drain() + narrowed: Final = _multiply_outcome_on_both(gateway, peer, key, alias) + assert all(call.error is not None and call.text != "6" for call in narrowed), [call.raw for call in narrowed] + assert tool_calls(upstream.drain()) == (), "a revoked call reached the peer" + + _update_tool_permissions(gateway, key, identity, {}) + _listing_on_both(gateway, peer, key, all_tools) + cleared: Final = _multiply_outcome_on_both(gateway, peer, key, alias) + assert [call.text for call in cleared] == ["6", "6"], [call.raw for call in cleared] + + _update_tool_permissions(gateway, key, identity, {identity: ["add"]}) + _listing_on_both(gateway, peer, key, add_only) + + _update_tool_permissions(gateway, key, identity, None) + _listing_on_both(gateway, peer, key, all_tools) + nulled: Final = _multiply_outcome_on_both(gateway, peer, key, alias) + assert [call.text for call in nulled] == ["6", "6"], [call.raw for call in nulled] 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 cd929dfeeb3..69f66ca3939 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 @@ -20469,3 +20469,272 @@ async def test_generate_service_account_key_generates_uuid_when_no_alias(monkeyp assert data.metadata is not None assert data.metadata["service_account_id"] + +@pytest.mark.asyncio +async def test_key_update_invalidates_cached_object_permission(monkeypatch): + """Regression: /key/update must drop the cached permission row, not just the cached key. + + The permission row is cached under its own id and the upsert keeps that id, so a key read + after the update re-attached the OLD grants until the management-object TTL expired, which + served revoked MCP tools and withheld newly granted ones. + """ + from litellm.proxy._types import LiteLLM_ObjectPermissionBase + from litellm.proxy.auth.auth_checks import get_object_permission + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + permission_id = "objperm-lit5479" + grants = {"old": ["tool_a"], "new": ["tool_a", "tool_b"]} + + def _row(tools): + row = MagicMock() + row.dict.return_value = { + "object_permission_id": permission_id, + "mcp_tool_permissions": {"server-1": tools}, + } + return row + + mock_prisma_client = AsyncMock() + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + side_effect=lambda **kwargs: _row(grants["old"]) + ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=MagicMock(object_permission_id=permission_id) + ) + existing_key_row = LiteLLM_VerificationToken( + token="hashed-sk-lit5479", + user_id="user-123", + object_permission_id=permission_id, + ) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=existing_key_row + ) + updated_key = MagicMock() + updated_key.model_dump.return_value = {"user_id": "user-123"} + mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_key}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + 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": grants["old"]} + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook" + ): + await _process_single_key_update( + update_key_request=UpdateKeyRequest( + key="sk-lit5479", + object_permission=LiteLLM_ObjectPermissionBase( + mcp_tool_permissions={"server-1": grants["new"]} + ), + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin-user", + ), + litellm_changed_by=None, + prisma_client=mock_prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=AsyncMock(), + llm_router=None, + existing_key_row=existing_key_row, + ) + + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.side_effect = ( + lambda **kwargs: _row(grants["new"]) + ) + 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": 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", + ] + + +@pytest.mark.asyncio +async def test_key_update_evicts_object_permission_before_key_object(monkeypatch): + """The permission row must be evicted before the key object. + + ``get_key_object`` embeds the permission row in the cached key object, so a request landing + between the two evictions would otherwise re-cache stale grants for a full key TTL. + """ + from litellm.proxy._types import LiteLLM_ObjectPermissionBase + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, object_permission_cache_key + from litellm.proxy.utils import _hash_token_if_needed + + deleted: list[str] = [] + + class _RecordingCache(UserApiKeyCache): + def delete_cache(self, key: str) -> None: + deleted.append(key) + super().delete_cache(key) + + async def async_delete_cache(self, key: str) -> None: + deleted.append(key) + await super().async_delete_cache(key) + + permission_id = "objperm-order" + mock_prisma_client = AsyncMock() + existing_permission_row = MagicMock() + existing_permission_row.model_dump.return_value = { + "object_permission_id": permission_id, + "mcp_tool_permissions": {"server-1": ["tool_a"]}, + } + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=existing_permission_row + ) + mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock( + return_value=MagicMock(object_permission_id=permission_id) + ) + existing_key_row = LiteLLM_VerificationToken( + token="hashed-sk-lit5479", + user_id="user-123", + object_permission_id=permission_id, + ) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key_row) + updated_key = MagicMock() + updated_key.model_dump.return_value = {"user_id": "user-123"} + mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_key}) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + with patch( + "litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook" + ): + await _process_single_key_update( + update_key_request=UpdateKeyRequest( + key="sk-lit5479", + object_permission=LiteLLM_ObjectPermissionBase( + mcp_tool_permissions={"server-1": ["tool_a", "tool_b"]} + ), + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin-user", + ), + litellm_changed_by=None, + prisma_client=mock_prisma_client, + user_api_key_cache=_RecordingCache(), + proxy_logging_obj=AsyncMock(), + llm_router=None, + existing_key_row=existing_key_row, + ) + + assert deleted.index(object_permission_cache_key(permission_id)) < deleted.index( + _hash_token_if_needed("sk-lit5479") + ), deleted