From c54fa59e3f9c2868e915397c4236fd351d7cbb73 Mon Sep 17 00:00:00 2001 From: Vaishnavi220506 Date: Thu, 1 Oct 2026 23:39:59 +0530 Subject: [PATCH 1/3] fix(routing): retain adaptive feedback after failed flushes --- .../router_strategy/adaptive_router/README.md | 14 +- .../adaptive_router/update_queue.py | 182 ++++++++++------- .../adaptive_router/test_update_queue.py | 190 ++++++++++++++++++ 3 files changed, 310 insertions(+), 76 deletions(-) diff --git a/litellm/router_strategy/adaptive_router/README.md b/litellm/router_strategy/adaptive_router/README.md index 09420a8dd9d..0bf6cae3bf7 100644 --- a/litellm/router_strategy/adaptive_router/README.md +++ b/litellm/router_strategy/adaptive_router/README.md @@ -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. diff --git a/litellm/router_strategy/adaptive_router/update_queue.py b/litellm/router_strategy/adaptive_router/update_queue.py index e28f2379f9c..3f29ce1bab2 100644 --- a/litellm/router_strategy/adaptive_router/update_queue.py +++ b/litellm/router_strategy/adaptive_router/update_queue.py @@ -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 --------------------------------------------------- diff --git a/tests/unit/router_strategy/adaptive_router/test_update_queue.py b/tests/unit/router_strategy/adaptive_router/test_update_queue.py index 9baa69a19e0..d7ab46d64d4 100644 --- a/tests/unit/router_strategy/adaptive_router/test_update_queue.py +++ b/tests/unit/router_strategy/adaptive_router/test_update_queue.py @@ -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 From 62b23f95ac582c61734ed5caaf774badcdc4ac1b Mon Sep 17 00:00:00 2001 From: Vaishnavi220506 Date: Fri, 2 Oct 2026 10:10:11 +0530 Subject: [PATCH 2/3] fix(routing): remove obsolete mutability suppressions --- litellm/router_strategy/adaptive_router/update_queue.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/litellm/router_strategy/adaptive_router/update_queue.py b/litellm/router_strategy/adaptive_router/update_queue.py index 3f29ce1bab2..23aa28c3df6 100644 --- a/litellm/router_strategy/adaptive_router/update_queue.py +++ b/litellm/router_strategy/adaptive_router/update_queue.py @@ -112,7 +112,7 @@ class AdaptiveRouterUpdateQueue: # Sort keys to give deterministic write order across writers and # reduce the chance of cross-row deadlocks when other workers race us. - pending: Final = dict(batch) # mutable-ok: tracks unacknowledged rows until the batch is restored + pending: Final = dict(batch) try: for key, payload in sorted(batch.items()): router, rt, model = key @@ -174,7 +174,7 @@ class AdaptiveRouterUpdateQueue: if not batch: return 0 - pending: Final = dict(batch) # mutable-ok: tracks unacknowledged rows until the batch is restored + pending: Final = dict(batch) try: for key, payload in sorted(batch.items()): session_id, router, model = key @@ -221,7 +221,7 @@ class AdaptiveRouterUpdateQueue: 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 + 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() } @@ -229,7 +229,7 @@ class AdaptiveRouterUpdateQueue: 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._session_agg = {**pending, **self._session_agg} self._max_session_size_seen = max(self._max_session_size_seen, len(self._session_agg)) # ---- Observability --------------------------------------------------- From 546984a2b219c3cafd15e4dc0f4217f28d276a34 Mon Sep 17 00:00:00 2001 From: Vaishnavi220506 Date: Fri, 2 Oct 2026 10:24:16 +0530 Subject: [PATCH 3/3] fix(routing): bound failed session snapshot retention --- .../router_strategy/adaptive_router/README.md | 11 +- .../adaptive_router/update_queue.py | 13 ++- .../adaptive_router/test_update_queue.py | 104 ++++++++++++++++++ 3 files changed, 124 insertions(+), 4 deletions(-) diff --git a/litellm/router_strategy/adaptive_router/README.md b/litellm/router_strategy/adaptive_router/README.md index 0bf6cae3bf7..9bde6272861 100644 --- a/litellm/router_strategy/adaptive_router/README.md +++ b/litellm/router_strategy/adaptive_router/README.md @@ -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. diff --git a/litellm/router_strategy/adaptive_router/update_queue.py b/litellm/router_strategy/adaptive_router/update_queue.py index 23aa28c3df6..8893e3d2f21 100644 --- a/litellm/router_strategy/adaptive_router/update_queue.py +++ b/litellm/router_strategy/adaptive_router/update_queue.py @@ -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 --------------------------------------------------- diff --git a/tests/unit/router_strategy/adaptive_router/test_update_queue.py b/tests/unit/router_strategy/adaptive_router/test_update_queue.py index d7ab46d64d4..4b3dbb52565 100644 --- a/tests/unit/router_strategy/adaptive_router/test_update_queue.py +++ b/tests/unit/router_strategy/adaptive_router/test_update_queue.py @@ -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