From 11a622aa5a486f846ac29c8f335c2d6d00c5165e Mon Sep 17 00:00:00 2001 From: Vedant Madane <6527493+VedantMadane@users.noreply.github.com> Date: Sat, 17 Jan 2026 10:35:40 +0530 Subject: [PATCH 01/11] Fix extract_cacheable_prefix to handle string content with message-level cache_control (fixes #19228) --- litellm/router_utils/prompt_caching_cache.py | 16 ++++ .../test_router_prompt_caching.py | 86 +++++++++++++++++++ 2 files changed, 102 insertions(+) diff --git a/litellm/router_utils/prompt_caching_cache.py b/litellm/router_utils/prompt_caching_cache.py index dbf8b8fcba8..69698f282b1 100644 --- a/litellm/router_utils/prompt_caching_cache.py +++ b/litellm/router_utils/prompt_caching_cache.py @@ -76,6 +76,22 @@ class PromptCachingCache: for msg_idx, message in enumerate(messages): content = message.get("content") + + # Check for cache_control at message level (when content is a string) + # This handles the case where cache_control is a sibling of string content: + # {"role": "user", "content": "...", "cache_control": {"type": "ephemeral"}} + message_level_cache_control = message.get("cache_control") + if ( + message_level_cache_control is not None + and isinstance(message_level_cache_control, dict) + and message_level_cache_control.get("type") == "ephemeral" + ): + last_cacheable_message_idx = msg_idx + # Set to None to indicate the entire message content is cacheable + # (not a specific content block index within a list) + last_cacheable_content_idx = None + + # Also check for cache_control within content blocks (when content is a list) if not isinstance(content, list): continue diff --git a/tests/router_unit_tests/test_router_prompt_caching.py b/tests/router_unit_tests/test_router_prompt_caching.py index 73469b8fe3b..7fbaf985b0f 100644 --- a/tests/router_unit_tests/test_router_prompt_caching.py +++ b/tests/router_unit_tests/test_router_prompt_caching.py @@ -185,3 +185,89 @@ async def test_router_prompt_caching_same_cacheable_prefix_routes_to_same_deploy assert ( model_id_1 == model_id_2 == model_id_3 ), f"All requests should route to same deployment, but got: {model_id_1}, {model_id_2}, {model_id_3}" + + +def test_extract_cacheable_prefix_with_string_content_and_message_level_cache_control(): + """ + Test that extract_cacheable_prefix correctly handles messages where: + - content is a string (not a list of content blocks) + - cache_control is a sibling key at the message level + + This is a valid message format per LiteLLM's ChatCompletionUserMessage type: + {"role": "user", "content": "...", "cache_control": {"type": "ephemeral"}} + + Regression test for issue #19228. + """ + # Test case 1: Single message with string content and message-level cache_control + messages_string_content = [ + {"role": "system", "content": "You are a helpful assistant"}, + { + "role": "user", + "content": "This is a large message that should be cached", + "cache_control": {"type": "ephemeral", "ttl": "5m"}, + }, + ] + + result = PromptCachingCache.extract_cacheable_prefix(messages_string_content) + + # Should return both messages (system + user with cache_control) + assert len(result) == 2, f"Expected 2 messages, got {len(result)}" + assert result[0]["role"] == "system" + assert result[1]["role"] == "user" + assert result[1]["content"] == "This is a large message that should be cached" + assert result[1].get("cache_control") == {"type": "ephemeral", "ttl": "5m"} + + +def test_extract_cacheable_prefix_with_string_content_no_cache_control(): + """ + Test that extract_cacheable_prefix returns empty list when: + - content is a string + - no cache_control is present + """ + messages_no_cache = [ + {"role": "system", "content": "You are a helpful assistant"}, + {"role": "user", "content": "Hello"}, + ] + + result = PromptCachingCache.extract_cacheable_prefix(messages_no_cache) + + # Should return empty list (no cacheable content) + assert len(result) == 0, f"Expected 0 messages, got {len(result)}" + + +def test_extract_cacheable_prefix_mixed_string_and_list_content(): + """ + Test that extract_cacheable_prefix handles messages with a mix of: + - String content with message-level cache_control + - List content with block-level cache_control + + The last cache_control (regardless of format) should determine the cacheable prefix. + """ + # Message with string content + cache_control, followed by message with list content + cache_control + messages_mixed = [ + {"role": "system", "content": "You are a helpful assistant"}, + { + "role": "user", + "content": "First cached message", + "cache_control": {"type": "ephemeral"}, + }, + { + "role": "user", + "content": [ + { + "type": "text", + "text": "Second cached message in list format", + "cache_control": {"type": "ephemeral"}, + } + ], + }, + {"role": "user", "content": "This should not be in the prefix"}, + ] + + result = PromptCachingCache.extract_cacheable_prefix(messages_mixed) + + # Should include first 3 messages (up to and including the last cache_control) + assert len(result) == 3, f"Expected 3 messages, got {len(result)}" + assert result[0]["role"] == "system" + assert result[1]["content"] == "First cached message" + assert isinstance(result[2]["content"], list) From 45eb35938bb3f7b873181bdd22f60ae48348f5a3 Mon Sep 17 00:00:00 2001 From: Chesars Date: Mon, 19 Jan 2026 08:49:03 -0300 Subject: [PATCH 02/11] fix: drop_params not dropping prompt_cache_key for non-OpenAI providers Fixes #19225 Add prompt_cache_key and other missing OpenAI Chat Completions params to DEFAULT_CHAT_COMPLETION_PARAM_VALUES so drop_params: true works. Also fix additional_drop_params to filter extra params for all providers, not just OpenAI/Azure. --- litellm/constants.py | 4 + litellm/utils.py | 2 + tests/test_litellm/test_utils.py | 159 +++++++++++++++++++++++++++++++ 3 files changed, 165 insertions(+) diff --git a/litellm/constants.py b/litellm/constants.py index 4ea0be247b3..c2c19fddf74 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -532,6 +532,10 @@ DEFAULT_CHAT_COMPLETION_PARAM_VALUES = { "web_search_options": None, "service_tier": None, "safety_identifier": None, + "prompt_cache_key": None, + "prompt_cache_retention": None, + "store": None, + "metadata": None, } openai_compatible_endpoints: List = [ diff --git a/litellm/utils.py b/litellm/utils.py index 3cf300802aa..4d2dec04fb8 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -4578,6 +4578,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 diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 6c9c1f31c09..f6c24d19df5 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -2945,3 +2945,162 @@ def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926(): and not any_assistant_message_has_thinking_blocks(messages) ) assert should_drop_thinking is False + + +class TestAdditionalDropParamsForNonOpenAIProviders: + """ + Test additional_drop_params functionality for non-OpenAI providers. + + Fixes https://github.com/BerriAI/litellm/issues/19225 + + The bug was that additional_drop_params only filtered params for OpenAI/Azure + providers, but not for other providers like Bedrock. This caused OpenAI-specific + params like prompt_cache_key to be passed to Bedrock, resulting in errors. + """ + + def test_additional_drop_params_filters_for_bedrock(self): + """ + Test that additional_drop_params correctly filters params for Bedrock provider. + + Before the fix, prompt_cache_key would be passed through to Bedrock even when + specified in additional_drop_params, causing: + 'BedrockException - {"message":"The model returned the following errors: + prompt_cache_key: Extra inputs are not permitted"}' + """ + from litellm.utils import add_provider_specific_params_to_optional_params + + optional_params = {} + passed_params = { + "prompt_cache_key": "test_key_123", + "temperature": 0.7, + "model": "bedrock/anthropic.claude-v2", + } + openai_params = ["temperature", "max_tokens", "top_p", "model"] + + result = add_provider_specific_params_to_optional_params( + optional_params=optional_params, + passed_params=passed_params, + custom_llm_provider="bedrock", + openai_params=openai_params, + additional_drop_params=["prompt_cache_key"], + ) + + # prompt_cache_key should be filtered out + assert "prompt_cache_key" not in result + # temperature should still be there (it's in openai_params, not filtered) + # Note: temperature is in openai_params so it won't be added by this function + # The function only adds params NOT in openai_params + + def test_additional_drop_params_filters_multiple_params_for_non_openai(self): + """Test filtering multiple params for non-OpenAI providers.""" + from litellm.utils import add_provider_specific_params_to_optional_params + + optional_params = {} + passed_params = { + "prompt_cache_key": "test_key", + "some_openai_only_param": "value1", + "another_openai_param": "value2", + "keep_this_param": "keep_me", + } + openai_params = ["temperature", "max_tokens"] + + result = add_provider_specific_params_to_optional_params( + optional_params=optional_params, + passed_params=passed_params, + custom_llm_provider="anthropic", + openai_params=openai_params, + additional_drop_params=["prompt_cache_key", "some_openai_only_param"], + ) + + # Filtered params should not be present + assert "prompt_cache_key" not in result + assert "some_openai_only_param" not in result + # Non-filtered params should be present + assert result.get("another_openai_param") == "value2" + assert result.get("keep_this_param") == "keep_me" + + def test_additional_drop_params_none_keeps_all_params(self): + """Test that when additional_drop_params is None, all params are kept.""" + from litellm.utils import add_provider_specific_params_to_optional_params + + optional_params = {} + passed_params = { + "prompt_cache_key": "test_key", + "custom_param": "value", + } + openai_params = ["temperature"] + + result = add_provider_specific_params_to_optional_params( + optional_params=optional_params, + passed_params=passed_params, + custom_llm_provider="bedrock", + openai_params=openai_params, + additional_drop_params=None, + ) + + # All params should be present when additional_drop_params is None + assert result.get("prompt_cache_key") == "test_key" + assert result.get("custom_param") == "value" + + def test_additional_drop_params_empty_list_keeps_all_params(self): + """Test that when additional_drop_params is empty list, all params are kept.""" + from litellm.utils import add_provider_specific_params_to_optional_params + + optional_params = {} + passed_params = { + "prompt_cache_key": "test_key", + "custom_param": "value", + } + openai_params = ["temperature"] + + result = add_provider_specific_params_to_optional_params( + optional_params=optional_params, + passed_params=passed_params, + custom_llm_provider="bedrock", + openai_params=openai_params, + additional_drop_params=[], + ) + + # All params should be present when additional_drop_params is empty + assert result.get("prompt_cache_key") == "test_key" + assert result.get("custom_param") == "value" + + +class TestDropParamsWithPromptCacheKey: + """ + Test that drop_params: true correctly drops prompt_cache_key for non-OpenAI providers. + + Fixes https://github.com/BerriAI/litellm/issues/19225 + + prompt_cache_key is an OpenAI-specific parameter that should be automatically + dropped when using providers like Bedrock that don't support it. + """ + + def test_prompt_cache_key_in_default_params(self): + """Verify prompt_cache_key is now in DEFAULT_CHAT_COMPLETION_PARAM_VALUES.""" + from litellm.constants import DEFAULT_CHAT_COMPLETION_PARAM_VALUES + + assert "prompt_cache_key" in DEFAULT_CHAT_COMPLETION_PARAM_VALUES + assert "prompt_cache_retention" in DEFAULT_CHAT_COMPLETION_PARAM_VALUES + + def test_drop_params_removes_prompt_cache_key_for_bedrock(self): + """ + Test that get_optional_params with drop_params=True removes prompt_cache_key + for Bedrock provider since it's not in Bedrock's supported params. + """ + from litellm.utils import get_optional_params + + # Call get_optional_params for Bedrock with prompt_cache_key + # drop_params=True should remove it since Bedrock doesn't support it + result = get_optional_params( + model="anthropic.claude-3-sonnet-20240229-v1:0", + custom_llm_provider="bedrock", + prompt_cache_key="test_cache_key", + temperature=0.7, + drop_params=True, + ) + + # prompt_cache_key should be dropped for Bedrock + assert "prompt_cache_key" not in result + # temperature should remain (it's supported by Bedrock) + assert result.get("temperature") == 0.7 From 13d887a275de1985d3b2f731678c4890aac00fe6 Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Mon, 19 Jan 2026 21:01:34 -0600 Subject: [PATCH 03/11] Fix queue persistence to Redis (#19304) * Fix queue persistence to Redis * add test --- litellm/scheduler.py | 1 + tests/local_testing/test_scheduler.py | 31 ++++++++++++++++++++++++++- 2 files changed, 31 insertions(+), 1 deletion(-) diff --git a/litellm/scheduler.py b/litellm/scheduler.py index 3225ba0451c..5f3dd4cbf61 100644 --- a/litellm/scheduler.py +++ b/litellm/scheduler.py @@ -84,6 +84,7 @@ class Scheduler: if queue[0][1] == id: # Remove the item from the queue heapq.heappop(queue) + await self.save_queue(queue=queue, model_name=model_name) print_verbose(f"Popped id: {id}") return True else: diff --git a/tests/local_testing/test_scheduler.py b/tests/local_testing/test_scheduler.py index 8a2a117e6ab..f5b44224853 100644 --- a/tests/local_testing/test_scheduler.py +++ b/tests/local_testing/test_scheduler.py @@ -10,7 +10,7 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path from litellm import Router -from litellm.scheduler import FlowItem, Scheduler +from litellm.scheduler import FlowItem, Scheduler, SchedulerCacheKeys from litellm import ModelResponse @@ -40,6 +40,35 @@ async def test_scheduler_diff_model_names(): ) +@pytest.mark.asyncio +async def test_scheduler_poll_persists_queue_to_cache(): + class StubRedisCache: + def __init__(self): + self.store = {} + + async def async_get_cache(self, key, **kwargs): + return self.store.get(key) + + async def async_set_cache(self, key, value, **kwargs): + self.store[key] = value + + redis_cache = StubRedisCache() + scheduler = Scheduler(redis_cache=redis_cache) + + item1 = FlowItem(priority=0, request_id="10", model_name="gpt-3.5-turbo") + item2 = FlowItem(priority=0, request_id="11", model_name="gpt-3.5-turbo") + await scheduler.add_request(item1) + await scheduler.add_request(item2) + + await scheduler.poll( + id="10", model_name="gpt-3.5-turbo", health_deployments=[] + ) + + queue_key = f"{SchedulerCacheKeys.queue.value}:{item1.model_name}" + updated_queue = redis_cache.store[queue_key] + assert updated_queue[0][1] == "11" + + @pytest.mark.parametrize("p0, p1", [(0, 0), (0, 1), (1, 0)]) @pytest.mark.parametrize("healthy_deployments", [[{"key": "value"}], []]) @pytest.mark.asyncio From 004bde2c45014eb10fc8b73edf0d62220f51d38d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=8D=97=E8=BE=B0=E7=87=8F=E7=82=9A?= <95487306+LingXuanYin@users.noreply.github.com> Date: Tue, 20 Jan 2026 11:02:29 +0800 Subject: [PATCH 04/11] feat (volcengine) : Support Volcengine responses api (#18508) * Add Volcengine responses adapter * fix llms/volcengine/responses/transformation.py:507:9: F841 Local variable `origin` is assigned to but never used fix llms/volcengine/responses/transformation.py:95: error: Argument "headers" to "VolcEngineError" has incompatible type add more supported optional params removed redundant manual logging/utils fallbacks so litellm/__init__.py uses the registry only. --- litellm/__init__.py | 41 +- litellm/_lazy_imports_registry.py | 3 +- litellm/llms/volcengine/__init__.py | 4 +- .../volcengine/responses/transformation.py | 557 ++++++++++++++++++ litellm/utils.py | 54 +- ...est_volcengine_responses_transformation.py | 274 +++++++++ 6 files changed, 886 insertions(+), 47 deletions(-) create mode 100644 litellm/llms/volcengine/responses/transformation.py create mode 100644 tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py diff --git a/litellm/__init__.py b/litellm/__init__.py index 9eb3f075d5e..134f8206d37 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index f37c4dc6d04..3025f1b0c51 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -198,6 +198,7 @@ LLM_CONFIG_NAMES = ( "AzureOpenAIOSeriesResponsesAPIConfig", "XAIResponsesAPIConfig", "LiteLLMProxyResponsesAPIConfig", + "VolcEngineResponsesAPIConfig", "GoogleAIStudioInteractionsConfig", "OpenAIOSeriesConfig", "AnthropicSkillsConfig", @@ -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", ] - diff --git a/litellm/llms/volcengine/__init__.py b/litellm/llms/volcengine/__init__.py index 0887937bed5..fc0098e84d9 100644 --- a/litellm/llms/volcengine/__init__.py +++ b/litellm/llms/volcengine/__init__.py @@ -1,6 +1,6 @@ """ Volcengine LLM Provider -Support for Volcengine (ByteDance) chat and embedding models +Support for Volcengine (ByteDance) chat, embedding, and responses models. """ from .chat.transformation import VolcEngineChatConfig @@ -10,6 +10,7 @@ from .common_utils import ( get_volcengine_headers, ) from .embedding import VolcEngineEmbeddingConfig +from .responses.transformation import VolcEngineResponsesAPIConfig # For backward compatibility, keep the old class name VolcEngineConfig = VolcEngineChatConfig @@ -18,6 +19,7 @@ __all__ = [ "VolcEngineChatConfig", "VolcEngineConfig", # backward compatibility "VolcEngineEmbeddingConfig", + "VolcEngineResponsesAPIConfig", "VolcEngineError", "get_volcengine_base_url", "get_volcengine_headers", diff --git a/litellm/llms/volcengine/responses/transformation.py b/litellm/llms/volcengine/responses/transformation.py new file mode 100644 index 00000000000..872c8dcf118 --- /dev/null +++ b/litellm/llms/volcengine/responses/transformation.py @@ -0,0 +1,557 @@ +from typing import ( + TYPE_CHECKING, + Any, + Dict, + List, + Literal, + Optional, + Tuple, + Union, + get_args, + get_origin, +) + +import httpx +from pydantic import fields as pyd_fields + +import litellm +from litellm._logging import verbose_logger +from litellm.types.llms.openai import ResponseInputParam, ResponsesAPIStreamingResponse +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig +from litellm.litellm_core_utils.core_helpers import process_response_headers +from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( + _safe_convert_created_field, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.openai import ( + ResponsesAPIOptionalRequestParams, + ResponsesAPIResponse, +) +from litellm.types.responses.main import DeleteResponseResult +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders + +from ..common_utils import ( + VolcEngineError, + get_volcengine_base_url, + get_volcengine_headers, +) + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj + + LiteLLMLoggingObj = _LiteLLMLoggingObj +else: + LiteLLMLoggingObj = Any + + +class VolcEngineResponsesAPIConfig(OpenAIResponsesAPIConfig): + _SUPPORTED_OPTIONAL_PARAMS: List[str] = [ + # Doc-listed knobs + "instructions", + "max_output_tokens", + "previous_response_id", + "store", + "reasoning", + "stream", + "temperature", + "top_p", + "text", + "tools", + "tool_choice", + "max_tool_calls", + "thinking", + "caching", + "expire_at", + "context_management", + # LiteLLM-internal metadata (not sent to provider) + "metadata", + # Request plumbing helpers + "extra_headers", + "extra_query", + "extra_body", + "timeout", + ] + + @property + def custom_llm_provider(self) -> LlmProviders: + return LlmProviders.VOLCENGINE + + def get_supported_openai_params(self, model: str) -> list: + """ + Volcengine Responses API: only documented parameters are supported. + """ + supported = ["input", "model"] + list(self._SUPPORTED_OPTIONAL_PARAMS) + # Do not advertise internal-only metadata to callers; we still accept and drop it before send. + if "metadata" in supported: + supported.remove("metadata") + return supported + + def get_error_class( + self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers] + ) -> VolcEngineError: + typed_headers: httpx.Headers = ( + headers if isinstance(headers, httpx.Headers) else httpx.Headers(headers or {}) + ) + return VolcEngineError( + status_code=status_code, + message=error_message, + headers=typed_headers, + ) + + def validate_environment( + self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams] + ) -> dict: + """ + Build auth headers for Volcengine Responses API. + """ + if litellm_params is None: + litellm_params = GenericLiteLLMParams() + elif isinstance(litellm_params, dict): + litellm_params = GenericLiteLLMParams(**litellm_params) + + api_key = ( + litellm_params.api_key + or litellm.api_key + or get_secret_str("ARK_API_KEY") + or get_secret_str("VOLCENGINE_API_KEY") + ) + + if api_key is None: + raise ValueError( + "Volcengine API key is required. Set ARK_API_KEY / VOLCENGINE_API_KEY or pass api_key." + ) + + return get_volcengine_headers(api_key=api_key, extra_headers=headers) + + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Construct Volcengine Responses API endpoint. + """ + base_url = ( + api_base + or litellm.api_base + or get_secret_str("VOLCENGINE_API_BASE") + or get_secret_str("ARK_API_BASE") + or get_volcengine_base_url() + ) + + base_url = base_url.rstrip("/") + + if base_url.endswith("/responses"): + return base_url + if base_url.endswith("/api/v3"): + return f"{base_url}/responses" + return f"{base_url}/api/v3/responses" + + def map_openai_params( + self, + response_api_optional_params: ResponsesAPIOptionalRequestParams, + model: str, + drop_params: bool, + ) -> Dict: + """ + Volcengine Responses API aligns with OpenAI parameters. + Remove parameters not supported by the public docs. + """ + params = { + key: value + for key, value in dict(response_api_optional_params).items() + if key in self._SUPPORTED_OPTIONAL_PARAMS + } + + # LiteLLM metadata is internal-only; don't send to provider + params.pop("metadata", None) + + # Volcengine docs do not list parallel_tool_calls; drop it to avoid backend errors. + if "parallel_tool_calls" in params: + verbose_logger.debug( + "Volcengine Responses API: dropping unsupported 'parallel_tool_calls' param." + ) + params.pop("parallel_tool_calls", None) + + return params + + def transform_responses_api_request( + self, + model: str, + input: Union[str, ResponseInputParam], + response_api_optional_request_params: Dict, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Dict: + """ + Volcengine rejects any undocumented fields (including extra_body). Fail fast + with clear errors and re-filter with the documented whitelist before delegating + to the OpenAI base transformer. + """ + allowed = set(self._SUPPORTED_OPTIONAL_PARAMS) + + sanitized_optional = { + k: v for k, v in response_api_optional_request_params.items() if k in allowed + } + # Ensure metadata never reaches provider + sanitized_optional.pop("metadata", None) + sanitized_optional.pop("parallel_tool_calls", None) + + # If extra_body is provided, filter its keys against the same allowlist to avoid + # leaking unsupported params to the provider. + if isinstance(sanitized_optional.get("extra_body"), dict): + filtered_body = { + k: v for k, v in sanitized_optional["extra_body"].items() if k in allowed + } + if filtered_body: + sanitized_optional["extra_body"] = filtered_body + else: + sanitized_optional.pop("extra_body", None) + + return super().transform_responses_api_request( + model=model, + input=input, + response_api_optional_request_params=sanitized_optional, + litellm_params=litellm_params, + headers=headers, + ) + + def transform_streaming_response( + self, + model: str, + parsed_chunk: dict, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIStreamingResponse: + """ + Volcengine may omit required fields; auto-fill them using event model defaults. + """ + chunk = parsed_chunk + + # Patch missing response.output on response.* events + if isinstance(chunk, dict): + resp = chunk.get("response") + if isinstance(resp, dict) and "output" not in resp: + patched_chunk = dict(chunk) + patched_resp = dict(resp) + patched_resp["output"] = [] + patched_chunk["response"] = patched_resp + chunk = patched_chunk + + event_type = str(chunk.get("type")) if isinstance(chunk, dict) else None + event_pydantic_model = OpenAIResponsesAPIConfig.get_event_model_class( + event_type=event_type + ) + + patched_chunk = self._fill_missing_fields(chunk, event_pydantic_model) + + return event_pydantic_model(**patched_chunk) + + def transform_response_api_response( + self, + model: str, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + try: + logging_obj.post_call( + original_response=raw_response.text, + additional_args={"complete_input_dict": {}}, + ) + raw_response_json = raw_response.json() + if "created_at" in raw_response_json: + raw_response_json["created_at"] = _safe_convert_created_field( + raw_response_json["created_at"] + ) + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + try: + response = ResponsesAPIResponse(**raw_response_json) + except Exception: + verbose_logger.debug( + "Volcengine Responses API: falling back to model_construct for response parsing." + ) + response = ResponsesAPIResponse.model_construct(**raw_response_json) + + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + ######################################################### + ########## DELETE RESPONSE API TRANSFORMATION ############## + ######################################################### + def transform_delete_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}" + data: Dict = {} + return url, data + + def transform_delete_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> DeleteResponseResult: + try: + raw_response_json = raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + try: + return DeleteResponseResult(**raw_response_json) + except Exception: + verbose_logger.debug( + "Volcengine Responses API: falling back to model_construct for delete response parsing." + ) + return DeleteResponseResult.model_construct(**raw_response_json) + + ######################################################### + ########## GET RESPONSE API TRANSFORMATION ############### + ######################################################### + def transform_get_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}" + data: Dict = {} + return url, data + + def transform_get_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + try: + raw_response_json = raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + response = ResponsesAPIResponse(**raw_response_json) + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + ######################################################### + ########## LIST INPUT ITEMS TRANSFORMATION ############# + ######################################################### + def transform_list_input_items_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + after: Optional[str] = None, + before: Optional[str] = None, + include: Optional[List[str]] = None, + limit: int = 20, + order: Literal["asc", "desc"] = "desc", + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}/input_items" + params: Dict[str, Any] = {} + if after is not None: + params["after"] = after + if before is not None: + params["before"] = before + if include: + params["include"] = ",".join(include) + if limit is not None: + params["limit"] = limit + if order is not None: + params["order"] = order + return url, params + + def transform_list_input_items_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> Dict: + try: + return raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + ######################################################### + ########## CANCEL RESPONSE API TRANSFORMATION ########## + ######################################################### + def transform_cancel_response_api_request( + self, + response_id: str, + api_base: str, + litellm_params: GenericLiteLLMParams, + headers: dict, + ) -> Tuple[str, Dict]: + url = f"{api_base}/{response_id}/cancel" + data: Dict = {} + return url, data + + def transform_cancel_response_api_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + ) -> ResponsesAPIResponse: + try: + raw_response_json = raw_response.json() + except Exception: + raise VolcEngineError( + message=raw_response.text, status_code=raw_response.status_code + ) + + raw_response_headers = dict(raw_response.headers) + processed_headers = process_response_headers(raw_response_headers) + + response = ResponsesAPIResponse(**raw_response_json) + response._hidden_params["additional_headers"] = processed_headers + response._hidden_params["headers"] = raw_response_headers + return response + + def should_fake_stream( + self, + model: Optional[str], + stream: Optional[bool], + custom_llm_provider: Optional[str] = None, + ) -> bool: + """ + Volcengine Responses API supports native streaming; never fall back to fake stream. + """ + return False + + @staticmethod + def _fill_missing_fields( + chunk: Any, event_model: Any + ) -> Dict[str, Any]: + """ + Heuristically fill missing required fields with safe defaults based on the + event model's field annotations. This keeps parsing tolerant of providers that + omit non-essential fields. + """ + if not isinstance(chunk, dict) or event_model is None: + return chunk + + patched: Dict[str, Any] = dict(chunk) + fields_map = getattr(event_model, "model_fields", {}) or {} + + for name, field in fields_map.items(): + if name in patched: + patched[name] = VolcEngineResponsesAPIConfig._maybe_fill_nested( + patched[name], field.annotation + ) + continue + + # Explicit default or factory + if field.default is not pyd_fields.PydanticUndefined and field.default is not None: + patched[name] = field.default + continue + if ( + field.default_factory is not None + and field.default_factory is not pyd_fields.PydanticUndefined + ): + patched[name] = field.default_factory() + continue + + # Heuristic defaults for missing required fields + patched[name] = VolcEngineResponsesAPIConfig._default_for_annotation( + field.annotation + ) + + return patched + + @staticmethod + def _default_for_annotation(annotation: Any) -> Any: + origin = get_origin(annotation) + args = get_args(annotation) + + if annotation is int: + return 0 + if annotation is list or origin is list: + return [] + if origin is Union: + # Prefer empty list when any option is a list + if any((arg is list or get_origin(arg) is list) for arg in args): + return [] + if type(None) in args: + return None + if origin is Union and type(None) in args: + return None + + # Fallback to None when no safer guess exists + return None + + @staticmethod + def _maybe_fill_nested(value: Any, annotation: Any) -> Any: + """ + Recursively fill nested dict/list structures based on the annotated model. + """ + model_cls = VolcEngineResponsesAPIConfig._pick_model_class(annotation, value) + args = get_args(annotation) + + if isinstance(value, dict) and model_cls is not None: + return VolcEngineResponsesAPIConfig._fill_missing_fields(value, model_cls) + + if isinstance(value, list): + # Attempt to fill list elements if we know the element annotation + elem_ann: Any = args[0] if args else None + if elem_ann is not None: + return [ + VolcEngineResponsesAPIConfig._maybe_fill_nested(v, elem_ann) + for v in value + ] + + return value + + @staticmethod + def _pick_model_class(annotation: Any, value: Any) -> Optional[Any]: + """ + Choose the best-matching Pydantic model class for a nested dict. + """ + candidates: List[Any] = [] + origin = get_origin(annotation) + + if hasattr(annotation, "model_fields"): + candidates.append(annotation) + if origin is Union: + for arg in get_args(annotation): + if hasattr(arg, "model_fields"): + candidates.append(arg) + + if not candidates: + return None + + # Try to match by literal "type" field when available + if isinstance(value, dict): + v_type = value.get("type") + for candidate in candidates: + try: + type_field = candidate.model_fields.get("type") + if type_field is None: + continue + literal_ann = type_field.annotation + if get_origin(literal_ann) is Literal: + literal_values = get_args(literal_ann) + if v_type in literal_values: + return candidate + except Exception: + continue + + # Fall back to the first candidate + return candidates[0] diff --git a/litellm/utils.py b/litellm/utils.py index ac194e4f339..71f8877aac0 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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) @@ -7084,7 +7084,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 +7618,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 +7784,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 +8015,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 +8032,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 +8491,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 +8929,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}") diff --git a/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py new file mode 100644 index 00000000000..16930bcb995 --- /dev/null +++ b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py @@ -0,0 +1,274 @@ +""" +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.router import GenericLiteLLMParams +from litellm.types.responses.main import DeleteResponseResult +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", + "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 From 58c8c2b7b120154b559ebc96570a89f23e82f76f Mon Sep 17 00:00:00 2001 From: Ryan Malloy Date: Mon, 19 Jan 2026 20:02:55 -0700 Subject: [PATCH 05/11] fix: HTTP client memory leaks in Presidio, OpenAI, and Gemini (#19190) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: prevent HTTP client memory leaks in Presidio and OpenAI wrappers Fixes multiple memory leak issues reported in #14540 and related tickets: **Presidio Guardrail Fix (#14540)** - Problem: Every guardrail check created a new aiohttp.ClientSession - Impact: High-traffic proxies accumulated thousands of unclosed sessions - Solution: Share a single session across all guardrail checks - Added `self._http_session` instance variable - Lazy session creation via `_get_http_session()` - Proper cleanup via `_close_http_session()` and `__del__()` - Files: litellm/proxy/guardrails/guardrail_hooks/presidio.py **OpenAI HTTP Client Caching (#14540)** - Problem: `_get_async_http_client()` created new httpx.AsyncClient on each call - Impact: OpenAI/Azure completions bypassed client caching system - Solution: Route through `get_async_httpx_client()` for TTL-based caching - Caches clients by provider and SSL config - Fallback to direct creation if caching fails - Applied to both async and sync client methods - Files: litellm/llms/openai/common_utils.py **Test Script** - Added validation script to demonstrate fixes - Counts file descriptors and unclosed session objects - Files: test_oom_fixes.py Related issues: #14384, #13251, #12443 * fix(oom): prevent memory leaks in Presidio guardrails and OpenAI client creation Fixes two high-impact memory leaks: 1. Presidio Guardrail Session Leak (issue #14540) - Problem: Created new aiohttp.ClientSession on every guardrail check - Impact: Runs on EVERY proxy request when PII masking enabled - Fix: Shared session pattern with lifecycle management - Files: litellm/proxy/guardrails/guardrail_hooks/presidio.py 2. OpenAI HTTP Client Cache Bypass (issue #14540) - Problem: _get_async_http_client() created new httpx.AsyncClient, bypassing TTL cache - Impact: Every completion created new client with own connection pool - Fix: Route through get_async_httpx_client() for proper caching - Critical: Include SSL config in cache key for correctness - Files: litellm/llms/openai/common_utils.py Validation: - Presidio: 100 requests → 0 new sessions (was 100) - OpenAI: 100 calls → 1 unique client (was 100) - test_oom_fixes.py: Automated validation script * fix(oom): resolve Gemini aiohttp session leak (issue #12443) Fixes persistent "Unclosed client session" warnings when using Gemini models. Root Causes: 1. Broken atexit cleanup - get_event_loop() fails at exit time 2. On-demand session creation without reliable cleanup Changes: 1. Fixed atexit Cleanup (async_client_cleanup.py) - OLD: Used get_event_loop() which fails when loop is closed - NEW: Always create fresh event loop at exit time - Ensures cleanup runs successfully even when main loop is closed 2. Added __del__ Cleanup (aiohttp_handler.py) - Defense-in-depth: cleanup on garbage collection - Handles abnormal termination cases - Similar pattern to Presidio guardrail fix 3. Enhanced Cleanup Scope (async_client_cleanup.py) - Now closes global base_llm_aiohttp_handler instance - Previously only checked cache, missed module-level handler Validation: - Test 1: __del__ cleanup → 0 sessions leaked ✓ - Test 2: atexit cleanup → 0 sessions leaked ✓ - test_gemini_session_leak.py: Automated validation Related: #14540 (broader OOM issue tracking) * fix(types): use LlmProviders enum for get_async_httpx_client MyPy was failing because llm_provider parameter expects Union[LlmProviders, httpxSpecialProvider], not a string. Changed from string "openai" to LlmProviders.OPENAI enum value. * test: move validation tests to proper CI directories - Move test_oom_fixes.py to tests/test_litellm/llms/ - Move test_gemini_session_leak.py to tests/test_litellm/llms/custom_httpx/ - Fix pytest warning: use pytest.skip() instead of return True This ensures CI actually runs our OOM fix validation tests. * fix(oom): add asyncio.Lock to prevent race conditions in Presidio session creation - Make _get_http_session() async with asyncio.Lock protection - Prevents multiple concurrent requests from creating orphaned sessions - Add concurrent load test (50 parallel requests) to validate fix - Test confirms only 1 session created under concurrent load Critical fix: Previous implementation had race condition where concurrent guardrail checks could create multiple sessions, defeating the shared session pattern and causing memory leaks. * fix(presidio): eliminate race condition in session lock initialization Move asyncio.Lock creation from lazy initialization in _get_http_session() to __init__. The previous lazy init had a race condition where concurrent coroutines could both see _session_lock as None, both create locks, and end up with different lock instances - defeating the synchronization. asyncio.Lock() can be safely created without an event loop; it only requires one when awaited. --- litellm/llms/custom_httpx/aiohttp_handler.py | 35 +++ .../llms/custom_httpx/async_client_cleanup.py | 46 ++- litellm/llms/openai/common_utils.py | 74 +++-- .../guardrails/guardrail_hooks/presidio.py | 228 ++++++++------ .../custom_httpx/test_gemini_session_leak.py | 185 +++++++++++ tests/test_litellm/llms/test_oom_fixes.py | 296 ++++++++++++++++++ 6 files changed, 741 insertions(+), 123 deletions(-) create mode 100755 tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py create mode 100644 tests/test_litellm/llms/test_oom_fixes.py diff --git a/litellm/llms/custom_httpx/aiohttp_handler.py b/litellm/llms/custom_httpx/aiohttp_handler.py index c7a04a49fc2..93b6c563dc1 100644 --- a/litellm/llms/custom_httpx/aiohttp_handler.py +++ b/litellm/llms/custom_httpx/aiohttp_handler.py @@ -134,6 +134,41 @@ class BaseLLMAIOHTTPHandler: # Ignore errors during transport cleanup pass + def __del__(self): + """ + Cleanup: close aiohttp session on instance destruction. + + Provides defense-in-depth for issue #12443 - ensures cleanup happens + even if atexit handler doesn't run (abnormal termination). + """ + if ( + self.client_session is not None + and not self.client_session.closed + and self._owns_session + ): + try: + import asyncio + + try: + loop = asyncio.get_event_loop() + if loop.is_running(): + # Event loop is running - schedule cleanup task + asyncio.create_task(self.close()) + else: + # Event loop exists but not running - run cleanup + loop.run_until_complete(self.close()) + except RuntimeError: + # No event loop available - create one for cleanup + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + loop.run_until_complete(self.close()) + finally: + loop.close() + except Exception: + # Silently ignore errors during __del__ to avoid issues + pass + async def _make_common_async_call( self, async_client_session: Optional[ClientSession], diff --git a/litellm/llms/custom_httpx/async_client_cleanup.py b/litellm/llms/custom_httpx/async_client_cleanup.py index 45602576764..abbc61dc96d 100644 --- a/litellm/llms/custom_httpx/async_client_cleanup.py +++ b/litellm/llms/custom_httpx/async_client_cleanup.py @@ -9,7 +9,8 @@ async def close_litellm_async_clients(): Close all cached async HTTP clients to prevent resource leaks. This function iterates through all cached clients in litellm's in-memory cache - and closes any aiohttp client sessions that are still open. + and closes any aiohttp client sessions that are still open. Also closes the + global base_llm_aiohttp_handler instance (issue #12443). """ # Import here to avoid circular import import litellm @@ -25,7 +26,7 @@ async def close_litellm_async_clients(): except Exception: # Silently ignore errors during cleanup pass - + # Handle AsyncHTTPHandler instances (used by Gemini and other providers) elif hasattr(handler, 'client'): client = handler.client @@ -43,7 +44,7 @@ async def close_litellm_async_clients(): except Exception: # Silently ignore errors during cleanup pass - + # Handle any other objects with aclose method elif hasattr(handler, 'aclose'): try: @@ -52,6 +53,17 @@ async def close_litellm_async_clients(): # Silently ignore errors during cleanup pass + # Close the global base_llm_aiohttp_handler instance (issue #12443) + # This is used by Gemini and other providers that use aiohttp + if hasattr(litellm, 'base_llm_aiohttp_handler'): + base_handler = getattr(litellm, 'base_llm_aiohttp_handler', None) + if isinstance(base_handler, BaseLLMAIOHTTPHandler) and hasattr(base_handler, 'close'): + try: + await base_handler.close() + except Exception: + # Silently ignore errors during cleanup + pass + def register_async_client_cleanup(): """ @@ -62,22 +74,24 @@ def register_async_client_cleanup(): import atexit def cleanup_wrapper(): + """ + Cleanup wrapper that creates a fresh event loop for atexit cleanup. + + At exit time, the main event loop is often already closed. Creating a new + event loop ensures cleanup runs successfully (fixes issue #12443). + """ try: - loop = asyncio.get_event_loop() - if loop.is_running(): - # Schedule the cleanup coroutine - loop.create_task(close_litellm_async_clients()) - else: - # Run the cleanup coroutine - loop.run_until_complete(close_litellm_async_clients()) - except Exception: - # If we can't get an event loop or it's already closed, try creating a new one + # Always create a fresh event loop at exit time + # Don't use get_event_loop() - it may be closed or unavailable + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) try: - loop = asyncio.new_event_loop() loop.run_until_complete(close_litellm_async_clients()) + finally: + # Clean up the loop we created loop.close() - except Exception: - # Silently ignore errors during cleanup - pass + except Exception: + # Silently ignore errors during cleanup to avoid exit handler failures + pass atexit.register(cleanup_wrapper) diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index ce470f04aca..9f1ec9250cf 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -15,12 +15,14 @@ if TYPE_CHECKING: from aiohttp import ClientSession import litellm +from litellm._logging import verbose_logger from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.http_handler import ( _DEFAULT_TTL_FOR_HTTPX_CLIENTS, AsyncHTTPHandler, get_ssl_configuration, ) +from litellm.types.utils import LlmProviders class OpenAIError(BaseLLMException): @@ -203,30 +205,66 @@ 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 None + + # Get a cached AsyncHTTPHandler which manages the httpx.AsyncClient + cached_handler = get_async_httpx_client( + llm_provider=LlmProviders.OPENAI, # Cache key includes provider + params=params, # Include SSL config in cache key shared_session=shared_session, - ), - follow_redirects=True, - ) + ) + # Return the underlying httpx client from the handler + return cached_handler.client + except (ImportError, AttributeError, KeyError) as e: + # Fallback to creating a client directly if caching system unavailable + # This preserves backwards compatibility + verbose_logger.debug( + f"Client caching unavailable ({type(e).__name__}), using direct client creation" + ) + ssl_config = get_ssl_configuration() + return httpx.AsyncClient( + verify=ssl_config, + transport=AsyncHTTPHandler._create_async_transport( + ssl_context=ssl_config + if isinstance(ssl_config, ssl.SSLContext) + else None, + ssl_verify=ssl_config if isinstance(ssl_config, bool) else None, + shared_session=shared_session, + ), + follow_redirects=True, + ) @staticmethod def _get_sync_http_client() -> Optional[httpx.Client]: if litellm.client_session is not None: return litellm.client_session - # Get unified SSL configuration - ssl_config = get_ssl_configuration() + # Use the global cached client system to prevent memory leaks (issue #14540) + from litellm.llms.custom_httpx.http_handler import _get_httpx_client - return httpx.Client( - verify=ssl_config, - follow_redirects=True, - ) + try: + # Get SSL config and include in params for proper cache key + ssl_config = get_ssl_configuration() + params = {"ssl_verify": ssl_config} if ssl_config is not None else None + + # Get a cached HTTPHandler which manages the httpx.Client + cached_handler = _get_httpx_client(params=params) + # Return the underlying httpx client from the handler + return cached_handler.client + except (ImportError, AttributeError, KeyError) as e: + # Fallback to creating a client directly if caching system unavailable + verbose_logger.debug( + f"Client caching unavailable ({type(e).__name__}), using direct client creation" + ) + ssl_config = get_ssl_configuration() + return httpx.Client( + verify=ssl_config, + follow_redirects=True, + ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 4d7f4a5b125..20df54b62ce 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -102,6 +102,11 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): presidio_score_thresholds or {} ) self.presidio_language = presidio_language or "en" + # Shared HTTP session to prevent memory leaks (issue #14540) + self._http_session: Optional[aiohttp.ClientSession] = None + # Lock to prevent race conditions when creating session under concurrent load + # Note: asyncio.Lock() can be created without an event loop; it only needs one when awaited + self._session_lock: asyncio.Lock = asyncio.Lock() if mock_testing is True: # for testing purposes only return @@ -167,6 +172,47 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): "http://" + self.presidio_anonymizer_api_base ) + async def _get_http_session(self) -> aiohttp.ClientSession: + """ + Get or create the shared HTTP session for Presidio API calls. + + Fixes memory leak (issue #14540) where every guardrail check created + a new aiohttp.ClientSession that was never properly closed. + + Thread-safe: Uses asyncio.Lock to prevent race conditions when + multiple concurrent requests try to create the session simultaneously. + """ + async with self._session_lock: + if self._http_session is None or self._http_session.closed: + self._http_session = aiohttp.ClientSession() + return self._http_session + + async def _close_http_session(self) -> None: + """Close the HTTP session if it exists.""" + if self._http_session is not None and not self._http_session.closed: + await self._http_session.close() + self._http_session = None + + def __del__(self): + """Cleanup: close HTTP session on instance destruction.""" + if self._http_session is not None and not self._http_session.closed: + try: + # Try to close the session, but don't fail if event loop is gone + import asyncio + try: + loop = asyncio.get_event_loop() + if loop.is_running(): + # Schedule cleanup, don't block __del__ + asyncio.create_task(self._close_http_session()) + else: + loop.run_until_complete(self._close_http_session()) + except RuntimeError: + # Event loop is closed, can't clean up - not ideal but better than crashing + pass + except Exception: + # Suppress all exceptions in __del__ to avoid issues during shutdown + pass + def _get_presidio_analyze_request_payload( self, text: str, @@ -223,67 +269,69 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): ) return [] - async with aiohttp.ClientSession() as session: - if self.mock_redacted_text is not None: - return self.mock_redacted_text + if self.mock_redacted_text is not None: + return self.mock_redacted_text - # Make the request to /analyze - analyze_url = f"{self.presidio_analyzer_api_base}analyze" + # Use shared session to prevent memory leak (issue #14540) + session = await self._get_http_session() - analyze_payload: PresidioAnalyzeRequest = ( - self._get_presidio_analyze_request_payload( - text=text, - presidio_config=presidio_config, - request_data=request_data, - ) + # Make the request to /analyze + analyze_url = f"{self.presidio_analyzer_api_base}analyze" + + analyze_payload: PresidioAnalyzeRequest = ( + self._get_presidio_analyze_request_payload( + text=text, + presidio_config=presidio_config, + request_data=request_data, ) + ) - verbose_proxy_logger.debug( - "Making request to: %s with payload: %s", - analyze_url, - analyze_payload, - ) + verbose_proxy_logger.debug( + "Making request to: %s with payload: %s", + analyze_url, + analyze_payload, + ) - async with session.post(analyze_url, json=analyze_payload) as response: - analyze_results = await response.json() - verbose_proxy_logger.debug("analyze_results: %s", analyze_results) + async with session.post(analyze_url, json=analyze_payload) as response: + analyze_results = await response.json() + verbose_proxy_logger.debug("analyze_results: %s", analyze_results) - # Handle error responses from Presidio (e.g., {'error': 'No text provided'}) - # Presidio may return a dict instead of a list when errors occur - if isinstance(analyze_results, dict): - if "error" in analyze_results: - verbose_proxy_logger.warning( - "Presidio analyzer returned error: %s, returning empty list", - analyze_results.get("error") - ) - return [] - # If it's a dict but not an error, try to process it as a single item - verbose_proxy_logger.debug( - "Presidio returned dict (not list), attempting to process as single item" + # Handle error responses from Presidio (e.g., {'error': 'No text provided'}) + # Presidio may return a dict instead of a list when errors occur + if isinstance(analyze_results, dict): + if "error" in analyze_results: + verbose_proxy_logger.warning( + "Presidio analyzer returned error: %s, returning empty list", + analyze_results.get("error") ) - try: - return [PresidioAnalyzeResponseItem(**analyze_results)] - except Exception as e: - verbose_proxy_logger.warning( - "Failed to parse Presidio dict response: %s, returning empty list", - e - ) - return [] + return [] + # If it's a dict but not an error, try to process it as a single item + verbose_proxy_logger.debug( + "Presidio returned dict (not list), attempting to process as single item" + ) + try: + return [PresidioAnalyzeResponseItem(**analyze_results)] + except Exception as e: + verbose_proxy_logger.warning( + "Failed to parse Presidio dict response: %s, returning empty list", + e + ) + return [] - # Normal case: list of results - final_results = [] - for item in analyze_results: - try: - final_results.append(PresidioAnalyzeResponseItem(**item)) - except TypeError as te: - # Handle case where item is not a dict (shouldn't happen, but be defensive) - verbose_proxy_logger.warning( - "Skipping invalid Presidio result item: %s (error: %s)", - item, - te, - ) - continue - return final_results + # Normal case: list of results + final_results = [] + for item in analyze_results: + try: + final_results.append(PresidioAnalyzeResponseItem(**item)) + except TypeError as te: + # Handle case where item is not a dict (shouldn't happen, but be defensive) + verbose_proxy_logger.warning( + "Skipping invalid Presidio result item: %s (error: %s)", + item, + te, + ) + continue + return final_results except Exception as e: raise e @@ -302,46 +350,48 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): if isinstance(analyze_results, list) and len(analyze_results) == 0: return text - async with aiohttp.ClientSession() as session: - # Make the request to /anonymize - anonymize_url = f"{self.presidio_anonymizer_api_base}anonymize" - verbose_proxy_logger.debug("Making request to: %s", anonymize_url) - anonymize_payload = { - "text": text, - "analyzer_results": analyze_results, - } + # Use shared session to prevent memory leak (issue #14540) + session = await self._get_http_session() - async with session.post( - anonymize_url, json=anonymize_payload - ) as response: - redacted_text = await response.json() + # Make the request to /anonymize + anonymize_url = f"{self.presidio_anonymizer_api_base}anonymize" + verbose_proxy_logger.debug("Making request to: %s", anonymize_url) + anonymize_payload = { + "text": text, + "analyzer_results": analyze_results, + } - new_text = text - if redacted_text is not None: - verbose_proxy_logger.debug("redacted_text: %s", redacted_text) - for item in redacted_text["items"]: - start = item["start"] - end = item["end"] - replacement = item["text"] # replacement token - if item["operator"] == "replace" and output_parse_pii is True: - # check if token in dict - # if exists, add a uuid to the replacement token for swapping back to the original text in llm response output parsing - if replacement in self.pii_tokens: - replacement = replacement + str(uuid.uuid4()) + async with session.post( + anonymize_url, json=anonymize_payload + ) as response: + redacted_text = await response.json() - self.pii_tokens[replacement] = new_text[ - start:end - ] # get text it'll replace + new_text = text + if redacted_text is not None: + verbose_proxy_logger.debug("redacted_text: %s", redacted_text) + for item in redacted_text["items"]: + start = item["start"] + end = item["end"] + replacement = item["text"] # replacement token + if item["operator"] == "replace" and output_parse_pii is True: + # check if token in dict + # if exists, add a uuid to the replacement token for swapping back to the original text in llm response output parsing + if replacement in self.pii_tokens: + replacement = replacement + str(uuid.uuid4()) - new_text = new_text[:start] + replacement + new_text[end:] - entity_type = item.get("entity_type", None) - if entity_type is not None: - masked_entity_count[entity_type] = ( - masked_entity_count.get(entity_type, 0) + 1 - ) - return redacted_text["text"] - else: - raise Exception(f"Invalid anonymizer response: {redacted_text}") + self.pii_tokens[replacement] = new_text[ + start:end + ] # get text it'll replace + + new_text = new_text[:start] + replacement + new_text[end:] + entity_type = item.get("entity_type", None) + if entity_type is not None: + masked_entity_count[entity_type] = ( + masked_entity_count.get(entity_type, 0) + 1 + ) + return redacted_text["text"] + else: + raise Exception(f"Invalid anonymizer response: {redacted_text}") except Exception as e: raise e diff --git a/tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py b/tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py new file mode 100755 index 00000000000..99a1eb427d7 --- /dev/null +++ b/tests/test_litellm/llms/custom_httpx/test_gemini_session_leak.py @@ -0,0 +1,185 @@ +#!/usr/bin/env python3 +""" +Test script for issue #12443: Gemini aiohttp session leak + +Validates that: +1. BaseLLMAIOHTTPHandler properly closes sessions via __del__ +2. atexit handler works with new event loop approach +3. No "Unclosed client session" warnings are generated +""" + +import asyncio +import gc +import sys +from pathlib import Path + +import pytest + +# Add litellm to path +sys.path.insert(0, str(Path(__file__).parent)) + + +def count_aiohttp_sessions(): + """Count unclosed aiohttp ClientSession objects""" + import aiohttp + + count = 0 + for obj in gc.get_objects(): + if isinstance(obj, aiohttp.ClientSession): + if not obj.closed: + count += 1 + return count + + +async def test_aiohttp_handler_cleanup(): + """Test BaseLLMAIOHTTPHandler session cleanup""" + print("\n" + "=" * 70) + print("TEST: BaseLLMAIOHTTPHandler Session Cleanup") + print("=" * 70) + + from litellm.llms.custom_httpx.aiohttp_handler import BaseLLMAIOHTTPHandler + + initial_sessions = count_aiohttp_sessions() + print(f"\nInitial unclosed sessions: {initial_sessions}") + + # Create handler and trigger session creation + print("\nCreating BaseLLMAIOHTTPHandler and triggering session creation...") + handler = BaseLLMAIOHTTPHandler() + + # This triggers session creation (line 111 of aiohttp_handler.py) + session = handler._get_async_client_session() + print(f"Session created: {session}") + + sessions_after_create = count_aiohttp_sessions() + print(f"Sessions after creation: {sessions_after_create}") + + # Delete handler - should trigger __del__ cleanup + print("\nDeleting handler (should trigger __del__)...") + del handler + del session + gc.collect() + await asyncio.sleep(0.1) # Let async cleanup finish + + final_sessions = count_aiohttp_sessions() + print(f"Final unclosed sessions: {final_sessions}") + + session_diff = final_sessions - initial_sessions + print(f"\nSession difference: {session_diff:+d}") + + if session_diff == 0: + print("\n✅ PASS: __del__ cleanup working correctly") + return True + else: + print(f"\n❌ FAIL: {session_diff} sessions leaked") + return False + + +async def test_atexit_cleanup(): + """Test that atexit cleanup works with new event loop approach""" + print("\n" + "=" * 70) + print("TEST: atexit Cleanup (new event loop approach)") + print("=" * 70) + + from litellm.llms.custom_httpx.async_client_cleanup import ( + close_litellm_async_clients, + ) + + initial_sessions = count_aiohttp_sessions() + print(f"\nInitial unclosed sessions: {initial_sessions}") + + # Use the actual global base_llm_aiohttp_handler from litellm.main + print("\nAccessing global base_llm_aiohttp_handler (like Gemini does)...") + import litellm + + handler = litellm.base_llm_aiohttp_handler + session = handler._get_async_client_session() + + sessions_after_create = count_aiohttp_sessions() + print(f"Sessions after creation: {sessions_after_create}") + + # Call cleanup function (simulates atexit) + print("\nCalling close_litellm_async_clients() (simulates atexit)...") + await close_litellm_async_clients() + + gc.collect() + await asyncio.sleep(0.1) + + final_sessions = count_aiohttp_sessions() + print(f"Final unclosed sessions: {final_sessions}") + + session_diff = final_sessions - initial_sessions + print(f"\nSession difference: {session_diff:+d}") + + if session_diff == 0: + print("\n✅ PASS: atexit cleanup working correctly") + return True + else: + print(f"\n❌ FAIL: {session_diff} sessions leaked") + return False + + +def test_new_event_loop_atexit(): + """Test that the new atexit handler can create a fresh event loop""" + print("\n" + "=" * 70) + print("TEST: atexit with Fresh Event Loop Creation") + print("=" * 70) + + from litellm.llms.custom_httpx.async_client_cleanup import ( + close_litellm_async_clients, + ) + + print("\nVerifying atexit handler can create fresh loop (no running loop)...") + print("Note: At atexit time, there's typically no running event loop") + + # Save current loop to restore later + try: + current_loop = asyncio.get_running_loop() + print("Warning: Found running loop - can't test atexit scenario accurately") + pytest.skip("Cannot test atexit scenario when event loop is running") + except RuntimeError: + pass # Good - no running loop + + # Create a new loop like the fixed atexit handler does + print("Creating new event loop (like fixed atexit handler)...") + new_loop = asyncio.new_event_loop() + asyncio.set_event_loop(new_loop) + + try: + new_loop.run_until_complete(close_litellm_async_clients()) + print("✅ Successfully ran cleanup with fresh event loop") + finally: + new_loop.close() + + +async def main(): + """Run all tests""" + print("\n" + "=" * 70) + print("Gemini aiohttp Session Leak Fix Validation (Issue #12443)") + print("=" * 70) + + results = [] + + # Test 1: __del__ cleanup + results.append(await test_aiohttp_handler_cleanup()) + + # Test 2: atexit cleanup function + results.append(await test_atexit_cleanup()) + + print("\n" + "=" * 70) + print("Test Results") + print("=" * 70) + passed = sum(results) + total = len(results) + print(f"\nPassed: {passed}/{total}") + + if passed == total: + print("\n✅ All tests PASSED - Issue #12443 is FIXED") + else: + print(f"\n❌ {total - passed} test(s) FAILED") + + return passed == total + + +if __name__ == "__main__": + success = asyncio.run(main()) + sys.exit(0 if success else 1) diff --git a/tests/test_litellm/llms/test_oom_fixes.py b/tests/test_litellm/llms/test_oom_fixes.py new file mode 100644 index 00000000000..3b0a2a16fd1 --- /dev/null +++ b/tests/test_litellm/llms/test_oom_fixes.py @@ -0,0 +1,296 @@ +#!/usr/bin/env python3 +""" +Memory Leak Fix Validation Script + +Tests the fixes for issues #14540 and related OOM problems: +1. Presidio guardrail aiohttp session leak (presidio.py) +2. OpenAI common_utils httpx.AsyncClient creation bypass + +This script demonstrates that the fixes prevent memory leaks by: +- Tracking open file descriptors (each HTTP client creates sockets) +- Monitoring aiohttp ClientSession objects +- Checking httpx.AsyncClient instances + +Run with: python test_oom_fixes.py +""" + +import asyncio +import gc +import os +import sys +import tracemalloc +from pathlib import Path + +# Add litellm to path +sys.path.insert(0, str(Path(__file__).parent)) + + +def count_open_fds(): + """Count open file descriptors (proxy for open connections)""" + try: + fd_dir = Path(f"/proc/{os.getpid()}/fd") + if fd_dir.exists(): + return len(list(fd_dir.iterdir())) + except Exception: + pass + return None + + +def count_aiohttp_sessions(): + """Count unclosed aiohttp ClientSession objects""" + import aiohttp + + count = 0 + for obj in gc.get_objects(): + if isinstance(obj, aiohttp.ClientSession): + if not obj.closed: + count += 1 + return count + + +def count_httpx_clients(): + """Count httpx AsyncClient instances""" + import httpx + + async_clients = 0 + sync_clients = 0 + for obj in gc.get_objects(): + if isinstance(obj, httpx.AsyncClient): + if not obj.is_closed: + async_clients += 1 + elif isinstance(obj, httpx.Client): + if not obj.is_closed: + sync_clients += 1 + return async_clients, sync_clients + + +async def test_presidio_fix(): + """ + Test that Presidio guardrail doesn't leak aiohttp sessions. + + Before fix: Each call to analyze_text() created a new aiohttp.ClientSession + After fix: Reuses a single session stored in self._http_session + """ + print("\n" + "=" * 70) + print("TEST 1: Presidio Guardrail Session Leak Fix (Sequential)") + print("=" * 70) + + from litellm.proxy.guardrails.guardrail_hooks.presidio import ( + _OPTIONAL_PresidioPIIMasking, + ) + + # Create Presidio instance with mock testing mode + presidio = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + mock_redacted_text={"text": "mocked"}, + ) + + initial_fds = count_open_fds() + initial_sessions = count_aiohttp_sessions() + + print(f"\nInitial state:") + print(f" - Open file descriptors: {initial_fds}") + print(f" - Unclosed aiohttp sessions: {initial_sessions}") + + # Simulate 100 sequential requests + print(f"\nSimulating 100 sequential guardrail checks...") + for i in range(100): + # This would previously create a new ClientSession on each call + result = await presidio.check_pii( + text="test@email.com", + output_parse_pii=False, + presidio_config=None, + request_data={}, + ) + + # Force garbage collection + gc.collect() + await asyncio.sleep(0.1) # Let async cleanup finish + + final_fds = count_open_fds() + final_sessions = count_aiohttp_sessions() + + print(f"\nAfter 100 sequential requests:") + print(f" - Open file descriptors: {final_fds}") + print(f" - Unclosed aiohttp sessions: {final_sessions}") + + if final_fds and initial_fds: + fd_diff = final_fds - initial_fds + print(f" - FD difference: {fd_diff:+d}") + + session_diff = final_sessions - initial_sessions + print(f" - Session difference: {session_diff:+d}") + + # Cleanup + await presidio._close_http_session() + + print(f"\n✅ RESULT: Session leak {'PREVENTED' if session_diff <= 1 else 'DETECTED'}") + print( + f" Expected: ≤1 new session (the shared one), Got: {session_diff} new sessions" + ) + + +async def test_presidio_concurrent_load(): + """ + Test that Presidio guardrail handles concurrent requests without race conditions. + + Critical test: Validates that asyncio.Lock prevents multiple concurrent requests + from creating multiple sessions, which would leak memory under production load. + """ + print("\n" + "=" * 70) + print("TEST 2: Presidio Concurrent Load (Race Condition Check)") + print("=" * 70) + + from litellm.proxy.guardrails.guardrail_hooks.presidio import ( + _OPTIONAL_PresidioPIIMasking, + ) + + # Create Presidio instance with mock testing mode + presidio = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + mock_redacted_text={"text": "mocked"}, + ) + + initial_sessions = count_aiohttp_sessions() + print(f"\nInitial unclosed sessions: {initial_sessions}") + + # Simulate 50 concurrent requests (realistic proxy load) + print(f"\nSimulating 50 CONCURRENT guardrail checks...") + tasks = [] + for i in range(50): + task = presidio.check_pii( + text=f"test{i}@email.com", + output_parse_pii=False, + presidio_config=None, + request_data={}, + ) + tasks.append(task) + + # Execute all 50 requests concurrently + await asyncio.gather(*tasks) + + # Force garbage collection + gc.collect() + await asyncio.sleep(0.1) + + final_sessions = count_aiohttp_sessions() + print(f"Final unclosed sessions: {final_sessions}") + + session_diff = final_sessions - initial_sessions + print(f"\nSession difference: {session_diff:+d}") + + # Cleanup + await presidio._close_http_session() + + # CRITICAL: Should only create 1 session even with 50 concurrent requests + if session_diff <= 1: + print("\n✅ PASS: Race condition prevented - only 1 session created") + return True + else: + print(f"\n❌ FAIL: Race condition detected - {session_diff} sessions created!") + print(" This indicates asyncio.Lock is not working correctly") + return False + + +async def test_openai_client_caching(): + """ + Test that OpenAI common_utils caches httpx clients instead of creating new ones. + + Before fix: Each call to _get_async_http_client() created a new httpx.AsyncClient + After fix: Routes through get_async_httpx_client() which provides TTL-based caching + """ + print("\n" + "=" * 70) + print("TEST 2: OpenAI HTTP Client Caching Fix") + print("=" * 70) + + from litellm.llms.openai.common_utils import BaseOpenAILLM + + initial_async, initial_sync = count_httpx_clients() + print(f"\nInitial state:") + print(f" - Unclosed httpx.AsyncClient instances: {initial_async}") + print(f" - Unclosed httpx.Client instances: {initial_sync}") + + # Simulate 100 calls to get HTTP client + print(f"\nSimulating 100 client retrievals...") + clients = [] + for i in range(100): + # This would previously create a new AsyncClient on each call + client = BaseOpenAILLM._get_async_http_client() + clients.append(client) + + # Force garbage collection + gc.collect() + + final_async, final_sync = count_httpx_clients() + + print(f"\nAfter 100 retrievals:") + print(f" - Unclosed httpx.AsyncClient instances: {final_async}") + print(f" - Unclosed httpx.Client instances: {final_sync}") + + async_diff = final_async - initial_async + print(f" - AsyncClient difference: {async_diff:+d}") + + # Check if we got the same client instance (caching works) + unique_clients = len(set(id(c) for c in clients if c is not None)) + print(f" - Unique client instances returned: {unique_clients}") + + print( + f"\n✅ RESULT: Client caching {'WORKING' if unique_clients <= 2 else 'BROKEN'}" + ) + print( + f" Expected: ≤2 unique clients (due to TTL), Got: {unique_clients} unique clients" + ) + + +async def main(): + """Run all memory leak tests""" + print("\n" + "=" * 70) + print("LiteLLM OOM Fixes Validation") + print("Testing fixes for issues #14540, #14384, #13251, #12443") + print("=" * 70) + + # Start memory tracking + tracemalloc.start() + + results = [] + + try: + # Test 1: Sequential Presidio + await test_presidio_fix() + results.append(True) # Sequential test always passes if no exception + + # Test 2: Concurrent Presidio (race condition check) + result = await test_presidio_concurrent_load() + results.append(result) + + # Test 3: OpenAI client caching + await test_openai_client_caching() + results.append(True) + + print("\n" + "=" * 70) + print("Test Results") + print("=" * 70) + passed = sum(results) + total = len(results) + print(f"\nPassed: {passed}/{total}") + + if passed == total: + print("\n✅ All tests PASSED") + else: + print(f"\n❌ {total - passed} test(s) FAILED") + + # Show memory stats + current, peak = tracemalloc.get_traced_memory() + print(f"\nMemory usage:") + print(f" - Current: {current / 1024 / 1024:.1f} MB") + print(f" - Peak: {peak / 1024 / 1024:.1f} MB") + + return passed == total + + finally: + tracemalloc.stop() + + +if __name__ == "__main__": + success = asyncio.run(main()) + sys.exit(0 if success else 1) From 0dfc3fad5a1b33f8d41af109aa6ee29e121b4dd3 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=9D=B4=EB=AA=85=ED=98=84?= <117622409+flex-myeonghyeon@users.noreply.github.com> Date: Tue, 20 Jan 2026 12:14:58 +0900 Subject: [PATCH 06/11] Fix: bedrock invoke claude 4 optional params #19318 (#19381) --- litellm/utils.py | 28 +- tests/llm_translation/test_optional_params.py | 254 ++++++++++++++++++ 2 files changed, 269 insertions(+), 13 deletions(-) diff --git a/litellm/utils.py b/litellm/utils.py index 71f8877aac0..ce4fddbe739 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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, diff --git a/tests/llm_translation/test_optional_params.py b/tests/llm_translation/test_optional_params.py index bc85b99eee7..95700eb29b9 100644 --- a/tests/llm_translation/test_optional_params.py +++ b/tests/llm_translation/test_optional_params.py @@ -1471,6 +1471,260 @@ def test_bedrock_invoke_anthropic_max_tokens(): assert optional_params["max_tokens"] == 1024 +def test_bedrock_invoke_claude_4_anthropic_max_tokens(): + passed_params = { + "model": "invoke/us.anthropic.claude-sonnet-4-5-20250929-v1:0", + "functions": None, + "function_call": None, + "temperature": 0.8, + "top_p": None, + "n": 1, + "stream": False, + "stream_options": None, + "stop": None, + "max_tokens": None, + "max_completion_tokens": 1024, + "modalities": None, + "prediction": None, + "audio": None, + "presence_penalty": None, + "frequency_penalty": None, + "logit_bias": None, + "user": None, + "custom_llm_provider": "bedrock", + "response_format": {"type": "text"}, + "seed": None, + "tools": [ + { + "type": "function", + "function": { + "name": "generate_plan", + "description": "Generate a plan to execute the task using only the tools outlined in your context.", + "input_schema": { + "type": "object", + "properties": { + "steps": { + "type": "array", + "items": { + "type": "object", + "properties": { + "type": { + "type": "string", + "description": "The type of step to execute", + }, + "tool_name": { + "type": "string", + "description": "The name of the tool to use for this step", + }, + "tool_input": { + "type": "object", + "description": "The input to pass to the tool. Make sure this complies with the schema for the tool.", + }, + "tool_output": { + "type": "object", + "description": "(Optional) The output from the tool if needed for future steps. Make sure this complies with the schema for the tool.", + }, + }, + "required": ["type"], + }, + } + }, + }, + }, + }, + { + "type": "function", + "function": { + "name": "generate_wire_tool", + "description": "Create a wire transfer with complete wire instructions", + "input_schema": { + "type": "object", + "properties": { + "company_id": { + "type": "integer", + "description": "The ID of the company receiving the investment", + }, + "investment_id": { + "type": "integer", + "description": "The ID of the investment memo", + }, + "dollar_amount": { + "type": "number", + "description": "The amount to wire in USD", + }, + "wiring_instructions": { + "type": "object", + "description": "Complete bank account and routing information for the wire", + "properties": { + "account_name": { + "type": "string", + "description": "Name on the bank account", + }, + "address_1": { + "type": "string", + "description": "Primary address line", + }, + "address_2": { + "type": "string", + "description": "Secondary address line (optional)", + }, + "city": {"type": "string"}, + "state": {"type": "string"}, + "zip": {"type": "string"}, + "country": {"type": "string", "default": "US"}, + "bank_name": {"type": "string"}, + "account_number": {"type": "string"}, + "routing_number": {"type": "string"}, + "account_type": { + "type": "string", + "enum": ["checking", "savings"], + "default": "checking", + }, + "swift_code": { + "type": "string", + "description": "Required for international wires", + }, + "iban": { + "type": "string", + "description": "Required for some international wires", + }, + "bank_city": {"type": "string"}, + "bank_state": {"type": "string"}, + "bank_country": {"type": "string", "default": "US"}, + "bank_to_bank_instructions": { + "type": "string", + "description": "Additional instructions for the bank (optional)", + }, + "intermediary_bank_name": { + "type": "string", + "description": "Name of intermediary bank if required (optional)", + }, + }, + "required": [ + "account_name", + "address_1", + "country", + "bank_name", + "account_number", + "routing_number", + "account_type", + "bank_country", + ], + }, + }, + "required": [ + "company_id", + "investment_id", + "dollar_amount", + "wiring_instructions", + ], + }, + }, + }, + { + "type": "function", + "function": { + "name": "search_companies", + "description": "Search for companies by name or other criteria to get their IDs", + "input_schema": { + "type": "object", + "properties": { + "query": { + "type": "string", + "description": "Name or part of name to search for", + }, + "batch": { + "type": "string", + "description": 'Optional batch filter (e.g., "W21", "S22")', + }, + "status": { + "type": "string", + "enum": [ + "live", + "dead", + "adrift", + "exited", + "went_public", + "all", + ], + "description": "Filter by company status", + "default": "live", + }, + "limit": { + "type": "integer", + "description": "Maximum number of results to return", + "default": 10, + }, + }, + "required": ["query"], + }, + "output_schema": { + "type": "object", + "properties": { + "status": { + "type": "string", + "description": "Success or error status", + }, + "results": { + "type": "array", + "description": "List of companies matching the search criteria", + "items": { + "type": "object", + "properties": { + "id": { + "type": "integer", + "description": "Company ID to use in other API calls", + }, + "name": {"type": "string"}, + "batch": {"type": "string"}, + "status": {"type": "string"}, + "valuation": {"type": "string"}, + "url": {"type": "string"}, + "description": {"type": "string"}, + "founders": {"type": "string"}, + }, + }, + }, + "results_count": { + "type": "integer", + "description": "Number of companies returned", + }, + "total_matches": { + "type": "integer", + "description": "Total number of matches found", + }, + }, + }, + }, + }, + ], + "tool_choice": None, + "max_retries": 0, + "logprobs": None, + "top_logprobs": None, + "extra_headers": None, + "api_version": None, + "parallel_tool_calls": None, + "drop_params": True, + "reasoning_effort": None, + "additional_drop_params": None, + "messages": [ + { + "role": "system", + "content": "You are an AI assistant that helps prepare a wire for a pro rata investment.", + }, + {"role": "user", "content": [{"type": "text", "text": "hi"}]}, + ], + "thinking": None, + "kwargs": {}, + } + optional_params = get_optional_params(**passed_params) + print(f"optional_params: {optional_params}") + + assert "max_tokens_to_sample" not in optional_params + assert optional_params["max_tokens"] == 1024 + + def test_azure_modalities_param(): optional_params = get_optional_params( model="chatgpt-v2", From 581d086c20fc391eec1f73df27393a8714ab23cd Mon Sep 17 00:00:00 2001 From: victorigualada <21220224+victorigualada@users.noreply.github.com> Date: Tue, 20 Jan 2026 05:29:50 +0100 Subject: [PATCH 07/11] fix(responses): stream tool call events in completion bridge (#19368) Emit Responses API streaming events for tool calls when the underlying chat stream contains tool_call deltas, and recover tool calls into the stream when they only appear in the final response. --- .../streaming_iterator.py | 179 +++++++++++++++++- ...test_tool_call_streaming_transformation.py | 116 ++++++++++++ 2 files changed, 294 insertions(+), 1 deletion(-) create mode 100644 tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index def2f72437d..dd7936059a0 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -21,6 +21,8 @@ from litellm.types.llms.openai import ( OutputTextAnnotationAddedEvent, OutputTextDeltaEvent, OutputTextDoneEvent, + FunctionCallArgumentsDeltaEvent, + FunctionCallArgumentsDoneEvent, ReasoningSummaryTextDeltaEvent, ResponseCompletedEvent, ResponseCreatedEvent, @@ -79,6 +81,161 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): Union[ModelResponse, TextCompletionResponse] ] = None self.final_text: str = "" + self._pending_tool_events: List[BaseLiteLLMOpenAIResponseObject] = [] + self._tool_output_index_by_call_id: dict[str, int] = {} + self._tool_args_by_call_id: dict[str, str] = {} + self._next_tool_output_index: int = 1 # output_index=0 reserved for the message item + self._final_tool_events_queued: bool = False + + def _get_or_assign_tool_output_index(self, call_id: str) -> int: + existing = self._tool_output_index_by_call_id.get(call_id) + if existing is not None: + return existing + idx = self._next_tool_output_index + self._next_tool_output_index += 1 + self._tool_output_index_by_call_id[call_id] = idx + return idx + + def _queue_tool_call_delta_events(self, tool_calls: object) -> None: + """ + Convert chat-completions streaming `tool_calls` deltas into Responses API streaming events. + + We emit: + - response.output_item.added (function_call) + - response.function_call_arguments.delta + """ + if not isinstance(tool_calls, list): + return + + for tc in tool_calls: + call_id_raw = tc.get("id") if isinstance(tc, dict) else getattr(tc, "id", None) + if not call_id_raw: + continue + call_id = str(call_id_raw) + + fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None) + fn_name = "" + fn_args_delta = "" + if isinstance(fn, dict): + fn_name = str(fn.get("name") or "") + fn_args_delta = str(fn.get("arguments") or "") + else: + fn_name = str(getattr(fn, "name", "") or "") + fn_args_delta = str(getattr(fn, "arguments", "") or "") + + output_index = self._get_or_assign_tool_output_index(call_id) + + if call_id not in self._tool_args_by_call_id: + self._tool_args_by_call_id[call_id] = "" + self._pending_tool_events.append( + OutputItemAddedEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=output_index, + item=BaseLiteLLMOpenAIResponseObject( + **{ + "type": "function_call", + "id": call_id, + "call_id": call_id, + "name": fn_name, + "arguments": "", + "status": "in_progress", + } + ), + ) + ) + + if fn_args_delta: + self._tool_args_by_call_id[call_id] += fn_args_delta + self._pending_tool_events.append( + FunctionCallArgumentsDeltaEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA, + item_id=call_id, + output_index=output_index, + delta=fn_args_delta, + ) + ) + + def _queue_final_tool_call_done_events(self, litellm_complete_object: ModelResponse) -> None: + """ + Ensure tool calls that were not streamed as deltas still get emitted before response.completed. + """ + if self._final_tool_events_queued: + return + self._final_tool_events_queued = True + + try: + message = litellm_complete_object.choices[0].message # type: ignore + tool_calls = getattr(message, "tool_calls", None) + except Exception: + tool_calls = None + + if not tool_calls or not isinstance(tool_calls, list): + return + + for tc in tool_calls: + call_id_raw = tc.get("id") if isinstance(tc, dict) else getattr(tc, "id", None) + if not call_id_raw: + continue + call_id = str(call_id_raw) + output_index = self._get_or_assign_tool_output_index(call_id) + + fn = tc.get("function") if isinstance(tc, dict) else getattr(tc, "function", None) + fn_name = "" + fn_args = "" + if isinstance(fn, dict): + fn_name = str(fn.get("name") or "") + fn_args = str(fn.get("arguments") or "") + else: + fn_name = str(getattr(fn, "name", "") or "") + fn_args = str(getattr(fn, "arguments", "") or "") + + # If we never sent output_item.added for this call_id, emit it now. + if call_id not in self._tool_args_by_call_id: + self._tool_args_by_call_id[call_id] = "" + self._pending_tool_events.append( + OutputItemAddedEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + output_index=output_index, + item=BaseLiteLLMOpenAIResponseObject( + **{ + "type": "function_call", + "id": call_id, + "call_id": call_id, + "name": fn_name, + "arguments": "", + "status": "in_progress", + } + ), + ) + ) + + final_args = fn_args or self._tool_args_by_call_id.get(call_id, "") + self._pending_tool_events.append( + FunctionCallArgumentsDoneEvent( + type=ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE, + item_id=call_id, + output_index=output_index, + arguments=final_args, + ) + ) + + self._pending_tool_events.append( + OutputItemDoneEvent( + type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, + output_index=output_index, + sequence_number=1, + item=BaseLiteLLMOpenAIResponseObject( + **{ + "type": "function_call", + "id": call_id, + "call_id": call_id, + "name": fn_name, + "arguments": final_args, + "status": "completed", + } + ), + ) + ) def _default_response_created_event_data(self) -> dict: response_created_event_data = { @@ -310,6 +467,12 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): ): self.litellm_model_response = self.create_litellm_model_response() if self.litellm_model_response: + # If tool calls exist, emit tool events before finishing/response.completed. + if isinstance(self.litellm_model_response, ModelResponse): + self._queue_final_tool_call_done_events(self.litellm_model_response) + if self._pending_tool_events: + return self._pending_tool_events.pop(0) + done_event = self.return_default_done_events(self.litellm_model_response) if done_event: return done_event @@ -462,13 +625,27 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): content_index=0, delta=delta_content, ) + + # Priority 3: Handle tool call deltas (if any) -> queue events and emit them + if ( + chunk.choices + and hasattr(chunk.choices[0].delta, "tool_calls") + and chunk.choices[0].delta.tool_calls + ): + self._queue_tool_call_delta_events(chunk.choices[0].delta.tool_calls) + if self._pending_tool_events: + return self._pending_tool_events.pop(0) - # Priority 3: If we have pending annotation events, emit the next one + # Priority 4: If we have pending annotation events, emit the next one # This happens when the current chunk has no text/reasoning content if hasattr(self, '_pending_annotation_events') and self._pending_annotation_events: event = self._pending_annotation_events.pop(0) return event + # Priority 5: If we have pending tool events (from earlier chunk), emit the next one + if self._pending_tool_events: + return self._pending_tool_events.pop(0) + return None def _get_delta_string_from_streaming_choices( diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py new file mode 100644 index 00000000000..51150383b01 --- /dev/null +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_call_streaming_transformation.py @@ -0,0 +1,116 @@ +""" +Tests for streaming tool-calls in Responses API transformation. + +Ensures that when the underlying chat-completions stream includes tool_calls deltas, +LiteLLM emits Responses API streaming events (output_item.added + function_call_arguments.*). + +Also ensures that tool calls that only appear in the final built response still get emitted +before response.completed. +""" + +from unittest.mock import AsyncMock + +from litellm.responses.litellm_completion_transformation.streaming_iterator import ( + LiteLLMCompletionStreamingIterator, +) +from litellm.types.llms.openai import ResponsesAPIStreamEvents +from litellm.types.utils import Delta, ModelResponse, ModelResponseStream, StreamingChoices + + +def test_tool_call_delta_is_emitted_as_responses_events(): + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + + # A streaming chunk with tool_calls delta but no text + chunk = ModelResponseStream( + id="chunk-1", + created=123, + model="test-model", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta( + role="assistant", + content="", + tool_calls=[ + { + "id": "call_1", + "type": "function", + "function": {"name": "do_thing", "arguments": '{"x":1}'}, + } + ], + ), + ) + ], + ) + + evt1 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk) + assert evt1 is not None + assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + assert evt1.output_index == 1 + + evt2 = iterator._transform_chat_completion_chunk_to_response_api_chunk(chunk) + assert evt2 is not None + assert evt2.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA + assert evt2.item_id == "call_1" + assert evt2.output_index == 1 + assert evt2.delta == '{"x":1}' + + +def test_tool_calls_present_only_in_final_response_are_emitted_before_completed(): + iterator = LiteLLMCompletionStreamingIterator( + model="test-model", + litellm_custom_stream_wrapper=AsyncMock(), + request_input="Test input", + responses_api_request={}, + ) + + # Construct a final ModelResponse with tool_calls on the message. + # We bypass the stream builder and directly set iterator.litellm_model_response. + response = ModelResponse( + id="resp-1", + created=123, + model="test-model", + object="chat.completion", + choices=[ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_2", + "type": "function", + "function": {"name": "do_thing", "arguments": '{"y":2}'}, + "index": 0, + } + ], + }, + } + ], + ) + iterator.litellm_model_response = response + + # First common_done_event_logic call should yield tool events, not response.completed. + evt1 = iterator.common_done_event_logic(sync_mode=True) + assert evt1.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED + assert evt1.output_index == 1 + + evt2 = iterator.common_done_event_logic(sync_mode=True) + assert evt2.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DONE + assert evt2.item_id == "call_2" + assert evt2.output_index == 1 + assert evt2.arguments == '{"y":2}' + + evt3 = iterator.common_done_event_logic(sync_mode=True) + assert evt3.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE + assert evt3.output_index == 1 + From 7d6d419a6704e46e97dacff7e4bbc13d22e1b74c Mon Sep 17 00:00:00 2001 From: victorigualada <21220224+victorigualada@users.noreply.github.com> Date: Tue, 20 Jan 2026 05:37:59 +0100 Subject: [PATCH 08/11] fix: preserve tool output ordering for gemini in responses bridge (#19360) * fix: preserve tool output ordering for gemini in responses bridge - Keep function_call_output adjacent to its function_call when building chat messages - Normalize function_call_output.output lists (input_* parts) into tool message content * fix test * small improvements --- .../transformation.py | 151 ++++++++++++++---- ...test_function_call_output_normalization.py | 40 +++++ ..._tool_output_order_preserved_for_gemini.py | 78 +++++++++ 3 files changed, 242 insertions(+), 27 deletions(-) create mode 100644 tests/test_litellm/responses/litellm_completion_transformation/test_function_call_output_normalization.py create mode 100644 tests/test_litellm/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index eaa80c6cfe4..3badbc50578 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -367,14 +367,6 @@ class LiteLLMCompletionResponsesConfig: ChatCompletionResponseMessage, ] ] = [] - tool_call_output_messages: List[ - Union[ - AllMessageValues, - GenericChatCompletionMessage, - ChatCompletionMessageToolCall, - ChatCompletionResponseMessage, - ] - ] = [] if isinstance(input, str): messages.append(ChatCompletionUserMessage(role="user", content=input)) @@ -385,15 +377,6 @@ class LiteLLMCompletionResponsesConfig: input_item=_input ) - ######################################################### - # If Input Item is a Tool Call Output, add it to the tool_call_output_messages list - ######################################################### - if LiteLLMCompletionResponsesConfig._is_input_item_tool_call_output( - input_item=_input - ): - tool_call_output_messages.extend(chat_completion_messages) - continue - if LiteLLMCompletionResponsesConfig._is_input_item_function_call( input_item=_input ): @@ -401,15 +384,57 @@ class LiteLLMCompletionResponsesConfig: if call_id_raw: existing_tool_call_ids.add(str(call_id_raw)) - messages.extend(chat_completion_messages) + ######################################################### + # If Input Item is a Tool Call Output, add it to the tool_call_output_messages list + # preserving the ordering of tool call outputs. Some models require the tool + # result to immediately follow the assistant tool call. + ######################################################### + if LiteLLMCompletionResponsesConfig._is_input_item_tool_call_output( + input_item=_input + ): + if not chat_completion_messages: + continue - deduped_tool_call_messages = ( - LiteLLMCompletionResponsesConfig._deduplicate_tool_call_output_messages( - tool_call_output_messages=tool_call_output_messages, - existing_tool_call_ids=existing_tool_call_ids, - ) - ) - messages.extend(deduped_tool_call_messages) + deduped_in_place: List[Any] = [] + for m in chat_completion_messages: + role = "" + if isinstance(m, dict): + role = str(m.get("role") or "") + else: + role = str(getattr(m, "role", "") or "") + + # Drop assistant tool_calls wrappers if we already have this call_id + if role == "assistant": + tool_calls: Any = ( + m.get("tool_calls") + if isinstance(m, dict) + else getattr(m, "tool_calls", None) + ) + call_id = "" + if ( + isinstance(tool_calls, Sequence) + and not isinstance(tool_calls, (str, bytes)) + and len(tool_calls) > 0 + ): + first_call = tool_calls[0] + call_id_raw = ( + first_call.get("id") + if isinstance(first_call, dict) + else getattr(first_call, "id", None) + ) + if call_id_raw: + call_id = str(call_id_raw) + if call_id and call_id in existing_tool_call_ids: + continue + if call_id: + existing_tool_call_ids.add(call_id) + + deduped_in_place.append(m) + + messages.extend(deduped_in_place) + continue + + messages.extend(chat_completion_messages) return messages @staticmethod @@ -821,10 +846,82 @@ class LiteLLMCompletionResponsesConfig: # Empty call_id means we can't create a valid tool message if not call_id: return [] - + + def _normalize_function_call_output_to_tool_content( + output: Any, + ) -> Any: + """ + Normalize Responses API function_call_output.output into a shape that downstream + chat adapters (esp. Gemini) can reliably consume. + + OpenAI Responses API typically uses: + - output: string + + Some clients/adapters send: + - output: [{"type": "input_text", "text": "..."}, {"type": "input_image", ...}] + + For chat tool messages we normalize to either: + - string (preferred) + - list of {"type": "text"|"image_url", ...} blocks (for multimodal tool outputs) + """ + if output is None: + return "" + if isinstance(output, str): + return output + + # Some adapters represent tool output as a list of "input_*" parts + if isinstance(output, list): + normalized_blocks: List[Dict[str, Any]] = [] + text_acc: List[str] = [] + for part in output: + if not isinstance(part, dict): + continue + part_type = part.get("type") + if part_type in ("input_text", "output_text", "text"): + txt = part.get("text") + if isinstance(txt, str) and txt: + text_acc.append(txt) + normalized_blocks.append({"type": "text", "text": txt}) + elif part_type in ("input_image", "image_url"): + image_url_val = part.get("image_url") or part.get("url") + if isinstance(image_url_val, dict): + url = image_url_val.get("url") + if isinstance(url, str) and url: + normalized_blocks.append( + {"type": "image_url", "image_url": {"url": url}} + ) + elif isinstance(image_url_val, str) and image_url_val: + normalized_blocks.append( + {"type": "image_url", "image_url": {"url": image_url_val}} + ) + + # Prefer structured blocks if we have images; otherwise return a string. + if any(b.get("type") == "image_url" for b in normalized_blocks): + # Ensure we include any accumulated text as text blocks too + return normalized_blocks + if text_acc: + return "".join(text_acc) + try: + # last resort: keep something meaningful for providers that require a string + import json as _json + + return _json.dumps(output) + except Exception: + return str(output) + + # Fallback for dict/number/etc. + try: + import json as _json + + return _json.dumps(output) + except Exception: + return str(output) + tool_output_message = ChatCompletionToolMessage( role="tool", - content=tool_call_output.get("output") or "", + content=_normalize_function_call_output_to_tool_content( + tool_call_output.get("output") + ), tool_call_id=str(call_id), ) diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_function_call_output_normalization.py b/tests/test_litellm/responses/litellm_completion_transformation/test_function_call_output_normalization.py new file mode 100644 index 00000000000..19aeba7f9cd --- /dev/null +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_function_call_output_normalization.py @@ -0,0 +1,40 @@ +""" +Tests for normalizing Responses API function_call_output into chat tool messages. + +This is important for Gemini/Vertex, which expects tool results to be represented +as tool/function response parts; if the tool output is passed as a list of input_* parts, +we normalize it to text/image blocks or a string. +""" + +from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, +) + + +def test_function_call_output_list_input_text_is_converted_to_tool_string_content(): + out = LiteLLMCompletionResponsesConfig._transform_responses_api_tool_call_output_to_chat_completion_message( + tool_call_output={ + "type": "function_call_output", + "call_id": "call_1", + "output": [{"type": "input_text", "text": "hello"}, {"type": "input_text", "text": " world"}], + } + ) + + assert len(out) == 1 + msg = out[0] + assert msg["role"] == "tool" + assert msg["tool_call_id"] == "call_1" + assert msg["content"] == "hello world" + + +def test_function_call_output_string_passthrough(): + out = LiteLLMCompletionResponsesConfig._transform_responses_api_tool_call_output_to_chat_completion_message( + tool_call_output={ + "type": "function_call_output", + "call_id": "call_1", + "output": '{"ok":true}', + } + ) + assert len(out) == 1 + assert out[0]["content"] == '{"ok":true}' + diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py new file mode 100644 index 00000000000..5cb01fbae61 --- /dev/null +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_tool_output_order_preserved_for_gemini.py @@ -0,0 +1,78 @@ +""" +Regression: preserve function_call_output ordering. + +Gemini/Vertex requires tool outputs to immediately follow the assistant tool call. +The ResponsesAPI->Chat conversion must not move tool outputs to the end. +""" + +from litellm.responses.litellm_completion_transformation.transformation import ( + LiteLLMCompletionResponsesConfig, +) + + +def test_function_call_output_stays_adjacent_to_tool_call(): + msgs = LiteLLMCompletionResponsesConfig._transform_response_input_param_to_chat_completion_message( + input=[ + { + "role": "user", + "type": "message", + "content": [{"type": "input_text", "text": "Call echo with 'hello'."}], + }, + { + "type": "function_call", + "name": "echo", + "call_id": "call_123", + "arguments": '{"text":"hello"}', + }, + { + "type": "function_call_output", + "call_id": "call_123", + "output": '{"text":"hello"}', + }, + { + "role": "assistant", + "type": "message", + "content": [{"type": "output_text", "text": "Done."}], + }, + { + "role": "user", + "type": "message", + "content": [{"type": "input_text", "text": "Now say hi."}], + }, + ] + ) + + # Find the assistant message that contains tool_calls + tool_call_idx = None + tool_msg_idx = None + assistant_ok_idx = None + + for i, m in enumerate(msgs): + if isinstance(m, dict) and m.get("role") == "assistant" and m.get("tool_calls"): + tool_call_idx = i + if isinstance(m, dict) and m.get("role") == "tool": + tool_msg_idx = i + + # Assistant "Done." can be either a plain string or a structured content list + if isinstance(m, dict) and m.get("role") == "assistant": + content = m.get("content") + if content == "Done.": + assistant_ok_idx = i + elif isinstance(content, list): + for block in content: + if ( + isinstance(block, dict) + and block.get("type") == "text" + and block.get("text") == "Done." + ): + assistant_ok_idx = i + break + + assert tool_call_idx is not None + assert tool_msg_idx is not None + assert assistant_ok_idx is not None + + # Tool output must be right after tool call, and before the assistant "Done." message. + assert tool_msg_idx == tool_call_idx + 1 + assert assistant_ok_idx > tool_msg_idx + From f0785d5a5154f07c68429d2d5067656f487f9beb Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 20 Jan 2026 17:26:40 +0530 Subject: [PATCH 09/11] Fix:test_supported_params_limited_to_docs --- .../responses/test_volcengine_responses_transformation.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py index 16930bcb995..823fd82d1ce 100644 --- a/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py +++ b/tests/test_litellm/llms/volcengine/responses/test_volcengine_responses_transformation.py @@ -14,8 +14,8 @@ from litellm.llms.volcengine.responses.transformation import ( VolcEngineResponsesAPIConfig, ) from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams -from litellm.types.router import GenericLiteLLMParams from litellm.types.responses.main import DeleteResponseResult +from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import LlmProviders from litellm.utils import ProviderConfigManager @@ -204,6 +204,7 @@ class TestVolcengineResponsesAPITransformation: "thinking", "caching", "expire_at", + "context_management", "extra_headers", "extra_query", "extra_body", From cd96c8cbb0b9a33ac06a8dab5a185dac61ff55ed Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 20 Jan 2026 17:39:08 +0530 Subject: [PATCH 10/11] Fix:test_aaaaazure_tenant_id_auth --- litellm/llms/azure/azure.py | 4 ++-- litellm/llms/custom_httpx/http_handler.py | 10 +++++++--- litellm/llms/openai/common_utils.py | 3 ++- tests/local_testing/test_azure_openai.py | 5 +++++ 4 files changed, 16 insertions(+), 6 deletions(-) diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 3ef0186ba0e..7250d6fc94c 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -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 += "/" diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 7fdb78c1670..57a6d04c995 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1168,8 +1168,10 @@ def get_async_httpx_client( return _cached_client if params is not None: - params["shared_session"] = shared_session - _new_client = AsyncHTTPHandler(**params) + # Filter out params that are only used for cache key, not for AsyncHTTPHandler.__init__ + handler_params = {k: v for k, v in params.items() if k != "disable_aiohttp_transport"} + handler_params["shared_session"] = shared_session + _new_client = AsyncHTTPHandler(**handler_params) else: _new_client = AsyncHTTPHandler( timeout=httpx.Timeout(timeout=600.0, connect=5.0), @@ -1215,7 +1217,9 @@ def _get_httpx_client(params: Optional[dict] = None) -> HTTPHandler: return _cached_client if params is not None: - _new_client = HTTPHandler(**params) + # Filter out params that are only used for cache key, not for HTTPHandler.__init__ + handler_params = {k: v for k, v in params.items() if k != "disable_aiohttp_transport"} + _new_client = HTTPHandler(**handler_params) else: _new_client = HTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0)) diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index 9f1ec9250cf..8bcecd35232 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -212,7 +212,8 @@ class BaseOpenAILLM: 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 + 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( diff --git a/tests/local_testing/test_azure_openai.py b/tests/local_testing/test_azure_openai.py index ed0dda3e15a..e95c1b6fcce 100644 --- a/tests/local_testing/test_azure_openai.py +++ b/tests/local_testing/test_azure_openai.py @@ -41,6 +41,11 @@ async def test_aaaaazure_tenant_id_auth(respx_mock: MockRouter): PROD Test """ litellm.disable_aiohttp_transport = True # since this uses respx, we need to set use_aiohttp_transport to False + + # Clear the HTTP client cache to ensure respx mocking works + # This is critical because respx only intercepts clients created AFTER mocking is active + if hasattr(litellm, 'in_memory_llm_clients_cache'): + litellm.in_memory_llm_clients_cache.flush_cache() router = Router( model_list=[ From 8b2472063849d559df57db9a1bdac555c7c85463 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 20 Jan 2026 18:20:22 +0530 Subject: [PATCH 11/11] fix: test_standard_logging_payload_includes_guardrail_information --- tests/guardrails_tests/test_tracing_guardrails.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/tests/guardrails_tests/test_tracing_guardrails.py b/tests/guardrails_tests/test_tracing_guardrails.py index 02ff7c0e4f6..8e7ce27bc28 100644 --- a/tests/guardrails_tests/test_tracing_guardrails.py +++ b/tests/guardrails_tests/test_tracing_guardrails.py @@ -74,7 +74,7 @@ async def test_standard_logging_payload_includes_guardrail_information(): class MockClientSession: def __init__(self): - pass + self.closed = False async def __aenter__(self): return self @@ -82,6 +82,9 @@ async def test_standard_logging_payload_includes_guardrail_information(): async def __aexit__(self, exc_type, exc_val, exc_tb): pass + async def close(self): + self.closed = True + def post(self, url, json=None): class MockResponse: def __init__(self, response_obj):