mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-10 22:24:51 +00:00
Merge pull request #19260 from BerriAI/litellm_deleted_keys_endpoint
[Feature] Status Query for Keys and Teams List
This commit is contained in:
@@ -2229,6 +2229,9 @@ class UserAPIKeyAuth(
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def check_api_key(cls, values):
|
||||
# If values is already an instance (not a dict), return it as-is
|
||||
if not isinstance(values, dict):
|
||||
return values
|
||||
if values.get("api_key") is not None:
|
||||
values.update(
|
||||
{"token": cls._safe_hash_litellm_api_key(values.get("api_key"))}
|
||||
@@ -3359,7 +3362,7 @@ class TeamListResponseObject(LiteLLM_TeamTable):
|
||||
|
||||
|
||||
class KeyListResponseObject(TypedDict, total=False):
|
||||
keys: List[Union[str, UserAPIKeyAuth]]
|
||||
keys: List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]]
|
||||
total_count: Optional[int]
|
||||
current_page: Optional[int]
|
||||
total_pages: Optional[int]
|
||||
|
||||
@@ -3150,12 +3150,14 @@ async def list_keys(
|
||||
),
|
||||
sort_order: str = Query(default="desc", description="Sort order ('asc' or 'desc')"),
|
||||
expand: Optional[List[str]] = Query(None, description="Expand related objects (e.g. 'user')"),
|
||||
status: Optional[str] = Query(None, description="Filter by status (e.g. 'deleted')"),
|
||||
) -> KeyListResponseObject:
|
||||
"""
|
||||
List all keys for a given user / team / organization.
|
||||
|
||||
Parameters:
|
||||
expand: Optional[List[str]] - Expand related objects (e.g. 'user' to include user information)
|
||||
status: Optional[str] - Filter by status. Currently supports "deleted" to query deleted keys.
|
||||
|
||||
Returns:
|
||||
{
|
||||
@@ -3177,6 +3179,15 @@ async def list_keys(
|
||||
verbose_proxy_logger.error("Database not connected")
|
||||
raise Exception("Database not connected")
|
||||
|
||||
# Validate status parameter
|
||||
if status is not None and status != "deleted":
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Invalid status value. Currently only 'deleted' is supported."
|
||||
},
|
||||
)
|
||||
|
||||
complete_user_info = await validate_key_list_check(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
user_id=user_id,
|
||||
@@ -3217,6 +3228,7 @@ async def list_keys(
|
||||
sort_by=sort_by,
|
||||
sort_order=sort_order,
|
||||
expand=expand,
|
||||
status=status,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug("Successfully prepared response")
|
||||
@@ -3230,7 +3242,7 @@ async def list_keys(
|
||||
message=getattr(e, "detail", f"error({str(e)})"),
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
param=getattr(e, "param", "None"),
|
||||
code=getattr(e, "status_code", status.HTTP_500_INTERNAL_SERVER_ERROR),
|
||||
code=getattr(e, "status_code", fastapi.status.HTTP_500_INTERNAL_SERVER_ERROR),
|
||||
)
|
||||
elif isinstance(e, ProxyException):
|
||||
raise e
|
||||
@@ -3238,7 +3250,7 @@ async def list_keys(
|
||||
message="Authentication Error, " + str(e),
|
||||
type=ProxyErrorTypes.internal_server_error,
|
||||
param=getattr(e, "param", "None"),
|
||||
code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
code=fastapi.status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
)
|
||||
|
||||
|
||||
@@ -3424,6 +3436,7 @@ async def _list_key_helper(
|
||||
sort_by: Optional[str] = None,
|
||||
sort_order: str = "desc",
|
||||
expand: Optional[List[str]] = None,
|
||||
status: Optional[str] = None,
|
||||
) -> KeyListResponseObject:
|
||||
"""
|
||||
Helper function to list keys
|
||||
@@ -3468,28 +3481,51 @@ async def _list_key_helper(
|
||||
else None
|
||||
)
|
||||
|
||||
# Determine which table to query based on status
|
||||
use_deleted_table = status == "deleted"
|
||||
|
||||
# Fetch keys with pagination
|
||||
keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where=where, # type: ignore
|
||||
skip=skip, # type: ignore
|
||||
take=size, # type: ignore
|
||||
order=(
|
||||
order_by
|
||||
if order_by
|
||||
else [
|
||||
{"created_at": "desc"},
|
||||
{"token": "desc"}, # fallback sort
|
||||
]
|
||||
),
|
||||
include={"object_permission": True},
|
||||
)
|
||||
if use_deleted_table:
|
||||
keys = await prisma_client.db.litellm_deletedverificationtoken.find_many(
|
||||
where=where, # type: ignore
|
||||
skip=skip, # type: ignore
|
||||
take=size, # type: ignore
|
||||
order=(
|
||||
order_by
|
||||
if order_by
|
||||
else [
|
||||
{"created_at": "desc"},
|
||||
{"token": "desc"}, # fallback sort
|
||||
]
|
||||
),
|
||||
)
|
||||
else:
|
||||
keys = await prisma_client.db.litellm_verificationtoken.find_many(
|
||||
where=where, # type: ignore
|
||||
skip=skip, # type: ignore
|
||||
take=size, # type: ignore
|
||||
order=(
|
||||
order_by
|
||||
if order_by
|
||||
else [
|
||||
{"created_at": "desc"},
|
||||
{"token": "desc"}, # fallback sort
|
||||
]
|
||||
),
|
||||
include={"object_permission": True},
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(f"Fetched {len(keys)} keys")
|
||||
|
||||
# Get total count of keys
|
||||
total_count = await prisma_client.db.litellm_verificationtoken.count(
|
||||
where=where # type: ignore
|
||||
)
|
||||
if use_deleted_table:
|
||||
total_count = await prisma_client.db.litellm_deletedverificationtoken.count(
|
||||
where=where # type: ignore
|
||||
)
|
||||
else:
|
||||
total_count = await prisma_client.db.litellm_verificationtoken.count(
|
||||
where=where # type: ignore
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(f"Total count of keys: {total_count}")
|
||||
|
||||
@@ -3507,18 +3543,31 @@ async def _list_key_helper(
|
||||
user_map = {user.user_id: user for user in users}
|
||||
|
||||
# Prepare response
|
||||
key_list: List[Union[str, UserAPIKeyAuth]] = []
|
||||
key_list: List[Union[str, UserAPIKeyAuth, LiteLLM_DeletedVerificationToken]] = []
|
||||
for key in keys:
|
||||
key_dict = key.dict()
|
||||
# Attach object_permission if object_permission_id is set
|
||||
key_dict = await attach_object_permission_to_dict(key_dict, prisma_client)
|
||||
# Convert Prisma model to dict (supports both Pydantic v1 and v2)
|
||||
try:
|
||||
key_dict = key.model_dump()
|
||||
except Exception:
|
||||
# Fallback for Pydantic v1 compatibility
|
||||
key_dict = key.dict()
|
||||
# Attach object_permission if object_permission_id is set (only for non-deleted keys)
|
||||
if not use_deleted_table:
|
||||
key_dict = await attach_object_permission_to_dict(key_dict, prisma_client)
|
||||
|
||||
# Include user information if expand includes "user"
|
||||
if expand and "user" in expand and key.user_id and key.user_id in user_map:
|
||||
key_dict["user"] = user_map[key.user_id].dict()
|
||||
try:
|
||||
key_dict["user"] = user_map[key.user_id].model_dump()
|
||||
except Exception:
|
||||
key_dict["user"] = user_map[key.user_id].dict()
|
||||
|
||||
if return_full_object is True or (expand and "user" in expand):
|
||||
key_list.append(UserAPIKeyAuth(**key_dict)) # Return full key object
|
||||
if use_deleted_table:
|
||||
# Use deleted key type to preserve deleted_at, deleted_by, etc.
|
||||
key_list.append(LiteLLM_DeletedVerificationToken(**key_dict))
|
||||
else:
|
||||
key_list.append(UserAPIKeyAuth(**key_dict)) # Return full key object
|
||||
else:
|
||||
_token = key_dict.get("token")
|
||||
key_list.append(cast(str, _token)) # Return only the token
|
||||
|
||||
@@ -2955,6 +2955,83 @@ async def list_available_teams(
|
||||
return available_teams_correct_type
|
||||
|
||||
|
||||
async def _build_team_list_where_conditions(
|
||||
prisma_client: PrismaClient,
|
||||
team_id: Optional[str],
|
||||
team_alias: Optional[str],
|
||||
organization_id: Optional[str],
|
||||
user_id: Optional[str],
|
||||
use_deleted_table: bool,
|
||||
) -> Dict[str, Any]:
|
||||
"""Build where conditions for team list query."""
|
||||
where_conditions: Dict[str, Any] = {}
|
||||
|
||||
if team_id:
|
||||
where_conditions["team_id"] = team_id
|
||||
|
||||
if team_alias:
|
||||
where_conditions["team_alias"] = {
|
||||
"contains": team_alias,
|
||||
"mode": "insensitive", # Case-insensitive search
|
||||
}
|
||||
|
||||
if organization_id:
|
||||
where_conditions["organization_id"] = organization_id
|
||||
|
||||
if user_id:
|
||||
try:
|
||||
user_object = await prisma_client.db.litellm_usertable.find_unique(
|
||||
where={"user_id": user_id}
|
||||
)
|
||||
except Exception:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"User not found, passed user_id={user_id}"},
|
||||
)
|
||||
if user_object is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"User not found, passed user_id={user_id}"},
|
||||
)
|
||||
user_object_correct_type = LiteLLM_UserTable(**user_object.model_dump())
|
||||
|
||||
if use_deleted_table:
|
||||
where_conditions["members"] = {"has": user_id}
|
||||
else:
|
||||
if team_id is None:
|
||||
where_conditions["team_id"] = {"in": user_object_correct_type.teams}
|
||||
elif team_id in user_object_correct_type.teams:
|
||||
where_conditions["team_id"] = team_id
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"User is not a member of team_id={team_id}"},
|
||||
)
|
||||
|
||||
return where_conditions
|
||||
|
||||
|
||||
def _convert_teams_to_response(
|
||||
teams: List[Any], use_deleted_table: bool
|
||||
) -> List[Union[LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]]:
|
||||
"""Convert Prisma models to Pydantic models."""
|
||||
team_list: List[Union[LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]] = []
|
||||
if teams:
|
||||
for team in teams:
|
||||
# Convert Prisma model to dict (supports both Pydantic v1 and v2)
|
||||
try:
|
||||
team_dict = team.model_dump()
|
||||
except Exception:
|
||||
# Fallback for Pydantic v1 compatibility
|
||||
team_dict = team.dict()
|
||||
if use_deleted_table:
|
||||
# Use deleted team type to preserve deleted_at, deleted_by, etc.
|
||||
team_list.append(LiteLLM_DeletedTeamTable(**team_dict))
|
||||
else:
|
||||
team_list.append(LiteLLM_TeamTable(**team_dict))
|
||||
return team_list
|
||||
|
||||
|
||||
@router.get(
|
||||
"/v2/team/list",
|
||||
tags=["team management"],
|
||||
@@ -2991,6 +3068,9 @@ async def list_team_v2(
|
||||
sort_order: str = fastapi.Query(
|
||||
default="asc", description="Sort order ('asc' or 'desc')"
|
||||
),
|
||||
status: Optional[str] = fastapi.Query(
|
||||
default=None, description="Filter by status (e.g. 'deleted')"
|
||||
),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
@@ -3013,6 +3093,8 @@ async def list_team_v2(
|
||||
Column to sort by (e.g. 'team_id', 'team_alias', 'created_at')
|
||||
sort_order: str
|
||||
Sort order ('asc' or 'desc')
|
||||
status: Optional[str]
|
||||
Filter by status. Currently supports "deleted" to query deleted teams.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
@@ -3037,50 +3119,28 @@ async def list_team_v2(
|
||||
if user_id is None and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
|
||||
user_id = user_api_key_dict.user_id
|
||||
|
||||
if status is not None and status != "deleted":
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"error": "Invalid status value. Currently only 'deleted' is supported."
|
||||
},
|
||||
)
|
||||
|
||||
use_deleted_table = status == "deleted"
|
||||
|
||||
# Calculate skip and take for pagination
|
||||
skip = (page - 1) * page_size
|
||||
|
||||
# Build where conditions based on provided parameters
|
||||
where_conditions: Dict[str, Any] = {}
|
||||
|
||||
if team_id:
|
||||
where_conditions["team_id"] = team_id
|
||||
|
||||
if team_alias:
|
||||
where_conditions["team_alias"] = {
|
||||
"contains": team_alias,
|
||||
"mode": "insensitive", # Case-insensitive search
|
||||
}
|
||||
|
||||
if organization_id:
|
||||
where_conditions["organization_id"] = organization_id
|
||||
|
||||
if user_id:
|
||||
try:
|
||||
user_object = await prisma_client.db.litellm_usertable.find_unique(
|
||||
where={"user_id": user_id}
|
||||
)
|
||||
except Exception:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"User not found, passed user_id={user_id}"},
|
||||
)
|
||||
if user_object is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"User not found, passed user_id={user_id}"},
|
||||
)
|
||||
user_object_correct_type = LiteLLM_UserTable(**user_object.model_dump())
|
||||
# Find teams where this user is a member by checking members_with_roles array
|
||||
if team_id is None:
|
||||
where_conditions["team_id"] = {"in": user_object_correct_type.teams}
|
||||
elif team_id in user_object_correct_type.teams:
|
||||
where_conditions["team_id"] = team_id
|
||||
else:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail={"error": f"User is not a member of team_id={team_id}"},
|
||||
)
|
||||
where_conditions = await _build_team_list_where_conditions(
|
||||
prisma_client=prisma_client,
|
||||
team_id=team_id,
|
||||
team_alias=team_alias,
|
||||
organization_id=organization_id,
|
||||
user_id=user_id,
|
||||
use_deleted_table=use_deleted_table,
|
||||
)
|
||||
|
||||
# Build order_by conditions
|
||||
valid_sort_columns = ["team_id", "team_alias", "created_at"]
|
||||
@@ -3091,20 +3151,35 @@ async def list_team_v2(
|
||||
order_by = {sort_by: sort_order.lower()}
|
||||
|
||||
# Get teams with pagination
|
||||
teams = await prisma_client.db.litellm_teamtable.find_many(
|
||||
where=where_conditions,
|
||||
skip=skip,
|
||||
take=page_size,
|
||||
order=order_by if order_by else {"created_at": "desc"}, # Default sort
|
||||
)
|
||||
# Get total count for pagination
|
||||
total_count = await prisma_client.db.litellm_teamtable.count(where=where_conditions)
|
||||
if use_deleted_table:
|
||||
teams = await prisma_client.db.litellm_deletedteamtable.find_many(
|
||||
where=where_conditions,
|
||||
skip=skip,
|
||||
take=page_size,
|
||||
order=order_by if order_by else {"created_at": "desc"}, # Default sort
|
||||
)
|
||||
# Get total count for pagination
|
||||
total_count = await prisma_client.db.litellm_deletedteamtable.count(
|
||||
where=where_conditions
|
||||
)
|
||||
else:
|
||||
teams = await prisma_client.db.litellm_teamtable.find_many(
|
||||
where=where_conditions,
|
||||
skip=skip,
|
||||
take=page_size,
|
||||
order=order_by if order_by else {"created_at": "desc"}, # Default sort
|
||||
)
|
||||
# Get total count for pagination
|
||||
total_count = await prisma_client.db.litellm_teamtable.count(where=where_conditions)
|
||||
|
||||
# Calculate total pages
|
||||
total_pages = -(-total_count // page_size) # Ceiling division
|
||||
|
||||
# Convert Prisma models to Pydantic models, preserving deleted fields when applicable
|
||||
team_list = _convert_teams_to_response(teams, use_deleted_table)
|
||||
|
||||
return {
|
||||
"teams": [team.model_dump() for team in teams] if teams else [],
|
||||
"teams": team_list,
|
||||
"total": total_count,
|
||||
"page": page,
|
||||
"page_size": page_size,
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from typing import Any, Dict, List, Optional
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_DeletedTeamTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_UserTable,
|
||||
@@ -45,7 +46,7 @@ class UpdateTeamMemberPermissionsRequest(BaseModel):
|
||||
class TeamListResponse(BaseModel):
|
||||
"""Response to get the list of teams"""
|
||||
|
||||
teams: List[LiteLLM_TeamTable]
|
||||
teams: List[Union[LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]]
|
||||
total: int
|
||||
page: int
|
||||
page_size: int
|
||||
|
||||
@@ -3517,6 +3517,7 @@ async def test_list_keys(prisma_client):
|
||||
sort_by=None,
|
||||
sort_order="desc",
|
||||
expand=None,
|
||||
status=None,
|
||||
)
|
||||
print("response=", response)
|
||||
assert "keys" in response
|
||||
@@ -3542,6 +3543,7 @@ async def test_list_keys(prisma_client):
|
||||
sort_by=None,
|
||||
sort_order="desc",
|
||||
expand=None,
|
||||
status=None,
|
||||
)
|
||||
print("pagination response=", response)
|
||||
assert len(response["keys"]) == 2
|
||||
@@ -3583,6 +3585,7 @@ async def test_list_keys(prisma_client):
|
||||
sort_by=None,
|
||||
sort_order="desc",
|
||||
expand=None,
|
||||
status=None,
|
||||
)
|
||||
print("filtered user_id response=", response)
|
||||
assert len(response["keys"]) == 1
|
||||
@@ -3605,6 +3608,7 @@ async def test_list_keys(prisma_client):
|
||||
sort_by=None,
|
||||
sort_order="desc",
|
||||
expand=None,
|
||||
status=None,
|
||||
)
|
||||
assert len(response["keys"]) == 1
|
||||
assert _key in response["keys"]
|
||||
|
||||
@@ -10,7 +10,7 @@ sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
||||
@@ -39,6 +39,7 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
|
||||
check_team_key_model_specific_limits,
|
||||
delete_verification_tokens,
|
||||
generate_key_helper_fn,
|
||||
list_keys,
|
||||
prepare_key_update_data,
|
||||
validate_key_team_change,
|
||||
)
|
||||
@@ -3799,47 +3800,61 @@ async def test_list_keys_with_expand_user():
|
||||
mock_prisma_client = AsyncMock()
|
||||
|
||||
# Create mock keys with user_ids
|
||||
mock_key1 = MagicMock()
|
||||
mock_key1.token = "token1"
|
||||
mock_key1.user_id = "user123"
|
||||
mock_key1.dict.return_value = {
|
||||
key1_dict = {
|
||||
"token": "token1",
|
||||
"user_id": "user123",
|
||||
"key_alias": "key1",
|
||||
"models": ["gpt-4"],
|
||||
}
|
||||
mock_key1 = MagicMock()
|
||||
mock_key1.token = "token1"
|
||||
mock_key1.user_id = "user123"
|
||||
# Set up model_dump() to raise AttributeError so it falls back to dict()
|
||||
mock_key1.model_dump = MagicMock(side_effect=AttributeError("model_dump not available"))
|
||||
mock_key1.dict = MagicMock(return_value=key1_dict)
|
||||
|
||||
mock_key2 = MagicMock()
|
||||
mock_key2.token = "token2"
|
||||
mock_key2.user_id = "user456"
|
||||
mock_key2.dict.return_value = {
|
||||
key2_dict = {
|
||||
"token": "token2",
|
||||
"user_id": "user456",
|
||||
"key_alias": "key2",
|
||||
"models": ["gpt-3.5-turbo"],
|
||||
}
|
||||
mock_key2 = MagicMock()
|
||||
mock_key2.token = "token2"
|
||||
mock_key2.user_id = "user456"
|
||||
# Set up model_dump() to raise AttributeError so it falls back to dict()
|
||||
mock_key2.model_dump = MagicMock(side_effect=AttributeError("model_dump not available"))
|
||||
mock_key2.dict = MagicMock(return_value=key2_dict)
|
||||
|
||||
mock_find_many_keys = AsyncMock(return_value=[mock_key1, mock_key2])
|
||||
mock_count_keys = AsyncMock(return_value=2)
|
||||
|
||||
# Create mock users
|
||||
mock_user1 = MagicMock()
|
||||
mock_user1.user_id = "user123"
|
||||
mock_user1.user_email = "user1@example.com"
|
||||
mock_user1.dict.return_value = {
|
||||
user1_dict = {
|
||||
"user_id": "user123",
|
||||
"user_email": "user1@example.com",
|
||||
"user_alias": "User One",
|
||||
}
|
||||
mock_user1 = MagicMock()
|
||||
# Set user_id as a real attribute (not a MagicMock)
|
||||
mock_user1.user_id = "user123"
|
||||
mock_user1.user_email = "user1@example.com"
|
||||
# Set up both model_dump() and dict() to return the same dict
|
||||
mock_user1.model_dump = MagicMock(return_value=user1_dict)
|
||||
mock_user1.dict = MagicMock(return_value=user1_dict)
|
||||
|
||||
mock_user2 = MagicMock()
|
||||
mock_user2.user_id = "user456"
|
||||
mock_user2.user_email = "user2@example.com"
|
||||
mock_user2.dict.return_value = {
|
||||
user2_dict = {
|
||||
"user_id": "user456",
|
||||
"user_email": "user2@example.com",
|
||||
"user_alias": "User Two",
|
||||
}
|
||||
mock_user2 = MagicMock()
|
||||
# Set user_id as a real attribute (not a MagicMock)
|
||||
mock_user2.user_id = "user456"
|
||||
mock_user2.user_email = "user2@example.com"
|
||||
# Set up both model_dump() and dict() to return the same dict
|
||||
mock_user2.model_dump = MagicMock(return_value=user2_dict)
|
||||
mock_user2.dict = MagicMock(return_value=user2_dict)
|
||||
|
||||
mock_find_many_users = AsyncMock(return_value=[mock_user1, mock_user2])
|
||||
|
||||
@@ -3847,6 +3862,108 @@ async def test_list_keys_with_expand_user():
|
||||
mock_prisma_client.db.litellm_verificationtoken.count = mock_count_keys
|
||||
mock_prisma_client.db.litellm_usertable.find_many = mock_find_many_users
|
||||
|
||||
# Patch attach_object_permission_to_dict to just return the dict unchanged
|
||||
async def mock_attach_object_permission(d, _):
|
||||
return d
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.attach_object_permission_to_dict",
|
||||
side_effect=mock_attach_object_permission,
|
||||
):
|
||||
args = {
|
||||
"prisma_client": mock_prisma_client,
|
||||
"page": 1,
|
||||
"size": 50,
|
||||
"user_id": None,
|
||||
"team_id": None,
|
||||
"organization_id": None,
|
||||
"key_alias": None,
|
||||
"key_hash": None,
|
||||
"exclude_team_id": None,
|
||||
"return_full_object": False, # This should be overridden by expand=user
|
||||
"admin_team_ids": None,
|
||||
"include_created_by_keys": False,
|
||||
"expand": ["user"], # Test the expand parameter
|
||||
}
|
||||
|
||||
result = await _list_key_helper(**args)
|
||||
|
||||
# Verify that keys were fetched
|
||||
mock_find_many_keys.assert_called_once()
|
||||
mock_count_keys.assert_called_once()
|
||||
|
||||
# Verify that users were fetched
|
||||
# Note: Order doesn't matter for the 'in' query, so we just check that both user_ids are present
|
||||
call_args = mock_find_many_users.call_args
|
||||
assert call_args is not None
|
||||
where_clause = call_args.kwargs["where"]
|
||||
assert "user_id" in where_clause
|
||||
assert "in" in where_clause["user_id"]
|
||||
user_ids_in_query = set(where_clause["user_id"]["in"])
|
||||
assert user_ids_in_query == {"user123", "user456"}
|
||||
|
||||
# Verify response structure
|
||||
assert len(result["keys"]) == 2
|
||||
assert result["total_count"] == 2
|
||||
assert result["current_page"] == 1
|
||||
assert result["total_pages"] == 1
|
||||
|
||||
# Verify that user data is included in the response
|
||||
# Since expand=user is specified, keys should be full objects
|
||||
assert isinstance(result["keys"][0], UserAPIKeyAuth)
|
||||
assert isinstance(result["keys"][1], UserAPIKeyAuth)
|
||||
|
||||
# Verify user data is attached to keys
|
||||
assert result["keys"][0].user == {
|
||||
"user_id": "user123",
|
||||
"user_email": "user1@example.com",
|
||||
"user_alias": "User One",
|
||||
}
|
||||
assert result["keys"][1].user == {
|
||||
"user_id": "user456",
|
||||
"user_email": "user2@example.com",
|
||||
"user_alias": "User Two",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_keys_with_status_deleted():
|
||||
"""
|
||||
Test that status="deleted" parameter correctly queries the deleted keys table.
|
||||
"""
|
||||
mock_prisma_client = AsyncMock()
|
||||
|
||||
# Mock deleted keys table
|
||||
mock_deleted_key1 = MagicMock()
|
||||
mock_deleted_key1.token = "deleted_token1"
|
||||
mock_deleted_key1.user_id = "user123"
|
||||
mock_deleted_key1.dict.return_value = {
|
||||
"token": "deleted_token1",
|
||||
"user_id": "user123",
|
||||
"key_alias": "deleted_key1",
|
||||
}
|
||||
|
||||
mock_deleted_key2 = MagicMock()
|
||||
mock_deleted_key2.token = "deleted_token2"
|
||||
mock_deleted_key2.user_id = "user456"
|
||||
mock_deleted_key2.dict.return_value = {
|
||||
"token": "deleted_token2",
|
||||
"user_id": "user456",
|
||||
"key_alias": "deleted_key2",
|
||||
}
|
||||
|
||||
mock_find_many_deleted = AsyncMock(return_value=[mock_deleted_key1, mock_deleted_key2])
|
||||
mock_count_deleted = AsyncMock(return_value=2)
|
||||
|
||||
# Mock regular keys table (should not be called)
|
||||
mock_find_many_regular = AsyncMock(return_value=[])
|
||||
mock_count_regular = AsyncMock(return_value=0)
|
||||
|
||||
mock_prisma_client.db.litellm_deletedverificationtoken.find_many = mock_find_many_deleted
|
||||
mock_prisma_client.db.litellm_deletedverificationtoken.count = mock_count_deleted
|
||||
mock_prisma_client.db.litellm_verificationtoken.find_many = mock_find_many_regular
|
||||
mock_prisma_client.db.litellm_verificationtoken.count = mock_count_regular
|
||||
|
||||
args = {
|
||||
"prisma_client": mock_prisma_client,
|
||||
"page": 1,
|
||||
@@ -3857,50 +3974,60 @@ async def test_list_keys_with_expand_user():
|
||||
"key_alias": None,
|
||||
"key_hash": None,
|
||||
"exclude_team_id": None,
|
||||
"return_full_object": False, # This should be overridden by expand=user
|
||||
"return_full_object": False,
|
||||
"admin_team_ids": None,
|
||||
"include_created_by_keys": False,
|
||||
"expand": ["user"], # Test the expand parameter
|
||||
"status": "deleted", # Test the status parameter
|
||||
}
|
||||
|
||||
|
||||
result = await _list_key_helper(**args)
|
||||
|
||||
# Verify that keys were fetched
|
||||
mock_find_many_keys.assert_called_once()
|
||||
mock_count_keys.assert_called_once()
|
||||
|
||||
# Verify that users were fetched
|
||||
# Note: Order doesn't matter for the 'in' query, so we just check that both user_ids are present
|
||||
call_args = mock_find_many_users.call_args
|
||||
assert call_args is not None
|
||||
where_clause = call_args.kwargs["where"]
|
||||
assert "user_id" in where_clause
|
||||
assert "in" in where_clause["user_id"]
|
||||
user_ids_in_query = set(where_clause["user_id"]["in"])
|
||||
assert user_ids_in_query == {"user123", "user456"}
|
||||
|
||||
|
||||
# Verify that deleted table was queried
|
||||
mock_find_many_deleted.assert_called_once()
|
||||
mock_count_deleted.assert_called_once()
|
||||
|
||||
# Verify that regular table was NOT queried
|
||||
mock_find_many_regular.assert_not_called()
|
||||
mock_count_regular.assert_not_called()
|
||||
|
||||
# Verify response structure
|
||||
assert len(result["keys"]) == 2
|
||||
assert result["total_count"] == 2
|
||||
assert result["current_page"] == 1
|
||||
assert result["total_pages"] == 1
|
||||
|
||||
# Verify that user data is included in the response
|
||||
# Since expand=user is specified, keys should be full objects
|
||||
assert isinstance(result["keys"][0], UserAPIKeyAuth)
|
||||
assert isinstance(result["keys"][1], UserAPIKeyAuth)
|
||||
|
||||
# Verify user data is attached to keys
|
||||
assert result["keys"][0].user == {
|
||||
"user_id": "user123",
|
||||
"user_email": "user1@example.com",
|
||||
"user_alias": "User One",
|
||||
}
|
||||
assert result["keys"][1].user == {
|
||||
"user_id": "user456",
|
||||
"user_email": "user2@example.com",
|
||||
"user_alias": "User Two",
|
||||
}
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_keys_with_invalid_status():
|
||||
"""
|
||||
Test that invalid status parameter raises ProxyException.
|
||||
"""
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
mock_prisma_client = AsyncMock()
|
||||
|
||||
# Mock the endpoint function directly to test validation
|
||||
from litellm.proxy.management_endpoints.key_management_endpoints import list_keys
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.utils import ProxyException
|
||||
|
||||
mock_request = Mock()
|
||||
mock_user_api_key_dict = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
|
||||
# Mock prisma_client to be non-None
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
|
||||
# Should raise ProxyException for invalid status (HTTPException is caught and re-raised as ProxyException)
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await list_keys(
|
||||
request=mock_request,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
status="invalid_status", # Invalid status value
|
||||
)
|
||||
|
||||
# Verify ProxyException properties
|
||||
assert exc_info.value.code == '400'
|
||||
assert "Invalid status value" in str(exc_info.value.message)
|
||||
assert "deleted" in str(exc_info.value.message)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
||||
@@ -2068,6 +2068,7 @@ async def test_list_team_v2_security_check_non_admin_user():
|
||||
http_request=mock_request,
|
||||
user_id=None, # Non-admin trying to query all teams
|
||||
user_api_key_dict=mock_user_api_key_dict_non_admin,
|
||||
status=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
@@ -2108,6 +2109,7 @@ async def test_list_team_v2_security_check_non_admin_user_other_user():
|
||||
http_request=mock_request,
|
||||
user_id="other_user_456", # Non-admin trying to query other user's teams
|
||||
user_api_key_dict=mock_user_api_key_dict_non_admin,
|
||||
status=None,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
@@ -2166,6 +2168,7 @@ async def test_list_team_v2_security_check_non_admin_user_own_teams():
|
||||
team_id=None,
|
||||
page=1,
|
||||
page_size=10,
|
||||
status=None,
|
||||
)
|
||||
|
||||
# Should return results without error
|
||||
@@ -2215,6 +2218,7 @@ async def test_list_team_v2_security_check_admin_user():
|
||||
user_api_key_dict=mock_user_api_key_dict_admin,
|
||||
page=1,
|
||||
page_size=10,
|
||||
status=None,
|
||||
)
|
||||
|
||||
# Should return results without error
|
||||
@@ -2223,6 +2227,110 @@ async def test_list_team_v2_security_check_admin_user():
|
||||
assert result["total"] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_team_v2_with_status_deleted():
|
||||
"""
|
||||
Test that status="deleted" parameter correctly queries the deleted teams table.
|
||||
"""
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.team_endpoints import list_team_v2
|
||||
|
||||
# Mock request
|
||||
mock_request = Mock(spec=Request)
|
||||
|
||||
# Mock admin user
|
||||
mock_user_api_key_dict_admin = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id="admin_user_123",
|
||||
)
|
||||
|
||||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client:
|
||||
# Mock prisma client and database operations
|
||||
mock_db = Mock()
|
||||
mock_prisma_client.db = mock_db
|
||||
|
||||
# Mock deleted teams
|
||||
mock_deleted_team1 = Mock(model_dump=lambda: {"team_id": "team_1", "team_alias": "Deleted Team 1"})
|
||||
mock_deleted_team2 = Mock(model_dump=lambda: {"team_id": "team_2", "team_alias": "Deleted Team 2"})
|
||||
|
||||
# Mock deleted teams table (should be called)
|
||||
mock_db.litellm_deletedteamtable.find_many = AsyncMock(return_value=[mock_deleted_team1, mock_deleted_team2])
|
||||
mock_db.litellm_deletedteamtable.count = AsyncMock(return_value=2)
|
||||
|
||||
# Mock regular teams table (should NOT be called)
|
||||
mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[])
|
||||
mock_db.litellm_teamtable.count = AsyncMock(return_value=0)
|
||||
|
||||
# Should NOT raise an exception
|
||||
result = await list_team_v2(
|
||||
http_request=mock_request,
|
||||
user_id=None, # Admin querying all teams
|
||||
user_api_key_dict=mock_user_api_key_dict_admin,
|
||||
page=1,
|
||||
page_size=10,
|
||||
status="deleted", # Test the status parameter
|
||||
)
|
||||
|
||||
# Verify that deleted table was queried
|
||||
mock_db.litellm_deletedteamtable.find_many.assert_called_once()
|
||||
mock_db.litellm_deletedteamtable.count.assert_called_once()
|
||||
|
||||
# Verify that regular table was NOT queried
|
||||
mock_db.litellm_teamtable.find_many.assert_not_called()
|
||||
mock_db.litellm_teamtable.count.assert_not_called()
|
||||
|
||||
# Should return results without error
|
||||
assert "teams" in result
|
||||
assert "total" in result
|
||||
assert result["total"] == 2
|
||||
assert len(result["teams"]) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_team_v2_with_invalid_status():
|
||||
"""
|
||||
Test that invalid status parameter raises HTTPException.
|
||||
"""
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.management_endpoints.team_endpoints import list_team_v2
|
||||
|
||||
# Mock request
|
||||
mock_request = Mock(spec=Request)
|
||||
|
||||
# Mock admin user
|
||||
mock_user_api_key_dict_admin = UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
user_id="admin_user_123",
|
||||
)
|
||||
|
||||
mock_prisma_client = Mock()
|
||||
|
||||
# Mock prisma_client to be non-None
|
||||
with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client):
|
||||
# Should raise HTTPException for invalid status
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await list_team_v2(
|
||||
http_request=mock_request,
|
||||
user_id=None,
|
||||
user_api_key_dict=mock_user_api_key_dict_admin,
|
||||
page=1,
|
||||
page_size=10,
|
||||
status="invalid_status", # Invalid status value
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "Invalid status value" in str(exc_info.value.detail)
|
||||
assert "deleted" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_member_delete_cleans_membership(mock_db_client, mock_admin_auth):
|
||||
"""
|
||||
|
||||
Reference in New Issue
Block a user