fix(least-busy): share in-flight request counts across proxy workers

Least-busy kept one dict of in-flight counts per model group in the
router cache, which reads in-memory first, so every worker and replica
routed on its own stale copy and each write overwrote the shared value.

Counts now live in one key per deployment, incremented and read through
Redis when the router has a Redis cache, and in the process-local cache
otherwise.
This commit is contained in:
mateo-berri 2026-09-05 20:37:17 -07:00
parent 11a0c0abf0
commit 48cc4efca3
7 changed files with 376 additions and 227 deletions

View file

@ -54,10 +54,10 @@
"limit": 0
},
"reportMissingParameterType": {
"limit": 5570
"limit": 5551
},
"reportMissingTypeArgument": {
"limit": 15281
"limit": 15277
},
"reportMissingTypeStubs": {
"limit": 40
@ -99,19 +99,19 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44358
"limit": 44357
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38271
"limit": 38240
},
"reportUnknownParameterType": {
"limit": 19584
"limit": 19558
},
"reportUnknownVariableType": {
"limit": 29814
"limit": 29781
},
"reportUnnecessaryCast": {
"limit": 110

View file

@ -680,7 +680,7 @@ class RedisCache(BaseCache):
# NON blocking - notify users Redis is throwing an exception
print_verbose(f"litellm.caching.caching: set() - Got exception from REDIS : {e}")
def increment_cache(self, key, value: int, ttl: float | None = None, **kwargs) -> int:
def increment_cache(self, key, value: int, ttl: float | None = None, refresh_ttl: bool = False, **kwargs) -> int:
_redis_client: Final = self.redis_client
start_time = time.time()
set_ttl: Final = self.get_ttl(ttl=ttl)
@ -701,7 +701,7 @@ class RedisCache(BaseCache):
if set_ttl is not None:
# check if key already has ttl, if not -> set ttl
start_time = time.time()
current_ttl: Final = _redis_client.ttl(key)
current_ttl: Final = -1 if refresh_ttl else _redis_client.ttl(key)
end_time = time.time()
_duration = end_time - start_time
self.service_logger_obj.service_success_hook(

View file

@ -1,17 +1,96 @@
#### What this does ####
# identifies least busy deployment
# How is this achieved?
# - Before each call, have the router print the state of requests {"deployment": "requests_in_flight"}
# - use litellm.input_callbacks to log when a request is just about to be made to a model - {"deployment-id": traffic}
# - use litellm.success + failure callbacks to log when a request completed
# - in get_available_deployment, for a given model group name -> pick based on traffic
import random
from collections.abc import Mapping, Sequence
from typing import Final
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_router_logger
from litellm.caching.caching import DualCache
from litellm.integrations.custom_logger import CustomLogger
IN_FLIGHT_COUNT_TTL_SECONDS: Final = 60 * 60
class _ModelInfo(TypedDict, total=False):
id: ReadOnly[str | int | None]
class _Metadata(TypedDict, total=False):
model_group: ReadOnly[str | None]
class _LitellmParams(TypedDict, total=False):
metadata: ReadOnly[_Metadata | None]
model_info: ReadOnly[_ModelInfo | None]
class _CallKwargs(TypedDict, total=False):
litellm_params: ReadOnly[_LitellmParams | None]
class _DeploymentModelInfo(TypedDict):
id: ReadOnly[str | int]
class _Deployment(TypedDict):
model_info: ReadOnly[_DeploymentModelInfo]
_CALL_KWARGS: Final = TypeAdapter(_CallKwargs)
_DEPLOYMENTS: Final = TypeAdapter(list[_Deployment])
_REDIS_COUNTS: Final = TypeAdapter(dict[str, float | None])
_MEMORY_COUNTS: Final = TypeAdapter(tuple[float | None, ...])
def _request_count_key(model_group: str, deployment_id: str) -> str:
return f"{model_group}_request_count:{deployment_id}"
def _deployment_ref(kwargs: Mapping[str, object]) -> tuple[str, str] | None:
try:
call: Final = _CALL_KWARGS.validate_python(kwargs)
except ValidationError:
return None
litellm_params: Final = call.get("litellm_params")
metadata: Final = litellm_params.get("metadata") if litellm_params else None
model_info: Final = litellm_params.get("model_info") if litellm_params else None
model_group: Final = metadata.get("model_group") if metadata else None
deployment_id: Final = model_info.get("id") if model_info else None
if model_group is None or deployment_id is None:
return None
return model_group, str(deployment_id)
def _request_count_keys(model_group: str, healthy_deployments: Sequence[Mapping[str, object]]) -> tuple[str, ...]:
return tuple(
_request_count_key(model_group, str(deployment["model_info"]["id"]))
for deployment in _DEPLOYMENTS.validate_python(healthy_deployments)
)
def _as_count(value: float | None) -> int:
return 0 if value is None else int(value)
def _least_busy(
healthy_deployments: Sequence[Mapping[str, object]], counts: tuple[int, ...]
) -> Mapping[str, object] | None:
if not healthy_deployments:
return None
return healthy_deployments[min(range(len(healthy_deployments)), key=lambda index: counts[index])]
def _warn_unreadable(model_group: str, error: Exception) -> None:
verbose_router_logger.warning(
"least-busy routing could not read the in-flight counts for %s, treating every deployment as idle: %s",
model_group,
error,
)
def _warn_unwritable(key: str, error: Exception) -> None:
verbose_router_logger.warning("least-busy routing could not update the in-flight count under %s: %s", key, error)
class LeastBusyLoggingHandler(CustomLogger):
test_flag: bool = False
@ -21,194 +100,101 @@ class LeastBusyLoggingHandler(CustomLogger):
def __init__(self, router_cache: DualCache):
self.router_cache = router_cache
def log_pre_api_call(self, model, messages, kwargs):
"""
Log when a model is being used.
def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None:
self._increment(kwargs, 1)
Caching based on model group.
"""
try:
if kwargs["litellm_params"].get("metadata") is None:
pass
else:
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
if model_group is None or id is None:
return
elif isinstance(id, int):
id = str(id)
def log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
self._increment(kwargs, -1)
if self.test_flag:
self.logged_success += 1
request_count_api_key: Final = f"{model_group}_request_count"
# update cache
request_count_dict: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
request_count_dict[id] = request_count_dict.get(id, 0) + 1
def log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
self._increment(kwargs, -1)
if self.test_flag:
self.logged_failure += 1
self.router_cache.set_cache(key=request_count_api_key, value=request_count_dict)
except Exception:
pass
async def async_log_success_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
await self._async_increment(kwargs, -1)
if self.test_flag:
self.logged_success += 1
def log_success_event(self, kwargs, response_obj, start_time, end_time):
try:
if kwargs["litellm_params"].get("metadata") is None:
pass
else:
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
if model_group is None or id is None:
return
elif isinstance(id, int):
id = str(id)
request_count_api_key: Final = f"{model_group}_request_count"
# decrement count in cache
request_count_dict: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
if request_count_value is None:
return
request_count_dict[id] = request_count_value - 1
self.router_cache.set_cache(key=request_count_api_key, value=request_count_dict)
### TESTING ###
if self.test_flag:
self.logged_success += 1
except Exception:
pass
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
try:
if kwargs["litellm_params"].get("metadata") is None:
pass
else:
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
if model_group is None or id is None:
return
elif isinstance(id, int):
id = str(id)
request_count_api_key: Final = f"{model_group}_request_count"
# decrement count in cache
request_count_dict: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
if request_count_value is None:
return
request_count_dict[id] = request_count_value - 1
self.router_cache.set_cache(key=request_count_api_key, value=request_count_dict)
### TESTING ###
if self.test_flag:
self.logged_failure += 1
except Exception:
pass
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
try:
if kwargs["litellm_params"].get("metadata") is None:
pass
else:
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
if model_group is None or id is None:
return
elif isinstance(id, int):
id = str(id)
request_count_api_key: Final = f"{model_group}_request_count"
# decrement count in cache
request_count_dict: Final = await self.router_cache.async_get_cache(key=request_count_api_key) or {}
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
if request_count_value is None:
return
request_count_dict[id] = request_count_value - 1
await self.router_cache.async_set_cache(key=request_count_api_key, value=request_count_dict)
### TESTING ###
if self.test_flag:
self.logged_success += 1
except Exception:
pass
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
try:
if kwargs["litellm_params"].get("metadata") is None:
pass
else:
model_group: Final = kwargs["litellm_params"]["metadata"].get("model_group", None)
id = kwargs["litellm_params"].get("model_info", {}).get("id", None)
if model_group is None or id is None:
return
elif isinstance(id, int):
id = str(id)
request_count_api_key: Final = f"{model_group}_request_count"
# decrement count in cache
request_count_dict: Final = await self.router_cache.async_get_cache(key=request_count_api_key) or {}
request_count_value: Final[int | None] = request_count_dict.get(id, 0)
if request_count_value is None:
return
request_count_dict[id] = request_count_value - 1
await self.router_cache.async_set_cache(key=request_count_api_key, value=request_count_dict)
### TESTING ###
if self.test_flag:
self.logged_failure += 1
except Exception:
pass
def _get_available_deployments(
self,
healthy_deployments: list,
all_deployments: dict,
):
"""
Helper to get deployments using least busy strategy
"""
for d in healthy_deployments:
## if healthy deployment not yet used
if d["model_info"]["id"] not in all_deployments:
all_deployments[d["model_info"]["id"]] = 0
# map deployment to id
# pick least busy deployment
min_traffic = float("inf")
min_deployment = None
for k, v in all_deployments.items():
if v < min_traffic:
min_traffic = v
min_deployment = k
if min_deployment is not None:
## check if min deployment is a string, if so, cast it to int
for m in healthy_deployments:
if m["model_info"]["id"] == min_deployment:
return m
min_deployment = random.choice(healthy_deployments)
else:
min_deployment = random.choice(healthy_deployments)
return min_deployment
async def async_log_failure_event(
self, kwargs: Mapping[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
await self._async_increment(kwargs, -1)
if self.test_flag:
self.logged_failure += 1
def get_available_deployments(
self,
model_group: str,
healthy_deployments: list,
):
"""
Sync helper to get deployments using least busy strategy
"""
request_count_api_key: Final = f"{model_group}_request_count"
all_deployments: Final = self.router_cache.get_cache(key=request_count_api_key) or {}
return self._get_available_deployments(
healthy_deployments=healthy_deployments,
all_deployments=all_deployments,
)
self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]]
) -> Mapping[str, object] | None:
keys: Final = _request_count_keys(model_group, healthy_deployments)
try:
counts: Final = tuple(_as_count(value) for value in self._read_counts(keys))
except Exception as e:
_warn_unreadable(model_group, e)
return _least_busy(healthy_deployments, (0,) * len(keys))
return _least_busy(healthy_deployments, counts)
async def async_get_available_deployments(self, model_group: str, healthy_deployments: list):
"""
Async helper to get deployments using least busy strategy
"""
request_count_api_key: Final = f"{model_group}_request_count"
all_deployments: Final = await self.router_cache.async_get_cache(key=request_count_api_key) or {}
return self._get_available_deployments(
healthy_deployments=healthy_deployments,
all_deployments=all_deployments,
)
async def async_get_available_deployments(
self, model_group: str, healthy_deployments: Sequence[Mapping[str, object]]
) -> Mapping[str, object] | None:
keys: Final = _request_count_keys(model_group, healthy_deployments)
try:
counts: Final = tuple(_as_count(value) for value in await self._async_read_counts(keys))
except Exception as e:
_warn_unreadable(model_group, e)
return _least_busy(healthy_deployments, (0,) * len(keys))
return _least_busy(healthy_deployments, counts)
def _increment(self, kwargs: Mapping[str, object], delta: int) -> None:
ref: Final = _deployment_ref(kwargs)
if ref is None:
return
key: Final = _request_count_key(*ref)
redis_cache: Final = self.router_cache.redis_cache
try:
if redis_cache is None:
self.router_cache.increment_cache(key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS)
else:
redis_cache.increment_cache(key, delta, ttl=IN_FLIGHT_COUNT_TTL_SECONDS, refresh_ttl=True)
except Exception as e:
_warn_unwritable(key, e)
async def _async_increment(self, kwargs: Mapping[str, object], delta: int) -> None:
ref: Final = _deployment_ref(kwargs)
if ref is None:
return
key: Final = _request_count_key(*ref)
redis_cache: Final = self.router_cache.redis_cache
try:
if redis_cache is None:
await self.router_cache.async_increment_cache(
key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS
)
else:
await redis_cache.async_increment(key, delta, ttl=IN_FLIGHT_COUNT_TTL_SECONDS, refresh_ttl=True)
except Exception as e:
_warn_unwritable(key, e)
def _read_counts(self, keys: tuple[str, ...]) -> tuple[float | None, ...]:
redis_cache: Final = self.router_cache.redis_cache
if redis_cache is None:
return _MEMORY_COUNTS.validate_python(self.router_cache.batch_get_cache(list(keys), local_only=True))
by_key: Final = _REDIS_COUNTS.validate_python(redis_cache.batch_get_cache(key_list=list(keys)))
return tuple(by_key.get(key) for key in keys)
async def _async_read_counts(self, keys: tuple[str, ...]) -> tuple[float | None, ...]:
redis_cache: Final = self.router_cache.redis_cache
if redis_cache is None:
return _MEMORY_COUNTS.validate_python(
await self.router_cache.async_batch_get_cache(list(keys), local_only=True)
)
by_key: Final = _REDIS_COUNTS.validate_python(await redis_cache.async_batch_get_cache(key_list=list(keys)))
return tuple(by_key.get(key) for key in keys)

View file

@ -1,6 +1,6 @@
{
"ANN001": {
"limit": 2956
"limit": 2937
},
"ANN002": {
"limit": 71
@ -9,10 +9,10 @@
"limit": 806
},
"ANN201": {
"limit": 1979
"limit": 1972
},
"ANN202": {
"limit": 831
"limit": 830
},
"ANN204": {
"limit": 683
@ -57,7 +57,7 @@
"limit": 3
},
"BLE001": {
"limit": 2916
"limit": 2915
},
"C401": {
"limit": 8
@ -189,7 +189,7 @@
"limit": 0
},
"S110": {
"limit": 217
"limit": 212
},
"S112": {
"limit": 22

View file

@ -33,8 +33,8 @@ def test_model_added():
}
}
least_busy_logger.log_pre_api_call(model="test", messages=[], kwargs=kwargs)
request_count_api_key = f"gpt-3.5-turbo_request_count"
assert test_cache.get_cache(key=request_count_api_key) is not None
request_count_api_key = "gpt-3.5-turbo_request_count:1234"
assert test_cache.get_cache(key=request_count_api_key) == 1
def test_get_available_deployments():
@ -52,8 +52,8 @@ def test_get_available_deployments():
}
}
least_busy_logger.log_pre_api_call(model="test", messages=[], kwargs=kwargs)
request_count_api_key = f"{model_group}_request_count"
assert test_cache.get_cache(key=request_count_api_key) is not None
request_count_api_key = f"{model_group}_request_count:1234"
assert test_cache.get_cache(key=request_count_api_key) == 1
# test_get_available_deployments()
@ -104,15 +104,20 @@ async def test_router_get_available_deployments(async_test):
router.leastbusy_logger.test_flag = True
model_group = "azure-model"
request_count_dict = {1: 10, 2: 54, 3: 100}
cache_key = f"{model_group}_request_count"
request_count_dict = {"1": 10, "2": 54, "3": 100}
cache_keys = {
deployment_id: f"{model_group}_request_count:{deployment_id}"
for deployment_id in request_count_dict
}
if async_test is True:
await router.cache.async_set_cache(key=cache_key, value=request_count_dict)
for deployment_id, count in request_count_dict.items():
await router.cache.async_set_cache(key=cache_keys[deployment_id], value=count)
deployment = await router.async_get_available_deployment(
model=model_group, messages=None, request_kwargs={}
)
else:
router.cache.set_cache(key=cache_key, value=request_count_dict)
for deployment_id, count in request_count_dict.items():
router.cache.set_cache(key=cache_keys[deployment_id], value=count)
deployment = router.get_available_deployment(model=model_group, messages=None)
print(f"deployment: {deployment}")
assert deployment["model_info"]["id"] == "1"
@ -124,15 +129,18 @@ async def test_router_get_available_deployments(async_test):
messages=[{"role": "user", "content": "Hey, how's it going?"}],
)
return_dict = router.cache.get_cache(key=cache_key)
# wait 2 seconds
time.sleep(2)
return_dict = {
deployment_id: router.cache.get_cache(key=cache_key)
for deployment_id, cache_key in cache_keys.items()
}
assert router.leastbusy_logger.logged_success == 1
assert return_dict[1] == 10
assert return_dict[2] == 54
assert return_dict[3] == 100
assert return_dict["1"] == 10
assert return_dict["2"] == 54
assert return_dict["3"] == 100
## Test with Real calls ##
@ -192,9 +200,11 @@ async def test_router_atext_completion_streaming():
await asyncio.sleep(random.uniform(0, 2))
await router.atext_completion(model=model, prompt=prompt, stream=True)
cache_key = f"{model}_request_count"
## check if calls equally distributed
cache_dict = router.cache.get_cache(key=cache_key)
cache_dict = {
deployment_id: router.cache.get_cache(key=f"{model}_request_count:{deployment_id}")
for deployment_id in ("1", "2", "3")
}
for k, v in cache_dict.items():
assert v == 1, f"Failed. K={k} called v={v} times, cache_dict={cache_dict}"
@ -259,8 +269,10 @@ async def test_router_completion_streaming():
await asyncio.sleep(random.uniform(0, 2))
await router.acompletion(model=model, messages=messages, stream=True)
cache_key = f"{model}_request_count"
## check if calls equally distributed
cache_dict = router.cache.get_cache(key=cache_key)
cache_dict = {
deployment_id: router.cache.get_cache(key=f"{model}_request_count:{deployment_id}")
for deployment_id in ("1", "2", "3")
}
for k, v in cache_dict.items():
assert v == 1, f"Failed. K={k} called v={v} times, cache_dict={cache_dict}"

View file

@ -0,0 +1,151 @@
import json
from typing import Final
import pytest
from litellm.caching.caching import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.router_strategy.least_busy import IN_FLIGHT_COUNT_TTL_SECONDS, LeastBusyLoggingHandler
GROUP: Final = "least-busy-group"
DEPLOYMENT_A: Final[dict[str, object]] = {"model_info": {"id": "dep-a"}}
DEPLOYMENT_B: Final[dict[str, object]] = {"model_info": {"id": "dep-b"}}
HEALTHY: Final = [DEPLOYMENT_A, DEPLOYMENT_B]
def _call_kwargs(deployment_id: str) -> dict[str, object]:
return {"litellm_params": {"metadata": {"model_group": GROUP}, "model_info": {"id": deployment_id}}}
class SharedRedisCounters:
"""Stores JSON strings and hands back a fresh object per read, the way a real Redis client does."""
def __init__(self) -> None:
self.encoded: dict[str, str] = {}
self.ttls: dict[str, float] = {}
def count(self, key: str) -> object:
raw: Final = self.encoded.get(key)
return None if raw is None else json.loads(raw)
def get_cache(self, key: str, **kwargs: object) -> object:
return self.count(key)
def set_cache(self, key: str, value: object, **kwargs: object) -> None:
self.encoded[key] = json.dumps(value)
async def async_get_cache(self, key: str, **kwargs: object) -> object:
return self.count(key)
async def async_set_cache(self, key: str, value: object, **kwargs: object) -> None:
self.set_cache(key, value)
def increment_cache(self, key: str, value: int, ttl: float | None = None, refresh_ttl: bool = False) -> int:
current: Final = self.count(key) or 0
assert isinstance(current, int)
incremented: Final = current + value
self.encoded[key] = json.dumps(incremented)
if ttl is not None and (refresh_ttl or key not in self.ttls):
self.ttls[key] = ttl
return incremented
async def async_increment(self, key: str, value: float, ttl: int | None = None, refresh_ttl: bool = False) -> float:
return self.increment_cache(key, int(value), ttl, refresh_ttl)
def batch_get_cache(self, key_list: list[str], **kwargs: object) -> dict[str, object]:
return {key: self.count(key) for key in key_list}
async def async_batch_get_cache(self, key_list: list[str], **kwargs: object) -> dict[str, object]:
return self.batch_get_cache(key_list)
def _worker(shared: SharedRedisCounters | None) -> LeastBusyLoggingHandler:
cache: Final = DualCache(in_memory_cache=InMemoryCache(), redis_cache=shared) # pyright: ignore[reportArgumentType] # duck-typed Redis double
return LeastBusyLoggingHandler(router_cache=cache)
@pytest.mark.asyncio
async def test_worker_routes_around_a_request_another_worker_started() -> None:
shared: Final = SharedRedisCounters()
streaming_worker: Final = _worker(shared)
picking_worker: Final = _worker(shared)
picking_worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
await picking_worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
streaming_worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
assert await picking_worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
await streaming_worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
assert await picking_worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A
def test_sync_pick_reads_the_shared_counts() -> None:
shared: Final = SharedRedisCounters()
streaming_worker: Final = _worker(shared)
picking_worker: Final = _worker(shared)
picking_worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
picking_worker.log_success_event(_call_kwargs("dep-a"), None, None, None)
streaming_worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
assert picking_worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
streaming_worker.log_failure_event(_call_kwargs("dep-a"), None, None, None)
assert picking_worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A
def test_redis_counts_keep_a_refreshed_ttl() -> None:
shared: Final = SharedRedisCounters()
worker: Final = _worker(shared)
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
worker.log_success_event(_call_kwargs("dep-a"), None, None, None)
assert shared.count(f"{GROUP}_request_count:dep-a") == 0
assert shared.ttls == {f"{GROUP}_request_count:dep-a": IN_FLIGHT_COUNT_TTL_SECONDS}
@pytest.mark.asyncio
async def test_counts_stay_in_memory_without_redis() -> None:
worker: Final = _worker(None)
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
assert await worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
await worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
assert await worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A
assert worker.router_cache.get_cache(f"{GROUP}_request_count:dep-a") == 0
class UnavailableRedis(SharedRedisCounters):
def batch_get_cache(self, key_list: list[str], **kwargs: object) -> dict[str, object]:
raise ConnectionError("redis is down")
def increment_cache(self, key: str, value: int, ttl: float | None = None, refresh_ttl: bool = False) -> int:
raise ConnectionError("redis is down")
def test_redis_outage_never_fails_the_request() -> None:
worker: Final = _worker(UnavailableRedis())
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A
def test_calls_without_a_deployment_are_ignored() -> None:
shared: Final = SharedRedisCounters()
worker: Final = _worker(shared)
worker.log_pre_api_call(model="m", messages=[], kwargs={"litellm_params": {"metadata": None}})
worker.log_pre_api_call(model="m", messages=[], kwargs={})
assert shared.encoded == {}

View file

@ -1,9 +1,9 @@
{
"LIT001": {
"limit": 22180
"limit": 22176
},
"LIT002": {
"limit": 26729
"limit": 26721
},
"LIT003": {
"limit": 261
@ -27,10 +27,10 @@
"limit": 0
},
"LIT010": {
"limit": 16426
"limit": 16412
},
"LIT011": {
"limit": 5506
"limit": 5505
},
"LIT012": {
"limit": 4486