mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge pull request #19386 from BerriAI/litellm_staging_01_20_2026
Litellm staging 01 20 2026
This commit is contained in:
commit
37ce6957ab
28 changed files with 2750 additions and 218 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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 += "/"
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
557
litellm/llms/volcengine/responses/transformation.py
Normal file
557
litellm/llms/volcengine/responses/transformation.py
Normal 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]
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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=[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
185
tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py
Executable file
185
tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py
Executable 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)
|
||||
296
tests/test_litellm/llms/test_oom_fixes.py
Normal file
296
tests/test_litellm/llms/test_oom_fixes.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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}'
|
||||
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -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
|
||||
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue