diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 9839c9be469..083c4f14e2b 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -5,7 +5,7 @@ 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 +from typing import TYPE_CHECKING, Any, Final, cast import httpx @@ -71,6 +71,13 @@ class PrefetchedUsage: return _active_prefetched_usage.get() +def _declares_tpm_limit(deployment: Mapping[str, object]) -> bool: + """Whether the deployment carries a ``tpm`` limit that selection has to weigh the prompt against.""" + nested: Final = (deployment.get("litellm_params"), deployment.get("model_info")) + sources: Final = (deployment, *(source for source in nested if isinstance(source, Mapping))) + return any(source.get("tpm") is not None for source in sources) + + class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): """ Updated version of TPM/RPM Logging. @@ -421,10 +428,13 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): for idx, key in enumerate(rpm_keys): rpm_dict[rpm_keys[idx].split(":")[0]] = rpm_values[idx] - try: - input_tokens = token_counter(messages=messages, text=input) - except Exception: - input_tokens = 0 + deployments: Final = cast(Sequence[Mapping[str, object]], healthy_deployments) # cast-ok: untyped dicts + input_tokens = 0 + if any(_declares_tpm_limit(deployment) for deployment in deployments): + try: + input_tokens = token_counter(messages=messages, text=input) + except Exception: + input_tokens = 0 verbose_router_logger.debug("input_tokens=%s", input_tokens) # ----------------------- # Find lowest used model diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index 40c6456431d..1abd273c0e3 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -88,7 +88,7 @@ RULES: Final[Rules] = ( RouteRule(Route.MESSAGES, Rollout.RUST_OPT_IN, providers=frozenset({"anthropic"})), RouteRule(Route.MESSAGES, Rollout.PYTHON_ONLY), RouteRule(Route.RESPONSES, Rollout.PYTHON_ONLY), - RouteRule(Route.TOKEN_COUNTER, Rollout.PYTHON_ONLY), + RouteRule(Route.TOKEN_COUNTER, Rollout.RUST_OPT_IN), RouteRule(Route.TOKENIZER, Rollout.PYTHON_ONLY), RouteRule(Route.TRANSCRIPTION, Rollout.RUST_REQUIRED, providers=frozenset({"bedrock"})), SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.GOOGLE_KMS.value})), diff --git a/litellm/rust_bridge/token_counter.py b/litellm/rust_bridge/token_counter.py index 250ad18d44c..f289357ccd7 100644 --- a/litellm/rust_bridge/token_counter.py +++ b/litellm/rust_bridge/token_counter.py @@ -79,7 +79,7 @@ def rust_tokenizer(model: str) -> RustTokenizer | None: @lru_cache(maxsize=4) def _counter(factory: RustTokenCounterFactory, tokenizer: RustTokenizer) -> RustTokenCounter: - return factory.from_tokenizer(_native_tokenizer(tokenizer)) + return factory.from_tokenizer(_native_tokenizer(tokenizer), fast=True) def _native_tokenizer(tokenizer: RustTokenizer) -> NativeTokenizer: diff --git a/tests/unit/router_strategy/test_lowest_tpm_rpm.py b/tests/unit/router_strategy/test_lowest_tpm_rpm.py index 625f648bec4..a3d183ff129 100644 --- a/tests/unit/router_strategy/test_lowest_tpm_rpm.py +++ b/tests/unit/router_strategy/test_lowest_tpm_rpm.py @@ -1,9 +1,10 @@ from datetime import datetime, timedelta from typing import Final -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, Mock import pytest +import litellm.router_strategy.lowest_tpm_rpm_v2 as strategy_module from litellm import Router from litellm.caching.dual_cache import DualCache from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2, PrefetchedUsage @@ -42,14 +43,11 @@ def test_usage_based_routing_v1_selects_the_lowest_recorded_tpm() -> None: } now: Final = datetime.now() cache_keys: Final = tuple( - f"{MODEL_GROUP}:tpm:{(now + timedelta(minutes=offset)).strftime('%H-%M')}" - for offset in range(60) + f"{MODEL_GROUP}:tpm:{(now + timedelta(minutes=offset)).strftime('%H-%M')}" for offset in range(60) ) for cache_key in cache_keys: - router.cache.set_cache( - key=cache_key, value=usage_by_deployment, ttl=float("inf") - ) + router.cache.set_cache(key=cache_key, value=usage_by_deployment, ttl=float("inf")) deployment: Final = router.get_available_deployment( model=MODEL_GROUP, @@ -109,11 +107,58 @@ async def test_v2_subclass_overriding_async_get_available_deployments_with_the_o ) router.lowesttpm_logger_v2 = OldSignatureV2(router_cache=router.cache, routing_args={}) - response: Final = await router.acompletion( - model=MODEL_GROUP, messages=[{"role": "user", "content": "x"}] - ) + 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}", } + + +@pytest.mark.asyncio +async def test_v2_counts_prompt_tokens_only_when_a_deployment_declares_a_tpm_limit( + monkeypatch: pytest.MonkeyPatch, +) -> None: + counter: Final = Mock(return_value=40_000) + monkeypatch.setattr(strategy_module, "token_counter", counter) + router_cache: Final = DualCache() + monkeypatch.setattr(router_cache, "async_batch_get_cache", AsyncMock(return_value=[10, 20, None, None])) + strategy: Final = LowestTPMLoggingHandler_v2(router_cache=router_cache) + messages: Final = [{"role": "user", "content": "a long prompt"}] + unlimited: Final = [ + {"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "a"}}, + {"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "b"}}, + ] + + chosen: Final = await strategy.async_get_available_deployments( + model_group="g", healthy_deployments=unlimited, messages=messages + ) + assert chosen["model_info"]["id"] == "a", "the lowest counter still wins without any tpm limit" + assert counter.call_count == 0, "no deployment has a tpm limit, so the prompt is never tokenized for selection" + + limited: Final = [ + {"model_name": "g", "litellm_params": {"model": "m", "tpm": 30_000}, "model_info": {"id": "a"}}, + {"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "b"}}, + ] + chosen_limited: Final = await strategy.async_get_available_deployments( + model_group="g", healthy_deployments=limited, messages=messages + ) + assert counter.call_count == 1, "a declared tpm limit needs the prompt size to be enforced" + assert chosen_limited["model_info"]["id"] == "b", "a's tpm limit cannot fit the counted prompt, so b is picked" + + +@pytest.mark.asyncio +async def test_v2_treats_a_failed_prompt_count_as_zero_tokens(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(strategy_module, "token_counter", Mock(side_effect=ValueError("no tokenizer"))) + router_cache: Final = DualCache() + monkeypatch.setattr(router_cache, "async_batch_get_cache", AsyncMock(return_value=[10, 20, None, None])) + strategy: Final = LowestTPMLoggingHandler_v2(router_cache=router_cache) + limited: Final = [ + {"model_name": "g", "litellm_params": {"model": "m", "tpm": 30_000}, "model_info": {"id": "a"}}, + {"model_name": "g", "litellm_params": {"model": "m"}, "model_info": {"id": "b"}}, + ] + + chosen: Final = await strategy.async_get_available_deployments( + model_group="g", healthy_deployments=limited, messages=[{"role": "user", "content": "a long prompt"}] + ) + assert chosen["model_info"]["id"] == "a", "a failed count weighs as zero tokens, so the lowest counter still wins" diff --git a/tests/unit/rust_bridge/test_catalog.py b/tests/unit/rust_bridge/test_catalog.py index 95e98fe98da..e7c50739e04 100644 --- a/tests/unit/rust_bridge/test_catalog.py +++ b/tests/unit/rust_bridge/test_catalog.py @@ -47,7 +47,7 @@ def test_shipped_decisions( if route is Route.OCR or (route is Route.TRANSCRIPTION and provider == "bedrock"): assert catalog.rollout(context) is Rollout.RUST_REQUIRED assert catalog.decision(context) is Decision.RUST_REQUIRED - elif route is Route.MESSAGES and provider == "anthropic": + elif (route is Route.MESSAGES and provider == "anthropic") or route is Route.TOKEN_COUNTER: assert catalog.rollout(context) is Rollout.RUST_OPT_IN opted_in: Final = environment == "1" or (environment is None and process is True) assert catalog.decision(context) is (Decision.RUST_WITH_FALLBACK if opted_in else Decision.PYTHON) diff --git a/tests/unit/rust_bridge/test_token_counter.py b/tests/unit/rust_bridge/test_token_counter.py index 3af7127a9b3..d6358bee721 100644 --- a/tests/unit/rust_bridge/test_token_counter.py +++ b/tests/unit/rust_bridge/test_token_counter.py @@ -94,7 +94,7 @@ async def test_native_count_returns_typed_count_and_reuses_one_counter(fake_toke assert second == first assert len(factory.counters) == 1 assert factory.counters[0].bodies == [BODY, BODY] - assert factory.counters[0].fast is False + assert factory.counters[0].fast is True assert factory.counters[0].tokenizer is tokenizer_dispatch.native_anthropic() assert json.loads(factory.counters[0].tokenizer.json or "")["model"]["type"] == "BPE" @@ -113,7 +113,7 @@ async def test_tiktoken_counter_is_built_over_the_shared_encoding_once( assert len(factory.counters) == 1 assert factory.counters[0].tokenizer.name == tokenizer assert factory.counters[0].tokenizer is tokenizer_dispatch.native_encoding(tokenizer) - assert factory.counters[0].fast is False + assert factory.counters[0].fast is True assert factory.counters[0].bodies == [BODY, BODY]