diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 162914fb292..3e6f7e4be7b 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -33,9 +33,18 @@ def make_sync_call( json_mode: Optional[bool] = False, fake_stream: bool = False, stream_chunk_size: int = 1024, + timeout: Optional[Union[float, httpx.Timeout]] = None, ): + if timeout is not None and isinstance(timeout, (float, int)): + timeout = httpx.Timeout(timeout) + if client is None: - client = _get_httpx_client() # Create a new client if none provided + _params: dict = {} + if timeout is not None: + _params["timeout"] = timeout + client = _get_httpx_client( + params=_params if _params else None + ) response = client.post( api_base, @@ -43,6 +52,7 @@ def make_sync_call( data=data, stream=not fake_stream, logging_obj=logging_obj, + timeout=timeout, ) if response.status_code != 200: @@ -469,6 +479,7 @@ class BedrockConverseLLM(BaseAWSLLM): json_mode=json_mode, fake_stream=fake_stream, stream_chunk_size=stream_chunk_size, + timeout=timeout, ) streaming_response = CustomStreamWrapper( completion_stream=completion_stream, diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 4a11e3eb158..518b21ac404 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -197,13 +197,14 @@ async def make_call( timeout: Optional[Union[float, httpx.Timeout]] = None, ): try: + if timeout is not None and isinstance(timeout, (float, int)): + timeout = httpx.Timeout(timeout) + if client is None: _params: dict = {} if logging_obj and logging_obj.litellm_params and logging_obj.litellm_params.get("ssl_verify"): _params["ssl_verify"] = logging_obj.litellm_params.get("ssl_verify") if timeout is not None: - if isinstance(timeout, (float, int)): - timeout = httpx.Timeout(timeout) _params["timeout"] = timeout client = get_async_httpx_client( llm_provider=litellm.LlmProviders.BEDROCK, diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 88fcbd62bcf..1258e2b72bd 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -2382,13 +2382,14 @@ async def make_call( logging_obj, timeout: Optional[Union[float, httpx.Timeout]] = None, ): + if timeout is not None and isinstance(timeout, (float, int)): + timeout = httpx.Timeout(timeout) + if gemini_client is not None: client = gemini_client if client is None: _params: dict = {} if timeout is not None: - if isinstance(timeout, (float, int)): - timeout = httpx.Timeout(timeout) _params["timeout"] = timeout client = get_async_httpx_client( llm_provider=litellm.LlmProviders.VERTEX_AI, diff --git a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py index b0dc7e17f5a..f178d8b69ae 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py @@ -251,3 +251,105 @@ def test_bedrock_invoke_async_streaming_passes_timeout_to_make_call(): assert make_call_partial.keywords.get("timeout") == timeout, ( "timeout must be forwarded via partial() to make_call()" ) + + +def test_bedrock_converse_async_streaming_passes_timeout_to_make_call(): + """ + BedrockConverseLLM.async_streaming() must forward timeout to make_call() + so streaming requests respect the user-configured timeout. + + Fixes https://github.com/BerriAI/litellm/issues/23375 + """ + import asyncio + from unittest.mock import AsyncMock, MagicMock, patch + + import httpx + + from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM + + handler = BedrockConverseLLM() + timeout = httpx.Timeout(7.0) + + mock_completion_stream = MagicMock() + + fake_prepped = MagicMock() + fake_prepped.headers = {"Authorization": "test"} + + credentials = MagicMock() + + with patch( + "litellm.llms.bedrock.chat.converse_handler.make_call", + new_callable=AsyncMock, + return_value=mock_completion_stream, + ) as mock_make_call, patch( + "litellm.AmazonConverseConfig", + ) as mock_converse_config, patch.object( + handler, "get_request_headers", return_value=fake_prepped, + ): + mock_converse_config.return_value._async_transform_request = AsyncMock( + return_value={"messages": []} + ) + + async def run(): + await handler.async_streaming( + model="anthropic.claude-3-sonnet", + messages=[{"role": "user", "content": "hi"}], + api_base="https://example.com", + model_response=MagicMock(), + timeout=timeout, + encoding=None, + logging_obj=MagicMock(), + stream=True, + optional_params={}, + litellm_params={"aws_region_name": "us-west-2"}, + credentials=credentials, + ) + + asyncio.run(run()) + + _, kwargs = mock_make_call.call_args + assert kwargs.get("timeout") == timeout, ( + "timeout must be forwarded to make_call() in BedrockConverseLLM.async_streaming()" + ) + + +def test_bedrock_converse_sync_make_sync_call_passes_timeout_to_client_post(): + """ + make_sync_call() in converse_handler must forward timeout to client.post() + so synchronous streaming requests also respect the user-configured timeout. + + Fixes https://github.com/BerriAI/litellm/issues/23375 + """ + from unittest.mock import MagicMock, patch + + import httpx + + from litellm.llms.bedrock.chat.converse_handler import make_sync_call + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.iter_bytes = MagicMock(return_value=iter([])) + + mock_client = MagicMock() + mock_client.post = MagicMock(return_value=mock_response) + + timeout = httpx.Timeout(4.0) + + with patch( + "litellm.llms.bedrock.chat.converse_handler.AWSEventStreamDecoder" + ): + make_sync_call( + client=mock_client, + api_base="https://example.com", + headers={}, + data="{}", + model="anthropic.claude-3-sonnet", + messages=[], + logging_obj=MagicMock(), + timeout=timeout, + ) + + _, kwargs = mock_client.post.call_args + assert kwargs.get("timeout") == timeout, ( + "timeout must be forwarded to client.post() in converse make_sync_call()" + ) diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 206ad5a0e6b..ecc9d645ab2 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -3887,3 +3887,50 @@ def test_vertex_make_call_creates_client_with_timeout_when_no_client_provided(): assert kwargs.get("params", {}).get("timeout") == timeout, ( "timeout must be passed to get_async_httpx_client() params" ) + + +def test_vertex_make_call_passes_timeout_with_gemini_client(): + """ + When a user-provided gemini_client is passed, make_call() must still + forward timeout to client.post() so the per-request timeout is respected. + + Fixes https://github.com/BerriAI/litellm/issues/23375 + """ + import asyncio + from unittest.mock import AsyncMock, MagicMock + + import httpx + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + make_call, + ) + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.aiter_lines = AsyncMock(return_value=iter([])) + mock_response.raise_for_status = MagicMock() + + mock_gemini_client = MagicMock() + mock_gemini_client.post = AsyncMock(return_value=mock_response) + + timeout = httpx.Timeout(10.0) + + async def run(): + await make_call( + client=None, + gemini_client=mock_gemini_client, + api_base="https://example.com", + headers={}, + data="{}", + model="gemini-2.0-flash", + messages=[], + logging_obj=MagicMock(), + timeout=timeout, + ) + + asyncio.run(run()) + + _, kwargs = mock_gemini_client.post.call_args + assert kwargs.get("timeout") == timeout, ( + "timeout must be forwarded to client.post() when gemini_client is provided" + )