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,
|
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,
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
|
|
@ -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",
|
||||||
|
|
|
||||||
|
|
@ -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()"
|
||||||
|
)
|
||||||
|
|
|
||||||
|
|
@ -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"
|
||||||
|
)
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue