mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(router): keep batch retrieves out of routing strategy state
Stamping model_group on a batch retrieve routed the whole batch job's token usage into the per-model-group counters that usage-based, latency-based, cost-based and least-busy routing read, so polling a finished batch could exhaust a group's TPM or RPM window and lock live chat traffic out with RouterRateLimitError. Polling also drove the least-busy in-flight counts negative once per poll per deployment, which pinned chat to whichever deployment had been polled most. The strategy callbacks now skip batch retrieve call types, so a retrieve still lands in spend logs under its model group while the numbers that pick a deployment for the next chat request stay driven by live traffic only.
This commit is contained in:
parent
ad2afe5e65
commit
63ea197433
7 changed files with 111 additions and 2 deletions
|
|
@ -6177,7 +6177,8 @@ class Router:
|
|||
function_name="aretrieve_batch",
|
||||
)
|
||||
model_group: Final = requested_model_group or model_name["model_name"]
|
||||
new_kwargs[metadata_variable_name].setdefault("model_group", model_group)
|
||||
if not new_kwargs[metadata_variable_name].get("model_group"):
|
||||
new_kwargs[metadata_variable_name]["model_group"] = model_group
|
||||
new_kwargs.pop("custom_llm_provider", None)
|
||||
data.pop("custom_llm_provider", None)
|
||||
return await litellm.aretrieve_batch(
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ from typing import Final
|
|||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
|
||||
|
||||
|
||||
class LeastBusyLoggingHandler(CustomLogger):
|
||||
|
|
@ -27,6 +28,8 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
|
||||
Caching based on model group.
|
||||
"""
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
|
|
@ -48,6 +51,8 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
pass
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
|
|
@ -76,6 +81,8 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
pass
|
||||
|
||||
def log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
|
|
@ -103,6 +110,8 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
pass
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
|
|
@ -131,6 +140,8 @@ class LeastBusyLoggingHandler(CustomLogger):
|
|||
pass
|
||||
|
||||
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
if kwargs["litellm_params"].get("metadata") is None:
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from litellm import ModelResponse, token_counter, verbose_logger
|
|||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
|
||||
|
||||
|
||||
class LowestCostLoggingHandler(CustomLogger):
|
||||
|
|
@ -19,6 +20,8 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
self.router_cache = router_cache
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
"""
|
||||
Update usage on success
|
||||
|
|
@ -96,6 +99,8 @@ class LowestCostLoggingHandler(CustomLogger):
|
|||
)
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
"""
|
||||
Update cost usage on success
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ from litellm import ModelResponse, token_counter, verbose_logger
|
|||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.core_helpers import _get_parent_otel_span_from_kwargs, safe_divide_seconds
|
||||
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
|
||||
from litellm.types.utils import LiteLLMPydanticObjectBase
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -35,6 +36,8 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
self.routing_args = RoutingArgs(**routing_args)
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
"""
|
||||
Update latency usage on success
|
||||
|
|
@ -167,6 +170,8 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
"""
|
||||
Check if Timeout Error, if timeout set deployment latency -> 100
|
||||
"""
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
metadata_field: Final = self._select_metadata_field(kwargs)
|
||||
_exception: Final = kwargs.get("exception", None)
|
||||
|
|
@ -221,6 +226,8 @@ class LowestLatencyLoggingHandler(CustomLogger):
|
|||
)
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
"""
|
||||
Update latency usage on success
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ from litellm import token_counter
|
|||
from litellm._logging import verbose_router_logger
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.router_utils.batch_utils import is_batch_retrieve_call_type
|
||||
from litellm.types.utils import LiteLLMPydanticObjectBase
|
||||
from litellm.utils import print_verbose
|
||||
|
||||
|
|
@ -27,6 +28,8 @@ class LowestTPMLoggingHandler(CustomLogger):
|
|||
self.routing_args = RoutingArgs(**routing_args)
|
||||
|
||||
def log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
"""
|
||||
Update TPM/RPM usage on success
|
||||
|
|
@ -79,6 +82,8 @@ class LowestTPMLoggingHandler(CustomLogger):
|
|||
verbose_router_logger.debug(traceback.format_exc())
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
if is_batch_retrieve_call_type(kwargs.get("call_type")):
|
||||
return
|
||||
try:
|
||||
"""
|
||||
Update TPM/RPM usage on success
|
||||
|
|
|
|||
|
|
@ -185,6 +185,7 @@ def is_batch_retrieve_call_type(call_type: object) -> bool:
|
|||
"""
|
||||
A batch retrieve reports the whole job's token usage, which the provider spent
|
||||
asynchronously over the life of the batch, and reports it again on every poll of the
|
||||
finished batch. Per-minute usage counters must not be fed from it.
|
||||
finished batch. The counters that measure live traffic, per-minute rate limits and the
|
||||
routing strategies' own state, must not be fed from it.
|
||||
"""
|
||||
return isinstance(call_type, str) and call_type in BATCH_RETRIEVE_CALL_TYPES
|
||||
|
|
|
|||
|
|
@ -1166,6 +1166,85 @@ async def test_arouter_aretrieve_batch_does_not_consume_deployment_rate_limits(m
|
|||
assert usage_keys == []
|
||||
|
||||
|
||||
_ROUTING_STRATEGY_CACHE_MARKERS = ("_map", "_request_count", ":tpm:", ":rpm:")
|
||||
|
||||
|
||||
async def _router_strategy_keys(router, timeout: float = 2.0) -> list[str]:
|
||||
loop = asyncio.get_event_loop()
|
||||
deadline = loop.time() + timeout
|
||||
while loop.time() < deadline:
|
||||
keys = sorted(
|
||||
key
|
||||
for key in router.cache.in_memory_cache.cache_dict
|
||||
if any(marker in key for marker in _ROUTING_STRATEGY_CACHE_MARKERS)
|
||||
)
|
||||
if keys:
|
||||
return keys
|
||||
await asyncio.sleep(0.05)
|
||||
return []
|
||||
|
||||
|
||||
def _batch_fan_out_router(routing_strategy: str):
|
||||
return litellm.Router(
|
||||
routing_strategy=routing_strategy,
|
||||
model_list=[
|
||||
{
|
||||
"model_name": _BATCH_GROUP,
|
||||
"litellm_params": {
|
||||
"model": _BATCH_DEPLOYMENT_MODEL,
|
||||
"api_base": _BATCH_API_BASE,
|
||||
"api_key": "sk-fake",
|
||||
},
|
||||
"model_info": {"id": "batch-dep"},
|
||||
},
|
||||
{
|
||||
"model_name": _UNRELATED_BATCH_GROUP,
|
||||
"litellm_params": {
|
||||
"model": _BATCH_DEPLOYMENT_MODEL,
|
||||
"api_base": _UNRELATED_BATCH_API_BASE,
|
||||
"api_key": "sk-fake",
|
||||
},
|
||||
"model_info": {"id": "unrelated-dep"},
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"routing_strategy",
|
||||
["usage-based-routing", "latency-based-routing", "cost-based-routing", "least-busy"],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_arouter_aretrieve_batch_does_not_feed_routing_strategies(
|
||||
monkeypatch: pytest.MonkeyPatch, routing_strategy: str
|
||||
):
|
||||
"""
|
||||
Every routing strategy picks a deployment from what recent live traffic did.
|
||||
A batch retrieve reports the whole job on every poll and probes deployments the
|
||||
caller never named, so polling a finished batch must not move the numbers that
|
||||
decide where the next chat request goes.
|
||||
"""
|
||||
import respx
|
||||
|
||||
collector = _BatchPayloadCollector()
|
||||
monkeypatch.setattr(litellm, "callbacks", [collector])
|
||||
monkeypatch.setattr(litellm, "input_callback", [])
|
||||
router = _batch_fan_out_router(routing_strategy)
|
||||
|
||||
with respx.mock(assert_all_called=True) as respx_mock:
|
||||
_mock_batch_provider(respx_mock)
|
||||
respx_mock.get(f"{_UNRELATED_BATCH_API_BASE}/batches/{_BATCH_ID}").mock(
|
||||
return_value=httpx.Response(404, json=_BATCH_NOT_FOUND)
|
||||
)
|
||||
for _ in range(3):
|
||||
response = await router.aretrieve_batch(batch_id=_BATCH_ID)
|
||||
await collector.retrieve_batch_payload()
|
||||
strategy_keys = await _router_strategy_keys(router)
|
||||
|
||||
assert response.id == _BATCH_ID
|
||||
assert strategy_keys == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arouter_aretrieve_file_content():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue