This commit is contained in:
yujonglee 2026-09-06 06:35:42 +00:00 • committed by GitHub
commit 1e1bf2a60d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
29 changed files with 2137 additions and 1747 deletions

View file

@ -1,9 +1,9 @@
{
"reportAny": {
"limit": 13429
"limit": 13426
},
"reportArgumentType": {
"limit": 2198
"limit": 2192
},
"reportAssignmentType": {
"limit": 319
@ -24,13 +24,13 @@
"limit": 19
},
"reportExplicitAny": {
"limit": 3369
"limit": 3368
},
"reportFunctionMemberAccess": {
"limit": 7
},
"reportGeneralTypeIssues": {
"limit": 101
"limit": 100
},
"reportIncompatibleMethodOverride": {
"limit": 56
@ -57,7 +57,7 @@
"limit": 5570
},
"reportMissingTypeArgument": {
"limit": 15281
"limit": 15279
},
"reportMissingTypeStubs": {
"limit": 40
@ -90,7 +90,7 @@
"limit": 8
},
"reportReturnType": {
"limit": 180
"limit": 178
},
"reportTypedDictNotRequiredAccess": {
"limit": 22
@ -99,31 +99,31 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44358
"limit": 44019
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38271
"limit": 38266
},
"reportUnknownParameterType": {
"limit": 19584
"limit": 19583
},
"reportUnknownVariableType": {
"limit": 29814
"limit": 29810
},
"reportUnnecessaryCast": {
"limit": 110
},
"reportUnnecessaryComparison": {
"limit": 687
"limit": 686
},
"reportUnnecessaryContains": {
"limit": 4
},
"reportUnnecessaryIsInstance": {
"limit": 816
"limit": 815
},
"reportUntypedBaseClass": {
"limit": 0

View file

@ -27,6 +27,8 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.rust_bridge.dispatch import anative_first, native_first
from litellm.rust_bridge.runtime import DispatchResult
from litellm.types.llms.anthropic import (
ContentBlockDelta,
ContentBlockStart,
@ -368,15 +370,7 @@ class AnthropicChatCompletion(BaseLLM):
if config is None:
raise ValueError(f"Provider config not found for model: {model} and provider: {custom_llm_provider}")
def build_request() -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream
"""Translate the request the Python way, returning `(headers, data)`.
The pair stays mutable because the streaming path rewrites it in
place (`data["stream"] = True`) before sending.
Shared by the normal path and by the Rust path's fallback, which
builds it only when the Rust call did not serve the request.
"""
def prepare_python() -> tuple[dict[str, str], dict[str, object]]: # mutable-ok: stream mutates data
request_data: Final = config.transform_request(
model=model,
messages=messages,
@ -384,12 +378,29 @@ class AnthropicChatCompletion(BaseLLM):
litellm_params=litellm_params,
headers=headers,
)
return update_request_with_filtered_beta(
python_headers, data = update_request_with_filtered_beta(
headers=headers,
request_data=request_data,
provider=custom_llm_provider,
)
## 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": python_headers,
},
)
print_verbose(f"_is_function_call: {_is_function_call}")
return python_headers, data
# 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
@ -406,67 +417,26 @@ class AnthropicChatCompletion(BaseLLM):
litellm_params=litellm_params,
stream=stream,
)
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
"model": model,
"messages": messages,
**rust_optional_params,
},
"api_base": api_base,
"headers": headers,
}
if serves_via_rust:
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
"model": model,
"messages": messages,
**rust_optional_params,
},
"api_base": api_base,
"headers": headers,
}
logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args)
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key=api_key,
additional_args=rust_logging_args,
)
if acompletion is True:
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key=api_key,
additional_args=rust_logging_args,
)
async def python_fallback() -> "ModelResponse | CustomStreamWrapper":
# pre_call already fired for this request above. The Rust
# path only declines before the provider is called, so this
# is the same attempt continuing, not a second one.
fallback_headers, fallback_data = build_request()
return await self.acompletion_function(
model=model,
messages=messages,
data=fallback_data,
api_base=api_base,
custom_prompt_dict=custom_prompt_dict,
model_response=model_response,
print_verbose=print_verbose,
encoding=encoding,
api_key=api_key,
provider_config=config,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
_is_function_call=_is_function_call,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=fallback_headers,
client=client,
json_mode=json_mode,
timeout=timeout,
)
return rust_chat_completions_bridge.achat_completions_or_fallback(
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
python_fallback=python_fallback,
)
rust_response: Final = rust_chat_completions_bridge.chat_completions(
def native_completion() -> DispatchResult[ModelResponse]:
return rust_chat_completions_bridge.chat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
@ -477,34 +447,37 @@ class AnthropicChatCompletion(BaseLLM):
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
eligible=serves_via_rust,
)
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,
async def native_acompletion() -> DispatchResult[ModelResponse]:
return await rust_chat_completions_bridge.achat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": api_base,
"headers": headers,
},
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
eligible=serves_via_rust,
)
print_verbose(f"_is_function_call: {_is_function_call}")
if acompletion is True:
@anative_first(
native=native_acompletion,
route="chat_completions",
errors=lambda: rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model),
)
async def execute_async() -> ModelResponse | CustomStreamWrapper:
headers, data = prepare_python()
if (
stream is True
): # if function call - fake the streaming (need complete blocks for output parsing in openai format)
print_verbose("makes async anthropic streaming POST request")
data["stream"] = stream
return self.acompletion_stream_function(
return await self.acompletion_stream_function(
model=model,
messages=messages,
data=data,
@ -526,7 +499,7 @@ class AnthropicChatCompletion(BaseLLM):
client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None),
)
else:
return self.acompletion_function(
return await self.acompletion_function(
model=model,
messages=messages,
data=data,
@ -548,7 +521,14 @@ class AnthropicChatCompletion(BaseLLM):
json_mode=json_mode,
timeout=timeout,
)
else:
@native_first(
native=native_completion,
route="chat_completions",
errors=lambda: rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model),
)
def execute_sync() -> ModelResponse | CustomStreamWrapper:
headers, data = prepare_python()
## COMPLETION CALL
if (
stream is True
@ -580,13 +560,12 @@ class AnthropicChatCompletion(BaseLLM):
)
else:
if client is None or not isinstance(client, HTTPHandler):
client = _get_httpx_client(params={"timeout": timeout})
else:
client = client
python_client: Final = (
client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout})
)
try:
response: Final = client.post(
response: Final = python_client.post(
api_base,
headers=headers,
data=json.dumps(data),
@ -607,20 +586,21 @@ class AnthropicChatCompletion(BaseLLM):
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 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 execute_async() if acompletion else execute_sync()
def embedding(self):
# logic for parsing in - calling - parsing out model embedding calls

View file

@ -1,13 +1,23 @@
import base64
from typing import Final
from typing import Final, NoReturn
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.dispatch import PROPAGATE, anative_first, native_first
from litellm.rust_bridge.runtime import DispatchResult, adapt_result
from litellm.types.utils import FileTypes, TranscriptionResponse
def _unavailable() -> NoReturn:
raise RuntimeError("Rust audio transcription bridge is unavailable")
async def _aunavailable() -> NoReturn:
_unavailable()
class BedrockAudioTranscriptionRustDispatch:
@staticmethod
def _audio_payload(audio_file: FileTypes) -> dict[str, object]:
@ -31,6 +41,37 @@ class BedrockAudioTranscriptionRustDispatch:
"filename": processed_audio.filename,
}
def _attempt_audio_transcriptions(
self,
*,
model: str,
audio_file: FileTypes,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> DispatchResult[TranscriptionResponse]:
result: Final = rust_transcription_bridge.transcription(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
)
return adapt_result(result, lambda response: TranscriptionResponse(**response))
@native_first(
native=_attempt_audio_transcriptions,
route="audio transcription",
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: (
PROPAGATE
),
)
def audio_transcriptions(
self,
*,
@ -43,7 +84,21 @@ class BedrockAudioTranscriptionRustDispatch:
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> TranscriptionResponse:
rust_response: Final = rust_transcription_bridge.transcription(
_unavailable()
async def _attempt_async_audio_transcriptions(
self,
*,
model: str,
audio_file: FileTypes,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> DispatchResult[TranscriptionResponse]:
result: Final = await rust_transcription_bridge.atranscription(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
@ -53,10 +108,15 @@ class BedrockAudioTranscriptionRustDispatch:
optional_params=optional_params,
timeout=timeout,
)
if rust_response is None:
raise RuntimeError("Rust audio transcription bridge is unavailable")
return TranscriptionResponse(**rust_response)
return adapt_result(result, lambda response: TranscriptionResponse(**response))
@anative_first(
native=_attempt_async_audio_transcriptions,
route="audio transcription",
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: (
PROPAGATE
),
)
async def async_audio_transcriptions(
self,
*,
@ -69,16 +129,4 @@ class BedrockAudioTranscriptionRustDispatch:
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> TranscriptionResponse:
rust_response: Final = await rust_transcription_bridge.atranscription(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
)
if rust_response is None:
raise RuntimeError("Rust audio transcription bridge is unavailable")
return TranscriptionResponse(**rust_response)
await _aunavailable()

View file

@ -18,6 +18,8 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.rust_bridge.dispatch import anative_first, native_first
from litellm.rust_bridge.runtime import DispatchResult
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
@ -406,54 +408,25 @@ class BedrockConverseLLM(BaseAWSLLM):
litellm_params=litellm_params,
stream=stream,
)
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
"messages": messages,
**optional_params,
},
"api_base": proxy_endpoint_url,
"headers": headers,
}
if serves_via_rust:
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
"messages": messages,
**optional_params,
},
"api_base": proxy_endpoint_url,
"headers": headers,
}
logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args)
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key="",
additional_args=rust_logging_args,
)
if acompletion:
return rust_chat_completions_bridge.achat_completions_or_fallback(
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
api_base=proxy_endpoint_url,
custom_llm_provider="bedrock",
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
python_fallback=lambda: self.async_completion(
model=model,
messages=messages,
api_base=proxy_endpoint_url,
model_response=model_response,
encoding=encoding,
logging_obj=logging_obj,
optional_params=optional_params,
stream=stream,
litellm_params=litellm_params,
logger_fn=logger_fn,
headers=headers,
timeout=timeout,
client=client,
credentials=credentials,
api_key=api_key,
skip_pre_call_logging=True,
),
)
rust_response: Final = rust_chat_completions_bridge.chat_completions(
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key="",
additional_args=rust_logging_args,
)
def native_completion() -> DispatchResult[ModelResponse]:
return rust_chat_completions_bridge.chat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
@ -464,16 +437,33 @@ class BedrockConverseLLM(BaseAWSLLM):
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
eligible=serves_via_rust,
)
if rust_response is not None:
return rust_response
### ROUTING (ASYNC, STREAMING, SYNC)
if acompletion:
if isinstance(client, HTTPHandler):
client = None
async def native_acompletion() -> DispatchResult[ModelResponse]:
return await rust_chat_completions_bridge.achat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
api_base=proxy_endpoint_url,
custom_llm_provider="bedrock",
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
eligible=serves_via_rust,
)
@anative_first(
native=native_acompletion,
route="chat_completions",
errors=lambda: rust_chat_completions_bridge.error_handling("bedrock", model),
)
async def execute_async() -> ModelResponse | CustomStreamWrapper:
python_client: Final = None if isinstance(client, HTTPHandler) else client
if stream is True:
return self.async_streaming(
return await self.async_streaming(
model=model,
messages=messages,
api_base=proxy_endpoint_url,
@ -486,7 +476,7 @@ class BedrockConverseLLM(BaseAWSLLM):
logger_fn=logger_fn,
headers=headers,
timeout=timeout,
client=client,
client=python_client,
json_mode=json_mode,
fake_stream=fake_stream,
credentials=credentials,
@ -494,7 +484,7 @@ class BedrockConverseLLM(BaseAWSLLM):
stream_chunk_size=stream_chunk_size,
)
### ASYNC COMPLETION
return self.async_completion(
return await self.async_completion(
model=model,
messages=messages,
api_base=proxy_endpoint_url,
@ -507,108 +497,112 @@ class BedrockConverseLLM(BaseAWSLLM):
logger_fn=logger_fn,
headers=headers,
timeout=timeout,
client=client,
client=python_client,
credentials=credentials,
api_key=api_key,
skip_pre_call_logging=serves_via_rust,
)
@native_first(
native=native_completion,
route="chat_completions",
errors=lambda: rust_chat_completions_bridge.error_handling("bedrock", model),
)
def execute_sync() -> ModelResponse | CustomStreamWrapper:
## 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,
)
## TRANSFORMATION ##
## 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.
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,
},
)
resolved_timeout: Final = httpx.Timeout(timeout) if isinstance(timeout, (float, int)) else timeout
python_client: Final = (
_get_httpx_client({"timeout": resolved_timeout} if resolved_timeout is not None else None)
if client is None or isinstance(client, AsyncHTTPHandler)
else client
)
_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)
if stream is not None and stream is True:
completion_stream, response_headers = make_sync_call(
client=python_client,
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,
)
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,
)
return streaming_response
## 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,
### COMPLETION
try:
response: Final = python_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="",
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,
optional_params=optional_params,
encoding=encoding,
)
sync_transformed_response.set_provider_response_headers(response.headers)
return sync_transformed_response
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 execute_async() if acompletion else execute_sync()

View file

@ -1,8 +1,8 @@
import asyncio
import json
import ssl
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence
from contextlib import asynccontextmanager
from collections.abc import AsyncGenerator, AsyncIterator, Coroutine, Iterator, Mapping, Sequence
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from functools import lru_cache
from types import MappingProxyType, ModuleType
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, TypeVar, Union, cast, get_type_hints
@ -92,6 +92,8 @@ from litellm.responses.streaming_iterator import (
ResponsesWebSocketStreaming,
SyncResponsesAPIStreamingIterator,
)
from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, anative_context, anative_first
from litellm.rust_bridge.runtime import DispatchResult, NativeSkipped, NativeSkipReason, adapt_result
from litellm.types.containers.main import (
ContainerFileListResponse,
ContainerListResponse,
@ -2225,116 +2227,111 @@ class BaseLLMHTTPHandler:
},
)
rust_messages_response: Final = await self._maybe_rust_anthropic_messages(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
model=model,
api_key=api_key,
api_base=api_base,
headers=headers,
request_body=request_body,
timeout=self._resolve_anthropic_messages_timeout(
async def native_messages() -> DispatchResult[AnthropicMessagesResponse | AsyncIterator[object]]:
result: Final = await self._attempt_rust_anthropic_messages(
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
)
if rust_messages_response is not None:
if stream:
return self._rust_anthropic_messages_fake_stream(rust_messages_response)
return await self._finalize_anthropic_messages_response(
initial_response=rust_messages_response,
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
api_key=api_key,
kwargs=kwargs,
)
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
async_httpx_client=async_httpx_client,
request_url=request_url,
headers=headers,
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
request_body=request_body,
stream=stream or False,
logging_obj=logging_obj,
provider_config=anthropic_messages_provider_config,
litellm_params=litellm_params,
api_key=api_key,
model=model,
timeout=self._resolve_anthropic_messages_timeout(
litellm_params=litellm_params,
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
)
# used for logging + cost tracking
logging_obj.model_call_details["httpx_response"] = response
initial_response: AsyncIterator | AnthropicMessagesResponse
if stream:
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
AnthropicMessagesStreamingResponse,
anthropic_messages_stream_hidden_params,
)
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
model=model,
httpx_response=response,
api_base=api_base,
headers=headers,
request_body=request_body,
litellm_logging_obj=logging_obj,
timeout=self._resolve_anthropic_messages_timeout(
litellm_params=litellm_params,
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
)
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
return adapt_result(result, self._rust_anthropic_messages_fake_stream) if stream else result
if not self._has_agentic_completion_hook(logging_obj):
# No callback overrides async_should_run_agentic_loop, so the
# agentic wrapper's only effect would be buffering every chunk
# and rebuilding the response from SSE at end-of-stream to call
# hooks that all return (False, {}). Stream through directly and
# skip that per-chunk + end-of-stream overhead.
return AnthropicMessagesStreamingResponse(
completion_stream=completion_stream,
hidden_params=stream_hidden_params,
@anative_first(native=native_messages, route="messages", errors=lambda: PYTHON_ON_ERROR)
async def execute_messages() -> AnthropicMessagesResponse | AsyncIterator[object]:
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
async_httpx_client=async_httpx_client,
request_url=request_url,
headers=headers,
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
request_body=request_body,
stream=stream or False,
logging_obj=logging_obj,
provider_config=anthropic_messages_provider_config,
litellm_params=litellm_params,
api_key=api_key,
model=model,
timeout=self._resolve_anthropic_messages_timeout(
litellm_params=litellm_params,
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
)
# used for logging + cost tracking
logging_obj.model_call_details["httpx_response"] = response
initial_response: AsyncIterator | AnthropicMessagesResponse
if stream:
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
AnthropicMessagesStreamingResponse,
anthropic_messages_stream_hidden_params,
)
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
AgenticAnthropicStreamingIterator,
)
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
model=model,
httpx_response=response,
request_body=request_body,
litellm_logging_obj=logging_obj,
)
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
logging_obj=logging_obj,
tools=anthropic_messages_optional_request_params.get("tools"),
)
initial_response = AgenticAnthropicStreamingIterator(
completion_stream=completion_stream,
http_handler=self,
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
hold_back=bool(held_back_tool_names),
server_fulfilled_tool_names=held_back_tool_names,
)
return AnthropicMessagesStreamingResponse(
completion_stream=initial_response,
hidden_params=stream_hidden_params,
)
else:
initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
if not self._has_agentic_completion_hook(logging_obj):
# No callback overrides async_should_run_agentic_loop, so the
# agentic wrapper's only effect would be buffering every chunk
# and rebuilding the response from SSE at end-of-stream to call
# hooks that all return (False, {}). Stream through directly and
# skip that per-chunk + end-of-stream overhead.
return AnthropicMessagesStreamingResponse(
completion_stream=completion_stream,
hidden_params=stream_hidden_params,
)
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
AgenticAnthropicStreamingIterator,
)
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
logging_obj=logging_obj,
tools=anthropic_messages_optional_request_params.get("tools"),
)
initial_response = AgenticAnthropicStreamingIterator(
completion_stream=completion_stream,
http_handler=self,
model=model,
messages=messages,
anthropic_messages_provider_config=anthropic_messages_provider_config,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
hold_back=bool(held_back_tool_names),
server_fulfilled_tool_names=held_back_tool_names,
)
return AnthropicMessagesStreamingResponse(
completion_stream=initial_response,
hidden_params=stream_hidden_params,
)
else:
initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
return initial_response
initial_response: Final = await execute_messages()
if stream:
return initial_response
return await self._finalize_anthropic_messages_response(
initial_response=initial_response,
model=model,
@ -2384,7 +2381,7 @@ class BaseLLMHTTPHandler:
)
@staticmethod
async def _maybe_rust_anthropic_messages(
async def _attempt_rust_anthropic_messages(
*,
custom_llm_provider: str,
litellm_params: GenericLiteLLMParams,
@ -2395,41 +2392,36 @@ class BaseLLMHTTPHandler:
headers: dict,
request_body: dict,
timeout: float | httpx.Timeout | None,
) -> AnthropicMessagesResponse | None:
) -> DispatchResult[AnthropicMessagesResponse]:
if custom_llm_provider not in ("azure_ai", "anthropic"):
return None
return NativeSkipped(NativeSkipReason.INELIGIBLE)
from litellm.rust_bridge.configuration import rust_enabled
if not rust_enabled():
return None
return NativeSkipped(NativeSkipReason.DISABLED)
if has_agentic_hook:
return None
return NativeSkipped(NativeSkipReason.INELIGIBLE)
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 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
verbose_logger.debug(
"Rust Anthropic messages bridge raised %s; falling back to Python path",
type(rust_error).__name__,
)
return None
if rust_response is None:
return None
result: 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,
)
response_obj: Final = cast(AnthropicMessagesResponse, dict(rust_response))
response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
return response_obj
def adapt(rust_response: dict[str, object]) -> AnthropicMessagesResponse:
return cast(
AnthropicMessagesResponse,
{**rust_response, "_hidden_params": {"additional_headers": {"x-litellm-rust": "true"}}},
)
return adapt_result(result, adapt)
@staticmethod
def _rust_anthropic_messages_fake_stream(
@ -6507,20 +6499,22 @@ class BaseLLMHTTPHandler:
},
)
from litellm.rust_bridge import responses_websocket as rust_responses_websocket
async def attempt_connection() -> DispatchResult[
AbstractAsyncContextManager[rust_responses_websocket.ConnectionAdapter]
]:
if not _rust_responses_websocket_enabled(custom_llm_provider):
return NativeSkipped(NativeSkipReason.INELIGIBLE)
return await rust_responses_websocket.managed_connect(
url=ws_url,
headers={str(key): str(value) for key, value in headers.items()},
timeout=timeout,
)
@anative_context(native=attempt_connection, route="responses_websocket", errors=lambda: PYTHON_ON_ERROR)
@asynccontextmanager
async def _backend_connection():
if _rust_responses_websocket_enabled(custom_llm_provider):
from litellm.rust_bridge import responses_websocket as rust_responses_websocket
rust_backend: Final = await rust_responses_websocket.connect(
url=ws_url,
headers={str(key): str(value) for key, value in headers.items()},
timeout=timeout,
)
if rust_backend is not None:
yield rust_backend
return
async def _backend_connection() -> AsyncGenerator[ClientConnection, None]:
async with websockets.connect(
ws_url,
additional_headers=headers,

View file

@ -8,7 +8,6 @@ import mimetypes
import os
import re
from collections.abc import Callable, Coroutine, Mapping
from dataclasses import dataclass
from io import IOBase
from typing import Any, Final, cast
@ -18,18 +17,16 @@ import litellm
from litellm._logging import verbose_logger
from litellm.constants import request_timeout
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.azure_ai.ocr.common_utils import (
is_azure_document_intelligence_model,
)
from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model
from litellm.llms.base_llm.ocr.transformation import (
OCR_REQUEST_FORMAT_PARAM,
BaseOCRConfig,
OCRResponse,
parse_ocr_request_format,
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.dispatch import PROPAGATE, anative_first, native_first
from litellm.rust_bridge.runtime import DispatchResult
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
@ -38,36 +35,6 @@ base_llm_http_handler = BaseLLMHTTPHandler()
#################################################
@dataclass
class _PreparedOCRRequest:
model: str
document: dict[str, Any]
api_key: str | None
api_base: str | None
custom_llm_provider: str
extra_headers: dict[str, object] | None
provider_config: BaseOCRConfig
optional_params: dict[str, object]
litellm_params: dict[str, object]
effective_timeout: float | httpx.Timeout
litellm_logging_obj: LiteLLMLoggingObj
@dataclass
class _PreparedRustOCRCall:
api_key: str | None
api_base: str | None
headers: dict[str, object]
optional_params: dict[str, object]
_RUST_OCR_PROVIDERS: Final = {
"mistral",
"azure_ai",
"vertex_ai",
}
def _prepare_ocr_request(
model: str,
document: Mapping[str, object],
@ -77,7 +44,7 @@ def _prepare_ocr_request(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
kwargs: dict[str, object],
) -> _PreparedOCRRequest:
) -> rust_ocr_bridge.PreparedOCRRequest:
litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
litellm_call_id: Final = cast(str | None, kwargs.get("litellm_call_id", None))
@ -174,7 +141,7 @@ def _prepare_ocr_request(
custom_llm_provider=custom_llm_provider,
)
return _PreparedOCRRequest(
return rust_ocr_bridge.PreparedOCRRequest(
model=model,
document=document,
api_key=api_key,
@ -189,146 +156,67 @@ def _prepare_ocr_request(
)
def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
return False
if not prepared_request.provider_config.supports_rust_bridge():
return False
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
def _rust_bridge_optional_params(
prepared_request: _PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
) -> dict[str, object]:
optional_params: Final = dict(prepared_request.optional_params)
if prepared_request.custom_llm_provider == "vertex_ai":
vertex_project: Final = (
prepared_request.litellm_params.get("vertex_project")
or prepared_request.litellm_params.get("vertex_ai_project")
or litellm.vertex_project
or resolve_secret("VERTEXAI_PROJECT")
)
vertex_location: Final = (
prepared_request.litellm_params.get("vertex_location")
or prepared_request.litellm_params.get("vertex_ai_location")
or litellm.vertex_location
or resolve_secret("VERTEXAI_LOCATION")
or resolve_secret("VERTEX_LOCATION")
)
if vertex_project is not None:
optional_params["vertex_project"] = vertex_project
if vertex_location is not None:
optional_params["vertex_location"] = vertex_location
return optional_params
def _rust_bridge_api_base(
prepared_request: _PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
) -> str | None:
if prepared_request.api_base is not None:
return prepared_request.api_base
if prepared_request.custom_llm_provider == "azure_ai":
if is_azure_document_intelligence_model(prepared_request.model):
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
return resolve_secret("AZURE_AI_API_BASE")
return None
def _prepare_rust_ocr_call(
prepared_request: _PreparedOCRRequest,
@anative_first(
native=rust_ocr_bridge.aattempt_ocr,
route="ocr",
errors=lambda prepared_request, resolve_api_key: PROPAGATE,
)
async def _execute_aocr(
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> _PreparedRustOCRCall:
provider_config: Final = prepared_request.provider_config
api_key_env_var: Final = provider_config.get_api_key_env_var()
resolved_api_key: Final = prepared_request.api_key or (
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
)
resolved_headers: Final = provider_config.validate_environment(
headers=prepared_request.extra_headers or {},
model=prepared_request.model,
api_key=resolved_api_key,
api_base=prepared_request.api_base,
litellm_params=prepared_request.litellm_params,
)
resolved_complete_url: Final = provider_config.get_complete_url(
api_base=prepared_request.api_base,
) -> OCRResponse:
pending: Final = base_llm_http_handler.ocr(
model=prepared_request.model,
document=prepared_request.document,
optional_params=prepared_request.optional_params,
timeout=prepared_request.effective_timeout,
logging_obj=prepared_request.litellm_logging_obj,
api_key=prepared_request.api_key,
api_base=prepared_request.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
aocr=True,
headers=prepared_request.extra_headers,
provider_config=prepared_request.provider_config,
litellm_params=prepared_request.litellm_params,
)
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
prepared_request.litellm_logging_obj.pre_call(
input="OCR document processing",
api_key=resolved_api_key,
additional_args={
"complete_input_dict": {
"model": prepared_request.model,
"document": prepared_request.document,
**rust_optional_params,
},
"api_base": resolved_complete_url,
"headers": resolved_headers,
},
)
return _PreparedRustOCRCall(
api_key=resolved_api_key,
api_base=rust_api_base,
headers=cast(dict[str, object], resolved_headers),
optional_params=rust_optional_params,
)
response: Final = await pending if asyncio.iscoroutine(pending) else pending
if response is None:
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
return response
def _run_rust_ocr(
prepared_request: _PreparedOCRRequest,
def _attempt_ocr(
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> OCRResponse | None:
if rust_ocr_bridge.load_rust_ocr() is None:
return None
prepared: Final = _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
)
rust_response: Final = rust_ocr_bridge.ocr(
is_async: bool,
) -> DispatchResult[OCRResponse]:
return rust_ocr_bridge.attempt_ocr(prepared_request=prepared_request, resolve_api_key=resolve_api_key)
@native_first(
native=_attempt_ocr,
route="ocr",
errors=lambda prepared_request, resolve_api_key, is_async: PROPAGATE,
)
def _execute_ocr(
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
is_async: bool,
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
return base_llm_http_handler.ocr(
model=prepared_request.model,
document=prepared_request.document,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
optional_params=prepared.optional_params,
optional_params=prepared_request.optional_params,
timeout=prepared_request.effective_timeout,
)
if rust_response is None:
return None
return OCRResponse.model_validate(rust_response)
async def _run_rust_aocr(
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> OCRResponse | None:
if rust_ocr_bridge.load_rust_aocr() is None:
return None
prepared: Final = _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
)
rust_response: Final = await rust_ocr_bridge.aocr(
model=prepared_request.model,
document=prepared_request.document,
api_key=prepared.api_key,
api_base=prepared.api_base,
logging_obj=prepared_request.litellm_logging_obj,
api_key=prepared_request.api_key,
api_base=prepared_request.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
optional_params=prepared.optional_params,
timeout=prepared_request.effective_timeout,
aocr=is_async,
headers=prepared_request.extra_headers,
provider_config=prepared_request.provider_config,
litellm_params=prepared_request.litellm_params,
)
if rust_response is None:
return None
return OCRResponse.model_validate(rust_response)
@client
@ -425,40 +313,9 @@ async def aocr(
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
if _rust_ocr_supported(prepared) and rust_enabled():
from litellm.secret_managers.main import get_secret_str
from litellm.secret_managers.main import get_secret_str
rust_response: Final = await _run_rust_aocr(
prepared_request=prepared,
resolve_api_key=get_secret_str,
)
if rust_response is None:
verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path")
else:
return rust_response
response = base_llm_http_handler.ocr(
model=prepared.model,
document=prepared.document,
optional_params=prepared.optional_params,
timeout=prepared.effective_timeout,
logging_obj=prepared.litellm_logging_obj,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared.custom_llm_provider,
aocr=True,
headers=prepared.extra_headers,
provider_config=prepared.provider_config,
litellm_params=prepared.litellm_params,
)
if asyncio.iscoroutine(response):
response = await response
if response is None:
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
return response
return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str)
except Exception as e:
raise litellm.exception_type(
model=model,
@ -697,34 +554,9 @@ def ocr(
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
if _rust_ocr_supported(prepared) and rust_enabled():
from litellm.secret_managers.main import get_secret_str
from litellm.secret_managers.main import get_secret_str
rust_response: Final = _run_rust_ocr(
prepared_request=prepared,
resolve_api_key=get_secret_str,
)
if rust_response is None:
verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path")
else:
return rust_response
response: Final = base_llm_http_handler.ocr(
model=prepared.model,
document=prepared.document,
optional_params=prepared.optional_params,
timeout=prepared.effective_timeout,
logging_obj=prepared.litellm_logging_obj,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared.custom_llm_provider,
aocr=_is_async,
headers=prepared.extra_headers,
provider_config=prepared.provider_config,
litellm_params=prepared.litellm_params,
)
return response
return _execute_ocr(prepared_request=prepared, resolve_api_key=get_secret_str, is_async=_is_async)
except Exception as e:
raise litellm.exception_type(
model=model,

View file

@ -1,9 +1,11 @@
from __future__ import annotations
from collections.abc import Callable
from typing import Final, Generic, TypeVar
from types import ModuleType
from typing import Final, Generic, TypeVar, cast # noqa: TID251 # PyO3 module boundary
from litellm.rust_bridge.loader import get_native_bridge
from litellm.rust_bridge.protocols import NativeModule
BindingT = TypeVar("BindingT")
@ -15,21 +17,38 @@ class _Unset:
_UNSET: Final = _Unset()
class Unchanged:
pass
UNCHANGED: Final = Unchanged()
class NativeBinding(Generic[BindingT]):
"""Resolve one native attribute with an explicit, resettable test override."""
def __init__(self, attribute: str, *, validate: Callable[[object], BindingT | None]) -> None:
self._attribute: Final = attribute
self._validate: Final = validate
def __init__(
self,
select: Callable[[NativeModule], BindingT],
*,
module_loader: Callable[[], ModuleType | None] | None = None,
) -> None:
self._select: Final = select
self._module_loader: Final = module_loader
self._override: BindingT | None | _Unset = _UNSET
def load(self) -> BindingT | None:
if not isinstance(self._override, _Unset):
return self._override
native: Final = get_native_bridge()
native: Final = self._module_loader() if self._module_loader is not None else get_native_bridge()
if native is None:
return None
return self._validate(getattr(native, self._attribute, None))
module: Final = cast(NativeModule, native) # cast-ok: PyO3 exports are validated individually below
try:
value: Final = self._select(module)
except AttributeError:
return None
return value if callable(value) else None
def override(self, value: BindingT | None) -> None:
self._override = value
@ -38,12 +57,21 @@ class NativeBinding(Generic[BindingT]):
self._override = _UNSET
def native_exception_types() -> 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
_DECLINED: Final = NativeBinding(lambda native: native.RustBridgeDeclined)
_UPSTREAM: Final = NativeBinding(lambda native: native.RustUpstreamError)
def _exception_class(value: object) -> type[BaseException] | None:
if isinstance(value, type) and issubclass(value, BaseException):
return value
return None
def native_upstream_types() -> tuple[type[BaseException], ...]:
upstream: Final = _exception_class(_UPSTREAM.load())
return () if upstream is None else (upstream,)
def native_declined_types() -> tuple[type[BaseException], ...]:
declined: Final = _exception_class(_DECLINED.load())
return () if declined is None else (declined,)

View file

@ -4,30 +4,31 @@ The Rust core owns the conversation translation, the provider call, and the
response normalization for the subset of `/chat/completions` requests it
accepts. This module only marshals inputs and hands the normalized result to
LiteLLM's existing `ModelResponse` builder.
``None`` means the provider was never called, so the caller is free to serve the
request on the Python path. A failure after the call was issued raises instead:
retrying it there would bill the customer for the same work twice.
"""
from __future__ import annotations
import json
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Final, Protocol
import httpx
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.exceptions import APIError
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
convert_to_model_response_object,
)
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.loader import get_native_bridge
from litellm.rust_bridge.dispatch import APIErrorMapping, ErrorAction, ErrorHandling
from litellm.rust_bridge.protocols import (
RustAchatCompletions,
RustChatCompletions,
RustChatCompletionsDecline,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.types.utils import ModelResponse
@ -45,47 +46,6 @@ _LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object])
RUST_RESPONSE_HEADER: Final = "x-litellm-rust"
class RustChatCompletions(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
) -> Mapping[str, object]:
raise NotImplementedError
class RustAchatCompletions(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
) -> Awaitable[Mapping[str, object]]:
raise NotImplementedError
class RustChatCompletionsDecline(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
custom_llm_provider: str | None,
) -> str | None:
raise NotImplementedError
class ResponseObserver(Protocol):
"""Invoked with the payload the core returned, on success only.
@ -126,67 +86,36 @@ 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: Final[NativeBinding[RustChatCompletions]] = NativeBinding(lambda native: native.chat_completions)
_ACHAT: Final[NativeBinding[RustAchatCompletions]] = NativeBinding(lambda native: native.achat_completions)
_CHAT_PREFLIGHT: Final[NativeBinding[RustChatCompletionsDecline]] = NativeBinding(
lambda native: native.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 | Unchanged = UNCHANGED,
achat_completions: RustAchatCompletions | None | Unchanged = UNCHANGED,
decline: RustChatCompletionsDecline | None | Unchanged = UNCHANGED,
) -> 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
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
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
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
if not isinstance(chat_completions, Unchanged):
if chat_completions is None:
_CHAT.reset()
else:
_CHAT.override(chat_completions)
if not isinstance(achat_completions, Unchanged):
if achat_completions is None:
_ACHAT.reset()
else:
_ACHAT.override(achat_completions)
if not isinstance(decline, Unchanged):
if decline is None:
_CHAT_PREFLIGHT.reset()
else:
_CHAT_PREFLIGHT.override(decline)
def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool:
@ -247,12 +176,12 @@ def rust_chat_completions_accepts(
return False
if stream:
return False
if not rust_enabled():
return False
if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params):
verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path")
return False
decline: Final = _load_rust_decline()
if not rust_enabled():
return False
decline: Final = _CHAT_PREFLIGHT.load()
if decline is None:
return False
try:
@ -262,67 +191,12 @@ def rust_chat_completions_accepts(
optional_params=optional_params,
custom_llm_provider=custom_llm_provider,
)
except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
verbose_logger.debug(
"Rust chat completions gate raised %s; staying on the Python path",
type(rust_error).__name__,
)
except Exception as error: # noqa: BLE001 # capability checks perform no provider I/O
verbose_logger.debug("Native chat acceptance check failed: %s", error)
return False
if reason is not None:
verbose_logger.debug("Rust chat completions declined (%s); using the Python path", reason)
return False
return True
def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None:
"""`(declined, upstream_failed)` from the native module, or None when absent."""
native_bridge: Final = get_native_bridge()
if native_bridge is None:
return None
declined: Final = getattr(native_bridge, "RustBridgeDeclined", None)
upstream: Final = getattr(native_bridge, "RustUpstreamError", None)
if declined is None or upstream is None:
return None
return declined, upstream
def _reraise_or_decline(
rust_error: BaseException,
*,
model: str,
custom_llm_provider: str | None,
) -> None:
"""Re-raise a failure the provider already saw, or return so the caller declines.
A request that never reached the provider is safe to serve on the Python
path. One that did is not: the provider has already done the work, so a
second attempt bills for it twice. Those surface as an `APIError` carrying
the upstream status, which LiteLLM's exception mapping already understands.
"""
exceptions: Final = _rust_bridge_exceptions()
if exceptions is None:
verbose_logger.debug(
"Rust chat completions bridge raised %s; falling back to Python path",
type(rust_error).__name__,
)
return
declined, upstream_failed = exceptions
if isinstance(rust_error, upstream_failed):
args: Final = rust_error.args
status: Final = args[0] if args else 0
message: Final = args[1] if len(args) > 1 else ""
raise APIError(
status_code=int(status) or 500,
message=f"litellm rust chat completions: {message}",
llm_provider=custom_llm_provider or "",
model=model,
)
if not isinstance(rust_error, declined):
raise rust_error
verbose_logger.debug(
"Rust chat completions declined before calling the provider (%s); using the Python path",
rust_error,
)
verbose_logger.debug("Native chat request is ineligible: %s", reason)
return reason is None
def _build_model_response(
@ -351,12 +225,14 @@ def chat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
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(
eligible: bool = True,
) -> DispatchResult[ModelResponse]:
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
on_response(rust_response)
return _build_model_response(rust_response, model_response)
def call(native: RustChatCompletions, timeout_seconds: float | None) -> Mapping[str, object]:
return native(
model=model,
messages=messages,
optional_params=optional_params,
@ -364,13 +240,17 @@ def chat_completions(
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
timeout_seconds=timeout_seconds,
)
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)
return attempt(
load=_CHAT.load,
enabled=rust_enabled(),
eligible=eligible,
prepare=lambda: timeout_to_seconds(timeout),
call=call,
adapt=adapt,
)
async def achat_completions(
@ -385,12 +265,14 @@ async def achat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
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(
eligible: bool = True,
) -> DispatchResult[ModelResponse]:
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
on_response(rust_response)
return _build_model_response(rust_response, model_response)
async def call(native: RustAchatCompletions, timeout_seconds: float | None) -> Mapping[str, object]:
return await native(
model=model,
messages=messages,
optional_params=optional_params,
@ -398,49 +280,22 @@ async def achat_completions(
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
timeout_seconds=timeout_seconds,
)
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)
async def achat_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[[], 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,
return await aattempt(
load=_ACHAT.load,
enabled=rust_enabled(),
eligible=eligible,
prepare=lambda: timeout_to_seconds(timeout),
call=call,
adapt=adapt,
)
def error_handling(provider: str, model: str) -> ErrorHandling:
return ErrorHandling(
declined=ErrorAction.SKIP,
upstream=APIErrorMapping(provider=provider, model=model),
missing_metadata=ErrorAction.SKIP,
)
if response is not None:
return response
return await python_fallback()

View file

@ -0,0 +1,186 @@
from __future__ import annotations
from collections.abc import AsyncGenerator, Awaitable, Callable
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from dataclasses import dataclass
from enum import Enum
from functools import wraps
from typing import Final, ParamSpec, TypeAlias, TypeVar
from litellm._logging import verbose_logger
from litellm.exceptions import APIError
from litellm.rust_bridge.bindings import native_declined_types, native_upstream_types
from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason
NativeT = TypeVar("NativeT")
PythonT = TypeVar("PythonT")
P = ParamSpec("P")
class ErrorAction(Enum):
RAISE = "raise"
SKIP = "skip"
@dataclass(frozen=True, slots=True)
class APIErrorMapping:
provider: str
model: str
FailureAction: TypeAlias = ErrorAction | APIErrorMapping
@dataclass(frozen=True, slots=True)
class ErrorHandling:
declined: FailureAction = ErrorAction.RAISE
upstream: FailureAction = ErrorAction.RAISE
unknown: FailureAction = ErrorAction.RAISE
missing_metadata: FailureAction = ErrorAction.RAISE
unexpected: FailureAction = ErrorAction.RAISE
PROPAGATE: Final = ErrorHandling()
PYTHON_ON_ERROR: Final = ErrorHandling(
declined=ErrorAction.SKIP,
upstream=ErrorAction.SKIP,
unknown=ErrorAction.SKIP,
missing_metadata=ErrorAction.SKIP,
unexpected=ErrorAction.SKIP,
)
def _handle_error(error: Exception, action: FailureAction, route: str, reason: NativeSkipReason) -> NativeSkipped:
match action:
case ErrorAction.SKIP:
return NativeSkipped(reason, str(error))
case ErrorAction.RAISE:
raise error
case APIErrorMapping(provider, model):
args: Final[tuple[object, ...]] = error.args
attribute_status: Final = getattr(error, "status_code", None)
attribute_message: Final = getattr(error, "message", None)
status_value: Final = attribute_status if isinstance(attribute_status, int) else (args[0] if args else 0)
message_value: Final = (
attribute_message if isinstance(attribute_message, str) else (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 {route}: {message}",
llm_provider=provider,
model=model,
) from error
def _resolve(result: DispatchResult[NativeT], errors: ErrorHandling, route: str) -> Handled[NativeT] | NativeSkipped:
if not isinstance(result, NativeFailed):
return result
declined: Final = native_declined_types()
upstream: Final = native_upstream_types()
if not declined or not upstream:
return _handle_error(result.error, errors.missing_metadata, route, NativeSkipReason.FAILED)
if isinstance(result.error, declined):
return _handle_error(result.error, errors.declined, route, NativeSkipReason.DECLINED)
if isinstance(result.error, upstream):
return _handle_error(result.error, errors.upstream, route, NativeSkipReason.FAILED)
return _handle_error(result.error, errors.unknown, route, NativeSkipReason.FAILED)
def _log_skip(route: str, skipped: NativeSkipped) -> None:
verbose_logger.debug("Native %s skipped (%s): %s", route, skipped.reason.value, skipped.detail or "")
def native_first(
*,
native: Callable[P, DispatchResult[NativeT]],
route: str,
errors: Callable[P, ErrorHandling],
) -> Callable[[Callable[P, PythonT]], Callable[P, NativeT | PythonT]]:
def wrap(implementation: Callable[P, PythonT]) -> Callable[P, NativeT | PythonT]:
@wraps(implementation)
def run(
*args: P.args,
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
) -> NativeT | PythonT:
rules: Final = errors(*args, **kwargs)
try:
attempted: Final = native(*args, **kwargs)
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED)
_log_skip(route, skipped)
else:
result: Final = _resolve(attempted, rules, route)
if isinstance(result, Handled):
return result.value
_log_skip(route, result)
return implementation(*args, **kwargs)
return run
return wrap
def anative_first(
*,
native: Callable[P, Awaitable[DispatchResult[NativeT]]],
route: str,
errors: Callable[P, ErrorHandling],
) -> Callable[[Callable[P, Awaitable[PythonT]]], Callable[P, Awaitable[NativeT | PythonT]]]:
def wrap(implementation: Callable[P, Awaitable[PythonT]]) -> Callable[P, Awaitable[NativeT | PythonT]]:
@wraps(implementation)
async def run(
*args: P.args,
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
) -> NativeT | PythonT:
rules: Final = errors(*args, **kwargs)
try:
attempted: Final = await native(*args, **kwargs)
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
skipped: Final = _handle_error(error, rules.unexpected, route, NativeSkipReason.FAILED)
_log_skip(route, skipped)
else:
result: Final = _resolve(attempted, rules, route)
if isinstance(result, Handled):
return result.value
_log_skip(route, result)
return await implementation(*args, **kwargs)
return run
return wrap
def anative_context(
*,
native: Callable[P, Awaitable[DispatchResult[AbstractAsyncContextManager[NativeT]]]],
route: str,
errors: Callable[P, ErrorHandling],
) -> Callable[
[Callable[P, AbstractAsyncContextManager[PythonT]]],
Callable[P, AbstractAsyncContextManager[NativeT | PythonT]],
]:
def wrap(
implementation: Callable[P, AbstractAsyncContextManager[PythonT]],
) -> Callable[P, AbstractAsyncContextManager[NativeT | PythonT]]:
@anative_first(native=native, route=route, errors=errors)
async def acquire(
*args: P.args,
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
) -> AbstractAsyncContextManager[PythonT]:
return implementation(*args, **kwargs)
@wraps(implementation)
@asynccontextmanager
async def run(
*args: P.args,
**kwargs: P.kwargs, # kwargs-ok: ParamSpec preserves the wrapped signature
) -> AsyncGenerator[NativeT | PythonT, None]:
manager: Final = await acquire(*args, **kwargs)
async with manager as connection:
yield connection
return run
return wrap

View file

@ -2,90 +2,42 @@
from __future__ import annotations
from collections.abc import Awaitable
from dataclasses import dataclass
from typing import Final, Protocol, cast
from typing import Final
import httpx
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.protocols import RustAmessages, RustMessages
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
from litellm.rust_bridge.timeouts import timeout_to_seconds
class RustMessages(Protocol):
def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
) -> dict[str, object]:
raise NotImplementedError
class RustAmessages(Protocol):
def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]:
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]] = NativeBinding(lambda native: native.messages)
_AMESSAGES: Final[NativeBinding[RustAmessages]] = NativeBinding(lambda native: native.amessages)
def set_rust_messages(
*,
messages: RustMessages | None | _Unset = _UNSET,
amessages: RustAmessages | None | _Unset = _UNSET,
messages: RustMessages | None | Unchanged = UNCHANGED,
amessages: RustAmessages | None | Unchanged = UNCHANGED,
) -> None:
if not isinstance(messages, _Unset):
_STATE.messages = messages
if not isinstance(amessages, _Unset):
_STATE.amessages = amessages
if not isinstance(messages, Unchanged):
if messages is None:
_MESSAGES.reset()
else:
_MESSAGES.override(messages)
if not isinstance(amessages, Unchanged):
if amessages is None:
_AMESSAGES.reset()
else:
_AMESSAGES.override(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(
@ -97,18 +49,22 @@ def messages(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
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),
) -> DispatchResult[dict[str, object]]:
return attempt(
load=_MESSAGES.load,
enabled=True,
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_messages, timeout_seconds: 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_seconds,
),
adapt=identity,
)
@ -121,16 +77,20 @@ async def amessages(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
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),
) -> DispatchResult[dict[str, object]]:
return await aattempt(
load=_AMESSAGES.load,
enabled=True,
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_amessages, timeout_seconds: 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_seconds,
),
adapt=identity,
)

View file

@ -1,56 +1,57 @@
"""Thin Python wrapper for the native Rust OCR bridge."""
from __future__ import annotations
from collections.abc import Awaitable
from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables
from collections.abc import Callable
from dataclasses import dataclass
from typing import Final
import httpx
from pydantic import TypeAdapter
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model
from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse
from litellm.rust_bridge import configuration as _configuration
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
from litellm.rust_bridge.protocols import RustAocr, RustOcr
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
from litellm.rust_bridge.timeouts import timeout_to_seconds
_OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr)
_AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr)
_HEADERS: Final = TypeAdapter(dict[str, object])
class RustOcr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
raise NotImplementedError
@dataclass(frozen=True, slots=True)
class PreparedOCRRequest:
model: str
document: dict[str, object]
api_key: str | None
api_base: str | None
custom_llm_provider: str
extra_headers: dict[str, object] | None
provider_config: BaseOCRConfig
optional_params: dict[str, object]
litellm_params: dict[str, object]
effective_timeout: float | httpx.Timeout
litellm_logging_obj: LiteLLMLoggingObj
class RustAocr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError
@dataclass(frozen=True, slots=True)
class _PreparedRustOCRCall:
api_key: str | None
api_base: str | None
headers: dict[str, object]
optional_params: dict[str, object]
def _as_ocr(value: object) -> RustOcr | None:
return cast(RustOcr, value) if callable(value) else None
def _as_aocr(value: object) -> RustAocr | None:
return cast(RustAocr, value) if callable(value) else None
_OCR: Final = NativeBinding("ocr", validate=_as_ocr)
_AOCR: Final = NativeBinding("aocr", validate=_as_aocr)
_RUST_OCR_PROVIDERS: Final = frozenset(
{
"mistral",
"azure_ai",
"vertex_ai",
}
)
def load_rust_ocr() -> RustOcr | None:
@ -61,53 +62,150 @@ def load_rust_aocr() -> RustAocr | None:
return _AOCR.load()
def ocr(
*,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
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),
def _rust_ocr_supported(prepared_request: PreparedOCRRequest) -> bool:
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
return False
if not prepared_request.provider_config.supports_rust_bridge():
return False
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
def _rust_bridge_optional_params(
prepared_request: PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
) -> dict[str, object]:
if prepared_request.custom_llm_provider != "vertex_ai":
return prepared_request.optional_params
vertex_project: Final = (
prepared_request.litellm_params.get("vertex_project")
or prepared_request.litellm_params.get("vertex_ai_project")
or litellm.vertex_project
or resolve_secret("VERTEXAI_PROJECT")
)
vertex_location: Final = (
prepared_request.litellm_params.get("vertex_location")
or prepared_request.litellm_params.get("vertex_ai_location")
or litellm.vertex_location
or resolve_secret("VERTEXAI_LOCATION")
or resolve_secret("VERTEX_LOCATION")
)
return {
**prepared_request.optional_params,
**{
name: value
for name, value in (("vertex_project", vertex_project), ("vertex_location", vertex_location))
if value is not None
},
}
def _rust_bridge_api_base(
prepared_request: PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
) -> str | None:
if prepared_request.api_base is not None:
return prepared_request.api_base
if prepared_request.custom_llm_provider == "azure_ai":
if is_azure_document_intelligence_model(prepared_request.model):
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
return resolve_secret("AZURE_AI_API_BASE")
return None
def _prepare_rust_ocr_call(
prepared_request: PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> _PreparedRustOCRCall:
provider_config: Final = prepared_request.provider_config
api_key_env_var: Final = provider_config.get_api_key_env_var()
resolved_api_key: Final = prepared_request.api_key or (
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
)
resolved_headers: Final = _HEADERS.validate_python(
provider_config.validate_environment(
headers=prepared_request.extra_headers or {},
model=prepared_request.model,
api_key=resolved_api_key,
api_base=prepared_request.api_base,
litellm_params=prepared_request.litellm_params,
)
)
resolved_complete_url: Final = provider_config.get_complete_url(
api_base=prepared_request.api_base,
model=prepared_request.model,
optional_params=prepared_request.optional_params,
litellm_params=prepared_request.litellm_params,
)
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
prepared_request.litellm_logging_obj.pre_call(
input="OCR document processing",
api_key=resolved_api_key,
additional_args={
"complete_input_dict": {
"model": prepared_request.model,
"document": prepared_request.document,
**rust_optional_params,
},
"api_base": resolved_complete_url,
"headers": resolved_headers,
},
)
return _PreparedRustOCRCall(
api_key=resolved_api_key,
api_base=rust_api_base,
headers=resolved_headers,
optional_params=rust_optional_params,
)
async def aocr(
*,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
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),
def attempt_ocr(
prepared_request: PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> DispatchResult[OCRResponse]:
return attempt(
load=_OCR.load,
enabled=_configuration.rust_enabled(),
prepare=lambda: _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
),
call=lambda native, prepared: native(
model=prepared_request.model,
document=prepared_request.document,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
optional_params=prepared.optional_params,
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
),
adapt=OCRResponse.model_validate,
eligible=_rust_ocr_supported(prepared_request),
)
async def aattempt_ocr(
prepared_request: PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> DispatchResult[OCRResponse]:
return await aattempt(
load=_AOCR.load,
enabled=_configuration.rust_enabled(),
prepare=lambda: _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
),
call=lambda native, prepared: native(
model=prepared_request.model,
document=prepared_request.document,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
optional_params=prepared.optional_params,
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
),
adapt=OCRResponse.model_validate,
eligible=_rust_ocr_supported(prepared_request),
)

View file

@ -0,0 +1,180 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping, Sequence
from typing import Protocol
class RustChatCompletions(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
) -> Mapping[str, object]: ...
class RustAchatCompletions(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
) -> Awaitable[Mapping[str, object]]: ...
class RustChatCompletionsDecline(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
custom_llm_provider: str | None,
) -> str | None: ...
class RustResponsesWebSocket(Protocol):
async def send_text(self, text: str) -> None: ...
async def recv_text(self) -> str | None: ...
async def close(self) -> None: ...
class RustResponsesWebSocketConnection(Protocol):
@classmethod
async def connect(
cls,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
) -> RustResponsesWebSocket: ...
class NativeModule(Protocol):
@property
def chat_completions(self) -> RustChatCompletions: ...
@property
def achat_completions(self) -> RustAchatCompletions: ...
@property
def chat_completions_decline(self) -> RustChatCompletionsDecline: ...
@property
def ResponsesWebSocketConnection(self) -> type[RustResponsesWebSocketConnection]: ...
@property
def RustBridgeDeclined(self) -> type[BaseException]: ...
@property
def RustUpstreamError(self) -> type[BaseException]: ...
@property
def messages(self) -> RustMessages: ...
@property
def amessages(self) -> RustAmessages: ...
@property
def ocr(self) -> RustOcr: ...
@property
def aocr(self) -> RustAocr: ...
@property
def transcription(self) -> RustTranscription: ...
@property
def atranscription(self) -> RustAtranscription: ...
class RustMessages(Protocol):
def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
) -> dict[str, object]: ...
class RustAmessages(Protocol):
def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]: ...
class RustOcr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]: ...
class RustAocr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]: ...
class RustTranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]: ...
class RustAtranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]: ...

View file

@ -2,70 +2,39 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Final, Protocol
from collections.abc import AsyncGenerator
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from typing import Final
import httpx
from websockets.exceptions import ConnectionClosedOK
from litellm.rust_bridge.loader import get_native_bridge
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.protocols import (
RustResponsesWebSocket,
RustResponsesWebSocketConnection,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result
from litellm.rust_bridge.timeouts import timeout_to_seconds
class RustResponsesWebSocket(Protocol):
async def send_text(self, text: str) -> None: ...
async def recv_text(self) -> str | None: ...
async def close(self) -> None: ...
class RustResponsesWebSocketConnection(Protocol):
@classmethod
async def connect(
cls,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
) -> RustResponsesWebSocket: ...
class _Unset:
pass
_UNSET: Final[_Unset] = _Unset()
@dataclass(slots=True)
class _RustResponsesWebSocketState:
connection: RustResponsesWebSocketConnection | None = None
_STATE: Final[_RustResponsesWebSocketState] = _RustResponsesWebSocketState()
_RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding(
lambda native: native.ResponsesWebSocketConnection,
)
def set_rust_responses_websocket(
*,
connection: RustResponsesWebSocketConnection | None | _Unset = _UNSET,
connection: RustResponsesWebSocketConnection | None | Unchanged = UNCHANGED,
) -> None:
if not isinstance(connection, _Unset):
_STATE.connection = connection
if not isinstance(connection, Unchanged):
if connection is None:
_RESPONSES_WEBSOCKET.reset()
else:
_RESPONSES_WEBSOCKET.override(connection)
def load_rust_responses_websocket() -> RustResponsesWebSocketConnection | None:
if _STATE.connection is not None:
return _STATE.connection
native_bridge: Final = get_native_bridge()
if native_bridge is None:
return None
connection_type: Final[RustResponsesWebSocketConnection | None] = getattr(
native_bridge, "ResponsesWebSocketConnection", None
)
return connection_type
class _ConnectionAdapter:
class ConnectionAdapter:
def __init__(self, connection: RustResponsesWebSocket):
self._connection: Final[RustResponsesWebSocket] = connection
@ -87,16 +56,34 @@ async def connect(
url: str,
headers: dict[str, str],
timeout: float | httpx.Timeout | None,
) -> _ConnectionAdapter | None:
connection_type: Final = load_rust_responses_websocket()
if connection_type is None:
return None
try:
connection: Final = await connection_type.connect(
) -> DispatchResult[ConnectionAdapter]:
return await aattempt(
load=_RESPONSES_WEBSOCKET.load,
enabled=rust_enabled(),
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda connection_type, timeout_seconds: connection_type.connect(
url=url,
headers=headers,
timeout_seconds=timeout_to_seconds(timeout),
)
except Exception: # noqa: BLE001 # bridge failures must fall back to Python
return None
return _ConnectionAdapter(connection)
timeout_seconds=timeout_seconds,
),
adapt=ConnectionAdapter,
)
@asynccontextmanager
async def _connection_context(connection: ConnectionAdapter) -> AsyncGenerator[ConnectionAdapter, None]:
try:
yield connection
finally:
await connection.close()
async def managed_connect(
*,
url: str,
headers: dict[str, str],
timeout: float | httpx.Timeout | None,
) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]:
result: Final = await connect(url=url, headers=headers, timeout=timeout)
return adapt_result(result, _connection_context)

View file

@ -3,148 +3,93 @@ from __future__ import annotations
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from enum import Enum
from typing import Final, Generic, NoReturn, TypeAlias, TypeVar
from litellm.exceptions import APIError
from litellm.rust_bridge.bindings import native_exception_types
from typing import Final, Generic, TypeAlias, TypeVar
BindingT = TypeVar("BindingT")
NativeT = TypeVar("NativeT")
RequestT = TypeVar("RequestT")
ResultT = TypeVar("ResultT")
class FallbackMode(Enum):
PYTHON = "python"
RUST_REQUIRED = "rust_required"
class NativeSkipReason(Enum):
DISABLED = "disabled"
INELIGIBLE = "ineligible"
UNAVAILABLE = "unavailable"
DECLINED = "declined"
FAILED = "failed"
@dataclass(frozen=True, slots=True)
class RustHandled(Generic[ResultT]):
class Handled(Generic[ResultT]):
value: ResultT
@dataclass(frozen=True, slots=True)
class RustDeclined:
reason: str
class NativeSkipped:
reason: NativeSkipReason
detail: str | None = None
@dataclass(frozen=True, slots=True)
class RustUnavailable:
pass
class NativeFailed:
error: Exception
RustAttempt: TypeAlias = RustHandled[ResultT] | RustDeclined | RustUnavailable
DispatchResult: TypeAlias = Handled[ResultT] | NativeSkipped | NativeFailed
@dataclass(frozen=True, slots=True)
class BridgeErrorContext:
route: str
provider: str
model: str
def invoke(
*,
native_call: Callable[[], NativeT] | None,
fallback: Callable[[], ResultT],
adapt: Callable[[NativeT], ResultT],
mode: FallbackMode,
context: BridgeErrorContext,
) -> ResultT:
result: Final = attempt(native_call=native_call, adapt=adapt, context=context)
if isinstance(result, RustHandled):
return result.value
if mode is FallbackMode.PYTHON:
return fallback()
_raise_required(result, context)
async def ainvoke(
*,
native_call: Callable[[], Awaitable[NativeT]] | None,
fallback: Callable[[], Awaitable[ResultT]],
adapt: Callable[[NativeT], ResultT],
mode: FallbackMode,
context: BridgeErrorContext,
) -> ResultT:
result: Final = await aattempt(native_call=native_call, adapt=adapt, context=context)
if isinstance(result, RustHandled):
return result.value
if mode is FallbackMode.PYTHON:
return await fallback()
_raise_required(result, context)
def _select(load: Callable[[], BindingT | None], enabled: bool, eligible: bool) -> BindingT | NativeSkipped:
if not enabled:
return NativeSkipped(NativeSkipReason.DISABLED)
if not eligible:
return NativeSkipped(NativeSkipReason.INELIGIBLE)
binding: Final = load()
return NativeSkipped(NativeSkipReason.UNAVAILABLE) if binding is None else binding
def attempt(
*,
native_call: Callable[[], NativeT] | None,
load: Callable[[], BindingT | None],
enabled: bool,
eligible: bool,
prepare: Callable[[], RequestT],
call: Callable[[BindingT, RequestT], NativeT],
adapt: Callable[[NativeT], ResultT],
context: BridgeErrorContext,
) -> RustAttempt[ResultT]:
if native_call is None:
return RustUnavailable()
exceptions: Final = native_exception_types()
if exceptions is None:
return RustHandled(adapt(native_call()))
declined, upstream = exceptions
) -> DispatchResult[ResultT]:
binding: Final = _select(load, enabled, eligible)
if isinstance(binding, NativeSkipped):
return binding
try:
value: Final = native_call()
except declined as error:
return RustDeclined(reason=_decline_reason(error))
except upstream as error:
_raise_upstream(error, context)
return RustHandled(adapt(value))
value: Final = call(binding, prepare())
except Exception as error: # noqa: BLE001 # orchestration applies the endpoint's declared error policy
return NativeFailed(error)
return Handled(adapt(value))
async def aattempt(
*,
native_call: Callable[[], Awaitable[NativeT]] | None,
load: Callable[[], BindingT | None],
enabled: bool,
eligible: bool,
prepare: Callable[[], RequestT],
call: Callable[[BindingT, RequestT], Awaitable[NativeT]],
adapt: Callable[[NativeT], ResultT],
context: BridgeErrorContext,
) -> RustAttempt[ResultT]:
if native_call is None:
return RustUnavailable()
exceptions: Final = native_exception_types()
if exceptions is None:
return RustHandled(adapt(await native_call()))
declined, upstream = exceptions
) -> DispatchResult[ResultT]:
binding: Final = _select(load, enabled, eligible)
if isinstance(binding, NativeSkipped):
return binding
try:
value: 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(value))
value: Final = await call(binding, prepare())
except Exception as error: # noqa: BLE001 # orchestration applies the endpoint's declared error policy
return NativeFailed(error)
return Handled(adapt(value))
def _decline_reason(error: BaseException) -> str:
reason: Final[object] = error.args[0] if error.args else str(error)
return reason if isinstance(reason, str) else str(reason)
def identity(value: ResultT) -> ResultT:
return value
def _raise_required(
result: RustDeclined | RustUnavailable,
context: BridgeErrorContext,
) -> NoReturn:
raise RuntimeError(f"Rust {context.route} bridge {_required_reason(result)}")
def _required_reason(result: RustDeclined | RustUnavailable) -> str:
match 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[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 {context.route}: {message}",
llm_provider=context.provider,
model=context.model,
) from error
def adapt_result(result: DispatchResult[NativeT], adapt: Callable[[NativeT], ResultT]) -> DispatchResult[ResultT]:
if isinstance(result, Handled):
return Handled(adapt(result.value))
return result

View file

@ -1,99 +1,41 @@
from __future__ import annotations
from collections.abc import Awaitable
from dataclasses import dataclass
from typing import Final, Protocol, cast
from typing import Final
import httpx
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
from litellm.rust_bridge.timeouts import timeout_to_seconds
class RustTranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
raise NotImplementedError
class RustAtranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]:
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]] = NativeBinding(lambda native: native.transcription)
_ATRANSCRIPTION: Final[NativeBinding[RustAtranscription]] = NativeBinding(lambda native: native.atranscription)
def configure_rust_transcription(
*,
transcription: RustTranscription | None | _Unset = _UNSET,
atranscription: RustAtranscription | None | _Unset = _UNSET,
transcription: RustTranscription | None | Unchanged = UNCHANGED,
atranscription: RustAtranscription | None | Unchanged = UNCHANGED,
) -> None:
if not isinstance(transcription, _Unset):
_STATE.transcription = transcription
if not isinstance(atranscription, _Unset):
_STATE.atranscription = atranscription
if not isinstance(transcription, Unchanged):
if transcription is None:
_TRANSCRIPTION.reset()
else:
_TRANSCRIPTION.override(transcription)
if not isinstance(atranscription, Unchanged):
if atranscription is None:
_ATRANSCRIPTION.reset()
else:
_ATRANSCRIPTION.override(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(
@ -106,19 +48,23 @@ def transcription(
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
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),
) -> DispatchResult[dict[str, object]]:
return attempt(
load=_TRANSCRIPTION.load,
enabled=True,
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_transcription, timeout_seconds: 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_seconds,
),
adapt=identity,
)
@ -132,17 +78,21 @@ async def atranscription(
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
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),
) -> DispatchResult[dict[str, object]]:
return await aattempt(
load=_ATRANSCRIPTION.load,
enabled=True,
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_atranscription, timeout_seconds: 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_seconds,
),
adapt=identity,
)

View file

@ -12,7 +12,7 @@
"limit": 1979
},
"ANN202": {
"limit": 831
"limit": 830
},
"ANN204": {
"limit": 683
@ -150,13 +150,13 @@
"limit": 253
},
"PLW0127": {
"limit": 57
"limit": 55
},
"PLW0602": {
"limit": 215
},
"PLW0603": {
"limit": 190
"limit": 184
},
"PLW1508": {
"limit": 190
@ -195,7 +195,7 @@
"limit": 22
},
"SIM101": {
"limit": 56
"limit": 55
},
"SIM102": {
"limit": 310
@ -231,7 +231,7 @@
"limit": 5
},
"TID251": {
"limit": 1035
"limit": 1029
},
"TRY002": {
"limit": 524
@ -246,7 +246,7 @@
"limit": 109
},
"TRY300": {
"limit": 852
"limit": 846
},
"UP028": {
"limit": 2

View file

@ -9,6 +9,7 @@ import pytest
import litellm
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import configuration
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
@ -133,9 +134,9 @@ def test_load_rust_amessages_returns_injected_impl():
assert rust_messages.load_rust_amessages() is bridge
def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
def test_messages_wrapper_reports_unavailable(monkeypatch):
monkeypatch.setattr(
importlib.import_module("litellm.rust_bridge"),
importlib.import_module("litellm.rust_bridge.bindings"),
"get_native_bridge",
lambda: None,
)
@ -150,7 +151,7 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
extra_headers={},
timeout=30.0,
)
assert result is None
assert result == NativeSkipped(NativeSkipReason.UNAVAILABLE)
def test_messages_wrapper_forwards_args_and_converts_timeout():
@ -168,7 +169,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout():
timeout=httpx.Timeout(600.0, read=42.0),
)
assert response == FAKE_MESSAGES_RESPONSE
assert response == Handled(FAKE_MESSAGES_RESPONSE)
assert bridge.calls[0] == {
"model": "claude-sonnet-4-5",
"body": REQUEST_BODY,
@ -196,7 +197,7 @@ async def test_amessages_wrapper_forwards_args():
timeout=12.5,
)
assert response == FAKE_MESSAGES_RESPONSE
assert response == Handled(FAKE_MESSAGES_RESPONSE)
assert bridge.calls[0]["model"] == "claude-sonnet-4-5"
assert bridge.calls[0]["timeout_seconds"] == 12.5
@ -214,7 +215,7 @@ def _gate(**overrides):
"timeout": 30.0,
}
kwargs.update(overrides)
return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs)
return BaseLLMHTTPHandler._attempt_rust_anthropic_messages(**kwargs)
@pytest.mark.asyncio
@ -225,7 +226,8 @@ async def test_gate_invokes_rust_and_marks_response_header():
response = await _gate()
assert response is not None
assert isinstance(response, Handled)
response = response.value
assert response["id"] == "msg_123"
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
call = bridge.calls[0]
@ -238,14 +240,13 @@ async def test_gate_invokes_rust_and_marks_response_header():
@pytest.mark.asyncio
async def test_gate_falls_back_to_python_when_bridge_raises():
async def test_gate_reports_failure_to_harness():
bridge = RaisingAsyncMessages()
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate()
assert response is None
assert isinstance(response, NativeFailed)
assert bridge.calls == 1
@ -256,7 +257,7 @@ async def test_gate_skips_rust_when_flag_absent():
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
assert response is None
assert isinstance(response, NativeSkipped)
assert bridge.calls == 0
@ -268,7 +269,8 @@ async def test_gate_uses_process_enable_without_request_override():
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
assert response is not None
assert isinstance(response, Handled)
response = response.value
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
@ -286,7 +288,8 @@ async def test_gate_invokes_rust_for_native_anthropic_provider():
headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"},
)
assert response is not None
assert isinstance(response, Handled)
response = response.value
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
assert bridge.calls[0]["api_key"] == "sk-ant"
@ -303,7 +306,8 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch):
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
)
assert response is not None
assert isinstance(response, Handled)
response = response.value
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
@ -318,7 +322,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
)
assert response is None
assert isinstance(response, NativeSkipped)
assert bridge.calls == 0
@ -330,7 +334,7 @@ async def test_gate_skips_rust_for_unsupported_provider():
response = await _gate(custom_llm_provider="openai")
assert response is None
assert isinstance(response, NativeSkipped)
assert bridge.calls == 0
@ -342,7 +346,7 @@ async def test_gate_skips_rust_for_agentic_hook():
response = await _gate(has_agentic_hook=True)
assert response is None
assert isinstance(response, NativeSkipped)
assert bridge.calls == 0
@ -358,7 +362,8 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag():
request_body=streaming_body,
)
assert response is not None
assert isinstance(response, Handled)
response = response.value
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
assert "stream" not in bridge.calls[0]["body"]
assert bridge.calls[0]["body"] == REQUEST_BODY
@ -383,7 +388,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"),
importlib.import_module("litellm.rust_bridge.bindings"),
"get_native_bridge",
lambda: None,
)
@ -391,4 +396,56 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch):
response = await _gate()
assert response is None
assert isinstance(response, NativeSkipped)
@pytest.mark.asyncio
@pytest.mark.parametrize("selection", ("native", "disabled", "failed"))
async def test_messages_handler_runs_selected_backend_once(selection: str) -> None:
from datetime import datetime
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import AnthropicMessagesConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
bridge = RaisingAsyncMessages() if selection == "failed" else RecordingAsyncMessages()
rust_messages.set_rust_messages(amessages=bridge)
litellm.rust(selection != "disabled")
requests: list[httpx.Request] = []
def respond(request: httpx.Request) -> httpx.Response:
requests.append(request)
return httpx.Response(200, json=FAKE_MESSAGES_RESPONSE)
logging_obj = Logging(
model=FAKE_MESSAGES_RESPONSE["model"],
messages=[],
stream=False,
call_type="anthropic_messages",
start_time=datetime.now(),
litellm_call_id="harness-test",
function_id="harness-test",
)
client = AsyncHTTPHandler()
await client.client.aclose()
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as transport:
client.client = transport
response = await BaseLLMHTTPHandler().async_anthropic_messages_handler(
model=FAKE_MESSAGES_RESPONSE["model"],
messages=[{"role": "user", "content": "hello"}],
anthropic_messages_provider_config=AnthropicMessagesConfig(),
anthropic_messages_optional_request_params={"max_tokens": 10},
custom_llm_provider="anthropic",
litellm_params=GenericLiteLLMParams(),
logging_obj=logging_obj,
api_key="sk-test",
api_base="https://example.test",
client=client,
)
assert response["id"] == FAKE_MESSAGES_RESPONSE["id"]
assert len(requests) == (0 if selection == "native" else 1)
assert (bridge.calls if isinstance(bridge, RaisingAsyncMessages) else len(bridge.calls)) == (
0 if selection == "disabled" else 1
)

View file

@ -2467,7 +2467,7 @@ class TestRustChatCompletionsHook:
def declining_native(**_kwargs):
raise _Declined("blank message text")
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
bridge.set_rust_chat_completions(
decline=lambda **_kwargs: None, chat_completions=declining_native
)
@ -2499,7 +2499,7 @@ class TestRustChatCompletionsHook:
RustBridgeDeclined = _Declined
RustUpstreamError = type("_Upstream", (Exception,), {})
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
async def declining_native(**_kwargs):
raise _Declined("blank message text")
@ -2559,7 +2559,7 @@ class TestRustChatCompletionsHook:
RustBridgeDeclined = _Declined
RustUpstreamError = type("_Upstream", (Exception,), {})
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
def declining_native(**_kwargs):
raise _Declined("blank message text")

View file

@ -205,7 +205,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("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
async def declining_native(**_kwargs):
raise _Declined("blank message text")
@ -282,7 +282,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("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()),
patch.object(
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
),
@ -389,7 +389,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("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()):
bridge.set_rust_chat_completions(
decline=lambda **_kwargs: None, chat_completions=declining_native
)
@ -475,7 +475,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("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()):
bridge.set_rust_chat_completions(
decline=lambda **_kwargs: None, chat_completions=declining_native
)

View file

@ -11,7 +11,8 @@ supplied api_base is always honoured.
from litellm.llms.azure_ai.ocr.common_utils import (
is_azure_document_intelligence_model,
)
from litellm.ocr.main import _prepare_ocr_request, _rust_bridge_api_base
from litellm.ocr.main import _prepare_ocr_request
from litellm.rust_bridge.ocr import _rust_bridge_api_base
_DOC = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
_DOC_INTELLIGENCE_ENDPOINT = "https://di.cognitiveservices.azure.com"

View file

@ -12,13 +12,13 @@ import pytest
import litellm
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
from litellm.ocr.main import _PreparedOCRRequest, _rust_ocr_supported
from litellm.rust_bridge.ocr import PreparedOCRRequest, _rust_ocr_supported
DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
def _prepared(optional_params: dict[str, object]) -> _PreparedOCRRequest:
return _PreparedOCRRequest(
def _prepared(optional_params: dict[str, object]) -> PreparedOCRRequest:
return PreparedOCRRequest(
model="doc-intelligence/prebuilt-layout",
document=dict(DOCUMENT),
api_key="fake-key",

View file

@ -3,7 +3,8 @@
import builtins
import importlib
import types
from typing import Any
from typing import Any, Final
from unittest.mock import AsyncMock, Mock
import httpx
import pytest
@ -11,13 +12,14 @@ import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge import configuration
from litellm.rust_bridge.runtime import Handled
from litellm.rust_bridge.timeouts import timeout_to_seconds
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
# function onto `litellm.ocr` and shadows the submodule, so import the modules
# explicitly via importlib rather than attribute traversal.
ocr_main = importlib.import_module("litellm.ocr.main")
rust_bridge = importlib.import_module("litellm.rust_bridge.ocr")
rust_bridge_bindings = importlib.import_module("litellm.rust_bridge.bindings")
rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader")
MODEL = "mistral/mistral-ocr-latest"
@ -162,6 +164,9 @@ class FakeOCRConfig:
def get_api_key_env_var(self) -> str:
return self.api_key_env_var
def supports_rust_bridge(self) -> bool:
return True
def validate_environment(
self,
*,
@ -198,7 +203,7 @@ def build_prepared_request(
litellm_params: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = 12.5,
) -> Any:
return ocr_main._PreparedOCRRequest(
return rust_bridge.PreparedOCRRequest(
model=model,
document=document,
api_key=api_key,
@ -334,7 +339,7 @@ def test_toggle_without_ocr_arg_preserves_injected_impl():
def test_explicit_ocr_none_clears_injected_impl(monkeypatch):
monkeypatch.setattr(
rust_bridge_bindings,
importlib.import_module("litellm.rust_bridge.bindings"),
"get_native_bridge",
lambda: None,
)
@ -354,7 +359,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(
rust_bridge_bindings,
importlib.import_module("litellm.rust_bridge.bindings"),
"get_native_bridge",
lambda: None,
)
@ -371,7 +376,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(
rust_bridge_bindings,
importlib.import_module("litellm.rust_bridge.bindings"),
"get_native_bridge",
lambda: fake_module,
)
@ -382,74 +387,9 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch):
def test_timeout_to_seconds_handles_float_timeout_and_none():
assert rust_bridge._timeout_to_seconds(12.5) == 12.5
assert rust_bridge._timeout_to_seconds(None) is None
assert rust_bridge._timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0
def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
bridge = RecordingBridge()
litellm.rust(True)
rust_bridge._OCR.override(bridge)
response = rust_bridge.ocr(
model="mistral-ocr-latest",
document=DOCUMENT,
api_key="sk-test",
api_base="https://proxy.internal",
custom_llm_provider="mistral",
extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"},
optional_params={"include_image_base64": True, "pages": [0]},
timeout=12.5,
)
assert response == FAKE_OCR_RESPONSE
call = bridge.calls[0]
assert call == {
"model": "mistral-ocr-latest",
"document": DOCUMENT,
"api_key": "sk-test",
"api_base": "https://proxy.internal",
"custom_llm_provider": "mistral",
"extra_headers": {
"Authorization": "Bearer sk-test",
"x-trace-id": "trace-1",
},
"optional_params": {"include_image_base64": True, "pages": [0]},
"timeout_seconds": 12.5,
}
@pytest.mark.asyncio
async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
bridge = RecordingAsyncBridge()
litellm.rust(True)
rust_bridge._AOCR.override(bridge)
response = await rust_bridge.aocr(
model="mistral-ocr-maas",
document=DOCUMENT,
api_key=None,
api_base=None,
custom_llm_provider="vertex_ai",
extra_headers=None,
optional_params={"vertex_project": "project-1"},
timeout=httpx.Timeout(30.0, read=42.0),
)
assert response == FAKE_OCR_RESPONSE
assert bridge.calls[0] == {
"model": "mistral-ocr-maas",
"document": DOCUMENT,
"api_key": None,
"api_base": None,
"custom_llm_provider": "vertex_ai",
"extra_headers": None,
"optional_params": {"vertex_project": "project-1"},
"timeout_seconds": 42.0,
}
assert timeout_to_seconds(12.5) == 12.5
assert timeout_to_seconds(None) is None
assert timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0
def test_run_rust_ocr_prepares_request_and_wraps_response():
@ -458,7 +398,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
litellm.rust(True)
rust_bridge._OCR.override(bridge)
response = ocr_main._run_rust_ocr(
response = rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
logging_obj=logging_obj,
api_base="https://proxy.internal",
@ -469,6 +409,8 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
resolve_api_key=lambda _name: None,
)
assert isinstance(response, Handled)
response = response.value
assert isinstance(response, OCRResponse)
assert response.pages[0].markdown == "hello world"
assert bridge.calls[0] == {
@ -491,7 +433,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
litellm.rust(True)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(api_key=None, timeout=None),
resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None,
)
@ -507,7 +449,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver():
def _resolver(name: str) -> str | None:
raise AssertionError(f"resolver should not be called for {name}")
ocr_main._run_rust_ocr(
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
api_key="sk-explicit",
timeout=None,
@ -528,7 +470,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
resolver_calls.append(name)
return "sk-provider-env"
ocr_main._run_rust_ocr(
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"),
model="provider-ocr-model",
@ -547,7 +489,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
litellm.rust(True)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
custom_llm_provider="vertex_ai",
model="mistral-ocr-maas",
@ -580,7 +522,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana
"VERTEXAI_LOCATION": "us-east5",
}.get(name)
ocr_main._run_rust_ocr(
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
custom_llm_provider="vertex_ai",
model="mistral-ocr-maas",
@ -598,7 +540,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
litellm.rust(True)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
custom_llm_provider="azure_ai",
model="pixtral-12b-2409",
@ -616,7 +558,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
litellm.rust(True)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
custom_llm_provider="azure_ai",
model="doc-intelligence/prebuilt-layout",
@ -637,7 +579,7 @@ def test_run_rust_ocr_runs_pre_call_logging():
litellm.rust(True)
rust_bridge._OCR.override(bridge)
ocr_main._run_rust_ocr(
rust_bridge.attempt_ocr(
prepared_request=build_prepared_request(
logging_obj=logging_obj,
api_base="https://api.mistral.ai/v1",
@ -661,30 +603,6 @@ def test_run_rust_ocr_runs_pre_call_logging():
}
def test_ocr_routes_to_rust_when_enabled(fake_bridge):
response = litellm.ocr(
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
extra_headers={"x-trace-id": "trace-1"},
include_image_base64=True,
)
assert isinstance(response, OCRResponse)
assert response.pages[0].markdown == "hello world"
assert len(fake_bridge.calls) == 1
call = fake_bridge.calls[0]
assert call["model"] == "mistral-ocr-latest"
assert call["document"] == DOCUMENT
assert call["api_key"] == "sk-test"
assert call["custom_llm_provider"] == "mistral"
assert call["extra_headers"] == {
"Authorization": "Bearer sk-test",
"x-trace-id": "trace-1",
}
assert call["optional_params"].get("include_image_base64") is True
def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge):
response = litellm.ocr(
model="azure_ai/pixtral-12b-2409",
@ -794,34 +712,50 @@ def test_ocr_passes_default_request_timeout_to_rust(fake_bridge):
assert fake_bridge.calls[0]["timeout_seconds"] == float(request_timeout)
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.rust(False)
rust_bridge._OCR.override(bridge)
# The impl stays available for injection, but the disabled flag gates usage,
# so ocr() never reaches the Rust path (asserted via the enabled-path test).
assert bridge.calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize("enabled", (False, True))
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_ocr_fallback_skips_native_preparation(
monkeypatch: pytest.MonkeyPatch, enabled: bool, asynchronous: bool
) -> None:
monkeypatch.setattr(importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", lambda: None)
litellm.rust(enabled)
expected: Final = OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr")
fallback: Final = AsyncMock(return_value=expected) if asynchronous else Mock(return_value=expected)
def unexpected_preparation(*_args: object, **_kwargs: object) -> None:
pytest.fail("Python fallback must not resolve native credentials or emit native pre_call")
monkeypatch.setattr(rust_bridge, "_prepare_rust_ocr_call", unexpected_preparation)
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fallback)
response: Final = (
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
if asynchronous
else litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
)
assert response is expected
fallback.assert_called_once()
def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch):
"""Rust enabled but no bridge available (no injected impl, no compiled wheel):
ocr() must degrade to the Python HTTP handler instead of raising."""
monkeypatch.setattr(rust_bridge, "load_rust_ocr", lambda: None)
litellm.rust(True) # enabled, but load_rust_ocr() returns None in CI
@pytest.mark.asyncio
async def test_aocr_rejects_empty_python_fallback_response(monkeypatch: pytest.MonkeyPatch) -> None:
captured: dict[str, object] = {}
captured = {}
def fake_exception_type(**kwargs: object) -> CapturedException:
captured.update(kwargs)
return CapturedException("wrapped")
def fake_handler_ocr(**kwargs):
captured["called"] = True
return OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr")
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", AsyncMock(return_value=None))
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fake_handler_ocr)
with pytest.raises(CapturedException, match="wrapped"):
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
assert captured.get("called") is True # Python path was used
assert isinstance(response, OCRResponse)
original: Final = captured["original_exception"]
assert isinstance(original, ValueError)
assert str(original) == "Got an unexpected None response from the OCR API: None"
def test_ocr_provider_configs_expose_api_key_env_vars():

View file

@ -4,6 +4,7 @@ import pytest
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
from litellm.rust_bridge import configuration, responses_websocket
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
class _FakeNativeConnection:
@ -57,31 +58,29 @@ def test_rust_websocket_bridge_uses_process_enablement() -> None:
@pytest.mark.asyncio
async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection())
adapter = responses_websocket.ConnectionAdapter(_ClosedNativeConnection())
with pytest.raises(responses_websocket.ConnectionClosedOK):
await adapter.recv()
@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)
async def test_bridge_reports_unavailable(monkeypatch: pytest.MonkeyPatch) -> None:
configuration.rust(True)
responses_websocket._RESPONSES_WEBSOCKET.override(None)
assert (
await responses_websocket.connect(
url="wss://example.test/responses",
headers={},
timeout=None,
)
is None
)
assert await responses_websocket.connect(
url="wss://example.test/responses",
headers={},
timeout=None,
) == NativeSkipped(NativeSkipReason.UNAVAILABLE)
@pytest.mark.asyncio
async def test_enabled_bridge_connects_and_adapts_socket(
monkeypatch: pytest.MonkeyPatch,
) -> None:
configuration.rust(True)
responses_websocket.set_rust_responses_websocket(connection=_FakeNativeBridge)
connection = await responses_websocket.connect(
@ -90,7 +89,56 @@ async def test_enabled_bridge_connects_and_adapts_socket(
timeout=1.0,
)
assert connection is not None
assert isinstance(connection, Handled)
connection = connection.value
await connection.send("response.create")
assert await connection.recv() == "response.completed"
await connection.close()
class _FailingNativeBridge:
@classmethod
async def connect(
cls,
*,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
) -> _FakeNativeConnection:
raise RuntimeError("connection failed")
@pytest.mark.asyncio
async def test_connection_failure_is_reported_to_orchestration() -> None:
configuration.rust(True)
responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge)
result = await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None)
assert isinstance(result, NativeFailed)
assert str(result.error) == "connection failed"
@pytest.mark.asyncio
async def test_managed_connection_closes_native_socket_on_consumer_failure() -> None:
configuration.rust(True)
socket = _FakeNativeConnection()
class Bridge:
@classmethod
async def connect(
cls, *, url: str, headers: dict[str, str], timeout_seconds: float | None
) -> _FakeNativeConnection:
return socket
responses_websocket.set_rust_responses_websocket(connection=Bridge)
result = await responses_websocket.managed_connect(url="wss://example.test/responses", headers={}, timeout=1.0)
assert isinstance(result, Handled)
async def use_connection() -> None:
async with result.value as connection:
await connection.send("hello")
raise ValueError("consumer failed")
with pytest.raises(ValueError, match="consumer failed"):
await use_connection()
assert socket.sent == ["hello"]
assert socket.closed

View file

@ -1,3 +1,7 @@
import json
import subprocess
import sys
from pathlib import Path
from types import SimpleNamespace
from typing import Final
@ -6,30 +10,104 @@ import pytest
from litellm.rust_bridge import bindings
def test_binding_distinguishes_disable_from_reset(monkeypatch) -> None:
native = SimpleNamespace(route=lambda: "native")
def test_binding_distinguishes_disable_from_reset(monkeypatch: pytest.MonkeyPatch) -> None:
native: Final = SimpleNamespace(chat_completions=lambda: "native")
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
binding: bindings.NativeBinding[object] = bindings.NativeBinding("route", validate=lambda value: value)
assert binding.load() is native.route
binding: Final = bindings.NativeBinding(lambda module: module.chat_completions)
assert binding.load() is native.chat_completions
binding.override(None)
assert binding.load() is None
replacement = object()
binding.override(replacement)
assert binding.load() is replacement
replacement: Final = SimpleNamespace(chat_completions=lambda: "replacement")
binding.override(replacement.chat_completions)
assert binding.load() is replacement.chat_completions
binding.reset()
assert binding.load() is native.route
assert binding.load() is native.chat_completions
@pytest.mark.parametrize(("value", "expected"), ((3, 3), ("invalid", None), (None, None)))
def test_binding_validates_native_attribute(
monkeypatch: pytest.MonkeyPatch, value: object, expected: int | None
) -> None:
native: Final = SimpleNamespace(route=value)
@pytest.mark.parametrize("native", (None, SimpleNamespace(), SimpleNamespace(chat_completions=3)))
def test_missing_or_invalid_export_is_unavailable(monkeypatch: pytest.MonkeyPatch, native: object) -> None:
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
binding: Final = bindings.NativeBinding("route", validate=lambda item: item if isinstance(item, int) else None)
binding: Final = bindings.NativeBinding(lambda module: module.chat_completions)
assert binding.load() == expected
assert binding.load() is None
def test_selection_is_lazy_and_preserves_other_exports(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(bindings, "get_native_bridge", lambda: pytest.fail("must not load during construction"))
binding: Final = bindings.NativeBinding(lambda module: module.chat_completions)
native: Final = SimpleNamespace(chat_completions=lambda: "native", achat_completions=None)
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
assert binding.load() is native.chat_completions
assert bindings.NativeBinding(lambda module: module.achat_completions).load() is None
@pytest.mark.parametrize("invalid", (None, str, lambda: None))
def test_native_exception_types_reject_non_exception_classes(monkeypatch: pytest.MonkeyPatch, invalid: object) -> None:
native: Final = SimpleNamespace(RustBridgeDeclined=invalid, RustUpstreamError=RuntimeError)
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
assert bindings.native_declined_types() == ()
assert bindings.native_upstream_types() == (RuntimeError,)
@pytest.mark.parametrize(
("expression", "expected_rule"),
(
("NativeBinding(lambda native: native.chat_completion)", "reportAttributeAccessIssue"),
(
"wrong: NativeBinding[RustAchatCompletions] = NativeBinding(lambda native: native.chat_completions)",
"reportAssignmentType",
),
("NativeBinding(lambda native: native.ocrr)", "reportAttributeAccessIssue"),
(
"wrong: NativeBinding[RustAmessages] = NativeBinding(lambda native: native.messages)",
"reportAssignmentType",
),
(
"wrong: NativeBinding[RustAocr] = NativeBinding(lambda native: native.ocr)",
"reportAssignmentType",
),
(
"wrong: NativeBinding[RustAtranscription] = NativeBinding(lambda native: native.transcription)",
"reportAssignmentType",
),
),
)
def test_selectors_are_checked_by_type_checker(tmp_path: Path, expression: str, expected_rule: str) -> None:
source: Final = tmp_path / "binding_contract.py"
source.write_text(
"from typing_extensions import assert_type\n"
"from litellm.rust_bridge.bindings import NativeBinding\n"
"from litellm.rust_bridge.protocols import RustChatCompletions, RustAchatCompletions, "
"RustMessages, RustAmessages, RustOcr, RustAocr, RustTranscription, RustAtranscription\n"
"binding = NativeBinding(lambda native: native.chat_completions)\n"
"assert_type(binding, NativeBinding[RustChatCompletions])\n"
"assert_type(NativeBinding(lambda native: native.messages), NativeBinding[RustMessages])\n"
"assert_type(NativeBinding(lambda native: native.ocr), NativeBinding[RustOcr])\n"
"assert_type(NativeBinding(lambda native: native.transcription), NativeBinding[RustTranscription])\n"
+ expression
+ "\n"
)
config: Final = tmp_path / "pyrightconfig.json"
config.write_text(
json.dumps(
{
"include": [str(source)],
"extraPaths": [str(Path(__file__).resolve().parents[3])],
"typeCheckingMode": "basic",
}
)
)
result: Final = subprocess.run(
[sys.executable, "-m", "basedpyright", "--project", str(config), "--outputjson"],
capture_output=True,
text=True,
check=False,
)
diagnostics: Final = json.loads(result.stdout)["generalDiagnostics"]
assert result.returncode == 1, result.stdout + result.stderr
assert [(item["rule"], item["range"]["start"]["line"]) for item in diagnostics] == [
(expected_rule, len(source.read_text().splitlines()) - 1)
]

View file

@ -10,8 +10,9 @@ from __future__ import annotations
import pytest
import litellm
from litellm.rust_bridge import configuration
from litellm.rust_bridge import bindings, configuration
from litellm.rust_bridge import chat_completions as bridge
from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeFailed
from litellm.types.utils import ModelResponse
RUST_RESPONSE = {
@ -54,7 +55,7 @@ class _FakeNative:
def _fake_native_bridge(monkeypatch):
"""Expose the bridge's exception classes without the compiled extension."""
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
monkeypatch.setattr(bindings, "get_native_bridge", lambda: _FakeNative())
def _hide_native_bridge(monkeypatch):
@ -63,7 +64,7 @@ def _hide_native_bridge(monkeypatch):
There is no injection seam for "the .so is absent", so the loader itself is
replaced; every other case here uses `set_rust_chat_completions`.
"""
monkeypatch.setattr(bridge, "get_native_bridge", lambda: None)
monkeypatch.setattr(bindings, "get_native_bridge", lambda: None)
@pytest.fixture(autouse=True)
@ -250,7 +251,8 @@ class TestSyncCall:
result = bridge.chat_completions(**_call_kwargs(model_response))
assert result is not None
assert isinstance(result, Handled)
result = result.value
assert result.choices[0].message.content == "hello from rust"
assert result.choices[0].finish_reason == "stop"
assert result.model == "claude-sonnet-4-5-20260101"
@ -266,14 +268,14 @@ class TestSyncCall:
bridge.chat_completions(**_call_kwargs(ModelResponse()))
assert native.calls[0]["timeout_seconds"] == 30.0
def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
def test_reports_unavailable_bridge(self, monkeypatch):
_hide_native_bridge(monkeypatch)
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeSkipped)
def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
def test_reports_native_decline_to_orchestration(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 isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeFailed)
class TestAsyncCall:
@ -281,115 +283,18 @@ class TestAsyncCall:
async def test_builds_a_model_response(self):
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
result = await bridge.achat_completions(**_call_kwargs(ModelResponse()))
assert result is not None
assert isinstance(result, Handled)
result = result.value
assert result.choices[0].message.content == "hello from rust"
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
@pytest.mark.asyncio
async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
async def test_reports_unavailable_bridge(self, monkeypatch):
_hide_native_bridge(monkeypatch)
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeSkipped)
@pytest.mark.asyncio
async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
async def test_reports_native_decline_to_orchestration(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
class TestAsyncFallbackWrapper:
@pytest.mark.asyncio
async def test_returns_the_rust_response_without_running_the_fallback(self):
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
ran = []
async def fallback():
ran.append(True)
return "python"
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result.choices[0].message.content == "hello from rust"
assert ran == []
@pytest.mark.asyncio
async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch):
_fake_native_bridge(monkeypatch)
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
async def fallback():
return "python"
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result == "python"
@pytest.mark.asyncio
async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch):
_hide_native_bridge(monkeypatch)
async def fallback():
return "python"
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result == "python"
class TestFailureClassification:
"""A failure the provider already saw must not be retried on the Python
path: it would bill the customer for the same work twice."""
@pytest.fixture(autouse=True)
def _native_exceptions(self, monkeypatch):
_fake_native_bridge(monkeypatch)
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
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()))
assert raised.value.status_code == 429
assert "rate limited" in str(raised.value)
def test_a_transport_failure_with_no_response_surfaces_as_a_500(self):
from litellm.exceptions import APIError
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()))
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()))
@pytest.mark.asyncio
async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self):
from litellm.exceptions import APIError
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom")))
ran = []
async def fallback():
ran.append(True)
return "python"
with pytest.raises(APIError):
await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert ran == [], "a request the provider already served must not be re-issued"
@pytest.mark.asyncio
async def test_the_async_wrapper_falls_back_on_a_decline(self):
bridge.set_rust_chat_completions(
achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text"))
)
async def fallback():
return "python"
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
assert result == "python"
assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeFailed)

View file

@ -0,0 +1,307 @@
from __future__ import annotations
import asyncio
import logging
from types import SimpleNamespace
from typing import Final
import pytest
from litellm.exceptions import APIError
from litellm.rust_bridge import bindings
from litellm.rust_bridge.chat_completions import error_handling
from litellm.rust_bridge.dispatch import PROPAGATE, PYTHON_ON_ERROR, anative_first, native_first
from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, NativeSkipped, NativeSkipReason
class Declined(Exception):
pass
class Upstream(Exception):
pass
@pytest.fixture(autouse=True)
def native_metadata(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(
bindings, "get_native_bridge", lambda: SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream)
)
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("reason", tuple(NativeSkipReason))
async def test_shared_dispatch_calls_python_once_and_logs_skip(
asynchronous: bool, reason: NativeSkipReason, caplog: pytest.LogCaptureFixture
) -> None:
caplog.set_level(logging.DEBUG, logger="LiteLLM")
calls: Final[list[str]] = []
def native() -> DispatchResult[str]:
calls.append("native")
return NativeSkipped(reason, "diagnostic detail")
async def anative() -> DispatchResult[str]:
return native()
def python() -> str:
calls.append("python")
return "python response"
async def apython() -> str:
return python()
result: Final = (
await anative_first(native=anative, route="test", errors=lambda: PROPAGATE)(apython)()
if asynchronous
else native_first(native=native, route="test", errors=lambda: PROPAGATE)(python)()
)
assert result == "python response"
assert calls == ["native", "python"]
assert f"Native test skipped ({reason.value}): diagnostic detail" in caplog.text
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_native_success_does_not_run_python_even_when_value_is_none(asynchronous: bool) -> None:
async def native() -> DispatchResult[None]:
return Handled(None)
def python() -> str:
pytest.fail("handled results must not run Python")
async def apython() -> str:
return python()
result: Final = (
await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)()
if asynchronous
else native_first(native=lambda: Handled(None), route="test", errors=lambda: PYTHON_ON_ERROR)(python)()
)
assert result is None
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("policy", ("chat", "propagate", "python"))
@pytest.mark.parametrize("kind", ("declined", "upstream", "unknown", "unexpected", "missing"))
async def test_declarations_preserve_endpoint_error_behavior(
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, policy: str, kind: str
) -> None:
if kind == "missing":
monkeypatch.setattr(bindings, "get_native_bridge", lambda: None)
error: Final = (
Declined("unsupported")
if kind == "declined"
else Upstream(429, "rate limited")
if kind == "upstream"
else RuntimeError("failed")
)
rules: Final = (
error_handling("anthropic", "model")
if policy == "chat"
else PYTHON_ON_ERROR
if policy == "python"
else PROPAGATE
)
calls: Final[list[str]] = []
def native() -> DispatchResult[str]:
if kind == "unexpected":
raise error
return NativeFailed(error)
async def anative() -> DispatchResult[str]:
return native()
def python() -> str:
calls.append("python")
return "python response"
async def apython() -> str:
return python()
async def run() -> str:
if asynchronous:
return await anative_first(native=anative, route="chat_completions", errors=lambda: rules)(apython)()
return native_first(native=native, route="chat_completions", errors=lambda: rules)(python)()
if policy == "python" or (policy == "chat" and kind in ("declined", "missing")):
assert await run() == "python response"
assert calls == ["python"]
elif policy == "chat" and kind == "upstream":
with pytest.raises(APIError) as caught:
await run()
assert caught.value.status_code == 429
assert caught.value.model == "model"
assert caught.value.llm_provider == "anthropic"
assert caught.value.__cause__ is error
assert calls == []
else:
with pytest.raises(type(error)) as caught_original:
await run()
assert caught_original.value is error
assert calls == []
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_python_failure_is_never_reclassified_as_native_failure(asynchronous: bool) -> None:
error: Final = RuntimeError("Python failed")
calls: Final[list[str]] = []
async def native() -> DispatchResult[str]:
return NativeSkipped(NativeSkipReason.UNAVAILABLE)
def python() -> str:
calls.append("python")
raise error
async def apython() -> str:
return python()
async def run() -> str:
if asynchronous:
return await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)()
return native_first(
native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE), route="test", errors=lambda: PYTHON_ON_ERROR
)(python)()
with pytest.raises(RuntimeError) as caught:
await run()
assert caught.value is error
assert calls == ["python"]
@pytest.mark.asyncio
async def test_cancellation_does_not_run_python() -> None:
async def native() -> DispatchResult[str]:
raise asyncio.CancelledError
async def python() -> str:
pytest.fail("cancellation must not dispatch Python")
with pytest.raises(asyncio.CancelledError):
await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(python)()
@pytest.mark.parametrize("status", (0, 401, 403, 429, 500, 503))
def test_chat_upstream_mapping_preserves_status_message_and_context(status: int) -> None:
error: Final = Upstream(status, "upstream failed")
with pytest.raises(APIError, match="upstream failed") as caught:
native_first(
native=lambda: NativeFailed(error),
route="chat_completions",
errors=lambda: error_handling("anthropic", "model"),
)(lambda: pytest.fail("upstream errors must not run Python"))()
assert caught.value.status_code == (status or 500)
assert caught.value.model == "model"
assert caught.value.llm_provider == "anthropic"
assert caught.value.__cause__ is error
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_registered_wrapper_preserves_arguments_and_request_error_context(asynchronous: bool) -> None:
calls: Final[list[tuple[str, str, str]]] = []
def native(provider: str, *, model: str) -> DispatchResult[str]:
calls.append(("native", provider, model))
return (
NativeFailed(Upstream(429, "limited"))
if model == "limited"
else NativeSkipped(NativeSkipReason.UNAVAILABLE)
)
async def anative(provider: str, *, model: str) -> DispatchResult[str]:
return native(provider, model=model)
def rules(provider: str, *, model: str):
return error_handling(provider, model)
@native_first(native=native, route="chat_completions", errors=rules)
def execute(provider: str, *, model: str) -> str:
calls.append(("python", provider, model))
return model
@anative_first(native=anative, route="chat_completions", errors=rules)
async def aexecute(provider: str, *, model: str) -> str:
calls.append(("python", provider, model))
return model
assert (await aexecute("first", model="ok") if asynchronous else execute("first", model="ok")) == "ok"
async def fail() -> None:
if asynchronous:
await aexecute("second", model="limited")
else:
execute("second", model="limited")
with pytest.raises(APIError) as caught:
await fail()
assert caught.value.llm_provider == "second"
assert caught.value.model == "limited"
assert calls == [("native", "first", "ok"), ("python", "first", "ok"), ("native", "second", "limited")]
@pytest.mark.asyncio
@pytest.mark.parametrize("selection", ("native", "unavailable", "failed"))
@pytest.mark.parametrize("failure", ("none", "body", "cleanup", "cancel"))
async def test_context_selection_and_lifetime_are_separate(selection: str, failure: str) -> None:
from collections.abc import AsyncGenerator
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from litellm.rust_bridge.dispatch import anative_context
events: Final[list[str]] = []
error: Final = RuntimeError("connection use failed")
@asynccontextmanager
async def connection(name: str) -> AsyncGenerator[str, None]:
events.append(f"{name}:enter")
try:
yield name
finally:
events.append(f"{name}:exit")
if failure == "cleanup":
raise error
async def native() -> DispatchResult[AbstractAsyncContextManager[str]]:
events.append("attempt")
if selection == "failed":
raise RuntimeError("connect failed")
if selection == "unavailable":
return NativeSkipped(NativeSkipReason.UNAVAILABLE)
return Handled(connection("native"))
@anative_context(native=native, route="websocket", errors=lambda: PYTHON_ON_ERROR)
def execute() -> AbstractAsyncContextManager[str]:
events.append("python")
return connection("python")
async def run() -> None:
async with execute() as name:
assert name == ("native" if selection == "native" else "python")
if failure == "body":
raise error
if failure == "cancel":
raise asyncio.CancelledError
if failure == "none":
await run()
elif failure == "cancel":
with pytest.raises(asyncio.CancelledError):
await run()
else:
with pytest.raises(RuntimeError) as caught:
await run()
assert caught.value is error
expected: Final = (
["attempt", "native:enter", "native:exit"]
if selection == "native"
else ["attempt", "python", "python:enter", "python:exit"]
)
assert events == expected

View file

@ -1,95 +1,117 @@
from __future__ import annotations
from types import SimpleNamespace
from typing import Final
import pytest
from litellm.exceptions import APIError
from litellm.rust_bridge import bindings, runtime
class RustBridgeDeclined(Exception):
pass
class RustUpstreamError(Exception):
pass
@pytest.fixture(autouse=True)
def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None:
native = SimpleNamespace(
RustBridgeDeclined=RustBridgeDeclined,
RustUpstreamError=RustUpstreamError,
)
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
def context() -> runtime.BridgeErrorContext:
return runtime.BridgeErrorContext(route="messages", provider="anthropic", model="model")
def test_invoke_tags_native_decline_before_running_fallback() -> None:
calls: list[str] = []
def decline() -> object:
calls.append("rust")
raise RustBridgeDeclined("unsupported")
value = runtime.invoke(
native_call=decline,
fallback=lambda: calls.append("python") or "fallback",
adapt=str,
mode=runtime.FallbackMode.PYTHON,
context=context(),
)
assert value == "fallback"
assert calls == ["rust", "python"]
def test_invoke_translates_upstream_without_fallback() -> None:
def fail() -> object:
raise RustUpstreamError(429, "rate limited")
with pytest.raises(APIError, match="rate limited") as caught:
runtime.invoke(
native_call=fail,
fallback=lambda: pytest.fail("fallback must not run"),
adapt=str,
mode=runtime.FallbackMode.PYTHON,
context=context(),
)
assert caught.value.status_code == 429
from litellm.rust_bridge import runtime
@pytest.mark.asyncio
async def test_ainvoke_handles_native_success() -> None:
async def native() -> int:
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("state", ("disabled", "ineligible", "unavailable", "handled"))
async def test_attempt_only_prepares_selected_requests(asynchronous: bool, state: str) -> None:
events: Final[list[str]] = []
def load() -> object | None:
events.append("load")
return None if state == "unavailable" else object()
def prepare() -> int:
events.append("prepare")
return 3
async def fallback() -> str:
pytest.fail("fallback must not run")
def call(_binding: object, request: int) -> int:
events.append("call")
return request * 2
assert (
await runtime.ainvoke(
native_call=native,
fallback=fallback,
adapt=str,
mode=runtime.FallbackMode.PYTHON,
context=context(),
async def acall(binding: object, request: int) -> int:
return call(binding, request)
def adapt(value: int) -> str:
events.append("adapt")
return str(value)
result: Final = (
await runtime.aattempt(
load=load,
enabled=state != "disabled",
eligible=state != "ineligible",
prepare=prepare,
call=acall,
adapt=adapt,
)
if asynchronous
else runtime.attempt(
load=load,
enabled=state != "disabled",
eligible=state != "ineligible",
prepare=prepare,
call=call,
adapt=adapt,
)
== "3"
)
if state == "handled":
assert result == runtime.Handled("6")
assert events == ["load", "prepare", "call", "adapt"]
else:
assert result == runtime.NativeSkipped(runtime.NativeSkipReason(state))
assert events == (["load"] if state == "unavailable" else [])
def test_required_mode_rejects_unavailable_bridge() -> None:
with pytest.raises(RuntimeError, match="is unavailable"):
runtime.invoke(
native_call=None,
fallback=lambda: pytest.fail("fallback must not run"),
adapt=str,
mode=runtime.FallbackMode.RUST_REQUIRED,
context=context(),
)
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
@pytest.mark.parametrize("phase", ("prepare", "call"))
async def test_attempt_reports_failure_without_deciding_retry(asynchronous: bool, phase: str) -> None:
error: Final = RuntimeError("native failure")
def prepare() -> int:
if phase == "prepare":
raise error
return 3
def call(_binding: object, request: int) -> int:
raise error
async def acall(binding: object, request: int) -> int:
return call(binding, request)
def adapt(value: int) -> str:
pytest.fail("failed attempts cannot be adapted")
result: Final = (
await runtime.aattempt(load=object, enabled=True, eligible=True, prepare=prepare, call=acall, adapt=adapt)
if asynchronous
else runtime.attempt(load=object, enabled=True, eligible=True, prepare=prepare, call=call, adapt=adapt)
)
assert isinstance(result, runtime.NativeFailed)
assert result.error is error
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", (False, True))
async def test_adaptation_failure_remains_distinct_from_native_failure(asynchronous: bool) -> None:
error: Final = ValueError("invalid response")
async def acall(_binding: object, request: int) -> int:
return request
def adapt(value: int) -> str:
raise error
async def run() -> None:
if asynchronous:
await runtime.aattempt(load=object, enabled=True, eligible=True, prepare=lambda: 3, call=acall, adapt=adapt)
else:
runtime.attempt(
load=object,
enabled=True,
eligible=True,
prepare=lambda: 3,
call=lambda binding, request: request,
adapt=adapt,
)
with pytest.raises(ValueError, match="invalid response") as caught:
await run()
assert caught.value is error

View file

@ -4,6 +4,7 @@ import pytest
import litellm
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
from litellm.rust_bridge.runtime import Handled
rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
@ -55,7 +56,8 @@ def test_enabled_sync_bridge_receives_audio() -> None:
optional_params={"temperature": 0},
timeout=5.0,
)
assert result == {"text": "hello"}
assert isinstance(result, Handled)
assert result.value == {"text": "hello"}
assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"}
@ -72,18 +74,19 @@ async def test_enabled_async_bridge() -> None:
optional_params={},
timeout=None,
)
assert result == {"text": "async"}
assert result == Handled({"text": "async"})
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("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
assert rust_bridge.load_rust_transcription() is None
assert rust_bridge.load_rust_atranscription() is None
def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None)
rust_bridge.configure_rust_transcription(transcription=None)
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
with pytest.raises(RuntimeError, match="bridge is unavailable"):
BedrockAudioTranscriptionRustDispatch().audio_transcriptions(
@ -100,10 +103,8 @@ def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) ->
@pytest.mark.asyncio
async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
async def unavailable(**_: object) -> None:
return None
monkeypatch.setattr(rust_bridge, "atranscription", unavailable)
rust_bridge.configure_rust_transcription(atranscription=None)
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
with pytest.raises(RuntimeError, match="bridge is unavailable"):
await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions(

View file

@ -1,6 +1,6 @@
{
"LIT001": {
"limit": 22180
"limit": 22165
},
"LIT002": {
"limit": 26729
@ -15,7 +15,7 @@
"limit": 0
},
"LIT006": {
"limit": 1035
"limit": 1022
},
"LIT007": {
"limit": 0
@ -27,10 +27,10 @@
"limit": 0
},
"LIT010": {
"limit": 16426
"limit": 16419
},
"LIT011": {
"limit": 5506
"limit": 5497
},
"LIT012": {
"limit": 4486