fix(lens): retry contended investigation updates with jittered backoff

This commit is contained in:
Ishaan Jaff 2026-10-03 16:45:46 -07:00
parent ac4d11abbc
commit 9669c2454d
No known key found for this signature in database
3 changed files with 79 additions and 3 deletions

View file

@ -494,6 +494,8 @@ DEFAULT_MAX_LRU_CACHE_SIZE: Final = int(os.getenv("DEFAULT_MAX_LRU_CACHE_SIZE",
_REALTIME_BODY_CACHE_SIZE = 1000 # Keep realtime helper caches bounded; workloads rarely exceed 1k models/intents
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))
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

@ -1,10 +1,13 @@
import asyncio
import json
import random
from collections.abc import AsyncIterator, Awaitable, Callable
from types import MappingProxyType
from typing import Final, Protocol
from pydantic import BaseModel, JsonValue, TypeAdapter
from litellm.constants import LENS_UPDATE_ATTEMPTS, LENS_UPDATE_BACKOFF_SECONDS
from litellm.proxy.db.prisma_client import PrismaWrapper
from litellm.proxy.lens.models import Job, Lens, Scope, Worker
@ -22,8 +25,9 @@ _ROWS: Final = TypeAdapter(tuple[Row, ...])
class LensRepository:
def __init__(self, db: Database) -> None:
def __init__(self, db: Database, sleep: Callable[[float], Awaitable[None]] = asyncio.sleep) -> None:
self.db: Final = db
self.sleep: Final = sleep
async def lenses(self) -> tuple[Lens, ...]:
rows: Final = _ROWS.validate_python(await self.db.query_raw('SELECT data FROM "LiteLLM_Lens" ORDER BY id'))
@ -47,12 +51,18 @@ class LensRepository:
return lens
async def update(
self, lens_id: str, transform: Callable[[Lens], Lens], attempts: int = 8, *, changed_only: bool = False
self,
lens_id: str,
transform: Callable[[Lens], Lens],
attempts: int = LENS_UPDATE_ATTEMPTS,
*,
changed_only: bool = False,
) -> Lens | None:
for _ in range(attempts):
for attempt in range(attempts):
completed, updated = await self._try_update(lens_id, transform, changed_only)
if completed:
return updated
await self.sleep(random.uniform(0, LENS_UPDATE_BACKOFF_SECONDS * min(attempt + 1, 8)))
return None
async def _try_update(

View file

@ -0,0 +1,64 @@
from datetime import datetime, timezone
from typing import Final
import pytest
from litellm.constants import LENS_UPDATE_ATTEMPTS
from litellm.proxy.lens.models import Check, Lens, LensSettings, Scope
from litellm.proxy.lens.repository import LensRepository, Row
NOW: Final = datetime(2026, 1, 15, tzinfo=timezone.utc)
STORED: Final = Lens(
id="lens",
scope=Scope(team_id="alpha"),
settings=LensSettings(name="Swarm", model="cerebras/gpt-oss-120b", checks=(Check(id="c", instruction="Find loops"),)),
created_at=NOW,
next_run_at=NOW,
budget_month=NOW.strftime("%Y-%m"),
)
class ContendedDatabase:
def __init__(self, losses: int) -> None:
self.losses: Final = losses
self.writes = 0 # rebind-ok: counts write attempts made under contention
async def query_raw(self, query: str, *args: object) -> object:
if query.startswith("SELECT data FROM"):
return (Row(data=STORED.model_dump(mode="json")),)
self.writes += 1
return (Row(data=1 if self.writes > self.losses else 0),)
async def execute_raw(self, query: str, *args: object) -> int:
return 0
async def no_wait(_: float) -> None:
return None
def renamed(lens: Lens) -> Lens:
return lens.model_copy(update={"settings": lens.settings.model_copy(update={"name": "Swarm (renamed)"})})
@pytest.mark.asyncio
async def test_update_survives_the_contention_of_a_fast_model_writing_every_review() -> None:
db: Final = ContendedDatabase(losses=12)
updated: Final = await LensRepository(db, sleep=no_wait).update("lens", renamed)
assert updated is not None
assert updated.settings.name == "Swarm (renamed)"
assert db.writes == 13
@pytest.mark.asyncio
async def test_update_backs_off_between_lost_writes_and_gives_up_after_the_limit() -> None:
waits: list[float] = [] # mutable-ok: records each backoff the repository requests
async def record(seconds: float) -> None:
waits.append(seconds)
db: Final = ContendedDatabase(losses=LENS_UPDATE_ATTEMPTS)
assert await LensRepository(db, sleep=record).update("lens", renamed) is None
assert db.writes == LENS_UPDATE_ATTEMPTS
assert len(waits) == LENS_UPDATE_ATTEMPTS
assert all(w >= 0 for w in waits)