From d79600987edb658e07e8d6e9f0ce83f63d36bbeb Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 30 Sep 2026 00:29:20 -0700 Subject: [PATCH] perf(router): honour the cooldown read interval in the routing prefetch (#43815) Resolves LIT-9043 Co-authored-by: yassin Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/caching/dual_cache.py | 14 + litellm/router_utils/routing_read_batch.py | 71 ++- tests/unit/caching/test_dual_cache.py | 30 ++ .../test_request_redis_batch_pre_call.py | 422 +++++++++++++++++- 4 files changed, 534 insertions(+), 3 deletions(-) diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index 4af1edae457..47ce1d35895 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -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(): diff --git a/litellm/router_utils/routing_read_batch.py b/litellm/router_utils/routing_read_batch.py index 4039d7b1508..e465f06640e 100644 --- a/litellm/router_utils/routing_read_batch.py +++ b/litellm/router_utils/routing_read_batch.py @@ -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 } diff --git a/tests/unit/caching/test_dual_cache.py b/tests/unit/caching/test_dual_cache.py index eb2f19ac377..521fda31b58 100644 --- a/tests/unit/caching/test_dual_cache.py +++ b/tests/unit/caching/test_dual_cache.py @@ -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) diff --git a/tests/unit/caching/test_request_redis_batch_pre_call.py b/tests/unit/caching/test_request_redis_batch_pre_call.py index 3569031a3e1..d4388110131 100644 --- a/tests/unit/caching/test_request_redis_batch_pre_call.py +++ b/tests/unit/caching/test_request_redis_batch_pre_call.py @@ -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)