From 592fd00504d0fbe33581bf864c183b2422d161c6 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 14 Sep 2026 17:28:27 -0700 Subject: [PATCH] fix(rust): apply token counter rollout at the public API boundary --- litellm/litellm_core_utils/token_counter.py | 36 +++++++++++++ litellm/rust_bridge/README.md | 8 +-- litellm/rust_bridge/catalog.py | 54 +++++++++---------- .../chat_completions/definition.py | 4 +- litellm/rust_bridge/configuration.py | 10 ++-- litellm/rust_bridge/embeddings/definition.py | 4 +- litellm/rust_bridge/image_edit/definition.py | 4 +- .../image_generation/definition.py | 4 +- litellm/rust_bridge/messages/definition.py | 4 +- litellm/rust_bridge/moderation/definition.py | 4 +- litellm/rust_bridge/ocr/definition.py | 4 +- litellm/rust_bridge/rerank/definition.py | 4 +- litellm/rust_bridge/responses/definition.py | 4 +- litellm/rust_bridge/route.py | 7 ++- litellm/rust_bridge/speech/definition.py | 4 +- litellm/rust_bridge/token_counter/__init__.py | 3 +- .../rust_bridge/token_counter/definition.py | 5 +- litellm/rust_bridge/token_counter/value.py | 18 +++---- .../rust_bridge/transcription/definition.py | 4 +- litellm/utils.py | 9 ---- .../spend_tracking/test_budget_reservation.py | 15 ++++-- .../rust_bridge/test_configuration_env.py | 39 ++++++++++++-- tests/test_litellm/rust_bridge/test_route.py | 19 +++---- .../test_litellm/rust_bridge/test_runtime.py | 5 +- .../rust_bridge/test_token_counter.py | 6 +-- .../test_route_foundation.py | 28 +++++----- 26 files changed, 181 insertions(+), 125 deletions(-) diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 4c61fac82bb..769a1da4956 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -422,6 +422,42 @@ def token_counter( Returns: int: The number of tokens in the text. """ + from functools import partial + + from litellm.rust_bridge.runtime import BridgeErrorContext, invoke + from litellm.rust_bridge.token_counter.definition import COMPONENT + + return invoke( + execution=COMPONENT.resolve(), + native_call=None, + python_fallback=partial( + _token_counter_python, + model=model, + custom_tokenizer=custom_tokenizer, + text=text, + messages=messages, + count_response_tokens=count_response_tokens, + tools=tools, + tool_choice=tool_choice, + use_default_image_token_count=use_default_image_token_count, + default_token_count=default_token_count, + ), + adapt=lambda count: count, + context=BridgeErrorContext(route=COMPONENT.name.value, provider="", model=model), + ) + + +def _token_counter_python( + model="", + custom_tokenizer: dict | SelectTokenizerResponse | None = None, + text: str | list[str] | None = None, + messages: Sequence[AllMessageValues | Message] | None = None, + count_response_tokens: bool | None = False, + tools: list[ChatCompletionToolParam] | None = None, + tool_choice: ChatCompletionNamedToolChoiceParam | None = None, + use_default_image_token_count: bool | None = False, + default_token_count: int | None = None, +) -> int: from litellm.utils import convert_list_message_to_dict ######################################################### diff --git a/litellm/rust_bridge/README.md b/litellm/rust_bridge/README.md index 641cad6dfc1..bc40449dc34 100644 --- a/litellm/rust_bridge/README.md +++ b/litellm/rust_bridge/README.md @@ -6,7 +6,7 @@ Every SDK API has one `NativeComponent` in the immutable `COMPONENTS` catalog. A `RustImplementationState` records whether Rust is unimplemented, experimental, or ready. `RolloutPolicy` independently selects unsupported, Python-only, Rust opt-in, Rust opt-out, or Rust-required execution. Optional Rust execution can fall back to Python. Rust-required execution cannot -OCR completed delivery is ready and default-on. Messages, chat completions, raw-request input token counting, and Responses WebSocket transport are experimental and opt-in. The public `litellm.token_counter()` and other completed APIs remain Python-only. Bedrock transcription requires Rust because it has no Python implementation; Python-backed transcription providers remain on Python +OCR completed delivery is ready and default-on. Messages, chat completions, and Responses WebSocket transport are experimental and opt-in. The public `litellm.token_counter()` and other completed APIs remain Python-only. Bedrock transcription requires Rust because it has no Python implementation; Python-backed transcription providers remain on Python ```python execution = COMPONENT.resolve( @@ -28,11 +28,11 @@ Provider failures, host callback failures, cancellation, conversion failures, an `invoke` and `ainvoke` return the native result or execute the supplied fallback directly. There is no public admission, prepare, accepts, or can-handle API -Token counting has two catalog entries. `UtilityName.TOKEN_COUNTER` declares the public `litellm.token_counter()` as Python-only, with no native exports. The public function continues to execute Python directly regardless of `litellm.rust(bool)` or `LITELLM_RUST` +`ComponentName` identifies every API in the catalog, including token counting. The public `litellm.token_counter()` resolves `ComponentName.TOKEN_COUNTER` and executes through `invoke`. Its policy is `PYTHON_ONLY`, so environment and process overrides keep public calls on the Python implementation -`UtilityName.REQUEST_INPUT_TOKEN_COUNTER` owns the experimental raw-request optimization. Budget reservation keeps its direct `litellm.rust_bridge.token_counter.count_input_tokens` import. That adapter uses `REQUEST_COMPONENT`; `COMPONENT` describes the public API +The direct `litellm.rust_bridge.token_counter.count_input_tokens` adapter bypasses public rollout policy and attempts its native binding even with `LITELLM_RUST=0` or `litellm.rust(False)`. Budget reservation keeps that direct import. Missing native bindings or request bodies, unsupported inputs, and unavailable tokenizer resources use the supplied Python fallback. Unexpected failures propagate without replay -The request optimization follows its component policy. Its one native counting entrypoint validates the tokenizer configuration and request body, obtains and caches the required tokenizer resource, then counts. Unsupported inputs decline, known resource loading failures report native unavailability, and unexpected counting failures propagate +The raw-body binding remains experimental. It does not implement the synchronous public signature, so public dispatch has no native callable yet. The catalog declares the existing native export and the public Python-only policy in one component ## Package layout diff --git a/litellm/rust_bridge/catalog.py b/litellm/rust_bridge/catalog.py index 741bca200aa..eb309245115 100644 --- a/litellm/rust_bridge/catalog.py +++ b/litellm/rust_bridge/catalog.py @@ -5,11 +5,10 @@ from litellm.rust_bridge.configuration import ( CapabilityContext, CapabilityDefinition, CapabilitySpec, + ComponentName, DeliveryMode, RolloutPolicy, - RouteName, RustImplementationState, - UtilityName, ) from litellm.rust_bridge.route import NativeComponent @@ -84,7 +83,7 @@ def _transcription_capability(context: CapabilityContext) -> CapabilityDefinitio def _component( - name: RouteName | UtilityName, + name: ComponentName, capability: CapabilitySpec, exports: tuple[str, ...], ) -> NativeComponent: @@ -93,8 +92,8 @@ def _component( COMPONENTS: Final = MappingProxyType( { - RouteName.OCR: _component( - RouteName.OCR, + ComponentName.OCR: _component( + ComponentName.OCR, _ocr_capability, ( "ocr", @@ -106,42 +105,41 @@ COMPONENTS: Final = MappingProxyType( "_ocr_lifecycle", ), ), - RouteName.MESSAGES: _component( - RouteName.MESSAGES, + ComponentName.MESSAGES: _component( + ComponentName.MESSAGES, _experimental_completed, ("messages", "amessages", "_messages_lifecycle"), ), - RouteName.CHAT_COMPLETIONS: _component( - RouteName.CHAT_COMPLETIONS, + ComponentName.CHAT_COMPLETIONS: _component( + ComponentName.CHAT_COMPLETIONS, _experimental_completed, ("chat_completions", "achat_completions", "_chat_completions_lifecycle"), ), - RouteName.TRANSCRIPTION: _component( - RouteName.TRANSCRIPTION, + ComponentName.TRANSCRIPTION: _component( + ComponentName.TRANSCRIPTION, _transcription_capability, ("transcription", "atranscription", "_transcription_lifecycle"), ), - RouteName.EMBEDDINGS: _component(RouteName.EMBEDDINGS, _python_completed, ("_embeddings_lifecycle",)), - RouteName.RERANK: _component(RouteName.RERANK, _python_completed, ("_rerank_lifecycle",)), - RouteName.IMAGE_GENERATION: _component( - RouteName.IMAGE_GENERATION, _python_completed, ("_image_generation_lifecycle",) + ComponentName.EMBEDDINGS: _component(ComponentName.EMBEDDINGS, _python_completed, ("_embeddings_lifecycle",)), + ComponentName.RERANK: _component(ComponentName.RERANK, _python_completed, ("_rerank_lifecycle",)), + ComponentName.IMAGE_GENERATION: _component( + ComponentName.IMAGE_GENERATION, _python_completed, ("_image_generation_lifecycle",) ), - RouteName.IMAGE_EDIT: _component(RouteName.IMAGE_EDIT, _python_completed, ("_image_edit_lifecycle",)), - RouteName.SPEECH: _component(RouteName.SPEECH, _python_completed, ("_speech_lifecycle",)), - RouteName.MODERATION: _component(RouteName.MODERATION, _python_completed, ("_moderation_lifecycle",)), - RouteName.RESPONSES: _component( - RouteName.RESPONSES, + ComponentName.IMAGE_EDIT: _component(ComponentName.IMAGE_EDIT, _python_completed, ("_image_edit_lifecycle",)), + ComponentName.SPEECH: _component(ComponentName.SPEECH, _python_completed, ("_speech_lifecycle",)), + ComponentName.MODERATION: _component(ComponentName.MODERATION, _python_completed, ("_moderation_lifecycle",)), + ComponentName.RESPONSES: _component( + ComponentName.RESPONSES, _responses_capability, ("ResponsesWebSocketConnection", "_responses_lifecycle"), ), - UtilityName.TOKEN_COUNTER: _component( - UtilityName.TOKEN_COUNTER, - _unimplemented(), - (), - ), - UtilityName.REQUEST_INPUT_TOKEN_COUNTER: _component( - UtilityName.REQUEST_INPUT_TOKEN_COUNTER, - _experimental_completed, + ComponentName.TOKEN_COUNTER: _component( + ComponentName.TOKEN_COUNTER, + CapabilityDefinition( + rust=RustImplementationState.EXPERIMENTAL, + python_available=True, + rollout=RolloutPolicy.PYTHON_ONLY, + ), ("count_input_tokens",), ), } diff --git a/litellm/rust_bridge/chat_completions/definition.py b/litellm/rust_bridge/chat_completions/definition.py index 8a26cfcfc55..c388d81c6fd 100644 --- a/litellm/rust_bridge/chat_completions/definition.py +++ b/litellm/rust_bridge/chat_completions/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.CHAT_COMPLETIONS] +COMPONENT: Final = COMPONENTS[ComponentName.CHAT_COMPLETIONS] diff --git a/litellm/rust_bridge/configuration.py b/litellm/rust_bridge/configuration.py index d5a231842a8..282c896ad23 100644 --- a/litellm/rust_bridge/configuration.py +++ b/litellm/rust_bridge/configuration.py @@ -9,7 +9,7 @@ DEFAULT_RUST_ENABLED: Final = False _GLOBAL_ENV_NAME: Final = "LITELLM_RUST" -class RouteName(str, Enum): +class ComponentName(str, Enum): OCR = "ocr" MESSAGES = "messages" CHAT_COMPLETIONS = "chat_completions" @@ -21,11 +21,7 @@ class RouteName(str, Enum): SPEECH = "speech" MODERATION = "moderation" RESPONSES = "responses" - - -class UtilityName(str, Enum): TOKEN_COUNTER = "token_counter" - REQUEST_INPUT_TOKEN_COUNTER = "request_input_token_counter" class RustImplementationState(str, Enum): @@ -78,7 +74,9 @@ class CapabilityDefinition: raise ValueError("an unimplemented Rust capability must use its Python implementation") return if self.rollout is RolloutPolicy.PYTHON_ONLY: - raise ValueError("an implemented Rust capability must declare a Rust rollout") + if not self.python_available: + raise ValueError("a Python-only capability requires a Python implementation") + return if self.rollout is RolloutPolicy.RUST_OPT_IN or self.rollout is RolloutPolicy.RUST_OPT_OUT: if not self.python_available: raise ValueError("an optional Rust capability requires a Python fallback") diff --git a/litellm/rust_bridge/embeddings/definition.py b/litellm/rust_bridge/embeddings/definition.py index 889bee08fae..d2892a451e9 100644 --- a/litellm/rust_bridge/embeddings/definition.py +++ b/litellm/rust_bridge/embeddings/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.EMBEDDINGS] +COMPONENT: Final = COMPONENTS[ComponentName.EMBEDDINGS] diff --git a/litellm/rust_bridge/image_edit/definition.py b/litellm/rust_bridge/image_edit/definition.py index f02440ef2db..f7108eeff6e 100644 --- a/litellm/rust_bridge/image_edit/definition.py +++ b/litellm/rust_bridge/image_edit/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.IMAGE_EDIT] +COMPONENT: Final = COMPONENTS[ComponentName.IMAGE_EDIT] diff --git a/litellm/rust_bridge/image_generation/definition.py b/litellm/rust_bridge/image_generation/definition.py index fd705d5f90b..ca93089c1a5 100644 --- a/litellm/rust_bridge/image_generation/definition.py +++ b/litellm/rust_bridge/image_generation/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.IMAGE_GENERATION] +COMPONENT: Final = COMPONENTS[ComponentName.IMAGE_GENERATION] diff --git a/litellm/rust_bridge/messages/definition.py b/litellm/rust_bridge/messages/definition.py index 8559bf7b022..04f14b9ed0f 100644 --- a/litellm/rust_bridge/messages/definition.py +++ b/litellm/rust_bridge/messages/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.MESSAGES] +COMPONENT: Final = COMPONENTS[ComponentName.MESSAGES] diff --git a/litellm/rust_bridge/moderation/definition.py b/litellm/rust_bridge/moderation/definition.py index fece2e5dd04..439f5753bcd 100644 --- a/litellm/rust_bridge/moderation/definition.py +++ b/litellm/rust_bridge/moderation/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.MODERATION] +COMPONENT: Final = COMPONENTS[ComponentName.MODERATION] diff --git a/litellm/rust_bridge/ocr/definition.py b/litellm/rust_bridge/ocr/definition.py index 9e2e03214ae..09d5901ab7d 100644 --- a/litellm/rust_bridge/ocr/definition.py +++ b/litellm/rust_bridge/ocr/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.OCR] +COMPONENT: Final = COMPONENTS[ComponentName.OCR] diff --git a/litellm/rust_bridge/rerank/definition.py b/litellm/rust_bridge/rerank/definition.py index e8107e8fa4f..135f867c3ab 100644 --- a/litellm/rust_bridge/rerank/definition.py +++ b/litellm/rust_bridge/rerank/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.RERANK] +COMPONENT: Final = COMPONENTS[ComponentName.RERANK] diff --git a/litellm/rust_bridge/responses/definition.py b/litellm/rust_bridge/responses/definition.py index 94a389b8ff1..df282163173 100644 --- a/litellm/rust_bridge/responses/definition.py +++ b/litellm/rust_bridge/responses/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.RESPONSES] +COMPONENT: Final = COMPONENTS[ComponentName.RESPONSES] diff --git a/litellm/rust_bridge/route.py b/litellm/rust_bridge/route.py index f7eaf249f1b..f3e932b6a3f 100644 --- a/litellm/rust_bridge/route.py +++ b/litellm/rust_bridge/route.py @@ -9,9 +9,8 @@ from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.configuration import ( CapabilityContext, CapabilitySpec, + ComponentName, ExecutionDecision, - RouteName, - UtilityName, capability_decision, ) from litellm.rust_bridge.errors import RustRouteUnavailableError, RustRouteUnsupportedError @@ -61,7 +60,7 @@ def _lifecycle(value: object) -> NativeLifecycle[object, object] | None: @dataclass(frozen=True, slots=True) class ComponentExecution: - route_name: RouteName | UtilityName + route_name: ComponentName decision: ExecutionDecision def require_supported(self) -> None: @@ -84,7 +83,7 @@ class ComponentExecution: @dataclass(frozen=True, slots=True) class NativeComponent: - name: RouteName | UtilityName + name: ComponentName capability: CapabilitySpec exports: tuple[str, ...] diff --git a/litellm/rust_bridge/speech/definition.py b/litellm/rust_bridge/speech/definition.py index 604a93920ff..cd5d73ca5e4 100644 --- a/litellm/rust_bridge/speech/definition.py +++ b/litellm/rust_bridge/speech/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.SPEECH] +COMPONENT: Final = COMPONENTS[ComponentName.SPEECH] diff --git a/litellm/rust_bridge/token_counter/__init__.py b/litellm/rust_bridge/token_counter/__init__.py index acd6e86aef8..664d71093fc 100644 --- a/litellm/rust_bridge/token_counter/__init__.py +++ b/litellm/rust_bridge/token_counter/__init__.py @@ -1,12 +1,11 @@ from typing import Final -from litellm.rust_bridge.token_counter.definition import COMPONENT, REQUEST_COMPONENT +from litellm.rust_bridge.token_counter.definition import COMPONENT from litellm.rust_bridge.token_counter.types import InputTokenCount, RustTokenizer from litellm.rust_bridge.token_counter.value import TOKEN_COUNTER, count_input_tokens, rust_tokenizer __all__: Final = ( "COMPONENT", - "REQUEST_COMPONENT", "TOKEN_COUNTER", "InputTokenCount", "RustTokenizer", diff --git a/litellm/rust_bridge/token_counter/definition.py b/litellm/rust_bridge/token_counter/definition.py index 7197e6ec3bb..6e53f62d208 100644 --- a/litellm/rust_bridge/token_counter/definition.py +++ b/litellm/rust_bridge/token_counter/definition.py @@ -1,7 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import UtilityName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[UtilityName.TOKEN_COUNTER] -REQUEST_COMPONENT: Final = COMPONENTS[UtilityName.REQUEST_INPUT_TOKEN_COUNTER] +COMPONENT: Final = COMPONENTS[ComponentName.TOKEN_COUNTER] diff --git a/litellm/rust_bridge/token_counter/value.py b/litellm/rust_bridge/token_counter/value.py index 1f2af468804..391204cdc06 100644 --- a/litellm/rust_bridge/token_counter/value.py +++ b/litellm/rust_bridge/token_counter/value.py @@ -7,8 +7,10 @@ from pydantic import TypeAdapter import litellm from litellm.litellm_core_utils.default_encoding import cl100k_base_rank_file, o200k_base_rank_file from litellm.litellm_core_utils.token_counter import openai_tokenizer_encoding, uses_legacy_message_accounting +from litellm.rust_bridge.configuration import ExecutionDecision +from litellm.rust_bridge.route import ComponentExecution from litellm.rust_bridge.runtime import BridgeErrorContext, ainvoke -from litellm.rust_bridge.token_counter.definition import REQUEST_COMPONENT +from litellm.rust_bridge.token_counter.definition import COMPONENT from litellm.rust_bridge.token_counter.types import InputTokenCount, RustTokenCounter, RustTokenizer from litellm.utils import claude_json_str, huggingface_tokenizer_kind @@ -19,13 +21,10 @@ def _as_counter(value: object) -> RustTokenCounter | None: return cast(RustTokenCounter, value) if callable(value) else None -TOKEN_COUNTER: Final = REQUEST_COMPONENT.bind("count_input_tokens", validate=_as_counter) +TOKEN_COUNTER: Final = COMPONENT.bind("count_input_tokens", validate=_as_counter) -def rust_tokenizer(model: str) -> RustTokenizer | None: - execution: Final = REQUEST_COMPONENT.resolve() - if execution.select(TOKEN_COUNTER) is None: - return None +def rust_tokenizer(model: str) -> RustTokenizer: kind: Final = None if litellm.disable_hf_tokenizer_download is True else huggingface_tokenizer_kind(model) encoding: Final = openai_tokenizer_encoding(model).name if kind is None else "" return RustTokenizer( @@ -49,14 +48,13 @@ def _tokenizer_resource(tokenizer: str) -> str: async def count_input_tokens(body: bytes, tokenizer: RustTokenizer) -> InputTokenCount | None: - execution: Final = REQUEST_COMPONENT.resolve() - counter: Final = execution.select(TOKEN_COUNTER) + counter: Final = TOKEN_COUNTER.load() async def python_fallback() -> None: return None return await ainvoke( - execution=execution, + execution=ComponentExecution(COMPONENT.name, ExecutionDecision.RUST_WITH_FALLBACK), native_call=( lambda: counter( body, @@ -71,5 +69,5 @@ async def count_input_tokens(body: bytes, tokenizer: RustTokenizer) -> InputToke else None, python_fallback=python_fallback, adapt=_INPUT_TOKEN_COUNT.validate_python, - context=BridgeErrorContext(route=REQUEST_COMPONENT.name.value, provider="", model=""), + context=BridgeErrorContext(route=COMPONENT.name.value, provider="", model=""), ) diff --git a/litellm/rust_bridge/transcription/definition.py b/litellm/rust_bridge/transcription/definition.py index 4307ee1395b..9428e7ba920 100644 --- a/litellm/rust_bridge/transcription/definition.py +++ b/litellm/rust_bridge/transcription/definition.py @@ -1,6 +1,6 @@ from typing import Final from litellm.rust_bridge.catalog import COMPONENTS -from litellm.rust_bridge.configuration import RouteName +from litellm.rust_bridge.configuration import ComponentName -COMPONENT: Final = COMPONENTS[RouteName.TRANSCRIPTION] +COMPONENT: Final = COMPONENTS[ComponentName.TRANSCRIPTION] diff --git a/litellm/utils.py b/litellm/utils.py index d4e3d58ba9f..d6b1e83aab3 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2361,15 +2361,6 @@ def token_counter( Kept for backwards compatibility. """ - ######################################################### - # Flag to disable token counter - # We've gotten reports of this consuming CPU cycles, - # exposing this flag to allow users to disable - # it to confirm if this is indeed the issue - ######################################################### - if litellm.disable_token_counter is True: - return 0 - return _get_token_counter_new()( model, custom_tokenizer, diff --git a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py index 3e0acf917aa..19bd16962cf 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py +++ b/tests/test_litellm/proxy/spend_tracking/test_budget_reservation.py @@ -404,7 +404,7 @@ async def test_rust_decline_falls_back_to_python_count(rust_counter: None, model @pytest.mark.asyncio -async def test_disabled_rust_never_sees_the_raw_body(rust_counter: None) -> None: +async def test_direct_budget_counter_ignores_disabled_public_rollout(rust_counter: None) -> None: factory: Final = _RecordingFactory() litellm.rust(False) rust_token_counter.TOKEN_COUNTER.override(factory) @@ -417,9 +417,16 @@ async def test_disabled_rust_never_sees_the_raw_body(rust_counter: None) -> None raw_body=json.dumps(body).encode(), ) - assert factory.calls == [] - assert set(counts) == {ANTHROPIC_TOKENIZER_MODEL, CL100K_MODEL, O200K_MODEL} - assert not set(counts.values()) & set(RUST_INPUT_TOKENS_BY_TOKENIZER.values()) + assert factory.calls == [ + ("anthropic", json.dumps(body).encode()), + ("cl100k_base", json.dumps(body).encode()), + ("o200k_base", json.dumps(body).encode()), + ] + assert counts == { + ANTHROPIC_TOKENIZER_MODEL: RUST_INPUT_TOKENS_BY_TOKENIZER["anthropic"], + CL100K_MODEL: RUST_INPUT_TOKENS_BY_TOKENIZER["cl100k_base"], + O200K_MODEL: RUST_INPUT_TOKENS_BY_TOKENIZER["o200k_base"], + } @pytest.mark.asyncio diff --git a/tests/test_litellm/rust_bridge/test_configuration_env.py b/tests/test_litellm/rust_bridge/test_configuration_env.py index fb8f67d5f9c..da2db6a5bdb 100644 --- a/tests/test_litellm/rust_bridge/test_configuration_env.py +++ b/tests/test_litellm/rust_bridge/test_configuration_env.py @@ -12,10 +12,16 @@ from litellm.litellm_core_utils import token_counter as python_counter from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge import token_counter as bridge from litellm.rust_bridge.configuration import ( + CapabilityDefinition, + ComponentName, ExecutionDecision, + RolloutPolicy, + RustImplementationState, _parse_env_bool, # pyright: ignore[reportPrivateUsage] # directly test env parsing contract ) -from litellm.rust_bridge.token_counter import COMPONENT +from litellm.rust_bridge.errors import RustRouteUnsupportedError +from litellm.rust_bridge.route import NativeComponent +from litellm.rust_bridge.token_counter import COMPONENT, definition @pytest.mark.parametrize(("value", "expected"), (("1", True), ("0", False), (" 1 ", True), (" 0 ", False))) @@ -64,10 +70,34 @@ def test_public_token_counter_stays_python_only( configuration.reset_rust_configuration() +@pytest.mark.parametrize("disabled", (False, True)) +@pytest.mark.parametrize("through_compatibility_wrapper", (False, True)) +def test_public_counter_enforces_catalog_decision( + monkeypatch: pytest.MonkeyPatch, disabled: bool, through_compatibility_wrapper: bool +) -> None: + monkeypatch.setattr(litellm, "disable_token_counter", disabled) + monkeypatch.setattr( + definition, + "COMPONENT", + NativeComponent( + name=ComponentName.TOKEN_COUNTER, + capability=CapabilityDefinition( + rust=RustImplementationState.UNIMPLEMENTED, + python_available=False, + rollout=RolloutPolicy.UNSUPPORTED, + ), + exports=(), + ), + ) + counter: Final = litellm.token_counter if through_compatibility_wrapper else python_counter.token_counter + with pytest.raises(RustRouteUnsupportedError, match="token_counter"): + counter(model="gpt-4o", text="hello") + + @pytest.mark.asyncio @pytest.mark.parametrize("environment", (None, "0", "1")) @pytest.mark.parametrize("override", (None, False, True)) -async def test_budget_direct_import_follows_request_rollout( +async def test_budget_direct_import_bypasses_public_rollout( monkeypatch: pytest.MonkeyPatch, environment: str | None, override: bool | None ) -> None: from litellm.proxy.spend_tracking.budget_reservation import count_request_input_tokens @@ -96,13 +126,12 @@ async def test_budget_direct_import_follows_request_rollout( return {"model": model, "input_tokens": 42} bridge.TOKEN_COUNTER.override(counter) - enabled: Final = override if override is not None else environment == "1" try: budget: Final = await count_request_input_tokens( request_body=json.loads(body), route="/v1/messages", llm_router=None, raw_body=body ) - assert budget == {model: 42 if enabled else litellm.token_counter(model=model, messages=messages)} - assert native_calls == ([body] if enabled else []) + assert budget == {model: 42} + assert native_calls == [body] finally: bridge.TOKEN_COUNTER.reset() configuration.reset_rust_configuration() diff --git a/tests/test_litellm/rust_bridge/test_route.py b/tests/test_litellm/rust_bridge/test_route.py index bc1e1e7fe1e..062fbc8cbf6 100644 --- a/tests/test_litellm/rust_bridge/test_route.py +++ b/tests/test_litellm/rust_bridge/test_route.py @@ -12,10 +12,10 @@ from litellm.rust_bridge.catalog import COMPONENTS from litellm.rust_bridge.configuration import ( CapabilityContext, CapabilityDefinition, + ComponentName, DeliveryMode, ExecutionDecision, RolloutPolicy, - RouteName, RustImplementationState, ) from litellm.rust_bridge.errors import RustRouteUnavailableError, RustRouteUnsupportedError @@ -39,14 +39,14 @@ def _unexpected_load() -> ModuleType: def test_python_decision_does_not_load_native(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.delenv("LITELLM_RUST", raising=False) - component: Final = COMPONENTS[RouteName.MESSAGES] + component: Final = COMPONENTS[ComponentName.MESSAGES] binding: Final = component.bind("messages", validate=_string, module_loader=_unexpected_load) assert component.resolve().select(binding) is None def test_delivery_mode_uses_one_component_policy(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("LITELLM_RUST", "1") - component: Final = COMPONENTS[RouteName.MESSAGES] + component: Final = COMPONENTS[ComponentName.MESSAGES] completed: Final = component.resolve(CapabilityContext(delivery=DeliveryMode.COMPLETED)) streaming: Final = component.resolve(CapabilityContext(delivery=DeliveryMode.STREAMING)) assert completed.decision is ExecutionDecision.RUST_WITH_FALLBACK @@ -56,7 +56,7 @@ def test_delivery_mode_uses_one_component_policy(monkeypatch: pytest.MonkeyPatch def test_binding_discovery_validation_and_override() -> None: module: Final = ModuleType("fake_native") setattr(module, "messages", "native") - component: Final = COMPONENTS[RouteName.MESSAGES] + component: Final = COMPONENTS[ComponentName.MESSAGES] binding: Final = component.bind("messages", validate=_string, module_loader=lambda: module) assert binding.load() == "native" binding.configure("override") @@ -72,7 +72,7 @@ def test_binding_discovery_validation_and_override() -> None: def test_component_rejects_undeclared_binding() -> None: with pytest.raises(ValueError, match="not declared"): - COMPONENTS[RouteName.MESSAGES].bind("typo", validate=_string) + COMPONENTS[ComponentName.MESSAGES].bind("typo", validate=_string) class ModelCapability: @@ -97,7 +97,7 @@ class ModelCapability: def test_dynamic_capability_resolves_once_before_binding_selection() -> None: resolver: Final = ModelCapability() component: Final = NativeComponent( - name=RouteName.TRANSCRIPTION, + name=ComponentName.TRANSCRIPTION, capability=resolver, exports=("transcription",), ) @@ -116,7 +116,7 @@ def test_bedrock_transcription_requires_rust_regardless_of_overrides( monkeypatch.setenv("LITELLM_RUST", environment_override) if process_override is not None: configuration.rust(process_override) - component: Final = COMPONENTS[RouteName.TRANSCRIPTION] + component: Final = COMPONENTS[ComponentName.TRANSCRIPTION] binding: Final = component.bind("transcription", validate=_string, module_loader=lambda: None) execution: Final = component.resolve(CapabilityContext(provider="bedrock", model="model")) assert execution.decision is ExecutionDecision.RUST_REQUIRED @@ -127,7 +127,7 @@ def test_bedrock_transcription_requires_rust_regardless_of_overrides( @pytest.mark.parametrize("provider", ("openai", "azure", "azure_ai", "groq", "mistral", "nvidia_riva", "soniox")) def test_python_transcription_providers_skip_native_discovery(provider: str) -> None: configuration.rust(True) - component: Final = COMPONENTS[RouteName.TRANSCRIPTION] + component: Final = COMPONENTS[ComponentName.TRANSCRIPTION] binding: Final = component.bind("transcription", validate=_string, module_loader=_unexpected_load) execution: Final = component.resolve(CapabilityContext(provider=provider, model="model")) assert execution.decision is ExecutionDecision.PYTHON @@ -136,7 +136,7 @@ def test_python_transcription_providers_skip_native_discovery(provider: str) -> @pytest.mark.parametrize("provider", ("unknown-provider", "anthropic")) def test_unsupported_transcription_never_selects_an_implementation(provider: str) -> None: - component: Final = COMPONENTS[RouteName.TRANSCRIPTION] + component: Final = COMPONENTS[ComponentName.TRANSCRIPTION] binding: Final = component.bind("transcription", validate=_string, module_loader=_unexpected_load) execution: Final = component.resolve(CapabilityContext(provider=provider, model="model")) assert execution.decision is ExecutionDecision.UNSUPPORTED @@ -154,6 +154,7 @@ def test_unsupported_transcription_never_selects_an_implementation(provider: str (RustImplementationState.UNIMPLEMENTED, False, RolloutPolicy.RUST_REQUIRED), (RustImplementationState.UNIMPLEMENTED, True, RolloutPolicy.UNSUPPORTED), (RustImplementationState.READY, False, RolloutPolicy.UNSUPPORTED), + (RustImplementationState.EXPERIMENTAL, False, RolloutPolicy.PYTHON_ONLY), ), ) def test_invalid_capability_definitions_rejected( diff --git a/tests/test_litellm/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py index c3b24546e5c..76a0d3c25b2 100644 --- a/tests/test_litellm/rust_bridge/test_runtime.py +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -9,10 +9,11 @@ import pytest from litellm.exceptions import APIError from litellm.rust_bridge import bindings, runtime -from litellm.rust_bridge.configuration import ExecutionDecision, RouteName +from litellm.rust_bridge.configuration import ComponentName, ExecutionDecision from litellm.rust_bridge.errors import RustRouteDeclinedError, RustRouteUnavailableError, RustRouteUnsupportedError from litellm.rust_bridge.route import ComponentExecution + class RustBridgeDeclined(Exception): pass @@ -47,7 +48,7 @@ async def _invoke( adapt: Callable[[object], object], fallback: Callable[[], object] = lambda: "python", ) -> object: - execution: Final = ComponentExecution(route_name=RouteName.MESSAGES, decision=decision) + execution: Final = ComponentExecution(route_name=ComponentName.MESSAGES, decision=decision) context: Final = runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model") if not asynchronous: return runtime.invoke( diff --git a/tests/test_litellm/rust_bridge/test_token_counter.py b/tests/test_litellm/rust_bridge/test_token_counter.py index fd9febad083..c32238cec22 100644 --- a/tests/test_litellm/rust_bridge/test_token_counter.py +++ b/tests/test_litellm/rust_bridge/test_token_counter.py @@ -69,12 +69,12 @@ def reset_bridge(monkeypatch: pytest.MonkeyPatch) -> Iterator[None]: @pytest.mark.asyncio @pytest.mark.parametrize("tokenizer", TOKENIZERS) -async def test_disabled_bridge_never_calls_native(tokenizer: bridge.RustTokenizer) -> None: +async def test_direct_bridge_bypasses_disabled_public_rollout(tokenizer: bridge.RustTokenizer) -> None: counter: Final = _RecordingCounter() bridge.TOKEN_COUNTER.override(counter) litellm.rust(False) - assert await bridge.count_input_tokens(BODY, tokenizer) is None - assert counter.calls == [] + assert await bridge.count_input_tokens(BODY, tokenizer) == bridge.InputTokenCount(model=MODEL, input_tokens=42) + assert counter.calls == [(BODY, tokenizer, tokenizer.kind or tokenizer.encoding)] @pytest.mark.asyncio diff --git a/tests/test_litellm_rust/test_route_foundation.py b/tests/test_litellm_rust/test_route_foundation.py index a7b483854f3..900c0fa3aee 100644 --- a/tests/test_litellm_rust/test_route_foundation.py +++ b/tests/test_litellm_rust/test_route_foundation.py @@ -8,7 +8,7 @@ from litellm.rust_bridge import _native from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.catalog import NATIVE_EXPORTS from litellm.rust_bridge.chat_completions.lifecycle import LIFECYCLE as CHAT_COMPLETIONS -from litellm.rust_bridge.configuration import ExecutionDecision, RouteName +from litellm.rust_bridge.configuration import ComponentName, ExecutionDecision from litellm.rust_bridge.embeddings.lifecycle import LIFECYCLE as EMBEDDINGS from litellm.rust_bridge.image_edit.lifecycle import LIFECYCLE as IMAGE_EDIT from litellm.rust_bridge.image_generation.lifecycle import LIFECYCLE as IMAGE_GENERATION @@ -23,17 +23,17 @@ from litellm.rust_bridge.transcription.lifecycle import LIFECYCLE as TRANSCRIPTI pytestmark = pytest.mark.requires_rust_extension -UNIMPLEMENTED: Final[dict[RouteName, NativeBinding[NativeLifecycle[object, object]]]] = { - RouteName.MESSAGES: MESSAGES, - RouteName.CHAT_COMPLETIONS: CHAT_COMPLETIONS, - RouteName.TRANSCRIPTION: TRANSCRIPTION, - RouteName.EMBEDDINGS: EMBEDDINGS, - RouteName.RERANK: RERANK, - RouteName.IMAGE_GENERATION: IMAGE_GENERATION, - RouteName.IMAGE_EDIT: IMAGE_EDIT, - RouteName.SPEECH: SPEECH, - RouteName.MODERATION: MODERATION, - RouteName.RESPONSES: RESPONSES, +UNIMPLEMENTED: Final[dict[ComponentName, NativeBinding[NativeLifecycle[object, object]]]] = { + ComponentName.MESSAGES: MESSAGES, + ComponentName.CHAT_COMPLETIONS: CHAT_COMPLETIONS, + ComponentName.TRANSCRIPTION: TRANSCRIPTION, + ComponentName.EMBEDDINGS: EMBEDDINGS, + ComponentName.RERANK: RERANK, + ComponentName.IMAGE_GENERATION: IMAGE_GENERATION, + ComponentName.IMAGE_EDIT: IMAGE_EDIT, + ComponentName.SPEECH: SPEECH, + ComponentName.MODERATION: MODERATION, + ComponentName.RESPONSES: RESPONSES, } @@ -49,7 +49,7 @@ def test_catalog_exports_are_registered() -> None: @pytest.mark.parametrize(("route_name", "binding"), tuple(UNIMPLEMENTED.items())) @pytest.mark.parametrize("asynchronous", (False, True)) def test_package_lifecycle_binding_declines_without_input_reads( - route_name: RouteName, + route_name: ComponentName, binding: NativeBinding[NativeLifecycle[object, object]], asynchronous: bool, ) -> None: @@ -62,7 +62,7 @@ def test_package_lifecycle_binding_declines_without_input_reads( @pytest.mark.parametrize(("route_name", "binding"), tuple(UNIMPLEMENTED.items())) def test_package_stub_decline_selects_python( - route_name: RouteName, + route_name: ComponentName, binding: NativeBinding[NativeLifecycle[object, object]], ) -> None: native: Final[NativeLifecycle[object, object] | None] = binding.load()