mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
fix(least-busy): fall back to per-worker counts when Redis is unreadable
Keep each worker's own in-flight counter up to date alongside the shared one, so a Redis outage routes on that worker's counts the way it did before this branch instead of treating every deployment as idle. Floor a counter at zero when a decrement finds the key gone, which happens when a request outlives the 1 hour TTL, so an expired counter cannot settle at -1 and win every pick.
This commit is contained in:
parent
48cc4efca3
commit
c5aa4f0718
2 changed files with 95 additions and 41 deletions
|
|
@ -39,7 +39,7 @@ class _Deployment(TypedDict):
|
|||
_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, ...])
|
||||
_MEMORY_COUNTS: Final = TypeAdapter(tuple[float | None, ...] | None)
|
||||
|
||||
|
||||
def _request_count_key(model_group: str, deployment_id: str) -> str:
|
||||
|
|
@ -68,8 +68,20 @@ def _request_count_keys(model_group: str, healthy_deployments: Sequence[Mapping[
|
|||
)
|
||||
|
||||
|
||||
def _as_count(value: float | None) -> int:
|
||||
return 0 if value is None else int(value)
|
||||
def _as_counts(values: Sequence[float | None]) -> tuple[int, ...]:
|
||||
return tuple(0 if value is None else int(value) for value in values)
|
||||
|
||||
|
||||
def _shared_counts(raw: object, keys: tuple[str, ...]) -> tuple[int, ...]:
|
||||
by_key: Final = _REDIS_COUNTS.validate_python(raw)
|
||||
return _as_counts([by_key.get(key) for key in keys])
|
||||
|
||||
|
||||
def _local_counts(raw: object, keys: tuple[str, ...]) -> tuple[int, ...]:
|
||||
values: Final = _MEMORY_COUNTS.validate_python(raw)
|
||||
if values is None or len(values) != len(keys):
|
||||
return (0,) * len(keys)
|
||||
return _as_counts(values)
|
||||
|
||||
|
||||
def _least_busy(
|
||||
|
|
@ -82,7 +94,8 @@ def _least_busy(
|
|||
|
||||
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",
|
||||
"least-busy routing could not read the shared in-flight counts for %s, "
|
||||
"falling back to this worker's own counts: %s",
|
||||
model_group,
|
||||
error,
|
||||
)
|
||||
|
|
@ -135,23 +148,31 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
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)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
shared: Final = _shared_counts(redis_cache.batch_get_cache(key_list=list(keys)), keys)
|
||||
except Exception as e:
|
||||
_warn_unreadable(model_group, e)
|
||||
else:
|
||||
return _least_busy(healthy_deployments, shared)
|
||||
local: Final = _local_counts(self.router_cache.batch_get_cache(list(keys), local_only=True), keys)
|
||||
return _least_busy(healthy_deployments, local)
|
||||
|
||||
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)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
shared: Final = _shared_counts(await redis_cache.async_batch_get_cache(key_list=list(keys)), keys)
|
||||
except Exception as e:
|
||||
_warn_unreadable(model_group, e)
|
||||
else:
|
||||
return _least_busy(healthy_deployments, shared)
|
||||
local: Final = _local_counts(await self.router_cache.async_batch_get_cache(list(keys), local_only=True), keys)
|
||||
return _least_busy(healthy_deployments, local)
|
||||
|
||||
def _increment(self, kwargs: Mapping[str, object], delta: int) -> None:
|
||||
ref: Final = _deployment_ref(kwargs)
|
||||
|
|
@ -160,10 +181,16 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
key: Final = _request_count_key(*ref)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
try:
|
||||
local: Final = self.router_cache.increment_cache(
|
||||
key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS
|
||||
)
|
||||
if local < 0:
|
||||
self.router_cache.set_cache(key, 0, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
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)
|
||||
return
|
||||
shared: Final = redis_cache.increment_cache(key, delta, ttl=IN_FLIGHT_COUNT_TTL_SECONDS, refresh_ttl=True)
|
||||
if shared < 0:
|
||||
redis_cache.set_cache(key, 0, ttl=IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
except Exception as e:
|
||||
_warn_unwritable(key, e)
|
||||
|
||||
|
|
@ -174,27 +201,17 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
key: Final = _request_count_key(*ref)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
try:
|
||||
local: Final = await self.router_cache.async_increment_cache(
|
||||
key, delta, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS
|
||||
)
|
||||
if local is not None and local < 0:
|
||||
await self.router_cache.async_set_cache(key, 0, local_only=True, ttl=IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
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)
|
||||
return
|
||||
shared: Final = await redis_cache.async_increment(
|
||||
key, delta, ttl=IN_FLIGHT_COUNT_TTL_SECONDS, refresh_ttl=True
|
||||
)
|
||||
if shared < 0:
|
||||
await redis_cache.async_set_cache(key, 0, ttl=IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
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)
|
||||
|
|
|
|||
|
|
@ -133,14 +133,51 @@ class UnavailableRedis(SharedRedisCounters):
|
|||
raise ConnectionError("redis is down")
|
||||
|
||||
|
||||
def test_redis_outage_never_fails_the_request() -> None:
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_redis_outage_falls_back_to_this_workers_own_counts() -> 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_B
|
||||
assert await worker.async_get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
|
||||
await worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_A
|
||||
|
||||
|
||||
def test_a_shared_counter_that_expired_mid_request_cannot_go_negative() -> None:
|
||||
shared: Final = SharedRedisCounters()
|
||||
worker: Final = _worker(shared)
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
shared.encoded.clear()
|
||||
worker.log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert shared.count(f"{GROUP}_request_count:dep-a") == 0
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
|
||||
assert worker.get_available_deployments(GROUP, HEALTHY) is DEPLOYMENT_B
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a_local_counter_that_expired_mid_request_cannot_go_negative() -> None:
|
||||
worker: Final = _worker(None)
|
||||
in_memory: Final = worker.router_cache.in_memory_cache
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
in_memory.delete_cache(f"{GROUP}_request_count:dep-a")
|
||||
await worker.async_log_success_event(_call_kwargs("dep-a"), None, None, None)
|
||||
|
||||
assert worker.router_cache.get_cache(f"{GROUP}_request_count:dep-a") == 0
|
||||
|
||||
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
|
||||
|
||||
|
||||
def test_calls_without_a_deployment_are_ignored() -> None:
|
||||
shared: Final = SharedRedisCounters()
|
||||
worker: Final = _worker(shared)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue