mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
Merge branch 'BerriAI:main' into LangfuseUsageDetails
This commit is contained in:
commit
c5aa6f540f
6 changed files with 649 additions and 41 deletions
|
|
@ -203,7 +203,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
await self._async_log_event_base(
|
||||
kwargs=kwargs,
|
||||
|
|
@ -212,7 +212,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
end_time=end_time,
|
||||
)
|
||||
pass
|
||||
|
||||
|
||||
async def _async_log_event_base(self, kwargs, response_obj, start_time, end_time):
|
||||
try:
|
||||
|
|
@ -242,7 +241,6 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
verbose_logger.exception(f"s3 Layer Error - {str(e)}")
|
||||
pass
|
||||
|
||||
|
||||
async def async_upload_data_to_s3(
|
||||
self, batch_logging_element: s3BatchLoggingElement
|
||||
):
|
||||
|
|
@ -277,8 +275,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
# Prepare the URL
|
||||
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}"
|
||||
|
||||
if self.s3_endpoint_url:
|
||||
url = self.s3_endpoint_url + "/" + batch_logging_element.s3_object_key
|
||||
if self.s3_endpoint_url and self.s3_bucket_name:
|
||||
url = (
|
||||
self.s3_endpoint_url
|
||||
+ "/"
|
||||
+ self.s3_bucket_name
|
||||
+ "/"
|
||||
+ batch_logging_element.s3_object_key
|
||||
)
|
||||
|
||||
# Convert JSON to string
|
||||
json_string = safe_dumps(batch_logging_element.payload)
|
||||
|
|
@ -420,8 +424,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
# Prepare the URL
|
||||
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{batch_logging_element.s3_object_key}"
|
||||
|
||||
if self.s3_endpoint_url:
|
||||
url = self.s3_endpoint_url + "/" + batch_logging_element.s3_object_key
|
||||
if self.s3_endpoint_url and self.s3_bucket_name:
|
||||
url = (
|
||||
self.s3_endpoint_url
|
||||
+ "/"
|
||||
+ self.s3_bucket_name
|
||||
+ "/"
|
||||
+ batch_logging_element.s3_object_key
|
||||
)
|
||||
|
||||
# Convert JSON to string
|
||||
json_string = safe_dumps(batch_logging_element.payload)
|
||||
|
|
@ -462,14 +472,13 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
except Exception as e:
|
||||
verbose_logger.exception(f"Error uploading to s3: {str(e)}")
|
||||
|
||||
|
||||
async def _download_object_from_s3(self, s3_object_key: str) -> Optional[dict]:
|
||||
"""
|
||||
Download and parse JSON object from S3.
|
||||
|
||||
|
||||
Args:
|
||||
s3_object_key: The S3 object key to download
|
||||
|
||||
|
||||
Returns:
|
||||
Optional[dict]: The parsed JSON object or None if not found/error
|
||||
"""
|
||||
|
|
@ -481,7 +490,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
from botocore.awsrequest import AWSRequest
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call S3. Run 'pip install boto3'.")
|
||||
|
||||
|
||||
try:
|
||||
from litellm.litellm_core_utils.asyncify import asyncify
|
||||
|
||||
|
|
@ -506,8 +515,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
# Prepare the URL
|
||||
url = f"https://{self.s3_bucket_name}.s3.{self.s3_region_name}.amazonaws.com/{s3_object_key}"
|
||||
|
||||
if self.s3_endpoint_url:
|
||||
url = self.s3_endpoint_url + "/" + s3_object_key
|
||||
if self.s3_endpoint_url and self.s3_bucket_name:
|
||||
url = (
|
||||
self.s3_endpoint_url
|
||||
+ "/"
|
||||
+ self.s3_bucket_name
|
||||
+ "/"
|
||||
+ s3_object_key
|
||||
)
|
||||
|
||||
# Prepare the request for GET operation
|
||||
# For GET requests, we need x-amz-content-sha256 with hash of empty string
|
||||
|
|
@ -533,12 +548,14 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
response = await self.async_httpx_client.get(url, headers=signed_headers)
|
||||
|
||||
if response.status_code != 200:
|
||||
verbose_logger.exception("S3 object not found, saw response=", response.text)
|
||||
verbose_logger.exception(
|
||||
"S3 object not found, saw response=", response.text
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
# Parse JSON response
|
||||
return response.json()
|
||||
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error downloading from S3: {str(e)}")
|
||||
return None
|
||||
|
|
@ -551,11 +568,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
Get the proxy server request from cold storage
|
||||
|
||||
Allows fetching a dict of the proxy server request from s3 or GCS bucket.
|
||||
|
||||
|
||||
Args:
|
||||
request_id: The unique request ID to search for
|
||||
start_time: The start time of the request (datetime or ISO string)
|
||||
|
||||
|
||||
Returns:
|
||||
Optional[dict]: The request data dictionary or None if not found
|
||||
"""
|
||||
|
|
@ -564,5 +581,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
|
|||
downloaded_object = await self._download_object_from_s3(object_key)
|
||||
return downloaded_object
|
||||
except Exception as e:
|
||||
verbose_logger.exception(f"Error retrieving object {object_key} from cold storage: {str(e)}")
|
||||
return None
|
||||
verbose_logger.exception(
|
||||
f"Error retrieving object {object_key} from cold storage: {str(e)}"
|
||||
)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from litellm.types.files import (
|
|||
get_file_type_from_extension,
|
||||
is_gemini_1_5_accepted_file_type,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionAssistantMessage,
|
||||
|
|
@ -492,7 +493,8 @@ def _transform_request_body(
|
|||
data["generationConfig"] = generation_config
|
||||
if cached_content is not None:
|
||||
data["cachedContent"] = cached_content
|
||||
if labels is not None:
|
||||
# Only add labels for Vertex AI endpoints (not Google GenAI/AI Studio) and only if non-empty
|
||||
if labels and custom_llm_provider != LlmProviders.GEMINI:
|
||||
data["labels"] = labels
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
@ -647,3 +649,5 @@ def _transform_system_message(
|
|||
return SystemInstructions(parts=system_content_blocks), messages
|
||||
|
||||
return None, messages
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from fastapi import (
|
|||
Response,
|
||||
)
|
||||
from typing_extensions import TypedDict
|
||||
from pydantic import BaseModel
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -29,6 +30,7 @@ from litellm.proxy._types import (
|
|||
Member,
|
||||
NewTeamRequest,
|
||||
NewUserRequest,
|
||||
NewUserResponse,
|
||||
TeamMemberAddRequest,
|
||||
TeamMemberDeleteRequest,
|
||||
UserAPIKeyAuth,
|
||||
|
|
@ -101,6 +103,13 @@ class ScimUserData(TypedDict):
|
|||
active: Optional[bool]
|
||||
|
||||
|
||||
class GroupMemberExtractionResult(BaseModel):
|
||||
"""Result of extracting and processing group members."""
|
||||
existing_member_ids: List[str]
|
||||
created_users: List[NewUserResponse]
|
||||
all_member_ids: List[str] # existing + newly created
|
||||
|
||||
|
||||
scim_router = APIRouter(
|
||||
prefix="/scim/v2",
|
||||
tags=["✨ SCIM v2 (Enterprise Only)"],
|
||||
|
|
@ -190,21 +199,47 @@ def _build_scim_metadata(given_name: Optional[str], family_name: Optional[str],
|
|||
return metadata
|
||||
|
||||
|
||||
async def _extract_group_member_ids(group: SCIMGroup) -> List[str]:
|
||||
"""Extract valid member IDs from SCIMGroup, verifying users exist."""
|
||||
async def _extract_group_member_ids(group: SCIMGroup) -> GroupMemberExtractionResult:
|
||||
"""
|
||||
Extract member IDs from SCIMGroup, creating users that don't exist.
|
||||
|
||||
Returns:
|
||||
GroupMemberExtractionResult with existing members, created users, and all member IDs
|
||||
"""
|
||||
prisma_client = await _get_prisma_client_or_raise_exception()
|
||||
member_ids = []
|
||||
existing_member_ids = []
|
||||
created_users = []
|
||||
all_member_ids = []
|
||||
|
||||
if group.members:
|
||||
for member in group.members:
|
||||
user_id = member.value
|
||||
|
||||
# Check if user exists
|
||||
user = await prisma_client.db.litellm_usertable.find_unique(
|
||||
where={"user_id": member.value}
|
||||
where={"user_id": user_id}
|
||||
)
|
||||
|
||||
if user:
|
||||
member_ids.append(member.value)
|
||||
existing_member_ids.append(user_id)
|
||||
all_member_ids.append(user_id)
|
||||
else:
|
||||
# Create the user if they don't exist using our helper
|
||||
created_user = await _create_user_if_not_exists(
|
||||
user_id=user_id,
|
||||
created_via="scim_group_membership"
|
||||
)
|
||||
|
||||
if created_user:
|
||||
created_users.append(created_user)
|
||||
all_member_ids.append(user_id)
|
||||
# If creation failed, user is skipped (logged in helper)
|
||||
|
||||
return member_ids
|
||||
return GroupMemberExtractionResult(
|
||||
existing_member_ids=existing_member_ids,
|
||||
created_users=created_users,
|
||||
all_member_ids=all_member_ids
|
||||
)
|
||||
|
||||
|
||||
async def _get_team_members_display(member_ids: List[str]) -> List[SCIMMember]:
|
||||
|
|
@ -239,6 +274,51 @@ async def _handle_team_membership_changes(user_id: str, existing_teams: List[str
|
|||
)
|
||||
|
||||
|
||||
async def _create_user_if_not_exists(user_id: str, created_via: str = "scim_group") -> Optional[NewUserResponse]:
|
||||
"""
|
||||
Helper function to create a user if they don't exist.
|
||||
|
||||
Args:
|
||||
user_id: The user ID to create
|
||||
created_via: Context for where the user was created from
|
||||
|
||||
Returns:
|
||||
LiteLLM_UserTable if user was created, None if creation failed
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
|
||||
|
||||
try:
|
||||
# Get default role for new internal users
|
||||
default_role: Optional[
|
||||
Literal[
|
||||
LitellmUserRoles.PROXY_ADMIN,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||||
LitellmUserRoles.INTERNAL_USER,
|
||||
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
|
||||
]
|
||||
] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
|
||||
if litellm.default_internal_user_params:
|
||||
default_role = litellm.default_internal_user_params.get("user_role")
|
||||
|
||||
new_user_request = NewUserRequest(
|
||||
user_id=user_id,
|
||||
user_email=user_id, # We don't have email from group membership
|
||||
user_alias=None,
|
||||
teams=[], # Teams will be added separately
|
||||
metadata={"created_via": created_via},
|
||||
auto_create_key=False,
|
||||
user_role=default_role,
|
||||
)
|
||||
|
||||
created_user = await new_user(data=new_user_request)
|
||||
verbose_proxy_logger.info(f"Created user {user_id} via {created_via}")
|
||||
return created_user
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Failed to create user {user_id}: {e}")
|
||||
return None
|
||||
|
||||
|
||||
async def _get_team_member_user_ids_from_team(team: LiteLLM_TeamTable) -> List[str]:
|
||||
"""
|
||||
Get the IDs of the members from a team.
|
||||
|
|
@ -256,6 +336,8 @@ async def _get_team_member_user_ids_from_team(team: LiteLLM_TeamTable) -> List[s
|
|||
member_user_ids.append(user_id)
|
||||
return member_user_ids
|
||||
|
||||
|
||||
|
||||
# Dependency to set the correct SCIM Content-Type
|
||||
async def set_scim_content_type(response: Response):
|
||||
"""Sets the Content-Type header to application/scim+json"""
|
||||
|
|
@ -914,9 +996,9 @@ async def create_group(
|
|||
detail={"error": f"Group already exists with ID: {team_id}"},
|
||||
)
|
||||
|
||||
# Extract valid member IDs
|
||||
member_ids = await _extract_group_member_ids(group)
|
||||
members_with_roles = [Member(user_id=member_id, role="user") for member_id in member_ids]
|
||||
# Extract and process group members (creating users that don't exist)
|
||||
member_result = await _extract_group_member_ids(group)
|
||||
members_with_roles = [Member(user_id=member_id, role="user") for member_id in member_result.all_member_ids]
|
||||
|
||||
# Create team in database
|
||||
created_team = await new_team(
|
||||
|
|
@ -959,9 +1041,10 @@ async def update_group(
|
|||
prisma_client = await _get_prisma_client_or_raise_exception()
|
||||
existing_team = await _check_team_exists(group_id)
|
||||
|
||||
# Extract valid member IDs
|
||||
member_ids = await _extract_group_member_ids(group)
|
||||
verbose_proxy_logger.debug(f"SCIM PUT GROUP member_ids: {member_ids}")
|
||||
# Extract and process group members (creating users that don't exist)
|
||||
member_result = await _extract_group_member_ids(group)
|
||||
verbose_proxy_logger.debug(f"SCIM PUT GROUP all_member_ids: {member_result.all_member_ids}")
|
||||
verbose_proxy_logger.debug(f"SCIM PUT GROUP created_users: {len(member_result.created_users)}")
|
||||
|
||||
# Prepare update data
|
||||
existing_metadata = existing_team.metadata if existing_team.metadata else {}
|
||||
|
|
@ -978,10 +1061,10 @@ async def update_group(
|
|||
data=update_data,
|
||||
)
|
||||
|
||||
# Handle user-team relationship changes using the same approach as patch_group
|
||||
# Handle user-team relationship changes
|
||||
current_members = set(await _get_team_member_user_ids_from_team(existing_team))
|
||||
verbose_proxy_logger.debug(f"SCIM PUT GROUP current_members: {current_members}")
|
||||
final_members = set(member_ids)
|
||||
final_members = set(member_result.all_member_ids)
|
||||
verbose_proxy_logger.debug(f"SCIM PUT GROUP final_members: {final_members}")
|
||||
|
||||
await _handle_group_membership_changes(
|
||||
|
|
@ -1075,7 +1158,7 @@ async def _process_group_patch_operations(
|
|||
elif path.startswith("members"):
|
||||
# Handle member operations
|
||||
member_values = _extract_group_values(value)
|
||||
# Validate that users exist
|
||||
# Create users that don't exist and get all valid member IDs
|
||||
valid_members = []
|
||||
for member_id in member_values:
|
||||
user = await prisma_client.db.litellm_usertable.find_unique(
|
||||
|
|
@ -1083,6 +1166,16 @@ async def _process_group_patch_operations(
|
|||
)
|
||||
if user:
|
||||
valid_members.append(member_id)
|
||||
else:
|
||||
# Create the user if they don't exist using our helper
|
||||
created_user = await _create_user_if_not_exists(
|
||||
user_id=member_id,
|
||||
created_via="scim_group_patch"
|
||||
)
|
||||
|
||||
if created_user:
|
||||
valid_members.append(member_id)
|
||||
# If creation failed, user is skipped (logged in helper)
|
||||
|
||||
if op_type == "replace":
|
||||
final_members = set(valid_members)
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from litellm.types.utils import StandardLoggingPayload
|
|||
|
||||
class TestS3V2UnitTests:
|
||||
"""Test that S3 v2 integration only uses safe_dumps and not json.dumps"""
|
||||
|
||||
def test_s3_v2_source_code_analysis(self):
|
||||
"""Test that S3 v2 source code only imports and uses safe_dumps"""
|
||||
import inspect
|
||||
|
|
@ -18,7 +19,139 @@ class TestS3V2UnitTests:
|
|||
|
||||
# Get the source code of the s3_v2 module
|
||||
source_code = inspect.getsource(s3_v2)
|
||||
|
||||
|
||||
# Verify that json.dumps is not used directly in the code
|
||||
assert "json.dumps(" not in source_code, \
|
||||
"S3 v2 should not use json.dumps directly"
|
||||
assert (
|
||||
"json.dumps(" not in source_code
|
||||
), "S3 v2 should not use json.dumps directly"
|
||||
|
||||
@patch('asyncio.create_task')
|
||||
@patch('litellm.integrations.s3_v2.CustomBatchLogger.periodic_flush')
|
||||
def test_s3_v2_endpoint_url(self, mock_periodic_flush, mock_create_task):
|
||||
"""testing s3 endpoint url"""
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from litellm.types.integrations.s3_v2 import s3BatchLoggingElement
|
||||
|
||||
# Mock periodic_flush and create_task to prevent async task creation during init
|
||||
mock_periodic_flush.return_value = None
|
||||
mock_create_task.return_value = None
|
||||
|
||||
# Mock response for all tests
|
||||
mock_response = MagicMock()
|
||||
mock_response.status_code = 200
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
# Create a test batch logging element
|
||||
test_element = s3BatchLoggingElement(
|
||||
s3_object_key="2025-09-14/test-key.json",
|
||||
payload={"test": "data"},
|
||||
s3_object_download_filename="test-file.json"
|
||||
)
|
||||
|
||||
# Test 1: Custom endpoint URL with bucket name
|
||||
s3_logger = S3Logger(
|
||||
s3_bucket_name="test-bucket",
|
||||
s3_endpoint_url="https://s3.amazonaws.com",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1"
|
||||
)
|
||||
|
||||
s3_logger.async_httpx_client = AsyncMock()
|
||||
s3_logger.async_httpx_client.put.return_value = mock_response
|
||||
|
||||
asyncio.run(s3_logger.async_upload_data_to_s3(test_element))
|
||||
|
||||
call_args = s3_logger.async_httpx_client.put.call_args
|
||||
assert call_args is not None
|
||||
url = call_args[0][0]
|
||||
expected_url = "https://s3.amazonaws.com/test-bucket/2025-09-14/test-key.json"
|
||||
assert url == expected_url, f"Expected URL {expected_url}, got {url}"
|
||||
|
||||
# Test 2: MinIO-compatible endpoint
|
||||
s3_logger_minio = S3Logger(
|
||||
s3_bucket_name="litellm-logs",
|
||||
s3_endpoint_url="https://minio.example.com:9000",
|
||||
s3_aws_access_key_id="minio-key",
|
||||
s3_aws_secret_access_key="minio-secret",
|
||||
s3_region_name="us-east-1"
|
||||
)
|
||||
|
||||
s3_logger_minio.async_httpx_client = AsyncMock()
|
||||
s3_logger_minio.async_httpx_client.put.return_value = mock_response
|
||||
|
||||
asyncio.run(s3_logger_minio.async_upload_data_to_s3(test_element))
|
||||
|
||||
call_args_minio = s3_logger_minio.async_httpx_client.put.call_args
|
||||
assert call_args_minio is not None
|
||||
url_minio = call_args_minio[0][0]
|
||||
expected_minio_url = "https://minio.example.com:9000/litellm-logs/2025-09-14/test-key.json"
|
||||
assert url_minio == expected_minio_url, f"Expected MinIO URL {expected_minio_url}, got {url_minio}"
|
||||
|
||||
# Test 3: Custom endpoint without bucket name (should fall back to default)
|
||||
s3_logger_no_bucket = S3Logger(
|
||||
s3_endpoint_url="https://s3.amazonaws.com",
|
||||
s3_aws_access_key_id="test-key",
|
||||
s3_aws_secret_access_key="test-secret",
|
||||
s3_region_name="us-east-1"
|
||||
)
|
||||
|
||||
s3_logger_no_bucket.async_httpx_client = AsyncMock()
|
||||
s3_logger_no_bucket.async_httpx_client.put.return_value = mock_response
|
||||
|
||||
asyncio.run(s3_logger_no_bucket.async_upload_data_to_s3(test_element))
|
||||
|
||||
call_args_no_bucket = s3_logger_no_bucket.async_httpx_client.put.call_args
|
||||
assert call_args_no_bucket is not None
|
||||
url_no_bucket = call_args_no_bucket[0][0]
|
||||
# Should use default S3 URL format when bucket is missing (bucket becomes None in URL)
|
||||
assert "s3.us-east-1.amazonaws.com" in url_no_bucket
|
||||
assert "https://" in url_no_bucket
|
||||
# Should not include the custom endpoint since bucket is missing
|
||||
assert "https://s3.amazonaws.com/" not in url_no_bucket
|
||||
|
||||
# Test 4: Sync upload method with custom endpoint
|
||||
s3_logger_sync = S3Logger(
|
||||
s3_bucket_name="sync-bucket",
|
||||
s3_endpoint_url="https://custom.s3.endpoint.com",
|
||||
s3_aws_access_key_id="sync-key",
|
||||
s3_aws_secret_access_key="sync-secret",
|
||||
s3_region_name="us-east-1"
|
||||
)
|
||||
|
||||
mock_sync_client = MagicMock()
|
||||
mock_sync_client.put.return_value = mock_response
|
||||
|
||||
with patch('litellm.integrations.s3_v2._get_httpx_client', return_value=mock_sync_client):
|
||||
s3_logger_sync.upload_data_to_s3(test_element)
|
||||
|
||||
call_args_sync = mock_sync_client.put.call_args
|
||||
assert call_args_sync is not None
|
||||
url_sync = call_args_sync[0][0]
|
||||
expected_sync_url = "https://custom.s3.endpoint.com/sync-bucket/2025-09-14/test-key.json"
|
||||
assert url_sync == expected_sync_url, f"Expected sync URL {expected_sync_url}, got {url_sync}"
|
||||
|
||||
# Test 5: Download method with custom endpoint
|
||||
s3_logger_download = S3Logger(
|
||||
s3_bucket_name="download-bucket",
|
||||
s3_endpoint_url="https://download.s3.endpoint.com",
|
||||
s3_aws_access_key_id="download-key",
|
||||
s3_aws_secret_access_key="download-secret",
|
||||
s3_region_name="us-east-1"
|
||||
)
|
||||
|
||||
mock_download_response = MagicMock()
|
||||
mock_download_response.status_code = 200
|
||||
mock_download_response.json = MagicMock(return_value={"downloaded": "data"})
|
||||
s3_logger_download.async_httpx_client = AsyncMock()
|
||||
s3_logger_download.async_httpx_client.get.return_value = mock_download_response
|
||||
|
||||
result = asyncio.run(s3_logger_download._download_object_from_s3("2025-09-14/download-test-key.json"))
|
||||
|
||||
call_args_download = s3_logger_download.async_httpx_client.get.call_args
|
||||
assert call_args_download is not None
|
||||
url_download = call_args_download[0][0]
|
||||
expected_download_url = "https://download.s3.endpoint.com/download-bucket/2025-09-14/download-test-key.json"
|
||||
assert url_download == expected_download_url, f"Expected download URL {expected_download_url}, got {url_download}"
|
||||
|
||||
assert result == {"downloaded": "data"}
|
||||
|
|
@ -1,4 +1,7 @@
|
|||
from litellm.llms.vertex_ai.gemini.transformation import check_if_part_exists_in_parts
|
||||
from litellm.llms.vertex_ai.gemini.transformation import (
|
||||
check_if_part_exists_in_parts,
|
||||
_transform_request_body,
|
||||
)
|
||||
|
||||
|
||||
def test_check_if_part_exists_in_parts():
|
||||
|
|
@ -73,3 +76,82 @@ def test_check_if_part_exists_in_parts_camel_case_snake_case():
|
|||
}
|
||||
|
||||
assert check_if_part_exists_in_parts(parts_mixed, part_mixed_casing)
|
||||
|
||||
|
||||
# Tests for issue #14556: Labels field provider-aware filtering
|
||||
def test_google_genai_excludes_labels():
|
||||
"""Test that Google GenAI/AI Studio endpoints exclude labels when custom_llm_provider='gemini'"""
|
||||
messages = [{"role": "user", "content": "test"}]
|
||||
optional_params = {"labels": {"project": "test", "team": "ai"}}
|
||||
litellm_params = {}
|
||||
|
||||
result = _transform_request_body(
|
||||
messages=messages,
|
||||
model="gemini-2.5-pro",
|
||||
optional_params=optional_params,
|
||||
custom_llm_provider="gemini",
|
||||
litellm_params=litellm_params,
|
||||
cached_content=None,
|
||||
)
|
||||
|
||||
# Google GenAI/AI Studio should NOT include labels
|
||||
assert "labels" not in result
|
||||
assert "contents" in result
|
||||
|
||||
|
||||
def test_vertex_ai_includes_labels():
|
||||
"""Test that Vertex AI endpoints include labels when custom_llm_provider='vertex_ai'"""
|
||||
messages = [{"role": "user", "content": "test"}]
|
||||
optional_params = {"labels": {"project": "test", "team": "ai"}}
|
||||
litellm_params = {}
|
||||
|
||||
result = _transform_request_body(
|
||||
messages=messages,
|
||||
model="gemini-2.5-pro",
|
||||
optional_params=optional_params,
|
||||
custom_llm_provider="vertex_ai",
|
||||
litellm_params=litellm_params,
|
||||
cached_content=None,
|
||||
)
|
||||
|
||||
# Vertex AI SHOULD include labels
|
||||
assert "labels" in result
|
||||
assert result["labels"] == {"project": "test", "team": "ai"}
|
||||
|
||||
|
||||
|
||||
def test_metadata_to_labels_vertex_only():
|
||||
"""Test that metadata->labels conversion only happens for Vertex AI"""
|
||||
messages = [{"role": "user", "content": "test"}]
|
||||
optional_params = {}
|
||||
litellm_params = {
|
||||
"metadata": {
|
||||
"requester_metadata": {
|
||||
"user": "john_doe",
|
||||
"project": "test-project"
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
# Google GenAI/AI Studio should not include labels from metadata
|
||||
result = _transform_request_body(
|
||||
messages=messages,
|
||||
model="gemini-2.5-pro",
|
||||
optional_params=optional_params.copy(),
|
||||
custom_llm_provider="gemini",
|
||||
litellm_params=litellm_params.copy(),
|
||||
cached_content=None,
|
||||
)
|
||||
assert "labels" not in result
|
||||
|
||||
# Vertex AI should include labels from metadata
|
||||
result = _transform_request_body(
|
||||
messages=messages,
|
||||
model="gemini-2.5-pro",
|
||||
optional_params=optional_params.copy(),
|
||||
custom_llm_provider="vertex_ai",
|
||||
litellm_params=litellm_params.copy(),
|
||||
cached_content=None,
|
||||
)
|
||||
assert "labels" in result
|
||||
assert result["labels"] == {"user": "john_doe", "project": "test-project"}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from litellm.proxy._types import LitellmUserRoles, NewUserRequest, ProxyExceptio
|
|||
from litellm.proxy.management_endpoints.scim.scim_v2 import (
|
||||
UserProvisionerHelpers,
|
||||
_handle_team_membership_changes,
|
||||
create_group,
|
||||
create_user,
|
||||
get_service_provider_config,
|
||||
patch_user,
|
||||
|
|
@ -910,4 +911,280 @@ async def test_update_group_e2e(mocker):
|
|||
assert len(result.members) == 3
|
||||
|
||||
# Verify SCIM transformation was called with updated team
|
||||
ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team)
|
||||
ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_group_with_nonexistent_users_creates_users(mocker):
|
||||
"""
|
||||
Test that creating a group with non-existent users creates those users.
|
||||
This tests the scenario: Group Push ['new user', existing users...]
|
||||
"""
|
||||
# Test data
|
||||
group_id = "test-group-123"
|
||||
scim_group = SCIMGroup(
|
||||
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
|
||||
id=group_id,
|
||||
displayName="Test Group",
|
||||
members=[
|
||||
SCIMMember(value="existing-user", display="Existing User"), # This user exists
|
||||
SCIMMember(value="new-user-1", display="New User 1"), # This user doesn't exist
|
||||
SCIMMember(value="new-user-2", display="New User 2"), # This user doesn't exist
|
||||
]
|
||||
)
|
||||
|
||||
#########################################################
|
||||
# We expect new-user-1 and new-user-2 to be created
|
||||
#########################################################
|
||||
|
||||
# Mock prisma client
|
||||
mock_prisma_client = mocker.MagicMock()
|
||||
mock_prisma_client.db = mocker.MagicMock()
|
||||
mock_prisma_client.db.litellm_teamtable = mocker.MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
|
||||
|
||||
# Mock team operations - team doesn't exist yet
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
|
||||
|
||||
# Mock user lookup - only existing-user exists
|
||||
def mock_user_lookup(where):
|
||||
user_id = where["user_id"]
|
||||
if user_id == "existing-user":
|
||||
mock_user = mocker.MagicMock()
|
||||
mock_user.user_id = user_id
|
||||
return mock_user
|
||||
return None # new-user-1 and new-user-2 don't exist
|
||||
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
|
||||
|
||||
# Mock dependencies
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
|
||||
AsyncMock(return_value=mock_prisma_client)
|
||||
)
|
||||
|
||||
# Mock new_user function to track user creation
|
||||
mock_new_user = mocker.patch(
|
||||
"litellm.proxy.management_endpoints.internal_user_endpoints.new_user",
|
||||
AsyncMock()
|
||||
)
|
||||
|
||||
# Mock created users return values
|
||||
def mock_new_user_side_effect(data):
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
return LiteLLM_UserTable(
|
||||
user_id=data.user_id,
|
||||
user_email=data.user_email,
|
||||
metadata=data.metadata,
|
||||
teams=data.teams,
|
||||
user_role=data.user_role
|
||||
)
|
||||
|
||||
mock_new_user.side_effect = mock_new_user_side_effect
|
||||
|
||||
# Mock new_team function
|
||||
mock_created_team = mocker.MagicMock()
|
||||
mock_created_team.team_id = group_id
|
||||
mock_created_team.team_alias = "Test Group"
|
||||
|
||||
mock_new_team = mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.new_team",
|
||||
AsyncMock(return_value=mock_created_team)
|
||||
)
|
||||
|
||||
# Mock SCIM transformation
|
||||
expected_scim_response = SCIMGroup(
|
||||
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
|
||||
id=group_id,
|
||||
displayName="Test Group",
|
||||
members=[
|
||||
SCIMMember(value="existing-user", display="existing-user"),
|
||||
SCIMMember(value="new-user-1", display="new-user-1"),
|
||||
SCIMMember(value="new-user-2", display="new-user-2")
|
||||
]
|
||||
)
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
|
||||
AsyncMock(return_value=expected_scim_response)
|
||||
)
|
||||
|
||||
# Execute the create_group function
|
||||
result = await create_group(group=scim_group)
|
||||
|
||||
#########################################################
|
||||
# Assert that new-user-1 and new-user-2 were created
|
||||
#########################################################
|
||||
|
||||
# Verify that new_user was called exactly twice (for new-user-1 and new-user-2)
|
||||
assert mock_new_user.call_count == 2
|
||||
|
||||
# Check the user creation calls
|
||||
created_user_ids = set()
|
||||
for call in mock_new_user.call_args_list:
|
||||
user_request = call.kwargs["data"]
|
||||
created_user_ids.add(user_request.user_id)
|
||||
assert user_request.metadata["created_via"] == "scim_group_membership"
|
||||
assert user_request.user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
|
||||
assert user_request.auto_create_key is False
|
||||
assert user_request.teams == [] # Teams added separately
|
||||
|
||||
assert created_user_ids == {"new-user-1", "new-user-2"}
|
||||
|
||||
# Verify team creation was called with all members (existing + created)
|
||||
mock_new_team.assert_called_once()
|
||||
team_request = mock_new_team.call_args.kwargs["data"]
|
||||
assert team_request.team_id == group_id
|
||||
assert team_request.team_alias == "Test Group"
|
||||
|
||||
# Verify all members are in the team (existing + newly created)
|
||||
member_user_ids = {member.user_id for member in team_request.members_with_roles}
|
||||
assert member_user_ids == {"existing-user", "new-user-1", "new-user-2"}
|
||||
|
||||
# Verify response
|
||||
assert result.id == group_id
|
||||
assert result.displayName == "Test Group"
|
||||
assert len(result.members) == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_group_with_nonexistent_users_creates_users(mocker):
|
||||
"""
|
||||
Test that updating a group with non-existent users creates those users.
|
||||
This tests the scenario where a group is updated with members that don't exist in user table.
|
||||
"""
|
||||
# Test data
|
||||
group_id = "existing-group-456"
|
||||
|
||||
# Mock existing team
|
||||
mock_existing_team = mocker.MagicMock()
|
||||
mock_existing_team.team_id = group_id
|
||||
mock_existing_team.team_alias = "Old Group Name"
|
||||
mock_existing_team.members = ["old-user"]
|
||||
mock_existing_team.members_with_roles = [{"user_id": "old-user", "role": "user"}]
|
||||
mock_existing_team.metadata = {"existing": "data"}
|
||||
|
||||
# SCIM group update request
|
||||
scim_group_update = SCIMGroup(
|
||||
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
|
||||
id=group_id,
|
||||
displayName="Updated Group Name",
|
||||
members=[
|
||||
SCIMMember(value="existing-user", display="Existing User"), # This user exists
|
||||
SCIMMember(value="new-user-3", display="New User 3"), # This user doesn't exist
|
||||
SCIMMember(value="new-user-4", display="New User 4"), # This user doesn't exist
|
||||
]
|
||||
)
|
||||
|
||||
# Mock prisma client
|
||||
mock_prisma_client = mocker.MagicMock()
|
||||
mock_prisma_client.db = mocker.MagicMock()
|
||||
mock_prisma_client.db.litellm_teamtable = mocker.MagicMock()
|
||||
mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
|
||||
|
||||
# Mock team operations
|
||||
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=mock_existing_team)
|
||||
|
||||
# Mock updated team response
|
||||
mock_updated_team = mocker.MagicMock()
|
||||
mock_updated_team.team_id = group_id
|
||||
mock_updated_team.team_alias = "Updated Group Name"
|
||||
mock_updated_team.members = ["existing-user", "new-user-3", "new-user-4"]
|
||||
mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=mock_updated_team)
|
||||
|
||||
# Mock user lookup - only existing-user exists
|
||||
def mock_user_lookup(where):
|
||||
user_id = where["user_id"]
|
||||
if user_id == "existing-user":
|
||||
mock_user = mocker.MagicMock()
|
||||
mock_user.user_id = user_id
|
||||
return mock_user
|
||||
return None # new-user-3 and new-user-4 don't exist
|
||||
|
||||
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(side_effect=mock_user_lookup)
|
||||
|
||||
# Mock dependencies
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
|
||||
AsyncMock(return_value=mock_prisma_client)
|
||||
)
|
||||
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._check_team_exists",
|
||||
AsyncMock(return_value=mock_existing_team)
|
||||
)
|
||||
|
||||
# Mock new_user function to track user creation
|
||||
mock_new_user = mocker.patch(
|
||||
"litellm.proxy.management_endpoints.internal_user_endpoints.new_user",
|
||||
AsyncMock()
|
||||
)
|
||||
|
||||
# Mock created users return values
|
||||
def mock_new_user_side_effect(data):
|
||||
from litellm.proxy._types import LiteLLM_UserTable
|
||||
return LiteLLM_UserTable(
|
||||
user_id=data.user_id,
|
||||
user_email=data.user_email,
|
||||
metadata=data.metadata,
|
||||
teams=data.teams,
|
||||
user_role=data.user_role
|
||||
)
|
||||
|
||||
mock_new_user.side_effect = mock_new_user_side_effect
|
||||
|
||||
# Mock group membership changes
|
||||
mock_handle_group_membership_changes = mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2._handle_group_membership_changes",
|
||||
AsyncMock()
|
||||
)
|
||||
|
||||
# Mock SCIM transformation
|
||||
expected_scim_response = SCIMGroup(
|
||||
schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
|
||||
id=group_id,
|
||||
displayName="Updated Group Name",
|
||||
members=[
|
||||
SCIMMember(value="existing-user", display="existing-user"),
|
||||
SCIMMember(value="new-user-3", display="new-user-3"),
|
||||
SCIMMember(value="new-user-4", display="new-user-4")
|
||||
]
|
||||
)
|
||||
mocker.patch(
|
||||
"litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group",
|
||||
AsyncMock(return_value=expected_scim_response)
|
||||
)
|
||||
|
||||
# Execute the update_group function
|
||||
result = await update_group(group_id=group_id, group=scim_group_update)
|
||||
|
||||
# Verify that new_user was called exactly twice (for new-user-3 and new-user-4)
|
||||
assert mock_new_user.call_count == 2
|
||||
|
||||
# Check the user creation calls
|
||||
created_user_ids = set()
|
||||
for call in mock_new_user.call_args_list:
|
||||
user_request = call.kwargs["data"]
|
||||
created_user_ids.add(user_request.user_id)
|
||||
assert user_request.metadata["created_via"] == "scim_group_membership"
|
||||
assert user_request.user_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
|
||||
assert user_request.auto_create_key is False
|
||||
assert user_request.teams == [] # Teams added separately
|
||||
|
||||
assert created_user_ids == {"new-user-3", "new-user-4"}
|
||||
|
||||
# Verify team update was called
|
||||
mock_prisma_client.db.litellm_teamtable.update.assert_called_once()
|
||||
update_call = mock_prisma_client.db.litellm_teamtable.update.call_args
|
||||
assert update_call[1]["where"]["team_id"] == group_id
|
||||
assert update_call[1]["data"]["team_alias"] == "Updated Group Name"
|
||||
|
||||
# Verify group membership changes were handled with all members (existing + created)
|
||||
mock_handle_group_membership_changes.assert_called_once()
|
||||
membership_call = mock_handle_group_membership_changes.call_args
|
||||
assert membership_call[1]["group_id"] == group_id
|
||||
assert membership_call[1]["final_members"] == {"existing-user", "new-user-3", "new-user-4"}
|
||||
|
||||
# Verify response
|
||||
assert result.id == group_id
|
||||
assert result.displayName == "Updated Group Name"
|
||||
assert len(result.members) == 3
|
||||
Loading…
Add table
Reference in a new issue