mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
0721cffab2
commit
4ed0267b5a
3 changed files with 163 additions and 6 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")),)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue