Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions libs/core/llmstudio_core/providers/bedrock_converse.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
48 changes: 36 additions & 12 deletions libs/core/llmstudio_core/providers/provider.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
import asyncio
import random
import time
import uuid
from abc import ABC, abstractmethod
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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")

Expand Down Expand Up @@ -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)
Expand All @@ -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")

Expand Down