feat(proxy): centralize Rust fallback and provenance

This commit is contained in:
Yujong Lee 2026-09-01 11:20:13 -07:00
parent b461ae6d63
commit 2395450714
26 changed files with 1726 additions and 1056 deletions

View file

@ -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<AppState> {
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);
}
}

View file

@ -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

View file

@ -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,
)

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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,

View file

@ -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)

View file

@ -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()

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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=""),
)

View file

@ -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

View file

@ -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=""),
)

View file

@ -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

View file

@ -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()

View file

@ -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")

View file

@ -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
)

View file

@ -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():

View file

@ -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.

View file

@ -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"}')

View file

@ -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):

View file

@ -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"},
}

View file

@ -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:

View file

@ -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",
}