diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 3a77325bd5a..9881482e02d 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2125,7 +2125,7 @@ def _validate_soft_budget_value(soft_budget: float | None) -> None: async def _update_key_soft_budget( - prisma_client: PrismaClient, + db: "Prisma", existing_key_row: LiteLLM_VerificationToken, soft_budget: float | None, changed_by: str, @@ -2137,7 +2137,7 @@ async def _update_key_soft_budget( "updated_by": changed_by, } budget_where: Final[prisma_types.LiteLLM_BudgetTableWhereUniqueInput] = {"budget_id": existing_budget_id} - await BudgetRepository(prisma_client).table.update(where=budget_where, data=budget_update) + await db.litellm_budgettable.update(where=budget_where, data=budget_update) return existing_budget_id if soft_budget is None: return None @@ -2146,24 +2146,20 @@ async def _update_key_soft_budget( "created_by": changed_by, "updated_by": changed_by, } - created_budget: Final[prisma_models.LiteLLM_BudgetTable] = await BudgetRepository(prisma_client).table.create( - data=budget_create - ) + created_budget: Final[prisma_models.LiteLLM_BudgetTable] = await db.litellm_budgettable.create(data=budget_create) return created_budget.budget_id async def _apply_soft_budget_update( data: UpdateKeyRequest, non_default_values: Mapping[str, object], - prisma_client: PrismaClient, + db: "Prisma", existing_key_row: LiteLLM_VerificationToken, changed_by: str, ) -> Mapping[str, object]: - if "soft_budget" not in data.model_fields_set: - return non_default_values remaining: Final = MappingProxyType({k: v for k, v in non_default_values.items() if k != "soft_budget"}) updated_budget_id: Final = await _update_key_soft_budget( - prisma_client=prisma_client, + db=db, existing_key_row=existing_key_row, soft_budget=data.soft_budget, changed_by=changed_by, @@ -2173,6 +2169,30 @@ async def _apply_soft_budget_update( return remaining +async def _update_key_row_with_soft_budget( + prisma_client: PrismaClient, + key: str, + data: UpdateKeyRequest, + non_default_values: Mapping[str, object], + existing_key_row: LiteLLM_VerificationToken, + changed_by: str, +) -> dict[str, object]: + hashed_token: Final = _hash_token_if_needed(key) + async with prisma_client.tx() as tx: + update_values: Final = await _apply_soft_budget_update( + data=data, + non_default_values=non_default_values, + db=tx, + existing_key_row=existing_key_row, + changed_by=changed_by, + ) + updated_row: Final = await tx.litellm_verificationtoken.update( + where={"token": hashed_token}, + data=with_settings_updated_at(prisma_client.jsonify_object({**update_values, "token": hashed_token})), + ) + return {"token": hashed_token, "data": updated_row.model_dump() if updated_row is not None else {}} + + async def prepare_key_update_data( data: UpdateKeyRequest | RegenerateKeyRequest, existing_key_row: LiteLLM_VerificationToken, @@ -3043,17 +3063,20 @@ async def update_key_fn( if prisma_client is None: raise Exception("Not connected to DB!") - update_values: Final = await _apply_soft_budget_update( - data=data, - non_default_values=non_default_values, - prisma_client=prisma_client, - existing_key_row=existing_key_row, - changed_by=user_api_key_dict.user_id or litellm_proxy_admin_name, + changed_by: Final = user_api_key_dict.user_id or litellm_proxy_admin_name + response: Final = ( + await _update_key_row_with_soft_budget( + prisma_client=prisma_client, + key=key, + data=data, + non_default_values=non_default_values, + existing_key_row=existing_key_row, + changed_by=changed_by, + ) + if "soft_budget" in data.model_fields_set + else await prisma_client.update_data(token=key, data={**non_default_values, "token": key}) ) - _data: Final = {**update_values, "token": key} - response: Final = await prisma_client.update_data(token=key, data=_data) - # 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 _delete_cache_key_object( diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index df0c9f953fc..ea5d4b711a9 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -17595,18 +17595,19 @@ async def test_update_key_soft_budget_updates_existing_budget_row(): ) existing_key = LiteLLM_VerificationToken(token="test-token", budget_id="budget-123") - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_budgettable.update = AsyncMock() + mock_db = MagicMock() + mock_db.litellm_budgettable.update = AsyncMock() + mock_db.litellm_budgettable.create = AsyncMock() result = await _update_key_soft_budget( - prisma_client=mock_prisma_client, + db=mock_db, existing_key_row=existing_key, soft_budget=25.0, changed_by="user-1", ) assert result == "budget-123" - mock_prisma_client.db.litellm_budgettable.update.assert_awaited_once_with( + mock_db.litellm_budgettable.update.assert_awaited_once_with( where={"budget_id": "budget-123"}, data={"soft_budget": 25.0, "updated_by": "user-1"}, ) @@ -17619,18 +17620,19 @@ async def test_update_key_soft_budget_clears_existing_budget_row(): ) existing_key = LiteLLM_VerificationToken(token="test-token", budget_id="budget-123") - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_budgettable.update = AsyncMock() + mock_db = MagicMock() + mock_db.litellm_budgettable.update = AsyncMock() + mock_db.litellm_budgettable.create = AsyncMock() result = await _update_key_soft_budget( - prisma_client=mock_prisma_client, + db=mock_db, existing_key_row=existing_key, soft_budget=None, changed_by="user-1", ) assert result == "budget-123" - mock_prisma_client.db.litellm_budgettable.update.assert_awaited_once_with( + mock_db.litellm_budgettable.update.assert_awaited_once_with( where={"budget_id": "budget-123"}, data={"soft_budget": None, "updated_by": "user-1"}, ) @@ -17645,18 +17647,19 @@ async def test_update_key_soft_budget_creates_budget_row_when_key_has_none(): existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None) created_row = MagicMock() created_row.budget_id = "budget-new" - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_budgettable.create = AsyncMock(return_value=created_row) + mock_db = MagicMock() + mock_db.litellm_budgettable.create = AsyncMock(return_value=created_row) + mock_db.litellm_budgettable.update = AsyncMock() result = await _update_key_soft_budget( - prisma_client=mock_prisma_client, + db=mock_db, existing_key_row=existing_key, soft_budget=10.5, changed_by="user-1", ) assert result == "budget-new" - mock_prisma_client.db.litellm_budgettable.create.assert_awaited_once_with( + mock_db.litellm_budgettable.create.assert_awaited_once_with( data={"soft_budget": 10.5, "created_by": "user-1", "updated_by": "user-1"} ) @@ -17668,20 +17671,20 @@ async def test_update_key_soft_budget_noop_when_clearing_without_budget_row(): ) existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None) - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_budgettable.create = AsyncMock() - mock_prisma_client.db.litellm_budgettable.update = AsyncMock() + mock_db = MagicMock() + mock_db.litellm_budgettable.create = AsyncMock() + mock_db.litellm_budgettable.update = AsyncMock() result = await _update_key_soft_budget( - prisma_client=mock_prisma_client, + db=mock_db, existing_key_row=existing_key, soft_budget=None, changed_by="user-1", ) assert result is None - mock_prisma_client.db.litellm_budgettable.create.assert_not_awaited() - mock_prisma_client.db.litellm_budgettable.update.assert_not_awaited() + mock_db.litellm_budgettable.create.assert_not_awaited() + mock_db.litellm_budgettable.update.assert_not_awaited() def test_update_key_request_accepts_soft_budget(): @@ -17712,28 +17715,6 @@ def test_validate_soft_budget_value_rejects_invalid_values(invalid_value): assert "soft_budget must be a non-negative finite number" in str(exc_info.value.detail) -@pytest.mark.asyncio -async def test_apply_soft_budget_update_noop_when_field_not_set(): - from litellm.proxy.management_endpoints.key_management_endpoints import ( - _apply_soft_budget_update, - ) - - existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None) - mock_prisma_client = MagicMock() - original_values = {"max_budget": 10.0} - - result = await _apply_soft_budget_update( - data=UpdateKeyRequest(key="sk-test", max_budget=10.0), - non_default_values=original_values, - prisma_client=mock_prisma_client, - existing_key_row=existing_key, - changed_by="user-1", - ) - - assert result is original_values - mock_prisma_client.db.litellm_budgettable.create.assert_not_called() - - @pytest.mark.asyncio async def test_apply_soft_budget_update_adds_budget_id_for_new_budget_row(): from litellm.proxy.management_endpoints.key_management_endpoints import ( @@ -17743,13 +17724,14 @@ async def test_apply_soft_budget_update_adds_budget_id_for_new_budget_row(): existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None) created_row = MagicMock() created_row.budget_id = "budget-created-456" - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_budgettable.create = AsyncMock(return_value=created_row) + mock_db = MagicMock() + mock_db.litellm_budgettable.create = AsyncMock(return_value=created_row) + mock_db.litellm_budgettable.update = AsyncMock() result = await _apply_soft_budget_update( data=UpdateKeyRequest(key="sk-test", soft_budget=25.0), non_default_values={"soft_budget": 25.0}, - prisma_client=mock_prisma_client, + db=mock_db, existing_key_row=existing_key, changed_by="user-1", ) @@ -17764,24 +17746,95 @@ async def test_apply_soft_budget_update_keeps_existing_budget_id_out_of_token_up ) existing_key = LiteLLM_VerificationToken(token="test-token", budget_id="budget-123") - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_budgettable.update = AsyncMock() + mock_db = MagicMock() + mock_db.litellm_budgettable.update = AsyncMock() + mock_db.litellm_budgettable.create = AsyncMock() result = await _apply_soft_budget_update( data=UpdateKeyRequest(key="sk-test", soft_budget=40.0), non_default_values={"soft_budget": 40.0, "max_budget": 100.0}, - prisma_client=mock_prisma_client, + db=mock_db, existing_key_row=existing_key, changed_by="user-1", ) assert dict(result) == {"max_budget": 100.0} - mock_prisma_client.db.litellm_budgettable.update.assert_awaited_once_with( + mock_db.litellm_budgettable.update.assert_awaited_once_with( where={"budget_id": "budget-123"}, data={"soft_budget": 40.0, "updated_by": "user-1"}, ) +@pytest.mark.asyncio +async def test_update_key_row_with_soft_budget_updates_budget_and_key_in_transaction(): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _update_key_row_with_soft_budget, + ) + + existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None) + created_row = MagicMock(budget_id="budget-new") + updated_row = MagicMock() + updated_row.model_dump.return_value = {"token": "hashed", "budget_id": "budget-new"} + tx = MagicMock() + tx.litellm_budgettable.create = AsyncMock(return_value=created_row) + tx.litellm_verificationtoken.update = AsyncMock(return_value=updated_row) + tx_context = MagicMock() + tx_context.__aenter__ = AsyncMock(return_value=tx) + tx_context.__aexit__ = AsyncMock(return_value=None) + prisma_client = MagicMock() + prisma_client.tx.return_value = tx_context + prisma_client.jsonify_object = lambda data: dict(data) + + result = await _update_key_row_with_soft_budget( + prisma_client=prisma_client, + key="sk-test", + data=UpdateKeyRequest(key="sk-test", soft_budget=25.0), + non_default_values={"soft_budget": 25.0}, + existing_key_row=existing_key, + changed_by="user-1", + ) + + assert set(result) == {"token", "data"} + assert result["data"] == {"token": "hashed", "budget_id": "budget-new"} + tx.litellm_verificationtoken.update.assert_awaited_once() + update_call = tx.litellm_verificationtoken.update.await_args + assert update_call.kwargs["where"] == {"token": result["token"]} + assert update_call.kwargs["data"]["budget_id"] == "budget-new" + assert "soft_budget" not in update_call.kwargs["data"] + + +@pytest.mark.asyncio +async def test_update_key_row_with_soft_budget_propagates_transaction_error(): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _update_key_row_with_soft_budget, + ) + + existing_key = LiteLLM_VerificationToken(token="test-token", budget_id=None) + created_row = MagicMock(budget_id="budget-new") + tx = MagicMock() + tx.litellm_budgettable.create = AsyncMock(return_value=created_row) + tx.litellm_verificationtoken.update = AsyncMock(side_effect=RuntimeError("update failed")) + tx_context = MagicMock() + tx_context.__aenter__ = AsyncMock(return_value=tx) + tx_context.__aexit__ = AsyncMock(return_value=None) + prisma_client = MagicMock() + prisma_client.tx.return_value = tx_context + prisma_client.jsonify_object = lambda data: dict(data) + + with pytest.raises(RuntimeError, match="update failed"): + await _update_key_row_with_soft_budget( + prisma_client=prisma_client, + key="sk-test", + data=UpdateKeyRequest(key="sk-test", soft_budget=25.0), + non_default_values={"soft_budget": 25.0}, + existing_key_row=existing_key, + changed_by="user-1", + ) + + tx_context.__aexit__.assert_awaited_once() + assert tx_context.__aexit__.await_args.args[0] is RuntimeError + + def test_generate_key_request_blank_team_id_is_personal(): """The UI Team-field clear submits team_id=""; it must count as no team (LIT-3925).""" from litellm.proxy._types import RegenerateKeyRequest