mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #40009 from BerriAI/litellm_lit_7039_least_busy_shared_counts
fix(least-busy): share in-flight request counts across proxy workers
This commit is contained in:
commit
9b74e6f34e
12 changed files with 759 additions and 223 deletions
|
|
@ -1440,6 +1440,7 @@ jobs:
|
|||
TEST_FILES=$(printf "%s\n" \
|
||||
tests/local_testing/test_dual_cache.py \
|
||||
tests/local_testing/test_redis_batch_optimizations.py \
|
||||
tests/local_testing/test_redis_increment_with_floor.py \
|
||||
tests/local_testing/test_router_utils.py)
|
||||
echo "$TEST_FILES" | circleci tests run \
|
||||
--verbose \
|
||||
|
|
|
|||
|
|
@ -20,6 +20,8 @@ from contextvars import ContextVar
|
|||
from datetime import timedelta
|
||||
from typing import TYPE_CHECKING, Any, Final, Protocol, TypeVar, cast
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.constants import (
|
||||
|
|
@ -80,11 +82,29 @@ class _AsyncRedisCommands(Protocol):
|
|||
|
||||
def pipeline(self, transaction: bool = True) -> "Pipeline[bytes]": ...
|
||||
|
||||
def eval(self, script: str, numkeys: int, *keys_and_args: str | bytes | float) -> Awaitable[object]: ...
|
||||
|
||||
|
||||
_BREAKER_GUARD_FRAME_NAMES: Final = frozenset(
|
||||
{"<lambda>", "wrapper", "_run_under_circuit_breaker", "_run_under_circuit_breaker_sync"}
|
||||
)
|
||||
|
||||
_INCREMENT_WITH_FLOOR_LUA: Final = (
|
||||
"local count = redis.call('INCRBY', KEYS[1], ARGV[1]) "
|
||||
"if count < 0 then count = redis.call('INCRBY', KEYS[1], -count) end "
|
||||
"if redis.call('TTL', KEYS[1]) < 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]) end "
|
||||
"return count"
|
||||
)
|
||||
|
||||
_LUA_COUNT: Final = TypeAdapter(int)
|
||||
_OPTIONAL_COUNTS: Final = TypeAdapter(tuple[int | None, ...])
|
||||
|
||||
|
||||
def _decoded_counts(values: Sequence[bytes | str | None]) -> tuple[int | None, ...]:
|
||||
return _OPTIONAL_COUNTS.validate_python(
|
||||
tuple(value.decode("utf-8") if isinstance(value, bytes) else value for value in values)
|
||||
)
|
||||
|
||||
|
||||
def _get_call_stack_info(num_frames: int = 2) -> str:
|
||||
"""
|
||||
|
|
@ -736,6 +756,43 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
raise e
|
||||
|
||||
@_redis_circuit_breaker_guard_sync
|
||||
def increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
"""Add ``value`` to ``key``, clamp the result at zero, and give a new key ``ttl``, in one Lua call.
|
||||
|
||||
A counter whose key expired while a request was still in flight would otherwise be
|
||||
recreated negative by that request's decrement. Clamping inside the same call is what
|
||||
keeps it safe: a separate corrective write could land after another pod's increment and
|
||||
erase it.
|
||||
|
||||
The TTL is set only on a key that has none, so a counter expires ``ttl`` after it was
|
||||
created rather than ``ttl`` after it was last touched. Refreshing it on every touch
|
||||
would keep a count a dead worker never decremented alive for as long as the group
|
||||
takes traffic. Returns the resulting count.
|
||||
"""
|
||||
namespaced_key: Final = self.check_and_fix_namespace(key=key)
|
||||
count: Final[object] = self.redis_client.eval( # pyright: ignore[reportAttributeAccessIssue] # stubs omit eval
|
||||
_INCREMENT_WITH_FLOOR_LUA, 1, namespaced_key, value, ttl
|
||||
)
|
||||
return _LUA_COUNT.validate_python(count)
|
||||
|
||||
@_redis_circuit_breaker_guard_sync
|
||||
def batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
"""Read integer counters for ``key_list``, in order, raising when Redis cannot answer.
|
||||
|
||||
``batch_get_cache`` swallows every failure and returns an empty dict, which the caller
|
||||
cannot tell apart from "every counter is unset". A caller that has to fall back to its
|
||||
own numbers when Redis is unreachable needs the failure, not a dict of zeros.
|
||||
"""
|
||||
namespaced_keys: Final = [self.check_and_fix_namespace(key=key) for key in key_list]
|
||||
return _decoded_counts(self._run_redis_mget_operation(keys=namespaced_keys))
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
"""Async twin of ``batch_get_counts``, raising on failure the same way."""
|
||||
namespaced_keys: Final = [self.check_and_fix_namespace(key=key) for key in key_list]
|
||||
return _decoded_counts(await self._async_run_redis_mget_operation(keys=namespaced_keys))
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_scan_iter(self, pattern: str, count: int = 100) -> list:
|
||||
start_time: Final = time.time()
|
||||
|
|
@ -1241,6 +1298,14 @@ class RedisCache(BaseCache):
|
|||
result = result.decode()
|
||||
return float(result)
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
"""Async twin of ``increment_with_floor``, sharing its Lua script and its guarantees."""
|
||||
_redis_client: Final = self._async_commands()
|
||||
namespaced_key: Final = self.check_and_fix_namespace(key=key)
|
||||
count: Final = await _redis_client.eval(_INCREMENT_WITH_FLOOR_LUA, 1, namespaced_key, value, ttl)
|
||||
return _LUA_COUNT.validate_python(count)
|
||||
|
||||
async def flush_cache_buffer(self):
|
||||
print_verbose(f"flushing to redis....reached size of buffer {len(self.redis_batch_writing_buffer)}")
|
||||
await self.async_set_cache_pipeline(self.redis_batch_writing_buffer)
|
||||
|
|
|
|||
|
|
@ -1248,7 +1248,7 @@ class Router:
|
|||
selector = LeastBusyLoggingHandler(router_cache=self.cache)
|
||||
if register_callbacks:
|
||||
if isinstance(litellm.input_callback, list):
|
||||
litellm.input_callback.append(selector)
|
||||
litellm.logging_callback_manager.add_litellm_input_callback(selector)
|
||||
else:
|
||||
litellm.input_callback = [selector]
|
||||
case RoutingStrategy.USAGE_BASED_ROUTING.value:
|
||||
|
|
@ -4214,10 +4214,12 @@ class Router:
|
|||
}
|
||||
)
|
||||
litellm_logging_object = cast(LiteLLMLogging, litellm_logging_object)
|
||||
prompt_management_deployment: Final = self.get_available_deployment(
|
||||
specific_deployment: Final = kwargs.pop("specific_deployment", None)
|
||||
prompt_management_deployment: Final = await self.async_get_available_deployment(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "prompt"}],
|
||||
specific_deployment=kwargs.pop("specific_deployment", None),
|
||||
messages=cast(list[dict[str, str]], messages), # cast-ok: selection reads messages structurally
|
||||
specific_deployment=specific_deployment,
|
||||
request_kwargs=kwargs,
|
||||
)
|
||||
|
||||
self._update_kwargs_with_deployment(deployment=prompt_management_deployment, kwargs=kwargs)
|
||||
|
|
|
|||
|
|
@ -1,17 +1,103 @@
|
|||
#### 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])
|
||||
_MEMORY_COUNTS: Final = TypeAdapter(tuple[float | None, ...] | 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_counts(values: Sequence[float | None]) -> tuple[int, ...]:
|
||||
return tuple(0 if value is None else int(value) for value in values)
|
||||
|
||||
|
||||
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(
|
||||
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 shared in-flight counts for %s, "
|
||||
"falling back to this worker's own counts: %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
|
||||
|
|
@ -20,195 +106,101 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
|
||||
def __init__(self, router_cache: DualCache):
|
||||
self.router_cache = router_cache
|
||||
self.router_cache_id = str(id(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)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
shared: Final = _as_counts(redis_cache.batch_get_counts(list(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: 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)
|
||||
redis_cache: Final = self.router_cache.redis_cache
|
||||
if redis_cache is not None:
|
||||
try:
|
||||
shared: Final = _as_counts(await redis_cache.async_batch_get_counts(list(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)
|
||||
if ref is None:
|
||||
return
|
||||
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:
|
||||
return
|
||||
redis_cache.increment_with_floor(key, delta, IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
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:
|
||||
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:
|
||||
return
|
||||
await redis_cache.async_increment_with_floor(key, delta, IN_FLIGHT_COUNT_TTL_SECONDS)
|
||||
except Exception as e:
|
||||
_warn_unwritable(key, e)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"ANN001": {
|
||||
"limit": 2956
|
||||
"limit": 2918
|
||||
},
|
||||
"ANN002": {
|
||||
"limit": 71
|
||||
|
|
@ -9,10 +9,10 @@
|
|||
"limit": 806
|
||||
},
|
||||
"ANN201": {
|
||||
"limit": 1979
|
||||
"limit": 1965
|
||||
},
|
||||
"ANN202": {
|
||||
"limit": 831
|
||||
"limit": 829
|
||||
},
|
||||
"ANN204": {
|
||||
"limit": 683
|
||||
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 3
|
||||
},
|
||||
"BLE001": {
|
||||
"limit": 2916
|
||||
"limit": 2914
|
||||
},
|
||||
"C401": {
|
||||
"limit": 8
|
||||
|
|
@ -189,7 +189,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"S110": {
|
||||
"limit": 217
|
||||
"limit": 207
|
||||
},
|
||||
"S112": {
|
||||
"limit": 22
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
80
tests/local_testing/test_redis_increment_with_floor.py
Normal file
80
tests/local_testing/test_redis_increment_with_floor.py
Normal file
|
|
@ -0,0 +1,80 @@
|
|||
"""Least-busy routing keeps its in-flight counters in Redis, and the clamp at zero plus the
|
||||
create-once TTL both live inside a Lua script. Nothing but a real Redis runs that script, so
|
||||
these are the only tests that fail when the script itself is wrong."""
|
||||
|
||||
import os
|
||||
import uuid
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from dotenv import load_dotenv
|
||||
|
||||
load_dotenv()
|
||||
|
||||
from litellm.caching.redis_cache import RedisCache
|
||||
|
||||
TTL: Final = 600
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def counter():
|
||||
cache: Final = RedisCache(host=os.getenv("REDIS_HOST"), port=os.getenv("REDIS_PORT"))
|
||||
key: Final = f"lit7039-{uuid.uuid4()}"
|
||||
yield cache, key, cache.check_and_fix_namespace(key=key)
|
||||
cache.delete_cache(key)
|
||||
|
||||
|
||||
def test_a_counter_adds_every_increment_and_reads_back_what_it_holds(counter):
|
||||
cache, key, _ = counter
|
||||
|
||||
assert cache.increment_with_floor(key, 3, TTL) == 3
|
||||
assert cache.increment_with_floor(key, 2, TTL) == 5
|
||||
assert cache.batch_get_counts([key]) == (5,)
|
||||
|
||||
|
||||
def test_a_decrement_past_zero_leaves_the_counter_at_zero(counter):
|
||||
"""A worker whose counter expired mid-request decrements a key that is no longer there.
|
||||
Without the clamp that deployment reads negative, and least-busy pins every later request
|
||||
on it until the count climbs back to zero."""
|
||||
cache, key, _ = counter
|
||||
|
||||
assert cache.increment_with_floor(key, 1, TTL) == 1
|
||||
assert cache.increment_with_floor(key, -5, TTL) == 0
|
||||
assert cache.batch_get_counts([key]) == (0,)
|
||||
|
||||
|
||||
def test_traffic_never_pushes_a_counters_expiry_back_out(counter):
|
||||
"""The TTL is what releases a count whose worker died mid-request. Rewriting it on every
|
||||
touch would keep that stuck count alive for as long as the group takes traffic."""
|
||||
cache, key, namespaced_key = counter
|
||||
|
||||
cache.increment_with_floor(key, 1, TTL)
|
||||
assert cache.redis_client.ttl(namespaced_key) > TTL - 60
|
||||
|
||||
cache.redis_client.expire(namespaced_key, 30)
|
||||
cache.increment_with_floor(key, 1, TTL)
|
||||
|
||||
assert cache.redis_client.ttl(namespaced_key) <= 30
|
||||
|
||||
|
||||
def test_clamping_to_zero_keeps_the_expiry_it_already_had(counter):
|
||||
cache, key, namespaced_key = counter
|
||||
|
||||
cache.increment_with_floor(key, 1, TTL)
|
||||
cache.redis_client.expire(namespaced_key, 30)
|
||||
|
||||
assert cache.increment_with_floor(key, -5, TTL) == 0
|
||||
assert cache.redis_client.ttl(namespaced_key) <= 30
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_counter_behaves_the_same_way(counter):
|
||||
cache, key, namespaced_key = counter
|
||||
|
||||
assert await cache.async_increment_with_floor(key, 2, TTL) == 2
|
||||
assert await cache.async_batch_get_counts([key]) == (2,)
|
||||
|
||||
cache.redis_client.expire(namespaced_key, 30)
|
||||
|
||||
assert await cache.async_increment_with_floor(key, -9, TTL) == 0
|
||||
assert cache.redis_client.ttl(namespaced_key) <= 30
|
||||
|
|
@ -525,6 +525,50 @@ def test_circuit_breaker_open_keeps_sync_batch_get_cache_as_a_miss(sync_batch_re
|
|||
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit6729"]) == {}
|
||||
|
||||
|
||||
def test_batch_get_counts_raises_where_batch_get_cache_reports_a_miss(sync_batch_redis_cache):
|
||||
"""A caller that must fall back when Redis is unreachable needs the failure, not zeros.
|
||||
|
||||
The batch read answers a dead Redis with an empty dict, which a counting caller cannot tell
|
||||
apart from "every counter is unset". Least-busy routing read that as an idle deployment and
|
||||
kept sending traffic to it instead of falling back to this worker's own in-flight counts.
|
||||
"""
|
||||
assert sync_batch_redis_cache.batch_get_cache(key_list=["lit7039"]) == {}
|
||||
|
||||
with pytest.raises(OSError, match="redis unavailable"):
|
||||
sync_batch_redis_cache.batch_get_counts(["lit7039"])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_batch_get_counts_raises_where_async_batch_get_cache_reports_a_miss(redis_no_ping: None):
|
||||
"""Async twin: the async batch read hides the same failure behind an empty dict."""
|
||||
failing_client = AsyncMock()
|
||||
failing_client.mget.side_effect = OSError("redis unavailable")
|
||||
with patch( # test-quality-ok: RedisCache.__init__ builds its client eagerly, with no injection point
|
||||
"litellm._redis.get_redis_client", return_value=MagicMock()
|
||||
):
|
||||
cache = RedisCache(host="127.0.0.1", port=6379)
|
||||
|
||||
with patch.object(cache, "init_async_client", return_value=failing_client):
|
||||
assert await cache.async_batch_get_cache(key_list=["lit7039"]) == {}
|
||||
|
||||
with pytest.raises(OSError, match="redis unavailable"):
|
||||
await cache.async_batch_get_counts(["lit7039"])
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stored", [b"3", "3"])
|
||||
def test_batch_get_counts_reads_counters_in_order_and_keeps_unset_keys_apart(stored, redis_no_ping: None):
|
||||
"""Counters come back positionally, so an unset key has to stay a hole rather than shift the
|
||||
rest of the row onto the wrong deployments, and a count has to survive whether the client
|
||||
hands it back as bytes or as text."""
|
||||
with patch( # test-quality-ok: RedisCache.__init__ builds its client eagerly, with no injection point
|
||||
"litellm._redis.get_redis_client", return_value=MagicMock()
|
||||
):
|
||||
cache = RedisCache(host="127.0.0.1", port=6379)
|
||||
cache.redis_client.mget.return_value = [stored, None, b"0"]
|
||||
|
||||
assert cache.batch_get_counts(["dep-a", "dep-b", "dep-c"]) == (3, None, 0)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sync_batch_cache_with_service_logger(redis_no_ping: None) -> Iterator[tuple[RedisCache, ServiceLogging]]:
|
||||
service_logger = ServiceLogging(mock_testing=True)
|
||||
|
|
|
|||
187
tests/test_litellm/router_strategy/test_least_busy.py
Normal file
187
tests/test_litellm/router_strategy/test_least_busy.py
Normal file
|
|
@ -0,0 +1,187 @@
|
|||
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:
|
||||
"""Mirrors what Redis gives the handler: increments clamped at zero, a TTL set once when
|
||||
the key is created, and ordered reads that raise rather than invent a value."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.counts: dict[str, int] = {}
|
||||
self.ttls: dict[str, int] = {}
|
||||
|
||||
def count(self, key: str) -> int | None:
|
||||
return self.counts.get(key)
|
||||
|
||||
def expire(self, key: str) -> None:
|
||||
self.counts.pop(key, None)
|
||||
self.ttls.pop(key, None)
|
||||
|
||||
def increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
incremented: Final = max(0, self.counts.get(key, 0) + value)
|
||||
self.counts[key] = incremented
|
||||
self.ttls.setdefault(key, ttl)
|
||||
return incremented
|
||||
|
||||
async def async_increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
return self.increment_with_floor(key, value, ttl)
|
||||
|
||||
def batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
return tuple(self.counts.get(key) for key in key_list)
|
||||
|
||||
async def async_batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
return self.batch_get_counts(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_the_handler_never_pushes_a_counters_ttl_forward() -> None:
|
||||
"""A worker that dies mid-request leaves a +1 nobody will ever decrement. Redis expires that
|
||||
stuck count an hour after the key was created, which only works while nothing writes the TTL
|
||||
again: a handler that refreshed it on every touch would keep the count alive for as long as
|
||||
the group takes traffic, and the deployment would read busier than it is forever."""
|
||||
shared: Final = SharedRedisCounters()
|
||||
worker: Final = _worker(shared)
|
||||
key: Final = f"{GROUP}_request_count:dep-a"
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs=_call_kwargs("dep-a"))
|
||||
|
||||
assert shared.ttls == {key: IN_FLIGHT_COUNT_TTL_SECONDS}
|
||||
|
||||
shared.ttls[key] = 5
|
||||
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(key) == 1
|
||||
assert shared.ttls == {key: 5}
|
||||
|
||||
|
||||
@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 increment_with_floor(self, key: str, value: int, ttl: int) -> int:
|
||||
raise ConnectionError("redis is down")
|
||||
|
||||
def batch_get_counts(self, key_list: list[str]) -> tuple[int | None, ...]:
|
||||
raise ConnectionError("redis is down")
|
||||
|
||||
|
||||
@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.expire(f"{GROUP}_request_count:dep-a")
|
||||
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 shared.count(f"{GROUP}_request_count:dep-a") == 1
|
||||
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)
|
||||
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs={"litellm_params": {"metadata": None}})
|
||||
worker.log_pre_api_call(model="m", messages=[], kwargs={})
|
||||
|
||||
assert shared.counts == {}
|
||||
|
|
@ -12,6 +12,7 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm import Router
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.router import RoutingGroup, RoutingStrategy
|
||||
|
||||
|
||||
|
|
@ -435,6 +436,81 @@ def test_update_settings_unregisters_group_selectors_when_groups_removed(monkeyp
|
|||
assert router._group_selectors == {}
|
||||
|
||||
|
||||
def test_two_least_busy_groups_count_a_request_once(monkeypatch):
|
||||
"""
|
||||
Least-busy counts a request up from the pre-call hooks on `litellm.input_callback` and
|
||||
back down from the success hooks on `litellm.callbacks`. The success list drops a second
|
||||
selector of the same class, so a pre-call list that kept both counted every request twice
|
||||
and released it once, and the deployment's in-flight count climbed until it looked pinned.
|
||||
"""
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "input_callback", [])
|
||||
|
||||
router = _build_router(
|
||||
routing_strategy="least-busy",
|
||||
routing_groups=[
|
||||
{
|
||||
"group_name": "fast",
|
||||
"models": ["filtered-model"],
|
||||
"routing_strategy": "least-busy",
|
||||
}
|
||||
],
|
||||
)
|
||||
kwargs = {
|
||||
"litellm_params": {
|
||||
"metadata": {"model_group": "filtered-model"},
|
||||
"model_info": {"id": "deploy-1"},
|
||||
}
|
||||
}
|
||||
|
||||
for callback in litellm.input_callback:
|
||||
if isinstance(callback, CustomLogger):
|
||||
callback.log_pre_api_call(model="filtered-model", messages=[], kwargs=kwargs)
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
callback.log_success_event(kwargs, None, None, None)
|
||||
|
||||
assert router.cache.get_cache("filtered-model_request_count:deploy-1") == 0
|
||||
|
||||
|
||||
def test_two_routers_in_one_process_each_count_their_own_requests(monkeypatch):
|
||||
"""
|
||||
Least-busy hangs its counting off litellm's global callback lists, and those lists keep one
|
||||
logger per class unless the instances differ in a plain attribute. Two routers in one process
|
||||
(a second Router, or a per-request `user_config` one) therefore have to register separately:
|
||||
a second router whose selector is dropped counts nothing, reads zero for every deployment,
|
||||
and sends every request to whichever one is listed first.
|
||||
"""
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
monkeypatch.setattr(litellm, "input_callback", [])
|
||||
|
||||
first = _build_router(routing_strategy="least-busy")
|
||||
second = _build_router(routing_strategy="least-busy")
|
||||
kwargs = {
|
||||
"litellm_params": {
|
||||
"metadata": {"model_group": "filtered-model"},
|
||||
"model_info": {"id": "deploy-1"},
|
||||
}
|
||||
}
|
||||
|
||||
for callback in litellm.input_callback:
|
||||
if isinstance(callback, CustomLogger):
|
||||
callback.log_pre_api_call(model="filtered-model", messages=[], kwargs=kwargs)
|
||||
|
||||
assert second.cache.get_cache("filtered-model_request_count:deploy-1") == 1
|
||||
assert (
|
||||
second.get_available_deployment(model="filtered-model", messages=[])["model_info"]["id"]
|
||||
== "deploy-2"
|
||||
)
|
||||
|
||||
for callback in litellm.callbacks:
|
||||
if isinstance(callback, CustomLogger):
|
||||
callback.log_success_event(kwargs, None, None, None)
|
||||
|
||||
assert first.cache.get_cache("filtered-model_request_count:deploy-1") == 0
|
||||
assert second.cache.get_cache("filtered-model_request_count:deploy-1") == 0
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Direct helper coverage
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -56,6 +56,18 @@ class BlockEverything:
|
|||
return context
|
||||
|
||||
|
||||
class MessageRecorder:
|
||||
"""Records what each plugin pass was handed, then blocks so the request stops there."""
|
||||
|
||||
def __init__(self):
|
||||
self.seen = []
|
||||
|
||||
async def run(self, context: RoutingContext) -> RoutingContext:
|
||||
self.seen.append(list(context.raw_messages))
|
||||
context.candidate_models = []
|
||||
return context
|
||||
|
||||
|
||||
def _smart_router_model_list():
|
||||
return [
|
||||
{
|
||||
|
|
@ -164,6 +176,71 @@ async def test_async_completion_with_unsupported_strategy_rejects_configured_plu
|
|||
await router.acompletion(model="smart-router", messages=[{"role": "user", "content": "hi"}])
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_management_model_still_runs_the_plugin_pipeline():
|
||||
"""
|
||||
A prompt-management model routes through its own factory, which picked the deployment
|
||||
on the synchronous path. Plugins never run there, so the guard turned every such request
|
||||
into an error message about the caller's own API choice, on an async call the caller made
|
||||
correctly. It also read the in-flight counts with a blocking call inside the event loop.
|
||||
"""
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "cached-claude",
|
||||
"litellm_params": {
|
||||
"model": "anthropic_cache_control_hook/claude-sonnet-5",
|
||||
"prompt_id": "cache-points",
|
||||
},
|
||||
}
|
||||
],
|
||||
routing_strategy="least-busy",
|
||||
plugins=[BlockEverything()],
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="No deployments left after routing-plugin filtering"):
|
||||
await router.acompletion(
|
||||
model="cached-claude",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
litellm_call_id="lit-7039",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_prompt_management_plugins_see_the_callers_own_messages():
|
||||
"""
|
||||
The prompt-management factory picks its deployment with a placeholder message, which was
|
||||
harmless while that pick ran on the synchronous path (plugins never ran there at all). Now
|
||||
that the pick runs the plugin pipeline, a plugin that classifies request content would score
|
||||
the placeholder instead of the conversation, and the narrowing it produces decides which
|
||||
deployments the real call is allowed to use.
|
||||
"""
|
||||
recorder = MessageRecorder()
|
||||
router = Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "cached-claude",
|
||||
"litellm_params": {
|
||||
"model": "anthropic_cache_control_hook/claude-sonnet-5",
|
||||
"prompt_id": "cache-points",
|
||||
},
|
||||
}
|
||||
],
|
||||
routing_strategy="least-busy",
|
||||
plugins=[recorder],
|
||||
)
|
||||
messages = [{"role": "user", "content": "wire me $40,000 to account 12345"}]
|
||||
|
||||
with pytest.raises(ValueError, match="No deployments left after routing-plugin filtering"):
|
||||
await router.acompletion(
|
||||
model="cached-claude",
|
||||
messages=messages,
|
||||
litellm_call_id="lit-7039",
|
||||
)
|
||||
|
||||
assert recorder.seen == [messages]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_router_without_plugins_is_unaffected():
|
||||
"""Regression guard: a Router with no `plugins` configured behaves exactly as before."""
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22180
|
||||
"limit": 22174
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26729
|
||||
"limit": 26715
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 261
|
||||
|
|
@ -27,10 +27,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16426
|
||||
"limit": 16398
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5506
|
||||
"limit": 5504
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4486
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue