diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 26986aab586..162914fb292 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -149,6 +149,7 @@ class BedrockConverseLLM(BaseAWSLLM): fake_stream=fake_stream, json_mode=json_mode, 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 9b06e198203..4a11e3eb158 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -194,16 +194,20 @@ async def make_call( json_mode: Optional[bool] = False, bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None, stream_chunk_size: int = 1024, + timeout: Optional[Union[float, httpx.Timeout]] = None, ): try: 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, - params={"ssl_verify": logging_obj.litellm_params.get("ssl_verify")} - if logging_obj - and logging_obj.litellm_params - and logging_obj.litellm_params.get("ssl_verify") - else None, + params=_params if _params else None, ) # Create a new client if none provided response = await client.post( @@ -212,6 +216,7 @@ async def make_call( data=data, stream=not fake_stream, logging_obj=logging_obj, + timeout=timeout, ) if response.status_code != 200: @@ -1240,6 +1245,7 @@ class BedrockLLM(BaseAWSLLM): logging_obj=logging_obj, fake_stream=True if "ai21" in api_base else False, stream_chunk_size=stream_chunk_size, + timeout=timeout, ), model=model, custom_llm_provider="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 df7a4a6511d..6108a1c20bb 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 @@ -2380,17 +2380,25 @@ async def make_call( model: str, messages: list, logging_obj, + timeout: Optional[Union[float, httpx.Timeout]] = None, ): if gemini_client is not None: client = gemini_client if client is None: + _async_client_params: dict = {} + if timeout is not None: + if isinstance(timeout, (float, int)): + timeout = httpx.Timeout(timeout) + _async_client_params["timeout"] = timeout client = get_async_httpx_client( llm_provider=litellm.LlmProviders.VERTEX_AI, + params=_async_client_params if _async_client_params else None, ) try: response = await client.post( - api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj + api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj, + timeout=timeout, ) response.raise_for_status() except httpx.HTTPStatusError as e: @@ -2565,6 +2573,7 @@ class VertexLLM(VertexBase): model=model, messages=messages, logging_obj=logging_obj, + timeout=timeout, ), model=model, custom_llm_provider="vertex_ai_beta", 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 a415d550215..b0dc7e17f5a 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_invoke_handler.py @@ -200,3 +200,54 @@ def test_bedrock_converse_streaming_consistent_id(): assert ( response.id == expected_id ), "All chunk IDs must match the one captured from the messageStart event" + + +def test_bedrock_invoke_async_streaming_passes_timeout_to_make_call(): + """ + async_streaming() in BedrockInvokeModelHandler must include timeout in the + partial() call 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.invoke_handler import BedrockLLM + + handler = BedrockLLM() + timeout = httpx.Timeout(5.0) + captured_partial = {} + + class FakeCustomStreamWrapper: + def __init__(self, *args, make_call=None, **kwargs): + captured_partial["make_call"] = make_call + + with patch( + "litellm.llms.bedrock.chat.invoke_handler.CustomStreamWrapper", + FakeCustomStreamWrapper, + ): + 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(), + print_verbose=MagicMock(), + data='{"prompt": "hi"}', + timeout=timeout, + encoding=None, + logging_obj=MagicMock(), + stream=True, + optional_params={}, + ) + + asyncio.run(run()) + + make_call_partial = captured_partial.get("make_call") + assert make_call_partial is not None + assert make_call_partial.keywords.get("timeout") == timeout, ( + "timeout must be forwarded via partial() to make_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 965fc03a33d..206ad5a0e6b 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 @@ -3789,3 +3789,101 @@ def test_sync_streaming_uses_custom_client(): # Verify that gemini_client is in the partial's keywords assert "gemini_client" in partial_make_sync_call.keywords assert partial_make_sync_call.keywords["gemini_client"] is mock_client + + +def test_vertex_make_call_passes_timeout_to_client_post(): + """ + make_call() in vertex_and_google_ai_studio_gemini must forward timeout + to client.post() so streaming requests respect the user-configured timeout. + + 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_client = MagicMock() + mock_client.post = AsyncMock(return_value=mock_response) + + timeout = httpx.Timeout(3.0) + + async def run(): + await make_call( + client=mock_client, + gemini_client=None, + api_base="https://example.com", + headers={}, + data="{}", + model="gemini-2.0-flash", + messages=[], + logging_obj=MagicMock(), + timeout=timeout, + ) + + asyncio.run(run()) + + _, kwargs = mock_client.post.call_args + assert kwargs.get("timeout") == timeout, ( + "timeout must be forwarded to client.post() for streaming to respect it" + ) + + +def test_vertex_make_call_creates_client_with_timeout_when_no_client_provided(): + """ + When no client is provided, make_call() must pass timeout to + get_async_httpx_client() so the created client has the correct timeout. + + Fixes https://github.com/BerriAI/litellm/issues/23375 + """ + import asyncio + from unittest.mock import AsyncMock, MagicMock, patch + + import httpx + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + make_call, + ) + + timeout = httpx.Timeout(3.0) + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.aiter_lines = AsyncMock(return_value=iter([])) + mock_response.raise_for_status = MagicMock() + + mock_created_client = MagicMock() + mock_created_client.post = AsyncMock(return_value=mock_response) + + with patch( + "litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.get_async_httpx_client", + return_value=mock_created_client, + ) as mock_get_client: + async def run(): + await make_call( + client=None, + gemini_client=None, + api_base="https://example.com", + headers={}, + data="{}", + model="gemini-2.0-flash", + messages=[], + logging_obj=MagicMock(), + timeout=timeout, + ) + + asyncio.run(run()) + + _, kwargs = mock_get_client.call_args + assert kwargs.get("params", {}).get("timeout") == timeout, ( + "timeout must be passed to get_async_httpx_client() params" + )