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 <apk@cognition.ai>

* test(integration): cover mcp tool permission widen, narrow and clear on both workers

Co-Authored-By: bot_apk <apk@cognition.ai>

* fix(key_management): evict object permission rows before the key object

Co-Authored-By: bot_apk <apk@cognition.ai>

---------

Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: bot_apk <apk@cognition.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-09-23 15:55:09 -07:00 • committed by GitHub
parent 34f5874b93
commit 6593566cdb
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 386 additions and 7 deletions

View file

@ -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),

View file

@ -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,

View file

@ -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]

View file

@ -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