mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
[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:
parent
1c1e41c51d
commit
c40580f892
5 changed files with 220 additions and 68 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue