mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(routing): retain adaptive feedback after failed flushes
This commit is contained in:
parent
5ccb1a143b
commit
c54fa59e3f
3 changed files with 310 additions and 76 deletions
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 ---------------------------------------------------
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue