From 49c6c070fbb4c304263cc6df800fe1159d2504b7 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 2 Oct 2026 09:52:09 +0000 Subject: [PATCH] fix(proxy): validate bulk key team changes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../key_management_endpoints.py | 40 +++++++-- .../management/test_project_lifecycle.py | 45 +++++++++- .../test_key_management_endpoints.py | 89 +++++++++++++++++++ 3 files changed, 163 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 60d0f292548..0f9965f5629 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2761,6 +2761,23 @@ def _resolve_token_to_update(data: UpdateKeyRequest, existing_key_row: LiteLLM_V return existing_key_row.token +async def _check_single_key_update_team_permissions( + user_api_key_dict: UserAPIKeyAuth, + prisma_client: PrismaClient | None, + existing_key_row: LiteLLM_VerificationToken, + user_api_key_cache: UserApiKeyCache, +) -> None: + if prisma_client is None: + return + await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( + user_api_key_dict=user_api_key_dict, + route=KeyManagementRoutes.KEY_UPDATE, + prisma_client=prisma_client, + existing_key_row=existing_key_row, + user_api_key_cache=user_api_key_cache, + ) + + async def _process_single_key_update( update_key_request: UpdateKeyRequest, user_api_key_dict: UserAPIKeyAuth, @@ -2824,15 +2841,12 @@ async def _process_single_key_update( entity="key", ) - # Check team member permissions - if prisma_client is not None: - await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint( - user_api_key_dict=user_api_key_dict, - route=KeyManagementRoutes.KEY_UPDATE, - prisma_client=prisma_client, - existing_key_row=existing_key_row, - user_api_key_cache=user_api_key_cache, - ) + await _check_single_key_update_team_permissions( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + existing_key_row=existing_key_row, + user_api_key_cache=user_api_key_cache, + ) # Custom key update hook if user_custom_key_update is not None: @@ -2886,6 +2900,14 @@ async def _process_single_key_update( llm_router=llm_router, ) + if prisma_client is not None: + await _check_key_project_team_on_mutation( + data=update_key_request, + existing_key_row=existing_key_row, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + key_request: Final = await _with_validated_object_permission( update_key_request=update_key_request, team_obj=team_obj, diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index 867bb604162..b13ed839fc1 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -253,18 +253,59 @@ def test_key_regenerate_rejects_foreign_project_without_changing_key(ownership_g ) assert generated.status_code == 200, generated.text key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"]) + scenario.cleanups.callback(scenario.delete_key, key) before: Final = _key_rows(key) assert len(before) == 1 response: Final = ownership_gateway.request( "POST", f"/key/{key}/regenerate", {"project_id": project_b} ) _discard_unexpected_key(ownership_gateway, response) - if response.status_code != 200: - scenario.cleanups.callback(scenario.delete_key, key) assert response.status_code == 400, response.text assert _key_rows(key) == before +def test_key_bulk_update_rejects_foreign_team_project_and_preserves_key(ownership_gateway: Gateway) -> None: + with ownership_gateway.scenario() as scenario: + model: Final = scenario.model() + team_a: Final = scenario.team(models=[model]) + team_b: Final = scenario.team(models=[model]) + project_a: Final = scenario.project(team_a, models=[model]) + key: Final = scenario.key(team_id=team_a, project_id=project_a, models=[model]) + before: Final = _key_rows(key) + assert len(before) == 1 + assert before[0]["team_id"] == team_a + assert before[0]["project_id"] == project_a + + response: Final = ownership_gateway.request( + "POST", + "/key/bulk_update", + { + "keys": [ + {"key": key, "team_id": team_b}, + {"key": key, "max_budget": 10, "tags": ["bulk-update"]}, + ] + }, + ) + assert response.status_code == 200, response.text + body: Final = JSON_OBJECT.validate_json(response.content) + failed_updates: Final = body["failed_updates"] + successful_updates: Final = body["successful_updates"] + assert isinstance(failed_updates, list), response.text + assert isinstance(successful_updates, list), response.text + assert len(failed_updates) == 1 + assert len(successful_updates) == 1 + failed_update: Final = object_value(failed_updates[0]) + successful_update: Final = object_value(successful_updates[0]) + assert string_value(failed_update["key"]) == key + assert f"Project {project_a} belongs to team {team_a}" in string_value(failed_update["failed_reason"]) + assert string_value(successful_update["key"]) == key + + after: Final = _key_rows(key) + assert len(after) == 1 + assert after[0]["team_id"] == before[0]["team_id"] + assert after[0]["project_id"] == before[0]["project_id"] + + def test_project_update_rejects_moving_project_with_attached_key(ownership_gateway: Gateway) -> None: with ownership_gateway.scenario() as scenario: model: Final = scenario.model() diff --git a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py index b4284d3d7a9..df4f159f156 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -20745,6 +20745,95 @@ async def test_key_update_rejects_team_change_for_project_bound_key(monkeypatch: assert expected_detail in str(error.value.detail) +@pytest.mark.asyncio +async def test_bulk_key_update_rejects_project_team_change_and_allows_other_fields( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy.management_endpoints.key_management_endpoints import bulk_update_keys + from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateKeyRequest, + BulkUpdateKeyRequestItem, + ) + + user_api_key_cache: Final = await _cache_with_project( + _OWNED_PROJECT, [], team_id=_OWNERSHIP_PROJECT_TEAM + ) + mock_prisma_client: Final = _configure_key_endpoints(monkeypatch, user_api_key_cache) + existing_key_row: Final = LiteLLM_VerificationToken( + token="hashed-bulk-key", + user_id=None, + models=[], + team_id=_OWNERSHIP_PROJECT_TEAM, + project_id=_OWNED_PROJECT, + ) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=existing_key_row) + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM) + ) + mock_prisma_client.get_data = AsyncMock(return_value=existing_key_row) + updated_key: Final = MagicMock() + updated_key.model_dump.return_value = { + "max_budget": 10.0, + "tags": ["bulk-update"], + "team_id": _OWNERSHIP_PROJECT_TEAM, + "project_id": _OWNED_PROJECT, + } + mock_prisma_client.update_data = AsyncMock(return_value={"data": updated_key}) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + + request: Final = BulkUpdateKeyRequest( + keys=[ + BulkUpdateKeyRequestItem(key="sk-bulk-key", team_id=_OWNERSHIP_DESTINATION_TEAM), + BulkUpdateKeyRequestItem(key="sk-bulk-key", max_budget=10.0, tags=["bulk-update"]), + ] + ) + user_api_key_dict: Final = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + user_id="admin", + ) + + with ( + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + new_callable=AsyncMock, + return_value=LiteLLM_TeamTable(team_id=_OWNERSHIP_DESTINATION_TEAM), + ), + patch("litellm.proxy.management_endpoints.common_utils._premium_user_check"), + patch("litellm.proxy.utils._premium_user_check"), + 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_updated_hook", + new_callable=AsyncMock, + ), + ): + response: Final = await bulk_update_keys( + data=request, + user_api_key_dict=user_api_key_dict, + litellm_changed_by=None, + ) + + assert response.total_requested == 2 + assert [(item.key, item.failed_reason) for item in response.failed_updates] == [ + ( + "sk-bulk-key", + ( + f"Project {_OWNED_PROJECT} belongs to team {_OWNERSHIP_PROJECT_TEAM}, but the key belongs to " + f"{_OWNERSHIP_DESTINATION_TEAM}. A key can only be attached to a project owned by its own team." + ), + ) + ] + assert len(response.successful_updates) == 1 + assert response.successful_updates[0].key == "sk-bulk-key" + assert response.successful_updates[0].key_info["max_budget"] == 10.0 + assert response.successful_updates[0].key_info["tags"] == ["bulk-update"] + mock_prisma_client.update_data.assert_awaited_once() + + @pytest.mark.parametrize( "request_fields", [{"key_alias": "renamed"}, {"team_id": "team-a"}],