mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
address review: normalize timeout consistently, fix sync converse path, add tests
This commit is contained in:
parent
bbb3fd2b2d
commit
64c50b86ba
5 changed files with 167 additions and 5 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue