diff --git a/litellm/router.py b/litellm/router.py index 8aaf58d5a3e..115faad000c 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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, diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 6e21d5d1f1f..6c9acbdfecb 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -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: diff --git a/litellm/router_utils/routing_read_batch.py b/litellm/router_utils/routing_read_batch.py index e465f06640e..17b7f537730 100644 --- a/litellm/router_utils/routing_read_batch.py +++ b/litellm/router_utils/routing_read_batch.py @@ -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 diff --git a/tests/unit/router_strategy/test_lowest_tpm_rpm.py b/tests/unit/router_strategy/test_lowest_tpm_rpm.py index 0fa11cda20c..625f648bec4 100644 --- a/tests/unit/router_strategy/test_lowest_tpm_rpm.py +++ b/tests/unit/router_strategy/test_lowest_tpm_rpm.py @@ -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}", + } diff --git a/tests/unit/router_utils/test_routing_read_batch.py b/tests/unit/router_utils/test_routing_read_batch.py index 73be5fd4a3a..a74e3e24f99 100644 --- a/tests/unit/router_utils/test_routing_read_batch.py +++ b/tests/unit/router_utils/test_routing_read_batch.py @@ -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({}) diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index e4e65f8904c..96dddf15869 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -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"