diff --git a/litellm/__init__.py b/litellm/__init__.py index 83cbcfda3f1..4c59c335ce1 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1272,7 +1272,7 @@ if TYPE_CHECKING: from litellm.types.utils import ModelInfo as _ModelInfoType from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.caching.caching import Cache - + # Type stubs for lazy-loaded configs to help mypy from .llms.bedrock.chat.converse_transformation import AmazonConverseConfig as AmazonConverseConfig from .llms.openai_like.chat.handler import OpenAILikeChatConfig as OpenAILikeChatConfig @@ -1378,6 +1378,7 @@ if TYPE_CHECKING: from .llms.azure.responses.o_series_transformation import AzureOpenAIOSeriesResponsesAPIConfig as AzureOpenAIOSeriesResponsesAPIConfig from .llms.xai.responses.transformation import XAIResponsesAPIConfig as XAIResponsesAPIConfig from .llms.litellm_proxy.responses.transformation import LiteLLMProxyResponsesAPIConfig as LiteLLMProxyResponsesAPIConfig + from .llms.volcengine.responses.transformation import VolcEngineResponsesAPIConfig as VolcEngineResponsesAPIConfig from .llms.manus.responses.transformation import ManusResponsesAPIConfig as ManusResponsesAPIConfig from .llms.gemini.interactions.transformation import GoogleAIStudioInteractionsConfig as GoogleAIStudioInteractionsConfig from .llms.openai.chat.o_series_transformation import OpenAIOSeriesConfig as OpenAIOSeriesConfig, OpenAIOSeriesConfig as OpenAIO1Config @@ -1391,7 +1392,7 @@ if TYPE_CHECKING: from .llms.openai.chat.gpt_audio_transformation import OpenAIGPTAudioConfig as OpenAIGPTAudioConfig from .llms.nvidia_nim.chat.transformation import NvidiaNimConfig as NvidiaNimConfig from .llms.nvidia_nim.embed import NvidiaNimEmbeddingConfig as NvidiaNimEmbeddingConfig - + # Type stubs for lazy-loaded config instances openaiOSeriesConfig: OpenAIOSeriesConfig openAIGPTConfig: OpenAIGPTConfig @@ -1399,7 +1400,7 @@ if TYPE_CHECKING: openAIGPT5Config: OpenAIGPT5Config nvidiaNimConfig: NvidiaNimConfig nvidiaNimEmbeddingConfig: NvidiaNimEmbeddingConfig - + # Import config classes that need type stubs (for mypy) - import with _ prefix to avoid circular reference from .llms.vllm.completion.transformation import VLLMConfig as _VLLMConfig from .llms.deepseek.chat.transformation import DeepSeekChatConfig as _DeepSeekChatConfig @@ -1417,7 +1418,7 @@ if TYPE_CHECKING: from .llms.lm_studio.embed.transformation import LmStudioEmbeddingConfig as _LmStudioEmbeddingConfig from .llms.watsonx.embed.transformation import IBMWatsonXEmbeddingConfig as _IBMWatsonXEmbeddingConfig from .llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig as _VertexGeminiConfig - + # Type stubs for lazy-loaded config classes (to help mypy understand types) VLLMConfig: Type[_VLLMConfig] DeepSeekChatConfig: Type[_DeepSeekChatConfig] @@ -1435,7 +1436,7 @@ if TYPE_CHECKING: LmStudioEmbeddingConfig: Type[_LmStudioEmbeddingConfig] IBMWatsonXEmbeddingConfig: Type[_IBMWatsonXEmbeddingConfig] VertexAIConfig: Type[_VertexGeminiConfig] # Alias for VertexGeminiConfig - + from .llms.featherless_ai.chat.transformation import FeatherlessAIConfig as FeatherlessAIConfig from .llms.cerebras.chat import CerebrasConfig as CerebrasConfig from .llms.baseten.chat import BasetenConfig as BasetenConfig @@ -1557,14 +1558,14 @@ if TYPE_CHECKING: # Custom logger class (lazy-loaded) from litellm.integrations.custom_logger import CustomLogger - + # Datadog LLM observability params (lazy-loaded) from litellm.types.integrations.datadog_llm_obs import DatadogLLMObsInitParams - + # Logging callback manager class and instance (lazy-loaded) from litellm.litellm_core_utils.logging_callback_manager import LoggingCallbackManager logging_callback_manager: LoggingCallbackManager - + # provider_list is lazy-loaded from litellm.types.utils import LlmProviders provider_list: List[Union[LlmProviders, str]] @@ -1594,12 +1595,12 @@ def __getattr__(name: str) -> Any: from litellm.llms.custom_httpx.async_client_cleanup import register_async_client_cleanup register_async_client_cleanup() _async_client_cleanup_registered = True - + # Use cached registry from _lazy_imports instead of importing tuples every time from ._lazy_imports import _get_lazy_import_registry - + registry = _get_lazy_import_registry() - + # Check if name is in registry and call the cached handler function if name in registry: handler_func = registry[name] @@ -1614,7 +1615,7 @@ def __getattr__(name: str) -> Any: from .main import encoding as _encoding _globals["encoding"] = _encoding return _globals["encoding"] - + # Lazy load bedrock_tool_name_mappings instance if name == "bedrock_tool_name_mappings": from ._lazy_imports import _get_litellm_globals @@ -1624,7 +1625,7 @@ def __getattr__(name: str) -> Any: from .llms.bedrock.chat.invoke_handler import bedrock_tool_name_mappings as _bedrock_tool_name_mappings _globals["bedrock_tool_name_mappings"] = _bedrock_tool_name_mappings return _globals["bedrock_tool_name_mappings"] - + # Lazy load AzureOpenAIError exception class if name == "AzureOpenAIError": from ._lazy_imports import _get_litellm_globals @@ -1634,7 +1635,7 @@ def __getattr__(name: str) -> Any: from .llms.azure.common_utils import AzureOpenAIError as _AzureOpenAIError _globals["AzureOpenAIError"] = _AzureOpenAIError return _globals["AzureOpenAIError"] - + # Lazy load openaiOSeriesConfig instance if name == "openaiOSeriesConfig": from ._lazy_imports import _get_litellm_globals @@ -1644,7 +1645,7 @@ def __getattr__(name: str) -> Any: config_class = __getattr__("OpenAIOSeriesConfig") _globals["openaiOSeriesConfig"] = config_class() return _globals["openaiOSeriesConfig"] - + # Lazy load other config instances _config_instances = { "openAIGPTConfig": "OpenAIGPTConfig", @@ -1661,11 +1662,11 @@ def __getattr__(name: str) -> Any: config_class = __getattr__(_config_instances[name]) _globals[name] = config_class() return _globals[name] - + # Handle OpenAIO1Config alias if name == "OpenAIO1Config": return __getattr__("OpenAIOSeriesConfig") - + # Lazy load provider_list if name == "provider_list": from ._lazy_imports import _get_litellm_globals @@ -1676,7 +1677,7 @@ def __getattr__(name: str) -> Any: from litellm.types.utils import LlmProviders _globals["provider_list"] = list(LlmProviders) return _globals["provider_list"] - + # Lazy load priority_reservation_settings instance if name == "priority_reservation_settings": from ._lazy_imports import _get_litellm_globals @@ -1687,7 +1688,7 @@ def __getattr__(name: str) -> Any: PriorityReservationSettings = __getattr__("PriorityReservationSettings") _globals["priority_reservation_settings"] = PriorityReservationSettings() return _globals["priority_reservation_settings"] - + # Lazy load logging_callback_manager instance if name == "logging_callback_manager": from ._lazy_imports import _get_litellm_globals @@ -1698,7 +1699,7 @@ def __getattr__(name: str) -> Any: LoggingCallbackManager = __getattr__("LoggingCallbackManager") _globals["logging_callback_manager"] = LoggingCallbackManager() return _globals["logging_callback_manager"] - + # Lazy load _service_logger module if name == "_service_logger": from ._lazy_imports import _get_litellm_globals diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 46abde4eb5e..a92c6f95b0e 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -198,6 +198,7 @@ LLM_CONFIG_NAMES = ( "AzureOpenAIOSeriesResponsesAPIConfig", "XAIResponsesAPIConfig", "LiteLLMProxyResponsesAPIConfig", + "VolcEngineResponsesAPIConfig", "GoogleAIStudioInteractionsConfig", "OpenAIOSeriesConfig", "AnthropicSkillsConfig", @@ -593,6 +594,7 @@ _LLM_CONFIGS_IMPORT_MAP = { "AzureOpenAIOSeriesResponsesAPIConfig": (".llms.azure.responses.o_series_transformation", "AzureOpenAIOSeriesResponsesAPIConfig"), "XAIResponsesAPIConfig": (".llms.xai.responses.transformation", "XAIResponsesAPIConfig"), "LiteLLMProxyResponsesAPIConfig": (".llms.litellm_proxy.responses.transformation", "LiteLLMProxyResponsesAPIConfig"), + "VolcEngineResponsesAPIConfig": (".llms.volcengine.responses.transformation", "VolcEngineResponsesAPIConfig"), "ManusResponsesAPIConfig": (".llms.manus.responses.transformation", "ManusResponsesAPIConfig"), "GoogleAIStudioInteractionsConfig": (".llms.gemini.interactions.transformation", "GoogleAIStudioInteractionsConfig"), "OpenAIOSeriesConfig": (".llms.openai.chat.o_series_transformation", "OpenAIOSeriesConfig"), @@ -778,4 +780,3 @@ __all__ = [ "_LLM_PROVIDER_LOGIC_IMPORT_MAP", "_UTILS_MODULE_IMPORT_MAP", ] - diff --git a/litellm/constants.py b/litellm/constants.py index f21e7517482..c98551fb1b6 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -543,6 +543,10 @@ DEFAULT_CHAT_COMPLETION_PARAM_VALUES = { "web_search_options": None, "service_tier": None, "safety_identifier": None, + "prompt_cache_key": None, + "prompt_cache_retention": None, + "store": None, + "metadata": None, } openai_compatible_endpoints: List = [ diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 65e77f014a3..785976ed319 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -354,7 +354,7 @@ class PromptTokensDetailsResult(TypedDict): image_tokens: int character_count: int image_count: int - video_length_seconds: int + video_length_seconds: float def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult: @@ -400,10 +400,10 @@ def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult: ) video_length_seconds = ( cast( - Optional[int], + Optional[float], getattr(usage.prompt_tokens_details, "video_length_seconds", 0), ) - or 0 + or 0.0 ) return PromptTokensDetailsResult( @@ -415,7 +415,7 @@ def _parse_prompt_tokens_details(usage: Usage) -> PromptTokensDetailsResult: image_tokens=image_tokens, character_count=character_count, image_count=image_count, - video_length_seconds=video_length_seconds, + video_length_seconds=float(video_length_seconds), ) @@ -561,7 +561,7 @@ def generic_cost_per_token( # noqa: PLR0915 image_tokens=0, character_count=0, image_count=0, - video_length_seconds=0, + video_length_seconds=0.0, ) if usage.prompt_tokens_details: prompt_tokens_details = _parse_prompt_tokens_details(usage) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index dced06b6b3a..cb9fe0aeb30 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -221,7 +221,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): ### CHECK IF CLOUDFLARE AI GATEWAY ### ### if so - set the model as part of the base url - if "gateway.ai.cloudflare.com" in api_base: + if api_base is not None and "gateway.ai.cloudflare.com" in api_base: client = self._init_azure_client_for_cloudflare_ai_gateway( api_base=api_base, model=model, @@ -1344,7 +1344,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): prompt: Optional[str] = None, ) -> dict: client_session = litellm.client_session or httpx.Client() - if "gateway.ai.cloudflare.com" in api_base: + if api_base is not None and "gateway.ai.cloudflare.com" in api_base: ## build base url - assume api base includes resource name if not api_base.endswith("/"): api_base += "/" diff --git a/litellm/llms/custom_httpx/aiohttp_handler.py b/litellm/llms/custom_httpx/aiohttp_handler.py index c7a04a49fc2..93b6c563dc1 100644 --- a/litellm/llms/custom_httpx/aiohttp_handler.py +++ b/litellm/llms/custom_httpx/aiohttp_handler.py @@ -134,6 +134,41 @@ class BaseLLMAIOHTTPHandler: # Ignore errors during transport cleanup pass + def __del__(self): + """ + Cleanup: close aiohttp session on instance destruction. + + Provides defense-in-depth for issue #12443 - ensures cleanup happens + even if atexit handler doesn't run (abnormal termination). + """ + if ( + self.client_session is not None + and not self.client_session.closed + and self._owns_session + ): + try: + import asyncio + + try: + loop = asyncio.get_event_loop() + if loop.is_running(): + # Event loop is running - schedule cleanup task + asyncio.create_task(self.close()) + else: + # Event loop exists but not running - run cleanup + loop.run_until_complete(self.close()) + except RuntimeError: + # No event loop available - create one for cleanup + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + loop.run_until_complete(self.close()) + finally: + loop.close() + except Exception: + # Silently ignore errors during __del__ to avoid issues + pass + async def _make_common_async_call( self, async_client_session: Optional[ClientSession], diff --git a/litellm/llms/custom_httpx/async_client_cleanup.py b/litellm/llms/custom_httpx/async_client_cleanup.py index 45602576764..abbc61dc96d 100644 --- a/litellm/llms/custom_httpx/async_client_cleanup.py +++ b/litellm/llms/custom_httpx/async_client_cleanup.py @@ -9,7 +9,8 @@ async def close_litellm_async_clients(): Close all cached async HTTP clients to prevent resource leaks. This function iterates through all cached clients in litellm's in-memory cache - and closes any aiohttp client sessions that are still open. + and closes any aiohttp client sessions that are still open. Also closes the + global base_llm_aiohttp_handler instance (issue #12443). """ # Import here to avoid circular import import litellm @@ -25,7 +26,7 @@ async def close_litellm_async_clients(): except Exception: # Silently ignore errors during cleanup pass - + # Handle AsyncHTTPHandler instances (used by Gemini and other providers) elif hasattr(handler, 'client'): client = handler.client @@ -43,7 +44,7 @@ async def close_litellm_async_clients(): except Exception: # Silently ignore errors during cleanup pass - + # Handle any other objects with aclose method elif hasattr(handler, 'aclose'): try: @@ -52,6 +53,17 @@ async def close_litellm_async_clients(): # Silently ignore errors during cleanup pass + # Close the global base_llm_aiohttp_handler instance (issue #12443) + # This is used by Gemini and other providers that use aiohttp + if hasattr(litellm, 'base_llm_aiohttp_handler'): + base_handler = getattr(litellm, 'base_llm_aiohttp_handler', None) + if isinstance(base_handler, BaseLLMAIOHTTPHandler) and hasattr(base_handler, 'close'): + try: + await base_handler.close() + except Exception: + # Silently ignore errors during cleanup + pass + def register_async_client_cleanup(): """ @@ -62,22 +74,24 @@ def register_async_client_cleanup(): import atexit def cleanup_wrapper(): + """ + Cleanup wrapper that creates a fresh event loop for atexit cleanup. + + At exit time, the main event loop is often already closed. Creating a new + event loop ensures cleanup runs successfully (fixes issue #12443). + """ try: - loop = asyncio.get_event_loop() - if loop.is_running(): - # Schedule the cleanup coroutine - loop.create_task(close_litellm_async_clients()) - else: - # Run the cleanup coroutine - loop.run_until_complete(close_litellm_async_clients()) - except Exception: - # If we can't get an event loop or it's already closed, try creating a new one + # Always create a fresh event loop at exit time + # Don't use get_event_loop() - it may be closed or unavailable + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) try: - loop = asyncio.new_event_loop() loop.run_until_complete(close_litellm_async_clients()) + finally: + # Clean up the loop we created loop.close() - except Exception: - # Silently ignore errors during cleanup - pass + except Exception: + # Silently ignore errors during cleanup to avoid exit handler failures + pass atexit.register(cleanup_wrapper) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 7fdb78c1670..57a6d04c995 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1168,8 +1168,10 @@ def get_async_httpx_client( return _cached_client if params is not None: - params["shared_session"] = shared_session - _new_client = AsyncHTTPHandler(**params) + # Filter out params that are only used for cache key, not for AsyncHTTPHandler.__init__ + handler_params = {k: v for k, v in params.items() if k != "disable_aiohttp_transport"} + handler_params["shared_session"] = shared_session + _new_client = AsyncHTTPHandler(**handler_params) else: _new_client = AsyncHTTPHandler( timeout=httpx.Timeout(timeout=600.0, connect=5.0), @@ -1215,7 +1217,9 @@ def _get_httpx_client(params: Optional[dict] = None) -> HTTPHandler: return _cached_client if params is not None: - _new_client = HTTPHandler(**params) + # Filter out params that are only used for cache key, not for HTTPHandler.__init__ + handler_params = {k: v for k, v in params.items() if k != "disable_aiohttp_transport"} + _new_client = HTTPHandler(**handler_params) else: _new_client = HTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0)) diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index ce470f04aca..8bcecd35232 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -15,12 +15,14 @@ if TYPE_CHECKING: from aiohttp import ClientSession import litellm +from litellm._logging import verbose_logger from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.http_handler import ( _DEFAULT_TTL_FOR_HTTPX_CLIENTS, AsyncHTTPHandler, get_ssl_configuration, ) +from litellm.types.utils import LlmProviders class OpenAIError(BaseLLMException): @@ -203,30 +205,67 @@ class BaseOpenAILLM: if litellm.aclient_session is not None: return litellm.aclient_session - # Get unified SSL configuration - ssl_config = get_ssl_configuration() + # Use the global cached client system to prevent memory leaks (issue #14540) + # This routes through get_async_httpx_client() which provides TTL-based caching + from litellm.llms.custom_httpx.http_handler import get_async_httpx_client - return httpx.AsyncClient( - verify=ssl_config, - transport=AsyncHTTPHandler._create_async_transport( - ssl_context=ssl_config - if isinstance(ssl_config, ssl.SSLContext) - else None, - ssl_verify=ssl_config if isinstance(ssl_config, bool) else None, + try: + # Get SSL config and include in params for proper cache key + ssl_config = get_ssl_configuration() + params = {"ssl_verify": ssl_config} if ssl_config is not None else {} + params["disable_aiohttp_transport"] = litellm.disable_aiohttp_transport + + # Get a cached AsyncHTTPHandler which manages the httpx.AsyncClient + cached_handler = get_async_httpx_client( + llm_provider=LlmProviders.OPENAI, # Cache key includes provider + params=params, # Include SSL config in cache key shared_session=shared_session, - ), - follow_redirects=True, - ) + ) + # Return the underlying httpx client from the handler + return cached_handler.client + except (ImportError, AttributeError, KeyError) as e: + # Fallback to creating a client directly if caching system unavailable + # This preserves backwards compatibility + verbose_logger.debug( + f"Client caching unavailable ({type(e).__name__}), using direct client creation" + ) + ssl_config = get_ssl_configuration() + return httpx.AsyncClient( + verify=ssl_config, + transport=AsyncHTTPHandler._create_async_transport( + ssl_context=ssl_config + if isinstance(ssl_config, ssl.SSLContext) + else None, + ssl_verify=ssl_config if isinstance(ssl_config, bool) else None, + shared_session=shared_session, + ), + follow_redirects=True, + ) @staticmethod def _get_sync_http_client() -> Optional[httpx.Client]: if litellm.client_session is not None: return litellm.client_session - # Get unified SSL configuration - ssl_config = get_ssl_configuration() + # Use the global cached client system to prevent memory leaks (issue #14540) + from litellm.llms.custom_httpx.http_handler import _get_httpx_client - return httpx.Client( - verify=ssl_config, - follow_redirects=True, - ) + try: + # Get SSL config and include in params for proper cache key + ssl_config = get_ssl_configuration() + params = {"ssl_verify": ssl_config} if ssl_config is not None else None + + # Get a cached HTTPHandler which manages the httpx.Client + cached_handler = _get_httpx_client(params=params) + # Return the underlying httpx client from the handler + return cached_handler.client + except (ImportError, AttributeError, KeyError) as e: + # Fallback to creating a client directly if caching system unavailable + verbose_logger.debug( + f"Client caching unavailable ({type(e).__name__}), using direct client creation" + ) + ssl_config = get_ssl_configuration() + return httpx.Client( + verify=ssl_config, + follow_redirects=True, + ) diff --git a/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py b/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py index 2cb2ac9ed8f..d82c2bebb7f 100644 --- a/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py +++ b/litellm/llms/vertex_ai/multimodal_embeddings/transformation.py @@ -265,7 +265,7 @@ class VertexAIMultimodalEmbeddingConfig(BaseEmbeddingConfig): image_count += 1 ## Calculate video embeddings usage - video_length_seconds = 0 + video_length_seconds = 0.0 for prediction in vertex_predictions["predictions"]: video_embeddings = prediction.get("videoEmbeddings") if video_embeddings: diff --git a/litellm/llms/volcengine/__init__.py b/litellm/llms/volcengine/__init__.py index 0887937bed5..fc0098e84d9 100644 --- a/litellm/llms/volcengine/__init__.py +++ b/litellm/llms/volcengine/__init__.py @@ -1,6 +1,6 @@ """ Volcengine LLM Provider -Support for Volcengine (ByteDance) chat and embedding models +Support for Volcengine (ByteDance) chat, embedding, and responses models. """ from .chat.transformation import VolcEngineChatConfig @@ -10,6 +10,7 @@ from .common_utils import ( get_volcengine_headers, ) from .embedding import VolcEngineEmbeddingConfig +from .responses.transformation import VolcEngineResponsesAPIConfig # For backward compatibility, keep the old class name VolcEngineConfig = VolcEngineChatConfig @@ -18,6 +19,7 @@ __all__ = [ "VolcEngineChatConfig", "VolcEngineConfig", # backward compatibility "VolcEngineEmbeddingConfig", + "VolcEngineResponsesAPIConfig", "VolcEngineError", "get_volcengine_base_url", "get_volcengine_headers", diff --git a/litellm/llms/volcengine/responses/transformation.py b/litellm/llms/volcengine/responses/transformation.py new file mode 100644 index 00000000000..872c8dcf118 --- /dev/null +++ b/litellm/llms/volcengine/responses/transformation.py @@ -0,0 +1,557 @@ +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Literal, + Optional, + Tuple, + Union, + get_args, + get_origin, +) + +import httpx +from pydantic import fields as pyd_fields + +import litellm +from litellm._logging import verbose_logger +from litellm.types.llms.openai import ResponseInputParam, ResponsesAPIStreamingResponse +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig +from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( + _safe_convert_created_field, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import ( + ResponsesAPIOptionalRequestParams, + ResponsesAPIResponse, +) +from litellm.types.responses.main import DeleteResponseResult +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders + +from ..common_utils import ( + VolcEngineError, + get_volcengine_base_url, + get_volcengine_headers, +) + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig): + _SUPPORTED_OPTIONAL_PARAMS: List[str] = [ + # Doc-listed knobs + "instructions", + "max_output_tokens", + "previous_response_id", + "store", + "reasoning", + "stream", + "temperature", + "top_p", + "text", + "tools", + "tool_choice", + "max_tool_calls", + "thinking", + "caching", + "expire_at", + "context_management", + # LiteLLM-internal metadata (not sent to provider) + "metadata", + # Request plumbing helpers + "extra_headers", + "extra_query", + "extra_body", + "timeout", + ] + + @property + def custom_llm_provider(self) -> LlmProviders: + return LlmProviders.VOLCENGINE + + def get_supported_openai_params(self, model: str) -> list: + """ + Volcengine Responses API: only documented parameters are supported. + """ + supported = ["input", "model"] + list(self._SUPPORTED_OPTIONAL_PARAMS) + # Do not advertise internal-only metadata to callers; we still accept and drop it before send. + if "metadata" in supported: + supported.remove("metadata") + return supported + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> VolcEngineError: + typed_headers: httpx.Headers = ( + headers if isinstance(headers, httpx.Headers) else httpx.Headers(headers or {}) + ) + return VolcEngineError( + status_code=status_code, + message=error_message, + headers=typed_headers, + ) + + def validate_environment( + self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams] + ) -> dict: + """ + Build auth headers for Volcengine Responses API. + """ + if litellm_params is None: + litellm_params = GenericLiteLLMParams() + elif isinstance(litellm_params, dict): + litellm_params = GenericLiteLLMParams(**litellm_params) + + api_key = ( + litellm_params.api_key + or litellm.api_key + or get_secret_str("ARK_API_KEY") + or get_secret_str("VOLCENGINE_API_KEY") + ) + + if api_key is None: + raise ValueError( + "Volcengine API key is required. Set ARK_API_KEY / VOLCENGINE_API_KEY or pass api_key." + ) + + return get_volcengine_headers(api_key=api_key, extra_headers=headers) + + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Construct Volcengine Responses API endpoint. + """ + base_url = ( + api_base + or litellm.api_base + or get_secret_str("VOLCENGINE_API_BASE") + or get_secret_str("ARK_API_BASE") + or get_volcengine_base_url() + ) + + base_url = base_url.rstrip("/") + + if base_url.endswith("/responses"): + return base_url + if base_url.endswith("/api/v3"): + return f"{base_url}/responses" + return f"{base_url}/api/v3/responses" + + def map_openai_params( + self, + response_api_optional_params: ResponsesAPIOptionalRequestParams, + model: str, + drop_params: bool, + ) -> Dict: + """ + Volcengine Responses API aligns with OpenAI parameters. + Remove parameters not supported by the public docs. + """ + params = { + key: value + for key, value in dict(response_api_optional_params).items() + if key in self._SUPPORTED_OPTIONAL_PARAMS + } + + # LiteLLM metadata is internal-only; don't send to provider + params.pop("metadata", None) + + # Volcengine docs do not list parallel_tool_calls; drop it to avoid backend errors. + if "parallel_tool_calls" in params: + verbose_logger.debug( + "Volcengine Responses API: dropping unsupported 'parallel_tool_calls' param." + ) + params.pop("parallel_tool_calls", None) + + return params + + def transform_responses_api_request( + self, + model: str, + input: Union[str, ResponseInputParam], + response_api_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Dict: + """ + Volcengine rejects any undocumented fields (including extra_body). Fail fast + with clear errors and re-filter with the documented whitelist before delegating + to the OpenAI base transformer. + """ + allowed = set(self._SUPPORTED_OPTIONAL_PARAMS) + + sanitized_optional = { + k: v for k, v in response_api_optional_request_params.items() if k in allowed + } + # Ensure metadata never reaches provider + sanitized_optional.pop("metadata", None) + sanitized_optional.pop("parallel_tool_calls", None) + + # If extra_body is provided, filter its keys against the same allowlist to avoid + # leaking unsupported params to the provider. + if isinstance(sanitized_optional.get("extra_body"), dict): + filtered_body = { + k: v for k, v in sanitized_optional["extra_body"].items() if k in allowed + } + if filtered_body: + sanitized_optional["extra_body"] = filtered_body + else: + sanitized_optional.pop("extra_body", None) + + return super().transform_responses_api_request( + model=model, + input=input, + response_api_optional_request_params=sanitized_optional, + litellm_params=litellm_params, + headers=headers, + ) + + def transform_streaming_response( + self, + model: str, + parsed_chunk: dict, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIStreamingResponse: + """ + Volcengine may omit required fields; auto-fill them using event model defaults. + """ + chunk = parsed_chunk + + # Patch missing response.output on response.* events + if isinstance(chunk, dict): + resp = chunk.get("response") + if isinstance(resp, dict) and "output" not in resp: + patched_chunk = dict(chunk) + patched_resp = dict(resp) + patched_resp["output"] = [] + patched_chunk["response"] = patched_resp + chunk = patched_chunk + + event_type = str(chunk.get("type")) if isinstance(chunk, dict) else None + event_pydantic_model = OpenAIResponsesAPIConfig.get_event_model_class( + event_type=event_type + ) + + patched_chunk = self._fill_missing_fields(chunk, event_pydantic_model) + + return event_pydantic_model(**patched_chunk) + + def transform_response_api_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + try: + logging_obj.post_call( + original_response=raw_response.text, + additional_args={"complete_input_dict": {}}, + ) + raw_response_json = raw_response.json() + if "created_at" in raw_response_json: + raw_response_json["created_at"] = _safe_convert_created_field( + raw_response_json["created_at"] + ) + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + try: + response = ResponsesAPIResponse(**raw_response_json) + except Exception: + verbose_logger.debug( + "Volcengine Responses API: falling back to model_construct for response parsing." + ) + response = ResponsesAPIResponse.model_construct(**raw_response_json) + + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + ######################################################### + ########## DELETE RESPONSE API TRANSFORMATION ############## + ######################################################### + def transform_delete_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}" + data: Dict = {} + return url, data + + def transform_delete_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> DeleteResponseResult: + try: + raw_response_json = raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + try: + return DeleteResponseResult(**raw_response_json) + except Exception: + verbose_logger.debug( + "Volcengine Responses API: falling back to model_construct for delete response parsing." + ) + return DeleteResponseResult.model_construct(**raw_response_json) + + ######################################################### + ########## GET RESPONSE API TRANSFORMATION ############### + ######################################################### + def transform_get_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}" + data: Dict = {} + return url, data + + def transform_get_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + try: + raw_response_json = raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + response = ResponsesAPIResponse(**raw_response_json) + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + ######################################################### + ########## LIST INPUT ITEMS TRANSFORMATION ############# + ######################################################### + def transform_list_input_items_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + after: Optional[str] = None, + before: Optional[str] = None, + include: Optional[List[str]] = None, + limit: int = 20, + order: Literal["asc", "desc"] = "desc", + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}/input_items" + params: Dict[str, Any] = {} + if after is not None: + params["after"] = after + if before is not None: + params["before"] = before + if include: + params["include"] = ",".join(include) + if limit is not None: + params["limit"] = limit + if order is not None: + params["order"] = order + return url, params + + def transform_list_input_items_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> Dict: + try: + return raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + ######################################################### + ########## CANCEL RESPONSE API TRANSFORMATION ########## + ######################################################### + def transform_cancel_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}/cancel" + data: Dict = {} + return url, data + + def transform_cancel_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + try: + raw_response_json = raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + response = ResponsesAPIResponse(**raw_response_json) + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + def should_fake_stream( + self, + model: Optional[str], + stream: Optional[bool], + custom_llm_provider: Optional[str] = None, + ) -> bool: + """ + Volcengine Responses API supports native streaming; never fall back to fake stream. + """ + return False + + @staticmethod + def _fill_missing_fields( + chunk: Any, event_model: Any + ) -> Dict[str, Any]: + """ + Heuristically fill missing required fields with safe defaults based on the + event model's field annotations. This keeps parsing tolerant of providers that + omit non-essential fields. + """ + if not isinstance(chunk, dict) or event_model is None: + return chunk + + patched: Dict[str, Any] = dict(chunk) + fields_map = getattr(event_model, "model_fields", {}) or {} + + for name, field in fields_map.items(): + if name in patched: + patched[name] = VolcEngineResponsesAPIConfig._maybe_fill_nested( + patched[name], field.annotation + ) + continue + + # Explicit default or factory + if field.default is not pyd_fields.PydanticUndefined and field.default is not None: + patched[name] = field.default + continue + if ( + field.default_factory is not None + and field.default_factory is not pyd_fields.PydanticUndefined + ): + patched[name] = field.default_factory() + continue + + # Heuristic defaults for missing required fields + patched[name] = VolcEngineResponsesAPIConfig._default_for_annotation( + field.annotation + ) + + return patched + + @staticmethod + def _default_for_annotation(annotation: Any) -> Any: + origin = get_origin(annotation) + args = get_args(annotation) + + if annotation is int: + return 0 + if annotation is list or origin is list: + return [] + if origin is Union: + # Prefer empty list when any option is a list + if any((arg is list or get_origin(arg) is list) for arg in args): + return [] + if type(None) in args: + return None + if origin is Union and type(None) in args: + return None + + # Fallback to None when no safer guess exists + return None + + @staticmethod + def _maybe_fill_nested(value: Any, annotation: Any) -> Any: + """ + Recursively fill nested dict/list structures based on the annotated model. + """ + model_cls = VolcEngineResponsesAPIConfig._pick_model_class(annotation, value) + args = get_args(annotation) + + if isinstance(value, dict) and model_cls is not None: + return VolcEngineResponsesAPIConfig._fill_missing_fields(value, model_cls) + + if isinstance(value, list): + # Attempt to fill list elements if we know the element annotation + elem_ann: Any = args[0] if args else None + if elem_ann is not None: + return [ + VolcEngineResponsesAPIConfig._maybe_fill_nested(v, elem_ann) + for v in value + ] + + return value + + @staticmethod + def _pick_model_class(annotation: Any, value: Any) -> Optional[Any]: + """ + Choose the best-matching Pydantic model class for a nested dict. + """ + candidates: List[Any] = [] + origin = get_origin(annotation) + + if hasattr(annotation, "model_fields"): + candidates.append(annotation) + if origin is Union: + for arg in get_args(annotation): + if hasattr(arg, "model_fields"): + candidates.append(arg) + + if not candidates: + return None + + # Try to match by literal "type" field when available + if isinstance(value, dict): + v_type = value.get("type") + for candidate in candidates: + try: + type_field = candidate.model_fields.get("type") + if type_field is None: + continue + literal_ann = type_field.annotation + if get_origin(literal_ann) is Literal: + literal_values = get_args(literal_ann) + if v_type in literal_values: + return candidate + except Exception: + continue + + # Fall back to the first candidate + return candidates[0] diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 0eea0568f4c..517e5c0d695 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -13465,6 +13465,31 @@ "supports_vision": true, "supports_web_search": true }, + "gemini-2.5-computer-use-preview-10-2025": { + "input_cost_per_token": 1.25e-06, + "input_cost_per_token_above_200k_tokens": 2.5e-06, + "litellm_provider": "vertex_ai-language-models", + "max_images_per_prompt": 3000, + "max_input_tokens": 128000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_200k_tokens": 1.5e-05, + "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/computer-use", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, "gemini-embedding-001": { "input_cost_per_token": 1.5e-07, "litellm_provider": "vertex_ai-embedding-models", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 1a7f05716b3..9b9a988c07a 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -311,6 +311,88 @@ def get_request_route(request: Request) -> str: return request.url.path +def normalize_request_route(route: str) -> str: + """ + Normalize request routes by replacing dynamic path parameters with placeholders. + + This prevents high cardinality in Prometheus metrics by collapsing routes like: + - /v1/responses/1234567890 -> /v1/responses/{response_id} + - /v1/threads/thread_123 -> /v1/threads/{thread_id} + + Args: + route: The request route path + + Returns: + Normalized route with dynamic parameters replaced by placeholders + + Examples: + >>> normalize_request_route("/v1/responses/abc123") + '/v1/responses/{response_id}' + >>> normalize_request_route("/v1/responses/abc123/cancel") + '/v1/responses/{response_id}/cancel' + >>> normalize_request_route("/chat/completions") + '/chat/completions' + """ + # Define patterns for routes with dynamic IDs + # Format: (regex_pattern, replacement_template) + patterns = [ + # Responses API - must come before generic patterns + (r'^(/(?:openai/)?v1/responses)/([^/]+)(/input_items)$', r'\1/{response_id}\3'), + (r'^(/(?:openai/)?v1/responses)/([^/]+)(/cancel)$', r'\1/{response_id}\3'), + (r'^(/(?:openai/)?v1/responses)/([^/]+)$', r'\1/{response_id}'), + (r'^(/responses)/([^/]+)(/input_items)$', r'\1/{response_id}\3'), + (r'^(/responses)/([^/]+)(/cancel)$', r'\1/{response_id}\3'), + (r'^(/responses)/([^/]+)$', r'\1/{response_id}'), + + # Threads API + (r'^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)(/steps)/([^/]+)$', r'\1/{thread_id}\3/{run_id}\5/{step_id}'), + (r'^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)(/steps)$', r'\1/{thread_id}\3/{run_id}\5'), + (r'^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)(/cancel)$', r'\1/{thread_id}\3/{run_id}\5'), + (r'^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)(/submit_tool_outputs)$', r'\1/{thread_id}\3/{run_id}\5'), + (r'^(/(?:openai/)?v1/threads)/([^/]+)(/runs)/([^/]+)$', r'\1/{thread_id}\3/{run_id}'), + (r'^(/(?:openai/)?v1/threads)/([^/]+)(/runs)$', r'\1/{thread_id}\3'), + (r'^(/(?:openai/)?v1/threads)/([^/]+)(/messages)/([^/]+)$', r'\1/{thread_id}\3/{message_id}'), + (r'^(/(?:openai/)?v1/threads)/([^/]+)(/messages)$', r'\1/{thread_id}\3'), + (r'^(/(?:openai/)?v1/threads)/([^/]+)$', r'\1/{thread_id}'), + + # Vector Stores API + (r'^(/(?:openai/)?v1/vector_stores)/([^/]+)(/files)/([^/]+)$', r'\1/{vector_store_id}\3/{file_id}'), + (r'^(/(?:openai/)?v1/vector_stores)/([^/]+)(/files)$', r'\1/{vector_store_id}\3'), + (r'^(/(?:openai/)?v1/vector_stores)/([^/]+)(/file_batches)/([^/]+)$', r'\1/{vector_store_id}\3/{batch_id}'), + (r'^(/(?:openai/)?v1/vector_stores)/([^/]+)(/file_batches)$', r'\1/{vector_store_id}\3'), + (r'^(/(?:openai/)?v1/vector_stores)/([^/]+)$', r'\1/{vector_store_id}'), + + # Assistants API + (r'^(/(?:openai/)?v1/assistants)/([^/]+)$', r'\1/{assistant_id}'), + + # Files API + (r'^(/(?:openai/)?v1/files)/([^/]+)(/content)$', r'\1/{file_id}\3'), + (r'^(/(?:openai/)?v1/files)/([^/]+)$', r'\1/{file_id}'), + + # Batches API + (r'^(/(?:openai/)?v1/batches)/([^/]+)(/cancel)$', r'\1/{batch_id}\3'), + (r'^(/(?:openai/)?v1/batches)/([^/]+)$', r'\1/{batch_id}'), + + # Fine-tuning API + (r'^(/(?:openai/)?v1/fine_tuning/jobs)/([^/]+)(/events)$', r'\1/{fine_tuning_job_id}\3'), + (r'^(/(?:openai/)?v1/fine_tuning/jobs)/([^/]+)(/cancel)$', r'\1/{fine_tuning_job_id}\3'), + (r'^(/(?:openai/)?v1/fine_tuning/jobs)/([^/]+)(/checkpoints)$', r'\1/{fine_tuning_job_id}\3'), + (r'^(/(?:openai/)?v1/fine_tuning/jobs)/([^/]+)$', r'\1/{fine_tuning_job_id}'), + + # Models API + (r'^(/(?:openai/)?v1/models)/([^/]+)$', r'\1/{model}'), + ] + + # Apply patterns in order + for pattern, replacement in patterns: + normalized = re.sub(pattern, replacement, route) + if normalized != route: + return normalized + + # Return original route if no pattern matched + return route + + async def check_if_request_size_is_safe(request: Request) -> bool: """ Enterprise Only: diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index bc0c164a0ad..7e7c7c8c90c 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -28,8 +28,8 @@ from litellm.proxy.auth.auth_checks import ( _delete_cache_key_object, _get_user_role, _is_user_proxy_admin, - _virtual_key_max_budget_check, _virtual_key_max_budget_alert_check, + _virtual_key_max_budget_check, _virtual_key_soft_budget_check, can_key_call_model, common_checks, @@ -45,6 +45,7 @@ from litellm.proxy.auth.auth_utils import ( get_end_user_id_from_request_body, get_model_from_request, get_request_route, + normalize_request_route, pre_db_read_auth_checks, route_in_additonal_public_routes, ) @@ -1261,7 +1262,7 @@ async def user_api_key_auth( if end_user_id is not None: user_api_key_auth_obj.end_user_id = end_user_id - user_api_key_auth_obj.request_route = route + user_api_key_auth_obj.request_route = normalize_request_route(route) return user_api_key_auth_obj diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 4d7f4a5b125..20df54b62ce 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -102,6 +102,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): presidio_score_thresholds or {} ) self.presidio_language = presidio_language or "en" + # Shared HTTP session to prevent memory leaks (issue #14540) + self._http_session: Optional[aiohttp.ClientSession] = None + # Lock to prevent race conditions when creating session under concurrent load + # Note: asyncio.Lock() can be created without an event loop; it only needs one when awaited + self._session_lock: asyncio.Lock = asyncio.Lock() if mock_testing is True: # for testing purposes only return @@ -167,6 +172,47 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): "http://" + self.presidio_anonymizer_api_base ) + async def _get_http_session(self) -> aiohttp.ClientSession: + """ + Get or create the shared HTTP session for Presidio API calls. + + Fixes memory leak (issue #14540) where every guardrail check created + a new aiohttp.ClientSession that was never properly closed. + + Thread-safe: Uses asyncio.Lock to prevent race conditions when + multiple concurrent requests try to create the session simultaneously. + """ + async with self._session_lock: + if self._http_session is None or self._http_session.closed: + self._http_session = aiohttp.ClientSession() + return self._http_session + + async def _close_http_session(self) -> None: + """Close the HTTP session if it exists.""" + if self._http_session is not None and not self._http_session.closed: + await self._http_session.close() + self._http_session = None + + def __del__(self): + """Cleanup: close HTTP session on instance destruction.""" + if self._http_session is not None and not self._http_session.closed: + try: + # Try to close the session, but don't fail if event loop is gone + import asyncio + try: + loop = asyncio.get_event_loop() + if loop.is_running(): + # Schedule cleanup, don't block __del__ + asyncio.create_task(self._close_http_session()) + else: + loop.run_until_complete(self._close_http_session()) + except RuntimeError: + # Event loop is closed, can't clean up - not ideal but better than crashing + pass + except Exception: + # Suppress all exceptions in __del__ to avoid issues during shutdown + pass + def _get_presidio_analyze_request_payload( self, text: str, @@ -223,67 +269,69 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): ) return [] - async with aiohttp.ClientSession() as session: - if self.mock_redacted_text is not None: - return self.mock_redacted_text + if self.mock_redacted_text is not None: + return self.mock_redacted_text - # Make the request to /analyze - analyze_url = f"{self.presidio_analyzer_api_base}analyze" + # Use shared session to prevent memory leak (issue #14540) + session = await self._get_http_session() - analyze_payload: PresidioAnalyzeRequest = ( - self._get_presidio_analyze_request_payload( - text=text, - presidio_config=presidio_config, - request_data=request_data, - ) + # Make the request to /analyze + analyze_url = f"{self.presidio_analyzer_api_base}analyze" + + analyze_payload: PresidioAnalyzeRequest = ( + self._get_presidio_analyze_request_payload( + text=text, + presidio_config=presidio_config, + request_data=request_data, ) + ) - verbose_proxy_logger.debug( - "Making request to: %s with payload: %s", - analyze_url, - analyze_payload, - ) + verbose_proxy_logger.debug( + "Making request to: %s with payload: %s", + analyze_url, + analyze_payload, + ) - async with session.post(analyze_url, json=analyze_payload) as response: - analyze_results = await response.json() - verbose_proxy_logger.debug("analyze_results: %s", analyze_results) + async with session.post(analyze_url, json=analyze_payload) as response: + analyze_results = await response.json() + verbose_proxy_logger.debug("analyze_results: %s", analyze_results) - # Handle error responses from Presidio (e.g., {'error': 'No text provided'}) - # Presidio may return a dict instead of a list when errors occur - if isinstance(analyze_results, dict): - if "error" in analyze_results: - verbose_proxy_logger.warning( - "Presidio analyzer returned error: %s, returning empty list", - analyze_results.get("error") - ) - return [] - # If it's a dict but not an error, try to process it as a single item - verbose_proxy_logger.debug( - "Presidio returned dict (not list), attempting to process as single item" + # Handle error responses from Presidio (e.g., {'error': 'No text provided'}) + # Presidio may return a dict instead of a list when errors occur + if isinstance(analyze_results, dict): + if "error" in analyze_results: + verbose_proxy_logger.warning( + "Presidio analyzer returned error: %s, returning empty list", + analyze_results.get("error") ) - try: - return [PresidioAnalyzeResponseItem(**analyze_results)] - except Exception as e: - verbose_proxy_logger.warning( - "Failed to parse Presidio dict response: %s, returning empty list", - e - ) - return [] + return [] + # If it's a dict but not an error, try to process it as a single item + verbose_proxy_logger.debug( + "Presidio returned dict (not list), attempting to process as single item" + ) + try: + return [PresidioAnalyzeResponseItem(**analyze_results)] + except Exception as e: + verbose_proxy_logger.warning( + "Failed to parse Presidio dict response: %s, returning empty list", + e + ) + return [] - # Normal case: list of results - final_results = [] - for item in analyze_results: - try: - final_results.append(PresidioAnalyzeResponseItem(**item)) - except TypeError as te: - # Handle case where item is not a dict (shouldn't happen, but be defensive) - verbose_proxy_logger.warning( - "Skipping invalid Presidio result item: %s (error: %s)", - item, - te, - ) - continue - return final_results + # Normal case: list of results + final_results = [] + for item in analyze_results: + try: + final_results.append(PresidioAnalyzeResponseItem(**item)) + except TypeError as te: + # Handle case where item is not a dict (shouldn't happen, but be defensive) + verbose_proxy_logger.warning( + "Skipping invalid Presidio result item: %s (error: %s)", + item, + te, + ) + continue + return final_results except Exception as e: raise e @@ -302,46 +350,48 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if isinstance(analyze_results, list) and len(analyze_results) == 0: return text - async with aiohttp.ClientSession() as session: - # Make the request to /anonymize - anonymize_url = f"{self.presidio_anonymizer_api_base}anonymize" - verbose_proxy_logger.debug("Making request to: %s", anonymize_url) - anonymize_payload = { - "text": text, - "analyzer_results": analyze_results, - } + # Use shared session to prevent memory leak (issue #14540) + session = await self._get_http_session() - async with session.post( - anonymize_url, json=anonymize_payload - ) as response: - redacted_text = await response.json() + # Make the request to /anonymize + anonymize_url = f"{self.presidio_anonymizer_api_base}anonymize" + verbose_proxy_logger.debug("Making request to: %s", anonymize_url) + anonymize_payload = { + "text": text, + "analyzer_results": analyze_results, + } - new_text = text - if redacted_text is not None: - verbose_proxy_logger.debug("redacted_text: %s", redacted_text) - for item in redacted_text["items"]: - start = item["start"] - end = item["end"] - replacement = item["text"] # replacement token - if item["operator"] == "replace" and output_parse_pii is True: - # check if token in dict - # if exists, add a uuid to the replacement token for swapping back to the original text in llm response output parsing - if replacement in self.pii_tokens: - replacement = replacement + str(uuid.uuid4()) + async with session.post( + anonymize_url, json=anonymize_payload + ) as response: + redacted_text = await response.json() - self.pii_tokens[replacement] = new_text[ - start:end - ] # get text it'll replace + new_text = text + if redacted_text is not None: + verbose_proxy_logger.debug("redacted_text: %s", redacted_text) + for item in redacted_text["items"]: + start = item["start"] + end = item["end"] + replacement = item["text"] # replacement token + if item["operator"] == "replace" and output_parse_pii is True: + # check if token in dict + # if exists, add a uuid to the replacement token for swapping back to the original text in llm response output parsing + if replacement in self.pii_tokens: + replacement = replacement + str(uuid.uuid4()) - new_text = new_text[:start] + replacement + new_text[end:] - entity_type = item.get("entity_type", None) - if entity_type is not None: - masked_entity_count[entity_type] = ( - masked_entity_count.get(entity_type, 0) + 1 - ) - return redacted_text["text"] - else: - raise Exception(f"Invalid anonymizer response: {redacted_text}") + self.pii_tokens[replacement] = new_text[ + start:end + ] # get text it'll replace + + new_text = new_text[:start] + replacement + new_text[end:] + entity_type = item.get("entity_type", None) + if entity_type is not None: + masked_entity_count[entity_type] = ( + masked_entity_count.get(entity_type, 0) + 1 + ) + return redacted_text["text"] + else: + raise Exception(f"Invalid anonymizer response: {redacted_text}") except Exception as e: raise e diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index def2f72437d..dd7936059a0 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -21,6 +21,8 @@ from litellm.types.llms.openai import ( OutputTextAnnotationAddedEvent, OutputTextDeltaEvent, OutputTextDoneEvent, + FunctionCallArgumentsDeltaEvent, + FunctionCallArgumentsDoneEvent, ReasoningSummaryTextDeltaEvent, ResponseCompletedEvent, ResponseCreatedEvent, @@ -79,6 +81,161 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): Union[ModelResponse, TextCompletionResponse] ] = None self.final_text: str = "" + self._pending_tool_events: List[BaseLiteLLMOpenAIResponseObject] = [] + self._tool_output_index_by_call_id: dict[str, int] = {} + self._tool_args_by_call_id: dict[str, str] = {} + self._next_tool_output_index: int = 1 # output_index=0 reserved for the message item + self._final_tool_events_queued: bool = False + + def _get_or_assign_tool_output_index(self, call_id: str) -> int: + existing = self._tool_output_index_by_call_id.get(call_id) + if existing is not None: + return existing + idx = self._next_tool_output_index + self._next_tool_output_index += 1 + self._tool_output_index_by_call_id[call_id] = idx + return idx + + def _queue_tool_call_delta_events(self, tool_calls: object) -> None: + """ + Convert chat-completions streaming `tool_calls` deltas into Responses API streaming events. + + We emit: + - response.output_item.added (function_call) + - response.function_call_arguments.delta + """ + if not isinstance(tool_calls, list): + return + + for tc in tool_calls: + call_id_raw = tc.get("id") if isinstance(tc, dict) else getattr(tc, "id", None) + if not call_id_raw: + continue + call_id = str(call_id_raw) + + fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None) + fn_name = "" + fn_args_delta = "" + if isinstance(fn, dict): + fn_name = str(fn.get("name") or "") + fn_args_delta = str(fn.get("arguments") or "") + else: + fn_name = str(getattr(fn, "name", "") or "") + fn_args_delta = str(getattr(fn, "arguments", "") or "") + + output_index = self._get_or_assign_tool_output_index(call_id) + + if call_id not in self._tool_args_by_call_id: + self._tool_args_by_call_id[call_id] = "" + self._pending_tool_events.append( + OutputItemAddedEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=output_index, + item=BaseLiteLLMOpenAIResponseObject( + **{ + "type": "function_call", + "id": call_id, + "call_id": call_id, + "name": fn_name, + "arguments": "", + "status": "in_progress", + } + ), + ) + ) + + if fn_args_delta: + self._tool_args_by_call_id[call_id] += fn_args_delta + self._pending_tool_events.append( + FunctionCallArgumentsDeltaEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + item_id=call_id, + output_index=output_index, + delta=fn_args_delta, + ) + ) + + def _queue_final_tool_call_done_events(self, litellm_complete_object: ModelResponse) -> None: + """ + Ensure tool calls that were not streamed as deltas still get emitted before response.completed. + """ + if self._final_tool_events_queued: + return + self._final_tool_events_queued = True + + try: + message = litellm_complete_object.choices[0].message # type: ignore + tool_calls = getattr(message, "tool_calls", None) + except Exception: + tool_calls = None + + if not tool_calls or not isinstance(tool_calls, list): + return + + for tc in tool_calls: + call_id_raw = tc.get("id") if isinstance(tc, dict) else getattr(tc, "id", None) + if not call_id_raw: + continue + call_id = str(call_id_raw) + output_index = self._get_or_assign_tool_output_index(call_id) + + fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None) + fn_name = "" + fn_args = "" + if isinstance(fn, dict): + fn_name = str(fn.get("name") or "") + fn_args = str(fn.get("arguments") or "") + else: + fn_name = str(getattr(fn, "name", "") or "") + fn_args = str(getattr(fn, "arguments", "") or "") + + # If we never sent output_item.added for this call_id, emit it now. + if call_id not in self._tool_args_by_call_id: + self._tool_args_by_call_id[call_id] = "" + self._pending_tool_events.append( + OutputItemAddedEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=output_index, + item=BaseLiteLLMOpenAIResponseObject( + **{ + "type": "function_call", + "id": call_id, + "call_id": call_id, + "name": fn_name, + "arguments": "", + "status": "in_progress", + } + ), + ) + ) + + final_args = fn_args or self._tool_args_by_call_id.get(call_id, "") + self._pending_tool_events.append( + FunctionCallArgumentsDoneEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, + item_id=call_id, + output_index=output_index, + arguments=final_args, + ) + ) + + self._pending_tool_events.append( + OutputItemDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + output_index=output_index, + sequence_number=1, + item=BaseLiteLLMOpenAIResponseObject( + **{ + "type": "function_call", + "id": call_id, + "call_id": call_id, + "name": fn_name, + "arguments": final_args, + "status": "completed", + } + ), + ) + ) def _default_response_created_event_data(self) -> dict: response_created_event_data = { @@ -310,6 +467,12 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): ): self.litellm_model_response = self.create_litellm_model_response() if self.litellm_model_response: + # If tool calls exist, emit tool events before finishing/response.completed. + if isinstance(self.litellm_model_response, ModelResponse): + self._queue_final_tool_call_done_events(self.litellm_model_response) + if self._pending_tool_events: + return self._pending_tool_events.pop(0) + done_event = self.return_default_done_events(self.litellm_model_response) if done_event: return done_event @@ -462,13 +625,27 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): content_index=0, delta=delta_content, ) + + # Priority 3: Handle tool call deltas (if any) -> queue events and emit them + if ( + chunk.choices + and hasattr(chunk.choices[0].delta, "tool_calls") + and chunk.choices[0].delta.tool_calls + ): + self._queue_tool_call_delta_events(chunk.choices[0].delta.tool_calls) + if self._pending_tool_events: + return self._pending_tool_events.pop(0) - # Priority 3: If we have pending annotation events, emit the next one + # Priority 4: If we have pending annotation events, emit the next one # This happens when the current chunk has no text/reasoning content if hasattr(self, '_pending_annotation_events') and self._pending_annotation_events: event = self._pending_annotation_events.pop(0) return event + # Priority 5: If we have pending tool events (from earlier chunk), emit the next one + if self._pending_tool_events: + return self._pending_tool_events.pop(0) + return None def _get_delta_string_from_streaming_choices( diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index eaa80c6cfe4..3badbc50578 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -367,14 +367,6 @@ class LiteLLMCompletionResponsesConfig: ChatCompletionResponseMessage, ] ] = [] - tool_call_output_messages: List[ - Union[ - AllMessageValues, - GenericChatCompletionMessage, - ChatCompletionMessageToolCall, - ChatCompletionResponseMessage, - ] - ] = [] if isinstance(input, str): messages.append(ChatCompletionUserMessage(role="user", content=input)) @@ -385,15 +377,6 @@ class LiteLLMCompletionResponsesConfig: input_item=_input ) - ######################################################### - # If Input Item is a Tool Call Output, add it to the tool_call_output_messages list - ######################################################### - if LiteLLMCompletionResponsesConfig._is_input_item_tool_call_output( - input_item=_input - ): - tool_call_output_messages.extend(chat_completion_messages) - continue - if LiteLLMCompletionResponsesConfig._is_input_item_function_call( input_item=_input ): @@ -401,15 +384,57 @@ class LiteLLMCompletionResponsesConfig: if call_id_raw: existing_tool_call_ids.add(str(call_id_raw)) - messages.extend(chat_completion_messages) + ######################################################### + # If Input Item is a Tool Call Output, add it to the tool_call_output_messages list + # preserving the ordering of tool call outputs. Some models require the tool + # result to immediately follow the assistant tool call. + ######################################################### + if LiteLLMCompletionResponsesConfig._is_input_item_tool_call_output( + input_item=_input + ): + if not chat_completion_messages: + continue - deduped_tool_call_messages = ( - LiteLLMCompletionResponsesConfig._deduplicate_tool_call_output_messages( - tool_call_output_messages=tool_call_output_messages, - existing_tool_call_ids=existing_tool_call_ids, - ) - ) - messages.extend(deduped_tool_call_messages) + deduped_in_place: List[Any] = [] + for m in chat_completion_messages: + role = "" + if isinstance(m, dict): + role = str(m.get("role") or "") + else: + role = str(getattr(m, "role", "") or "") + + # Drop assistant tool_calls wrappers if we already have this call_id + if role == "assistant": + tool_calls: Any = ( + m.get("tool_calls") + if isinstance(m, dict) + else getattr(m, "tool_calls", None) + ) + call_id = "" + if ( + isinstance(tool_calls, Sequence) + and not isinstance(tool_calls, (str, bytes)) + and len(tool_calls) > 0 + ): + first_call = tool_calls[0] + call_id_raw = ( + first_call.get("id") + if isinstance(first_call, dict) + else getattr(first_call, "id", None) + ) + if call_id_raw: + call_id = str(call_id_raw) + if call_id and call_id in existing_tool_call_ids: + continue + if call_id: + existing_tool_call_ids.add(call_id) + + deduped_in_place.append(m) + + messages.extend(deduped_in_place) + continue + + messages.extend(chat_completion_messages) return messages @staticmethod @@ -821,10 +846,82 @@ class LiteLLMCompletionResponsesConfig: # Empty call_id means we can't create a valid tool message if not call_id: return [] - + + def _normalize_function_call_output_to_tool_content( + output: Any, + ) -> Any: + """ + Normalize Responses API function_call_output.output into a shape that downstream + chat adapters (esp. Gemini) can reliably consume. + + OpenAI Responses API typically uses: + - output: string + + Some clients/adapters send: + - output: [{"type": "input_text", "text": "..."}, {"type": "input_image", ...}] + + For chat tool messages we normalize to either: + - string (preferred) + - list of {"type": "text"|"image_url", ...} blocks (for multimodal tool outputs) + """ + if output is None: + return "" + if isinstance(output, str): + return output + + # Some adapters represent tool output as a list of "input_*" parts + if isinstance(output, list): + normalized_blocks: List[Dict[str, Any]] = [] + text_acc: List[str] = [] + for part in output: + if not isinstance(part, dict): + continue + part_type = part.get("type") + if part_type in ("input_text", "output_text", "text"): + txt = part.get("text") + if isinstance(txt, str) and txt: + text_acc.append(txt) + normalized_blocks.append({"type": "text", "text": txt}) + elif part_type in ("input_image", "image_url"): + image_url_val = part.get("image_url") or part.get("url") + if isinstance(image_url_val, dict): + url = image_url_val.get("url") + if isinstance(url, str) and url: + normalized_blocks.append( + {"type": "image_url", "image_url": {"url": url}} + ) + elif isinstance(image_url_val, str) and image_url_val: + normalized_blocks.append( + {"type": "image_url", "image_url": {"url": image_url_val}} + ) + + # Prefer structured blocks if we have images; otherwise return a string. + if any(b.get("type") == "image_url" for b in normalized_blocks): + # Ensure we include any accumulated text as text blocks too + return normalized_blocks + if text_acc: + return "".join(text_acc) + try: + # last resort: keep something meaningful for providers that require a string + import json as _json + + return _json.dumps(output) + except Exception: + return str(output) + + # Fallback for dict/number/etc. + try: + import json as _json + + return _json.dumps(output) + except Exception: + return str(output) + tool_output_message = ChatCompletionToolMessage( role="tool", - content=tool_call_output.get("output") or "", + content=_normalize_function_call_output_to_tool_content( + tool_call_output.get("output") + ), tool_call_id=str(call_id), ) diff --git a/litellm/router_utils/prompt_caching_cache.py b/litellm/router_utils/prompt_caching_cache.py index dbf8b8fcba8..69698f282b1 100644 --- a/litellm/router_utils/prompt_caching_cache.py +++ b/litellm/router_utils/prompt_caching_cache.py @@ -76,6 +76,22 @@ class PromptCachingCache: for msg_idx, message in enumerate(messages): content = message.get("content") + + # Check for cache_control at message level (when content is a string) + # This handles the case where cache_control is a sibling of string content: + # {"role": "user", "content": "...", "cache_control": {"type": "ephemeral"}} + message_level_cache_control = message.get("cache_control") + if ( + message_level_cache_control is not None + and isinstance(message_level_cache_control, dict) + and message_level_cache_control.get("type") == "ephemeral" + ): + last_cacheable_message_idx = msg_idx + # Set to None to indicate the entire message content is cacheable + # (not a specific content block index within a list) + last_cacheable_content_idx = None + + # Also check for cache_control within content blocks (when content is a list) if not isinstance(content, list): continue diff --git a/litellm/scheduler.py b/litellm/scheduler.py index 3225ba0451c..5f3dd4cbf61 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -84,6 +84,7 @@ class Scheduler: if queue[0][1] == id: # Remove the item from the queue heapq.heappop(queue) + await self.save_queue(queue=queue, model_name=model_name) print_verbose(f"Popped id: {id}") return True else: diff --git a/litellm/utils.py b/litellm/utils.py index ee31d8846bc..3e88c6fe9e3 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4186,14 +4186,21 @@ def get_optional_params( # noqa: PLR0915 ), ) elif "anthropic" in bedrock_base_model and bedrock_route == "invoke": - # Check for Claude 3+ models (Messages API) including regional prefixes and Claude 4 - # Models like eu.anthropic.claude-opus-4-5, us.anthropic.claude-3-5-sonnet, etc. - bedrock_base_model_lower = bedrock_base_model.lower() - is_messages_api_model = any( - indicator in bedrock_base_model_lower - for indicator in ["claude-3", "claude-opus-4", "claude-sonnet-4", "claude-haiku-4"] - ) - if is_messages_api_model: + if ( + bedrock_base_model + in litellm.AmazonAnthropicConfig.get_legacy_anthropic_model_names() + ): + optional_params = litellm.AmazonAnthropicConfig().map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=model, + drop_params=( + drop_params + if drop_params is not None and isinstance(drop_params, bool) + else False + ), + ) + else: optional_params = ( litellm.AmazonAnthropicClaudeConfig().map_openai_params( non_default_params=non_default_params, @@ -4206,18 +4213,6 @@ def get_optional_params( # noqa: PLR0915 ), ) ) - - else: - optional_params = litellm.AmazonAnthropicConfig().map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=model, - drop_params=( - drop_params - if drop_params is not None and isinstance(drop_params, bool) - else False - ), - ) elif provider_config is not None: optional_params = provider_config.map_openai_params( non_default_params=non_default_params, @@ -4652,6 +4647,8 @@ def add_provider_specific_params_to_optional_params( else: for k in passed_params.keys(): if k not in openai_params and passed_params[k] is not None: + if _should_drop_param(k=k, additional_drop_params=additional_drop_params): + continue optional_params[k] = passed_params[k] return optional_params @@ -5592,6 +5589,13 @@ def _get_model_info_helper( # noqa: PLR0915 input_cost_per_image_token=_model_info.get( "input_cost_per_image_token", None ), + input_cost_per_image=_model_info.get("input_cost_per_image", None), + input_cost_per_audio_per_second=_model_info.get( + "input_cost_per_audio_per_second", None + ), + input_cost_per_video_per_second=_model_info.get( + "input_cost_per_video_per_second", None + ), input_cost_per_token_batches=_model_info.get( "input_cost_per_token_batches" ), @@ -8147,6 +8151,8 @@ class ProviderConfigManager: # Note: GPT models (gpt-3.5, gpt-4, gpt-5, etc.) support temperature parameter # O-series models (o1, o3) do not contain "gpt" and have different parameter restrictions is_gpt_model = model and "gpt" in model.lower() + is_o_series = model and ("o_series" in model.lower() or (supports_reasoning(model) and not is_gpt_model)) + is_o_series = model and ( "o_series" in model.lower() or (supports_reasoning(model) and not is_gpt_model) @@ -8164,6 +8170,8 @@ class ProviderConfigManager: return litellm.ChatGPTResponsesAPIConfig() elif litellm.LlmProviders.LITELLM_PROXY == provider: return litellm.LiteLLMProxyResponsesAPIConfig() + elif litellm.LlmProviders.VOLCENGINE == provider: + return litellm.VolcEngineResponsesAPIConfig() elif litellm.LlmProviders.MANUS == provider: return litellm.ManusResponsesAPIConfig() return None diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index ac2689a74cf..135b0d46ed0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -13465,6 +13465,31 @@ "supports_vision": true, "supports_web_search": true }, + "gemini-2.5-computer-use-preview-10-2025": { + "input_cost_per_token": 1.25e-06, + "input_cost_per_token_above_200k_tokens": 2.5e-06, + "litellm_provider": "vertex_ai-language-models", + "max_images_per_prompt": 3000, + "max_input_tokens": 128000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_200k_tokens": 1.5e-05, + "source": "https://cloud.google.com/vertex-ai/generative-ai/docs/computer-use", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, "gemini-embedding-001": { "input_cost_per_token": 1.5e-07, "litellm_provider": "vertex_ai-embedding-models", diff --git a/tests/guardrails_tests/test_tracing_guardrails.py b/tests/guardrails_tests/test_tracing_guardrails.py index 02ff7c0e4f6..8e7ce27bc28 100644 --- a/tests/guardrails_tests/test_tracing_guardrails.py +++ b/tests/guardrails_tests/test_tracing_guardrails.py @@ -74,7 +74,7 @@ async def test_standard_logging_payload_includes_guardrail_information(): class MockClientSession: def __init__(self): - pass + self.closed = False async def __aenter__(self): return self @@ -82,6 +82,9 @@ async def test_standard_logging_payload_includes_guardrail_information(): async def __aexit__(self, exc_type, exc_val, exc_tb): pass + async def close(self): + self.closed = True + def post(self, url, json=None): class MockResponse: def __init__(self, response_obj): diff --git a/tests/llm_translation/test_optional_params.py b/tests/llm_translation/test_optional_params.py index bc85b99eee7..95700eb29b9 100644 --- a/tests/llm_translation/test_optional_params.py +++ b/tests/llm_translation/test_optional_params.py @@ -1471,6 +1471,260 @@ def test_bedrock_invoke_anthropic_max_tokens(): assert optional_params["max_tokens"] == 1024 +def test_bedrock_invoke_claude_4_anthropic_max_tokens(): + passed_params = { + "model": "invoke/us.anthropic.claude-sonnet-4-5-20250929-v1:0", + "functions": None, + "function_call": None, + "temperature": 0.8, + "top_p": None, + "n": 1, + "stream": False, + "stream_options": None, + "stop": None, + "max_tokens": None, + "max_completion_tokens": 1024, + "modalities": None, + "prediction": None, + "audio": None, + "presence_penalty": None, + "frequency_penalty": None, + "logit_bias": None, + "user": None, + "custom_llm_provider": "bedrock", + "response_format": {"type": "text"}, + "seed": None, + "tools": [ + { + "type": "function", + "function": { + "name": "generate_plan", + "description": "Generate a plan to execute the task using only the tools outlined in your context.", + "input_schema": { + "type": "object", + "properties": { + "steps": { + "type": "array", + "items": { + "type": "object", + "properties": { + "type": { + "type": "string", + "description": "The type of step to execute", + }, + "tool_name": { + "type": "string", + "description": "The name of the tool to use for this step", + }, + "tool_input": { + "type": "object", + "description": "The input to pass to the tool. Make sure this complies with the schema for the tool.", + }, + "tool_output": { + "type": "object", + "description": "(Optional) The output from the tool if needed for future steps. Make sure this complies with the schema for the tool.", + }, + }, + "required": ["type"], + }, + } + }, + }, + }, + }, + { + "type": "function", + "function": { + "name": "generate_wire_tool", + "description": "Create a wire transfer with complete wire instructions", + "input_schema": { + "type": "object", + "properties": { + "company_id": { + "type": "integer", + "description": "The ID of the company receiving the investment", + }, + "investment_id": { + "type": "integer", + "description": "The ID of the investment memo", + }, + "dollar_amount": { + "type": "number", + "description": "The amount to wire in USD", + }, + "wiring_instructions": { + "type": "object", + "description": "Complete bank account and routing information for the wire", + "properties": { + "account_name": { + "type": "string", + "description": "Name on the bank account", + }, + "address_1": { + "type": "string", + "description": "Primary address line", + }, + "address_2": { + "type": "string", + "description": "Secondary address line (optional)", + }, + "city": {"type": "string"}, + "state": {"type": "string"}, + "zip": {"type": "string"}, + "country": {"type": "string", "default": "US"}, + "bank_name": {"type": "string"}, + "account_number": {"type": "string"}, + "routing_number": {"type": "string"}, + "account_type": { + "type": "string", + "enum": ["checking", "savings"], + "default": "checking", + }, + "swift_code": { + "type": "string", + "description": "Required for international wires", + }, + "iban": { + "type": "string", + "description": "Required for some international wires", + }, + "bank_city": {"type": "string"}, + "bank_state": {"type": "string"}, + "bank_country": {"type": "string", "default": "US"}, + "bank_to_bank_instructions": { + "type": "string", + "description": "Additional instructions for the bank (optional)", + }, + "intermediary_bank_name": { + "type": "string", + "description": "Name of intermediary bank if required (optional)", + }, + }, + "required": [ + "account_name", + "address_1", + "country", + "bank_name", + "account_number", + "routing_number", + "account_type", + "bank_country", + ], + }, + }, + "required": [ + "company_id", + "investment_id", + "dollar_amount", + "wiring_instructions", + ], + }, + }, + }, + { + "type": "function", + "function": { + "name": "search_companies", + "description": "Search for companies by name or other criteria to get their IDs", + "input_schema": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "Name or part of name to search for", + }, + "batch": { + "type": "string", + "description": 'Optional batch filter (e.g., "W21", "S22")', + }, + "status": { + "type": "string", + "enum": [ + "live", + "dead", + "adrift", + "exited", + "went_public", + "all", + ], + "description": "Filter by company status", + "default": "live", + }, + "limit": { + "type": "integer", + "description": "Maximum number of results to return", + "default": 10, + }, + }, + "required": ["query"], + }, + "output_schema": { + "type": "object", + "properties": { + "status": { + "type": "string", + "description": "Success or error status", + }, + "results": { + "type": "array", + "description": "List of companies matching the search criteria", + "items": { + "type": "object", + "properties": { + "id": { + "type": "integer", + "description": "Company ID to use in other API calls", + }, + "name": {"type": "string"}, + "batch": {"type": "string"}, + "status": {"type": "string"}, + "valuation": {"type": "string"}, + "url": {"type": "string"}, + "description": {"type": "string"}, + "founders": {"type": "string"}, + }, + }, + }, + "results_count": { + "type": "integer", + "description": "Number of companies returned", + }, + "total_matches": { + "type": "integer", + "description": "Total number of matches found", + }, + }, + }, + }, + }, + ], + "tool_choice": None, + "max_retries": 0, + "logprobs": None, + "top_logprobs": None, + "extra_headers": None, + "api_version": None, + "parallel_tool_calls": None, + "drop_params": True, + "reasoning_effort": None, + "additional_drop_params": None, + "messages": [ + { + "role": "system", + "content": "You are an AI assistant that helps prepare a wire for a pro rata investment.", + }, + {"role": "user", "content": [{"type": "text", "text": "hi"}]}, + ], + "thinking": None, + "kwargs": {}, + } + optional_params = get_optional_params(**passed_params) + print(f"optional_params: {optional_params}") + + assert "max_tokens_to_sample" not in optional_params + assert optional_params["max_tokens"] == 1024 + + def test_azure_modalities_param(): optional_params = get_optional_params( model="chatgpt-v2", diff --git a/tests/local_testing/test_azure_openai.py b/tests/local_testing/test_azure_openai.py index ed0dda3e15a..e95c1b6fcce 100644 --- a/tests/local_testing/test_azure_openai.py +++ b/tests/local_testing/test_azure_openai.py @@ -41,6 +41,11 @@ async def test_aaaaazure_tenant_id_auth(respx_mock: MockRouter): PROD Test """ litellm.disable_aiohttp_transport = True # since this uses respx, we need to set use_aiohttp_transport to False + + # Clear the HTTP client cache to ensure respx mocking works + # This is critical because respx only intercepts clients created AFTER mocking is active + if hasattr(litellm, 'in_memory_llm_clients_cache'): + litellm.in_memory_llm_clients_cache.flush_cache() router = Router( model_list=[ diff --git a/tests/local_testing/test_scheduler.py b/tests/local_testing/test_scheduler.py index 8a2a117e6ab..f5b44224853 100644 --- a/tests/local_testing/test_scheduler.py +++ b/tests/local_testing/test_scheduler.py @@ -10,7 +10,7 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path from litellm import Router -from litellm.scheduler import FlowItem, Scheduler +from litellm.scheduler import FlowItem, Scheduler, SchedulerCacheKeys from litellm import ModelResponse @@ -40,6 +40,35 @@ async def test_scheduler_diff_model_names(): ) +@pytest.mark.asyncio +async def test_scheduler_poll_persists_queue_to_cache(): + class StubRedisCache: + def __init__(self): + self.store = {} + + async def async_get_cache(self, key, **kwargs): + return self.store.get(key) + + async def async_set_cache(self, key, value, **kwargs): + self.store[key] = value + + redis_cache = StubRedisCache() + scheduler = Scheduler(redis_cache=redis_cache) + + item1 = FlowItem(priority=0, request_id="10", model_name="gpt-3.5-turbo") + item2 = FlowItem(priority=0, request_id="11", model_name="gpt-3.5-turbo") + await scheduler.add_request(item1) + await scheduler.add_request(item2) + + await scheduler.poll( + id="10", model_name="gpt-3.5-turbo", health_deployments=[] + ) + + queue_key = f"{SchedulerCacheKeys.queue.value}:{item1.model_name}" + updated_queue = redis_cache.store[queue_key] + assert updated_queue[0][1] == "11" + + @pytest.mark.parametrize("p0, p1", [(0, 0), (0, 1), (1, 0)]) @pytest.mark.parametrize("healthy_deployments", [[{"key": "value"}], []]) @pytest.mark.asyncio diff --git a/tests/router_unit_tests/test_router_prompt_caching.py b/tests/router_unit_tests/test_router_prompt_caching.py index 73469b8fe3b..7fbaf985b0f 100644 --- a/tests/router_unit_tests/test_router_prompt_caching.py +++ b/tests/router_unit_tests/test_router_prompt_caching.py @@ -185,3 +185,89 @@ async def test_router_prompt_caching_same_cacheable_prefix_routes_to_same_deploy assert ( model_id_1 == model_id_2 == model_id_3 ), f"All requests should route to same deployment, but got: {model_id_1}, {model_id_2}, {model_id_3}" + + +def test_extract_cacheable_prefix_with_string_content_and_message_level_cache_control(): + """ + Test that extract_cacheable_prefix correctly handles messages where: + - content is a string (not a list of content blocks) + - cache_control is a sibling key at the message level + + This is a valid message format per LiteLLM's ChatCompletionUserMessage type: + {"role": "user", "content": "...", "cache_control": {"type": "ephemeral"}} + + Regression test for issue #19228. + """ + # Test case 1: Single message with string content and message-level cache_control + messages_string_content = [ + {"role": "system", "content": "You are a helpful assistant"}, + { + "role": "user", + "content": "This is a large message that should be cached", + "cache_control": {"type": "ephemeral", "ttl": "5m"}, + }, + ] + + result = PromptCachingCache.extract_cacheable_prefix(messages_string_content) + + # Should return both messages (system + user with cache_control) + assert len(result) == 2, f"Expected 2 messages, got {len(result)}" + assert result[0]["role"] == "system" + assert result[1]["role"] == "user" + assert result[1]["content"] == "This is a large message that should be cached" + assert result[1].get("cache_control") == {"type": "ephemeral", "ttl": "5m"} + + +def test_extract_cacheable_prefix_with_string_content_no_cache_control(): + """ + Test that extract_cacheable_prefix returns empty list when: + - content is a string + - no cache_control is present + """ + messages_no_cache = [ + {"role": "system", "content": "You are a helpful assistant"}, + {"role": "user", "content": "Hello"}, + ] + + result = PromptCachingCache.extract_cacheable_prefix(messages_no_cache) + + # Should return empty list (no cacheable content) + assert len(result) == 0, f"Expected 0 messages, got {len(result)}" + + +def test_extract_cacheable_prefix_mixed_string_and_list_content(): + """ + Test that extract_cacheable_prefix handles messages with a mix of: + - String content with message-level cache_control + - List content with block-level cache_control + + The last cache_control (regardless of format) should determine the cacheable prefix. + """ + # Message with string content + cache_control, followed by message with list content + cache_control + messages_mixed = [ + {"role": "system", "content": "You are a helpful assistant"}, + { + "role": "user", + "content": "First cached message", + "cache_control": {"type": "ephemeral"}, + }, + { + "role": "user", + "content": [ + { + "type": "text", + "text": "Second cached message in list format", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + {"role": "user", "content": "This should not be in the prefix"}, + ] + + result = PromptCachingCache.extract_cacheable_prefix(messages_mixed) + + # Should include first 3 messages (up to and including the last cache_control) + assert len(result) == 3, f"Expected 3 messages, got {len(result)}" + assert result[0]["role"] == "system" + assert result[1]["content"] == "First cached message" + assert isinstance(result[2]["content"], list) diff --git a/tests/test_litellm/integrations/test_prometheus_labels.py b/tests/test_litellm/integrations/test_prometheus_labels.py index 8a4295e98fc..c0b863ef6ee 100644 --- a/tests/test_litellm/integrations/test_prometheus_labels.py +++ b/tests/test_litellm/integrations/test_prometheus_labels.py @@ -3,7 +3,7 @@ Unit tests for prometheus metric labels configuration """ from litellm.types.integrations.prometheus import ( PrometheusMetricLabels, - UserAPIKeyLabelNames + UserAPIKeyLabelNames, ) @@ -42,9 +42,10 @@ def test_user_email_label_exists(): def test_prometheus_metric_labels_structure(): """Test that all required prometheus metrics have proper label structure""" - from litellm.types.integrations.prometheus import DEFINED_PROMETHEUS_METRICS from typing import get_args + from litellm.types.integrations.prometheus import DEFINED_PROMETHEUS_METRICS + # Test a few key metrics to ensure they have proper label structure test_metrics = [ "litellm_proxy_total_requests_metric", @@ -69,8 +70,161 @@ def test_prometheus_metric_labels_structure(): print(f"✅ {metric_name} has proper label structure with user_email") +def test_route_normalization_for_responses_api(): + """ + Test that route normalization prevents high cardinality in Prometheus metrics + for the /v1/responses/{response_id} endpoint. + + Issue: https://github.com/BerriAI/litellm/issues/XXXX + Each unique response ID was creating a separate metric line, causing the + /metrics endpoint to grow to ~30MB and take ~40 seconds to respond. + + Fix: Routes are normalized to collapse dynamic IDs into placeholders. + """ + from litellm.proxy.auth.auth_utils import normalize_request_route + + # Test responses API routes + responses_routes = [ + ("/v1/responses/1234567890", "/v1/responses/{response_id}"), + ("/v1/responses/9876543210", "/v1/responses/{response_id}"), + ("/v1/responses/abcdefghij", "/v1/responses/{response_id}"), + ("/v1/responses/resp_abc123", "/v1/responses/{response_id}"), + ("/v1/responses/litellm_poll_xyz", "/v1/responses/{response_id}"), + ] + + for original, expected in responses_routes: + normalized = normalize_request_route(original) + assert normalized == expected, \ + f"Failed: {original} -> {normalized} (expected {expected})" + + # Verify cardinality reduction + unique_normalized = set(normalize_request_route(route) for route, _ in responses_routes) + assert len(unique_normalized) == 1, \ + f"Expected 1 unique normalized route, got {len(unique_normalized)}: {unique_normalized}" + + print(f"✅ Responses API routes: {len(responses_routes)} different IDs normalized to 1 metric label") + + +def test_route_normalization_for_sub_routes(): + """Test that sub-routes like /cancel and /input_items are normalized correctly""" + from litellm.proxy.auth.auth_utils import normalize_request_route + + sub_routes = [ + ("/v1/responses/id1/cancel", "/v1/responses/{response_id}/cancel"), + ("/v1/responses/id2/cancel", "/v1/responses/{response_id}/cancel"), + ("/v1/responses/id3/input_items", "/v1/responses/{response_id}/input_items"), + ("/openai/v1/responses/id4/input_items", "/openai/v1/responses/{response_id}/input_items"), + ] + + for original, expected in sub_routes: + normalized = normalize_request_route(original) + assert normalized == expected, \ + f"Failed: {original} -> {normalized} (expected {expected})" + + print("✅ Sub-routes normalized correctly") + + +def test_route_normalization_preserves_static_routes(): + """Test that static routes are not affected by normalization""" + from litellm.proxy.auth.auth_utils import normalize_request_route + + static_routes = [ + "/chat/completions", + "/v1/chat/completions", + "/v1/embeddings", + "/health", + "/metrics", + "/v1/models", + "/v1/responses", # List endpoint without ID + ] + + for route in static_routes: + normalized = normalize_request_route(route) + assert normalized == route, \ + f"Static route should not be modified: {route} -> {normalized}" + + print(f"✅ {len(static_routes)} static routes preserved") + + +def test_route_normalization_other_dynamic_apis(): + """Test normalization for other OpenAI-compatible APIs with dynamic IDs""" + from litellm.proxy.auth.auth_utils import normalize_request_route + + test_cases = [ + # Threads API + ("/v1/threads/thread_123", "/v1/threads/{thread_id}"), + ("/v1/threads/thread_abc/messages", "/v1/threads/{thread_id}/messages"), + ("/v1/threads/thread_abc/runs/run_123", "/v1/threads/{thread_id}/runs/{run_id}"), + + # Vector Stores API + ("/v1/vector_stores/vs_123", "/v1/vector_stores/{vector_store_id}"), + ("/v1/vector_stores/vs_123/files", "/v1/vector_stores/{vector_store_id}/files"), + + # Assistants API + ("/v1/assistants/asst_123", "/v1/assistants/{assistant_id}"), + + # Files API + ("/v1/files/file_123", "/v1/files/{file_id}"), + ("/v1/files/file_123/content", "/v1/files/{file_id}/content"), + + # Batches API + ("/v1/batches/batch_123", "/v1/batches/{batch_id}"), + ("/v1/batches/batch_123/cancel", "/v1/batches/{batch_id}/cancel"), + ] + + for original, expected in test_cases: + normalized = normalize_request_route(original) + assert normalized == expected, \ + f"Failed: {original} -> {normalized} (expected {expected})" + + print(f"✅ {len(test_cases)} other API routes normalized correctly") + + +def test_prometheus_metrics_use_normalized_routes(): + """ + Test that Prometheus metrics use the normalized route in labels + to prevent high cardinality. + """ + from unittest.mock import MagicMock + + from litellm.integrations.prometheus import ( + PrometheusLogger, + UserAPIKeyLabelValues, + prometheus_label_factory, + ) + + # Create a mock PrometheusLogger + prometheus_logger = MagicMock() + prometheus_logger.get_labels_for_metric = PrometheusLogger.get_labels_for_metric.__get__(prometheus_logger) + + # Test with a normalized route + enum_values = UserAPIKeyLabelValues( + route="/v1/responses/{response_id}", # Normalized route + status_code="200", + requested_model="gpt-4", + ) + + labels = prometheus_label_factory( + supported_enum_labels=prometheus_logger.get_labels_for_metric( + metric_name="litellm_proxy_total_requests_metric" + ), + enum_values=enum_values, + ) + + # Verify the route is normalized in labels + assert labels["route"] == "/v1/responses/{response_id}", \ + f"Expected normalized route in labels, got: {labels.get('route')}" + + print("✅ Prometheus metrics use normalized routes in labels") + + if __name__ == "__main__": test_user_email_in_required_metrics() test_user_email_label_exists() test_prometheus_metric_labels_structure() - print("All prometheus label tests passed!") \ No newline at end of file + test_route_normalization_for_responses_api() + test_route_normalization_for_sub_routes() + test_route_normalization_preserves_static_routes() + test_route_normalization_other_dynamic_apis() + test_prometheus_metrics_use_normalized_routes() + print("\n✅ All prometheus label tests passed!") \ No newline at end of file diff --git a/tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py b/tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py new file mode 100755 index 00000000000..99a1eb427d7 --- /dev/null +++ b/tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py @@ -0,0 +1,185 @@ +#!/usr/bin/env python3 +""" +Test script for issue #12443: Gemini aiohttp session leak + +Validates that: +1. BaseLLMAIOHTTPHandler properly closes sessions via __del__ +2. atexit handler works with new event loop approach +3. No "Unclosed client session" warnings are generated +""" + +import asyncio +import gc +import sys +from pathlib import Path + +import pytest + +# Add litellm to path +sys.path.insert(0, str(Path(__file__).parent)) + + +def count_aiohttp_sessions(): + """Count unclosed aiohttp ClientSession objects""" + import aiohttp + + count = 0 + for obj in gc.get_objects(): + if isinstance(obj, aiohttp.ClientSession): + if not obj.closed: + count += 1 + return count + + +async def test_aiohttp_handler_cleanup(): + """Test BaseLLMAIOHTTPHandler session cleanup""" + print("\n" + "=" * 70) + print("TEST: BaseLLMAIOHTTPHandler Session Cleanup") + print("=" * 70) + + from litellm.llms.custom_httpx.aiohttp_handler import BaseLLMAIOHTTPHandler + + initial_sessions = count_aiohttp_sessions() + print(f"\nInitial unclosed sessions: {initial_sessions}") + + # Create handler and trigger session creation + print("\nCreating BaseLLMAIOHTTPHandler and triggering session creation...") + handler = BaseLLMAIOHTTPHandler() + + # This triggers session creation (line 111 of aiohttp_handler.py) + session = handler._get_async_client_session() + print(f"Session created: {session}") + + sessions_after_create = count_aiohttp_sessions() + print(f"Sessions after creation: {sessions_after_create}") + + # Delete handler - should trigger __del__ cleanup + print("\nDeleting handler (should trigger __del__)...") + del handler + del session + gc.collect() + await asyncio.sleep(0.1) # Let async cleanup finish + + final_sessions = count_aiohttp_sessions() + print(f"Final unclosed sessions: {final_sessions}") + + session_diff = final_sessions - initial_sessions + print(f"\nSession difference: {session_diff:+d}") + + if session_diff == 0: + print("\n✅ PASS: __del__ cleanup working correctly") + return True + else: + print(f"\n❌ FAIL: {session_diff} sessions leaked") + return False + + +async def test_atexit_cleanup(): + """Test that atexit cleanup works with new event loop approach""" + print("\n" + "=" * 70) + print("TEST: atexit Cleanup (new event loop approach)") + print("=" * 70) + + from litellm.llms.custom_httpx.async_client_cleanup import ( + close_litellm_async_clients, + ) + + initial_sessions = count_aiohttp_sessions() + print(f"\nInitial unclosed sessions: {initial_sessions}") + + # Use the actual global base_llm_aiohttp_handler from litellm.main + print("\nAccessing global base_llm_aiohttp_handler (like Gemini does)...") + import litellm + + handler = litellm.base_llm_aiohttp_handler + session = handler._get_async_client_session() + + sessions_after_create = count_aiohttp_sessions() + print(f"Sessions after creation: {sessions_after_create}") + + # Call cleanup function (simulates atexit) + print("\nCalling close_litellm_async_clients() (simulates atexit)...") + await close_litellm_async_clients() + + gc.collect() + await asyncio.sleep(0.1) + + final_sessions = count_aiohttp_sessions() + print(f"Final unclosed sessions: {final_sessions}") + + session_diff = final_sessions - initial_sessions + print(f"\nSession difference: {session_diff:+d}") + + if session_diff == 0: + print("\n✅ PASS: atexit cleanup working correctly") + return True + else: + print(f"\n❌ FAIL: {session_diff} sessions leaked") + return False + + +def test_new_event_loop_atexit(): + """Test that the new atexit handler can create a fresh event loop""" + print("\n" + "=" * 70) + print("TEST: atexit with Fresh Event Loop Creation") + print("=" * 70) + + from litellm.llms.custom_httpx.async_client_cleanup import ( + close_litellm_async_clients, + ) + + print("\nVerifying atexit handler can create fresh loop (no running loop)...") + print("Note: At atexit time, there's typically no running event loop") + + # Save current loop to restore later + try: + current_loop = asyncio.get_running_loop() + print("Warning: Found running loop - can't test atexit scenario accurately") + pytest.skip("Cannot test atexit scenario when event loop is running") + except RuntimeError: + pass # Good - no running loop + + # Create a new loop like the fixed atexit handler does + print("Creating new event loop (like fixed atexit handler)...") + new_loop = asyncio.new_event_loop() + asyncio.set_event_loop(new_loop) + + try: + new_loop.run_until_complete(close_litellm_async_clients()) + print("✅ Successfully ran cleanup with fresh event loop") + finally: + new_loop.close() + + +async def main(): + """Run all tests""" + print("\n" + "=" * 70) + print("Gemini aiohttp Session Leak Fix Validation (Issue #12443)") + print("=" * 70) + + results = [] + + # Test 1: __del__ cleanup + results.append(await test_aiohttp_handler_cleanup()) + + # Test 2: atexit cleanup function + results.append(await test_atexit_cleanup()) + + print("\n" + "=" * 70) + print("Test Results") + print("=" * 70) + passed = sum(results) + total = len(results) + print(f"\nPassed: {passed}/{total}") + + if passed == total: + print("\n✅ All tests PASSED - Issue #12443 is FIXED") + else: + print(f"\n❌ {total - passed} test(s) FAILED") + + return passed == total + + +if __name__ == "__main__": + success = asyncio.run(main()) + sys.exit(0 if success else 1) diff --git a/tests/test_litellm/llms/test_oom_fixes.py b/tests/test_litellm/llms/test_oom_fixes.py new file mode 100644 index 00000000000..3b0a2a16fd1 --- /dev/null +++ b/tests/test_litellm/llms/test_oom_fixes.py @@ -0,0 +1,296 @@ +#!/usr/bin/env python3 +""" +Memory Leak Fix Validation Script + +Tests the fixes for issues #14540 and related OOM problems: +1. Presidio guardrail aiohttp session leak (presidio.py) +2. OpenAI common_utils httpx.AsyncClient creation bypass + +This script demonstrates that the fixes prevent memory leaks by: +- Tracking open file descriptors (each HTTP client creates sockets) +- Monitoring aiohttp ClientSession objects +- Checking httpx.AsyncClient instances + +Run with: python test_oom_fixes.py +""" + +import asyncio +import gc +import os +import sys +import tracemalloc +from pathlib import Path + +# Add litellm to path +sys.path.insert(0, str(Path(__file__).parent)) + + +def count_open_fds(): + """Count open file descriptors (proxy for open connections)""" + try: + fd_dir = Path(f"/proc/{os.getpid()}/fd") + if fd_dir.exists(): + return len(list(fd_dir.iterdir())) + except Exception: + pass + return None + + +def count_aiohttp_sessions(): + """Count unclosed aiohttp ClientSession objects""" + import aiohttp + + count = 0 + for obj in gc.get_objects(): + if isinstance(obj, aiohttp.ClientSession): + if not obj.closed: + count += 1 + return count + + +def count_httpx_clients(): + """Count httpx AsyncClient instances""" + import httpx + + async_clients = 0 + sync_clients = 0 + for obj in gc.get_objects(): + if isinstance(obj, httpx.AsyncClient): + if not obj.is_closed: + async_clients += 1 + elif isinstance(obj, httpx.Client): + if not obj.is_closed: + sync_clients += 1 + return async_clients, sync_clients + + +async def test_presidio_fix(): + """ + Test that Presidio guardrail doesn't leak aiohttp sessions. + + Before fix: Each call to analyze_text() created a new aiohttp.ClientSession + After fix: Reuses a single session stored in self._http_session + """ + print("\n" + "=" * 70) + print("TEST 1: Presidio Guardrail Session Leak Fix (Sequential)") + print("=" * 70) + + from litellm.proxy.guardrails.guardrail_hooks.presidio import ( + _OPTIONAL_PresidioPIIMasking, + ) + + # Create Presidio instance with mock testing mode + presidio = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + mock_redacted_text={"text": "mocked"}, + ) + + initial_fds = count_open_fds() + initial_sessions = count_aiohttp_sessions() + + print(f"\nInitial state:") + print(f" - Open file descriptors: {initial_fds}") + print(f" - Unclosed aiohttp sessions: {initial_sessions}") + + # Simulate 100 sequential requests + print(f"\nSimulating 100 sequential guardrail checks...") + for i in range(100): + # This would previously create a new ClientSession on each call + result = await presidio.check_pii( + text="test@email.com", + output_parse_pii=False, + presidio_config=None, + request_data={}, + ) + + # Force garbage collection + gc.collect() + await asyncio.sleep(0.1) # Let async cleanup finish + + final_fds = count_open_fds() + final_sessions = count_aiohttp_sessions() + + print(f"\nAfter 100 sequential requests:") + print(f" - Open file descriptors: {final_fds}") + print(f" - Unclosed aiohttp sessions: {final_sessions}") + + if final_fds and initial_fds: + fd_diff = final_fds - initial_fds + print(f" - FD difference: {fd_diff:+d}") + + session_diff = final_sessions - initial_sessions + print(f" - Session difference: {session_diff:+d}") + + # Cleanup + await presidio._close_http_session() + + print(f"\n✅ RESULT: Session leak {'PREVENTED' if session_diff <= 1 else 'DETECTED'}") + print( + f" Expected: ≤1 new session (the shared one), Got: {session_diff} new sessions" + ) + + +async def test_presidio_concurrent_load(): + """ + Test that Presidio guardrail handles concurrent requests without race conditions. + + Critical test: Validates that asyncio.Lock prevents multiple concurrent requests + from creating multiple sessions, which would leak memory under production load. + """ + print("\n" + "=" * 70) + print("TEST 2: Presidio Concurrent Load (Race Condition Check)") + print("=" * 70) + + from litellm.proxy.guardrails.guardrail_hooks.presidio import ( + _OPTIONAL_PresidioPIIMasking, + ) + + # Create Presidio instance with mock testing mode + presidio = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + mock_redacted_text={"text": "mocked"}, + ) + + initial_sessions = count_aiohttp_sessions() + print(f"\nInitial unclosed sessions: {initial_sessions}") + + # Simulate 50 concurrent requests (realistic proxy load) + print(f"\nSimulating 50 CONCURRENT guardrail checks...") + tasks = [] + for i in range(50): + task = presidio.check_pii( + text=f"test{i}@email.com", + output_parse_pii=False, + presidio_config=None, + request_data={}, + ) + tasks.append(task) + + # Execute all 50 requests concurrently + await asyncio.gather(*tasks) + + # Force garbage collection + gc.collect() + await asyncio.sleep(0.1) + + final_sessions = count_aiohttp_sessions() + print(f"Final unclosed sessions: {final_sessions}") + + session_diff = final_sessions - initial_sessions + print(f"\nSession difference: {session_diff:+d}") + + # Cleanup + await presidio._close_http_session() + + # CRITICAL: Should only create 1 session even with 50 concurrent requests + if session_diff <= 1: + print("\n✅ PASS: Race condition prevented - only 1 session created") + return True + else: + print(f"\n❌ FAIL: Race condition detected - {session_diff} sessions created!") + print(" This indicates asyncio.Lock is not working correctly") + return False + + +async def test_openai_client_caching(): + """ + Test that OpenAI common_utils caches httpx clients instead of creating new ones. + + Before fix: Each call to _get_async_http_client() created a new httpx.AsyncClient + After fix: Routes through get_async_httpx_client() which provides TTL-based caching + """ + print("\n" + "=" * 70) + print("TEST 2: OpenAI HTTP Client Caching Fix") + print("=" * 70) + + from litellm.llms.openai.common_utils import BaseOpenAILLM + + initial_async, initial_sync = count_httpx_clients() + print(f"\nInitial state:") + print(f" - Unclosed httpx.AsyncClient instances: {initial_async}") + print(f" - Unclosed httpx.Client instances: {initial_sync}") + + # Simulate 100 calls to get HTTP client + print(f"\nSimulating 100 client retrievals...") + clients = [] + for i in range(100): + # This would previously create a new AsyncClient on each call + client = BaseOpenAILLM._get_async_http_client() + clients.append(client) + + # Force garbage collection + gc.collect() + + final_async, final_sync = count_httpx_clients() + + print(f"\nAfter 100 retrievals:") + print(f" - Unclosed httpx.AsyncClient instances: {final_async}") + print(f" - Unclosed httpx.Client instances: {final_sync}") + + async_diff = final_async - initial_async + print(f" - AsyncClient difference: {async_diff:+d}") + + # Check if we got the same client instance (caching works) + unique_clients = len(set(id(c) for c in clients if c is not None)) + print(f" - Unique client instances returned: {unique_clients}") + + print( + f"\n✅ RESULT: Client caching {'WORKING' if unique_clients <= 2 else 'BROKEN'}" + ) + print( + f" Expected: ≤2 unique clients (due to TTL), Got: {unique_clients} unique clients" + ) + + +async def main(): + """Run all memory leak tests""" + print("\n" + "=" * 70) + print("LiteLLM OOM Fixes Validation") + print("Testing fixes for issues #14540, #14384, #13251, #12443") + print("=" * 70) + + # Start memory tracking + tracemalloc.start() + + results = [] + + try: + # Test 1: Sequential Presidio + await test_presidio_fix() + results.append(True) # Sequential test always passes if no exception + + # Test 2: Concurrent Presidio (race condition check) + result = await test_presidio_concurrent_load() + results.append(result) + + # Test 3: OpenAI client caching + await test_openai_client_caching() + results.append(True) + + print("\n" + "=" * 70) + print("Test Results") + print("=" * 70) + passed = sum(results) + total = len(results) + print(f"\nPassed: {passed}/{total}") + + if passed == total: + print("\n✅ All tests PASSED") + else: + print(f"\n❌ {total - passed} test(s) FAILED") + + # Show memory stats + current, peak = tracemalloc.get_traced_memory() + print(f"\nMemory usage:") + print(f" - Current: {current / 1024 / 1024:.1f} MB") + print(f" - Peak: {peak / 1024 / 1024:.1f} MB") + + return passed == total + + finally: + tracemalloc.stop() + + +if __name__ == "__main__": + success = asyncio.run(main()) + sys.exit(0 if success else 1) diff --git a/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py new file mode 100644 index 00000000000..823fd82d1ce --- /dev/null +++ b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py @@ -0,0 +1,275 @@ +""" +Tests for Volcengine Responses API transformation. +""" +import os +import sys + +import httpx +import pytest + +sys.path.insert(0, os.path.abspath("../../../../..")) + +import litellm +from litellm.llms.volcengine.responses.transformation import ( + VolcEngineResponsesAPIConfig, +) +from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams +from litellm.types.responses.main import DeleteResponseResult +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + + +class TestVolcengineResponsesAPITransformation: + """Test Volcengine Responses API configuration and transformations.""" + + def test_provider_config_registration(self): + """Provider registry should return VolcEngineResponsesAPIConfig.""" + config = ProviderConfigManager.get_provider_responses_api_config( + model="volcengine/demo-model", + provider=LlmProviders.VOLCENGINE, + ) + + assert config is not None, "Config should not be None for Volcengine provider" + assert isinstance( + config, VolcEngineResponsesAPIConfig + ), f"Expected VolcEngineResponsesAPIConfig, got {type(config)}" + assert ( + config.custom_llm_provider == LlmProviders.VOLCENGINE + ), "custom_llm_provider should be VOLCENGINE" + + def test_parallel_tool_calls_dropped(self): + """Volcengine does not list parallel_tool_calls; ensure it is removed.""" + config = VolcEngineResponsesAPIConfig() + params = ResponsesAPIOptionalRequestParams( + parallel_tool_calls=True, + temperature=0.5, + metadata={"k": "v"}, + ) + + mapped = config.map_openai_params( + response_api_optional_params=params, + model="volcengine/demo-model", + drop_params=False, + ) + + assert "parallel_tool_calls" not in mapped, "parallel_tool_calls must be dropped" + assert mapped.get("temperature") == 0.5 + assert "metadata" not in mapped, "Undocumented params should not be included" + + def test_unsupported_params_are_dropped(self): + """Unknown fields should be dropped before send, including nested extra_body.""" + config = VolcEngineResponsesAPIConfig() + + request = config.transform_responses_api_request( + model="volcengine/demo-model", + input="hi", + response_api_optional_request_params={ + "unsupported_custom_param": 0.1, + "temperature": 0.2, + "metadata": {"k": "v"}, + "extra_body": {"unsupported_custom_param": 1, "temperature": 0.3}, + }, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert "unsupported_custom_param" not in request + assert request["temperature"] == 0.2 + assert "metadata" not in request + assert "extra_body" in request + assert "unsupported_custom_param" not in request["extra_body"] + assert request["extra_body"]["temperature"] == 0.3 + + def test_get_complete_url_variants(self): + """Ensure Volcengine endpoint construction handles different bases.""" + config = VolcEngineResponsesAPIConfig() + + default_url = config.get_complete_url(api_base=None, litellm_params={}) + assert default_url == "https://ark.cn-beijing.volces.com/api/v3/responses" + + api_base_with_api = config.get_complete_url( + api_base="https://custom.volc.com/api/v3", litellm_params={} + ) + assert api_base_with_api == "https://custom.volc.com/api/v3/responses" + + api_base_full = config.get_complete_url( + api_base="https://custom.volc.com/api/v3/responses", litellm_params={} + ) + assert api_base_full == "https://custom.volc.com/api/v3/responses" + + @pytest.mark.parametrize( + "litellm_params, expected_key", + [ + ({"api_key": "dict-key"}, "dict-key"), + (GenericLiteLLMParams(api_key="attr-key"), "attr-key"), + ], + ) + def test_validate_environment_uses_api_key( + self, monkeypatch, litellm_params, expected_key + ): + """validate_environment should pull api key from params/env and attach headers.""" + config = VolcEngineResponsesAPIConfig() + + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.delenv("ARK_API_KEY", raising=False) + monkeypatch.delenv("VOLCENGINE_API_KEY", raising=False) + + headers = config.validate_environment( + headers={}, model="volcengine/demo-model", litellm_params=litellm_params + ) + + assert headers.get("Authorization") == f"Bearer {expected_key}" + assert headers.get("Content-Type") == "application/json" + + def test_validate_environment_raises_without_key(self, monkeypatch): + """validate_environment should error when no key is available.""" + config = VolcEngineResponsesAPIConfig() + + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.delenv("ARK_API_KEY", raising=False) + monkeypatch.delenv("VOLCENGINE_API_KEY", raising=False) + + with pytest.raises(ValueError): + config.validate_environment( + headers={}, model="volcengine/demo", litellm_params={} + ) + + def test_unsupported_params_are_dropped_with_extra_body(self): + """Unknown fields (including extra_body) should be dropped before send.""" + config = VolcEngineResponsesAPIConfig() + + request = config.transform_responses_api_request( + model="volcengine/demo-model", + input="hi", + response_api_optional_request_params={ + "unsupported_custom_param": 0.1, + "temperature": 0.2, + "metadata": {"k": "v"}, + "extra_body": {"unsupported_custom_param": 1, "temperature": 0.3}, + }, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert "unsupported_custom_param" not in request + assert "metadata" not in request + assert request["temperature"] == 0.2 + assert "extra_body" in request + assert "unsupported_custom_param" not in request["extra_body"] + assert request["extra_body"]["temperature"] == 0.3 + + def test_valid_thinking_caching_and_expire_at_pass(self): + """Documented params should pass through without validation errors.""" + config = VolcEngineResponsesAPIConfig() + request = config.transform_responses_api_request( + model="volcengine/demo-model", + input="hi", + response_api_optional_request_params={ + "instructions": "do X", + "thinking": {"type": "enabled"}, + "caching": {"type": "enabled"}, + "expire_at": 1234567890, + "temperature": 0.5, + }, + litellm_params=GenericLiteLLMParams(), + headers={}, + ) + + assert request["thinking"]["type"] == "enabled" + assert request["caching"]["type"] == "enabled" + assert request["expire_at"] == 1234567890 + assert request["instructions"] == "do X" + + def test_supported_params_limited_to_docs(self): + """Supported params should match documented Volcengine surface.""" + config = VolcEngineResponsesAPIConfig() + supported = set(config.get_supported_openai_params("volcengine/demo-model")) + + expected = { + "input", + "model", + "instructions", + "max_output_tokens", + "previous_response_id", + "store", + "reasoning", + "stream", + "temperature", + "top_p", + "text", + "tools", + "tool_choice", + "max_tool_calls", + "thinking", + "caching", + "expire_at", + "context_management", + "extra_headers", + "extra_query", + "extra_body", + "timeout", + } + + assert supported == expected + + def test_error_class_returns_volcengine_error(self): + """Errors should be wrapped with VolcEngineError for consistent handling.""" + config = VolcEngineResponsesAPIConfig() + error = config.get_error_class("bad request", 400, headers={"x": "y"}) + from litellm.llms.volcengine.common_utils import VolcEngineError + + assert isinstance(error, VolcEngineError) + assert error.status_code == 400 + assert error.message == "bad request" + assert error.headers.get("x") == "y" + + def test_transform_response_api_response_sets_headers_and_created_at(self): + """Responses should include processed headers and keep created_at intact.""" + config = VolcEngineResponsesAPIConfig() + response_payload = { + "id": "resp_123", + "object": "response", + "created_at": 123, + "status": "completed", + "output": [], + "model": "demo-model", + "usage": {"input_tokens": 0, "output_tokens": 0, "total_tokens": 0}, + } + http_response = httpx.Response( + status_code=200, + json=response_payload, + request=httpx.Request("POST", "https://example.com/responses"), + headers={"x-test": "1"}, + ) + + result = config.transform_response_api_response( + model="volcengine/demo-model", + raw_response=http_response, + logging_obj=type( + "Logger", + (), + {"post_call": staticmethod(lambda **kwargs: None)}, + ), + ) + + assert result.created_at == 123 + assert result._hidden_params["headers"].get("x-test") == "1" + assert "additional_headers" in result._hidden_params + + def test_transform_delete_response_api_response_parses_json(self): + """DELETE response parsing should return DeleteResponseResult.""" + config = VolcEngineResponsesAPIConfig() + http_response = httpx.Response( + status_code=200, + json={"id": "resp_123", "deleted": True}, + request=httpx.Request("DELETE", "https://example.com/responses/resp_123"), + ) + + result = config.transform_delete_response_api_response( + raw_response=http_response, + logging_obj=None, + ) + + assert isinstance(result, DeleteResponseResult) + assert result.deleted is True diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_function_call_output_normalization.py b/tests/test_litellm/responses/litellm_completion_transformation/test_function_call_output_normalization.py new file mode 100644 index 00000000000..19aeba7f9cd --- /dev/null +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_function_call_output_normalization.py @@ -0,0 +1,40 @@ +""" +Tests for normalizing Responses API function_call_output into chat tool messages. + +This is important for Gemini/Vertex, which expects tool results to be represented +as tool/function response parts; if the tool output is passed as a list of input_* parts, +we normalize it to text/image blocks or a string. +""" + +from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, +) + + +def test_function_call_output_list_input_text_is_converted_to_tool_string_content(): + out = LiteLLMCompletionResponsesConfig._transform_responses_api_tool_call_output_to_chat_completion_message( + tool_call_output={ + "type": "function_call_output", + "call_id": "call_1", + "output": [{"type": "input_text", "text": "hello"}, {"type": "input_text", "text": " world"}], + } + ) + + assert len(out) == 1 + msg = out[0] + assert msg["role"] == "tool" + assert msg["tool_call_id"] == "call_1" + assert msg["content"] == "hello world" + + +def test_function_call_output_string_passthrough(): + out = LiteLLMCompletionResponsesConfig._transform_responses_api_tool_call_output_to_chat_completion_message( + tool_call_output={ + "type": "function_call_output", + "call_id": "call_1", + "output": '{"ok":true}', + } + ) + assert len(out) == 1 + assert out[0]["content"] == '{"ok":true}' + diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py new file mode 100644 index 00000000000..51150383b01 --- /dev/null +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py @@ -0,0 +1,116 @@ +""" +Tests for streaming tool-calls in Responses API transformation. + +Ensures that when the underlying chat-completions stream includes tool_calls deltas, +LiteLLM emits Responses API streaming events (output_item.added + function_call_arguments.*). + +Also ensures that tool calls that only appear in the final built response still get emitted +before response.completed. +""" + +from unittest.mock import AsyncMock + +from litellm.responses.litellm_completion_transformation.streaming_iterator import ( + LiteLLMCompletionStreamingIterator, +) +from litellm.types.llms.openai import ResponsesAPIStreamEvents +from litellm.types.utils import Delta, ModelResponse, ModelResponseStream, StreamingChoices + + +def test_tool_call_delta_is_emitted_as_responses_events(): + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + + # A streaming chunk with tool_calls delta but no text + chunk = ModelResponseStream( + id="chunk-1", + created=123, + model="test-model", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + role="assistant", + content="", + tool_calls=[ + { + "id": "call_1", + "type": "function", + "function": {"name": "do_thing", "arguments": '{"x":1}'}, + } + ], + ), + ) + ], + ) + + evt1 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk) + assert evt1 is not None + assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + assert evt1.output_index == 1 + + evt2 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk) + assert evt2 is not None + assert evt2.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA + assert evt2.item_id == "call_1" + assert evt2.output_index == 1 + assert evt2.delta == '{"x":1}' + + +def test_tool_calls_present_only_in_final_response_are_emitted_before_completed(): + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + + # Construct a final ModelResponse with tool_calls on the message. + # We bypass the stream builder and directly set iterator.litellm_model_response. + response = ModelResponse( + id="resp-1", + created=123, + model="test-model", + object="chat.completion", + choices=[ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_2", + "type": "function", + "function": {"name": "do_thing", "arguments": '{"y":2}'}, + "index": 0, + } + ], + }, + } + ], + ) + iterator.litellm_model_response = response + + # First common_done_event_logic call should yield tool events, not response.completed. + evt1 = iterator.common_done_event_logic(sync_mode=True) + assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + assert evt1.output_index == 1 + + evt2 = iterator.common_done_event_logic(sync_mode=True) + assert evt2.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE + assert evt2.item_id == "call_2" + assert evt2.output_index == 1 + assert evt2.arguments == '{"y":2}' + + evt3 = iterator.common_done_event_logic(sync_mode=True) + assert evt3.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE + assert evt3.output_index == 1 + diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py new file mode 100644 index 00000000000..5cb01fbae61 --- /dev/null +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py @@ -0,0 +1,78 @@ +""" +Regression: preserve function_call_output ordering. + +Gemini/Vertex requires tool outputs to immediately follow the assistant tool call. +The ResponsesAPI->Chat conversion must not move tool outputs to the end. +""" + +from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, +) + + +def test_function_call_output_stays_adjacent_to_tool_call(): + msgs = LiteLLMCompletionResponsesConfig._transform_response_input_param_to_chat_completion_message( + input=[ + { + "role": "user", + "type": "message", + "content": [{"type": "input_text", "text": "Call echo with 'hello'."}], + }, + { + "type": "function_call", + "name": "echo", + "call_id": "call_123", + "arguments": '{"text":"hello"}', + }, + { + "type": "function_call_output", + "call_id": "call_123", + "output": '{"text":"hello"}', + }, + { + "role": "assistant", + "type": "message", + "content": [{"type": "output_text", "text": "Done."}], + }, + { + "role": "user", + "type": "message", + "content": [{"type": "input_text", "text": "Now say hi."}], + }, + ] + ) + + # Find the assistant message that contains tool_calls + tool_call_idx = None + tool_msg_idx = None + assistant_ok_idx = None + + for i, m in enumerate(msgs): + if isinstance(m, dict) and m.get("role") == "assistant" and m.get("tool_calls"): + tool_call_idx = i + if isinstance(m, dict) and m.get("role") == "tool": + tool_msg_idx = i + + # Assistant "Done." can be either a plain string or a structured content list + if isinstance(m, dict) and m.get("role") == "assistant": + content = m.get("content") + if content == "Done.": + assistant_ok_idx = i + elif isinstance(content, list): + for block in content: + if ( + isinstance(block, dict) + and block.get("type") == "text" + and block.get("text") == "Done." + ): + assistant_ok_idx = i + break + + assert tool_call_idx is not None + assert tool_msg_idx is not None + assert assistant_ok_idx is not None + + # Tool output must be right after tool call, and before the assistant "Done." message. + assert tool_msg_idx == tool_call_idx + 1 + assert assistant_ok_idx > tool_msg_idx + diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 6c9c1f31c09..f6c24d19df5 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -2945,3 +2945,162 @@ def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926(): and not any_assistant_message_has_thinking_blocks(messages) ) assert should_drop_thinking is False + + +class TestAdditionalDropParamsForNonOpenAIProviders: + """ + Test additional_drop_params functionality for non-OpenAI providers. + + Fixes https://github.com/BerriAI/litellm/issues/19225 + + The bug was that additional_drop_params only filtered params for OpenAI/Azure + providers, but not for other providers like Bedrock. This caused OpenAI-specific + params like prompt_cache_key to be passed to Bedrock, resulting in errors. + """ + + def test_additional_drop_params_filters_for_bedrock(self): + """ + Test that additional_drop_params correctly filters params for Bedrock provider. + + Before the fix, prompt_cache_key would be passed through to Bedrock even when + specified in additional_drop_params, causing: + 'BedrockException - {"message":"The model returned the following errors: + prompt_cache_key: Extra inputs are not permitted"}' + """ + from litellm.utils import add_provider_specific_params_to_optional_params + + optional_params = {} + passed_params = { + "prompt_cache_key": "test_key_123", + "temperature": 0.7, + "model": "bedrock/anthropic.claude-v2", + } + openai_params = ["temperature", "max_tokens", "top_p", "model"] + + result = add_provider_specific_params_to_optional_params( + optional_params=optional_params, + passed_params=passed_params, + custom_llm_provider="bedrock", + openai_params=openai_params, + additional_drop_params=["prompt_cache_key"], + ) + + # prompt_cache_key should be filtered out + assert "prompt_cache_key" not in result + # temperature should still be there (it's in openai_params, not filtered) + # Note: temperature is in openai_params so it won't be added by this function + # The function only adds params NOT in openai_params + + def test_additional_drop_params_filters_multiple_params_for_non_openai(self): + """Test filtering multiple params for non-OpenAI providers.""" + from litellm.utils import add_provider_specific_params_to_optional_params + + optional_params = {} + passed_params = { + "prompt_cache_key": "test_key", + "some_openai_only_param": "value1", + "another_openai_param": "value2", + "keep_this_param": "keep_me", + } + openai_params = ["temperature", "max_tokens"] + + result = add_provider_specific_params_to_optional_params( + optional_params=optional_params, + passed_params=passed_params, + custom_llm_provider="anthropic", + openai_params=openai_params, + additional_drop_params=["prompt_cache_key", "some_openai_only_param"], + ) + + # Filtered params should not be present + assert "prompt_cache_key" not in result + assert "some_openai_only_param" not in result + # Non-filtered params should be present + assert result.get("another_openai_param") == "value2" + assert result.get("keep_this_param") == "keep_me" + + def test_additional_drop_params_none_keeps_all_params(self): + """Test that when additional_drop_params is None, all params are kept.""" + from litellm.utils import add_provider_specific_params_to_optional_params + + optional_params = {} + passed_params = { + "prompt_cache_key": "test_key", + "custom_param": "value", + } + openai_params = ["temperature"] + + result = add_provider_specific_params_to_optional_params( + optional_params=optional_params, + passed_params=passed_params, + custom_llm_provider="bedrock", + openai_params=openai_params, + additional_drop_params=None, + ) + + # All params should be present when additional_drop_params is None + assert result.get("prompt_cache_key") == "test_key" + assert result.get("custom_param") == "value" + + def test_additional_drop_params_empty_list_keeps_all_params(self): + """Test that when additional_drop_params is empty list, all params are kept.""" + from litellm.utils import add_provider_specific_params_to_optional_params + + optional_params = {} + passed_params = { + "prompt_cache_key": "test_key", + "custom_param": "value", + } + openai_params = ["temperature"] + + result = add_provider_specific_params_to_optional_params( + optional_params=optional_params, + passed_params=passed_params, + custom_llm_provider="bedrock", + openai_params=openai_params, + additional_drop_params=[], + ) + + # All params should be present when additional_drop_params is empty + assert result.get("prompt_cache_key") == "test_key" + assert result.get("custom_param") == "value" + + +class TestDropParamsWithPromptCacheKey: + """ + Test that drop_params: true correctly drops prompt_cache_key for non-OpenAI providers. + + Fixes https://github.com/BerriAI/litellm/issues/19225 + + prompt_cache_key is an OpenAI-specific parameter that should be automatically + dropped when using providers like Bedrock that don't support it. + """ + + def test_prompt_cache_key_in_default_params(self): + """Verify prompt_cache_key is now in DEFAULT_CHAT_COMPLETION_PARAM_VALUES.""" + from litellm.constants import DEFAULT_CHAT_COMPLETION_PARAM_VALUES + + assert "prompt_cache_key" in DEFAULT_CHAT_COMPLETION_PARAM_VALUES + assert "prompt_cache_retention" in DEFAULT_CHAT_COMPLETION_PARAM_VALUES + + def test_drop_params_removes_prompt_cache_key_for_bedrock(self): + """ + Test that get_optional_params with drop_params=True removes prompt_cache_key + for Bedrock provider since it's not in Bedrock's supported params. + """ + from litellm.utils import get_optional_params + + # Call get_optional_params for Bedrock with prompt_cache_key + # drop_params=True should remove it since Bedrock doesn't support it + result = get_optional_params( + model="anthropic.claude-3-sonnet-20240229-v1:0", + custom_llm_provider="bedrock", + prompt_cache_key="test_cache_key", + temperature=0.7, + drop_params=True, + ) + + # prompt_cache_key should be dropped for Bedrock + assert "prompt_cache_key" not in result + # temperature should remain (it's supported by Bedrock) + assert result.get("temperature") == 0.7