diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index e9f8c477f36..e966741e776 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -4,9 +4,10 @@ mod service; use axum::Router; use axum::body::Body; -use axum::extract::{Json, State}; +use axum::extract::{Json, Request, State}; use axum::http::StatusCode; use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE, HeaderMap, HeaderValue}; +use axum::middleware::{self, Next}; use axum::response::{IntoResponse, Response}; use axum::routing::post; use litellm_core::Error; @@ -16,9 +17,23 @@ use crate::auth::RequireMasterKey; use crate::constants::{MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH}; use crate::state::AppState; +const CORE_ENGINE_HEADER: &str = "x-litellm-core"; +const RUST_CORE_ENGINE: &str = "rust"; + /// This route's contribution to the app router. pub fn router() -> Router { - Router::new().route(MESSAGES_ROUTE_PATH, post(handle)) + Router::new() + .route(MESSAGES_ROUTE_PATH, post(handle)) + .route_layer(middleware::from_fn(core_engine_header)) +} + +async fn core_engine_header(request: Request, next: Next) -> Response { + let mut response = next.run(request).await; + response.headers_mut().insert( + CORE_ENGINE_HEADER, + HeaderValue::from_static(RUST_CORE_ENGINE), + ); + response } async fn handle( @@ -143,6 +158,7 @@ mod tests { use tower::ServiceExt; use super::super::app; + use super::{CORE_ENGINE_HEADER, RUST_CORE_ENGINE}; use crate::io::realtime_pool::RealtimePool; use crate::state::AppState; @@ -290,6 +306,7 @@ mod tests { .await .expect("route responds"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE); let body = axum::body::to_bytes(response.into_body(), usize::MAX) .await .expect("response body reads"); @@ -340,6 +357,7 @@ mod tests { .await .expect("route responds"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE); let upstream_request = server.await.expect("upstream task completes"); let (_, upstream_body) = upstream_request .split_once("\r\n\r\n") @@ -378,6 +396,7 @@ mod tests { .await .expect("route responds"); assert_eq!(response.status(), StatusCode::OK); + assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE); assert_eq!( response .headers() @@ -443,6 +462,7 @@ mod tests { .await .expect("route responds"); assert_eq!(response.status(), StatusCode::BAD_GATEWAY); + assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE); let response_body = axum::body::to_bytes(response.into_body(), usize::MAX) .await .expect("response body reads"); @@ -473,6 +493,7 @@ mod tests { .await .expect("route responds"); assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE); } #[tokio::test] @@ -495,6 +516,7 @@ mod tests { .await .expect("route responds"); assert_eq!(response.status(), StatusCode::UNAUTHORIZED); + assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE); } #[tokio::test] @@ -517,5 +539,6 @@ mod tests { .await .expect("route responds"); assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE); } } diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 877aab0ed00..c89c2838128 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 ( @@ -578,48 +621,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/anthropic/experimental_pass_through/messages/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py index 45c7825344b..3d2c5f48bab 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/streaming_iterator.py @@ -6,7 +6,7 @@ from typing import Any, Final, Protocol, runtime_checkable import httpx from pydantic import TypeAdapter -from typing_extensions import TypedDict +from typing_extensions import ReadOnly, TypedDict from litellm.constants import ( ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS, @@ -19,6 +19,7 @@ from litellm.llms.anthropic.common_utils import ANTHROPIC_ERROR_STATUS_CODE_MAP from litellm.proxy.pass_through_endpoints.success_handler import ( PassThroughEndpointLogging, ) +from litellm.rust_bridge.runtime import CoreEngine, execution_additional_headers from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType from litellm.types.utils import GenericStreamingChunk, ModelResponseStream @@ -306,7 +307,8 @@ def _anthropic_content_block_events(index: int, block: Mapping[str, object]) -> class AnthropicMessagesStreamHiddenParams(TypedDict): - additional_headers: dict[str, str] + additional_headers: ReadOnly[Mapping[str, str]] + core_engine: ReadOnly[str] @runtime_checkable @@ -325,8 +327,10 @@ _RESPONSE_HEADERS_ADAPTER: Final[TypeAdapter[dict[str, str]]] = TypeAdapter(dict def anthropic_messages_stream_hidden_params( response_headers: httpx.Headers, ) -> AnthropicMessagesStreamHiddenParams: + additional_headers: Final = _RESPONSE_HEADERS_ADAPTER.validate_python(process_response_headers(response_headers)) return AnthropicMessagesStreamHiddenParams( - additional_headers=_RESPONSE_HEADERS_ADAPTER.validate_python(process_response_headers(response_headers)) + additional_headers=execution_additional_headers(additional_headers, CoreEngine.PYTHON), + core_engine=CoreEngine.PYTHON.value, ) diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index b1f8c957ff4..e1504a18f71 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -5,6 +5,7 @@ import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file from litellm.rust_bridge import transcription as rust_transcription_bridge +from litellm.rust_bridge.runtime import CoreEngine, execution_hidden_params from litellm.types.utils import FileTypes, TranscriptionResponse @@ -55,7 +56,9 @@ class BedrockAudioTranscriptionRustDispatch: ) if rust_response is None: raise RuntimeError("Rust audio transcription bridge is unavailable") - return TranscriptionResponse(**rust_response) + response: Final = TranscriptionResponse(**rust_response) + response["_hidden_params"] = execution_hidden_params(None, CoreEngine.RUST) + return response async def async_audio_transcriptions( self, @@ -81,4 +84,6 @@ class BedrockAudioTranscriptionRustDispatch: ) if rust_response is None: raise RuntimeError("Rust audio transcription bridge is unavailable") - return TranscriptionResponse(**rust_response) + response: Final = TranscriptionResponse(**rust_response) + response["_hidden_params"] = execution_hidden_params(None, CoreEngine.RUST) + return response diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index ca5f1298360..50ae00e8331 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -365,6 +365,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 @@ -437,7 +522,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, @@ -448,9 +533,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: @@ -496,103 +580,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/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 4dd7f7d9a0b..ac7a1907021 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -19,7 +19,6 @@ import litellm.types.utils from litellm._logging import _redact_string, verbose_logger from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES -from litellm.exceptions import APIError from litellm.litellm_core_utils.agentic_loop_settings import ( DEFAULT_MAX_AGENTIC_LOOPS, validated_max_agentic_loops, @@ -89,6 +88,7 @@ from litellm.responses.streaming_iterator import ( ResponsesWebSocketStreaming, SyncResponsesAPIStreamingIterator, ) +from litellm.rust_bridge.runtime import CoreEngine, execution_additional_headers, execution_hidden_params from litellm.types.containers.main import ( ContainerFileListResponse, ContainerListResponse, @@ -235,6 +235,23 @@ def _google_genai_streaming_hidden_params( } +def _anthropic_messages_with_core_engine( + response: AnthropicMessagesResponse, + source: CoreEngine, +) -> AnthropicMessagesResponse: + existing_hidden_params: Final = response.get("_hidden_params") + return cast( # cast-ok: Anthropic response TypedDict is preserved while adding its documented hidden field + AnthropicMessagesResponse, + { # mutable-ok: Anthropic response contract is a mutable TypedDict + **response, + "_hidden_params": execution_hidden_params( + existing_hidden_params if isinstance(existing_hidden_params, dict) else None, + source, + ), + }, + ) + + @lru_cache(maxsize=None) def _responses_api_optional_request_param_names() -> frozenset[str]: return frozenset(get_type_hints(ResponsesAPIOptionalRequestParams).keys()) @@ -2361,8 +2378,17 @@ class BaseLLMHTTPHandler: kwargs=kwargs_for_agentic, ) - return self._maybe_wrap_in_fake_stream( + selected_response: Final = cast( # cast-ok: hook result is validated by the Anthropic response pipeline + AnthropicMessagesResponse, final_response if final_response is not None else initial_response, + ) + existing_hidden_params: Final = selected_response.get("_hidden_params") + existing_source: Final = ( + existing_hidden_params.get("core_engine") if isinstance(existing_hidden_params, dict) else None + ) + source: Final = CoreEngine.RUST if existing_source == CoreEngine.RUST.value else CoreEngine.PYTHON + return self._maybe_wrap_in_fake_stream( + _anthropic_messages_with_core_engine(selected_response, source), logging_obj, "anthropic_messages", ) @@ -2394,45 +2420,20 @@ class BaseLLMHTTPHandler: from litellm.rust_bridge import messages as rust_messages_bridge upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"} - try: - rust_response: Final = await rust_messages_bridge.amessages( - model=model, - body=upstream_body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - ) - except Exception as rust_error: # noqa: BLE001 - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - declined: Final = getattr(native_bridge, "RustBridgeDeclined", None) - upstream_failed: Final = getattr(native_bridge, "RustUpstreamError", None) - if isinstance(upstream_failed, type) and 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 messages: {message}", - llm_provider=custom_llm_provider, - model=model, - ) from rust_error - if not isinstance(declined, type) or not isinstance(rust_error, declined): - raise - verbose_logger.debug( - "Rust Anthropic messages bridge declined before calling the provider (%s); falling back to Python path", - rust_error, - ) - return None + rust_response: Final = await rust_messages_bridge.amessages( + model=model, + body=upstream_body, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + ) if rust_response is None: return None response_obj: Final = cast(AnthropicMessagesResponse, dict(rust_response)) - response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}} - return response_obj + return _anthropic_messages_with_core_engine(response_obj, CoreEngine.RUST) @staticmethod def _rust_anthropic_messages_fake_stream( @@ -2447,7 +2448,10 @@ class BaseLLMHTTPHandler: ) completion_stream = cast(AsyncIterator[bytes], FakeAnthropicMessagesStreamIterator(response=rust_response)) - hidden_params: Final = AnthropicMessagesStreamHiddenParams(additional_headers={"x-litellm-rust": "true"}) + hidden_params: Final = AnthropicMessagesStreamHiddenParams( + additional_headers=execution_additional_headers(None, CoreEngine.RUST), + core_engine=CoreEngine.RUST.value, + ) return AnthropicMessagesStreamingResponse( completion_stream=completion_stream, hidden_params=hidden_params, @@ -6508,13 +6512,14 @@ class BaseLLMHTTPHandler: if _rust_responses_websocket_enabled(custom_llm_provider, litellm_params): from litellm.rust_bridge import responses_websocket as rust_responses_websocket - rust_backend: Final = await rust_responses_websocket.connect( + rust_execution: Final = await rust_responses_websocket.connect( provider="openai", api_key=api_key, api_base=api_base, headers={str(key): str(value) for key, value in headers.items()}, timeout=timeout, ) + rust_backend: Final = rust_execution.value if rust_backend is not None: yield rust_backend return diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index b260ec6e06f..5404e77d761 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -29,6 +29,7 @@ from litellm.llms.base_llm.ocr.transformation import ( ) from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import ocr as rust_ocr_bridge +from litellm.rust_bridge.runtime import CoreEngine, execution_hidden_params from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -60,6 +61,24 @@ class _PreparedRustOCRCall: optional_params: dict[str, object] +def _with_core_engine(response: OCRResponse, source: CoreEngine) -> OCRResponse: + existing_hidden_params: Final = cast( # cast-ok: OCRResponse initializes hidden params as a mapping + Mapping[str, object], + response._hidden_params, # pyright: ignore[reportPrivateUsage] # provenance uses the response's established metadata channel + ) + response._hidden_params = execution_hidden_params( # pyright: ignore[reportPrivateUsage] # OCRResponse has no public hidden-params setter + existing_hidden_params, source + ) + return response + + +async def _await_with_core_engine( + response: Coroutine[object, object, OCRResponse], + source: CoreEngine, +) -> OCRResponse: + return _with_core_engine(await response, source) + + _RUST_OCR_PROVIDERS: Final = { "mistral", "azure_ai", @@ -306,7 +325,7 @@ def _run_rust_ocr( ) if rust_response is None: return None - return OCRResponse.model_validate(rust_response) + return _with_core_engine(OCRResponse.model_validate(rust_response), CoreEngine.RUST) async def _run_rust_aocr( @@ -331,7 +350,7 @@ async def _run_rust_aocr( ) if rust_response is None: return None - return OCRResponse.model_validate(rust_response) + return _with_core_engine(OCRResponse.model_validate(rust_response), CoreEngine.RUST) @client @@ -461,7 +480,7 @@ async def aocr( if response is None: raise ValueError(f"Got an unexpected None response from the OCR API: {response}") - return response + return _with_core_engine(response, CoreEngine.PYTHON) except Exception as e: raise litellm.exception_type( model=model, @@ -727,7 +746,9 @@ def ocr( litellm_params=prepared.litellm_params, ) - return response + if asyncio.iscoroutine(response): + return _await_with_core_engine(response, CoreEngine.PYTHON) + return _with_core_engine(response, CoreEngine.PYTHON) except Exception as e: raise litellm.exception_type( model=model, diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 1d860b02875..a3351063c7c 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1550,6 +1550,16 @@ class ProxyBaseLLMRequestProcessing: logging_obj=litellm_logging_obj, use_logging_obj=read_timing_from_logging_obj, ) + core_engine_value: Final = hidden_params.get("core_engine") + core_engine: Final = ( + core_engine_value + if isinstance(core_engine_value, str) and core_engine_value in ("python", "rust") + else "python" + ) + reserved_core_headers: Final = frozenset({"x-litellm-core", "x-litellm-rust"}) + forwarded_headers: Final = { # mutable-ok: response headers are assembled into a mutable transport dict + key: value for key, value in kwargs.items() if key.lower() not in reserved_core_headers + } cost_breakdown: Final = _get_cost_breakdown_from_logging_obj( litellm_logging_obj=litellm_logging_obj, response_cost=response_cost @@ -1634,7 +1644,12 @@ class ProxyBaseLLMRequestProcessing: str(fastest_response_batch_completion) if fastest_response_batch_completion is not None else None ), "x-litellm-timeout": str(timeout) if timeout is not None else None, - **{k: str(v) for k, v in kwargs.items()}, + **{ # mutable-ok: response header assembly requires a mutable dict + key: str(value) for key, value in forwarded_headers.items() + }, + "x-litellm-core": core_engine, + # mutable-ok: conditionally extending the mutable response header dict + **({"x-litellm-rust": "true"} if core_engine == "rust" else {}), } if request_data: remaining_tokens_header: Final = get_remaining_tokens_and_requests_from_request_data(request_data) diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index c599667ab17..c4952c1e87e 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -14,20 +14,33 @@ 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 +from typing import TYPE_CHECKING, Final, Protocol, cast 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.configuration import rust_enabled -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.runtime import ( + LEGACY_RUST_HEADER, + UNSET, + BridgeErrorContext, + CoreEngine, + ExecutionResult, + FallbackMode, + NativeBinding, + RustHandled, + Unset, + aattempt, + ainvoke, + attempt, + execution_hidden_params, + invoke, +) from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.utils import ModelResponse @@ -42,7 +55,7 @@ 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" +RUST_RESPONSE_HEADER: Final = LEGACY_RUST_HEADER class RustChatCompletions(Protocol): @@ -126,67 +139,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: @@ -275,57 +255,6 @@ def rust_chat_completions_accepts( 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,13 +262,33 @@ 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 + hidden_params=execution_hidden_params(None, CoreEngine.RUST), # 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__}") return built +def _unwrap_execution(result: ExecutionResult[object]) -> object: + value: Final = result.value + if isinstance(value, ModelResponse): + raw_hidden_params: Final = cast( # cast-ok: runtime validation below narrows this metadata before use + object, + value._hidden_params, # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params getter + ) + hidden_params: Final = ( + cast( # cast-ok: the isinstance check validates the mapping boundary + Mapping[str, object], raw_hidden_params + ) + if isinstance(raw_hidden_params, Mapping) + else None + ) + value._hidden_params = execution_hidden_params( # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params setter + hidden_params, result.source + ) + return value + + def chat_completions( *, model: str, @@ -354,10 +303,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 +316,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 +344,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 +357,63 @@ 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 _unwrap_execution( + 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 +430,32 @@ 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 _unwrap_execution( + 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..24beab60f8b 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.runtime 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 40d0ddf622b..9d037a838fe 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -3,11 +3,21 @@ 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.runtime import ( + UNSET, + BridgeErrorContext, + FallbackMode, + NativeBinding, + Unset, + ainvoke, + async_none, + identity, + invoke, +) from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -39,53 +49,25 @@ class RustAmessages(Protocol): raise NotImplementedError -class _Unset: - pass - - -_UNSET: Final[_Unset] = _Unset() - - -@dataclass(slots=True) -class _RustMessagesState: - messages: RustMessages | None = None - amessages: RustAmessages | None = None - - -_STATE: Final[_RustMessagesState] = _RustMessagesState() +_MESSAGES: Final = NativeBinding[RustMessages]("messages") +_AMESSAGES: Final = NativeBinding[RustAmessages]("amessages") 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): - _STATE.messages = messages - if not isinstance(amessages, _Unset): - _STATE.amessages = amessages + _MESSAGES.update(messages) + _AMESSAGES.update(amessages) def load_rust_messages() -> RustMessages | None: - if _STATE.messages is not None: - return _STATE.messages - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustMessages, getattr(native_bridge, "messages", None)) + return _MESSAGES.load() def load_rust_amessages() -> RustAmessages | None: - if _STATE.amessages is not None: - return _STATE.amessages - from litellm.rust_bridge import get_native_bridge - - native_bridge: Final = get_native_bridge() - if native_bridge is None: - return None - return cast(RustAmessages, getattr(native_bridge, "amessages", None)) + return _AMESSAGES.load() def messages( @@ -99,17 +81,27 @@ def messages( timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_messages: Final = load_rust_messages() - if rust_messages is None: - return None - return rust_messages( - model=model, - body=body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), + native_call: Final = ( + None + if rust_messages is None + else lambda: rust_messages( + model=model, + body=body, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ) ) + result: Final = invoke( + native_call=native_call, + fallback=lambda: None, + adapt=identity, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route="messages", provider=custom_llm_provider or "", model=model), + ) + return result.value async def amessages( @@ -123,14 +115,24 @@ async def amessages( timeout: float | httpx.Timeout | None, ) -> dict[str, object] | None: rust_amessages: Final = load_rust_amessages() - if rust_amessages is None: - return None - return await rust_amessages( - model=model, - body=body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), + native_call: Final = ( + None + if rust_amessages is None + else lambda: rust_amessages( + model=model, + body=body, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ) ) + result: Final = await ainvoke( + native_call=native_call, + fallback=async_none, + adapt=identity, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route="messages", provider=custom_llm_provider or "", model=model), + ) + return result.value diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index cbe444d44a6..133405ae3f2 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -3,11 +3,22 @@ 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.runtime import ( + UNSET, + BridgeErrorContext, + FallbackMode, + NativeBinding, + Unset, + 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 +55,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,18 +88,28 @@ 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), + ) ) + result: Final = invoke( + native_call=native_call, + fallback=lambda: None, + adapt=identity, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route="ocr", provider=custom_llm_provider or "", model=model), + ) + return result.value async def aocr( @@ -127,15 +124,25 @@ 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), + ) ) + result: Final = 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), + ) + return result.value diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index b4079b3d37e..0aa65f03846 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -2,15 +2,24 @@ from __future__ import annotations import json from collections.abc import Mapping -from dataclasses import dataclass -from types import MappingProxyType from typing import Final, Protocol import httpx from pydantic import TypeAdapter from websockets.exceptions import ConnectionClosedOK -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.runtime import ( + UNSET, + BridgeErrorContext, + CoreEngine, + ExecutionResult, + FallbackMode, + NativeBinding, + Unset, + acall, + ainvoke, + async_none, +) from litellm.rust_bridge.timeouts import timeout_to_seconds _EVENT_ADAPTER: Final = TypeAdapter(Mapping[str, object]) @@ -36,57 +45,46 @@ class RustResponsesWebSocketConnection(Protocol): ) -> RustResponsesWebSocket: ... -class _Unset: - pass - - -_UNSET: Final = _Unset() - - -@dataclass(slots=True) -class _RustResponsesWebSocketState: - connection: type[RustResponsesWebSocketConnection] | None = None - - -_STATE: Final = _RustResponsesWebSocketState() +_CONNECTION: Final = NativeBinding[type[RustResponsesWebSocketConnection]]("ResponsesWebSocketSession") def set_rust_responses_websocket( *, - connection: type[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() -> type[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[type[RustResponsesWebSocketConnection] | None] = getattr( - native_bridge, "ResponsesWebSocketSession", None - ) - return connection_type + return _CONNECTION.load() class _ConnectionAdapter: def __init__(self, connection: RustResponsesWebSocket): self._connection: Final = connection + self.core_engine: Final = CoreEngine.RUST async def send(self, text: str) -> None: event: Final = _EVENT_ADAPTER.validate_json(text) - await self._connection.send_event(event) + await acall( + lambda: self._connection.send_event(event), + BridgeErrorContext(route="responses websocket", provider="openai", model=""), + ) async def recv(self) -> str: - event: Final = await self._connection.recv_event() + event: Final = await acall( + self._connection.recv_event, + BridgeErrorContext(route="responses websocket", provider="openai", model=""), + ) if event is None: raise ConnectionClosedOK(None, None) return json.dumps(dict(event), separators=(",", ":")) # mutable-ok: JSON requires a concrete dict async def close(self) -> None: - await self._connection.close() + await acall( + self._connection.close, + BridgeErrorContext(route="responses websocket", provider="openai", model=""), + ) async def connect( @@ -96,23 +94,28 @@ async def connect( api_base: str | None, headers: Mapping[str, str], timeout: float | httpx.Timeout | None, -) -> _ConnectionAdapter | None: +) -> ExecutionResult[_ConnectionAdapter | None]: connection_type: Final = load_rust_responses_websocket() - if connection_type is None: - return None - credentials: Final = None if api_key is None else MappingProxyType({"api_key": api_key}) - try: - connection: Final = await connection_type.connect( + credentials: Final = ( + None + if api_key is None + else {"api_key": api_key} # mutable-ok: native bridge serialization requires a plain dict + ) + native_call: Final = ( + None + if connection_type is None + else lambda: connection_type.connect( provider, credentials, api_base, headers, timeout_to_seconds(timeout), ) - except Exception as error: - native: Final = get_native_bridge() - declined: Final = None if native is None else getattr(native, "RustBridgeDeclined", None) - if isinstance(declined, type) and isinstance(error, declined): - return None - raise - return _ConnectionAdapter(connection) + ) + return await ainvoke( + native_call=native_call, + fallback=async_none, + adapt=_ConnectionAdapter, + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route="responses websocket", provider=provider, model=""), + ) diff --git a/litellm/rust_bridge/runtime.py b/litellm/rust_bridge/runtime.py new file mode 100644 index 00000000000..5468a901ab3 --- /dev/null +++ b/litellm/rust_bridge/runtime.py @@ -0,0 +1,289 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Mapping +from dataclasses import dataclass +from enum import Enum +from typing import Final, Generic, NoReturn, TypeAlias, TypeVar, cast + +from litellm.exceptions import APIError +from litellm.rust_bridge.loader import get_native_bridge + +NativeT = TypeVar("NativeT") +ResultT = TypeVar("ResultT") +BindingT = TypeVar("BindingT") + +CORE_ENGINE_HIDDEN_PARAM: Final = "core_engine" +CORE_ENGINE_HEADER: Final = "x-litellm-core" +LEGACY_RUST_HEADER: Final = "x-litellm-rust" + + +class Unset: + pass + + +UNSET: Final = Unset() + + +class FallbackMode(Enum): + PYTHON = "python" + RUST_REQUIRED = "rust_required" + + +class CoreEngine(str, Enum): + PYTHON = "python" + RUST = "rust" + + +@dataclass(frozen=True, slots=True) +class ExecutionResult(Generic[ResultT]): + value: ResultT + source: CoreEngine + + +@dataclass(frozen=True, slots=True) +class RustHandled(Generic[ResultT]): + value: ResultT + + +@dataclass(frozen=True, slots=True) +class RustDeclined: + reason: str + + +@dataclass(frozen=True, slots=True) +class RustUnavailable: + pass + + +RustAttempt: TypeAlias = RustHandled[ResultT] | RustDeclined | RustUnavailable + + +def execution_headers(source: CoreEngine) -> dict[str, str]: + if source is CoreEngine.RUST: + return { # mutable-ok: response adapters require a mutable header dict + CORE_ENGINE_HEADER: source.value, + LEGACY_RUST_HEADER: "true", + } + return {CORE_ENGINE_HEADER: source.value} # mutable-ok: response adapters require a mutable header dict + + +def execution_additional_headers( + additional_headers: Mapping[str, object] | None, + source: CoreEngine, +) -> dict[str, str]: + existing: Final = additional_headers or {} # mutable-ok: empty default is local and never mutated + reserved: Final = frozenset({CORE_ENGINE_HEADER, LEGACY_RUST_HEADER}) + preserved: Final = { # mutable-ok: response adapters require a mutable header dict + str(name): str(value) for name, value in existing.items() if str(name).lower() not in reserved + } + return { # mutable-ok: response adapters require a mutable header dict + **preserved, + **execution_headers(source), + } + + +def execution_hidden_params( + hidden_params: Mapping[str, object] | None, + source: CoreEngine, +) -> dict[str, object]: + existing: Final = hidden_params or {} # mutable-ok: empty default is local and never mutated + raw_headers: Final = existing.get("additional_headers") + additional_headers: Final[Mapping[str, object]] = ( + cast( # cast-ok: the isinstance check validates the mapping boundary + Mapping[str, object], raw_headers + ) + if isinstance(raw_headers, Mapping) + else {} # mutable-ok: empty default is local and never mutated + ) + return { # mutable-ok: response objects require mutable hidden params + **existing, + CORE_ENGINE_HIDDEN_PARAM: source.value, + "additional_headers": execution_additional_headers(additional_headers, source), + } + + +@dataclass(frozen=True, slots=True) +class BridgeErrorContext: + route: str + provider: str + model: str + + +class NativeBinding(Generic[BindingT]): + def __init__(self, attribute: str) -> None: + self._attribute: Final = attribute + self._override: BindingT | Unset = UNSET + + def load(self) -> BindingT | None: + if not isinstance(self._override, Unset): + return self._override + native: Final = get_native_bridge() + return ( + None + if native is None + else cast( # cast-ok: each route validates the callable against its binding protocol + BindingT | None, getattr(native, self._attribute, None) + ) + ) + + def set(self, value: BindingT | None) -> None: + self._override = UNSET if value is None else value + + def update(self, value: BindingT | None | Unset) -> None: + if not isinstance(value, Unset): + self.set(value) + + +def invoke( + *, + native_call: Callable[[], NativeT] | None, + fallback: Callable[[], ResultT], + adapt: Callable[[NativeT], ResultT], + mode: FallbackMode, + context: BridgeErrorContext, +) -> ExecutionResult[ResultT]: + attempt_result: Final = attempt(native_call=native_call, adapt=adapt, context=context) + if isinstance(attempt_result, RustHandled): + return ExecutionResult(value=attempt_result.value, source=CoreEngine.RUST) + if mode is FallbackMode.PYTHON: + return ExecutionResult(value=fallback(), source=CoreEngine.PYTHON) + _raise_required(attempt_result, context) + + +async def ainvoke( + *, + native_call: Callable[[], Awaitable[NativeT]] | None, + fallback: Callable[[], Awaitable[ResultT]], + adapt: Callable[[NativeT], ResultT], + mode: FallbackMode, + context: BridgeErrorContext, +) -> ExecutionResult[ResultT]: + attempt_result: Final = await aattempt(native_call=native_call, adapt=adapt, context=context) + if isinstance(attempt_result, RustHandled): + return ExecutionResult(value=attempt_result.value, source=CoreEngine.RUST) + if mode is FallbackMode.PYTHON: + return ExecutionResult(value=await fallback(), source=CoreEngine.PYTHON) + _raise_required(attempt_result, context) + + +def attempt( + *, + native_call: Callable[[], NativeT] | None, + adapt: Callable[[NativeT], ResultT], + context: BridgeErrorContext, +) -> RustAttempt[ResultT]: + if native_call is None: + return RustUnavailable() + exceptions: Final = _native_exceptions() + if exceptions is None: + return RustHandled(adapt(native_call())) + declined, upstream = exceptions + try: + native_result: Final = native_call() + except declined as error: + return RustDeclined(reason=_decline_reason(error)) + except upstream as error: + _raise_upstream(error, context) + return RustHandled(adapt(native_result)) + + +async def aattempt( + *, + native_call: Callable[[], Awaitable[NativeT]] | None, + adapt: Callable[[NativeT], ResultT], + context: BridgeErrorContext, +) -> RustAttempt[ResultT]: + if native_call is None: + return RustUnavailable() + exceptions: Final = _native_exceptions() + if exceptions is None: + return RustHandled(adapt(await native_call())) + declined, upstream = exceptions + try: + native_result: Final = await native_call() + except declined as error: + return RustDeclined(reason=_decline_reason(error)) + except upstream as error: + _raise_upstream(error, context) + return RustHandled(adapt(native_result)) + + +def _native_exceptions() -> tuple[type[BaseException], type[BaseException]] | None: + native: Final = get_native_bridge() + if native is None: + return None + declined: Final = getattr(native, "RustBridgeDeclined", None) + upstream: Final = getattr(native, "RustUpstreamError", None) + if not isinstance(declined, type) or not isinstance(upstream, type): + return None + return declined, upstream + + +def _decline_reason(error: BaseException) -> str: + reason: Final = error.args[0] if error.args else str(error) + return reason if isinstance(reason, str) else str(reason) + + +def _raise_required( + attempt_result: RustDeclined | RustUnavailable, + context: BridgeErrorContext, +) -> NoReturn: + raise RuntimeError(f"Rust {context.route} bridge {_required_reason(attempt_result)}") + + +def _required_reason(attempt_result: RustDeclined | RustUnavailable) -> str: + match attempt_result: + case RustUnavailable(): + return "is unavailable" + case RustDeclined(reason=reason): + return f"declined the request: {reason}" + + +def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn: + args: Final = cast( # cast-ok: BaseException.args is always a tuple at runtime + tuple[object, ...], error.args + ) + status_value: Final = args[0] if args else 0 + message_value: Final = args[1] if len(args) > 1 else str(error) + status: Final = status_value if isinstance(status_value, int) else 0 + message: Final = message_value if isinstance(message_value, str) else str(message_value) + api_error: Final = APIError( + status_code=status or 500, + message=f"litellm rust {context.route}: {message}", + llm_provider=context.provider, + model=context.model, + ) + api_error.headers = execution_headers( # pyright: ignore[reportAttributeAccessIssue] # proxy reads exception headers + CoreEngine.RUST + ) + raise api_error from error + + +def call(operation: Callable[[], ResultT], context: BridgeErrorContext) -> ResultT: + exceptions: Final = _native_exceptions() + if exceptions is None: + return operation() + upstream: Final = exceptions[1] + try: + return operation() + except upstream as error: + _raise_upstream(error, context) + + +async def acall(operation: Callable[[], Awaitable[ResultT]], context: BridgeErrorContext) -> ResultT: + exceptions: Final = _native_exceptions() + if exceptions is None: + return await operation() + upstream: Final = exceptions[1] + try: + return await operation() + except upstream as error: + _raise_upstream(error, context) + + +def identity(value: ResultT) -> ResultT: + return value + + +async def async_none() -> None: + return None diff --git a/litellm/rust_bridge/streaming.py b/litellm/rust_bridge/streaming.py index 33ca3186d13..18bcdfa4ebe 100644 --- a/litellm/rust_bridge/streaming.py +++ b/litellm/rust_bridge/streaming.py @@ -1,16 +1,27 @@ from __future__ import annotations import json -from collections.abc import AsyncIterator, Generator, Iterator, Mapping -from dataclasses import dataclass +from collections.abc import AsyncIterator, Iterator, Mapping from functools import lru_cache from types import MappingProxyType -from typing import TYPE_CHECKING, Final, Literal, NoReturn, Protocol, TypeAlias, runtime_checkable +from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable import httpx -from litellm.exceptions import APIError -from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.runtime import ( + UNSET, + BridgeErrorContext, + CoreEngine, + ExecutionResult, + FallbackMode, + NativeBinding, + Unset, + acall, + ainvoke, + async_none, + call, + invoke, +) from litellm.rust_bridge.timeouts import timeout_to_seconds if TYPE_CHECKING: @@ -59,171 +70,49 @@ class RustAsyncStreamOpen(Protocol): ) -> RustEventStream: ... -@runtime_checkable -class ObjectAwaitable(Protocol): - def __await__(self) -> Generator[object, None, object]: ... - - -class _Unset: - pass - - -_UNSET: Final = _Unset() - - -@dataclass(slots=True) -class _RustStreamingState: - chat: RustStreamOpen | None = None - achat: RustAsyncStreamOpen | None = None - messages: RustStreamOpen | None = None - amessages: RustAsyncStreamOpen | None = None - responses: RustStreamOpen | None = None - aresponses: RustAsyncStreamOpen | None = None - - -_STATE: Final = _RustStreamingState() +_CHAT: Final = NativeBinding[RustStreamOpen]("chat_completions_stream") +_ACHAT: Final = NativeBinding[RustAsyncStreamOpen]("achat_completions_stream") +_MESSAGES: Final = NativeBinding[RustStreamOpen]("messages_stream") +_AMESSAGES: Final = NativeBinding[RustAsyncStreamOpen]("amessages_stream") +_RESPONSES: Final = NativeBinding[RustStreamOpen]("responses_stream") +_ARESPONSES: Final = NativeBinding[RustAsyncStreamOpen]("aresponses_stream") def set_rust_streaming( *, - chat: RustStreamOpen | None | _Unset = _UNSET, - achat: RustAsyncStreamOpen | None | _Unset = _UNSET, - messages: RustStreamOpen | None | _Unset = _UNSET, - amessages: RustAsyncStreamOpen | None | _Unset = _UNSET, - responses: RustStreamOpen | None | _Unset = _UNSET, - aresponses: RustAsyncStreamOpen | None | _Unset = _UNSET, + chat: RustStreamOpen | None | Unset = UNSET, + achat: RustAsyncStreamOpen | None | Unset = UNSET, + messages: RustStreamOpen | None | Unset = UNSET, + amessages: RustAsyncStreamOpen | None | Unset = UNSET, + responses: RustStreamOpen | None | Unset = UNSET, + aresponses: RustAsyncStreamOpen | None | Unset = UNSET, ) -> None: - if not isinstance(chat, _Unset): - _STATE.chat = chat - if not isinstance(achat, _Unset): - _STATE.achat = achat - if not isinstance(messages, _Unset): - _STATE.messages = messages - if not isinstance(amessages, _Unset): - _STATE.amessages = amessages - if not isinstance(responses, _Unset): - _STATE.responses = responses - if not isinstance(aresponses, _Unset): - _STATE.aresponses = aresponses - - -def _native_attribute(name: str) -> object | None: - native: Final = get_native_bridge() - return None if native is None else getattr(native, name, None) - - -def _native_sync_opener(name: str) -> RustStreamOpen | None: - opener: Final = _native_attribute(name) - if not callable(opener): - return None - - def open_stream( - request: Mapping[str, object], - provider: str, - credentials: Mapping[str, str] | None, - api_base: str | None, - extra_headers: Mapping[str, str] | None, - timeout_seconds: float | None, - ) -> RustEventStream: - stream: Final = opener( - request, - provider, - credentials, - api_base, - extra_headers, - timeout_seconds, - ) - if not isinstance(stream, RustEventStream): - raise TypeError("native stream opener returned an invalid stream") - return stream - - return open_stream - - -def _native_async_opener(name: str) -> RustAsyncStreamOpen | None: - opener: Final = _native_attribute(name) - if not callable(opener): - return None - - async def open_stream( - request: Mapping[str, object], - provider: str, - credentials: Mapping[str, str] | None, - api_base: str | None, - extra_headers: Mapping[str, str] | None, - timeout_seconds: float | None, - ) -> RustEventStream: - pending: Final = opener( - request, - provider, - credentials, - api_base, - extra_headers, - timeout_seconds, - ) - if not isinstance(pending, ObjectAwaitable): - raise TypeError("native async stream opener returned a non-awaitable") - stream: Final[object] = await pending - if not isinstance(stream, RustEventStream): - raise TypeError("native async stream opener returned an invalid stream") - return stream - - return open_stream + _CHAT.update(chat) + _ACHAT.update(achat) + _MESSAGES.update(messages) + _AMESSAGES.update(amessages) + _RESPONSES.update(responses) + _ARESPONSES.update(aresponses) def _sync_opener(api: StreamApi) -> RustStreamOpen | None: match api: case "chat_completions": - return _STATE.chat or _native_sync_opener("chat_completions_stream") + return _CHAT.load() case "messages": - return _STATE.messages or _native_sync_opener("messages_stream") + return _MESSAGES.load() case "responses": - return _STATE.responses or _native_sync_opener("responses_stream") + return _RESPONSES.load() def _async_opener(api: StreamApi) -> RustAsyncStreamOpen | None: match api: case "chat_completions": - return _STATE.achat or _native_async_opener("achat_completions_stream") + return _ACHAT.load() case "messages": - return _STATE.amessages or _native_async_opener("amessages_stream") + return _AMESSAGES.load() case "responses": - return _STATE.aresponses or _native_async_opener("aresponses_stream") - - -def _native_exceptions() -> tuple[type[BaseException], type[BaseException]] | None: - native: Final = get_native_bridge() - if native is None: - return None - declined: Final = getattr(native, "RustBridgeDeclined", None) - upstream: Final = getattr(native, "RustUpstreamError", None) - if not isinstance(declined, type) or not isinstance(upstream, type): - return None - return declined, upstream - - -def _handle_open_error(error: Exception, provider: str) -> None: - exceptions: Final = _native_exceptions() - if exceptions is not None and isinstance(error, exceptions[0]): - return - _raise_stream_error(error, provider) - - -def _raise_stream_error(error: Exception, provider: str) -> NoReturn: - exceptions: Final = _native_exceptions() - if exceptions is None or not isinstance(error, exceptions[1]): - raise error - args: Final[tuple[object, ...]] = error.args - status_value: Final = args[0] if args else 0 - message_value: Final = args[1] if len(args) > 1 else str(error) - status: Final = status_value if isinstance(status_value, int) else 0 - message: Final = message_value if isinstance(message_value, str) else str(message_value) - raise APIError( - status_code=status or 500, - message=f"litellm rust typed stream: {message}", - llm_provider=provider, - model="", - ) from error + return _ARESPONSES.load() class TypedEventStreamAdapter: @@ -231,6 +120,7 @@ class TypedEventStreamAdapter: self._stream: Final = stream self._provider: Final = provider self.metadata: Final = stream.metadata + self.core_engine: Final = CoreEngine.RUST self._mode: Literal["sync", "async"] | None = None def _claim(self, mode: Literal["sync", "async"]) -> None: @@ -252,10 +142,10 @@ class TypedEventStreamAdapter: return event def _next_event(self) -> Event | None: - try: - return self._stream.next_event() - except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary - _raise_stream_error(error, self._provider) + return call( + self._stream.next_event, + BridgeErrorContext(route="typed stream", provider=self._provider, model=""), + ) def __aiter__(self) -> AsyncIterator[Event]: self._claim("async") @@ -263,10 +153,10 @@ class TypedEventStreamAdapter: async def __anext__(self) -> Event: self._claim("async") - try: - event: Final = await self._stream.anext_event() - except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary - _raise_stream_error(error, self._provider) + event: Final = await acall( + self._stream.anext_event, + BridgeErrorContext(route="typed stream", provider=self._provider, model=""), + ) if event is None: raise StopAsyncIteration return event @@ -384,12 +274,12 @@ def open_stream( api_base: str | None, extra_headers: Mapping[str, str] | None, timeout: float | httpx.Timeout | None, -) -> TypedEventStreamAdapter | None: +) -> ExecutionResult[TypedEventStreamAdapter | None]: opener: Final = _sync_opener(api) - if opener is None: - return None - try: - stream: Final = opener( + native_call: Final = ( + None + if opener is None + else lambda: opener( request, provider, credentials, @@ -397,10 +287,14 @@ def open_stream( extra_headers, timeout_to_seconds(timeout), ) - except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary - _handle_open_error(error, provider) - return None - return TypedEventStreamAdapter(stream, provider) + ) + return invoke( + native_call=native_call, + fallback=lambda: None, + adapt=lambda stream: TypedEventStreamAdapter(stream, provider), + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route=f"{api} stream", provider=provider, model=""), + ) async def aopen_stream( @@ -412,12 +306,12 @@ async def aopen_stream( api_base: str | None, extra_headers: Mapping[str, str] | None, timeout: float | httpx.Timeout | None, -) -> TypedEventStreamAdapter | None: +) -> ExecutionResult[TypedEventStreamAdapter | None]: opener: Final = _async_opener(api) - if opener is None: - return None - try: - stream: Final = await opener( + native_call: Final = ( + None + if opener is None + else lambda: opener( request, provider, credentials, @@ -425,7 +319,11 @@ async def aopen_stream( extra_headers, timeout_to_seconds(timeout), ) - except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary - _handle_open_error(error, provider) - return None - return TypedEventStreamAdapter(stream, provider) + ) + return await ainvoke( + native_call=native_call, + fallback=async_none, + adapt=lambda stream: TypedEventStreamAdapter(stream, provider), + mode=FallbackMode.PYTHON, + context=BridgeErrorContext(route=f"{api} stream", provider=provider, model=""), + ) diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 3d71f6f8a50..80ea450a21a 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -1,11 +1,21 @@ 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.runtime import ( + UNSET, + BridgeErrorContext, + FallbackMode, + NativeBinding, + Unset, + ainvoke, + async_none, + identity, + invoke, +) from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -39,62 +49,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,18 +84,28 @@ 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), + ) ) + result: Final = 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), + ) + return result.value async def atranscription( @@ -135,15 +120,25 @@ 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), + ) ) + result: Final = 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), + ) + return result.value 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 c8d5db28a1e..321e5173be1 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -1,16 +1,22 @@ """Tests for the optional Rust-backed Anthropic Messages path.""" import importlib -from types import ModuleType from typing import cast import httpx import pytest import litellm -from litellm.exceptions import APIError -from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( + anthropic_messages_stream_hidden_params, +) +from litellm.llms.custom_httpx.llm_http_handler import ( + BaseLLMHTTPHandler, + _anthropic_messages_with_core_engine, +) from litellm.rust_bridge import configuration +from litellm.rust_bridge.messages import RustAmessages, RustMessages +from litellm.rust_bridge.runtime import UNSET, CoreEngine, Unset from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) @@ -18,6 +24,7 @@ from litellm.types.router import GenericLiteLLMParams rust_messages = importlib.import_module("litellm.rust_bridge.messages") rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader") +rust_bridge_runtime = importlib.import_module("litellm.rust_bridge.runtime") FAKE_MESSAGES_RESPONSE: dict[str, object] = { "id": "msg_123", @@ -36,6 +43,39 @@ REQUEST_BODY: dict[str, object] = { } +def _use_test_rust( + enabled: bool = True, + *, + messages: RustMessages | None | Unset = UNSET, + amessages: RustAmessages | None | Unset = UNSET, +) -> None: + rust_messages.set_rust_messages(messages=messages, amessages=amessages) + litellm.use_litellm_rust(enabled) + + +def test_python_messages_response_has_symmetric_provenance(): + response = _anthropic_messages_with_core_engine( + cast(AnthropicMessagesResponse, dict(FAKE_MESSAGES_RESPONSE)), + CoreEngine.PYTHON, + ) + + assert response["_hidden_params"]["core_engine"] == "python" + assert response["_hidden_params"]["additional_headers"] == {"x-litellm-core": "python"} + + +def test_python_messages_stream_provenance_precedes_iteration_and_cannot_be_spoofed(): + hidden_params = anthropic_messages_stream_hidden_params( + httpx.Headers({"x-litellm-core": "rust", "x-litellm-rust": "true", "x-provider": "kept"}) + ) + + assert hidden_params["core_engine"] == "python" + assert hidden_params["additional_headers"]["x-litellm-core"] == "python" + assert "x-litellm-rust" not in hidden_params["additional_headers"] + assert hidden_params["additional_headers"]["llm_provider-x-litellm-core"] == "rust" + assert hidden_params["additional_headers"]["llm_provider-x-litellm-rust"] == "true" + assert hidden_params["additional_headers"]["llm_provider-x-provider"] == "kept" + + class RecordingMessages: def __init__(self) -> None: self.calls: list[dict[str, object]] = [] @@ -102,30 +142,37 @@ class ExplodingAsyncMessages: class RaisingAsyncMessages: - def __init__(self, error: Exception) -> None: + def __init__(self) -> None: self.calls = 0 + + async def __call__(self, **kwargs: object) -> dict[str, object]: + self.calls += 1 + raise RuntimeError("upstream request failed with status 400: bad request") + + +class _FakeDeclined(Exception): + pass + + +class _FakeUpstream(Exception): + pass + + +class _FakeNativeErrors: + RustBridgeDeclined = _FakeDeclined + RustUpstreamError = _FakeUpstream + + +class RaisingTypedAsyncMessages: + def __init__(self, error: Exception) -> None: self.error = error + self.calls = 0 async def __call__(self, **kwargs: object) -> dict[str, object]: self.calls += 1 raise self.error -class FakeBridgeDeclined(Exception): - pass - - -class FakeUpstreamError(Exception): - pass - - -def _install_fake_bridge_exceptions(monkeypatch) -> None: - native_bridge = ModuleType("_native") - native_bridge.RustBridgeDeclined = FakeBridgeDeclined - native_bridge.RustUpstreamError = FakeUpstreamError - monkeypatch.setattr(rust_bridge_loader, "_cached_bridge", native_bridge) - - @pytest.fixture(autouse=True) def _reset_rust_flag(): rust_messages.set_rust_messages(messages=None, amessages=None) @@ -139,7 +186,7 @@ def _reset_rust_flag(): def test_load_rust_messages_returns_injected_impl(): bridge = RecordingMessages() - litellm.use_litellm_rust(True, messages=bridge) + _use_test_rust(True, messages=bridge) assert rust_messages.load_rust_messages() is bridge @@ -155,16 +202,12 @@ def test_bare_use_litellm_rust_still_toggles_ocr(): def test_load_rust_amessages_returns_injected_impl(): bridge = RecordingAsyncMessages() - litellm.use_litellm_rust(True, amessages=bridge) + _use_test_rust(True, amessages=bridge) assert rust_messages.load_rust_amessages() is bridge 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(rust_bridge_runtime, "get_native_bridge", lambda: None) litellm.use_litellm_rust(True) assert rust_messages.load_rust_messages() is None result = rust_messages.messages( @@ -181,7 +224,7 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): def test_messages_wrapper_forwards_args_and_converts_timeout(): bridge = RecordingMessages() - litellm.use_litellm_rust(True, messages=bridge) + _use_test_rust(True, messages=bridge) response = rust_messages.messages( model="claude-sonnet-4-5", @@ -208,7 +251,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout(): @pytest.mark.asyncio async def test_amessages_wrapper_forwards_args(): bridge = RecordingAsyncMessages() - litellm.use_litellm_rust(True, amessages=bridge) + _use_test_rust(True, amessages=bridge) response = await rust_messages.amessages( model="claude-sonnet-4-5", @@ -244,13 +287,17 @@ def _gate(**overrides): @pytest.mark.asyncio async def test_gate_invokes_rust_and_marks_response_header(): bridge = RecordingAsyncMessages() - litellm.use_litellm_rust(True, amessages=bridge) + _use_test_rust(True, amessages=bridge) response = await _gate() assert response is not None assert response["id"] == "msg_123" - assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} + assert response["_hidden_params"]["core_engine"] == "rust" + assert response["_hidden_params"]["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } call = bridge.calls[0] assert call["model"] == "claude-sonnet-4-5" assert call["body"] == REQUEST_BODY @@ -261,53 +308,33 @@ async def test_gate_invokes_rust_and_marks_response_header(): @pytest.mark.asyncio -async def test_gate_falls_back_only_when_bridge_declines(monkeypatch): - _install_fake_bridge_exceptions(monkeypatch) - bridge = RaisingAsyncMessages(FakeBridgeDeclined("unsupported request")) - litellm.use_litellm_rust(True, amessages=bridge) +async def test_gate_does_not_fall_back_for_unknown_bridge_error(): + bridge = RaisingAsyncMessages() + _use_test_rust(True, amessages=bridge) - response = await _gate() - - assert response is None + with pytest.raises(RuntimeError, match="upstream request failed"): + await _gate() assert bridge.calls == 1 @pytest.mark.asyncio -async def test_gate_surfaces_an_upstream_failure_without_fallback(monkeypatch): - _install_fake_bridge_exceptions(monkeypatch) - bridge = RaisingAsyncMessages(FakeUpstreamError(429, "429: rate limited")) - litellm.use_litellm_rust(True, amessages=bridge) +async def test_gate_does_not_fall_back_after_upstream_send(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(rust_bridge_runtime, "get_native_bridge", lambda: _FakeNativeErrors()) + bridge = RaisingTypedAsyncMessages(_FakeUpstream(429, "rate limited")) + _use_test_rust(True, amessages=bridge) - with pytest.raises(APIError) as exc_info: + with pytest.raises(litellm.APIError, match="rate limited"): await _gate() - - assert exc_info.value.status_code == 429 - assert "429: rate limited" in str(exc_info.value) assert bridge.calls == 1 @pytest.mark.asyncio -async def test_gate_maps_statusless_upstream_failure_to_500_without_fallback(monkeypatch): - _install_fake_bridge_exceptions(monkeypatch) - bridge = RaisingAsyncMessages(FakeUpstreamError(0, "request timed out")) - litellm.use_litellm_rust(True, amessages=bridge) - - with pytest.raises(APIError) as exc_info: - await _gate() - - assert exc_info.value.status_code == 500 - assert "request timed out" in str(exc_info.value) - assert bridge.calls == 1 - - -@pytest.mark.asyncio -async def test_gate_reraises_an_unknown_bridge_failure(): - bridge = RaisingAsyncMessages(RuntimeError("unknown bridge failure")) - litellm.use_litellm_rust(True, amessages=bridge) - - with pytest.raises(RuntimeError, match="unknown bridge failure"): - await _gate() +async def test_gate_falls_back_when_rust_declines(monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(rust_bridge_runtime, "get_native_bridge", lambda: _FakeNativeErrors()) + bridge = RaisingTypedAsyncMessages(_FakeDeclined("unsupported")) + _use_test_rust(True, amessages=bridge) + assert await _gate() is None assert bridge.calls == 1 @@ -322,22 +349,10 @@ async def test_gate_skips_rust_when_flag_absent(): assert bridge.calls == 0 -@pytest.mark.asyncio -async def test_gate_uses_process_enable_without_request_override(): - bridge = RecordingAsyncMessages() - rust_messages.set_rust_messages(amessages=bridge) - litellm.use_litellm_rust(True) - - response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure")) - - assert response is not None - assert bridge.calls[0]["custom_llm_provider"] == "azure_ai" - - @pytest.mark.asyncio async def test_gate_skips_rust_when_flag_false(): bridge = ExplodingAsyncMessages() - litellm.use_litellm_rust(True, amessages=bridge) + _use_test_rust(True, amessages=bridge) response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False)) @@ -348,7 +363,7 @@ async def test_gate_skips_rust_when_flag_false(): @pytest.mark.asyncio async def test_gate_invokes_rust_for_native_anthropic_provider(): bridge = RecordingAsyncMessages() - litellm.use_litellm_rust(True, amessages=bridge) + _use_test_rust(True, amessages=bridge) response = await _gate( custom_llm_provider="anthropic", @@ -359,7 +374,11 @@ async def test_gate_invokes_rust_for_native_anthropic_provider(): ) assert response is not None - assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} + assert response["_hidden_params"]["core_engine"] == "rust" + assert response["_hidden_params"]["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } assert bridge.calls[0]["custom_llm_provider"] == "anthropic" assert bridge.calls[0]["api_key"] == "sk-ant" @@ -397,7 +416,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch): @pytest.mark.asyncio async def test_gate_skips_rust_for_unsupported_provider(): bridge = ExplodingAsyncMessages() - litellm.use_litellm_rust(True, amessages=bridge) + _use_test_rust(True, amessages=bridge) response = await _gate(custom_llm_provider="openai") @@ -408,7 +427,7 @@ async def test_gate_skips_rust_for_unsupported_provider(): @pytest.mark.asyncio async def test_gate_skips_rust_for_agentic_hook(): bridge = ExplodingAsyncMessages() - litellm.use_litellm_rust(True, amessages=bridge) + _use_test_rust(True, amessages=bridge) response = await _gate(has_agentic_hook=True) @@ -419,7 +438,7 @@ async def test_gate_skips_rust_for_agentic_hook(): @pytest.mark.asyncio async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag(): bridge = RecordingAsyncMessages() - litellm.use_litellm_rust(True, amessages=bridge) + _use_test_rust(True, amessages=bridge) streaming_body = {**REQUEST_BODY, "stream": True} response = await _gate( @@ -428,7 +447,11 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag(): ) assert response is not None - assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} + assert response["_hidden_params"]["core_engine"] == "rust" + assert response["_hidden_params"]["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } assert "stream" not in bridge.calls[0]["body"] assert bridge.calls[0]["body"] == REQUEST_BODY @@ -438,7 +461,11 @@ async def test_fake_stream_wraps_rust_response_as_anthropic_sse(): response = cast(AnthropicMessagesResponse, dict(FAKE_MESSAGES_RESPONSE)) stream = BaseLLMHTTPHandler._rust_anthropic_messages_fake_stream(response) - assert stream._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} + assert stream._hidden_params["core_engine"] == "rust" + assert stream._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } chunks = [chunk async for chunk in stream] joined = b"".join(chunks) @@ -451,11 +478,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(rust_bridge_runtime, "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 bd750a47f63..11ee5ec4893 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 @@ -8,6 +8,7 @@ import pytest import litellm from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.llms.anthropic.chat.handler import ModelResponseIterator, make_call +from litellm.rust_bridge import runtime as bridge_runtime from litellm.types.llms.openai import ( ChatCompletionToolCallChunk, ChatCompletionToolCallFunctionChunk, @@ -2286,7 +2287,11 @@ class TestRustChatCompletionsHook: 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 response._hidden_params["core_engine"] == "rust" + assert response._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } assert len(seen["call"]) == 1 def test_the_core_receives_the_untranslated_openai_messages(self): @@ -2427,7 +2432,7 @@ class TestRustChatCompletionsHook: def declining_native(**_kwargs): raise _Declined("blank message text") - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr(bridge_runtime, "get_native_bridge", lambda: _FakeNative()) bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) @@ -2459,7 +2464,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr(bridge_runtime, "get_native_bridge", lambda: _FakeNative()) async def declining_native(**_kwargs): raise _Declined("blank message text") @@ -2501,7 +2506,11 @@ class TestRustChatCompletionsHook: ) assert result.choices[0].message.content == "hello from rust" - assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} + assert result._hidden_params["core_engine"] == "rust" + assert result._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } assert not python_call.called @@ -2519,7 +2528,7 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + monkeypatch.setattr(bridge_runtime, "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 8e67a7e3438..22d22b8eebb 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 @@ -15,6 +15,7 @@ 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 chat_completions as bridge +from litellm.rust_bridge import runtime as bridge_runtime from litellm.types.utils import ModelResponse RUST_RESPONSE = { @@ -118,7 +119,11 @@ def test_rust_true_serves_the_call_and_stamps_the_header(): response = _run() assert response.choices[0].message.content == "hello from rust" - assert response._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} + assert response._hidden_params["core_engine"] == "rust" + assert response._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } assert len(seen["call"]) == 1 @@ -206,7 +211,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_runtime, "get_native_bridge", lambda: _FakeNative()) async def declining_native(**_kwargs): raise _Declined("blank message text") @@ -256,7 +261,11 @@ 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 result._hidden_params["core_engine"] == "rust" + assert result._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } assert not python_call.called @@ -283,7 +292,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_runtime, "get_native_bridge", lambda: _FakeNative()), patch.object( BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS ), @@ -390,7 +399,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_runtime, "get_native_bridge", lambda: _FakeNative()): bridge.set_rust_chat_completions( decline=lambda **_kwargs: None, chat_completions=declining_native ) @@ -475,7 +484,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_runtime, "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..7839f2557a8 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -11,6 +11,8 @@ import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge import configuration +from litellm.rust_bridge.ocr import RustAocr, RustOcr +from litellm.rust_bridge.runtime import UNSET, Unset # `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"), + importlib.import_module("litellm.rust_bridge.runtime"), "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"), + importlib.import_module("litellm.rust_bridge.runtime"), "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"), + importlib.import_module("litellm.rust_bridge.runtime"), "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( @@ -460,6 +450,11 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): assert isinstance(response, OCRResponse) assert response.pages[0].markdown == "hello world" + assert response._hidden_params["core_engine"] == "rust" + assert response._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } assert bridge.calls[0] == { "model": "mistral-ocr-latest", "document": DOCUMENT, @@ -477,7 +472,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 +484,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 +503,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 +525,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 +551,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 +574,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 +591,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 +611,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 +698,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") @@ -724,6 +719,7 @@ async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge): assert isinstance(response, OCRResponse) assert response.pages[0].markdown == "hello world" + assert response._hidden_params["core_engine"] == "rust" assert len(fake_async_bridge.calls) == 1 call = fake_async_bridge.calls[0] assert call["model"] == "mistral-ocr-latest" @@ -748,7 +744,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 +772,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, @@ -802,6 +798,29 @@ def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch): assert captured.get("called") is True # Python path was used assert isinstance(response, OCRResponse) + assert response._hidden_params["core_engine"] == "python" + assert response._hidden_params["additional_headers"] == {"x-litellm-core": "python"} + + +def test_ocr_unsupported_provider_skips_rust_and_reports_python(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 response._hidden_params["core_engine"] == "python" + assert response._hidden_params["additional_headers"] == {"x-litellm-core": "python"} + assert bridge.calls == [] def test_ocr_provider_configs_expose_api_key_env_vars(): diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index df14224af5c..d0926c5df08 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -1241,6 +1241,28 @@ class TestProxyBaseLLMRequestProcessing: breakdown_none = _get_cost_breakdown_from_logging_obj(None) assert all(value is None for value in breakdown_none) + def test_get_custom_headers_sets_and_reserves_core_provenance(self): + mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth) + mock_user_api_key_dict.tpm_limit = None + mock_user_api_key_dict.rpm_limit = None + mock_user_api_key_dict.max_budget = None + mock_user_api_key_dict.spend = 0 + + python_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=mock_user_api_key_dict, + **{"x-litellm-core": "rust", "x-litellm-rust": "true"}, + ) + rust_headers = ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=mock_user_api_key_dict, + hidden_params={"core_engine": "rust"}, + **{"x-litellm-core": "python", "x-litellm-rust": "false"}, + ) + + assert python_headers["x-litellm-core"] == "python" + assert "x-litellm-rust" not in python_headers + assert rust_headers["x-litellm-core"] == "rust" + assert rust_headers["x-litellm-rust"] == "true" + def test_get_custom_headers_key_spend_includes_response_cost(self): """ Test that x-litellm-key-spend header includes the current request's response_cost. diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index b7f5c73d2a4..9cffb2b16af 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -4,8 +4,10 @@ from collections.abc import Mapping import pytest +from litellm.exceptions import APIError 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 configuration, responses_websocket, runtime +from litellm.rust_bridge.runtime import CoreEngine from litellm.types.router import GenericLiteLLMParams @@ -35,6 +37,24 @@ class _ClosedNativeConnection: return None +class _FailingNativeConnection(_ClosedNativeConnection): + async def send_event(self, event: Mapping[str, object]) -> None: + raise _Upstream(503, "session send failed") + + +class _Declined(Exception): + pass + + +class _Upstream(Exception): + pass + + +class _NativeErrors: + RustBridgeDeclined = _Declined + RustUpstreamError = _Upstream + + class _FakeNativeBridge: @classmethod async def connect( @@ -85,19 +105,17 @@ 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(runtime, "get_native_bridge", lambda: None) - assert ( - await responses_websocket.connect( - provider="openai", - api_key=None, - api_base="https://example.test", - headers={}, - timeout=None, - ) - is None + result = await responses_websocket.connect( + provider="openai", + api_key=None, + api_base="https://example.test", + headers={}, + timeout=None, ) + assert result.value is None + assert result.source is CoreEngine.PYTHON @pytest.mark.asyncio @@ -106,7 +124,7 @@ async def test_enabled_bridge_connects_and_adapts_socket( ) -> None: responses_websocket.set_rust_responses_websocket(connection=_FakeNativeBridge) - connection = await responses_websocket.connect( + result = await responses_websocket.connect( provider="openai", api_key="key", api_base="https://example.test", @@ -114,7 +132,20 @@ async def test_enabled_bridge_connects_and_adapts_socket( timeout=1.0, ) + assert result.source is CoreEngine.RUST + connection = result.value assert connection is not None + assert connection.core_engine is CoreEngine.RUST await connection.send('{"type":"response.create","model":"gpt-5"}') assert await connection.recv() == '{"type":"response.completed"}' await connection.close() + + +@pytest.mark.asyncio +async def test_connected_rust_session_never_falls_back(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(runtime, "get_native_bridge", lambda: _NativeErrors()) + connection = responses_websocket._ConnectionAdapter(_FailingNativeConnection()) + + assert connection.core_engine is CoreEngine.RUST + with pytest.raises(APIError, match="session send failed"): + await connection.send('{"type":"response.create"}') diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index 03921133c77..668bf73cc1e 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -11,7 +11,7 @@ import pytest import litellm from litellm.rust_bridge import chat_completions as bridge -from litellm.rust_bridge import configuration +from litellm.rust_bridge import configuration, runtime from litellm.types.utils import ModelResponse RUST_RESPONSE = { @@ -54,7 +54,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(runtime, "get_native_bridge", lambda: _FakeNative()) def _hide_native_bridge(monkeypatch): @@ -63,31 +63,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(runtime, "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 +98,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 +108,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 +151,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 +174,30 @@ 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 exploding(**_kwargs): - raise RuntimeError("boom") - - bridge.set_rust_chat_completions(decline=exploding) - assert _accepts() is 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 _call_kwargs(model_response: ModelResponse) -> dict: @@ -247,6 +215,18 @@ 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): native = _RecordingCall() @@ -254,7 +234,7 @@ class TestSyncCall: 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 +243,66 @@ 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._hidden_params["core_engine"] == "rust" + assert result._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "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_stamped_python(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 + assert result._hidden_params["core_engine"] == "python" + assert result._hidden_params["additional_headers"] == {"x-litellm-core": "python"} 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"} + assert result._hidden_params["core_engine"] == "rust" + assert result._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "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 +351,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 +367,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/rust_bridge/test_runtime.py b/tests/test_litellm/rust_bridge/test_runtime.py new file mode 100644 index 00000000000..d165b08edb8 --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_runtime.py @@ -0,0 +1,238 @@ +from __future__ import annotations + +from collections.abc import Callable + +import pytest + +from litellm.exceptions import APIError +from litellm.rust_bridge import runtime +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + CoreEngine, + FallbackMode, + RustDeclined, + RustHandled, + RustUnavailable, +) + + +class Declined(Exception): + pass + + +class Upstream(Exception): + pass + + +class Native: + RustBridgeDeclined = Declined + RustUpstreamError = Upstream + + +CONTEXT = BridgeErrorContext(route="messages", provider="anthropic", model="claude") + + +@pytest.fixture(autouse=True) +def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(runtime, "get_native_bridge", lambda: Native) + + +@pytest.mark.parametrize( + ("native_call", "expected", "source", "fallback_calls"), + [ + (lambda: "native", "adapted native", CoreEngine.RUST, 0), + (None, "fallback", CoreEngine.PYTHON, 1), + (lambda: (_ for _ in ()).throw(Declined("not sent")), "fallback", CoreEngine.PYTHON, 1), + ], +) +def test_invoke_uses_fallback_only_when_safe( + native_call: Callable[[], str] | None, + expected: str, + source: CoreEngine, + fallback_calls: int, +) -> None: + calls = 0 + + def fallback() -> str: + nonlocal calls + calls += 1 + return "fallback" + + result = runtime.invoke( + native_call=native_call, + fallback=fallback, + adapt=lambda value: f"adapted {value}", + mode=FallbackMode.PYTHON, + context=CONTEXT, + ) + + assert result.value == expected + assert result.source is source + assert calls == fallback_calls + + +def test_invoke_converts_upstream_error_without_fallback() -> None: + fallback_calls = 0 + + def fallback() -> str: + nonlocal fallback_calls + fallback_calls += 1 + return "fallback" + + with pytest.raises(APIError) as exc_info: + runtime.invoke( + native_call=lambda: (_ for _ in ()).throw(Upstream(429, "rate limited")), + fallback=fallback, + adapt=str, + mode=FallbackMode.PYTHON, + context=CONTEXT, + ) + + assert exc_info.value.status_code == 429 + assert "rate limited" in str(exc_info.value) + assert exc_info.value.headers == {"x-litellm-core": "rust", "x-litellm-rust": "true"} + assert fallback_calls == 0 + + +def test_invoke_propagates_unknown_error_without_fallback() -> None: + with pytest.raises(ValueError, match="bug"): + runtime.invoke( + native_call=lambda: (_ for _ in ()).throw(ValueError("bug")), + fallback=lambda: pytest.fail("fallback must not run"), + adapt=str, + mode=FallbackMode.PYTHON, + context=CONTEXT, + ) + + +@pytest.mark.parametrize( + "native_call", + [None, lambda: (_ for _ in ()).throw(Declined("unsupported"))], +) +def test_rust_required_route_never_falls_back( + native_call: Callable[[], str] | None, +) -> None: + with pytest.raises(RuntimeError): + runtime.invoke( + native_call=native_call, + fallback=lambda: pytest.fail("fallback must not run"), + adapt=str, + mode=FallbackMode.RUST_REQUIRED, + context=CONTEXT, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", [None, Declined("not sent")]) +async def test_ainvoke_success_or_safe_fallback(failure: Exception | None) -> None: + fallback_calls = 0 + + async def native_call() -> str: + if failure is not None: + raise failure + return "native" + + async def fallback() -> str: + nonlocal fallback_calls + fallback_calls += 1 + return "fallback" + + result = await runtime.ainvoke( + native_call=native_call, + fallback=fallback, + adapt=lambda value: f"adapted {value}", + mode=FallbackMode.PYTHON, + context=CONTEXT, + ) + + assert result.value == ("fallback" if failure else "adapted native") + assert result.source is (CoreEngine.PYTHON if failure else CoreEngine.RUST) + assert fallback_calls == (1 if failure else 0) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", [Upstream(503, "unavailable"), ValueError("bug")]) +async def test_ainvoke_never_falls_back_after_unsafe_failure(failure: Exception) -> None: + async def native_call() -> str: + raise failure + + async def fallback() -> str: + pytest.fail("fallback must not run") + + expected = APIError if isinstance(failure, Upstream) else ValueError + with pytest.raises(expected): + await runtime.ainvoke( + native_call=native_call, + fallback=fallback, + adapt=str, + mode=FallbackMode.PYTHON, + context=CONTEXT, + ) + + +@pytest.mark.parametrize( + ("native_call", "expected_type"), + [ + (lambda: "native", RustHandled), + (lambda: (_ for _ in ()).throw(Declined("unsupported")), RustDeclined), + (None, RustUnavailable), + ], +) +def test_attempt_classifies_bridge_control_flow( + native_call: Callable[[], str] | None, + expected_type: type[RustHandled[str]] | type[RustDeclined] | type[RustUnavailable], +) -> None: + result = runtime.attempt(native_call=native_call, adapt=str, context=CONTEXT) + + assert isinstance(result, expected_type) + if isinstance(result, RustDeclined): + assert result.reason == "unsupported" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("failure", "expected_type"), + [ + (None, RustHandled), + (Declined("unsupported"), RustDeclined), + ], +) +async def test_aattempt_classifies_bridge_control_flow( + failure: Exception | None, + expected_type: type[RustHandled[str]] | type[RustDeclined], +) -> None: + async def native_call() -> str: + if failure is not None: + raise failure + return "native" + + result = await runtime.aattempt(native_call=native_call, adapt=str, context=CONTEXT) + + assert isinstance(result, expected_type) + if isinstance(result, RustDeclined): + assert result.reason == "unsupported" + + +@pytest.mark.asyncio +async def test_aattempt_classifies_unavailable_bridge() -> None: + result = await runtime.aattempt(native_call=None, adapt=str, context=CONTEXT) + + assert isinstance(result, RustUnavailable) + + +def test_execution_hidden_params_overwrites_reserved_provider_headers() -> None: + hidden_params = runtime.execution_hidden_params( + { + "additional_headers": { + "X-LiteLLM-Core": "rust", + "x-litellm-rust": "true", + "x-provider": "kept", + } + }, + CoreEngine.PYTHON, + ) + + assert hidden_params == { + "core_engine": "python", + "additional_headers": {"x-provider": "kept", "x-litellm-core": "python"}, + } diff --git a/tests/test_litellm/rust_bridge/test_streaming.py b/tests/test_litellm/rust_bridge/test_streaming.py index 04624425dc4..2407a25bada 100644 --- a/tests/test_litellm/rust_bridge/test_streaming.py +++ b/tests/test_litellm/rust_bridge/test_streaming.py @@ -7,7 +7,8 @@ from unittest.mock import MagicMock import pytest from litellm.exceptions import APIError -from litellm.rust_bridge import streaming +from litellm.rust_bridge import runtime, streaming +from litellm.rust_bridge.runtime import CoreEngine class _FakeEventStream: @@ -101,6 +102,24 @@ class _FailingOpen: raise self._error +class _MidstreamFailingStream(_FakeEventStream): + def next_event(self) -> Mapping[str, object] | None: + raise _Declined("decoder rejected event") + + +class _MidstreamFailingOpen: + def __call__( + self, + request: Mapping[str, object], + provider: str, + credentials: Mapping[str, str] | None, + api_base: str | None, + extra_headers: Mapping[str, str] | None, + timeout_seconds: float | None, + ) -> _FakeEventStream: + return _MidstreamFailingStream(()) + + @pytest.fixture(autouse=True) def reset_bridge() -> Iterator[None]: streaming.set_rust_streaming( @@ -133,7 +152,7 @@ def _chat_event(text: str) -> Mapping[str, object]: def test_unavailable_native_bridge_falls_back(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(streaming, "get_native_bridge", lambda: None) + monkeypatch.setattr(runtime, "get_native_bridge", lambda: None) result: Final = streaming.open_stream( api="chat_completions", provider="anthropic", @@ -144,15 +163,16 @@ def test_unavailable_native_bridge_falls_back(monkeypatch: pytest.MonkeyPatch) - timeout=None, ) - assert result is None + assert result.value is None + assert result.source is CoreEngine.PYTHON -def test_declined_open_failure_falls_back( +def test_guaranteed_not_sent_open_failure_falls_back( monkeypatch: pytest.MonkeyPatch, ) -> None: python_calls = 0 - monkeypatch.setattr(streaming, "get_native_bridge", lambda: _NativeErrors()) + monkeypatch.setattr(runtime, "get_native_bridge", lambda: _NativeErrors()) streaming.set_rust_streaming( chat=_FailingOpen(_Declined("unsupported request")), ) @@ -166,16 +186,17 @@ def test_declined_open_failure_falls_back( extra_headers=None, timeout=None, ) - if result is None: + if result.value is None: python_calls += 1 assert python_calls == 1 + assert result.source is CoreEngine.PYTHON -def test_upstream_open_failure_never_falls_back( +def test_possibly_sent_open_failure_never_falls_back( monkeypatch: pytest.MonkeyPatch, ) -> None: - monkeypatch.setattr(streaming, "get_native_bridge", lambda: _NativeErrors()) + monkeypatch.setattr(runtime, "get_native_bridge", lambda: _NativeErrors()) streaming.set_rust_streaming( chat=_FailingOpen(_Upstream(503, "connection closed after request")), ) @@ -206,10 +227,12 @@ def test_sync_typed_events_preserve_shape_metadata_and_close() -> None: timeout=1.0, ) - assert result is not None - assert tuple(event["text"] for event in result) == ("one", "two") - assert result.metadata["provider"] == "anthropic" - result.close() + assert result.source is CoreEngine.RUST + assert result.value is not None + assert result.value.core_engine is CoreEngine.RUST + assert tuple(event["text"] for event in result.value) == ("one", "two") + assert result.value.metadata["provider"] == "anthropic" + result.value.close() assert opener.calls == 1 @@ -231,6 +254,25 @@ def test_chat_events_flow_through_custom_stream_wrapper() -> None: assert chunk.choices[0].delta.content == "hello" +def test_rust_stream_never_falls_back_after_open(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(runtime, "get_native_bridge", lambda: _NativeErrors()) + streaming.set_rust_streaming(chat=_MidstreamFailingOpen()) + result: Final = streaming.open_stream( + api="chat_completions", + provider="anthropic", + request={"model": "claude", "messages": []}, + credentials=None, + api_base=None, + extra_headers=None, + timeout=None, + ) + + assert result.source is CoreEngine.RUST + assert result.value is not None + with pytest.raises(_Declined, match="decoder rejected event"): + next(result.value) + + @pytest.mark.asyncio async def test_async_typed_events_and_cancellation() -> None: opener: Final = _RecordingAsyncOpen((_chat_event("one"), _chat_event("two"))) @@ -245,10 +287,11 @@ async def test_async_typed_events_and_cancellation() -> None: timeout=None, ) - assert result is not None - collected: Final = tuple([event async for event in result]) + assert result.source is CoreEngine.RUST + assert result.value is not None + collected: Final = tuple([event async for event in result.value]) assert tuple(event["text"] for event in collected) == ("one", "two") - await result.aclose() + await result.value.aclose() def test_messages_events_are_wrapped_in_existing_sse_bytes() -> None: diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index bbeb6c38f78..9aa2dc822d0 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_runtime = importlib.import_module("litellm.rust_bridge.runtime") 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_runtime, "get_native_bridge", lambda: None) assert rust_bridge.load_rust_transcription() is None assert rust_bridge.load_rust_atranscription() is None @@ -132,6 +133,11 @@ def test_bedrock_transcription_uses_rust_only_path() -> None: rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) assert response.text == "rust" + assert response._hidden_params["core_engine"] == "rust" + assert response._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + } @pytest.mark.asyncio @@ -149,3 +155,8 @@ async def test_bedrock_atranscription_uses_rust_only_path() -> None: rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) assert response.text == "rust" + assert response._hidden_params["core_engine"] == "rust" + assert response._hidden_params["additional_headers"] == { + "x-litellm-core": "rust", + "x-litellm-rust": "true", + }