mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
forward timeout to make_call() for Bedrock and Vertex AI streaming
This commit is contained in:
parent
1f721fd133
commit
5eecd2df94
5 changed files with 171 additions and 6 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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()"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue