diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index cd9b2c071a5..410f2a84c6e 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -149,7 +149,7 @@ class AdaptiveRouter: ) async def load_state_from_db(self, prisma_client: Any) -> None: - """Override cold-start cells with persisted state. Called once at startup.""" + """Add persisted online deltas to configured priors. Called once at startup.""" if prisma_client is None: return try: @@ -165,7 +165,8 @@ class AdaptiveRouter: continue if row.model_name not in self.config.available_models: continue - self._cells[(rt, row.model_name)] = BanditCell(alpha=row.alpha, beta=row.beta) + key = (rt, row.model_name) # rebind-ok: each persisted row has its own cell + self._cells[key] = apply_delta(self._cells[key], row.alpha, row.beta) loaded += 1 verbose_router_logger.info( "AdaptiveRouter[%s]: loaded %d cells from DB", 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 2de3bc6b04f..88339b13eeb 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 @@ -297,7 +297,7 @@ async def test_record_turn_bounds_feedback_contexts_and_evicts_least_recent_sess @pytest.mark.asyncio -async def test_load_state_from_db_overrides_cold_start(): +async def test_load_state_from_db_adds_online_deltas_to_cold_start(): r = _make_router() cold = r._cells[(RequestType.GENERAL, "fast")] @@ -312,8 +312,38 @@ async def test_load_state_from_db_overrides_cold_start(): await r.load_state_from_db(prisma) 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, cold.beta) + assert (new_cell.alpha, new_cell.beta) == (cold.alpha + 42.0, cold.beta + 13.0) + + +@pytest.mark.asyncio +async def test_load_state_from_db_preserves_evaluation_priors(): + cfg: Final = AdaptiveRouterConfig( + available_models=["fast"], + evaluation_priors=( + AdaptiveRouterEvaluationPrior( + request_type=RequestType.GENERAL, + model="fast", + successes=18, + failures=2, + ), + ), + ) + router: Final = AdaptiveRouter( + router_name="seeded", + config=cfg, + model_to_prefs={"fast": AdaptiveRouterPreferences(quality_tier=2)}, + model_to_cost={"fast": 0.001}, + ) + seeded: Final = router._cells[(RequestType.GENERAL, "fast")] + row: Final = MagicMock(request_type="general", model_name="fast", alpha=2.0, beta=3.0) + prisma: Final = MagicMock() + prisma.db.litellm_adaptiverouterstate.find_many = AsyncMock(return_value=[row]) + + await router.load_state_from_db(prisma) + + restored: Final = router._cells[(RequestType.GENERAL, "fast")] + assert restored.alpha == seeded.alpha + 2.0 + assert restored.beta == seeded.beta + 3.0 @pytest.mark.asyncio @@ -338,7 +368,7 @@ 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 + 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 3071f916ef1..ff79b35e607 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 @@ -186,8 +186,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_adds_online_deltas_to_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" @@ -200,8 +201,8 @@ async def test_load_state_from_db_overrides_cold_start(): await router.load_state_from_db(prisma) 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