From 1d24f7fdbc7d92977485b9521ef08d8880952caf Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 3 Oct 2026 00:43:46 +0000 Subject: [PATCH] fix(proxy): run key project ownership right before the key row write Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../key_management_endpoints.py | 28 +++++------ .../management/test_project_lifecycle.py | 22 ++++++++- .../test_key_management_endpoints.py | 48 +++++++++++++++++-- 3 files changed, 78 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index df429b5acbc..c68316871f4 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1553,13 +1553,6 @@ async def _common_key_generation_helper( prisma_client=prisma_client, ) - if data.project_id is not None and prisma_client is not None: - await _check_key_project_team( - project_id=data.project_id, - key_team_id=data.team_id, - prisma_client=prisma_client, - ) - response = await generate_key_helper_fn(request_type="key", **data_json, table_name="key", llm_router=llm_router) response["soft_budget"] = data.soft_budget # include the user-input soft budget in the response @@ -4922,6 +4915,13 @@ async def generate_key_helper_fn( # the LiteLLM_VerificationToken table will increase in size if we don't do this check return user_data + if project_id is not None: + await _check_key_project_team( + project_id=project_id, + key_team_id=team_id, + prisma_client=prisma_client, + ) + ## CREATE KEY verbose_proxy_logger.debug( "prisma_client: Creating Key= %s", @@ -5742,6 +5742,13 @@ async def _execute_virtual_key_regeneration( prisma_client=prisma_client, ) update_data.update(update_values) + if data is not None: + await _check_key_project_team_on_mutation( + data=data, + existing_key_row=key_in_db, + prisma_client=prisma_client, + ) + jsonified_update_data: Final[Mapping[str, object]] = prisma_client.jsonify_object(data=update_data) # Snapshot before the token update: the FK cascade rewrites mapping rows to the new hash, @@ -5766,13 +5773,6 @@ async def _execute_virtual_key_regeneration( grace_period=data.grace_period if data else None, ) - if data is not None: - await _check_key_project_team_on_mutation( - data=data, - existing_key_row=key_in_db, - prisma_client=prisma_client, - ) - updated_token: Final[LiteLLM_VerificationToken | None] = await _prisma_table( VerificationTokenRepository(prisma_client) ).update( diff --git a/tests/integration/management/test_project_lifecycle.py b/tests/integration/management/test_project_lifecycle.py index 84cfebc72a4..e86a989271a 100644 --- a/tests/integration/management/test_project_lifecycle.py +++ b/tests/integration/management/test_project_lifecycle.py @@ -33,6 +33,20 @@ def _key_rows(key: str) -> list[dict[str, JsonValue]]: ) +def _deleted_key_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT token FROM "LiteLLM_DeletedVerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + +def _deprecated_key_rows(key: str) -> list[dict[str, JsonValue]]: + return read_rows( + 'SELECT token FROM "LiteLLM_DeprecatedVerificationToken" WHERE token = %s', + (sha256(key.encode()).hexdigest(),), + ) + + def _cli_session_token( user_id: str, team_id: str | None, @@ -289,13 +303,19 @@ def test_key_regenerate_rejects_foreign_project_without_changing_key(ownership_g key: Final = string_value(JSON_OBJECT.validate_json(generated.content)["key"]) scenario.cleanups.callback(scenario.delete_key, key) before: Final = _key_rows(key) + before_deleted: Final = _deleted_key_rows(key) + before_deprecated: Final = _deprecated_key_rows(key) assert len(before) == 1 + assert before_deleted == [] + assert before_deprecated == [] response: Final = ownership_gateway.request( - "POST", f"/key/{key}/regenerate", {"project_id": project_b} + "POST", f"/key/{key}/regenerate", {"project_id": project_b, "grace_period": "1h"} ) _discard_unexpected_key(ownership_gateway, response) assert response.status_code == 400, response.text assert _key_rows(key) == before + assert _deleted_key_rows(key) == before_deleted + assert _deprecated_key_rows(key) == before_deprecated def test_key_generation_rejects_missing_project_without_writing_key(ownership_gateway: Gateway) -> None: 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 8127d095755..7e5b4987953 100644 --- a/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_key_management_endpoints.py @@ -20671,6 +20671,37 @@ async def test_key_generation_organization_membership_error_precedes_project_own assert error.value.detail == "Caller is not a member of organization_id=org-not-member" +@pytest.mark.asyncio +async def test_key_generation_premium_permission_error_precedes_project_ownership( + monkeypatch: pytest.MonkeyPatch, +) -> None: + 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) + monkeypatch.setattr(litellm, "default_key_generate_params", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) + + with pytest.raises(HTTPException) as error: + await _common_key_generation_helper( + data=GenerateKeyRequest( + project_id=_OWNED_PROJECT, + team_id=_OWNERSHIP_KEY_TEAM, + permissions={"get_spend_routes": True}, + ), + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + ), + litellm_changed_by=None, + team_table=None, + ) + + assert error.value.status_code == 500 + assert error.value.detail == {"error": "Internal Server Error."} + mock_prisma_client.insert_data.assert_not_awaited() + + @pytest.mark.asyncio async def test_key_generation_duplicate_alias_error_precedes_project_ownership( monkeypatch: pytest.MonkeyPatch, @@ -21175,6 +21206,12 @@ async def test_regenerate_checks_project_team_ownership( mock_prisma_client.writer_db.litellm_projecttable.find_unique = AsyncMock( return_value=LiteLLM_ProjectTable(project_id=_OWNED_PROJECT, team_id="team-b") ) + deleted_history_table: Final = MagicMock() + deleted_history_table.create_many = AsyncMock() + mock_prisma_client.db.litellm_deletedverificationtoken = deleted_history_table + deprecated_key_table: Final = MagicMock() + deprecated_key_table.upsert = AsyncMock() + mock_prisma_client.db.litellm_deprecatedverificationtoken = deprecated_key_table user_api_key_cache: Final = await _cache_with_project(_OWNED_PROJECT, [], team_id="team-b") async def regenerate() -> None: @@ -21184,10 +21221,6 @@ async def test_regenerate_checks_project_team_ownership( 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, @@ -21198,7 +21231,10 @@ async def test_regenerate_checks_project_team_ownership( key_in_db=existing_key, hashed_api_key="abc123", key="abc123", - data=RegenerateKeyRequest(project_id=_OWNED_PROJECT), + data=RegenerateKeyRequest( + project_id=_OWNED_PROJECT, + grace_period="1h" if expected_status is not None else None, + ), user_api_key_dict=_make_regenerate_user_api_key_dict(), litellm_changed_by=None, user_api_key_cache=user_api_key_cache, @@ -21210,6 +21246,8 @@ async def test_regenerate_checks_project_team_ownership( await regenerate() assert error.value.status_code == expected_status assert "belongs to team team-b, but the key belongs to team-a" in str(error.value.detail) + deleted_history_table.create_many.assert_not_awaited() + deprecated_key_table.upsert.assert_not_awaited() else: await regenerate()