forward timeout to make_call() for Bedrock and Vertex AI streaming

This commit is contained in:
pradyyadav 2026-03-12 08:49:45 +05:30
parent 1f721fd133
commit 5eecd2df94
5 changed files with 171 additions and 6 deletions

View file

@ -149,6 +149,7 @@ class BedrockConverseLLM(BaseAWSLLM):
fake_stream=fake_stream, fake_stream=fake_stream,
json_mode=json_mode, json_mode=json_mode,
stream_chunk_size=stream_chunk_size, stream_chunk_size=stream_chunk_size,
timeout=timeout,
) )
streaming_response = CustomStreamWrapper( streaming_response = CustomStreamWrapper(
completion_stream=completion_stream, completion_stream=completion_stream,

View file

@ -194,16 +194,20 @@ async def make_call(
json_mode: Optional[bool] = False, json_mode: Optional[bool] = False,
bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None, bedrock_invoke_provider: Optional[litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL] = None,
stream_chunk_size: int = 1024, stream_chunk_size: int = 1024,
timeout: Optional[Union[float, httpx.Timeout]] = None,
): ):
try: try:
if client is None: 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( client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.BEDROCK, llm_provider=litellm.LlmProviders.BEDROCK,
params={"ssl_verify": logging_obj.litellm_params.get("ssl_verify")} params=_params if _params else None,
if logging_obj
and logging_obj.litellm_params
and logging_obj.litellm_params.get("ssl_verify")
else None,
) # Create a new client if none provided ) # Create a new client if none provided
response = await client.post( response = await client.post(
@ -212,6 +216,7 @@ async def make_call(
data=data, data=data,
stream=not fake_stream, stream=not fake_stream,
logging_obj=logging_obj, logging_obj=logging_obj,
timeout=timeout,
) )
if response.status_code != 200: if response.status_code != 200:
@ -1240,6 +1245,7 @@ class BedrockLLM(BaseAWSLLM):
logging_obj=logging_obj, logging_obj=logging_obj,
fake_stream=True if "ai21" in api_base else False, fake_stream=True if "ai21" in api_base else False,
stream_chunk_size=stream_chunk_size, stream_chunk_size=stream_chunk_size,
timeout=timeout,
), ),
model=model, model=model,
custom_llm_provider="bedrock", custom_llm_provider="bedrock",

View file

@ -2380,17 +2380,25 @@ async def make_call(
model: str, model: str,
messages: list, messages: list,
logging_obj, logging_obj,
timeout: Optional[Union[float, httpx.Timeout]] = None,
): ):
if gemini_client is not None: if gemini_client is not None:
client = gemini_client client = gemini_client
if client is None: 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( client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.VERTEX_AI, llm_provider=litellm.LlmProviders.VERTEX_AI,
params=_async_client_params if _async_client_params else None,
) )
try: try:
response = await client.post( 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() response.raise_for_status()
except httpx.HTTPStatusError as e: except httpx.HTTPStatusError as e:
@ -2565,6 +2573,7 @@ class VertexLLM(VertexBase):
model=model, model=model,
messages=messages, messages=messages,
logging_obj=logging_obj, logging_obj=logging_obj,
timeout=timeout,
), ),
model=model, model=model,
custom_llm_provider="vertex_ai_beta", custom_llm_provider="vertex_ai_beta",

View file

@ -200,3 +200,54 @@ def test_bedrock_converse_streaming_consistent_id():
assert ( assert (
response.id == expected_id response.id == expected_id
), "All chunk IDs must match the one captured from the messageStart event" ), "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()"
)

View file

@ -3789,3 +3789,101 @@ def test_sync_streaming_uses_custom_client():
# Verify that gemini_client is in the partial's keywords # Verify that gemini_client is in the partial's keywords
assert "gemini_client" in partial_make_sync_call.keywords assert "gemini_client" in partial_make_sync_call.keywords
assert partial_make_sync_call.keywords["gemini_client"] is mock_client 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"
)