This commit is contained in:
Devulapalli Naga Sri Vaishnavi 2026-10-04 05:30:09 -04:00 • committed by GitHub
commit 0cbb1d491d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 430 additions and 76 deletions

View file

@ -67,10 +67,27 @@ Callers may pass header `x-litellm-min-quality-tier: 3` (or metadata key
(uptime, not quality). Skipped if conversation has fewer than
`SIGNAL_GATE_MIN_MESSAGES` messages.
- **Persistence.** Bandit cells: aggregated deltas, eventually consistent.
Session rows: last-write-wins snapshots.
Session rows: last-write-wins snapshots. Failed or cancelled flushes retain
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
## Known v0 limitations
- **Retries are at-least-once.** If a state increment commits but its response
is lost, retrying can count that feedback twice. There is no durable receipt
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. 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.
- **Hard sample cap at 200.** Once `α + β > 200`, deltas are silently dropped.

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:
@ -42,6 +43,8 @@ class AdaptiveRouterUpdateQueue:
self._state_agg: dict[StateKey, dict[str, float]] = {}
self._session_agg: dict[SessionKey, Mapping[str, object]] = {}
self._lock = asyncio.Lock()
self._state_flush_lock = asyncio.Lock()
self._session_flush_lock = asyncio.Lock()
self._max_state_size_seen = 0
self._max_session_size_seen = 0
@ -85,14 +88,22 @@ 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))
# ---- Flushers (called by background task) ----------------------------
async def flush_state_to_db(self, prisma_client: object) -> int:
async with self._state_flush_lock:
return await self._flush_state_to_db(prisma_client)
async def _flush_state_to_db(self, prisma_client: object) -> int:
"""
Drain state aggregator and apply to LiteLLM_AdaptiveRouterState.
Returns number of cells flushed.
@ -106,49 +117,57 @@ class AdaptiveRouterUpdateQueue:
# Sort keys to give deterministic write order across writers and
# reduce the chance of cross-row deadlocks when other workers race us.
for key in sorted(batch.keys()):
router, rt, model = key
payload = batch[key]
try:
# Atomic increment: push the delta directly into the DB so
# concurrent flushers from multiple pods don't overwrite each
# other. The upsert creates the row with the delta as the
# initial value on first write, then increments on subsequent
# writes — no read-modify-write race.
await AdaptiveRouterStateRepository(prisma_client).table.upsert(
where={
"router_name_request_type_model_name": {
"router_name": router,
"request_type": rt,
"model_name": model,
}
},
data={
"create": {
"router_name": router,
"request_type": rt,
"model_name": model,
"alpha": payload["delta_alpha"],
"beta": payload["delta_beta"],
"total_samples": int(payload["samples_added"]),
pending: Final = dict(batch)
try:
for key, payload in sorted(batch.items()):
router, rt, model = key
try:
# Atomic increment: push the delta directly into the DB so
# concurrent flushers from multiple pods don't overwrite each
# other. The upsert creates the row with the delta as the
# initial value on first write, then increments on subsequent
# writes — no read-modify-write race.
await AdaptiveRouterStateRepository(prisma_client).table.upsert(
where={
"router_name_request_type_model_name": {
"router_name": router,
"request_type": rt,
"model_name": model,
}
},
"update": {
"alpha": {"increment": payload["delta_alpha"]},
"beta": {"increment": payload["delta_beta"]},
"total_samples": {"increment": int(payload["samples_added"])},
data={
"create": {
"router_name": router,
"request_type": rt,
"model_name": model,
"alpha": payload["delta_alpha"],
"beta": payload["delta_beta"],
"total_samples": int(payload["samples_added"]),
},
"update": {
"alpha": {"increment": payload["delta_alpha"]},
"beta": {"increment": payload["delta_beta"]},
"total_samples": {"increment": int(payload["samples_added"])},
},
},
},
)
except Exception as e:
verbose_router_logger.exception(
"AdaptiveRouterUpdateQueue: failed to flush state for %s: %s",
key,
e,
)
)
pending.pop(key)
except Exception as e:
verbose_router_logger.exception(
"AdaptiveRouterUpdateQueue: failed to flush state for %s: %s",
key,
e,
)
finally:
await self._restore_state_batch(pending)
return len(batch)
return len(batch) - len(pending)
async def flush_session_to_db(self, prisma_client: object) -> int:
async with self._session_flush_lock:
return await self._flush_session_to_db(prisma_client)
async def _flush_session_to_db(self, prisma_client: object) -> int:
"""
Drain session aggregator and upsert into LiteLLM_AdaptiveRouterSession.
Returns number of session rows flushed.
@ -160,45 +179,69 @@ class AdaptiveRouterUpdateQueue:
if not batch:
return 0
for key in sorted(batch.keys()):
session_id, router, model = key
payload = batch[key]
try:
# NOTE: Prisma client lower-cases model names, so
# `LiteLLM_AdaptiveRouterSession` -> `litellm_adaptiveroutersession`
# (single 's', not 'litellm_adaptiverouterssession').
# Strip PK fields from the update payload — Prisma rejects
# writes to fields that are part of the @@id. asdict(state)
# always carries them, so build a separate update dict.
update_payload = {
k: v for k, v in payload.items() if k not in ("session_id", "router_name", "model_name")
}
await AdaptiveRouterSessionRepository(prisma_client).table.upsert(
where={
"session_id_router_name_model_name": {
"session_id": session_id,
"router_name": router,
"model_name": model,
}
},
data={
"create": {
"session_id": session_id,
"router_name": router,
"model_name": model,
**update_payload,
pending: Final = dict(batch)
try:
for key, payload in sorted(batch.items()):
session_id, router, model = key
try:
# NOTE: Prisma client lower-cases model names, so
# `LiteLLM_AdaptiveRouterSession` -> `litellm_adaptiveroutersession`
# (single 's', not 'litellm_adaptiverouterssession').
# Strip PK fields from the update payload — Prisma rejects
# writes to fields that are part of the @@id. asdict(state)
# always carries them, so build a separate update dict.
update_payload = {
k: v for k, v in payload.items() if k not in ("session_id", "router_name", "model_name")
}
await AdaptiveRouterSessionRepository(prisma_client).table.upsert(
where={
"session_id_router_name_model_name": {
"session_id": session_id,
"router_name": router,
"model_name": model,
}
},
"update": update_payload,
},
)
except Exception as e:
verbose_router_logger.exception(
"AdaptiveRouterUpdateQueue: failed to flush session for %s: %s",
key,
e,
)
data={
"create": {
"session_id": session_id,
"router_name": router,
"model_name": model,
**update_payload,
},
"update": update_payload,
},
)
pending.pop(key)
except Exception as e:
verbose_router_logger.exception(
"AdaptiveRouterUpdateQueue: failed to flush session for %s: %s",
key,
e,
)
finally:
await self._restore_session_batch(pending)
return len(batch)
return len(batch) - len(pending)
async def _restore_state_batch(self, pending: Mapping[StateKey, Mapping[str, float]]) -> None:
async with self._lock:
for key, payload in pending.items():
self._state_agg[key] = {
field: value + (self._state_agg[key][field] if key in self._state_agg else 0)
for field, value in payload.items()
}
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:
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

@ -115,3 +115,297 @@ async def test_max_size_observability(queue):
await queue.add_state_delta("r1", "code_generation", "gpt-4", 1.0, 0.0)
sizes = await queue.queue_size()
assert sizes["max_state_seen"] >= 3
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["state", "session"])
async def test_failed_flush_retains_only_failed_rows_and_reports_successes(queue, mock_prisma, kind):
for model in ("a", "b", "c"):
if kind == "state":
await queue.add_state_delta("router", "general", model, 1.0, 0.5)
else:
await queue.add_session_state("session", "router", model, {"turn_count": 1})
table = getattr(mock_prisma.db, f"litellm_adaptiverouter{kind}")
table.upsert.side_effect = [None, RuntimeError("database unavailable"), None, None]
flush = getattr(queue, f"flush_{kind}_to_db")
assert await flush(mock_prisma) == 2
assert (await queue.queue_size())[f"{kind}_pending"] == 1
assert await flush(mock_prisma) == 1
assert (await queue.queue_size())[f"{kind}_pending"] == 0
assert [call.kwargs["data"]["create"]["model_name"] for call in table.upsert.call_args_list] == ["a", "b", "c", "b"]
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["state", "session"])
async def test_failed_flush_merges_concurrent_updates_without_replaying_successes(queue, mock_prisma, kind):
started = asyncio.Event()
release = asyncio.Event()
async def fail_write(**kwargs):
started.set()
await release.wait()
raise RuntimeError("database unavailable")
table = getattr(mock_prisma.db, f"litellm_adaptiverouter{kind}")
table.upsert.side_effect = fail_write
flush = getattr(queue, f"flush_{kind}_to_db")
if kind == "state":
await queue.add_state_delta("router", "general", "model", 1.0, 0.5)
else:
await queue.add_session_state("session", "router", "model", {"turn_count": 1})
task = asyncio.create_task(flush(mock_prisma))
await started.wait()
if kind == "state":
await queue.add_state_delta("router", "general", "model", 2.0, 1.0)
else:
await queue.add_session_state("session", "router", "model", {"turn_count": 2})
release.set()
assert await task == 0
table.upsert.side_effect = None
assert await flush(mock_prisma) == 1
update = table.upsert.call_args.kwargs["data"]["update"]
assert update == (
{"alpha": {"increment": 3.0}, "beta": {"increment": 1.5}, "total_samples": {"increment": 2}}
if kind == "state"
else {"turn_count": 2}
)
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["state", "session"])
async def test_cancelled_flush_retains_unfinished_rows_without_replaying_acknowledged_rows(queue, mock_prisma, kind):
started = asyncio.Event()
blocked = asyncio.Event()
async def write(**kwargs):
if kwargs["data"]["create"]["model_name"] == "b":
started.set()
await blocked.wait()
for model in ("a", "b", "c"):
if kind == "state":
await queue.add_state_delta("router", "general", model, 1.0, 0.5)
else:
await queue.add_session_state("session", "router", model, {"turn_count": 1})
table = getattr(mock_prisma.db, f"litellm_adaptiverouter{kind}")
table.upsert.side_effect = write
flush = getattr(queue, f"flush_{kind}_to_db")
task = asyncio.create_task(flush(mock_prisma))
await started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert (await queue.queue_size())[f"{kind}_pending"] == 2
table.upsert.side_effect = None
assert await flush(mock_prisma) == 2
assert [call.kwargs["data"]["create"]["model_name"] for call in table.upsert.call_args_list] == ["a", "b", "b", "c"]
@pytest.mark.asyncio
async def test_overlapping_session_flushes_cannot_persist_an_older_snapshot_last(queue, mock_prisma):
started = asyncio.Event()
release = asyncio.Event()
stored = []
async def write(**kwargs):
turn_count = kwargs["data"]["update"]["turn_count"]
if turn_count == 1:
started.set()
await release.wait()
stored.append(turn_count)
mock_prisma.db.litellm_adaptiveroutersession.upsert.side_effect = write
await queue.add_session_state("session", "router", "model", {"turn_count": 1})
first = asyncio.create_task(queue.flush_session_to_db(mock_prisma))
await started.wait()
await queue.add_session_state("session", "router", "model", {"turn_count": 2})
second = asyncio.create_task(queue.flush_session_to_db(mock_prisma))
await asyncio.sleep(0)
release.set()
assert await asyncio.gather(first, second) == [1, 1]
assert stored == [1, 2]
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["state", "session"])
async def test_repeated_write_failures_keep_the_queue_until_recovery(queue, mock_prisma, kind):
if kind == "state":
await queue.add_state_delta("router", "general", "model", 1.0, 0.5)
else:
await queue.add_session_state("session", "router", "model", {"turn_count": 1})
table = getattr(mock_prisma.db, f"litellm_adaptiverouter{kind}")
table.upsert.side_effect = RuntimeError("database unavailable")
flush = getattr(queue, f"flush_{kind}_to_db")
for _ in range(3):
assert await flush(mock_prisma) == 0
assert (await queue.queue_size())[f"{kind}_pending"] == 1
table.upsert.side_effect = None
assert await flush(mock_prisma) == 1
assert await flush(mock_prisma) == 0
update = table.upsert.call_args.kwargs["data"]["update"]
assert update == (
{"alpha": {"increment": 1.0}, "beta": {"increment": 0.5}, "total_samples": {"increment": 1}}
if kind == "state"
else {"turn_count": 1}
)
@pytest.mark.asyncio
async def test_unacknowledged_state_commit_is_retried_with_at_least_once_semantics(queue, mock_prisma):
persisted = []
async def commit_then_lose_response(**kwargs):
persisted.append(kwargs["data"]["update"])
if len(persisted) == 1:
raise ConnectionError("commit response lost")
mock_prisma.db.litellm_adaptiverouterstate.upsert.side_effect = commit_then_lose_response
await queue.add_state_delta("router", "general", "model", 1.0, 0.5)
assert await queue.flush_state_to_db(mock_prisma) == 0
assert (await queue.queue_size())["state_pending"] == 1
assert await queue.flush_state_to_db(mock_prisma) == 1
assert (
persisted == [{"alpha": {"increment": 1.0}, "beta": {"increment": 0.5}, "total_samples": {"increment": 1}}] * 2
)
@pytest.mark.asyncio
@pytest.mark.parametrize("kind", ["state", "session"])
async def test_queue_high_water_mark_includes_restored_rows_and_new_keys(queue, mock_prisma, kind):
started = asyncio.Event()
release = asyncio.Event()
async def fail_write(**kwargs):
started.set()
await release.wait()
raise RuntimeError("database unavailable")
table = getattr(mock_prisma.db, f"litellm_adaptiverouter{kind}")
table.upsert.side_effect = fail_write
if kind == "state":
await queue.add_state_delta("router", "general", "a", 1.0, 0.5)
else:
await queue.add_session_state("session", "router", "a", {"turn_count": 1})
task = asyncio.create_task(getattr(queue, f"flush_{kind}_to_db")(mock_prisma))
await started.wait()
for model in ("b", "c"):
if kind == "state":
await queue.add_state_delta("router", "general", model, 1.0, 0.5)
else:
await queue.add_session_state("session", "router", model, {"turn_count": 1})
release.set()
assert await task == 0
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