diff --git a/litellm/timeout.py b/litellm/timeout.py index 33d1ce92e53..8a27f7c3fe5 100644 --- a/litellm/timeout.py +++ b/litellm/timeout.py @@ -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"] diff --git a/tests/unit/test_timeout.py b/tests/unit/test_timeout.py new file mode 100644 index 00000000000..4f8a247f90c --- /dev/null +++ b/tests/unit/test_timeout.py @@ -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)