mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
feat(proxy): centralize Rust fallback and provenance
This commit is contained in:
parent
b461ae6d63
commit
2395450714
26 changed files with 1726 additions and 1056 deletions
|
|
@ -4,9 +4,10 @@ mod service;
|
|||
|
||||
use axum::Router;
|
||||
use axum::body::Body;
|
||||
use axum::extract::{Json, State};
|
||||
use axum::extract::{Json, Request, State};
|
||||
use axum::http::StatusCode;
|
||||
use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE, HeaderMap, HeaderValue};
|
||||
use axum::middleware::{self, Next};
|
||||
use axum::response::{IntoResponse, Response};
|
||||
use axum::routing::post;
|
||||
use litellm_core::Error;
|
||||
|
|
@ -16,9 +17,23 @@ use crate::auth::RequireMasterKey;
|
|||
use crate::constants::{MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH};
|
||||
use crate::state::AppState;
|
||||
|
||||
const CORE_ENGINE_HEADER: &str = "x-litellm-core";
|
||||
const RUST_CORE_ENGINE: &str = "rust";
|
||||
|
||||
/// This route's contribution to the app router.
|
||||
pub fn router() -> Router<AppState> {
|
||||
Router::new().route(MESSAGES_ROUTE_PATH, post(handle))
|
||||
Router::new()
|
||||
.route(MESSAGES_ROUTE_PATH, post(handle))
|
||||
.route_layer(middleware::from_fn(core_engine_header))
|
||||
}
|
||||
|
||||
async fn core_engine_header(request: Request, next: Next) -> Response {
|
||||
let mut response = next.run(request).await;
|
||||
response.headers_mut().insert(
|
||||
CORE_ENGINE_HEADER,
|
||||
HeaderValue::from_static(RUST_CORE_ENGINE),
|
||||
);
|
||||
response
|
||||
}
|
||||
|
||||
async fn handle(
|
||||
|
|
@ -143,6 +158,7 @@ mod tests {
|
|||
use tower::ServiceExt;
|
||||
|
||||
use super::super::app;
|
||||
use super::{CORE_ENGINE_HEADER, RUST_CORE_ENGINE};
|
||||
use crate::io::realtime_pool::RealtimePool;
|
||||
use crate::state::AppState;
|
||||
|
||||
|
|
@ -290,6 +306,7 @@ mod tests {
|
|||
.await
|
||||
.expect("route responds");
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE);
|
||||
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("response body reads");
|
||||
|
|
@ -340,6 +357,7 @@ mod tests {
|
|||
.await
|
||||
.expect("route responds");
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE);
|
||||
let upstream_request = server.await.expect("upstream task completes");
|
||||
let (_, upstream_body) = upstream_request
|
||||
.split_once("\r\n\r\n")
|
||||
|
|
@ -378,6 +396,7 @@ mod tests {
|
|||
.await
|
||||
.expect("route responds");
|
||||
assert_eq!(response.status(), StatusCode::OK);
|
||||
assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE);
|
||||
assert_eq!(
|
||||
response
|
||||
.headers()
|
||||
|
|
@ -443,6 +462,7 @@ mod tests {
|
|||
.await
|
||||
.expect("route responds");
|
||||
assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
|
||||
assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE);
|
||||
let response_body = axum::body::to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
.expect("response body reads");
|
||||
|
|
@ -473,6 +493,7 @@ mod tests {
|
|||
.await
|
||||
.expect("route responds");
|
||||
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
|
||||
assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
@ -495,6 +516,7 @@ mod tests {
|
|||
.await
|
||||
.expect("route responds");
|
||||
assert_eq!(response.status(), StatusCode::UNAUTHORIZED);
|
||||
assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
|
@ -517,5 +539,6 @@ mod tests {
|
|||
.await
|
||||
.expect("route responds");
|
||||
assert_eq!(response.status(), StatusCode::BAD_REQUEST);
|
||||
assert_eq!(response.headers()[CORE_ENGINE_HEADER], RUST_CORE_ENGINE);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ Calling + translation logic for anthropic's `/v1/messages` endpoint
|
|||
|
||||
import copy
|
||||
import json
|
||||
from collections.abc import Callable
|
||||
from collections.abc import Callable, Mapping
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast
|
||||
|
||||
import httpx
|
||||
|
|
@ -380,7 +380,10 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
request_data: Final = config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params={**optional_params, "is_vertex_request": is_vertex_request},
|
||||
optional_params={ # mutable-ok: provider transforms require a plain mutable request dict
|
||||
**optional_params,
|
||||
"is_vertex_request": is_vertex_request,
|
||||
},
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
@ -390,6 +393,46 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
def sync_python_path(
|
||||
request_headers: Mapping[str, object], request_data: Mapping[str, object]
|
||||
) -> ModelResponse:
|
||||
request_client: Final = (
|
||||
_get_httpx_client(params={"timeout": timeout}) # mutable-ok: HTTP client factory requires a dict
|
||||
if client is None or not isinstance(client, HTTPHandler)
|
||||
else client
|
||||
)
|
||||
try:
|
||||
response: Final = request_client.post(
|
||||
api_base,
|
||||
headers=dict(request_headers), # mutable-ok: HTTP client requires mutable headers
|
||||
data=json.dumps(request_data),
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 # provider exceptions are normalized below
|
||||
status_code: Final = getattr(error, "status_code", 500)
|
||||
error_headers = getattr(error, "headers", None)
|
||||
error_text = getattr(error, "text", str(error))
|
||||
error_response: Final = getattr(error, "response", None)
|
||||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
if error_response and hasattr(error_response, "text"):
|
||||
error_text = getattr(error_response, "text", error_text)
|
||||
raise AnthropicError(message=error_text, status_code=status_code, headers=error_headers)
|
||||
return config.transform_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
model_response=model_response,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
request_data=dict(request_data), # mutable-ok: response transform mutates request metadata
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
encoding=encoding,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
|
||||
# The Rust core owns the whole call for the subset it accepts, so ask
|
||||
# before transforming: whichever path runs emits pre_call exactly once.
|
||||
# `get_config` merges the class-level defaults (Anthropic's required
|
||||
|
|
@ -466,7 +509,12 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
on_response=log_rust_post_call,
|
||||
python_fallback=python_fallback,
|
||||
)
|
||||
rust_response: Final = rust_chat_completions_bridge.chat_completions(
|
||||
|
||||
def sync_python_fallback() -> ModelResponse:
|
||||
fallback_headers, fallback_data = build_request()
|
||||
return sync_python_path(fallback_headers, fallback_data)
|
||||
|
||||
return rust_chat_completions_bridge.chat_completions_or_fallback(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
|
|
@ -477,26 +525,21 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=sync_python_fallback,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
|
||||
headers, data = build_request()
|
||||
|
||||
## LOGGING
|
||||
# Reaching here with `serves_via_rust` set means the Rust attempt
|
||||
# declined at call time, before the provider was called, and already
|
||||
# logged this request. That is the same attempt continuing.
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={ # mutable-ok: logging callback contract requires a mutable dict
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
print_verbose(f"_is_function_call: {_is_function_call}")
|
||||
if acompletion is True:
|
||||
if (
|
||||
|
|
@ -578,48 +621,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
_response_headers=process_anthropic_headers(headers),
|
||||
)
|
||||
|
||||
else:
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
client = _get_httpx_client(params={"timeout": timeout})
|
||||
else:
|
||||
client = client
|
||||
|
||||
try:
|
||||
response: Final = client.post(
|
||||
api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
timeout=timeout,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
status_code: Final = getattr(e, "status_code", 500)
|
||||
error_headers = getattr(e, "headers", None)
|
||||
error_text = getattr(e, "text", str(e))
|
||||
error_response: Final[object] = getattr(e, "response", None)
|
||||
if error_headers is None and error_response:
|
||||
error_headers = getattr(error_response, "headers", None)
|
||||
if error_response and hasattr(error_response, "text"):
|
||||
error_text = getattr(error_response, "text", error_text)
|
||||
raise AnthropicError(
|
||||
message=error_text,
|
||||
status_code=status_code,
|
||||
headers=error_headers,
|
||||
)
|
||||
|
||||
return config.transform_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
model_response=model_response,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
request_data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
encoding=encoding,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
return sync_python_path(headers, data)
|
||||
|
||||
def embedding(self):
|
||||
# logic for parsing in - calling - parsing out model embedding calls
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from typing import Any, Final, Protocol, runtime_checkable
|
|||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import TypedDict
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.constants import (
|
||||
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS,
|
||||
|
|
@ -19,6 +19,7 @@ from litellm.llms.anthropic.common_utils import ANTHROPIC_ERROR_STATUS_CODE_MAP
|
|||
from litellm.proxy.pass_through_endpoints.success_handler import (
|
||||
PassThroughEndpointLogging,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import CoreEngine, execution_additional_headers
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType
|
||||
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
|
||||
|
|
@ -306,7 +307,8 @@ def _anthropic_content_block_events(index: int, block: Mapping[str, object]) ->
|
|||
|
||||
|
||||
class AnthropicMessagesStreamHiddenParams(TypedDict):
|
||||
additional_headers: dict[str, str]
|
||||
additional_headers: ReadOnly[Mapping[str, str]]
|
||||
core_engine: ReadOnly[str]
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
|
|
@ -325,8 +327,10 @@ _RESPONSE_HEADERS_ADAPTER: Final[TypeAdapter[dict[str, str]]] = TypeAdapter(dict
|
|||
def anthropic_messages_stream_hidden_params(
|
||||
response_headers: httpx.Headers,
|
||||
) -> AnthropicMessagesStreamHiddenParams:
|
||||
additional_headers: Final = _RESPONSE_HEADERS_ADAPTER.validate_python(process_response_headers(response_headers))
|
||||
return AnthropicMessagesStreamHiddenParams(
|
||||
additional_headers=_RESPONSE_HEADERS_ADAPTER.validate_python(process_response_headers(response_headers))
|
||||
additional_headers=execution_additional_headers(additional_headers, CoreEngine.PYTHON),
|
||||
core_engine=CoreEngine.PYTHON.value,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ import httpx
|
|||
|
||||
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
|
||||
from litellm.rust_bridge import transcription as rust_transcription_bridge
|
||||
from litellm.rust_bridge.runtime import CoreEngine, execution_hidden_params
|
||||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
|
||||
|
|
@ -55,7 +56,9 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
)
|
||||
if rust_response is None:
|
||||
raise RuntimeError("Rust audio transcription bridge is unavailable")
|
||||
return TranscriptionResponse(**rust_response)
|
||||
response: Final = TranscriptionResponse(**rust_response)
|
||||
response["_hidden_params"] = execution_hidden_params(None, CoreEngine.RUST)
|
||||
return response
|
||||
|
||||
async def async_audio_transcriptions(
|
||||
self,
|
||||
|
|
@ -81,4 +84,6 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
)
|
||||
if rust_response is None:
|
||||
raise RuntimeError("Rust audio transcription bridge is unavailable")
|
||||
return TranscriptionResponse(**rust_response)
|
||||
response: Final = TranscriptionResponse(**rust_response)
|
||||
response["_hidden_params"] = execution_hidden_params(None, CoreEngine.RUST)
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -365,6 +365,91 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
# Filter beta headers in HTTP headers before making the request
|
||||
headers = update_headers_with_filtered_beta(headers=headers, provider="bedrock_converse")
|
||||
|
||||
def sync_python_path(*, skip_pre_call_logging: bool) -> ModelResponse | CustomStreamWrapper:
|
||||
request_data: Final = litellm.AmazonConverseConfig()._transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=extra_headers,
|
||||
)
|
||||
serialized_data: Final = json.dumps(request_data)
|
||||
prepped: Final = self.get_request_headers(
|
||||
credentials=credentials,
|
||||
aws_region_name=aws_region_name,
|
||||
extra_headers=extra_headers,
|
||||
endpoint_url=proxy_endpoint_url,
|
||||
data=serialized_data,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
)
|
||||
if not skip_pre_call_logging:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={ # mutable-ok: logging callback contract requires a mutable dict
|
||||
"complete_input_dict": serialized_data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
request_timeout: Final = (
|
||||
httpx.Timeout(timeout) if isinstance(timeout, float) or isinstance(timeout, int) else timeout
|
||||
)
|
||||
request_client: Final = (
|
||||
_get_httpx_client( # mutable-ok: HTTP client factory requires a mutable options dict
|
||||
{"timeout": request_timeout} if request_timeout is not None else {}
|
||||
)
|
||||
if client is None or isinstance(client, AsyncHTTPHandler)
|
||||
else client
|
||||
)
|
||||
if stream is not None and stream is True:
|
||||
completion_stream, response_headers = make_sync_call(
|
||||
client=request_client if isinstance(request_client, HTTPHandler) else None,
|
||||
api_base=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=serialized_data,
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
return CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
_response_headers=response_headers,
|
||||
)
|
||||
try:
|
||||
response: Final = request_client.post(
|
||||
url=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=serialized_data,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as error:
|
||||
raise BedrockError(status_code=error.response.status_code, message=error.response.text)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
transformed: Final = litellm.AmazonConverseConfig()._transform_response(
|
||||
model=model,
|
||||
response=response,
|
||||
model_response=model_response,
|
||||
stream=stream if isinstance(stream, bool) else False,
|
||||
logging_obj=logging_obj,
|
||||
api_key="",
|
||||
data=serialized_data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
encoding=encoding,
|
||||
)
|
||||
transformed.set_provider_response_headers(response.headers)
|
||||
return transformed
|
||||
|
||||
# The Rust core owns the whole call for the subset it accepts. Ask
|
||||
# before transforming so whichever path runs emits pre_call once, and
|
||||
# hand down the credentials, region and endpoint this handler already
|
||||
|
|
@ -437,7 +522,7 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
skip_pre_call_logging=True,
|
||||
),
|
||||
)
|
||||
rust_response: Final = rust_chat_completions_bridge.chat_completions(
|
||||
return rust_chat_completions_bridge.chat_completions_or_fallback(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
|
|
@ -448,9 +533,8 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
python_fallback=lambda: sync_python_path(skip_pre_call_logging=True),
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
|
||||
### ROUTING (ASYNC, STREAMING, SYNC)
|
||||
if acompletion:
|
||||
|
|
@ -496,103 +580,4 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
api_key=api_key,
|
||||
)
|
||||
|
||||
## TRANSFORMATION ##
|
||||
|
||||
_data: Final = litellm.AmazonConverseConfig()._transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=extra_headers,
|
||||
)
|
||||
data: Final = json.dumps(_data)
|
||||
|
||||
prepped: Final = self.get_request_headers(
|
||||
credentials=credentials,
|
||||
aws_region_name=aws_region_name,
|
||||
extra_headers=extra_headers,
|
||||
endpoint_url=proxy_endpoint_url,
|
||||
data=data,
|
||||
headers=headers,
|
||||
api_key=api_key,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
# Reaching here with `serves_via_rust` set means the synchronous Rust
|
||||
# attempt declined at call time, before the provider was called, and
|
||||
# already logged this request. That is the same attempt continuing.
|
||||
# The asynchronous branch above returns before this point, and hands
|
||||
# its own fallback `skip_pre_call_logging=True` for the same reason.
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
if client is None or isinstance(client, AsyncHTTPHandler):
|
||||
_params: Final = {}
|
||||
if timeout is not None:
|
||||
if isinstance(timeout, float) or isinstance(timeout, int):
|
||||
timeout = httpx.Timeout(timeout)
|
||||
_params["timeout"] = timeout
|
||||
client = _get_httpx_client(_params)
|
||||
else:
|
||||
client = client
|
||||
|
||||
if stream is not None and stream is True:
|
||||
completion_stream, response_headers = make_sync_call(
|
||||
client=(client if client is not None and isinstance(client, HTTPHandler) else None),
|
||||
api_base=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
streaming_response: Final = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
_response_headers=response_headers,
|
||||
)
|
||||
|
||||
return streaming_response
|
||||
|
||||
### COMPLETION
|
||||
|
||||
try:
|
||||
response: Final = client.post(
|
||||
url=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
|
||||
model=model,
|
||||
response=response,
|
||||
model_response=model_response,
|
||||
stream=stream if isinstance(stream, bool) else False,
|
||||
logging_obj=logging_obj,
|
||||
api_key="",
|
||||
data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
encoding=encoding,
|
||||
)
|
||||
sync_transformed_response.set_provider_response_headers(response.headers)
|
||||
return sync_transformed_response
|
||||
return sync_python_path(skip_pre_call_logging=False)
|
||||
|
|
|
|||
|
|
@ -19,7 +19,6 @@ import litellm.types.utils
|
|||
from litellm._logging import _redact_string, verbose_logger
|
||||
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.litellm_core_utils.agentic_loop_settings import (
|
||||
DEFAULT_MAX_AGENTIC_LOOPS,
|
||||
validated_max_agentic_loops,
|
||||
|
|
@ -89,6 +88,7 @@ from litellm.responses.streaming_iterator import (
|
|||
ResponsesWebSocketStreaming,
|
||||
SyncResponsesAPIStreamingIterator,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import CoreEngine, execution_additional_headers, execution_hidden_params
|
||||
from litellm.types.containers.main import (
|
||||
ContainerFileListResponse,
|
||||
ContainerListResponse,
|
||||
|
|
@ -235,6 +235,23 @@ def _google_genai_streaming_hidden_params(
|
|||
}
|
||||
|
||||
|
||||
def _anthropic_messages_with_core_engine(
|
||||
response: AnthropicMessagesResponse,
|
||||
source: CoreEngine,
|
||||
) -> AnthropicMessagesResponse:
|
||||
existing_hidden_params: Final = response.get("_hidden_params")
|
||||
return cast( # cast-ok: Anthropic response TypedDict is preserved while adding its documented hidden field
|
||||
AnthropicMessagesResponse,
|
||||
{ # mutable-ok: Anthropic response contract is a mutable TypedDict
|
||||
**response,
|
||||
"_hidden_params": execution_hidden_params(
|
||||
existing_hidden_params if isinstance(existing_hidden_params, dict) else None,
|
||||
source,
|
||||
),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@lru_cache(maxsize=None)
|
||||
def _responses_api_optional_request_param_names() -> frozenset[str]:
|
||||
return frozenset(get_type_hints(ResponsesAPIOptionalRequestParams).keys())
|
||||
|
|
@ -2361,8 +2378,17 @@ class BaseLLMHTTPHandler:
|
|||
kwargs=kwargs_for_agentic,
|
||||
)
|
||||
|
||||
return self._maybe_wrap_in_fake_stream(
|
||||
selected_response: Final = cast( # cast-ok: hook result is validated by the Anthropic response pipeline
|
||||
AnthropicMessagesResponse,
|
||||
final_response if final_response is not None else initial_response,
|
||||
)
|
||||
existing_hidden_params: Final = selected_response.get("_hidden_params")
|
||||
existing_source: Final = (
|
||||
existing_hidden_params.get("core_engine") if isinstance(existing_hidden_params, dict) else None
|
||||
)
|
||||
source: Final = CoreEngine.RUST if existing_source == CoreEngine.RUST.value else CoreEngine.PYTHON
|
||||
return self._maybe_wrap_in_fake_stream(
|
||||
_anthropic_messages_with_core_engine(selected_response, source),
|
||||
logging_obj,
|
||||
"anthropic_messages",
|
||||
)
|
||||
|
|
@ -2394,45 +2420,20 @@ class BaseLLMHTTPHandler:
|
|||
from litellm.rust_bridge import messages as rust_messages_bridge
|
||||
|
||||
upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"}
|
||||
try:
|
||||
rust_response: Final = await rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
body=upstream_body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as rust_error: # noqa: BLE001
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
native_bridge: Final = get_native_bridge()
|
||||
declined: Final = getattr(native_bridge, "RustBridgeDeclined", None)
|
||||
upstream_failed: Final = getattr(native_bridge, "RustUpstreamError", None)
|
||||
if isinstance(upstream_failed, type) and isinstance(rust_error, upstream_failed):
|
||||
args: Final = rust_error.args
|
||||
status: Final = args[0] if args else 0
|
||||
message: Final = args[1] if len(args) > 1 else ""
|
||||
raise APIError(
|
||||
status_code=int(status) or 500,
|
||||
message=f"litellm rust messages: {message}",
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
) from rust_error
|
||||
if not isinstance(declined, type) or not isinstance(rust_error, declined):
|
||||
raise
|
||||
verbose_logger.debug(
|
||||
"Rust Anthropic messages bridge declined before calling the provider (%s); falling back to Python path",
|
||||
rust_error,
|
||||
)
|
||||
return None
|
||||
rust_response: Final = await rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
body=upstream_body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_response is None:
|
||||
return None
|
||||
|
||||
response_obj: Final = cast(AnthropicMessagesResponse, dict(rust_response))
|
||||
response_obj["_hidden_params"] = {"additional_headers": {"x-litellm-rust": "true"}}
|
||||
return response_obj
|
||||
return _anthropic_messages_with_core_engine(response_obj, CoreEngine.RUST)
|
||||
|
||||
@staticmethod
|
||||
def _rust_anthropic_messages_fake_stream(
|
||||
|
|
@ -2447,7 +2448,10 @@ class BaseLLMHTTPHandler:
|
|||
)
|
||||
|
||||
completion_stream = cast(AsyncIterator[bytes], FakeAnthropicMessagesStreamIterator(response=rust_response))
|
||||
hidden_params: Final = AnthropicMessagesStreamHiddenParams(additional_headers={"x-litellm-rust": "true"})
|
||||
hidden_params: Final = AnthropicMessagesStreamHiddenParams(
|
||||
additional_headers=execution_additional_headers(None, CoreEngine.RUST),
|
||||
core_engine=CoreEngine.RUST.value,
|
||||
)
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=completion_stream,
|
||||
hidden_params=hidden_params,
|
||||
|
|
@ -6508,13 +6512,14 @@ class BaseLLMHTTPHandler:
|
|||
if _rust_responses_websocket_enabled(custom_llm_provider, litellm_params):
|
||||
from litellm.rust_bridge import responses_websocket as rust_responses_websocket
|
||||
|
||||
rust_backend: Final = await rust_responses_websocket.connect(
|
||||
rust_execution: Final = await rust_responses_websocket.connect(
|
||||
provider="openai",
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers={str(key): str(value) for key, value in headers.items()},
|
||||
timeout=timeout,
|
||||
)
|
||||
rust_backend: Final = rust_execution.value
|
||||
if rust_backend is not None:
|
||||
yield rust_backend
|
||||
return
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from litellm.llms.base_llm.ocr.transformation import (
|
|||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
from litellm.rust_bridge.runtime import CoreEngine, execution_hidden_params
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
|
|
@ -60,6 +61,24 @@ class _PreparedRustOCRCall:
|
|||
optional_params: dict[str, object]
|
||||
|
||||
|
||||
def _with_core_engine(response: OCRResponse, source: CoreEngine) -> OCRResponse:
|
||||
existing_hidden_params: Final = cast( # cast-ok: OCRResponse initializes hidden params as a mapping
|
||||
Mapping[str, object],
|
||||
response._hidden_params, # pyright: ignore[reportPrivateUsage] # provenance uses the response's established metadata channel
|
||||
)
|
||||
response._hidden_params = execution_hidden_params( # pyright: ignore[reportPrivateUsage] # OCRResponse has no public hidden-params setter
|
||||
existing_hidden_params, source
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
async def _await_with_core_engine(
|
||||
response: Coroutine[object, object, OCRResponse],
|
||||
source: CoreEngine,
|
||||
) -> OCRResponse:
|
||||
return _with_core_engine(await response, source)
|
||||
|
||||
|
||||
_RUST_OCR_PROVIDERS: Final = {
|
||||
"mistral",
|
||||
"azure_ai",
|
||||
|
|
@ -306,7 +325,7 @@ def _run_rust_ocr(
|
|||
)
|
||||
if rust_response is None:
|
||||
return None
|
||||
return OCRResponse.model_validate(rust_response)
|
||||
return _with_core_engine(OCRResponse.model_validate(rust_response), CoreEngine.RUST)
|
||||
|
||||
|
||||
async def _run_rust_aocr(
|
||||
|
|
@ -331,7 +350,7 @@ async def _run_rust_aocr(
|
|||
)
|
||||
if rust_response is None:
|
||||
return None
|
||||
return OCRResponse.model_validate(rust_response)
|
||||
return _with_core_engine(OCRResponse.model_validate(rust_response), CoreEngine.RUST)
|
||||
|
||||
|
||||
@client
|
||||
|
|
@ -461,7 +480,7 @@ async def aocr(
|
|||
if response is None:
|
||||
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
|
||||
|
||||
return response
|
||||
return _with_core_engine(response, CoreEngine.PYTHON)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
@ -727,7 +746,9 @@ def ocr(
|
|||
litellm_params=prepared.litellm_params,
|
||||
)
|
||||
|
||||
return response
|
||||
if asyncio.iscoroutine(response):
|
||||
return _await_with_core_engine(response, CoreEngine.PYTHON)
|
||||
return _with_core_engine(response, CoreEngine.PYTHON)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -1550,6 +1550,16 @@ class ProxyBaseLLMRequestProcessing:
|
|||
logging_obj=litellm_logging_obj,
|
||||
use_logging_obj=read_timing_from_logging_obj,
|
||||
)
|
||||
core_engine_value: Final = hidden_params.get("core_engine")
|
||||
core_engine: Final = (
|
||||
core_engine_value
|
||||
if isinstance(core_engine_value, str) and core_engine_value in ("python", "rust")
|
||||
else "python"
|
||||
)
|
||||
reserved_core_headers: Final = frozenset({"x-litellm-core", "x-litellm-rust"})
|
||||
forwarded_headers: Final = { # mutable-ok: response headers are assembled into a mutable transport dict
|
||||
key: value for key, value in kwargs.items() if key.lower() not in reserved_core_headers
|
||||
}
|
||||
|
||||
cost_breakdown: Final = _get_cost_breakdown_from_logging_obj(
|
||||
litellm_logging_obj=litellm_logging_obj, response_cost=response_cost
|
||||
|
|
@ -1634,7 +1644,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
str(fastest_response_batch_completion) if fastest_response_batch_completion is not None else None
|
||||
),
|
||||
"x-litellm-timeout": str(timeout) if timeout is not None else None,
|
||||
**{k: str(v) for k, v in kwargs.items()},
|
||||
**{ # mutable-ok: response header assembly requires a mutable dict
|
||||
key: str(value) for key, value in forwarded_headers.items()
|
||||
},
|
||||
"x-litellm-core": core_engine,
|
||||
# mutable-ok: conditionally extending the mutable response header dict
|
||||
**({"x-litellm-rust": "true"} if core_engine == "rust" else {}),
|
||||
}
|
||||
if request_data:
|
||||
remaining_tokens_header: Final = get_remaining_tokens_and_requests_from_request_data(request_data)
|
||||
|
|
|
|||
|
|
@ -14,20 +14,33 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
from typing import TYPE_CHECKING, Final, Protocol, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
convert_to_model_response_object,
|
||||
)
|
||||
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.loader import get_native_bridge
|
||||
from litellm.rust_bridge.runtime import (
|
||||
LEGACY_RUST_HEADER,
|
||||
UNSET,
|
||||
BridgeErrorContext,
|
||||
CoreEngine,
|
||||
ExecutionResult,
|
||||
FallbackMode,
|
||||
NativeBinding,
|
||||
RustHandled,
|
||||
Unset,
|
||||
aattempt,
|
||||
ainvoke,
|
||||
attempt,
|
||||
execution_hidden_params,
|
||||
invoke,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -42,7 +55,7 @@ RUST_CHAT_COMPLETIONS_PROVIDERS: Final = frozenset({"anthropic", "bedrock"})
|
|||
# rather than narrowing an unparameterized `Mapping` and typing the result Any.
|
||||
_LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
RUST_RESPONSE_HEADER: Final = "x-litellm-rust"
|
||||
RUST_RESPONSE_HEADER: Final = LEGACY_RUST_HEADER
|
||||
|
||||
|
||||
class RustChatCompletions(Protocol):
|
||||
|
|
@ -126,67 +139,34 @@ def response_logger(
|
|||
return log
|
||||
|
||||
|
||||
class _Unset:
|
||||
pass
|
||||
|
||||
|
||||
_UNSET: Final[_Unset] = _Unset()
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _RustChatCompletionsState:
|
||||
chat_completions: RustChatCompletions | None = None
|
||||
achat_completions: RustAchatCompletions | None = None
|
||||
decline: RustChatCompletionsDecline | None = None
|
||||
|
||||
|
||||
_STATE: Final[_RustChatCompletionsState] = _RustChatCompletionsState()
|
||||
_CHAT_COMPLETIONS: Final = NativeBinding[RustChatCompletions]("chat_completions")
|
||||
_ACHAT_COMPLETIONS: Final = NativeBinding[RustAchatCompletions]("achat_completions")
|
||||
_DECLINE: Final = NativeBinding[RustChatCompletionsDecline]("chat_completions_decline")
|
||||
|
||||
|
||||
def set_rust_chat_completions(
|
||||
*,
|
||||
chat_completions: RustChatCompletions | None | _Unset = _UNSET,
|
||||
achat_completions: RustAchatCompletions | None | _Unset = _UNSET,
|
||||
decline: RustChatCompletionsDecline | None | _Unset = _UNSET,
|
||||
chat_completions: RustChatCompletions | None | Unset = UNSET,
|
||||
achat_completions: RustAchatCompletions | None | Unset = UNSET,
|
||||
decline: RustChatCompletionsDecline | None | Unset = UNSET,
|
||||
) -> None:
|
||||
"""Inject the native callables, so tests can supply a double instead of
|
||||
patching module attributes."""
|
||||
if not isinstance(chat_completions, _Unset):
|
||||
_STATE.chat_completions = chat_completions
|
||||
if not isinstance(achat_completions, _Unset):
|
||||
_STATE.achat_completions = achat_completions
|
||||
if not isinstance(decline, _Unset):
|
||||
_STATE.decline = decline
|
||||
_CHAT_COMPLETIONS.update(chat_completions)
|
||||
_ACHAT_COMPLETIONS.update(achat_completions)
|
||||
_DECLINE.update(decline)
|
||||
|
||||
|
||||
def load_rust_chat_completions() -> RustChatCompletions | None:
|
||||
if _STATE.chat_completions is not None:
|
||||
return _STATE.chat_completions
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
loaded: RustChatCompletions | None = getattr(native_bridge, "chat_completions", None)
|
||||
return loaded
|
||||
return _CHAT_COMPLETIONS.load()
|
||||
|
||||
|
||||
def load_rust_achat_completions() -> RustAchatCompletions | None:
|
||||
if _STATE.achat_completions is not None:
|
||||
return _STATE.achat_completions
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
loaded: RustAchatCompletions | None = getattr(native_bridge, "achat_completions", None)
|
||||
return loaded
|
||||
return _ACHAT_COMPLETIONS.load()
|
||||
|
||||
|
||||
def _load_rust_decline() -> RustChatCompletionsDecline | None:
|
||||
if _STATE.decline is not None:
|
||||
return _STATE.decline
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
loaded: RustChatCompletionsDecline | None = getattr(native_bridge, "chat_completions_decline", None)
|
||||
return loaded
|
||||
return _DECLINE.load()
|
||||
|
||||
|
||||
def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool:
|
||||
|
|
@ -275,57 +255,6 @@ def rust_chat_completions_accepts(
|
|||
return True
|
||||
|
||||
|
||||
def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None:
|
||||
"""`(declined, upstream_failed)` from the native module, or None when absent."""
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
declined: Final = getattr(native_bridge, "RustBridgeDeclined", None)
|
||||
upstream: Final = getattr(native_bridge, "RustUpstreamError", None)
|
||||
if declined is None or upstream is None:
|
||||
return None
|
||||
return declined, upstream
|
||||
|
||||
|
||||
def _reraise_or_decline(
|
||||
rust_error: BaseException,
|
||||
*,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
) -> None:
|
||||
"""Re-raise a failure the provider already saw, or return so the caller declines.
|
||||
|
||||
A request that never reached the provider is safe to serve on the Python
|
||||
path. One that did is not: the provider has already done the work, so a
|
||||
second attempt bills for it twice. Those surface as an `APIError` carrying
|
||||
the upstream status, which LiteLLM's exception mapping already understands.
|
||||
"""
|
||||
exceptions: Final = _rust_bridge_exceptions()
|
||||
if exceptions is None:
|
||||
verbose_logger.debug(
|
||||
"Rust chat completions bridge raised %s; falling back to Python path",
|
||||
type(rust_error).__name__,
|
||||
)
|
||||
return
|
||||
declined, upstream_failed = exceptions
|
||||
if isinstance(rust_error, upstream_failed):
|
||||
args: Final = rust_error.args
|
||||
status: Final = args[0] if args else 0
|
||||
message: Final = args[1] if len(args) > 1 else ""
|
||||
raise APIError(
|
||||
status_code=int(status) or 500,
|
||||
message=f"litellm rust chat completions: {message}",
|
||||
llm_provider=custom_llm_provider or "",
|
||||
model=model,
|
||||
)
|
||||
if not isinstance(rust_error, declined):
|
||||
raise rust_error
|
||||
verbose_logger.debug(
|
||||
"Rust chat completions declined before calling the provider (%s); using the Python path",
|
||||
rust_error,
|
||||
)
|
||||
|
||||
|
||||
def _build_model_response(
|
||||
rust_response: Mapping[str, object],
|
||||
model_response: ModelResponse,
|
||||
|
|
@ -333,13 +262,33 @@ def _build_model_response(
|
|||
built: Final = convert_to_model_response_object(
|
||||
response_object=dict(rust_response), # mutable-ok: the converter takes a real dict and rewrites it
|
||||
model_response_object=model_response,
|
||||
hidden_params={"additional_headers": {RUST_RESPONSE_HEADER: "true"}}, # mutable-ok: rewritten by the converter
|
||||
hidden_params=execution_hidden_params(None, CoreEngine.RUST), # mutable-ok: rewritten by the converter
|
||||
)
|
||||
if not isinstance(built, ModelResponse):
|
||||
raise TypeError(f"expected a ModelResponse from the rust path, got {type(built).__name__}")
|
||||
return built
|
||||
|
||||
|
||||
def _unwrap_execution(result: ExecutionResult[object]) -> object:
|
||||
value: Final = result.value
|
||||
if isinstance(value, ModelResponse):
|
||||
raw_hidden_params: Final = cast( # cast-ok: runtime validation below narrows this metadata before use
|
||||
object,
|
||||
value._hidden_params, # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params getter
|
||||
)
|
||||
hidden_params: Final = (
|
||||
cast( # cast-ok: the isinstance check validates the mapping boundary
|
||||
Mapping[str, object], raw_hidden_params
|
||||
)
|
||||
if isinstance(raw_hidden_params, Mapping)
|
||||
else None
|
||||
)
|
||||
value._hidden_params = execution_hidden_params( # pyright: ignore[reportPrivateUsage] # ModelResponse exposes no public hidden-params setter
|
||||
hidden_params, result.source
|
||||
)
|
||||
return value
|
||||
|
||||
|
||||
def chat_completions(
|
||||
*,
|
||||
model: str,
|
||||
|
|
@ -354,10 +303,10 @@ def chat_completions(
|
|||
on_response: ResponseObserver,
|
||||
) -> ModelResponse | None:
|
||||
rust_chat_completions: Final = load_rust_chat_completions()
|
||||
if rust_chat_completions is None:
|
||||
return None
|
||||
try:
|
||||
rust_response: Final = rust_chat_completions(
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_chat_completions is None
|
||||
else lambda: rust_chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -367,11 +316,18 @@ def chat_completions(
|
|||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw
|
||||
_reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider)
|
||||
return None
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
)
|
||||
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
result: Final = attempt(
|
||||
native_call=native_call,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(route="chat completions", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
return result.value if isinstance(result, RustHandled) else None
|
||||
|
||||
|
||||
async def achat_completions(
|
||||
|
|
@ -388,10 +344,10 @@ async def achat_completions(
|
|||
on_response: ResponseObserver,
|
||||
) -> ModelResponse | None:
|
||||
rust_achat_completions: Final = load_rust_achat_completions()
|
||||
if rust_achat_completions is None:
|
||||
return None
|
||||
try:
|
||||
rust_response: Final = await rust_achat_completions(
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_achat_completions is None
|
||||
else lambda: rust_achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
|
|
@ -401,11 +357,63 @@ async def achat_completions(
|
|||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw
|
||||
_reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider)
|
||||
return None
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
)
|
||||
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
result: Final = await aattempt(
|
||||
native_call=native_call,
|
||||
adapt=adapt,
|
||||
context=BridgeErrorContext(route="chat completions", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
return result.value if isinstance(result, RustHandled) else None
|
||||
|
||||
|
||||
def chat_completions_or_fallback(
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object],
|
||||
model_response: ModelResponse,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
python_fallback: Callable[[], object],
|
||||
) -> object:
|
||||
rust_chat_completions: Final = load_rust_chat_completions()
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_chat_completions is None
|
||||
else lambda: rust_chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
|
||||
def adapt(rust_response: Mapping[str, object]) -> object:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
return _unwrap_execution(
|
||||
invoke(
|
||||
native_call=native_call,
|
||||
fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=BridgeErrorContext(route="chat completions", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
async def achat_completions_or_fallback(
|
||||
|
|
@ -422,26 +430,32 @@ async def achat_completions_or_fallback(
|
|||
on_response: ResponseObserver,
|
||||
python_fallback: Callable[[], Awaitable[object]],
|
||||
) -> object:
|
||||
"""Await the Rust path, falling back to the caller's own Python path when
|
||||
the bridge is unavailable or the call fails.
|
||||
|
||||
The caller supplies the fallback, so the bridge stays free of provider
|
||||
dispatch. This exists because a caller that dispatches asynchronously has
|
||||
already returned a coroutine by the time a Rust failure surfaces, and so
|
||||
cannot fall back on its own.
|
||||
"""
|
||||
response: Final = await achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout=timeout,
|
||||
on_response=on_response,
|
||||
rust_achat_completions: Final = load_rust_achat_completions()
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_achat_completions is None
|
||||
else lambda: rust_achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
|
||||
def adapt(rust_response: Mapping[str, object]) -> object:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
return _unwrap_execution(
|
||||
await ainvoke(
|
||||
native_call=native_call,
|
||||
fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=BridgeErrorContext(route="chat completions", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
)
|
||||
if response is not None:
|
||||
return response
|
||||
return await python_fallback()
|
||||
|
|
|
|||
|
|
@ -4,6 +4,8 @@ import os
|
|||
import warnings
|
||||
from typing import TYPE_CHECKING, Final
|
||||
|
||||
from litellm.rust_bridge.runtime import UNSET, Unset
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.rust_bridge.messages import RustAmessages, RustMessages
|
||||
from litellm.rust_bridge.ocr import RustAocr, RustOcr
|
||||
|
|
@ -17,13 +19,6 @@ _GLOBAL_ENV_NAME: Final = "LITELLM_RUST"
|
|||
_LEGACY_OCR_ENV_NAME: Final = "LITELLM_USE_RUST_OCR"
|
||||
|
||||
|
||||
class _Unset:
|
||||
pass
|
||||
|
||||
|
||||
_UNSET: Final = _Unset()
|
||||
|
||||
|
||||
class _RustConfiguration:
|
||||
def __init__(self) -> None:
|
||||
self.override: bool | None = None
|
||||
|
|
@ -107,13 +102,13 @@ def reset_rust_configuration() -> None:
|
|||
def use_litellm_rust(
|
||||
enabled: bool = True,
|
||||
*,
|
||||
ocr: RustOcr | None | _Unset = _UNSET,
|
||||
aocr: RustAocr | None | _Unset = _UNSET,
|
||||
messages: RustMessages | None | _Unset = _UNSET,
|
||||
amessages: RustAmessages | None | _Unset = _UNSET,
|
||||
responses_websocket: type[RustResponsesWebSocketConnection] | None | _Unset = _UNSET,
|
||||
transcription: RustTranscription | None | _Unset = _UNSET,
|
||||
atranscription: RustAtranscription | None | _Unset = _UNSET,
|
||||
ocr: RustOcr | None | Unset = UNSET,
|
||||
aocr: RustAocr | None | Unset = UNSET,
|
||||
messages: RustMessages | None | Unset = UNSET,
|
||||
amessages: RustAmessages | None | Unset = UNSET,
|
||||
responses_websocket: type[RustResponsesWebSocketConnection] | None | Unset = UNSET,
|
||||
transcription: RustTranscription | None | Unset = UNSET,
|
||||
atranscription: RustAtranscription | None | Unset = UNSET,
|
||||
) -> None:
|
||||
"""Set the process override for optional Rust paths.
|
||||
|
||||
|
|
@ -121,7 +116,7 @@ def use_litellm_rust(
|
|||
"""
|
||||
_CONFIGURATION.override = enabled
|
||||
bindings: Final = (ocr, aocr, messages, amessages, responses_websocket, transcription, atranscription)
|
||||
if all(isinstance(binding, _Unset) for binding in bindings):
|
||||
if all(isinstance(binding, Unset) for binding in bindings):
|
||||
return
|
||||
warnings.warn(
|
||||
"Injecting Rust bridge implementations through use_litellm_rust() is deprecated; "
|
||||
|
|
@ -130,28 +125,28 @@ def use_litellm_rust(
|
|||
stacklevel=2,
|
||||
)
|
||||
|
||||
if not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset):
|
||||
if not isinstance(ocr, Unset) or not isinstance(aocr, Unset):
|
||||
from litellm.rust_bridge.ocr import set_rust_ocr
|
||||
|
||||
if not isinstance(ocr, _Unset):
|
||||
if not isinstance(ocr, Unset):
|
||||
set_rust_ocr(ocr=ocr)
|
||||
if not isinstance(aocr, _Unset):
|
||||
if not isinstance(aocr, Unset):
|
||||
set_rust_ocr(aocr=aocr)
|
||||
if not isinstance(messages, _Unset) or not isinstance(amessages, _Unset):
|
||||
if not isinstance(messages, Unset) or not isinstance(amessages, Unset):
|
||||
from litellm.rust_bridge.messages import set_rust_messages
|
||||
|
||||
if not isinstance(messages, _Unset):
|
||||
if not isinstance(messages, Unset):
|
||||
set_rust_messages(messages=messages)
|
||||
if not isinstance(amessages, _Unset):
|
||||
if not isinstance(amessages, Unset):
|
||||
set_rust_messages(amessages=amessages)
|
||||
if not isinstance(responses_websocket, _Unset):
|
||||
if not isinstance(responses_websocket, Unset):
|
||||
from litellm.rust_bridge.responses_websocket import set_rust_responses_websocket
|
||||
|
||||
set_rust_responses_websocket(connection=responses_websocket)
|
||||
if not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset):
|
||||
if not isinstance(transcription, Unset) or not isinstance(atranscription, Unset):
|
||||
from litellm.rust_bridge.transcription import configure_rust_transcription
|
||||
|
||||
if not isinstance(transcription, _Unset):
|
||||
if not isinstance(transcription, Unset):
|
||||
configure_rust_transcription(transcription=transcription)
|
||||
if not isinstance(atranscription, _Unset):
|
||||
if not isinstance(atranscription, Unset):
|
||||
configure_rust_transcription(atranscription=atranscription)
|
||||
|
|
|
|||
|
|
@ -3,11 +3,21 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Protocol, cast
|
||||
from typing import Final, Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.rust_bridge.runtime import (
|
||||
UNSET,
|
||||
BridgeErrorContext,
|
||||
FallbackMode,
|
||||
NativeBinding,
|
||||
Unset,
|
||||
ainvoke,
|
||||
async_none,
|
||||
identity,
|
||||
invoke,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
|
||||
|
|
@ -39,53 +49,25 @@ class RustAmessages(Protocol):
|
|||
raise NotImplementedError
|
||||
|
||||
|
||||
class _Unset:
|
||||
pass
|
||||
|
||||
|
||||
_UNSET: Final[_Unset] = _Unset()
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _RustMessagesState:
|
||||
messages: RustMessages | None = None
|
||||
amessages: RustAmessages | None = None
|
||||
|
||||
|
||||
_STATE: Final[_RustMessagesState] = _RustMessagesState()
|
||||
_MESSAGES: Final = NativeBinding[RustMessages]("messages")
|
||||
_AMESSAGES: Final = NativeBinding[RustAmessages]("amessages")
|
||||
|
||||
|
||||
def set_rust_messages(
|
||||
*,
|
||||
messages: RustMessages | None | _Unset = _UNSET,
|
||||
amessages: RustAmessages | None | _Unset = _UNSET,
|
||||
messages: RustMessages | None | Unset = UNSET,
|
||||
amessages: RustAmessages | None | Unset = UNSET,
|
||||
) -> None:
|
||||
if not isinstance(messages, _Unset):
|
||||
_STATE.messages = messages
|
||||
if not isinstance(amessages, _Unset):
|
||||
_STATE.amessages = amessages
|
||||
_MESSAGES.update(messages)
|
||||
_AMESSAGES.update(amessages)
|
||||
|
||||
|
||||
def load_rust_messages() -> RustMessages | None:
|
||||
if _STATE.messages is not None:
|
||||
return _STATE.messages
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
return cast(RustMessages, getattr(native_bridge, "messages", None))
|
||||
return _MESSAGES.load()
|
||||
|
||||
|
||||
def load_rust_amessages() -> RustAmessages | None:
|
||||
if _STATE.amessages is not None:
|
||||
return _STATE.amessages
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
return cast(RustAmessages, getattr(native_bridge, "amessages", None))
|
||||
return _AMESSAGES.load()
|
||||
|
||||
|
||||
def messages(
|
||||
|
|
@ -99,17 +81,27 @@ def messages(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_messages: Final = load_rust_messages()
|
||||
if rust_messages is None:
|
||||
return None
|
||||
return rust_messages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_messages is None
|
||||
else lambda: rust_messages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
result: Final = invoke(
|
||||
native_call=native_call,
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=BridgeErrorContext(route="messages", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
return result.value
|
||||
|
||||
|
||||
async def amessages(
|
||||
|
|
@ -123,14 +115,24 @@ async def amessages(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_amessages: Final = load_rust_amessages()
|
||||
if rust_amessages is None:
|
||||
return None
|
||||
return await rust_amessages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_amessages is None
|
||||
else lambda: rust_amessages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
result: Final = await ainvoke(
|
||||
native_call=native_call,
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=BridgeErrorContext(route="messages", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
return result.value
|
||||
|
|
|
|||
|
|
@ -3,11 +3,22 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable
|
||||
from typing import Final, Protocol, cast
|
||||
from typing import Final, Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.rust_bridge import configuration as _configuration
|
||||
from litellm.rust_bridge.runtime import (
|
||||
UNSET,
|
||||
BridgeErrorContext,
|
||||
FallbackMode,
|
||||
NativeBinding,
|
||||
Unset,
|
||||
ainvoke,
|
||||
async_none,
|
||||
identity,
|
||||
invoke,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
|
||||
|
||||
rust_ocr_enabled = _configuration.rust_ocr_enabled
|
||||
|
|
@ -44,49 +55,25 @@ class RustAocr(Protocol):
|
|||
raise NotImplementedError
|
||||
|
||||
|
||||
class _Unset:
|
||||
pass
|
||||
|
||||
|
||||
_UNSET: Final[_Unset] = _Unset()
|
||||
|
||||
|
||||
_rust_ocr_impl: RustOcr | None = None
|
||||
_rust_aocr_impl: RustAocr | None = None
|
||||
_OCR: Final = NativeBinding[RustOcr]("ocr")
|
||||
_AOCR: Final = NativeBinding[RustAocr]("aocr")
|
||||
|
||||
|
||||
def set_rust_ocr(
|
||||
*,
|
||||
ocr: RustOcr | None | _Unset = _UNSET,
|
||||
aocr: RustAocr | None | _Unset = _UNSET,
|
||||
ocr: RustOcr | None | Unset = UNSET,
|
||||
aocr: RustAocr | None | Unset = UNSET,
|
||||
) -> None:
|
||||
global _rust_ocr_impl, _rust_aocr_impl
|
||||
if not isinstance(ocr, _Unset):
|
||||
_rust_ocr_impl = ocr
|
||||
if not isinstance(aocr, _Unset):
|
||||
_rust_aocr_impl = aocr
|
||||
_OCR.update(ocr)
|
||||
_AOCR.update(aocr)
|
||||
|
||||
|
||||
def load_rust_ocr() -> RustOcr | None:
|
||||
if _rust_ocr_impl is not None:
|
||||
return _rust_ocr_impl
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
return cast(RustOcr, native_bridge.ocr)
|
||||
return _OCR.load()
|
||||
|
||||
|
||||
def load_rust_aocr() -> RustAocr | None:
|
||||
if _rust_aocr_impl is not None:
|
||||
return _rust_aocr_impl
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
return cast(RustAocr, getattr(native_bridge, "aocr", None))
|
||||
return _AOCR.load()
|
||||
|
||||
|
||||
def ocr(
|
||||
|
|
@ -101,18 +88,28 @@ def ocr(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_ocr: Final = load_rust_ocr()
|
||||
if rust_ocr is None:
|
||||
return None
|
||||
return rust_ocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_ocr is None
|
||||
else lambda: rust_ocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
result: Final = invoke(
|
||||
native_call=native_call,
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=BridgeErrorContext(route="ocr", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
return result.value
|
||||
|
||||
|
||||
async def aocr(
|
||||
|
|
@ -127,15 +124,25 @@ async def aocr(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_aocr: Final = load_rust_aocr()
|
||||
if rust_aocr is None:
|
||||
return None
|
||||
return await rust_aocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_aocr is None
|
||||
else lambda: rust_aocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
result: Final = await ainvoke(
|
||||
native_call=native_call,
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=BridgeErrorContext(route="ocr", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
return result.value
|
||||
|
|
|
|||
|
|
@ -2,15 +2,24 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Protocol
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
from websockets.exceptions import ConnectionClosedOK
|
||||
|
||||
from litellm.rust_bridge.loader import get_native_bridge
|
||||
from litellm.rust_bridge.runtime import (
|
||||
UNSET,
|
||||
BridgeErrorContext,
|
||||
CoreEngine,
|
||||
ExecutionResult,
|
||||
FallbackMode,
|
||||
NativeBinding,
|
||||
Unset,
|
||||
acall,
|
||||
ainvoke,
|
||||
async_none,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
_EVENT_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
|
|
@ -36,57 +45,46 @@ class RustResponsesWebSocketConnection(Protocol):
|
|||
) -> RustResponsesWebSocket: ...
|
||||
|
||||
|
||||
class _Unset:
|
||||
pass
|
||||
|
||||
|
||||
_UNSET: Final = _Unset()
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _RustResponsesWebSocketState:
|
||||
connection: type[RustResponsesWebSocketConnection] | None = None
|
||||
|
||||
|
||||
_STATE: Final = _RustResponsesWebSocketState()
|
||||
_CONNECTION: Final = NativeBinding[type[RustResponsesWebSocketConnection]]("ResponsesWebSocketSession")
|
||||
|
||||
|
||||
def set_rust_responses_websocket(
|
||||
*,
|
||||
connection: type[RustResponsesWebSocketConnection] | None | _Unset = _UNSET,
|
||||
connection: type[RustResponsesWebSocketConnection] | None | Unset = UNSET,
|
||||
) -> None:
|
||||
if not isinstance(connection, _Unset):
|
||||
_STATE.connection = connection
|
||||
_CONNECTION.update(connection)
|
||||
|
||||
|
||||
def load_rust_responses_websocket() -> type[RustResponsesWebSocketConnection] | None:
|
||||
if _STATE.connection is not None:
|
||||
return _STATE.connection
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
connection_type: Final[type[RustResponsesWebSocketConnection] | None] = getattr(
|
||||
native_bridge, "ResponsesWebSocketSession", None
|
||||
)
|
||||
return connection_type
|
||||
return _CONNECTION.load()
|
||||
|
||||
|
||||
class _ConnectionAdapter:
|
||||
def __init__(self, connection: RustResponsesWebSocket):
|
||||
self._connection: Final = connection
|
||||
self.core_engine: Final = CoreEngine.RUST
|
||||
|
||||
async def send(self, text: str) -> None:
|
||||
event: Final = _EVENT_ADAPTER.validate_json(text)
|
||||
await self._connection.send_event(event)
|
||||
await acall(
|
||||
lambda: self._connection.send_event(event),
|
||||
BridgeErrorContext(route="responses websocket", provider="openai", model=""),
|
||||
)
|
||||
|
||||
async def recv(self) -> str:
|
||||
event: Final = await self._connection.recv_event()
|
||||
event: Final = await acall(
|
||||
self._connection.recv_event,
|
||||
BridgeErrorContext(route="responses websocket", provider="openai", model=""),
|
||||
)
|
||||
if event is None:
|
||||
raise ConnectionClosedOK(None, None)
|
||||
return json.dumps(dict(event), separators=(",", ":")) # mutable-ok: JSON requires a concrete dict
|
||||
|
||||
async def close(self) -> None:
|
||||
await self._connection.close()
|
||||
await acall(
|
||||
self._connection.close,
|
||||
BridgeErrorContext(route="responses websocket", provider="openai", model=""),
|
||||
)
|
||||
|
||||
|
||||
async def connect(
|
||||
|
|
@ -96,23 +94,28 @@ async def connect(
|
|||
api_base: str | None,
|
||||
headers: Mapping[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> _ConnectionAdapter | None:
|
||||
) -> ExecutionResult[_ConnectionAdapter | None]:
|
||||
connection_type: Final = load_rust_responses_websocket()
|
||||
if connection_type is None:
|
||||
return None
|
||||
credentials: Final = None if api_key is None else MappingProxyType({"api_key": api_key})
|
||||
try:
|
||||
connection: Final = await connection_type.connect(
|
||||
credentials: Final = (
|
||||
None
|
||||
if api_key is None
|
||||
else {"api_key": api_key} # mutable-ok: native bridge serialization requires a plain dict
|
||||
)
|
||||
native_call: Final = (
|
||||
None
|
||||
if connection_type is None
|
||||
else lambda: connection_type.connect(
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
headers,
|
||||
timeout_to_seconds(timeout),
|
||||
)
|
||||
except Exception as error:
|
||||
native: Final = get_native_bridge()
|
||||
declined: Final = None if native is None else getattr(native, "RustBridgeDeclined", None)
|
||||
if isinstance(declined, type) and isinstance(error, declined):
|
||||
return None
|
||||
raise
|
||||
return _ConnectionAdapter(connection)
|
||||
)
|
||||
return await ainvoke(
|
||||
native_call=native_call,
|
||||
fallback=async_none,
|
||||
adapt=_ConnectionAdapter,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=BridgeErrorContext(route="responses websocket", provider=provider, model=""),
|
||||
)
|
||||
|
|
|
|||
289
litellm/rust_bridge/runtime.py
Normal file
289
litellm/rust_bridge/runtime.py
Normal file
|
|
@ -0,0 +1,289 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Final, Generic, NoReturn, TypeAlias, TypeVar, cast
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge.loader import get_native_bridge
|
||||
|
||||
NativeT = TypeVar("NativeT")
|
||||
ResultT = TypeVar("ResultT")
|
||||
BindingT = TypeVar("BindingT")
|
||||
|
||||
CORE_ENGINE_HIDDEN_PARAM: Final = "core_engine"
|
||||
CORE_ENGINE_HEADER: Final = "x-litellm-core"
|
||||
LEGACY_RUST_HEADER: Final = "x-litellm-rust"
|
||||
|
||||
|
||||
class Unset:
|
||||
pass
|
||||
|
||||
|
||||
UNSET: Final = Unset()
|
||||
|
||||
|
||||
class FallbackMode(Enum):
|
||||
PYTHON = "python"
|
||||
RUST_REQUIRED = "rust_required"
|
||||
|
||||
|
||||
class CoreEngine(str, Enum):
|
||||
PYTHON = "python"
|
||||
RUST = "rust"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ExecutionResult(Generic[ResultT]):
|
||||
value: ResultT
|
||||
source: CoreEngine
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RustHandled(Generic[ResultT]):
|
||||
value: ResultT
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RustDeclined:
|
||||
reason: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class RustUnavailable:
|
||||
pass
|
||||
|
||||
|
||||
RustAttempt: TypeAlias = RustHandled[ResultT] | RustDeclined | RustUnavailable
|
||||
|
||||
|
||||
def execution_headers(source: CoreEngine) -> dict[str, str]:
|
||||
if source is CoreEngine.RUST:
|
||||
return { # mutable-ok: response adapters require a mutable header dict
|
||||
CORE_ENGINE_HEADER: source.value,
|
||||
LEGACY_RUST_HEADER: "true",
|
||||
}
|
||||
return {CORE_ENGINE_HEADER: source.value} # mutable-ok: response adapters require a mutable header dict
|
||||
|
||||
|
||||
def execution_additional_headers(
|
||||
additional_headers: Mapping[str, object] | None,
|
||||
source: CoreEngine,
|
||||
) -> dict[str, str]:
|
||||
existing: Final = additional_headers or {} # mutable-ok: empty default is local and never mutated
|
||||
reserved: Final = frozenset({CORE_ENGINE_HEADER, LEGACY_RUST_HEADER})
|
||||
preserved: Final = { # mutable-ok: response adapters require a mutable header dict
|
||||
str(name): str(value) for name, value in existing.items() if str(name).lower() not in reserved
|
||||
}
|
||||
return { # mutable-ok: response adapters require a mutable header dict
|
||||
**preserved,
|
||||
**execution_headers(source),
|
||||
}
|
||||
|
||||
|
||||
def execution_hidden_params(
|
||||
hidden_params: Mapping[str, object] | None,
|
||||
source: CoreEngine,
|
||||
) -> dict[str, object]:
|
||||
existing: Final = hidden_params or {} # mutable-ok: empty default is local and never mutated
|
||||
raw_headers: Final = existing.get("additional_headers")
|
||||
additional_headers: Final[Mapping[str, object]] = (
|
||||
cast( # cast-ok: the isinstance check validates the mapping boundary
|
||||
Mapping[str, object], raw_headers
|
||||
)
|
||||
if isinstance(raw_headers, Mapping)
|
||||
else {} # mutable-ok: empty default is local and never mutated
|
||||
)
|
||||
return { # mutable-ok: response objects require mutable hidden params
|
||||
**existing,
|
||||
CORE_ENGINE_HIDDEN_PARAM: source.value,
|
||||
"additional_headers": execution_additional_headers(additional_headers, source),
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BridgeErrorContext:
|
||||
route: str
|
||||
provider: str
|
||||
model: str
|
||||
|
||||
|
||||
class NativeBinding(Generic[BindingT]):
|
||||
def __init__(self, attribute: str) -> None:
|
||||
self._attribute: Final = attribute
|
||||
self._override: BindingT | Unset = UNSET
|
||||
|
||||
def load(self) -> BindingT | None:
|
||||
if not isinstance(self._override, Unset):
|
||||
return self._override
|
||||
native: Final = get_native_bridge()
|
||||
return (
|
||||
None
|
||||
if native is None
|
||||
else cast( # cast-ok: each route validates the callable against its binding protocol
|
||||
BindingT | None, getattr(native, self._attribute, None)
|
||||
)
|
||||
)
|
||||
|
||||
def set(self, value: BindingT | None) -> None:
|
||||
self._override = UNSET if value is None else value
|
||||
|
||||
def update(self, value: BindingT | None | Unset) -> None:
|
||||
if not isinstance(value, Unset):
|
||||
self.set(value)
|
||||
|
||||
|
||||
def invoke(
|
||||
*,
|
||||
native_call: Callable[[], NativeT] | None,
|
||||
fallback: Callable[[], ResultT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
mode: FallbackMode,
|
||||
context: BridgeErrorContext,
|
||||
) -> ExecutionResult[ResultT]:
|
||||
attempt_result: Final = attempt(native_call=native_call, adapt=adapt, context=context)
|
||||
if isinstance(attempt_result, RustHandled):
|
||||
return ExecutionResult(value=attempt_result.value, source=CoreEngine.RUST)
|
||||
if mode is FallbackMode.PYTHON:
|
||||
return ExecutionResult(value=fallback(), source=CoreEngine.PYTHON)
|
||||
_raise_required(attempt_result, context)
|
||||
|
||||
|
||||
async def ainvoke(
|
||||
*,
|
||||
native_call: Callable[[], Awaitable[NativeT]] | None,
|
||||
fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
mode: FallbackMode,
|
||||
context: BridgeErrorContext,
|
||||
) -> ExecutionResult[ResultT]:
|
||||
attempt_result: Final = await aattempt(native_call=native_call, adapt=adapt, context=context)
|
||||
if isinstance(attempt_result, RustHandled):
|
||||
return ExecutionResult(value=attempt_result.value, source=CoreEngine.RUST)
|
||||
if mode is FallbackMode.PYTHON:
|
||||
return ExecutionResult(value=await fallback(), source=CoreEngine.PYTHON)
|
||||
_raise_required(attempt_result, context)
|
||||
|
||||
|
||||
def attempt(
|
||||
*,
|
||||
native_call: Callable[[], NativeT] | None,
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
context: BridgeErrorContext,
|
||||
) -> RustAttempt[ResultT]:
|
||||
if native_call is None:
|
||||
return RustUnavailable()
|
||||
exceptions: Final = _native_exceptions()
|
||||
if exceptions is None:
|
||||
return RustHandled(adapt(native_call()))
|
||||
declined, upstream = exceptions
|
||||
try:
|
||||
native_result: Final = native_call()
|
||||
except declined as error:
|
||||
return RustDeclined(reason=_decline_reason(error))
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
return RustHandled(adapt(native_result))
|
||||
|
||||
|
||||
async def aattempt(
|
||||
*,
|
||||
native_call: Callable[[], Awaitable[NativeT]] | None,
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
context: BridgeErrorContext,
|
||||
) -> RustAttempt[ResultT]:
|
||||
if native_call is None:
|
||||
return RustUnavailable()
|
||||
exceptions: Final = _native_exceptions()
|
||||
if exceptions is None:
|
||||
return RustHandled(adapt(await native_call()))
|
||||
declined, upstream = exceptions
|
||||
try:
|
||||
native_result: Final = await native_call()
|
||||
except declined as error:
|
||||
return RustDeclined(reason=_decline_reason(error))
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
return RustHandled(adapt(native_result))
|
||||
|
||||
|
||||
def _native_exceptions() -> tuple[type[BaseException], type[BaseException]] | None:
|
||||
native: Final = get_native_bridge()
|
||||
if native is None:
|
||||
return None
|
||||
declined: Final = getattr(native, "RustBridgeDeclined", None)
|
||||
upstream: Final = getattr(native, "RustUpstreamError", None)
|
||||
if not isinstance(declined, type) or not isinstance(upstream, type):
|
||||
return None
|
||||
return declined, upstream
|
||||
|
||||
|
||||
def _decline_reason(error: BaseException) -> str:
|
||||
reason: Final = error.args[0] if error.args else str(error)
|
||||
return reason if isinstance(reason, str) else str(reason)
|
||||
|
||||
|
||||
def _raise_required(
|
||||
attempt_result: RustDeclined | RustUnavailable,
|
||||
context: BridgeErrorContext,
|
||||
) -> NoReturn:
|
||||
raise RuntimeError(f"Rust {context.route} bridge {_required_reason(attempt_result)}")
|
||||
|
||||
|
||||
def _required_reason(attempt_result: RustDeclined | RustUnavailable) -> str:
|
||||
match attempt_result:
|
||||
case RustUnavailable():
|
||||
return "is unavailable"
|
||||
case RustDeclined(reason=reason):
|
||||
return f"declined the request: {reason}"
|
||||
|
||||
|
||||
def _raise_upstream(error: BaseException, context: BridgeErrorContext) -> NoReturn:
|
||||
args: Final = cast( # cast-ok: BaseException.args is always a tuple at runtime
|
||||
tuple[object, ...], error.args
|
||||
)
|
||||
status_value: Final = args[0] if args else 0
|
||||
message_value: Final = args[1] if len(args) > 1 else str(error)
|
||||
status: Final = status_value if isinstance(status_value, int) else 0
|
||||
message: Final = message_value if isinstance(message_value, str) else str(message_value)
|
||||
api_error: Final = APIError(
|
||||
status_code=status or 500,
|
||||
message=f"litellm rust {context.route}: {message}",
|
||||
llm_provider=context.provider,
|
||||
model=context.model,
|
||||
)
|
||||
api_error.headers = execution_headers( # pyright: ignore[reportAttributeAccessIssue] # proxy reads exception headers
|
||||
CoreEngine.RUST
|
||||
)
|
||||
raise api_error from error
|
||||
|
||||
|
||||
def call(operation: Callable[[], ResultT], context: BridgeErrorContext) -> ResultT:
|
||||
exceptions: Final = _native_exceptions()
|
||||
if exceptions is None:
|
||||
return operation()
|
||||
upstream: Final = exceptions[1]
|
||||
try:
|
||||
return operation()
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
|
||||
|
||||
async def acall(operation: Callable[[], Awaitable[ResultT]], context: BridgeErrorContext) -> ResultT:
|
||||
exceptions: Final = _native_exceptions()
|
||||
if exceptions is None:
|
||||
return await operation()
|
||||
upstream: Final = exceptions[1]
|
||||
try:
|
||||
return await operation()
|
||||
except upstream as error:
|
||||
_raise_upstream(error, context)
|
||||
|
||||
|
||||
def identity(value: ResultT) -> ResultT:
|
||||
return value
|
||||
|
||||
|
||||
async def async_none() -> None:
|
||||
return None
|
||||
|
|
@ -1,16 +1,27 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Generator, Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from collections.abc import AsyncIterator, Iterator, Mapping
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal, NoReturn, Protocol, TypeAlias, runtime_checkable
|
||||
from typing import TYPE_CHECKING, Final, Literal, Protocol, TypeAlias, runtime_checkable
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge.loader import get_native_bridge
|
||||
from litellm.rust_bridge.runtime import (
|
||||
UNSET,
|
||||
BridgeErrorContext,
|
||||
CoreEngine,
|
||||
ExecutionResult,
|
||||
FallbackMode,
|
||||
NativeBinding,
|
||||
Unset,
|
||||
acall,
|
||||
ainvoke,
|
||||
async_none,
|
||||
call,
|
||||
invoke,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -59,171 +70,49 @@ class RustAsyncStreamOpen(Protocol):
|
|||
) -> RustEventStream: ...
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ObjectAwaitable(Protocol):
|
||||
def __await__(self) -> Generator[object, None, object]: ...
|
||||
|
||||
|
||||
class _Unset:
|
||||
pass
|
||||
|
||||
|
||||
_UNSET: Final = _Unset()
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _RustStreamingState:
|
||||
chat: RustStreamOpen | None = None
|
||||
achat: RustAsyncStreamOpen | None = None
|
||||
messages: RustStreamOpen | None = None
|
||||
amessages: RustAsyncStreamOpen | None = None
|
||||
responses: RustStreamOpen | None = None
|
||||
aresponses: RustAsyncStreamOpen | None = None
|
||||
|
||||
|
||||
_STATE: Final = _RustStreamingState()
|
||||
_CHAT: Final = NativeBinding[RustStreamOpen]("chat_completions_stream")
|
||||
_ACHAT: Final = NativeBinding[RustAsyncStreamOpen]("achat_completions_stream")
|
||||
_MESSAGES: Final = NativeBinding[RustStreamOpen]("messages_stream")
|
||||
_AMESSAGES: Final = NativeBinding[RustAsyncStreamOpen]("amessages_stream")
|
||||
_RESPONSES: Final = NativeBinding[RustStreamOpen]("responses_stream")
|
||||
_ARESPONSES: Final = NativeBinding[RustAsyncStreamOpen]("aresponses_stream")
|
||||
|
||||
|
||||
def set_rust_streaming(
|
||||
*,
|
||||
chat: RustStreamOpen | None | _Unset = _UNSET,
|
||||
achat: RustAsyncStreamOpen | None | _Unset = _UNSET,
|
||||
messages: RustStreamOpen | None | _Unset = _UNSET,
|
||||
amessages: RustAsyncStreamOpen | None | _Unset = _UNSET,
|
||||
responses: RustStreamOpen | None | _Unset = _UNSET,
|
||||
aresponses: RustAsyncStreamOpen | None | _Unset = _UNSET,
|
||||
chat: RustStreamOpen | None | Unset = UNSET,
|
||||
achat: RustAsyncStreamOpen | None | Unset = UNSET,
|
||||
messages: RustStreamOpen | None | Unset = UNSET,
|
||||
amessages: RustAsyncStreamOpen | None | Unset = UNSET,
|
||||
responses: RustStreamOpen | None | Unset = UNSET,
|
||||
aresponses: RustAsyncStreamOpen | None | Unset = UNSET,
|
||||
) -> None:
|
||||
if not isinstance(chat, _Unset):
|
||||
_STATE.chat = chat
|
||||
if not isinstance(achat, _Unset):
|
||||
_STATE.achat = achat
|
||||
if not isinstance(messages, _Unset):
|
||||
_STATE.messages = messages
|
||||
if not isinstance(amessages, _Unset):
|
||||
_STATE.amessages = amessages
|
||||
if not isinstance(responses, _Unset):
|
||||
_STATE.responses = responses
|
||||
if not isinstance(aresponses, _Unset):
|
||||
_STATE.aresponses = aresponses
|
||||
|
||||
|
||||
def _native_attribute(name: str) -> object | None:
|
||||
native: Final = get_native_bridge()
|
||||
return None if native is None else getattr(native, name, None)
|
||||
|
||||
|
||||
def _native_sync_opener(name: str) -> RustStreamOpen | None:
|
||||
opener: Final = _native_attribute(name)
|
||||
if not callable(opener):
|
||||
return None
|
||||
|
||||
def open_stream(
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> RustEventStream:
|
||||
stream: Final = opener(
|
||||
request,
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
)
|
||||
if not isinstance(stream, RustEventStream):
|
||||
raise TypeError("native stream opener returned an invalid stream")
|
||||
return stream
|
||||
|
||||
return open_stream
|
||||
|
||||
|
||||
def _native_async_opener(name: str) -> RustAsyncStreamOpen | None:
|
||||
opener: Final = _native_attribute(name)
|
||||
if not callable(opener):
|
||||
return None
|
||||
|
||||
async def open_stream(
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> RustEventStream:
|
||||
pending: Final = opener(
|
||||
request,
|
||||
provider,
|
||||
credentials,
|
||||
api_base,
|
||||
extra_headers,
|
||||
timeout_seconds,
|
||||
)
|
||||
if not isinstance(pending, ObjectAwaitable):
|
||||
raise TypeError("native async stream opener returned a non-awaitable")
|
||||
stream: Final[object] = await pending
|
||||
if not isinstance(stream, RustEventStream):
|
||||
raise TypeError("native async stream opener returned an invalid stream")
|
||||
return stream
|
||||
|
||||
return open_stream
|
||||
_CHAT.update(chat)
|
||||
_ACHAT.update(achat)
|
||||
_MESSAGES.update(messages)
|
||||
_AMESSAGES.update(amessages)
|
||||
_RESPONSES.update(responses)
|
||||
_ARESPONSES.update(aresponses)
|
||||
|
||||
|
||||
def _sync_opener(api: StreamApi) -> RustStreamOpen | None:
|
||||
match api:
|
||||
case "chat_completions":
|
||||
return _STATE.chat or _native_sync_opener("chat_completions_stream")
|
||||
return _CHAT.load()
|
||||
case "messages":
|
||||
return _STATE.messages or _native_sync_opener("messages_stream")
|
||||
return _MESSAGES.load()
|
||||
case "responses":
|
||||
return _STATE.responses or _native_sync_opener("responses_stream")
|
||||
return _RESPONSES.load()
|
||||
|
||||
|
||||
def _async_opener(api: StreamApi) -> RustAsyncStreamOpen | None:
|
||||
match api:
|
||||
case "chat_completions":
|
||||
return _STATE.achat or _native_async_opener("achat_completions_stream")
|
||||
return _ACHAT.load()
|
||||
case "messages":
|
||||
return _STATE.amessages or _native_async_opener("amessages_stream")
|
||||
return _AMESSAGES.load()
|
||||
case "responses":
|
||||
return _STATE.aresponses or _native_async_opener("aresponses_stream")
|
||||
|
||||
|
||||
def _native_exceptions() -> tuple[type[BaseException], type[BaseException]] | None:
|
||||
native: Final = get_native_bridge()
|
||||
if native is None:
|
||||
return None
|
||||
declined: Final = getattr(native, "RustBridgeDeclined", None)
|
||||
upstream: Final = getattr(native, "RustUpstreamError", None)
|
||||
if not isinstance(declined, type) or not isinstance(upstream, type):
|
||||
return None
|
||||
return declined, upstream
|
||||
|
||||
|
||||
def _handle_open_error(error: Exception, provider: str) -> None:
|
||||
exceptions: Final = _native_exceptions()
|
||||
if exceptions is not None and isinstance(error, exceptions[0]):
|
||||
return
|
||||
_raise_stream_error(error, provider)
|
||||
|
||||
|
||||
def _raise_stream_error(error: Exception, provider: str) -> NoReturn:
|
||||
exceptions: Final = _native_exceptions()
|
||||
if exceptions is None or not isinstance(error, exceptions[1]):
|
||||
raise error
|
||||
args: Final[tuple[object, ...]] = error.args
|
||||
status_value: Final = args[0] if args else 0
|
||||
message_value: Final = args[1] if len(args) > 1 else str(error)
|
||||
status: Final = status_value if isinstance(status_value, int) else 0
|
||||
message: Final = message_value if isinstance(message_value, str) else str(message_value)
|
||||
raise APIError(
|
||||
status_code=status or 500,
|
||||
message=f"litellm rust typed stream: {message}",
|
||||
llm_provider=provider,
|
||||
model="",
|
||||
) from error
|
||||
return _ARESPONSES.load()
|
||||
|
||||
|
||||
class TypedEventStreamAdapter:
|
||||
|
|
@ -231,6 +120,7 @@ class TypedEventStreamAdapter:
|
|||
self._stream: Final = stream
|
||||
self._provider: Final = provider
|
||||
self.metadata: Final = stream.metadata
|
||||
self.core_engine: Final = CoreEngine.RUST
|
||||
self._mode: Literal["sync", "async"] | None = None
|
||||
|
||||
def _claim(self, mode: Literal["sync", "async"]) -> None:
|
||||
|
|
@ -252,10 +142,10 @@ class TypedEventStreamAdapter:
|
|||
return event
|
||||
|
||||
def _next_event(self) -> Event | None:
|
||||
try:
|
||||
return self._stream.next_event()
|
||||
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
|
||||
_raise_stream_error(error, self._provider)
|
||||
return call(
|
||||
self._stream.next_event,
|
||||
BridgeErrorContext(route="typed stream", provider=self._provider, model=""),
|
||||
)
|
||||
|
||||
def __aiter__(self) -> AsyncIterator[Event]:
|
||||
self._claim("async")
|
||||
|
|
@ -263,10 +153,10 @@ class TypedEventStreamAdapter:
|
|||
|
||||
async def __anext__(self) -> Event:
|
||||
self._claim("async")
|
||||
try:
|
||||
event: Final = await self._stream.anext_event()
|
||||
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
|
||||
_raise_stream_error(error, self._provider)
|
||||
event: Final = await acall(
|
||||
self._stream.anext_event,
|
||||
BridgeErrorContext(route="typed stream", provider=self._provider, model=""),
|
||||
)
|
||||
if event is None:
|
||||
raise StopAsyncIteration
|
||||
return event
|
||||
|
|
@ -384,12 +274,12 @@ def open_stream(
|
|||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> TypedEventStreamAdapter | None:
|
||||
) -> ExecutionResult[TypedEventStreamAdapter | None]:
|
||||
opener: Final = _sync_opener(api)
|
||||
if opener is None:
|
||||
return None
|
||||
try:
|
||||
stream: Final = opener(
|
||||
native_call: Final = (
|
||||
None
|
||||
if opener is None
|
||||
else lambda: opener(
|
||||
request,
|
||||
provider,
|
||||
credentials,
|
||||
|
|
@ -397,10 +287,14 @@ def open_stream(
|
|||
extra_headers,
|
||||
timeout_to_seconds(timeout),
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
|
||||
_handle_open_error(error, provider)
|
||||
return None
|
||||
return TypedEventStreamAdapter(stream, provider)
|
||||
)
|
||||
return invoke(
|
||||
native_call=native_call,
|
||||
fallback=lambda: None,
|
||||
adapt=lambda stream: TypedEventStreamAdapter(stream, provider),
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=BridgeErrorContext(route=f"{api} stream", provider=provider, model=""),
|
||||
)
|
||||
|
||||
|
||||
async def aopen_stream(
|
||||
|
|
@ -412,12 +306,12 @@ async def aopen_stream(
|
|||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> TypedEventStreamAdapter | None:
|
||||
) -> ExecutionResult[TypedEventStreamAdapter | None]:
|
||||
opener: Final = _async_opener(api)
|
||||
if opener is None:
|
||||
return None
|
||||
try:
|
||||
stream: Final = await opener(
|
||||
native_call: Final = (
|
||||
None
|
||||
if opener is None
|
||||
else lambda: opener(
|
||||
request,
|
||||
provider,
|
||||
credentials,
|
||||
|
|
@ -425,7 +319,11 @@ async def aopen_stream(
|
|||
extra_headers,
|
||||
timeout_to_seconds(timeout),
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 # native stream errors cross a dynamic extension boundary
|
||||
_handle_open_error(error, provider)
|
||||
return None
|
||||
return TypedEventStreamAdapter(stream, provider)
|
||||
)
|
||||
return await ainvoke(
|
||||
native_call=native_call,
|
||||
fallback=async_none,
|
||||
adapt=lambda stream: TypedEventStreamAdapter(stream, provider),
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=BridgeErrorContext(route=f"{api} stream", provider=provider, model=""),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,11 +1,21 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Protocol, cast
|
||||
from typing import Final, Protocol
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.rust_bridge.runtime import (
|
||||
UNSET,
|
||||
BridgeErrorContext,
|
||||
FallbackMode,
|
||||
NativeBinding,
|
||||
Unset,
|
||||
ainvoke,
|
||||
async_none,
|
||||
identity,
|
||||
invoke,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
|
||||
|
|
@ -39,62 +49,27 @@ class RustAtranscription(Protocol):
|
|||
raise NotImplementedError
|
||||
|
||||
|
||||
class _Unset:
|
||||
pass
|
||||
|
||||
|
||||
_UNSET: Final[_Unset] = _Unset()
|
||||
|
||||
|
||||
@dataclass
|
||||
class _RustTranscriptionState:
|
||||
transcription: RustTranscription | None = None
|
||||
atranscription: RustAtranscription | None = None
|
||||
|
||||
|
||||
_STATE: Final = _RustTranscriptionState()
|
||||
_TRANSCRIPTION: Final = NativeBinding[RustTranscription]("transcription")
|
||||
_ATRANSCRIPTION: Final = NativeBinding[RustAtranscription]("atranscription")
|
||||
|
||||
|
||||
def configure_rust_transcription(
|
||||
enabled: bool = True,
|
||||
*,
|
||||
transcription: RustTranscription | None | _Unset = _UNSET,
|
||||
atranscription: RustAtranscription | None | _Unset = _UNSET,
|
||||
transcription: RustTranscription | None | Unset = UNSET,
|
||||
atranscription: RustAtranscription | None | Unset = UNSET,
|
||||
) -> None:
|
||||
if not isinstance(transcription, _Unset):
|
||||
_STATE.transcription = transcription
|
||||
if not isinstance(atranscription, _Unset):
|
||||
_STATE.atranscription = atranscription
|
||||
_ = enabled
|
||||
_TRANSCRIPTION.update(transcription)
|
||||
_ATRANSCRIPTION.update(atranscription)
|
||||
|
||||
|
||||
def load_rust_transcription() -> RustTranscription | None:
|
||||
if _STATE.transcription is not None:
|
||||
return _STATE.transcription
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
native_bridge: Final = get_native_bridge()
|
||||
return (
|
||||
None
|
||||
if native_bridge is None
|
||||
else cast( # cast-ok: native extension protocol is runtime-defined
|
||||
RustTranscription, getattr(native_bridge, "transcription", None)
|
||||
)
|
||||
)
|
||||
return _TRANSCRIPTION.load()
|
||||
|
||||
|
||||
def load_rust_atranscription() -> RustAtranscription | None:
|
||||
if _STATE.atranscription is not None:
|
||||
return _STATE.atranscription
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
native_bridge: Final = get_native_bridge()
|
||||
return (
|
||||
None
|
||||
if native_bridge is None
|
||||
else cast( # cast-ok: native extension protocol is runtime-defined
|
||||
RustAtranscription, getattr(native_bridge, "atranscription", None)
|
||||
)
|
||||
)
|
||||
return _ATRANSCRIPTION.load()
|
||||
|
||||
|
||||
def transcription(
|
||||
|
|
@ -109,18 +84,28 @@ def transcription(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_transcription: Final = load_rust_transcription()
|
||||
if rust_transcription is None:
|
||||
return None
|
||||
return rust_transcription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_transcription is None
|
||||
else lambda: rust_transcription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
result: Final = invoke(
|
||||
native_call=native_call,
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
mode=FallbackMode.RUST_REQUIRED,
|
||||
context=BridgeErrorContext(route="audio transcription", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
return result.value
|
||||
|
||||
|
||||
async def atranscription(
|
||||
|
|
@ -135,15 +120,25 @@ async def atranscription(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
rust_atranscription: Final = load_rust_atranscription()
|
||||
if rust_atranscription is None:
|
||||
return None
|
||||
return await rust_atranscription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
native_call: Final = (
|
||||
None
|
||||
if rust_atranscription is None
|
||||
else lambda: rust_atranscription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
)
|
||||
)
|
||||
result: Final = await ainvoke(
|
||||
native_call=native_call,
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
mode=FallbackMode.RUST_REQUIRED,
|
||||
context=BridgeErrorContext(route="audio transcription", provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
return result.value
|
||||
|
|
|
|||
|
|
@ -1,16 +1,22 @@
|
|||
"""Tests for the optional Rust-backed Anthropic Messages path."""
|
||||
|
||||
import importlib
|
||||
from types import ModuleType
|
||||
from typing import cast
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
anthropic_messages_stream_hidden_params,
|
||||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import (
|
||||
BaseLLMHTTPHandler,
|
||||
_anthropic_messages_with_core_engine,
|
||||
)
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.messages import RustAmessages, RustMessages
|
||||
from litellm.rust_bridge.runtime import UNSET, CoreEngine, Unset
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
|
|
@ -18,6 +24,7 @@ from litellm.types.router import GenericLiteLLMParams
|
|||
|
||||
rust_messages = importlib.import_module("litellm.rust_bridge.messages")
|
||||
rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader")
|
||||
rust_bridge_runtime = importlib.import_module("litellm.rust_bridge.runtime")
|
||||
|
||||
FAKE_MESSAGES_RESPONSE: dict[str, object] = {
|
||||
"id": "msg_123",
|
||||
|
|
@ -36,6 +43,39 @@ REQUEST_BODY: dict[str, object] = {
|
|||
}
|
||||
|
||||
|
||||
def _use_test_rust(
|
||||
enabled: bool = True,
|
||||
*,
|
||||
messages: RustMessages | None | Unset = UNSET,
|
||||
amessages: RustAmessages | None | Unset = UNSET,
|
||||
) -> None:
|
||||
rust_messages.set_rust_messages(messages=messages, amessages=amessages)
|
||||
litellm.use_litellm_rust(enabled)
|
||||
|
||||
|
||||
def test_python_messages_response_has_symmetric_provenance():
|
||||
response = _anthropic_messages_with_core_engine(
|
||||
cast(AnthropicMessagesResponse, dict(FAKE_MESSAGES_RESPONSE)),
|
||||
CoreEngine.PYTHON,
|
||||
)
|
||||
|
||||
assert response["_hidden_params"]["core_engine"] == "python"
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-core": "python"}
|
||||
|
||||
|
||||
def test_python_messages_stream_provenance_precedes_iteration_and_cannot_be_spoofed():
|
||||
hidden_params = anthropic_messages_stream_hidden_params(
|
||||
httpx.Headers({"x-litellm-core": "rust", "x-litellm-rust": "true", "x-provider": "kept"})
|
||||
)
|
||||
|
||||
assert hidden_params["core_engine"] == "python"
|
||||
assert hidden_params["additional_headers"]["x-litellm-core"] == "python"
|
||||
assert "x-litellm-rust" not in hidden_params["additional_headers"]
|
||||
assert hidden_params["additional_headers"]["llm_provider-x-litellm-core"] == "rust"
|
||||
assert hidden_params["additional_headers"]["llm_provider-x-litellm-rust"] == "true"
|
||||
assert hidden_params["additional_headers"]["llm_provider-x-provider"] == "kept"
|
||||
|
||||
|
||||
class RecordingMessages:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
|
|
@ -102,30 +142,37 @@ class ExplodingAsyncMessages:
|
|||
|
||||
|
||||
class RaisingAsyncMessages:
|
||||
def __init__(self, error: Exception) -> None:
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, **kwargs: object) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise RuntimeError("upstream request failed with status 400: bad request")
|
||||
|
||||
|
||||
class _FakeDeclined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeUpstream(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _FakeNativeErrors:
|
||||
RustBridgeDeclined = _FakeDeclined
|
||||
RustUpstreamError = _FakeUpstream
|
||||
|
||||
|
||||
class RaisingTypedAsyncMessages:
|
||||
def __init__(self, error: Exception) -> None:
|
||||
self.error = error
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, **kwargs: object) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise self.error
|
||||
|
||||
|
||||
class FakeBridgeDeclined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class FakeUpstreamError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _install_fake_bridge_exceptions(monkeypatch) -> None:
|
||||
native_bridge = ModuleType("_native")
|
||||
native_bridge.RustBridgeDeclined = FakeBridgeDeclined
|
||||
native_bridge.RustUpstreamError = FakeUpstreamError
|
||||
monkeypatch.setattr(rust_bridge_loader, "_cached_bridge", native_bridge)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_rust_flag():
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None)
|
||||
|
|
@ -139,7 +186,7 @@ def _reset_rust_flag():
|
|||
|
||||
def test_load_rust_messages_returns_injected_impl():
|
||||
bridge = RecordingMessages()
|
||||
litellm.use_litellm_rust(True, messages=bridge)
|
||||
_use_test_rust(True, messages=bridge)
|
||||
assert rust_messages.load_rust_messages() is bridge
|
||||
|
||||
|
||||
|
|
@ -155,16 +202,12 @@ def test_bare_use_litellm_rust_still_toggles_ocr():
|
|||
|
||||
def test_load_rust_amessages_returns_injected_impl():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
assert rust_messages.load_rust_amessages() is bridge
|
||||
|
||||
|
||||
def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge"),
|
||||
"get_native_bridge",
|
||||
lambda: None,
|
||||
)
|
||||
monkeypatch.setattr(rust_bridge_runtime, "get_native_bridge", lambda: None)
|
||||
litellm.use_litellm_rust(True)
|
||||
assert rust_messages.load_rust_messages() is None
|
||||
result = rust_messages.messages(
|
||||
|
|
@ -181,7 +224,7 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
|
|||
|
||||
def test_messages_wrapper_forwards_args_and_converts_timeout():
|
||||
bridge = RecordingMessages()
|
||||
litellm.use_litellm_rust(True, messages=bridge)
|
||||
_use_test_rust(True, messages=bridge)
|
||||
|
||||
response = rust_messages.messages(
|
||||
model="claude-sonnet-4-5",
|
||||
|
|
@ -208,7 +251,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout():
|
|||
@pytest.mark.asyncio
|
||||
async def test_amessages_wrapper_forwards_args():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
response = await rust_messages.amessages(
|
||||
model="claude-sonnet-4-5",
|
||||
|
|
@ -244,13 +287,17 @@ def _gate(**overrides):
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_invokes_rust_and_marks_response_header():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
response = await _gate()
|
||||
|
||||
assert response is not None
|
||||
assert response["id"] == "msg_123"
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert response["_hidden_params"]["core_engine"] == "rust"
|
||||
assert response["_hidden_params"]["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
call = bridge.calls[0]
|
||||
assert call["model"] == "claude-sonnet-4-5"
|
||||
assert call["body"] == REQUEST_BODY
|
||||
|
|
@ -261,53 +308,33 @@ async def test_gate_invokes_rust_and_marks_response_header():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_falls_back_only_when_bridge_declines(monkeypatch):
|
||||
_install_fake_bridge_exceptions(monkeypatch)
|
||||
bridge = RaisingAsyncMessages(FakeBridgeDeclined("unsupported request"))
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
async def test_gate_does_not_fall_back_for_unknown_bridge_error():
|
||||
bridge = RaisingAsyncMessages()
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
response = await _gate()
|
||||
|
||||
assert response is None
|
||||
with pytest.raises(RuntimeError, match="upstream request failed"):
|
||||
await _gate()
|
||||
assert bridge.calls == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_surfaces_an_upstream_failure_without_fallback(monkeypatch):
|
||||
_install_fake_bridge_exceptions(monkeypatch)
|
||||
bridge = RaisingAsyncMessages(FakeUpstreamError(429, "429: rate limited"))
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
async def test_gate_does_not_fall_back_after_upstream_send(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(rust_bridge_runtime, "get_native_bridge", lambda: _FakeNativeErrors())
|
||||
bridge = RaisingTypedAsyncMessages(_FakeUpstream(429, "rate limited"))
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
with pytest.raises(APIError) as exc_info:
|
||||
with pytest.raises(litellm.APIError, match="rate limited"):
|
||||
await _gate()
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert "429: rate limited" in str(exc_info.value)
|
||||
assert bridge.calls == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_maps_statusless_upstream_failure_to_500_without_fallback(monkeypatch):
|
||||
_install_fake_bridge_exceptions(monkeypatch)
|
||||
bridge = RaisingAsyncMessages(FakeUpstreamError(0, "request timed out"))
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
|
||||
with pytest.raises(APIError) as exc_info:
|
||||
await _gate()
|
||||
|
||||
assert exc_info.value.status_code == 500
|
||||
assert "request timed out" in str(exc_info.value)
|
||||
assert bridge.calls == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_reraises_an_unknown_bridge_failure():
|
||||
bridge = RaisingAsyncMessages(RuntimeError("unknown bridge failure"))
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
|
||||
with pytest.raises(RuntimeError, match="unknown bridge failure"):
|
||||
await _gate()
|
||||
async def test_gate_falls_back_when_rust_declines(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(rust_bridge_runtime, "get_native_bridge", lambda: _FakeNativeErrors())
|
||||
bridge = RaisingTypedAsyncMessages(_FakeDeclined("unsupported"))
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
assert await _gate() is None
|
||||
assert bridge.calls == 1
|
||||
|
||||
|
||||
|
|
@ -322,22 +349,10 @@ async def test_gate_skips_rust_when_flag_absent():
|
|||
assert bridge.calls == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_uses_process_enable_without_request_override():
|
||||
bridge = RecordingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
litellm.use_litellm_rust(True)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
|
||||
assert response is not None
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_when_flag_false():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False))
|
||||
|
||||
|
|
@ -348,7 +363,7 @@ async def test_gate_skips_rust_when_flag_false():
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_invokes_rust_for_native_anthropic_provider():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
response = await _gate(
|
||||
custom_llm_provider="anthropic",
|
||||
|
|
@ -359,7 +374,11 @@ async def test_gate_invokes_rust_for_native_anthropic_provider():
|
|||
)
|
||||
|
||||
assert response is not None
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert response["_hidden_params"]["core_engine"] == "rust"
|
||||
assert response["_hidden_params"]["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
|
||||
assert bridge.calls[0]["api_key"] == "sk-ant"
|
||||
|
||||
|
|
@ -397,7 +416,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_unsupported_provider():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
response = await _gate(custom_llm_provider="openai")
|
||||
|
||||
|
|
@ -408,7 +427,7 @@ async def test_gate_skips_rust_for_unsupported_provider():
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_agentic_hook():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
response = await _gate(has_agentic_hook=True)
|
||||
|
||||
|
|
@ -419,7 +438,7 @@ async def test_gate_skips_rust_for_agentic_hook():
|
|||
@pytest.mark.asyncio
|
||||
async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
_use_test_rust(True, amessages=bridge)
|
||||
|
||||
streaming_body = {**REQUEST_BODY, "stream": True}
|
||||
response = await _gate(
|
||||
|
|
@ -428,7 +447,11 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag():
|
|||
)
|
||||
|
||||
assert response is not None
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert response["_hidden_params"]["core_engine"] == "rust"
|
||||
assert response["_hidden_params"]["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
assert "stream" not in bridge.calls[0]["body"]
|
||||
assert bridge.calls[0]["body"] == REQUEST_BODY
|
||||
|
||||
|
|
@ -438,7 +461,11 @@ async def test_fake_stream_wraps_rust_response_as_anthropic_sse():
|
|||
response = cast(AnthropicMessagesResponse, dict(FAKE_MESSAGES_RESPONSE))
|
||||
stream = BaseLLMHTTPHandler._rust_anthropic_messages_fake_stream(response)
|
||||
|
||||
assert stream._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert stream._hidden_params["core_engine"] == "rust"
|
||||
assert stream._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
|
||||
chunks = [chunk async for chunk in stream]
|
||||
joined = b"".join(chunks)
|
||||
|
|
@ -451,11 +478,7 @@ async def test_fake_stream_wraps_rust_response_as_anthropic_sse():
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_falls_back_when_bridge_unavailable(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge"),
|
||||
"get_native_bridge",
|
||||
lambda: None,
|
||||
)
|
||||
monkeypatch.setattr(rust_bridge_runtime, "get_native_bridge", lambda: None)
|
||||
litellm.use_litellm_rust(True)
|
||||
|
||||
response = await _gate()
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import pytest
|
|||
import litellm
|
||||
from litellm.constants import RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.llms.anthropic.chat.handler import ModelResponseIterator, make_call
|
||||
from litellm.rust_bridge import runtime as bridge_runtime
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolCallFunctionChunk,
|
||||
|
|
@ -2286,7 +2287,11 @@ class TestRustChatCompletionsHook:
|
|||
response = AnthropicChatCompletion().completion(**self._completion_kwargs())
|
||||
|
||||
assert response.choices[0].message.content == "hello from rust"
|
||||
assert response._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert response._hidden_params["core_engine"] == "rust"
|
||||
assert response._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
assert len(seen["call"]) == 1
|
||||
|
||||
def test_the_core_receives_the_untranslated_openai_messages(self):
|
||||
|
|
@ -2427,7 +2432,7 @@ class TestRustChatCompletionsHook:
|
|||
def declining_native(**_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
||||
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
|
||||
monkeypatch.setattr(bridge_runtime, "get_native_bridge", lambda: _FakeNative())
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, chat_completions=declining_native
|
||||
)
|
||||
|
|
@ -2459,7 +2464,7 @@ class TestRustChatCompletionsHook:
|
|||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
|
||||
monkeypatch.setattr(bridge_runtime, "get_native_bridge", lambda: _FakeNative())
|
||||
|
||||
async def declining_native(**_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
|
@ -2501,7 +2506,11 @@ class TestRustChatCompletionsHook:
|
|||
)
|
||||
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert result._hidden_params["core_engine"] == "rust"
|
||||
assert result._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
assert not python_call.called
|
||||
|
||||
|
||||
|
|
@ -2519,7 +2528,7 @@ class TestRustChatCompletionsHook:
|
|||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
|
||||
monkeypatch.setattr(bridge_runtime, "get_native_bridge", lambda: _FakeNative())
|
||||
|
||||
def declining_native(**_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from botocore.credentials import Credentials
|
|||
from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
from litellm.rust_bridge import runtime as bridge_runtime
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
RUST_RESPONSE = {
|
||||
|
|
@ -118,7 +119,11 @@ def test_rust_true_serves_the_call_and_stamps_the_header():
|
|||
response = _run()
|
||||
|
||||
assert response.choices[0].message.content == "hello from rust"
|
||||
assert response._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert response._hidden_params["core_engine"] == "rust"
|
||||
assert response._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
assert len(seen["call"]) == 1
|
||||
|
||||
|
||||
|
|
@ -206,7 +211,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
|
|||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
|
||||
monkeypatch.setattr(bridge_runtime, "get_native_bridge", lambda: _FakeNative())
|
||||
|
||||
async def declining_native(**_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
|
@ -256,7 +261,11 @@ async def test_the_async_path_serves_the_rust_response_without_the_fallback():
|
|||
)
|
||||
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert result._hidden_params["core_engine"] == "rust"
|
||||
assert result._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
assert not python_call.called
|
||||
|
||||
|
||||
|
|
@ -283,7 +292,7 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines():
|
|||
return ModelResponse()
|
||||
|
||||
with (
|
||||
patch.object(bridge, "get_native_bridge", lambda: _FakeNative()),
|
||||
patch.object(bridge_runtime, "get_native_bridge", lambda: _FakeNative()),
|
||||
patch.object(
|
||||
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
|
||||
),
|
||||
|
|
@ -390,7 +399,7 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines():
|
|||
|
||||
logging_obj = MagicMock()
|
||||
|
||||
with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()):
|
||||
with patch.object(bridge_runtime, "get_native_bridge", lambda: _FakeNative()):
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, chat_completions=declining_native
|
||||
)
|
||||
|
|
@ -475,7 +484,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines():
|
|||
|
||||
logging_obj, calls = _recording_logging_obj()
|
||||
|
||||
with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()):
|
||||
with patch.object(bridge_runtime, "get_native_bridge", lambda: _FakeNative()):
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, chat_completions=declining_native
|
||||
)
|
||||
|
|
|
|||
|
|
@ -11,6 +11,8 @@ import pytest
|
|||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.ocr import RustAocr, RustOcr
|
||||
from litellm.rust_bridge.runtime import UNSET, Unset
|
||||
|
||||
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
|
||||
# function onto `litellm.ocr` and shadows the submodule, so import the modules
|
||||
|
|
@ -34,6 +36,16 @@ FAKE_OCR_RESPONSE: dict[str, object] = {
|
|||
}
|
||||
|
||||
|
||||
def _use_test_rust(
|
||||
enabled: bool = True,
|
||||
*,
|
||||
ocr: RustOcr | None | Unset = UNSET,
|
||||
aocr: RustAocr | None | Unset = UNSET,
|
||||
) -> None:
|
||||
rust_bridge.set_rust_ocr(ocr=ocr, aocr=aocr)
|
||||
litellm.use_litellm_rust(enabled)
|
||||
|
||||
|
||||
class CapturedException(Exception):
|
||||
pass
|
||||
|
||||
|
|
@ -228,7 +240,8 @@ def _reset_rust_flag():
|
|||
def fake_bridge():
|
||||
"""Enable the Rust path with an injected recording bridge (no native wheel)."""
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
litellm.use_litellm_rust(True)
|
||||
return bridge
|
||||
|
||||
|
||||
|
|
@ -236,7 +249,8 @@ def fake_bridge():
|
|||
def fake_async_bridge():
|
||||
"""Enable the async Rust path with an injected recording bridge."""
|
||||
bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, aocr=bridge)
|
||||
rust_bridge.set_rust_ocr(aocr=bridge)
|
||||
litellm.use_litellm_rust(True)
|
||||
return bridge
|
||||
|
||||
|
||||
|
|
@ -248,21 +262,16 @@ def test_use_litellm_rust_toggles_flag():
|
|||
assert rust_bridge.rust_ocr_enabled() is False
|
||||
|
||||
|
||||
def test_env_var_enables_rust_ocr(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_USE_RUST_OCR", "1")
|
||||
with pytest.warns(DeprecationWarning, match="LITELLM_USE_RUST_OCR is deprecated"):
|
||||
assert rust_bridge.rust_ocr_enabled() is True
|
||||
|
||||
|
||||
def test_explicit_false_overrides_process_enable():
|
||||
def test_explicit_false_overrides_the_process_switch():
|
||||
litellm.use_litellm_rust(True)
|
||||
prepared = build_prepared_request(litellm_params={"rust": False})
|
||||
|
||||
assert ocr_main._rust_ocr_enabled(build_prepared_request(litellm_params={"rust": False})) is False
|
||||
assert ocr_main._rust_ocr_enabled(prepared) is False
|
||||
|
||||
|
||||
def test_load_rust_ocr_returns_injected_impl():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
|
||||
|
||||
|
|
@ -306,40 +315,21 @@ def test_native_bridge_available_reflects_loader(monkeypatch):
|
|||
|
||||
def test_load_rust_aocr_returns_injected_impl():
|
||||
bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, aocr=bridge)
|
||||
_use_test_rust(True, aocr=bridge)
|
||||
assert rust_bridge.load_rust_aocr() is bridge
|
||||
|
||||
|
||||
def test_toggle_without_ocr_arg_preserves_injected_impl():
|
||||
"""Regression: routine enable/disable calls must not clobber a prior injection.
|
||||
|
||||
Earlier, ``use_litellm_rust()`` unconditionally assigned the keyword default
|
||||
of ``None`` to ``_rust_ocr_impl``, silently dropping a custom bridge whenever
|
||||
a caller toggled the flag without re-passing ``ocr=``.
|
||||
"""
|
||||
bridge = RecordingBridge()
|
||||
async_bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge, aocr=async_bridge)
|
||||
|
||||
litellm.use_litellm_rust(False)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
assert rust_bridge.load_rust_aocr() is async_bridge
|
||||
litellm.use_litellm_rust(True)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
assert rust_bridge.load_rust_aocr() is async_bridge
|
||||
|
||||
|
||||
def test_explicit_ocr_none_clears_injected_impl(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge"),
|
||||
importlib.import_module("litellm.rust_bridge.runtime"),
|
||||
"get_native_bridge",
|
||||
lambda: None,
|
||||
)
|
||||
bridge = RecordingBridge()
|
||||
async_bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge, aocr=async_bridge)
|
||||
rust_bridge.set_rust_ocr(ocr=bridge, aocr=async_bridge)
|
||||
|
||||
litellm.use_litellm_rust(True, ocr=None, aocr=None)
|
||||
rust_bridge.set_rust_ocr(ocr=None, aocr=None)
|
||||
assert rust_bridge.load_rust_ocr() is None
|
||||
assert rust_bridge.load_rust_aocr() is None
|
||||
|
||||
|
|
@ -348,7 +338,7 @@ def test_load_rust_ocr_none_when_extension_absent(monkeypatch):
|
|||
"""With no injected impl and no compiled wheel, the loader returns None so the
|
||||
caller degrades to the Python path instead of raising ImportError."""
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge"),
|
||||
importlib.import_module("litellm.rust_bridge.runtime"),
|
||||
"get_native_bridge",
|
||||
lambda: None,
|
||||
)
|
||||
|
|
@ -365,7 +355,7 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch):
|
|||
fake_module.ocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined]
|
||||
fake_module.aocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined]
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge"),
|
||||
importlib.import_module("litellm.rust_bridge.runtime"),
|
||||
"get_native_bridge",
|
||||
lambda: fake_module,
|
||||
)
|
||||
|
|
@ -384,7 +374,7 @@ def test_timeout_to_seconds_handles_float_timeout_and_none():
|
|||
def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
|
||||
bridge = RecordingBridge()
|
||||
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
response = rust_bridge.ocr(
|
||||
model="mistral-ocr-latest",
|
||||
document=DOCUMENT,
|
||||
|
|
@ -417,7 +407,7 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
|
|||
async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
|
||||
bridge = RecordingAsyncBridge()
|
||||
|
||||
litellm.use_litellm_rust(True, aocr=bridge)
|
||||
_use_test_rust(True, aocr=bridge)
|
||||
response = await rust_bridge.aocr(
|
||||
model="mistral-ocr-maas",
|
||||
document=DOCUMENT,
|
||||
|
|
@ -445,7 +435,7 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
|
|||
def test_run_rust_ocr_prepares_request_and_wraps_response():
|
||||
bridge = RecordingBridge()
|
||||
logging_obj = RecordingLogging()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
response = ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -460,6 +450,11 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
|
|||
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert response.pages[0].markdown == "hello world"
|
||||
assert response._hidden_params["core_engine"] == "rust"
|
||||
assert response._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
assert bridge.calls[0] == {
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": DOCUMENT,
|
||||
|
|
@ -477,7 +472,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
|
|||
|
||||
def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(api_key=None, timeout=None),
|
||||
|
|
@ -489,7 +484,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
|||
|
||||
def test_run_rust_ocr_prefers_explicit_key_over_resolver():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
def _resolver(name: str) -> str | None:
|
||||
raise AssertionError(f"resolver should not be called for {name}")
|
||||
|
|
@ -508,7 +503,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver():
|
|||
def test_run_rust_ocr_uses_provider_api_key_env_var():
|
||||
bridge = RecordingBridge()
|
||||
resolver_calls = []
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
def _resolver(name):
|
||||
resolver_calls.append(name)
|
||||
|
|
@ -530,7 +525,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
|
|||
|
||||
def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -556,7 +551,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
|||
|
||||
def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
def _resolver(name: str) -> str | None:
|
||||
return {
|
||||
|
|
@ -579,7 +574,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana
|
|||
|
||||
def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -596,7 +591,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
|
|||
|
||||
def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -616,7 +611,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
|
|||
def test_run_rust_ocr_runs_pre_call_logging():
|
||||
logging_obj = RecordingLogging()
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
|
|
@ -703,7 +698,7 @@ def test_ocr_exception_type_uses_resolved_provider_context(
|
|||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.use_litellm_rust(True, ocr=RaisingBridge())
|
||||
_use_test_rust(True, ocr=RaisingBridge())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
|
@ -724,6 +719,7 @@ async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge):
|
|||
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert response.pages[0].markdown == "hello world"
|
||||
assert response._hidden_params["core_engine"] == "rust"
|
||||
assert len(fake_async_bridge.calls) == 1
|
||||
call = fake_async_bridge.calls[0]
|
||||
assert call["model"] == "mistral-ocr-latest"
|
||||
|
|
@ -748,7 +744,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context(
|
|||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.use_litellm_rust(True, aocr=RaisingAsyncBridge())
|
||||
_use_test_rust(True, aocr=RaisingAsyncBridge())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
|
@ -776,7 +772,7 @@ def test_ocr_passes_default_request_timeout_to_rust(fake_bridge):
|
|||
def test_ocr_does_not_route_to_rust_when_disabled():
|
||||
"""With the flag off, the bridge must not be consulted even if an impl exists."""
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(False, ocr=bridge)
|
||||
_use_test_rust(False, ocr=bridge)
|
||||
|
||||
assert rust_bridge.rust_ocr_enabled() is False
|
||||
# The impl stays available for injection, but the disabled flag gates usage,
|
||||
|
|
@ -802,6 +798,29 @@ def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch):
|
|||
|
||||
assert captured.get("called") is True # Python path was used
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert response._hidden_params["core_engine"] == "python"
|
||||
assert response._hidden_params["additional_headers"] == {"x-litellm-core": "python"}
|
||||
|
||||
|
||||
def test_ocr_unsupported_provider_skips_rust_and_reports_python(monkeypatch):
|
||||
bridge = RecordingBridge()
|
||||
_use_test_rust(True, ocr=bridge)
|
||||
|
||||
def fake_handler_ocr(**kwargs):
|
||||
return OCRResponse(pages=[], model="parse-v3", object="ocr")
|
||||
|
||||
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fake_handler_ocr)
|
||||
|
||||
response = litellm.ocr(
|
||||
model="reducto/parse-v3",
|
||||
document={"type": "document_url", "document_url": "reducto://document-id"},
|
||||
api_key="test-key",
|
||||
)
|
||||
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert response._hidden_params["core_engine"] == "python"
|
||||
assert response._hidden_params["additional_headers"] == {"x-litellm-core": "python"}
|
||||
assert bridge.calls == []
|
||||
|
||||
|
||||
def test_ocr_provider_configs_expose_api_key_env_vars():
|
||||
|
|
|
|||
|
|
@ -1241,6 +1241,28 @@ class TestProxyBaseLLMRequestProcessing:
|
|||
breakdown_none = _get_cost_breakdown_from_logging_obj(None)
|
||||
assert all(value is None for value in breakdown_none)
|
||||
|
||||
def test_get_custom_headers_sets_and_reserves_core_provenance(self):
|
||||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||||
mock_user_api_key_dict.tpm_limit = None
|
||||
mock_user_api_key_dict.rpm_limit = None
|
||||
mock_user_api_key_dict.max_budget = None
|
||||
mock_user_api_key_dict.spend = 0
|
||||
|
||||
python_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
**{"x-litellm-core": "rust", "x-litellm-rust": "true"},
|
||||
)
|
||||
rust_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
hidden_params={"core_engine": "rust"},
|
||||
**{"x-litellm-core": "python", "x-litellm-rust": "false"},
|
||||
)
|
||||
|
||||
assert python_headers["x-litellm-core"] == "python"
|
||||
assert "x-litellm-rust" not in python_headers
|
||||
assert rust_headers["x-litellm-core"] == "rust"
|
||||
assert rust_headers["x-litellm-rust"] == "true"
|
||||
|
||||
def test_get_custom_headers_key_spend_includes_response_cost(self):
|
||||
"""
|
||||
Test that x-litellm-key-spend header includes the current request's response_cost.
|
||||
|
|
|
|||
|
|
@ -4,8 +4,10 @@ from collections.abc import Mapping
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
|
||||
from litellm.rust_bridge import configuration, responses_websocket
|
||||
from litellm.rust_bridge import configuration, responses_websocket, runtime
|
||||
from litellm.rust_bridge.runtime import CoreEngine
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
|
||||
|
|
@ -35,6 +37,24 @@ class _ClosedNativeConnection:
|
|||
return None
|
||||
|
||||
|
||||
class _FailingNativeConnection(_ClosedNativeConnection):
|
||||
async def send_event(self, event: Mapping[str, object]) -> None:
|
||||
raise _Upstream(503, "session send failed")
|
||||
|
||||
|
||||
class _Declined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _Upstream(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class _NativeErrors:
|
||||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = _Upstream
|
||||
|
||||
|
||||
class _FakeNativeBridge:
|
||||
@classmethod
|
||||
async def connect(
|
||||
|
|
@ -85,19 +105,17 @@ async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(responses_websocket, "_STATE", responses_websocket._RustResponsesWebSocketState())
|
||||
monkeypatch.setattr(responses_websocket, "get_native_bridge", lambda: None)
|
||||
monkeypatch.setattr(runtime, "get_native_bridge", lambda: None)
|
||||
|
||||
assert (
|
||||
await responses_websocket.connect(
|
||||
provider="openai",
|
||||
api_key=None,
|
||||
api_base="https://example.test",
|
||||
headers={},
|
||||
timeout=None,
|
||||
)
|
||||
is None
|
||||
result = await responses_websocket.connect(
|
||||
provider="openai",
|
||||
api_key=None,
|
||||
api_base="https://example.test",
|
||||
headers={},
|
||||
timeout=None,
|
||||
)
|
||||
assert result.value is None
|
||||
assert result.source is CoreEngine.PYTHON
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -106,7 +124,7 @@ async def test_enabled_bridge_connects_and_adapts_socket(
|
|||
) -> None:
|
||||
responses_websocket.set_rust_responses_websocket(connection=_FakeNativeBridge)
|
||||
|
||||
connection = await responses_websocket.connect(
|
||||
result = await responses_websocket.connect(
|
||||
provider="openai",
|
||||
api_key="key",
|
||||
api_base="https://example.test",
|
||||
|
|
@ -114,7 +132,20 @@ async def test_enabled_bridge_connects_and_adapts_socket(
|
|||
timeout=1.0,
|
||||
)
|
||||
|
||||
assert result.source is CoreEngine.RUST
|
||||
connection = result.value
|
||||
assert connection is not None
|
||||
assert connection.core_engine is CoreEngine.RUST
|
||||
await connection.send('{"type":"response.create","model":"gpt-5"}')
|
||||
assert await connection.recv() == '{"type":"response.completed"}'
|
||||
await connection.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connected_rust_session_never_falls_back(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(runtime, "get_native_bridge", lambda: _NativeErrors())
|
||||
connection = responses_websocket._ConnectionAdapter(_FailingNativeConnection())
|
||||
|
||||
assert connection.core_engine is CoreEngine.RUST
|
||||
with pytest.raises(APIError, match="session send failed"):
|
||||
await connection.send('{"type":"response.create"}')
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge import configuration, runtime
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
RUST_RESPONSE = {
|
||||
|
|
@ -54,7 +54,7 @@ class _FakeNative:
|
|||
|
||||
def _fake_native_bridge(monkeypatch):
|
||||
"""Expose the bridge's exception classes without the compiled extension."""
|
||||
monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative())
|
||||
monkeypatch.setattr(runtime, "get_native_bridge", lambda: _FakeNative())
|
||||
|
||||
|
||||
def _hide_native_bridge(monkeypatch):
|
||||
|
|
@ -63,31 +63,23 @@ def _hide_native_bridge(monkeypatch):
|
|||
There is no injection seam for "the .so is absent", so the loader itself is
|
||||
replaced; every other case here uses `set_rust_chat_completions`.
|
||||
"""
|
||||
monkeypatch.setattr(bridge, "get_native_bridge", lambda: None)
|
||||
monkeypatch.setattr(runtime, "get_native_bridge", lambda: None)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_bridge():
|
||||
"""Every test starts with no injected callables, and leaves none behind."""
|
||||
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
|
||||
bridge.set_rust_chat_completions(
|
||||
chat_completions=None,
|
||||
achat_completions=None,
|
||||
decline=lambda **_kwargs: None,
|
||||
)
|
||||
configuration.reset_rust_configuration()
|
||||
yield
|
||||
bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None)
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
class _RecordingDecline:
|
||||
"""A stand-in for the native gate that records what it was asked."""
|
||||
|
||||
def __init__(self, reason: str | None = None):
|
||||
self.reason = reason
|
||||
self.calls: list[dict] = []
|
||||
|
||||
def __call__(self, **kwargs):
|
||||
self.calls.append(kwargs)
|
||||
return self.reason
|
||||
|
||||
|
||||
class _RecordingCall:
|
||||
def __init__(self, result=None, error: Exception | None = None):
|
||||
self.result = result if result is not None else dict(RUST_RESPONSE)
|
||||
|
|
@ -106,7 +98,7 @@ class _RecordingAsyncCall(_RecordingCall):
|
|||
return _RecordingCall.__call__(self, **kwargs)
|
||||
|
||||
|
||||
def _accepts(**overrides) -> bool:
|
||||
def _should_attempt(**overrides) -> bool:
|
||||
kwargs = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": MESSAGES,
|
||||
|
|
@ -116,52 +108,40 @@ def _accepts(**overrides) -> bool:
|
|||
"stream": None,
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
kwargs.pop("asynchronous", None)
|
||||
return bridge.rust_chat_completions_accepts(**kwargs)
|
||||
|
||||
|
||||
class TestGate:
|
||||
class TestEligibility:
|
||||
def test_declines_when_the_deployment_did_not_opt_in(self, monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
assert _accepts(litellm_params={}) is False
|
||||
assert _accepts(litellm_params=None) is False
|
||||
assert _accepts(litellm_params={"rust": False}) is False
|
||||
assert gate.calls == [], "the gate must not be consulted before opt-in"
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall())
|
||||
assert _should_attempt(litellm_params={}) is False
|
||||
assert _should_attempt(litellm_params=None) is False
|
||||
assert _should_attempt(litellm_params={"rust": False}) is False
|
||||
|
||||
def test_accepts_when_the_deployment_opted_in_and_the_core_agrees(self, monkeypatch):
|
||||
def test_attempts_when_the_deployment_opted_in(self, monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
assert _accepts() is True
|
||||
assert gate.calls[0]["model"] == "claude-sonnet-4-5"
|
||||
assert gate.calls[0]["custom_llm_provider"] == "anthropic"
|
||||
|
||||
def test_explicit_false_overrides_process_enable(self):
|
||||
bridge.set_rust_chat_completions(decline=_RecordingDecline())
|
||||
configuration.use_litellm_rust(True)
|
||||
|
||||
assert _accepts(litellm_params={"rust": False}) is False
|
||||
|
||||
def test_process_enable_applies_without_request_override(self):
|
||||
bridge.set_rust_chat_completions(decline=_RecordingDecline())
|
||||
configuration.use_litellm_rust(True)
|
||||
|
||||
assert _accepts(litellm_params={}) is True
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall())
|
||||
assert _should_attempt() is True
|
||||
|
||||
def test_the_env_var_opts_in_without_a_per_model_flag(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "true")
|
||||
bridge.set_rust_chat_completions(decline=_RecordingDecline())
|
||||
assert _accepts(litellm_params={}) is True
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall())
|
||||
assert _should_attempt(litellm_params={}) is True
|
||||
|
||||
def test_explicit_false_overrides_the_process_switch(self):
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall())
|
||||
configuration.use_litellm_rust(True)
|
||||
|
||||
assert _should_attempt(litellm_params={"rust": False}) is False
|
||||
|
||||
def test_declines_streaming_and_providers_off_the_path(self, monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
assert _accepts(stream=True) is False
|
||||
assert _accepts(custom_llm_provider="openai") is False
|
||||
assert _accepts(custom_llm_provider=None) is False
|
||||
assert gate.calls == []
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall())
|
||||
assert _should_attempt(stream=True) is False
|
||||
assert _should_attempt(custom_llm_provider="openai") is False
|
||||
assert _should_attempt(custom_llm_provider=None) is False
|
||||
|
||||
def test_declines_an_anthropic_request_carrying_a_litellm_metadata_user_id(self, monkeypatch):
|
||||
"""`AnthropicConfig.transform_request` copies a valid `user_id` into the Messages body.
|
||||
|
|
@ -171,24 +151,21 @@ class TestGate:
|
|||
to Anthropic with the abuse-detection attribution silently missing.
|
||||
"""
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
assert _accepts(litellm_params={"rust": True, "metadata": {"user_id": "u-123"}}) is False
|
||||
assert gate.calls == [], "the core must not be consulted for a request it cannot see the key of"
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall())
|
||||
assert _should_attempt(litellm_params={"rust": True, "metadata": {"user_id": "u-123"}}) is False
|
||||
|
||||
# Bedrock's Converse transform reads no `user_id`, and an Anthropic request
|
||||
# whose metadata carries none is one Python would not attribute either.
|
||||
assert (
|
||||
_accepts(
|
||||
_should_attempt(
|
||||
custom_llm_provider="bedrock",
|
||||
model="bedrock/us-east-1/anthropic.claude-v2",
|
||||
litellm_params={"rust": True, "metadata": {"user_id": "u-123"}},
|
||||
)
|
||||
is True
|
||||
)
|
||||
assert _accepts(litellm_params={"rust": True, "metadata": {"trace_id": "t-1"}}) is True
|
||||
assert _accepts(litellm_params={"rust": True, "metadata": {"user_id": None}}) is True
|
||||
assert _accepts(litellm_params={"rust": True, "metadata": None}) is True
|
||||
assert _should_attempt(litellm_params={"rust": True, "metadata": {"trace_id": "t-1"}}) is True
|
||||
assert _should_attempt(litellm_params={"rust": True, "metadata": {"user_id": None}}) is True
|
||||
assert _should_attempt(litellm_params={"rust": True, "metadata": None}) is True
|
||||
|
||||
def test_declines_a_bedrock_request_while_the_proxy_owns_request_metadata(self, monkeypatch):
|
||||
"""`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the
|
||||
|
|
@ -197,39 +174,30 @@ class TestGate:
|
|||
who armed `bedrock_request_metadata_fields` keeps the Python path.
|
||||
"""
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall())
|
||||
bedrock = {
|
||||
"custom_llm_provider": "bedrock",
|
||||
"model": "bedrock/us-east-1/anthropic.claude-v2",
|
||||
}
|
||||
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["user_api_key_team_id"])
|
||||
assert _accepts(**bedrock) is False
|
||||
assert gate.calls == [], "the core must not be consulted for a field it cannot write"
|
||||
assert _accepts() is True, "arming Bedrock attribution must not decline Anthropic"
|
||||
assert _should_attempt(**bedrock) is False
|
||||
assert _should_attempt() is True, "arming Bedrock attribution must not decline Anthropic"
|
||||
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", None)
|
||||
assert _accepts(**bedrock) is True, "the decline follows the operator's opt-in alone"
|
||||
|
||||
def test_declines_when_the_core_declines(self, monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
bridge.set_rust_chat_completions(decline=_RecordingDecline("streaming"))
|
||||
assert _accepts() is False
|
||||
assert _should_attempt(**bedrock) is True, "the decline follows the operator's opt-in alone"
|
||||
|
||||
def test_declines_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
bridge.set_rust_chat_completions(decline=None)
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert _accepts() is False
|
||||
assert _should_attempt() is False
|
||||
|
||||
def test_declines_when_the_gate_itself_raises(self, monkeypatch):
|
||||
monkeypatch.delenv("LITELLM_RUST", raising=False)
|
||||
|
||||
def exploding(**_kwargs):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
bridge.set_rust_chat_completions(decline=exploding)
|
||||
assert _accepts() is False
|
||||
def test_checks_the_native_capability_gate(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(decline=None)
|
||||
assert _should_attempt() is False
|
||||
bridge.set_rust_chat_completions(decline=lambda **_kwargs: None)
|
||||
assert _should_attempt() is True
|
||||
|
||||
|
||||
def _call_kwargs(model_response: ModelResponse) -> dict:
|
||||
|
|
@ -247,6 +215,18 @@ def _call_kwargs(model_response: ModelResponse) -> dict:
|
|||
}
|
||||
|
||||
|
||||
def _sync_call_kwargs(model_response: ModelResponse) -> dict:
|
||||
return {**_call_kwargs(model_response), "python_fallback": lambda: "python"}
|
||||
|
||||
|
||||
async def _async_python_fallback() -> object:
|
||||
return "python"
|
||||
|
||||
|
||||
def _async_call_kwargs(model_response: ModelResponse) -> dict:
|
||||
return {**_call_kwargs(model_response), "python_fallback": _async_python_fallback}
|
||||
|
||||
|
||||
class TestSyncCall:
|
||||
def test_builds_a_model_response_and_stamps_the_rust_header(self):
|
||||
native = _RecordingCall()
|
||||
|
|
@ -254,7 +234,7 @@ class TestSyncCall:
|
|||
model_response = ModelResponse()
|
||||
original_id = model_response.id
|
||||
|
||||
result = bridge.chat_completions(**_call_kwargs(model_response))
|
||||
result = bridge.chat_completions_or_fallback(**_sync_call_kwargs(model_response))
|
||||
|
||||
assert result is not None
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
|
|
@ -263,44 +243,66 @@ class TestSyncCall:
|
|||
assert result.usage.prompt_tokens == 11
|
||||
assert result.usage.completion_tokens == 4
|
||||
assert result.usage.total_tokens == 15
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert result._hidden_params["core_engine"] == "rust"
|
||||
assert result._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
assert result.id == original_id, "the rust path must keep the chatcmpl id litellm already minted"
|
||||
|
||||
def test_passes_the_timeout_through_as_seconds(self):
|
||||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse()))
|
||||
assert native.calls[0]["timeout_seconds"] == 30.0
|
||||
|
||||
def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
assert bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) == "python"
|
||||
|
||||
def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
assert bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) == "python"
|
||||
|
||||
def test_model_response_fallback_is_stamped_python(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("unsupported")))
|
||||
fallback_response = ModelResponse()
|
||||
|
||||
result = bridge.chat_completions_or_fallback(
|
||||
**_call_kwargs(ModelResponse()),
|
||||
python_fallback=lambda: fallback_response,
|
||||
)
|
||||
|
||||
assert result is fallback_response
|
||||
assert result._hidden_params["core_engine"] == "python"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-core": "python"}
|
||||
|
||||
|
||||
class TestAsyncCall:
|
||||
@pytest.mark.asyncio
|
||||
async def test_builds_a_model_response(self):
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
|
||||
result = await bridge.achat_completions(**_call_kwargs(ModelResponse()))
|
||||
result = await bridge.achat_completions_or_fallback(**_async_call_kwargs(ModelResponse()))
|
||||
assert result is not None
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert result._hidden_params["core_engine"] == "rust"
|
||||
assert result._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
assert await bridge.achat_completions_or_fallback(**_async_call_kwargs(ModelResponse())) == "python"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
|
||||
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
assert await bridge.achat_completions_or_fallback(**_async_call_kwargs(ModelResponse())) == "python"
|
||||
|
||||
|
||||
class TestAsyncFallbackWrapper:
|
||||
|
|
@ -349,14 +351,14 @@ class TestFailureClassification:
|
|||
|
||||
def test_a_decline_falls_back_because_nothing_was_sent(self):
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
assert bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse())) == "python"
|
||||
|
||||
def test_an_upstream_failure_is_surfaced_with_its_status(self):
|
||||
from litellm.exceptions import APIError
|
||||
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited")))
|
||||
with pytest.raises(APIError) as raised:
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse()))
|
||||
assert raised.value.status_code == 429
|
||||
assert "rate limited" in str(raised.value)
|
||||
|
||||
|
|
@ -365,13 +367,13 @@ class TestFailureClassification:
|
|||
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset")))
|
||||
with pytest.raises(APIError) as raised:
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse()))
|
||||
assert raised.value.status_code == 500
|
||||
|
||||
def test_an_unrecognized_error_is_not_swallowed(self):
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else")))
|
||||
with pytest.raises(RuntimeError):
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
bridge.chat_completions_or_fallback(**_sync_call_kwargs(ModelResponse()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self):
|
||||
|
|
|
|||
238
tests/test_litellm/rust_bridge/test_runtime.py
Normal file
238
tests/test_litellm/rust_bridge/test_runtime.py
Normal file
|
|
@ -0,0 +1,238 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge import runtime
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
CoreEngine,
|
||||
FallbackMode,
|
||||
RustDeclined,
|
||||
RustHandled,
|
||||
RustUnavailable,
|
||||
)
|
||||
|
||||
|
||||
class Declined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class Upstream(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class Native:
|
||||
RustBridgeDeclined = Declined
|
||||
RustUpstreamError = Upstream
|
||||
|
||||
|
||||
CONTEXT = BridgeErrorContext(route="messages", provider="anthropic", model="claude")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(runtime, "get_native_bridge", lambda: Native)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("native_call", "expected", "source", "fallback_calls"),
|
||||
[
|
||||
(lambda: "native", "adapted native", CoreEngine.RUST, 0),
|
||||
(None, "fallback", CoreEngine.PYTHON, 1),
|
||||
(lambda: (_ for _ in ()).throw(Declined("not sent")), "fallback", CoreEngine.PYTHON, 1),
|
||||
],
|
||||
)
|
||||
def test_invoke_uses_fallback_only_when_safe(
|
||||
native_call: Callable[[], str] | None,
|
||||
expected: str,
|
||||
source: CoreEngine,
|
||||
fallback_calls: int,
|
||||
) -> None:
|
||||
calls = 0
|
||||
|
||||
def fallback() -> str:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
return "fallback"
|
||||
|
||||
result = runtime.invoke(
|
||||
native_call=native_call,
|
||||
fallback=fallback,
|
||||
adapt=lambda value: f"adapted {value}",
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=CONTEXT,
|
||||
)
|
||||
|
||||
assert result.value == expected
|
||||
assert result.source is source
|
||||
assert calls == fallback_calls
|
||||
|
||||
|
||||
def test_invoke_converts_upstream_error_without_fallback() -> None:
|
||||
fallback_calls = 0
|
||||
|
||||
def fallback() -> str:
|
||||
nonlocal fallback_calls
|
||||
fallback_calls += 1
|
||||
return "fallback"
|
||||
|
||||
with pytest.raises(APIError) as exc_info:
|
||||
runtime.invoke(
|
||||
native_call=lambda: (_ for _ in ()).throw(Upstream(429, "rate limited")),
|
||||
fallback=fallback,
|
||||
adapt=str,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=CONTEXT,
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert "rate limited" in str(exc_info.value)
|
||||
assert exc_info.value.headers == {"x-litellm-core": "rust", "x-litellm-rust": "true"}
|
||||
assert fallback_calls == 0
|
||||
|
||||
|
||||
def test_invoke_propagates_unknown_error_without_fallback() -> None:
|
||||
with pytest.raises(ValueError, match="bug"):
|
||||
runtime.invoke(
|
||||
native_call=lambda: (_ for _ in ()).throw(ValueError("bug")),
|
||||
fallback=lambda: pytest.fail("fallback must not run"),
|
||||
adapt=str,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=CONTEXT,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"native_call",
|
||||
[None, lambda: (_ for _ in ()).throw(Declined("unsupported"))],
|
||||
)
|
||||
def test_rust_required_route_never_falls_back(
|
||||
native_call: Callable[[], str] | None,
|
||||
) -> None:
|
||||
with pytest.raises(RuntimeError):
|
||||
runtime.invoke(
|
||||
native_call=native_call,
|
||||
fallback=lambda: pytest.fail("fallback must not run"),
|
||||
adapt=str,
|
||||
mode=FallbackMode.RUST_REQUIRED,
|
||||
context=CONTEXT,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure", [None, Declined("not sent")])
|
||||
async def test_ainvoke_success_or_safe_fallback(failure: Exception | None) -> None:
|
||||
fallback_calls = 0
|
||||
|
||||
async def native_call() -> str:
|
||||
if failure is not None:
|
||||
raise failure
|
||||
return "native"
|
||||
|
||||
async def fallback() -> str:
|
||||
nonlocal fallback_calls
|
||||
fallback_calls += 1
|
||||
return "fallback"
|
||||
|
||||
result = await runtime.ainvoke(
|
||||
native_call=native_call,
|
||||
fallback=fallback,
|
||||
adapt=lambda value: f"adapted {value}",
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=CONTEXT,
|
||||
)
|
||||
|
||||
assert result.value == ("fallback" if failure else "adapted native")
|
||||
assert result.source is (CoreEngine.PYTHON if failure else CoreEngine.RUST)
|
||||
assert fallback_calls == (1 if failure else 0)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("failure", [Upstream(503, "unavailable"), ValueError("bug")])
|
||||
async def test_ainvoke_never_falls_back_after_unsafe_failure(failure: Exception) -> None:
|
||||
async def native_call() -> str:
|
||||
raise failure
|
||||
|
||||
async def fallback() -> str:
|
||||
pytest.fail("fallback must not run")
|
||||
|
||||
expected = APIError if isinstance(failure, Upstream) else ValueError
|
||||
with pytest.raises(expected):
|
||||
await runtime.ainvoke(
|
||||
native_call=native_call,
|
||||
fallback=fallback,
|
||||
adapt=str,
|
||||
mode=FallbackMode.PYTHON,
|
||||
context=CONTEXT,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("native_call", "expected_type"),
|
||||
[
|
||||
(lambda: "native", RustHandled),
|
||||
(lambda: (_ for _ in ()).throw(Declined("unsupported")), RustDeclined),
|
||||
(None, RustUnavailable),
|
||||
],
|
||||
)
|
||||
def test_attempt_classifies_bridge_control_flow(
|
||||
native_call: Callable[[], str] | None,
|
||||
expected_type: type[RustHandled[str]] | type[RustDeclined] | type[RustUnavailable],
|
||||
) -> None:
|
||||
result = runtime.attempt(native_call=native_call, adapt=str, context=CONTEXT)
|
||||
|
||||
assert isinstance(result, expected_type)
|
||||
if isinstance(result, RustDeclined):
|
||||
assert result.reason == "unsupported"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
("failure", "expected_type"),
|
||||
[
|
||||
(None, RustHandled),
|
||||
(Declined("unsupported"), RustDeclined),
|
||||
],
|
||||
)
|
||||
async def test_aattempt_classifies_bridge_control_flow(
|
||||
failure: Exception | None,
|
||||
expected_type: type[RustHandled[str]] | type[RustDeclined],
|
||||
) -> None:
|
||||
async def native_call() -> str:
|
||||
if failure is not None:
|
||||
raise failure
|
||||
return "native"
|
||||
|
||||
result = await runtime.aattempt(native_call=native_call, adapt=str, context=CONTEXT)
|
||||
|
||||
assert isinstance(result, expected_type)
|
||||
if isinstance(result, RustDeclined):
|
||||
assert result.reason == "unsupported"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aattempt_classifies_unavailable_bridge() -> None:
|
||||
result = await runtime.aattempt(native_call=None, adapt=str, context=CONTEXT)
|
||||
|
||||
assert isinstance(result, RustUnavailable)
|
||||
|
||||
|
||||
def test_execution_hidden_params_overwrites_reserved_provider_headers() -> None:
|
||||
hidden_params = runtime.execution_hidden_params(
|
||||
{
|
||||
"additional_headers": {
|
||||
"X-LiteLLM-Core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
"x-provider": "kept",
|
||||
}
|
||||
},
|
||||
CoreEngine.PYTHON,
|
||||
)
|
||||
|
||||
assert hidden_params == {
|
||||
"core_engine": "python",
|
||||
"additional_headers": {"x-provider": "kept", "x-litellm-core": "python"},
|
||||
}
|
||||
|
|
@ -7,7 +7,8 @@ from unittest.mock import MagicMock
|
|||
import pytest
|
||||
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.rust_bridge import streaming
|
||||
from litellm.rust_bridge import runtime, streaming
|
||||
from litellm.rust_bridge.runtime import CoreEngine
|
||||
|
||||
|
||||
class _FakeEventStream:
|
||||
|
|
@ -101,6 +102,24 @@ class _FailingOpen:
|
|||
raise self._error
|
||||
|
||||
|
||||
class _MidstreamFailingStream(_FakeEventStream):
|
||||
def next_event(self) -> Mapping[str, object] | None:
|
||||
raise _Declined("decoder rejected event")
|
||||
|
||||
|
||||
class _MidstreamFailingOpen:
|
||||
def __call__(
|
||||
self,
|
||||
request: Mapping[str, object],
|
||||
provider: str,
|
||||
credentials: Mapping[str, str] | None,
|
||||
api_base: str | None,
|
||||
extra_headers: Mapping[str, str] | None,
|
||||
timeout_seconds: float | None,
|
||||
) -> _FakeEventStream:
|
||||
return _MidstreamFailingStream(())
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_bridge() -> Iterator[None]:
|
||||
streaming.set_rust_streaming(
|
||||
|
|
@ -133,7 +152,7 @@ def _chat_event(text: str) -> Mapping[str, object]:
|
|||
|
||||
|
||||
def test_unavailable_native_bridge_falls_back(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(streaming, "get_native_bridge", lambda: None)
|
||||
monkeypatch.setattr(runtime, "get_native_bridge", lambda: None)
|
||||
result: Final = streaming.open_stream(
|
||||
api="chat_completions",
|
||||
provider="anthropic",
|
||||
|
|
@ -144,15 +163,16 @@ def test_unavailable_native_bridge_falls_back(monkeypatch: pytest.MonkeyPatch) -
|
|||
timeout=None,
|
||||
)
|
||||
|
||||
assert result is None
|
||||
assert result.value is None
|
||||
assert result.source is CoreEngine.PYTHON
|
||||
|
||||
|
||||
def test_declined_open_failure_falls_back(
|
||||
def test_guaranteed_not_sent_open_failure_falls_back(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
python_calls = 0
|
||||
|
||||
monkeypatch.setattr(streaming, "get_native_bridge", lambda: _NativeErrors())
|
||||
monkeypatch.setattr(runtime, "get_native_bridge", lambda: _NativeErrors())
|
||||
streaming.set_rust_streaming(
|
||||
chat=_FailingOpen(_Declined("unsupported request")),
|
||||
)
|
||||
|
|
@ -166,16 +186,17 @@ def test_declined_open_failure_falls_back(
|
|||
extra_headers=None,
|
||||
timeout=None,
|
||||
)
|
||||
if result is None:
|
||||
if result.value is None:
|
||||
python_calls += 1
|
||||
|
||||
assert python_calls == 1
|
||||
assert result.source is CoreEngine.PYTHON
|
||||
|
||||
|
||||
def test_upstream_open_failure_never_falls_back(
|
||||
def test_possibly_sent_open_failure_never_falls_back(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setattr(streaming, "get_native_bridge", lambda: _NativeErrors())
|
||||
monkeypatch.setattr(runtime, "get_native_bridge", lambda: _NativeErrors())
|
||||
streaming.set_rust_streaming(
|
||||
chat=_FailingOpen(_Upstream(503, "connection closed after request")),
|
||||
)
|
||||
|
|
@ -206,10 +227,12 @@ def test_sync_typed_events_preserve_shape_metadata_and_close() -> None:
|
|||
timeout=1.0,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
assert tuple(event["text"] for event in result) == ("one", "two")
|
||||
assert result.metadata["provider"] == "anthropic"
|
||||
result.close()
|
||||
assert result.source is CoreEngine.RUST
|
||||
assert result.value is not None
|
||||
assert result.value.core_engine is CoreEngine.RUST
|
||||
assert tuple(event["text"] for event in result.value) == ("one", "two")
|
||||
assert result.value.metadata["provider"] == "anthropic"
|
||||
result.value.close()
|
||||
assert opener.calls == 1
|
||||
|
||||
|
||||
|
|
@ -231,6 +254,25 @@ def test_chat_events_flow_through_custom_stream_wrapper() -> None:
|
|||
assert chunk.choices[0].delta.content == "hello"
|
||||
|
||||
|
||||
def test_rust_stream_never_falls_back_after_open(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(runtime, "get_native_bridge", lambda: _NativeErrors())
|
||||
streaming.set_rust_streaming(chat=_MidstreamFailingOpen())
|
||||
result: Final = streaming.open_stream(
|
||||
api="chat_completions",
|
||||
provider="anthropic",
|
||||
request={"model": "claude", "messages": []},
|
||||
credentials=None,
|
||||
api_base=None,
|
||||
extra_headers=None,
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert result.source is CoreEngine.RUST
|
||||
assert result.value is not None
|
||||
with pytest.raises(_Declined, match="decoder rejected event"):
|
||||
next(result.value)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_typed_events_and_cancellation() -> None:
|
||||
opener: Final = _RecordingAsyncOpen((_chat_event("one"), _chat_event("two")))
|
||||
|
|
@ -245,10 +287,11 @@ async def test_async_typed_events_and_cancellation() -> None:
|
|||
timeout=None,
|
||||
)
|
||||
|
||||
assert result is not None
|
||||
collected: Final = tuple([event async for event in result])
|
||||
assert result.source is CoreEngine.RUST
|
||||
assert result.value is not None
|
||||
collected: Final = tuple([event async for event in result.value])
|
||||
assert tuple(event["text"] for event in collected) == ("one", "two")
|
||||
await result.aclose()
|
||||
await result.value.aclose()
|
||||
|
||||
|
||||
def test_messages_events_are_wrapped_in_existing_sse_bytes() -> None:
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ import litellm
|
|||
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
|
||||
|
||||
rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
|
||||
rust_bridge_runtime = importlib.import_module("litellm.rust_bridge.runtime")
|
||||
|
||||
|
||||
class SyncBridge:
|
||||
|
|
@ -77,7 +78,7 @@ async def test_enabled_async_bridge() -> None:
|
|||
|
||||
def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None)
|
||||
monkeypatch.setattr("litellm.rust_bridge.get_native_bridge", lambda: None)
|
||||
monkeypatch.setattr(rust_bridge_runtime, "get_native_bridge", lambda: None)
|
||||
assert rust_bridge.load_rust_transcription() is None
|
||||
assert rust_bridge.load_rust_atranscription() is None
|
||||
|
||||
|
|
@ -132,6 +133,11 @@ def test_bedrock_transcription_uses_rust_only_path() -> None:
|
|||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None)
|
||||
|
||||
assert response.text == "rust"
|
||||
assert response._hidden_params["core_engine"] == "rust"
|
||||
assert response._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -149,3 +155,8 @@ async def test_bedrock_atranscription_uses_rust_only_path() -> None:
|
|||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None)
|
||||
|
||||
assert response.text == "rust"
|
||||
assert response._hidden_params["core_engine"] == "rust"
|
||||
assert response._hidden_params["additional_headers"] == {
|
||||
"x-litellm-core": "rust",
|
||||
"x-litellm-rust": "true",
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue