diff --git a/Makefile b/Makefile index 14e5340d885..f2753b09ff5 100644 --- a/Makefile +++ b/Makefile @@ -103,7 +103,7 @@ format-check: install-dev # Single fetch of the PR base so the delta-based gates below share one network round # trip instead of each re-fetching when chained from `lint`. lint-fetch-base: - git fetch origin litellm_internal_staging 2>/dev/null || true + git fetch origin litellm_internal_staging # Mirror test-linting.yml's lint job environment: the proxy-dev group plus a generated # Prisma client, so basedpyright resolves the same modules CI does (without the generated @@ -162,7 +162,7 @@ lint-ruff-FULL-dev: install-dev else echo "No changed .py files to check."; fi lint-basedpyright: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) - ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py + ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging # Type-discipline budget (mutable collections / casts / type guards / kwargs / # unexplained suppressions), the test-linting.yml step `make lint` used to omit. @@ -176,9 +176,6 @@ lint-basedpyright-budget-update: install-dev lint-fetch-base lint-format: format-check -lint-strict-budget: install-dev - $(UV_RUN) python scripts/ruff_strict_gate.py - lint-ruff-budget: install-dev $(UV_RUN) python scripts/ruff_strict_gate.py @@ -230,7 +227,7 @@ test: install-test-deps $(UV_RUN) pytest tests/ test-unit: install-test-deps - $(UV_RUN) pytest tests/test_litellm -x -vv -n 4 --reruns 2 --reruns-delay 1 + $(UV_RUN) pytest tests/test_litellm -x -vv -n 4 # Matrix test targets (matching CI workflow groups) test-unit-llms: install-test-deps diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index 82e4cf9c841..b67f7d42127 100644 --- a/backend/routes/allowlist.py +++ b/backend/routes/allowlist.py @@ -39,7 +39,6 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = ( "/model_access_group/", "/model_hub/", "/v1/access_group", - "/v1/unified_access_group", "/access_group/", "/router/", "/router_settings", @@ -49,7 +48,6 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = ( "/cache_settings", "/cost_tracking", "/cost/", - "/config_overrides/", "/credentials", "/credential", "/provider/budgets", diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py b/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py index 7353b995d2a..8824f4c02de 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/callback_controls.py @@ -14,74 +14,53 @@ from litellm.types.utils import StandardCallbackDynamicParams class EnterpriseCallbackControls: @staticmethod def is_callback_disabled_dynamically( - callback: litellm.CALLBACK_TYPES, - litellm_params: dict, - standard_callback_dynamic_params: StandardCallbackDynamicParams, - ) -> bool: - """ - Check if a callback is disabled via the x-litellm-disable-callbacks header or via `litellm_disabled_callbacks` in standard_callback_dynamic_params. - - Args: - callback: The callback to check (can be string, CustomLogger instance, or callable) - litellm_params: Parameters containing proxy server request info - - Returns: - bool: True if the callback should be disabled, False otherwise - """ - from litellm.litellm_core_utils.custom_logger_registry import ( - CustomLoggerRegistry, - ) - - try: - disabled_callbacks = EnterpriseCallbackControls.get_disabled_callbacks( - litellm_params, standard_callback_dynamic_params + callback: litellm.CALLBACK_TYPES, + litellm_params: dict, + standard_callback_dynamic_params: StandardCallbackDynamicParams + ) -> bool: + """ + Check if a callback is disabled via the x-litellm-disable-callbacks header or via `litellm_disabled_callbacks` in standard_callback_dynamic_params. + + Args: + callback: The callback to check (can be string, CustomLogger instance, or callable) + litellm_params: Parameters containing proxy server request info + + Returns: + bool: True if the callback should be disabled, False otherwise + """ + from litellm.litellm_core_utils.custom_logger_registry import ( + CustomLoggerRegistry, ) - verbose_logger.debug( - f"Dynamically disabled callbacks from {X_LITELLM_DISABLE_CALLBACKS}: {disabled_callbacks}" - ) - verbose_logger.debug( - f"Checking if {callback} is disabled via headers. Disable callbacks from headers: {disabled_callbacks}" - ) - if disabled_callbacks is not None: - ######################################################### - # premium user check - ######################################################### - if ( - not EnterpriseCallbackControls._should_allow_dynamic_callback_disabling() - ): - return False - ######################################################### - if isinstance(callback, str): - if callback.lower() in disabled_callbacks: - verbose_logger.debug( - f"Not logging to {callback} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}" - ) - return True - elif isinstance(callback, CustomLogger): - # get the string name of the callback - callback_str = ( - CustomLoggerRegistry.get_callback_str_from_class_type( - callback.__class__ - ) - ) - if ( - callback_str is not None - and callback_str.lower() in disabled_callbacks - ): - verbose_logger.debug( - f"Not logging to {callback_str} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}" - ) - return True - return False - except Exception as e: - verbose_logger.debug(f"Error checking disabled callbacks header: {str(e)}") - return False + try: + disabled_callbacks = EnterpriseCallbackControls.get_disabled_callbacks(litellm_params, standard_callback_dynamic_params) + verbose_logger.debug(f"Dynamically disabled callbacks from {X_LITELLM_DISABLE_CALLBACKS}: {disabled_callbacks}") + verbose_logger.debug(f"Checking if {callback} is disabled via headers. Disable callbacks from headers: {disabled_callbacks}") + if disabled_callbacks is not None: + ######################################################### + # premium user check + ######################################################### + if not EnterpriseCallbackControls._should_allow_dynamic_callback_disabling(): + return False + ######################################################### + if isinstance(callback, str): + if callback.lower() in disabled_callbacks: + verbose_logger.debug(f"Not logging to {callback} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}") + return True + elif isinstance(callback, CustomLogger): + # get the string name of the callback + callback_str = CustomLoggerRegistry.get_callback_str_from_class_type(callback.__class__) + if callback_str is not None and callback_str.lower() in disabled_callbacks: + verbose_logger.debug(f"Not logging to {callback_str} because it is disabled via {X_LITELLM_DISABLE_CALLBACKS}") + return True + return False + except Exception as e: + verbose_logger.debug( + f"Error checking disabled callbacks header: {str(e)}" + ) + return False @staticmethod - def get_disabled_callbacks( - litellm_params: dict, - standard_callback_dynamic_params: StandardCallbackDynamicParams, - ) -> Optional[List[str]]: + def get_disabled_callbacks(litellm_params: dict, standard_callback_dynamic_params: StandardCallbackDynamicParams) -> Optional[List[str]]: """ Get the disabled callbacks from the standard callback dynamic params. """ @@ -92,24 +71,18 @@ class EnterpriseCallbackControls: request_headers = get_proxy_server_request_headers(litellm_params) disabled_callbacks = request_headers.get(X_LITELLM_DISABLE_CALLBACKS, None) if disabled_callbacks is not None: - disabled_callbacks = set( - [cb.strip().lower() for cb in disabled_callbacks.split(",")] - ) + disabled_callbacks = set([cb.strip().lower() for cb in disabled_callbacks.split(",")]) return list(disabled_callbacks) + ######################################################### # check if disabled via request body ######################################################### - if ( - standard_callback_dynamic_params.get("litellm_disabled_callbacks", None) - is not None - ): - return standard_callback_dynamic_params.get( - "litellm_disabled_callbacks", None - ) - + if standard_callback_dynamic_params.get("litellm_disabled_callbacks", None) is not None: + return standard_callback_dynamic_params.get("litellm_disabled_callbacks", None) + return None - + @staticmethod def _should_allow_dynamic_callback_disabling(): import litellm @@ -117,14 +90,10 @@ class EnterpriseCallbackControls: # Check if admin has disabled this feature if litellm.allow_dynamic_callback_disabling is not True: - verbose_logger.debug( - "Dynamic callback disabling is disabled by admin via litellm.allow_dynamic_callback_disabling" - ) + verbose_logger.debug("Dynamic callback disabling is disabled by admin via litellm.allow_dynamic_callback_disabling") return False - + if premium_user: return True - verbose_logger.warning( - f"Disabling callbacks using request headers is an enterprise feature. {CommonProxyErrors.not_premium_user.value}" - ) - return False + verbose_logger.warning(f"Disabling callbacks using request headers is an enterprise feature. {CommonProxyErrors.not_premium_user.value}") + return False \ No newline at end of file diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py index a1e8def2bb2..8fc2d66d531 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/sendgrid_email.py @@ -15,6 +15,7 @@ from litellm.llms.custom_httpx.http_handler import ( from .base_email import BaseEmailLogger + SENDGRID_API_ENDPOINT = "https://api.sendgrid.com/v3/mail/send" @@ -78,4 +79,4 @@ class SendGridEmailLogger(BaseEmailLogger): verbose_logger.debug( f"SendGrid response status={response.status_code}, body={response.text}" ) - return + return \ No newline at end of file diff --git a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/smtp_email.py b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/smtp_email.py index 8e4dbde437b..8efdaf231b7 100644 --- a/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/smtp_email.py +++ b/enterprise/litellm_enterprise/enterprise_callbacks/send_emails/smtp_email.py @@ -1,7 +1,6 @@ """ This is the litellm SMTP email integration """ - import asyncio from typing import List diff --git a/enterprise/litellm_enterprise/litellm_core_utils/litellm_logging.py b/enterprise/litellm_enterprise/litellm_core_utils/litellm_logging.py index 24941e90ab8..44ba0063ffe 100644 --- a/enterprise/litellm_enterprise/litellm_core_utils/litellm_logging.py +++ b/enterprise/litellm_enterprise/litellm_core_utils/litellm_logging.py @@ -1,7 +1,6 @@ """ Enterprise specific logging utils """ - from litellm.litellm_core_utils.litellm_logging import StandardLoggingMetadata diff --git a/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py b/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py index 4f2eaa3c468..18ac29b9781 100644 --- a/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/audit_logging_endpoints.py @@ -153,11 +153,11 @@ async def get_audit_logs( # Return paginated response return PaginatedAuditLogResponse( - audit_logs=( - [AuditLogResponse(**audit_log.model_dump()) for audit_log in audit_logs] - if audit_logs - else [] - ), + audit_logs=[ + AuditLogResponse(**audit_log.model_dump()) for audit_log in audit_logs + ] + if audit_logs + else [], total=total_count, page=page, page_size=page_size, diff --git a/enterprise/litellm_enterprise/proxy/auth/__init__.py b/enterprise/litellm_enterprise/proxy/auth/__init__.py index dc70b57ab55..f67826ca7fa 100644 --- a/enterprise/litellm_enterprise/proxy/auth/__init__.py +++ b/enterprise/litellm_enterprise/proxy/auth/__init__.py @@ -7,4 +7,4 @@ including custom SSO handlers and advanced authentication features. from .custom_sso_handler import EnterpriseCustomSSOHandler -__all__ = ["EnterpriseCustomSSOHandler"] +__all__ = ["EnterpriseCustomSSOHandler"] \ No newline at end of file diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index 8b2d15c1578..dc0168683c8 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -33,7 +33,9 @@ class CheckResponsesCost: self.prisma_client: PrismaClient = prisma_client self.llm_router: Router = llm_router - async def _expire_stale_rows(self, cutoff: datetime, batch_size: int) -> int: + async def _expire_stale_rows( + self, cutoff: datetime, batch_size: int + ) -> int: """Execute the bounded UPDATE that marks stale rows as 'stale_expired'. Isolated so it can be swapped / mocked in tests without touching the @@ -72,9 +74,7 @@ class CheckResponsesCost: rows per invocation to avoid overwhelming the DB when there is a large backlog. """ - cutoff = datetime.now(timezone.utc) - timedelta( - days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS - ) + cutoff = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS) result = await self._expire_stale_rows(cutoff, STALE_OBJECT_CLEANUP_BATCH_SIZE) if result > 0: verbose_proxy_logger.warning( @@ -105,7 +105,7 @@ class CheckResponsesCost: take=MAX_OBJECTS_PER_POLL_CYCLE, order={"created_at": "asc"}, ) - + verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check") completed_jobs = [] @@ -120,33 +120,29 @@ class CheckResponsesCost: # Get the stored response object to extract model information stored_response = job.file_object model_name = stored_response.get("model", None) - + # Decrypt the response ID - responses_id_security, _, _ = ( - ResponsesIDSecurity()._decrypt_response_id(unified_object_id) - ) - + responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(unified_object_id) + # Prepare metadata with model information for cost tracking litellm_metadata = { "user_api_key_user_id": job.created_by or "default-user-id", } - + # Add model information if available if model_name: litellm_metadata["model"] = model_name - litellm_metadata["model_group"] = ( - model_name # Use same value for model_group - ) - + litellm_metadata["model_group"] = model_name # Use same value for model_group + response = await litellm.aget_responses( response_id=responses_id_security, litellm_metadata=litellm_metadata, ) - + verbose_proxy_logger.debug( f"Response {unified_object_id} status: {response.status}, model: {model_name}" ) - + except Exception as e: verbose_proxy_logger.info( f"Skipping job {unified_object_id} due to error: {e}" @@ -159,7 +155,7 @@ class CheckResponsesCost: f"Response {unified_object_id} is complete. Cost automatically tracked by aget_responses." ) completed_jobs.append(job) - + elif response.status in ["failed", "cancelled"]: verbose_proxy_logger.info( f"Response {unified_object_id} has status {response.status}, marking as complete" @@ -175,3 +171,4 @@ class CheckResponsesCost: verbose_proxy_logger.info( f"Marked {len(completed_jobs)} response jobs as completed" ) + diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py index 70634537c55..254d816039c 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_vector_stores.py @@ -41,7 +41,7 @@ class _PROXY_LiteLLMManagedVectorStores( ): """ Managed vector stores with target_model_names support. - + This class provides functionality to: - Create vector stores across multiple models - Retrieve vector stores by unified ID @@ -77,14 +77,14 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> str: """ Generate the format string for the unified vector store ID. - + Format: litellm_proxy:vector_store;unified_id,;target_model_names,;resource_id,;model_id, """ # VectorStoreCreateResponse is a TypedDict, so resource_object is a dictionary # Extract provider resource ID from the response provider_resource_id = resource_object.get("id", "") - + # Model ID is stored in hidden params if the response object supports it # For TypedDict responses, we need to check if _hidden_params was added hidden_params: Dict[str, Any] = {} @@ -109,18 +109,20 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> VectorStoreCreateResponse: """ Create a vector store for a specific model. - + Args: llm_router: LiteLLM router instance model: Model name to create vector store for request_data: Request data for vector store creation litellm_parent_otel_span: OpenTelemetry span for tracing - + Returns: VectorStoreCreateResponse from the provider """ # Use the router to create the vector store - response = await llm_router.avector_store_create(model=model, **request_data) + response = await llm_router.avector_store_create( + model=model, **request_data + ) return response # ============================================================================ @@ -137,14 +139,14 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> VectorStoreCreateResponse: """ Create a vector store across multiple models. - + Args: create_request: Vector store creation request parameters llm_router: LiteLLM router instance target_model_names_list: List of target model names litellm_parent_otel_span: OpenTelemetry span for tracing user_api_key_dict: User API key authentication details - + Returns: VectorStoreCreateResponse with unified ID """ @@ -194,7 +196,7 @@ class _PROXY_LiteLLMManagedVectorStores( # VectorStoreCreateResponse is a TypedDict, so we need to create a new dict with the unified ID response = responses[0].copy() response["id"] = unified_id - + verbose_logger.info( f"Successfully created managed vector store with unified ID: {unified_id}" ) @@ -210,13 +212,13 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> Dict[str, Any]: """ List vector stores created by a user. - + Args: user_api_key_dict: User API key authentication details limit: Maximum number of vector stores to return after: Cursor for pagination order: Sort order ('asc' or 'desc') - + Returns: Dictionary with list of vector stores and pagination info """ @@ -236,23 +238,23 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> bool: """ Check if user has access to a vector store. - + Args: vector_store_id: The unified vector store ID user_api_key_dict: User API key authentication details - + Returns: True if user has access, False otherwise """ is_unified_id = is_base64_encoded_unified_id(vector_store_id) - + if is_unified_id: # Check access for managed vector store return await self.can_user_access_unified_resource_id( vector_store_id, user_api_key_dict, ) - + # Not a managed vector store, allow access return True @@ -261,22 +263,24 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> bool: """ Check if user has access to a managed vector store in request data. - + Args: data: Request data containing vector_store_id user_api_key_dict: User API key authentication details - + Returns: True if this is a managed vector store and user has access - + Raises: HTTPException: If user doesn't have access """ vector_store_id = cast(Optional[str], data.get("vector_store_id")) is_unified_id = ( - is_base64_encoded_unified_id(vector_store_id) if vector_store_id else False + is_base64_encoded_unified_id(vector_store_id) + if vector_store_id + else False ) - + if is_unified_id and vector_store_id: if await self.can_user_access_unified_resource_id( vector_store_id, user_api_key_dict @@ -287,7 +291,7 @@ class _PROXY_LiteLLMManagedVectorStores( status_code=403, detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}", ) - + return False # ============================================================================ @@ -303,18 +307,18 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> Union[Exception, str, Dict, None]: """ Pre-call hook to handle vector store operations. - + This hook intercepts vector store requests and: - Validates access for managed vector stores - Transforms unified IDs to provider-specific IDs - Adds model routing information - + Args: user_api_key_dict: User API key authentication details cache: Cache instance data: Request data call_type: Type of call being made - + Returns: Modified request data or None """ @@ -326,40 +330,40 @@ class _PROXY_LiteLLMManagedVectorStores( # Handle vector store search operations if call_type == "avector_store_search": vector_store_id = data.get("vector_store_id") - + if vector_store_id: # Check if it's a managed vector store ID decoded_id = is_base64_encoded_unified_id(vector_store_id) - + if decoded_id: verbose_logger.debug( f"Processing managed vector store search: {vector_store_id}" ) - + # Check access has_access = await self.can_user_access_unified_resource_id( vector_store_id, user_api_key_dict ) - + if not has_access: raise HTTPException( status_code=403, detail=f"User {user_api_key_dict.user_id} does not have access to vector store {vector_store_id}", ) - + # Parse the unified ID to extract components parsed_id = parse_unified_id(vector_store_id) - + if parsed_id: # Extract the model ID and provider resource ID model_id = parsed_id.get("model_id") provider_resource_id = parsed_id.get("provider_resource_id") target_model_names = parsed_id.get("target_model_names", []) - + verbose_logger.debug( f"Decoded vector store - model_id: {model_id}, provider_resource_id: {provider_resource_id}, target_model_names: {target_model_names}" ) - + # Determine which model to use for routing # Priority: model_id (deployment ID) > first target_model_name routing_model = None @@ -367,28 +371,28 @@ class _PROXY_LiteLLMManagedVectorStores( routing_model = model_id elif target_model_names and len(target_model_names) > 0: routing_model = target_model_names[0] - + # Set the model for routing if routing_model: data["model"] = routing_model verbose_logger.info( f"Routing vector store search to model: {routing_model}" ) - + # Replace the unified ID with the provider-specific ID if provider_resource_id: data["vector_store_id"] = provider_resource_id verbose_logger.debug( f"Replaced unified ID with provider resource ID: {provider_resource_id}" ) - + # Handle vector store retrieve/delete operations elif call_type in ("avector_store_retrieve", "avector_store_delete"): await self.check_managed_vector_store_access(data, user_api_key_dict) - + # If it's a managed vector store, we'll handle it in the endpoint # No need to transform here as the endpoint will route to the hook - + return data # ============================================================================ @@ -403,15 +407,15 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> Any: """ Post-call hook to transform responses. - + This hook can be used to transform responses if needed. For now, it just passes through the response. - + Args: data: Request data user_api_key_dict: User API key authentication details response: Response from the provider - + Returns: Potentially modified response """ @@ -432,21 +436,21 @@ class _PROXY_LiteLLMManagedVectorStores( ) -> List[Dict]: """ Filter deployments based on vector store availability. - + This is used by the router to select only deployments that have the vector store available. - + Note: This method signature is a compromise between CustomLogger and BaseManagedResource parent classes which have incompatible signatures. The type: ignore[override] is necessary due to this multiple inheritance conflict. - + Args: model: Model name healthy_deployments: List of healthy deployments messages: Messages (unused for vector stores, required by CustomLogger interface) request_kwargs: Request kwargs containing vector_store_id and mappings parent_otel_span: OpenTelemetry span for tracing - + Returns: Filtered list of deployments """ diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py index 979f42361ff..1d3268da9a0 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/internal_user_endpoints.py @@ -2,6 +2,7 @@ Enterprise internal user management endpoints """ + from fastapi import APIRouter, Depends, HTTPException from litellm.proxy._types import UserAPIKeyAuth diff --git a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py index cf8c38719d7..5e799599862 100644 --- a/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py +++ b/enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py @@ -147,12 +147,12 @@ async def list_vector_stores( vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db( prisma_client=prisma_client ) - + # Also clean up in-memory registry to remove any deleted vector stores if litellm.vector_store_registry is not None: db_vector_store_ids = { - vs.get("vector_store_id") - for vs in vector_stores_from_db + vs.get("vector_store_id") + for vs in vector_stores_from_db if vs.get("vector_store_id") } # Remove any in-memory vector stores that no longer exist in database diff --git a/enterprise/litellm_enterprise/types/enterprise_callbacks/send_emails.py b/enterprise/litellm_enterprise/types/enterprise_callbacks/send_emails.py index d9d5a989abb..380b0a6facb 100644 --- a/enterprise/litellm_enterprise/types/enterprise_callbacks/send_emails.py +++ b/enterprise/litellm_enterprise/types/enterprise_callbacks/send_emails.py @@ -39,23 +39,15 @@ class EmailEvent(str, enum.Enum): soft_budget_crossed = "Soft Budget Crossed" max_budget_alert = "Max Budget Alert" - class EmailEventSettings(BaseModel): event: EmailEvent enabled: bool - - class EmailEventSettingsUpdateRequest(BaseModel): settings: List[EmailEventSettings] - - class EmailEventSettingsResponse(BaseModel): settings: List[EmailEventSettings] - - class DefaultEmailSettings(BaseModel): """Default settings for email events""" - settings: Dict[EmailEvent, bool] = Field( default_factory=lambda: { EmailEvent.virtual_key_created: True, # On by default @@ -65,12 +57,10 @@ class DefaultEmailSettings(BaseModel): EmailEvent.max_budget_alert: True, # On by default } ) - def to_dict(self) -> Dict[str, bool]: """Convert to dictionary with string keys for storage""" return {event.value: enabled for event, enabled in self.settings.items()} - @classmethod def get_defaults(cls) -> Dict[str, bool]: """Get the default settings as a dictionary with string keys""" - return cls().to_dict() + return cls().to_dict() \ No newline at end of file diff --git a/litellm/build-and-push.sh b/litellm/build-and-push.sh deleted file mode 100644 index 4c2bcc00155..00000000000 --- a/litellm/build-and-push.sh +++ /dev/null @@ -1,37 +0,0 @@ -#!/bin/bash -set -e - -# Configuration -ECR_REPO="420049223852.dkr.ecr.eu-central-1.amazonaws.com/litellm-deepkeep" -TAG="${1:-dev-$(date +%Y%m%d-%H%M%S)}" -FULL_IMAGE="${ECR_REPO}:${TAG}" - -echo "================================================" -echo "Building LiteLLM Docker image with UI" -echo "Image: ${FULL_IMAGE}" -echo "================================================" - -# Login to ECR -echo "Logging in to ECR..." -aws ecr get-login-password --region eu-central-1 | docker login --username AWS --password-stdin 420049223852.dkr.ecr.eu-central-1.amazonaws.com - -# Build the image -echo "Building Docker image..." -docker build -t "${FULL_IMAGE}" -f Dockerfile . - -# Push the image -echo "Pushing image to ECR..." -docker push "${FULL_IMAGE}" - -echo "================================================" -echo "✅ Image pushed successfully!" -echo "" -echo "To use in Tilt, update your tilt_config.yaniv-v2.yaml:" -echo "" -echo "litellm:" -echo " image:" -echo " repository: ${ECR_REPO}" -echo " tag: ${TAG}" -echo "" -echo "Or run: tilt trigger litellm" -echo "================================================" diff --git a/litellm/deepkeep_tilt_config.yaml b/litellm/deepkeep_tilt_config.yaml deleted file mode 100644 index 121190b52b7..00000000000 --- a/litellm/deepkeep_tilt_config.yaml +++ /dev/null @@ -1,22 +0,0 @@ -model_list: - - model_name: fake-openai-endpoint - litellm_params: - model: openai/fake-model - api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ - -general_settings: - master_key: sk-1234 - -litellm_settings: - drop_params: True - telemetry: False - -guardrails: - - guardrail_name: deepkeep-firewall - litellm_params: - guardrail: deepkeep - mode: [pre_call, post_call] - api_key: os.environ/DEEPKEEP_API_KEY - api_base: "http://localhost:8081/api" - deepkeep_firewall_id: "063e4e5ba5c2be55" diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index 6f79a39c883..4d51295f8dc 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -130,7 +130,9 @@ def classify_value(value: object, key: str = "scan") -> ValueClass: return "plaintext" if value.startswith(_V2_GCM_PREFIX): return "migrated" - decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) + decrypted = decrypt_value_helper( + value=value, key=key, exception_type="debug", return_original_value=False + ) if decrypted is None: # Did not decrypt under nacl and has no v2 marker: legacy plaintext. return "plaintext" @@ -149,7 +151,9 @@ def reencrypt_value(value: object, key: str = "migrate") -> object: return value if value.startswith(_V2_GCM_PREFIX): return value # idempotent: already migrated - decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False) + decrypted = decrypt_value_helper( + value=value, key=key, exception_type="debug", return_original_value=False + ) if decrypted is None: # Either legacy plaintext (no ciphertext to migrate) or corrupt. Either # way, do not overwrite — preserve the value as stored. @@ -157,7 +161,9 @@ def reencrypt_value(value: object, key: str = "migrate") -> object: return encrypt_value_helper(decrypted) -def reencrypt_selective_dict(data: dict[str, object], sensitive_keys: list[str]) -> dict[str, object]: +def reencrypt_selective_dict( + data: dict[str, object], sensitive_keys: list[str] +) -> dict[str, object]: """Return a copy of ``data`` with only ``sensitive_keys`` re-encrypted. Non-sensitive fields (e.g. ``base_url``, ``connection_id``) are left as-is. @@ -206,7 +212,9 @@ async def _migrate_config_settings_row( dict with selected sensitive fields (vantage_settings / cloudzero_settings). """ report = LocationReport(location=param_name) - record = await prisma_client.db.litellm_config.find_unique(where={"param_name": param_name}) + record = await prisma_client.db.litellm_config.find_unique( + where={"param_name": param_name} + ) if record is None or record.param_value is None: return report @@ -258,7 +266,9 @@ async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationR every present string field. """ report = LocationReport(location="sso_config") - record = await prisma_client.db.litellm_ssoconfig.find_unique(where={"id": "sso_config"}) + record = await prisma_client.db.litellm_ssoconfig.find_unique( + where={"id": "sso_config"} + ) if record is None or record.sso_settings is None: return report @@ -334,7 +344,9 @@ async def _migrate_callback_vars_table( rows = await table.find_many() for row in rows or []: metadata = getattr(row, "metadata", None) - if not isinstance(metadata, dict) or ("logging" not in metadata and "callback_settings" not in metadata): + if not isinstance(metadata, dict) or ( + "logging" not in metadata and "callback_settings" not in metadata + ): continue # Classify every callback-var value directly (strip the litellm_enc:: @@ -522,7 +534,9 @@ async def _scan_config_env_vars(prisma_client: object) -> LocationReport: """Scan the ``environment_variables`` config row (``param_value`` dict).""" report = LocationReport(location="config_environment_variables") try: - record = await prisma_client.db.litellm_config.find_unique(where={"param_name": "environment_variables"}) + record = await prisma_client.db.litellm_config.find_unique( + where={"param_name": "environment_variables"} + ) except Exception as e: # pragma: no cover - defensive verbose_proxy_logger.debug("scan: config env vars unavailable: %s", str(e)) return report @@ -543,7 +557,11 @@ async def _scan_covered_tables(prisma_client: object) -> list[LocationReport]: """Read-only classification of every rotation-covered table. No writes.""" reports: list[LocationReport] = [] for location, db_attr, json_cols, scalar_cols in _COVERED_TABLE_SPECS: - reports.append(await _scan_one_table(prisma_client, location, db_attr, json_cols, scalar_cols)) + reports.append( + await _scan_one_table( + prisma_client, location, db_attr, json_cols, scalar_cols + ) + ) reports.append(await _scan_config_env_vars(prisma_client)) return reports @@ -557,7 +575,9 @@ _VANTAGE_SENSITIVE = ["api_key", "integration_token"] _CLOUDZERO_SENSITIVE = ["api_key"] -async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: object) -> list[LocationReport]: +async def _migrate_covered_tables( + prisma_client: object, user_api_key_dict: object +) -> list[LocationReport]: """Re-encrypt the tables already covered by ``_rotate_master_key`` (model table, credentials, MCP credential/env tables, config environment_variables) by running that orchestrator in *same-key* mode. With the AES gate on, the @@ -577,7 +597,8 @@ async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: obje current_key = _get_salt_key() if current_key is None: raise RuntimeError( - "Cannot migrate covered tables: no salt key / master key is set. Set LITELLM_SALT_KEY before migrating." + "Cannot migrate covered tables: no salt key / master key is set. " + "Set LITELLM_SALT_KEY before migrating." ) await _rotate_master_key( prisma_client=cast("PrismaClient", prisma_client), @@ -627,9 +648,19 @@ async def migrate_encryption( # Net-new walkers (items 3, 4, 11, 12, 13). report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run)) - report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run)) - report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run)) - report.add(await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run)) + report.add( + await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run) + ) + report.add( + await _migrate_config_settings_row( + prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run + ) + ) + report.add( + await _migrate_config_settings_row( + prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run + ) + ) report.add(await _migrate_sso_config(prisma_client, dry_run)) return report @@ -652,10 +683,20 @@ async def check_encryption(prisma_client: object) -> MigrationReport: # Net-new walker locations, in dry-run (read-only) mode. report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run=True)) - report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run=True)) - report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True)) report.add( - await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True) + await _migrate_callback_vars_table( + prisma_client, "verification_token", dry_run=True + ) + ) + report.add( + await _migrate_config_settings_row( + prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True + ) + ) + report.add( + await _migrate_config_settings_row( + prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True + ) ) report.add(await _migrate_sso_config(prisma_client, dry_run=True)) return report diff --git a/scripts/ruff_strict_gate.py b/scripts/ruff_strict_gate.py index c9180d825ae..25f6c4d29ba 100644 --- a/scripts/ruff_strict_gate.py +++ b/scripts/ruff_strict_gate.py @@ -62,7 +62,12 @@ def head_violations() -> list: out = [] for item in _ruff_json(REPO_ROOT, STRICT_CONFIG): name = Path(item["filename"]) - rel = (name if name.is_absolute() else REPO_ROOT / name).resolve().relative_to(REPO_ROOT).as_posix() + rel = ( + (name if name.is_absolute() else REPO_ROOT / name) + .resolve() + .relative_to(REPO_ROOT) + .as_posix() + ) out.append(Violation(rel, item["location"]["row"], item["code"])) return out @@ -137,11 +142,15 @@ def cmd_check(base: str) -> None: return new = introduced( head, - parse_changed_lines(_run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET])), + parse_changed_lines( + _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET]) + ), ) print(f"FAIL: strict-rule totals exceed their limit (base {base}):") for breach in breaches: - print(f" {breach.rule}: total {breach.total} over limit {breach.cap} (this change added {breach.added})") + print( + f" {breach.rule}: total {breach.total} over limit {breach.cap} (this change added {breach.added})" + ) for violation in sorted(v for v in new if v.code == breach.rule): print(f" {violation.file}:{violation.line}") print( @@ -158,7 +167,9 @@ def ratcheted_budget(budget: dict, current: dict, base: dict) -> dict: stays put), so the limit only ever falls. """ return { - rule: {"limit": max(0, spec["limit"] - max(0, base.get(rule, 0) - current.get(rule, 0)))} + rule: { + "limit": max(0, spec["limit"] - max(0, base.get(rule, 0) - current.get(rule, 0))) + } for rule, spec in sorted(budget.items()) } @@ -172,34 +183,20 @@ def cmd_update(base_ref: str = DEFAULT_BASE) -> None: """ budget = json.loads(BUDGET_PATH.read_text()) base_point = _run(["git", "merge-base", base_ref, "HEAD"]).strip() or base_ref - updated = ratcheted_budget(budget, count_by_rule(head_violations()), base_counts(base_point)) + updated = ratcheted_budget( + budget, count_by_rule(head_violations()), base_counts(base_point) + ) BUDGET_PATH.write_text(json.dumps(updated, indent=2, sort_keys=True) + "\n") cleared = sum(budget[rule]["limit"] - updated[rule]["limit"] for rule in updated) print(f"Ratcheted strict-rule limits down by {cleared} violations this branch fixed") -def _resolve_base(ref: str) -> str: - """Return *ref* if it resolves; fall back to the upstream/ equivalent.""" - if subprocess.run(["git", "rev-parse", "--verify", ref], cwd=REPO_ROOT, capture_output=True).returncode == 0: - return ref - fallback = ref.replace("origin/", "upstream/", 1) - if ( - fallback != ref - and subprocess.run(["git", "rev-parse", "--verify", fallback], cwd=REPO_ROOT, capture_output=True).returncode - == 0 - ): - print(f"Note: '{ref}' not found, using '{fallback}' as base.", file=sys.stderr) - return fallback - return ref - - def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--base", default=DEFAULT_BASE) parser.add_argument("--update", action="store_true") args = parser.parse_args() - resolved = _resolve_base(args.base) - cmd_update(resolved) if args.update else cmd_check(resolved) + cmd_update(args.base) if args.update else cmd_check(args.base) if __name__ == "__main__": diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py index b18858a76b4..2c5306cec7d 100644 --- a/scripts/type_check_gate.py +++ b/scripts/type_check_gate.py @@ -131,7 +131,9 @@ def base_counts(ref: str) -> dict[str, int]: exe = shutil.which("basedpyright") or "basedpyright" with _temp_worktree(ref) as worktree: shutil.copy(PYRIGHT_CONFIG, worktree / "pyrightconfig.json") - proc = subprocess.run([exe, "--outputjson"], cwd=worktree, capture_output=True, text=True) + proc = subprocess.run( + [exe, "--outputjson"], cwd=worktree, capture_output=True, text=True + ) return count_basedpyright(proc.stdout, root=worktree) @@ -253,7 +255,9 @@ def evaluate( return sorted(breaches) -def is_vacuous_run(counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]]) -> bool: +def is_vacuous_run( + counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]] +) -> bool: """True when nothing was parsed but the budget expects errors -- the signature of a type checker that crashed or produced no output. The CI pipe swallows the tool's exit code (`tool || true`), so without this guard an @@ -275,7 +279,9 @@ def ratcheted_budget( not on update. """ return { - code: {"limit": max(0, spec["limit"] - max(0, base.get(code, 0) - current.get(code, 0)))} + code: { + "limit": max(0, spec["limit"] - max(0, base.get(code, 0) - current.get(code, 0))) + } for code, spec in sorted(budget.items()) } @@ -293,7 +299,10 @@ def cmd_update(current: Mapping[str, int], base_ref: str = DEFAULT_BASE) -> None updated = ratcheted_budget(budget, current, base_counts_cached(base_point)) BUDGET_PATH.write_text(json.dumps(updated, indent=2, sort_keys=True) + "\n") cleared = sum(budget[code]["limit"] - updated[code]["limit"] for code in updated) - print(f"Ratcheted basedpyright limits down by {cleared} errors this branch fixed across {len(updated)} rules") + print( + f"Ratcheted basedpyright limits down by {cleared} errors this branch fixed " + f"across {len(updated)} rules" + ) def cmd_check(base_ref: str) -> None: @@ -329,7 +338,9 @@ def cmd_check(base_ref: str) -> None: return print("FAIL: basedpyright errors exceed the per-rule limit:") for breach in breaches: - print(f" {breach.code}: total {breach.total} over limit {breach.cap} (this change added {breach.added})") + print( + f" {breach.code}: total {breach.total} over limit {breach.cap} (this change added {breach.added})" + ) print( "Reduce the new errors or remove an equal number elsewhere; the ceiling is " "the limit in basedpyright-code-budget.json." @@ -339,31 +350,15 @@ def cmd_check(base_ref: str) -> None: raise SystemExit(1) -def _resolve_base(ref: str) -> str: - """Return *ref* if it resolves; fall back to the upstream/ equivalent.""" - if subprocess.run(["git", "rev-parse", "--verify", ref], cwd=REPO_ROOT, capture_output=True).returncode == 0: - return ref - fallback = ref.replace("origin/", "upstream/", 1) - if ( - fallback != ref - and subprocess.run(["git", "rev-parse", "--verify", fallback], cwd=REPO_ROOT, capture_output=True).returncode - == 0 - ): - print(f"Note: '{ref}' not found, using '{fallback}' as base.", file=sys.stderr) - return fallback - return ref - - def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--base", default=DEFAULT_BASE) parser.add_argument("--update", action="store_true") args = parser.parse_args() - resolved = _resolve_base(args.base) if args.update: - cmd_update(count_basedpyright(sys.stdin.read()), resolved) + cmd_update(count_basedpyright(sys.stdin.read()), args.base) else: - cmd_check(resolved) + cmd_check(args.base) if __name__ == "__main__": diff --git a/tests/test_litellm/conftest.py b/tests/test_litellm/conftest.py index d44b417995a..f4aa1926d21 100644 --- a/tests/test_litellm/conftest.py +++ b/tests/test_litellm/conftest.py @@ -117,53 +117,6 @@ def _run_coroutine_if_needed(result): pass -def _shutdown_opentelemetry_providers(): - """ - Shutdown OpenTelemetry providers to stop background metric collection threads. - - This prevents the "Cannot call collect on a MetricReader until it is - registered on a MeterProvider" error that occurs when a PeriodicExportingMetricReader - is created but not properly registered/shutdown. - """ - try: - from opentelemetry import metrics, trace, _logs - from opentelemetry.sdk.metrics import MeterProvider as SDKMeterProvider - from opentelemetry.sdk.trace import TracerProvider as SDKTracerProvider - from opentelemetry.sdk._logs import LoggerProvider as SDKLoggerProvider - - # Shutdown MeterProvider if it's an SDK provider - meter_provider = metrics.get_meter_provider() - if isinstance(meter_provider, SDKMeterProvider): - try: - meter_provider.shutdown() - except Exception: - pass - - # Shutdown TracerProvider if it's an SDK provider - tracer_provider = trace.get_tracer_provider() - if isinstance(tracer_provider, SDKTracerProvider): - try: - tracer_provider.shutdown() - except Exception: - pass - - # Shutdown LoggerProvider if it's an SDK provider - try: - logger_provider = _logs.get_logger_provider() - if isinstance(logger_provider, SDKLoggerProvider): - try: - logger_provider.shutdown() - except Exception: - pass - except Exception: - pass - except ImportError: - # OpenTelemetry not installed - pass - except Exception: - pass - - def _close_handler_if_needed(handler): if handler is None: return @@ -509,7 +462,7 @@ def strict_isolation(): def pytest_sessionfinish(session, exitstatus): - """Close any globally cached HTTP clients and shutdown OpenTelemetry providers so xdist workers exit cleanly.""" + """Close any globally cached HTTP clients so xdist workers exit cleanly.""" _close_handler_if_needed(litellm.__dict__.get("module_level_client")) _close_handler_if_needed(litellm.__dict__.get("module_level_aclient")) litellm.__dict__.pop("module_level_client", None) @@ -519,5 +472,3 @@ def pytest_sessionfinish(session, exitstatus): _close_handler_if_needed(getattr(litellm, "aclient", None)) _close_handler_if_needed(getattr(litellm, "client", None)) _run_coroutine_if_needed(close_litellm_async_clients()) - # Shutdown OpenTelemetry providers to stop background metric collection threads - _shutdown_opentelemetry_providers() diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index 51a3ca9efa0..7ffd09b931f 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -2386,7 +2386,6 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): from opentelemetry.exporter.otlp.proto.http.metric_exporter import ( OTLPMetricExporter, ) - from opentelemetry.sdk.metrics import MeterProvider from opentelemetry.sdk.metrics.export import PeriodicExportingMetricReader config = OpenTelemetryConfig( @@ -2399,11 +2398,6 @@ class TestOpenTelemetryProtocolSelection(unittest.TestCase): self.assertIsInstance(reader, PeriodicExportingMetricReader) self.assertIsInstance(reader._exporter, OTLPMetricExporter) - # Properly shut down the reader to avoid background thread issues. - # The reader must be registered with a MeterProvider before shutdown. - meter_provider = MeterProvider(metric_readers=[reader]) - meter_provider.shutdown() - class TestOpenTelemetryExternalSpan(unittest.TestCase): """ diff --git a/tests/test_litellm/interactions/test_litellm_responses_bridge.py b/tests/test_litellm/interactions/test_litellm_responses_bridge.py index 33105cda088..17e7f9fc4ff 100644 --- a/tests/test_litellm/interactions/test_litellm_responses_bridge.py +++ b/tests/test_litellm/interactions/test_litellm_responses_bridge.py @@ -6,11 +6,7 @@ the litellm_responses bridge provider, which calls litellm.responses() internall """ import os -from unittest.mock import patch -import pytest - -from litellm.types.interactions.generated import InteractionsAPIResponse from tests.test_litellm.interactions.base_interactions_test import ( BaseInteractionsTest, ) @@ -30,25 +26,3 @@ class TestLiteLLMResponsesBridge(BaseInteractionsTest): def get_api_key(self) -> str: """Return the OpenAI API key from environment.""" return os.getenv("OPENAI_API_KEY", "") - - @pytest.mark.asyncio - async def test_acreate_simple(self): - """Test async interaction creation with mocked API call.""" - mock_response = InteractionsAPIResponse( - id="interaction-abc123", - status="completed", - model="gpt-4o", - outputs=[{"type": "text", "text": "The speed of light is approximately 299,792,458 meters per second."}], - usage={"input_tokens": 10, "output_tokens": 20}, - ) - - import litellm.interactions as interactions - - with patch("litellm.interactions.main.create", return_value=mock_response): - response = await interactions.acreate( - model=self.get_model(), - input="What is the speed of light?", - api_key="sk-fake-key-for-unit-test", - ) - assert response is not None - assert response.id is not None or response.status is not None diff --git a/tests/test_litellm/interactions/test_openapi_compliance.py b/tests/test_litellm/interactions/test_openapi_compliance.py index 7174e86069c..209e99895db 100644 --- a/tests/test_litellm/interactions/test_openapi_compliance.py +++ b/tests/test_litellm/interactions/test_openapi_compliance.py @@ -32,7 +32,8 @@ def _load_openapi_spec_dict() -> Dict[str, Any]: return response.json() except Exception as e: # pragma: no cover - defensive, env-dependent pytest.skip( - f"Skipping Google Interactions OpenAPI compliance tests - unable to load spec from {OPENAPI_SPEC_URL}: {e}" + f"Skipping Google Interactions OpenAPI compliance tests - " + f"unable to load spec from {OPENAPI_SPEC_URL}: {e}" ) @@ -120,7 +121,9 @@ class TestRequestCompliance: # Discriminator without explicit mapping — verify via oneOf one_of = content_schema.get("oneOf", []) ref_names = [opt["$ref"].split("/")[-1] for opt in one_of if "$ref" in opt] - assert "TextContent" in ref_names, f"TextContent not found in oneOf refs: {ref_names}" + assert ( + "TextContent" in ref_names + ), f"TextContent not found in oneOf refs: {ref_names}" print(f"Content type discriminator (no mapping), oneOf refs: {ref_names}") def test_text_content_schema(self, spec_dict): @@ -203,7 +206,9 @@ class TestResponseCompliance: expected_fields = ["total_input_tokens", "total_output_tokens", "total_tokens"] for field in expected_fields: - assert field in usage_schema["properties"], f"Usage field '{field}' not in spec" + assert ( + field in usage_schema["properties"] + ), f"Usage field '{field}' not in spec" print(f"✓ Usage field '{field}' exists") @@ -222,7 +227,9 @@ class TestToolsCompliance: """Verify FunctionDeclaration schema for function tools.""" if "FunctionDeclaration" in spec_dict["components"]["schemas"]: func_schema = spec_dict["components"]["schemas"]["FunctionDeclaration"] - assert "name" in func_schema.get("properties", {}) or "name" in func_schema.get("required", []) + assert "name" in func_schema.get( + "properties", {} + ) or "name" in func_schema.get("required", []) print("✓ FunctionDeclaration schema found") else: print("⚠ FunctionDeclaration schema not found (may be nested)") @@ -288,4 +295,6 @@ if __name__ == "__main__": if method in ["get", "post", "delete", "put", "patch"]: print(f" {method.upper()} {path}") - print(f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}...") + print( + f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}..." + ) diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py index 951b889dcd4..e430ce3b084 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_handler.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_handler.py @@ -2642,7 +2642,7 @@ def test_gemini_legacy_vertex_stop_finish_reason_normalised(): # Ensure the chunk is not treated as a ModelResponseStream mock_chunk.__class__ = type("FakeProtoChunk", (), {}) - with patch.dict("sys.modules", {"proto": MagicMock()}): + with patch("litellm.litellm_core_utils.streaming_handler.proto", create=True): wrapper.chunk_creator(chunk=mock_chunk) assert wrapper.received_finish_reason == "stop", ( @@ -2673,7 +2673,7 @@ def test_gemini_legacy_vertex_tool_calls_finish_reason_with_stop_enum(): mock_chunk.candidates = [mock_candidate] mock_chunk.__class__ = type("FakeProtoChunk", (), {}) - with patch.dict("sys.modules", {"proto": MagicMock()}): + with patch("litellm.litellm_core_utils.streaming_handler.proto", create=True): wrapper.chunk_creator(chunk=mock_chunk) # Signal that tool_calls were present in the stream diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py index 797a4e3a2bd..71e686563a5 100644 --- a/tests/test_litellm/litellm_core_utils/test_token_counter.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter.py @@ -363,7 +363,7 @@ def test_load_test_token_counter(model): total_time = end_time - start_time print("model={}, total test time={}".format(model, total_time)) - assert total_time < 30, f"Total encoding time > 30s, {total_time}" + assert total_time < 10, f"Total encoding time > 10s, {total_time}" def test_openai_token_with_image_and_text(): diff --git a/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py b/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py index 680905c1434..0a3bf403bf8 100644 --- a/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py +++ b/tests/test_litellm/llms/custom_httpx/test_credential_leak_prevention.py @@ -200,32 +200,8 @@ class TestSafeResponseHelpers: assert result == b"" -def _http_handler_symbols(): - """ - Re-import the relevant symbols from the current (possibly reloaded) version - of http_handler so that class-identity checks (isinstance / pytest.raises) - always use the same class object that _raise_masked_*_error will raise. - - Background: test_huggingface_embedding_handler.py calls - ``importlib.reload(litellm.llms.custom_httpx.http_handler)`` which - replaces the module-level ``MaskedHTTPStatusError`` class with a fresh - object. Functions defined in that module look up ``MaskedHTTPStatusError`` - via their ``__globals__`` (the module dict) at call time, so they raise the - NEW class. But tests that captured the OLD class at import time cannot - match it via ``pytest.raises`` or ``isinstance``. Fetching the class fresh - here avoids the mismatch. - """ - import litellm.llms.custom_httpx.http_handler as _m - return ( - _m.MaskedHTTPStatusError, - _m._raise_masked_sync_error, - _m._raise_masked_async_error, - ) - - class TestRaiseMaskedError: def test_sync_non_stream(self): - MaskedHTTPStatusError, _raise_masked_sync_error, _ = _http_handler_symbols() orig = _make_httpx_status_error( url="https://api.example.com?key=LEAKED_KEY", body="error body" ) @@ -238,7 +214,6 @@ class TestRaiseMaskedError: assert err.text == "error body" def test_sync_stream(self): - MaskedHTTPStatusError, _raise_masked_sync_error, _ = _http_handler_symbols() orig = _make_httpx_status_error( url="https://api.example.com?key=LEAKED_KEY", body="stream body" ) @@ -250,7 +225,6 @@ class TestRaiseMaskedError: assert err.message is not None def test_sync_breaks_exception_chain(self): - MaskedHTTPStatusError, _raise_masked_sync_error, _ = _http_handler_symbols() orig = _make_httpx_status_error() with pytest.raises(MaskedHTTPStatusError) as exc_info: _raise_masked_sync_error(orig, stream=False) @@ -259,7 +233,6 @@ class TestRaiseMaskedError: @pytest.mark.asyncio async def test_async_non_stream(self): - MaskedHTTPStatusError, _, _raise_masked_async_error = _http_handler_symbols() orig = _make_httpx_status_error( url="https://api.example.com?key=LEAKED_KEY", body="async error" ) @@ -273,7 +246,6 @@ class TestRaiseMaskedError: @pytest.mark.asyncio async def test_async_stream(self): - MaskedHTTPStatusError, _, _raise_masked_async_error = _http_handler_symbols() orig = _make_httpx_status_error( url="https://api.example.com?key=LEAKED_KEY", body="async stream" ) @@ -286,7 +258,6 @@ class TestRaiseMaskedError: @pytest.mark.asyncio async def test_async_breaks_chain(self): - MaskedHTTPStatusError, _, _raise_masked_async_error = _http_handler_symbols() orig = _make_httpx_status_error() with pytest.raises(MaskedHTTPStatusError) as exc_info: await _raise_masked_async_error(orig, stream=False) @@ -311,7 +282,6 @@ class TestHTTPHandlerErrorPaths: @pytest.mark.parametrize("method", ["post", "put", "patch", "delete"]) def test_sync_raises_masked_error(self, sync_handler, method): - MaskedHTTPStatusError, _, __ = _http_handler_symbols() with patch.object( sync_handler.client, "send", @@ -328,7 +298,6 @@ class TestHTTPHandlerErrorPaths: @pytest.mark.parametrize("method", ["post", "put", "patch", "delete"]) @pytest.mark.asyncio async def test_async_raises_masked_error(self, async_handler, method): - MaskedHTTPStatusError, _, __ = _http_handler_symbols() with patch.object( async_handler.client, "send", diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index 5021dd4bd2e..10539dd2fab 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -1024,16 +1024,7 @@ def test_sync_delete_responses_omits_body_for_azure(): captured: dict = {} _, fake_sync_delete = _build_delete_response_mock(captured) - # Patch via the full module path so we always target the *current* - # HTTPHandler class in sys.modules, not the one bound at import time. - # Using patch.object(HTTPHandler, ...) is fragile when another xdist - # worker test (e.g. test_huggingface_embedding_handler) reloads the - # http_handler module, creating a new class object: _get_httpx_client - # then returns instances of the new class, bypassing the old-class patch. - with patch( - "litellm.llms.custom_httpx.http_handler.HTTPHandler.delete", - new=fake_sync_delete, - ): + with patch.object(HTTPHandler, "delete", new=fake_sync_delete): litellm.delete_responses( response_id="resp_xyz", custom_llm_provider="azure", diff --git a/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py b/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py index aee7aa5e52b..8a072fa5097 100644 --- a/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py +++ b/tests/test_litellm/llms/huggingface/embedding/test_huggingface_embedding_handler.py @@ -21,20 +21,7 @@ def reload_huggingface_modules(): Reload modules to ensure fresh references after conftest reloads litellm. This ensures the HTTPHandler class being patched is the same one used by the embedding handler during parallel test execution. - - NOTE: The conftest module-reload only runs in *serial* mode (no xdist - worker). Reloading here in xdist mode is therefore unnecessary and - actively harmful: it replaces class objects inside http_handler with new - ones while other modules still hold references to the old ones, causing - isinstance / pytest.raises mismatches in sibling tests running in the - same worker. """ - import os - # Only reload when NOT running under xdist (matches the conftest guard). - if os.environ.get("PYTEST_XDIST_WORKER"): - yield - return - import litellm.llms.custom_httpx.http_handler as http_handler_module import litellm.llms.huggingface.embedding.handler as hf_embedding_handler_module diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py index cefd8eaa432..7e13459bca1 100644 --- a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py +++ b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py @@ -6,8 +6,6 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest -pytest.importorskip("botocore", reason="botocore is required for sagemaker tests") - sys.path.insert(0, os.path.abspath("../../../../..")) from litellm.llms.sagemaker.common_utils import AWSEventStreamDecoder from litellm.llms.sagemaker.completion.transformation import SagemakerConfig diff --git a/tests/test_litellm/models/__init__.py b/tests/test_litellm/models/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py b/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py index ce983782e42..b0d2595e48c 100644 --- a/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py +++ b/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py @@ -86,10 +86,6 @@ def _assert_not_role_blocked(response) -> None: err = detail.get("error", "") else: err = str(detail) - # err may itself be a nested dict (e.g. when the endpoint raises an - # unexpected exception and returns {"error": {"message": ...}}) - if not isinstance(err, str): - err = str(err) err_lower = err.lower() role_block_signals = ( "your role=", diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 54638f35401..90f46152837 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -1,7 +1,7 @@ import json import os import sys -from datetime import datetime, timedelta, timezone +from datetime import datetime, timedelta from types import SimpleNamespace from unittest.mock import ANY, AsyncMock, MagicMock, patch @@ -249,7 +249,7 @@ async def test_custom_auth_does_not_enforce_key_model_access_by_default(): async def test_post_custom_auth_expired_key_returns_unauthorized(): expired_token = UserAPIKeyAuth( token="test_token", - expires=datetime.now(timezone.utc) - timedelta(minutes=1), + expires=datetime.now() - timedelta(minutes=1), ) with pytest.raises(ProxyException) as exc_info: diff --git a/tests/test_litellm/proxy/client/__init__.py b/tests/test_litellm/proxy/client/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/test_litellm/proxy/client/test_credentials.py b/tests/test_litellm/proxy/client/test_credentials.py index bca113957b7..72c643467b2 100644 --- a/tests/test_litellm/proxy/client/test_credentials.py +++ b/tests/test_litellm/proxy/client/test_credentials.py @@ -13,6 +13,7 @@ import responses from litellm.proxy.client.credentials import CredentialsManagementClient from litellm.proxy.client.exceptions import UnauthorizedError +from litellm.proxy.credential_endpoints.endpoints import CredentialHelperUtils from litellm.types.utils import CredentialItem @@ -268,18 +269,7 @@ def test_get_unauthorized_error(client): def test_encrypt_credential_values_does_not_mutate_original(monkeypatch): """Ensure encrypt_credential_values returns a new encrypted object""" - try: - from litellm.proxy.credential_endpoints.endpoints import ( - CredentialHelperUtils, - ) - except ImportError as e: - pytest.skip(f"Proxy dependencies not available: {e}") - monkeypatch.setenv("LITELLM_SALT_KEY", "test-key") - monkeypatch.setattr( - "litellm.proxy.credential_endpoints.endpoints.encrypt_value_helper", - lambda value, key=None: f"encrypted_{value}", - ) credential = CredentialItem( credential_name="azure1", credential_values={"api_key": "sk-123"}, diff --git a/tests/test_litellm/proxy/conftest.py b/tests/test_litellm/proxy/conftest.py index 426e23c8d6f..1d71035b67f 100644 --- a/tests/test_litellm/proxy/conftest.py +++ b/tests/test_litellm/proxy/conftest.py @@ -17,31 +17,8 @@ from fastapi.testclient import TestClient _PROXY_MODULE_GLOBALS_TO_ISOLATE = ( "master_key", "prisma_client", - "user_config_file_path", - "general_settings", - "llm_router", - "llm_model_list", - "premium_user", - "store_model_in_db", - "proxy_logging_obj", - "redis_usage_cache", ) -# Known-good defaults to set at the START of every test so a corrupt -# snapshot from a previous worker run cannot propagate further. -_PROXY_MODULE_GLOBALS_RESET = { - "master_key": None, - "prisma_client": None, - "user_config_file_path": None, - "general_settings": {}, - "llm_router": None, - "llm_model_list": [], - "store_model_in_db": False, - # Prevent _FakeRedisCache leakage from test_redis_auth_cache_flag tests - # that call _init_cache with a mock Redis and don't restore redis_usage_cache. - "redis_usage_cache": None, -} - class StubClientNotConnectedError(Exception): pass @@ -74,10 +51,9 @@ def _isolate_proxy_module_globals(): Snapshot and restore module-level globals on litellm.proxy.proxy_server that tests sometimes mutate via raw setattr (not monkeypatch). - We also force a known-good initial state *before* yielding so that a - corrupt snapshot from a previous worker test cannot cascade into the - current test (snapshot/restore alone does not help when the pre-test - state is already wrong). + Without this, a leaked value — e.g. master_key set by a sibling test — + flips the auth short-circuit in user_api_key_auth and causes unrelated + tests in the same xdist worker to return 401 instead of 200. """ from litellm.proxy import proxy_server @@ -86,26 +62,6 @@ def _isolate_proxy_module_globals(): name: getattr(proxy_server, name, sentinel) for name in _PROXY_MODULE_GLOBALS_TO_ISOLATE } - - # Force clean defaults for the subset we know how to reset safely. - for name, default in _PROXY_MODULE_GLOBALS_RESET.items(): - setattr(proxy_server, name, default) - - # Reset litellm_config_cache redis backend — a _FakeRedisCache from - # test_redis_auth_cache_flag can leak here via _init_cache's direct - # litellm_config_cache.redis_cache = redis_usage_cache assignment, even - # though _init_cache is only called when the config has cache:true. - try: - from litellm.proxy.utils import litellm_config_cache - litellm_config_cache.redis_cache = None - except Exception: - pass - - # Also clear any stale FastAPI dependency overrides left by previous tests. - from litellm.proxy.proxy_server import app - saved_overrides = dict(app.dependency_overrides) - app.dependency_overrides.clear() - try: yield finally: @@ -115,9 +71,6 @@ def _isolate_proxy_module_globals(): delattr(proxy_server, name) else: setattr(proxy_server, name, value) - # Restore dependency overrides to pre-test state. - app.dependency_overrides.clear() - app.dependency_overrides.update(saved_overrides) @pytest.fixture(autouse=True) diff --git a/tests/test_litellm/proxy/hooks/test_proxy_hooks_init.py b/tests/test_litellm/proxy/hooks/test_proxy_hooks_init.py index aadc549b6bb..a6edd3db944 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_hooks_init.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_hooks_init.py @@ -39,31 +39,13 @@ def test_isolation_module_does_not_pull_in_proxy_utils(): import importlib import sys - # Save and restore the modules we evict so that subsequent tests in the - # same xdist worker are not left with a freshly-reimported module object. - # If we leave these entries absent, the next import creates a *new* module - # object, orphaning any function's __globals__ dict that still points at - # the old one — which silently breaks unittest.mock.patch on those names. - saved = { - mod: sys.modules.get(mod) - for mod in [ - "litellm.proxy.utils", - "litellm.proxy.management_endpoints.common_utils", - "litellm.llms.base_llm.managed_resources.isolation", - ] - } - for mod in saved: + for mod in [ + "litellm.proxy.utils", + "litellm.proxy.management_endpoints.common_utils", + "litellm.llms.base_llm.managed_resources.isolation", + ]: sys.modules.pop(mod, None) - try: - importlib.import_module("litellm.llms.base_llm.managed_resources.isolation") - assert "litellm.proxy.utils" not in sys.modules - assert "litellm.proxy.management_endpoints.common_utils" not in sys.modules - finally: - # Restore original module objects so later tests see a consistent - # sys.modules and their pre-imported names remain valid. - for mod, original in saved.items(): - if original is not None: - sys.modules[mod] = original - else: - sys.modules.pop(mod, None) + importlib.import_module("litellm.llms.base_llm.managed_resources.isolation") + assert "litellm.proxy.utils" not in sys.modules + assert "litellm.proxy.management_endpoints.common_utils" not in sys.modules diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 925bc72bb34..1482937ab3b 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -2254,11 +2254,10 @@ async def test_pass_through_request_query_params_forwarding(): @pytest.mark.asyncio async def test_pass_through_with_httpbin_redirect(): """ - Tests redirect handling in pass_through_request using a mocked HTTP transport. - The mock simulates: GET /redirect/1 -> 302 Location: /get -> 200 {url: ...} - No real network calls are made. + Integration test using httpbin.org redirect endpoint to test real redirect handling. + This tests the actual redirect handling capability end-to-end using the full pass_through_request function. """ - import json + from unittest.mock import MagicMock from fastapi import Request from starlette.datastructures import Headers, QueryParams @@ -2267,50 +2266,24 @@ async def test_pass_through_with_httpbin_redirect(): pass_through_request, ) - # Build the two responses the mock transport will return in order: - # 1. 302 redirect from /redirect/1 -> /get - redirect_response = httpx.Response( - status_code=302, - headers={"location": "https://httpbin.org/get"}, - content=b"", - request=httpx.Request("GET", "https://httpbin.org/redirect/1"), - ) - # 2. 200 final response from /get - final_body = json.dumps({"url": "https://httpbin.org/get"}).encode() - final_response = httpx.Response( - status_code=200, - headers={"content-type": "application/json"}, - content=final_body, - request=httpx.Request("GET", "https://httpbin.org/get"), - ) - - responses = iter([redirect_response, final_response]) - - class _MockTransport(httpx.AsyncBaseTransport): - async def handle_async_request(self, request: httpx.Request) -> httpx.Response: - return next(responses) - - # pass_through_request accesses get_async_httpx_client(...).client, so the - # mock must expose a .client attribute wrapping our transport. - mock_handler = MagicMock() - mock_handler.client = httpx.AsyncClient(transport=_MockTransport(), follow_redirects=True) - - # Create mock FastAPI request + # Create mock request mock_request = MagicMock(spec=Request) mock_request.method = "GET" mock_request.headers = Headers({}) mock_request.query_params = QueryParams("") + # Mock the body method to return empty bytes for GET request async def mock_body(): return b"" mock_request.body = mock_body + + # Mock user API key dict mock_user_api_key_dict = MagicMock() - with patch( - "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client", - return_value=mock_handler, - ): + try: + # Test with httpbin.org redirect endpoint + # This will redirect to httpbin.org/get response = await pass_through_request( request=mock_request, target="https://httpbin.org/redirect/1", @@ -2318,9 +2291,19 @@ async def test_pass_through_with_httpbin_redirect(): user_api_key_dict=mock_user_api_key_dict, ) - assert response.status_code == 200 - response_content = bytes(response.body).decode("utf-8") - assert '"url": "https://httpbin.org/get"' in response_content + # Should get the final response (200) from /get endpoint, not the redirect (302) + assert response.status_code == 200 + + # The response should be from the /get endpoint + response_content = bytes(response.body).decode("utf-8") + + # httpbin.org/get returns JSON with info about the request + assert '"url": "https://httpbin.org/get"' in response_content + except Exception as e: + # If httpbin.org is not accessible, skip the test + import pytest + + pytest.skip(f"Could not reach httpbin.org for integration test: {e}") @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py index 2979bd3feb4..a2f7476abd1 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py @@ -294,24 +294,9 @@ class TestUnifiedGuardrailCallTypeResolution: response_body = {"candidates": [{"content": {"parts": [{"text": "hello"}]}}]} - _UNIFIED_MOD = "litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail" - - # The module caches the translation mappings in a module-level global - # (`endpoint_guardrail_translation_mappings`). When tests run in parallel - # under xdist, a previously executed test in the same worker can populate - # that global with the *real* mapping before this test runs. The cache - # guard (`if … is None`) then skips calling `load_guardrail_translation_mappings` - # entirely, so the patch below would have no effect and - # `process_output_response` would never be awaited. - # - # Resetting the global to `None` inside the patch context forces the - # guard to fire and ensures the mock mapping is always used. with patch( - f"{_UNIFIED_MOD}.load_guardrail_translation_mappings" - ) as mock_load, patch( - f"{_UNIFIED_MOD}.endpoint_guardrail_translation_mappings", - None, - ): + "litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail.load_guardrail_translation_mappings" + ) as mock_load: mock_handler_instance = AsyncMock() mock_handler_instance.process_output_response = AsyncMock( return_value=response_body diff --git a/tests/test_litellm/proxy/test_update_llm_router_resilience.py b/tests/test_litellm/proxy/test_update_llm_router_resilience.py index d27fe9ce26a..fd0df4805e6 100644 --- a/tests/test_litellm/proxy/test_update_llm_router_resilience.py +++ b/tests/test_litellm/proxy/test_update_llm_router_resilience.py @@ -45,7 +45,6 @@ class TestUpdateLlmRouterResilience: mock_proxy_logging = MagicMock() with ( - patch("litellm.proxy.proxy_server.proxy_config", proxy_config), patch.object( proxy_config, "get_config", @@ -63,7 +62,6 @@ class TestUpdateLlmRouterResilience: patch("litellm.proxy.proxy_server.master_key", "sk-test"), patch("litellm.proxy.proxy_server.llm_model_list", []), patch("litellm.proxy.proxy_server.general_settings", {}), - patch("litellm.proxy.proxy_server.prisma_client", None), ): await proxy_config._update_llm_router( new_models=db_models, @@ -87,7 +85,6 @@ class TestUpdateLlmRouterResilience: mock_proxy_logging = MagicMock() with ( - patch("litellm.proxy.proxy_server.proxy_config", proxy_config), patch.object( proxy_config, "get_config", @@ -105,7 +102,6 @@ class TestUpdateLlmRouterResilience: patch("litellm.proxy.proxy_server.master_key", "sk-test"), patch("litellm.proxy.proxy_server.llm_model_list", []), patch("litellm.proxy.proxy_server.general_settings", {}), - patch("litellm.proxy.proxy_server.prisma_client", None), ): await proxy_config._update_llm_router( new_models=db_models, diff --git a/tests/test_litellm/test_compression.py b/tests/test_litellm/test_compression.py index 75bbff360c7..4fbcd4ed30d 100644 --- a/tests/test_litellm/test_compression.py +++ b/tests/test_litellm/test_compression.py @@ -443,19 +443,8 @@ def test_embedding_scorer_forwards_embedding_model_params(monkeypatch): # --------------------------------------------------------------------------- -def test_embedding_scorer(monkeypatch): - class _MockResponse: - data = [ - {"embedding": [1.0, 0.0, 0.0]}, - {"embedding": [0.0, 0.0, 1.0]}, - {"embedding": [0.9, 0.1, 0.0]}, - ] - - def fake_embedding(**kwargs): - return _MockResponse() - - monkeypatch.setattr(litellm, "embedding", fake_embedding) - +@pytest.mark.skipif(not os.environ.get("OPENAI_API_KEY"), reason="Needs OPENAI_API_KEY") +def test_embedding_scorer(): result = litellm.compress( messages=[ {"role": "user", "content": "Authentication code " * 2000}, diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 7004fa611b3..28cf4fa0744 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -12,7 +12,7 @@ sys.path.insert( ) # Adds the parent directory to the system path import urllib.parse -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import MagicMock, patch import litellm from litellm import main as litellm_main @@ -22,93 +22,18 @@ async def _async_fake_bedrock_image_details(image_url): return "ZmFrZS1pbWFnZQ==", "image/png" -def _check_call_args(mock_client, model: str) -> None: - """Inspect the captured HTTP call args and assert request-shaping rules.""" - print(mock_client.call_args.kwargs) - - if "data" in mock_client.call_args.kwargs: - json_str = mock_client.call_args.kwargs["data"] - else: - json_str = json.dumps(mock_client.call_args.kwargs.get("json", {})) - - if isinstance(json_str, bytes): - json_str = json_str.decode("utf-8") - - print(f"type of json_str: {type(json_str)}") - - # Bedrock models convert URLs to base64, while direct Anthropic models support URLs - # bedrock/invoke models use Anthropic messages API which supports URLs - if model.startswith("bedrock/invoke/"): - # bedrock/invoke should convert URLs to base64 (doesn't support URL references) - # URL should NOT be in the JSON (it should be converted to base64) - assert "https://awsmp-logos.s3.amazonaws.com" not in json_str - # Should have base64 data in the source (type="base64", not type="url") - assert '"type":"base64"' in json_str or '"type": "base64"' in json_str - # Should have "data" field containing base64 content - assert '"data"' in json_str - elif model.startswith("bedrock/"): - # Regular Bedrock models should convert URLs to base64 (uses "bytes" field) - # URL should NOT be in the JSON (it should be converted to base64) - assert "https://awsmp-logos.s3.amazonaws.com" not in json_str - # Should have "bytes" field (Bedrock uses "bytes" not "base64" in the field name) - assert '"bytes"' in json_str or '"bytes":' in json_str - elif model.startswith("anthropic/"): - # Direct Anthropic models should pass HTTPS URLs directly (HTTP URLs are converted to base64) - # Since we're using HTTPS URL, it should be passed as-is - assert "https://awsmp-logos.s3.amazonaws.com" in json_str - # For Anthropic, URL references use "url" type, not base64 - assert '"type":"url"' in json_str or '"type": "url"' in json_str - else: - # For other models, check format parameter is respected - assert "png" in json_str - assert "jpeg" not in json_str - - @pytest.fixture(autouse=True) def clear_client_cache(): """ - Clear (and replace) the HTTP client cache before each test to ensure mocks - are used. This prevents cached real clients from being reused across tests. - - Under xdist, tests from *other* files can run in the same worker process - and leave litellm.in_memory_llm_clients_cache in a broken state (e.g. - a previous test may have replaced it with a MagicMock or otherwise - corrupted it). To guard against this we always replace the cache with a - fresh LLMClientCache instance rather than just flushing the existing one. - We also clear the image-URL in-memory cache so stale URL→base64 mappings - from other tests cannot interfere. + Clear the HTTP client cache before each test to ensure mocks are used. + This prevents cached real clients from being reused across tests. """ - from litellm.caching.llm_caching_handler import LLMClientCache - - # Always install a fresh, known-good cache object. - fresh_cache = LLMClientCache() - setattr(litellm, "in_memory_llm_clients_cache", fresh_cache) - - # Also flush the image-URL cache that lives in image_handling.py. - try: - from litellm.litellm_core_utils.prompt_templates.image_handling import ( - in_memory_cache as _image_cache, - ) - - _image_cache.flush_cache() - except Exception: - pass - + cache = getattr(litellm, "in_memory_llm_clients_cache", None) + if cache is not None: + cache.flush_cache() yield - - # Flush both caches on teardown too. - try: - fresh_cache.flush_cache() - except Exception: - pass - try: - from litellm.litellm_core_utils.prompt_templates.image_handling import ( - in_memory_cache as _image_cache, - ) - - _image_cache.flush_cache() - except Exception: - pass + if cache is not None: + cache.flush_cache() @pytest.fixture(autouse=True) @@ -242,63 +167,27 @@ def test_completion_missing_role(openai_api_response): ) @pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio -async def test_url_with_format_param(model, sync_mode, monkeypatch): # noqa: PLR0912 +async def test_url_with_format_param(model, sync_mode, monkeypatch): from litellm import acompletion, completion from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.litellm_core_utils.prompt_templates import factory as prompt_factory - # This test is about request shaping, not live image downloads. Stub the - # URL->image conversion helpers at *every* level so that no matter which - # code path is taken (factory.py, image_handling.py, or the bedrock - # transformation module) we never make a real network call. This is - # especially important under xdist where a previous test in the same - # worker process may have left any of these module-level attributes in an - # unexpected state. - fake_base64_image = "data:image/png;base64,ZmFrZS1pbWFnZQ==" + if sync_mode: + client = HTTPHandler() + else: + client = AsyncHTTPHandler() - # Patch the reference that lives in factory.py (used by anthropic_messages_pt - # -> create_anthropic_image_param for the bedrock/invoke code path). + # This test is about request shaping, not live image downloads. Stub the + # URL->image conversion helpers so suite-level network/client state from + # earlier tests cannot prevent the mocked provider client from being hit. + fake_base64_image = "data:image/png;base64,ZmFrZS1pbWFnZQ==" monkeypatch.setattr( prompt_factory, "convert_url_to_base64", lambda url: fake_base64_image ) - - # Define the async stub here so it is always in scope for both try-blocks - # that follow (even if either import fails). - async def _fake_async_convert(url: str) -> str: - return fake_base64_image - - # Patch the original functions in image_handling so that any direct import - # of those names (e.g. in anthropic_claude3_transformation.py) is also - # intercepted. We do this defensively even for code paths that should not - # be reached with the current test input. - try: - import litellm.litellm_core_utils.prompt_templates.image_handling as _ih - - monkeypatch.setattr(_ih, "convert_url_to_base64", lambda url: fake_base64_image) - monkeypatch.setattr(_ih, "async_convert_url_to_base64", _fake_async_convert) - except Exception: - pass - - # Patch the references that were imported at module load time inside the - # bedrock anthropic transformation module. - try: - import litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation as _act - - monkeypatch.setattr( - _act, "convert_url_to_base64", lambda url: fake_base64_image - ) - monkeypatch.setattr( - _act, "async_convert_url_to_base64", _fake_async_convert - ) - except Exception: - pass - - # Patch BedrockImageProcessor helpers (used by the converse/bedrock-native - # code paths for the bedrock/us.* and bedrock/converse/* models). monkeypatch.setattr( prompt_factory.BedrockImageProcessor, "get_image_details", - staticmethod(lambda image_url: ("ZmFrZS1tbWFnZQ==", "image/png")), + staticmethod(lambda image_url: ("ZmFrZS1pbWFnZQ==", "image/png")), ) monkeypatch.setattr( prompt_factory.BedrockImageProcessor, @@ -326,59 +215,56 @@ async def test_url_with_format_param(model, sync_mode, monkeypatch): # noqa: PL } if model.startswith("gemini/"): args["api_key"] = "test-api-key" + with patch.object(client, "post", new=MagicMock()) as mock_client: + try: + if sync_mode: + response = completion(**args, client=client) + else: + response = await acompletion(**args, client=client) + print(response) + except Exception as e: + pass - if sync_mode: - # ------------------------------------------------------------------ # - # Sync path: patch HTTPHandler.post at the *class* level. - # - # Rationale: same as the async path below — if a previous xdist - # worker test left the HTTPHandler class reference in a broken state, - # the isinstance(client, HTTPHandler) check inside the provider - # handler fails and the framework creates a fresh internal client, - # bypassing the instance-level mock entirely. A class-level patch - # covers both the explicitly-passed client and any freshly-created - # internal client, making the test robust to xdist pollution. - # ------------------------------------------------------------------ # - with patch.object(HTTPHandler, "post", new_callable=MagicMock) as mock_client: - try: - response = completion(**args) - print(response) - except Exception as e: - pass + mock_client.assert_called() - mock_client.assert_called() - _check_call_args(mock_client, model) - else: - # ------------------------------------------------------------------ # - # Async path: patch AsyncHTTPHandler.post at the *class* level using - # AsyncMock instead of patching a single instance. - # - # Rationale: BaseLLMHTTPHandler.completion checks - # isinstance(client, AsyncHTTPHandler) - # before deciding whether to forward the caller-supplied client. If - # that check fails for any reason (e.g. another xdist worker test left - # the AsyncHTTPHandler name in a different state), the framework falls - # back to creating a fresh client via get_async_httpx_client(). A - # class-level patch covers *both* the explicit client and any freshly - # created internal client, so the mock is always called regardless of - # what happened to the AsyncHTTPHandler reference. - # - # Using AsyncMock (instead of MagicMock) also avoids the - # TypeError: object MagicMock can't be used in 'await' expression - # that would otherwise propagate through the error-handling stack in - # ways that differ across Python / asyncio versions. - # ------------------------------------------------------------------ # - with patch.object( - AsyncHTTPHandler, "post", new_callable=AsyncMock - ) as mock_client: - try: - response = await acompletion(**args) - print(response) - except Exception as e: - pass + print(mock_client.call_args.kwargs) - mock_client.assert_called() - _check_call_args(mock_client, model) + if "data" in mock_client.call_args.kwargs: + json_str = mock_client.call_args.kwargs["data"] + else: + json_str = json.dumps(mock_client.call_args.kwargs["json"]) + + if isinstance(json_str, bytes): + json_str = json_str.decode("utf-8") + + print(f"type of json_str: {type(json_str)}") + + # Bedrock models convert URLs to base64, while direct Anthropic models support URLs + # bedrock/invoke models use Anthropic messages API which supports URLs + if model.startswith("bedrock/invoke/"): + # bedrock/invoke should convert URLs to base64 (doesn't support URL references) + # URL should NOT be in the JSON (it should be converted to base64) + assert "https://awsmp-logos.s3.amazonaws.com" not in json_str + # Should have base64 data in the source (type="base64", not type="url") + assert '"type":"base64"' in json_str or '"type": "base64"' in json_str + # Should have "data" field containing base64 content + assert '"data"' in json_str + elif model.startswith("bedrock/"): + # Regular Bedrock models should convert URLs to base64 (uses "bytes" field) + # URL should NOT be in the JSON (it should be converted to base64) + assert "https://awsmp-logos.s3.amazonaws.com" not in json_str + # Should have "bytes" field (Bedrock uses "bytes" not "base64" in the field name) + assert '"bytes"' in json_str or '"bytes":' in json_str + elif model.startswith("anthropic/"): + # Direct Anthropic models should pass HTTPS URLs directly (HTTP URLs are converted to base64) + # Since we're using HTTPS URL, it should be passed as-is + assert "https://awsmp-logos.s3.amazonaws.com" in json_str + # For Anthropic, URL references use "url" type, not base64 + assert '"type":"url"' in json_str or '"type": "url"' in json_str + else: + # For other models, check format parameter is respected + assert "png" in json_str + assert "jpeg" not in json_str @pytest.mark.parametrize("model", ["gpt-4o-mini"]) diff --git a/ui/litellm-dashboard/src/components/OldTeams.test.tsx b/ui/litellm-dashboard/src/components/OldTeams.test.tsx index 2b8f36e684f..0b6e5786aaf 100644 --- a/ui/litellm-dashboard/src/components/OldTeams.test.tsx +++ b/ui/litellm-dashboard/src/components/OldTeams.test.tsx @@ -17,8 +17,6 @@ vi.mock("./networking", () => ({ v2TeamListCall: vi.fn(), getGuardrailsList: vi.fn().mockResolvedValue({ guardrails: [] }), getPoliciesList: vi.fn().mockResolvedValue({ policies: [] }), - vectorStoreListCall: vi.fn().mockResolvedValue({ data: [] }), - getAgentsList: vi.fn().mockResolvedValue({ agents: [] }), })); vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({ @@ -793,9 +791,7 @@ describe("OldTeams - access_group_ids in team create", () => { const createTeamSubmitButtons = screen.getAllByRole("button", { name: /create team/i }); const createTeamSubmitButton = createTeamSubmitButtons[createTeamSubmitButtons.length - 1]; - await act(async () => { - fireEvent.click(createTeamSubmitButton); - }); + fireEvent.click(createTeamSubmitButton); await waitFor(() => { expect(teamCreateCall).toHaveBeenCalledWith( diff --git a/ui/litellm-dashboard/tests/UsageViewSelect.adminFiltering.test.tsx b/ui/litellm-dashboard/tests/UsageViewSelect.adminFiltering.test.tsx deleted file mode 100644 index 4a7f4065e72..00000000000 --- a/ui/litellm-dashboard/tests/UsageViewSelect.adminFiltering.test.tsx +++ /dev/null @@ -1,169 +0,0 @@ -import React from "react"; -import { describe, it, expect, vi, beforeEach } from "vitest"; -import { render, screen, fireEvent, act } from "@testing-library/react"; -import { UsageViewSelect } from "../src/components/UsagePage/components/UsageViewSelect/UsageViewSelect"; - -// ── Types for antd mocks ───────────────────────────────────────────────────── - -type SelectOption = { - value: string; - label: React.ReactNode; -}; - -type SelectProps = { - value?: string; - onChange?: (value: string) => void; - options?: SelectOption[]; -}; - -type BadgeProps = { - count?: React.ReactNode; - children?: React.ReactNode; -}; - -// ── Mocks (mirrors the pattern from UsageViewSelect.test.tsx in src/) ────────── - -vi.mock("antd", async () => { - const React = await import("react"); - - function Select(props: SelectProps) { - const { value, onChange, options } = props; - return React.createElement( - "select", - { - value, - onChange: (e: React.ChangeEvent) => onChange?.(e.target.value), - role: "combobox", - }, - options?.map((opt) => React.createElement("option", { key: opt.value, value: opt.value }, opt.label)), - ); - } - Select.displayName = "AntdSelect"; - - function Badge(props: BadgeProps) { - return React.createElement("span", { "data-testid": "antd-badge" }, props.count, props.children); - } - Badge.displayName = "AntdBadge"; - - return { Select, Badge }; -}); - -vi.mock("@ant-design/icons", async () => { - const React = await import("react"); - const Icon = () => React.createElement("span", { "data-testid": "icon" }); - return { - GlobalOutlined: Icon, - BankOutlined: Icon, - TeamOutlined: Icon, - ShoppingCartOutlined: Icon, - TagsOutlined: Icon, - RobotOutlined: Icon, - UserOutlined: Icon, - LineChartOutlined: Icon, - BarChartOutlined: Icon, - }; -}); - -// ── Admin-only option values ─────────────────────────────────────────────────── - -const ADMIN_ONLY_VALUES = ["customer", "tag", "agent", "user", "user-agent-activity"]; -const ALL_VALUES = ["global", "organization", "team", "customer", "tag", "agent", "user", "user-agent-activity"]; -const NON_ADMIN_VISIBLE = ["global", "organization", "team"]; - -// ── Helpers ──────────────────────────────────────────────────────────────────── - -function getOptionValues(): string[] { - return Array.from(screen.getByRole("combobox").querySelectorAll("option")).map((o) => (o as HTMLOptionElement).value); -} - -// ── Tests ────────────────────────────────────────────────────────────────────── - -describe("UsageViewSelect — admin vs non-admin option filtering", () => { - const mockOnChange = vi.fn(); - - beforeEach(() => { - mockOnChange.mockClear(); - }); - - it("should render without crashing", () => { - render(); - expect(screen.getByRole("combobox")).toBeInTheDocument(); - }); - - it("should expose all 8 options to admin users", () => { - render(); - const values = getOptionValues(); - expect(values).toHaveLength(8); - ALL_VALUES.forEach((v) => expect(values).toContain(v)); - }); - - it("should hide adminOnly options from non-admin users", () => { - render(); - const values = getOptionValues(); - ADMIN_ONLY_VALUES.forEach((v) => expect(values).not.toContain(v)); - }); - - it("should show only 3 options (global, organization, team) for non-admin users", () => { - render(); - const values = getOptionValues(); - expect(values).toHaveLength(NON_ADMIN_VISIBLE.length); - NON_ADMIN_VISIBLE.forEach((v) => expect(values).toContain(v)); - }); - - it("should show 'Your Usage' instead of 'Global Usage' for non-admin users", () => { - render(); - const options = Array.from(screen.getByRole("combobox").querySelectorAll("option")); - const globalOption = options.find((o) => (o as HTMLOptionElement).value === "global") as HTMLOptionElement; - expect(globalOption.textContent).toBe("Your Usage"); - }); - - it("should show 'Global Usage' for admin users", () => { - render(); - const options = Array.from(screen.getByRole("combobox").querySelectorAll("option")); - const globalOption = options.find((o) => (o as HTMLOptionElement).value === "global") as HTMLOptionElement; - expect(globalOption.textContent).toBe("Global Usage"); - }); - - it("should show 'Your Organization Usage' instead of 'Organization Usage' for non-admin users", () => { - render(); - const options = Array.from(screen.getByRole("combobox").querySelectorAll("option")); - const orgOption = options.find((o) => (o as HTMLOptionElement).value === "organization") as HTMLOptionElement; - expect(orgOption.textContent).toBe("Your Organization Usage"); - }); - - it("should show 'Organization Usage' label for admin users", () => { - render(); - const options = Array.from(screen.getByRole("combobox").querySelectorAll("option")); - const orgOption = options.find((o) => (o as HTMLOptionElement).value === "organization") as HTMLOptionElement; - expect(orgOption.textContent).toBe("Organization Usage"); - }); - - it("should keep 'Team Usage' label unchanged for both admin and non-admin", () => { - render(); - const options = Array.from(screen.getByRole("combobox").querySelectorAll("option")); - const teamOption = options.find((o) => (o as HTMLOptionElement).value === "team") as HTMLOptionElement; - expect(teamOption.textContent).toBe("Team Usage"); - }); - - it("should call onChange with the correct option value when user changes selection", () => { - render(); - act(() => { - fireEvent.change(screen.getByRole("combobox"), { target: { value: "team" } }); - }); - expect(mockOnChange).toHaveBeenCalledWith("team"); - }); - - it("should use custom title and description when provided", () => { - render( - , - ); - expect(screen.getByText("My Custom Title")).toBeInTheDocument(); - expect(screen.getByText("My custom description")).toBeInTheDocument(); - }); -}); diff --git a/ui/litellm-dashboard/tests/useLogDetails.test.ts b/ui/litellm-dashboard/tests/useLogDetails.test.ts deleted file mode 100644 index 89dbbceb60a..00000000000 --- a/ui/litellm-dashboard/tests/useLogDetails.test.ts +++ /dev/null @@ -1,136 +0,0 @@ -import { describe, it, expect, vi, beforeEach } from "vitest"; -import { renderHook, waitFor } from "@testing-library/react"; -import React from "react"; -import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; - -// ── Mocks ───────────────────────────────────────────────────────────────────── - -vi.mock("@/components/networking", () => ({ - uiSpendLogDetailsCall: vi.fn(), -})); - -vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ - default: vi.fn(), -})); - -import { useLogDetails } from "../src/app/(dashboard)/hooks/logDetails/useLogDetails"; -import { uiSpendLogDetailsCall } from "@/components/networking"; -import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; - -// ── Helpers ──────────────────────────────────────────────────────────────────── - -const DEFAULT_AUTH = { - token: "mock-token", - accessToken: "mock-access-token", - userId: "user-1", - userEmail: "user@example.com", - userRole: "Admin", - premiumUser: false, - disabledPersonalKeyCreation: null, - showSSOBanner: false, -}; - -function makeWrapper() { - const qc = new QueryClient({ - defaultOptions: { queries: { retry: false } }, - }); - - const Wrapper = ({ children }: { children: React.ReactNode }) => - React.createElement(QueryClientProvider, { client: qc }, children); - - Wrapper.displayName = "QueryClientWrapper"; - return Wrapper; -} - -// ── Tests ────────────────────────────────────────────────────────────────────── - -describe("useLogDetails — conditional lazy loading", () => { - const mockUseAuthorized = vi.mocked(useAuthorized); - const mockApiCall = vi.mocked(uiSpendLogDetailsCall); - - beforeEach(() => { - vi.clearAllMocks(); - mockUseAuthorized.mockReturnValue(DEFAULT_AUTH); - mockApiCall.mockResolvedValue({ messages: [], response: {} }); - }); - - it("should not call the API when enabled is false", () => { - renderHook(() => useLogDetails("req-123", "2025-01-01 00:00:00", false), { wrapper: makeWrapper() }); - expect(mockApiCall).not.toHaveBeenCalled(); - }); - - it("should not call the API when requestId is undefined", () => { - renderHook(() => useLogDetails(undefined, "2025-01-01 00:00:00", true), { wrapper: makeWrapper() }); - expect(mockApiCall).not.toHaveBeenCalled(); - }); - - it("should not call the API when startTime is undefined", () => { - renderHook(() => useLogDetails("req-123", undefined, true), { wrapper: makeWrapper() }); - expect(mockApiCall).not.toHaveBeenCalled(); - }); - - it("should not call the API when accessToken is null", () => { - mockUseAuthorized.mockReturnValue({ ...DEFAULT_AUTH, accessToken: null }); - - renderHook(() => useLogDetails("req-123", "2025-01-01 00:00:00", true), { wrapper: makeWrapper() }); - expect(mockApiCall).not.toHaveBeenCalled(); - }); - - it("should call the API with accessToken, requestId and startTime when all conditions met", async () => { - const { result } = renderHook(() => useLogDetails("req-123", "2025-01-01 00:00:00", true), { - wrapper: makeWrapper(), - }); - - await waitFor(() => expect(result.current.isSuccess).toBe(true)); - - expect(mockApiCall).toHaveBeenCalledWith("mock-access-token", "req-123", "2025-01-01 00:00:00"); - }); - - it("should return the data from the API response", async () => { - const mockData = { messages: [{ role: "user", content: "hello" }], response: { id: "resp-1" } }; - mockApiCall.mockResolvedValue(mockData); - - const { result } = renderHook(() => useLogDetails("req-456", "2025-01-02 12:00:00", true), { - wrapper: makeWrapper(), - }); - - await waitFor(() => expect(result.current.isSuccess).toBe(true)); - - expect(result.current.data).toEqual(mockData); - }); - - it("should transition from disabled to enabled and trigger the API call", async () => { - const { result, rerender } = renderHook( - ({ enabled }: { enabled: boolean }) => useLogDetails("req-789", "2025-01-03 00:00:00", enabled), - { wrapper: makeWrapper(), initialProps: { enabled: false } }, - ); - - expect(mockApiCall).not.toHaveBeenCalled(); - - rerender({ enabled: true }); - - await waitFor(() => expect(result.current.isSuccess).toBe(true)); - - expect(mockApiCall).toHaveBeenCalledTimes(1); - expect(mockApiCall).toHaveBeenCalledWith("mock-access-token", "req-789", "2025-01-03 00:00:00"); - }); - - it("should expose isLoading=true while the API call is in progress", async () => { - // Make the API call never resolve during this check - let resolveCall!: (v: { messages: unknown[]; response: Record }) => void; - mockApiCall.mockReturnValue( - new Promise((res) => { - resolveCall = res; - }), - ); - - const { result } = renderHook(() => useLogDetails("req-loading", "2025-01-04 00:00:00", true), { - wrapper: makeWrapper(), - }); - - await waitFor(() => expect(result.current.isLoading).toBe(true)); - - // Clean up — resolve the pending promise to avoid open handles - resolveCall({ messages: [], response: {} }); - }); -}); diff --git a/ui/litellm-dashboard/tests/usePaginatedDailyActivity.test.ts b/ui/litellm-dashboard/tests/usePaginatedDailyActivity.test.ts deleted file mode 100644 index e2a61c0f65e..00000000000 --- a/ui/litellm-dashboard/tests/usePaginatedDailyActivity.test.ts +++ /dev/null @@ -1,316 +0,0 @@ -import React from "react"; -import { describe, it, expect, vi, beforeEach, afterEach } from "vitest"; -import { renderHook, act } from "@testing-library/react"; -import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -import { usePaginatedDailyActivity } from "../src/components/UsagePage/hooks/usePaginatedDailyActivity"; - -function makeWrapper() { - const qc = new QueryClient({ defaultOptions: { queries: { retry: false } } }); - - const Wrapper = ({ children }: { children: React.ReactNode }) => - React.createElement(QueryClientProvider, { client: qc }, children); - - Wrapper.displayName = "QueryClientWrapper"; - return Wrapper; -} - -/** Build a mock page response with controllable totals. */ -function mockPage(page: number, totalPages: number, extra: Record = {}) { - return { - results: [{ date: `2025-01-0${page}`, spend: page }], - metadata: { - total_pages: totalPages, - has_more: page < totalPages, - page, - total_spend: page * 10, - total_api_requests: page * 5, - total_prompt_tokens: 0, - total_completion_tokens: 0, - total_tokens: 0, - total_successful_requests: 0, - total_failed_requests: 0, - total_cache_read_input_tokens: 0, - total_cache_creation_input_tokens: 0, - ...extra, - }, - }; -} - -describe("usePaginatedDailyActivity", () => { - beforeEach(() => { - vi.useFakeTimers(); - }); - - afterEach(() => { - vi.useRealTimers(); - vi.clearAllMocks(); - }); - - it("should return EMPTY_DATA and not call fetchFn when enabled=false", () => { - const fetchFn = vi.fn(); - const { result } = renderHook( - () => - usePaginatedDailyActivity({ - fetchFn, - args: ["token", "2025-01-01", "2025-01-07"], - enabled: false, - }), - { wrapper: makeWrapper() }, - ); - - expect(fetchFn).not.toHaveBeenCalled(); - expect(result.current.loading).toBe(false); - expect(result.current.isFetchingMore).toBe(false); - expect(result.current.data.results).toHaveLength(0); - expect(result.current.data.metadata.total_pages).toBe(1); - }); - - it("should fetch page 1 and mark loading=false when total_pages=1", async () => { - const fetchFn = vi.fn().mockResolvedValue(mockPage(1, 1)); - const { result } = renderHook( - () => - usePaginatedDailyActivity({ - fetchFn, - args: ["token", "2025-01-01", "2025-01-07"], - enabled: true, - }), - { wrapper: makeWrapper() }, - ); - - // Initially loading - expect(result.current.loading).toBe(true); - - // Flush microtasks so the first page resolves - await act(async () => { - await vi.runAllTimersAsync(); - }); - - expect(fetchFn).toHaveBeenCalledTimes(1); - // Page is injected at index 3 - expect(fetchFn).toHaveBeenCalledWith("token", "2025-01-01", "2025-01-07", 1); - expect(result.current.loading).toBe(false); - expect(result.current.isFetchingMore).toBe(false); - expect(result.current.data.results).toHaveLength(1); - expect(result.current.progress.currentPage).toBe(1); - expect(result.current.progress.totalPages).toBe(1); - }); - - it("should auto-fetch pages 2..N and accumulate results", async () => { - const fetchFn = vi - .fn() - .mockResolvedValueOnce(mockPage(1, 3)) - .mockResolvedValueOnce(mockPage(2, 3)) - .mockResolvedValueOnce(mockPage(3, 3)); - - const { result } = renderHook( - () => - usePaginatedDailyActivity({ - fetchFn, - args: ["token", "2025-01-01", "2025-01-07"], - enabled: true, - }), - { wrapper: makeWrapper() }, - ); - - // Flush all timers (including the 300 ms delay between pages) and promises - await act(async () => { - await vi.runAllTimersAsync(); - }); - - expect(fetchFn).toHaveBeenCalledTimes(3); - expect(fetchFn).toHaveBeenNthCalledWith(1, "token", "2025-01-01", "2025-01-07", 1); - expect(fetchFn).toHaveBeenNthCalledWith(2, "token", "2025-01-01", "2025-01-07", 2); - expect(fetchFn).toHaveBeenNthCalledWith(3, "token", "2025-01-01", "2025-01-07", 3); - - expect(result.current.loading).toBe(false); - expect(result.current.isFetchingMore).toBe(false); - // All three pages' results should be accumulated - expect(result.current.data.results).toHaveLength(3); - expect(result.current.progress.currentPage).toBe(3); - expect(result.current.progress.totalPages).toBe(3); - }); - - it("should sum total_spend and total_api_requests across pages", async () => { - // page 1: spend=10, requests=5 - // page 2: spend=20, requests=10 - // page 3: spend=30, requests=15 - // Expected totals: spend=60, requests=30 - const fetchFn = vi - .fn() - .mockResolvedValueOnce(mockPage(1, 3)) - .mockResolvedValueOnce(mockPage(2, 3)) - .mockResolvedValueOnce(mockPage(3, 3)); - - const { result } = renderHook( - () => - usePaginatedDailyActivity({ - fetchFn, - args: ["token", "2025-01-01", "2025-01-07"], - enabled: true, - }), - { wrapper: makeWrapper() }, - ); - - await act(async () => { - await vi.runAllTimersAsync(); - }); - - expect(result.current.data.metadata.total_spend).toBe(60); - expect(result.current.data.metadata.total_api_requests).toBe(30); - }); - - it("should only flush state at batch boundaries (every 3 pages), not on every page", async () => { - // RENDER_BATCH_SIZE = 3, so with 6 pages we expect exactly 2 batch flushes - // at pages 3 and 6 (plus the initial page-1 setData). - const fetchFn = vi - .fn() - .mockImplementation((token: string, start: string, end: string, page: number) => - Promise.resolve(mockPage(page, 6)), - ); - - const setDataSpy = vi.fn(); - const { result } = renderHook( - () => - usePaginatedDailyActivity({ - fetchFn, - args: ["token", "2025-01-01", "2025-01-07"], - enabled: true, - }), - { wrapper: makeWrapper() }, - ); - - await act(async () => { - await vi.runAllTimersAsync(); - }); - - // The final accumulated result should have all 6 pages' results - expect(result.current.data.results).toHaveLength(6); - expect(result.current.progress.currentPage).toBe(6); - }); - - it("should set cancelled=true and stop fetching when cancel() is called", async () => { - // Provide 5 pages but cancel after page 1 resolves - const fetchFn = vi - .fn() - .mockImplementation((token: string, start: string, end: string, page: number) => - Promise.resolve(mockPage(page, 5)), - ); - - const { result } = renderHook( - () => - usePaginatedDailyActivity({ - fetchFn, - args: ["token", "2025-01-01", "2025-01-07"], - enabled: true, - }), - { wrapper: makeWrapper() }, - ); - - // Let page 1 complete - await act(async () => { - await Promise.resolve(); // flush microtasks for page 1 - }); - - // Cancel before pages 2-5 are fetched - act(() => { - result.current.cancel(); - }); - - // Advance timers to confirm no more fetches happen - await act(async () => { - await vi.runAllTimersAsync(); - }); - - expect(result.current.cancelled).toBe(true); - expect(result.current.isFetchingMore).toBe(false); - // fetchFn should have been called for page 1 and at most page 2 - // (depending on timing), but NOT for all 5 pages - expect(fetchFn.mock.calls.length).toBeLessThan(5); - }); - - it("should restart (increment fetchId) when args change", async () => { - const fetchFn = vi - .fn() - .mockImplementation((token: string, start: string, end: string, page: number) => - Promise.resolve(mockPage(page, 1)), - ); - - const initialArgs = ["token", "2025-01-01", "2025-01-07"]; - const { result, rerender } = renderHook( - ({ args }: { args: string[] }) => - usePaginatedDailyActivity({ - fetchFn, - args, - enabled: true, - }), - { wrapper: makeWrapper(), initialProps: { args: initialArgs } }, - ); - - await act(async () => { - await vi.runAllTimersAsync(); - }); - - const firstCallCount = fetchFn.mock.calls.length; - expect(firstCallCount).toBeGreaterThanOrEqual(1); - - // Change args to trigger a new fetch run - rerender({ args: ["token", "2025-01-08", "2025-01-14"] }); - - await act(async () => { - await vi.runAllTimersAsync(); - }); - - // New fetch should have fired with the updated date range - expect(fetchFn.mock.calls.length).toBeGreaterThan(firstCallCount); - const lastCall = fetchFn.mock.calls[fetchFn.mock.calls.length - 1]; - expect(lastCall[1]).toBe("2025-01-08"); - expect(lastCall[2]).toBe("2025-01-14"); - }); - - it("should set loading=false and isFetchingMore=false when fetchFn throws", async () => { - const fetchFn = vi.fn().mockRejectedValue(new Error("API error")); - - const { result } = renderHook( - () => - usePaginatedDailyActivity({ - fetchFn, - args: ["token", "2025-01-01", "2025-01-07"], - enabled: true, - }), - { wrapper: makeWrapper() }, - ); - - await act(async () => { - await vi.runAllTimersAsync(); - }); - - expect(result.current.loading).toBe(false); - expect(result.current.isFetchingMore).toBe(false); - }); - - it("should transition from enabled=false to enabled=true and start fetching", async () => { - const fetchFn = vi.fn().mockResolvedValue(mockPage(1, 1)); - - const { result, rerender } = renderHook( - ({ enabled }: { enabled: boolean }) => - usePaginatedDailyActivity({ - fetchFn, - args: ["token", "2025-01-01", "2025-01-07"], - enabled, - }), - { wrapper: makeWrapper(), initialProps: { enabled: false } }, - ); - - expect(fetchFn).not.toHaveBeenCalled(); - - rerender({ enabled: true }); - - await act(async () => { - await vi.runAllTimersAsync(); - }); - - expect(fetchFn).toHaveBeenCalledTimes(1); - expect(result.current.loading).toBe(false); - expect(result.current.data.results).toHaveLength(1); - }); -});