perf(router): honour the cooldown read interval in the routing prefetch (#43815)

Resolves LIT-9043

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-30 00:29:20 -07:00 • committed by GitHub
parent 6cf51383bf
commit d79600987e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 534 additions and 3 deletions

View file

@ -326,6 +326,20 @@ class DualCache(BaseCache):
return sublist_keys, previous_access_times
def reserve_redis_batch_reads(self, keys: Sequence[str]) -> tuple[list[str], dict[str, float | None]]:
"""Reserve the memory-missed keys whose throttled Redis reads are due, as a batch read would."""
if self.redis_cache is None:
return [], {} # mutable-ok: API contract returns an empty list and dictionary
key_list: Final = list(keys) # mutable-ok: batch_get_cache takes a list
memory: Final = self.in_memory_cache
in_memory_result: Final = (
None
if memory is None # pyright: ignore[reportUnnecessaryComparison] # handle an absent in-memory tier
else memory.batch_get_cache(key_list)
)
result: Final = in_memory_result if in_memory_result is not None else tuple(None for _ in key_list)
return self._reserve_redis_batch_keys(time.time(), key_list, result)
def _rollback_redis_batch_key_reservations(self, previous_access_times: dict[str, float | None]) -> None:
with self._last_redis_batch_access_time_lock:
for key, previous_time in previous_access_times.items():

View file

@ -8,6 +8,7 @@ different objects. `RoutingReadBatch` fetches both key sets in one
the usage slice to the strategy, so selection does not read again.
"""
import asyncio
import itertools
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
@ -35,13 +36,54 @@ else:
_PREFETCH_SLOT: Final = "routing_read"
async def _backfill_prefetched_cache(
cache: DualCache,
due_keys: tuple[str, ...],
values: Mapping[str, object],
) -> None:
cache_keys: Final = list(due_keys) # mutable-ok: _prepare_batch_get takes a list
prepare_batch_get: Final = cache._prepare_batch_get # pyright: ignore[reportPrivateUsage] # memory backfill
pending: Final = await prepare_batch_get(cache_keys, local_only=True)
redis_values: Final = { # mutable-ok: _apply_batch_get accepts a dictionary
key: values[key]
for key, local in zip(due_keys, pending.result)
if local is None and values.get(key) is not None
}
apply_batch_get: Final = cache._apply_batch_get # pyright: ignore[reportPrivateUsage] # cache backfill
await apply_batch_get(pending, redis_values)
@dataclass(frozen=True, slots=True)
class RoutingPrefetch:
"""The cooldown and usage keys of a model group, declared on the request's Redis batch before admission
flushes it, so the routing read rides the same round trip as the rate limiter's Lua calls."""
keys: frozenset[str]
fetched: frozenset[str]
result: BatchResult[Mapping[str, object]]
reservations: tuple[tuple[DualCache, tuple[str, ...], dict[str, float | None]], ...]
def release(self) -> None:
for cache, _, previous_access_times in self.reservations:
cache._rollback_redis_batch_key_reservations( # pyright: ignore[reportPrivateUsage] # rollback
previous_access_times
)
async def _settle(self, future: asyncio.Future[Mapping[str, object]]) -> None:
if future.cancelled():
self.release()
return
if future.exception() is not None:
self.release()
return
values: Final = future.result()
try:
for cache, due_keys, _ in self.reservations:
await _backfill_prefetched_cache(cache, due_keys, values)
except Exception:
self.release()
raise
@staticmethod
def arm(
@ -60,9 +102,28 @@ class RoutingPrefetch:
() if usage_selector is None else tuple(itertools.chain(*usage_selector.usage_counter_keys(deployments)))
)
keys: Final = (*cooldown_keys, *usage_keys)
request.prefetched[_PREFETCH_SLOT] = RoutingPrefetch(
keys=frozenset(keys), result=request.batch(redis_cache).mget(keys)
cooldown_store: Final = litellm_router_instance.cooldown_cache.cooldown_store
cooldown_due, cooldown_previous = cooldown_store.reserve_redis_batch_reads(cooldown_keys)
usage_cache: Final = None if usage_selector is None else usage_selector.router_cache
usage_reservation: Final = None if usage_cache is None else usage_cache.reserve_redis_batch_reads(usage_keys)
usage_due: Final = () if usage_reservation is None else tuple(usage_reservation[0])
due: Final = (*cooldown_due, *usage_due)
reservations: Final = (
(cooldown_store, tuple(cooldown_due), cooldown_previous),
*(
()
if usage_cache is None or usage_reservation is None
else ((usage_cache, usage_due, usage_reservation[1]),)
),
)
if not due:
return
result: Final = request.batch(redis_cache).mget(due)
prefetch: Final = RoutingPrefetch(
keys=frozenset(keys), fetched=frozenset(due), result=result, reservations=reservations
)
result.on_settled(prefetch._settle)
request.prefetched[_PREFETCH_SLOT] = prefetch
@staticmethod
def armed() -> bool:
@ -78,6 +139,8 @@ class RoutingPrefetch:
armed: Final = request.prefetched.pop(_PREFETCH_SLOT, None)
if isinstance(armed, RoutingPrefetch) and armed.keys.issuperset(needed):
return armed
if isinstance(armed, RoutingPrefetch):
armed.release()
return None
@ -149,6 +212,10 @@ class RoutingReadBatch:
results: Final[list[list[object | None] | None]] = [] # mutable-ok: filled per read below
for cache, keys in reads:
pending = await cache._prepare_batch_get(keys, local_only=True) # pyright: ignore[reportPrivateUsage] # same two-step read as async_batch_get_cache_shared
if any(
key not in prefetch.fetched for key, local_value in zip(keys, pending.result) if local_value is None
):
return None
missed = { # mutable-ok: _apply_batch_get takes a dict
key: values.get(key) for key, local in zip(keys, pending.result) if local is None
}

View file

@ -2,6 +2,7 @@ import asyncio
import logging
import time
import uuid
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -136,6 +137,35 @@ def test_dual_cache_batch_get_cache_rolls_back_redis_reservation_on_error():
assert "shared_a" not in dual_cache.last_redis_batch_access_time
def test_reserve_redis_batch_reads_reserves_memory_misses_and_can_be_rolled_back():
mock_redis: Final = MagicMock(spec=RedisCache)
dual_cache: Final = DualCache(
in_memory_cache=InMemoryCache(),
redis_cache=mock_redis,
default_redis_batch_cache_expiry=10,
)
dual_cache.in_memory_cache.set_cache("memory_key", "memory_value")
reserved, previous_access_times = dual_cache.reserve_redis_batch_reads(["memory_key", "missing_key"])
assert reserved == ["missing_key"]
assert previous_access_times == {"missing_key": None}
assert dual_cache.reserve_redis_batch_reads(["memory_key", "missing_key"]) == ([], {})
dual_cache._rollback_redis_batch_key_reservations(previous_access_times)
assert dual_cache.reserve_redis_batch_reads(["memory_key", "missing_key"]) == (
["missing_key"],
{"missing_key": None},
)
def test_reserve_redis_batch_reads_returns_empty_without_redis():
dual_cache: Final = DualCache(in_memory_cache=InMemoryCache(), redis_cache=None)
assert dual_cache.reserve_redis_batch_reads(["missing_key"]) == ([], {})
def test_dual_cache_batch_get_cache_returns_memory_only_when_redis_read_is_throttled():
mock_redis = _redis_mock_for_sync_batch({"throttled_key": "redis_value"})
dual_cache = DualCache(in_memory_cache=InMemoryCache(), redis_cache=mock_redis, default_redis_batch_cache_expiry=10)

View file

@ -6,12 +6,14 @@ from __future__ import annotations
import asyncio
import hashlib
import json
from itertools import chain
from typing import Any, Final
from unittest.mock import AsyncMock, MagicMock
import pytest
from litellm import Router
import litellm.caching.dual_cache as dual_cache_module
from litellm.caching.dual_cache import DualCache
from litellm.caching.redis_batch import active_request_redis_batches, request_redis_batch_scope
from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LiteLLM_UserTable
@ -353,7 +355,12 @@ async def test_a_prefetch_that_does_not_cover_the_routing_keys_is_ignored_and_ro
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
armed = request.prefetched["routing_read"]
assert isinstance(armed, RoutingPrefetch)
request.prefetched["routing_read"] = RoutingPrefetch(keys=frozenset({"other"}), result=armed.result)
request.prefetched["routing_read"] = RoutingPrefetch(
keys=frozenset({"other"}),
fetched=armed.fetched,
result=armed.result,
reservations=armed.reservations,
)
deployment = await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
@ -363,6 +370,32 @@ async def test_a_prefetch_that_does_not_cover_the_routing_keys_is_ignored_and_ro
assert len(redis_cache.alone) == 1 # the shared cooldown+usage read, one round trip as in P1
@pytest.mark.asyncio
async def test_a_prefetch_with_incomplete_usage_keys_releases_cooldown_reservations():
client: Final = FakeClient(_lua_ok_replies)
redis_cache: Final = FakeRedisCache(client)
router: Final = _router(redis_cache)
cooldown_keys: Final = frozenset(
{
CooldownCache.get_cooldown_cache_key("dep-a"),
CooldownCache.get_cooldown_cache_key("dep-b"),
}
)
with request_redis_batch_scope():
RoutingPrefetch.arm(router, router.lowesttpm_logger_v2, router.model_list[:1])
deployment: Final = await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
fallback_cooldown_mgets: Final = tuple(
keys for command, keys in redis_cache.alone if command == "MGET" and cooldown_keys.issubset(keys)
)
assert deployment["model_info"]["id"] in {"dep-a", "dep-b"}
assert len(fallback_cooldown_mgets) == 1
@pytest.mark.asyncio
async def test_a_failed_prefetch_falls_back_to_the_shared_read():
client = FakeClient(_lua_ok_replies, fail=ConnectionError("redis down"))
@ -377,6 +410,217 @@ async def test_a_failed_prefetch_falls_back_to_the_shared_read():
assert deployment["model_info"]["id"] in {"dep-a", "dep-b"}
assert len(redis_cache.alone) == 1
cooldown_keys: Final = frozenset(
{
CooldownCache.get_cooldown_cache_key("dep-a"),
CooldownCache.get_cooldown_cache_key("dep-b"),
}
)
fallback_cooldown_mgets: Final = tuple(
keys for command, keys in redis_cache.alone if command == "MGET" and cooldown_keys.issubset(keys)
)
assert len(fallback_cooldown_mgets) == 1
@pytest.mark.asyncio
async def test_an_abandoned_prefetch_still_backfills_the_cooldown_it_read(monkeypatch):
clock: Final = 1_000_000.0
monkeypatch.setattr(dual_cache_module.time, "time", lambda: clock)
cooldown_key: Final = CooldownCache.get_cooldown_cache_key("dep-a")
active_cooldown: Final = {
"exception_received": "429",
"status_code": "429",
"timestamp": _FAR_FUTURE,
"cooldown_time": 60,
}
def replies(command: tuple[Any, ...]) -> Any:
if command[0] == "MGET":
return [json.dumps(active_cooldown) if key == cooldown_key else None for key in command[1:]]
return _lua_ok_replies(command)
client: Final = FakeClient(replies)
redis_cache: Final = FakeRedisCache(client)
redis_cache.store[cooldown_key] = active_cooldown
router: Final = _router(redis_cache, routing_strategy="simple-shuffle")
limiter: Final = _limiter(redis_cache)
with request_redis_batch_scope():
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
await limiter.atomic_check_and_increment_by_n(
descriptors=[_descriptor("api_key", "k1", 10)],
increments=[{"requests": 1}],
)
pipeline_count: Final = len(client.pipelines)
first_cooldown_mgets: Final = tuple(
command
for command in chain.from_iterable(pipeline.commands for pipeline in client.pipelines[:pipeline_count])
if command[0] == "MGET" and cooldown_key in command[1:]
)
with request_redis_batch_scope():
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
deployment: Final = await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
cooldowns: Final = await router.cooldown_cache.async_get_active_cooldowns(["dep-a"], parent_otel_span=None)
second_cooldown_mgets: Final = tuple(
command
for command in chain.from_iterable(pipeline.commands for pipeline in client.pipelines[pipeline_count:])
if command[0] == "MGET" and cooldown_key in command[1:]
)
assert deployment["model_info"]["id"] == "dep-b"
assert [model_id for model_id, _ in cooldowns] == ["dep-a"]
assert len(first_cooldown_mgets) == 1
assert second_cooldown_mgets == ()
assert redis_cache.alone == []
@pytest.mark.asyncio
async def test_prefetch_settlement_keeps_newer_memory_values_and_backfills_misses(monkeypatch):
clock: Final = 1_000_000.0
monkeypatch.setattr(dual_cache_module.time, "time", lambda: clock)
dep_a_key: Final = CooldownCache.get_cooldown_cache_key("dep-a")
dep_b_key: Final = CooldownCache.get_cooldown_cache_key("dep-b")
old_cooldown: Final = {
"exception_received": "429",
"status_code": "429",
"timestamp": _FAR_FUTURE,
"cooldown_time": 60,
}
newer_memory_cooldown: Final = {
"exception_received": "429",
"status_code": "429",
"timestamp": _FAR_FUTURE + 1,
"cooldown_time": 60,
}
redis_only_cooldown: Final = {
"exception_received": "429",
"status_code": "429",
"timestamp": _FAR_FUTURE + 2,
"cooldown_time": 60,
}
def replies(command: tuple[Any, ...]) -> Any:
if command[0] == "MGET":
return [json.dumps(redis_cache.store[key]) if key in redis_cache.store else None for key in command[1:]]
return _lua_ok_replies(command)
client: Final = FakeClient(replies)
redis_cache: Final = FakeRedisCache(client)
redis_cache.store[dep_a_key] = old_cooldown
redis_cache.store[dep_b_key] = redis_only_cooldown
router: Final = _router(redis_cache, routing_strategy="simple-shuffle")
memory_cache: Final = router.cooldown_cache.cooldown_store.in_memory_cache
assert memory_cache is not None
with request_redis_batch_scope() as request:
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
memory_cache.set_cache(dep_a_key, newer_memory_cooldown)
await request.flush_all()
prefetched_mgets: Final = tuple(command for command in client.pipelines[0].commands if command[0] == "MGET")
assert len(prefetched_mgets) == 1
assert frozenset(prefetched_mgets[0][1:]) == frozenset({dep_a_key, dep_b_key})
assert memory_cache.get_cache(dep_a_key) == newer_memory_cooldown
assert memory_cache.get_cache(dep_b_key) == redis_only_cooldown
@pytest.mark.asyncio
async def test_an_abandoned_prefetch_whose_mget_fails_releases_its_reservation(monkeypatch):
clock: Final = 1_000_000.0
monkeypatch.setattr(dual_cache_module.time, "time", lambda: clock)
cooldown_key: Final = CooldownCache.get_cooldown_cache_key("dep-a")
active_cooldown: Final = {
"exception_received": "429",
"status_code": "429",
"timestamp": _FAR_FUTURE,
"cooldown_time": 60,
}
mget_replies: Final = iter((ConnectionError("redis down"), None))
def replies(command: tuple[Any, ...]) -> Any:
if command[0] == "MGET":
response: Final = next(mget_replies)
if isinstance(response, Exception):
return response
return [json.dumps(active_cooldown) if key == cooldown_key else None for key in command[1:]]
return _lua_ok_replies(command)
client: Final = FakeClient(replies)
redis_cache: Final = FakeRedisCache(client)
redis_cache.store[cooldown_key] = active_cooldown
router: Final = _router(redis_cache, routing_strategy="simple-shuffle")
with request_redis_batch_scope() as request:
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
await request.flush_all()
pipeline_count: Final = len(client.pipelines)
first_cooldown_mgets: Final = tuple(
command
for command in chain.from_iterable(pipeline.commands for pipeline in client.pipelines[:pipeline_count])
if command[0] == "MGET" and cooldown_key in command[1:]
)
with request_redis_batch_scope():
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
deployment: Final = await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
cooldowns: Final = await router.cooldown_cache.async_get_active_cooldowns(["dep-a"], parent_otel_span=None)
second_cooldown_mgets: Final = tuple(
command
for command in chain.from_iterable(pipeline.commands for pipeline in client.pipelines[pipeline_count:])
if command[0] == "MGET" and cooldown_key in command[1:]
)
assert deployment["model_info"]["id"] == "dep-b"
assert [model_id for model_id, _ in cooldowns] == ["dep-a"]
assert len(first_cooldown_mgets) == 1
assert len(second_cooldown_mgets) == 1
@pytest.mark.asyncio
async def test_a_cooldown_that_leaves_memory_before_routing_is_read_again(monkeypatch):
clock: Final = 1_000_000.0
monkeypatch.setattr(dual_cache_module.time, "time", lambda: clock)
cooldown_key: Final = CooldownCache.get_cooldown_cache_key("dep-a")
active_cooldown: Final = {
"exception_received": "429",
"status_code": "429",
"timestamp": _FAR_FUTURE,
"cooldown_time": 60,
}
client: Final = FakeClient(_lua_ok_replies)
redis_cache: Final = FakeRedisCache(client)
redis_cache.store[cooldown_key] = active_cooldown
router: Final = _router(redis_cache, routing_strategy="simple-shuffle")
cooldown_store: Final = router.cooldown_cache.cooldown_store
memory_cache: Final = cooldown_store.in_memory_cache
assert memory_cache is not None
memory_cache.set_cache(cooldown_key, active_cooldown)
with request_redis_batch_scope() as request:
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
await request.flush_all()
memory_cache.delete_cache(cooldown_key)
deployment: Final = await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
prefetched_mgets: Final = tuple(command for command in client.pipelines[0].commands if command[0] == "MGET")
fallback_cooldown_mgets: Final = tuple(
keys for command, keys in redis_cache.alone if command == "MGET" and cooldown_key in keys
)
assert len(prefetched_mgets) == 1
assert prefetched_mgets[0][1:] == (CooldownCache.get_cooldown_cache_key("dep-b"),)
assert deployment["model_info"]["id"] == "dep-b"
assert fallback_cooldown_mgets == ((cooldown_key,),)
@pytest.mark.asyncio
@ -423,6 +667,182 @@ async def test_simple_shuffle_prefetches_only_its_cooldown_read_into_the_admissi
assert armed.keys == {CooldownCache.get_cooldown_cache_key("dep-a")} # no usage counters for shuffle
@pytest.mark.asyncio
@pytest.mark.parametrize("routing_strategy", ["simple-shuffle", "usage-based-routing-v2"])
@pytest.mark.parametrize("with_limiter", [True, False])
async def test_requests_within_the_cooldown_read_interval_read_cooldowns_from_redis_once(
routing_strategy: str, with_limiter: bool
):
client = FakeClient(_lua_ok_replies)
redis_cache = FakeRedisCache(client)
router = _router(redis_cache, routing_strategy=routing_strategy)
limiter = _limiter(redis_cache)
request_round_trips: list[tuple[int, int]] = []
for _ in range(3):
pipeline_count = len(client.pipelines)
alone_count = len(redis_cache.alone)
with request_redis_batch_scope():
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
if with_limiter:
await limiter.atomic_check_and_increment_by_n(
descriptors=[_descriptor("api_key", "k1", 10)],
increments=[{"requests": 1}],
)
await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
request_round_trips.append((len(client.pipelines) - pipeline_count, len(redis_cache.alone) - alone_count))
pipeline_mgets = [command for pipeline in client.pipelines for command in pipeline.commands if command[0] == "MGET"]
alone_mgets = [keys for command, keys in redis_cache.alone if command == "MGET"]
cooldown_keys = {
CooldownCache.get_cooldown_cache_key("dep-a"),
CooldownCache.get_cooldown_cache_key("dep-b"),
}
cooldown_mgets = [command[1:] for command in pipeline_mgets if cooldown_keys.intersection(command[1:])] + [
keys for keys in alone_mgets if cooldown_keys.intersection(keys)
]
assert len(cooldown_mgets) == 1
if not with_limiter:
assert request_round_trips[1:] == [(0, 0), (0, 0)]
@pytest.mark.asyncio
async def test_concurrent_requests_share_one_cooldown_read_per_interval():
client: Final = FakeClient(_lua_ok_replies)
redis_cache: Final = FakeRedisCache(client)
router: Final = _router(redis_cache)
first_armed: Final = asyncio.Event()
both_armed: Final = asyncio.Event()
async def route_after_both_requests_arm():
with request_redis_batch_scope():
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
if first_armed.is_set():
both_armed.set()
else:
first_armed.set()
await both_armed.wait()
return await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
deployments: Final = await asyncio.gather(route_after_both_requests_arm(), route_after_both_requests_arm())
cooldown_keys: Final = frozenset(
{
CooldownCache.get_cooldown_cache_key("dep-a"),
CooldownCache.get_cooldown_cache_key("dep-b"),
}
)
cooldown_mgets: Final = tuple(
command
for pipeline in client.pipelines
for command in pipeline.commands
if command[0] == "MGET" and cooldown_keys.intersection(command[1:])
)
assert all(deployment["model_info"]["id"] in {"dep-a", "dep-b"} for deployment in deployments)
assert len(cooldown_mgets) == 1
@pytest.mark.asyncio
async def test_the_prefetch_reads_cooldowns_again_once_the_read_interval_elapses(monkeypatch):
first_time: Final = 1_000_000.0
monkeypatch.setattr(dual_cache_module.time, "time", lambda: first_time)
active_cooldown = {
"exception_received": "429",
"status_code": "429",
"timestamp": _FAR_FUTURE,
"cooldown_time": 60,
}
mget_results = iter((None, active_cooldown))
def replies(command: tuple[Any, ...]) -> Any:
if command[0] == "MGET":
result = next(mget_results)
return [
None if result is None or key != CooldownCache.get_cooldown_cache_key("dep-a") else json.dumps(result)
for key in command[1:]
]
return _lua_ok_replies(command)
client = FakeClient(replies)
redis_cache = FakeRedisCache(client)
router = _router(redis_cache, routing_strategy="simple-shuffle")
cooldown_store = router.cooldown_cache.cooldown_store
with request_redis_batch_scope():
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
monkeypatch.setattr(
dual_cache_module.time,
"time",
lambda: first_time + cooldown_store.redis_batch_cache_expiry + 1,
)
with request_redis_batch_scope():
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
deployment = await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
cooldown_keys = {
CooldownCache.get_cooldown_cache_key("dep-a"),
CooldownCache.get_cooldown_cache_key("dep-b"),
}
cooldown_mgets = [
command
for pipeline in client.pipelines
for command in pipeline.commands
if command[0] == "MGET" and cooldown_keys.intersection(command[1:])
]
assert len(cooldown_mgets) == 2
assert deployment["model_info"]["id"] == "dep-b"
@pytest.mark.asyncio
async def test_the_prefetch_mget_carries_only_the_keys_whose_read_is_due(monkeypatch):
first_time: Final = 1_000_000.0
monkeypatch.setattr(dual_cache_module.time, "time", lambda: first_time)
client = FakeClient(_lua_ok_replies)
redis_cache = FakeRedisCache(client)
router = _router(redis_cache, routing_strategy="usage-based-routing-v2")
cooldown_store = router.cooldown_cache.cooldown_store
usage_cache = router.lowesttpm_logger_v2.router_cache
time_offset = cooldown_store.redis_batch_cache_expiry + 0.5
assert time_offset < usage_cache.redis_batch_cache_expiry
with request_redis_batch_scope():
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
monkeypatch.setattr(dual_cache_module.time, "time", lambda: first_time + time_offset)
with request_redis_batch_scope():
router.arm_routing_read_prefetch(_MODEL_GROUP, {})
await router.async_get_available_deployment(
model=_MODEL_GROUP, messages=[{"role": "user", "content": "ping"}], request_kwargs={}
)
cooldown_keys = frozenset(
{
CooldownCache.get_cooldown_cache_key("dep-a"),
CooldownCache.get_cooldown_cache_key("dep-b"),
}
)
second_pipeline_mgets = tuple(command for command in client.pipelines[1].commands if command[0] == "MGET")
assert len(client.pipelines) == 2
assert len(second_pipeline_mgets) == 1
assert frozenset(second_pipeline_mgets[0][1:]) == cooldown_keys
@pytest.mark.asyncio
async def test_two_backends_flush_concurrently_one_pipeline_each():
a_client, b_client = FakeClient(_lua_ok_replies), FakeClient(_lua_ok_replies)