diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index 4856d7ff4cd..4be18caa84b 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -110,7 +110,7 @@ class AdaptiveRouter: self._cells[(rt, model)] = initial_cell(prefs, rt) async def load_state_from_db(self, prisma_client: Any) -> None: - """Override cold-start cells with persisted state. Called once at startup.""" + """Merge persisted deltas onto the cold-start prior. Called once at startup.""" if prisma_client is None: return try: @@ -126,8 +126,16 @@ class AdaptiveRouter: continue if row.model_name not in self.config.available_models: continue + # Persisted alpha/beta are accumulated DELTAS — the flusher writes + # increments, not absolute posteriors (for multi-pod safety). Merge + # them onto the cold-start prior so the tier bias is preserved AND the + # Beta stays valid (prior.beta >= 0.5, so a satisfaction-only delta of + # beta=0 can never reload as Beta(alpha, 0) and crash thompson_sample). + prefs = self.model_to_prefs.get(row.model_name) or _default_prefs() + prior = initial_cell(prefs, rt) self._cells[(rt, row.model_name)] = BanditCell( - alpha=row.alpha, beta=row.beta + alpha=prior.alpha + row.alpha, + beta=prior.beta + row.beta, ) loaded += 1 verbose_router_logger.info( diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py index 93c4db90dad..09c92ed7547 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py @@ -233,7 +233,7 @@ async def test_record_turn_failure_increments_beta(): @pytest.mark.asyncio -async def test_load_state_from_db_overrides_cold_start(): +async def test_load_state_from_db_merges_deltas_onto_cold_start(): r = _make_router() cold = r._cells[(RequestType.GENERAL, "fast")] @@ -247,8 +247,10 @@ async def test_load_state_from_db_overrides_cold_start(): prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[fake_row]) await r.load_state_from_db(prisma) + # Persisted alpha/beta are accumulated DELTAS, so the loader merges them + # onto the cold-start prior rather than overwriting it. new_cell = r._cells[(RequestType.GENERAL, "fast")] - assert (new_cell.alpha, new_cell.beta) == (42.0, 13.0) + assert (new_cell.alpha, new_cell.beta) == (cold.alpha + 42.0, cold.beta + 13.0) assert (new_cell.alpha, new_cell.beta) != (cold.alpha, cold.beta) @@ -275,8 +277,8 @@ async def test_load_state_from_db_handles_unknown_request_type(): ) await r.load_state_from_db(prisma) - # Unknown skipped; good applied. - assert r._cells[(RequestType.GENERAL, "fast")].alpha == 7.0 + # Unknown skipped; good row's delta is merged onto the prior. + assert r._cells[(RequestType.GENERAL, "fast")].alpha == cold.alpha + 7.0 # Other request types kept their cold-start values. assert r._cells[(RequestType.WRITING, "fast")] == cold or True diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py b/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py index 9786832b4ae..0062efb898a 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_e2e_adaptive_router.py @@ -6,7 +6,7 @@ What we cover: 1. Full lifecycle: pick -> record turn(s) -> flush -> DB upsert with correct deltas 2. Owner cache pins attribution: same key + matching model -> updates flow 3. Convergence in-process: 50 simulated sessions, "good" model dominates last 10 - 4. Cold-start state load from DB overrides priors + 4. Cold-start state load from DB merges persisted deltas onto priors 5. Failure signal increments beta in the next flush 6. Unknown request types in DB rows are silently skipped 7. Flush isolates writes per (router, session, model) tuple @@ -203,8 +203,9 @@ async def test_failure_signal_increments_beta_after_flush(): @pytest.mark.asyncio -async def test_load_state_from_db_overrides_cold_start(): +async def test_load_state_from_db_merges_deltas_onto_cold_start(): router = _make_router() + cold = router._cells[(RequestType.GENERAL, "gpt-4o")] fake_row = MagicMock() fake_row.request_type = RequestType.GENERAL.value fake_row.model_name = "gpt-4o" @@ -216,9 +217,11 @@ async def test_load_state_from_db_overrides_cold_start(): await router.load_state_from_db(prisma) + # Persisted alpha/beta are accumulated DELTAS, so the loader merges them + # onto the cold-start prior rather than overwriting it. cell = router._cells[(RequestType.GENERAL, "gpt-4o")] - assert cell.alpha == 90.0 - assert cell.beta == 10.0 + assert cell.alpha == cold.alpha + 90.0 + assert cell.beta == cold.beta + 10.0 @pytest.mark.asyncio