mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
ac83c50137
commit
e4ad56ab46
3 changed files with 177 additions and 8 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue