From e4ad56ab461fafd5409dd5b32efeae8109c95451 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Mon, 21 Jul 2025 22:06:29 -0700 Subject: [PATCH] Fix team_member_budget update logic (#12843) * fix(team_endpoints.py): always remove team member budget from updated_kv this is not a field for the litellm team table Prevents startup issue * test(test_team_endpoints.py): add unit test to ensure 'team_member_budget' is never in update to table - separate logic * refactor: cleanup --- litellm/proxy/_types.py | 22 ++- .../management_endpoints/team_endpoints.py | 3 +- .../test_team_endpoints.py | 160 ++++++++++++++++++ 3 files changed, 177 insertions(+), 8 deletions(-) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b281c20b25b..5c66194d453 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -303,6 +303,11 @@ class LiteLLMRoutes(enum.Enum): "/v1/responses/{response_id}", "/responses/{response_id}/input_items", "/v1/responses/{response_id}/input_items", + # vector stores + "/vector_stores", + "/v1/vector_stores", + "/vector_stores/{vector_store_id}/search", + "/v1/vector_stores/{vector_store_id}/search", ] mapped_pass_through_routes = [ @@ -854,7 +859,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): command: Optional[str] = None args: List[str] = Field(default_factory=list) env: Dict[str, str] = Field(default_factory=dict) - + @model_validator(mode="before") @classmethod def validate_transport_fields(cls, values): @@ -871,7 +876,6 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase): return values - class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): server_id: str alias: Optional[str] = None @@ -886,7 +890,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): command: Optional[str] = None args: List[str] = Field(default_factory=list) env: Dict[str, str] = Field(default_factory=dict) - + @model_validator(mode="before") @classmethod def validate_transport_fields(cls, values): @@ -903,7 +907,6 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase): return values - class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): """Represents a LiteLLM_MCPServerTable record""" @@ -1762,11 +1765,15 @@ class UserAPIKeyAuth( @classmethod def check_api_key(cls, values): if values.get("api_key") is not None: - values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))}) + values.update( + {"token": cls._safe_hash_litellm_api_key(values.get("api_key"))} + ) if isinstance(values.get("api_key"), str): - values.update({"api_key": cls._safe_hash_litellm_api_key(values.get("api_key"))}) + values.update( + {"api_key": cls._safe_hash_litellm_api_key(values.get("api_key"))} + ) return values - + @classmethod def _safe_hash_litellm_api_key(cls, api_key: str) -> str: """ @@ -1778,6 +1785,7 @@ class UserAPIKeyAuth( if api_key.startswith("sk-"): return hash_token(api_key) from litellm.proxy.auth.handle_jwt import JWTHandler + if JWTHandler.is_jwt(token=api_key): return f"hashed-jwt-{hash_token(token=api_key)}" return api_key diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 1ba363d1fdb..40d37a3d937 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -216,7 +216,6 @@ async def _upsert_team_member_budget_table( user_api_key_dict=user_api_key_dict, team_member_budget=team_member_budget, ) - updated_kv.pop("team_member_budget", None) return updated_kv @@ -789,6 +788,8 @@ async def update_team( team_member_budget=data.team_member_budget, user_api_key_dict=user_api_key_dict, ) + else: + updated_kv.pop("team_member_budget", None) # Check object permission if data.object_permission is not None: diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index ec455a54176..aab70201ae6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -287,6 +287,9 @@ async def test_new_team_with_object_permission(mock_db_client, mock_admin_auth): mock_db_client.db.litellm_teamtable = MagicMock() mock_db_client.db.litellm_teamtable.create = mock_team_create mock_db_client.db.litellm_teamtable.count = mock_team_count + mock_db_client.db.litellm_teamtable.update = AsyncMock( + return_value=team_create_result + ) # 4. Mock user table update behaviour (called for each member) mock_db_client.db.litellm_usertable = MagicMock() @@ -1008,3 +1011,160 @@ def test_add_new_models_to_team_with_existing_models(): ) assert updated_models.sort() == ["model1", "model2", "model3", "model4"].sort() + + +@pytest.mark.asyncio +async def test_update_team_team_member_budget_not_passed_to_db(): + """ + Test that 'team_member_budget' is never passed to prisma_client.db.litellm_teamtable.update + regardless of whether the value is set or None. + + This ensures that team_member_budget is properly handled via the separate budget table + and not accidentally passed to the team table update operation. + """ + from unittest.mock import AsyncMock, MagicMock, Mock, patch + + from fastapi import Request + + from litellm.proxy._types import LitellmUserRoles, UpdateTeamRequest, UserAPIKeyAuth + from litellm.proxy.management_endpoints.team_endpoints import update_team + + # Mock dependencies + mock_request = Mock(spec=Request) + mock_user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="test_user_id" + ) + + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client, patch( + "litellm.proxy.proxy_server.llm_router" + ) as mock_llm_router, patch( + "litellm.proxy.proxy_server.user_api_key_cache" + ) as mock_cache, patch( + "litellm.proxy.proxy_server.proxy_logging_obj" + ) as mock_logging, patch( + "litellm.proxy.proxy_server.litellm_proxy_admin_name", "admin" + ), patch( + "litellm.proxy.auth.auth_checks._cache_team_object" + ) as mock_cache_team, patch( + "litellm.proxy.management_endpoints.team_endpoints._upsert_team_member_budget_table" + ) as mock_upsert_budget: + + # Setup mock prisma client + mock_existing_team = MagicMock() + mock_existing_team.model_dump.return_value = { + "team_id": "test_team_id", + "team_alias": "test_team", + "metadata": {"team_member_budget_id": "budget_123"}, + } + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock( + return_value=mock_existing_team + ) + + # Mock the update return value + mock_updated_team = MagicMock() + mock_updated_team.team_id = "test_team_id" + mock_updated_team.model_dump.return_value = {"team_id": "test_team_id"} + mock_prisma_client.db.litellm_teamtable.update = AsyncMock( + return_value=mock_updated_team + ) + mock_prisma_client.jsonify_team_object = MagicMock( + side_effect=lambda db_data: db_data + ) + + # Mock budget upsert to return updated_kv without team_member_budget + def mock_upsert_side_effect( + team_table, updated_kv, team_member_budget, user_api_key_dict + ): + # Remove team_member_budget from updated_kv as the real function does + result_kv = updated_kv.copy() + result_kv.pop("team_member_budget", None) + return result_kv + + mock_upsert_budget.side_effect = mock_upsert_side_effect + + # Test Case 1: team_member_budget is set (not None) + update_request_with_budget = UpdateTeamRequest( + team_id="test_team_id", team_member_budget=100.0, team_alias="updated_alias" + ) + + result = await update_team( + data=update_request_with_budget, + http_request=mock_request, + user_api_key_dict=mock_user_api_key_dict, + ) + + # Verify update was called + assert mock_prisma_client.db.litellm_teamtable.update.called + + # Get the call arguments + call_args = mock_prisma_client.db.litellm_teamtable.update.call_args + update_data = call_args[1]["data"] # data parameter from the update call + + # Verify team_member_budget is NOT in the update data + assert ( + "team_member_budget" not in update_data + ), f"team_member_budget should not be in update data, but found: {update_data}" + + # Verify other fields are present (team_alias should be there) + assert "team_alias" in update_data or "team_id" in str( + call_args + ), "Expected team update fields should be present" + + # Reset mock for second test + mock_prisma_client.db.litellm_teamtable.update.reset_mock() + + # Test Case 2: team_member_budget is None + update_request_without_budget = UpdateTeamRequest( + team_id="test_team_id", + team_member_budget=None, + team_alias="updated_alias_2", + ) + + result = await update_team( + data=update_request_without_budget, + http_request=mock_request, + user_api_key_dict=mock_user_api_key_dict, + ) + + # Verify update was called again + assert mock_prisma_client.db.litellm_teamtable.update.called + + # Get the call arguments for second call + call_args = mock_prisma_client.db.litellm_teamtable.update.call_args + update_data = call_args[1]["data"] # data parameter from the update call + + # Verify team_member_budget is NOT in the update data + assert ( + "team_member_budget" not in update_data + ), f"team_member_budget should not be in update data, but found: {update_data}" + + # Test Case 3: No team_member_budget field at all (excluded from request) + mock_prisma_client.db.litellm_teamtable.update.reset_mock() + + update_request_no_budget_field = UpdateTeamRequest( + team_id="test_team_id", + team_alias="updated_alias_3", + # team_member_budget not specified at all + ) + + result = await update_team( + data=update_request_no_budget_field, + http_request=mock_request, + user_api_key_dict=mock_user_api_key_dict, + ) + + # Verify update was called again + assert mock_prisma_client.db.litellm_teamtable.update.called + + # Get the call arguments for third call + call_args = mock_prisma_client.db.litellm_teamtable.update.call_args + update_data = call_args[1]["data"] # data parameter from the update call + + # Verify team_member_budget is NOT in the update data + assert ( + "team_member_budget" not in update_data + ), f"team_member_budget should not be in update data, but found: {update_data}" + + print( + "✅ All test cases passed: team_member_budget is properly excluded from database update operations" + )