mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(lens): retry contended investigation updates with jittered backoff
This commit is contained in:
parent
ac4d11abbc
commit
9669c2454d
3 changed files with 79 additions and 3 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
64
tests/unit/proxy/lens/test_repository.py
Normal file
64
tests/unit/proxy/lens/test_repository.py
Normal 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)
|
||||
Loading…
Add table
Reference in a new issue