diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index c82be07a5c5..b5e18bd8e7e 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -4,7 +4,7 @@ Calling + translation logic for anthropic's `/v1/messages` endpoint import copy import json -from collections.abc import Callable +from collections.abc import Callable, Mapping from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast import httpx @@ -380,7 +380,10 @@ class AnthropicChatCompletion(BaseLLM): request_data: Final = config.transform_request( model=model, messages=messages, - optional_params={**optional_params, "is_vertex_request": is_vertex_request}, + optional_params={ # mutable-ok: provider transforms require a plain mutable request dict + **optional_params, + "is_vertex_request": is_vertex_request, + }, litellm_params=litellm_params, headers=headers, ) @@ -390,6 +393,46 @@ class AnthropicChatCompletion(BaseLLM): provider=custom_llm_provider, ) + def sync_python_path( + request_headers: Mapping[str, object], request_data: Mapping[str, object] + ) -> ModelResponse: + request_client: Final = ( + _get_httpx_client(params={"timeout": timeout}) # mutable-ok: HTTP client factory requires a dict + if client is None or not isinstance(client, HTTPHandler) + else client + ) + try: + response: Final = request_client.post( + api_base, + headers=dict(request_headers), # mutable-ok: HTTP client requires mutable headers + data=json.dumps(request_data), + timeout=timeout, + logging_obj=logging_obj, + ) + except Exception as error: # noqa: BLE001 # provider exceptions are normalized below + status_code: Final = getattr(error, "status_code", 500) + error_headers = getattr(error, "headers", None) + error_text = getattr(error, "text", str(error)) + error_response: Final = getattr(error, "response", None) + if error_headers is None and error_response: + error_headers = getattr(error_response, "headers", None) + if error_response and hasattr(error_response, "text"): + error_text = getattr(error_response, "text", error_text) + raise AnthropicError(message=error_text, status_code=status_code, headers=error_headers) + return config.transform_response( + model=model, + raw_response=response, + model_response=model_response, + logging_obj=logging_obj, + api_key=api_key, + request_data=dict(request_data), # mutable-ok: response transform mutates request metadata + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + encoding=encoding, + json_mode=json_mode, + ) + # The Rust core owns the whole call for the subset it accepts, so ask # before transforming: whichever path runs emits pre_call exactly once. # `get_config` merges the class-level defaults (Anthropic's required @@ -466,7 +509,12 @@ class AnthropicChatCompletion(BaseLLM): on_response=log_rust_post_call, python_fallback=python_fallback, ) - rust_response: Final = rust_chat_completions_bridge.chat_completions( + + def sync_python_fallback() -> ModelResponse: + fallback_headers, fallback_data = build_request() + return sync_python_path(fallback_headers, fallback_data) + + return rust_chat_completions_bridge.chat_completions_or_fallback( model=model, messages=messages, optional_params=rust_optional_params, @@ -477,26 +525,21 @@ class AnthropicChatCompletion(BaseLLM): extra_headers=headers, timeout=timeout, on_response=log_rust_post_call, + python_fallback=sync_python_fallback, ) - if rust_response is not None: - return rust_response headers, data = build_request() ## LOGGING - # Reaching here with `serves_via_rust` set means the Rust attempt - # declined at call time, before the provider was called, and already - # logged this request. That is the same attempt continuing. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, - api_key=api_key, - additional_args={ - "complete_input_dict": data, - "api_base": api_base, - "headers": headers, - }, - ) + logging_obj.pre_call( + input=messages, + api_key=api_key, + additional_args={ # mutable-ok: logging callback contract requires a mutable dict + "complete_input_dict": data, + "api_base": api_base, + "headers": headers, + }, + ) print_verbose(f"_is_function_call: {_is_function_call}") if acompletion is True: if ( @@ -579,48 +622,7 @@ class AnthropicChatCompletion(BaseLLM): _response_headers=process_anthropic_headers(headers), ) - else: - if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client(params={"timeout": timeout}) - else: - client = client - - try: - response: Final = client.post( - api_base, - headers=headers, - data=json.dumps(data), - timeout=timeout, - logging_obj=logging_obj, - ) - except Exception as e: - status_code: Final = getattr(e, "status_code", 500) - error_headers = getattr(e, "headers", None) - error_text = getattr(e, "text", str(e)) - error_response: Final[object] = getattr(e, "response", None) - if error_headers is None and error_response: - error_headers = getattr(error_response, "headers", None) - if error_response and hasattr(error_response, "text"): - error_text = getattr(error_response, "text", error_text) - raise AnthropicError( - message=error_text, - status_code=status_code, - headers=error_headers, - ) - - return config.transform_response( - model=model, - raw_response=response, - model_response=model_response, - logging_obj=logging_obj, - api_key=api_key, - request_data=data, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - encoding=encoding, - json_mode=json_mode, - ) + return sync_python_path(headers, data) def embedding(self): # logic for parsing in - calling - parsing out model embedding calls diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 7d5f99ca893..7d6a4d3e6a1 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -383,6 +383,91 @@ class BedrockConverseLLM(BaseAWSLLM): # Filter beta headers in HTTP headers before making the request headers = update_headers_with_filtered_beta(headers=headers, provider="bedrock_converse") + def sync_python_path(*, skip_pre_call_logging: bool) -> ModelResponse | CustomStreamWrapper: + request_data: Final = litellm.AmazonConverseConfig()._transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=extra_headers, + ) + serialized_data: Final = json.dumps(request_data) + prepped: Final = self.get_request_headers( + credentials=credentials, + aws_region_name=aws_region_name, + extra_headers=extra_headers, + endpoint_url=proxy_endpoint_url, + data=serialized_data, + headers=headers, + api_key=api_key, + ) + if not skip_pre_call_logging: + logging_obj.pre_call( + input=messages, + api_key="", + additional_args={ # mutable-ok: logging callback contract requires a mutable dict + "complete_input_dict": serialized_data, + "api_base": proxy_endpoint_url, + "headers": prepped.headers, + }, + ) + request_timeout: Final = ( + httpx.Timeout(timeout) if isinstance(timeout, float) or isinstance(timeout, int) else timeout + ) + request_client: Final = ( + _get_httpx_client( # mutable-ok: HTTP client factory requires a mutable options dict + {"timeout": request_timeout} if request_timeout is not None else {} + ) + if client is None or isinstance(client, AsyncHTTPHandler) + else client + ) + if stream is not None and stream is True: + completion_stream, response_headers = make_sync_call( + client=request_client if isinstance(request_client, HTTPHandler) else None, + api_base=proxy_endpoint_url, + headers=prepped.headers, + data=serialized_data, + model=model, + messages=messages, + logging_obj=logging_obj, + json_mode=json_mode, + fake_stream=fake_stream, + stream_chunk_size=stream_chunk_size, + ) + return CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="bedrock", + logging_obj=logging_obj, + _response_headers=response_headers, + ) + try: + response: Final = request_client.post( + url=proxy_endpoint_url, + headers=prepped.headers, + data=serialized_data, + logging_obj=logging_obj, + ) + response.raise_for_status() + except httpx.HTTPStatusError as error: + raise BedrockError(status_code=error.response.status_code, message=error.response.text) + except httpx.TimeoutException: + raise BedrockError(status_code=408, message="Timeout error occurred.") + transformed: Final = litellm.AmazonConverseConfig()._transform_response( + model=model, + response=response, + model_response=model_response, + stream=stream if isinstance(stream, bool) else False, + logging_obj=logging_obj, + api_key="", + data=serialized_data, + messages=messages, + optional_params=optional_params, + encoding=encoding, + ) + transformed.set_provider_response_headers(response.headers) + return transformed + # The Rust core owns the whole call for the subset it accepts. Ask # before transforming so whichever path runs emits pre_call once, and # hand down the credentials, region and endpoint this handler already @@ -449,7 +534,7 @@ class BedrockConverseLLM(BaseAWSLLM): skip_pre_call_logging=True, ), ) - rust_response: Final = rust_chat_completions_bridge.chat_completions( + return rust_chat_completions_bridge.chat_completions_or_fallback( model=model, messages=messages, optional_params=rust_optional_params, @@ -460,9 +545,8 @@ class BedrockConverseLLM(BaseAWSLLM): extra_headers=headers, timeout=timeout, on_response=log_rust_post_call, + python_fallback=lambda: sync_python_path(skip_pre_call_logging=True), ) - if rust_response is not None: - return rust_response ### ROUTING (ASYNC, STREAMING, SYNC) if acompletion: @@ -508,103 +592,4 @@ class BedrockConverseLLM(BaseAWSLLM): api_key=api_key, ) - ## TRANSFORMATION ## - - _data: Final = litellm.AmazonConverseConfig()._transform_request( - model=model, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - headers=extra_headers, - ) - data: Final = json.dumps(_data) - - prepped: Final = self.get_request_headers( - credentials=credentials, - aws_region_name=aws_region_name, - extra_headers=extra_headers, - endpoint_url=proxy_endpoint_url, - data=data, - headers=headers, - api_key=api_key, - ) - - ## LOGGING - # Reaching here with `serves_via_rust` set means the synchronous Rust - # attempt declined at call time, before the provider was called, and - # already logged this request. That is the same attempt continuing. - # The asynchronous branch above returns before this point, and hands - # its own fallback `skip_pre_call_logging=True` for the same reason. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) - if client is None or isinstance(client, AsyncHTTPHandler): - _params: Final = {} - if timeout is not None: - if isinstance(timeout, float) or isinstance(timeout, int): - timeout = httpx.Timeout(timeout) - _params["timeout"] = timeout - client = _get_httpx_client(_params) - else: - client = client - - if stream is not None and stream is True: - completion_stream, response_headers = make_sync_call( - client=(client if client is not None and isinstance(client, HTTPHandler) else None), - api_base=proxy_endpoint_url, - headers=prepped.headers, - data=data, - model=model, - messages=messages, - logging_obj=logging_obj, - json_mode=json_mode, - fake_stream=fake_stream, - stream_chunk_size=stream_chunk_size, - ) - streaming_response: Final = CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="bedrock", - logging_obj=logging_obj, - _response_headers=response_headers, - ) - - return streaming_response - - ### COMPLETION - - try: - response: Final = client.post( - url=proxy_endpoint_url, - headers=prepped.headers, - data=data, - logging_obj=logging_obj, - ) - response.raise_for_status() - except httpx.HTTPStatusError as err: - error_code: Final = err.response.status_code - raise BedrockError(status_code=error_code, message=err.response.text) - except httpx.TimeoutException: - raise BedrockError(status_code=408, message="Timeout error occurred.") - - sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response( - model=model, - response=response, - model_response=model_response, - stream=stream if isinstance(stream, bool) else False, - logging_obj=logging_obj, - api_key="", - data=data, - messages=messages, - optional_params=optional_params, - encoding=encoding, - ) - sync_transformed_response.set_provider_response_headers(response.headers) - return sync_transformed_response + return sync_python_path(skip_pre_call_logging=False) diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index a4c529f7deb..25b0e1a17e7 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -7,11 +7,11 @@ from litellm.rust_bridge.loader import get_native_bridge BindingT = TypeVar("BindingT") -class _Unset: +class Unset: pass -_UNSET: Final = _Unset() +UNSET: Final = Unset() class NativeBinding(Generic[BindingT]): @@ -19,10 +19,10 @@ class NativeBinding(Generic[BindingT]): def __init__(self, attribute: str) -> None: self._attribute: Final = attribute - self._override: BindingT | None | _Unset = _UNSET + self._override: BindingT | None | Unset = UNSET def load(self) -> BindingT | None: - if not isinstance(self._override, _Unset): + if not isinstance(self._override, Unset): return self._override native: Final = get_native_bridge() if native is None: @@ -33,7 +33,15 @@ class NativeBinding(Generic[BindingT]): self._override = value def reset(self) -> None: - self._override = _UNSET + self._override = UNSET + + def update(self, value: BindingT | None | Unset) -> None: + if isinstance(value, Unset): + return + if value is None: + self.reset() + else: + self.override(value) def native_exception_types() -> tuple[type[BaseException], type[BaseException]] | None: diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index c599667ab17..efdb6eaa9c5 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -14,20 +14,29 @@ from __future__ import annotations import json from collections.abc import Awaitable, Callable, Mapping, Sequence -from dataclasses import dataclass from typing import TYPE_CHECKING, Final, Protocol import httpx from pydantic import TypeAdapter, ValidationError from litellm._logging import verbose_logger -from litellm.exceptions import APIError from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( convert_to_model_response_object, ) from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned +from litellm.rust_bridge.bindings import UNSET, NativeBinding, Unset from litellm.rust_bridge.configuration import rust_enabled -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + FallbackMode, + RustDeclined, + RustHandled, + aattempt, + ainvoke, + attempt, + identity, + invoke, +) from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.utils import ModelResponse @@ -42,9 +51,6 @@ RUST_CHAT_COMPLETIONS_PROVIDERS: Final = frozenset({"anthropic", "bedrock"}) # rather than narrowing an unparameterized `Mapping` and typing the result Any. _LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object]) -RUST_RESPONSE_HEADER: Final = "x-litellm-rust" - - class RustChatCompletions(Protocol): def __call__( self, @@ -126,67 +132,34 @@ def response_logger( return log -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustChatCompletionsState: - chat_completions: RustChatCompletions | None = None - achat_completions: RustAchatCompletions | None = None - decline: RustChatCompletionsDecline | None = None - - -_STATE: Final[_RustChatCompletionsState] = _RustChatCompletionsState() +_CHAT_COMPLETIONS: Final = NativeBinding[RustChatCompletions]("chat_completions") +_ACHAT_COMPLETIONS: Final = NativeBinding[RustAchatCompletions]("achat_completions") +_DECLINE: Final = NativeBinding[RustChatCompletionsDecline]("chat_completions_decline") def set_rust_chat_completions( *, - chat_completions: RustChatCompletions | None | _Unset = _UNSET, - achat_completions: RustAchatCompletions | None | _Unset = _UNSET, - decline: RustChatCompletionsDecline | None | _Unset = _UNSET, + chat_completions: RustChatCompletions | None | Unset = UNSET, + achat_completions: RustAchatCompletions | None | Unset = UNSET, + decline: RustChatCompletionsDecline | None | Unset = UNSET, ) -> None: """Inject the native callables, so tests can supply a double instead of patching module attributes.""" - if not isinstance(chat_completions, _Unset): - _STATE.chat_completions = chat_completions - if not isinstance(achat_completions, _Unset): - _STATE.achat_completions = achat_completions - if not isinstance(decline, _Unset): - _STATE.decline = decline + _CHAT_COMPLETIONS.update(chat_completions) + _ACHAT_COMPLETIONS.update(achat_completions) + _DECLINE.update(decline) def load_rust_chat_completions() -> RustChatCompletions | None: - if _STATE.chat_completions is not None: - return _STATE.chat_completions - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - loaded: RustChatCompletions | None = getattr(native_bridge, "chat_completions", None) - return loaded + return _CHAT_COMPLETIONS.load() def load_rust_achat_completions() -> RustAchatCompletions | None: - if _STATE.achat_completions is not None: - return _STATE.achat_completions - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - loaded: RustAchatCompletions | None = getattr(native_bridge, "achat_completions", None) - return loaded + return _ACHAT_COMPLETIONS.load() def _load_rust_decline() -> RustChatCompletionsDecline | None: - if _STATE.decline is not None: - return _STATE.decline - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - loaded: RustChatCompletionsDecline | None = getattr(native_bridge, "chat_completions_decline", None) - return loaded + return _DECLINE.load() def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool: @@ -256,76 +229,35 @@ def rust_chat_completions_accepts( decline: Final = _load_rust_decline() if decline is None: return False - try: - reason: Final = decline( + gate_result: Final = attempt( + native_call=lambda: decline( model=model, messages=messages, optional_params=optional_params, custom_llm_provider=custom_llm_provider, - ) - except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path + ), + adapt=identity, + context=BridgeErrorContext( + route="chat completions capability check", + provider=custom_llm_provider or "", + model=model, + ), + ) + if isinstance(gate_result, RustDeclined): verbose_logger.debug( - "Rust chat completions gate raised %s; staying on the Python path", - type(rust_error).__name__, + "Rust chat completions declined (%s); using the Python path", + gate_result.reason, ) return False + if not isinstance(gate_result, RustHandled): + return False + reason: Final = gate_result.value if reason is not None: verbose_logger.debug("Rust chat completions declined (%s); using the Python path", reason) return False return True -def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None: - """`(declined, upstream_failed)` from the native module, or None when absent.""" - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - declined: Final = getattr(native_bridge, "RustBridgeDeclined", None) - upstream: Final = getattr(native_bridge, "RustUpstreamError", None) - if declined is None or upstream is None: - return None - return declined, upstream - - -def _reraise_or_decline( - rust_error: BaseException, - *, - model: str, - custom_llm_provider: str | None, -) -> None: - """Re-raise a failure the provider already saw, or return so the caller declines. - - A request that never reached the provider is safe to serve on the Python - path. One that did is not: the provider has already done the work, so a - second attempt bills for it twice. Those surface as an `APIError` carrying - the upstream status, which LiteLLM's exception mapping already understands. - """ - exceptions: Final = _rust_bridge_exceptions() - if exceptions is None: - verbose_logger.debug( - "Rust chat completions bridge raised %s; falling back to Python path", - type(rust_error).__name__, - ) - return - declined, upstream_failed = exceptions - if isinstance(rust_error, upstream_failed): - args: Final = rust_error.args - status: Final = args[0] if args else 0 - message: Final = args[1] if len(args) > 1 else "" - raise APIError( - status_code=int(status) or 500, - message=f"litellm rust chat completions: {message}", - llm_provider=custom_llm_provider or "", - model=model, - ) - if not isinstance(rust_error, declined): - raise rust_error - verbose_logger.debug( - "Rust chat completions declined before calling the provider (%s); using the Python path", - rust_error, - ) - - def _build_model_response( rust_response: Mapping[str, object], model_response: ModelResponse, @@ -333,7 +265,6 @@ def _build_model_response( built: Final = convert_to_model_response_object( response_object=dict(rust_response), # mutable-ok: the converter takes a real dict and rewrites it model_response_object=model_response, - hidden_params={"additional_headers": {RUST_RESPONSE_HEADER: "true"}}, # mutable-ok: rewritten by the converter ) if not isinstance(built, ModelResponse): raise TypeError(f"expected a ModelResponse from the rust path, got {type(built).__name__}") @@ -354,10 +285,10 @@ def chat_completions( on_response: ResponseObserver, ) -> ModelResponse | None: rust_chat_completions: Final = load_rust_chat_completions() - if rust_chat_completions is None: - return None - try: - rust_response: Final = rust_chat_completions( + native_call: Final = ( + None + if rust_chat_completions is None + else lambda: rust_chat_completions( model=model, messages=messages, optional_params=optional_params, @@ -367,11 +298,18 @@ def chat_completions( extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), ) - except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw - _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) - return None - on_response(rust_response) - return _build_model_response(rust_response, model_response) + ) + + def adapt(rust_response: Mapping[str, object]) -> ModelResponse: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + result: Final = attempt( + native_call=native_call, + adapt=adapt, + context=BridgeErrorContext(route="chat completions", provider=custom_llm_provider or "", model=model), + ) + return result.value if isinstance(result, RustHandled) else None async def achat_completions( @@ -388,10 +326,10 @@ async def achat_completions( on_response: ResponseObserver, ) -> ModelResponse | None: rust_achat_completions: Final = load_rust_achat_completions() - if rust_achat_completions is None: - return None - try: - rust_response: Final = await rust_achat_completions( + native_call: Final = ( + None + if rust_achat_completions is None + else lambda: rust_achat_completions( model=model, messages=messages, optional_params=optional_params, @@ -401,11 +339,61 @@ async def achat_completions( extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), ) - except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw - _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) - return None - on_response(rust_response) - return _build_model_response(rust_response, model_response) + ) + + def adapt(rust_response: Mapping[str, object]) -> ModelResponse: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + result: Final = await aattempt( + native_call=native_call, + adapt=adapt, + context=BridgeErrorContext(route="chat completions", provider=custom_llm_provider or "", model=model), + ) + return result.value if isinstance(result, RustHandled) else None + + +def chat_completions_or_fallback( + *, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object], + model_response: ModelResponse, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout: float | httpx.Timeout | None, + on_response: ResponseObserver, + python_fallback: Callable[[], object], +) -> object: + rust_chat_completions: Final = load_rust_chat_completions() + native_call: Final = ( + None + if rust_chat_completions is None + else lambda: rust_chat_completions( + model=model, + messages=messages, + optional_params=optional_params, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ) + ) + + def adapt(rust_response: Mapping[str, object]) -> object: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + return invoke( + native_call=native_call, + fallback=python_fallback, + adapt=adapt, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route="chat completions", provider=custom_llm_provider or "", model=model), + ) async def achat_completions_or_fallback( @@ -422,26 +410,30 @@ async def achat_completions_or_fallback( on_response: ResponseObserver, python_fallback: Callable[[], Awaitable[object]], ) -> object: - """Await the Rust path, falling back to the caller's own Python path when - the bridge is unavailable or the call fails. - - The caller supplies the fallback, so the bridge stays free of provider - dispatch. This exists because a caller that dispatches asynchronously has - already returned a coroutine by the time a Rust failure surfaces, and so - cannot fall back on its own. - """ - response: Final = await achat_completions( - model=model, - messages=messages, - optional_params=optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout=timeout, - on_response=on_response, + rust_achat_completions: Final = load_rust_achat_completions() + native_call: Final = ( + None + if rust_achat_completions is None + else lambda: rust_achat_completions( + model=model, + messages=messages, + optional_params=optional_params, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ) + ) + + def adapt(rust_response: Mapping[str, object]) -> object: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + return await ainvoke( + native_call=native_call, + fallback=python_fallback, + adapt=adapt, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route="chat completions", provider=custom_llm_provider or "", model=model), ) - if response is not None: - return response - return await python_fallback() diff --git a/litellm/rust_bridge/configuration.py b/litellm/rust_bridge/configuration.py index a8eecd022d6..8ae208c836a 100644 --- a/litellm/rust_bridge/configuration.py +++ b/litellm/rust_bridge/configuration.py @@ -4,6 +4,8 @@ import os import warnings from typing import TYPE_CHECKING, Final +from litellm.rust_bridge.bindings import UNSET, Unset + if TYPE_CHECKING: from litellm.rust_bridge.messages import RustAmessages, RustMessages from litellm.rust_bridge.ocr import RustAocr, RustOcr @@ -17,13 +19,6 @@ _GLOBAL_ENV_NAME: Final = "LITELLM_RUST" _LEGACY_OCR_ENV_NAME: Final = "LITELLM_USE_RUST_OCR" -class _Unset: - pass - - -_UNSET: Final = _Unset() - - class _RustConfiguration: def __init__(self) -> None: self.override: bool | None = None @@ -107,13 +102,13 @@ def reset_rust_configuration() -> None: def use_litellm_rust( enabled: bool = True, *, - ocr: RustOcr | None | _Unset = _UNSET, - aocr: RustAocr | None | _Unset = _UNSET, - messages: RustMessages | None | _Unset = _UNSET, - amessages: RustAmessages | None | _Unset = _UNSET, - responses_websocket: type[RustResponsesWebSocketConnection] | None | _Unset = _UNSET, - transcription: RustTranscription | None | _Unset = _UNSET, - atranscription: RustAtranscription | None | _Unset = _UNSET, + ocr: RustOcr | None | Unset = UNSET, + aocr: RustAocr | None | Unset = UNSET, + messages: RustMessages | None | Unset = UNSET, + amessages: RustAmessages | None | Unset = UNSET, + responses_websocket: type[RustResponsesWebSocketConnection] | None | Unset = UNSET, + transcription: RustTranscription | None | Unset = UNSET, + atranscription: RustAtranscription | None | Unset = UNSET, ) -> None: """Set the process override for optional Rust paths. @@ -121,7 +116,7 @@ def use_litellm_rust( """ _CONFIGURATION.override = enabled bindings: Final = (ocr, aocr, messages, amessages, responses_websocket, transcription, atranscription) - if all(isinstance(binding, _Unset) for binding in bindings): + if all(isinstance(binding, Unset) for binding in bindings): return warnings.warn( "Injecting Rust bridge implementations through use_litellm_rust() is deprecated; " @@ -130,28 +125,28 @@ def use_litellm_rust( stacklevel=2, ) - if not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset): + if not isinstance(ocr, Unset) or not isinstance(aocr, Unset): from litellm.rust_bridge.ocr import set_rust_ocr - if not isinstance(ocr, _Unset): + if not isinstance(ocr, Unset): set_rust_ocr(ocr=ocr) - if not isinstance(aocr, _Unset): + if not isinstance(aocr, Unset): set_rust_ocr(aocr=aocr) - if not isinstance(messages, _Unset) or not isinstance(amessages, _Unset): + if not isinstance(messages, Unset) or not isinstance(amessages, Unset): from litellm.rust_bridge.messages import set_rust_messages - if not isinstance(messages, _Unset): + if not isinstance(messages, Unset): set_rust_messages(messages=messages) - if not isinstance(amessages, _Unset): + if not isinstance(amessages, Unset): set_rust_messages(amessages=amessages) - if not isinstance(responses_websocket, _Unset): + if not isinstance(responses_websocket, Unset): from litellm.rust_bridge.responses_websocket import set_rust_responses_websocket set_rust_responses_websocket(connection=responses_websocket) - if not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset): + if not isinstance(transcription, Unset) or not isinstance(atranscription, Unset): from litellm.rust_bridge.transcription import configure_rust_transcription - if not isinstance(transcription, _Unset): + if not isinstance(transcription, Unset): configure_rust_transcription(transcription=transcription) - if not isinstance(atranscription, _Unset): + if not isinstance(atranscription, Unset): configure_rust_transcription(atranscription=atranscription) diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index a4890a626d8..e0bdef20dfe 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -7,7 +7,7 @@ from typing import Final, Protocol import httpx -from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.bindings import UNSET, NativeBinding, Unset from litellm.rust_bridge.runtime import ( BridgeErrorContext, FallbackMode, @@ -51,22 +51,13 @@ _MESSAGES: Final = NativeBinding[RustMessages]("messages") _AMESSAGES: Final = NativeBinding[RustAmessages]("amessages") -class _Unset: - pass - - -_UNSET: Final = _Unset() - - def set_rust_messages( *, - messages: RustMessages | None | _Unset = _UNSET, - amessages: RustAmessages | None | _Unset = _UNSET, + messages: RustMessages | None | Unset = UNSET, + amessages: RustAmessages | None | Unset = UNSET, ) -> None: - if not isinstance(messages, _Unset): - _MESSAGES.reset() if messages is None else _MESSAGES.override(messages) - if not isinstance(amessages, _Unset): - _AMESSAGES.reset() if amessages is None else _AMESSAGES.override(amessages) + _MESSAGES.update(messages) + _AMESSAGES.update(amessages) def load_rust_messages() -> RustMessages | None: diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index cbe444d44a6..355254fb237 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -3,11 +3,20 @@ from __future__ import annotations from collections.abc import Awaitable -from typing import Final, Protocol, cast +from typing import Final, Protocol import httpx from litellm.rust_bridge import configuration as _configuration +from litellm.rust_bridge.bindings import UNSET, NativeBinding, Unset +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + FallbackMode, + ainvoke, + async_none, + identity, + invoke, +) from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds rust_ocr_enabled = _configuration.rust_ocr_enabled @@ -44,49 +53,25 @@ class RustAocr(Protocol): raise NotImplementedError -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -_rust_ocr_impl: RustOcr | None = None -_rust_aocr_impl: RustAocr | None = None +_OCR: Final = NativeBinding[RustOcr]("ocr") +_AOCR: Final = NativeBinding[RustAocr]("aocr") def set_rust_ocr( *, - ocr: RustOcr | None | _Unset = _UNSET, - aocr: RustAocr | None | _Unset = _UNSET, + ocr: RustOcr | None | Unset = UNSET, + aocr: RustAocr | None | Unset = UNSET, ) -> None: - global _rust_ocr_impl, _rust_aocr_impl - if not isinstance(ocr, _Unset): - _rust_ocr_impl = ocr - if not isinstance(aocr, _Unset): - _rust_aocr_impl = aocr + _OCR.update(ocr) + _AOCR.update(aocr) def load_rust_ocr() -> RustOcr | None: - if _rust_ocr_impl is not None: - return _rust_ocr_impl - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustOcr, native_bridge.ocr) + return _OCR.load() def load_rust_aocr() -> RustAocr | None: - if _rust_aocr_impl is not None: - return _rust_aocr_impl - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustAocr, getattr(native_bridge, "aocr", None)) + return _AOCR.load() def ocr( @@ -101,17 +86,26 @@ def ocr( timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_ocr: Final = load_rust_ocr() - if rust_ocr is None: - return None - return rust_ocr( - model=model, - document=document, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=_timeout_to_seconds(timeout), + native_call: Final = ( + None + if rust_ocr is None + else lambda: rust_ocr( + model=model, + document=document, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout_seconds=_timeout_to_seconds(timeout), + ) + ) + return invoke( + native_call=native_call, + fallback=lambda: None, + adapt=identity, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route="ocr", provider=custom_llm_provider or "", model=model), ) @@ -127,15 +121,24 @@ async def aocr( timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_aocr: Final = load_rust_aocr() - if rust_aocr is None: - return None - return await rust_aocr( - model=model, - document=document, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=_timeout_to_seconds(timeout), + native_call: Final = ( + None + if rust_aocr is None + else lambda: rust_aocr( + model=model, + document=document, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout_seconds=_timeout_to_seconds(timeout), + ) + ) + return await ainvoke( + native_call=native_call, + fallback=async_none, + adapt=identity, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route="ocr", provider=custom_llm_provider or "", model=model), ) diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 0634867af1c..4d505a757c9 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -2,13 +2,19 @@ from __future__ import annotations -from dataclasses import dataclass from typing import Final, Protocol import httpx from websockets.exceptions import ConnectionClosedOK -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.bindings import UNSET, NativeBinding, Unset +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + FallbackMode, + acall, + ainvoke, + async_none, +) from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -30,56 +36,46 @@ class RustResponsesWebSocketConnection(Protocol): ) -> RustResponsesWebSocket: ... -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustResponsesWebSocketState: - connection: RustResponsesWebSocketConnection | None = None - - -_STATE: Final[_RustResponsesWebSocketState] = _RustResponsesWebSocketState() +_CONNECTION: Final = NativeBinding[type[RustResponsesWebSocketConnection]]( + "ResponsesWebSocketConnection" +) def set_rust_responses_websocket( *, - connection: RustResponsesWebSocketConnection | None | _Unset = _UNSET, + connection: type[RustResponsesWebSocketConnection] | None | Unset = UNSET, ) -> None: - if not isinstance(connection, _Unset): - _STATE.connection = connection + _CONNECTION.update(connection) -def load_rust_responses_websocket() -> RustResponsesWebSocketConnection | None: - if _STATE.connection is not None: - return _STATE.connection - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - connection_type: Final[RustResponsesWebSocketConnection | None] = getattr( - native_bridge, "ResponsesWebSocketConnection", None - ) - return connection_type +def load_rust_responses_websocket() -> type[RustResponsesWebSocketConnection] | None: + return _CONNECTION.load() class _ConnectionAdapter: def __init__(self, connection: RustResponsesWebSocket): - self._connection: Final[RustResponsesWebSocket] = connection + self._connection: Final = connection async def send(self, text: str) -> None: - await self._connection.send_text(text) + await acall( + lambda: self._connection.send_text(text), + BridgeErrorContext(route="responses websocket", provider="openai", model=""), + ) async def recv(self) -> str: - message: Final = await self._connection.recv_text() + message: Final = await acall( + self._connection.recv_text, + BridgeErrorContext(route="responses websocket", provider="openai", model=""), + ) if message is None: raise ConnectionClosedOK(None, None) return message async def close(self) -> None: - await self._connection.close() + await acall( + self._connection.close, + BridgeErrorContext(route="responses websocket", provider="openai", model=""), + ) async def connect( @@ -89,14 +85,19 @@ async def connect( timeout: float | httpx.Timeout | None, ) -> _ConnectionAdapter | None: connection_type: Final = load_rust_responses_websocket() - if connection_type is None: - return None - try: - connection: Final = await connection_type.connect( + native_call: Final = ( + None + if connection_type is None + else lambda: connection_type.connect( url=url, headers=headers, timeout_seconds=timeout_to_seconds(timeout), ) - except Exception: # noqa: BLE001 # bridge failures must fall back to Python - return None - return _ConnectionAdapter(connection) + ) + return await ainvoke( + native_call=native_call, + fallback=async_none, + adapt=_ConnectionAdapter, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route="responses websocket", provider="openai", model=""), + ) diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 3d71f6f8a50..bf1e3fb60d8 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -1,11 +1,19 @@ from __future__ import annotations from collections.abc import Awaitable -from dataclasses import dataclass -from typing import Final, Protocol, cast +from typing import Final, Protocol import httpx +from litellm.rust_bridge.bindings import UNSET, NativeBinding, Unset +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + FallbackMode, + ainvoke, + async_none, + identity, + invoke, +) from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -39,62 +47,27 @@ class RustAtranscription(Protocol): raise NotImplementedError -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass -class _RustTranscriptionState: - transcription: RustTranscription | None = None - atranscription: RustAtranscription | None = None - - -_STATE: Final = _RustTranscriptionState() +_TRANSCRIPTION: Final = NativeBinding[RustTranscription]("transcription") +_ATRANSCRIPTION: Final = NativeBinding[RustAtranscription]("atranscription") def configure_rust_transcription( enabled: bool = True, *, - transcription: RustTranscription | None | _Unset = _UNSET, - atranscription: RustAtranscription | None | _Unset = _UNSET, + transcription: RustTranscription | None | Unset = UNSET, + atranscription: RustAtranscription | None | Unset = UNSET, ) -> None: - if not isinstance(transcription, _Unset): - _STATE.transcription = transcription - if not isinstance(atranscription, _Unset): - _STATE.atranscription = atranscription + _ = enabled + _TRANSCRIPTION.update(transcription) + _ATRANSCRIPTION.update(atranscription) def load_rust_transcription() -> RustTranscription | None: - if _STATE.transcription is not None: - return _STATE.transcription - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - return ( - None - if native_bridge is None - else cast( # cast-ok: native extension protocol is runtime-defined - RustTranscription, getattr(native_bridge, "transcription", None) - ) - ) + return _TRANSCRIPTION.load() def load_rust_atranscription() -> RustAtranscription | None: - if _STATE.atranscription is not None: - return _STATE.atranscription - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - return ( - None - if native_bridge is None - else cast( # cast-ok: native extension protocol is runtime-defined - RustAtranscription, getattr(native_bridge, "atranscription", None) - ) - ) + return _ATRANSCRIPTION.load() def transcription( @@ -109,17 +82,26 @@ def transcription( timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_transcription: Final = load_rust_transcription() - if rust_transcription is None: - return None - return rust_transcription( - model=model, - audio=audio, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_to_seconds(timeout), + native_call: Final = ( + None + if rust_transcription is None + else lambda: rust_transcription( + model=model, + audio=audio, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout_seconds=timeout_to_seconds(timeout), + ) + ) + return invoke( + native_call=native_call, + fallback=lambda: None, + adapt=identity, + mode=FallbackMode.RUST_REQUIRED, + context=BridgeErrorContext(route="audio transcription", provider=custom_llm_provider or "", model=model), ) @@ -135,15 +117,24 @@ async def atranscription( timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_atranscription: Final = load_rust_atranscription() - if rust_atranscription is None: - return None - return await rust_atranscription( - model=model, - audio=audio, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_to_seconds(timeout), + native_call: Final = ( + None + if rust_atranscription is None + else lambda: rust_atranscription( + model=model, + audio=audio, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout_seconds=timeout_to_seconds(timeout), + ) + ) + return await ainvoke( + native_call=native_call, + fallback=async_none, + adapt=identity, + mode=FallbackMode.RUST_REQUIRED, + context=BridgeErrorContext(route="audio transcription", provider=custom_llm_provider or "", model=model), ) diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index 90a5c8c35f3..362d3a5f5a0 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -10,7 +10,7 @@ import pytest import litellm from litellm.exceptions import APIError from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler -from litellm.rust_bridge import configuration +from litellm.rust_bridge import bindings, configuration from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) @@ -160,11 +160,7 @@ def test_load_rust_amessages_returns_injected_impl(): def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): - monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), - "get_native_bridge", - lambda: None, - ) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) litellm.use_litellm_rust(True) assert rust_messages.load_rust_messages() is None result = rust_messages.messages( @@ -448,11 +444,7 @@ async def test_fake_stream_wraps_rust_response_as_anthropic_sse(): @pytest.mark.asyncio async def test_gate_falls_back_when_bridge_unavailable(monkeypatch): - monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), - "get_native_bridge", - lambda: None, - ) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) litellm.use_litellm_rust(True) response = await _gate() diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py index 043537f8c1f..1d9a426283a 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py @@ -10,6 +10,7 @@ import litellm from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.llms.anthropic.chat.handler import ModelResponseIterator, make_call from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.rust_bridge import bindings as bridge_bindings from litellm.types.llms.openai import ( ChatCompletionToolCallChunk, ChatCompletionToolCallFunctionChunk, @@ -2318,14 +2319,13 @@ class TestRustChatCompletionsHook: bridge.set_rust_chat_completions(decline=gate, chat_completions=native) return seen - def test_rust_true_serves_the_call_and_stamps_the_header(self): + def test_rust_true_serves_the_call(self): from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion seen = self._inject() response = AnthropicChatCompletion().completion(**self._completion_kwargs()) assert response.choices[0].message.content == "hello from rust" - assert response._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} assert len(seen["call"]) == 1 def test_the_core_receives_the_untranslated_openai_messages(self): @@ -2466,7 +2466,7 @@ class TestRustChatCompletionsHook: def declining_native(**_kwargs): raise _Declined("blank message text") - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr(bridge_bindings, "get_native_bridge", lambda: _FakeNative()) bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) @@ -2498,7 +2498,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr(bridge_bindings, "get_native_bridge", lambda: _FakeNative()) async def declining_native(**_kwargs): raise _Declined("blank message text") @@ -2540,7 +2540,6 @@ class TestRustChatCompletionsHook: ) assert result.choices[0].message.content == "hello from rust" - assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} assert not python_call.called @@ -2558,7 +2557,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr(bridge_bindings, "get_native_bridge", lambda: _FakeNative()) def declining_native(**_kwargs): raise _Declined("blank message text") diff --git a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py index 21e3239f623..8ac88a8172b 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py @@ -10,10 +10,11 @@ from unittest.mock import MagicMock, patch import httpx import pytest - from botocore.credentials import Credentials + from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.rust_bridge import bindings as bridge_bindings from litellm.rust_bridge import chat_completions as bridge from litellm.types.utils import ModelResponse @@ -111,12 +112,11 @@ def _recording_logging_obj(): return logging_obj, calls -def test_rust_true_serves_the_call_and_stamps_the_header(): +def test_rust_true_serves_the_call(): seen = _inject() response = _run() assert response.choices[0].message.content == "hello from rust" - assert response._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} assert len(seen["call"]) == 1 @@ -204,7 +204,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr(bridge_bindings, "get_native_bridge", lambda: _FakeNative()) async def declining_native(**_kwargs): raise _Declined("blank message text") @@ -254,7 +254,6 @@ async def test_the_async_path_serves_the_rust_response_without_the_fallback(): ) assert result.choices[0].message.content == "hello from rust" - assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} assert not python_call.called @@ -281,7 +280,7 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines(): return ModelResponse() with ( - patch.object(bridge, "get_native_bridge", lambda: _FakeNative()), + patch.object(bridge_bindings, "get_native_bridge", lambda: _FakeNative()), patch.object( BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS ), @@ -388,7 +387,7 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines(): logging_obj = MagicMock() - with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()): + with patch.object(bridge_bindings, "get_native_bridge", lambda: _FakeNative()): bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) @@ -473,7 +472,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(): logging_obj, calls = _recording_logging_obj() - with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()): + with patch.object(bridge_bindings, "get_native_bridge", lambda: _FakeNative()): bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 0764aec7185..e8f3ea3a0a9 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -10,7 +10,9 @@ import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse -from litellm.rust_bridge import configuration +from litellm.rust_bridge import bindings, configuration +from litellm.rust_bridge.bindings import UNSET, Unset +from litellm.rust_bridge.ocr import RustAocr, RustOcr # `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` # function onto `litellm.ocr` and shadows the submodule, so import the modules @@ -34,6 +36,16 @@ FAKE_OCR_RESPONSE: dict[str, object] = { } +def _use_test_rust( + enabled: bool = True, + *, + ocr: RustOcr | None | Unset = UNSET, + aocr: RustAocr | None | Unset = UNSET, +) -> None: + rust_bridge.set_rust_ocr(ocr=ocr, aocr=aocr) + litellm.use_litellm_rust(enabled) + + class CapturedException(Exception): pass @@ -228,7 +240,8 @@ def _reset_rust_flag(): def fake_bridge(): """Enable the Rust path with an injected recording bridge (no native wheel).""" bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + rust_bridge.set_rust_ocr(ocr=bridge) + litellm.use_litellm_rust(True) return bridge @@ -236,7 +249,8 @@ def fake_bridge(): def fake_async_bridge(): """Enable the async Rust path with an injected recording bridge.""" bridge = RecordingAsyncBridge() - litellm.use_litellm_rust(True, aocr=bridge) + rust_bridge.set_rust_ocr(aocr=bridge) + litellm.use_litellm_rust(True) return bridge @@ -248,21 +262,16 @@ def test_use_litellm_rust_toggles_flag(): assert rust_bridge.rust_ocr_enabled() is False -def test_env_var_enables_rust_ocr(monkeypatch): - monkeypatch.setenv("LITELLM_USE_RUST_OCR", "1") - with pytest.warns(DeprecationWarning, match="LITELLM_USE_RUST_OCR is deprecated"): - assert rust_bridge.rust_ocr_enabled() is True - - -def test_explicit_false_overrides_process_enable(): +def test_explicit_false_overrides_the_process_switch(): litellm.use_litellm_rust(True) + prepared = build_prepared_request(litellm_params={"rust": False}) - assert ocr_main._rust_ocr_enabled(build_prepared_request(litellm_params={"rust": False})) is False + assert ocr_main._rust_ocr_enabled(prepared) is False def test_load_rust_ocr_returns_injected_impl(): bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) assert rust_bridge.load_rust_ocr() is bridge @@ -306,40 +315,21 @@ def test_native_bridge_available_reflects_loader(monkeypatch): def test_load_rust_aocr_returns_injected_impl(): bridge = RecordingAsyncBridge() - litellm.use_litellm_rust(True, aocr=bridge) + _use_test_rust(True, aocr=bridge) assert rust_bridge.load_rust_aocr() is bridge -def test_toggle_without_ocr_arg_preserves_injected_impl(): - """Regression: routine enable/disable calls must not clobber a prior injection. - - Earlier, ``use_litellm_rust()`` unconditionally assigned the keyword default - of ``None`` to ``_rust_ocr_impl``, silently dropping a custom bridge whenever - a caller toggled the flag without re-passing ``ocr=``. - """ - bridge = RecordingBridge() - async_bridge = RecordingAsyncBridge() - litellm.use_litellm_rust(True, ocr=bridge, aocr=async_bridge) - - litellm.use_litellm_rust(False) - assert rust_bridge.load_rust_ocr() is bridge - assert rust_bridge.load_rust_aocr() is async_bridge - litellm.use_litellm_rust(True) - assert rust_bridge.load_rust_ocr() is bridge - assert rust_bridge.load_rust_aocr() is async_bridge - - def test_explicit_ocr_none_clears_injected_impl(monkeypatch): monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), + bindings, "get_native_bridge", lambda: None, ) bridge = RecordingBridge() async_bridge = RecordingAsyncBridge() - litellm.use_litellm_rust(True, ocr=bridge, aocr=async_bridge) + rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge) - litellm.use_litellm_rust(True, ocr=None, aocr=None) + rust_bridge.set_rust_ocr(ocr=None, aocr=None) assert rust_bridge.load_rust_ocr() is None assert rust_bridge.load_rust_aocr() is None @@ -348,7 +338,7 @@ def test_load_rust_ocr_none_when_extension_absent(monkeypatch): """With no injected impl and no compiled wheel, the loader returns None so the caller degrades to the Python path instead of raising ImportError.""" monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), + bindings, "get_native_bridge", lambda: None, ) @@ -365,7 +355,7 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch): fake_module.ocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined] fake_module.aocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined] monkeypatch.setattr( - importlib.import_module("litellm.rust_bridge"), + bindings, "get_native_bridge", lambda: fake_module, ) @@ -384,7 +374,7 @@ def test_timeout_to_seconds_handles_float_timeout_and_none(): def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) response = rust_bridge.ocr( model="mistral-ocr-latest", document=DOCUMENT, @@ -417,7 +407,7 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): bridge = RecordingAsyncBridge() - litellm.use_litellm_rust(True, aocr=bridge) + _use_test_rust(True, aocr=bridge) response = await rust_bridge.aocr( model="mistral-ocr-maas", document=DOCUMENT, @@ -445,7 +435,7 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): def test_run_rust_ocr_prepares_request_and_wraps_response(): bridge = RecordingBridge() logging_obj = RecordingLogging() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) response = ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -477,7 +467,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request(api_key=None, timeout=None), @@ -489,7 +479,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): def test_run_rust_ocr_prefers_explicit_key_over_resolver(): bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) def _resolver(name: str) -> str | None: raise AssertionError(f"resolver should not be called for {name}") @@ -508,7 +498,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver(): def test_run_rust_ocr_uses_provider_api_key_env_var(): bridge = RecordingBridge() resolver_calls = [] - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) def _resolver(name): resolver_calls.append(name) @@ -530,7 +520,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -556,7 +546,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager(): bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) def _resolver(name: str) -> str | None: return { @@ -579,7 +569,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -596,7 +586,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -616,7 +606,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): def test_run_rust_ocr_runs_pre_call_logging(): logging_obj = RecordingLogging() bridge = RecordingBridge() - litellm.use_litellm_rust(True, ocr=bridge) + _use_test_rust(True, ocr=bridge) ocr_main._run_rust_ocr( prepared_request=build_prepared_request( @@ -703,7 +693,7 @@ def test_ocr_exception_type_uses_resolved_provider_context( return CapturedException("wrapped") monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) - litellm.use_litellm_rust(True, ocr=RaisingBridge()) + _use_test_rust(True, ocr=RaisingBridge()) with pytest.raises(CapturedException): litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test") @@ -748,7 +738,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context( return CapturedException("wrapped") monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) - litellm.use_litellm_rust(True, aocr=RaisingAsyncBridge()) + _use_test_rust(True, aocr=RaisingAsyncBridge()) with pytest.raises(CapturedException): await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") @@ -776,7 +766,7 @@ def test_ocr_passes_default_request_timeout_to_rust(fake_bridge): def test_ocr_does_not_route_to_rust_when_disabled(): """With the flag off, the bridge must not be consulted even if an impl exists.""" bridge = RecordingBridge() - litellm.use_litellm_rust(False, ocr=bridge) + _use_test_rust(False, ocr=bridge) assert rust_bridge.rust_ocr_enabled() is False # The impl stays available for injection, but the disabled flag gates usage, @@ -804,6 +794,25 @@ def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch): assert isinstance(response, OCRResponse) +def test_ocr_unsupported_provider_skips_rust(monkeypatch): + bridge = RecordingBridge() + _use_test_rust(True, ocr=bridge) + + def fake_handler_ocr(**kwargs): + return OCRResponse(pages=[], model="parse-v3", object="ocr") + + monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fake_handler_ocr) + + response = litellm.ocr( + model="reducto/parse-v3", + document={"type": "document_url", "document_url": "reducto://document-id"}, + api_key="test-key", + ) + + assert isinstance(response, OCRResponse) + assert bridge.calls == [] + + def test_ocr_provider_configs_expose_api_key_env_vars(): from litellm.llms.azure_ai.ocr.document_intelligence.transformation import ( AzureDocumentIntelligenceOCRConfig, diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 1233ddf1785..9aa9c4acd3d 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -3,7 +3,7 @@ from __future__ import annotations import pytest from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled -from litellm.rust_bridge import configuration, responses_websocket +from litellm.rust_bridge import bindings, configuration, responses_websocket from litellm.types.router import GenericLiteLLMParams @@ -76,8 +76,7 @@ async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None: @pytest.mark.asyncio async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(responses_websocket, "_STATE", responses_websocket._RustResponsesWebSocketState()) - monkeypatch.setattr(responses_websocket, "get_native_bridge", lambda: None) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) assert ( await responses_websocket.connect( diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index 03921133c77..a7c8570e78d 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -10,8 +10,9 @@ from __future__ import annotations import pytest import litellm +from litellm.exceptions import APIError +from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge import chat_completions as bridge -from litellm.rust_bridge import configuration from litellm.types.utils import ModelResponse RUST_RESPONSE = { @@ -54,7 +55,7 @@ class _FakeNative: def _fake_native_bridge(monkeypatch): """Expose the bridge's exception classes without the compiled extension.""" - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative()) def _hide_native_bridge(monkeypatch): @@ -63,31 +64,23 @@ def _hide_native_bridge(monkeypatch): There is no injection seam for "the .so is absent", so the loader itself is replaced; every other case here uses `set_rust_chat_completions`. """ - monkeypatch.setattr(bridge, "get_native_bridge", lambda: None) + monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) @pytest.fixture(autouse=True) def reset_bridge(): """Every test starts with no injected callables, and leaves none behind.""" - bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) + bridge.set_rust_chat_completions( + chat_completions=None, + achat_completions=None, + decline=lambda **_kwargs: None, + ) configuration.reset_rust_configuration() yield bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) configuration.reset_rust_configuration() -class _RecordingDecline: - """A stand-in for the native gate that records what it was asked.""" - - def __init__(self, reason: str | None = None): - self.reason = reason - self.calls: list[dict] = [] - - def __call__(self, **kwargs): - self.calls.append(kwargs) - return self.reason - - class _RecordingCall: def __init__(self, result=None, error: Exception | None = None): self.result = result if result is not None else dict(RUST_RESPONSE) @@ -106,7 +99,7 @@ class _RecordingAsyncCall(_RecordingCall): return _RecordingCall.__call__(self, **kwargs) -def _accepts(**overrides) -> bool: +def _should_attempt(**overrides) -> bool: kwargs = { "model": "claude-sonnet-4-5", "messages": MESSAGES, @@ -116,52 +109,40 @@ def _accepts(**overrides) -> bool: "stream": None, } kwargs.update(overrides) + kwargs.pop("asynchronous", None) return bridge.rust_chat_completions_accepts(**kwargs) -class TestGate: +class TestEligibility: def test_declines_when_the_deployment_did_not_opt_in(self, monkeypatch): monkeypatch.delenv("LITELLM_RUST", raising=False) - gate = _RecordingDecline() - bridge.set_rust_chat_completions(decline=gate) - assert _accepts(litellm_params={}) is False - assert _accepts(litellm_params=None) is False - assert _accepts(litellm_params={"rust": False}) is False - assert gate.calls == [], "the gate must not be consulted before opt-in" + bridge.set_rust_chat_completions(chat_completions=_RecordingCall()) + assert _should_attempt(litellm_params={}) is False + assert _should_attempt(litellm_params=None) is False + assert _should_attempt(litellm_params={"rust": False}) is False - def test_accepts_when_the_deployment_opted_in_and_the_core_agrees(self, monkeypatch): + def test_attempts_when_the_deployment_opted_in(self, monkeypatch): monkeypatch.delenv("LITELLM_RUST", raising=False) - gate = _RecordingDecline() - bridge.set_rust_chat_completions(decline=gate) - assert _accepts() is True - assert gate.calls[0]["model"] == "claude-sonnet-4-5" - assert gate.calls[0]["custom_llm_provider"] == "anthropic" - - def test_explicit_false_overrides_process_enable(self): - bridge.set_rust_chat_completions(decline=_RecordingDecline()) - configuration.use_litellm_rust(True) - - assert _accepts(litellm_params={"rust": False}) is False - - def test_process_enable_applies_without_request_override(self): - bridge.set_rust_chat_completions(decline=_RecordingDecline()) - configuration.use_litellm_rust(True) - - assert _accepts(litellm_params={}) is True + bridge.set_rust_chat_completions(chat_completions=_RecordingCall()) + assert _should_attempt() is True def test_the_env_var_opts_in_without_a_per_model_flag(self, monkeypatch): monkeypatch.setenv("LITELLM_RUST", "true") - bridge.set_rust_chat_completions(decline=_RecordingDecline()) - assert _accepts(litellm_params={}) is True + bridge.set_rust_chat_completions(chat_completions=_RecordingCall()) + assert _should_attempt(litellm_params={}) is True + + def test_explicit_false_overrides_the_process_switch(self): + bridge.set_rust_chat_completions(chat_completions=_RecordingCall()) + configuration.use_litellm_rust(True) + + assert _should_attempt(litellm_params={"rust": False}) is False def test_declines_streaming_and_providers_off_the_path(self, monkeypatch): monkeypatch.delenv("LITELLM_RUST", raising=False) - gate = _RecordingDecline() - bridge.set_rust_chat_completions(decline=gate) - assert _accepts(stream=True) is False - assert _accepts(custom_llm_provider="openai") is False - assert _accepts(custom_llm_provider=None) is False - assert gate.calls == [] + bridge.set_rust_chat_completions(chat_completions=_RecordingCall()) + assert _should_attempt(stream=True) is False + assert _should_attempt(custom_llm_provider="openai") is False + assert _should_attempt(custom_llm_provider=None) is False def test_declines_an_anthropic_request_carrying_a_litellm_metadata_user_id(self, monkeypatch): """`AnthropicConfig.transform_request` copies a valid `user_id` into the Messages body. @@ -171,24 +152,21 @@ class TestGate: to Anthropic with the abuse-detection attribution silently missing. """ monkeypatch.delenv("LITELLM_RUST", raising=False) - gate = _RecordingDecline() - bridge.set_rust_chat_completions(decline=gate) - assert _accepts(litellm_params={"rust": True, "metadata": {"user_id": "u-123"}}) is False - assert gate.calls == [], "the core must not be consulted for a request it cannot see the key of" + bridge.set_rust_chat_completions(chat_completions=_RecordingCall()) + assert _should_attempt(litellm_params={"rust": True, "metadata": {"user_id": "u-123"}}) is False # Bedrock's Converse transform reads no `user_id`, and an Anthropic request # whose metadata carries none is one Python would not attribute either. assert ( - _accepts( + _should_attempt( custom_llm_provider="bedrock", - model="bedrock/us-east-1/anthropic.claude-v2", litellm_params={"rust": True, "metadata": {"user_id": "u-123"}}, ) is True ) - assert _accepts(litellm_params={"rust": True, "metadata": {"trace_id": "t-1"}}) is True - assert _accepts(litellm_params={"rust": True, "metadata": {"user_id": None}}) is True - assert _accepts(litellm_params={"rust": True, "metadata": None}) is True + assert _should_attempt(litellm_params={"rust": True, "metadata": {"trace_id": "t-1"}}) is True + assert _should_attempt(litellm_params={"rust": True, "metadata": {"user_id": None}}) is True + assert _should_attempt(litellm_params={"rust": True, "metadata": None}) is True def test_declines_a_bedrock_request_while_the_proxy_owns_request_metadata(self, monkeypatch): """`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the @@ -197,39 +175,49 @@ class TestGate: who armed `bedrock_request_metadata_fields` keeps the Python path. """ monkeypatch.delenv("LITELLM_RUST", raising=False) - gate = _RecordingDecline() - bridge.set_rust_chat_completions(decline=gate) + bridge.set_rust_chat_completions(chat_completions=_RecordingCall()) bedrock = { "custom_llm_provider": "bedrock", - "model": "bedrock/us-east-1/anthropic.claude-v2", } monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["user_api_key_team_id"]) - assert _accepts(**bedrock) is False - assert gate.calls == [], "the core must not be consulted for a field it cannot write" - assert _accepts() is True, "arming Bedrock attribution must not decline Anthropic" + assert _should_attempt(**bedrock) is False + assert _should_attempt() is True, "arming Bedrock attribution must not decline Anthropic" monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", None) - assert _accepts(**bedrock) is True, "the decline follows the operator's opt-in alone" - - def test_declines_when_the_core_declines(self, monkeypatch): - monkeypatch.delenv("LITELLM_RUST", raising=False) - bridge.set_rust_chat_completions(decline=_RecordingDecline("streaming")) - assert _accepts() is False + assert _should_attempt(**bedrock) is True, "the decline follows the operator's opt-in alone" def test_declines_when_the_bridge_is_unavailable(self, monkeypatch): monkeypatch.delenv("LITELLM_RUST", raising=False) + bridge.set_rust_chat_completions(decline=None) _hide_native_bridge(monkeypatch) - assert _accepts() is False + assert _should_attempt() is False - def test_declines_when_the_gate_itself_raises(self, monkeypatch): - monkeypatch.delenv("LITELLM_RUST", raising=False) + def test_checks_the_native_capability_gate(self, monkeypatch): + _hide_native_bridge(monkeypatch) + bridge.set_rust_chat_completions(decline=None) + assert _should_attempt() is False + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None) + assert _should_attempt() is True - def exploding(**_kwargs): - raise RuntimeError("boom") + def test_native_unsupported_capability_declines(self, monkeypatch): + _fake_native_bridge(monkeypatch) - bridge.set_rust_chat_completions(decline=exploding) - assert _accepts() is False + def decline(**_kwargs): + raise _FakeDeclined("unsupported request") + + bridge.set_rust_chat_completions(decline=decline) + assert _should_attempt() is False + + def test_native_capability_failure_does_not_fall_back(self, monkeypatch): + _fake_native_bridge(monkeypatch) + + def fail(**_kwargs): + raise _FakeUpstream(502, "capability check failed") + + bridge.set_rust_chat_completions(decline=fail) + with pytest.raises(APIError, match="capability check failed"): + _should_attempt() def _call_kwargs(model_response: ModelResponse) -> dict: @@ -247,14 +235,26 @@ def _call_kwargs(model_response: ModelResponse) -> dict: } +def _sync_call_kwargs(model_response: ModelResponse) -> dict: + return {**_call_kwargs(model_response), "python_fallback": lambda: "python"} + + +async def _async_python_fallback() -> object: + return "python" + + +def _async_call_kwargs(model_response: ModelResponse) -> dict: + return {**_call_kwargs(model_response), "python_fallback": _async_python_fallback} + + class TestSyncCall: - def test_builds_a_model_response_and_stamps_the_rust_header(self): + def test_builds_a_model_response(self): native = _RecordingCall() bridge.set_rust_chat_completions(chat_completions=native) model_response = ModelResponse() original_id = model_response.id - result = bridge.chat_completions(**_call_kwargs(model_response)) + result = bridge.chat_completions_or_fallback(**_sync_call_kwargs(model_response)) assert result is not None assert result.choices[0].message.content == "hello from rust" @@ -263,44 +263,54 @@ class TestSyncCall: assert result.usage.prompt_tokens == 11 assert result.usage.completion_tokens == 4 assert result.usage.total_tokens == 15 - assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} assert result.id == original_id, "the rust path must keep the chatcmpl id litellm already minted" def test_passes_the_timeout_through_as_seconds(self): native = _RecordingCall() bridge.set_rust_chat_completions(chat_completions=native) - bridge.chat_completions(**_call_kwargs(ModelResponse())) + bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) assert native.calls[0]["timeout_seconds"] == 30.0 def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + assert bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) == "python" def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + assert bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) == "python" + + def test_model_response_fallback_is_returned_unchanged(self, monkeypatch): + _fake_native_bridge(monkeypatch) + bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("unsupported"))) + fallback_response = ModelResponse() + + result = bridge.chat_completions_or_fallback( + **_call_kwargs(ModelResponse()), + python_fallback=lambda: fallback_response, + ) + + assert result is fallback_response class TestAsyncCall: @pytest.mark.asyncio async def test_builds_a_model_response(self): bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) - result = await bridge.achat_completions(**_call_kwargs(ModelResponse())) + result = await bridge.achat_completions_or_fallback(**_async_call_kwargs(ModelResponse())) assert result is not None assert result.choices[0].message.content == "hello from rust" - assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} @pytest.mark.asyncio async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None + assert await bridge.achat_completions_or_fallback(**_async_call_kwargs(ModelResponse())) == "python" @pytest.mark.asyncio async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))) - assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None + assert await bridge.achat_completions_or_fallback(**_async_call_kwargs(ModelResponse())) == "python" class TestAsyncFallbackWrapper: @@ -349,14 +359,14 @@ class TestFailureClassification: def test_a_decline_falls_back_because_nothing_was_sent(self): bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + assert bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) == "python" def test_an_upstream_failure_is_surfaced_with_its_status(self): from litellm.exceptions import APIError bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited"))) with pytest.raises(APIError) as raised: - bridge.chat_completions(**_call_kwargs(ModelResponse())) + bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) assert raised.value.status_code == 429 assert "rate limited" in str(raised.value) @@ -365,13 +375,13 @@ class TestFailureClassification: bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset"))) with pytest.raises(APIError) as raised: - bridge.chat_completions(**_call_kwargs(ModelResponse())) + bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) assert raised.value.status_code == 500 def test_an_unrecognized_error_is_not_swallowed(self): bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else"))) with pytest.raises(RuntimeError): - bridge.chat_completions(**_call_kwargs(ModelResponse())) + bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) @pytest.mark.asyncio async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self): diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index bbeb6c38f78..92e86f48c31 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -6,6 +6,7 @@ import litellm from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") +rust_bridge_bindings = importlib.import_module("litellm.rust_bridge.bindings") class SyncBridge: @@ -77,7 +78,7 @@ async def test_enabled_async_bridge() -> None: def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None: rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) - monkeypatch.setattr("litellm.rust_bridge.get_native_bridge", lambda: None) + monkeypatch.setattr(rust_bridge_bindings, "get_native_bridge", lambda: None) assert rust_bridge.load_rust_transcription() is None assert rust_bridge.load_rust_atranscription() is None