fix(router): keep batch retrieves out of the sync success counter and type the metadata helper

This commit is contained in:
mateo-berri 2026-09-19 02:39:38 -07:00
parent 12bc9ad7b4
commit 2d13ca06d0
3 changed files with 45 additions and 5 deletions

View file

@ -6,7 +6,7 @@
######################################################################
import asyncio
import os
from collections.abc import Mapping
from collections.abc import Mapping, MutableMapping
from types import MappingProxyType
from typing import Any, Final, cast
@ -57,17 +57,17 @@ from litellm.types.llms.openai import LiteLLMBatchCreateRequest
router: Final = APIRouter()
def _litellm_metadata_of(data: dict) -> dict:
def _litellm_metadata_of(data: MutableMapping[str, object]) -> MutableMapping[str, object]:
"""The request's litellm_metadata mapping, created on the request when it carries none.
The success handler reads this mapping, so a flag or a model group set here has to live
inside it rather than beside it.
"""
existing: Final = data.get("litellm_metadata")
if isinstance(existing, dict):
if isinstance(existing, MutableMapping):
return existing
created: Final = {} # mutable-ok: the logging layer copies and extends this mapping, so it cannot be a read-only view
data["litellm_metadata"] = created
created: Final[dict[str, object]] = {} # mutable-ok: the logging layer copies and extends this mapping
data["litellm_metadata"] = created # rebind-ok: the success handler reads the request's own mapping
return created

View file

@ -8097,6 +8097,8 @@ class Router:
- key: str - The key used to increment the cache
- None: if no key is found
"""
if is_batch_retrieve_call_type(kwargs.get("call_type")):
return None
id = None
if kwargs["litellm_params"].get("metadata") is None:
pass

View file

@ -49,6 +49,7 @@ from litellm.router import (
from litellm.router_strategy import simple_shuffle
from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit
from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments
from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute
from litellm.types.llms.openai import ChatCompletionRequest
from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy
@ -1217,6 +1218,43 @@ async def test_arouter_aretrieve_batch_does_not_consume_deployment_rate_limits(m
assert usage_keys == []
@pytest.mark.parametrize(
("call_type", "expected_key", "expected_successes"),
[
("aretrieve_batch", None, 0),
("retrieve_batch", None, 0),
("acompletion", "batch-dep:successes", 1),
],
)
def test_sync_deployment_callback_on_success_skips_batch_retrieves(
call_type: str, expected_key: str | None, expected_successes: int
):
router = litellm.Router(
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"},
}
]
)
key = router.sync_deployment_callback_on_success(
kwargs={
"call_type": call_type,
"litellm_params": {"metadata": {"model_group": _BATCH_GROUP}, "model_info": {"id": "batch-dep"}},
},
completion_response=None,
start_time=datetime.now(),
end_time=datetime.now(),
)
assert key == expected_key
assert (
get_deployment_successes_for_current_minute(litellm_router_instance=router, deployment_id="batch-dep")
== expected_successes
)
_ROUTING_STRATEGY_CACHE_MARKERS = ("_map", "_request_count", ":tpm:", ":rpm:")