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
This commit is contained in:
moe-berri 2026-10-08 14:08:55 -07:00 • committed by GitHub
parent 0721cffab2
commit 4ed0267b5a
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 163 additions and 6 deletions

View file

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

View file

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

View file

@ -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")),)