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