fix(router): carry per-request routing reads on context variables instead of public method kwargs (#43814)

Co-authored-by: yassin <yassin@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-30 07:34:06 +00:00 • committed by GitHub
parent d79600987e
commit 314ff111e5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 193 additions and 37 deletions

View file

@ -1813,7 +1813,6 @@ class Router:
messages: list[dict[str, str]] | None,
input: str | list | None,
request_kwargs: dict | None,
prefetched_usage: PrefetchedUsage | None = None,
) -> Any | None:
"""
Asks the strategy selector for a deployment. Caller handles
@ -1839,14 +1838,6 @@ class Router:
messages=messages,
input=input,
)
case "usage-based-routing-v2" if isinstance(selector, LowestTPMLoggingHandler_v2):
return await selector.async_get_available_deployments(
model_group=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
prefetched_usage=prefetched_usage,
)
case "usage-based-routing-v2" | "cost-based-routing":
return await selector.async_get_available_deployments(
model_group=model,
@ -12958,7 +12949,6 @@ class Router:
specific_deployment: bool | None = False,
parent_otel_span: Span | None = None,
health_check_probe: bool = False,
routing_read_batch: RoutingReadBatch | None = None,
) -> list[dict] | dict:
"""
Get the healthy deployments for a model.
@ -13011,6 +13001,7 @@ class Router:
health_check_probe=health_check_probe,
)
routing_read_batch: Final = RoutingReadBatch.active()
cooldown_deployments: Final = (
await _async_get_cooldown_deployments(litellm_router_instance=self, parent_otel_span=parent_otel_span)
if routing_read_batch is None
@ -13298,15 +13289,15 @@ class Router:
strategy, strategy_selector = self._get_routing_context(model, request_kwargs)
routing_read_batch: Final = RoutingReadBatch.for_strategy(strategy, strategy_selector)
healthy_deployments: Final = await self.async_get_healthy_deployments(
model=model,
request_kwargs=request_kwargs,
messages=messages,
input=input,
specific_deployment=specific_deployment,
parent_otel_span=parent_otel_span,
routing_read_batch=routing_read_batch,
)
with RoutingReadBatch.scoped(routing_read_batch):
healthy_deployments: Final = await self.async_get_healthy_deployments(
model=model,
request_kwargs=request_kwargs,
messages=messages,
input=input,
specific_deployment=specific_deployment,
parent_otel_span=parent_otel_span,
)
if isinstance(healthy_deployments, dict):
await self._async_override_selector_pre_call_check(
strategy, strategy_selector, healthy_deployments, parent_otel_span
@ -13328,16 +13319,18 @@ class Router:
model=model,
request_kwargs=request_kwargs,
)
deployment: Final = await self._select_deployment_async(
strategy=strategy,
selector=strategy_selector,
model=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
request_kwargs=request_kwargs,
prefetched_usage=routing_read_batch.prefetched_usage if routing_read_batch is not None else None,
)
with PrefetchedUsage.scoped(
routing_read_batch.prefetched_usage if routing_read_batch is not None else None
):
deployment: Final = await self._select_deployment_async(
strategy=strategy,
selector=strategy_selector,
model=model,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
request_kwargs=request_kwargs,
)
if deployment is None:
exception: Final = await async_raise_no_deployment_exception(
litellm_router_instance=self,

View file

@ -1,7 +1,9 @@
#### What this does ####
# identifies lowest tpm deployment
import random
from collections.abc import Mapping, Sequence
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, Final
@ -32,6 +34,9 @@ class RoutingArgs(LiteLLMPydanticObjectBase):
ttl: int = 1 * 60 # 1min (RPM/TPM expire key)
_active_prefetched_usage: Final[ContextVar["PrefetchedUsage | None"]] = ContextVar("prefetched_usage", default=None)
@dataclass(frozen=True)
class PrefetchedUsage:
"""
@ -51,6 +56,19 @@ class PrefetchedUsage:
return None
return [self.values.get(key) for key in keys]
@staticmethod
@contextmanager
def scoped(usage: "PrefetchedUsage | None") -> Iterator[None]:
token: Final = _active_prefetched_usage.set(usage)
try:
yield
finally:
_active_prefetched_usage.reset(token)
@staticmethod
def active() -> "PrefetchedUsage | None":
return _active_prefetched_usage.get()
class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
"""
@ -455,13 +473,13 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
healthy_deployments: list,
messages: list[dict[str, str]] | None = None,
input: str | list | None = None,
prefetched_usage: PrefetchedUsage | None = None,
):
"""
Async implementation of get deployments.
Reduces time to retrieve the tpm/rpm values from cache. `prefetched_usage` skips the cache
read when it already holds this request's counters (see `RoutingReadBatch`).
Reduces time to retrieve the tpm/rpm values from cache. A `PrefetchedUsage` scoped
to this request skips the cache read when it already holds its counters (see
`RoutingReadBatch`).
"""
# get list of potential deployments
verbose_router_logger.debug(
@ -473,6 +491,7 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger):
tpm_keys, rpm_keys = self.usage_counter_keys(healthy_deployments)
combined_tpm_rpm_keys: Final = tpm_keys + rpm_keys
prefetched_usage: Final = PrefetchedUsage.active()
if prefetched_usage is not None and prefetched_usage.covers(combined_tpm_rpm_keys):
combined_tpm_rpm_values = prefetched_usage.values_for(combined_tpm_rpm_keys)
else:

View file

@ -10,7 +10,9 @@ the usage slice to the strategy, so selection does not read again.
import asyncio
import itertools
from collections.abc import Mapping, Sequence
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from contextvars import ContextVar
from dataclasses import dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final
@ -144,11 +146,29 @@ class RoutingPrefetch:
return None
_active_routing_read_batch: Final[ContextVar["RoutingReadBatch | None"]] = ContextVar(
"routing_read_batch", default=None
)
class RoutingReadBatch:
def __init__(self, usage_selector: LowestTPMLoggingHandler_v2 | None) -> None:
self.usage_selector: Final = usage_selector
self.prefetched_usage: PrefetchedUsage | None = None
@staticmethod
@contextmanager
def scoped(batch: "RoutingReadBatch | None") -> Iterator[None]:
token: Final = _active_routing_read_batch.set(batch)
try:
yield
finally:
_active_routing_read_batch.reset(token)
@staticmethod
def active() -> "RoutingReadBatch | None":
return _active_routing_read_batch.get()
@staticmethod
def for_strategy(strategy: str | None, selector: object) -> "RoutingReadBatch | None":
"""Usage-based routing reads its counters with the cooldown state; every other strategy reads only the

View file

@ -72,11 +72,48 @@ async def test_v2_async_selection_uses_prefetched_counters_only_when_they_cover_
keys = tpm_keys + rpm_keys
covering = PrefetchedUsage(keys=frozenset(keys), values=dict(zip(keys, [10, 100, None, None])))
chosen = await strategy.async_get_available_deployments(model_group="g", healthy_deployments=deployments, prefetched_usage=covering)
with PrefetchedUsage.scoped(covering):
chosen: Final = await strategy.async_get_available_deployments(model_group="g", healthy_deployments=deployments)
assert chosen["model_info"]["id"] == "a", "the prefetched counters say a is the lowest"
router_cache.async_batch_get_cache.assert_not_awaited()
stale = PrefetchedUsage(keys=frozenset(keys[:1]), values={keys[0]: 10})
chosen = await strategy.async_get_available_deployments(model_group="g", healthy_deployments=deployments, prefetched_usage=stale)
assert chosen["model_info"]["id"] == "b", "counters that do not cover this minute's keys are read again"
with PrefetchedUsage.scoped(stale):
chosen_stale: Final = await strategy.async_get_available_deployments(
model_group="g", healthy_deployments=deployments
)
assert chosen_stale["model_info"]["id"] == "b", "counters that do not cover this minute's keys are read again"
router_cache.async_batch_get_cache.assert_awaited_once_with(keys=keys)
@pytest.mark.asyncio
async def test_v2_subclass_overriding_async_get_available_deployments_with_the_old_signature_still_routes() -> None:
class OldSignatureV2(LowestTPMLoggingHandler_v2):
async def async_get_available_deployments(
self,
model_group: str,
healthy_deployments: list,
messages: list[dict[str, str]] | None = None,
input: str | list | None = None,
):
return await super().async_get_available_deployments(
model_group=model_group,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
)
router: Final = Router(
model_list=[_deployment(HIGH_USAGE_DEPLOYMENT_ID), _deployment(LOW_USAGE_DEPLOYMENT_ID)],
routing_strategy="usage-based-routing-v2",
)
router.lowesttpm_logger_v2 = OldSignatureV2(router_cache=router.cache, routing_args={})
response: Final = await router.acompletion(
model=MODEL_GROUP, messages=[{"role": "user", "content": "x"}]
)
assert response.choices[0].message.content in {
f"from {HIGH_USAGE_DEPLOYMENT_ID}",
f"from {LOW_USAGE_DEPLOYMENT_ID}",
}

View file

@ -6,6 +6,7 @@ Before `RoutingReadBatch`, `async_get_available_deployment` issued one MGET for
"""
import time
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -13,6 +14,7 @@ import pytest
import litellm
from litellm import Router
from litellm.caching.redis_cache import RedisCache
from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2
_MODEL_GROUP = "claude"
_MESSAGES = [{"role": "user", "content": "ping"}]
@ -79,6 +81,46 @@ async def test_usage_based_routing_reads_cooldowns_and_counters_in_one_redis_rou
], "cooldown state and usage counters must arrive in one MGET"
@pytest.mark.asyncio
async def test_usage_based_routing_still_batches_when_the_strategy_is_a_fixed_signature_subclass():
class OldSignatureV2(LowestTPMLoggingHandler_v2):
async def async_get_available_deployments(
self,
model_group: str,
healthy_deployments: list,
messages: list[dict[str, str]] | None = None,
input: str | list | None = None,
):
return await super().async_get_available_deployments(
model_group=model_group,
healthy_deployments=healthy_deployments,
messages=messages,
input=input,
)
redis: Final = _redis_answering({})
router: Final = _router(redis, "usage-based-routing-v2")
router.lowesttpm_logger_v2 = OldSignatureV2(router_cache=router.cache)
router.cache.async_batch_get_cache = AsyncMock(wraps=router.cache.async_batch_get_cache)
deployment: Final = await router.async_get_available_deployment(
model=_MODEL_GROUP, request_kwargs={}, messages=_MESSAGES
)
assert deployment["model_info"]["id"] in {"dep-a", "dep-b"}
assert _redis_key_families(redis) == [
[
"dep-a:anthropic/claude-x:rpm",
"dep-a:anthropic/claude-x:tpm",
"dep-b:anthropic/claude-x:rpm",
"dep-b:anthropic/claude-x:tpm",
"deployment:dep-a:cooldown",
"deployment:dep-b:cooldown",
]
], "the subclassed strategy must still get the batched read, not a second MGET"
router.cache.async_batch_get_cache.assert_not_awaited()
@pytest.mark.asyncio
async def test_simple_shuffle_still_reads_only_cooldowns():
redis = _redis_answering({})

View file

@ -48,6 +48,7 @@ from litellm.router import (
_is_retriable_anthropic_status,
_responses_stream_holds_event,
_without_line_breaks,
Span,
)
from litellm.router_strategy import simple_shuffle
from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit
@ -18837,3 +18838,47 @@ def test_a_failed_routing_read_prefetch_logs_the_request_model_without_its_line_
assert messages == [
"routing read prefetch not armed for gpt-4ERROR forged entry: no deployments for gpt-4ERROR forged entry"
]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"routing_strategy",
["simple-shuffle", "usage-based-routing-v2", "least-busy", "latency-based-routing"],
)
async def test_router_subclass_overriding_async_get_healthy_deployments_with_the_old_signature_still_routes(
routing_strategy: str,
) -> None:
class OldSignatureRouter(litellm.Router):
async def async_get_healthy_deployments(
self,
model: str,
request_kwargs: dict,
messages: list[dict[str, str]] | None = None,
input: str | list | None = None,
specific_deployment: bool | None = False,
parent_otel_span: Span | None = None,
health_check_probe: bool = False,
):
return await super().async_get_healthy_deployments(
model=model,
request_kwargs=request_kwargs,
messages=messages,
input=input,
specific_deployment=specific_deployment,
parent_otel_span=parent_otel_span,
health_check_probe=health_check_probe,
)
router: Final = OldSignatureRouter(
model_list=[
{
"model_name": "m",
"litellm_params": {"model": "openai/gpt-4o", "api_key": "x", "mock_response": "hi"},
}
],
routing_strategy=routing_strategy,
)
response: Final = await router.acompletion(model="m", messages=[{"role": "user", "content": "x"}])
assert response.choices[0].message.content == "hi"