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:
mateo-berri 2026-09-06 01:02:53 -07:00
parent ad2afe5e65
commit 63ea197433
7 changed files with 111 additions and 2 deletions

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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():
"""