refactor(native): wrap endpoint execution in a shared harness

This commit is contained in:
Yujong Lee 2026-09-05 20:46:26 -07:00
parent 3fe0d809d1
commit 57c21b6898
16 changed files with 899 additions and 619 deletions

View file

@ -3,7 +3,7 @@
"limit": 13426
},
"reportArgumentType": {
"limit": 2194
"limit": 2192
},
"reportAssignmentType": {
"limit": 319
@ -30,7 +30,7 @@
"limit": 7
},
"reportGeneralTypeIssues": {
"limit": 101
"limit": 100
},
"reportIncompatibleMethodOverride": {
"limit": 56
@ -57,7 +57,7 @@
"limit": 5570
},
"reportMissingTypeArgument": {
"limit": 15281
"limit": 15279
},
"reportMissingTypeStubs": {
"limit": 40
@ -99,31 +99,31 @@
"limit": 0
},
"reportUnknownArgumentType": {
"limit": 44025
"limit": 44019
},
"reportUnknownLambdaType": {
"limit": 109
},
"reportUnknownMemberType": {
"limit": 38269
"limit": 38266
},
"reportUnknownParameterType": {
"limit": 19584
"limit": 19583
},
"reportUnknownVariableType": {
"limit": 29813
"limit": 29810
},
"reportUnnecessaryCast": {
"limit": 110
},
"reportUnnecessaryComparison": {
"limit": 687
"limit": 686
},
"reportUnnecessaryContains": {
"limit": 4
},
"reportUnnecessaryIsInstance": {
"limit": 816
"limit": 815
},
"reportUntypedBaseClass": {
"limit": 0

View file

@ -27,7 +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 adispatch, dispatch
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,
@ -369,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,
@ -385,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
@ -407,114 +417,67 @@ 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,
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,
)
def native_completion() -> DispatchResult[ModelResponse]:
return rust_chat_completions_bridge.chat_completions(
model=model,
messages=messages,
optional_params=rust_optional_params,
model_response=model_response,
api_key=api_key,
additional_args=rust_logging_args,
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,
)
if acompletion is True:
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 adispatch(
native=lambda: 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=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
),
python=python_fallback,
route="chat_completions",
errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model),
)
rust_response: Final = dispatch(
native=lambda: rust_chat_completions_bridge.chat_completions(
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=lambda: None,
route="chat_completions",
errors=rust_chat_completions_bridge.error_handling(custom_llm_provider or "", model),
)
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,
@ -536,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,
@ -558,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
@ -590,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),
@ -617,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

@ -5,7 +5,8 @@ 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, adispatch, dispatch
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
@ -40,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,
*,
@ -52,23 +84,39 @@ class BedrockAudioTranscriptionRustDispatch:
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> TranscriptionResponse:
rust_response: Final = dispatch(
native=lambda: 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,
),
python=_unavailable,
route="audio transcription",
errors=PROPAGATE,
)
return TranscriptionResponse(**rust_response)
_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,
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))
@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,
*,
@ -81,19 +129,4 @@ class BedrockAudioTranscriptionRustDispatch:
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> TranscriptionResponse:
rust_response: Final = await adispatch(
native=lambda: 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,
),
python=_aunavailable,
route="audio transcription",
errors=PROPAGATE,
)
return TranscriptionResponse(**rust_response)
await _aunavailable()

View file

@ -18,7 +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 adispatch, dispatch
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
@ -407,83 +408,62 @@ 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 adispatch(
native=lambda: 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,
),
python=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,
),
route="chat_completions",
errors=rust_chat_completions_bridge.error_handling("bedrock", model),
)
rust_response: Final = dispatch(
native=lambda: rust_chat_completions_bridge.chat_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,
),
python=lambda: None,
route="chat_completions",
errors=rust_chat_completions_bridge.error_handling("bedrock", model),
)
if rust_response is not None:
return rust_response
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
logging_obj=logging_obj,
messages=messages,
api_key="",
additional_args=rust_logging_args,
)
### ROUTING (ASYNC, STREAMING, SYNC)
if acompletion:
if isinstance(client, HTTPHandler):
client = None
def native_completion() -> DispatchResult[ModelResponse]:
return rust_chat_completions_bridge.chat_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,
)
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,
@ -496,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,
@ -504,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,
@ -517,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"}
from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, adispatch, async_none
rust_response: Final = await adispatch(
native=lambda: 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,
),
python=async_none,
route="messages",
errors=PYTHON_ON_ERROR,
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,
)
if rust_response is None:
return None
response_obj: Final = cast(AnthropicMessagesResponse, dict(rust_response))
response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
return response_obj
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,26 +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
from litellm.rust_bridge.dispatch import PYTHON_ON_ERROR, adispatch, async_none
rust_backend: Final = await adispatch(
native=lambda: rust_responses_websocket.connect(
url=ws_url,
headers={str(key): str(value) for key, value in headers.items()},
timeout=timeout,
),
python=async_none,
route="responses_websocket",
errors=PYTHON_ON_ERROR,
)
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

@ -7,7 +7,7 @@ import base64
import mimetypes
import os
import re
from collections.abc import Coroutine, Mapping
from collections.abc import Callable, Coroutine, Mapping
from io import IOBase
from typing import Any, Final, cast
@ -25,7 +25,8 @@ from litellm.llms.base_llm.ocr.transformation import (
)
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.rust_bridge.dispatch import PROPAGATE, adispatch, dispatch
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
@ -155,6 +156,69 @@ def _prepare_ocr_request(
)
@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],
) -> 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,
)
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 _attempt_ocr(
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
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,
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=is_async,
headers=prepared_request.extra_headers,
provider_config=prepared_request.provider_config,
litellm_params=prepared_request.litellm_params,
)
@client
async def aocr(
model: str,
@ -251,32 +315,7 @@ async def aocr(
from litellm.secret_managers.main import get_secret_str
async def python_fallback() -> OCRResponse:
pending: 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=True,
headers=prepared.extra_headers,
provider_config=prepared.provider_config,
litellm_params=prepared.litellm_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
return await adispatch(
native=lambda: rust_ocr_bridge.aattempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str),
python=python_fallback,
route="ocr",
errors=PROPAGATE,
)
return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str)
except Exception as e:
raise litellm.exception_type(
model=model,
@ -517,28 +556,7 @@ def ocr(
from litellm.secret_managers.main import get_secret_str
def python_fallback() -> OCRResponse | Coroutine[object, object, OCRResponse]:
return 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 dispatch(
native=lambda: rust_ocr_bridge.attempt_ocr(prepared_request=prepared, resolve_api_key=get_secret_str),
python=python_fallback,
route="ocr",
errors=PROPAGATE,
)
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

@ -225,6 +225,7 @@ def chat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
on_response: ResponseObserver,
eligible: bool = True,
) -> DispatchResult[ModelResponse]:
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
on_response(rust_response)
@ -245,7 +246,7 @@ def chat_completions(
return attempt(
load=_CHAT.load,
enabled=rust_enabled(),
eligible=True,
eligible=eligible,
prepare=lambda: timeout_to_seconds(timeout),
call=call,
adapt=adapt,
@ -264,6 +265,7 @@ async def achat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
on_response: ResponseObserver,
eligible: bool = True,
) -> DispatchResult[ModelResponse]:
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
on_response(rust_response)
@ -284,7 +286,7 @@ async def achat_completions(
return await aattempt(
load=_ACHAT.load,
enabled=rust_enabled(),
eligible=True,
eligible=eligible,
prepare=lambda: timeout_to_seconds(timeout),
call=call,
adapt=adapt,

View file

@ -1,9 +1,11 @@
from __future__ import annotations
from collections.abc import Awaitable, Callable
from collections.abc import AsyncGenerator, Awaitable, Callable
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from dataclasses import dataclass
from enum import Enum
from typing import Final, TypeAlias, TypeVar
from functools import wraps
from typing import Final, ParamSpec, TypeAlias, TypeVar
from litellm._logging import verbose_logger
from litellm.exceptions import APIError
@ -12,6 +14,7 @@ from litellm.rust_bridge.runtime import DispatchResult, Handled, NativeFailed, N
NativeT = TypeVar("NativeT")
PythonT = TypeVar("PythonT")
P = ParamSpec("P")
class ErrorAction(Enum):
@ -89,49 +92,95 @@ def _log_skip(route: str, skipped: NativeSkipped) -> None:
verbose_logger.debug("Native %s skipped (%s): %s", route, skipped.reason.value, skipped.detail or "")
def dispatch(
def native_first(
*,
native: Callable[[], DispatchResult[NativeT]],
python: Callable[[], PythonT],
native: Callable[P, DispatchResult[NativeT]],
route: str,
errors: ErrorHandling,
) -> NativeT | PythonT:
try:
attempted: Final = native()
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED)
_log_skip(route, unexpected)
return python()
result: Final = _resolve(attempted, errors, route)
match result:
case Handled(value):
return value
case NativeSkipped():
_log_skip(route, result)
return python()
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
async def adispatch(
def anative_first(
*,
native: Callable[[], Awaitable[DispatchResult[NativeT]]],
python: Callable[[], Awaitable[PythonT]],
native: Callable[P, Awaitable[DispatchResult[NativeT]]],
route: str,
errors: ErrorHandling,
) -> NativeT | PythonT:
try:
attempted: Final = await native()
except Exception as error: # noqa: BLE001 # preserve declared handling of loading and adaptation failures
unexpected: Final = _handle_error(error, errors.unexpected, route, NativeSkipReason.FAILED)
_log_skip(route, unexpected)
return await python()
result: Final = _resolve(attempted, errors, route)
match result:
case Handled(value):
return value
case NativeSkipped():
_log_skip(route, result)
return await python()
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
async def async_none() -> None:
return None
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,6 +2,8 @@
from __future__ import annotations
from collections.abc import AsyncGenerator
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from typing import Final
import httpx
@ -13,7 +15,7 @@ from litellm.rust_bridge.protocols import (
RustResponsesWebSocket,
RustResponsesWebSocketConnection,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt
from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result
from litellm.rust_bridge.timeouts import timeout_to_seconds
_RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding(
@ -32,7 +34,7 @@ def set_rust_responses_websocket(
_RESPONSES_WEBSOCKET.override(connection)
class _ConnectionAdapter:
class ConnectionAdapter:
def __init__(self, connection: RustResponsesWebSocket):
self._connection: Final[RustResponsesWebSocket] = connection
@ -54,7 +56,7 @@ async def connect(
url: str,
headers: dict[str, str],
timeout: float | httpx.Timeout | None,
) -> DispatchResult[_ConnectionAdapter]:
) -> DispatchResult[ConnectionAdapter]:
return await aattempt(
load=_RESPONSES_WEBSOCKET.load,
enabled=rust_enabled(),
@ -65,5 +67,23 @@ async def connect(
headers=headers,
timeout_seconds=timeout_seconds,
),
adapt=_ConnectionAdapter,
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

@ -87,3 +87,9 @@ async def aattempt(
def identity(value: ResultT) -> ResultT:
return value
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

@ -12,7 +12,7 @@
"limit": 1979
},
"ANN202": {
"limit": 831
"limit": 830
},
"ANN204": {
"limit": 683
@ -150,7 +150,7 @@
"limit": 253
},
"PLW0127": {
"limit": 57
"limit": 55
},
"PLW0602": {
"limit": 215
@ -195,7 +195,7 @@
"limit": 22
},
"SIM101": {
"limit": 56
"limit": 55
},
"SIM102": {
"limit": 310

View file

@ -9,7 +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, NativeSkipped, NativeSkipReason
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
@ -127,6 +127,16 @@ def test_load_rust_messages_returns_injected_impl():
assert rust_messages.load_rust_messages() is bridge
def test_bare_rust_still_toggles_ocr():
from litellm.rust_bridge.ocr import rust_ocr_enabled
litellm.rust(True)
assert rust_ocr_enabled() is True
litellm.rust(False)
assert rust_ocr_enabled() is False
def test_load_rust_amessages_returns_injected_impl():
bridge = RecordingAsyncMessages()
litellm.rust(True)
@ -215,7 +225,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
@ -226,7 +236,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]
@ -239,13 +250,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 +267,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,10 +279,24 @@ 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"
@pytest.mark.asyncio
async def test_gate_ignores_request_flag_when_process_enabled():
bridge = RecordingAsyncMessages()
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False))
assert isinstance(response, Handled)
response = response.value
assert len(bridge.calls) == 1
@pytest.mark.asyncio
async def test_gate_invokes_rust_for_native_anthropic_provider():
bridge = RecordingAsyncMessages()
@ -286,7 +311,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 +329,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 +345,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 +357,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 +369,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 +385,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
@ -391,4 +419,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

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

@ -4,7 +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, NativeSkipped, NativeSkipReason, NativeFailed
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
class _FakeNativeConnection:
@ -58,7 +58,7 @@ 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()
@ -115,3 +115,30 @@ async def test_connection_failure_is_reported_to_orchestration() -> None:
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

@ -10,7 +10,7 @@ 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, adispatch, dispatch
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
@ -53,9 +53,9 @@ async def test_shared_dispatch_calls_python_once_and_logs_skip(
return python()
result: Final = (
await adispatch(native=anative, python=apython, route="test", errors=PROPAGATE)
await anative_first(native=anative, route="test", errors=lambda: PROPAGATE)(apython)()
if asynchronous
else dispatch(native=native, python=python, route="test", errors=PROPAGATE)
else native_first(native=native, route="test", errors=lambda: PROPAGATE)(python)()
)
assert result == "python response"
assert calls == ["native", "python"]
@ -65,6 +65,7 @@ async def test_shared_dispatch_calls_python_once_and_logs_skip(
@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)
@ -75,9 +76,9 @@ async def test_native_success_does_not_run_python_even_when_value_is_none(asynch
return python()
result: Final = (
await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR)
await anative_first(native=native, route="test", errors=lambda: PYTHON_ON_ERROR)(apython)()
if asynchronous
else dispatch(native=lambda: Handled(None), python=python, route="test", errors=PYTHON_ON_ERROR)
else native_first(native=lambda: Handled(None), route="test", errors=lambda: PYTHON_ON_ERROR)(python)()
)
assert result is None
@ -124,8 +125,8 @@ async def test_declarations_preserve_endpoint_error_behavior(
async def run() -> str:
if asynchronous:
return await adispatch(native=anative, python=apython, route="chat_completions", errors=rules)
return dispatch(native=native, python=python, route="chat_completions", errors=rules)
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"
@ -163,13 +164,10 @@ async def test_python_failure_is_never_reclassified_as_native_failure(asynchrono
async def run() -> str:
if asynchronous:
return await adispatch(native=native, python=apython, route="test", errors=PYTHON_ON_ERROR)
return dispatch(
native=lambda: NativeSkipped(NativeSkipReason.UNAVAILABLE),
python=python,
route="test",
errors=PYTHON_ON_ERROR,
)
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()
@ -179,6 +177,7 @@ async def test_python_failure_is_never_reclassified_as_native_failure(asynchrono
@pytest.mark.asyncio
async def test_cancellation_does_not_run_python() -> None:
async def native() -> DispatchResult[str]:
raise asyncio.CancelledError
@ -186,20 +185,123 @@ async def test_cancellation_does_not_run_python() -> None:
pytest.fail("cancellation must not dispatch Python")
with pytest.raises(asyncio.CancelledError):
await adispatch(native=native, python=python, route="test", errors=PYTHON_ON_ERROR)
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:
dispatch(
native_first(
native=lambda: NativeFailed(error),
python=lambda: pytest.fail("upstream errors must not run Python"),
route="chat_completions",
errors=error_handling("anthropic", "model"),
)
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

@ -30,7 +30,7 @@
"limit": 16419
},
"LIT011": {
"limit": 5506
"limit": 5497
},
"LIT012": {
"limit": 4486