Merge branch 'main' into litellm_create-character-endpoint-fixes

This commit is contained in:
Sameer Kankute 2026-03-16 17:58:16 +05:30 • committed by GitHub
commit 9beec825d4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
60 changed files with 1971 additions and 1654 deletions

File diff suppressed because it is too large Load diff

View file

@ -131,6 +131,10 @@ class CheckBatchCost:
# every subsequent poll cycle.
if self._has_batch_processed_column:
try:
# Include "complete"/"completed" batches: the retrieve_batch
# endpoint may transition a batch to "complete" before
# CheckBatchCost runs. The batch_processed=False filter
# already prevents reprocessing finished batches.
jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many(
where={
"file_purpose": "batch",
@ -140,8 +144,6 @@ class CheckBatchCost:
"failed",
"expired",
"cancelled",
"complete",
"completed",
"stale_expired",
]
},

View file

@ -199,7 +199,8 @@ def create_batch( # noqa: PLR0915
)
### TIMEOUT LOGIC ###
timeout = _resolve_timeout(optional_params, kwargs, custom_llm_provider)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=None,
optional_params=optional_params.model_dump(),
@ -207,7 +208,6 @@ def create_batch( # noqa: PLR0915
"litellm_call_id": litellm_call_id,
"proxy_server_request": proxy_server_request,
"model_info": model_info,
"metadata": metadata,
"preset_cache_key": None,
"stream_response": {},
**optional_params.model_dump(exclude_unset=True),
@ -584,7 +584,8 @@ def retrieve_batch(
**kwargs,
)
if litellm_logging_obj is not None:
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
user=None,
optional_params=optional_params.model_dump(),

View file

@ -91,7 +91,8 @@ def create_sync_endpoint_function(endpoint_config: Dict) -> Callable:
optional_params = {k: kwargs.get(k) for k in path_params if k in kwargs}
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
optional_params=optional_params,
litellm_params={"litellm_call_id": litellm_call_id},

View file

@ -233,7 +233,8 @@ def create_container(
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
optional_params=dict(container_create_request_params),
litellm_params={
@ -438,7 +439,8 @@ def list_containers(
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
optional_params=dict(container_list_optional_params),
litellm_params={
@ -626,7 +628,8 @@ def retrieve_container(
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
optional_params={},
litellm_params={
@ -811,7 +814,8 @@ def delete_container(
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
optional_params={},
litellm_params={
@ -1010,7 +1014,8 @@ def list_container_files(
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
optional_params={
"container_id": container_id,
@ -1255,7 +1260,8 @@ def upload_container_file(
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
optional_params={"container_id": container_id},
litellm_params={

View file

@ -193,7 +193,8 @@ def create_eval(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params=request_body,
litellm_params={
@ -382,7 +383,8 @@ def list_evals(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params=query_params,
litellm_params={
@ -536,7 +538,8 @@ def get_eval(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"eval_id": eval_id},
litellm_params={
@ -760,7 +763,8 @@ def update_eval(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params=request_body,
litellm_params={
@ -914,7 +918,8 @@ def delete_eval(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"eval_id": eval_id},
litellm_params={
@ -1071,7 +1076,8 @@ def cancel_eval(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"eval_id": eval_id},
litellm_params={
@ -1262,7 +1268,8 @@ def create_run(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params=request_body,
litellm_params={
@ -1450,7 +1457,8 @@ def list_runs(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"eval_id": eval_id, **query_params},
litellm_params={
@ -1610,7 +1618,8 @@ def get_run(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"eval_id": eval_id, "run_id": run_id},
litellm_params={
@ -1773,7 +1782,8 @@ def cancel_run(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"eval_id": eval_id, "run_id": run_id},
litellm_params={
@ -1941,7 +1951,8 @@ def delete_run(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"eval_id": eval_id, "run_id": run_id},
litellm_params={

View file

@ -185,7 +185,8 @@ class GenerateContentHelper:
if litellm_logging_obj is None:
raise ValueError("litellm_logging_obj is required, but got None")
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=dict(generate_content_config_dict),
litellm_params={

View file

@ -40,6 +40,9 @@ from litellm.utils import exception_type, get_litellm_params
llm_http_handler: BaseLLMHTTPHandler = BaseLLMHTTPHandler()
from openai.types.audio.transcription_create_params import FileTypes # type: ignore
# BFL handlers
from litellm.llms.black_forest_labs.image_edit.handler import bfl_image_edit
from litellm.llms.black_forest_labs.image_generation.handler import bfl_image_generation
from litellm.main import (
azure_chat_completions,
base_llm_aiohttp_handler,
@ -50,10 +53,6 @@ from litellm.main import (
openai_image_variations,
)
# BFL handlers
from litellm.llms.black_forest_labs.image_edit.handler import bfl_image_edit
from litellm.llms.black_forest_labs.image_generation.handler import bfl_image_generation
###########################################
from litellm.secret_managers.main import get_secret_str
from litellm.types.images.main import ImageEditOptionalRequestParams
@ -297,7 +296,8 @@ def image_generation( # noqa: PLR0915
litellm_params_dict = get_litellm_params(**kwargs)
logging: Logging = litellm_logging_obj
logging.update_environment_variables(
logging.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params=optional_params,
@ -308,7 +308,6 @@ def image_generation( # noqa: PLR0915
"logger_fn": logger_fn,
"proxy_server_request": proxy_server_request,
"model_info": model_info,
"metadata": metadata,
"preset_cache_key": None,
"stream_response": {},
},
@ -894,7 +893,8 @@ def image_edit( # noqa: PLR0915
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params=dict(image_edit_request_params),
@ -902,7 +902,6 @@ def image_edit( # noqa: PLR0915
**image_edit_request_params,
"litellm_call_id": litellm_call_id,
"model_info": model_info,
"metadata": metadata,
},
custom_llm_provider=custom_llm_provider,
)

View file

@ -34,16 +34,7 @@ Usage:
import asyncio
import contextvars
from functools import partial
from typing import (
Any,
AsyncIterator,
Coroutine,
Dict,
Iterator,
List,
Optional,
Union,
)
from typing import Any, AsyncIterator, Coroutine, Dict, Iterator, List, Optional, Union
import httpx
@ -306,7 +297,8 @@ def create(
**kwargs,
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=dict(optional_params),
litellm_params={"litellm_call_id": litellm_call_id},
@ -416,7 +408,8 @@ def get(
f"Interactions API not supported for: {custom_llm_provider}"
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"interaction_id": interaction_id},
litellm_params={"litellm_call_id": litellm_call_id},
@ -519,7 +512,8 @@ def delete(
f"Interactions API not supported for: {custom_llm_provider}"
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"interaction_id": interaction_id},
litellm_params={"litellm_call_id": litellm_call_id},
@ -622,7 +616,8 @@ def cancel(
f"Interactions API not supported for: {custom_llm_provider}"
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"interaction_id": interaction_id},
litellm_params={"litellm_call_id": litellm_call_id},

View file

@ -568,6 +568,42 @@ class Logging(LiteLLMLoggingBaseClass):
if "custom_llm_provider" in self.model_call_details:
self.custom_llm_provider = self.model_call_details["custom_llm_provider"]
def update_from_kwargs(
self,
kwargs: Dict,
litellm_params: Optional[Dict] = None,
optional_params: Optional[Dict] = None,
model: Optional[str] = None,
user: Optional[str] = None,
**additional_params,
):
"""
Convenience wrapper around update_environment_variables that
automatically extracts metadata/litellm_metadata from kwargs,
so callers don't need to manually plumb them into litellm_params.
"""
base_litellm_params: Dict[str, Any] = {}
if "metadata" in kwargs:
base_litellm_params["metadata"] = kwargs["metadata"]
if "litellm_metadata" in kwargs and isinstance(
kwargs["litellm_metadata"], dict
):
base_litellm_params["litellm_metadata"] = kwargs["litellm_metadata"]
if "metadata" not in base_litellm_params:
base_litellm_params["metadata"] = kwargs["litellm_metadata"].copy()
if litellm_params:
base_litellm_params.update(litellm_params)
self.update_environment_variables(
litellm_params=base_litellm_params,
optional_params=optional_params or {},
model=model,
user=user,
**additional_params,
)
def update_messages(self, messages: List[AllMessageValues]):
"""
Update the logged value of the messages in the model_call_details

View file

@ -1882,12 +1882,11 @@ class BaseLLMHTTPHandler:
headers=headers, provider=custom_llm_provider
)
logging_obj.update_environment_variables(
logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=dict(anthropic_messages_optional_request_params),
litellm_params={
"metadata": kwargs.get("metadata", {}),
"litellm_metadata": kwargs.get("litellm_metadata", {}),
"preset_cache_key": None,
"stream_response": {},
**anthropic_messages_optional_request_params,

View file

@ -69,7 +69,8 @@ class LiteLLMSkillsTransformationHandler:
"""
# Pre-call logging
if logging_obj:
logging_obj.update_environment_variables(
logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"display_title": display_title},
litellm_params={"litellm_call_id": litellm_call_id},
@ -172,7 +173,8 @@ class LiteLLMSkillsTransformationHandler:
"""
# Pre-call logging
if logging_obj:
logging_obj.update_environment_variables(
logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"limit": limit, "offset": offset},
litellm_params={"litellm_call_id": litellm_call_id},
@ -231,7 +233,8 @@ class LiteLLMSkillsTransformationHandler:
"""
# Pre-call logging
if logging_obj:
logging_obj.update_environment_variables(
logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"skill_id": skill_id},
litellm_params={"litellm_call_id": litellm_call_id},
@ -277,7 +280,8 @@ class LiteLLMSkillsTransformationHandler:
"""
# Pre-call logging
if logging_obj:
logging_obj.update_environment_variables(
logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"skill_id": skill_id},
litellm_params={"litellm_call_id": litellm_call_id},

View file

@ -298,7 +298,8 @@ def ocr(
verbose_logger.debug(f"OCR optional_params after mapping: {optional_params}")
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params={

View file

@ -32,6 +32,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
get_original_file_id,
prepare_data_with_credentials,
resolve_input_file_id_to_unified,
resolve_output_file_ids_to_unified,
update_batch_in_database,
)
from litellm.proxy.utils import handle_exception_on_proxy, is_known_model
@ -405,9 +406,11 @@ async def retrieve_batch( # noqa: PLR0915
verbose_proxy_logger=verbose_proxy_logger,
)
# If batch is in a terminal state, return immediately
# If batch is in a terminal state, return immediately.
# Include "complete" (DB-normalized form of "completed").
if response is not None and response.status in [
"completed",
"complete",
"failed",
"cancelled",
"expired",
@ -417,10 +420,11 @@ async def retrieve_batch( # noqa: PLR0915
data=data, user_api_key_dict=user_api_key_dict, response=response
)
# async_post_call_success_hook replaces batch.id and output_file_id with unified IDs
# but not input_file_id. Resolve raw provider ID to unified ID.
# The DB may store raw provider file IDs (before hooks translate them).
# Resolve any raw input/output/error file IDs to unified IDs.
if unified_batch_id:
await resolve_input_file_id_to_unified(response, prisma_client)
await resolve_output_file_ids_to_unified(response, prisma_client)
asyncio.create_task(
proxy_logging_obj.update_request_status(

View file

@ -697,6 +697,28 @@ async def resolve_input_file_id_to_unified(response, prisma_client) -> None:
pass
async def resolve_output_file_ids_to_unified(response, prisma_client) -> None:
"""
If the batch response contains raw provider output_file_id or error_file_id
(not already unified IDs), look up the corresponding unified file IDs from
the managed file table and replace them in-place.
"""
if not prisma_client:
return
for attr in ("output_file_id", "error_file_id"):
raw_id = getattr(response, attr, None)
if not raw_id or _is_base64_encoded_unified_file_id(raw_id):
continue
try:
managed_file = await prisma_client.db.litellm_managedfiletable.find_first(
where={"flat_model_file_ids": {"has": raw_id}}
)
if managed_file:
setattr(response, attr, managed_file.unified_file_id)
except Exception:
pass
async def get_batch_from_database(
batch_id: str,
unified_batch_id: Union[str, Literal[False]],
@ -809,14 +831,43 @@ async def update_batch_in_database(
# Normalize status for database storage
db_status = response.status if response.status != "completed" else "complete"
await prisma_client.db.litellm_managedobjecttable.update(
where={"unified_object_id": batch_id},
data={
"status": db_status,
"file_object": response.model_dump_json(),
"updated_at": litellm.utils.get_utc_datetime(),
},
)
update_data: dict = {
"status": db_status,
"file_object": response.model_dump_json(),
"updated_at": litellm.utils.get_utc_datetime(),
}
# When a batch reaches completion, also mark batch_processed=True.
# The cost callback is enqueued asynchronously during the
# aretrieve_batch call that detected completion (via the @client
# decorator). It is not awaited, so there is a theoretical window
# where the callback hasn't executed yet. In practice the callback
# completes reliably. Setting the flag here unblocks file deletion
# which queries batch_processed=False. CheckBatchCost acts as a
# safety net for the rare case where the callback fails.
if db_status == "complete":
update_data["batch_processed"] = True
try:
await prisma_client.db.litellm_managedobjecttable.update(
where={"unified_object_id": batch_id},
data=update_data,
)
except Exception as col_err:
# If the batch_processed column doesn't exist (old schema),
# retry without it so the status update still succeeds.
err_str = str(col_err).lower()
if "batch_processed" in err_str and update_data.get("batch_processed") is not None:
verbose_proxy_logger.warning(
f"batch_processed column not found, retrying update without it: {col_err}"
)
update_data.pop("batch_processed", None)
await prisma_client.db.litellm_managedobjecttable.update(
where={"unified_object_id": batch_id},
data=update_data,
)
else:
raise
except Exception as e:
verbose_proxy_logger.error(
f"Failed to update batch status in ManagedObjectTable: {e}"

View file

@ -557,11 +557,11 @@ class ProxyInitializationHelpers:
envvar="MAX_REQUESTS_BEFORE_RESTART",
)
@click.option(
"--skip_db_migration_check",
"--enforce_prisma_migration_check",
is_flag=True,
default=False,
help="Warn and continue instead of exiting when database migration fails.",
envvar="SKIP_DB_MIGRATION_CHECK",
help="Exit with error if database migration fails on startup.",
envvar="ENFORCE_PRISMA_MIGRATION_CHECK",
)
def run_server( # noqa: PLR0915
host,
@ -602,7 +602,7 @@ def run_server( # noqa: PLR0915
skip_server_startup,
keepalive_timeout,
max_requests_before_restart,
skip_db_migration_check: bool,
enforce_prisma_migration_check: bool,
):
args = locals()
if local:
@ -716,6 +716,7 @@ def run_server( # noqa: PLR0915
for k, v in new_env_var.items():
os.environ[k] = v
litellm_settings = None
if config is not None:
"""
Allow user to pass in db url via config
@ -830,7 +831,9 @@ def run_server( # noqa: PLR0915
"pool_timeout": db_connection_timeout,
}
database_url = get_secret("DATABASE_URL", default_value=None)
modified_url = append_query_params(database_url, params)
modified_url = append_query_params(
str(database_url) if database_url else None, params
)
os.environ["DATABASE_URL"] = modified_url
if os.getenv("DIRECT_URL", None) is not None:
### add connection pool + pool timeout args
@ -865,17 +868,17 @@ def run_server( # noqa: PLR0915
if not PrismaManager.setup_database(
use_migrate=not use_prisma_db_push
):
if skip_db_migration_check:
print( # noqa
"\033[1;33mLiteLLM Proxy: Database migration failed but continuing startup. "
"Pass --skip_db_migration_check to allow this.\033[0m"
)
else:
if enforce_prisma_migration_check:
print( # noqa
"\033[1;31mLiteLLM Proxy: Database setup failed after multiple retries. "
"The proxy cannot start safely. Please check your database connection and migration status.\033[0m"
)
sys.exit(1)
else:
print( # noqa
"\033[1;33mLiteLLM Proxy: Database migration failed but continuing startup. "
"Set --enforce_prisma_migration_check or ENFORCE_PRISMA_MIGRATION_CHECK=true to exit on failure.\033[0m"
)
else:
print( # noqa
f"Unable to connect to DB. DATABASE_URL found in environment, but prisma package not found." # noqa

View file

@ -136,7 +136,8 @@ async def acreate_realtime_client_secret(
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model_name,
optional_params={"expires_after": expires_after, "session": session},
litellm_params={"api_base": resolved_api_base},
@ -186,7 +187,8 @@ async def arealtime_calls(
dynamic_api_key=dynamic_api_key,
litellm_params=litellm_params,
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model_name,
optional_params={"realtime_calls": True, "session": session},
litellm_params={"api_base": resolved_api_base},
@ -247,7 +249,8 @@ async def _arealtime( # noqa: PLR0915
if query_params is not None:
query_params = {**query_params, "model": model}
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params={},

View file

@ -108,7 +108,6 @@ def rerank( # noqa: PLR0915
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
proxy_server_request = kwargs.get("proxy_server_request", None)
model_info = kwargs.get("model_info", None)
metadata = kwargs.get("metadata", {})
user = kwargs.get("user", None)
client = kwargs.get("client", None)
try:
@ -164,7 +163,8 @@ def rerank( # noqa: PLR0915
model_response = RerankResponse()
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params=dict(optional_rerank_params),
@ -172,7 +172,6 @@ def rerank( # noqa: PLR0915
"litellm_call_id": litellm_call_id,
"proxy_server_request": proxy_server_request,
"model_info": model_info,
"metadata": metadata,
"preset_cache_key": None,
"stream_response": {},
**optional_params.model_dump(exclude_unset=True),

View file

@ -692,11 +692,11 @@ def responses(
return run_async_function(aresponses_api_with_mcp, **mcp_call_kwargs)
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=custom_llm_provider,
)
)
local_vars.update(kwargs)
@ -738,11 +738,9 @@ def responses(
)
)
# Pre Call logging - preserve metadata for custom callbacks
# When called from completion bridge (codex models), metadata is in litellm_metadata
metadata_for_callbacks = metadata or kwargs.get("litellm_metadata") or {}
litellm_logging_obj.update_environment_variables(
# Pre Call logging
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params=dict(responses_api_request_params),
@ -750,8 +748,6 @@ def responses(
**responses_api_request_params,
"aresponses": _is_async,
"litellm_call_id": litellm_call_id,
"metadata": metadata_for_callbacks,
"litellm_metadata": kwargs.get("litellm_metadata", {}),
},
custom_llm_provider=custom_llm_provider,
)
@ -912,11 +908,11 @@ def delete_responses(
raise ValueError("custom_llm_provider is required but passed as None")
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None:
@ -927,7 +923,8 @@ def delete_responses(
local_vars.update(kwargs)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=local_vars,
model=None,
optional_params={
"response_id": response_id,
@ -1092,11 +1089,11 @@ def get_responses(
raise ValueError("custom_llm_provider is required but passed as None")
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None:
@ -1107,7 +1104,8 @@ def get_responses(
local_vars.update(kwargs)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=local_vars,
model=None,
optional_params={
"response_id": response_id,
@ -1249,11 +1247,11 @@ def list_input_items(
if custom_llm_provider is None:
raise ValueError("custom_llm_provider is required but passed as None")
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None:
@ -1263,7 +1261,8 @@ def list_input_items(
local_vars.update(kwargs)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=local_vars,
model=None,
optional_params={"response_id": response_id},
litellm_params={"litellm_call_id": litellm_call_id},
@ -1407,11 +1406,11 @@ def cancel_responses(
raise ValueError("custom_llm_provider is required but passed as None")
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=None,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None:
@ -1422,7 +1421,8 @@ def cancel_responses(
local_vars.update(kwargs)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=local_vars,
model=None,
optional_params={
"response_id": response_id,
@ -1594,11 +1594,11 @@ def compact_responses(
raise ValueError("custom_llm_provider is required but passed as None")
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=custom_llm_provider,
responses_api_provider_config: Optional[BaseResponsesAPIConfig] = (
ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=custom_llm_provider,
)
)
if responses_api_provider_config is None:
@ -1626,7 +1626,8 @@ def compact_responses(
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=local_vars,
model=model,
optional_params=dict(responses_api_request_params),
litellm_params={
@ -1729,7 +1730,8 @@ async def _aresponses_websocket(
api_key=api_key,
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params={},

View file

@ -286,7 +286,8 @@ def search(
# Pre Call logging
model_name = f"{search_provider}/search"
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model_name,
optional_params=optional_params,
litellm_params={

View file

@ -204,7 +204,8 @@ def create_skill(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params=request_body,
litellm_params={
@ -389,7 +390,8 @@ def list_skills(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params=query_params,
litellm_params={
@ -556,7 +558,8 @@ def get_skill(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"skill_id": skill_id},
litellm_params={
@ -722,7 +725,8 @@ def delete_skill(
)
# Pre-call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"skill_id": skill_id},
litellm_params={

View file

@ -146,7 +146,8 @@ def create(
)
create_request["file_id"] = file_id
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={
"vector_store_id": vector_store_id,
@ -279,7 +280,8 @@ def list(
VectorStoreFileRequestUtils.get_list_query_params(local_vars)
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"vector_store_id": vector_store_id, **list_query},
litellm_params={
@ -387,7 +389,8 @@ def retrieve(
f"Vector store file retrieve is not supported for {custom_llm_provider}"
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={
"vector_store_id": vector_store_id,
@ -498,7 +501,8 @@ def retrieve_content(
f"Vector store file content retrieve is not supported for {custom_llm_provider}"
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={
"vector_store_id": vector_store_id,
@ -619,7 +623,8 @@ def update(
)
update_request["attributes"] = attributes
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={
"vector_store_id": vector_store_id,
@ -733,7 +738,8 @@ def delete(
f"Vector store file delete is not supported for {custom_llm_provider}"
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={
"vector_store_id": vector_store_id,

View file

@ -3,6 +3,7 @@ LiteLLM SDK Functions for Creating and Searching Vector Stores
"""
import asyncio
import builtins
import contextvars
from functools import partial
from typing import Any, Coroutine, Dict, List, Optional, Union
@ -233,7 +234,8 @@ def create(
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={
"name": name,
@ -395,11 +397,11 @@ def search(
## MOCK RESPONSE LOGIC
if litellm_params.mock_response and isinstance(
litellm_params.mock_response, (str, list)
litellm_params.mock_response, (str, builtins.list)
):
mock_results = None
if isinstance(litellm_params.mock_response, list):
mock_results = litellm_params.mock_response
if isinstance(litellm_params.mock_response, builtins.list):
mock_results = litellm_params.mock_response # type: ignore[assignment]
return mock_vector_store_search_response(mock_results=mock_results)
# Default to OpenAI for vector stores
@ -440,7 +442,8 @@ def search(
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=api_type,
optional_params={
"vector_store_id": vector_store_id,
@ -585,7 +588,8 @@ def retrieve(
f"Vector store retrieve is not supported for {custom_llm_provider}"
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"vector_store_id": vector_store_id},
litellm_params={"litellm_call_id": litellm_call_id},
@ -732,7 +736,8 @@ def list(
f"Vector store list is not supported for {custom_llm_provider}"
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={
"after": after,
@ -895,7 +900,8 @@ def update(
)
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={
"vector_store_id": vector_store_id,
@ -1035,7 +1041,8 @@ def delete(
f"Vector store delete is not supported for {custom_llm_provider}"
)
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=None,
optional_params={"vector_store_id": vector_store_id},
litellm_params={"litellm_call_id": litellm_call_id},

View file

@ -120,17 +120,18 @@ async def avideo_generation(
def video_generation(
prompt: str,
model: Optional[str] = None,
input_reference: Optional[str] = None,
input_reference: Optional[FileTypes] = None,
seconds: Optional[str] = None,
size: Optional[str] = None,
user: Optional[str] = None,
timeout=600, # default to 10 minutes
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_version: Optional[str] = None,
custom_llm_provider=None,
timeout: int = 600,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
*,
avideo_generation: Literal[True],
**kwargs,
**kwargs: Any,
) -> Coroutine[Any, Any, VideoObject]:
...
@ -139,18 +140,18 @@ def video_generation(
def video_generation(
prompt: str,
model: Optional[str] = None,
input_reference: Optional[str] = None,
input_reference: Optional[FileTypes] = None,
seconds: Optional[str] = None,
size: Optional[str] = None,
user: Optional[str] = None,
timeout=600, # default to 10 minutes
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_version: Optional[str] = None,
custom_llm_provider=None,
timeout: int = 600,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
*,
avideo_generation: Literal[False] = False,
**kwargs,
**kwargs: Any,
) -> VideoObject:
...
@ -232,7 +233,8 @@ def video_generation( # noqa: PLR0915
)
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params=dict(video_generation_request_params),
@ -349,7 +351,8 @@ def video_content(
}
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
user=kwargs.get("user"),
optional_params=dict(video_content_request_params),
@ -529,14 +532,14 @@ async def avideo_remix(
def video_remix(
video_id: str,
prompt: str,
timeout=600, # default to 10 minutes
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_version: Optional[str] = None,
custom_llm_provider=None,
timeout: int = 600,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
*,
avideo_remix: Literal[True],
**kwargs,
**kwargs: Any,
) -> Coroutine[Any, Any, VideoObject]:
...
@ -545,14 +548,14 @@ def video_remix(
def video_remix(
video_id: str,
prompt: str,
timeout=600, # default to 10 minutes
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_version: Optional[str] = None,
custom_llm_provider=None,
timeout: int = 600,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
*,
avideo_remix: Literal[False] = False,
**kwargs,
**kwargs: Any,
) -> VideoObject:
...
@ -619,7 +622,8 @@ def video_remix( # noqa: PLR0915
}
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
user=kwargs.get("user"),
optional_params=dict(video_remix_request_params),
@ -745,14 +749,14 @@ def video_list(
after: Optional[str] = None,
limit: Optional[int] = None,
order: Optional[str] = None,
timeout=600, # default to 10 minutes
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_version: Optional[str] = None,
custom_llm_provider=None,
timeout: int = 600,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
*,
avideo_list: Literal[True],
**kwargs,
**kwargs: Any,
) -> Coroutine[Any, Any, List[VideoObject]]:
...
@ -762,14 +766,14 @@ def video_list(
after: Optional[str] = None,
limit: Optional[int] = None,
order: Optional[str] = None,
timeout=600, # default to 10 minutes
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_version: Optional[str] = None,
custom_llm_provider=None,
timeout: int = 600,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
*,
avideo_list: Literal[False] = False,
**kwargs,
**kwargs: Any,
) -> List[VideoObject]:
...
@ -835,7 +839,8 @@ def video_list( # noqa: PLR0915
}
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
user=kwargs.get("user"),
optional_params=dict(video_list_request_params),
@ -850,7 +855,7 @@ def video_list( # noqa: PLR0915
litellm_logging_obj.call_type = CallTypes.video_list.value
# Call the handler with _is_async flag instead of directly calling the async handler
return base_llm_http_handler.video_list_handler(
return base_llm_http_handler.video_list_handler( # type: ignore[return-value]
after=after,
limit=limit,
order=order,
@ -946,14 +951,14 @@ async def avideo_status(
@overload
def video_status(
video_id: str,
timeout=600, # default to 10 minutes
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_version: Optional[str] = None,
custom_llm_provider=None,
timeout: int = 600,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
*,
avideo_status: Literal[True],
**kwargs,
**kwargs: Any,
) -> Coroutine[Any, Any, VideoObject]:
...
@ -961,14 +966,14 @@ def video_status(
@overload
def video_status(
video_id: str,
timeout=600, # default to 10 minutes
api_key: Optional[str] = None,
api_base: Optional[str] = None,
api_version: Optional[str] = None,
custom_llm_provider=None,
timeout: int = 600,
custom_llm_provider: Optional[str] = None,
extra_headers: Optional[Dict[str, Any]] = None,
extra_query: Optional[Dict[str, Any]] = None,
extra_body: Optional[Dict[str, Any]] = None,
*,
avideo_status: Literal[False] = False,
**kwargs,
**kwargs: Any,
) -> VideoObject:
...
@ -1055,7 +1060,8 @@ def video_status( # noqa: PLR0915
}
# Pre Call logging
litellm_logging_obj.update_environment_variables(
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model="",
user=kwargs.get("user"),
optional_params=dict(video_status_request_params),

View file

@ -29,10 +29,26 @@ verbose_logger.setLevel(logging.DEBUG)
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import StandardLoggingPayload
import random
import socket
import httpx
from unittest.mock import patch, MagicMock
def _can_resolve_openai():
"""Check if api.openai.com is reachable (DNS resolves)."""
try:
socket.getaddrinfo("api.openai.com", 443, socket.AF_UNSPEC, socket.SOCK_STREAM)
return True
except socket.gaierror:
return False
skip_if_no_openai_network = pytest.mark.skipif(
not _can_resolve_openai(),
reason="Cannot resolve api.openai.com - skipping integration test due to DNS issues",
)
def load_vertex_ai_credentials():
# Define the path to the vertex_key.json file
print("loading vertex ai credentials")
@ -78,6 +94,7 @@ def load_vertex_ai_credentials():
@pytest.mark.parametrize("provider", ["openai"]) # , "azure"
@pytest.mark.asyncio
@skip_if_no_openai_network
async def test_create_batch(provider):
"""
1. Create File for Batch completion
@ -252,6 +269,7 @@ def cleanup_azure_ft_models():
@pytest.mark.parametrize("provider", ["openai"])
@pytest.mark.asyncio()
@pytest.mark.flaky(retries=3, delay=1)
@skip_if_no_openai_network
async def test_async_create_batch(provider):
"""
1. Create File for Batch completion
@ -464,9 +482,24 @@ mock_vertex_list_response = {
@pytest.mark.asyncio
async def test_avertex_batch_prediction(monkeypatch):
monkeypatch.setenv("GCS_BUCKET_NAME", "litellm-local")
monkeypatch.setenv("VERTEXAI_PROJECT", "mock-project")
monkeypatch.setenv("VERTEXAI_LOCATION", "us-central1")
# Mock Google auth so the test doesn't need real credentials
mock_creds = MagicMock()
mock_creds.token = "mock-token"
mock_creds.valid = True
mock_creds.expiry = None
monkeypatch.setattr(
"google.auth.default",
lambda *args, **kwargs: (mock_creds, "mock-project"),
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
client = AsyncHTTPHandler()
# Configure mock response object
mock_response = MagicMock()
mock_response.raise_for_status.return_value = None
async def mock_side_effect(*args, **kwargs):
print("args", args, "kwargs", kwargs)
@ -478,21 +511,10 @@ async def test_avertex_batch_prediction(monkeypatch):
mock_response.status_code = 200
return mock_response
with patch.object(
client, "post", side_effect=mock_side_effect
) as mock_post, patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post"
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
side_effect=mock_side_effect,
) as mock_global_post:
# Configure mock responses
mock_response = MagicMock()
mock_response.raise_for_status.return_value = None
# Set up different responses for different API calls
mock_post.side_effect = mock_side_effect
mock_global_post.side_effect = mock_side_effect
# load_vertex_ai_credentials()
litellm.set_verbose = True
litellm._turn_on_debug()
file_name = "vertex_batch_completions.jsonl"
@ -504,7 +526,6 @@ async def test_avertex_batch_prediction(monkeypatch):
file=open(file_path, "rb"),
purpose="batch",
custom_llm_provider="vertex_ai",
client=client
)
print("Response from creating file=", file_obj)
@ -623,6 +644,7 @@ async def test_vertex_async_create_batch_logs_error_body_on_http_error():
@pytest.mark.asyncio
@skip_if_no_openai_network
async def test_delete_batch_output_file():
"""
Test that deleting a batch output file works correctly.

View file

@ -1,4 +1,9 @@
# conftest.py
#
# xdist-compatible test isolation for guardrails tests.
# Pattern matches tests/test_litellm/conftest.py:
# - Function-scoped fixture saves/restores litellm globals (no reload)
# - Module-scoped fixture reloads only in single-process mode
import importlib
import os
@ -10,58 +15,85 @@ sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
import asyncio
@pytest.fixture(scope="session")
def event_loop():
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
yield loop
loop.close()
@pytest.fixture(scope="function", autouse=True)
def setup_and_teardown():
def isolate_litellm_state():
"""
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
Per-function isolation fixture.
Saves and restores litellm callback/global state so tests don't leak
side effects. Works safely under pytest-xdist parallel execution.
"""
curr_dir = os.getcwd() # Get the current working directory
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the project directory to the system path
# Save original callback state
original_state = {}
for attr in (
"callbacks",
"success_callback",
"failure_callback",
"_async_success_callback",
"_async_failure_callback",
):
if hasattr(litellm, attr):
val = getattr(litellm, attr)
original_state[attr] = val.copy() if val else []
import litellm
from litellm import Router
import asyncio
# Save other globals that tests commonly mutate
for attr in ("set_verbose", "cache", "num_retries"):
if hasattr(litellm, attr):
original_state[attr] = getattr(litellm, attr)
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
# flush all logs
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
# Flush cache before test
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
# Clear callbacks before test
for attr in (
"success_callback",
"failure_callback",
"_async_success_callback",
"_async_failure_callback",
):
if hasattr(litellm, attr):
setattr(litellm, attr, [])
importlib.reload(litellm)
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
import litellm.proxy.proxy_server
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
import asyncio
loop = asyncio.get_event_loop_policy().new_event_loop()
asyncio.set_event_loop(loop)
print(litellm)
# from litellm import Router, completion, aembedding, acompletion, embedding
yield
# Teardown code (executes after the yield point)
loop.close() # Close the loop created earlier
asyncio.set_event_loop(None) # Remove the reference to the loop
# Restore all saved state
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
for attr, original_value in original_state.items():
if hasattr(litellm, attr):
setattr(litellm, attr, original_value)
@pytest.fixture(scope="module", autouse=True)
def setup_and_teardown():
"""
Module-scoped setup. Reloads litellm only in single-process mode
(skipped under xdist to avoid cross-worker interference).
"""
sys.path.insert(0, os.path.abspath("../.."))
import litellm
worker_id = os.environ.get("PYTEST_XDIST_WORKER", None)
if worker_id is None:
importlib.reload(litellm)
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
import litellm.proxy.proxy_server
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
yield
def pytest_collection_modifyitems(config, items):

View file

@ -24,6 +24,7 @@ print("Python Path:", sys.path)
print("Current Working Directory:", os.getcwd())
import functools
from typing import Optional
from unittest.mock import MagicMock, patch
@ -34,6 +35,19 @@ from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2
from litellm.types.secret_managers.main import KeyManagementSettings
def skip_on_throttling(func):
"""Skip async test on AWS ThrottlingException instead of failing."""
@functools.wraps(func)
async def wrapper(*args, **kwargs):
try:
return await func(*args, **kwargs)
except Exception as e:
if "ThrottlingException" in str(e):
pytest.skip(f"AWS throttling: {e}")
raise
return wrapper
def check_aws_credentials():
"""Helper function to check if AWS credentials are set"""
required_vars = ["AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION_NAME"]
@ -43,6 +57,7 @@ def check_aws_credentials():
@pytest.mark.asyncio
@skip_on_throttling
async def test_write_and_read_simple_secret():
"""Test writing and reading a simple string secret"""
check_aws_credentials()
@ -84,6 +99,7 @@ async def test_write_and_read_simple_secret():
@pytest.mark.asyncio
@skip_on_throttling
async def test_write_and_read_json_secret():
"""Test writing and reading a JSON structured secret"""
check_aws_credentials()
@ -128,6 +144,7 @@ async def test_write_and_read_json_secret():
@pytest.mark.asyncio
@skip_on_throttling
async def test_read_nonexistent_secret():
"""Test reading a secret that doesn't exist"""
check_aws_credentials()
@ -141,6 +158,7 @@ async def test_read_nonexistent_secret():
@pytest.mark.asyncio
@skip_on_throttling
async def test_primary_secret_functionality():
"""Test storing and retrieving secrets from a primary secret"""
check_aws_credentials()
@ -196,6 +214,7 @@ async def test_primary_secret_functionality():
assert delete_response is not None
@pytest.mark.asyncio
@skip_on_throttling
async def test_write_secret_with_description_and_tags():
"""Test writing a secret with description and tags"""
check_aws_credentials()
@ -402,6 +421,7 @@ def test_load_aws_secret_manager_with_settings():
@pytest.mark.asyncio
@skip_on_throttling
async def test_end_to_end_iam_role_secret_write():
"""
Test writing a secret using IAM role assumption (integration test)

View file

@ -2,8 +2,10 @@ import json
import os
import sys
import time
from contextlib import asynccontextmanager, contextmanager
from datetime import datetime
from unittest.mock import AsyncMock, patch, MagicMock
import httpx
import pytest
import asyncio
@ -13,6 +15,63 @@ sys.path.insert(
import litellm
# Fake Vertex AI Gemini response for mocking
FAKE_VERTEX_GEMINI_RESPONSE = {
"candidates": [
{
"content": {
"parts": [{"text": "Hello! How can I help you today?"}],
"role": "model",
},
"finishReason": "STOP",
}
],
"usageMetadata": {
"promptTokenCount": 5,
"candidatesTokenCount": 8,
"totalTokenCount": 13,
},
}
def _make_fake_httpx_response(url: str) -> httpx.Response:
"""Create a fake httpx.Response that looks like a Vertex AI Gemini response."""
response = httpx.Response(
status_code=200,
json=FAKE_VERTEX_GEMINI_RESPONSE,
request=httpx.Request("POST", url),
)
return response
@asynccontextmanager
async def _vertex_ai_mocks():
"""Context manager that mocks Vertex AI auth and HTTP calls.
Mocks at the httpx.AsyncClient.send level so that the
@track_llm_api_timing decorator on AsyncHTTPHandler.post still runs,
preserving the overhead measurement.
"""
fake_response = _make_fake_httpx_response(
"https://fake-vertex-endpoint/v1/models/gemini-1.5-flash:generateContent"
)
async def fake_send(self, request, **kwargs):
await asyncio.sleep(0.2) # simulate ~200ms network latency
return fake_response
with patch(
"litellm.llms.vertex_ai.vertex_llm_base.VertexBase._ensure_access_token_async",
new_callable=AsyncMock,
return_value=("Bearer fake-token", "fake-project"),
), patch.object(
httpx.AsyncClient,
"send",
new=fake_send,
):
yield
@pytest.mark.asyncio
@pytest.mark.parametrize(
"model",
@ -39,16 +98,19 @@ async def test_litellm_overhead_non_streaming(model):
# Specific cases for models
#########################################################
if model == "vertex_ai/gemini-1.5-flash":
kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/v1/projects/pathrise-convert-1606954137718/locations/us-central1/publishers/google/models/gemini-1.0-pro-vision-001"
# warmup call for auth validation on vertex_ai models
await litellm.acompletion(**kwargs)
kwargs["vertex_project"] = "fake-project"
kwargs["vertex_location"] = "us-central1"
if model == "openai/self_hosted":
kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/"
async def _run():
return await litellm.acompletion(**kwargs)
response = await litellm.acompletion(
**kwargs
)
if model == "vertex_ai/gemini-1.5-flash":
async with _vertex_ai_mocks():
response = await _run()
else:
response = await _run()
#########################################################
# End of specific cases for models
#########################################################

View file

@ -1,9 +1,14 @@
# conftest.py
#
# xdist-compatible test isolation for llm_translation tests.
# Mirrors the pattern in tests/local_testing/conftest.py:
# - Function-scoped fixture resets litellm globals to true defaults
# - Module-scoped reload only in single-process mode
import importlib
import os
import sys
import asyncio
import pytest
sys.path.insert(
@ -13,6 +18,24 @@ import litellm
import asyncio
# ---------------------------------------------------------------------------
# Capture TRUE defaults at conftest import time (before test modules pollute).
# ---------------------------------------------------------------------------
_SCALAR_DEFAULTS = {
"num_retries": getattr(litellm, "num_retries", None),
"set_verbose": getattr(litellm, "set_verbose", False),
"cache": getattr(litellm, "cache", None),
"allowed_fails": getattr(litellm, "allowed_fails", 3),
"disable_aiohttp_transport": getattr(litellm, "disable_aiohttp_transport", False),
"force_ipv4": getattr(litellm, "force_ipv4", False),
"drop_params": getattr(litellm, "drop_params", None),
"modify_params": getattr(litellm, "modify_params", False),
"api_base": getattr(litellm, "api_base", None),
"api_key": getattr(litellm, "api_key", None),
"cohere_key": getattr(litellm, "cohere_key", None),
}
@pytest.fixture(scope="session")
def event_loop():
try:
@ -29,20 +52,39 @@ def setup_and_teardown(event_loop): # Add event_loop as a dependency
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm import Router
# ---- Save current state (for teardown restore) ----
original_state = {}
for attr in (
"callbacks",
"success_callback",
"failure_callback",
"_async_success_callback",
"_async_failure_callback",
):
if hasattr(litellm, attr):
val = getattr(litellm, attr)
original_state[attr] = val.copy() if val else []
for attr in _SCALAR_DEFAULTS:
if hasattr(litellm, attr):
original_state[attr] = getattr(litellm, attr)
# ---- Reset to true defaults before the test ----
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
# flush all logs
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
importlib.reload(litellm)
# Set the event loop from the fixture
asyncio.set_event_loop(event_loop)
print(litellm)
yield
# ---- Teardown ----
for attr, original_value in original_state.items():
if hasattr(litellm, attr):
setattr(litellm, attr, original_value)
# Clean up any pending tasks
pending = asyncio.all_tasks(event_loop)
for task in pending:

View file

@ -838,7 +838,7 @@ async def test_gemini_image_generation_async():
IMAGE_URL = response.choices[0].message.images[0]["image_url"]
print("IMAGE_URL: ", IMAGE_URL)
assert CONTENT is not None, "CONTENT is not None"
# content may be None when the model returns only an image with no text
assert IMAGE_URL is not None, "IMAGE_URL is not None"
assert IMAGE_URL["url"] is not None, "IMAGE_URL['url'] is not None"
assert IMAGE_URL["url"].startswith("data:image/png;base64,")

View file

@ -1,4 +1,15 @@
# conftest.py
#
# xdist-compatible test isolation for local_testing tests.
# Pattern matches tests/test_litellm/conftest.py:
# - Function-scoped fixture saves/restores litellm globals (no reload)
# - Module-scoped fixture reloads only in single-process mode
#
# IMPORTANT: True defaults are captured at conftest import time (before any
# test module can pollute them via module-level assignments like
# `litellm.num_retries = 3`). The function-scoped fixture resets globals to
# these true defaults before every test, preventing cross-test contamination
# under xdist where module reload is skipped.
import importlib
import os
@ -11,60 +22,126 @@ sys.path.insert(
) # Adds the parent directory to the system path
import litellm
import asyncio
@pytest.fixture(scope="session")
def event_loop():
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
yield loop
loop.close()
# ---------------------------------------------------------------------------
# Capture TRUE defaults at conftest import time. This runs before any test
# module's top-level code (e.g. `litellm.num_retries = 3`) executes, so
# the values here are guaranteed to be the real package defaults.
# ---------------------------------------------------------------------------
_SCALAR_DEFAULTS = {
"num_retries": getattr(litellm, "num_retries", None),
"num_retries_per_request": getattr(litellm, "num_retries_per_request", None),
"request_timeout": getattr(litellm, "request_timeout", None),
"set_verbose": getattr(litellm, "set_verbose", False),
"cache": getattr(litellm, "cache", None),
"allowed_fails": getattr(litellm, "allowed_fails", 3),
"default_fallbacks": getattr(litellm, "default_fallbacks", None),
"enable_azure_ad_token_refresh": getattr(litellm, "enable_azure_ad_token_refresh", None),
"tag_budget_config": getattr(litellm, "tag_budget_config", None),
"model_cost": getattr(litellm, "model_cost", None),
"token_counter": getattr(litellm, "token_counter", None),
"disable_aiohttp_transport": getattr(litellm, "disable_aiohttp_transport", False),
"force_ipv4": getattr(litellm, "force_ipv4", False),
"drop_params": getattr(litellm, "drop_params", None),
"modify_params": getattr(litellm, "modify_params", False),
"api_base": getattr(litellm, "api_base", None),
"api_key": getattr(litellm, "api_key", None),
}
@pytest.fixture(scope="function", autouse=True)
def setup_and_teardown():
def isolate_litellm_state():
"""
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
Per-function isolation fixture.
Resets litellm globals to their true defaults before each test and
restores them afterward, so tests don't leak side effects.
Works safely under pytest-xdist parallel execution.
"""
curr_dir = os.getcwd() # Get the current working directory
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the project directory to the system path
# ---- Save current callback state (for teardown restore) ----
original_state = {}
for attr in (
"callbacks",
"success_callback",
"failure_callback",
"_async_success_callback",
"_async_failure_callback",
):
if hasattr(litellm, attr):
val = getattr(litellm, attr)
original_state[attr] = val.copy() if val else []
import litellm
from litellm import Router
import asyncio
# Save list-type globals
for attr in ("pre_call_rules", "post_call_rules"):
if hasattr(litellm, attr):
val = getattr(litellm, attr)
original_state[attr] = val.copy() if val else []
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
# flush all logs
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
# Save scalar globals
for attr in _SCALAR_DEFAULTS:
if hasattr(litellm, attr):
original_state[attr] = getattr(litellm, attr)
# ---- Reset to true defaults before the test ----
# Flush HTTP client cache
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
importlib.reload(litellm)
# Clear callbacks and rules
for attr in (
"callbacks",
"success_callback",
"failure_callback",
"_async_success_callback",
"_async_failure_callback",
"pre_call_rules",
"post_call_rules",
):
if hasattr(litellm, attr):
setattr(litellm, attr, [])
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
import litellm.proxy.proxy_server
# Reset scalar globals to true defaults (prevents contamination from
# module-level code like `litellm.num_retries = 3` in test files)
for attr, default_val in _SCALAR_DEFAULTS.items():
if hasattr(litellm, attr):
setattr(litellm, attr, default_val)
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
import asyncio
loop = asyncio.get_event_loop_policy().new_event_loop()
asyncio.set_event_loop(loop)
print(litellm)
# from litellm import Router, completion, aembedding, acompletion, embedding
yield
# Teardown code (executes after the yield point)
loop.close() # Close the loop created earlier
asyncio.set_event_loop(None) # Remove the reference to the loop
# ---- Teardown: restore saved state ----
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
for attr, original_value in original_state.items():
if hasattr(litellm, attr):
setattr(litellm, attr, original_value)
@pytest.fixture(scope="module", autouse=True)
def setup_and_teardown():
"""
Module-scoped setup. Reloads litellm only in single-process mode
(skipped under xdist to avoid cross-worker interference).
"""
sys.path.insert(0, os.path.abspath("../.."))
import litellm
worker_id = os.environ.get("PYTEST_XDIST_WORKER", None)
if worker_id is None:
importlib.reload(litellm)
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
import litellm.proxy.proxy_server
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
yield
def pytest_collection_modifyitems(config, items):

View file

@ -22,33 +22,37 @@ from litellm import Router
load_dotenv()
model_list = [
{ # list of model deployments
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
{
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
},
"tpm": 1000000,
"rpm": 9000,
},
]
kwargs = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hey, how's it going?"}],
}
def _make_model_list():
return [
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
},
"tpm": 1000000,
"rpm": 9000,
},
]
def _make_kwargs():
return {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "Hey, how's it going?"}],
}
@pytest.mark.flaky(retries=3, delay=1)
@ -58,8 +62,9 @@ def test_multiple_deployments_sync():
litellm.set_verbose = False
results = []
kwargs = _make_kwargs()
router = Router(
model_list=model_list,
model_list=_make_model_list(),
redis_host=os.getenv("REDIS_HOST"),
redis_password=os.getenv("REDIS_PASSWORD"),
redis_port=int(os.getenv("REDIS_PORT")), # type: ignore
@ -85,9 +90,10 @@ def test_multiple_deployments_parallel():
litellm.set_verbose = False # Corrected the syntax for setting verbose to False
results = []
futures = {}
kwargs = _make_kwargs()
start_time = time.time()
router = Router(
model_list=model_list,
model_list=_make_model_list(),
redis_host=os.getenv("REDIS_HOST"),
redis_password=os.getenv("REDIS_PASSWORD"),
redis_port=int(os.getenv("REDIS_PORT")), # type: ignore

View file

@ -3691,6 +3691,8 @@ def test_vertex_ai_llama_tool_calling():
response = completion(**args)
except litellm.RateLimitError:
pytest.skip("Rate limit error")
except litellm.NotFoundError:
pytest.skip("Model not found / resource unavailable")
print(response)
assert response.choices[0].message.tool_calls is not None

View file

@ -147,7 +147,7 @@ def test_caching_dynamic_args(): # test in memory cache
port=_redis_port_env,
password=_redis_password_env,
)
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
print(f"response1: {response1}")
print(f"response2: {response2}")
@ -173,7 +173,7 @@ def test_caching_v2(): # test in memory cache
try:
litellm.set_verbose = True
litellm.cache = Cache()
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
print(f"response1: {response1}")
print(f"response2: {response2}")
@ -200,9 +200,9 @@ def test_caching_with_ttl():
litellm.set_verbose = True
litellm.cache = Cache()
response1 = completion(
model="gpt-3.5-turbo", messages=messages, caching=True, ttl=0
model="gpt-3.5-turbo", messages=messages, caching=True, ttl=0, mock_response="Hello world from cache test 1"
)
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test 2")
print(f"response1: {response1}")
print(f"response2: {response2}")
litellm.cache = None # disable cache
@ -221,8 +221,8 @@ def test_caching_with_default_ttl():
try:
litellm.set_verbose = True
litellm.cache = Cache(ttl=0)
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
print(f"response1: {response1}")
print(f"response2: {response2}")
litellm.cache = None # disable cache
@ -247,10 +247,10 @@ async def test_caching_with_cache_controls(sync_flag):
if sync_flag:
## TTL = 0
response1 = completion(
model="gpt-3.5-turbo", messages=messages, cache={"ttl": 0}
model="gpt-3.5-turbo", messages=messages, cache={"ttl": 0}, mock_response="Hello world"
)
response2 = completion(
model="gpt-3.5-turbo", messages=messages, cache={"s-maxage": 10}
model="gpt-3.5-turbo", messages=messages, cache={"s-maxage": 10}, mock_response="Hello world"
)
assert response2["id"] != response1["id"]
@ -315,7 +315,6 @@ async def test_caching_with_cache_controls(sync_flag):
# test_caching_with_cache_controls()
@pytest.mark.flaky(retries=3, delay=1)
def test_caching_with_models_v2():
messages = [
{"role": "user", "content": "who is ishaan CTO of litellm from litellm 2023"}
@ -323,9 +322,9 @@ def test_caching_with_models_v2():
litellm.cache = Cache()
print("test2 for caching")
litellm.set_verbose = True
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
response1 = completion(model="gpt-3.5-turbo", messages=messages, caching=True, mock_response="Hello world from cache test")
response2 = completion(model="gpt-3.5-turbo", messages=messages, caching=True)
response3 = completion(model="gpt-4.1-nano", messages=messages, caching=True)
response3 = completion(model="gpt-4.1-nano", messages=messages, caching=True, mock_response="Different model response")
print(f"response1: {response1}")
print(f"response2: {response2}")
print(f"response3: {response3}")
@ -424,7 +423,7 @@ def test_embedding_caching():
text_to_embed = [embedding_large_text]
start_time = time.time()
embedding1 = embedding(
model="text-embedding-ada-002", input=text_to_embed, caching=True
model="text-embedding-ada-002", input=text_to_embed, caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
)
end_time = time.time()
print(f"Embedding 1 response time: {end_time - start_time} seconds")
@ -460,12 +459,12 @@ async def test_embedding_caching_individual_items_and_then_list():
"world",
]
embedding1 = await aembedding(
model="text-embedding-ada-002", input=text_to_embed[0], caching=True
model="text-embedding-ada-002", input=text_to_embed[0], caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
)
initial_prompt_tokens = embedding1.usage.prompt_tokens
await asyncio.sleep(1)
embedding2 = await aembedding(
model="text-embedding-ada-002", input=text_to_embed[1], caching=True
model="text-embedding-ada-002", input=text_to_embed[1], caching=True, mock_response="0.6,0.7,0.8,0.9,1.0"
)
await asyncio.sleep(1)
embedding3 = await aembedding(
@ -481,7 +480,7 @@ async def test_embedding_caching_individual_items_and_then_list():
additional_text = "this is a new text"
text_to_embed.append(additional_text)
embedding4 = await aembedding(
model="text-embedding-ada-002", input=text_to_embed, caching=True
model="text-embedding-ada-002", input=text_to_embed, caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
)
assert embedding4.usage.prompt_tokens > embedding3.usage.prompt_tokens
@ -491,7 +490,7 @@ async def test_embedding_caching_individual_items():
litellm.cache = Cache()
text_to_embed = "hello"
embedding1 = await aembedding(
model="text-embedding-ada-002", input=text_to_embed, caching=True
model="text-embedding-ada-002", input=text_to_embed, caching=True, mock_response="0.1,0.2,0.3,0.4,0.5"
)
await asyncio.sleep(1)
@ -533,6 +532,7 @@ def test_embedding_caching_azure():
api_base=api_base,
api_version=api_version,
caching=True,
mock_response="0.1,0.2,0.3,0.4,0.5",
)
end_time = time.time()
print(f"Embedding 1 response time: {end_time - start_time} seconds")
@ -762,6 +762,7 @@ async def test_redis_cache_basic():
response1 = completion(
model="gpt-3.5-turbo",
messages=messages,
mock_response="Hello world from cache test",
)
cache_key = litellm.cache.get_cache_key(
@ -803,6 +804,7 @@ async def test_redis_batch_cache_write():
response1 = await litellm.acompletion(
model="gpt-3.5-turbo",
messages=messages,
mock_response="Hello world from cache test",
)
response2 = await litellm.acompletion(
@ -843,14 +845,15 @@ def test_redis_cache_completion():
messages=messages,
caching=True,
max_tokens=20,
mock_response="Hello world from cache test",
)
response2 = completion(
model="gpt-3.5-turbo", messages=messages, caching=True, max_tokens=20
)
response3 = completion(
model="gpt-3.5-turbo", messages=messages, caching=True, temperature=0.5
model="gpt-3.5-turbo", messages=messages, caching=True, temperature=0.5, mock_response="Different params response"
)
response4 = completion(model="gpt-4o-mini", messages=messages, caching=True)
response4 = completion(model="gpt-4o-mini", messages=messages, caching=True, mock_response="Different model response")
print("\nresponse 1", response1)
print("\nresponse 2", response2)
@ -928,12 +931,13 @@ def test_redis_cache_completion_stream():
max_tokens=40,
temperature=0.2,
stream=True,
mock_response="In the stillness of numbers, the world turns quietly.",
)
response_1_id = ""
for chunk in response1:
print(chunk)
response_1_id = chunk.id
time.sleep(0.5)
time.sleep(1)
response2 = completion(
model="gpt-3.5-turbo",
messages=messages,
@ -1072,12 +1076,13 @@ async def test_redis_cache_acompletion_stream():
max_tokens=40,
temperature=1,
stream=True,
mock_response="In the stillness of numbers, the world turns quietly.",
)
async for chunk in response1:
response_1_content += chunk.choices[0].delta.content or ""
print(response_1_content)
await asyncio.sleep(0.5)
await asyncio.sleep(1)
print("\n\n Response 1 content: ", response_1_content, "\n\n")
response2 = await litellm.acompletion(
@ -1122,7 +1127,7 @@ async def test_redis_cache_atext_completion():
print("test for caching, atext_completion")
response1 = await litellm.atext_completion(
model="gpt-3.5-turbo-instruct", prompt=prompt, max_tokens=40, temperature=1
model="gpt-3.5-turbo-instruct", prompt=prompt, max_tokens=40, temperature=1, mock_response="Hello world from cache test"
)
await asyncio.sleep(0.5)
@ -1164,6 +1169,7 @@ async def test_redis_cache_acompletion_stream_bedrock():
max_tokens=40,
temperature=1,
stream=True,
mock_response="In the stillness of numbers, the world turns quietly.",
)
async for chunk in response1:
print(chunk)
@ -1231,6 +1237,7 @@ async def test_s3_cache_stream_azure(sync_mode):
max_tokens=40,
temperature=1,
stream=True,
mock_response="In the stillness of numbers, the world turns quietly.",
)
for chunk in response1:
print(chunk)
@ -1244,6 +1251,7 @@ async def test_s3_cache_stream_azure(sync_mode):
max_tokens=40,
temperature=1,
stream=True,
mock_response="In the stillness of numbers, the world turns quietly.",
)
async for chunk in response1:
print(chunk)
@ -1406,6 +1414,7 @@ def test_custom_redis_cache_with_key():
temperature=1,
caching=True,
num_retries=3,
mock_response="Hello world from cache test",
)
response2 = completion(
model="gpt-3.5-turbo",
@ -1420,6 +1429,7 @@ def test_custom_redis_cache_with_key():
temperature=1,
caching=False,
num_retries=3,
mock_response="Different uncached response",
)
print(f"response1: {response1}")
@ -1448,21 +1458,15 @@ def test_cache_override():
# test embedding
response1 = embedding(
model="text-embedding-ada-002", input=["hello who are you"], caching=False
model="text-embedding-ada-002", input=["hello who are you"], caching=False, mock_response="0.1,0.2,0.3,0.4,0.5"
)
start_time = time.time()
response2 = embedding(
model="text-embedding-ada-002", input=["hello who are you"], caching=False
model="text-embedding-ada-002", input=["hello who are you"], caching=False, mock_response="0.6,0.7,0.8,0.9,1.0"
)
end_time = time.time()
print(f"Embedding 2 response time: {end_time - start_time} seconds")
assert (
end_time - start_time > 0.05
) # ensure 2nd response comes in over 0.05s. This should not be cached.
# When caching=False, responses should have different IDs
assert response1.data[0].embedding != response2.data[0].embedding
# test_cache_override()
@ -1494,6 +1498,7 @@ async def test_cache_control_overrides():
}
],
caching=True,
mock_response="Hello world from cache test",
)
print(response1)
@ -1510,6 +1515,7 @@ async def test_cache_control_overrides():
],
caching=True,
cache={"no-cache": True},
mock_response="Hello world from cache test",
)
print(response2)
@ -1542,6 +1548,7 @@ def test_sync_cache_control_overrides():
}
],
caching=True,
mock_response="Hello world from cache test",
)
print(response1)
@ -1558,6 +1565,7 @@ def test_sync_cache_control_overrides():
],
caching=True,
cache={"no-cache": True},
mock_response="Hello world from cache test",
)
print(response2)
@ -1770,6 +1778,7 @@ def test_redis_semantic_cache_completion():
}
],
max_tokens=20,
mock_response="Summer sun shines bright and warm.",
)
print(f"response1: {response1}")
@ -1815,6 +1824,7 @@ async def test_redis_semantic_cache_acompletion():
}
],
max_tokens=5,
mock_response="Summer sun shines bright and warm.",
)
print(f"response1: {response1}")
@ -1850,11 +1860,14 @@ def test_caching_redis_simple(caplog, capsys):
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": f"Hello, how are you? Wink {uuid_str}"}],
stream=True,
mock_response="Hello world from cache test",
)
for m in x:
print(m)
print(time.time() - s)
time.sleep(1) # wait for cache write to propagate
s2 = time.time()
x = completion(
model="gpt-3.5-turbo",
@ -2634,7 +2647,6 @@ def test_redis_caching_multiple_namespaces():
), f"Expected different response ID for no namespace vs namespaced. Got {response_1.id} and {response_4.id}"
@pytest.mark.flaky(retries=3, delay=1)
def test_caching_with_reasoning_content():
"""
Test that reasoning content is cached
@ -2650,6 +2662,7 @@ def test_caching_with_reasoning_content():
model="anthropic/claude-sonnet-4-5-20250929",
messages=messages,
thinking={"type": "enabled", "budget_tokens": 1024},
mock_response="LiteLLM is a unified API interface for LLMs.",
)
response_2 = completion(
@ -2660,7 +2673,6 @@ def test_caching_with_reasoning_content():
print(f"response 2: {response_2.model_dump_json(indent=4)}")
assert response_2._hidden_params["cache_hit"] == True
assert response_2.choices[0].message.reasoning_content is not None
except litellm.InternalServerError as e:
pytest.skip(f"Anthropic API returned InternalServerError - {str(e)}")

View file

@ -2937,7 +2937,7 @@ def test_completion_together_ai_mixtral():
def test_completion_together_ai_llama():
litellm.set_verbose = True
model_name = "together_ai/meta-llama/Meta-Llama-3.1-8B-Instruct-Turbo"
model_name = "together_ai/meta-llama/Llama-3.3-70B-Instruct-Turbo"
try:
messages = [
{"role": "user", "content": "What llm are you?"},

View file

@ -490,6 +490,7 @@ async def test_cost_tracking_with_caching():
assert response_cost_2 == 0
@pytest.mark.flaky(retries=3, delay=3)
def test_redis_cache_completion_stream():
# Important Test - This tests if we can add to streaming cache, when custom callbacks are set
import random
@ -522,6 +523,7 @@ def test_redis_cache_completion_stream():
temperature=0.2,
stream=True,
caching=True,
mock_response="In the stillness of numbers, the world turns quietly.",
)
response_1_content = ""
response_1_id = None
@ -531,7 +533,7 @@ def test_redis_cache_completion_stream():
response_1_content += chunk.choices[0].delta.content or ""
print(response_1_content)
time.sleep(5) # sleep for cache write to propagate
time.sleep(1) # sleep for cache write to propagate
response2 = completion(
model="gpt-3.5-turbo",
messages=messages,
@ -553,9 +555,9 @@ def test_redis_cache_completion_stream():
assert (
response_1_id == response_2_id
), f"Response 1 != Response 2. Same params, Response 1{response_1_content} != Response 2{response_2_content}"
# assert (
# response_1_content == response_2_content
# ), f"Response 1 != Response 2. Same params, Response 1{response_1_content} != Response 2{response_2_content}"
assert (
response_1_content == response_2_content
), f"Response 1 != Response 2. Same params, Response 1{response_1_content} != Response 2{response_2_content}"
litellm.success_callback = []
litellm._async_success_callback = []
litellm.cache = None

View file

@ -333,6 +333,10 @@ def test_parallel_function_call_anthropic_error_msg(
Reference Issue: https://github.com/BerriAI/litellm/issues/5747, https://github.com/BerriAI/litellm/issues/5388
"""
# Ensure modify_params is False so UnsupportedParamsError is raised
# (other tests in this file set it to True and don't reset it)
original_modify_params = litellm.modify_params
litellm.modify_params = False
try:
litellm.set_verbose = True
@ -363,6 +367,8 @@ def test_parallel_function_call_anthropic_error_msg(
print(e)
except Exception as e:
pytest.fail(f"Error occurred: {e}")
finally:
litellm.modify_params = original_modify_params
def test_parallel_function_call_stream():

View file

@ -825,8 +825,9 @@ def test_router_context_window_check_pre_call_check_out_group():
{
"model_name": "gpt-3.5-turbo-large", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "gpt-3.5-turbo-1106",
"model": "gpt-4.1-mini",
"api_key": os.getenv("OPENAI_API_KEY"),
"mock_response": "Alexander was a great conqueror.",
},
},
]
@ -2107,11 +2108,13 @@ async def test_aaarouter_dynamic_cooldown_message_retry_time(sync_mode):
User feedback: litellm says "No deployments available for selected model, Try again in 60 seconds"
but Azure says to retry in at most 9s
```
{"message": "litellm.proxy.proxy_server.embeddings(): Exception occured - No deployments available for selected model, Try again in 60 seconds. Passed model=text-embedding-ada-002. pre-call-checks=False, allowed_model_region=n/a, cooldown_list=[('b49cbc9314273db7181fe69b1b19993f04efb88f2c1819947c538bac08097e4c', {'Exception Received': 'litellm.RateLimitError: AzureException RateLimitError - Requests to the Embeddings_Create Operation under Azure OpenAI API version 2023-09-01-preview have exceeded call rate limit of your current OpenAI S0 pricing tier. Please retry after 9 seconds. Please go here: https://aka.ms/oai/quotaincrease if you would like to further increase the default rate limit.', 'Status Code': '429'})]", "level": "ERROR", "timestamp": "2024-08-22T03:25:36.900476"}
```
Tests that:
1. deployment_callback_on_failure reads retry-after header and uses it as cooldown time
2. Cooled-down deployments appear in get_cooldown_deployments
3. RouterRateLimitError is raised with the correct cooldown_time when all deployments are cooled down
"""
litellm.set_verbose = True
from httpx import Headers, Request, Response
cooldown_time = 30.0
router = Router(
model_list=[
@ -2128,104 +2131,75 @@ async def test_aaarouter_dynamic_cooldown_message_retry_time(sync_mode):
},
},
],
set_verbose=True,
debug_level="DEBUG",
cooldown_time=cooldown_time,
)
openai_client = openai.OpenAI(api_key="")
def _return_exception(*args, **kwargs):
from httpx import Headers, Request, Response
kwargs = {
"request": Request("POST", "https://www.google.com"),
"message": "Error code: 429 - Rate Limit Error!",
"body": {"detail": "Rate Limit Error!"},
"code": None,
"param": None,
"type": None,
"response": Response(
status_code=429,
headers=Headers(
{
"date": "Sat, 21 Sep 2024 22:56:53 GMT",
"server": "uvicorn",
"retry-after": f"{cooldown_time}",
"content-length": "30",
"content-type": "application/json",
}
),
request=Request("POST", "http://0.0.0.0:9000/chat/completions"),
# Build a 429 exception with retry-after header, matching what the OpenAI SDK raises
mock_exception = litellm.RateLimitError(
message="Rate Limit Error!",
llm_provider="openai",
model="text-embedding-ada-002",
response=Response(
status_code=429,
headers=Headers(
{
"retry-after": f"{cooldown_time}",
"content-type": "application/json",
}
),
"status_code": 429,
"request_id": None,
request=Request("POST", "https://api.openai.com/v1/embeddings"),
),
)
# Directly invoke the Router's failure callback for each deployment,
# simulating what the logging framework would do on failure.
# This tests the cooldown logic without depending on the global customLogger state.
model_ids = router.get_model_ids()
for model_id in model_ids:
deployment_kwargs = {
"exception": mock_exception,
"litellm_params": {
"model_info": {"id": model_id},
},
}
exception = Exception()
for k, v in kwargs.items():
setattr(exception, k, v)
raise exception
with patch.object(
openai_client.embeddings.with_raw_response,
"create",
side_effect=_return_exception,
):
for _ in range(1):
try:
if sync_mode:
router.embedding(
model="text-embedding-ada-002",
input="Hello world!",
client=openai_client,
)
else:
await router.aembedding(
model="text-embedding-ada-002",
input="Hello world!",
client=openai_client,
)
except litellm.RateLimitError:
pass
await asyncio.sleep(5)
if sync_mode:
cooldown_deployments = _get_cooldown_deployments(
litellm_router_instance=router, parent_otel_span=None
)
else:
cooldown_deployments = await _async_get_cooldown_deployments(
litellm_router_instance=router, parent_otel_span=None
)
print(
"Cooldown deployments - {}\n{}".format(
cooldown_deployments, len(cooldown_deployments)
)
router.deployment_callback_on_failure(
kwargs=deployment_kwargs,
completion_response=None,
start_time=None,
end_time=None,
)
assert len(cooldown_deployments) > 0
exception_raised = False
try:
if sync_mode:
router.embedding(
model="text-embedding-ada-002",
input="Hello world!",
client=openai_client,
)
else:
await router.aembedding(
model="text-embedding-ada-002",
input="Hello world!",
client=openai_client,
)
except litellm.types.router.RouterRateLimitError as e:
print(e)
exception_raised = True
assert e.cooldown_time == cooldown_time
if sync_mode:
cooldown_deployments = _get_cooldown_deployments(
litellm_router_instance=router, parent_otel_span=None
)
else:
cooldown_deployments = await _async_get_cooldown_deployments(
litellm_router_instance=router, parent_otel_span=None
)
assert exception_raised
assert len(cooldown_deployments) > 0
# Verify that a subsequent call raises RouterRateLimitError with correct cooldown_time
exception_raised = False
try:
if sync_mode:
router.embedding(
model="text-embedding-ada-002",
input="Hello world!",
mock_response=[0.1, 0.2, 0.3],
)
else:
await router.aembedding(
model="text-embedding-ada-002",
input="Hello world!",
mock_response=[0.1, 0.2, 0.3],
)
except litellm.types.router.RouterRateLimitError as e:
exception_raised = True
assert e.cooldown_time == cooldown_time
assert exception_raised
@pytest.mark.parametrize("sync_mode", [True, False])

View file

@ -376,7 +376,11 @@ async def test_single_deployment_cooldown_with_allowed_fails():
except litellm.Timeout:
pass
await asyncio.sleep(2)
# Poll until the mock is called (or timeout)
for _ in range(40):
if mock_client.call_count >= 1:
break
await asyncio.sleep(0.1)
mock_client.assert_called_once()
@ -426,7 +430,11 @@ async def test_single_deployment_cooldown_with_allowed_fail_policy():
except litellm.Timeout:
pass
await asyncio.sleep(2)
# Poll until the mock is called (or timeout)
for _ in range(40):
if mock_client.call_count >= 1:
break
await asyncio.sleep(0.1)
mock_client.assert_called_once()

View file

@ -1,16 +1,11 @@
import asyncio
import os
import random
import sys
import time
import traceback
from datetime import datetime, timedelta
from dotenv import load_dotenv
load_dotenv()
import copy
import os
sys.path.insert(
0, os.path.abspath("../..")
@ -21,36 +16,40 @@ import pytest
import litellm
from litellm import Router
router = Router(
model_list=[
{
"model_name": "azure-model",
"litellm_params": {
"model": "openai/very-special-endpoint",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/", # If you are Krrish, this is OpenAI Endpoint3 on our Railway endpoint :)
"api_key": "fake-key",
},
"model_info": {"id": "very-special-endpoint"},
},
{
"model_name": "azure-model",
"litellm_params": {
"model": "openai/fast-endpoint",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
"api_key": "fake-key",
},
"model_info": {"id": "fast-endpoint"},
},
],
set_verbose=True,
debug_level="DEBUG",
)
from litellm.router import CustomRoutingStrategyBase
def _create_router():
return Router(
model_list=[
{
"model_name": "azure-model",
"litellm_params": {
"model": "openai/very-special-endpoint",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
"api_key": "fake-key",
},
"model_info": {"id": "very-special-endpoint"},
},
{
"model_name": "azure-model",
"litellm_params": {
"model": "openai/fast-endpoint",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
"api_key": "fake-key",
},
"model_info": {"id": "fast-endpoint"},
},
],
set_verbose=True,
debug_level="DEBUG",
)
class CustomRoutingStrategy(CustomRoutingStrategyBase):
def __init__(self, router_instance: Router):
self._router = router_instance
async def async_get_available_deployment(
self,
model: str,
@ -59,22 +58,8 @@ class CustomRoutingStrategy(CustomRoutingStrategyBase):
specific_deployment: Optional[bool] = False,
request_kwargs: Optional[Dict] = None,
):
"""
Asynchronously retrieves the available deployment based on the given parameters.
Args:
model (str): The name of the model.
messages (Optional[List[Dict[str, str]]], optional): The list of messages for a given request. Defaults to None.
input (Optional[Union[str, List]], optional): The input for a given embedding request. Defaults to None.
specific_deployment (Optional[bool], optional): Whether to retrieve a specific deployment. Defaults to False.
request_kwargs (Optional[Dict], optional): Additional request keyword arguments. Defaults to None.
Returns:
Returns an element from litellm.router.model_list
"""
print("In CUSTOM async get available deployment")
model_list = router.model_list
model_list = self._router.model_list
print("router model list=", model_list)
for model in model_list:
if isinstance(model, dict):
@ -90,29 +75,15 @@ class CustomRoutingStrategy(CustomRoutingStrategyBase):
specific_deployment: Optional[bool] = False,
request_kwargs: Optional[Dict] = None,
):
"""
Synchronously retrieves the available deployment based on the given parameters.
Args:
model (str): The name of the model.
messages (Optional[List[Dict[str, str]]], optional): The list of messages for a given request. Defaults to None.
input (Optional[Union[str, List]], optional): The input for a given embedding request. Defaults to None.
specific_deployment (Optional[bool], optional): Whether to retrieve a specific deployment. Defaults to False.
request_kwargs (Optional[Dict], optional): Additional request keyword arguments. Defaults to None.
Returns:
Returns an element from litellm.router.model_list
"""
pass
@pytest.mark.asyncio
async def test_custom_routing():
import litellm
litellm.set_verbose = True
router.set_custom_routing_strategy(CustomRoutingStrategy())
router = _create_router()
router.set_custom_routing_strategy(CustomRoutingStrategy(router))
# make 4 requests
for _ in range(4):
@ -126,11 +97,6 @@ async def test_custom_routing():
await asyncio.sleep(1)
print("done sending initial requests to collect latency")
"""
Note: for debugging
- By this point: slow-endpoint should have timed out 3-4 times and should be heavily penalized :)
- The next 10 requests should all be routed to the fast-endpoint
"""
deployments = {}
# make 10 requests
@ -145,6 +111,3 @@ async def test_custom_routing():
else:
deployments[_picked_model_id] += 1
print("deployments", deployments)
# ALL the Requests should have been routed to the fast-endpoint
# assert deployments["fast-endpoint"] == 10

View file

@ -83,6 +83,7 @@ def test_async_fallbacks(caplog):
log
for log in captured_logs
if "Task exception was never retrieved" not in log
and "Task was destroyed but it is pending" not in log
and "get_available_deployment" not in log
and "in the Langfuse queue" not in log
]

View file

@ -14,14 +14,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import litellm
from litellm import Router
from litellm.integrations.custom_logger import CustomLogger
from typing import Any, Dict
import sys
import os
from typing import List, Dict
sys.path.insert(0, os.path.abspath("../.."))
from typing import Any, Dict, List
from litellm.router_utils.fallback_event_handlers import (
run_async_fallback,
@ -53,18 +46,47 @@ def create_test_router():
)
router: Router = create_test_router()
def create_test_router_2():
return Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
},
},
{
"model_name": "gpt-4",
"litellm_params": {
"model": "gpt-4",
"api_key": "very-fake-key",
},
},
{
"model_name": "fake-openai-endpoint-2",
"litellm_params": {
"model": "openai/fake-openai-endpoint-2",
"api_key": "working-key-since-this-is-fake-endpoint",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
},
},
],
)
@pytest.mark.parametrize(
"original_function",
[router._acompletion, router._atext_completion, router._aembedding],
"function_name",
["_acompletion", "_atext_completion", "_aembedding"],
)
@pytest.mark.asyncio
async def test_run_async_fallback(original_function):
async def test_run_async_fallback(function_name):
"""
Basic test - given a list of fallback models, run the original function with the fallback models
"""
router = create_test_router()
original_function = getattr(router, function_name)
litellm.set_verbose = True
fallback_model_group = ["gpt-4"]
original_model_group = "gpt-3.5-turbo"
@ -79,11 +101,11 @@ async def test_run_async_fallback(original_function):
"metadata": {"previous_models": ["gpt-3.5-turbo"]},
}
if original_function == router._aembedding:
if function_name == "_aembedding":
request_kwargs["input"] = "hello this is a test for run_async_fallback"
elif original_function == router._atext_completion:
elif function_name == "_atext_completion":
request_kwargs["prompt"] = "hello this is a test for run_async_fallback"
elif original_function == router._acompletion:
elif function_name == "_acompletion":
request_kwargs["messages"] = [{"role": "user", "content": "Hello, world!"}]
result = await run_async_fallback(
@ -100,11 +122,11 @@ async def test_run_async_fallback(original_function):
assert result is not None
if original_function == router._acompletion:
if function_name == "_acompletion":
assert isinstance(result, litellm.ModelResponse)
elif original_function == router._atext_completion:
elif function_name == "_atext_completion":
assert isinstance(result, litellm.TextCompletionResponse)
elif original_function == router._aembedding:
elif function_name == "_aembedding":
assert isinstance(result, litellm.EmbeddingResponse)
@ -198,14 +220,17 @@ async def test_log_failure_fallback_event():
@pytest.mark.asyncio
@pytest.mark.parametrize(
"original_function", [router._acompletion, router._atext_completion]
"function_name", ["_acompletion", "_atext_completion"]
)
async def test_failed_fallbacks_raise_most_recent_exception(original_function):
async def test_failed_fallbacks_raise_most_recent_exception(function_name):
"""
Tests that if all fallbacks fail, the most recent occuring exception is raised
meaning the exception from the last fallback model is raised
"""
router = create_test_router()
original_function = getattr(router, function_name)
fallback_model_group = ["gpt-4"]
original_model_group = "gpt-3.5-turbo"
original_exception = litellm.exceptions.InternalServerError(
@ -218,11 +243,11 @@ async def test_failed_fallbacks_raise_most_recent_exception(original_function):
"metadata": {"previous_models": ["gpt-3.5-turbo"]}
}
if original_function == router._aembedding:
if function_name == "_aembedding":
request_kwargs["input"] = "hello this is a test for run_async_fallback"
elif original_function == router._atext_completion:
elif function_name == "_atext_completion":
request_kwargs["prompt"] = "hello this is a test for run_async_fallback"
elif original_function == router._acompletion:
elif function_name == "_acompletion":
request_kwargs["messages"] = [{"role": "user", "content": "Hello, world!"}]
with pytest.raises(litellm.exceptions.RateLimitError):
@ -240,39 +265,11 @@ async def test_failed_fallbacks_raise_most_recent_exception(original_function):
)
router_2 = Router(
model_list=[
{
"model_name": "gpt-3.5-turbo",
"litellm_params": {
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
},
},
{
"model_name": "gpt-4",
"litellm_params": {
"model": "gpt-4",
"api_key": "very-fake-key",
},
},
{
"model_name": "fake-openai-endpoint-2",
"litellm_params": {
"model": "openai/fake-openai-endpoint-2",
"api_key": "working-key-since-this-is-fake-endpoint",
"api_base": "https://exampleopenaiendpoint-production.up.railway.app/",
},
},
],
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"original_function", [router_2._acompletion, router_2._atext_completion]
"function_name", ["_acompletion", "_atext_completion"]
)
async def test_multiple_fallbacks(original_function):
async def test_multiple_fallbacks(function_name):
"""
Tests that if multiple fallbacks passed:
- fallback 1 = bad configured deployment / failing endpoint
@ -281,6 +278,9 @@ async def test_multiple_fallbacks(original_function):
Assert that:
- a success response is received from the working endpoint (fallback 2)
"""
router_2 = create_test_router_2()
original_function = getattr(router_2, function_name)
fallback_model_group = ["gpt-4", "fake-openai-endpoint-2"]
original_model_group = "gpt-3.5-turbo"
original_exception = Exception("Simulated error")
@ -289,11 +289,11 @@ async def test_multiple_fallbacks(original_function):
"metadata": {"previous_models": ["gpt-3.5-turbo"]}
}
if original_function == router_2._aembedding:
if function_name == "_aembedding":
request_kwargs["input"] = "hello this is a test for run_async_fallback"
elif original_function == router_2._atext_completion:
elif function_name == "_atext_completion":
request_kwargs["prompt"] = "hello this is a test for run_async_fallback"
elif original_function == router_2._acompletion:
elif function_name == "_acompletion":
request_kwargs["messages"] = [{"role": "user", "content": "Hello, world!"}]
result = await run_async_fallback(

View file

@ -500,55 +500,25 @@ async def test_dynamic_fallbacks_async():
@pytest.mark.asyncio
async def test_async_fallbacks_streaming():
"""Test that router.acompletion with stream=True and mock_response works correctly."""
litellm.set_verbose = False
model_list = [
{ # list of model deployments
"model_name": "azure/gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
{
"model_name": "azure/gpt-3.5-turbo",
"litellm_params": {
"model": "azure/gpt-4.1-mini",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
{ # list of model deployments
"model_name": "azure/gpt-3.5-turbo-context-fallback", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/gpt-4.1-mini",
"api_key": os.getenv("AZURE_API_KEY"),
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_API_BASE"),
"api_key": "fake-key",
"api_version": "2024-01-01",
"api_base": "https://fake.openai.azure.com",
},
"tpm": 240000,
"rpm": 1800,
},
{
"model_name": "azure/gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "azure/chatgpt-functioncalling",
"api_key": "bad-key",
"api_version": os.getenv("AZURE_API_VERSION"),
"api_base": os.getenv("AZURE_API_BASE"),
},
"tpm": 240000,
"rpm": 1800,
},
{
"model_name": "gpt-3.5-turbo", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "gpt-3.5-turbo",
"api_key": os.getenv("OPENAI_API_KEY"),
},
"tpm": 1000000,
"rpm": 9000,
},
{
"model_name": "gpt-3.5-turbo-16k", # openai model name
"litellm_params": { # params for litellm completion/embedding call
"model": "gpt-3.5-turbo-16k",
"api_key": os.getenv("OPENAI_API_KEY"),
"model_name": "gpt-4o-mini",
"litellm_params": {
"model": "gpt-4o-mini",
"api_key": "fake-key",
},
"tpm": 1000000,
"rpm": 9000,
@ -557,24 +527,23 @@ async def test_async_fallbacks_streaming():
router = Router(
model_list=model_list,
fallbacks=[{"azure/gpt-3.5-turbo": ["gpt-3.5-turbo"]}],
context_window_fallbacks=[
{"azure/gpt-3.5-turbo-context-fallback": ["gpt-3.5-turbo-16k"]},
{"gpt-3.5-turbo": ["gpt-3.5-turbo-16k"]},
],
fallbacks=[{"azure/gpt-3.5-turbo": ["gpt-4o-mini"]}],
set_verbose=False,
)
customHandler = MyCustomHandler()
litellm.callbacks = [customHandler]
user_message = "Hello, how are you?"
messages = [{"content": user_message, "role": "user"}]
try:
response = await router.acompletion(**kwargs, stream=True)
print(f"customHandler.previous_models: {customHandler.previous_models}")
await asyncio.sleep(
0.05
) # allow a delay as success_callbacks are on a separate thread
assert customHandler.previous_models == 3 # 1 init call + 2 retries (fallback not counted as previous)
response = await router.acompletion(
model="azure/gpt-3.5-turbo",
messages=[{"role": "user", "content": user_message}],
stream=True,
mock_response="This is a mock streaming response",
)
chunks = []
async for chunk in response:
chunks.append(chunk)
assert len(chunks) > 0, "Expected at least one streaming chunk"
router.reset()
except litellm.Timeout as e:
pass
@ -840,8 +809,6 @@ def test_ausage_based_routing_fallbacks():
set_verbose=True,
debug_level="DEBUG",
routing_strategy="usage-based-routing-v2",
redis_host=os.environ["REDIS_HOST"],
redis_port=int(os.environ["REDIS_PORT"]),
num_retries=0,
)

View file

@ -134,6 +134,7 @@ async def test_completion_sagemaker_messages_api(sync_mode):
],
temperature=0.2,
max_tokens=80,
num_retries=0,
client=client,
)
except Exception as e:

View file

@ -94,8 +94,15 @@ def test_bedrock_timeout():
def test_hanging_request_azure():
"""
Test that a slow Azure request properly raises APITimeoutError via the Router.
Uses a mock to simulate a slow HTTP response so the timeout fires reliably,
rather than racing against real network latency.
"""
litellm.set_verbose = True
import asyncio
from unittest.mock import AsyncMock, patch
try:
router = litellm.Router(
@ -103,7 +110,7 @@ def test_hanging_request_azure():
{
"model_name": "azure-gpt",
"litellm_params": {
"model": "azure/gpt-4o-new-test",
"model": "azure/gpt-4.1-mini",
"api_base": os.environ["AZURE_API_BASE"],
"api_key": os.environ["AZURE_API_KEY"],
},
@ -118,17 +125,27 @@ def test_hanging_request_azure():
encoded = litellm.utils.encode(model="gpt-3.5-turbo", text="blue")[0]
original_send = httpx.AsyncClient.send
async def _slow_send(self, request, *args, **kwargs):
await asyncio.sleep(5)
return await original_send(self, request, *args, **kwargs)
async def _test():
response = await router.acompletion(
model="azure-gpt",
messages=[
{"role": "user", "content": f"what color is red {uuid.uuid4()}"}
],
logit_bias={encoded: 100},
timeout=0.01,
)
print(response)
return response
with patch.object(httpx.AsyncClient, "send", new=_slow_send):
response = await router.acompletion(
model="azure-gpt",
messages=[
{
"role": "user",
"content": f"what color is red {uuid.uuid4()}",
}
],
logit_bias={encoded: 100},
timeout=0.01,
)
print(response)
return response
response = asyncio.run(_test())

View file

@ -1,4 +1,12 @@
# conftest.py
#
# xdist-compatible test isolation for logging callback tests.
#
# Key design: capture litellm's true default values at conftest import time
# (BEFORE test modules are imported) so we can reset to clean defaults before
# each test. This is necessary because some test modules set module-level
# globals like `litellm.num_retries = 3` which pollute state for all tests
# in the same xdist worker.
import importlib
import os
@ -10,58 +18,118 @@ sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
import asyncio
@pytest.fixture(scope="session")
def event_loop():
try:
loop = asyncio.get_running_loop()
except RuntimeError:
loop = asyncio.new_event_loop()
yield loop
loop.close()
_LIST_ATTRS = (
"callbacks",
"success_callback",
"failure_callback",
"_async_success_callback",
"_async_failure_callback",
"service_callback",
"pre_call_rules",
"post_call_rules",
)
_SCALAR_ATTRS = (
"set_verbose",
"cache",
"num_retries",
"num_retries_per_request",
"turn_off_message_logging",
"redact_messages_in_exceptions",
"redact_user_api_key_info",
"s3_callback_params",
"datadog_params",
"vector_store_registry",
)
# ---- Capture true defaults at conftest import time ----
# This runs BEFORE any test modules are imported, so values are clean.
_DEFAULTS: dict = {}
for _attr in _LIST_ATTRS:
if hasattr(litellm, _attr):
_val = getattr(litellm, _attr)
_DEFAULTS[_attr] = _val.copy() if isinstance(_val, list) else _val
for _attr in _SCALAR_ATTRS:
if hasattr(litellm, _attr):
_DEFAULTS[_attr] = getattr(litellm, _attr)
@pytest.fixture(scope="function", autouse=True)
def setup_and_teardown():
def isolate_litellm_state():
"""
This fixture reloads litellm before every function. To speed up testing by removing callbacks being chained.
Per-function isolation fixture.
Resets litellm state to the true defaults captured at conftest import time,
then restores after the test. This prevents module-level mutations (e.g.
`litellm.num_retries = 3` at the top of test_langfuse_e2e_test.py) from
leaking across tests within the same xdist worker.
"""
curr_dir = os.getcwd() # Get the current working directory
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the project directory to the system path
from litellm.litellm_core_utils import litellm_logging as ll_logging
import litellm
from litellm import Router
import asyncio
# Flush cache and clear internal logger instances before test
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
# flush all logs
asyncio.run(GLOBAL_LOGGING_WORKER.clear_queue())
# Clear cached logger instances (LangsmithLogger, SlackAlerting, etc.)
ll_logging._in_memory_loggers.clear()
# Reset ALL attrs to their true defaults before the test runs.
# This undoes any module-level mutations from test file imports.
for attr in _LIST_ATTRS:
if attr in _DEFAULTS:
default = _DEFAULTS[attr]
setattr(litellm, attr, default.copy() if isinstance(default, list) else default)
importlib.reload(litellm)
for attr in _SCALAR_ATTRS:
if attr in _DEFAULTS:
setattr(litellm, attr, _DEFAULTS[attr])
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
import litellm.proxy.proxy_server
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
import asyncio
loop = asyncio.get_event_loop_policy().new_event_loop()
asyncio.set_event_loop(loop)
print(litellm)
# from litellm import Router, completion, aembedding, acompletion, embedding
yield
# Teardown code (executes after the yield point)
loop.close() # Close the loop created earlier
asyncio.set_event_loop(None) # Remove the reference to the loop
# Teardown: reset back to defaults again (belt-and-suspenders)
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
ll_logging._in_memory_loggers.clear()
for attr in _LIST_ATTRS:
if attr in _DEFAULTS:
default = _DEFAULTS[attr]
setattr(litellm, attr, default.copy() if isinstance(default, list) else default)
for attr in _SCALAR_ATTRS:
if attr in _DEFAULTS:
setattr(litellm, attr, _DEFAULTS[attr])
@pytest.fixture(scope="module", autouse=True)
def setup_and_teardown():
"""
Module-scoped setup. Reloads litellm only in single-process mode
(skipped under xdist to avoid cross-worker interference).
"""
sys.path.insert(0, os.path.abspath("../.."))
import litellm
worker_id = os.environ.get("PYTEST_XDIST_WORKER", None)
if worker_id is None:
importlib.reload(litellm)
try:
if hasattr(litellm, "proxy") and hasattr(litellm.proxy, "proxy_server"):
import litellm.proxy.proxy_server
importlib.reload(litellm.proxy.proxy_server)
except Exception as e:
print(f"Error reloading litellm.proxy.proxy_server: {e}")
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
yield
def pytest_collection_modifyitems(config, items):

View file

@ -475,7 +475,11 @@ async def test_langsmith_queue_logging():
mock_response="This is a mock response",
)
await asyncio.sleep(3)
# Poll for async callbacks to complete (up to 10s)
for _ in range(20):
if len(test_langsmith_logger.log_queue) >= 5:
break
await asyncio.sleep(0.5)
# Check that logs are in the queue
assert len(test_langsmith_logger.log_queue) == 5
@ -490,8 +494,11 @@ async def test_langsmith_queue_logging():
mock_response="This is a mock response",
)
# Wait a short time for any asynchronous operations to complete
await asyncio.sleep(1)
# Poll for flush to complete (up to 10s)
for _ in range(20):
if len(test_langsmith_logger.log_queue) < 5:
break
await asyncio.sleep(0.5)
print(
"Length of langsmith log queue: {}".format(

View file

@ -400,6 +400,8 @@ async def test_batch_status_sync_from_provider_to_database():
assert update_call_args.kwargs["data"]["status"] == "complete" # "completed" normalized to "complete"
assert "file_object" in update_call_args.kwargs["data"]
assert "updated_at" in update_call_args.kwargs["data"]
# batch_processed must be set to True when batch transitions to complete
assert update_call_args.kwargs["data"]["batch_processed"] is True
# Verify logger was called with status change message
mock_logger.info.assert_called()

View file

@ -109,17 +109,25 @@ async def test_basic_vertex_ai_pass_through_with_spendlog():
print("response", response)
await asyncio.sleep(40)
spend_after = await call_spend_logs_endpoint()
print("spend_after", spend_after)
# Poll for spend update instead of fixed sleep - spend logging is async/batched
max_wait = 120 # total seconds to wait
poll_interval = 10 # seconds between checks
elapsed = 0
spend_after = spend_before
while elapsed < max_wait:
await asyncio.sleep(poll_interval)
elapsed += poll_interval
spend_after = await call_spend_logs_endpoint() or 0.0
print(f"spend_after (elapsed={elapsed}s)", spend_after)
if spend_after > spend_before:
break
assert (
spend_after > spend_before
), "Spend should be greater than before. spend_before: {}, spend_after: {}".format(
spend_before, spend_after
), "Spend should be greater than before after {}s. spend_before: {}, spend_after: {}".format(
elapsed, spend_before, spend_after
)
pass
@pytest.mark.asyncio()
@pytest.mark.skip(reason="skip flaky test - vertex pass through streaming is flaky")

View file

@ -41,12 +41,65 @@ def litellm_proxy_config():
}
MAX_RETRIES = 3
async def _run_streaming_test(model_name: str) -> tuple[list[str], str]:
"""
Run a single streaming test attempt for the given model.
Returns (received_chunks, full_response).
"""
options = ClaudeAgentOptions(
system_prompt=(
"You are a helpful AI assistant. "
"Always follow the user's instructions exactly."
),
model=model_name,
max_turns=5,
)
test_query = (
"Respond with exactly the following text and nothing else:\n"
"Hello from LiteLLM!"
)
received_chunks: list[str] = []
full_response = ""
async with ClaudeSDKClient(options=options) as client:
await client.query(test_query)
async for msg in client.receive_response():
if hasattr(msg, 'type'):
if msg.type == 'content_block_delta':
if hasattr(msg, 'delta') and hasattr(msg.delta, 'text'):
chunk_text = msg.delta.text
received_chunks.append(chunk_text)
full_response += chunk_text
elif msg.type == 'content_block_start':
if hasattr(msg, 'content_block') and hasattr(msg.content_block, 'text'):
chunk_text = msg.content_block.text
received_chunks.append(chunk_text)
full_response += chunk_text
# Fallback to content handling
if hasattr(msg, 'content'):
for content_block in msg.content:
if hasattr(content_block, 'text'):
chunk_text = content_block.text
received_chunks.append(chunk_text)
full_response += chunk_text
return received_chunks, full_response
@pytest.mark.asyncio
@pytest.mark.parametrize("model_name,model_description", TEST_MODELS)
async def test_claude_agent_sdk_streaming(litellm_proxy_config, model_name, model_description):
"""
Test streaming messages with Claude Agent SDK through LiteLLM proxy.
This validates:
1. Claude Agent SDK can connect to LiteLLM proxy
2. Streaming works correctly
@ -55,25 +108,53 @@ async def test_claude_agent_sdk_streaming(litellm_proxy_config, model_name, mode
print(f"\n{'='*60}")
print(f"Testing: {model_name} ({model_description})")
print(f"{'='*60}")
# Configure agent options
options = ClaudeAgentOptions(
system_prompt="You are a helpful AI assistant. Be concise.",
model=model_name,
max_turns=5,
last_error: Exception | None = None
for attempt in range(1, MAX_RETRIES + 1):
try:
received_chunks, full_response = await _run_streaming_test(model_name)
# Assertions
print(f"\n✅ Received {len(received_chunks)} chunks")
print(f"📝 Full response: {full_response[:100]}...")
# Verify we got a response
assert len(full_response) > 0, f"No response received from {model_name}"
# Verify streaming (should have multiple chunks for most responses)
# Note: Very short responses might come in 1 chunk, so we just verify we got content
assert len(received_chunks) > 0, f"No chunks received from {model_name}"
# Verify response contains expected content (case insensitive)
assert "hello" in full_response.lower(), (
f"Response doesn't contain expected greeting: {full_response}"
)
print(f"✅ Test passed for {model_name} (attempt {attempt})")
return # Success
except Exception as e:
last_error = e
print(f"⚠️ Attempt {attempt}/{MAX_RETRIES} failed for {model_name}: {e}")
if attempt < MAX_RETRIES:
await asyncio.sleep(2)
pytest.fail(
f"Test failed for {model_name} ({model_description}) after {MAX_RETRIES} attempts: {last_error}"
)
# Test query
test_query = "Say 'Hello from LiteLLM!' and nothing else."
# Track streaming
received_chunks = []
full_response = ""
try:
async with ClaudeSDKClient(options=options) as client:
await client.query(test_query)
# Collect streaming response
async for msg in client.receive_response():
# Handle different message types
@ -90,7 +171,7 @@ async def test_claude_agent_sdk_streaming(litellm_proxy_config, model_name, mode
chunk_text = msg.content_block.text
received_chunks.append(chunk_text)
full_response += chunk_text
# Fallback to content handling
if hasattr(msg, 'content'):
for content_block in msg.content:
@ -98,23 +179,23 @@ async def test_claude_agent_sdk_streaming(litellm_proxy_config, model_name, mode
chunk_text = content_block.text
received_chunks.append(chunk_text)
full_response += chunk_text
# Assertions
print(f"\n✅ Received {len(received_chunks)} chunks")
print(f"📝 Full response: {full_response[:100]}...")
# Verify we got a response
assert len(full_response) > 0, f"No response received from {model_name}"
# Verify streaming (should have multiple chunks for most responses)
# Note: Very short responses might come in 1 chunk, so we just verify we got content
assert len(received_chunks) > 0, f"No chunks received from {model_name}"
# Verify response is non-empty (don't assert on specific LLM content — it's non-deterministic)
assert len(full_response.strip()) > 0, f"Empty response received from {model_name}"
print(f"✅ Test passed for {model_name}")
except Exception as e:
pytest.fail(f"Test failed for {model_name} ({model_description}): {str(e)}")

View file

@ -205,7 +205,7 @@ class TestManagedFilesAPI(ManagedFilesBase, UserKeyTestMixin):
return metadata
def _delete_file(self, file_id, label, max_retries=9, retry_delay=20):
def _delete_file(self, file_id, label, max_retries=10, retry_delay=5):
print(f"\nDeleting {label}: {self.shorten_id(file_id)}")
for attempt in range(max_retries):
try:
@ -235,7 +235,7 @@ class TestManagedFilesAPI(ManagedFilesBase, UserKeyTestMixin):
# Tests
# ------------------------------------------------------------------
@pytest.mark.flaky(reruns=5)
@pytest.mark.flaky(reruns=2)
@pytest.mark.parametrize(
"model_name",
get_batch_model_names(),

View file

@ -84,8 +84,14 @@ class TestCheckBatchCost:
assert find_call[1]["order"] == {"created_at": "asc"}
not_in = find_call[1]["where"]["status"]["not_in"]
assert "stale_expired" in not_in
assert "complete" in not_in
assert "completed" in not_in
# "complete"/"completed" are intentionally NOT excluded from the
# primary query — the batch_processed=False filter is sufficient.
# This allows CheckBatchCost to pick up batches that were
# transitioned to "complete" by the retrieve_batch endpoint
# before CheckBatchCost had a chance to process them.
assert "complete" not in not_in
assert "completed" not in not_in
assert find_call[1]["where"]["batch_processed"] is False
@pytest.mark.asyncio
async def test_fallback_query_used_when_batch_processed_missing(

View file

@ -49,6 +49,11 @@ def isolate_litellm_state():
if hasattr(litellm, '_async_failure_callback'):
original_state['_async_failure_callback'] = litellm._async_failure_callback.copy() if litellm._async_failure_callback else []
# Store routing globals — leaked model_fallbacks causes tests to route
# through async_completion_with_fallbacks / Router, bypassing HTTP mocks
if hasattr(litellm, 'model_fallbacks'):
original_state['model_fallbacks'] = litellm.model_fallbacks
# Store transport/network globals — many tests set these without restoring,
# causing subsequent tests to get None from _create_async_transport()
for _attr in ('disable_aiohttp_transport', 'force_ipv4'):
@ -59,7 +64,9 @@ def isolate_litellm_state():
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
# Clear success/failure callbacks to prevent chaining
# Clear all callback lists to prevent cross-test contamination
if hasattr(litellm, 'callbacks'):
litellm.callbacks = []
if hasattr(litellm, 'success_callback'):
litellm.success_callback = []
if hasattr(litellm, 'failure_callback'):
@ -69,6 +76,10 @@ def isolate_litellm_state():
if hasattr(litellm, '_async_failure_callback'):
litellm._async_failure_callback = []
# Clear routing globals
if hasattr(litellm, 'model_fallbacks'):
litellm.model_fallbacks = None
yield
# Cleanup after test

View file

@ -202,13 +202,16 @@ class TestImageEditCustomPricing:
mock_logging_obj = MagicMock()
mock_logging_obj.model_call_details = {}
original_update = mock_logging_obj.update_environment_variables
original_update = mock_logging_obj.update_from_kwargs
def capturing_update(**kwargs):
captured_litellm_params.update(kwargs.get("litellm_params", {}))
return original_update(**kwargs)
def capturing_update(**update_kwargs):
captured_litellm_params.update(update_kwargs.get("litellm_params", {}))
inner_kwargs = update_kwargs.get("kwargs", {})
if "metadata" in inner_kwargs:
captured_litellm_params["metadata"] = inner_kwargs["metadata"]
return original_update(**update_kwargs)
mock_logging_obj.update_environment_variables = capturing_update
mock_logging_obj.update_from_kwargs = capturing_update
with patch(
"litellm.images.main.get_llm_provider",

View file

@ -180,6 +180,110 @@ def test_use_custom_pricing_not_detected_litellm_metadata_no_pricing():
assert use_custom_pricing_for_model(litellm_params) is False
class TestUpdateFromKwargs:
"""Tests for the update_from_kwargs convenience wrapper."""
def test_extracts_metadata_from_kwargs(self, logging_obj):
metadata = {"user_api_key": "sk-test", "model_info": {"id": "abc"}}
kwargs = {"metadata": metadata, "other_key": "ignored"}
logging_obj.update_from_kwargs(
kwargs=kwargs,
litellm_params={"litellm_call_id": "call-1"},
)
assert logging_obj.litellm_params["metadata"] == metadata
assert logging_obj.litellm_params["litellm_call_id"] == "call-1"
def test_extracts_litellm_metadata_from_kwargs(self, logging_obj):
lm_meta = {
"model_info": {
"id": "deploy-1",
"input_cost_per_token": 0.001,
"output_cost_per_token": 0.002,
}
}
kwargs = {"litellm_metadata": lm_meta}
logging_obj.update_from_kwargs(
kwargs=kwargs,
litellm_params={"litellm_call_id": "call-2"},
)
assert logging_obj.litellm_params["litellm_metadata"] == lm_meta
assert logging_obj.litellm_params["litellm_call_id"] == "call-2"
def test_backfills_metadata_from_litellm_metadata(self, logging_obj):
"""When only litellm_metadata is present, metadata should be backfilled."""
lm_meta = {"model_info": {"id": "deploy-1"}}
kwargs = {"litellm_metadata": lm_meta}
logging_obj.update_from_kwargs(kwargs=kwargs)
assert logging_obj.litellm_params["metadata"] == lm_meta
def test_no_backfill_when_metadata_already_present(self, logging_obj):
metadata = {"user_api_key": "sk-real"}
lm_meta = {"model_info": {"id": "deploy-1"}}
kwargs = {"metadata": metadata, "litellm_metadata": lm_meta}
logging_obj.update_from_kwargs(kwargs=kwargs)
assert logging_obj.litellm_params["metadata"] == metadata
assert logging_obj.litellm_params["litellm_metadata"] == lm_meta
def test_caller_litellm_params_win_over_kwargs(self, logging_obj):
"""Explicit litellm_params from the caller should override auto-extracted values."""
kwargs = {"metadata": {"from_kwargs": True}}
logging_obj.update_from_kwargs(
kwargs=kwargs,
litellm_params={"metadata": {"from_caller": True}, "litellm_call_id": "x"},
)
assert logging_obj.litellm_params["metadata"] == {"from_caller": True}
def test_custom_pricing_detected_via_litellm_metadata(self, logging_obj):
"""Custom pricing in litellm_metadata.model_info should set custom_pricing flag."""
from litellm.litellm_core_utils.litellm_logging import (
use_custom_pricing_for_model,
)
lm_meta = {
"model_info": {
"id": "deploy-custom",
"input_cost_per_token": 0.005,
"output_cost_per_token": 0.015,
}
}
kwargs = {"litellm_metadata": lm_meta}
logging_obj.update_from_kwargs(kwargs=kwargs)
assert use_custom_pricing_for_model(logging_obj.litellm_params) is True
def test_additional_params_forwarded(self, logging_obj):
kwargs = {"metadata": {}}
logging_obj.update_from_kwargs(
kwargs=kwargs,
model="gpt-5",
user="test-user",
optional_params={"temperature": 0.7},
custom_llm_provider="openai",
)
assert logging_obj.model == "gpt-5"
assert logging_obj.user == "test-user"
assert logging_obj.model_call_details["custom_llm_provider"] == "openai"
def test_empty_kwargs_no_error(self, logging_obj):
logging_obj.update_from_kwargs(
kwargs={},
litellm_params={"litellm_call_id": "call-empty"},
)
assert logging_obj.litellm_params["litellm_call_id"] == "call-empty"
def test_logging_prevent_double_logging(logging_obj):
"""
When using a bridge, log only once from the underlying bridge call.

View file

@ -156,12 +156,11 @@ async def test_async_anthropic_messages_handler_extra_headers():
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_passes_litellm_metadata():
"""Ensure litellm_metadata from kwargs is included in litellm_params
passed to update_environment_variables.
"""Ensure litellm_metadata from kwargs is forwarded via update_from_kwargs.
Routes like /messages store model_info under kwargs['litellm_metadata'].
The handler must forward this into litellm_params so that
use_custom_pricing_for_model can detect custom pricing. Regression test for #23185.
The handler must forward this so that use_custom_pricing_for_model can
detect custom pricing. Regression test for #23185.
"""
handler = BaseLLMHTTPHandler()
@ -187,7 +186,7 @@ async def test_async_anthropic_messages_handler_passes_litellm_metadata():
mock_client.post = AsyncMock(return_value=mock_response)
mock_logging_obj = Mock()
mock_logging_obj.update_environment_variables = Mock()
mock_logging_obj.update_from_kwargs = Mock()
mock_logging_obj.model_call_details = {}
mock_logging_obj.stream = False
@ -218,14 +217,14 @@ async def test_async_anthropic_messages_handler_passes_litellm_metadata():
except Exception:
pass
mock_logging_obj.update_environment_variables.assert_called_once()
call_kwargs = mock_logging_obj.update_environment_variables.call_args
litellm_params_arg = call_kwargs.kwargs.get(
"litellm_params", call_kwargs[1].get("litellm_params", {})
) if call_kwargs.kwargs else call_kwargs[1].get("litellm_params", {})
mock_logging_obj.update_from_kwargs.assert_called_once()
call_kwargs = mock_logging_obj.update_from_kwargs.call_args
kwargs_arg = call_kwargs.kwargs.get(
"kwargs", call_kwargs[1].get("kwargs", {})
) if call_kwargs.kwargs else call_kwargs[1].get("kwargs", {})
assert "litellm_metadata" in litellm_params_arg
assert litellm_params_arg["litellm_metadata"]["model_info"] == custom_model_info
assert "litellm_metadata" in kwargs_arg
assert kwargs_arg["litellm_metadata"]["model_info"] == custom_model_info
@pytest.mark.asyncio

View file

@ -92,30 +92,33 @@ async def test_metadata_passed_to_custom_callback_codex_models():
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
litellm.callbacks = [callback]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _make_mock_http_response(
mock_response.model_dump()
)
# gpt-5.1-codex has mode=responses - routes through responses bridge
await litellm.acompletion(
model="gpt-5.1-codex",
messages=[{"role": "user", "content": "Hello"}],
metadata=test_metadata,
)
try:
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _make_mock_http_response(
mock_response.model_dump()
)
# gpt-5.1-codex has mode=responses - routes through responses bridge
await litellm.acompletion(
model="gpt-5.1-codex",
messages=[{"role": "user", "content": "Hello"}],
metadata=test_metadata,
)
await asyncio.wait_for(callback.event.wait(), timeout=5.0)
await asyncio.wait_for(callback.event.wait(), timeout=5.0)
assert callback.captured_kwargs is not None, "Callback should have been invoked"
assert callback.captured_kwargs is not None, "Callback should have been invoked"
litellm_params = callback.captured_kwargs.get("litellm_params", {})
metadata = litellm_params.get("metadata") or {}
litellm_params = callback.captured_kwargs.get("litellm_params", {})
metadata = litellm_params.get("metadata") or {}
assert "foo" in metadata, "metadata['foo'] should be accessible in callback"
assert metadata["foo"] == "bar"
assert metadata.get("trace_id") == "test-123"
assert "foo" in metadata, "metadata['foo'] should be accessible in callback"
assert metadata["foo"] == "bar"
assert metadata.get("trace_id") == "test-123"
finally:
litellm.callbacks = original_callbacks
@pytest.mark.asyncio
@ -152,27 +155,31 @@ async def test_metadata_passed_via_litellm_metadata_responses_api():
test_metadata = {"request_id": "req-456"}
callback = MetadataCaptureCallback()
original_callbacks = litellm.callbacks.copy() if litellm.callbacks else []
litellm.callbacks = [callback]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _make_mock_http_response(
mock_response.model_dump()
)
await litellm.aresponses(
model="gpt-4o",
input="hi",
litellm_metadata=test_metadata,
)
try:
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = _make_mock_http_response(
mock_response.model_dump()
)
await litellm.aresponses(
model="gpt-4o",
input="hi",
litellm_metadata=test_metadata,
)
await asyncio.wait_for(callback.event.wait(), timeout=5.0)
await asyncio.wait_for(callback.event.wait(), timeout=5.0)
assert callback.captured_kwargs is not None
assert callback.captured_kwargs is not None
litellm_params = callback.captured_kwargs.get("litellm_params", {})
metadata = litellm_params.get("metadata") or {}
litellm_params = callback.captured_kwargs.get("litellm_params", {})
metadata = litellm_params.get("metadata") or {}
assert "request_id" in metadata
assert metadata["request_id"] == "req-456"
assert "request_id" in metadata
assert metadata["request_id"] == "req-456"
finally:
litellm.callbacks = original_callbacks

View file

@ -23,25 +23,26 @@ import pytest
sys.path.insert(0, os.path.abspath("../.."))
import json
import litellm
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponsesAPIResponse
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
class MockResponse:
def __init__(self, json_data, status_code):
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
self.headers = {}
def json(self):
return self._json_data
def _build_mock_response(output_items, response_id="resp_mock-123"):
"""Build a ResponsesAPIResponse that ``async_response_api_handler`` would return."""
return ResponsesAPIResponse(
id=response_id,
created_at=1741476542,
status="completed",
model="openai/gpt-5.1-codex",
output=output_items,
usage={"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
)
def _get_item_id(item) -> str:
@ -51,18 +52,8 @@ def _get_item_id(item) -> str:
return getattr(item, "id", "") or ""
def _has_encrypted_content(item) -> bool:
"""Check whether an output item carries encrypted_content."""
if isinstance(item, dict):
return "encrypted_content" in item
return hasattr(item, "encrypted_content") and getattr(item, "encrypted_content") is not None
def _extract_encoded_item_id(response) -> str:
"""
Walk the response output and return the first litellm-encoded item ID
(i.e. one that starts with ``encitem_``).
"""
"""Return the first ``encitem_``-prefixed item ID from the response output."""
for item in response.output or []:
item_id = _get_item_id(item)
if item_id.startswith("encitem_"):
@ -254,14 +245,14 @@ async def test_encrypted_content_affinity_tracks_and_routes():
"""
The first response rewrites encrypted-content item IDs to encoded form.
The follow-up request with those encoded IDs is pinned to the same deployment.
Mocks ``async_response_api_handler`` (the method that makes the HTTP call)
so the test is deterministic regardless of the HTTP transport in use.
The ``@client`` decorator and ``_update_responses_api_response_id_with_model_id``
post-processing still run, so item-ID rewriting is exercised end-to-end.
"""
mock_response_data = {
"id": "resp_mock-123",
"object": "response",
"created_at": 1741476542,
"status": "completed",
"model": "openai/gpt-5.1-codex",
"output": [
mock_resp = _build_mock_response(
output_items=[
{
"type": "message",
"id": "msg_abc123",
@ -276,10 +267,7 @@ async def test_encrypted_content_affinity_tracks_and_routes():
"encrypted_content": "gAAAAABpnW_yEYmSNEyOG...",
},
],
"parallel_tool_calls": True,
"usage": {"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
"error": None,
}
)
router = litellm.Router(
model_list=[
@ -301,6 +289,7 @@ async def test_encrypted_content_affinity_tracks_and_routes():
},
],
optional_pre_call_checks=["encrypted_content_affinity"],
num_retries=0,
)
selected_deployments = []
@ -311,14 +300,13 @@ async def test_encrypted_content_affinity_tracks_and_routes():
return seq[1] if len(seq) > 1 else seq[0]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
"litellm.llms.custom_httpx.llm_http_handler.BaseLLMHTTPHandler.async_response_api_handler",
new_callable=AsyncMock,
) as mock_post, patch(
return_value=mock_resp,
), patch(
"litellm.router_strategy.simple_shuffle.random.choice",
side_effect=deterministic_choice,
):
mock_post.return_value = MockResponse(mock_response_data, 200)
# First request — goes to deployment-1 via deterministic_choice
first_response = await router.aresponses(
model="openai.gpt-5.1-codex",
@ -376,6 +364,7 @@ async def test_encrypted_content_affinity_no_effect_on_chat_completions():
},
],
optional_pre_call_checks=["encrypted_content_affinity"],
num_retries=0,
)
response1 = await router.acompletion(
@ -394,15 +383,10 @@ async def test_encrypted_content_affinity_no_effect_on_chat_completions():
async def test_encrypted_content_affinity_bypasses_rpm_limits():
"""
When encrypted content affinity pins to a deployment, the request
goes through even if normal routing would avoid it.
goes through even if normal routing would avoid it (usage-based-routing-v2).
"""
mock_response_data = {
"id": "resp_mock-rpm-test",
"object": "response",
"created_at": 1741476542,
"status": "completed",
"model": "openai/gpt-5.1-codex",
"output": [
mock_resp = _build_mock_response(
output_items=[
{
"type": "reasoning",
"id": "rs_encrypted_must_pin",
@ -410,9 +394,8 @@ async def test_encrypted_content_affinity_bypasses_rpm_limits():
"encrypted_content": "gAAAAABpnW_yEYmSNEyOG...",
},
],
"usage": {"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
"error": None,
}
response_id="resp_mock-rpm-test",
)
router = litellm.Router(
model_list=[
@ -435,6 +418,7 @@ async def test_encrypted_content_affinity_bypasses_rpm_limits():
],
optional_pre_call_checks=["encrypted_content_affinity"],
routing_strategy="usage-based-routing-v2",
num_retries=0,
)
selected_deployments = []
@ -445,14 +429,13 @@ async def test_encrypted_content_affinity_bypasses_rpm_limits():
return seq[1] if len(seq) > 1 else seq[0]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
"litellm.llms.custom_httpx.llm_http_handler.BaseLLMHTTPHandler.async_response_api_handler",
new_callable=AsyncMock,
) as mock_post, patch(
return_value=mock_resp,
), patch(
"litellm.router_strategy.simple_shuffle.random.choice",
side_effect=deterministic_choice,
):
mock_post.return_value = MockResponse(mock_response_data, 200)
first_response = await router.aresponses(
model="openai.gpt-5.1-codex",
input="Initial request",
@ -488,13 +471,8 @@ async def test_encrypted_content_affinity_no_match_normal_routing():
Input items with non-encoded IDs (no encitem_ prefix) fall through to
normal load balancing.
"""
mock_response_data = {
"id": "resp_mock-no-match",
"object": "response",
"created_at": 1741476542,
"status": "completed",
"model": "openai/gpt-5.1-codex",
"output": [
mock_resp = _build_mock_response(
output_items=[
{
"type": "message",
"id": "msg_new",
@ -503,9 +481,8 @@ async def test_encrypted_content_affinity_no_match_normal_routing():
"content": [{"type": "output_text", "text": "Response"}],
},
],
"usage": {"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
"error": None,
}
response_id="resp_mock-no-match",
)
router = litellm.Router(
model_list=[
@ -527,14 +504,14 @@ async def test_encrypted_content_affinity_no_match_normal_routing():
},
],
optional_pre_call_checks=["encrypted_content_affinity"],
num_retries=0,
)
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
"litellm.llms.custom_httpx.llm_http_handler.BaseLLMHTTPHandler.async_response_api_handler",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = MockResponse(mock_response_data, 200)
return_value=mock_resp,
):
# Non-encoded item ID — no affinity should kick in
response = await router.aresponses(
model="openai.gpt-5.1-codex",
@ -551,22 +528,16 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
Test affinity routing when items have wrapped encrypted_content but no ID.
This simulates Codex client behavior where IDs are omitted.
"""
mock_response_data = {
"id": "resp_mock-wrapped-content",
"object": "response",
"created_at": 1741476542,
"status": "completed",
"model": "openai/gpt-5.1-codex",
"output": [
mock_resp = _build_mock_response(
output_items=[
{
"type": "reasoning",
"status": "completed",
"encrypted_content": "gAAAAABpnW_yEYmSNEyOG_original_content",
},
],
"usage": {"input_tokens": 5, "output_tokens": 10, "total_tokens": 15},
"error": None,
}
response_id="resp_mock-wrapped-content",
)
router = litellm.Router(
model_list=[
@ -588,6 +559,7 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
},
],
optional_pre_call_checks=["encrypted_content_affinity"],
num_retries=0,
)
selected_deployments = []
@ -598,14 +570,13 @@ async def test_encrypted_content_affinity_with_wrapped_content_no_id():
return seq[1] if len(seq) > 1 else seq[0]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
"litellm.llms.custom_httpx.llm_http_handler.BaseLLMHTTPHandler.async_response_api_handler",
new_callable=AsyncMock,
) as mock_post, patch(
return_value=mock_resp,
), patch(
"litellm.router_strategy.simple_shuffle.random.choice",
side_effect=deterministic_choice,
):
mock_post.return_value = MockResponse(mock_response_data, 200)
# First request — goes to deployment-1
first_response = await router.aresponses(
model="openai.gpt-5.1-codex",

View file

@ -6,76 +6,83 @@ encoding is loaded at import time (pre-#18070 behavior) instead of lazy loading.
This addresses issue #18659: VCR cassette creation broken by lazy loading.
For now, this only affects encoding as it was the only reported issue.
Tests that need to clear sys.modules and re-import litellm run in subprocesses
to avoid contaminating the test process's module graph (which breaks mock.patch
for all subsequent tests on the same xdist worker).
"""
import os
import subprocess
import sys
import textwrap
import pytest
def _run_python(script: str, env_override: dict | None = None) -> subprocess.CompletedProcess:
"""Run a Python script in a subprocess and return the result."""
import os
env = os.environ.copy()
# Remove the var so each test controls it explicitly
env.pop("LITELLM_DISABLE_LAZY_LOADING", None)
env.pop("TIKTOKEN_CACHE_DIR", None)
if env_override:
env.update(env_override)
return subprocess.run(
[sys.executable, "-c", textwrap.dedent(script)],
capture_output=True,
text=True,
env=env,
timeout=60,
)
def test_eager_loading_enabled():
"""Test that encoding is loaded at import time when env var is set"""
# Set environment variable
os.environ["LITELLM_DISABLE_LAZY_LOADING"] = "1"
# Clear any cached modules to ensure fresh import
modules_to_clear = [k for k in sys.modules.keys() if k.startswith("litellm")]
for module in modules_to_clear:
del sys.modules[module]
# Import litellm - encoding should be loaded immediately
import litellm
# Check that encoding is available (not lazy loaded)
assert hasattr(litellm, "encoding"), "Encoding should be available when eager loading is enabled"
# Verify it's actually the encoding object
encoding = litellm.encoding
assert encoding is not None, "Encoding should not be None"
# Test that it works
tokens = encoding.encode("Hello, world!")
assert len(tokens) > 0, "Encoding should work"
result = _run_python(
"""
import litellm
assert hasattr(litellm, "encoding"), "Encoding should be available when eager loading is enabled"
encoding = litellm.encoding
assert encoding is not None, "Encoding should not be None"
tokens = encoding.encode("Hello, world!")
assert len(tokens) > 0, "Encoding should work"
""",
env_override={"LITELLM_DISABLE_LAZY_LOADING": "1"},
)
assert result.returncode == 0, f"Subprocess failed:\nstdout: {result.stdout}\nstderr: {result.stderr}"
def test_eager_loading_env_var_values():
"""Test that various env var values enable eager loading"""
values = ["1", "true", "True", "TRUE", "yes", "Yes", "YES", "on", "On", "ON"]
for value in values:
os.environ["LITELLM_DISABLE_LAZY_LOADING"] = value
# Clear modules
modules_to_clear = [k for k in sys.modules.keys() if k.startswith("litellm")]
for module in modules_to_clear:
del sys.modules[module]
import litellm
assert hasattr(litellm, "encoding"), f"Encoding should be available for value: {value}"
encoding = litellm.encoding
tokens = encoding.encode("test")
assert len(tokens) > 0
result = _run_python(
"""
import litellm
assert hasattr(litellm, "encoding"), "Encoding should be available"
encoding = litellm.encoding
tokens = encoding.encode("test")
assert len(tokens) > 0
""",
env_override={"LITELLM_DISABLE_LAZY_LOADING": value},
)
assert result.returncode == 0, (
f"Failed for value {value!r}:\nstdout: {result.stdout}\nstderr: {result.stderr}"
)
def test_lazy_loading_default():
"""Test that encoding is lazy loaded by default (when env var is not set)"""
# Remove environment variable if set
if "LITELLM_DISABLE_LAZY_LOADING" in os.environ:
del os.environ["LITELLM_DISABLE_LAZY_LOADING"]
# Clear any cached modules
modules_to_clear = [k for k in sys.modules.keys() if k.startswith("litellm")]
for module in modules_to_clear:
del sys.modules[module]
# Import litellm - encoding should NOT be loaded yet
import litellm
# Encoding should be accessible via __getattr__ (lazy loading)
encoding = litellm.encoding # This triggers lazy loading
# Verify it works
tokens = encoding.encode("Hello, world!")
assert len(tokens) > 0, "Encoding should work"
result = _run_python(
"""
import litellm
# Encoding should be accessible via __getattr__ (lazy loading)
encoding = litellm.encoding
tokens = encoding.encode("Hello, world!")
assert len(tokens) > 0, "Encoding should work"
""",
)
assert result.returncode == 0, f"Subprocess failed:\nstdout: {result.stdout}\nstderr: {result.stderr}"
def test_tiktoken_cache_dir_set_on_lazy_load():
@ -84,33 +91,15 @@ def test_tiktoken_cache_dir_set_on_lazy_load():
This ensures the local tiktoken cache is used instead of downloading
from the internet. Regression test for issue #19768.
"""
# Remove environment variables to ensure clean state
if "LITELLM_DISABLE_LAZY_LOADING" in os.environ:
del os.environ["LITELLM_DISABLE_LAZY_LOADING"]
if "TIKTOKEN_CACHE_DIR" in os.environ:
del os.environ["TIKTOKEN_CACHE_DIR"]
# Clear any cached modules
modules_to_clear = [k for k in sys.modules.keys() if k.startswith("litellm")]
for module in modules_to_clear:
del sys.modules[module]
# Import litellm fresh
import litellm
# Access encoding (triggers lazy load)
_ = litellm.encoding
# Verify TIKTOKEN_CACHE_DIR is now set and points to local tokenizers
assert "TIKTOKEN_CACHE_DIR" in os.environ, "TIKTOKEN_CACHE_DIR should be set after lazy loading encoding"
cache_dir = os.environ["TIKTOKEN_CACHE_DIR"]
assert "tokenizers" in cache_dir, f"TIKTOKEN_CACHE_DIR should point to tokenizers directory, got: {cache_dir}"
@pytest.fixture(autouse=True)
def cleanup_env():
"""Clean up environment variable after each test"""
yield
if "LITELLM_DISABLE_LAZY_LOADING" in os.environ:
del os.environ["LITELLM_DISABLE_LAZY_LOADING"]
result = _run_python(
"""
import os
import litellm
# Access encoding (triggers lazy load)
_ = litellm.encoding
assert "TIKTOKEN_CACHE_DIR" in os.environ, "TIKTOKEN_CACHE_DIR should be set after lazy loading encoding"
cache_dir = os.environ["TIKTOKEN_CACHE_DIR"]
assert "tokenizers" in cache_dir, f"TIKTOKEN_CACHE_DIR should point to tokenizers directory, got: {cache_dir}"
""",
)
assert result.returncode == 0, f"Subprocess failed:\nstdout: {result.stdout}\nstderr: {result.stderr}"

View file

@ -24,6 +24,10 @@ export default defineConfig({
/* Collect trace when retrying the failed test. See https://playwright.dev/docs/trace-viewer */
trace: "on-first-retry",
/* Action timeout for clicks, fills, waitForSelector, etc. */
actionTimeout: 15 * 1000,
navigationTimeout: 30 * 1000,
},
/* Configure projects for major browsers */
@ -40,7 +44,7 @@ export default defineConfig({
],
/* Timeout settings */
timeout: 4 * 60 * 1000,
timeout: 3 * 60 * 1000,
expect: {
timeout: 10 * 1000,
},