address review: normalize timeout consistently, fix sync converse path, add tests

This commit is contained in:
pradyyadav 2026-03-13 00:31:30 +05:30
parent bbb3fd2b2d
commit 64c50b86ba
5 changed files with 167 additions and 5 deletions

View file

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

View file

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

View file

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

View file

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

View file

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