mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-26 01:12:21 +00:00
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:
parent
34f5874b93
commit
6593566cdb
4 changed files with 386 additions and 7 deletions
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue