fix(proxy): validate bulk object_permission against the key's team as /key/update does

This commit is contained in:
mateo-berri 2026-09-19 04:59:26 -07:00
parent f982d3e046
commit df6a222cb8
2 changed files with 78 additions and 14 deletions

View file

@ -2798,9 +2798,18 @@ async def _process_single_key_update(
llm_router=llm_router,
)
key_request: Final = await _with_validated_object_permission(
update_key_request=update_key_request,
team_obj=team_obj,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_api_key_dict=user_api_key_dict,
)
# Prepare update data
non_default_values = await prepare_key_update_data(
data=update_key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router
data=key_request, existing_key_row=existing_key_row, prisma_client=prisma_client, llm_router=llm_router
)
await _enforce_custom_key_policy(
@ -2809,7 +2818,7 @@ async def _process_single_key_update(
operation="update",
existing_key_row=existing_key_row,
non_default_values=non_default_values,
request=update_key_request,
request=key_request,
),
)
@ -2825,15 +2834,15 @@ async def _process_single_key_update(
existing_key_row=existing_key_row,
prisma_client=prisma_client,
)
_data: Final = {**update_values, "token": update_key_request.key}
_data: Final = {**update_values, "token": key_request.key}
response: Final[Mapping[str, object] | None] = cast( # cast-ok: every update_data branch returns a str-keyed dict
"Mapping[str, object] | None",
await prisma_client.update_data(token=update_key_request.key, data=_data),
await prisma_client.update_data(token=key_request.key, data=_data),
)
# Delete cache
await _delete_cache_key_object(
hashed_token=_hash_token_if_needed(update_key_request.key),
hashed_token=_hash_token_if_needed(key_request.key),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
@ -2842,17 +2851,15 @@ async def _process_single_key_update(
# authenticating against the access groups it just lost.
await sync_key_update_access_group_membership(
prisma_client=prisma_client,
key_token=_hash_token_if_needed(
_resolve_token_to_update(data=update_key_request, existing_key_row=existing_key_row)
),
data=update_key_request,
key_token=_hash_token_if_needed(_resolve_token_to_update(data=key_request, existing_key_row=existing_key_row)),
data=key_request,
existing_key_row=existing_key_row,
)
# Trigger async hook
asyncio.create_task(
KeyManagementEventHooks.async_key_updated_hook(
data=update_key_request,
data=key_request,
existing_key_row=existing_key_row,
response=response,
user_api_key_dict=user_api_key_dict,
@ -2875,6 +2882,31 @@ async def _process_single_key_update(
return updated_key_info
async def _with_validated_object_permission(
update_key_request: UpdateKeyRequest,
team_obj: LiteLLM_TeamTableCachedObj | None,
existing_key_row: LiteLLM_VerificationToken,
prisma_client: PrismaClient | None,
user_api_key_cache: UserApiKeyCache,
user_api_key_dict: UserAPIKeyAuth,
) -> UpdateKeyRequest:
if update_key_request.object_permission is None:
return update_key_request
normalized_object_permission: Final = await _validate_mcp_servers_for_key_update(
data=update_key_request,
team_obj=team_obj,
existing_key_row=existing_key_row,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
is_proxy_admin=user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value,
)
if normalized_object_permission is None:
return update_key_request
return update_key_request.model_copy(
update=MappingProxyType({"object_permission": LiteLLM_ObjectPermissionBase(**normalized_object_permission)})
)
async def _validate_mcp_servers_for_key_update(
data: "UpdateKeyRequest",
team_obj: Optional["LiteLLM_TeamTableCachedObj"],

View file

@ -80,7 +80,11 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
validate_key_team_change,
)
from litellm.proxy.proxy_server import app
from litellm.types.proxy.management_endpoints.key_management_endpoints import CustomKeyPolicyRequest
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyRequest,
BulkUpdateKeyResponse,
CustomKeyPolicyRequest,
)
client = TestClient(app)
@ -7098,23 +7102,29 @@ async def test_list_key_helper_applies_search_to_prisma_where():
_BULK_UPDATE_TOKEN: Final = "1f2e3d4c5b6a79880123456789abcdef0123456789abcdef0123456789abcdef"
_BULK_UPDATE_TEAM: Final = LiteLLM_TeamTableCachedObj(team_id="team-1")
async def _bulk_update_one_key(monkeypatch, item_payload: Mapping[str, object]) -> AsyncMock:
async def _run_bulk_update_on_one_key(
monkeypatch, item_payload: Mapping[str, object], team: LiteLLM_TeamTableCachedObj = _BULK_UPDATE_TEAM
) -> tuple[BulkUpdateKeyResponse, AsyncMock]:
from litellm.proxy.management_endpoints.key_management_endpoints import bulk_update_keys
from litellm.types.proxy.management_endpoints.key_management_endpoints import BulkUpdateKeyRequest
key_in_db = LiteLLM_VerificationToken(
token=_BULK_UPDATE_TOKEN, user_id="test-user", team_id="team-1", max_budget=100.0, budget_id="budget-1"
)
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=key_in_db)
mock_prisma_client.get_data = AsyncMock(return_value=key_in_db)
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_objectpermissiontable.upsert = AsyncMock(
return_value=MagicMock(object_permission_id="objperm-bulk")
)
mock_prisma_client.update_data = AsyncMock(return_value={"data": {"token": _BULK_UPDATE_TOKEN}})
_setup_update_key_mocks(monkeypatch, mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", AsyncMock(return_value=team)
)
with (
patch( # test-quality-ok: the handler reads the cache and hook singletons from module globals, no injection seam
@ -7138,8 +7148,13 @@ async def _bulk_update_one_key(monkeypatch, item_payload: Mapping[str, object])
litellm_changed_by=None,
)
return response, mock_prisma_client
async def _bulk_update_one_key(monkeypatch, item_payload: Mapping[str, object]) -> AsyncMock:
response, prisma = await _run_bulk_update_on_one_key(monkeypatch, item_payload)
assert response.failed_updates == []
return mock_prisma_client
return prisma
def _written_key_row(prisma: AsyncMock) -> Mapping[str, object]:
@ -7178,6 +7193,23 @@ async def test_bulk_update_keys_object_permission_is_granted_not_dropped(monkeyp
assert not {"max_budget", "team_id", "budget_id"} & written.keys()
@pytest.mark.asyncio
async def test_bulk_update_keys_object_permission_outside_the_team_allowlist_is_refused(monkeypatch):
"""A bulk item's object_permission is checked against the key's team exactly as /key/update
checks it, so a team key cannot be granted a search tool its team does not allow."""
team = LiteLLM_TeamTableCachedObj(
team_id="team-1",
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-team-1", search_tools=["team-search"]),
)
response, prisma = await _run_bulk_update_on_one_key(
monkeypatch, {"object_permission": {"search_tools": ["other-search"]}}, team=team
)
assert response.successful_updates == []
assert "not allowed by team 'team-1'" in response.failed_updates[0].failed_reason
prisma.update_data.assert_not_called()
@pytest.mark.asyncio
async def test_generate_key_negative_max_budget():
"""