mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(routing): bound failed session snapshot retention
This commit is contained in:
parent
62b23f95ac
commit
546984a2b2
3 changed files with 124 additions and 4 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 ---------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue