This commit is contained in:
devin-ai-integration[bot] 2026-10-04 23:19:19 +08:00 • committed by GitHub
commit ded1ba7306
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 74 additions and 19 deletions

View file

@ -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

View file

@ -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})),

View file

@ -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:

View file

@ -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"

View file

@ -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)

View file

@ -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]