diff --git a/litellm/constants.py b/litellm/constants.py index 49514fc4d0e..a59e1ca6fcd 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -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( diff --git a/litellm/proxy/lens/repository.py b/litellm/proxy/lens/repository.py index 6e1e2da112a..8875cf3512b 100644 --- a/litellm/proxy/lens/repository.py +++ b/litellm/proxy/lens/repository.py @@ -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( diff --git a/tests/unit/proxy/lens/test_repository.py b/tests/unit/proxy/lens/test_repository.py new file mode 100644 index 00000000000..a5cb60eaa3e --- /dev/null +++ b/tests/unit/proxy/lens/test_repository.py @@ -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)