This commit is contained in:
sha-ir 2026-06-27 14:58:33 +03:00 • committed by GitHub
commit ede4159676
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 23 additions and 10 deletions

View file

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

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