mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 3447c3529a into b781d157d7
This commit is contained in:
commit
4a106efc86
2 changed files with 71 additions and 6 deletions
|
|
@ -65,13 +65,15 @@ def timeout(timeout_duration: float = 0.0, exception_to_raise=Timeout):
|
|||
|
||||
@wraps(func)
|
||||
async def async_wrapper(*args, **kwargs):
|
||||
local_timeout_duration = timeout_duration
|
||||
if "force_timeout" in kwargs:
|
||||
local_timeout_duration = kwargs["force_timeout"]
|
||||
elif "request_timeout" in kwargs and kwargs["request_timeout"] is not None:
|
||||
local_timeout_duration = kwargs["request_timeout"]
|
||||
local_timeout_duration: Final = (
|
||||
kwargs["force_timeout"]
|
||||
if kwargs.get("force_timeout") is not None
|
||||
else kwargs["request_timeout"]
|
||||
if kwargs.get("request_timeout") is not None
|
||||
else timeout_duration
|
||||
)
|
||||
try:
|
||||
value: Final = await asyncio.wait_for(func(*args, **kwargs), timeout=timeout_duration)
|
||||
value: Final = await asyncio.wait_for(func(*args, **kwargs), timeout=local_timeout_duration)
|
||||
return value
|
||||
except asyncio.TimeoutError:
|
||||
model: Final = args[0] if len(args) > 0 else kwargs["model"]
|
||||
|
|
|
|||
63
tests/unit/test_timeout.py
Normal file
63
tests/unit/test_timeout.py
Normal file
|
|
@ -0,0 +1,63 @@
|
|||
"""Unit tests for litellm.timeout decorator."""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.exceptions import Timeout
|
||||
from litellm.timeout import timeout
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("timeout_arg", ["request_timeout", "force_timeout"])
|
||||
async def test_async_timeout_decorator_enforces_per_call_timeout(timeout_arg: str) -> None:
|
||||
@timeout(timeout_duration=60.0)
|
||||
async def hung_func(**kwargs):
|
||||
await asyncio.sleep(10.0)
|
||||
return "never"
|
||||
|
||||
with pytest.raises(Timeout) as exc_info:
|
||||
await hung_func(**{timeout_arg: 0.001, "model": "test-model"})
|
||||
|
||||
assert "0.001 second(s)" in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("timeout_arg", ["request_timeout", "force_timeout"])
|
||||
async def test_async_timeout_decorator_extends_duration_with_per_call_timeout(
|
||||
timeout_arg: str,
|
||||
) -> None:
|
||||
@timeout(timeout_duration=0.001)
|
||||
async def delayed_func(**kwargs):
|
||||
await asyncio.sleep(0.01)
|
||||
return "completed"
|
||||
|
||||
result: Final = await delayed_func(**{timeout_arg: 60.0, "model": "test-model"})
|
||||
assert result == "completed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_timeout_decorator_handles_none_force_timeout() -> None:
|
||||
@timeout(timeout_duration=0.001)
|
||||
async def hung_func(**kwargs):
|
||||
await asyncio.sleep(10.0)
|
||||
return "never"
|
||||
|
||||
with pytest.raises(Timeout) as exc_info:
|
||||
await hung_func(force_timeout=None, model="test-model")
|
||||
|
||||
assert "0.001 second(s)" in str(exc_info.value)
|
||||
|
||||
|
||||
def test_sync_timeout_decorator_enforces_per_call_timeout() -> None:
|
||||
@timeout(timeout_duration=60.0)
|
||||
def hung_sync_func(**kwargs):
|
||||
time.sleep(10.0)
|
||||
return "never"
|
||||
|
||||
with pytest.raises(Timeout) as exc_info:
|
||||
hung_sync_func(request_timeout=0.001, model="test-model")
|
||||
|
||||
assert "0.001 second(s)" in str(exc_info.value)
|
||||
Loading…
Add table
Reference in a new issue