fix(routing): retain adaptive feedback after failed flushes

This commit is contained in:
Vaishnavi220506 2026-10-01 23:39:59 +05:30
parent 5ccb1a143b
commit c54fa59e3f
3 changed files with 310 additions and 76 deletions

View file

@ -67,10 +67,22 @@ 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. 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. A prolonged database outage retains pending session keys
in memory until a later flush succeeds
- **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

@ -42,6 +42,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
@ -93,6 +95,10 @@ class AdaptiveRouterUpdateQueue:
# ---- 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 +112,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) # mutable-ok: tracks unacknowledged rows until the batch is restored
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 +174,63 @@ 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) # mutable-ok: tracks unacknowledged rows until the batch is restored
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] = { # mutable-ok: queue-owned accumulator merges unacknowledged deltas
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:
async with self._lock:
self._session_agg = {**pending, **self._session_agg} # mutable-ok: newer queued snapshots win over retries
self._max_session_size_seen = max(self._max_session_size_seen, len(self._session_agg))
# ---- Observability ---------------------------------------------------

View file

@ -115,3 +115,193 @@ 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