diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index c05f98d22ae..75d4d13eb71 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -24,7 +24,7 @@ "limit": 42 }, "reportExplicitAny": { - "limit": 10393 + "limit": 10389 }, "reportFunctionMemberAccess": { "limit": 11 diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index d51689f5f3b..5abb21879d3 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -2,7 +2,6 @@ from __future__ import annotations -import os from dataclasses import dataclass from typing import Awaitable, Final, Protocol, Union, cast @@ -46,43 +45,26 @@ class _Unset: _UNSET: Final[_Unset] = _Unset() -def _env_enables_rust_messages() -> bool: - return os.getenv("LITELLM_USE_RUST_MESSAGES", "").strip().lower() in { - "1", - "true", - "yes", - "on", - } - - @dataclass(slots=True) class _RustMessagesState: - enabled: bool messages: RustMessages | None = None amessages: RustAmessages | None = None -_STATE: Final[_RustMessagesState] = _RustMessagesState(enabled=_env_enables_rust_messages()) +_STATE: Final[_RustMessagesState] = _RustMessagesState() def set_rust_messages( - enabled: bool | _Unset = _UNSET, *, messages: RustMessages | None | _Unset = _UNSET, amessages: RustAmessages | None | _Unset = _UNSET, ) -> None: - if not isinstance(enabled, _Unset): - _STATE.enabled = enabled if not isinstance(messages, _Unset): _STATE.messages = messages if not isinstance(amessages, _Unset): _STATE.amessages = amessages -def rust_messages_enabled() -> bool: - return _STATE.enabled - - def load_rust_messages() -> RustMessages | None: if _STATE.messages is not None: return _STATE.messages diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 3654b9fb232..35de2eb9727 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -73,21 +73,24 @@ def use_litellm_rust( amessages: RustAmessages | None | _Unset = _UNSET, ) -> None: global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl - _rust_ocr_enabled = enabled + configuring_ocr = not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset) + configuring_messages = not isinstance(messages, _Unset) or not isinstance(amessages, _Unset) + if configuring_ocr or not configuring_messages: + _rust_ocr_enabled = enabled if not isinstance(ocr, _Unset): _rust_ocr_impl = ocr if not isinstance(aocr, _Unset): _rust_aocr_impl = aocr - if isinstance(messages, _Unset) and isinstance(amessages, _Unset): + if not configuring_messages: return from litellm.rust_bridge.messages import set_rust_messages if not isinstance(messages, _Unset) and not isinstance(amessages, _Unset): - set_rust_messages(enabled, messages=messages, amessages=amessages) + set_rust_messages(messages=messages, amessages=amessages) elif not isinstance(messages, _Unset): - set_rust_messages(enabled, messages=messages) + set_rust_messages(messages=messages) else: - set_rust_messages(enabled, amessages=amessages) + set_rust_messages(amessages=amessages) def rust_ocr_enabled() -> bool: diff --git a/tests/documentation_tests/test_env_keys.py b/tests/documentation_tests/test_env_keys.py index 1086d67b279..60fbd505d67 100644 --- a/tests/documentation_tests/test_env_keys.py +++ b/tests/documentation_tests/test_env_keys.py @@ -28,7 +28,6 @@ EXCLUDED_GUARD_ONLY_VARS = { # environment settings docs until the feature is ready for broad use. EXCLUDED_ROLLOUT_FLAGS = { "LITELLM_USE_RUST_OCR", - "LITELLM_USE_RUST_MESSAGES", } EXCLUDED_TERMINAL_VARS = { diff --git a/tests/test_litellm/rust_bridge/test_messages.py b/tests/test_litellm/rust_bridge/test_messages.py index 6f0e1e190a1..f5be34ae0b6 100644 --- a/tests/test_litellm/rust_bridge/test_messages.py +++ b/tests/test_litellm/rust_bridge/test_messages.py @@ -109,6 +109,27 @@ def test_load_rust_messages_returns_injected_impl(): assert rust_messages.load_rust_messages() is bridge +def test_configuring_messages_does_not_enable_ocr(): + from litellm.rust_bridge.ocr import rust_ocr_enabled + + litellm.use_litellm_rust(False) + assert rust_ocr_enabled() is False + + litellm.use_litellm_rust(True, messages=RecordingMessages()) + + assert rust_ocr_enabled() is False + + +def test_bare_use_litellm_rust_still_toggles_ocr(): + from litellm.rust_bridge.ocr import rust_ocr_enabled + + litellm.use_litellm_rust(True) + assert rust_ocr_enabled() is True + + litellm.use_litellm_rust(False) + assert rust_ocr_enabled() is False + + def test_load_rust_amessages_returns_injected_impl(): bridge = RecordingAsyncMessages() litellm.use_litellm_rust(True, amessages=bridge) diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 8fe7c588cf1..0482f47e5bc 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1112 + "limit": 1111 }, "LIT007": { "limit": 0