mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
chore: remove lint/format-only changes and non-feature files
Revert all lint-infra and black/ruff-reformat-only changes back to upstream/litellm_internal_staging so the PR diff shows only the DeepKeep guardrail feature: - Makefile, scripts/ruff_strict_gate.py, scripts/type_check_gate.py (lint-gate infra) - credential_migration.py + enterprise/* + assorted test files (black-reformat / xdist test-isolation drift) - backend/routes/allowlist.py (merge glue) Remove non-feature local artifacts: build-and-push.sh, deepkeep_tilt_config.yaml, stray __init__.py collision shims, and unrelated UI test files.
This commit is contained in:
parent
c71d7b4077
commit
5221441b8e
44 changed files with 353 additions and 1416 deletions
9
Makefile
9
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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -1,7 +1,6 @@
|
|||
"""
|
||||
This is the litellm SMTP email integration
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import List
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
"""
|
||||
Enterprise specific logging utils
|
||||
"""
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import StandardLoggingMetadata
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -7,4 +7,4 @@ including custom SSO handlers and advanced authentication features.
|
|||
|
||||
from .custom_sso_handler import EnterpriseCustomSSOHandler
|
||||
|
||||
__all__ = ["EnterpriseCustomSSOHandler"]
|
||||
__all__ = ["EnterpriseCustomSSOHandler"]
|
||||
|
|
@ -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"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,<uuid>;target_model_names,<models>;resource_id,<vs_id>;model_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
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@
|
|||
Enterprise internal user management endpoints
|
||||
"""
|
||||
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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 "================================================"
|
||||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]}..."
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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=",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"},
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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},
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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<HTMLSelectElement>) => 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(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={false} />);
|
||||
expect(screen.getByRole("combobox")).toBeInTheDocument();
|
||||
});
|
||||
|
||||
it("should expose all 8 options to admin users", () => {
|
||||
render(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={true} />);
|
||||
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(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={false} />);
|
||||
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(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={false} />);
|
||||
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(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={false} />);
|
||||
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(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={true} />);
|
||||
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(<UsageViewSelect value="organization" onChange={mockOnChange} isAdmin={false} />);
|
||||
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(<UsageViewSelect value="organization" onChange={mockOnChange} isAdmin={true} />);
|
||||
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(<UsageViewSelect value="team" onChange={mockOnChange} isAdmin={false} />);
|
||||
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(<UsageViewSelect value="global" onChange={mockOnChange} isAdmin={true} />);
|
||||
act(() => {
|
||||
fireEvent.change(screen.getByRole("combobox"), { target: { value: "team" } });
|
||||
});
|
||||
expect(mockOnChange).toHaveBeenCalledWith("team");
|
||||
});
|
||||
|
||||
it("should use custom title and description when provided", () => {
|
||||
render(
|
||||
<UsageViewSelect
|
||||
value="global"
|
||||
onChange={mockOnChange}
|
||||
isAdmin={false}
|
||||
title="My Custom Title"
|
||||
description="My custom description"
|
||||
/>,
|
||||
);
|
||||
expect(screen.getByText("My Custom Title")).toBeInTheDocument();
|
||||
expect(screen.getByText("My custom description")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
|
@ -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<string, unknown> }) => 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: {} });
|
||||
});
|
||||
});
|
||||
|
|
@ -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<string, unknown> = {}) {
|
||||
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);
|
||||
});
|
||||
});
|
||||
Loading…
Add table
Reference in a new issue