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:
Tin Chi Lo 2026-08-04 12:45:52 -07:00
parent 7d902b387c
commit 547a0c6d1e
2 changed files with 104 additions and 9 deletions

View file

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

View file

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