fix(lens): wait out provider rate limits and retry model calls four times

This commit is contained in:
Ishaan Jaff 2026-10-03 17:11:24 -07:00
parent 68cc5484d1
commit 5b89b1384a
No known key found for this signature in database
3 changed files with 57 additions and 5 deletions

View file

@ -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(

View file

@ -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:

View file

@ -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