mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
parent
6cf51383bf
commit
d79600987e
4 changed files with 534 additions and 3 deletions
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue