From 4ed0267b5a081ac4884431fd7f5ef41fce20a4d1 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Thu, 8 Oct 2026 14:08:55 -0700 Subject: [PATCH] fix(lens): prevent progress updates from starving analysis budget reservations (#45432) * fix(lens): serialize analysis budget reservations with progress updates * test(lens): verify budget reservations block competing progress writes --- litellm/proxy/lens/inference.py | 4 +- .../database/test_lens_repository.py | 156 +++++++++++++++++- tests/unit/proxy/lens/test_endpoints.py | 9 +- 3 files changed, 163 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/lens/inference.py b/litellm/proxy/lens/inference.py index 599c88cf4d2..fbb661c027c 100644 --- a/litellm/proxy/lens/inference.py +++ b/litellm/proxy/lens/inference.py @@ -325,7 +325,7 @@ def renew_reservation(lens: Lens, reservation_id: str, now: datetime) -> Lens: async def wait_for_reservation( repo: LensRepository, lens_id: str, reservation_id: str, reserve: Callable[[Lens], Lens] ) -> None: - while (reserved := await repo.update(lens_id, reserve)) is not None: + while (reserved := await repo.update_locked(lens_id, reserve)) is not None: if any(held.id == reservation_id for held in reserved.reservations): return await asyncio.sleep(0.25) @@ -341,7 +341,7 @@ async def renew_budget_reservation( try: async with timeout(BUDGET_RENEW_INTERVAL): if ( - await repo.update( + await repo.update_locked( lens_id, lambda e: renew_reservation(e, reservation_id, datetime.now(timezone.utc)) ) is None diff --git a/tests/integration/database/test_lens_repository.py b/tests/integration/database/test_lens_repository.py index ff46944b911..6349da7ab2a 100644 --- a/tests/integration/database/test_lens_repository.py +++ b/tests/integration/database/test_lens_repository.py @@ -1,6 +1,7 @@ import asyncio import os -from collections.abc import AsyncIterator +from collections.abc import AsyncGenerator, AsyncIterator +from contextlib import asynccontextmanager from datetime import datetime, timedelta, timezone from pathlib import Path from types import SimpleNamespace @@ -15,6 +16,7 @@ from fastapi import HTTPException from prisma import Prisma from psycopg import sql from pydantic import TypeAdapter +from typing_extensions import LiteralString from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.db.prisma_client import PrismaWrapper @@ -35,7 +37,7 @@ from litellm.proxy.lens.models import ( TraceIdentity, Worker, ) -from litellm.proxy.lens.repository import LensRepository, WriterDatabase +from litellm.proxy.lens.repository import Database, LensRepository, Row, WriterDatabase from litellm.proxy.lens.state import cancel_job, claim_job, current_job, due_at, end_job, queue_job, replace_job @@ -770,6 +772,156 @@ async def test_delayed_progress_cannot_replace_a_newer_checkpoint(lens_db: Prism await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', claimed.id) +class ProgressInterleavingDatabase: + def __init__( + self, + database: Database, + read: asyncio.Future[int], + resume: asyncio.Event, + committed: asyncio.Event, + ) -> None: + self.database: Final = database + self.read: Final = read + self.resume: Final = resume + self.committed: Final = committed + + async def query_raw(self, query: LiteralString, *args: object) -> object: + rows: Final = await self.database.query_raw(query, *args) + if query == 'SELECT data FROM "LiteLLM_Lens" WHERE id=$1' and not self.read.done(): + backend: Final = TypeAdapter(tuple[Row, ...]).validate_python( + await self.database.query_raw("SELECT to_jsonb(pg_backend_pid()) AS data") + ) + self.read.set_result(TypeAdapter(int).validate_python(backend[0].data)) + await self.resume.wait() + return rows + + async def execute_raw(self, query: LiteralString, *args: object) -> int: + return await self.database.execute_raw(query, *args) + + @asynccontextmanager + async def transaction(self) -> AsyncGenerator[Database]: + async with self.database.transaction() as database: + yield ProgressInterleavingDatabase(database, self.read, self.resume, self.committed) + self.committed.set() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("renewing", (False, True)) +async def test_budget_reservation_survives_competing_progress(lens_db: Prisma, renewing: bool) -> None: + from litellm.proxy.lens.inference import ( + BUDGET_LEASE, + renew_budget_reservation, + reserve_attempt, + wait_for_reservation, + ) + from litellm.proxy.lens.models import BudgetReservation + from tests.unit.proxy.lens.test_state import lens, worker + + now: Final = datetime.now(timezone.utc) + claimed: Final = claim_job(queue_job(lens(), now, uuid4().hex), worker(), now).model_copy( + update={"id": uuid4().hex, "budget_month": now.strftime("%Y-%m")} + ) + job: Final = claimed.jobs[0] + hold: Final = BudgetReservation( + id=uuid4().hex, job_id=job.id, amount=1, month=claimed.budget_month, expires_at=now + BUDGET_LEASE + ) + database: Final = WriterDatabase(PrismaWrapper(lens_db)) + repo: Final = LensRepository(database) + read: Final[asyncio.Future[int]] = asyncio.get_running_loop().create_future() + resume: Final = asyncio.Event() + committed: Final = asyncio.Event() + competing: Final = ProgressInterleavingDatabase(database, read, resume, committed) + await repo.create(claimed.model_copy(update={"reservations": (hold,) if renewing else ()})) + admitted: Final = asyncio.Event() + admitted.set() + operation: Final = asyncio.create_task( + renew_budget_reservation(LensRepository(competing), claimed.id, hold.id, admitted) + if renewing + else wait_for_reservation( + LensRepository(competing), + claimed.id, + hold.id, + lambda current: reserve_attempt(current, job, worker().id, hold, now), + ) + ) + + async def write_progress() -> Lens | None: + await read + return await repo.progress(claimed.id, job, Progress(stage="Reviewing traces concurrently")) + + progress: Final = asyncio.create_task(write_progress()) + try: + async with asyncio.timeout(45): + blocker: Final = await read + async with asyncio.timeout(5): + while not await lens_db.query_raw( + "SELECT pid FROM pg_stat_activity WHERE $1::int=ANY(pg_blocking_pids(pid))", blocker + ): + assert not progress.done(), "Progress committed before the reservation released its lock" + await asyncio.sleep(0.01) + assert not progress.done() + assert not committed.is_set() + resume.set() + await committed.wait() + updated: Final = await progress + assert updated is not None + assert updated.jobs[0].stage == "Reviewing traces concurrently" + assert tuple(reservation.id for reservation in updated.reservations) == (hold.id,) + if not renewing: + await operation + stored: Final = await repo.get(claimed.id) + assert stored is not None + assert tuple(reservation.id for reservation in stored.reservations) == (hold.id,) + assert stored.spent == claimed.spent + assert stored.jobs[0].cost == 0 + assert stored.jobs[0].stage == "Reviewing traces concurrently" + if renewing: + assert stored.reservations[0].expires_at is not None + assert stored.reservations[0].expires_at > now + BUDGET_LEASE + finally: + operation.cancel() + progress.cancel() + await asyncio.gather(operation, progress, return_exceptions=True) + await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', claimed.id) + + +@pytest.mark.asyncio +async def test_parallel_reservations_release_the_lock_while_waiting_for_budget(lens_db: Prisma) -> None: + from litellm.proxy.lens.inference import reserve_amount, settle_amount, wait_for_reservation + from litellm.proxy.lens.models import BudgetReservation + from tests.unit.proxy.lens.test_state import lens + + now: Final = datetime.now(timezone.utc) + original: Final = lens() + queued: Final = queue_job(original, now, uuid4().hex).model_copy( + update={"id": uuid4().hex, "settings": original.settings.model_copy(update={"monthly_budget": 2})} + ) + repo: Final = LensRepository(WriterDatabase(PrismaWrapper(lens_db))) + holds: Final = tuple( + BudgetReservation(id=uuid4().hex, job_id=queued.jobs[0].id, amount=1, month=queued.budget_month) + for _ in range(8) + ) + await repo.create(queued) + try: + + async def analyze(hold: BudgetReservation) -> None: + await wait_for_reservation(repo, queued.id, hold.id, lambda current: reserve_amount(current, hold)) + stored: Final = await repo.get(queued.id) + assert stored is not None and hold in stored.reservations + assert stored.spent + sum(reservation.amount for reservation in stored.reservations) <= 2 + assert await repo.update_locked(queued.id, lambda current: settle_amount(current, hold.id, 0.125, None)) + + async with asyncio.timeout(15): + await asyncio.gather(*(analyze(hold) for hold in holds)) + stored: Final = await repo.get(queued.id) + assert stored is not None + assert stored.spent == 1 + assert stored.jobs[0].cost == 1 + assert stored.reservations == () + finally: + await lens_db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', queued.id) + + @pytest.mark.asyncio async def test_locked_settlement_charges_every_concurrent_call_exactly_once(lens_db: Prisma) -> None: from litellm.proxy.lens.inference import settle_amount diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index fee53320802..e9ad63a701e 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -1,5 +1,6 @@ import asyncio -from collections.abc import Callable, Mapping +from collections.abc import AsyncGenerator, Callable, Mapping +from contextlib import asynccontextmanager from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final @@ -49,7 +50,7 @@ from litellm.proxy.lens.models import ( TraceIdentity, Worker, ) -from litellm.proxy.lens.repository import DueLens, Row +from litellm.proxy.lens.repository import Database, DueLens, Row from litellm.proxy.lens.signals import SignalConfig, StoredTraceSignal from litellm.proxy.lens.state import claim_job, queue_job, replace_job from litellm.rust_bridge.trace.generated.models import ExecutionRow, LensSampleParams @@ -69,6 +70,10 @@ class ResultDatabase: self.stored = stored self.completed: tuple[ReviewVersion, ...] = () + @asynccontextmanager + async def transaction(self) -> AsyncGenerator[Database]: + yield self + async def query_raw(self, query: str, *args: object) -> tuple[Row, ...]: if query.startswith("SELECT data FROM"): return (Row(data=self.stored.model_dump(mode="json")),)