mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
parent
d79600987e
commit
314ff111e5
6 changed files with 193 additions and 37 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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({})
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue