mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(rust): apply token counter rollout at the public API boundary
This commit is contained in:
parent
6a30c58a17
commit
592fd00504
26 changed files with 181 additions and 125 deletions
|
|
@ -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
|
||||
|
||||
#########################################################
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",),
|
||||
),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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, ...]
|
||||
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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=""),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue