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
This commit is contained in:
Krish Dholakia 2025-07-21 22:06:29 -07:00 • committed by Krrish Dholakia
parent ac83c50137
commit e4ad56ab46
3 changed files with 177 additions and 8 deletions

View file

@ -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

View file

@ -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:

View file

@ -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"
)