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:
Yaniv Israel 2026-07-08 01:26:51 +03:00
parent c71d7b4077
commit 5221441b8e
44 changed files with 353 additions and 1416 deletions

View file

@ -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

View file

@ -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",

View file

@ -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

View file

@ -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

View file

@ -1,7 +1,6 @@
"""
This is the litellm SMTP email integration
"""
import asyncio
from typing import List

View file

@ -1,7 +1,6 @@
"""
Enterprise specific logging utils
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingMetadata

View file

@ -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,

View file

@ -7,4 +7,4 @@ including custom SSO handlers and advanced authentication features.
from .custom_sso_handler import EnterpriseCustomSSOHandler
__all__ = ["EnterpriseCustomSSOHandler"]
__all__ = ["EnterpriseCustomSSOHandler"]

View file

@ -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"
)

View file

@ -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
"""

View file

@ -2,6 +2,7 @@
Enterprise internal user management endpoints
"""
from fastapi import APIRouter, Depends, HTTPException
from litellm.proxy._types import UserAPIKeyAuth

View file

@ -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

View file

@ -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()

View file

@ -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 "================================================"

View file

@ -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"

View file

@ -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

View file

@ -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__":

View file

@ -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__":

View file

@ -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()

View file

@ -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):
"""

View file

@ -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

View file

@ -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]}..."
)

View file

@ -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

View file

@ -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():

View file

@ -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",

View file

@ -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",

View file

@ -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

View file

@ -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

View file

@ -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=",

View file

@ -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:

View file

@ -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"},

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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},

View file

@ -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"])

View file

@ -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(

View file

@ -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();
});
});

View file

@ -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: {} });
});
});

View file

@ -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);
});
});