fix(rust): apply token counter rollout at the public API boundary

This commit is contained in:
Yujong Lee 2026-09-14 17:28:27 -07:00
parent 6a30c58a17
commit 592fd00504
26 changed files with 181 additions and 125 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -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=""),
)

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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