Merge pull request #19386 from BerriAI/litellm_staging_01_20_2026

Litellm staging 01 20 2026
This commit is contained in:
Sameer Kankute 2026-01-20 18:53:59 +05:30 • committed by GitHub
commit 37ce6957ab
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
28 changed files with 2750 additions and 218 deletions

View file

@ -1268,7 +1268,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
@ -1374,6 +1374,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
@ -1387,7 +1388,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
@ -1395,7 +1396,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
@ -1413,7 +1414,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]
@ -1431,7 +1432,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
@ -1551,14 +1552,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]]
@ -1588,12 +1589,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]
@ -1608,7 +1609,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
@ -1618,7 +1619,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
@ -1628,7 +1629,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
@ -1638,7 +1639,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",
@ -1655,11 +1656,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
@ -1670,7 +1671,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
@ -1681,7 +1682,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
@ -1692,7 +1693,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

View file

@ -198,6 +198,7 @@ LLM_CONFIG_NAMES = (
"AzureOpenAIOSeriesResponsesAPIConfig",
"XAIResponsesAPIConfig",
"LiteLLMProxyResponsesAPIConfig",
"VolcEngineResponsesAPIConfig",
"GoogleAIStudioInteractionsConfig",
"OpenAIOSeriesConfig",
"AnthropicSkillsConfig",
@ -591,6 +592,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"),
@ -774,4 +776,3 @@ __all__ = [
"_LLM_PROVIDER_LOGIC_IMPORT_MAP",
"_UTILS_MODULE_IMPORT_MAP",
]

View file

@ -542,6 +542,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 = [

View file

@ -215,7 +215,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,
@ -1338,7 +1338,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 += "/"

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -619,7 +619,7 @@ def load_credentials_from_list(kwargs: dict):
"""
# Access CredentialAccessor via module to trigger lazy loading if needed
CredentialAccessor = getattr(sys.modules[__name__], 'CredentialAccessor')
credential_name = kwargs.get("litellm_credential_name")
if credential_name and litellm.credential_list:
credential_accessor = CredentialAccessor.get_credential_values(credential_name)
@ -646,7 +646,7 @@ def _is_gemini_model(model: Optional[str], custom_llm_provider: Optional[str]) -
if custom_llm_provider in ["vertex_ai", "vertex_ai_beta"]:
return model is not None and "gemini" in model.lower()
return True
# Check if model name contains gemini
return model is not None and "gemini" in model.lower()
@ -668,7 +668,7 @@ def _process_assistant_message_tool_calls(
"""
role = msg_copy.get("role")
tool_calls = msg_copy.get("tool_calls")
if role == "assistant" and isinstance(tool_calls, list):
new_tool_calls = []
for tc in tool_calls:
@ -681,17 +681,17 @@ def _process_assistant_message_tool_calls(
else:
new_tool_calls.append(tc)
continue
# Remove thought signature from ID if present
if isinstance(tc_dict.get("id"), str):
if thought_signature_separator in tc_dict["id"]:
tc_dict["id"] = _remove_thought_signature_from_id(
tc_dict["id"], thought_signature_separator
)
new_tool_calls.append(tc_dict)
msg_copy["tool_calls"] = new_tool_calls
return msg_copy
@ -706,7 +706,7 @@ def _process_tool_message_id(msg_copy: dict, thought_signature_separator: str) -
msg_copy["tool_call_id"] = _remove_thought_signature_from_id(
msg_copy["tool_call_id"], thought_signature_separator
)
return msg_copy
@ -717,7 +717,7 @@ def _remove_thought_signatures_from_messages(
Remove thought signatures from tool call IDs in all messages.
"""
processed_messages = []
for msg in messages:
# Handle Pydantic models (convert to dict)
if hasattr(msg, "model_dump"):
@ -728,17 +728,17 @@ def _remove_thought_signatures_from_messages(
# Unknown type, keep as is
processed_messages.append(msg)
continue
# Process assistant messages with tool_calls
msg_dict = _process_assistant_message_tool_calls(
msg_dict, thought_signature_separator
)
# Process tool messages with tool_call_id
msg_dict = _process_tool_message_id(msg_dict, thought_signature_separator)
processed_messages.append(msg_dict)
return processed_messages
@ -958,7 +958,7 @@ def function_setup( # noqa: PLR0915
input=buffer.getvalue(),
model=model,
)
### REMOVE THOUGHT SIGNATURES FROM TOOL CALL IDS FOR NON-GEMINI MODELS ###
# Gemini models embed thought signatures in tool call IDs. When sending
# messages with tool calls to non-Gemini providers, we need to remove these
@ -974,7 +974,7 @@ def function_setup( # noqa: PLR0915
# Get custom_llm_provider to determine target provider
custom_llm_provider = kwargs.get("custom_llm_provider")
# If custom_llm_provider not in kwargs, try to determine it from the model
if not custom_llm_provider and model:
try:
@ -985,18 +985,18 @@ def function_setup( # noqa: PLR0915
except Exception:
# If we can't determine the provider, skip this processing
pass
# Only process if target is NOT a Gemini model
if not _is_gemini_model(model, custom_llm_provider):
verbose_logger.debug(
"Removing thought signatures from tool call IDs for non-Gemini model"
)
# Process messages to remove thought signatures
processed_messages = _remove_thought_signatures_from_messages(
messages, THOUGHT_SIGNATURE_SEPARATOR
)
# Update messages in kwargs or args
if "messages" in kwargs:
kwargs["messages"] = processed_messages
@ -3035,7 +3035,7 @@ def get_optional_params_embeddings( # noqa: PLR0915
):
# Lazy load get_supported_openai_params
get_supported_openai_params = getattr(sys.modules[__name__], 'get_supported_openai_params')
# retrieve all parameters passed to the function
passed_params = locals()
custom_llm_provider = passed_params.pop("custom_llm_provider", None)
@ -4121,7 +4121,21 @@ def get_optional_params( # noqa: PLR0915
),
)
elif "anthropic" in bedrock_base_model and bedrock_route == "invoke":
if bedrock_base_model.startswith("anthropic.claude-3"):
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,
@ -4134,18 +4148,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,
@ -4578,6 +4580,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
@ -7084,7 +7088,7 @@ def get_valid_models(
# init litellm_params
#################################
from litellm.types.router import LiteLLM_Params
if litellm_params is None:
litellm_params = LiteLLM_Params(model="")
if api_key is not None:
@ -7618,7 +7622,7 @@ class ProviderConfigManager:
@staticmethod
def _build_provider_config_map() -> dict[LlmProviders, tuple[Callable, bool]]:
"""Build the provider-to-config mapping dictionary.
Returns a dict mapping provider to (factory_function, needs_model_parameter).
This avoids expensive inspect.signature() calls at runtime.
"""
@ -7784,7 +7788,7 @@ class ProviderConfigManager:
) -> Optional[BaseConfig]:
"""
Returns the provider config for a given provider.
Uses O(1) dictionary lookup for fast provider resolution.
"""
# Check JSON providers FIRST (these override standard mappings)
@ -8015,6 +8019,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)
@ -8030,6 +8036,8 @@ class ProviderConfigManager:
return litellm.GithubCopilotResponsesAPIConfig()
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
@ -8487,7 +8495,7 @@ class ProviderConfigManager:
from litellm.llms.vertex_ai.ocr.common_utils import get_vertex_ai_ocr_config
return get_vertex_ai_ocr_config(model=model)
MistralOCRConfig = getattr(sys.modules[__name__], 'MistralOCRConfig')
PROVIDER_TO_CONFIG_MAP = {
litellm.LlmProviders.MISTRAL: MistralOCRConfig,
@ -8925,12 +8933,12 @@ def __getattr__(name: str) -> Any:
"""Lazy import handler for utils module with cached registry for improved performance."""
# Use cached registry from _lazy_imports instead of importing tuples every time
from litellm._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]
return handler_func(name)
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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