[Fix] JWT - Fix error when team member already part of team (#11735)

* fix _check_member_duplication

* fix map_user_to_teams

* test_map_user_to_teams_handles_already_in_team_exception

* test_team_endpoints.py
This commit is contained in:
Ishaan Jaff 2025-06-14 15:50:16 -07:00 • committed by GitHub
parent 1c1e41c51d
commit c40580f892
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 220 additions and 68 deletions

View file

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

View file

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

View file

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

View file

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

View file

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