mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(spend): read session progress from the fold and bound every path into staging
Two defects with one shape: a fact the folded state already owned was re-derived from the staged turns, and an invariant of the staging buffer was enforced somewhere other than where the buffer grows. The row's last activity came from the maximum start time of the staged turns. fold_turn deliberately refuses to advance a session's state for a turn that started before its last, so the two disagree exactly when a late turn is in the batch, and the staged turns are the wrong one of the pair: the column moved backwards, which shortens the session on the dashboard, moves its retention cutoff earlier, and changes what the next turn is classified against after a reload. It now comes off the folded state, which is the only thing that decides whether a session advanced. The staging ceiling was checked in record_turn while the buffer grew in _stage, which the replay path also calls. Turns keep arriving between the drain that zeroes the count and the replay that puts a failed batch back, so a database fault that persisted grew the staging by an interval each flush, without limit. The ceiling moved into _stage, so every path into the buffer is bounded by it. Once full the replay is refused, which under a sustained fault is an undercounted dashboard rather than the out-of-memory kill the cap exists for. Both regressions are pinned. The staging one needs the arrival to land inside the flush to reproduce, so the test drives a turn in from the read the flush performs; a version recording turns after the flush returned passed against the bug and was worthless
This commit is contained in:
parent
7d902b387c
commit
547a0c6d1e
2 changed files with 104 additions and 9 deletions
|
|
@ -110,11 +110,20 @@ def _increment(value: float) -> Mapping[str, float]:
|
|||
|
||||
|
||||
def _upsert_data(key: SessionKey, pending: _Pending, delta: TurnDelta) -> Mapping[str, Mapping[str, object]]:
|
||||
"""One session's write: create its row, or add this interval onto the row already there."""
|
||||
"""One session's write: create its row, or add this interval onto the row already there.
|
||||
|
||||
Every field the row carries about session progress is read off the folded
|
||||
state rather than re-derived from the staged turns. ``fold_turn`` refuses to
|
||||
advance the state for a turn that started before the session's last, so the
|
||||
two disagree exactly when a late turn is in the batch, and the staged turns
|
||||
are the wrong one of the pair: taking their maximum would move the row's last
|
||||
activity backwards, shortening the session, moving its retention cutoff, and
|
||||
changing what the next turn is classified against after a reload.
|
||||
"""
|
||||
counters = counters_of(delta)
|
||||
increments = MappingProxyType({name: _increment(value) for name, value in counters.items()})
|
||||
shared = { # mutable-ok: prisma's write API takes dict payloads
|
||||
"last_turn_at": _epoch_to_datetime(max(turn.started_at for turn in pending.turns)),
|
||||
"last_turn_at": _epoch_to_datetime(delta.state.last_turn_at),
|
||||
"last_model": delta.state.last_model,
|
||||
"model_state": state_column(delta.state),
|
||||
"baseline_model": pending.baseline_model,
|
||||
|
|
@ -147,7 +156,8 @@ class AutoRouterSessionQueue:
|
|||
``session_id`` is caller-controlled, so the staging is capped on turns
|
||||
held rather than on sessions seen, which is the quantity that actually
|
||||
bounds the memory. Past the cap a turn is dropped and logged, because
|
||||
benchmark rows are not worth an out-of-memory kill.
|
||||
benchmark rows are not worth an out-of-memory kill. ``_stage`` owns that
|
||||
cap, so every path into the staging is bounded by it and not just this one.
|
||||
|
||||
That cap is counted here rather than inherited from ``BaseUpdateQueue``,
|
||||
which ``SpendUpdateQueue`` and ``DailySpendUpdateQueue`` both use, on
|
||||
|
|
@ -157,25 +167,41 @@ class AutoRouterSessionQueue:
|
|||
number on one tab; a stalled spend write is money that went unbilled.
|
||||
"""
|
||||
async with self._lock:
|
||||
if self._staged_turns >= self._max_staged_turns:
|
||||
_warn_staging_full(self._max_staged_turns)
|
||||
return
|
||||
self._stage(key, router_kind, baseline_model, [turn]) # mutable-ok: becomes this session's staging buffer
|
||||
|
||||
def _stage(self, key: SessionKey, router_kind: str, baseline_model: str | None, turns: list[TurnFacts]) -> None:
|
||||
"""Add turns to whatever this session already has staged; callers hold the lock.
|
||||
|
||||
The ceiling is applied here because this is the only point the staging
|
||||
grows. Enforcing it in one caller instead would leave every other caller
|
||||
unbounded, and the replay path is exactly such a caller: turns keep
|
||||
arriving between the drain that zeroes the count and the replay that puts
|
||||
a failed batch back, so a database fault that persists would grow the
|
||||
staging by an interval each time, which is the out-of-memory kill the cap
|
||||
exists to prevent.
|
||||
|
||||
Once the staging is full the replay is what gets refused, since it stages
|
||||
last. Under a sustained fault something has to be dropped, and losing the
|
||||
older interval's counters is no worse than refusing the traffic still
|
||||
arriving; both are an undercounted dashboard rather than a lost spend row.
|
||||
|
||||
Arriving turns and a replayed batch stage identically, because the fold
|
||||
sorts by start time and so does not care which of the two came first.
|
||||
"""
|
||||
self._staged_turns += len(turns)
|
||||
room = self._max_staged_turns - self._staged_turns
|
||||
admitted = turns if len(turns) <= room else turns[: max(room, 0)]
|
||||
if len(admitted) < len(turns):
|
||||
_warn_staging_full(self._max_staged_turns)
|
||||
if not admitted:
|
||||
return
|
||||
self._staged_turns += len(admitted)
|
||||
current = self._pending.get(key)
|
||||
if current is None:
|
||||
self._pending[key] = _Pending(router_kind=router_kind, baseline_model=baseline_model, turns=turns)
|
||||
self._pending[key] = _Pending(router_kind=router_kind, baseline_model=baseline_model, turns=admitted)
|
||||
return
|
||||
current.router_kind = router_kind
|
||||
current.baseline_model = baseline_model or current.baseline_model
|
||||
current.turns.extend(turns)
|
||||
current.turns.extend(admitted)
|
||||
|
||||
async def flush(self, prisma_client: "PrismaClient") -> int:
|
||||
"""Fold and write every staged session. Returns rows written.
|
||||
|
|
|
|||
|
|
@ -407,6 +407,18 @@ class _RecordingTable:
|
|||
self.upserts.extend(statements)
|
||||
|
||||
|
||||
class _RefillingTable(_RecordingTable):
|
||||
"""Lets turns arrive in the middle of a flush, between the drain and the replay."""
|
||||
|
||||
def __init__(self, arrive, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self._arrive = arrive
|
||||
|
||||
async def find_many(self, where):
|
||||
await self._arrive()
|
||||
return await super().find_many(where)
|
||||
|
||||
|
||||
class _RecordingBatchActions:
|
||||
"""Prisma's batcher: statements queue up and land only when the batch commits."""
|
||||
|
||||
|
|
@ -579,6 +591,38 @@ class TestFlushDurability:
|
|||
assert await queue.flush(prisma) == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestTheRowRecordsTheFoldsAnswer:
|
||||
"""Session progress on the row comes from the folded state, not the staged turns.
|
||||
|
||||
They agree until a turn arrives late, at which point `fold_turn` refuses to
|
||||
advance the state and the staged turns still carry the late timestamp.
|
||||
"""
|
||||
|
||||
async def test_a_late_turn_does_not_rewind_the_rows_last_activity(self):
|
||||
from litellm.proxy.spend_tracking.auto_router_session_queue import AutoRouterSessionQueue
|
||||
|
||||
table = _RecordingTable(rows=(_stored_session(),))
|
||||
queue = AutoRouterSessionQueue()
|
||||
await queue.record_turn(("s1", "g"), "complexity", None, _turn_at(30))
|
||||
|
||||
assert await queue.flush(_RecordingPrisma(table)) == 1
|
||||
update = table.upserts[0][1]["update"]
|
||||
assert update["last_turn_at"] == datetime.fromtimestamp(60.0, tz=timezone.utc)
|
||||
|
||||
async def test_a_turn_that_did_advance_the_session_moves_last_activity_forward(self):
|
||||
"""The guard must not freeze the column for ordinary in-order traffic."""
|
||||
from litellm.proxy.spend_tracking.auto_router_session_queue import AutoRouterSessionQueue
|
||||
|
||||
table = _RecordingTable(rows=(_stored_session(),))
|
||||
queue = AutoRouterSessionQueue()
|
||||
await queue.record_turn(("s1", "g"), "complexity", None, _turn_at(300))
|
||||
|
||||
assert await queue.flush(_RecordingPrisma(table)) == 1
|
||||
update = table.upserts[0][1]["update"]
|
||||
assert update["last_turn_at"] == datetime.fromtimestamp(300.0, tz=timezone.utc)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
class TestStagingCostsTheSamePerTurn:
|
||||
"""Staging must not re-copy the turns already staged for that session.
|
||||
|
|
@ -653,6 +697,31 @@ class TestStagingIsBounded:
|
|||
assert await queue.flush(prisma) == 1
|
||||
assert table.upserts[1][1]["update"]["turns"] == {"increment": 1}
|
||||
|
||||
async def test_a_replayed_batch_is_bounded_by_the_same_ceiling(self):
|
||||
"""A database that stays down must not grow the staging one failed flush at a time.
|
||||
|
||||
The window that matters is inside the flush: the drain has already zeroed
|
||||
the counter, turns keep arriving against it, and only then is the failed
|
||||
batch put back. A ceiling checked by the arriving path alone lets the
|
||||
replay land on top of a staging that is already full.
|
||||
"""
|
||||
from litellm.proxy.spend_tracking.auto_router_session_queue import AutoRouterSessionQueue
|
||||
|
||||
queue = AutoRouterSessionQueue(max_staged_turns=2)
|
||||
|
||||
async def arrive_mid_flush():
|
||||
for at in (200, 260):
|
||||
await queue.record_turn(("s2", "g"), "complexity", None, _turn_at(at))
|
||||
|
||||
table = _RefillingTable(arrive_mid_flush, fail_write=True)
|
||||
for at in (0, 60):
|
||||
await queue.record_turn(("s1", "g"), "complexity", None, _turn_at(at))
|
||||
|
||||
await queue.flush(_RecordingPrisma(table))
|
||||
|
||||
held = sum(len(pending.turns) for pending in queue._pending.values())
|
||||
assert held <= 2, f"staging grew past its ceiling to {held} turns"
|
||||
|
||||
|
||||
def test_chunking_caps_what_one_statement_carries_and_keeps_it_in_key_order():
|
||||
"""Key order is the lock order two pods draining the same sessions have to agree on."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue