mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(lens): wait out provider rate limits and retry model calls four times
This commit is contained in:
parent
68cc5484d1
commit
5b89b1384a
3 changed files with 57 additions and 5 deletions
|
|
@ -496,6 +496,8 @@ INITIAL_RETRY_DELAY: Final = float(os.getenv("INITIAL_RETRY_DELAY", 0.5))
|
|||
MAX_RETRY_DELAY: Final = float(os.getenv("MAX_RETRY_DELAY", 8.0))
|
||||
LENS_UPDATE_ATTEMPTS: Final = get_env_int("LENS_UPDATE_ATTEMPTS", 40)
|
||||
LENS_UPDATE_BACKOFF_SECONDS: Final = float(os.getenv("LENS_UPDATE_BACKOFF_SECONDS", 0.02))
|
||||
LENS_MODEL_RETRIES: Final = get_env_int("LENS_MODEL_RETRIES", 4)
|
||||
LENS_MODEL_RETRY_MAX_SECONDS: Final = float(os.getenv("LENS_MODEL_RETRY_MAX_SECONDS", 60))
|
||||
JITTER: Final = float(os.getenv("JITTER", 0.75))
|
||||
DEFAULT_IN_MEMORY_TTL = int(os.getenv("DEFAULT_IN_MEMORY_TTL", 5)) # default time to live for the in-memory cache
|
||||
DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE: Final = int(
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ from typing import Final
|
|||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, ValidationError
|
||||
|
||||
from litellm.constants import LENS_MODEL_RETRIES, LENS_MODEL_RETRY_MAX_SECONDS
|
||||
|
||||
from .analysis import AnalysisResponseError, analyze_sample, validation_details
|
||||
from .models import Claim, Coverage, ExecutionContent, ModelRequest, ModelResult, Progress, Result, Review, Sample
|
||||
|
||||
|
|
@ -36,6 +38,17 @@ class ModelErrorEnvelope(BaseModel):
|
|||
detail: PublicModelError
|
||||
|
||||
|
||||
def retry_delay(error: httpx.TransportError | httpx.HTTPStatusError, attempt: int) -> float:
|
||||
backoff: Final = float(min(2**attempt, LENS_MODEL_RETRY_MAX_SECONDS))
|
||||
if not isinstance(error, httpx.HTTPStatusError):
|
||||
return backoff
|
||||
requested: Final = error.response.headers.get("retry-after", "")
|
||||
try:
|
||||
return min(max(float(requested), backoff), LENS_MODEL_RETRY_MAX_SECONDS)
|
||||
except ValueError:
|
||||
return backoff
|
||||
|
||||
|
||||
def failure_message(error: Exception) -> str:
|
||||
if isinstance(error, AnalysisResponseError):
|
||||
return str(error)
|
||||
|
|
@ -115,9 +128,9 @@ class LensWorker:
|
|||
503,
|
||||
504,
|
||||
)
|
||||
if not retryable or attempt >= 2:
|
||||
if not retryable or attempt >= LENS_MODEL_RETRIES:
|
||||
raise
|
||||
await self.sleep(2**attempt)
|
||||
await self.sleep(retry_delay(exc, attempt))
|
||||
return await self.model_request(path, body, attempt + 1)
|
||||
|
||||
async def run_once(self) -> bool:
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import httpx
|
|||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.constants import LENS_MODEL_RETRIES, LENS_MODEL_RETRY_MAX_SECONDS
|
||||
from litellm.proxy.lens.models import (
|
||||
Claim,
|
||||
Execution,
|
||||
|
|
@ -18,7 +19,7 @@ from litellm.proxy.lens.models import (
|
|||
TracePart,
|
||||
)
|
||||
from litellm.proxy.lens.state import queue_job
|
||||
from litellm.proxy.lens.worker import LensWorker, failure_message
|
||||
from litellm.proxy.lens.worker import LensWorker, failure_message, retry_delay
|
||||
from tests.unit.proxy.lens.test_state import NOW, lens
|
||||
|
||||
|
||||
|
|
@ -70,8 +71,44 @@ async def test_transient_retries_are_bounded() -> None:
|
|||
await LensWorker(client, sleep=sleep).model_request(
|
||||
"/model", ModelRequest(purpose="extract", prompt="review")
|
||||
)
|
||||
assert attempts.qsize() == 3
|
||||
assert tuple(delays.get_nowait() for _ in range(delays.qsize())) == (1, 2)
|
||||
assert attempts.qsize() == LENS_MODEL_RETRIES + 1
|
||||
assert tuple(delays.get_nowait() for _ in range(delays.qsize())) == tuple(
|
||||
float(min(2**n, LENS_MODEL_RETRY_MAX_SECONDS)) for n in range(LENS_MODEL_RETRIES)
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rate_limited_model_waits_as_long_as_the_provider_asks_then_completes() -> None:
|
||||
attempts: Final = SimpleQueue[str]()
|
||||
delays: Final = SimpleQueue[float]()
|
||||
expected: Final = ModelResult(content='{"observations":[]}', cost=0.01)
|
||||
|
||||
def handle(request: httpx.Request) -> httpx.Response:
|
||||
attempts.put(request.url.path)
|
||||
if attempts.qsize() <= 3:
|
||||
return httpx.Response(429, headers={"retry-after": "30"})
|
||||
return httpx.Response(200, json=expected.model_dump())
|
||||
|
||||
async def sleep(delay: float) -> None:
|
||||
delays.put(delay)
|
||||
|
||||
async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client:
|
||||
result: Final = await LensWorker(client, sleep=sleep).model_request(
|
||||
"/model", ModelRequest(purpose="extract", prompt="review")
|
||||
)
|
||||
assert result == expected
|
||||
assert tuple(delays.get_nowait() for _ in range(delays.qsize())) == (30, 30, 30)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("retry_after", "attempt", "expected"),
|
||||
(("", 1, 2), ("5", 0, 5), ("1", 3, 8), ("9999", 0, LENS_MODEL_RETRY_MAX_SECONDS), ("soon", 2, 4)),
|
||||
)
|
||||
def test_retry_delay_prefers_the_providers_wait_within_bounds(retry_after: str, attempt: int, expected: float) -> None:
|
||||
request: Final = httpx.Request("POST", "https://proxy.test/model")
|
||||
headers: Final = {"retry-after": retry_after} if retry_after else {}
|
||||
error: Final = httpx.HTTPStatusError("limited", request=request, response=httpx.Response(429, headers=headers))
|
||||
assert retry_delay(error, attempt) == expected
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue