diff --git a/evoagentx/models/litellm_model.py b/evoagentx/models/litellm_model.py index 9f4fb755..6bc1b723 100644 --- a/evoagentx/models/litellm_model.py +++ b/evoagentx/models/litellm_model.py @@ -162,6 +162,7 @@ def single_generate(self, messages: List[dict], **kwargs) -> str: return output + @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(5)) async def single_generate_async(self, messages: List[dict], **kwargs) -> str: """ Generate a single response using the async LiteLLM completion function. diff --git a/evoagentx/models/openai_model.py b/evoagentx/models/openai_model.py index 75708ed2..30c47938 100644 --- a/evoagentx/models/openai_model.py +++ b/evoagentx/models/openai_model.py @@ -247,6 +247,7 @@ def single_generate(self, messages: List[dict], **kwargs) -> str: def batch_generate(self, batch_messages: List[List[dict]], **kwargs) -> List[str]: return [self.single_generate(messages=one_messages, **kwargs) for one_messages in batch_messages] + @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(5)) async def single_generate_async(self, messages: List[dict], **kwargs) -> str: stream = kwargs.get("stream", self.config.stream) diff --git a/evoagentx/models/openrouter_model.py b/evoagentx/models/openrouter_model.py index e31e140d..c3d13423 100644 --- a/evoagentx/models/openrouter_model.py +++ b/evoagentx/models/openrouter_model.py @@ -331,6 +331,7 @@ def single_generate(self, messages: List[dict], **kwargs) -> str: def batch_generate(self, batch_messages: List[List[dict]], **kwargs) -> List[str]: return [self.single_generate(messages=one_messages, **kwargs) for one_messages in batch_messages] + @retry(wait=wait_random_exponential(min=1, max=60), stop=stop_after_attempt(5)) async def single_generate_async(self, messages: List[dict], **kwargs) -> str: stream = kwargs.get("stream", self.config.stream) output_response = kwargs.get("output_response", self.config.output_response) diff --git a/tests/src/models/test_async_retry.py b/tests/src/models/test_async_retry.py new file mode 100644 index 00000000..593f4363 --- /dev/null +++ b/tests/src/models/test_async_retry.py @@ -0,0 +1,98 @@ +from collections.abc import Callable +from typing import Any +from unittest.mock import AsyncMock + +import pytest +from tenacity import wait_none + +from evoagentx.models.litellm_model import LiteLLM +from evoagentx.models.model_configs import ( + LiteLLMConfig, + OpenAILLMConfig, + OpenRouterConfig, +) +from evoagentx.models.openai_model import OpenAILLM +from evoagentx.models.openrouter_model import OpenRouterLLM +from tests.src.models.mock_response import ( + get_openai_chat_completion, + get_openrouter_chat_completion, +) + + +def _make_openai_llm() -> OpenAILLM: + return OpenAILLM( + config=OpenAILLMConfig( + model="gpt-4o-mini", + openai_key="mock_openai_key", + output_response=False, + ) + ) + + +def _make_litellm() -> LiteLLM: + return LiteLLM( + config=LiteLLMConfig( + model="gpt-4o-mini", + openai_key="mock_openai_key", + output_response=False, + ) + ) + + +def _make_openrouter_llm() -> OpenRouterLLM: + return OpenRouterLLM( + config=OpenRouterConfig( + model="openai/gpt-4o-mini", + openrouter_key="mock_openrouter_key", + output_response=False, + ) + ) + + +@pytest.mark.parametrize( + ("llm_factory", "create_target", "response", "expected"), + [ + ( + _make_openai_llm, + "openai.resources.chat.completions.AsyncCompletions.create", + get_openai_chat_completion(), + "Beijing", + ), + ( + _make_litellm, + "evoagentx.models.litellm_model.acompletion", + get_openai_chat_completion(), + "Beijing", + ), + ( + _make_openrouter_llm, + "openai.resources.chat.completions.AsyncCompletions.create", + get_openrouter_chat_completion(), + "Paris", + ), + ], +) +async def test_single_generate_async_retries_transient_failure( + mocker, + monkeypatch, + llm_factory: Callable[[], Any], + create_target: str, + response: Any, + expected: str, +) -> None: + llm = llm_factory() + create = mocker.patch( + create_target, + new_callable=AsyncMock, + side_effect=[RuntimeError("transient error"), response], + ) + retry_controller = getattr(type(llm).single_generate_async, "retry", None) + assert retry_controller is not None + monkeypatch.setattr(retry_controller, "wait", wait_none()) + + result = await llm.single_generate_async( + messages=[{"role": "user", "content": "hello"}] + ) + + assert result == expected + assert create.await_count == 2