fix(routing): bound failed session snapshot retention

This commit is contained in:
Vaishnavi220506 2026-10-02 10:24:16 +05:30
parent 62b23f95ac
commit 546984a2b2
3 changed files with 124 additions and 4 deletions

View file

@ -68,7 +68,8 @@ Callers may pass header `x-litellm-min-quality-tier: 3` (or metadata key
`SIGNAL_GATE_MIN_MESSAGES` messages.
- **Persistence.** Bandit cells: aggregated deltas, eventually consistent.
Session rows: last-write-wins snapshots. Failed or cancelled flushes retain
unacknowledged rows for the next flush. Retried deltas merge with new feedback;
unacknowledged rows for the next flush within the retention limit below.
Retried deltas merge with new feedback;
a newer queued session snapshot replaces a failed older snapshot. Flushes of
the same queue are serialized within one process, and their return values
count acknowledged writes only
@ -80,8 +81,12 @@ Callers may pass header `x-litellm-min-quality-tier: 3` (or metadata key
to distinguish that outcome from a failed write. The queues are in memory,
so process termination still loses pending updates. Session write ordering
across workers, or server queries continuing after client cancellation, is
not guaranteed. A prolonged database outage retains pending session keys
in memory until a later flush succeeds
not guaranteed. Failed or cancelled session batches retain at most 1,024
snapshots, prioritizing snapshots queued during the flush, then the most
recently queued failed rows. Older snapshots beyond that bound are dropped.
Successful flushes do not impose this retention limit on newly queued rows.
Session keys containing NUL characters are discarded before enqueueing because
PostgreSQL cannot store them in text columns
- **Latency is not in the score.** Quality + cost only. A pathologically slow
model can still be picked.

View file

@ -30,6 +30,7 @@ from litellm.repositories.table_repositories import (
StateKey = tuple[str, str, str] # (router_name, request_type, model_name)
SessionKey = tuple[str, str, str] # (session_id, router_name, model_name)
_MAX_SESSION_RETRY_ENTRIES: Final[int] = 1024
class AdaptiveRouterUpdateQueue:
@ -87,8 +88,12 @@ class AdaptiveRouterUpdateQueue:
SessionState (signals counts + bookkeeping fields). The flusher will
upsert this into LiteLLM_AdaptiveRouterSession.
"""
if any("\0" in part for part in (session_id, router_name, model_name)):
verbose_router_logger.warning("AdaptiveRouterUpdateQueue: session key cannot be stored in PostgreSQL")
return
key: Final[SessionKey] = (session_id, router_name, model_name)
async with self._lock:
self._session_agg.pop(key, None)
self._session_agg[key] = state_dict
self._max_session_size_seen = max(self._max_session_size_seen, len(self._session_agg))
@ -228,8 +233,14 @@ class AdaptiveRouterUpdateQueue:
self._max_state_size_seen = max(self._max_state_size_seen, len(self._state_agg))
async def _restore_session_batch(self, pending: Mapping[SessionKey, Mapping[str, object]]) -> None:
if not pending:
return
async with self._lock:
self._session_agg = {**pending, **self._session_agg}
restored: Final = dict(pending)
for key in self._session_agg:
restored.pop(key, None)
merged: Final = {**restored, **self._session_agg}
self._session_agg = dict(tuple(merged.items())[-_MAX_SESSION_RETRY_ENTRIES:])
self._max_session_size_seen = max(self._max_session_size_seen, len(self._session_agg))
# ---- Observability ---------------------------------------------------

View file

@ -305,3 +305,107 @@ async def test_queue_high_water_mark_includes_restored_rows_and_new_keys(queue,
sizes = await queue.queue_size()
assert sizes[f"{kind}_pending"] == 3
assert sizes[f"max_{kind}_seen"] == 3
@pytest.mark.asyncio
@pytest.mark.parametrize("invalid_field", ["session_id", "router_name", "model_name"])
async def test_session_keys_that_postgres_cannot_store_are_not_queued(queue, mock_prisma, invalid_field):
key = {"session_id": "session", "router_name": "router", "model_name": "model"}
key[invalid_field] += "\0"
await queue.add_session_state(**key, state_dict={"turn_count": 1})
assert (await queue.queue_size())["session_pending"] == 0
assert await queue.flush_session_to_db(mock_prisma) == 0
mock_prisma.db.litellm_adaptiveroutersession.upsert.assert_not_awaited()
@pytest.mark.asyncio
async def test_failed_session_retention_keeps_only_the_newest_rows(queue, mock_prisma, monkeypatch):
monkeypatch.setattr(
"litellm.router_strategy.adaptive_router.update_queue._MAX_SESSION_RETRY_ENTRIES", 3, raising=False
)
table = mock_prisma.db.litellm_adaptiveroutersession
table.upsert.side_effect = RuntimeError("database unavailable")
for index in range(6):
await queue.add_session_state(f"session-{index}", "router", "model", {"turn_count": index})
await queue.add_session_state("session-0", "router", "model", {"turn_count": 9})
assert await queue.flush_session_to_db(mock_prisma) == 0
assert (await queue.queue_size())["session_pending"] == 3
table.reset_mock()
table.upsert.side_effect = None
assert await queue.flush_session_to_db(mock_prisma) == 3
assert [
(call.kwargs["data"]["create"]["session_id"], call.kwargs["data"]["update"]["turn_count"])
for call in table.upsert.await_args_list
] == [("session-0", 9), ("session-4", 4), ("session-5", 5)]
@pytest.mark.asyncio
@pytest.mark.parametrize("cancelled", [False, True])
async def test_session_retention_prioritizes_new_snapshots_during_a_failed_flush(
queue, mock_prisma, monkeypatch, cancelled
):
monkeypatch.setattr(
"litellm.router_strategy.adaptive_router.update_queue._MAX_SESSION_RETRY_ENTRIES", 2, raising=False
)
started = asyncio.Event()
release = asyncio.Event()
async def fail_write(**kwargs):
started.set()
await release.wait()
raise RuntimeError("database unavailable")
table = mock_prisma.db.litellm_adaptiveroutersession
table.upsert.side_effect = fail_write
for session in ("a", "b", "c"):
await queue.add_session_state(session, "router", "model", {"turn_count": 1})
task = asyncio.create_task(queue.flush_session_to_db(mock_prisma))
await started.wait()
await queue.add_session_state("a", "router", "model", {"turn_count": 2})
await queue.add_session_state("d", "router", "model", {"turn_count": 3})
if cancelled:
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
else:
release.set()
assert await task == 0
assert (await queue.queue_size())["session_pending"] == 2
table.reset_mock()
table.upsert.side_effect = None
assert await queue.flush_session_to_db(mock_prisma) == 2
assert [
(call.kwargs["data"]["create"]["session_id"], call.kwargs["data"]["update"]["turn_count"])
for call in table.upsert.await_args_list
] == [("a", 2), ("d", 3)]
@pytest.mark.asyncio
async def test_successful_flush_does_not_cap_newly_queued_sessions(queue, mock_prisma, monkeypatch):
monkeypatch.setattr(
"litellm.router_strategy.adaptive_router.update_queue._MAX_SESSION_RETRY_ENTRIES", 1, raising=False
)
started = asyncio.Event()
release = asyncio.Event()
async def write(**kwargs):
started.set()
await release.wait()
table = mock_prisma.db.litellm_adaptiveroutersession
table.upsert.side_effect = write
await queue.add_session_state("old", "router", "model", {"turn_count": 1})
task = asyncio.create_task(queue.flush_session_to_db(mock_prisma))
await started.wait()
for session in ("a", "b"):
await queue.add_session_state(session, "router", "model", {"turn_count": 1})
release.set()
assert await task == 1
assert (await queue.queue_size())["session_pending"] == 2
table.upsert.side_effect = None
assert await queue.flush_session_to_db(mock_prisma) == 2