test(adaptive_router): align load_state_from_db tests with merge-delta contract

load_state_from_db now merges persisted deltas onto the cold-start prior instead
of overwriting it (flush_state_to_db writes Prisma increments, not absolute
posteriors), so the three tests that pinned the old overwrite behavior were
asserting the raw delta row back. Update them to assert prior + delta, and rename
the two `test_load_state_from_db_overrides_cold_start` cases to
`test_load_state_from_db_merges_deltas_onto_cold_start` to match the new semantics.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
sha-ir 2026-05-31 16:18:48 -04:00
parent 431bff3168
commit 9c155ae5a8
2 changed files with 13 additions and 8 deletions

View file

@ -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

View file

@ -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