diff --git a/libs/core/llmstudio_core/providers/bedrock_converse.py b/libs/core/llmstudio_core/providers/bedrock_converse.py index 3503956..8878e0c 100644 --- a/libs/core/llmstudio_core/providers/bedrock_converse.py +++ b/libs/core/llmstudio_core/providers/bedrock_converse.py @@ -17,6 +17,7 @@ import boto3 import requests +from botocore.config import Config from llmstudio_core.exceptions import ProviderError from llmstudio_core.providers.provider import ChatRequest, ProviderCore, provider from llmstudio_core.utils import OpenAIToolFunction @@ -46,6 +47,7 @@ def __init__(self, config, **kwargs): aws_secret_access_key=self.secret_key if self.secret_key else os.getenv("BEDROCK_SECRET_KEY"), + config=Config(retries={"mode": "adaptive", "max_attempts": 6}), ) @staticmethod diff --git a/libs/core/llmstudio_core/providers/provider.py b/libs/core/llmstudio_core/providers/provider.py index 9646919..0b978de 100644 --- a/libs/core/llmstudio_core/providers/provider.py +++ b/libs/core/llmstudio_core/providers/provider.py @@ -1,3 +1,5 @@ +import asyncio +import random import time import uuid from abc import ABC, abstractmethod @@ -34,6 +36,30 @@ provider_registry = {} +def _is_retryable_error(error: Exception) -> bool: + error_str = str(error).lower() + error_msg = str(error) + retryable_keywords = [ + "throttling", + "ratelimit", + "rate limit", + "toomanyrequests", + "too many requests", + "429", + ] + for keyword in retryable_keywords: + if keyword in error_str: + return True + try: + if hasattr(error, "response"): + error_code = error.response.get("Error", {}).get("Code", "") + if error_code == "ThrottlingException": + return True + except Exception: + pass + return False + + def provider(cls): """Decorator to register a new provider.""" provider_registry[cls._provider_config_name()] = cls @@ -209,7 +235,7 @@ async def achat( self.validate_model(request) - for _ in range(request.retries + 1): + for attempt in range(request.retries + 1): try: start_time = time.time() response = await self.agenerate_client(request) @@ -219,12 +245,11 @@ async def achat( return response_handler else: return await response_handler.__anext__() - # except HTTPException as e: - # if e.status_code == 429: - # continue # Retry on rate limit error - # else: - # raise e # Raise other HTTP exceptions except Exception as e: + if _is_retryable_error(e) and attempt < request.retries: + backoff = min(2 ** attempt, 32) + random.uniform(0, 1) + await asyncio.sleep(backoff) + continue raise ProviderError(str(e)) raise ProviderError("Too many requests") @@ -285,7 +310,7 @@ def chat( self.validate_model(request) - for _ in range(request.retries + 1): + for attempt in range(request.retries + 1): try: start_time = time.time() response = self.generate_client(request) @@ -295,12 +320,11 @@ def chat( return response_handler else: return response_handler.__next__() - # except HTTPExceptio as e: - # if e.status_code == 429: - # continue # Retry on rate limit error - # else: - # raise e # Raise other HTTP exceptions except Exception as e: + if _is_retryable_error(e) and attempt < request.retries: + backoff = min(2 ** attempt, 32) + random.uniform(0, 1) + time.sleep(backoff) + continue raise ProviderError(str(e)) raise ProviderError("Too many requests")