diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 706f239a531..14414e1c5ea 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -686,9 +686,9 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase): allowed_cache_controls: Optional[list] = [] config: Optional[dict] = {} permissions: Optional[dict] = {} - model_max_budget: Optional[ - dict - ] = {} # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {} + model_max_budget: Optional[dict] = ( + {} + ) # {"gpt-4": 5.0, "gpt-3.5-turbo": 5.0}, defaults to {} model_config = ConfigDict(protected_namespaces=()) model_rpm_limit: Optional[dict] = None @@ -1027,12 +1027,12 @@ class NewCustomerRequest(BudgetNewRequest): blocked: bool = False # allow/disallow requests for this end-user budget_id: Optional[str] = None # give either a budget_id or max_budget spend: Optional[float] = None - allowed_model_region: Optional[ - AllowedModelRegion - ] = None # require all user requests to use models in this specific region - default_model: Optional[ - str - ] = None # if no equivalent model in allowed region - default all requests to this model + allowed_model_region: Optional[AllowedModelRegion] = ( + None # require all user requests to use models in this specific region + ) + default_model: Optional[str] = ( + None # if no equivalent model in allowed region - default all requests to this model + ) @model_validator(mode="before") @classmethod @@ -1054,12 +1054,12 @@ class UpdateCustomerRequest(LiteLLMPydanticObjectBase): blocked: bool = False # allow/disallow requests for this end-user max_budget: Optional[float] = None budget_id: Optional[str] = None # give either a budget_id or max_budget - allowed_model_region: Optional[ - AllowedModelRegion - ] = None # require all user requests to use models in this specific region - default_model: Optional[ - str - ] = None # if no equivalent model in allowed region - default all requests to this model + allowed_model_region: Optional[AllowedModelRegion] = ( + None # require all user requests to use models in this specific region + ) + default_model: Optional[str] = ( + None # if no equivalent model in allowed region - default all requests to this model + ) class DeleteCustomerRequest(LiteLLMPydanticObjectBase): @@ -1197,9 +1197,9 @@ class BlockKeyRequest(LiteLLMPydanticObjectBase): class AddTeamCallback(LiteLLMPydanticObjectBase): callback_name: str - callback_type: Optional[ - Literal["success", "failure", "success_and_failure"] - ] = "success_and_failure" + callback_type: Optional[Literal["success", "failure", "success_and_failure"]] = ( + "success_and_failure" + ) callback_vars: Dict[str, str] @model_validator(mode="before") @@ -1484,9 +1484,9 @@ class ConfigList(LiteLLMPydanticObjectBase): stored_in_db: Optional[bool] field_default_value: Any premium_field: bool = False - nested_fields: Optional[ - List[FieldDetail] - ] = None # For nested dictionary or Pydantic fields + nested_fields: Optional[List[FieldDetail]] = ( + None # For nested dictionary or Pydantic fields + ) class ConfigGeneralSettings(LiteLLMPydanticObjectBase): @@ -1756,9 +1756,9 @@ class LiteLLM_OrganizationMembershipTable(LiteLLMPydanticObjectBase): budget_id: Optional[str] = None created_at: datetime updated_at: datetime - user: Optional[ - Any - ] = None # You might want to replace 'Any' with a more specific type if available + user: Optional[Any] = ( + None # You might want to replace 'Any' with a more specific type if available + ) litellm_budget_table: Optional[LiteLLM_BudgetTable] = None model_config = ConfigDict(protected_namespaces=()) @@ -2423,6 +2423,11 @@ class ProxyErrorTypes(str, enum.Enum): Organization does not have access to the vector store """ + team_member_already_in_team = "team_member_already_in_team" + """ + Team member is already in team + """ + @classmethod def get_model_access_error_type_for_object( cls, object_type: Literal["key", "user", "team", "org"] @@ -2599,9 +2604,9 @@ class TeamModelDeleteRequest(BaseModel): # Organization Member Requests class OrganizationMemberAddRequest(OrgMemberAddRequest): organization_id: str - max_budget_in_organization: Optional[ - float - ] = None # Users max budget within the organization + max_budget_in_organization: Optional[float] = ( + None # Users max budget within the organization + ) class OrganizationMemberDeleteRequest(MemberDeleteRequest): @@ -2792,9 +2797,9 @@ class ProviderBudgetResponse(LiteLLMPydanticObjectBase): Maps provider names to their budget configs. """ - providers: Dict[ - str, ProviderBudgetResponseObject - ] = {} # Dictionary mapping provider names to their budget configurations + providers: Dict[str, ProviderBudgetResponseObject] = ( + {} + ) # Dictionary mapping provider names to their budget configurations class ProxyStateVariables(TypedDict): @@ -2922,9 +2927,9 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): enforce_rbac: bool = False roles_jwt_field: Optional[str] = None # v2 on role mappings role_mappings: Optional[List[RoleMapping]] = None - object_id_jwt_field: Optional[ - str - ] = None # can be either user / team, inferred from the role mapping + object_id_jwt_field: Optional[str] = ( + None # can be either user / team, inferred from the role mapping + ) scope_mappings: Optional[List[ScopeMapping]] = None enforce_scope_based_access: bool = False enforce_team_based_model_access: bool = False diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 80f05c7b0cb..3f7e7ec0f64 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -1,9 +1,9 @@ """ -Supports using JWT's for authenticating into the proxy. +Supports using JWT's for authenticating into the proxy. -Currently only supports admin. +Currently only supports admin. -JWT token must have 'litellm_proxy_admin' in scope. +JWT token must have 'litellm_proxy_admin' in scope. """ import json @@ -31,6 +31,8 @@ from litellm.proxy._types import ( LiteLLM_UserTable, LitellmUserRoles, Member, + ProxyErrorTypes, + ProxyException, ScopeMapping, Span, TeamMemberAddRequest, @@ -888,13 +890,25 @@ class JWTAuthManager: ), team_id=team_object.team_id, ) - # add user to team - await team_member_add( - data=data, - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN - ), # [TODO]: expose an internal service role, for better tracking - ) + # add user to team - make this non-blocking to avoid authentication failures + try: + await team_member_add( + data=data, + user_api_key_dict=UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN + ), # [TODO]: expose an internal service role, for better tracking + ) + verbose_proxy_logger.debug( + f"Successfully added user {user_object.user_id} to team {team_object.team_id}" + ) + except ProxyException as e: + if e.type == ProxyErrorTypes.team_member_already_in_team: + verbose_proxy_logger.debug( + f"User {user_object.user_id} is already a member of team {team_object.team_id}" + ) + return None + else: + raise e return None @staticmethod diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 52e8a3b7f4f..3c24036ebeb 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -65,8 +65,8 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.management_endpoints.common_utils import ( _is_user_team_admin, _set_object_metadata_field, - _user_has_admin_view, _upsert_budget_and_membership, + _user_has_admin_view, ) from litellm.proxy.management_endpoints.tag_management_endpoints import ( get_daily_activity, @@ -112,9 +112,7 @@ async def get_all_team_memberships( ) -> List[LiteLLM_TeamMembership]: """Get all team memberships for a given user""" ## GET ALL MEMBERSHIPS ## - where_obj: Dict[str, Dict[str, List[str]]] = { - "team_id": {"in": team_ids} - } + where_obj: Dict[str, Dict[str, List[str]]] = {"team_id": {"in": team_ids}} if user_id is not None: where_obj["user_id"] = {"in": [user_id]} # if user_id is None: @@ -699,12 +697,12 @@ async def update_team( updated_kv["model_id"] = _model_id updated_kv = prisma_client.jsonify_team_object(db_data=updated_kv) - team_row: Optional[ - LiteLLM_TeamTable - ] = await prisma_client.db.litellm_teamtable.update( - where={"team_id": data.team_id}, - data=updated_kv, - include={"litellm_model_table": True}, # type: ignore + team_row: Optional[LiteLLM_TeamTable] = ( + await prisma_client.db.litellm_teamtable.update( + where={"team_id": data.team_id}, + data=updated_kv, + include={"litellm_model_table": True}, # type: ignore + ) ) if team_row is None or team_row.team_id is None: @@ -825,11 +823,11 @@ def team_member_add_duplication_check( ): def _check_member_duplication(member: Member): if member.user_id in [m.user_id for m in existing_team_row.members_with_roles]: - raise HTTPException( - status_code=400, - detail={ - "error": f"User={member.user_id} already in team. Existing members={existing_team_row.members_with_roles}" - }, + raise ProxyException( + message=f"User={member.user_id} already in team. Existing members={existing_team_row.members_with_roles}", + type=ProxyErrorTypes.team_member_already_in_team, + param="user_id", + code="400", ) if isinstance(data.member, Member): @@ -1351,10 +1349,10 @@ async def delete_team( team_rows: List[LiteLLM_TeamTable] = [] for team_id in data.team_ids: try: - team_row_base: Optional[ - BaseModel - ] = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} + team_row_base: Optional[BaseModel] = ( + await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id} + ) ) if team_row_base is None: raise Exception @@ -1512,11 +1510,11 @@ async def team_info( ) try: - team_info: Optional[ - BaseModel - ] = await prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id}, - include={"object_permission": True}, + team_info: Optional[BaseModel] = ( + await prisma_client.db.litellm_teamtable.find_unique( + where={"team_id": team_id}, + include={"object_permission": True}, + ) ) if team_info is None: raise Exception diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index f2e33a57801..ca61b1c3fc5 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -2,7 +2,13 @@ from unittest.mock import AsyncMock, patch import pytest -from litellm.proxy._types import LiteLLM_TeamTable, LiteLLM_UserTable, Member +from litellm.proxy._types import ( + LiteLLM_TeamTable, + LiteLLM_UserTable, + Member, + ProxyErrorTypes, + ProxyException, +) from litellm.proxy.auth.handle_jwt import JWTAuthManager @@ -46,6 +52,64 @@ async def test_map_user_to_teams_add_new_user(): assert call_args.team_id == "test_team_1" +@pytest.mark.asyncio +async def test_map_user_to_teams_handles_already_in_team_exception(): + """Test that team_member_already_in_team exception is handled gracefully""" + # Setup test data + user = LiteLLM_UserTable(user_id="test_user_1") + team = LiteLLM_TeamTable(team_id="test_team_1", members_with_roles=[]) + + # Create a ProxyException with team_member_already_in_team error type + already_in_team_exception = ProxyException( + message="User test_user_1 already in team", + type=ProxyErrorTypes.team_member_already_in_team, + param="user_id", + code="400", + ) + + # Mock team_member_add to raise the exception + with patch( + "litellm.proxy.management_endpoints.team_endpoints.team_member_add", + new_callable=AsyncMock, + side_effect=already_in_team_exception, + ) as mock_add: + with patch("litellm.proxy.auth.handle_jwt.verbose_proxy_logger") as mock_logger: + # This should not raise an exception + result = await JWTAuthManager.map_user_to_teams( + user_object=user, team_object=team + ) + + # Verify the method completed successfully + assert result is None + mock_add.assert_called_once() + + +@pytest.mark.asyncio +async def test_map_user_to_teams_reraises_other_proxy_exceptions(): + """Test that other ProxyException types are re-raised""" + # Setup test data + user = LiteLLM_UserTable(user_id="test_user_1") + team = LiteLLM_TeamTable(team_id="test_team_1", members_with_roles=[]) + + # Create a ProxyException with a different error type + other_exception = ProxyException( + message="Some other error", + type=ProxyErrorTypes.internal_server_error, + param="some_param", + code="500", + ) + + # Mock team_member_add to raise the exception + with patch( + "litellm.proxy.management_endpoints.team_endpoints.team_member_add", + new_callable=AsyncMock, + side_effect=other_exception, + ) as mock_add: + # This should re-raise the exception + with pytest.raises(ProxyException) as exc_info: + await JWTAuthManager.map_user_to_teams(user_object=user, team_object=team) + + @pytest.mark.asyncio async def test_map_user_to_teams_null_inputs(): """Test that method handles null inputs gracefully""" 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 c7addf5deda..be4504fba99 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -18,6 +18,10 @@ from litellm.proxy._types import ( LiteLLM_OrganizationTable, LiteLLM_TeamTable, LitellmUserRoles, + Member, + ProxyErrorTypes, + ProxyException, + TeamMemberAddRequest, ) from litellm.proxy.management_endpoints.team_endpoints import ( user_api_key_auth, # Assuming this dependency is needed @@ -26,6 +30,7 @@ from litellm.proxy.management_endpoints.team_endpoints import ( GetTeamMemberPermissionsResponse, UpdateTeamMemberPermissionsRequest, router, + team_member_add_duplication_check, validate_team_org_change, ) from litellm.proxy.management_helpers.team_member_permission_checks import ( @@ -510,3 +515,69 @@ async def test_team_update_object_permissions_missing_permission_record(monkeypa # Verify upsert was called to create new record mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() + + +def test_team_member_add_duplication_check_raises_proxy_exception(): + """ + Test that team_member_add_duplication_check raises ProxyException when a user is already in the team + """ + # Create a mock team with existing members + existing_team_row = MagicMock(spec=LiteLLM_TeamTable) + existing_team_row.team_id = "test-team-123" + existing_team_row.members_with_roles = [ + Member(user_id="existing-user-id", role="user"), + Member(user_id="another-user-id", role="admin"), + ] + + # Create a request to add a member who is already in the team + duplicate_member = Member(user_id="existing-user-id", role="user") + data = TeamMemberAddRequest( + team_id="test-team-123", + member=duplicate_member, + ) + + # Test that ProxyException is raised with the correct error type + with pytest.raises(ProxyException) as exc_info: + team_member_add_duplication_check( + data=data, + existing_team_row=existing_team_row, + ) + + # Verify the exception details + assert exc_info.value.type == ProxyErrorTypes.team_member_already_in_team + assert exc_info.value.param == "user_id" + assert exc_info.value.code == "400" + assert "existing-user-id" in str(exc_info.value.message) + assert "already in team" in str(exc_info.value.message) + + +def test_team_member_add_duplication_check_allows_new_member(): + """ + Test that team_member_add_duplication_check allows adding a new member who is not already in the team + """ + # Create a mock team with existing members + existing_team_row = MagicMock(spec=LiteLLM_TeamTable) + existing_team_row.team_id = "test-team-123" + existing_team_row.members_with_roles = [ + Member(user_id="existing-user-id", role="user"), + Member(user_id="another-user-id", role="admin"), + ] + + # Create a request to add a member who is NOT already in the team + new_member = Member(user_id="new-user-id", role="user") + data = TeamMemberAddRequest( + team_id="test-team-123", + member=new_member, + ) + + # Test that no exception is raised for a new member + try: + team_member_add_duplication_check( + data=data, + existing_team_row=existing_team_row, + ) + # If we reach here, no exception was raised, which is expected + assert True + except ProxyException: + # If a ProxyException is raised, the test should fail + pytest.fail("ProxyException should not be raised for a new member")