From 27bf5d9364e2e7dc06073fe5817a74344a56503e Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Tue, 8 Sep 2026 09:15:13 -0700 Subject: [PATCH] fix(rust): repair retained callback CI gates --- litellm/rust_bridge/_lifecycle.py | 4 +- litellm/rust_bridge/messages.py | 8 +- litellm/rust_bridge/ocr.py | 7 +- .../rust_bridge/native_route_wheel_test.py | 104 +++++++++++++++--- 4 files changed, 99 insertions(+), 24 deletions(-) diff --git a/litellm/rust_bridge/_lifecycle.py b/litellm/rust_bridge/_lifecycle.py index 20632d99e8d..d3f4636789d 100644 --- a/litellm/rust_bridge/_lifecycle.py +++ b/litellm/rust_bridge/_lifecycle.py @@ -20,9 +20,7 @@ TerminalAction = Literal[ "sync_failure", "async_failure", ] -_OPTIONAL_ARGUMENTS_ADAPTER: Final[TypeAdapter[dict[str, object] | None]] = TypeAdapter( - dict[str, object] | None -) +_OPTIONAL_ARGUMENTS_ADAPTER: Final[TypeAdapter[dict[str, object] | None]] = TypeAdapter(dict[str, object] | None) class NativeOutcome(IntEnum): diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 1d3d2694c2f..2f7e25a3770 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -56,8 +56,8 @@ _UNSET: Final[_Unset] = _Unset() @dataclass(slots=True) class _RustMessagesState: - messages: RustMessages | None = None - amessages: RustAmessages | None = None + messages: RustMessages | None | _Unset = _UNSET + amessages: RustAmessages | None | _Unset = _UNSET _STATE: Final = _RustMessagesState() @@ -87,13 +87,13 @@ def set_rust_messages( def load_rust_messages() -> RustMessages | None: - if _STATE.messages is not None: + if not isinstance(_STATE.messages, _Unset): return _STATE.messages return _MESSAGES.load() def load_rust_amessages() -> RustAmessages | None: - if _STATE.amessages is not None: + if not isinstance(_STATE.amessages, _Unset): return _STATE.amessages return _AMESSAGES.load() diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index e5ef4e73e3d..e8391739c9c 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -257,7 +257,12 @@ def invoke_terminal( return None if action == "sync_success_if_needed": if logging._should_run_sync_callbacks_for_async_calls(): # pyright: ignore[reportPrivateUsage] # preserves Logging's async callback policy - return invoke_terminal("sync_success", roots, logger, record, value, fallback_start_time, fallback_end_time) + + def run() -> None: + _retained: Final = roots + logging.success_handler(value, start_time, end_time) + + return utils.executor.submit(copy_context().run, run) return None exception: Final = cast(Exception, value) # cast-ok: Rust routes terminal failure values as Python exceptions trace: Final = "".join(traceback.format_exception(type(exception), exception, exception.__traceback__)) diff --git a/tests/test_litellm/rust_bridge/native_route_wheel_test.py b/tests/test_litellm/rust_bridge/native_route_wheel_test.py index 2896117c5d0..86e58b2e579 100644 --- a/tests/test_litellm/rust_bridge/native_route_wheel_test.py +++ b/tests/test_litellm/rust_bridge/native_route_wheel_test.py @@ -229,6 +229,51 @@ def restore_ocr_context(logger: object) -> None: pass +def drive_ocr_sync(arguments: dict[str, object], bindings: object) -> object: + logger: Final = initialize_ocr_logging(arguments, False) + try: + state: Final = bindings.prepare(arguments, logger, False) + except RuntimeError as error: + raise NotImplementedError(str(error)) from error + bindings.pre_call(state) + return bindings.send_sync(state) + + +async def drive_ocr_async(arguments: dict[str, object], bindings: object) -> object: + logger: Final = initialize_ocr_logging(arguments, True) + try: + state: Final = bindings.prepare(arguments, logger, True) + except RuntimeError as error: + raise NotImplementedError(str(error)) from error + bindings.pre_call(state) + return bindings.finish(await bindings.send(state)) + + +class MessagesLogging: + def pre_call(self, **_kwargs: object) -> None: + pass + + +def drive_messages_sync(arguments: dict[str, object], bindings: object) -> object: + state: Final = bindings.prepare(arguments, MessagesLogging()) + return bindings.send_sync(state) + + +async def drive_messages_async(arguments: dict[str, object], bindings: object) -> object: + state: Final = bindings.prepare(arguments, MessagesLogging()) + return await bindings.send(state) + + +def drive_chat_sync(arguments: dict[str, object], bindings: object) -> object: + state: Final = bindings.prepare(arguments, MessagesLogging()) + return bindings.send_sync(state) + + +async def drive_chat_async(arguments: dict[str, object], bindings: object) -> object: + state: Final = bindings.prepare(arguments, MessagesLogging()) + return await bindings.send(state) + + class WheelCallTypes: aocr = "aocr" @@ -265,22 +310,28 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]: }, } if route == "messages": - return common | { - "model": "claude-sonnet-4-5", - "body": { + return { + "arguments": common + | { "model": "claude-sonnet-4-5", - "max_tokens": 16, - "messages": [{"role": "user", "content": "hello-from-messages"}], - }, - "api_key": "sk-native", - "custom_llm_provider": "anthropic", + "body": { + "model": "claude-sonnet-4-5", + "max_tokens": 16, + "messages": [{"role": "user", "content": "hello-from-messages"}], + }, + "api_key": "sk-native", + "custom_llm_provider": "anthropic", + } } if route == "chat_completions": - return common | { - "model": "anthropic/claude-sonnet-4-5", - "messages": [{"role": "user", "content": "hello-from-chat"}], - "optional_params": {"max_tokens": 17}, - "api_key": "sk-native", + return { + "arguments": common + | { + "model": "anthropic/claude-sonnet-4-5", + "messages": [{"role": "user", "content": "hello-from-chat"}], + "optional_params": {"max_tokens": 17}, + "api_key": "sk-native", + } } raise AssertionError(f"unknown route: {route}") @@ -318,7 +369,7 @@ def assert_rate_limit(native: object, route: str, error: BaseException) -> None: assert error.llm_provider == "mistral" assert str(error) == "OCR provider request failed (HTTP 429)" return - if route == "chat_completions": + if route in {"messages", "chat_completions"}: upstream_error: Final = native.RustUpstreamError if not isinstance(error, upstream_error) or error.args[0] != 429: raise AssertionError(f"{route} returned the wrong 429 error: {error!r}") @@ -367,7 +418,8 @@ async def exercise_unsupported_ocr(native: object, api_base: str) -> None: else: native.ocr(arguments) except NotImplementedError as error: - assert operation in str(error) # noqa: PT017 # the selected native entry point is parametrized at runtime + if operation not in str(error): + raise AssertionError(f"native OCR returned the wrong decline: {error}") from error else: raise AssertionError(f"native OCR accepted unsupported operation: {operation}") assert arguments["litellm_logging_obj"].calls == () @@ -396,6 +448,14 @@ def exercise_routes(native_path: Path, api_base: str) -> object: ocr_bridge: Final = ModuleType("litellm.rust_bridge.ocr") ocr_bridge.initialize_logging = initialize_ocr_logging ocr_bridge.invoke_terminal = invoke_ocr_terminal + ocr_bridge._drive_sync = drive_ocr_sync + ocr_bridge._drive_async = drive_ocr_async + messages_bridge: Final = ModuleType("litellm.rust_bridge.messages") + messages_bridge._drive_sync = drive_messages_sync + messages_bridge._drive_async = drive_messages_async + chat_bridge: Final = ModuleType("litellm.rust_bridge.chat_completions") + chat_bridge._drive_sync = drive_chat_sync + chat_bridge._drive_async = drive_chat_async utils: Final = ModuleType("litellm.utils") utils.is_internal_call = ContextVar("wheel_internal_call", default=False) utils.async_pre_call_deployment_hook = pre_ocr_deployment @@ -418,7 +478,19 @@ def exercise_routes(native_path: Path, api_base: str) -> object: with patch.dict( sys.modules, packages - | {module.__name__: module for module in (transformation, exceptions, httpx, ocr_bridge, utils, types_utils)}, + | { + module.__name__: module + for module in ( + transformation, + exceptions, + httpx, + ocr_bridge, + messages_bridge, + chat_bridge, + utils, + types_utils, + ) + }, ): exercise_sync(native, api_base) asyncio.run(exercise_async(native, api_base))