mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge 152d041736 into e2dd92d18e
This commit is contained in:
commit
6cba73a2c1
50 changed files with 4231 additions and 3354 deletions
|
|
@ -1 +1 @@
|
|||
pub use crate::ocr::{OcrRequest, ocr};
|
||||
pub use crate::ocr::{OcrRequest, ocr, ocr_provider_supported};
|
||||
|
|
|
|||
|
|
@ -40,6 +40,11 @@ impl ResponsesWebSocketConnection {
|
|||
options: &RequestOptions,
|
||||
_context: &LiteLlmRequestContext,
|
||||
) -> Result<Self, Error> {
|
||||
if !litellm_core::responses::websocket::native_websocket_supported(
|
||||
options.custom_llm_provider.as_deref().unwrap_or("openai"),
|
||||
) {
|
||||
return Err(Error::Unsupported("unsupported native WebSocket provider"));
|
||||
}
|
||||
let headers = string_headers("Responses WebSocket", options.extra_headers.clone())?;
|
||||
let mut request = input
|
||||
.url
|
||||
|
|
|
|||
|
|
@ -35,6 +35,12 @@ pub async fn ocr(
|
|||
.await
|
||||
}
|
||||
|
||||
pub fn ocr_provider_supported(model: &str, provider: &str, request_format: Option<&str>) -> bool {
|
||||
common_utils::ocr_provider_config(provider, model).is_some_and(|config| {
|
||||
request_format != Some("native") || config.supported_ocr_params().contains(&"req_format")
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::integrations::types::RequestHooks;
|
||||
|
|
|
|||
|
|
@ -26,5 +26,9 @@ pub async fn audio_transcription(
|
|||
.await
|
||||
}
|
||||
|
||||
pub fn transcription_provider_supported(provider: &str) -> bool {
|
||||
prepare::provider_config(provider).is_some()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -9,7 +9,9 @@ use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderCo
|
|||
use super::types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
|
||||
pub(super) fn provider_config(
|
||||
provider: &str,
|
||||
) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
if provider == "bedrock" {
|
||||
return Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG);
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@
|
|||
//! calls the provider, and returns a typed OpenAI-shaped response.
|
||||
|
||||
use crate::Error;
|
||||
use crate::eligibility::native_route_decline;
|
||||
use crate::request_context::LiteLlmRequestContext;
|
||||
use crate::request_options::RequestOptions;
|
||||
mod client;
|
||||
|
|
@ -22,7 +23,8 @@ use serde_json::{Map, Value};
|
|||
|
||||
use handler::execute_chat_completions_provider_call;
|
||||
use prepare::{parse_messages, resolve_provider_config, resolve_request};
|
||||
use types::{ChatCompletionsRequest, ChatCompletionsResponse};
|
||||
use transformation::{ChatCompletionsProviderConfig, Unsupported};
|
||||
use types::{ChatCompletionsRequest, ChatCompletionsResponse, ChatMessage};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub async fn chat_completions(
|
||||
|
|
@ -45,8 +47,10 @@ pub fn chat_completions_decline_reason(
|
|||
custom_llm_provider: Option<&str>,
|
||||
messages: Value,
|
||||
optional_params: &Map<String, Value>,
|
||||
options: &RequestOptions,
|
||||
context: &LiteLlmRequestContext,
|
||||
) -> Option<&'static str> {
|
||||
let Ok((_, config)) = resolve_provider_config(model, custom_llm_provider) else {
|
||||
let Ok((_, provider, config)) = resolve_provider_config(model, custom_llm_provider) else {
|
||||
return Some("provider is not on the rust chat completions path");
|
||||
};
|
||||
let Ok(messages) = parse_messages(messages) else {
|
||||
|
|
@ -55,9 +59,41 @@ pub fn chat_completions_decline_reason(
|
|||
if messages.is_empty() {
|
||||
return Some("empty message list");
|
||||
}
|
||||
config
|
||||
.unsupported_reason(&messages, optional_params)
|
||||
.map(|reason| reason.0)
|
||||
unsupported_reason(
|
||||
provider,
|
||||
config,
|
||||
&messages,
|
||||
optional_params,
|
||||
options,
|
||||
context,
|
||||
)
|
||||
.map(|reason| reason.0)
|
||||
}
|
||||
|
||||
fn unsupported_reason(
|
||||
provider: &str,
|
||||
config: &dyn ChatCompletionsProviderConfig,
|
||||
messages: &[ChatMessage],
|
||||
optional_params: &Map<String, Value>,
|
||||
options: &RequestOptions,
|
||||
context: &LiteLlmRequestContext,
|
||||
) -> Option<Unsupported> {
|
||||
native_route_decline(true, &context.capabilities)
|
||||
.map(|reason| Unsupported(reason.reason()))
|
||||
.or_else(|| match provider {
|
||||
"anthropic" => options
|
||||
.anthropic
|
||||
.as_ref()
|
||||
.is_some_and(|anthropic| anthropic.has_user_id)
|
||||
.then_some(Unsupported("LiteLLM user metadata")),
|
||||
"bedrock" => options
|
||||
.bedrock
|
||||
.as_ref()
|
||||
.is_some_and(|bedrock| !bedrock.request_metadata_fields.is_empty())
|
||||
.then_some(Unsupported("LiteLLM request metadata forwarding")),
|
||||
_ => None,
|
||||
})
|
||||
.or_else(|| config.unsupported_reason(messages, optional_params))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
|
|||
|
|
@ -16,7 +16,7 @@ use super::types::{
|
|||
pub(super) fn resolve_provider_config<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> Result<(String, &'static dyn ChatCompletionsProviderConfig), Error> {
|
||||
) -> Result<(String, &'a str, &'static dyn ChatCompletionsProviderConfig), Error> {
|
||||
let provider_info = get_custom_llm_provider(model, custom_llm_provider)
|
||||
.or_else(|| {
|
||||
custom_llm_provider.map(|provider| CustomLlmProvider {
|
||||
|
|
@ -31,7 +31,11 @@ pub(super) fn resolve_provider_config<'a>(
|
|||
})?;
|
||||
let config = chat_completions_provider_config(provider_info.custom_llm_provider)
|
||||
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
|
||||
Ok((provider_info.model.to_string(), config))
|
||||
Ok((
|
||||
provider_info.model.to_string(),
|
||||
provider_info.custom_llm_provider,
|
||||
config,
|
||||
))
|
||||
}
|
||||
|
||||
pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error> {
|
||||
|
|
@ -42,9 +46,9 @@ pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error>
|
|||
pub(super) fn resolve_request(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
options: RequestOptions,
|
||||
_context: &LiteLlmRequestContext,
|
||||
context: &LiteLlmRequestContext,
|
||||
) -> Result<ResolvedChatCompletionsRequest, Error> {
|
||||
let (model, config) =
|
||||
let (model, provider, config) =
|
||||
resolve_provider_config(request.model, options.custom_llm_provider.as_deref())
|
||||
.map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?;
|
||||
let messages =
|
||||
|
|
@ -52,7 +56,14 @@ pub(super) fn resolve_request(
|
|||
if messages.is_empty() {
|
||||
return Err(Error::Declined("empty message list"));
|
||||
}
|
||||
if let Some(reason) = config.unsupported_reason(&messages, &request.optional_params) {
|
||||
if let Some(reason) = super::unsupported_reason(
|
||||
provider,
|
||||
config,
|
||||
&messages,
|
||||
&request.optional_params,
|
||||
&options,
|
||||
context,
|
||||
) {
|
||||
return Err(Error::Declined(reason.0));
|
||||
}
|
||||
Ok(ResolvedChatCompletionsRequest {
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
use crate::request_context::LiteLlmRequestContext;
|
||||
use crate::request_options::{BedrockOptions, RequestOptions};
|
||||
use crate::request_options::{AnthropicOptions, BedrockOptions, RequestOptions};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::error::Error;
|
||||
|
|
@ -513,7 +513,14 @@ fn decline_reason(
|
|||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
};
|
||||
super::chat_completions_decline_reason(model, provider, messages, ¶ms)
|
||||
super::chat_completions_decline_reason(
|
||||
model,
|
||||
provider,
|
||||
messages,
|
||||
¶ms,
|
||||
&RequestOptions::default(),
|
||||
&LiteLlmRequestContext::default(),
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -877,3 +884,74 @@ mod round_trip {
|
|||
));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn preflight_and_execution_share_provider_metadata_eligibility() {
|
||||
let messages = json!([{"role": "user", "content": "hi"}]);
|
||||
let cases = [
|
||||
(
|
||||
"anthropic",
|
||||
RequestOptions {
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
anthropic: Some(AnthropicOptions {
|
||||
user_id: Some("u-123".into()),
|
||||
has_user_id: true,
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
true,
|
||||
),
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
(
|
||||
"bedrock",
|
||||
RequestOptions {
|
||||
custom_llm_provider: Some("bedrock".into()),
|
||||
bedrock: Some(BedrockOptions {
|
||||
request_metadata_fields: vec!["user_api_key_team_id".into()],
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
true,
|
||||
),
|
||||
(
|
||||
"anthropic",
|
||||
RequestOptions {
|
||||
custom_llm_provider: Some("anthropic".into()),
|
||||
bedrock: Some(BedrockOptions {
|
||||
request_metadata_fields: vec!["user_api_key_team_id".into()],
|
||||
..Default::default()
|
||||
}),
|
||||
..Default::default()
|
||||
},
|
||||
false,
|
||||
),
|
||||
];
|
||||
for (provider, options, expected_decline) in cases {
|
||||
let context = LiteLlmRequestContext::default();
|
||||
let params = Map::new();
|
||||
let preflight = super::chat_completions_decline_reason(
|
||||
"claude-sonnet-4-5",
|
||||
Some(provider),
|
||||
messages.clone(),
|
||||
¶ms,
|
||||
&options,
|
||||
&context,
|
||||
);
|
||||
let execution = resolve_request(
|
||||
ChatCompletionsRequest {
|
||||
model: "claude-sonnet-4-5",
|
||||
messages: messages.clone(),
|
||||
optional_params: params,
|
||||
},
|
||||
options,
|
||||
&context,
|
||||
);
|
||||
assert_eq!(
|
||||
preflight.is_some(),
|
||||
expected_decline,
|
||||
"{provider} preflight"
|
||||
);
|
||||
assert_eq!(execution.is_err(), expected_decline, "{provider} execution");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
90
litellm-rust/crates/core/src/eligibility.rs
Normal file
90
litellm-rust/crates/core/src/eligibility.rs
Normal file
|
|
@ -0,0 +1,90 @@
|
|||
use crate::request_context::RequestCapabilities;
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum NativeRouteDecline {
|
||||
UnsupportedProvider,
|
||||
Streaming,
|
||||
AgenticHook,
|
||||
CustomClient,
|
||||
}
|
||||
|
||||
impl NativeRouteDecline {
|
||||
pub const fn reason(self) -> &'static str {
|
||||
match self {
|
||||
Self::UnsupportedProvider => "unsupported native provider",
|
||||
Self::Streaming => "native streaming is unavailable",
|
||||
Self::AgenticHook => "native agentic hooks are unavailable",
|
||||
Self::CustomClient => "native custom clients are unavailable",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn native_route_decline(
|
||||
provider_supported: bool,
|
||||
capabilities: &RequestCapabilities,
|
||||
) -> Option<NativeRouteDecline> {
|
||||
if !provider_supported {
|
||||
return Some(NativeRouteDecline::UnsupportedProvider);
|
||||
}
|
||||
if capabilities.stream {
|
||||
return Some(NativeRouteDecline::Streaming);
|
||||
}
|
||||
if capabilities.has_agentic_hook {
|
||||
return Some(NativeRouteDecline::AgenticHook);
|
||||
}
|
||||
if capabilities.has_custom_client {
|
||||
return Some(NativeRouteDecline::CustomClient);
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn admission_preserves_precedence_and_accepts_supported_unary_calls() {
|
||||
let all_unsupported = RequestCapabilities {
|
||||
stream: true,
|
||||
has_agentic_hook: true,
|
||||
has_custom_client: true,
|
||||
request_format: Some("native".into()),
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(
|
||||
native_route_decline(false, &all_unsupported),
|
||||
Some(NativeRouteDecline::UnsupportedProvider)
|
||||
);
|
||||
assert_eq!(
|
||||
native_route_decline(true, &all_unsupported),
|
||||
Some(NativeRouteDecline::Streaming)
|
||||
);
|
||||
assert_eq!(
|
||||
native_route_decline(true, &RequestCapabilities::default()),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn admission_rejects_each_unsupported_capability() {
|
||||
let cases = [
|
||||
(
|
||||
RequestCapabilities {
|
||||
has_agentic_hook: true,
|
||||
..Default::default()
|
||||
},
|
||||
NativeRouteDecline::AgenticHook,
|
||||
),
|
||||
(
|
||||
RequestCapabilities {
|
||||
has_custom_client: true,
|
||||
..Default::default()
|
||||
},
|
||||
NativeRouteDecline::CustomClient,
|
||||
),
|
||||
];
|
||||
for (capabilities, expected) in cases {
|
||||
assert_eq!(native_route_decline(true, &capabilities), Some(expected));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -3,6 +3,7 @@ pub mod caching;
|
|||
pub mod call_lifecycle;
|
||||
pub mod chat_completions;
|
||||
pub mod constants;
|
||||
pub mod eligibility;
|
||||
pub mod error;
|
||||
pub mod http_utils;
|
||||
pub mod messages;
|
||||
|
|
|
|||
|
|
@ -37,5 +37,9 @@ pub async fn messages_stream(
|
|||
execute_messages_provider_stream(request, options.clone()).await
|
||||
}
|
||||
|
||||
pub fn messages_provider_supported(provider: &str) -> bool {
|
||||
common_utils::messages_provider_config(provider).is_some()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
|
|
|||
|
|
@ -49,6 +49,7 @@ impl BedrockOptions {
|
|||
#[derive(Clone, Debug, Default)]
|
||||
pub struct AnthropicOptions {
|
||||
pub user_id: Option<String>,
|
||||
pub has_user_id: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default)]
|
||||
|
|
|
|||
|
|
@ -2,6 +2,14 @@ use crate::Error;
|
|||
use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH};
|
||||
use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult};
|
||||
|
||||
pub fn native_websocket_supported(provider: &str) -> bool {
|
||||
match provider {
|
||||
"openai" => crate::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG
|
||||
.supports_native_websocket(),
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
pub trait ResponsesWebSocketProviderConfig: Sync {
|
||||
fn supports_native_websocket(&self) -> bool {
|
||||
false
|
||||
|
|
|
|||
|
|
@ -35,8 +35,14 @@ impl ResponsesWebSocketConnection {
|
|||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let options: litellm_core::request_options::RequestOptions = options.into();
|
||||
let provider_supported = litellm_core::responses::websocket::native_websocket_supported(
|
||||
options.provider("openai"),
|
||||
);
|
||||
let context: litellm_core::request_context::LiteLlmRequestContext = context.into();
|
||||
if let Some(reason) = routes::definition::request_decline(provider_supported, &context) {
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
}
|
||||
let options: litellm_core::request_options::RequestOptions = options.into();
|
||||
let request = ResponsesWebSocketRequest { url: request.url };
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
let inner = RustResponsesWebSocketConnection::connect(request, &options, &context)
|
||||
|
|
@ -68,6 +74,20 @@ impl ResponsesWebSocketConnection {
|
|||
}
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (_model, custom_llm_provider, *, context))]
|
||||
fn responses_websocket_decline(
|
||||
_model: &str,
|
||||
custom_llm_provider: &str,
|
||||
context: NativeRequestContext,
|
||||
) -> Option<String> {
|
||||
let context: litellm_core::request_context::LiteLlmRequestContext = context.into();
|
||||
routes::definition::request_decline(
|
||||
litellm_core::responses::websocket::native_websocket_supported(custom_llm_provider),
|
||||
&context,
|
||||
)
|
||||
}
|
||||
|
||||
#[pymodule(gil_used = false)]
|
||||
mod _native {
|
||||
use pyo3::prelude::*;
|
||||
|
|
@ -77,6 +97,10 @@ mod _native {
|
|||
super::errors::register(module)?;
|
||||
super::routes::register(module)?;
|
||||
module.add_class::<super::ResponsesWebSocketConnection>()?;
|
||||
module.add_function(wrap_pyfunction!(
|
||||
super::responses_websocket_decline,
|
||||
module
|
||||
)?)?;
|
||||
super::diagnostics::register(module)
|
||||
}
|
||||
}
|
||||
|
|
@ -101,16 +125,20 @@ mod tests {
|
|||
let expected = [
|
||||
"RustBridgeDeclined",
|
||||
"RustUpstreamError",
|
||||
"ocr_decline",
|
||||
"ocr",
|
||||
"aocr",
|
||||
"transcription_decline",
|
||||
"transcription",
|
||||
"atranscription",
|
||||
"messages_decline",
|
||||
"messages",
|
||||
"amessages",
|
||||
"chat_completions_decline",
|
||||
"chat_completions",
|
||||
"achat_completions",
|
||||
"ResponsesWebSocketConnection",
|
||||
"responses_websocket_decline",
|
||||
"gil_stats",
|
||||
];
|
||||
|
||||
|
|
|
|||
|
|
@ -44,12 +44,14 @@ impl From<NativeBedrockOptions> for litellm_core::request_options::BedrockOption
|
|||
#[derive(FromPyObject)]
|
||||
struct NativeAnthropicOptions {
|
||||
user_id: Option<String>,
|
||||
has_user_id: bool,
|
||||
}
|
||||
|
||||
impl From<NativeAnthropicOptions> for litellm_core::request_options::AnthropicOptions {
|
||||
fn from(input: NativeAnthropicOptions) -> Self {
|
||||
Self {
|
||||
user_id: input.user_id,
|
||||
has_user_id: input.has_user_id,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -84,6 +86,12 @@ pub(crate) struct NativeRequestOptions {
|
|||
vertex: Option<NativeVertexOptions>,
|
||||
}
|
||||
|
||||
impl NativeRequestOptions {
|
||||
pub(crate) fn provider(&self, default: &'static str) -> &str {
|
||||
self.custom_llm_provider.as_deref().unwrap_or(default)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<NativeRequestOptions> for litellm_core::request_options::RequestOptions {
|
||||
fn from(input: NativeRequestOptions) -> Self {
|
||||
Self {
|
||||
|
|
|
|||
|
|
@ -22,7 +22,13 @@ fn prepare_transcription(
|
|||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
|
||||
let provider_supported = litellm_core::audio_transcription::transcription_provider_supported(
|
||||
options.provider("bedrock"),
|
||||
);
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
if let Some(reason) = super::definition::request_decline(provider_supported, &context) {
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
}
|
||||
let audio = input.audio;
|
||||
Ok(async move {
|
||||
run_route(
|
||||
|
|
@ -38,10 +44,25 @@ fn prepare_transcription(
|
|||
})
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (_model, custom_llm_provider, *, context))]
|
||||
fn transcription_decline(
|
||||
_model: &str,
|
||||
custom_llm_provider: &str,
|
||||
context: NativeRequestContext,
|
||||
) -> Option<String> {
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
super::definition::request_decline(
|
||||
litellm_core::audio_transcription::transcription_provider_supported(custom_llm_provider),
|
||||
&context,
|
||||
)
|
||||
}
|
||||
|
||||
bridge_route! {
|
||||
sync = transcription,
|
||||
asynchronous = atranscription,
|
||||
request = AudioTranscriptionInputs,
|
||||
prepare = prepare_transcription,
|
||||
errors = core_error_to_pyerr,
|
||||
extra = [transcription_decline],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ use litellm_core::chat_completions::chat_completions as run_route;
|
|||
use litellm_core::chat_completions::chat_completions_decline_reason;
|
||||
use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse};
|
||||
use litellm_core::request_context::LiteLlmRequestContext;
|
||||
use litellm_core::request_options::RequestOptions;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
use std::future::Future;
|
||||
|
|
@ -40,13 +41,21 @@ fn prepare_chat_completions(
|
|||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None))]
|
||||
#[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None, *, options, context))]
|
||||
#[allow(
|
||||
clippy::too_many_arguments,
|
||||
reason = "PyO3 preserves chat preflight inputs alongside separated options and context"
|
||||
)]
|
||||
fn chat_completions_decline(
|
||||
model: String,
|
||||
#[pyo3(from_py_with = litellm_python_interop::from_py)] messages: Value,
|
||||
#[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option<Value>,
|
||||
custom_llm_provider: Option<String>,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
) -> PyResult<Option<String>> {
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
let options: RequestOptions = options.into();
|
||||
let optional_params = match optional_params {
|
||||
None | Some(Value::Null) => Map::new(),
|
||||
Some(Value::Object(params)) => params,
|
||||
|
|
@ -61,6 +70,8 @@ fn chat_completions_decline(
|
|||
custom_llm_provider.as_deref(),
|
||||
messages,
|
||||
&optional_params,
|
||||
&options,
|
||||
&context,
|
||||
)
|
||||
.map(str::to_string))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -98,6 +98,14 @@ pub(super) fn add_function(
|
|||
module.add_function(function)
|
||||
}
|
||||
|
||||
pub(crate) fn request_decline(
|
||||
provider_supported: bool,
|
||||
context: &litellm_core::request_context::LiteLlmRequestContext,
|
||||
) -> Option<String> {
|
||||
litellm_core::eligibility::native_route_decline(provider_supported, &context.capabilities)
|
||||
.map(|reason| reason.reason().to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::ffi::CString;
|
||||
|
|
@ -262,6 +270,81 @@ for field in ('litellm_call_id', 'trace_id', 'request_model'):
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn acceptance_and_execution_decline_unsupported_requests_before_io() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = PyModule::new(py, "routes").expect("module should be created");
|
||||
crate::routes::register(&module).expect("routes should register");
|
||||
module
|
||||
.add_class::<crate::ResponsesWebSocketConnection>()
|
||||
.unwrap();
|
||||
module
|
||||
.add_function(
|
||||
wrap_pyfunction!(crate::responses_websocket_decline, &module).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
let locals = crate::marshal::request_fixtures(py);
|
||||
locals.set_item("routes", module).unwrap();
|
||||
py.run(
|
||||
c"
|
||||
for route, provider in (
|
||||
('messages', 'anthropic'),
|
||||
('transcription', 'bedrock'),
|
||||
('ocr', 'mistral'),
|
||||
('responses_websocket', 'openai'),
|
||||
):
|
||||
decline = getattr(routes, route + '_decline')
|
||||
assert decline('model', provider, context=context) is None, route
|
||||
for flag in ('stream', 'has_agentic_hook', 'has_custom_client'):
|
||||
flagged_context = replace(
|
||||
context,
|
||||
capabilities=replace(context.capabilities, **{flag: True}),
|
||||
)
|
||||
assert decline('model', provider, context=flagged_context) is not None, (route, flag)
|
||||
reason = decline('model', 'unsupported-native-provider', context=context)
|
||||
assert reason is not None, route
|
||||
request = Request(
|
||||
messages=[], body={}, audio={}, document={}, optional_params={},
|
||||
url='invalid-url-must-not-be-used',
|
||||
)
|
||||
unsupported_options = Options(custom_llm_provider='unsupported-native-provider')
|
||||
functions = (
|
||||
(routes.ResponsesWebSocketConnection.connect,)
|
||||
if route == 'responses_websocket'
|
||||
else (getattr(routes, route), getattr(routes, 'a' + route))
|
||||
)
|
||||
for execute in functions:
|
||||
try:
|
||||
execute(request, options=unsupported_options, context=context)
|
||||
except Exception as error:
|
||||
assert type(error).__name__ == 'RustBridgeDeclined', (route, error)
|
||||
assert str(error) == reason, (route, reason, error)
|
||||
else:
|
||||
raise AssertionError('unsupported request reached provider execution')
|
||||
native_context = replace(
|
||||
context,
|
||||
capabilities=replace(context.capabilities, request_format='native'),
|
||||
)
|
||||
litellm_context = replace(
|
||||
context,
|
||||
capabilities=replace(context.capabilities, request_format='litellm'),
|
||||
)
|
||||
assert routes.ocr_decline('model', 'mistral', context=native_context) is not None
|
||||
assert routes.ocr_decline('model', 'mistral', context=litellm_context) is None
|
||||
assert routes.ocr_decline(
|
||||
'doc-intelligence/prebuilt-layout',
|
||||
'azure_ai',
|
||||
context=native_context,
|
||||
) is None
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.expect("acceptance must match execution eligibility without I/O");
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_routes_execute_sync_and_async_contracts() {
|
||||
Python::initialize();
|
||||
|
|
|
|||
|
|
@ -20,7 +20,12 @@ fn prepare_messages(
|
|||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
) -> PyResult<impl Future<Output = Result<AnthropicMessagesResponse, Error>> + Send + 'static> {
|
||||
let provider_supported =
|
||||
litellm_core::messages::messages_provider_supported(options.provider("anthropic"));
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
if let Some(reason) = super::definition::request_decline(provider_supported, &context) {
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
}
|
||||
let body = required_value("body", input.body, Value::is_object, "dict")?;
|
||||
Ok(async move {
|
||||
run_route(
|
||||
|
|
@ -35,10 +40,25 @@ fn prepare_messages(
|
|||
})
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (_model, custom_llm_provider, *, context))]
|
||||
fn messages_decline(
|
||||
_model: &str,
|
||||
custom_llm_provider: &str,
|
||||
context: NativeRequestContext,
|
||||
) -> Option<String> {
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
super::definition::request_decline(
|
||||
litellm_core::messages::messages_provider_supported(custom_llm_provider),
|
||||
&context,
|
||||
)
|
||||
}
|
||||
|
||||
bridge_route! {
|
||||
sync = messages,
|
||||
asynchronous = amessages,
|
||||
request = MessagesInputs,
|
||||
prepare = prepare_messages,
|
||||
errors = core_error_to_pyerr,
|
||||
extra = [messages_decline],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use pyo3::prelude::*;
|
||||
|
||||
#[macro_use]
|
||||
mod definition;
|
||||
pub(crate) mod definition;
|
||||
|
||||
#[cfg(feature = "trace-parity")]
|
||||
mod gateway_messages;
|
||||
|
|
|
|||
|
|
@ -24,6 +24,14 @@ fn prepare_ocr(
|
|||
context: NativeRequestContext,
|
||||
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
let provider_supported = litellm_ai_gateway::io::ocr::ocr_provider_supported(
|
||||
&input.model,
|
||||
options.provider("mistral"),
|
||||
context.capabilities.request_format.as_deref(),
|
||||
);
|
||||
if let Some(reason) = super::definition::request_decline(provider_supported, &context) {
|
||||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
}
|
||||
let document = input.document;
|
||||
Ok(async move {
|
||||
run_route(
|
||||
|
|
@ -43,10 +51,27 @@ fn prepare_ocr(
|
|||
})
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (model, custom_llm_provider, *, context))]
|
||||
fn ocr_decline(
|
||||
model: &str,
|
||||
custom_llm_provider: &str,
|
||||
context: NativeRequestContext,
|
||||
) -> Option<String> {
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
let provider_supported = litellm_ai_gateway::io::ocr::ocr_provider_supported(
|
||||
model,
|
||||
custom_llm_provider,
|
||||
context.capabilities.request_format.as_deref(),
|
||||
);
|
||||
super::definition::request_decline(provider_supported, &context)
|
||||
}
|
||||
|
||||
bridge_route! {
|
||||
sync = ocr,
|
||||
asynchronous = aocr,
|
||||
request = OcrInputs,
|
||||
prepare = prepare_ocr,
|
||||
errors = ocr_error_to_pyerr,
|
||||
extra = [ocr_decline],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -10,7 +10,8 @@ import subprocess
|
|||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence
|
||||
from collections.abc import Awaitable, Callable, Generator, Iterator, Mapping, Sequence
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime as dt_object
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType, TracebackType
|
||||
|
|
@ -459,6 +460,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
log_raw_request_response: bool = False,
|
||||
supports_correlation_logging: bool = True,
|
||||
):
|
||||
self._suppress_next_pre_call: bool = False
|
||||
_input: Final[str | None] = messages # save original value of messages
|
||||
if messages is not None:
|
||||
if isinstance(messages, str):
|
||||
|
|
@ -1188,7 +1190,19 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
additional_args.get("api_base", "")
|
||||
)
|
||||
|
||||
@contextmanager
|
||||
def suppress_next_pre_call(self) -> Generator[None]:
|
||||
previous: Final = self._suppress_next_pre_call
|
||||
self._suppress_next_pre_call = True
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self._suppress_next_pre_call = previous
|
||||
|
||||
def pre_call(self, input, api_key, model=None, additional_args={}):
|
||||
if self._suppress_next_pre_call:
|
||||
self._suppress_next_pre_call = False
|
||||
return
|
||||
# Log the exact input to the LLM API
|
||||
try:
|
||||
self._pre_call(
|
||||
|
|
|
|||
|
|
@ -25,11 +25,6 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
_get_httpx_client,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
|
||||
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
|
||||
from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors
|
||||
from litellm.rust_bridge.request import anthropic_options, request_context
|
||||
from litellm.rust_bridge.runtime import DispatchResult
|
||||
from litellm.types.llms.anthropic import (
|
||||
ContentBlockDelta,
|
||||
ContentBlockStart,
|
||||
|
|
@ -371,7 +366,12 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
if config is None:
|
||||
raise ValueError(f"Provider config not found for model: {model} and provider: {custom_llm_provider}")
|
||||
|
||||
def prepare_python() -> tuple[dict[str, str], dict[str, object]]: # mutable-ok: stream mutates data
|
||||
def build_request() -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream
|
||||
"""Translate the request the Python way, returning `(headers, data)`.
|
||||
|
||||
The pair stays mutable because the streaming path rewrites it in
|
||||
place (`data["stream"] = True`) before sending.
|
||||
"""
|
||||
request_data: Final = config.transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
|
|
@ -379,119 +379,32 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
python_headers, data = update_request_with_filtered_beta(
|
||||
return update_request_with_filtered_beta(
|
||||
headers=headers,
|
||||
request_data=request_data,
|
||||
provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
## LOGGING
|
||||
# Reaching here with `serves_via_rust` set means the Rust attempt
|
||||
# declined at call time, before the provider was called, and already
|
||||
# logged this request. That is the same attempt continuing.
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": python_headers,
|
||||
},
|
||||
)
|
||||
print_verbose(f"_is_function_call: {_is_function_call}")
|
||||
return python_headers, data
|
||||
headers, data = build_request()
|
||||
|
||||
# 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
|
||||
# `max_tokens` among them) that `transform_request` would have applied.
|
||||
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
|
||||
**AnthropicConfig.get_config(model=model),
|
||||
**optional_params,
|
||||
}
|
||||
serves_via_rust: Final = rust_chat_completions_accepts(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
stream=stream,
|
||||
)
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
}
|
||||
if serves_via_rust:
|
||||
logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args)
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
## LOGGING
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
additional_args=rust_logging_args,
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": api_base,
|
||||
"headers": headers,
|
||||
},
|
||||
)
|
||||
rust_context: Final = request_context(
|
||||
logging_obj=logging_obj,
|
||||
request_model=logging_obj.model,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
def native_completion() -> DispatchResult[ModelResponse]:
|
||||
return rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
anthropic=anthropic_options(litellm_params),
|
||||
stream=bool(stream),
|
||||
has_custom_client=client is not None,
|
||||
eligible=serves_via_rust,
|
||||
context=rust_context,
|
||||
)
|
||||
|
||||
async def native_acompletion() -> DispatchResult[ModelResponse]:
|
||||
return await rust_chat_completions_bridge.achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
anthropic=anthropic_options(litellm_params),
|
||||
stream=bool(stream),
|
||||
has_custom_client=client is not None,
|
||||
eligible=serves_via_rust,
|
||||
context=rust_context,
|
||||
)
|
||||
|
||||
@anative_first(
|
||||
native=native_acompletion,
|
||||
route="chat_completions",
|
||||
errors=lambda: provider_errors(custom_llm_provider or "", model),
|
||||
)
|
||||
async def execute_async() -> ModelResponse | CustomStreamWrapper:
|
||||
headers, data = prepare_python()
|
||||
print_verbose(f"_is_function_call: {_is_function_call}")
|
||||
if acompletion is True:
|
||||
if (
|
||||
stream is True
|
||||
): # if function call - fake the streaming (need complete blocks for output parsing in openai format)
|
||||
print_verbose("makes async anthropic streaming POST request")
|
||||
data["stream"] = stream
|
||||
return await self.acompletion_stream_function(
|
||||
return self.acompletion_stream_function(
|
||||
model=model,
|
||||
messages=messages,
|
||||
data=data,
|
||||
|
|
@ -513,7 +426,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None),
|
||||
)
|
||||
else:
|
||||
return await self.acompletion_function(
|
||||
return self.acompletion_function(
|
||||
model=model,
|
||||
messages=messages,
|
||||
data=data,
|
||||
|
|
@ -535,14 +448,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
json_mode=json_mode,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
@native_first(
|
||||
native=native_completion,
|
||||
route="chat_completions",
|
||||
errors=lambda: provider_errors(custom_llm_provider or "", model),
|
||||
)
|
||||
def execute_sync() -> ModelResponse | CustomStreamWrapper:
|
||||
headers, data = prepare_python()
|
||||
else:
|
||||
## COMPLETION CALL
|
||||
if (
|
||||
stream is True
|
||||
|
|
@ -574,12 +480,11 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
)
|
||||
|
||||
else:
|
||||
python_client: Final = (
|
||||
client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout})
|
||||
)
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
client = _get_httpx_client(params={"timeout": timeout})
|
||||
|
||||
try:
|
||||
response: Final = python_client.post(
|
||||
response: Final = client.post(
|
||||
api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
|
|
@ -600,21 +505,20 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
status_code=status_code,
|
||||
headers=error_headers,
|
||||
)
|
||||
return config.transform_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
model_response=model_response,
|
||||
logging_obj=logging_obj,
|
||||
api_key=api_key,
|
||||
request_data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
encoding=encoding,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
|
||||
return execute_async() if acompletion else execute_sync()
|
||||
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,
|
||||
)
|
||||
|
||||
def embedding(self):
|
||||
# logic for parsing in - calling - parsing out model embedding calls
|
||||
|
|
|
|||
|
|
@ -315,6 +315,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def _resolved_provider(self) -> str:
|
||||
return self.custom_llm_provider or "anthropic"
|
||||
|
||||
@classmethod
|
||||
def get_config_for_model(cls, model: str) -> dict[str, object]:
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
return TypeAdapter(dict[str, object]).validate_python(cls.get_config(model=model))
|
||||
|
||||
@classmethod
|
||||
def get_config(cls, *, model: str | None = None):
|
||||
config: Final = super().get_config()
|
||||
|
|
|
|||
|
|
@ -559,27 +559,54 @@ def anthropic_messages_handler(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
anthropic_messages_provider_config: BaseAnthropicMessagesConfig | None = None
|
||||
def python_fallback():
|
||||
anthropic_messages_provider_config: BaseAnthropicMessagesConfig | None = None
|
||||
|
||||
if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]:
|
||||
anthropic_messages_provider_config = ProviderConfigManager.get_provider_anthropic_messages_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
if anthropic_messages_provider_config is None and _deployment_passes_through_anthropic_messages(
|
||||
kwargs.get("model_info")
|
||||
):
|
||||
from litellm.llms.openai_like.messages.transformation import (
|
||||
OpenAILikeAnthropicMessagesConfig,
|
||||
)
|
||||
if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]:
|
||||
anthropic_messages_provider_config = ProviderConfigManager.get_provider_anthropic_messages_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
)
|
||||
if anthropic_messages_provider_config is None and _deployment_passes_through_anthropic_messages(
|
||||
kwargs.get("model_info")
|
||||
):
|
||||
from litellm.llms.openai_like.messages.transformation import (
|
||||
OpenAILikeAnthropicMessagesConfig,
|
||||
)
|
||||
|
||||
anthropic_messages_provider_config = OpenAILikeAnthropicMessagesConfig(
|
||||
cache_control_ttl=_deployment_supports_cache_control_ttl(kwargs.get("model_info")),
|
||||
)
|
||||
if anthropic_messages_provider_config is None:
|
||||
# Route to Responses API for OpenAI / Azure, chat/completions for everything else.
|
||||
if _should_route_to_responses_api(custom_llm_provider, original_model, model):
|
||||
return LiteLLMMessagesToResponsesAPIHandler.anthropic_messages_handler(
|
||||
anthropic_messages_provider_config = OpenAILikeAnthropicMessagesConfig(
|
||||
cache_control_ttl=_deployment_supports_cache_control_ttl(kwargs.get("model_info")),
|
||||
)
|
||||
if anthropic_messages_provider_config is None:
|
||||
# Route to Responses API for OpenAI / Azure, chat/completions for everything else.
|
||||
if _should_route_to_responses_api(custom_llm_provider, original_model, model):
|
||||
return LiteLLMMessagesToResponsesAPIHandler.anthropic_messages_handler(
|
||||
max_tokens=max_tokens,
|
||||
messages=messages,
|
||||
model=original_model,
|
||||
metadata=metadata,
|
||||
stop_sequences=stop_sequences,
|
||||
stream=stream,
|
||||
system=system,
|
||||
temperature=temperature,
|
||||
thinking=thinking,
|
||||
tool_choice=tool_choice,
|
||||
tools=tools,
|
||||
top_k=top_k,
|
||||
top_p=top_p,
|
||||
_is_async=is_async,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
client=client,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
# The in-gateway context_management polyfill runs inside
|
||||
# ``async_anthropic_messages_handler`` so it can ``await`` the
|
||||
# summarization model for ``compact_20260112``. ``context_management``
|
||||
# is passed through as a regular kwarg.
|
||||
return LiteLLMMessagesToCompletionTransformationHandler.anthropic_messages_handler(
|
||||
max_tokens=max_tokens,
|
||||
messages=messages,
|
||||
model=original_model,
|
||||
|
|
@ -601,66 +628,87 @@ def anthropic_messages_handler(
|
|||
**kwargs,
|
||||
)
|
||||
|
||||
# The in-gateway context_management polyfill runs inside
|
||||
# ``async_anthropic_messages_handler`` so it can ``await`` the
|
||||
# summarization model for ``compact_20260112``. ``context_management``
|
||||
# is passed through as a regular kwarg.
|
||||
return LiteLLMMessagesToCompletionTransformationHandler.anthropic_messages_handler(
|
||||
max_tokens=max_tokens,
|
||||
messages=messages,
|
||||
model=original_model,
|
||||
metadata=metadata,
|
||||
stop_sequences=stop_sequences,
|
||||
stream=stream,
|
||||
system=system,
|
||||
temperature=temperature,
|
||||
thinking=thinking,
|
||||
tool_choice=tool_choice,
|
||||
tools=tools,
|
||||
top_k=top_k,
|
||||
top_p=top_p,
|
||||
if custom_llm_provider is None:
|
||||
raise ValueError(
|
||||
f"custom_llm_provider is required for Anthropic messages, passed in model={model}, custom_llm_provider={custom_llm_provider}"
|
||||
)
|
||||
|
||||
local_vars.update(kwargs)
|
||||
anthropic_messages_optional_request_params: Final = (
|
||||
AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
|
||||
params=local_vars,
|
||||
model=model,
|
||||
drop_params=litellm_params.get("drop_params") is True,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
)
|
||||
if is_reasoning_auto_summary_enabled():
|
||||
thinking_param: Final = anthropic_messages_optional_request_params.get("thinking")
|
||||
if isinstance(thinking_param, dict) and thinking_param.get("type") != "disabled":
|
||||
anthropic_messages_optional_request_params["thinking"] = {
|
||||
**thinking_param,
|
||||
"display": "summarized",
|
||||
}
|
||||
|
||||
return base_llm_http_handler.anthropic_messages_handler(
|
||||
model=model,
|
||||
messages=strip_provider_specific_fields_from_anthropic_messages(messages),
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=dict(anthropic_messages_optional_request_params),
|
||||
_is_async=is_async,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
client=client,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
stream=stream,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
if custom_llm_provider is None:
|
||||
raise ValueError(
|
||||
f"custom_llm_provider is required for Anthropic messages, passed in model={model}, custom_llm_provider={custom_llm_provider}"
|
||||
)
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
local_vars.update(kwargs)
|
||||
anthropic_messages_optional_request_params: Final = (
|
||||
AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
|
||||
params=local_vars,
|
||||
from litellm.rust_bridge.messages import dispatch_messages
|
||||
|
||||
return dispatch_messages(
|
||||
model=model,
|
||||
provider=custom_llm_provider,
|
||||
messages=TypeAdapter(list[dict[str, object]]).validate_python(
|
||||
strip_provider_specific_fields_from_anthropic_messages(messages)
|
||||
),
|
||||
body=lambda: _native_messages_body(
|
||||
params={**local_vars, **kwargs},
|
||||
model=model,
|
||||
drop_params=litellm_params.get("drop_params") is True,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
provider=custom_llm_provider,
|
||||
),
|
||||
params=litellm_params,
|
||||
logging=litellm_logging_obj,
|
||||
api_key=api_key or dynamic_api_key,
|
||||
api_base=api_base or dynamic_api_base,
|
||||
stream=bool(stream),
|
||||
asynchronous=bool(is_async),
|
||||
has_custom_client=client is not None,
|
||||
fallback=python_fallback,
|
||||
)
|
||||
|
||||
|
||||
def _native_messages_body(params: dict[str, object], model: str, drop_params: bool, provider: str) -> dict[str, object]:
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
requested: Final = TypeAdapter(dict[str, object]).validate_python(
|
||||
AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
|
||||
params=params,
|
||||
model=model,
|
||||
drop_params=drop_params,
|
||||
custom_llm_provider=provider,
|
||||
)
|
||||
)
|
||||
if is_reasoning_auto_summary_enabled():
|
||||
thinking_param: Final = anthropic_messages_optional_request_params.get("thinking")
|
||||
if isinstance(thinking_param, dict) and thinking_param.get("type") != "disabled":
|
||||
anthropic_messages_optional_request_params["thinking"] = {
|
||||
**thinking_param,
|
||||
"display": "summarized",
|
||||
}
|
||||
|
||||
return base_llm_http_handler.anthropic_messages_handler(
|
||||
model=model,
|
||||
messages=strip_provider_specific_fields_from_anthropic_messages(messages),
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=dict(anthropic_messages_optional_request_params),
|
||||
_is_async=is_async,
|
||||
client=client,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
litellm_params=litellm_params,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
stream=stream,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
thinking_param: Final = requested.get("thinking")
|
||||
if (
|
||||
is_reasoning_auto_summary_enabled()
|
||||
and isinstance(thinking_param, dict)
|
||||
and thinking_param.get("type") != "disabled"
|
||||
):
|
||||
return {**requested, "thinking": {**thinking_param, "display": "summarized"}}
|
||||
return requested
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from collections.abc import AsyncIterator, Iterator
|
|||
from typing import TYPE_CHECKING, Any, Final, Union
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
|
||||
from litellm.constants import DEFAULT_MAX_TOKENS, RESPONSE_FORMAT_TOOL_NAME
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
|
|
@ -70,6 +70,10 @@ class BaseConfig(ABC):
|
|||
def __init__(self):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
def get_config_for_model(cls, model: str) -> dict[str, object]:
|
||||
return TypeAdapter(dict[str, object]).validate_python(cls.get_config())
|
||||
|
||||
@classmethod
|
||||
def get_config(cls):
|
||||
return {
|
||||
|
|
|
|||
|
|
@ -1,38 +1,14 @@
|
|||
import base64
|
||||
from io import IOBase
|
||||
from typing import Final, NoReturn
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.rust_bridge import transcription as rust_transcription_bridge
|
||||
from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors
|
||||
from litellm.rust_bridge.request import request_context
|
||||
from litellm.rust_bridge.runtime import DispatchResult, adapt_result
|
||||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
|
||||
def _unavailable() -> NoReturn:
|
||||
raise RuntimeError("Rust audio transcription bridge is unavailable")
|
||||
|
||||
|
||||
async def _aunavailable() -> NoReturn:
|
||||
_unavailable()
|
||||
|
||||
|
||||
class BedrockAudioTranscriptionRustDispatch:
|
||||
@staticmethod
|
||||
def _input_source_kind(audio_file: FileTypes) -> str:
|
||||
content: Final = audio_file[1] if isinstance(audio_file, tuple) else audio_file
|
||||
if isinstance(content, (bytes, bytearray, memoryview)):
|
||||
return "bytes"
|
||||
if isinstance(content, IOBase):
|
||||
return "file"
|
||||
if isinstance(content, str):
|
||||
return "path"
|
||||
return "opaque"
|
||||
|
||||
@staticmethod
|
||||
def _audio_payload(audio_file: FileTypes) -> dict[str, object]:
|
||||
processed_audio: Final = process_audio_file(audio_file)
|
||||
|
|
@ -55,44 +31,6 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
"filename": processed_audio.filename,
|
||||
}
|
||||
|
||||
def _attempt_audio_transcriptions(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
audio_file: FileTypes,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
logging_obj: Logging | None = None,
|
||||
) -> DispatchResult[TranscriptionResponse]:
|
||||
result: Final = rust_transcription_bridge.transcription(
|
||||
model=model,
|
||||
audio=self._audio_payload(audio_file),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
input_source_kind=self._input_source_kind(audio_file),
|
||||
context=request_context(
|
||||
logging_obj=logging_obj,
|
||||
request_model=logging_obj.model if logging_obj is not None else model,
|
||||
litellm_params=logging_obj.litellm_params if logging_obj is not None else None,
|
||||
),
|
||||
)
|
||||
return adapt_result(result, lambda response: TranscriptionResponse(**response))
|
||||
|
||||
@native_first(
|
||||
native=_attempt_audio_transcriptions,
|
||||
route="audio transcription",
|
||||
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout, logging_obj=None: (
|
||||
provider_errors(custom_llm_provider, model)
|
||||
),
|
||||
)
|
||||
def audio_transcriptions(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -104,24 +42,8 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
logging_obj: Logging | None = None,
|
||||
) -> TranscriptionResponse:
|
||||
_unavailable()
|
||||
|
||||
async def _attempt_async_audio_transcriptions(
|
||||
self,
|
||||
*,
|
||||
model: str,
|
||||
audio_file: FileTypes,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
logging_obj: Logging | None = None,
|
||||
) -> DispatchResult[TranscriptionResponse]:
|
||||
result: Final = await rust_transcription_bridge.atranscription(
|
||||
rust_response: Final = rust_transcription_bridge.transcription(
|
||||
model=model,
|
||||
audio=self._audio_payload(audio_file),
|
||||
api_key=api_key,
|
||||
|
|
@ -130,22 +52,11 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
input_source_kind=self._input_source_kind(audio_file),
|
||||
context=request_context(
|
||||
logging_obj=logging_obj,
|
||||
request_model=logging_obj.model if logging_obj is not None else model,
|
||||
litellm_params=logging_obj.litellm_params if logging_obj is not None else None,
|
||||
),
|
||||
)
|
||||
return adapt_result(result, lambda response: TranscriptionResponse(**response))
|
||||
if rust_response is None:
|
||||
raise RuntimeError("Rust audio transcription bridge is unavailable")
|
||||
return TranscriptionResponse(**rust_response)
|
||||
|
||||
@anative_first(
|
||||
native=_attempt_async_audio_transcriptions,
|
||||
route="audio transcription",
|
||||
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout, logging_obj=None: (
|
||||
provider_errors(custom_llm_provider, model)
|
||||
),
|
||||
)
|
||||
async def async_audio_transcriptions(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -157,6 +68,17 @@ class BedrockAudioTranscriptionRustDispatch:
|
|||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
logging_obj: Logging | None = None,
|
||||
) -> TranscriptionResponse:
|
||||
await _aunavailable()
|
||||
rust_response: Final = await rust_transcription_bridge.atranscription(
|
||||
model=model,
|
||||
audio=self._audio_payload(audio_file),
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_response is None:
|
||||
raise RuntimeError("Rust audio transcription bridge is unavailable")
|
||||
return TranscriptionResponse(**rust_response)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,4 @@
|
|||
import json
|
||||
from collections.abc import Mapping
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Final
|
||||
|
||||
import httpx
|
||||
|
|
@ -16,11 +14,6 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
_get_httpx_client,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
|
||||
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
|
||||
from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors
|
||||
from litellm.rust_bridge.request import bedrock_options, request_context
|
||||
from litellm.rust_bridge.runtime import DispatchResult
|
||||
from litellm.types.utils import ModelResponse
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
|
|
@ -29,22 +22,6 @@ from ..common_utils import BedrockError, _get_all_bedrock_regions
|
|||
from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call
|
||||
|
||||
|
||||
def _sigv4_principal(credentials: Credentials | None) -> Mapping[str, str]:
|
||||
if credentials is None:
|
||||
return MappingProxyType({})
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: value
|
||||
for key, value in (
|
||||
("aws_access_key_id", credentials.access_key),
|
||||
("aws_secret_access_key", credentials.secret_key),
|
||||
("aws_session_token", credentials.token),
|
||||
)
|
||||
if value is not None
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def make_sync_call(
|
||||
client: HTTPHandler | None,
|
||||
api_base: str,
|
||||
|
|
@ -215,9 +192,6 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
)
|
||||
|
||||
## LOGGING
|
||||
# The Rust path already logged this request's pre_call before handing
|
||||
# it here, and it only declines before the provider is called, so this
|
||||
# is the same attempt continuing rather than a second one.
|
||||
if not skip_pre_call_logging:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
|
|
@ -390,94 +364,12 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
# Filter beta headers in HTTP headers before making the request
|
||||
headers = update_headers_with_filtered_beta(headers=headers, provider="bedrock_converse")
|
||||
|
||||
# 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
|
||||
# resolved so both paths sign as the same principal. Bearer-token auth
|
||||
# resolves no SigV4 principal at all, and each path reads that token
|
||||
# itself.
|
||||
rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy
|
||||
**optional_params,
|
||||
**_sigv4_principal(credentials),
|
||||
"aws_region_name": aws_region_name,
|
||||
}
|
||||
serves_via_rust: Final = rust_chat_completions_accepts(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
custom_llm_provider="bedrock",
|
||||
litellm_params=litellm_params,
|
||||
stream=stream,
|
||||
)
|
||||
rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict
|
||||
"complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent
|
||||
"messages": messages,
|
||||
**optional_params,
|
||||
},
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": headers,
|
||||
}
|
||||
if serves_via_rust:
|
||||
logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args)
|
||||
log_rust_post_call: Final = rust_chat_completions_bridge.response_logger(
|
||||
logging_obj=logging_obj,
|
||||
messages=messages,
|
||||
api_key="",
|
||||
additional_args=rust_logging_args,
|
||||
)
|
||||
rust_context: Final = request_context(
|
||||
logging_obj=logging_obj,
|
||||
request_model=logging_obj.model,
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
def native_completion() -> DispatchResult[ModelResponse]:
|
||||
return rust_chat_completions_bridge.chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
bedrock=bedrock_options(rust_optional_params),
|
||||
stream=bool(stream),
|
||||
has_custom_client=client is not None,
|
||||
eligible=serves_via_rust,
|
||||
context=rust_context,
|
||||
)
|
||||
|
||||
async def native_acompletion() -> DispatchResult[ModelResponse]:
|
||||
return await rust_chat_completions_bridge.achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=rust_optional_params,
|
||||
model_response=model_response,
|
||||
api_key=api_key,
|
||||
api_base=proxy_endpoint_url,
|
||||
custom_llm_provider="bedrock",
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
on_response=log_rust_post_call,
|
||||
bedrock=bedrock_options(rust_optional_params),
|
||||
stream=bool(stream),
|
||||
has_custom_client=client is not None,
|
||||
eligible=serves_via_rust,
|
||||
context=rust_context,
|
||||
)
|
||||
|
||||
@anative_first(
|
||||
native=native_acompletion,
|
||||
route="chat_completions",
|
||||
errors=lambda: provider_errors("bedrock", model),
|
||||
)
|
||||
async def execute_async() -> ModelResponse | CustomStreamWrapper:
|
||||
python_client: Final = None if isinstance(client, HTTPHandler) else client
|
||||
### ROUTING (ASYNC, STREAMING, SYNC)
|
||||
if acompletion:
|
||||
if isinstance(client, HTTPHandler):
|
||||
client = None
|
||||
if stream is True:
|
||||
return await self.async_streaming(
|
||||
return self.async_streaming(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=proxy_endpoint_url,
|
||||
|
|
@ -490,7 +382,7 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
logger_fn=logger_fn,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
client=python_client,
|
||||
client=client,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
credentials=credentials,
|
||||
|
|
@ -498,7 +390,7 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
### ASYNC COMPLETION
|
||||
return await self.async_completion(
|
||||
return self.async_completion(
|
||||
model=model,
|
||||
messages=messages,
|
||||
api_base=proxy_endpoint_url,
|
||||
|
|
@ -511,112 +403,100 @@ class BedrockConverseLLM(BaseAWSLLM):
|
|||
logger_fn=logger_fn,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
client=python_client,
|
||||
client=client,
|
||||
credentials=credentials,
|
||||
api_key=api_key,
|
||||
skip_pre_call_logging=serves_via_rust,
|
||||
)
|
||||
|
||||
@native_first(
|
||||
native=native_completion,
|
||||
route="chat_completions",
|
||||
errors=lambda: provider_errors("bedrock", model),
|
||||
## TRANSFORMATION ##
|
||||
|
||||
_data: Final = litellm.AmazonConverseConfig()._transform_request(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=extra_headers,
|
||||
)
|
||||
def execute_sync() -> ModelResponse | CustomStreamWrapper:
|
||||
## TRANSFORMATION ##
|
||||
data: Final = json.dumps(_data)
|
||||
|
||||
_data: Final = litellm.AmazonConverseConfig()._transform_request(
|
||||
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
|
||||
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, int)):
|
||||
timeout = httpx.Timeout(timeout)
|
||||
_params["timeout"] = timeout
|
||||
client = _get_httpx_client(_params)
|
||||
|
||||
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,
|
||||
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.
|
||||
if not serves_via_rust:
|
||||
logging_obj.pre_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
additional_args={
|
||||
"complete_input_dict": data,
|
||||
"api_base": proxy_endpoint_url,
|
||||
"headers": prepped.headers,
|
||||
},
|
||||
)
|
||||
resolved_timeout: Final = httpx.Timeout(timeout) if isinstance(timeout, (float, int)) else timeout
|
||||
python_client: Final = (
|
||||
_get_httpx_client({"timeout": resolved_timeout} if resolved_timeout is not None else None)
|
||||
if client is None or isinstance(client, AsyncHTTPHandler)
|
||||
else client
|
||||
)
|
||||
|
||||
if stream is not None and stream is True:
|
||||
completion_stream, response_headers = make_sync_call(
|
||||
client=python_client,
|
||||
api_base=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
model=model,
|
||||
messages=messages,
|
||||
logging_obj=logging_obj,
|
||||
json_mode=json_mode,
|
||||
fake_stream=fake_stream,
|
||||
stream_chunk_size=stream_chunk_size,
|
||||
)
|
||||
streaming_response: Final = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
logging_obj=logging_obj,
|
||||
_response_headers=response_headers,
|
||||
)
|
||||
|
||||
return streaming_response
|
||||
|
||||
### COMPLETION
|
||||
|
||||
try:
|
||||
response: Final = python_client.post(
|
||||
url=proxy_endpoint_url,
|
||||
headers=prepped.headers,
|
||||
data=data,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
response.raise_for_status()
|
||||
except httpx.HTTPStatusError as err:
|
||||
error_code: Final = err.response.status_code
|
||||
raise BedrockError(status_code=error_code, message=err.response.text)
|
||||
except httpx.TimeoutException:
|
||||
raise BedrockError(status_code=408, message="Timeout error occurred.")
|
||||
|
||||
sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response(
|
||||
model=model,
|
||||
response=response,
|
||||
model_response=model_response,
|
||||
stream=stream if isinstance(stream, bool) else False,
|
||||
logging_obj=logging_obj,
|
||||
api_key="",
|
||||
data=data,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
encoding=encoding,
|
||||
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,
|
||||
)
|
||||
sync_transformed_response.set_provider_response_headers(response.headers)
|
||||
return sync_transformed_response
|
||||
|
||||
return execute_async() if acompletion else execute_sync()
|
||||
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
|
||||
|
|
|
|||
|
|
@ -1,8 +1,7 @@
|
|||
import asyncio
|
||||
import json
|
||||
import ssl
|
||||
from collections.abc import AsyncGenerator, AsyncIterator, Coroutine, Iterator, Mapping, Sequence
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence
|
||||
from functools import lru_cache
|
||||
from types import MappingProxyType, ModuleType
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, TypeVar, Union, cast, get_type_hints
|
||||
|
|
@ -92,8 +91,6 @@ from litellm.responses.streaming_iterator import (
|
|||
ResponsesWebSocketStreaming,
|
||||
SyncResponsesAPIStreamingIterator,
|
||||
)
|
||||
from litellm.rust_bridge.dispatch import anative_context, anative_first, provider_errors
|
||||
from litellm.rust_bridge.runtime import DispatchResult, NativeSkipped, NativeSkipReason, adapt_result
|
||||
from litellm.types.containers.main import (
|
||||
ContainerFileListResponse,
|
||||
ContainerListResponse,
|
||||
|
|
@ -162,15 +159,6 @@ from litellm.utils import (
|
|||
async_pre_call_deployment_hook,
|
||||
)
|
||||
|
||||
|
||||
def _rust_responses_websocket_enabled(
|
||||
custom_llm_provider: str | None,
|
||||
) -> bool:
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
|
||||
return custom_llm_provider == "openai" and rust_enabled()
|
||||
|
||||
|
||||
from .http_handler import get_shared_realtime_ssl_context
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -183,9 +171,6 @@ if TYPE_CHECKING:
|
|||
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
|
||||
FakeAnthropicMessagesStreamIterator,
|
||||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
AnthropicMessagesStreamingResponse,
|
||||
)
|
||||
from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig
|
||||
from litellm.types.llms.openai_evals import (
|
||||
CancelEvalResponse,
|
||||
|
|
@ -245,7 +230,7 @@ def _responses_api_optional_request_param_names() -> frozenset[str]:
|
|||
return frozenset(get_type_hints(ResponsesAPIOptionalRequestParams).keys())
|
||||
|
||||
|
||||
def _custom_logger_callbacks(logging_obj: LiteLLMLoggingObj) -> list["CustomLogger"]:
|
||||
def _custom_logger_callbacks(logging_obj: LiteLLMLoggingObj | None) -> list["CustomLogger"]:
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_custom_logger_compatible_class,
|
||||
|
|
@ -2055,7 +2040,7 @@ class BaseLLMHTTPHandler:
|
|||
raise RuntimeError("unreachable: anthropic messages HTTP retry loop exited without return")
|
||||
|
||||
@staticmethod
|
||||
def _resolve_anthropic_messages_timeout(
|
||||
def resolve_anthropic_messages_timeout(
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
stream: bool,
|
||||
custom_llm_provider: str,
|
||||
|
|
@ -2227,116 +2212,86 @@ class BaseLLMHTTPHandler:
|
|||
},
|
||||
)
|
||||
|
||||
async def native_messages() -> DispatchResult[AnthropicMessagesResponse | AsyncIterator[object]]:
|
||||
result: Final = await self._attempt_rust_anthropic_messages(
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
|
||||
async_httpx_client=async_httpx_client,
|
||||
request_url=request_url,
|
||||
headers=headers,
|
||||
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
|
||||
request_body=request_body,
|
||||
stream=stream or False,
|
||||
logging_obj=logging_obj,
|
||||
provider_config=anthropic_messages_provider_config,
|
||||
litellm_params=litellm_params,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
timeout=self.resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
|
||||
stream=bool(stream),
|
||||
has_custom_client=client is not None,
|
||||
model=model,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=headers,
|
||||
request_body=request_body,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
return adapt_result(result, self._rust_anthropic_messages_fake_stream) if stream else result
|
||||
|
||||
@anative_first(
|
||||
native=native_messages, route="messages", errors=lambda: provider_errors(custom_llm_provider, model)
|
||||
)
|
||||
async def execute_messages() -> AnthropicMessagesResponse | AsyncIterator[object]:
|
||||
response: Final = await self._async_post_anthropic_messages_with_http_error_retry(
|
||||
async_httpx_client=async_httpx_client,
|
||||
request_url=request_url,
|
||||
headers=headers,
|
||||
signed_json_body=(signed_json_body if signed_json_body is not None else request_body_json),
|
||||
request_body=request_body,
|
||||
stream=stream or False,
|
||||
logging_obj=logging_obj,
|
||||
provider_config=anthropic_messages_provider_config,
|
||||
litellm_params=litellm_params,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
timeout=self._resolve_anthropic_messages_timeout(
|
||||
litellm_params=litellm_params,
|
||||
stream=stream or False,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
),
|
||||
)
|
||||
|
||||
# used for logging + cost tracking
|
||||
logging_obj.model_call_details["httpx_response"] = response
|
||||
|
||||
initial_response: AsyncIterator | AnthropicMessagesResponse
|
||||
if stream:
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
AnthropicMessagesStreamingResponse,
|
||||
anthropic_messages_stream_hidden_params,
|
||||
)
|
||||
|
||||
# used for logging + cost tracking
|
||||
logging_obj.model_call_details["httpx_response"] = response
|
||||
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
|
||||
model=model,
|
||||
httpx_response=response,
|
||||
request_body=request_body,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
|
||||
|
||||
initial_response: AsyncIterator | AnthropicMessagesResponse
|
||||
if stream:
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
AnthropicMessagesStreamingResponse,
|
||||
anthropic_messages_stream_hidden_params,
|
||||
)
|
||||
|
||||
completion_stream: Final = anthropic_messages_provider_config.get_async_streaming_response_iterator(
|
||||
model=model,
|
||||
httpx_response=response,
|
||||
request_body=request_body,
|
||||
litellm_logging_obj=logging_obj,
|
||||
)
|
||||
stream_hidden_params: Final = anthropic_messages_stream_hidden_params(response.headers)
|
||||
|
||||
if not self._has_agentic_completion_hook(logging_obj):
|
||||
# No callback overrides async_should_run_agentic_loop, so the
|
||||
# agentic wrapper's only effect would be buffering every chunk
|
||||
# and rebuilding the response from SSE at end-of-stream to call
|
||||
# hooks that all return (False, {}). Stream through directly and
|
||||
# skip that per-chunk + end-of-stream overhead.
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=completion_stream,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
AgenticAnthropicStreamingIterator,
|
||||
)
|
||||
|
||||
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
|
||||
logging_obj=logging_obj,
|
||||
tools=anthropic_messages_optional_request_params.get("tools"),
|
||||
)
|
||||
initial_response = AgenticAnthropicStreamingIterator(
|
||||
completion_stream=completion_stream,
|
||||
http_handler=self,
|
||||
model=model,
|
||||
messages=messages,
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
|
||||
hold_back=bool(held_back_tool_names),
|
||||
server_fulfilled_tool_names=held_back_tool_names,
|
||||
)
|
||||
if not self.has_agentic_completion_hook(logging_obj):
|
||||
# No callback overrides async_should_run_agentic_loop, so the
|
||||
# agentic wrapper's only effect would be buffering every chunk
|
||||
# and rebuilding the response from SSE at end-of-stream to call
|
||||
# hooks that all return (False, {}). Stream through directly and
|
||||
# skip that per-chunk + end-of-stream overhead.
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=initial_response,
|
||||
completion_stream=completion_stream,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
else:
|
||||
initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
return initial_response
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.agentic_streaming_iterator import (
|
||||
AgenticAnthropicStreamingIterator,
|
||||
)
|
||||
|
||||
held_back_tool_names: Final = self._server_fulfilled_tools_in_request(
|
||||
logging_obj=logging_obj,
|
||||
tools=anthropic_messages_optional_request_params.get("tools"),
|
||||
)
|
||||
initial_response = AgenticAnthropicStreamingIterator(
|
||||
completion_stream=completion_stream,
|
||||
http_handler=self,
|
||||
model=model,
|
||||
messages=messages,
|
||||
anthropic_messages_provider_config=anthropic_messages_provider_config,
|
||||
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
|
||||
logging_obj=logging_obj,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
kwargs={**kwargs, "api_key": api_key} if api_key else kwargs,
|
||||
hold_back=bool(held_back_tool_names),
|
||||
server_fulfilled_tool_names=held_back_tool_names,
|
||||
)
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=initial_response,
|
||||
hidden_params=stream_hidden_params,
|
||||
)
|
||||
else:
|
||||
initial_response = anthropic_messages_provider_config.transform_anthropic_messages_response(
|
||||
model=model,
|
||||
raw_response=response,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
initial_response: Final = await execute_messages()
|
||||
if stream:
|
||||
return initial_response
|
||||
return await self._finalize_anthropic_messages_response(
|
||||
initial_response=initial_response,
|
||||
model=model,
|
||||
|
|
@ -2385,80 +2340,6 @@ class BaseLLMHTTPHandler:
|
|||
"anthropic_messages",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
async def _attempt_rust_anthropic_messages(
|
||||
*,
|
||||
custom_llm_provider: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
has_agentic_hook: bool,
|
||||
stream: bool,
|
||||
has_custom_client: bool,
|
||||
model: str,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
headers: dict,
|
||||
request_body: dict,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
logging_obj: LiteLLMLoggingObj | None = None,
|
||||
) -> DispatchResult[AnthropicMessagesResponse]:
|
||||
if custom_llm_provider not in ("azure_ai", "anthropic"):
|
||||
return NativeSkipped(NativeSkipReason.INELIGIBLE)
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
|
||||
if not rust_enabled():
|
||||
return NativeSkipped(NativeSkipReason.DISABLED)
|
||||
if has_agentic_hook:
|
||||
return NativeSkipped(NativeSkipReason.INELIGIBLE)
|
||||
|
||||
from litellm.rust_bridge import messages as rust_messages_bridge
|
||||
from litellm.rust_bridge.request import request_context
|
||||
|
||||
upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"}
|
||||
result: Final = await rust_messages_bridge.amessages(
|
||||
model=model,
|
||||
body=upstream_body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
stream=stream,
|
||||
has_custom_client=has_custom_client,
|
||||
has_agentic_hook=has_agentic_hook,
|
||||
context=request_context(
|
||||
logging_obj=logging_obj,
|
||||
request_model=logging_obj.model if logging_obj is not None else model,
|
||||
litellm_params=litellm_params.model_dump(),
|
||||
),
|
||||
)
|
||||
|
||||
def adapt(rust_response: dict[str, object]) -> AnthropicMessagesResponse:
|
||||
return cast(
|
||||
AnthropicMessagesResponse,
|
||||
{**rust_response, "_hidden_params": {"additional_headers": {"x-litellm-rust": "true"}}},
|
||||
)
|
||||
|
||||
return adapt_result(result, adapt)
|
||||
|
||||
@staticmethod
|
||||
def _rust_anthropic_messages_fake_stream(
|
||||
rust_response: AnthropicMessagesResponse,
|
||||
) -> "AnthropicMessagesStreamingResponse":
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
|
||||
FakeAnthropicMessagesStreamIterator,
|
||||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
AnthropicMessagesStreamHiddenParams,
|
||||
AnthropicMessagesStreamingResponse,
|
||||
)
|
||||
|
||||
completion_stream = cast(AsyncIterator[bytes], FakeAnthropicMessagesStreamIterator(response=rust_response))
|
||||
hidden_params: Final = AnthropicMessagesStreamHiddenParams(additional_headers={"x-litellm-rust": "true"})
|
||||
return AnthropicMessagesStreamingResponse(
|
||||
completion_stream=completion_stream,
|
||||
hidden_params=hidden_params,
|
||||
)
|
||||
|
||||
def anthropic_messages_handler(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -2767,7 +2648,7 @@ class BaseLLMHTTPHandler:
|
|||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
if self._has_agentic_completion_hook(logging_obj):
|
||||
if self.has_agentic_completion_hook(logging_obj):
|
||||
agentic_kwargs: Final = dict(litellm_params) # mutable-ok: agentic hooks mutate kwargs in place
|
||||
final_response: Final = run_async_function(
|
||||
self._call_agentic_completion_hooks,
|
||||
|
|
@ -5158,7 +5039,7 @@ class BaseLLMHTTPHandler:
|
|||
return depth, max_loops, fingerprints
|
||||
|
||||
@staticmethod
|
||||
def _has_agentic_completion_hook(logging_obj: LiteLLMLoggingObj) -> bool:
|
||||
def has_agentic_completion_hook(logging_obj: LiteLLMLoggingObj | None) -> bool:
|
||||
"""
|
||||
True if any registered callback actually overrides
|
||||
``async_should_run_agentic_loop`` (the gate every agentic hook goes
|
||||
|
|
@ -6516,41 +6397,27 @@ class BaseLLMHTTPHandler:
|
|||
},
|
||||
)
|
||||
|
||||
from litellm.rust_bridge import responses_websocket as rust_responses_websocket
|
||||
from litellm.rust_bridge.request import request_context
|
||||
from litellm.rust_bridge.responses_websocket import open_connection
|
||||
|
||||
async def attempt_connection() -> DispatchResult[
|
||||
AbstractAsyncContextManager[rust_responses_websocket.ConnectionAdapter]
|
||||
]:
|
||||
if not _rust_responses_websocket_enabled(custom_llm_provider):
|
||||
return NativeSkipped(NativeSkipReason.INELIGIBLE)
|
||||
return await rust_responses_websocket.managed_connect(
|
||||
url=ws_url,
|
||||
headers={str(key): str(value) for key, value in headers.items()},
|
||||
timeout=timeout,
|
||||
context=request_context(
|
||||
logging_obj=logging_obj,
|
||||
request_model=logging_obj.model,
|
||||
litellm_params=litellm_params.model_dump(),
|
||||
),
|
||||
)
|
||||
|
||||
@anative_context(
|
||||
native=attempt_connection,
|
||||
route="responses_websocket",
|
||||
errors=lambda: provider_errors("openai", "responses websocket"),
|
||||
)
|
||||
@asynccontextmanager
|
||||
async def _backend_connection() -> AsyncGenerator[ClientConnection, None]:
|
||||
async with websockets.connect(
|
||||
async with open_connection(
|
||||
url=ws_url,
|
||||
headers={str(key): str(value) for key, value in headers.items()},
|
||||
timeout=timeout,
|
||||
model=model,
|
||||
provider=custom_llm_provider,
|
||||
context=request_context(
|
||||
logging_obj=logging_obj,
|
||||
request_model=logging_obj.model,
|
||||
litellm_params=litellm_params.model_dump(),
|
||||
),
|
||||
fallback=lambda: websockets.connect(
|
||||
ws_url,
|
||||
additional_headers=headers,
|
||||
max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES,
|
||||
ssl=ssl_context,
|
||||
) as backend:
|
||||
yield backend
|
||||
|
||||
async with _backend_connection() as backend_ws:
|
||||
),
|
||||
) as backend_ws:
|
||||
_request_data: Final[dict[str, object]] = {}
|
||||
if litellm_metadata:
|
||||
_request_data["litellm_metadata"] = litellm_metadata
|
||||
|
|
|
|||
694
litellm/main.py
694
litellm/main.py
|
|
@ -19,7 +19,7 @@ import random
|
|||
import sys
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import AsyncIterator, Coroutine, Iterable, Mapping, Sequence
|
||||
from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Mapping, Sequence
|
||||
from concurrent import futures
|
||||
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
|
||||
from copy import deepcopy
|
||||
|
|
@ -4958,6 +4958,183 @@ def _complete_langflow(ctx: _CompletionDispatchContext) -> _CompletionDispatchRe
|
|||
)
|
||||
|
||||
|
||||
_PYTHON_PRIMARY_COMPLETION_HANDLERS: Final[
|
||||
Mapping[str, Callable[[_CompletionDispatchContext], _CompletionDispatchResult]]
|
||||
] = MappingProxyType(
|
||||
{
|
||||
"azure": _complete_azure,
|
||||
"azure_text": _complete_azure_text,
|
||||
"deepseek": _complete_deepseek,
|
||||
"azure_ai": _complete_azure_ai,
|
||||
}
|
||||
)
|
||||
|
||||
_PYTHON_COMPATIBLE_COMPLETION_HANDLERS: Final[
|
||||
Mapping[str, Callable[[_CompletionDispatchContext], _CompletionDispatchResult]]
|
||||
] = MappingProxyType(
|
||||
{
|
||||
"fireworks_ai": _complete_fireworks_ai,
|
||||
"together_ai": _complete_together_ai,
|
||||
"heroku": _complete_heroku,
|
||||
"ragflow": _complete_ragflow,
|
||||
"xai": _complete_xai,
|
||||
"groq": _complete_groq,
|
||||
"bedrock_mantle": _complete_bedrock_mantle,
|
||||
"a2a": _complete_a2a,
|
||||
"gigachat": _complete_gigachat,
|
||||
"sap": _complete_sap,
|
||||
"aiohttp_openai": _complete_aiohttp_openai,
|
||||
"cometapi": _complete_cometapi,
|
||||
"minimax": _complete_minimax,
|
||||
"hosted_vllm": _complete_hosted_vllm,
|
||||
}
|
||||
)
|
||||
|
||||
_PYTHON_LEGACY_COMPLETION_HANDLERS: Final[
|
||||
Mapping[str, Callable[[_CompletionDispatchContext], _CompletionDispatchResult]]
|
||||
] = MappingProxyType(
|
||||
{
|
||||
"anthropic_text": _complete_anthropic_text,
|
||||
"anthropic": _complete_anthropic,
|
||||
"nlp_cloud": _complete_nlp_cloud,
|
||||
"aleph_alpha": _complete_aleph_alpha,
|
||||
"cohere_chat": _complete_cohere_chat,
|
||||
"cohere": _complete_cohere_chat,
|
||||
"maritalk": _complete_maritalk,
|
||||
"amazon_nova": _complete_amazon_nova,
|
||||
"huggingface": _complete_huggingface,
|
||||
"oci": _complete_oci,
|
||||
"compactifai": _complete_compactifai,
|
||||
"oobabooga": _complete_oobabooga,
|
||||
"databricks": _complete_databricks,
|
||||
"datarobot": _complete_datarobot,
|
||||
"openrouter": _complete_openrouter,
|
||||
"vercel_ai_gateway": _complete_vercel_ai_gateway,
|
||||
}
|
||||
)
|
||||
|
||||
_PYTHON_EXTENDED_COMPLETION_HANDLERS: Final[
|
||||
Mapping[str, Callable[[_CompletionDispatchContext], _CompletionDispatchResult]]
|
||||
] = MappingProxyType(
|
||||
{
|
||||
"vertex_ai_beta": _complete_vertex_ai_beta,
|
||||
"gemini": _complete_vertex_ai_beta,
|
||||
"vertex_ai": _complete_vertex_ai,
|
||||
"predibase": _complete_predibase,
|
||||
"text-completion-codestral": _complete_text_completion_codestral,
|
||||
"text-completion-inception": _complete_text_completion_inception,
|
||||
"sagemaker_chat": _complete_sagemaker_chat,
|
||||
"sagemaker_nova": _complete_sagemaker_chat,
|
||||
"sagemaker": _complete_sagemaker,
|
||||
"bedrock": _complete_bedrock,
|
||||
"watsonx": _complete_watsonx,
|
||||
"watsonx_text": _complete_watsonx_text,
|
||||
"vllm": _complete_vllm,
|
||||
"ollama": _complete_ollama,
|
||||
"ollama_chat": _complete_ollama_chat,
|
||||
"triton": _complete_triton,
|
||||
"cloudflare": _complete_cloudflare,
|
||||
}
|
||||
)
|
||||
|
||||
_PYTHON_ADDITIONAL_COMPLETION_HANDLERS: Final[
|
||||
Mapping[str, Callable[[_CompletionDispatchContext], _CompletionDispatchResult]]
|
||||
] = MappingProxyType(
|
||||
{
|
||||
"gradient_ai": _complete_gradient_ai,
|
||||
"gdc": _complete_gdc,
|
||||
"bytez": _complete_bytez,
|
||||
"lemonade": _complete_lemonade,
|
||||
}
|
||||
)
|
||||
|
||||
_PYTHON_CUSTOM_COMPLETION_HANDLERS: Final[
|
||||
Mapping[str, Callable[[_CompletionDispatchContext], _CompletionDispatchResult]]
|
||||
] = MappingProxyType(
|
||||
{
|
||||
"langgraph": _complete_langgraph,
|
||||
"langflow": _complete_langflow,
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _complete_python(_dispatch_ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
|
||||
custom_llm_provider: Final = _dispatch_ctx.custom_llm_provider
|
||||
model: Final = _dispatch_ctx.model
|
||||
kwargs: Final = _dispatch_ctx.kwargs
|
||||
if (primary_handler := _PYTHON_PRIMARY_COMPLETION_HANDLERS.get(custom_llm_provider)) is not None:
|
||||
return primary_handler(_dispatch_ctx)
|
||||
elif (
|
||||
custom_llm_provider == "text-completion-openai"
|
||||
or "ft:babbage-002" in model
|
||||
or "ft:davinci-002" in model
|
||||
or (
|
||||
custom_llm_provider in litellm.openai_text_completion_compatible_providers
|
||||
and kwargs.get("text_completion") is True
|
||||
)
|
||||
):
|
||||
return _complete_text_completion_openai(_dispatch_ctx)
|
||||
elif (compatible_handler := _PYTHON_COMPATIBLE_COMPLETION_HANDLERS.get(custom_llm_provider)) is not None:
|
||||
return compatible_handler(_dispatch_ctx)
|
||||
elif (
|
||||
model in litellm.open_ai_chat_completion_models
|
||||
and custom_llm_provider in (None, "openai")
|
||||
or custom_llm_provider == "custom_openai"
|
||||
or custom_llm_provider == "deepinfra"
|
||||
or (custom_llm_provider == "perplexity")
|
||||
or (custom_llm_provider == "nvidia_nim")
|
||||
or (custom_llm_provider == "cerebras")
|
||||
or (custom_llm_provider == "baseten")
|
||||
or (custom_llm_provider == "sambanova")
|
||||
or (custom_llm_provider == "volcengine")
|
||||
or (custom_llm_provider == "anyscale")
|
||||
or (custom_llm_provider == "openai")
|
||||
or (custom_llm_provider == "nebius")
|
||||
or (custom_llm_provider == "wandb")
|
||||
or (custom_llm_provider == "clarifai")
|
||||
or (custom_llm_provider in litellm.openai_compatible_providers)
|
||||
or JSONProviderRegistry.exists(custom_llm_provider)
|
||||
or ("ft:gpt-3.5-turbo" in model)
|
||||
):
|
||||
return _complete_custom_openai(_dispatch_ctx)
|
||||
return _complete_python_provider(_dispatch_ctx)
|
||||
|
||||
|
||||
def _complete_python_provider(_dispatch_ctx: _CompletionDispatchContext) -> _CompletionDispatchResult:
|
||||
custom_llm_provider: Final = _dispatch_ctx.custom_llm_provider
|
||||
model: Final = _dispatch_ctx.model
|
||||
if custom_llm_provider == "mistral":
|
||||
return _complete_mistral(_dispatch_ctx)
|
||||
elif "replicate" in model or custom_llm_provider == "replicate" or model in litellm.replicate_models:
|
||||
return _complete_replicate(_dispatch_ctx)
|
||||
elif "clarifai" in model or custom_llm_provider == "clarifai" or model in litellm.clarifai_models:
|
||||
return _complete_custom_openai(_dispatch_ctx)
|
||||
elif (legacy_handler := _PYTHON_LEGACY_COMPLETION_HANDLERS.get(custom_llm_provider)) is not None:
|
||||
return legacy_handler(_dispatch_ctx)
|
||||
elif custom_llm_provider == "palm":
|
||||
raise ValueError(
|
||||
"Palm was decommisioned on October 2024. Please use the `gemini/` route for Gemini Google AI Studio Models. Announcement: https://ai.google.dev/palm_docs/palm?hl=en"
|
||||
)
|
||||
elif (extended_handler := _PYTHON_EXTENDED_COMPLETION_HANDLERS.get(custom_llm_provider)) is not None:
|
||||
return extended_handler(_dispatch_ctx)
|
||||
elif custom_llm_provider == "petals" or model in litellm.petals_models:
|
||||
return _complete_petals(_dispatch_ctx)
|
||||
elif custom_llm_provider == "snowflake" or model in litellm.snowflake_models:
|
||||
return _complete_snowflake(_dispatch_ctx)
|
||||
elif (additional_handler := _PYTHON_ADDITIONAL_COMPLETION_HANDLERS.get(custom_llm_provider)) is not None:
|
||||
return additional_handler(_dispatch_ctx)
|
||||
elif custom_llm_provider == "ovhcloud" or model in litellm.ovhcloud_models:
|
||||
return _complete_ovhcloud(_dispatch_ctx)
|
||||
elif custom_llm_provider == "custom":
|
||||
return _complete_custom(_dispatch_ctx)
|
||||
elif custom_llm_provider in litellm._custom_providers:
|
||||
return _complete_custom_providers(_dispatch_ctx)
|
||||
elif (custom_handler := _PYTHON_CUSTOM_COMPLETION_HANDLERS.get(custom_llm_provider)) is not None:
|
||||
return custom_handler(_dispatch_ctx)
|
||||
else:
|
||||
raise LiteLLMUnknownProvider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
|
||||
@tracer.wrap()
|
||||
@client
|
||||
def completion(
|
||||
|
|
@ -5640,209 +5817,9 @@ def completion(
|
|||
timeout=timeout,
|
||||
top_p=top_p,
|
||||
)
|
||||
if custom_llm_provider == "azure":
|
||||
# azure configs
|
||||
## check dynamic params ##
|
||||
response = _complete_azure(_dispatch_ctx)
|
||||
elif custom_llm_provider == "azure_text":
|
||||
# azure configs
|
||||
response = _complete_azure_text(_dispatch_ctx)
|
||||
elif custom_llm_provider == "deepseek":
|
||||
## COMPLETION CALL
|
||||
from litellm.rust_bridge.chat_completions import dispatch_completion
|
||||
|
||||
response = _complete_deepseek(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "azure_ai":
|
||||
response = _complete_azure_ai(_dispatch_ctx)
|
||||
elif (
|
||||
custom_llm_provider == "text-completion-openai"
|
||||
or "ft:babbage-002" in model
|
||||
or "ft:davinci-002" in model # support for finetuned completion models
|
||||
or custom_llm_provider in litellm.openai_text_completion_compatible_providers
|
||||
and kwargs.get("text_completion") is True
|
||||
):
|
||||
response = _complete_text_completion_openai(_dispatch_ctx)
|
||||
elif custom_llm_provider == "fireworks_ai":
|
||||
## COMPLETION CALL
|
||||
response = _complete_fireworks_ai(_dispatch_ctx)
|
||||
elif custom_llm_provider == "together_ai":
|
||||
response = _complete_together_ai(_dispatch_ctx)
|
||||
elif custom_llm_provider == "heroku":
|
||||
response = _complete_heroku(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "ragflow":
|
||||
## COMPLETION CALL - RAGFlow uses HTTP handler to support custom URL paths
|
||||
response = _complete_ragflow(_dispatch_ctx)
|
||||
elif custom_llm_provider == "xai":
|
||||
## COMPLETION CALL
|
||||
response = _complete_xai(_dispatch_ctx)
|
||||
elif custom_llm_provider == "groq":
|
||||
response = _complete_groq(_dispatch_ctx)
|
||||
elif custom_llm_provider == "bedrock_mantle":
|
||||
response = _complete_bedrock_mantle(_dispatch_ctx)
|
||||
elif custom_llm_provider == "a2a":
|
||||
# A2A (Agent-to-Agent) Protocol
|
||||
# Resolve agent configuration from registry if model format is "a2a/<agent-name>"
|
||||
response = _complete_a2a(_dispatch_ctx)
|
||||
elif custom_llm_provider == "gigachat":
|
||||
# GigaChat - Sber AI's LLM (Russia)
|
||||
response = _complete_gigachat(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "sap":
|
||||
response = _complete_sap(_dispatch_ctx)
|
||||
elif custom_llm_provider == "aiohttp_openai":
|
||||
# NEW aiohttp provider for 10-100x higher RPS
|
||||
response = _complete_aiohttp_openai(_dispatch_ctx)
|
||||
elif custom_llm_provider == "cometapi":
|
||||
response = _complete_cometapi(_dispatch_ctx)
|
||||
elif custom_llm_provider == "minimax":
|
||||
response = _complete_minimax(_dispatch_ctx)
|
||||
elif custom_llm_provider == "hosted_vllm":
|
||||
response = _complete_hosted_vllm(_dispatch_ctx)
|
||||
elif (
|
||||
# A known OpenAI model name only decides the route when nothing else
|
||||
# resolved a provider. get_llm_provider() already maps these names to
|
||||
# "openai", so a different value here was asked for explicitly (or came
|
||||
# from a register_model entry), and the provider config built for it
|
||||
# would be handed to the OpenAI handler.
|
||||
(model in litellm.open_ai_chat_completion_models and custom_llm_provider in (None, "openai"))
|
||||
or custom_llm_provider == "custom_openai"
|
||||
or custom_llm_provider == "deepinfra"
|
||||
or custom_llm_provider == "perplexity"
|
||||
or custom_llm_provider == "nvidia_nim"
|
||||
or custom_llm_provider == "cerebras"
|
||||
or custom_llm_provider == "baseten"
|
||||
or custom_llm_provider == "sambanova"
|
||||
or custom_llm_provider == "volcengine"
|
||||
or custom_llm_provider == "anyscale"
|
||||
or custom_llm_provider == "openai"
|
||||
or custom_llm_provider == "nebius"
|
||||
or custom_llm_provider == "wandb"
|
||||
or custom_llm_provider == "clarifai"
|
||||
or custom_llm_provider in litellm.openai_compatible_providers
|
||||
or JSONProviderRegistry.exists(custom_llm_provider) # JSON-configured providers
|
||||
or "ft:gpt-3.5-turbo" in model # finetune gpt-3.5-turbo
|
||||
): # allow user to make an openai call with a custom base
|
||||
# note: if a user sets a custom base - we should ensure this works
|
||||
# allow for the setting of dynamic and stateful api-bases
|
||||
response = _complete_custom_openai(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "mistral":
|
||||
response = _complete_mistral(_dispatch_ctx)
|
||||
elif "replicate" in model or custom_llm_provider == "replicate" or model in litellm.replicate_models:
|
||||
# Setting the relevant API KEY for replicate, replicate defaults to using os.environ.get("REPLICATE_API_TOKEN")
|
||||
response = _complete_replicate(_dispatch_ctx)
|
||||
elif "clarifai" in model or custom_llm_provider == "clarifai" or model in litellm.clarifai_models:
|
||||
pass # Deprecated - handled in the openai compatible provider section above
|
||||
elif custom_llm_provider == "anthropic_text":
|
||||
response = _complete_anthropic_text(_dispatch_ctx)
|
||||
elif custom_llm_provider == "anthropic":
|
||||
response = _complete_anthropic(_dispatch_ctx)
|
||||
elif custom_llm_provider == "nlp_cloud":
|
||||
response = _complete_nlp_cloud(_dispatch_ctx)
|
||||
elif custom_llm_provider == "aleph_alpha":
|
||||
response = _complete_aleph_alpha(_dispatch_ctx)
|
||||
elif custom_llm_provider == "cohere_chat" or custom_llm_provider == "cohere":
|
||||
response = _complete_cohere_chat(_dispatch_ctx)
|
||||
elif custom_llm_provider == "maritalk":
|
||||
response = _complete_maritalk(_dispatch_ctx)
|
||||
elif custom_llm_provider == "amazon_nova":
|
||||
response = _complete_amazon_nova(_dispatch_ctx)
|
||||
elif custom_llm_provider == "huggingface":
|
||||
response = _complete_huggingface(_dispatch_ctx)
|
||||
elif custom_llm_provider == "oci":
|
||||
response = _complete_oci(_dispatch_ctx)
|
||||
elif custom_llm_provider == "compactifai":
|
||||
response = _complete_compactifai(_dispatch_ctx)
|
||||
elif custom_llm_provider == "oobabooga":
|
||||
response = _complete_oobabooga(_dispatch_ctx)
|
||||
elif custom_llm_provider == "databricks":
|
||||
response = _complete_databricks(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "datarobot":
|
||||
response = _complete_datarobot(_dispatch_ctx)
|
||||
elif custom_llm_provider == "openrouter":
|
||||
response = _complete_openrouter(_dispatch_ctx)
|
||||
elif custom_llm_provider == "vercel_ai_gateway":
|
||||
response = _complete_vercel_ai_gateway(_dispatch_ctx)
|
||||
elif custom_llm_provider == "palm":
|
||||
raise ValueError(
|
||||
"Palm was decommisioned on October 2024. Please use the `gemini/` route for Gemini Google AI Studio Models. Announcement: https://ai.google.dev/palm_docs/palm?hl=en"
|
||||
)
|
||||
elif custom_llm_provider == "vertex_ai_beta" or custom_llm_provider == "gemini":
|
||||
response = _complete_vertex_ai_beta(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
response = _complete_vertex_ai(_dispatch_ctx)
|
||||
elif custom_llm_provider == "predibase":
|
||||
response = _complete_predibase(_dispatch_ctx)
|
||||
elif custom_llm_provider == "text-completion-codestral":
|
||||
response = _complete_text_completion_codestral(_dispatch_ctx)
|
||||
elif custom_llm_provider == "text-completion-inception":
|
||||
response = _complete_text_completion_inception(_dispatch_ctx)
|
||||
elif custom_llm_provider in ("sagemaker_chat", "sagemaker_nova"):
|
||||
# boto3 reads keys from .env
|
||||
# sagemaker_chat: HF Messages API endpoints
|
||||
# sagemaker_nova: Nova models on SageMaker (OpenAI-compatible)
|
||||
response = _complete_sagemaker_chat(_dispatch_ctx)
|
||||
elif custom_llm_provider == "sagemaker":
|
||||
# boto3 reads keys from .env
|
||||
response = _complete_sagemaker(_dispatch_ctx)
|
||||
elif custom_llm_provider == "bedrock":
|
||||
# boto3 reads keys from .env
|
||||
response = _complete_bedrock(_dispatch_ctx)
|
||||
elif custom_llm_provider == "watsonx":
|
||||
response = _complete_watsonx(_dispatch_ctx)
|
||||
elif custom_llm_provider == "watsonx_text":
|
||||
response = _complete_watsonx_text(_dispatch_ctx)
|
||||
elif custom_llm_provider == "vllm":
|
||||
response = _complete_vllm(_dispatch_ctx)
|
||||
elif custom_llm_provider == "ollama":
|
||||
response = _complete_ollama(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "ollama_chat":
|
||||
response = _complete_ollama_chat(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "triton":
|
||||
response = _complete_triton(_dispatch_ctx)
|
||||
elif custom_llm_provider == "cloudflare":
|
||||
response = _complete_cloudflare(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "petals" or model in litellm.petals_models:
|
||||
response = _complete_petals(_dispatch_ctx)
|
||||
elif custom_llm_provider == "snowflake" or model in litellm.snowflake_models:
|
||||
response = _complete_snowflake(_dispatch_ctx)
|
||||
elif custom_llm_provider == "gradient_ai":
|
||||
response = _complete_gradient_ai(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "gdc":
|
||||
response = _complete_gdc(_dispatch_ctx)
|
||||
elif custom_llm_provider == "bytez":
|
||||
response = _complete_bytez(_dispatch_ctx)
|
||||
elif custom_llm_provider == "lemonade":
|
||||
response = _complete_lemonade(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "ovhcloud" or model in litellm.ovhcloud_models:
|
||||
response = _complete_ovhcloud(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "custom":
|
||||
response = _complete_custom(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider in litellm._custom_providers: # Assume custom LLM provider
|
||||
# Get the Custom Handler
|
||||
response = _complete_custom_providers(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "langgraph":
|
||||
# LangGraph - Agent Runtime Provider
|
||||
response = _complete_langgraph(_dispatch_ctx)
|
||||
|
||||
elif custom_llm_provider == "langflow":
|
||||
# LangFlow - Visual AI Agent Platform
|
||||
response = _complete_langflow(_dispatch_ctx)
|
||||
|
||||
else:
|
||||
raise LiteLLMUnknownProvider(model=model, custom_llm_provider=custom_llm_provider)
|
||||
return response
|
||||
return dispatch_completion(_dispatch_ctx, lambda: _complete_python(_dispatch_ctx))
|
||||
except Exception as e:
|
||||
## Map to OpenAI Exception
|
||||
raise exception_type(
|
||||
|
|
@ -7771,179 +7748,168 @@ def transcription(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
response: TranscriptionResponse | Coroutine[object, object, TranscriptionResponse] | None = None
|
||||
def python_fallback(
|
||||
file: FileTypes,
|
||||
api_key: str | None = api_key,
|
||||
api_base: str | None = api_base,
|
||||
api_version: str | None = api_version,
|
||||
) -> TranscriptionResponse | Coroutine[object, object, TranscriptionResponse]:
|
||||
response: TranscriptionResponse | Coroutine[object, object, TranscriptionResponse] | None = None
|
||||
|
||||
provider_config: Final = ProviderConfigManager.get_provider_audio_transcription_config(
|
||||
provider_config: Final = ProviderConfigManager.get_provider_audio_transcription_config(
|
||||
model=model,
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
)
|
||||
|
||||
if custom_llm_provider in AZURE_OPENAI_AUDIO_PROVIDERS and provider_config is None:
|
||||
# azure configs
|
||||
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
||||
api_version = api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
azure_ad_token: Final = kwargs.pop("azure_ad_token", None) or get_secret_str("AZURE_AD_TOKEN")
|
||||
|
||||
api_key = api_key or litellm.api_key or litellm.azure_key or get_secret_str("AZURE_API_KEY")
|
||||
|
||||
optional_params["extra_headers"] = extra_headers
|
||||
|
||||
response = azure_audio_transcriptions.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
api_version=api_version,
|
||||
azure_ad_token=azure_ad_token,
|
||||
max_retries=max_retries,
|
||||
litellm_params=litellm_params_dict,
|
||||
)
|
||||
elif custom_llm_provider == "openai" or (custom_llm_provider in litellm.openai_compatible_providers):
|
||||
api_base = (
|
||||
api_base
|
||||
or litellm.api_base
|
||||
or get_secret("OPENAI_BASE_URL")
|
||||
or get_secret("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
openai.organization = (
|
||||
litellm.organization
|
||||
or get_secret("OPENAI_ORGANIZATION")
|
||||
or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105
|
||||
)
|
||||
# set API KEY
|
||||
|
||||
api_key = api_key or litellm.api_key or litellm.openai_key or get_secret("OPENAI_API_KEY")
|
||||
response = openai_audio_transcriptions.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
logging_obj=litellm_logging_obj,
|
||||
max_retries=max_retries,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params_dict,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
elif custom_llm_provider == "nvidia_riva":
|
||||
# NVIDIA Riva is gRPC-based, not HTTP. It has its own dedicated handler
|
||||
# rather than going through base_llm_http_handler.
|
||||
response = nvidia_riva_audio_transcriptions.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params_dict,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
timeout=timeout,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
provider_config=(
|
||||
provider_config if isinstance(provider_config, NvidiaRivaAudioTranscriptionConfig) else None
|
||||
),
|
||||
)
|
||||
elif custom_llm_provider == "soniox":
|
||||
from litellm.llms.soniox.audio_transcription.handler import (
|
||||
SonioxAudioTranscriptionHandler,
|
||||
)
|
||||
|
||||
response = SonioxAudioTranscriptionHandler().audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params_dict,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
client=(client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None),
|
||||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
headers=extra_headers,
|
||||
provider_config=provider_config,
|
||||
)
|
||||
elif provider_config is not None:
|
||||
response = base_llm_http_handler.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params_dict,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
client=(client if client is not None and isinstance(client, (HTTPHandler, AsyncHTTPHandler)) else None),
|
||||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
headers={},
|
||||
provider_config=provider_config,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
|
||||
if response is None:
|
||||
raise ValueError("Unmapped provider passed in. Unable to get the response.")
|
||||
return response
|
||||
|
||||
from litellm.rust_bridge.transcription import dispatch_transcription
|
||||
|
||||
response: Final = dispatch_transcription(
|
||||
model=model,
|
||||
provider=LlmProviders(custom_llm_provider),
|
||||
provider=custom_llm_provider,
|
||||
file=file,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
logging=litellm_logging_obj,
|
||||
asynchronous=bool(atranscription),
|
||||
has_custom_client=client is not None or shared_session is not None,
|
||||
fallback=python_fallback,
|
||||
)
|
||||
|
||||
if custom_llm_provider in AZURE_OPENAI_AUDIO_PROVIDERS and provider_config is None:
|
||||
# azure configs
|
||||
api_base = api_base or litellm.api_base or get_secret_str("AZURE_API_BASE")
|
||||
|
||||
api_version = api_version or litellm.api_version or get_secret_str("AZURE_API_VERSION")
|
||||
|
||||
azure_ad_token: Final = kwargs.pop("azure_ad_token", None) or get_secret_str("AZURE_AD_TOKEN")
|
||||
|
||||
api_key = api_key or litellm.api_key or litellm.azure_key or get_secret_str("AZURE_API_KEY")
|
||||
|
||||
optional_params["extra_headers"] = extra_headers
|
||||
|
||||
response = azure_audio_transcriptions.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
api_version=api_version,
|
||||
azure_ad_token=azure_ad_token,
|
||||
max_retries=max_retries,
|
||||
litellm_params=litellm_params_dict,
|
||||
)
|
||||
elif custom_llm_provider == "openai" or (custom_llm_provider in litellm.openai_compatible_providers):
|
||||
api_base = (
|
||||
api_base
|
||||
or litellm.api_base
|
||||
or get_secret("OPENAI_BASE_URL")
|
||||
or get_secret("OPENAI_API_BASE")
|
||||
or "https://api.openai.com/v1"
|
||||
)
|
||||
openai.organization = (
|
||||
litellm.organization
|
||||
or get_secret("OPENAI_ORGANIZATION")
|
||||
or None # default - https://github.com/openai/openai-python/blob/284c1799070c723c6a553337134148a7ab088dd8/openai/util.py#L105
|
||||
)
|
||||
# set API KEY
|
||||
|
||||
api_key = api_key or litellm.api_key or litellm.openai_key or get_secret("OPENAI_API_KEY")
|
||||
response = openai_audio_transcriptions.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
client=client,
|
||||
timeout=timeout,
|
||||
logging_obj=litellm_logging_obj,
|
||||
max_retries=max_retries,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
provider_config=provider_config,
|
||||
litellm_params=litellm_params_dict,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
elif custom_llm_provider == "nvidia_riva":
|
||||
# NVIDIA Riva is gRPC-based, not HTTP. It has its own dedicated handler
|
||||
# rather than going through base_llm_http_handler.
|
||||
response = nvidia_riva_audio_transcriptions.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params_dict,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
timeout=timeout,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
provider_config=(
|
||||
provider_config if isinstance(provider_config, NvidiaRivaAudioTranscriptionConfig) else None
|
||||
),
|
||||
)
|
||||
elif custom_llm_provider == "soniox":
|
||||
from litellm.llms.soniox.audio_transcription.handler import (
|
||||
SonioxAudioTranscriptionHandler,
|
||||
)
|
||||
|
||||
response = SonioxAudioTranscriptionHandler().audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params_dict,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
client=(
|
||||
client
|
||||
if client is not None and (isinstance(client, HTTPHandler) or isinstance(client, AsyncHTTPHandler))
|
||||
else None
|
||||
),
|
||||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
headers=extra_headers,
|
||||
provider_config=provider_config,
|
||||
)
|
||||
elif custom_llm_provider == "bedrock":
|
||||
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
|
||||
|
||||
dispatch: Final = BedrockAudioTranscriptionRustDispatch()
|
||||
if atranscription:
|
||||
response = dispatch.async_audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
else:
|
||||
response = dispatch.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
elif provider_config is not None:
|
||||
response = base_llm_http_handler.audio_transcriptions(
|
||||
model=model,
|
||||
audio_file=file,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params_dict,
|
||||
model_response=model_response,
|
||||
atranscription=atranscription,
|
||||
client=(
|
||||
client
|
||||
if client is not None and (isinstance(client, HTTPHandler) or isinstance(client, AsyncHTTPHandler))
|
||||
else None
|
||||
),
|
||||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
logging_obj=litellm_logging_obj,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
headers={},
|
||||
provider_config=provider_config,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
|
||||
# Store duration in _hidden_params for cost calculation without
|
||||
# exposing it in the response body (see sync path comment above).
|
||||
if response is not None and not isinstance(response, Coroutine):
|
||||
if not isinstance(response, Coroutine):
|
||||
existing_duration: Final = getattr(response, "duration", None)
|
||||
if existing_duration is None:
|
||||
calculated_duration: Final = calculate_request_duration(file)
|
||||
if calculated_duration is not None:
|
||||
response._hidden_params["audio_transcription_duration"] = calculated_duration
|
||||
|
||||
if response is None:
|
||||
raise ValueError("Unmapped provider passed in. Unable to get the response.")
|
||||
return response
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -8,6 +8,8 @@ import mimetypes
|
|||
import os
|
||||
import re
|
||||
from collections.abc import Callable, Coroutine, Mapping
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from io import IOBase
|
||||
from typing import Any, Final, cast
|
||||
|
||||
|
|
@ -17,16 +19,25 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import request_timeout
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model
|
||||
from litellm.llms.azure_ai.ocr.common_utils import (
|
||||
is_azure_document_intelligence_model,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import (
|
||||
OCR_REQUEST_FORMAT_PARAM,
|
||||
BaseOCRConfig,
|
||||
OCRResponse,
|
||||
parse_ocr_request_format,
|
||||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors
|
||||
from litellm.rust_bridge.runtime import DispatchResult
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestCapabilities,
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
request_context,
|
||||
vertex_options,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
|
|
@ -35,6 +46,22 @@ base_llm_http_handler = BaseLLMHTTPHandler()
|
|||
#################################################
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PreparedOCRRequest:
|
||||
model: str
|
||||
document: dict[str, Any]
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
custom_llm_provider: str
|
||||
extra_headers: dict[str, object] | None
|
||||
provider_config: BaseOCRConfig
|
||||
optional_params: dict[str, object]
|
||||
litellm_params: dict[str, object]
|
||||
effective_timeout: float | httpx.Timeout
|
||||
litellm_logging_obj: LiteLLMLoggingObj
|
||||
execution_mode: str = "sync"
|
||||
|
||||
|
||||
def _prepare_ocr_request(
|
||||
model: str,
|
||||
document: Mapping[str, object],
|
||||
|
|
@ -44,7 +71,8 @@ def _prepare_ocr_request(
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
kwargs: dict[str, object],
|
||||
) -> rust_ocr_bridge.PreparedOCRRequest:
|
||||
execution_mode: str = "sync",
|
||||
) -> _PreparedOCRRequest:
|
||||
litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
|
||||
litellm_call_id: Final = cast(str | None, kwargs.get("litellm_call_id", None))
|
||||
|
||||
|
|
@ -141,7 +169,7 @@ def _prepare_ocr_request(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
return rust_ocr_bridge.PreparedOCRRequest(
|
||||
return _PreparedOCRRequest(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
|
|
@ -153,73 +181,175 @@ def _prepare_ocr_request(
|
|||
litellm_params=dict(litellm_params),
|
||||
effective_timeout=effective_timeout,
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
execution_mode=execution_mode,
|
||||
)
|
||||
|
||||
|
||||
@anative_first(
|
||||
native=rust_ocr_bridge.aattempt_ocr,
|
||||
route="ocr",
|
||||
errors=lambda prepared_request, resolve_api_key: provider_errors(
|
||||
prepared_request.custom_llm_provider, prepared_request.model
|
||||
),
|
||||
)
|
||||
async def _execute_aocr(
|
||||
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
|
||||
def _rust_bridge_optional_params(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> dict[str, object]:
|
||||
optional_params: Final = dict(prepared_request.optional_params)
|
||||
if prepared_request.custom_llm_provider == "vertex_ai":
|
||||
vertex_project: Final = (
|
||||
prepared_request.litellm_params.get("vertex_project")
|
||||
or prepared_request.litellm_params.get("vertex_ai_project")
|
||||
or litellm.vertex_project
|
||||
or resolve_secret("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_location: Final = (
|
||||
prepared_request.litellm_params.get("vertex_location")
|
||||
or prepared_request.litellm_params.get("vertex_ai_location")
|
||||
or litellm.vertex_location
|
||||
or resolve_secret("VERTEXAI_LOCATION")
|
||||
or resolve_secret("VERTEX_LOCATION")
|
||||
)
|
||||
if vertex_project is not None:
|
||||
optional_params["vertex_project"] = vertex_project
|
||||
if vertex_location is not None:
|
||||
optional_params["vertex_location"] = vertex_location
|
||||
return optional_params
|
||||
|
||||
|
||||
def _rust_bridge_api_base(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> str | None:
|
||||
if prepared_request.api_base is not None:
|
||||
return prepared_request.api_base
|
||||
if prepared_request.custom_llm_provider == "azure_ai":
|
||||
if is_azure_document_intelligence_model(prepared_request.model):
|
||||
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
return resolve_secret("AZURE_AI_API_BASE")
|
||||
return None
|
||||
|
||||
|
||||
def _prepare_rust_ocr_call(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> OCRResponse:
|
||||
pending: Final = base_llm_http_handler.ocr(
|
||||
) -> PreparedNativeCall[rust_ocr_bridge.NativeOCRRequest]:
|
||||
provider_config: Final = prepared_request.provider_config
|
||||
api_key_env_var: Final = provider_config.get_api_key_env_var()
|
||||
resolved_api_key: Final = prepared_request.api_key or (
|
||||
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
|
||||
)
|
||||
resolved_headers: Final = provider_config.validate_environment(
|
||||
headers=prepared_request.extra_headers or {},
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=prepared_request.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
logging_obj=prepared_request.litellm_logging_obj,
|
||||
api_key=prepared_request.api_key,
|
||||
api_key=resolved_api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
aocr=True,
|
||||
headers=prepared_request.extra_headers,
|
||||
provider_config=prepared_request.provider_config,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
response: Final = await pending if asyncio.iscoroutine(pending) else pending
|
||||
if response is None:
|
||||
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
|
||||
return response
|
||||
resolved_complete_url: Final = provider_config.get_complete_url(
|
||||
api_base=prepared_request.api_base,
|
||||
model=prepared_request.model,
|
||||
optional_params=prepared_request.optional_params,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
|
||||
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
|
||||
prepared_request.litellm_logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=resolved_api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": {
|
||||
"model": prepared_request.model,
|
||||
"document": prepared_request.document,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": resolved_complete_url,
|
||||
"headers": resolved_headers,
|
||||
},
|
||||
)
|
||||
return PreparedNativeCall(
|
||||
request=rust_ocr_bridge.NativeOCRRequest(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=prepared_request.optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
vertex=vertex_options(rust_optional_params),
|
||||
api_key=resolved_api_key,
|
||||
api_base=rust_api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=cast( # cast-ok: provider header normalization returns string-object pairs
|
||||
dict[str, object], resolved_headers
|
||||
),
|
||||
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
|
||||
),
|
||||
context=request_context(
|
||||
logging_obj=prepared_request.litellm_logging_obj,
|
||||
request_model=getattr(prepared_request.litellm_logging_obj, "model", prepared_request.model),
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
capabilities=NativeRequestCapabilities(
|
||||
execution_mode=prepared_request.execution_mode,
|
||||
input_source_kind=str(prepared_request.document.get("type") or "unknown"),
|
||||
request_format=(
|
||||
value
|
||||
if isinstance((value := prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM)), str)
|
||||
else None
|
||||
),
|
||||
native_response_format=(prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _attempt_ocr(
|
||||
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
|
||||
@dataclass
|
||||
class _OCROperation:
|
||||
request: _PreparedOCRRequest
|
||||
resolve_api_key: Callable[[str], str | None]
|
||||
python: Callable[[], OCRResponse | Coroutine[object, object, OCRResponse]]
|
||||
logged: bool = False
|
||||
|
||||
def prepare(self) -> PreparedNativeCall[rust_ocr_bridge.NativeOCRRequest]:
|
||||
prepared: Final = _prepare_rust_ocr_call(self.request, self.resolve_api_key)
|
||||
self.logged = True
|
||||
return prepared
|
||||
|
||||
def fallback(self) -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
with self.request.litellm_logging_obj.suppress_next_pre_call() if self.logged else nullcontext():
|
||||
return self.python()
|
||||
|
||||
async def afallback(self) -> OCRResponse:
|
||||
with self.request.litellm_logging_obj.suppress_next_pre_call() if self.logged else nullcontext():
|
||||
result: Final = self.python()
|
||||
return await result if isinstance(result, Coroutine) else result
|
||||
|
||||
|
||||
def _run_rust_ocr(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
is_async: bool,
|
||||
) -> DispatchResult[OCRResponse]:
|
||||
return rust_ocr_bridge.attempt_ocr(prepared_request=prepared_request, resolve_api_key=resolve_api_key)
|
||||
|
||||
|
||||
@native_first(
|
||||
native=_attempt_ocr,
|
||||
route="ocr",
|
||||
errors=lambda prepared_request, resolve_api_key, is_async: provider_errors(
|
||||
prepared_request.custom_llm_provider, prepared_request.model
|
||||
),
|
||||
)
|
||||
def _execute_ocr(
|
||||
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
is_async: bool,
|
||||
fallback: Callable[[], OCRResponse | Coroutine[object, object, OCRResponse]],
|
||||
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
return base_llm_http_handler.ocr(
|
||||
operation: Final = _OCROperation(prepared_request, resolve_api_key, fallback)
|
||||
return rust_ocr_bridge.dispatch_ocr(
|
||||
prepare=operation.prepare,
|
||||
fallback=operation.fallback,
|
||||
adapt=OCRResponse.model_validate,
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=prepared_request.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
logging_obj=prepared_request.litellm_logging_obj,
|
||||
api_key=prepared_request.api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
aocr=is_async,
|
||||
headers=prepared_request.extra_headers,
|
||||
provider_config=prepared_request.provider_config,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
provider=prepared_request.custom_llm_provider,
|
||||
request_format=(
|
||||
"native" if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" else None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def _run_rust_aocr(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
fallback: Callable[[], Coroutine[object, object, OCRResponse]],
|
||||
) -> OCRResponse:
|
||||
operation: Final = _OCROperation(prepared_request, resolve_api_key, fallback)
|
||||
return await rust_ocr_bridge.adispatch_ocr(
|
||||
prepare=operation.prepare,
|
||||
fallback=operation.afallback,
|
||||
adapt=OCRResponse.model_validate,
|
||||
model=prepared_request.model,
|
||||
provider=prepared_request.custom_llm_provider,
|
||||
request_format=(
|
||||
"native" if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" else None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -312,6 +442,7 @@ async def aocr(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
kwargs=kwargs,
|
||||
execution_mode="async",
|
||||
)
|
||||
model = prepared.model
|
||||
custom_llm_provider = prepared.custom_llm_provider
|
||||
|
|
@ -319,7 +450,31 @@ async def aocr(
|
|||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str)
|
||||
async def python_fallback() -> OCRResponse:
|
||||
pending: Final = base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared.effective_timeout,
|
||||
logging_obj=prepared.litellm_logging_obj,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared.custom_llm_provider,
|
||||
aocr=True,
|
||||
headers=prepared.extra_headers,
|
||||
provider_config=prepared.provider_config,
|
||||
litellm_params=prepared.litellm_params,
|
||||
)
|
||||
response: Final = await pending if asyncio.iscoroutine(pending) else pending
|
||||
if response is None:
|
||||
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
|
||||
return response
|
||||
|
||||
return await _run_rust_aocr(
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
fallback=python_fallback,
|
||||
)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
@ -560,7 +715,27 @@ def ocr(
|
|||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
return _execute_ocr(prepared_request=prepared, resolve_api_key=get_secret_str, is_async=_is_async)
|
||||
def python_fallback() -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
return base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared.effective_timeout,
|
||||
logging_obj=prepared.litellm_logging_obj,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared.custom_llm_provider,
|
||||
aocr=_is_async,
|
||||
headers=prepared.extra_headers,
|
||||
provider_config=prepared.provider_config,
|
||||
litellm_params=prepared.litellm_params,
|
||||
)
|
||||
|
||||
return _run_rust_ocr(
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
fallback=python_fallback,
|
||||
)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -4,23 +4,31 @@ The Rust core owns the conversation translation, the provider call, and the
|
|||
response normalization for the subset of `/chat/completions` requests it
|
||||
accepts. This module only marshals inputs and hands the normalized result to
|
||||
LiteLLM's existing `ModelResponse` builder.
|
||||
|
||||
``None`` means the provider was never called, so the caller is free to serve the
|
||||
request on the Python path. A failure after the call was issued raises instead:
|
||||
retrying it there would bill the customer for the same work twice.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Awaitable, Callable, Coroutine, Mapping, Sequence
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass, replace
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import (
|
||||
convert_to_model_response_object,
|
||||
)
|
||||
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
|
||||
from litellm.llms.bedrock.request_metadata import get_bedrock_request_metadata_fields
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.protocols import (
|
||||
RustAchatCompletions,
|
||||
|
|
@ -31,28 +39,35 @@ from litellm.rust_bridge.request import (
|
|||
NativeAnthropicOptions,
|
||||
NativeBedrockOptions,
|
||||
NativeChatCompletionsRequest,
|
||||
NativePreCallDetails,
|
||||
NativeRequestCapabilities,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
anthropic_options,
|
||||
bedrock_options,
|
||||
call_native,
|
||||
request_context,
|
||||
with_capabilities,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointBinding,
|
||||
EndpointDispatch,
|
||||
PythonFallback,
|
||||
async_none,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.completion import (
|
||||
_CompletionDispatchContext, # pyright: ignore[reportPrivateUsage] # shared internal SDK dispatch context
|
||||
_CompletionDispatchResult, # pyright: ignore[reportPrivateUsage] # shared internal SDK dispatch result
|
||||
)
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
# Providers whose `/chat/completions` deployments the Rust core can serve. A
|
||||
# provider outside this set never reaches the bridge.
|
||||
RUST_CHAT_COMPLETIONS_PROVIDERS: Final = frozenset({"anthropic", "bedrock"})
|
||||
|
||||
# `litellm_params` values are `object`, so validate the one this module reads
|
||||
# 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"
|
||||
|
||||
|
||||
|
|
@ -96,10 +111,16 @@ def response_logger(
|
|||
return log
|
||||
|
||||
|
||||
_CHAT: Final[NativeBinding[RustChatCompletions]] = NativeBinding(lambda native: native.chat_completions)
|
||||
_ACHAT: Final[NativeBinding[RustAchatCompletions]] = NativeBinding(lambda native: native.achat_completions)
|
||||
_CHAT_PREFLIGHT: Final[NativeBinding[RustChatCompletionsDecline]] = NativeBinding(
|
||||
lambda native: native.chat_completions_decline
|
||||
_CHAT: Final[EndpointDispatch[RustChatCompletions, RustAchatCompletions]] = EndpointDispatch.native(
|
||||
route="chat_completions",
|
||||
sync=lambda native: native.chat_completions,
|
||||
asynchronous=lambda native: native.achat_completions,
|
||||
enabled=rust_enabled,
|
||||
)
|
||||
_CHAT_PREFLIGHT: Final[EndpointBinding[RustChatCompletionsDecline]] = EndpointBinding.native(
|
||||
route="chat_completions",
|
||||
select=lambda native: native.chat_completions_decline,
|
||||
enabled=rust_enabled,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -113,14 +134,14 @@ def set_rust_chat_completions(
|
|||
patching module attributes."""
|
||||
if not isinstance(chat_completions, Unchanged):
|
||||
if chat_completions is None:
|
||||
_CHAT.reset()
|
||||
_CHAT.sync.reset()
|
||||
else:
|
||||
_CHAT.override(chat_completions)
|
||||
_CHAT.sync.override(chat_completions)
|
||||
if not isinstance(achat_completions, Unchanged):
|
||||
if achat_completions is None:
|
||||
_ACHAT.reset()
|
||||
_CHAT.asynchronous.reset()
|
||||
else:
|
||||
_ACHAT.override(achat_completions)
|
||||
_CHAT.asynchronous.override(achat_completions)
|
||||
if not isinstance(decline, Unchanged):
|
||||
if decline is None:
|
||||
_CHAT_PREFLIGHT.reset()
|
||||
|
|
@ -128,42 +149,43 @@ def set_rust_chat_completions(
|
|||
_CHAT_PREFLIGHT.override(decline)
|
||||
|
||||
|
||||
def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool:
|
||||
metadata: Final = litellm_params.get("metadata") if litellm_params is not None else None
|
||||
try:
|
||||
entries: Final = _LITELLM_METADATA_ADAPTER.validate_python(metadata)
|
||||
except ValidationError:
|
||||
return False
|
||||
return entries.get("user_id") is not None
|
||||
def _provider_eligibility_options(
|
||||
provider: str | None,
|
||||
litellm_params: Mapping[str, object] | None,
|
||||
optional_params: Mapping[str, object],
|
||||
) -> NativeRequestOptions:
|
||||
bedrock: Final = (
|
||||
replace(
|
||||
bedrock_options(optional_params),
|
||||
request_metadata_fields=get_bedrock_request_metadata_fields(),
|
||||
)
|
||||
if provider == "bedrock"
|
||||
else None
|
||||
)
|
||||
anthropic: Final = anthropic_options(litellm_params) if provider == "anthropic" else None
|
||||
return NativeRequestOptions(custom_llm_provider=provider, bedrock=bedrock, anthropic=anthropic)
|
||||
|
||||
|
||||
def _litellm_metadata_reaches_the_provider(
|
||||
custom_llm_provider: str | None, litellm_params: Mapping[str, object] | None
|
||||
) -> bool:
|
||||
"""Whether the Python transform would promote proxy-owned attribution into the
|
||||
provider request, below this gate and inside the function the Rust route replaces.
|
||||
def _eligibility_context(
|
||||
*,
|
||||
execution_mode: str | None = None,
|
||||
stream: bool,
|
||||
has_custom_client: bool = False,
|
||||
has_agentic_hook: bool = False,
|
||||
) -> NativeRequestContext:
|
||||
return NativeRequestContext(
|
||||
capabilities=NativeRequestCapabilities(
|
||||
execution_mode=execution_mode,
|
||||
stream=stream,
|
||||
has_custom_client=has_custom_client,
|
||||
has_agentic_hook=has_agentic_hook,
|
||||
)
|
||||
)
|
||||
|
||||
`AnthropicConfig.transform_request` promotes a valid `metadata["user_id"]`
|
||||
into the Messages body, so the core never sees the key and would send the
|
||||
request to Anthropic with the abuse-detection attribution missing.
|
||||
|
||||
`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the
|
||||
Converse body whenever the operator armed `bedrock_request_metadata_fields`.
|
||||
Owning that field also means evicting a caller-supplied one, which the core
|
||||
cannot do either, so ownership alone is the condition rather than whether
|
||||
anything resolved.
|
||||
|
||||
Deliberately a superset of Python's condition in both cases: declining a
|
||||
request Python would not have attributed anyway costs only the Rust path,
|
||||
while missing one loses the attribution silently.
|
||||
"""
|
||||
match custom_llm_provider:
|
||||
case "anthropic":
|
||||
return _anthropic_user_id_reaches_the_body(litellm_params)
|
||||
case "bedrock":
|
||||
return bedrock_request_metadata_is_owned()
|
||||
case _:
|
||||
return False
|
||||
def _execution_context(context: NativeRequestContext | None, mode: str) -> NativeRequestContext:
|
||||
current = context or NativeRequestContext()
|
||||
return with_capabilities(current, replace(current.capabilities, execution_mode=mode))
|
||||
|
||||
|
||||
def rust_chat_completions_accepts(
|
||||
|
|
@ -182,31 +204,16 @@ def rust_chat_completions_accepts(
|
|||
capability gate answers the second half; it resolves no credentials and
|
||||
performs no I/O.
|
||||
"""
|
||||
if custom_llm_provider not in RUST_CHAT_COMPLETIONS_PROVIDERS:
|
||||
return False
|
||||
if stream:
|
||||
return False
|
||||
if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params):
|
||||
verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path")
|
||||
return False
|
||||
if not rust_enabled():
|
||||
return False
|
||||
decline: Final = _CHAT_PREFLIGHT.load()
|
||||
if decline is None:
|
||||
return False
|
||||
try:
|
||||
reason: Final = decline(
|
||||
return _CHAT_PREFLIGHT.accepts(
|
||||
check=lambda decline: decline(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 # capability checks perform no provider I/O
|
||||
verbose_logger.debug("Native chat acceptance check failed: %s", error)
|
||||
return False
|
||||
if reason is not None:
|
||||
verbose_logger.debug("Native chat request is ineligible: %s", reason)
|
||||
return reason is None
|
||||
options=_provider_eligibility_options(custom_llm_provider, litellm_params, optional_params),
|
||||
context=_eligibility_context(stream=bool(stream)),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _build_model_response(
|
||||
|
|
@ -237,26 +244,19 @@ def chat_completions(
|
|||
on_response: ResponseObserver,
|
||||
bedrock: NativeBedrockOptions | None = None,
|
||||
anthropic: NativeAnthropicOptions | None = None,
|
||||
stream: bool = False,
|
||||
has_custom_client: bool = False,
|
||||
eligible: bool = True,
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> DispatchResult[ModelResponse]:
|
||||
) -> ModelResponse | None:
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
def call(
|
||||
native: RustChatCompletions, prepared: PreparedNativeCall[NativeChatCompletionsRequest]
|
||||
) -> Mapping[str, object]:
|
||||
return call_native(native, prepared)
|
||||
|
||||
return attempt(
|
||||
load=_CHAT.load,
|
||||
enabled=rust_enabled(),
|
||||
eligible=eligible,
|
||||
return _CHAT.invoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeChatCompletionsRequest(model=model, messages=messages, optional_params=optional_params),
|
||||
NativeChatCompletionsRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
|
|
@ -266,17 +266,12 @@ def chat_completions(
|
|||
bedrock=bedrock,
|
||||
anthropic=anthropic,
|
||||
),
|
||||
context=with_capabilities(
|
||||
context or NativeRequestContext(),
|
||||
NativeRequestCapabilities(
|
||||
execution_mode="sync",
|
||||
stream=stream,
|
||||
has_custom_client=has_custom_client,
|
||||
),
|
||||
),
|
||||
context=_execution_context(context, "sync"),
|
||||
),
|
||||
call=call,
|
||||
call=call_native,
|
||||
fallback=lambda: None,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -294,27 +289,19 @@ async def achat_completions(
|
|||
on_response: ResponseObserver,
|
||||
bedrock: NativeBedrockOptions | None = None,
|
||||
anthropic: NativeAnthropicOptions | None = None,
|
||||
stream: bool = False,
|
||||
has_custom_client: bool = False,
|
||||
eligible: bool = True,
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> DispatchResult[ModelResponse]:
|
||||
) -> ModelResponse | None:
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
async def call(
|
||||
native: RustAchatCompletions,
|
||||
prepared: PreparedNativeCall[NativeChatCompletionsRequest],
|
||||
) -> Mapping[str, object]:
|
||||
return await call_native(native, prepared)
|
||||
|
||||
return await aattempt(
|
||||
load=_ACHAT.load,
|
||||
enabled=rust_enabled(),
|
||||
eligible=eligible,
|
||||
return await _CHAT.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeChatCompletionsRequest(model=model, messages=messages, optional_params=optional_params),
|
||||
NativeChatCompletionsRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
|
|
@ -324,15 +311,209 @@ async def achat_completions(
|
|||
bedrock=bedrock,
|
||||
anthropic=anthropic,
|
||||
),
|
||||
context=with_capabilities(
|
||||
context or NativeRequestContext(),
|
||||
NativeRequestCapabilities(
|
||||
execution_mode="async",
|
||||
stream=stream,
|
||||
has_custom_client=has_custom_client,
|
||||
context=_execution_context(context, "async"),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=async_none,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
async def achat_completions_or_fallback(
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object],
|
||||
model_response: ModelResponse,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
python_fallback: Callable[[], Awaitable[object]],
|
||||
bedrock: NativeBedrockOptions | None = None,
|
||||
anthropic: NativeAnthropicOptions | None = None,
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> 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.
|
||||
"""
|
||||
|
||||
def adapt(rust_response: Mapping[str, object]) -> object:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
return await _CHAT.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeChatCompletionsRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
bedrock=bedrock,
|
||||
anthropic=anthropic,
|
||||
),
|
||||
context=_execution_context(context, "async"),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
_PARAMS_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_STR_ADAPTER: Final = TypeAdapter(str | None)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _ChatOperation:
|
||||
context: _CompletionDispatchContext
|
||||
python: Callable[[], _CompletionDispatchResult]
|
||||
pre_call_logged: bool = False
|
||||
|
||||
def assess(self) -> PythonFallback | None:
|
||||
ctx: Final = self.context
|
||||
return _CHAT_PREFLIGHT.assess(
|
||||
check=lambda decline: decline(
|
||||
model=ctx.model,
|
||||
messages=ctx.messages,
|
||||
optional_params=ctx.optional_params,
|
||||
custom_llm_provider=ctx.custom_llm_provider,
|
||||
options=_provider_eligibility_options(ctx.custom_llm_provider, ctx.litellm_params, ctx.optional_params),
|
||||
context=_eligibility_context(
|
||||
execution_mode="async" if ctx.acompletion else "sync",
|
||||
stream=bool(ctx.stream),
|
||||
has_custom_client=ctx.client is not None or ctx.shared_session is not None,
|
||||
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(ctx.logging),
|
||||
),
|
||||
),
|
||||
),
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
)
|
||||
|
||||
def prepare(self) -> PreparedNativeCall[NativeChatCompletionsRequest]:
|
||||
ctx: Final = self.context
|
||||
config: Final = ctx.provider_config
|
||||
defaults: Final = (
|
||||
_PARAMS_ADAPTER.validate_python(config.get_config_for_model(ctx.model))
|
||||
if config is not None
|
||||
else MappingProxyType({})
|
||||
)
|
||||
params: Final = _PARAMS_ADAPTER.validate_python(MappingProxyType({**defaults, **ctx.optional_params}))
|
||||
key: Final = (
|
||||
ctx.api_key
|
||||
or _STR_ADAPTER.validate_python(getattr(litellm, f"{ctx.custom_llm_provider}_key", None))
|
||||
or litellm.api_key
|
||||
or get_secret_str(f"{ctx.custom_llm_provider.upper()}_API_KEY")
|
||||
)
|
||||
base: Final = (
|
||||
ctx.api_base
|
||||
or litellm.api_base
|
||||
or get_secret_str(f"{ctx.custom_llm_provider.upper()}_API_BASE")
|
||||
or get_secret_str(f"{ctx.custom_llm_provider.upper()}_BASE_URL")
|
||||
)
|
||||
initial_headers: Final = _PARAMS_ADAPTER.validate_python(
|
||||
MappingProxyType({**(ctx.headers or MappingProxyType({})), **(ctx.extra_headers or MappingProxyType({}))})
|
||||
)
|
||||
headers: Final = (
|
||||
_PARAMS_ADAPTER.validate_python(
|
||||
config.validate_environment(
|
||||
api_key=key,
|
||||
api_base=base,
|
||||
headers=initial_headers,
|
||||
model=ctx.model,
|
||||
messages=ctx.messages,
|
||||
optional_params=params,
|
||||
litellm_params=ctx.litellm_params,
|
||||
)
|
||||
)
|
||||
if config is not None
|
||||
else initial_headers
|
||||
)
|
||||
log_details: Final[NativePreCallDetails] = {
|
||||
"complete_input_dict": {"model": ctx.model, "messages": ctx.messages, **params},
|
||||
"api_base": base or "",
|
||||
"headers": headers,
|
||||
}
|
||||
ctx.logging.pre_call(input=ctx.messages, api_key=key, additional_args=log_details)
|
||||
self.pre_call_logged = True
|
||||
provider_options: Final = _provider_eligibility_options(ctx.custom_llm_provider, ctx.litellm_params, params)
|
||||
return PreparedNativeCall(
|
||||
NativeChatCompletionsRequest(
|
||||
model=ctx.model,
|
||||
messages=ctx.messages,
|
||||
optional_params=params,
|
||||
),
|
||||
options=replace(
|
||||
provider_options,
|
||||
api_key=key,
|
||||
api_base=base,
|
||||
extra_headers=headers,
|
||||
timeout_seconds=timeout_to_seconds(float(ctx.timeout) if isinstance(ctx.timeout, str) else ctx.timeout),
|
||||
),
|
||||
context=request_context(
|
||||
logging_obj=ctx.logging,
|
||||
request_model=ctx.logging.model,
|
||||
litellm_params=ctx.litellm_params,
|
||||
capabilities=NativeRequestCapabilities(
|
||||
execution_mode="async" if ctx.acompletion else "sync",
|
||||
stream=bool(ctx.stream),
|
||||
has_custom_client=ctx.client is not None or ctx.shared_session is not None,
|
||||
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(ctx.logging),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
def fallback(self) -> _CompletionDispatchResult:
|
||||
with self.context.logging.suppress_next_pre_call() if self.pre_call_logged else nullcontext():
|
||||
return self.python()
|
||||
|
||||
async def afallback(self) -> ModelResponse | litellm.CustomStreamWrapper:
|
||||
with self.context.logging.suppress_next_pre_call() if self.pre_call_logged else nullcontext():
|
||||
result: Final = self.python()
|
||||
return await result if isinstance(result, Coroutine) else result
|
||||
|
||||
def adapt(self, response: Mapping[str, object]) -> ModelResponse:
|
||||
self.context.logging.post_call(
|
||||
input=self.context.messages,
|
||||
api_key=self.context.api_key,
|
||||
original_response=json.dumps(response),
|
||||
)
|
||||
return _build_model_response(response, self.context.model_response)
|
||||
|
||||
|
||||
def dispatch_completion(
|
||||
context: _CompletionDispatchContext,
|
||||
fallback: Callable[[], _CompletionDispatchResult],
|
||||
) -> _CompletionDispatchResult:
|
||||
operation: Final = _ChatOperation(context, fallback)
|
||||
error_context: Final = BridgeErrorContext(provider=context.custom_llm_provider, model=context.model)
|
||||
if context.acompletion:
|
||||
return _CHAT.ainvoke(
|
||||
prepare=operation.prepare,
|
||||
call=call_native,
|
||||
fallback=operation.afallback,
|
||||
adapt=operation.adapt,
|
||||
error_context=error_context,
|
||||
preflight=operation.assess,
|
||||
)
|
||||
return _CHAT.invoke(
|
||||
prepare=operation.prepare,
|
||||
call=call_native,
|
||||
fallback=operation.fallback,
|
||||
adapt=operation.adapt,
|
||||
error_context=error_context,
|
||||
preflight=operation.assess,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2,51 +2,97 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Final
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Callable, Coroutine, Iterator, Sequence
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from functools import reduce
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, TypeAdapter
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
|
||||
from litellm.rust_bridge.protocols import RustAmessages, RustMessages
|
||||
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
|
||||
from litellm.litellm_core_utils.dot_notation_indexing import delete_nested_value
|
||||
from litellm.litellm_core_utils.get_provider_specific_headers import ProviderSpecificHeaderUtils
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.protocols import RustAmessages, RustMessages, RustRouteDecline
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeMessagesRequest,
|
||||
NativePreCallDetails,
|
||||
NativeRequestCapabilities,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
request_context,
|
||||
with_capabilities,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointBinding,
|
||||
EndpointDispatch,
|
||||
PythonFallback,
|
||||
assess_route,
|
||||
async_none,
|
||||
identity,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
|
||||
from litellm.types.llms.openai import ChatCompletionUserMessage
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders, ProviderSpecificHeader
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
_MESSAGES: Final[NativeBinding[RustMessages]] = NativeBinding(lambda native: native.messages)
|
||||
_AMESSAGES: Final[NativeBinding[RustAmessages]] = NativeBinding(lambda native: native.amessages)
|
||||
_MESSAGES: Final[EndpointDispatch[RustMessages, RustAmessages]] = EndpointDispatch.native(
|
||||
route="messages",
|
||||
sync=lambda native: native.messages,
|
||||
asynchronous=lambda native: native.amessages,
|
||||
enabled=rust_enabled,
|
||||
)
|
||||
|
||||
|
||||
_PREFLIGHT: Final[EndpointBinding[RustRouteDecline]] = EndpointBinding.native(
|
||||
route="messages",
|
||||
select=lambda native: native.messages_decline,
|
||||
enabled=rust_enabled,
|
||||
)
|
||||
|
||||
|
||||
def set_rust_messages(
|
||||
*,
|
||||
messages: RustMessages | None | Unchanged = UNCHANGED,
|
||||
amessages: RustAmessages | None | Unchanged = UNCHANGED,
|
||||
decline: RustRouteDecline | None | Unchanged = UNCHANGED,
|
||||
) -> None:
|
||||
if not isinstance(decline, Unchanged):
|
||||
if decline is None:
|
||||
_PREFLIGHT.reset()
|
||||
else:
|
||||
_PREFLIGHT.override(decline)
|
||||
if not isinstance(messages, Unchanged):
|
||||
if messages is None:
|
||||
_MESSAGES.reset()
|
||||
_MESSAGES.sync.reset()
|
||||
else:
|
||||
_MESSAGES.override(messages)
|
||||
_MESSAGES.sync.override(messages)
|
||||
if not isinstance(amessages, Unchanged):
|
||||
if amessages is None:
|
||||
_AMESSAGES.reset()
|
||||
_MESSAGES.asynchronous.reset()
|
||||
else:
|
||||
_AMESSAGES.override(amessages)
|
||||
_MESSAGES.asynchronous.override(amessages)
|
||||
|
||||
|
||||
def load_rust_messages() -> RustMessages | None:
|
||||
return _MESSAGES.load()
|
||||
return _MESSAGES.sync.load()
|
||||
|
||||
|
||||
def load_rust_amessages() -> RustAmessages | None:
|
||||
return _AMESSAGES.load()
|
||||
return _MESSAGES.asynchronous.load()
|
||||
|
||||
|
||||
def messages(
|
||||
|
|
@ -62,13 +108,13 @@ def messages(
|
|||
has_custom_client: bool = False,
|
||||
has_agentic_hook: bool = False,
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> DispatchResult[dict[str, object]]:
|
||||
return attempt(
|
||||
load=_MESSAGES.load,
|
||||
enabled=True,
|
||||
eligible=True,
|
||||
) -> dict[str, object] | None:
|
||||
return _MESSAGES.invoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeMessagesRequest(model=model, body=body),
|
||||
NativeMessagesRequest(
|
||||
model=model,
|
||||
body=body,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
|
|
@ -87,7 +133,10 @@ def messages(
|
|||
),
|
||||
),
|
||||
call=call_native,
|
||||
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -104,13 +153,13 @@ async def amessages(
|
|||
has_custom_client: bool = False,
|
||||
has_agentic_hook: bool = False,
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> DispatchResult[dict[str, object]]:
|
||||
return await aattempt(
|
||||
load=_AMESSAGES.load,
|
||||
enabled=True,
|
||||
eligible=True,
|
||||
) -> dict[str, object] | None:
|
||||
return await _MESSAGES.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeMessagesRequest(model=model, body=body),
|
||||
NativeMessagesRequest(
|
||||
model=model,
|
||||
body=body,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
|
|
@ -129,5 +178,221 @@ async def amessages(
|
|||
),
|
||||
),
|
||||
call=call_native,
|
||||
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
MessagesResponse = AnthropicMessagesResponse | Iterator[bytes] | AsyncIterator[object]
|
||||
MessagesResult = MessagesResponse | Coroutine[object, None, MessagesResponse]
|
||||
_BODY_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
||||
class _NativeMessagesResponse(BaseModel):
|
||||
id: str
|
||||
type: Literal["message"]
|
||||
role: Literal["assistant"]
|
||||
model: str
|
||||
content: list[dict[str, object]]
|
||||
usage: dict[str, object]
|
||||
|
||||
|
||||
class _BridgedMessagesResponse(AnthropicMessagesResponse):
|
||||
_hidden_params: ReadOnly[dict[str, dict[str, str]]]
|
||||
|
||||
|
||||
_RESPONSE_ADAPTER: Final = TypeAdapter(AnthropicMessagesResponse)
|
||||
|
||||
|
||||
@dataclass
|
||||
class _MessagesOperation:
|
||||
model: str
|
||||
provider: str
|
||||
messages: list[dict[str, object]]
|
||||
body: Callable[[], dict[str, object]]
|
||||
params: GenericLiteLLMParams
|
||||
logging: Logging | None
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
python: Callable[[], MessagesResult]
|
||||
stream: bool = False
|
||||
asynchronous: bool = False
|
||||
has_custom_client: bool = False
|
||||
logged: bool = False
|
||||
|
||||
def prepare(self) -> PreparedNativeCall[NativeMessagesRequest]:
|
||||
requested: Final = self.body()
|
||||
body: Final = _BODY_ADAPTER.validate_python(
|
||||
reduce(
|
||||
delete_nested_value,
|
||||
TypeAdapter(tuple[str, ...]).validate_python(self.params.get("additional_drop_params") or ()),
|
||||
requested,
|
||||
)
|
||||
)
|
||||
provider_headers: Final = ProviderSpecificHeaderUtils.get_provider_specific_headers(
|
||||
TypeAdapter(ProviderSpecificHeader | Sequence[ProviderSpecificHeader] | None).validate_python(
|
||||
self.params.get("provider_specific_header")
|
||||
),
|
||||
self.provider,
|
||||
)
|
||||
config: Final = ProviderConfigManager.get_provider_anthropic_messages_config(
|
||||
model=self.model, provider=LlmProviders(self.provider)
|
||||
)
|
||||
initial_headers: Final = _BODY_ADAPTER.validate_python(
|
||||
MappingProxyType(
|
||||
{
|
||||
**(self.params.get("headers") or MappingProxyType({})),
|
||||
**(self.params.get("extra_headers") or MappingProxyType({})),
|
||||
**provider_headers,
|
||||
}
|
||||
)
|
||||
)
|
||||
validated_headers, base = (
|
||||
config.validate_anthropic_messages_environment(
|
||||
headers=initial_headers,
|
||||
model=self.model,
|
||||
messages=self.messages,
|
||||
optional_params=body,
|
||||
litellm_params=self.params.model_dump(),
|
||||
api_key=self.api_key,
|
||||
api_base=self.api_base,
|
||||
)
|
||||
if config is not None
|
||||
else (initial_headers, self.api_base)
|
||||
)
|
||||
headers: Final = (
|
||||
update_headers_with_filtered_beta(headers=validated_headers, provider=self.provider)
|
||||
if config is not None and config.should_filter_anthropic_beta_headers()
|
||||
else validated_headers
|
||||
)
|
||||
request_body: Final = _BODY_ADAPTER.validate_python(
|
||||
MappingProxyType({**body, "model": self.model, "messages": self.messages})
|
||||
)
|
||||
if self.logging is not None:
|
||||
self.logging.update_from_kwargs(
|
||||
kwargs=self.params.model_dump(),
|
||||
model=self.model,
|
||||
optional_params=body,
|
||||
litellm_params=self.params.model_dump(),
|
||||
custom_llm_provider=self.provider,
|
||||
)
|
||||
self.logging.model_call_details.update(request_body)
|
||||
log_details: Final[NativePreCallDetails] = {
|
||||
"complete_input_dict": request_body,
|
||||
"api_base": base or "",
|
||||
"headers": headers,
|
||||
}
|
||||
log_input: Final[ChatCompletionUserMessage] = {"role": "user", "content": json.dumps(request_body)}
|
||||
self.logging.pre_call(
|
||||
input=[log_input], # mutable-ok: logging callbacks expect a concrete message list
|
||||
api_key=self.api_key,
|
||||
additional_args=log_details,
|
||||
)
|
||||
self.logged = True
|
||||
return PreparedNativeCall(
|
||||
NativeMessagesRequest(
|
||||
model=self.model,
|
||||
body=request_body,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=self.api_key,
|
||||
api_base=base,
|
||||
custom_llm_provider=self.provider,
|
||||
extra_headers=headers,
|
||||
timeout_seconds=timeout_to_seconds(
|
||||
BaseLLMHTTPHandler.resolve_anthropic_messages_timeout(self.params, False, self.provider)
|
||||
),
|
||||
),
|
||||
context=request_context(
|
||||
logging_obj=self.logging,
|
||||
request_model=self.logging.model if self.logging is not None else self.model,
|
||||
litellm_params=self.params.model_dump(),
|
||||
capabilities=NativeRequestCapabilities(
|
||||
execution_mode="async" if self.asynchronous else "sync",
|
||||
stream=self.stream,
|
||||
has_custom_client=self.has_custom_client,
|
||||
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(self.logging),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
def fallback(self) -> MessagesResult:
|
||||
with self.logging.suppress_next_pre_call() if self.logged and self.logging is not None else nullcontext():
|
||||
return self.python()
|
||||
|
||||
async def afallback(self) -> MessagesResponse:
|
||||
with self.logging.suppress_next_pre_call() if self.logged and self.logging is not None else nullcontext():
|
||||
response: Final = self.python()
|
||||
return await response if isinstance(response, Coroutine) else response
|
||||
|
||||
def adapt(self, response: dict[str, object]) -> AnthropicMessagesResponse:
|
||||
_NativeMessagesResponse.model_validate(response)
|
||||
parsed: Final[_BridgedMessagesResponse] = {
|
||||
**_RESPONSE_ADAPTER.validate_python(response),
|
||||
"_hidden_params": {"additional_headers": {"x-litellm-rust": "true"}},
|
||||
}
|
||||
if self.logging is not None:
|
||||
self.logging.post_call(input=self.messages, api_key=self.api_key, original_response=json.dumps(response))
|
||||
return parsed
|
||||
|
||||
|
||||
def dispatch_messages(
|
||||
*,
|
||||
model: str,
|
||||
provider: str,
|
||||
messages: list[dict[str, object]],
|
||||
body: Callable[[], dict[str, object]],
|
||||
params: GenericLiteLLMParams,
|
||||
logging: Logging | None,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
stream: bool,
|
||||
asynchronous: bool,
|
||||
has_custom_client: bool,
|
||||
fallback: Callable[[], MessagesResult],
|
||||
) -> MessagesResult:
|
||||
operation: Final = _MessagesOperation(
|
||||
model,
|
||||
provider,
|
||||
messages,
|
||||
body,
|
||||
params,
|
||||
logging,
|
||||
api_key,
|
||||
api_base,
|
||||
fallback,
|
||||
stream,
|
||||
asynchronous,
|
||||
has_custom_client,
|
||||
)
|
||||
|
||||
def preflight() -> PythonFallback | None:
|
||||
return assess_route(
|
||||
_PREFLIGHT,
|
||||
model,
|
||||
provider,
|
||||
stream=stream,
|
||||
has_custom_client=has_custom_client,
|
||||
has_agentic_hook=BaseLLMHTTPHandler.has_agentic_completion_hook(logging),
|
||||
)
|
||||
|
||||
error_context: Final = BridgeErrorContext(provider=provider, model=model)
|
||||
if asynchronous:
|
||||
return _MESSAGES.ainvoke(
|
||||
prepare=operation.prepare,
|
||||
call=call_native,
|
||||
adapt=operation.adapt,
|
||||
fallback=operation.afallback,
|
||||
preflight=preflight,
|
||||
error_context=error_context,
|
||||
)
|
||||
return _MESSAGES.invoke(
|
||||
prepare=operation.prepare,
|
||||
call=call_native,
|
||||
adapt=operation.adapt,
|
||||
fallback=operation.fallback,
|
||||
preflight=preflight,
|
||||
error_context=error_context,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,277 +1,83 @@
|
|||
"""Thin Python wrapper for the native Rust OCR bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import Final, TypeVar
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model
|
||||
from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse
|
||||
from litellm.rust_bridge import configuration as _configuration
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.protocols import RustAocr, RustOcr
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeOCRRequest,
|
||||
NativeRequestCapabilities,
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
request_context,
|
||||
vertex_options,
|
||||
from . import configuration as _configuration
|
||||
from .protocols import RustAocr, RustOcr, RustRouteDecline
|
||||
from .request import NativeOCRRequest, PreparedNativeCall, call_native
|
||||
from .runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointBinding,
|
||||
EndpointDispatch,
|
||||
assess_route,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
_OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr)
|
||||
_AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr)
|
||||
_HEADERS: Final = TypeAdapter(dict[str, object])
|
||||
ResultT = TypeVar("ResultT")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PreparedOCRRequest:
|
||||
model: str
|
||||
document: dict[str, object]
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
custom_llm_provider: str
|
||||
extra_headers: dict[str, object] | None
|
||||
provider_config: BaseOCRConfig
|
||||
optional_params: dict[str, object]
|
||||
litellm_params: dict[str, object]
|
||||
effective_timeout: float | httpx.Timeout
|
||||
litellm_logging_obj: LiteLLMLoggingObj
|
||||
_OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native(
|
||||
route="ocr",
|
||||
sync=lambda native: native.ocr,
|
||||
asynchronous=lambda native: native.aocr,
|
||||
enabled=_configuration.rust_enabled,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PreparedRustOCRCall:
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
headers: dict[str, object]
|
||||
optional_params: dict[str, object]
|
||||
|
||||
|
||||
_RUST_OCR_PROVIDERS: Final = frozenset(
|
||||
{
|
||||
"mistral",
|
||||
"azure_ai",
|
||||
"vertex_ai",
|
||||
}
|
||||
_PREFLIGHT: Final[EndpointBinding[RustRouteDecline]] = EndpointBinding.native(
|
||||
route="ocr",
|
||||
select=lambda native: native.ocr_decline,
|
||||
enabled=_configuration.rust_enabled,
|
||||
)
|
||||
|
||||
|
||||
def load_rust_ocr() -> RustOcr | None:
|
||||
return _OCR.load()
|
||||
return _OCR.sync.load()
|
||||
|
||||
|
||||
def load_rust_aocr() -> RustAocr | None:
|
||||
return _AOCR.load()
|
||||
return _OCR.asynchronous.load()
|
||||
|
||||
|
||||
def _rust_ocr_supported(prepared_request: PreparedOCRRequest) -> bool:
|
||||
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
|
||||
return False
|
||||
if not prepared_request.provider_config.supports_rust_bridge():
|
||||
return False
|
||||
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
|
||||
|
||||
|
||||
def _ocr_input_source_kind(document: dict[str, object]) -> str:
|
||||
if "document_url" in document:
|
||||
return "document_url"
|
||||
if "image_url" in document:
|
||||
return "image_url"
|
||||
if "file" in document:
|
||||
return "file"
|
||||
return "inline"
|
||||
|
||||
|
||||
def _ocr_request_format(optional_params: dict[str, object]) -> str | None:
|
||||
value = optional_params.get(OCR_REQUEST_FORMAT_PARAM)
|
||||
return value if isinstance(value, str) else None
|
||||
|
||||
|
||||
def _rust_bridge_optional_params(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> dict[str, object]:
|
||||
if prepared_request.custom_llm_provider != "vertex_ai":
|
||||
return prepared_request.optional_params
|
||||
vertex_project: Final = (
|
||||
prepared_request.litellm_params.get("vertex_project")
|
||||
or prepared_request.litellm_params.get("vertex_ai_project")
|
||||
or litellm.vertex_project
|
||||
or resolve_secret("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_location: Final = (
|
||||
prepared_request.litellm_params.get("vertex_location")
|
||||
or prepared_request.litellm_params.get("vertex_ai_location")
|
||||
or litellm.vertex_location
|
||||
or resolve_secret("VERTEXAI_LOCATION")
|
||||
or resolve_secret("VERTEX_LOCATION")
|
||||
)
|
||||
return {
|
||||
**prepared_request.optional_params,
|
||||
**{
|
||||
name: value
|
||||
for name, value in (("vertex_project", vertex_project), ("vertex_location", vertex_location))
|
||||
if value is not None
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _rust_bridge_api_base(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> str | None:
|
||||
if prepared_request.api_base is not None:
|
||||
return prepared_request.api_base
|
||||
if prepared_request.custom_llm_provider == "azure_ai":
|
||||
if is_azure_document_intelligence_model(prepared_request.model):
|
||||
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
return resolve_secret("AZURE_AI_API_BASE")
|
||||
return None
|
||||
|
||||
|
||||
def _prepare_rust_ocr_call(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> _PreparedRustOCRCall:
|
||||
provider_config: Final = prepared_request.provider_config
|
||||
api_key_env_var: Final = provider_config.get_api_key_env_var()
|
||||
resolved_api_key: Final = prepared_request.api_key or (
|
||||
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
|
||||
)
|
||||
resolved_headers: Final = _HEADERS.validate_python(
|
||||
provider_config.validate_environment(
|
||||
headers=prepared_request.extra_headers or {},
|
||||
model=prepared_request.model,
|
||||
api_key=resolved_api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
)
|
||||
resolved_complete_url: Final = provider_config.get_complete_url(
|
||||
api_base=prepared_request.api_base,
|
||||
model=prepared_request.model,
|
||||
optional_params=prepared_request.optional_params,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
|
||||
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
|
||||
prepared_request.litellm_logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=resolved_api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": {
|
||||
"model": prepared_request.model,
|
||||
"document": prepared_request.document,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": resolved_complete_url,
|
||||
"headers": resolved_headers,
|
||||
},
|
||||
)
|
||||
return _PreparedRustOCRCall(
|
||||
api_key=resolved_api_key,
|
||||
api_base=rust_api_base,
|
||||
headers=resolved_headers,
|
||||
optional_params=rust_optional_params,
|
||||
def dispatch_ocr(
|
||||
*,
|
||||
prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]],
|
||||
fallback: Callable[[], ResultT],
|
||||
adapt: Callable[[Mapping[str, object]], ResultT],
|
||||
model: str,
|
||||
provider: str,
|
||||
eligible: bool = True,
|
||||
request_format: str | None = None,
|
||||
) -> ResultT:
|
||||
return _OCR.invoke(
|
||||
prepare=prepare,
|
||||
call=call_native,
|
||||
fallback=fallback,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=provider, model=model),
|
||||
eligible=eligible,
|
||||
preflight=lambda: assess_route(_PREFLIGHT, model, provider, request_format=request_format),
|
||||
)
|
||||
|
||||
|
||||
def attempt_ocr(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> DispatchResult[OCRResponse]:
|
||||
return attempt(
|
||||
load=_OCR.load,
|
||||
enabled=_configuration.rust_enabled(),
|
||||
prepare=lambda: _prepare_rust_ocr_call(
|
||||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
),
|
||||
call=lambda native, prepared: call_native(
|
||||
native,
|
||||
PreparedNativeCall(
|
||||
request=NativeOCRRequest(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=prepared.optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
|
||||
vertex=vertex_options(prepared.optional_params),
|
||||
),
|
||||
context=request_context(
|
||||
logging_obj=prepared_request.litellm_logging_obj,
|
||||
request_model=prepared_request.model,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
capabilities=NativeRequestCapabilities(
|
||||
execution_mode="sync",
|
||||
input_source_kind=_ocr_input_source_kind(prepared_request.document),
|
||||
request_format=_ocr_request_format(prepared_request.optional_params),
|
||||
native_response_format=(
|
||||
prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
adapt=OCRResponse.model_validate,
|
||||
eligible=_rust_ocr_supported(prepared_request),
|
||||
)
|
||||
|
||||
|
||||
async def aattempt_ocr(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> DispatchResult[OCRResponse]:
|
||||
return await aattempt(
|
||||
load=_AOCR.load,
|
||||
enabled=_configuration.rust_enabled(),
|
||||
prepare=lambda: _prepare_rust_ocr_call(
|
||||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
),
|
||||
call=lambda native, prepared: call_native(
|
||||
native,
|
||||
PreparedNativeCall(
|
||||
request=NativeOCRRequest(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=prepared.optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
|
||||
vertex=vertex_options(prepared.optional_params),
|
||||
),
|
||||
context=request_context(
|
||||
logging_obj=prepared_request.litellm_logging_obj,
|
||||
request_model=prepared_request.model,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
capabilities=NativeRequestCapabilities(
|
||||
execution_mode="async",
|
||||
input_source_kind=_ocr_input_source_kind(prepared_request.document),
|
||||
request_format=_ocr_request_format(prepared_request.optional_params),
|
||||
native_response_format=(
|
||||
prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
),
|
||||
adapt=OCRResponse.model_validate,
|
||||
eligible=_rust_ocr_supported(prepared_request),
|
||||
async def adispatch_ocr(
|
||||
*,
|
||||
prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]],
|
||||
fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[Mapping[str, object]], ResultT],
|
||||
model: str,
|
||||
provider: str,
|
||||
eligible: bool = True,
|
||||
request_format: str | None = None,
|
||||
) -> ResultT:
|
||||
return await _OCR.ainvoke(
|
||||
prepare=prepare,
|
||||
call=call_native,
|
||||
fallback=fallback,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=provider, model=model),
|
||||
eligible=eligible,
|
||||
preflight=lambda: assess_route(_PREFLIGHT, model, provider, request_format=request_format),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -31,6 +31,9 @@ class RustChatCompletionsDecline(Protocol):
|
|||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object] | None,
|
||||
custom_llm_provider: str | None,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
) -> str | None: ...
|
||||
|
||||
|
||||
|
|
@ -53,6 +56,16 @@ class RustResponsesWebSocketConnection(Protocol):
|
|||
) -> RustResponsesWebSocket: ...
|
||||
|
||||
|
||||
class RustRouteDecline(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> str | None: ...
|
||||
|
||||
|
||||
class NativeModule(Protocol):
|
||||
@property
|
||||
def chat_completions(self) -> RustChatCompletions: ...
|
||||
|
|
@ -89,3 +102,15 @@ class NativeModule(Protocol):
|
|||
|
||||
@property
|
||||
def atranscription(self) -> RustAtranscription: ...
|
||||
|
||||
@property
|
||||
def ocr_decline(self) -> RustRouteDecline: ...
|
||||
|
||||
@property
|
||||
def messages_decline(self) -> RustRouteDecline: ...
|
||||
|
||||
@property
|
||||
def transcription_decline(self) -> RustRouteDecline: ...
|
||||
|
||||
@property
|
||||
def responses_websocket_decline(self) -> RustRouteDecline: ...
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ class NativeBedrockOptions:
|
|||
@dataclass(frozen=True, slots=True)
|
||||
class NativeAnthropicOptions:
|
||||
user_id: str | None = None
|
||||
has_user_id: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -57,7 +58,10 @@ def bedrock_options(params: Mapping[str, object]) -> NativeBedrockOptions:
|
|||
def anthropic_options(litellm_params: Mapping[str, object] | None) -> NativeAnthropicOptions:
|
||||
metadata = None if litellm_params is None else litellm_params.get("metadata")
|
||||
user_id = metadata.get("user_id") if isinstance(metadata, Mapping) else None
|
||||
return NativeAnthropicOptions(user_id=user_id if isinstance(user_id, str) else None)
|
||||
return NativeAnthropicOptions(
|
||||
user_id=user_id if isinstance(user_id, str) else None,
|
||||
has_user_id=user_id is not None,
|
||||
)
|
||||
|
||||
|
||||
def vertex_options(params: Mapping[str, object]) -> NativeVertexOptions:
|
||||
|
|
@ -69,6 +73,15 @@ def vertex_options(params: Mapping[str, object]) -> NativeVertexOptions:
|
|||
)
|
||||
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
|
||||
class NativePreCallDetails(TypedDict):
|
||||
complete_input_dict: ReadOnly[Mapping[str, object]]
|
||||
api_base: ReadOnly[str]
|
||||
headers: ReadOnly[Mapping[str, object] | None]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NativeRequestOptions:
|
||||
api_key: str | None = None
|
||||
|
|
|
|||
|
|
@ -2,18 +2,19 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
from typing import Final
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable
|
||||
from contextlib import AbstractAsyncContextManager, AsyncExitStack, asynccontextmanager
|
||||
from typing import Final, Protocol
|
||||
|
||||
import httpx
|
||||
from websockets.exceptions import ConnectionClosedOK
|
||||
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.protocols import (
|
||||
RustResponsesWebSocket,
|
||||
RustResponsesWebSocketConnection,
|
||||
RustRouteDecline,
|
||||
)
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestCapabilities,
|
||||
|
|
@ -24,18 +25,38 @@ from litellm.rust_bridge.request import (
|
|||
call_native,
|
||||
with_capabilities,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointBinding,
|
||||
assess_route,
|
||||
async_none,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
_RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding(
|
||||
lambda native: native.ResponsesWebSocketConnection,
|
||||
_RESPONSES_WEBSOCKET: Final[EndpointBinding[RustResponsesWebSocketConnection]] = EndpointBinding.native(
|
||||
route="responses_websocket",
|
||||
select=lambda native: native.ResponsesWebSocketConnection,
|
||||
enabled=rust_enabled,
|
||||
)
|
||||
|
||||
|
||||
_PREFLIGHT: Final[EndpointBinding[RustRouteDecline]] = EndpointBinding.native(
|
||||
route="responses_websocket",
|
||||
select=lambda native: native.responses_websocket_decline,
|
||||
enabled=rust_enabled,
|
||||
)
|
||||
|
||||
|
||||
def set_rust_responses_websocket(
|
||||
*,
|
||||
connection: RustResponsesWebSocketConnection | None | Unchanged = UNCHANGED,
|
||||
decline: RustRouteDecline | None | Unchanged = UNCHANGED,
|
||||
) -> None:
|
||||
if not isinstance(decline, Unchanged):
|
||||
if decline is None:
|
||||
_PREFLIGHT.reset()
|
||||
else:
|
||||
_PREFLIGHT.override(decline)
|
||||
if not isinstance(connection, Unchanged):
|
||||
if connection is None:
|
||||
_RESPONSES_WEBSOCKET.reset()
|
||||
|
|
@ -43,7 +64,13 @@ def set_rust_responses_websocket(
|
|||
_RESPONSES_WEBSOCKET.override(connection)
|
||||
|
||||
|
||||
class ConnectionAdapter:
|
||||
class Connection(Protocol):
|
||||
async def send(self, text: str) -> None: ...
|
||||
async def recv(self) -> str | bytes: ...
|
||||
async def close(self) -> None: ...
|
||||
|
||||
|
||||
class _ConnectionAdapter:
|
||||
def __init__(self, connection: RustResponsesWebSocket):
|
||||
self._connection: Final[RustResponsesWebSocket] = connection
|
||||
|
||||
|
|
@ -65,54 +92,65 @@ async def connect(
|
|||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
websocket_mode: str = "native",
|
||||
requires_connection: bool = True,
|
||||
model: str = "responses websocket",
|
||||
provider: str = "openai",
|
||||
fallback: Callable[[], Awaitable[Connection | None]] = async_none,
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> DispatchResult[ConnectionAdapter]:
|
||||
return await aattempt(
|
||||
load=_RESPONSES_WEBSOCKET.load,
|
||||
enabled=rust_enabled(),
|
||||
eligible=True,
|
||||
) -> Connection | None:
|
||||
return await _RESPONSES_WEBSOCKET.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeResponsesWebSocketRequest(url=url),
|
||||
options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)),
|
||||
NativeResponsesWebSocketRequest(
|
||||
url=url,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
extra_headers=headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
custom_llm_provider=provider,
|
||||
),
|
||||
context=with_capabilities(
|
||||
context or NativeRequestContext(),
|
||||
NativeRequestCapabilities(
|
||||
execution_mode="async",
|
||||
websocket_mode=websocket_mode,
|
||||
requires_connection=requires_connection,
|
||||
websocket_mode="native",
|
||||
requires_connection=True,
|
||||
),
|
||||
),
|
||||
),
|
||||
call=lambda connection_type, prepared: call_native(connection_type.connect, prepared),
|
||||
adapt=ConnectionAdapter,
|
||||
call=lambda connection_type, request: call_native(connection_type.connect, request),
|
||||
preflight=lambda: assess_route(_PREFLIGHT, model, provider),
|
||||
fallback=fallback,
|
||||
adapt=_ConnectionAdapter,
|
||||
error_context=BridgeErrorContext(provider=provider, model=model),
|
||||
)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _connection_context(connection: ConnectionAdapter) -> AsyncGenerator[ConnectionAdapter, None]:
|
||||
try:
|
||||
yield connection
|
||||
finally:
|
||||
await connection.close()
|
||||
|
||||
|
||||
async def managed_connect(
|
||||
async def open_connection(
|
||||
*,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
websocket_mode: str = "managed",
|
||||
requires_connection: bool = True,
|
||||
model: str,
|
||||
provider: str,
|
||||
fallback: Callable[[], AbstractAsyncContextManager[Connection]],
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]:
|
||||
result: Final = await connect(
|
||||
url=url,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
websocket_mode=websocket_mode,
|
||||
requires_connection=requires_connection,
|
||||
context=context,
|
||||
)
|
||||
return adapt_result(result, _connection_context)
|
||||
) -> AsyncGenerator[Connection]:
|
||||
async with AsyncExitStack() as stack:
|
||||
|
||||
async def python_connection() -> Connection:
|
||||
return await stack.enter_async_context(fallback())
|
||||
|
||||
backend: Final = await connect(
|
||||
url=url,
|
||||
headers=headers,
|
||||
timeout=timeout,
|
||||
model=model,
|
||||
provider=provider,
|
||||
fallback=python_connection,
|
||||
context=context,
|
||||
)
|
||||
if backend is None:
|
||||
raise RuntimeError("WebSocket connection returned no connection")
|
||||
if isinstance(backend, _ConnectionAdapter):
|
||||
stack.push_async_callback(backend.close)
|
||||
yield backend
|
||||
|
|
|
|||
|
|
@ -1,14 +1,36 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from typing import Final, Generic, TypeAlias, TypeVar
|
||||
from typing import Final, Generic, NoReturn, Protocol, TypeAlias, TypeVar, assert_never
|
||||
|
||||
from litellm.exceptions import APIError, AuthenticationError, InternalServerError, RateLimitError
|
||||
from litellm.rust_bridge.bindings import (
|
||||
UNCHANGED,
|
||||
NativeBinding,
|
||||
Unchanged,
|
||||
native_declined_types,
|
||||
native_upstream_types,
|
||||
)
|
||||
from litellm.rust_bridge.protocols import NativeModule, RustRouteDecline
|
||||
from litellm.rust_bridge.request import NativeRequestCapabilities, NativeRequestContext
|
||||
|
||||
BindingT = TypeVar("BindingT")
|
||||
SelectedT = TypeVar("SelectedT")
|
||||
SelectedSyncT = TypeVar("SelectedSyncT")
|
||||
SelectedAsyncT = TypeVar("SelectedAsyncT")
|
||||
NativeT = TypeVar("NativeT")
|
||||
RequestT = TypeVar("RequestT")
|
||||
ResultT = TypeVar("ResultT")
|
||||
SyncBindingT = TypeVar("SyncBindingT")
|
||||
AsyncBindingT = TypeVar("AsyncBindingT")
|
||||
|
||||
|
||||
class PythonFallbackReason(Enum):
|
||||
NATIVE_DISABLED = "native_disabled"
|
||||
NATIVE_UNAVAILABLE = "native_unavailable"
|
||||
NATIVE_DECLINED = "native_declined"
|
||||
|
||||
|
||||
class NativeSkipReason(Enum):
|
||||
|
|
@ -24,6 +46,12 @@ class Handled(Generic[ResultT]):
|
|||
value: ResultT
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PythonFallback:
|
||||
reason: PythonFallbackReason
|
||||
detail: str | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NativeSkipped:
|
||||
reason: NativeSkipReason
|
||||
|
|
@ -35,54 +63,470 @@ class NativeFailed:
|
|||
error: Exception
|
||||
|
||||
|
||||
DispatchResult: TypeAlias = Handled[ResultT] | NativeSkipped | NativeFailed
|
||||
DispatchResult: TypeAlias = Handled[ResultT] | PythonFallback | NativeSkipped | NativeFailed
|
||||
|
||||
|
||||
def _select(load: Callable[[], BindingT | None], enabled: bool, eligible: bool) -> BindingT | NativeSkipped:
|
||||
if not enabled:
|
||||
return NativeSkipped(NativeSkipReason.DISABLED)
|
||||
if not eligible:
|
||||
return NativeSkipped(NativeSkipReason.INELIGIBLE)
|
||||
binding: Final = load()
|
||||
return NativeSkipped(NativeSkipReason.UNAVAILABLE) if binding is None else binding
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class BridgeErrorContext:
|
||||
provider: str
|
||||
model: str
|
||||
|
||||
|
||||
def attempt(
|
||||
*,
|
||||
load: Callable[[], BindingT | None],
|
||||
enabled: bool,
|
||||
eligible: bool,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[BindingT, RequestT], NativeT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
) -> DispatchResult[ResultT]:
|
||||
binding: Final = _select(load, enabled, eligible)
|
||||
if isinstance(binding, NativeSkipped):
|
||||
class RustEnablement(Protocol):
|
||||
def __call__(self) -> bool: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EndpointBinding(Generic[BindingT]):
|
||||
route: str
|
||||
load: Callable[[], BindingT | None]
|
||||
enabled: RustEnablement
|
||||
_native_binding: NativeBinding[BindingT] | None = field(default=None, repr=False)
|
||||
|
||||
@staticmethod
|
||||
def native(
|
||||
*,
|
||||
route: str,
|
||||
select: Callable[[NativeModule], SelectedT],
|
||||
enabled: RustEnablement,
|
||||
) -> EndpointBinding[SelectedT]:
|
||||
binding: Final = NativeBinding(select)
|
||||
return EndpointBinding(
|
||||
route=route,
|
||||
load=binding.load,
|
||||
enabled=enabled,
|
||||
_native_binding=binding,
|
||||
)
|
||||
|
||||
def override(self, value: BindingT | None) -> None:
|
||||
if self._native_binding is None:
|
||||
raise RuntimeError("only native Rust bridges support binding overrides")
|
||||
self._native_binding.override(value)
|
||||
|
||||
def reset(self) -> None:
|
||||
if self._native_binding is None:
|
||||
raise RuntimeError("only native Rust bridges support binding resets")
|
||||
self._native_binding.reset()
|
||||
|
||||
def _attempt(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[BindingT, RequestT], NativeT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> DispatchResult[ResultT]:
|
||||
binding_or_fallback: Final = self._binding_or_python_fallback(
|
||||
eligible=eligible,
|
||||
)
|
||||
if isinstance(binding_or_fallback, PythonFallback):
|
||||
return binding_or_fallback
|
||||
preflight_result: Final = preflight() if preflight is not None else None
|
||||
if preflight_result is not None:
|
||||
return preflight_result
|
||||
return self._attempt_call(
|
||||
call=lambda: call(binding_or_fallback, prepare()),
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
)
|
||||
|
||||
async def _aattempt(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[BindingT, RequestT], Awaitable[NativeT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> DispatchResult[ResultT]:
|
||||
binding_or_fallback: Final = self._binding_or_python_fallback(
|
||||
eligible=eligible,
|
||||
)
|
||||
if isinstance(binding_or_fallback, PythonFallback):
|
||||
return binding_or_fallback
|
||||
preflight_result: Final = preflight() if preflight is not None else None
|
||||
if preflight_result is not None:
|
||||
return preflight_result
|
||||
return await self._attempt_acall(
|
||||
call=lambda: call(binding_or_fallback, prepare()),
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
)
|
||||
|
||||
def invoke(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[BindingT, RequestT], NativeT],
|
||||
fallback: Callable[[], ResultT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
result: Final = self._attempt(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
match result:
|
||||
case Handled(value=value):
|
||||
return value
|
||||
case PythonFallback():
|
||||
return fallback()
|
||||
case _ as unreachable:
|
||||
assert_never(unreachable)
|
||||
|
||||
async def ainvoke(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[BindingT, RequestT], Awaitable[NativeT]],
|
||||
fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
result: Final = await self._aattempt(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
match result:
|
||||
case Handled(value=value):
|
||||
return value
|
||||
case PythonFallback():
|
||||
return await fallback()
|
||||
case _ as unreachable:
|
||||
assert_never(unreachable)
|
||||
|
||||
def assess(
|
||||
self,
|
||||
*,
|
||||
check: Callable[[BindingT], str | None],
|
||||
) -> PythonFallback | None:
|
||||
binding: Final = self._binding_or_python_fallback(eligible=True)
|
||||
if isinstance(binding, PythonFallback):
|
||||
return binding
|
||||
reason: Final = check(binding)
|
||||
return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, reason) if reason is not None else None
|
||||
|
||||
def accepts(
|
||||
self,
|
||||
*,
|
||||
check: Callable[[BindingT], str | None],
|
||||
eligible: bool = True,
|
||||
) -> bool:
|
||||
binding_or_fallback: Final = self._binding_or_python_fallback(
|
||||
eligible=eligible,
|
||||
)
|
||||
if isinstance(binding_or_fallback, PythonFallback):
|
||||
return False
|
||||
try:
|
||||
reason: Final = check(binding_or_fallback)
|
||||
except Exception: # noqa: BLE001 # preflight performs no provider I/O, so Python handoff is safe
|
||||
return False
|
||||
return reason is None
|
||||
|
||||
def require(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[BindingT, RequestT], NativeT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
result: Final = self._attempt(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
match result:
|
||||
case Handled(value=value):
|
||||
return value
|
||||
case PythonFallback():
|
||||
self._raise_required(result)
|
||||
case _ as unreachable:
|
||||
assert_never(unreachable)
|
||||
|
||||
async def arequire(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[BindingT, RequestT], Awaitable[NativeT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
result: Final = await self._aattempt(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
match result:
|
||||
case Handled(value=value):
|
||||
return value
|
||||
case PythonFallback():
|
||||
self._raise_required(result)
|
||||
case _ as unreachable:
|
||||
assert_never(unreachable)
|
||||
|
||||
def can_attempt(
|
||||
self,
|
||||
*,
|
||||
eligible: bool = True,
|
||||
) -> bool:
|
||||
return not isinstance(
|
||||
self._binding_or_python_fallback(eligible=eligible),
|
||||
PythonFallback,
|
||||
)
|
||||
|
||||
def _raise_required(self, fallback: PythonFallback) -> NoReturn:
|
||||
detail: Final = f": {fallback.detail}" if fallback.detail else ""
|
||||
reason: Final = _required_reason(fallback.reason)
|
||||
raise RuntimeError(f"native {self.route} endpoint {reason}{detail}")
|
||||
|
||||
def _binding_or_python_fallback(
|
||||
self,
|
||||
*,
|
||||
eligible: bool,
|
||||
) -> BindingT | PythonFallback:
|
||||
if not eligible or not self.enabled():
|
||||
return PythonFallback(PythonFallbackReason.NATIVE_DISABLED)
|
||||
binding: Final = self.load()
|
||||
if binding is None:
|
||||
return PythonFallback(PythonFallbackReason.NATIVE_UNAVAILABLE)
|
||||
return binding
|
||||
try:
|
||||
value: Final = call(binding, prepare())
|
||||
except Exception as error: # noqa: BLE001 # orchestration applies the endpoint's declared error policy
|
||||
return NativeFailed(error)
|
||||
return Handled(adapt(value))
|
||||
|
||||
def _attempt_call(
|
||||
self,
|
||||
*,
|
||||
call: Callable[[], NativeT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
) -> DispatchResult[ResultT]:
|
||||
declined: Final = native_declined_types()
|
||||
upstream: Final = native_upstream_types()
|
||||
if not declined or not upstream:
|
||||
return Handled(adapt(call()))
|
||||
try:
|
||||
value: Final = call()
|
||||
except declined as error:
|
||||
return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error))
|
||||
except upstream as error:
|
||||
self._raise_upstream(error, error_context)
|
||||
return Handled(adapt(value))
|
||||
|
||||
async def _attempt_acall(
|
||||
self,
|
||||
*,
|
||||
call: Callable[[], Awaitable[NativeT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
) -> DispatchResult[ResultT]:
|
||||
declined: Final = native_declined_types()
|
||||
upstream: Final = native_upstream_types()
|
||||
if not declined or not upstream:
|
||||
return Handled(adapt(await call()))
|
||||
try:
|
||||
value: Final = await call()
|
||||
except declined as error:
|
||||
return PythonFallback(PythonFallbackReason.NATIVE_DECLINED, _error_message(error))
|
||||
except upstream as error:
|
||||
self._raise_upstream(error, error_context)
|
||||
return Handled(adapt(value))
|
||||
|
||||
def _raise_upstream(self, error: BaseException, error_context: BridgeErrorContext) -> NoReturn:
|
||||
args: Final[tuple[object, ...]] = error.args
|
||||
attribute_status: Final = getattr(error, "status_code", None)
|
||||
attribute_message: Final = getattr(error, "message", None)
|
||||
status_value: Final = attribute_status if isinstance(attribute_status, int) else (args[0] if args else 0)
|
||||
message_value: Final = (
|
||||
attribute_message if isinstance(attribute_message, str) else (args[1] if len(args) > 1 else str(error))
|
||||
)
|
||||
status: Final = status_value if isinstance(status_value, int) else 0
|
||||
message: Final = message_value if isinstance(message_value, str) else str(message_value)
|
||||
error_message: Final = f"litellm rust {self.route}: {message}"
|
||||
if status == 401:
|
||||
raise AuthenticationError(
|
||||
message=error_message,
|
||||
llm_provider=error_context.provider,
|
||||
model=error_context.model,
|
||||
) from error
|
||||
if status == 429:
|
||||
raise RateLimitError(
|
||||
message=error_message,
|
||||
llm_provider=error_context.provider,
|
||||
model=error_context.model,
|
||||
) from error
|
||||
if status == 500:
|
||||
raise InternalServerError(
|
||||
message=error_message,
|
||||
llm_provider=error_context.provider,
|
||||
model=error_context.model,
|
||||
) from error
|
||||
raise APIError(
|
||||
status_code=status or 500,
|
||||
message=error_message,
|
||||
llm_provider=error_context.provider,
|
||||
model=error_context.model,
|
||||
) from error
|
||||
|
||||
|
||||
async def aattempt(
|
||||
*,
|
||||
load: Callable[[], BindingT | None],
|
||||
enabled: bool,
|
||||
eligible: bool,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[BindingT, RequestT], Awaitable[NativeT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
) -> DispatchResult[ResultT]:
|
||||
binding: Final = _select(load, enabled, eligible)
|
||||
if isinstance(binding, NativeSkipped):
|
||||
return binding
|
||||
try:
|
||||
value: Final = await call(binding, prepare())
|
||||
except Exception as error: # noqa: BLE001 # orchestration applies the endpoint's declared error policy
|
||||
return NativeFailed(error)
|
||||
return Handled(adapt(value))
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class EndpointDispatch(Generic[SyncBindingT, AsyncBindingT]):
|
||||
sync: EndpointBinding[SyncBindingT]
|
||||
asynchronous: EndpointBinding[AsyncBindingT]
|
||||
|
||||
@staticmethod
|
||||
def native(
|
||||
*,
|
||||
route: str,
|
||||
sync: Callable[[NativeModule], SelectedSyncT],
|
||||
asynchronous: Callable[[NativeModule], SelectedAsyncT],
|
||||
enabled: RustEnablement,
|
||||
) -> EndpointDispatch[SelectedSyncT, SelectedAsyncT]:
|
||||
return EndpointDispatch(
|
||||
sync=EndpointBinding.native(route=route, select=sync, enabled=enabled),
|
||||
asynchronous=EndpointBinding.native(
|
||||
route=route,
|
||||
select=asynchronous,
|
||||
enabled=enabled,
|
||||
),
|
||||
)
|
||||
|
||||
def override(
|
||||
self,
|
||||
*,
|
||||
sync: SyncBindingT | None | Unchanged = UNCHANGED,
|
||||
asynchronous: AsyncBindingT | None | Unchanged = UNCHANGED,
|
||||
) -> None:
|
||||
if not isinstance(sync, Unchanged):
|
||||
self.sync.override(sync)
|
||||
if not isinstance(asynchronous, Unchanged):
|
||||
self.asynchronous.override(asynchronous)
|
||||
|
||||
def reset(self) -> None:
|
||||
self.sync.reset()
|
||||
self.asynchronous.reset()
|
||||
|
||||
def invoke(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[SyncBindingT, RequestT], NativeT],
|
||||
fallback: Callable[[], ResultT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
return self.sync.invoke(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
fallback=fallback,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
|
||||
async def ainvoke(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[AsyncBindingT, RequestT], Awaitable[NativeT]],
|
||||
fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
return await self.asynchronous.ainvoke(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
fallback=fallback,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
|
||||
def require(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[SyncBindingT, RequestT], NativeT],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
return self.sync.require(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
|
||||
async def arequire(
|
||||
self,
|
||||
*,
|
||||
prepare: Callable[[], RequestT],
|
||||
call: Callable[[AsyncBindingT, RequestT], Awaitable[NativeT]],
|
||||
adapt: Callable[[NativeT], ResultT],
|
||||
error_context: BridgeErrorContext,
|
||||
eligible: bool = True,
|
||||
preflight: Callable[[], PythonFallback | None] | None = None,
|
||||
) -> ResultT:
|
||||
return await self.asynchronous.arequire(
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
error_context=error_context,
|
||||
eligible=eligible,
|
||||
preflight=preflight,
|
||||
)
|
||||
|
||||
|
||||
def _error_message(error: BaseException) -> str:
|
||||
reason: Final[object] = error.args[0] if error.args else str(error)
|
||||
return reason if isinstance(reason, str) else str(reason)
|
||||
|
||||
|
||||
def _required_reason(reason: PythonFallbackReason) -> str:
|
||||
match reason:
|
||||
case PythonFallbackReason.NATIVE_DISABLED:
|
||||
return "is disabled"
|
||||
case PythonFallbackReason.NATIVE_UNAVAILABLE:
|
||||
return "is unavailable"
|
||||
case PythonFallbackReason.NATIVE_DECLINED:
|
||||
return "declined the request"
|
||||
|
||||
|
||||
def always_enabled() -> bool:
|
||||
return True
|
||||
|
||||
|
||||
def identity(value: ResultT) -> ResultT:
|
||||
|
|
@ -93,3 +537,34 @@ def adapt_result(result: DispatchResult[NativeT], adapt: Callable[[NativeT], Res
|
|||
if isinstance(result, Handled):
|
||||
return Handled(adapt(result.value))
|
||||
return result
|
||||
|
||||
|
||||
async def async_none() -> None:
|
||||
return None
|
||||
|
||||
|
||||
def assess_route(
|
||||
binding: EndpointBinding[RustRouteDecline],
|
||||
model: str,
|
||||
provider: str,
|
||||
*,
|
||||
stream: bool = False,
|
||||
has_agentic_hook: bool = False,
|
||||
has_custom_client: bool = False,
|
||||
request_format: str | None = None,
|
||||
) -> PythonFallback | None:
|
||||
context: Final = NativeRequestContext(
|
||||
capabilities=NativeRequestCapabilities(
|
||||
stream=stream,
|
||||
has_agentic_hook=has_agentic_hook,
|
||||
has_custom_client=has_custom_client,
|
||||
request_format=request_format,
|
||||
)
|
||||
)
|
||||
return binding.assess(
|
||||
check=lambda decline: decline(
|
||||
model,
|
||||
provider,
|
||||
context=context,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,12 +1,25 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
from collections.abc import Callable, Coroutine
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from io import IOBase
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
|
||||
from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.protocols import RustAtranscription, RustRouteDecline, RustTranscription
|
||||
from litellm.rust_bridge.request import (
|
||||
NativePreCallDetails,
|
||||
NativeRequestCapabilities,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
|
|
@ -14,38 +27,67 @@ from litellm.rust_bridge.request import (
|
|||
PreparedNativeCall,
|
||||
bedrock_options,
|
||||
call_native,
|
||||
request_context,
|
||||
with_capabilities,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointBinding,
|
||||
EndpointDispatch,
|
||||
PythonFallback,
|
||||
always_enabled,
|
||||
assess_route,
|
||||
async_none,
|
||||
identity,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
_TRANSCRIPTION: Final[NativeBinding[RustTranscription]] = NativeBinding(lambda native: native.transcription)
|
||||
_ATRANSCRIPTION: Final[NativeBinding[RustAtranscription]] = NativeBinding(lambda native: native.atranscription)
|
||||
_TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] = EndpointDispatch.native(
|
||||
route="audio transcription",
|
||||
sync=lambda native: native.transcription,
|
||||
asynchronous=lambda native: native.atranscription,
|
||||
enabled=always_enabled,
|
||||
)
|
||||
|
||||
|
||||
_PREFLIGHT: Final[EndpointBinding[RustRouteDecline]] = EndpointBinding.native(
|
||||
route="transcription",
|
||||
select=lambda native: native.transcription_decline,
|
||||
enabled=always_enabled,
|
||||
)
|
||||
|
||||
|
||||
def configure_rust_transcription(
|
||||
*,
|
||||
transcription: RustTranscription | None | Unchanged = UNCHANGED,
|
||||
atranscription: RustAtranscription | None | Unchanged = UNCHANGED,
|
||||
decline: RustRouteDecline | None | Unchanged = UNCHANGED,
|
||||
) -> None:
|
||||
if not isinstance(decline, Unchanged):
|
||||
if decline is None:
|
||||
_PREFLIGHT.reset()
|
||||
else:
|
||||
_PREFLIGHT.override(decline)
|
||||
if not isinstance(transcription, Unchanged):
|
||||
if transcription is None:
|
||||
_TRANSCRIPTION.reset()
|
||||
_TRANSCRIPTION.sync.reset()
|
||||
else:
|
||||
_TRANSCRIPTION.override(transcription)
|
||||
_TRANSCRIPTION.sync.override(transcription)
|
||||
if not isinstance(atranscription, Unchanged):
|
||||
if atranscription is None:
|
||||
_ATRANSCRIPTION.reset()
|
||||
_TRANSCRIPTION.asynchronous.reset()
|
||||
else:
|
||||
_ATRANSCRIPTION.override(atranscription)
|
||||
_TRANSCRIPTION.asynchronous.override(atranscription)
|
||||
|
||||
|
||||
def load_rust_transcription() -> RustTranscription | None:
|
||||
return _TRANSCRIPTION.load()
|
||||
return _TRANSCRIPTION.sync.load()
|
||||
|
||||
|
||||
def load_rust_atranscription() -> RustAtranscription | None:
|
||||
return _ATRANSCRIPTION.load()
|
||||
return _TRANSCRIPTION.asynchronous.load()
|
||||
|
||||
|
||||
def transcription(
|
||||
|
|
@ -62,13 +104,14 @@ def transcription(
|
|||
has_custom_client: bool = False,
|
||||
input_source_kind: str | None = None,
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> DispatchResult[dict[str, object]]:
|
||||
return attempt(
|
||||
load=_TRANSCRIPTION.load,
|
||||
enabled=True,
|
||||
eligible=True,
|
||||
) -> dict[str, object] | None:
|
||||
return _TRANSCRIPTION.invoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeTranscriptionRequest(model=model, audio=audio, optional_params=optional_params),
|
||||
NativeTranscriptionRequest(
|
||||
model=model,
|
||||
audio=audio,
|
||||
optional_params=optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
|
|
@ -88,7 +131,10 @@ def transcription(
|
|||
),
|
||||
),
|
||||
call=call_native,
|
||||
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -106,13 +152,14 @@ async def atranscription(
|
|||
has_custom_client: bool = False,
|
||||
input_source_kind: str | None = None,
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> DispatchResult[dict[str, object]]:
|
||||
return await aattempt(
|
||||
load=_ATRANSCRIPTION.load,
|
||||
enabled=True,
|
||||
eligible=True,
|
||||
) -> dict[str, object] | None:
|
||||
return await _TRANSCRIPTION.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeTranscriptionRequest(model=model, audio=audio, optional_params=optional_params),
|
||||
NativeTranscriptionRequest(
|
||||
model=model,
|
||||
audio=audio,
|
||||
optional_params=optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
|
|
@ -132,5 +179,198 @@ async def atranscription(
|
|||
),
|
||||
),
|
||||
call=call_native,
|
||||
preflight=lambda: assess_route(_PREFLIGHT, model, custom_llm_provider or ""),
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
TranscriptionResult = TranscriptionResponse | Coroutine[object, object, TranscriptionResponse]
|
||||
|
||||
|
||||
def _input_source_kind(file: FileTypes) -> str:
|
||||
content: Final = file[1] if isinstance(file, tuple) else file
|
||||
if isinstance(content, (bytes, bytearray, memoryview)):
|
||||
return "bytes"
|
||||
if isinstance(content, IOBase):
|
||||
return "file"
|
||||
if isinstance(content, str):
|
||||
return "path"
|
||||
return "opaque"
|
||||
|
||||
|
||||
@dataclass
|
||||
class _TranscriptionOperation:
|
||||
model: str
|
||||
provider: str
|
||||
file: FileTypes
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
headers: dict[str, object] | None
|
||||
optional_params: dict[str, object]
|
||||
timeout: float | httpx.Timeout | None
|
||||
logging: Logging
|
||||
python: Callable[[FileTypes], TranscriptionResult]
|
||||
asynchronous: bool = False
|
||||
has_custom_client: bool = False
|
||||
fallback_file: FileTypes | None = None
|
||||
logged: bool = False
|
||||
|
||||
def prepare(self) -> PreparedNativeCall[NativeTranscriptionRequest]:
|
||||
key: Final = (
|
||||
self.api_key
|
||||
or litellm.api_key
|
||||
or TypeAdapter(str | None).validate_python(getattr(litellm, f"{self.provider}_key", None))
|
||||
or get_secret_str(f"{self.provider.upper()}_API_KEY")
|
||||
)
|
||||
base: Final = (
|
||||
self.api_base
|
||||
or litellm.api_base
|
||||
or get_secret_str(f"{self.provider.upper()}_BASE_URL")
|
||||
or get_secret_str(f"{self.provider.upper()}_API_BASE")
|
||||
)
|
||||
content: Final = self.file[1] if isinstance(self.file, tuple) else self.file
|
||||
position: Final = content.tell() if isinstance(content, IOBase) and content.seekable() else None
|
||||
try:
|
||||
processed: Final = process_audio_file(self.file)
|
||||
finally:
|
||||
if position is not None and isinstance(content, IOBase):
|
||||
content.seek(position)
|
||||
self.fallback_file = (processed.filename, processed.file_content, processed.content_type)
|
||||
audio: Final = TypeAdapter(dict[str, object]).validate_python(
|
||||
MappingProxyType(
|
||||
{
|
||||
"data": base64.b64encode(processed.file_content).decode("ascii"),
|
||||
"format": processed.filename.rsplit(".", 1)[-1].lower() if "." in processed.filename else "wav",
|
||||
"filename": processed.filename,
|
||||
}
|
||||
)
|
||||
)
|
||||
log_details: Final[NativePreCallDetails] = {
|
||||
"api_base": base or "",
|
||||
"headers": self.headers,
|
||||
"complete_input_dict": {"model": self.model, **self.optional_params},
|
||||
}
|
||||
self.logging.pre_call(input="audio transcription", api_key=key, additional_args=log_details)
|
||||
self.logged = True
|
||||
return PreparedNativeCall(
|
||||
NativeTranscriptionRequest(
|
||||
model=self.model,
|
||||
audio=audio,
|
||||
optional_params=self.optional_params,
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key=key,
|
||||
api_base=base,
|
||||
custom_llm_provider=self.provider,
|
||||
extra_headers=self.headers,
|
||||
timeout_seconds=timeout_to_seconds(self.timeout),
|
||||
bedrock=bedrock_options(self.optional_params),
|
||||
),
|
||||
context=request_context(
|
||||
logging_obj=self.logging,
|
||||
request_model=self.logging.model,
|
||||
litellm_params=self.logging.litellm_params,
|
||||
capabilities=NativeRequestCapabilities(
|
||||
execution_mode="async" if self.asynchronous else "sync",
|
||||
stream=self.optional_params.get("stream") is True,
|
||||
has_custom_client=self.has_custom_client,
|
||||
input_source_kind=_input_source_kind(self.file),
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
def fallback(self) -> TranscriptionResult:
|
||||
with self.logging.suppress_next_pre_call() if self.logged else nullcontext():
|
||||
return self.python(self.fallback_file if self.fallback_file is not None else self.file)
|
||||
|
||||
async def afallback(self) -> TranscriptionResponse:
|
||||
with self.logging.suppress_next_pre_call() if self.logged else nullcontext():
|
||||
result: Final = self.python(self.fallback_file if self.fallback_file is not None else self.file)
|
||||
return await result if isinstance(result, Coroutine) else result
|
||||
|
||||
def adapt(self, response: dict[str, object]) -> TranscriptionResponse:
|
||||
text: Final = TypeAdapter(str).validate_python(response["text"])
|
||||
parsed: Final = TranscriptionResponse(text=text)
|
||||
self.logging.post_call(
|
||||
input="audio transcription", api_key=self.api_key, original_response=json.dumps(response)
|
||||
)
|
||||
return parsed
|
||||
|
||||
|
||||
def dispatch_transcription(
|
||||
*,
|
||||
model: str,
|
||||
provider: str,
|
||||
file: FileTypes,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
logging: Logging,
|
||||
asynchronous: bool,
|
||||
has_custom_client: bool,
|
||||
fallback: Callable[[FileTypes], TranscriptionResult],
|
||||
) -> TranscriptionResult:
|
||||
operation: Final = _TranscriptionOperation(
|
||||
model,
|
||||
provider,
|
||||
file,
|
||||
api_key,
|
||||
api_base,
|
||||
headers,
|
||||
optional_params,
|
||||
timeout,
|
||||
logging,
|
||||
fallback,
|
||||
asynchronous,
|
||||
has_custom_client,
|
||||
)
|
||||
|
||||
def preflight() -> PythonFallback | None:
|
||||
return assess_route(
|
||||
_PREFLIGHT,
|
||||
model,
|
||||
provider,
|
||||
stream=optional_params.get("stream") is True,
|
||||
has_custom_client=has_custom_client,
|
||||
)
|
||||
|
||||
error_context: Final = BridgeErrorContext(provider=provider, model=model)
|
||||
if provider == "bedrock":
|
||||
if asynchronous:
|
||||
return _TRANSCRIPTION.arequire(
|
||||
prepare=operation.prepare,
|
||||
call=call_native,
|
||||
adapt=operation.adapt,
|
||||
error_context=error_context,
|
||||
preflight=preflight,
|
||||
)
|
||||
return _TRANSCRIPTION.require(
|
||||
prepare=operation.prepare,
|
||||
call=call_native,
|
||||
adapt=operation.adapt,
|
||||
error_context=error_context,
|
||||
preflight=preflight,
|
||||
)
|
||||
if asynchronous:
|
||||
return _TRANSCRIPTION.ainvoke(
|
||||
prepare=operation.prepare,
|
||||
call=call_native,
|
||||
adapt=operation.adapt,
|
||||
fallback=operation.afallback,
|
||||
error_context=error_context,
|
||||
eligible=rust_enabled(),
|
||||
preflight=preflight,
|
||||
)
|
||||
return _TRANSCRIPTION.invoke(
|
||||
prepare=operation.prepare,
|
||||
call=call_native,
|
||||
adapt=operation.adapt,
|
||||
fallback=operation.fallback,
|
||||
error_context=error_context,
|
||||
eligible=rust_enabled(),
|
||||
preflight=preflight,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,20 +1,14 @@
|
|||
"""Tests for the optional Rust-backed Anthropic Messages path."""
|
||||
|
||||
import importlib
|
||||
from typing import cast
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.request import NativeMessagesRequest, NativeRequestContext, NativeRequestOptions
|
||||
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.rust_bridge.request import NativeMessagesRequest, NativeRequestContext
|
||||
|
||||
rust_messages = importlib.import_module("litellm.rust_bridge.messages")
|
||||
rust_bridge_loader = importlib.import_module("litellm.rust_bridge.loader")
|
||||
|
|
@ -39,13 +33,12 @@ REQUEST_BODY: dict[str, object] = {
|
|||
class RecordingMessages:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
self.contexts: list[NativeRequestContext] = []
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
request: NativeMessagesRequest,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
|
|
@ -59,20 +52,18 @@ class RecordingMessages:
|
|||
"timeout_seconds": options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
self.contexts.append(context)
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
||||
|
||||
class RecordingAsyncMessages:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
self.contexts: list[NativeRequestContext] = []
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
request: NativeMessagesRequest,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
|
|
@ -86,7 +77,6 @@ class RecordingAsyncMessages:
|
|||
"timeout_seconds": options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
self.contexts.append(context)
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
||||
|
||||
|
|
@ -94,7 +84,9 @@ class ExplodingAsyncMessages:
|
|||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, *args: object, **kwargs: object) -> dict[str, object]:
|
||||
async def __call__(
|
||||
self, request: NativeMessagesRequest, *, options: object, context: NativeRequestContext
|
||||
) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise AssertionError("bridge must not be called")
|
||||
|
||||
|
|
@ -103,18 +95,28 @@ class RaisingAsyncMessages:
|
|||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, *args: object, **kwargs: object) -> dict[str, object]:
|
||||
async def __call__(
|
||||
self, request: NativeMessagesRequest, *, options: object, context: NativeRequestContext
|
||||
) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise RuntimeError("upstream request failed with status 400: bad request")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_rust_flag():
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None)
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None, decline=None)
|
||||
configuration.reset_rust_configuration()
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
rust_messages.set_rust_messages(
|
||||
decline=lambda model, custom_llm_provider, *, context: (
|
||||
"unsupported feature"
|
||||
if any(getattr(context.capabilities, key) for key in ("stream", "has_agentic_hook", "has_custom_client"))
|
||||
or context.capabilities.request_format == "native"
|
||||
else None
|
||||
)
|
||||
)
|
||||
yield
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None)
|
||||
rust_messages.set_rust_messages(messages=None, amessages=None, decline=None)
|
||||
configuration.reset_rust_configuration()
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
|
||||
|
|
@ -133,7 +135,7 @@ def test_load_rust_amessages_returns_injected_impl():
|
|||
assert rust_messages.load_rust_amessages() is bridge
|
||||
|
||||
|
||||
def test_messages_wrapper_reports_unavailable(monkeypatch):
|
||||
def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge.bindings"),
|
||||
"get_native_bridge",
|
||||
|
|
@ -150,7 +152,7 @@ def test_messages_wrapper_reports_unavailable(monkeypatch):
|
|||
extra_headers={},
|
||||
timeout=30.0,
|
||||
)
|
||||
assert result == NativeSkipped(NativeSkipReason.UNAVAILABLE)
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_messages_wrapper_forwards_args_and_converts_timeout():
|
||||
|
|
@ -168,7 +170,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout():
|
|||
timeout=httpx.Timeout(600.0, read=42.0),
|
||||
)
|
||||
|
||||
assert response == Handled(FAKE_MESSAGES_RESPONSE)
|
||||
assert response == FAKE_MESSAGES_RESPONSE
|
||||
assert bridge.calls[0] == {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"body": REQUEST_BODY,
|
||||
|
|
@ -196,331 +198,147 @@ async def test_amessages_wrapper_forwards_args():
|
|||
timeout=12.5,
|
||||
)
|
||||
|
||||
assert response == Handled(FAKE_MESSAGES_RESPONSE)
|
||||
assert response == FAKE_MESSAGES_RESPONSE
|
||||
assert bridge.calls[0]["model"] == "claude-sonnet-4-5"
|
||||
assert bridge.calls[0]["timeout_seconds"] == 12.5
|
||||
|
||||
|
||||
class PythonMessages:
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
def anthropic_messages_handler(self, **kwargs: object) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
return {**FAKE_MESSAGES_RESPONSE, "id": "python"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize("provider", ["anthropic", "azure_ai", "openai"])
|
||||
@pytest.mark.asyncio
|
||||
async def test_amessages_wrapper_preserves_capability_facts():
|
||||
bridge = RecordingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
await rust_messages.amessages(
|
||||
model="claude-sonnet-4-5",
|
||||
body=REQUEST_BODY,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider="anthropic",
|
||||
extra_headers=None,
|
||||
timeout=None,
|
||||
stream=True,
|
||||
has_custom_client=True,
|
||||
has_agentic_hook=True,
|
||||
)
|
||||
|
||||
capabilities = bridge.contexts[0].capabilities
|
||||
assert capabilities.execution_mode == "async"
|
||||
assert capabilities.stream is True
|
||||
assert capabilities.has_custom_client is True
|
||||
assert capabilities.has_agentic_hook is True
|
||||
|
||||
|
||||
def _gate(**overrides):
|
||||
kwargs = {
|
||||
"custom_llm_provider": "azure_ai",
|
||||
"litellm_params": GenericLiteLLMParams(api_key="sk-azure"),
|
||||
"has_agentic_hook": False,
|
||||
"stream": False,
|
||||
"has_custom_client": False,
|
||||
"model": "claude-sonnet-4-5",
|
||||
"api_key": "sk-azure",
|
||||
"api_base": "https://resource.services.ai.azure.com/anthropic",
|
||||
"headers": {"x-api-key": "sk-azure", "anthropic-version": "2023-06-01"},
|
||||
"request_body": dict(REQUEST_BODY),
|
||||
"timeout": 30.0,
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return BaseLLMHTTPHandler._attempt_rust_anthropic_messages(**kwargs)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_invokes_rust_and_marks_response_header():
|
||||
bridge = RecordingAsyncMessages()
|
||||
async def test_public_messages_routes_provider_acceptance(monkeypatch, asynchronous, provider):
|
||||
bridge = RecordingMessages()
|
||||
async_bridge = RecordingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate()
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
rust_messages.set_rust_messages(messages=bridge, amessages=async_bridge)
|
||||
if asynchronous:
|
||||
response = await litellm.anthropic.messages.acreate(
|
||||
model=f"{provider}/test-model",
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key="key",
|
||||
api_base="https://example.test",
|
||||
)
|
||||
else:
|
||||
response = litellm.anthropic.messages.create(
|
||||
model=f"{provider}/test-model",
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key="key",
|
||||
api_base="https://example.test",
|
||||
)
|
||||
assert response["id"] == "msg_123"
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
call = bridge.calls[0]
|
||||
assert call["model"] == "claude-sonnet-4-5"
|
||||
assert call["body"] == REQUEST_BODY
|
||||
assert call["api_key"] == "sk-azure"
|
||||
assert call["api_base"] == "https://resource.services.ai.azure.com/anthropic"
|
||||
assert call["extra_headers"] == {"x-api-key": "sk-azure", "anthropic-version": "2023-06-01"}
|
||||
assert call["timeout_seconds"] == 30.0
|
||||
assert response["_hidden_params"]["additional_headers"]["x-litellm-rust"] == "true"
|
||||
calls = async_bridge.calls if asynchronous else bridge.calls
|
||||
assert len(calls) == 1
|
||||
assert calls[0]["custom_llm_provider"] == provider
|
||||
assert calls[0]["body"]["max_tokens"] == 64
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_reports_failure_to_harness():
|
||||
bridge = RaisingAsyncMessages()
|
||||
def test_public_messages_strips_provider_specific_fields_before_native_dispatch():
|
||||
native = RecordingMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
rust_messages.set_rust_messages(messages=native)
|
||||
messages = [
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{
|
||||
"type": "tool_use",
|
||||
"id": "tool_1",
|
||||
"name": "lookup",
|
||||
"input": {},
|
||||
"provider_specific_fields": {"thought_signature": "signature"},
|
||||
}
|
||||
],
|
||||
}
|
||||
]
|
||||
|
||||
response = await _gate()
|
||||
assert isinstance(response, NativeFailed)
|
||||
assert bridge.calls == 1
|
||||
litellm.anthropic.messages.create(
|
||||
model="anthropic/test-model",
|
||||
max_tokens=64,
|
||||
messages=messages,
|
||||
api_key="key",
|
||||
)
|
||||
|
||||
content = native.calls[0]["body"]["messages"][0]["content"][0]
|
||||
assert "provider_specific_fields" not in content
|
||||
assert "provider_specific_fields" in messages[0]["content"][0]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_when_flag_absent():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert bridge.calls == 0
|
||||
@pytest.mark.parametrize("condition", ["disabled", "declined", "missing_binding", "missing_preflight", "stream"])
|
||||
def test_public_messages_fallback_once(monkeypatch, condition):
|
||||
module = importlib.import_module("litellm.llms.anthropic.experimental_pass_through.messages.handler")
|
||||
python = PythonMessages()
|
||||
monkeypatch.setattr(module, "base_llm_http_handler", python)
|
||||
bridge = RecordingMessages()
|
||||
litellm.rust(condition != "disabled")
|
||||
rust_messages.set_rust_messages(messages=bridge)
|
||||
if condition == "declined":
|
||||
rust_messages.set_rust_messages(decline=lambda model, custom_llm_provider, **features: "unsupported provider")
|
||||
elif condition == "missing_binding":
|
||||
rust_messages._MESSAGES.sync.override(None)
|
||||
elif condition == "missing_preflight":
|
||||
rust_messages._PREFLIGHT.override(None)
|
||||
litellm.anthropic.messages.create(
|
||||
model="anthropic/test-model",
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key="key",
|
||||
stream=condition == "stream",
|
||||
)
|
||||
assert python.calls == 1
|
||||
assert bridge.calls == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_uses_process_enable_without_request_override():
|
||||
bridge = RecordingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
@pytest.mark.parametrize("response", [{}, {"content": "invalid"}])
|
||||
def test_public_messages_invalid_response_does_not_fallback(monkeypatch, response):
|
||||
module = importlib.import_module("litellm.llms.anthropic.experimental_pass_through.messages.handler")
|
||||
python = PythonMessages()
|
||||
monkeypatch.setattr(module, "base_llm_http_handler", python)
|
||||
litellm.rust(True)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
|
||||
rust_messages.set_rust_messages(messages=lambda request, *, options, context: response)
|
||||
with pytest.raises(ValidationError):
|
||||
litellm.anthropic.messages.create(
|
||||
model="anthropic/test-model",
|
||||
max_tokens=64,
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key="key",
|
||||
)
|
||||
assert python.calls == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_invokes_rust_for_native_anthropic_provider():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
api_key="sk-ant",
|
||||
api_base="https://api.anthropic.com",
|
||||
headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"},
|
||||
def test_public_messages_preserves_headers_and_optional_parameters(monkeypatch):
|
||||
handler = importlib.import_module("litellm.llms.anthropic.experimental_pass_through.messages.handler")
|
||||
monkeypatch.setattr(handler, "is_reasoning_auto_summary_enabled", lambda: True)
|
||||
configuration.rust(True)
|
||||
native = RecordingMessages()
|
||||
rust_messages.set_rust_messages(messages=native)
|
||||
response = litellm.anthropic.messages.create(
|
||||
model="anthropic/claude-sonnet-4-5",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
max_tokens=2048,
|
||||
thinking={"type": "enabled", "budget_tokens": 1024},
|
||||
temperature=0.8,
|
||||
additional_drop_params=["temperature"],
|
||||
api_key="key",
|
||||
headers={"x-source": "forwarded"},
|
||||
extra_headers={"x-source": "extra"},
|
||||
provider_specific_header={
|
||||
"custom_llm_provider": "anthropic",
|
||||
"extra_headers": {"x-source": "scoped"},
|
||||
},
|
||||
)
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
|
||||
assert bridge.calls[0]["api_key"] == "sk-ant"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_invokes_rust_when_env_var_set(monkeypatch):
|
||||
bridge = RecordingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
|
||||
response = await _gate(
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
)
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
|
||||
bridge = ExplodingAsyncMessages()
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
|
||||
response = await _gate(
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
)
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_unsupported_provider():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(custom_llm_provider="openai")
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_for_agentic_hook():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(has_agentic_hook=True)
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
streaming_body = {**REQUEST_BODY, "stream": True}
|
||||
response = await _gate(
|
||||
has_agentic_hook=False,
|
||||
request_body=streaming_body,
|
||||
)
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert "stream" not in bridge.calls[0]["body"]
|
||||
assert bridge.calls[0]["body"] == REQUEST_BODY
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
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"}
|
||||
|
||||
chunks = [chunk async for chunk in stream]
|
||||
joined = b"".join(chunks)
|
||||
|
||||
assert b"event: message_start" in joined
|
||||
assert b"event: content_block_delta" in joined
|
||||
assert b"hello world" in joined
|
||||
assert b"event: message_stop" in joined
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_falls_back_when_bridge_unavailable(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge.bindings"),
|
||||
"get_native_bridge",
|
||||
lambda: None,
|
||||
)
|
||||
litellm.rust(True)
|
||||
|
||||
response = await _gate()
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("selection", ("native", "disabled", "failed", "declined", "upstream"))
|
||||
async def test_messages_handler_runs_selected_backend_once(selection: str, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.exceptions import RateLimitError
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import AnthropicMessagesConfig
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.rust_bridge import bindings
|
||||
|
||||
class Declined(Exception):
|
||||
pass
|
||||
|
||||
class Upstream(Exception):
|
||||
pass
|
||||
|
||||
error = (
|
||||
Upstream(429, "rate limited")
|
||||
if selection == "upstream"
|
||||
else Declined("unsupported")
|
||||
if selection == "declined"
|
||||
else RuntimeError("native failed")
|
||||
if selection == "failed"
|
||||
else None
|
||||
)
|
||||
|
||||
class Native:
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, *args: object, **kwargs: object) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
if error is not None:
|
||||
raise error
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
||||
bridge = Native()
|
||||
monkeypatch.setattr(
|
||||
bindings,
|
||||
"get_native_bridge",
|
||||
lambda: SimpleNamespace(
|
||||
RustBridgeDeclined=Declined,
|
||||
RustUpstreamError=Upstream,
|
||||
),
|
||||
)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
litellm.rust(selection != "disabled")
|
||||
requests: list[httpx.Request] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
requests.append(request)
|
||||
return httpx.Response(200, json=FAKE_MESSAGES_RESPONSE)
|
||||
|
||||
logging_obj = Logging(
|
||||
model=FAKE_MESSAGES_RESPONSE["model"],
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="anthropic_messages",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="harness-test",
|
||||
function_id="harness-test",
|
||||
)
|
||||
client = AsyncHTTPHandler()
|
||||
await client.client.aclose()
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as transport:
|
||||
client.client = transport
|
||||
|
||||
async def run():
|
||||
return await BaseLLMHTTPHandler().async_anthropic_messages_handler(
|
||||
model=FAKE_MESSAGES_RESPONSE["model"],
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
anthropic_messages_provider_config=AnthropicMessagesConfig(),
|
||||
anthropic_messages_optional_request_params={"max_tokens": 10},
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
logging_obj=logging_obj,
|
||||
api_key="sk-test",
|
||||
api_base="https://example.test",
|
||||
client=client,
|
||||
)
|
||||
|
||||
if selection in ("failed", "upstream"):
|
||||
with pytest.raises(RateLimitError if selection == "upstream" else RuntimeError) as caught:
|
||||
await run()
|
||||
if selection == "upstream":
|
||||
assert caught.value.__cause__ is error
|
||||
assert caught.value.llm_provider == "anthropic"
|
||||
assert caught.value.model == FAKE_MESSAGES_RESPONSE["model"]
|
||||
else:
|
||||
assert caught.value is error
|
||||
else:
|
||||
response = await run()
|
||||
assert response["id"] == FAKE_MESSAGES_RESPONSE["id"]
|
||||
assert len(requests) == (1 if selection in ("disabled", "declined") else 0)
|
||||
assert bridge.calls == (0 if selection == "disabled" else 1)
|
||||
assert response["id"] == "msg_123"
|
||||
assert native.calls[0]["extra_headers"]["x-source"] == "scoped"
|
||||
assert native.calls[0]["body"]["thinking"]["display"] == "summarized"
|
||||
assert "temperature" not in native.calls[0]["body"]
|
||||
|
|
|
|||
|
|
@ -2223,389 +2223,3 @@ def test_non_bash_tool_result_skipped():
|
|||
), f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}"
|
||||
|
||||
|
||||
class TestRustChatCompletionsHook:
|
||||
"""The `rust: true` opt-in on `/chat/completions` for the Anthropic provider.
|
||||
|
||||
The native callables are dependency-injected, so these run without the
|
||||
compiled extension.
|
||||
"""
|
||||
|
||||
RUST_RESPONSE = {
|
||||
"created": 1_700_000_000,
|
||||
"model": "claude-sonnet-4-5-20260101",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "hello from rust"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 11,
|
||||
"completion_tokens": 4,
|
||||
"total_tokens": 15,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 0,
|
||||
"cache_creation_tokens": 0,
|
||||
"text_tokens": 11,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_bridge(self, monkeypatch):
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
bridge.set_rust_chat_completions(
|
||||
chat_completions=None, achat_completions=None, decline=None
|
||||
)
|
||||
yield
|
||||
bridge.set_rust_chat_completions(
|
||||
chat_completions=None, achat_completions=None, decline=None
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _completion_kwargs(**overrides):
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
kwargs = {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"api_base": "https://api.anthropic.com/v1/messages",
|
||||
"custom_llm_provider": "anthropic",
|
||||
"custom_prompt_dict": {},
|
||||
"model_response": ModelResponse(),
|
||||
"print_verbose": lambda *_args, **_kwargs: None,
|
||||
"encoding": None,
|
||||
"api_key": "sk-ant-test",
|
||||
"logging_obj": MagicMock(),
|
||||
"optional_params": {"max_tokens": 16},
|
||||
"timeout": 30.0,
|
||||
"litellm_params": {},
|
||||
"acompletion": False,
|
||||
"headers": {},
|
||||
"client": None,
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return kwargs
|
||||
|
||||
@staticmethod
|
||||
def _recording_logging_obj():
|
||||
"""A logging object that keeps each hook's payload in a real list, so a
|
||||
test can assert which path logged and what it carried."""
|
||||
calls = {"pre_call": [], "post_call": []}
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.pre_call.side_effect = lambda **kwargs: calls["pre_call"].append(kwargs)
|
||||
logging_obj.post_call.side_effect = lambda **kwargs: calls["post_call"].append(kwargs)
|
||||
return logging_obj, calls
|
||||
|
||||
def _inject(self, *, decline_reason=None, sync_result=None, sync_error=None):
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
|
||||
seen = {"gate": [], "call": []}
|
||||
|
||||
def gate(**kwargs):
|
||||
seen["gate"].append(kwargs)
|
||||
return decline_reason
|
||||
|
||||
def native(**kwargs):
|
||||
seen["call"].append(kwargs)
|
||||
if sync_error is not None:
|
||||
raise sync_error
|
||||
return dict(sync_result if sync_result is not None else self.RUST_RESPONSE)
|
||||
|
||||
bridge.set_rust_chat_completions(decline=gate, chat_completions=native)
|
||||
return seen
|
||||
|
||||
def test_rust_true_serves_the_call_and_stamps_the_header(self):
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
|
||||
seen = self._inject()
|
||||
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 len(seen["call"]) == 1
|
||||
|
||||
def test_the_core_receives_the_untranslated_openai_messages(self):
|
||||
"""Rust owns the translation, so the handler must not pre-translate."""
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
|
||||
seen = self._inject()
|
||||
AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(
|
||||
messages=[
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"},
|
||||
]
|
||||
)
|
||||
)
|
||||
assert seen["call"][0]["messages"] == [
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"},
|
||||
]
|
||||
|
||||
def test_the_anthropic_max_tokens_default_is_merged_in_before_the_gate(self):
|
||||
"""`transform_request` applies `AnthropicConfig.get_config`; the Rust
|
||||
path skips it, so the handler has to merge it or Anthropic 400s on a
|
||||
request that omits `max_tokens`."""
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
|
||||
seen = self._inject()
|
||||
AnthropicChatCompletion().completion(**self._completion_kwargs(optional_params={}))
|
||||
assert "max_tokens" in seen["gate"][0]["optional_params"]
|
||||
assert seen["call"][0]["optional_params"]["max_tokens"] > 0
|
||||
|
||||
def test_a_caller_supplied_max_tokens_outranks_the_default(self):
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
|
||||
seen = self._inject()
|
||||
AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(optional_params={"max_tokens": 7})
|
||||
)
|
||||
assert seen["call"][0]["optional_params"]["max_tokens"] == 7
|
||||
|
||||
def test_without_the_opt_in_the_core_is_never_consulted(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
seen = self._inject()
|
||||
with patch.object(
|
||||
AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}
|
||||
) as transform, patch.object(
|
||||
AnthropicChatCompletion, "acompletion_function"
|
||||
):
|
||||
try:
|
||||
AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(litellm_params={})
|
||||
)
|
||||
except Exception:
|
||||
# The Python path goes on to make an HTTP call; reaching it is
|
||||
# the assertion, so the network failure below is expected.
|
||||
pass
|
||||
assert seen["gate"] == []
|
||||
assert seen["call"] == []
|
||||
assert transform.called
|
||||
|
||||
def test_a_declined_request_never_reaches_the_native_call(self):
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
seen = self._inject(decline_reason="unrecognized request parameter")
|
||||
with patch.object(
|
||||
AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}
|
||||
):
|
||||
try:
|
||||
AnthropicChatCompletion().completion(**self._completion_kwargs())
|
||||
except Exception:
|
||||
pass
|
||||
assert len(seen["gate"]) == 1
|
||||
assert seen["call"] == []
|
||||
|
||||
def test_streaming_stays_on_the_python_path(self):
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
seen = self._inject()
|
||||
with patch.object(
|
||||
AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}
|
||||
):
|
||||
try:
|
||||
AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(optional_params={"max_tokens": 16, "stream": True})
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
assert seen["gate"] == []
|
||||
|
||||
def test_pre_call_logging_fires_exactly_once_on_the_rust_path(self):
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
|
||||
seen = self._inject()
|
||||
logging_obj = MagicMock()
|
||||
AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(logging_obj=logging_obj)
|
||||
)
|
||||
assert logging_obj.pre_call.call_count == 1
|
||||
assert len(seen["call"]) == 1
|
||||
|
||||
def test_post_call_logging_fires_on_the_rust_path(self):
|
||||
"""The Rust core owns the provider call, so the Python transform that
|
||||
normally raises `post_call` never runs. Without the bridge hook every
|
||||
post_call callback goes silent and `original_response` stays unset."""
|
||||
import json
|
||||
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
|
||||
self._inject()
|
||||
logging_obj = MagicMock()
|
||||
AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(logging_obj=logging_obj)
|
||||
)
|
||||
|
||||
assert logging_obj.post_call.call_count == 1
|
||||
logged = logging_obj.post_call.call_args.kwargs["original_response"]
|
||||
assert json.loads(logged)["choices"][0]["message"]["content"] == "hello from rust"
|
||||
|
||||
def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(self, monkeypatch):
|
||||
"""A decline never reached the provider, so the Python path serves the
|
||||
request and owns the only post_call. Firing the hook there too would
|
||||
double every post_call callback for one request."""
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
|
||||
class _Declined(Exception):
|
||||
pass
|
||||
|
||||
class _FakeNative:
|
||||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
def declining_native(**_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, chat_completions=declining_native
|
||||
)
|
||||
|
||||
logging_obj, calls = self._recording_logging_obj()
|
||||
with patch.object(
|
||||
AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}
|
||||
):
|
||||
try:
|
||||
AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(logging_obj=logging_obj)
|
||||
)
|
||||
except Exception:
|
||||
# The Python path goes on to make an HTTP call; the log count is
|
||||
# the assertion, so a failure past this point is expected.
|
||||
pass
|
||||
|
||||
assert calls["post_call"] == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_path_falls_back_when_the_core_declines(self, monkeypatch):
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
|
||||
class _Declined(Exception):
|
||||
pass
|
||||
|
||||
class _FakeNative:
|
||||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
|
||||
|
||||
async def declining_native(**_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=declining_native
|
||||
)
|
||||
|
||||
sentinel = object()
|
||||
|
||||
async def python_path(**_kwargs):
|
||||
return sentinel
|
||||
|
||||
with patch.object(
|
||||
AnthropicChatCompletion, "acompletion_function", side_effect=python_path
|
||||
) as python_call:
|
||||
result = await AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(acompletion=True)
|
||||
)
|
||||
|
||||
assert result is sentinel
|
||||
assert python_call.called, "a failing rust call must re-enter the python path"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_path_serves_the_rust_response_without_the_fallback(self):
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
|
||||
async def native(**_kwargs):
|
||||
return dict(self.RUST_RESPONSE)
|
||||
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=native
|
||||
)
|
||||
|
||||
with patch.object(AnthropicChatCompletion, "acompletion_function") as python_call:
|
||||
result = await AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(acompletion=True)
|
||||
)
|
||||
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert not python_call.called
|
||||
|
||||
|
||||
def test_pre_call_logging_fires_once_when_the_sync_rust_call_declines(self, monkeypatch):
|
||||
"""One request, one pre_call, on the synchronous path too. Without the
|
||||
suppression the Python path logs a second time for the same attempt."""
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
|
||||
class _Declined(Exception):
|
||||
pass
|
||||
|
||||
class _FakeNative:
|
||||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
|
||||
|
||||
def declining_native(**_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, chat_completions=declining_native
|
||||
)
|
||||
|
||||
logging_obj, calls = self._recording_logging_obj()
|
||||
with patch.object(
|
||||
AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}
|
||||
):
|
||||
try:
|
||||
AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(logging_obj=logging_obj)
|
||||
)
|
||||
except Exception:
|
||||
# The Python path goes on to make an HTTP call; the log count is
|
||||
# the assertion, so a failure past this point is expected.
|
||||
pass
|
||||
|
||||
assert len(calls["pre_call"]) == 1
|
||||
assert calls["pre_call"][0]["additional_args"]["complete_input_dict"]["model"] == (
|
||||
"claude-sonnet-4-5"
|
||||
)
|
||||
|
||||
def test_pre_call_logging_still_fires_when_rust_is_not_involved(self, monkeypatch):
|
||||
"""The suppression must not swallow the log on the ordinary path."""
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
|
||||
self._inject()
|
||||
logging_obj, calls = self._recording_logging_obj()
|
||||
with patch.object(
|
||||
AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}
|
||||
):
|
||||
try:
|
||||
AnthropicChatCompletion().completion(
|
||||
**self._completion_kwargs(litellm_params={}, logging_obj=logging_obj)
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
assert len(calls["pre_call"]) == 1
|
||||
assert calls["pre_call"][0]["additional_args"]["complete_input_dict"] == {
|
||||
"model": "m",
|
||||
"messages": [],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,552 +1,75 @@
|
|||
"""Tests for `BedrockConverseLLM.completion`'s Rust chat completions hook.
|
||||
|
||||
The native callables are dependency-injected, so these run without the compiled
|
||||
extension, and AWS credential resolution is stubbed so nothing reaches STS.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from botocore.credentials import Credentials
|
||||
from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
import litellm
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
RUST_RESPONSE = {
|
||||
"created": 1_700_000_000,
|
||||
"model": "anthropic.claude-sonnet-4-5-v1:0",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "hello from rust"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 11,
|
||||
"completion_tokens": 4,
|
||||
"total_tokens": 15,
|
||||
"prompt_tokens_details": {
|
||||
"cached_tokens": 0,
|
||||
"cache_creation_tokens": 0,
|
||||
"text_tokens": 11,
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
RESOLVED_CREDENTIALS = Credentials(
|
||||
access_key="AKIARESOLVED",
|
||||
secret_key="resolved-secret",
|
||||
token="resolved-token",
|
||||
)
|
||||
from litellm.rust_bridge.request import NativeChatCompletionsRequest, NativeRequestContext, NativeRequestOptions
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_bridge(monkeypatch):
|
||||
def native_bridge(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
bridge.set_rust_chat_completions(
|
||||
chat_completions=None, achat_completions=None, decline=None
|
||||
)
|
||||
bridge.set_rust_chat_completions(decline=lambda **features: None)
|
||||
yield
|
||||
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=None)
|
||||
|
||||
|
||||
def _inject(*, decline_reason=None, error: Exception | None = None):
|
||||
seen: dict[str, list[dict]] = {"gate": [], "call": []}
|
||||
def test_native_bedrock_receives_explicit_auth_and_endpoint():
|
||||
requests = []
|
||||
|
||||
def gate(**kwargs):
|
||||
seen["gate"].append(kwargs)
|
||||
return decline_reason
|
||||
|
||||
def native(request, *, options, context):
|
||||
seen["call"].append(
|
||||
{
|
||||
"model": request.model,
|
||||
"messages": request.messages,
|
||||
"optional_params": request.optional_params,
|
||||
"api_key": options.api_key,
|
||||
"api_base": options.api_base,
|
||||
"context": context,
|
||||
}
|
||||
)
|
||||
if error is not None:
|
||||
raise error
|
||||
return dict(RUST_RESPONSE)
|
||||
|
||||
bridge.set_rust_chat_completions(decline=gate, chat_completions=native)
|
||||
return seen
|
||||
|
||||
|
||||
def _completion_kwargs(**overrides):
|
||||
kwargs = {
|
||||
"model": "bedrock/us-east-1/anthropic.claude-sonnet-4-5-v1:0",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"api_base": None,
|
||||
"custom_prompt_dict": {},
|
||||
"model_response": ModelResponse(),
|
||||
"encoding": None,
|
||||
"logging_obj": MagicMock(),
|
||||
"optional_params": {"maxTokens": 16},
|
||||
"acompletion": False,
|
||||
"timeout": 30.0,
|
||||
"litellm_params": {},
|
||||
"extra_headers": None,
|
||||
"client": None,
|
||||
"api_key": None,
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return kwargs
|
||||
|
||||
|
||||
def _run(*, credentials: Credentials | None = RESOLVED_CREDENTIALS, **overrides):
|
||||
with patch.object(BedrockConverseLLM, "get_credentials", return_value=credentials):
|
||||
return BedrockConverseLLM().completion(**_completion_kwargs(**overrides))
|
||||
|
||||
|
||||
def _recording_logging_obj():
|
||||
"""A logging object that keeps each hook's payload in a real list, so a test
|
||||
can assert which path logged and what it carried."""
|
||||
calls = {"pre_call": [], "post_call": []}
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.pre_call.side_effect = lambda **kwargs: calls["pre_call"].append(kwargs)
|
||||
logging_obj.post_call.side_effect = lambda **kwargs: calls["post_call"].append(kwargs)
|
||||
return logging_obj, calls
|
||||
|
||||
|
||||
def test_rust_true_serves_the_call_and_stamps_the_header():
|
||||
seen = _inject()
|
||||
response = _run()
|
||||
|
||||
assert response.choices[0].message.content == "hello from rust"
|
||||
assert response._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert len(seen["call"]) == 1
|
||||
|
||||
|
||||
def test_the_core_receives_the_credentials_this_handler_already_resolved():
|
||||
"""Both paths must sign as the same principal, so the resolved credentials
|
||||
are handed down rather than re-derived from ambient AWS state."""
|
||||
seen = _inject()
|
||||
_run()
|
||||
|
||||
params = seen["call"][0]["optional_params"]
|
||||
assert params["aws_access_key_id"] == "AKIARESOLVED"
|
||||
assert params["aws_secret_access_key"] == "resolved-secret"
|
||||
assert params["aws_session_token"] == "resolved-token"
|
||||
assert params["aws_region_name"] == "us-east-1"
|
||||
|
||||
|
||||
def test_the_core_receives_the_converse_url_this_handler_already_built():
|
||||
seen = _inject()
|
||||
_run()
|
||||
|
||||
assert seen["call"][0]["api_base"].endswith(
|
||||
"/model/anthropic.claude-sonnet-4-5-v1%3A0/converse"
|
||||
)
|
||||
assert "bedrock-runtime.us-east-1.amazonaws.com" in seen["call"][0]["api_base"]
|
||||
|
||||
|
||||
def test_the_core_receives_the_untranslated_openai_messages():
|
||||
seen = _inject()
|
||||
_run(
|
||||
messages=[
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"},
|
||||
]
|
||||
)
|
||||
assert seen["call"][0]["messages"] == [
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"},
|
||||
]
|
||||
|
||||
|
||||
def test_without_the_opt_in_the_core_is_never_consulted(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
seen = _inject()
|
||||
try:
|
||||
_run(litellm_params={})
|
||||
except Exception:
|
||||
# The Python path goes on to make an HTTP call; not reaching the gate
|
||||
# is the assertion, so a failure past this point is expected.
|
||||
pass
|
||||
assert seen["gate"] == []
|
||||
assert seen["call"] == []
|
||||
|
||||
|
||||
def test_streaming_stays_on_the_python_path():
|
||||
seen = _inject()
|
||||
try:
|
||||
_run(optional_params={"maxTokens": 16, "stream": True})
|
||||
except Exception:
|
||||
pass
|
||||
assert seen["gate"] == []
|
||||
|
||||
|
||||
def test_a_declined_request_never_reaches_the_native_call():
|
||||
seen = _inject(decline_reason="unrecognized request parameter")
|
||||
try:
|
||||
_run()
|
||||
except Exception:
|
||||
pass
|
||||
assert len(seen["gate"]) == 1
|
||||
assert seen["call"] == []
|
||||
|
||||
|
||||
def test_pre_call_logging_fires_exactly_once_on_the_rust_path():
|
||||
_inject()
|
||||
logging_obj = MagicMock()
|
||||
_run(logging_obj=logging_obj)
|
||||
assert logging_obj.pre_call.call_count == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
|
||||
class _Declined(Exception):
|
||||
pass
|
||||
|
||||
class _FakeNative:
|
||||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
|
||||
|
||||
async def declining_native(_request, **_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=declining_native
|
||||
)
|
||||
|
||||
sentinel = object()
|
||||
|
||||
async def python_path(**_kwargs):
|
||||
return sentinel
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
|
||||
),
|
||||
patch.object(
|
||||
BedrockConverseLLM, "async_completion", side_effect=python_path
|
||||
) as python_call,
|
||||
def native(
|
||||
request: NativeChatCompletionsRequest,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter=None,
|
||||
):
|
||||
result = await BedrockConverseLLM().completion(
|
||||
**_completion_kwargs(acompletion=True)
|
||||
)
|
||||
requests.append((request, options))
|
||||
return {
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "native"}, "finish_reason": "stop"}]
|
||||
}
|
||||
|
||||
assert result is sentinel
|
||||
assert python_call.called, "a failing rust call must re-enter the python path"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_path_serves_the_rust_response_without_the_fallback():
|
||||
async def native(_request, **_kwargs):
|
||||
return dict(RUST_RESPONSE)
|
||||
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=native
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
|
||||
),
|
||||
patch.object(BedrockConverseLLM, "async_completion") as python_call,
|
||||
):
|
||||
result = await BedrockConverseLLM().completion(
|
||||
**_completion_kwargs(acompletion=True)
|
||||
)
|
||||
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert not python_call.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines():
|
||||
"""One request, one pre_call. Without the suppression the Python fallback
|
||||
logs a second one and non-idempotent callbacks run twice."""
|
||||
|
||||
class _Declined(Exception):
|
||||
pass
|
||||
|
||||
class _FakeNative:
|
||||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
async def declining_native(_request, **_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
||||
logging_obj = MagicMock()
|
||||
served = []
|
||||
|
||||
async def python_path(**kwargs):
|
||||
served.append(kwargs)
|
||||
return ModelResponse()
|
||||
|
||||
with (
|
||||
patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()),
|
||||
patch.object(
|
||||
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
|
||||
),
|
||||
patch.object(
|
||||
BedrockConverseLLM, "async_completion", side_effect=python_path
|
||||
),
|
||||
):
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=declining_native
|
||||
)
|
||||
await BedrockConverseLLM().completion(
|
||||
**_completion_kwargs(acompletion=True, logging_obj=logging_obj)
|
||||
)
|
||||
|
||||
assert logging_obj.pre_call.call_count == 1
|
||||
assert served and served[0]["skip_pre_call_logging"] is True
|
||||
|
||||
|
||||
CONVERSE_RESPONSE = {
|
||||
"output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}},
|
||||
"stopReason": "end_turn",
|
||||
"usage": {"inputTokens": 5, "outputTokens": 2, "totalTokens": 7},
|
||||
}
|
||||
|
||||
|
||||
async def _drive_async_completion(*, skip_pre_call_logging: bool, logging_obj):
|
||||
"""Run the real `async_completion` with a stubbed transport."""
|
||||
import httpx as _httpx
|
||||
|
||||
client = MagicMock()
|
||||
|
||||
async def post(**_kwargs):
|
||||
return _httpx.Response(
|
||||
200,
|
||||
json=CONVERSE_RESPONSE,
|
||||
request=_httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com"),
|
||||
)
|
||||
|
||||
client.post = post
|
||||
client.__class__ = AsyncHTTPHandler
|
||||
|
||||
return await BedrockConverseLLM().async_completion(
|
||||
model="anthropic.claude-sonnet-4-5-v1:0",
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
response = litellm.completion(
|
||||
model="bedrock/anthropic.claude-sonnet-4-5-v1:0",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_base="https://bedrock-runtime.us-west-2.amazonaws.com/model/m/converse",
|
||||
model_response=ModelResponse(),
|
||||
timeout=30.0,
|
||||
encoding=None,
|
||||
logging_obj=logging_obj,
|
||||
stream=None,
|
||||
optional_params={"maxTokens": 16},
|
||||
litellm_params={"aws_region_name": "us-west-2"},
|
||||
credentials=RESOLVED_CREDENTIALS,
|
||||
headers={},
|
||||
client=client,
|
||||
skip_pre_call_logging=skip_pre_call_logging,
|
||||
aws_access_key_id="explicit-id",
|
||||
aws_secret_access_key="explicit-secret",
|
||||
aws_session_token="explicit-token",
|
||||
aws_region_name="us-east-1",
|
||||
api_base="https://example.test",
|
||||
max_tokens=7,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_completion_honors_the_pre_call_suppression():
|
||||
logging_obj = MagicMock()
|
||||
await _drive_async_completion(skip_pre_call_logging=True, logging_obj=logging_obj)
|
||||
assert logging_obj.pre_call.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_completion_logs_pre_call_by_default():
|
||||
"""The suppression must be opt-in, so every existing caller keeps its log."""
|
||||
logging_obj = MagicMock()
|
||||
await _drive_async_completion(skip_pre_call_logging=False, logging_obj=logging_obj)
|
||||
assert logging_obj.pre_call.call_count == 1
|
||||
|
||||
|
||||
def _sync_client_returning_converse_response():
|
||||
client = MagicMock()
|
||||
client.post.side_effect = lambda **_kwargs: httpx.Response(
|
||||
200,
|
||||
json=CONVERSE_RESPONSE,
|
||||
request=httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com"),
|
||||
assert response.choices[0].message.content == "native"
|
||||
assert len(requests) == 1
|
||||
assert requests[0][1].bedrock == bridge.NativeBedrockOptions(
|
||||
aws_access_key_id="explicit-id",
|
||||
aws_secret_access_key="explicit-secret",
|
||||
aws_session_token="explicit-token",
|
||||
aws_region_name="us-east-1",
|
||||
)
|
||||
client.__class__ = HTTPHandler
|
||||
return client
|
||||
assert requests[0][1].api_base == "https://example.test"
|
||||
|
||||
|
||||
def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines():
|
||||
"""One request, one pre_call, on the synchronous path too.
|
||||
@pytest.mark.parametrize("through_environment", [False, True])
|
||||
def test_native_bedrock_preserves_bearer_auth(monkeypatch, through_environment):
|
||||
requests = []
|
||||
|
||||
The gate accepts and logs, then the native call declines before the
|
||||
provider is reached, so execution continues into the Python path below.
|
||||
That is the same attempt continuing; without the suppression it logs a
|
||||
second pre_call and non-idempotent callbacks run twice for one request.
|
||||
"""
|
||||
def native(request, *, options, context, callback_adapter=None):
|
||||
requests.append((request, options))
|
||||
return {
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "native"}, "finish_reason": "stop"}]
|
||||
}
|
||||
|
||||
class _Declined(Exception):
|
||||
pass
|
||||
|
||||
class _FakeNative:
|
||||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
def declining_native(_request, **_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
||||
logging_obj = MagicMock()
|
||||
|
||||
with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()):
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, chat_completions=declining_native
|
||||
)
|
||||
response = _run(
|
||||
logging_obj=logging_obj,
|
||||
client=_sync_client_returning_converse_response(),
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "hi"
|
||||
assert logging_obj.pre_call.call_count == 1
|
||||
|
||||
|
||||
def test_the_sync_python_path_still_logs_pre_call_without_the_opt_in(monkeypatch):
|
||||
"""The suppression must not swallow the log on a request the gate declined,
|
||||
so a deployment with no `rust` flag keeps exactly the log it always had."""
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
logging_obj = MagicMock()
|
||||
response = _run(
|
||||
logging_obj=logging_obj,
|
||||
litellm_params={},
|
||||
client=_sync_client_returning_converse_response(),
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
if through_environment:
|
||||
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "bedrock-token")
|
||||
litellm.completion(
|
||||
model="bedrock/anthropic.claude-sonnet-4-5-v1:0",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key=None if through_environment else "bedrock-token",
|
||||
max_tokens=7,
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "hi"
|
||||
assert logging_obj.pre_call.call_count == 1
|
||||
|
||||
|
||||
def test_post_call_logging_fires_on_the_sync_rust_path():
|
||||
"""The Rust core owns the provider call, so the Converse transform that
|
||||
normally raises `post_call` never runs. Without the bridge hook every
|
||||
post_call callback goes silent and `original_response` stays unset."""
|
||||
import json
|
||||
|
||||
_inject()
|
||||
logging_obj = MagicMock()
|
||||
_run(logging_obj=logging_obj)
|
||||
|
||||
assert logging_obj.post_call.call_count == 1
|
||||
logged = logging_obj.post_call.call_args.kwargs["original_response"]
|
||||
assert json.loads(logged)["choices"][0]["message"]["content"] == "hello from rust"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_call_logging_fires_on_the_async_rust_path():
|
||||
"""The asynchronous path runs through the same hook, so the two paths
|
||||
cannot drift apart the way the pre_call suppression once did."""
|
||||
import json
|
||||
|
||||
async def native(_request, **_kwargs):
|
||||
return dict(RUST_RESPONSE)
|
||||
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, achat_completions=native
|
||||
)
|
||||
logging_obj = MagicMock()
|
||||
|
||||
with patch.object(
|
||||
BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS
|
||||
):
|
||||
await BedrockConverseLLM().completion(
|
||||
**_completion_kwargs(acompletion=True, logging_obj=logging_obj)
|
||||
)
|
||||
|
||||
assert logging_obj.post_call.call_count == 1
|
||||
logged = logging_obj.post_call.call_args.kwargs["original_response"]
|
||||
assert json.loads(logged)["choices"][0]["message"]["content"] == "hello from rust"
|
||||
|
||||
|
||||
def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines():
|
||||
"""A decline never reached the provider, so the Python path serves the
|
||||
request and owns the only post_call. Firing the hook there too would double
|
||||
every post_call callback for one request."""
|
||||
|
||||
class _Declined(Exception):
|
||||
pass
|
||||
|
||||
class _FakeNative:
|
||||
RustBridgeDeclined = _Declined
|
||||
RustUpstreamError = type("_Upstream", (Exception,), {})
|
||||
|
||||
def declining_native(_request, **_kwargs):
|
||||
raise _Declined("blank message text")
|
||||
|
||||
logging_obj, calls = _recording_logging_obj()
|
||||
|
||||
with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()):
|
||||
bridge.set_rust_chat_completions(
|
||||
decline=lambda **_kwargs: None, chat_completions=declining_native
|
||||
)
|
||||
response = _run(
|
||||
logging_obj=logging_obj,
|
||||
client=_sync_client_returning_converse_response(),
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "hi"
|
||||
assert len(calls["post_call"]) == 1
|
||||
assert "hi" in calls["post_call"][0]["original_response"]
|
||||
|
||||
|
||||
def test_bearer_token_auth_serves_when_boto3_resolves_no_sigv4_credentials(monkeypatch):
|
||||
"""With only `AWS_BEARER_TOKEN_BEDROCK` configured boto3 resolves no
|
||||
credentials at all. Preparing the Rust handoff must not dereference that
|
||||
None: the bearer token signs the request on its own."""
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "bedrock-bearer-token")
|
||||
client = _sync_client_returning_converse_response()
|
||||
|
||||
response = _run(credentials=None, litellm_params={}, client=client)
|
||||
|
||||
assert response.choices[0].message.content == "hi"
|
||||
sent_headers = client.post.call_args.kwargs["headers"]
|
||||
assert sent_headers["Authorization"] == "Bearer bedrock-bearer-token"
|
||||
|
||||
|
||||
def test_the_rust_opt_in_needs_no_sigv4_principal():
|
||||
"""The core resolves the bearer token itself, so a bearer-only deployment
|
||||
keeps its opt-in and the gate sees no aws_* credential keys to sign with."""
|
||||
seen = _inject()
|
||||
|
||||
response = _run(credentials=None, api_key="bedrock-bearer-token")
|
||||
|
||||
assert response.choices[0].message.content == "hello from rust"
|
||||
params = seen["call"][0]["optional_params"]
|
||||
assert not {"aws_access_key_id", "aws_secret_access_key", "aws_session_token"} & params.keys()
|
||||
assert params["aws_region_name"] == "us-east-1"
|
||||
assert seen["call"][0]["api_key"] == "bedrock-bearer-token"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("configured_through", ["env_var", "api_key"])
|
||||
def test_bearer_token_auth_never_runs_the_sigv4_credential_chain(monkeypatch, configured_through):
|
||||
"""The deployment's AWS profile does not exist, so resolving SigV4 credentials
|
||||
raises; a bearer-token deployment must still serve the request, since the
|
||||
bearer token alone signs it."""
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
if configured_through == "env_var":
|
||||
monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "bedrock-bearer-token")
|
||||
else:
|
||||
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
|
||||
client = _sync_client_returning_converse_response()
|
||||
|
||||
response = BedrockConverseLLM().completion(
|
||||
**_completion_kwargs(
|
||||
optional_params={"maxTokens": 16, "aws_profile_name": "litellm-no-such-aws-profile"},
|
||||
litellm_params={},
|
||||
client=client,
|
||||
api_key="bedrock-bearer-token" if configured_through == "api_key" else None,
|
||||
)
|
||||
)
|
||||
|
||||
assert response.choices[0].message.content == "hi"
|
||||
assert client.post.call_args.kwargs["headers"]["Authorization"] == "Bearer bedrock-bearer-token"
|
||||
assert len(requests) == 1
|
||||
assert requests[0][1].api_key == (None if through_environment else "bedrock-token")
|
||||
|
|
|
|||
|
|
@ -10,9 +10,10 @@ import pytest
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.code_interpreter_interception.handler import (
|
||||
CodeInterpreterInterceptionLogger,
|
||||
LITELLM_CODE_EXECUTION_TOOL_NAME,
|
||||
CodeInterpreterInterceptionLogger,
|
||||
)
|
||||
from litellm.llms.azure.videos.transformation import AzureVideoConfig
|
||||
from litellm.llms.base_llm.audio_transcription.transformation import (
|
||||
AudioTranscriptionRequestData,
|
||||
BaseAudioTranscriptionConfig,
|
||||
|
|
@ -24,9 +25,7 @@ from litellm.llms.custom_httpx.llm_http_handler import (
|
|||
_collect_ws_project_quota_callbacks,
|
||||
_google_genai_streaming_hidden_params,
|
||||
_has_pre_call_deployment_hook,
|
||||
_rust_responses_websocket_enabled,
|
||||
)
|
||||
from litellm.llms.azure.videos.transformation import AzureVideoConfig
|
||||
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -151,7 +150,7 @@ def test_response_api_handler_runs_agentic_hooks_in_sync_path(monkeypatch):
|
|||
)
|
||||
logging_obj = Mock()
|
||||
|
||||
monkeypatch.setattr(handler, "_has_agentic_completion_hook", Mock(return_value=True))
|
||||
monkeypatch.setattr(handler, "has_agentic_completion_hook", Mock(return_value=True))
|
||||
hook_mock = AsyncMock(return_value=final_response)
|
||||
monkeypatch.setattr(handler, "_call_agentic_completion_hooks", hook_mock)
|
||||
|
||||
|
|
@ -370,7 +369,7 @@ def test_get_agentic_loop_settings_defaults_and_overrides():
|
|||
assert fingerprints == ["fp-1", "fp-2"]
|
||||
|
||||
|
||||
def test_has_agentic_completion_hook_detection(monkeypatch):
|
||||
def testhas_agentic_completion_hook_detection(monkeypatch):
|
||||
"""The streaming path skips the agentic wrapper only when no callback
|
||||
overrides async_should_run_agentic_loop. Verify both directions."""
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
|
@ -381,7 +380,7 @@ def test_has_agentic_completion_hook_detection(monkeypatch):
|
|||
|
||||
# No callbacks at all -> no agentic hook.
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
assert handler._has_agentic_completion_hook(logging_obj) is False
|
||||
assert handler.has_agentic_completion_hook(logging_obj) is False
|
||||
|
||||
# A plain CustomLogger that does NOT override the gate -> still no hook
|
||||
# (so the wrapper is safely skipped).
|
||||
|
|
@ -389,7 +388,7 @@ def test_has_agentic_completion_hook_detection(monkeypatch):
|
|||
pass
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_PlainLogger()])
|
||||
assert handler._has_agentic_completion_hook(logging_obj) is False
|
||||
assert handler.has_agentic_completion_hook(logging_obj) is False
|
||||
|
||||
# A logger that overrides the gate (directly) -> hook present.
|
||||
class _AgenticLogger(CustomLogger):
|
||||
|
|
@ -399,7 +398,7 @@ def test_has_agentic_completion_hook_detection(monkeypatch):
|
|||
return True, {}
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_AgenticLogger()])
|
||||
assert handler._has_agentic_completion_hook(logging_obj) is True
|
||||
assert handler.has_agentic_completion_hook(logging_obj) is True
|
||||
|
||||
# Override inherited through an intermediate class is still detected
|
||||
# (function-identity check, not a leaf __dict__ check).
|
||||
|
|
@ -407,12 +406,12 @@ def test_has_agentic_completion_hook_detection(monkeypatch):
|
|||
pass
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_DerivedAgenticLogger()])
|
||||
assert handler._has_agentic_completion_hook(logging_obj) is True
|
||||
assert handler.has_agentic_completion_hook(logging_obj) is True
|
||||
|
||||
# Hook supplied via logging_obj.dynamic_success_callbacks is detected too.
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
logging_obj.dynamic_success_callbacks = [_AgenticLogger()]
|
||||
assert handler._has_agentic_completion_hook(logging_obj) is True
|
||||
assert handler.has_agentic_completion_hook(logging_obj) is True
|
||||
|
||||
# String-named callback entry (e.g. "datadog") must be resolved to its
|
||||
# CustomLogger instance via get_custom_logger_compatible_class -- the same
|
||||
|
|
@ -426,7 +425,7 @@ def test_has_agentic_completion_hook_detection(monkeypatch):
|
|||
"litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class",
|
||||
lambda name: agentic_via_string if name == "fake_string_callback" else None,
|
||||
)
|
||||
assert handler._has_agentic_completion_hook(logging_obj) is True
|
||||
assert handler.has_agentic_completion_hook(logging_obj) is True
|
||||
|
||||
# Unresolvable string (returns None) is skipped, no false positive.
|
||||
monkeypatch.setattr(litellm, "callbacks", ["unknown_callback"])
|
||||
|
|
@ -434,7 +433,7 @@ def test_has_agentic_completion_hook_detection(monkeypatch):
|
|||
"litellm.litellm_core_utils.litellm_logging.get_custom_logger_compatible_class",
|
||||
lambda name: None,
|
||||
)
|
||||
assert handler._has_agentic_completion_hook(logging_obj) is False
|
||||
assert handler.has_agentic_completion_hook(logging_obj) is False
|
||||
|
||||
|
||||
def test_fingerprint_agentic_tools_is_deterministic():
|
||||
|
|
@ -1181,7 +1180,7 @@ def test_sync_delete_responses_sets_json_content_type():
|
|||
({}, True, None, None),
|
||||
],
|
||||
)
|
||||
def test_resolve_anthropic_messages_timeout(
|
||||
def testresolve_anthropic_messages_timeout(
|
||||
monkeypatch, litellm_params_kwargs, stream, global_timeout, expected
|
||||
):
|
||||
from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS
|
||||
|
|
@ -1203,7 +1202,7 @@ def test_resolve_anthropic_messages_timeout(
|
|||
"litellm.request_timeout_explicitly_set", True, raising=False
|
||||
)
|
||||
|
||||
resolved = BaseLLMHTTPHandler._resolve_anthropic_messages_timeout(
|
||||
resolved = BaseLLMHTTPHandler.resolve_anthropic_messages_timeout(
|
||||
litellm_params=GenericLiteLLMParams(**litellm_params_kwargs),
|
||||
stream=stream,
|
||||
custom_llm_provider="anthropic",
|
||||
|
|
@ -1467,6 +1466,7 @@ def _make_responses_handler_call(signed_body):
|
|||
signing provider (e.g. Bedrock Mantle).
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -1522,6 +1522,7 @@ def test_responses_handler_signs_after_fake_stream_prep_strips_stream():
|
|||
We snapshot request_data at sign time and assert "stream" is already gone.
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
|
@ -1585,6 +1586,7 @@ def _make_compact_handler_call(signed_body, is_async):
|
|||
signing provider (e.g. Bedrock Mantle SigV4 / bearer).
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -2688,21 +2690,6 @@ async def test_generic_http_handler_async_streaming_forwards_provider_response_h
|
|||
assert "".join([chunk.choices[0].delta.content or "" for chunk in collected]) == "hi"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider, enabled, expected",
|
||||
[("openai", True, True), ("openai", False, False), ("azure", True, False),
|
||||
("hosted_vllm", True, False), (None, True, False)],
|
||||
)
|
||||
def test_the_rust_responses_websocket_needs_openai_and_process_enablement(
|
||||
custom_llm_provider, enabled, expected, monkeypatch
|
||||
):
|
||||
from litellm.rust_bridge import configuration
|
||||
|
||||
configuration.reset_rust_configuration()
|
||||
monkeypatch.setenv("LITELLM_RUST", "1" if enabled else "0")
|
||||
assert _rust_responses_websocket_enabled(custom_llm_provider) is expected
|
||||
|
||||
|
||||
def test_a_plain_callback_does_not_advertise_a_pre_call_deployment_hook(monkeypatch):
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
|
|
|||
|
|
@ -12,10 +12,8 @@ from litellm.llms.azure_ai.ocr.common_utils import (
|
|||
is_azure_document_intelligence_model,
|
||||
)
|
||||
from litellm.ocr.main import _prepare_ocr_request
|
||||
from litellm.rust_bridge.ocr import _rust_bridge_api_base
|
||||
|
||||
_DOC = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
|
||||
_DOC_INTELLIGENCE_ENDPOINT = "https://di.cognitiveservices.azure.com"
|
||||
_AZURE_AI_API_BASE = "https://generic-azure-ai.example.com"
|
||||
|
||||
|
||||
|
|
@ -24,13 +22,6 @@ class _FakeLogging:
|
|||
return None
|
||||
|
||||
|
||||
def _resolve_secret(name: str) -> str | None:
|
||||
return {
|
||||
"AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT": _DOC_INTELLIGENCE_ENDPOINT,
|
||||
"AZURE_AI_API_BASE": _AZURE_AI_API_BASE,
|
||||
}.get(name)
|
||||
|
||||
|
||||
def _prepare(model: str, api_base: str | None):
|
||||
return _prepare_ocr_request(
|
||||
model=model,
|
||||
|
|
@ -57,15 +48,13 @@ class TestIsAzureDocumentIntelligenceModel:
|
|||
|
||||
class TestDocIntelligenceApiBaseResolution:
|
||||
def test_generic_azure_ai_base_does_not_hijack_doc_intelligence(self, monkeypatch):
|
||||
"""Without an explicit api_base, the AZURE_AI_API_BASE fallback must not
|
||||
overwrite the endpoint, so it resolves to the Document Intelligence one."""
|
||||
"""The generic Azure AI fallback must stay out of native admission."""
|
||||
monkeypatch.setenv("AZURE_AI_API_BASE", _AZURE_AI_API_BASE)
|
||||
monkeypatch.delenv("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT", raising=False)
|
||||
|
||||
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", None)
|
||||
|
||||
assert prepared.api_base is None
|
||||
assert _rust_bridge_api_base(prepared, _resolve_secret) == _DOC_INTELLIGENCE_ENDPOINT
|
||||
|
||||
def test_explicit_api_base_is_honoured_for_doc_intelligence(self, monkeypatch):
|
||||
"""A caller-supplied api_base must always win, even for doc-intelligence."""
|
||||
|
|
@ -75,7 +64,6 @@ class TestDocIntelligenceApiBaseResolution:
|
|||
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", custom)
|
||||
|
||||
assert prepared.api_base == custom
|
||||
assert _rust_bridge_api_base(prepared, _resolve_secret) == custom
|
||||
|
||||
def test_generic_azure_ai_base_still_applies_to_mistral_ocr(self, monkeypatch):
|
||||
"""Non doc-intelligence azure_ai models keep using AZURE_AI_API_BASE."""
|
||||
|
|
|
|||
|
|
@ -1,54 +1,14 @@
|
|||
"""
|
||||
Tests for the OCR `req_format` option in the SDK request path:
|
||||
providers that don't support a native response must reject it, and the Rust
|
||||
bridge (which only returns the normalized shape) must not serve native requests.
|
||||
Tests for the OCR `req_format` option in the SDK request path.
|
||||
"""
|
||||
|
||||
import dataclasses
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
|
||||
from litellm.rust_bridge.ocr import PreparedOCRRequest, _rust_ocr_supported
|
||||
|
||||
DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
|
||||
|
||||
|
||||
def _prepared(optional_params: dict[str, object]) -> PreparedOCRRequest:
|
||||
return PreparedOCRRequest(
|
||||
model="doc-intelligence/prebuilt-layout",
|
||||
document=dict(DOCUMENT),
|
||||
api_key="fake-key",
|
||||
api_base="https://example.cognitiveservices.azure.com",
|
||||
custom_llm_provider="azure_ai",
|
||||
extra_headers=None,
|
||||
provider_config=MagicMock(),
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
effective_timeout=60.0,
|
||||
litellm_logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("optional_params", [{}, {"req_format": "litellm"}])
|
||||
def test_rust_ocr_serves_default_format(optional_params):
|
||||
assert _rust_ocr_supported(_prepared(optional_params)) is True
|
||||
|
||||
|
||||
def test_rust_ocr_skipped_for_native_format():
|
||||
assert _rust_ocr_supported(_prepared({"req_format": "native"})) is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider_config", [CohereParseConfig(), AzureAICohereParseConfig()])
|
||||
def test_rust_ocr_skipped_for_configs_without_bridge_support(provider_config):
|
||||
prepared = dataclasses.replace(_prepared({}), provider_config=provider_config)
|
||||
|
||||
assert _rust_ocr_supported(prepared) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_format_rejected_for_provider_without_support_as_bad_request():
|
||||
with pytest.raises(litellm.BadRequestError, match="not supported for provider") as exc_info:
|
||||
|
|
|
|||
|
|
@ -12,8 +12,13 @@ import pytest
|
|||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.request import NativeOCRRequest, NativeRequestContext, NativeRequestOptions
|
||||
from litellm.rust_bridge.runtime import Handled
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeOCRRequest,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
NativeVertexOptions,
|
||||
PreparedNativeCall,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
|
||||
|
|
@ -47,7 +52,6 @@ class RecordingBridge:
|
|||
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
self.contexts: list[NativeRequestContext] = []
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
|
|
@ -65,10 +69,10 @@ class RecordingBridge:
|
|||
"custom_llm_provider": options.custom_llm_provider,
|
||||
"extra_headers": options.extra_headers,
|
||||
"optional_params": request.optional_params,
|
||||
"vertex": options.vertex,
|
||||
"timeout_seconds": options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
self.contexts.append(context)
|
||||
return dict(FAKE_OCR_RESPONSE)
|
||||
|
||||
|
||||
|
|
@ -77,7 +81,6 @@ class RecordingAsyncBridge:
|
|||
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
self.contexts: list[NativeRequestContext] = []
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
|
|
@ -95,10 +98,10 @@ class RecordingAsyncBridge:
|
|||
"custom_llm_provider": options.custom_llm_provider,
|
||||
"extra_headers": options.extra_headers,
|
||||
"optional_params": request.optional_params,
|
||||
"vertex": options.vertex,
|
||||
"timeout_seconds": options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
self.contexts.append(context)
|
||||
return dict(FAKE_OCR_RESPONSE)
|
||||
|
||||
|
||||
|
|
@ -192,7 +195,7 @@ def build_prepared_request(
|
|||
litellm_params: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = 12.5,
|
||||
) -> Any:
|
||||
return rust_bridge.PreparedOCRRequest(
|
||||
return ocr_main._PreparedOCRRequest(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
|
|
@ -210,13 +213,26 @@ def build_prepared_request(
|
|||
@pytest.fixture(autouse=True)
|
||||
def _reset_rust_flag():
|
||||
"""Keep the global toggle isolated between tests."""
|
||||
rust_bridge._OCR.reset()
|
||||
rust_bridge._AOCR.reset()
|
||||
rust_bridge._OCR.sync.reset()
|
||||
rust_bridge._OCR.asynchronous.reset()
|
||||
rust_bridge._PREFLIGHT.reset()
|
||||
configuration.reset_rust_configuration()
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
rust_bridge._PREFLIGHT.override(
|
||||
lambda model, custom_llm_provider, *, context: (
|
||||
"unsupported feature"
|
||||
if any(getattr(context.capabilities, key) for key in ("stream", "has_agentic_hook", "has_custom_client"))
|
||||
or (
|
||||
context.capabilities.request_format == "native"
|
||||
and not (custom_llm_provider == "azure_ai" and "doc-intelligence" in model)
|
||||
)
|
||||
else None
|
||||
)
|
||||
)
|
||||
yield
|
||||
rust_bridge._OCR.reset()
|
||||
rust_bridge._AOCR.reset()
|
||||
rust_bridge._OCR.sync.reset()
|
||||
rust_bridge._OCR.asynchronous.reset()
|
||||
rust_bridge._PREFLIGHT.reset()
|
||||
configuration.reset_rust_configuration()
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
|
||||
|
|
@ -226,7 +242,7 @@ def fake_bridge():
|
|||
"""Enable the Rust path with an injected recording bridge (no native wheel)."""
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
return bridge
|
||||
|
||||
|
||||
|
|
@ -235,14 +251,27 @@ def fake_async_bridge():
|
|||
"""Enable the async Rust path with an injected recording bridge."""
|
||||
bridge = RecordingAsyncBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._AOCR.override(bridge)
|
||||
rust_bridge._OCR.asynchronous.override(bridge)
|
||||
return bridge
|
||||
|
||||
|
||||
def test_rust_toggles_flag():
|
||||
assert configuration.rust_enabled() is False
|
||||
litellm.rust(True)
|
||||
assert configuration.rust_enabled() is True
|
||||
litellm.rust(False)
|
||||
assert configuration.rust_enabled() is False
|
||||
|
||||
|
||||
def test_env_var_enables_rust_ocr(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
assert configuration.rust_enabled() is True
|
||||
|
||||
|
||||
def test_load_rust_ocr_returns_injected_impl():
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
|
||||
|
||||
|
|
@ -306,7 +335,7 @@ def test_native_bridge_available_reflects_loader(monkeypatch):
|
|||
def test_load_rust_aocr_returns_injected_impl():
|
||||
bridge = RecordingAsyncBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._AOCR.override(bridge)
|
||||
rust_bridge._OCR.asynchronous.override(bridge)
|
||||
assert rust_bridge.load_rust_aocr() is bridge
|
||||
|
||||
|
||||
|
|
@ -315,8 +344,8 @@ def test_toggle_without_ocr_arg_preserves_injected_impl():
|
|||
bridge = RecordingBridge()
|
||||
async_bridge = RecordingAsyncBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._AOCR.override(async_bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
rust_bridge._OCR.asynchronous.override(async_bridge)
|
||||
|
||||
litellm.rust(False)
|
||||
assert rust_bridge.load_rust_ocr() is bridge
|
||||
|
|
@ -335,11 +364,12 @@ def test_explicit_ocr_none_clears_injected_impl(monkeypatch):
|
|||
bridge = RecordingBridge()
|
||||
async_bridge = RecordingAsyncBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._AOCR.override(async_bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
rust_bridge._OCR.asynchronous.override(async_bridge)
|
||||
|
||||
rust_bridge._OCR.override(None)
|
||||
rust_bridge._AOCR.override(None)
|
||||
rust_bridge._OCR.sync.reset()
|
||||
rust_bridge._OCR.asynchronous.reset()
|
||||
rust_bridge._PREFLIGHT.reset()
|
||||
assert rust_bridge.load_rust_ocr() is None
|
||||
assert rust_bridge.load_rust_aocr() is None
|
||||
|
||||
|
|
@ -381,13 +411,105 @@ def test_timeout_to_seconds_handles_float_timeout_and_none():
|
|||
assert timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0
|
||||
|
||||
|
||||
def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
|
||||
bridge = RecordingBridge()
|
||||
|
||||
litellm.rust(True)
|
||||
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
response = rust_bridge.dispatch_ocr(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeOCRRequest(
|
||||
model="mistral-ocr-latest",
|
||||
document=DOCUMENT,
|
||||
optional_params={"include_image_base64": True, "pages": [0]},
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
api_key="sk-test",
|
||||
api_base="https://proxy.internal",
|
||||
custom_llm_provider="mistral",
|
||||
extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"},
|
||||
timeout_seconds=12.5,
|
||||
),
|
||||
),
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
adapt=dict,
|
||||
model="mistral-ocr-latest",
|
||||
provider="mistral",
|
||||
eligible=True,
|
||||
)
|
||||
|
||||
assert response == FAKE_OCR_RESPONSE
|
||||
call = bridge.calls[0]
|
||||
assert call == {
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": DOCUMENT,
|
||||
"api_key": "sk-test",
|
||||
"api_base": "https://proxy.internal",
|
||||
"custom_llm_provider": "mistral",
|
||||
"extra_headers": {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
},
|
||||
"optional_params": {"include_image_base64": True, "pages": [0]},
|
||||
"vertex": None,
|
||||
"timeout_seconds": 12.5,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
|
||||
bridge = RecordingAsyncBridge()
|
||||
|
||||
litellm.rust(True)
|
||||
|
||||
rust_bridge._OCR.asynchronous.override(bridge)
|
||||
|
||||
async def unexpected_fallback():
|
||||
pytest.fail("unexpected Python fallback")
|
||||
|
||||
response = await rust_bridge.adispatch_ocr(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeOCRRequest(
|
||||
model="mistral-ocr-maas",
|
||||
document=DOCUMENT,
|
||||
optional_params={},
|
||||
),
|
||||
options=NativeRequestOptions(
|
||||
custom_llm_provider="vertex_ai",
|
||||
vertex=NativeVertexOptions(project="project-1"),
|
||||
timeout_seconds=42.0,
|
||||
),
|
||||
),
|
||||
fallback=unexpected_fallback,
|
||||
adapt=dict,
|
||||
model="mistral-ocr-maas",
|
||||
provider="vertex_ai",
|
||||
eligible=True,
|
||||
)
|
||||
|
||||
assert response == FAKE_OCR_RESPONSE
|
||||
assert bridge.calls[0] == {
|
||||
"model": "mistral-ocr-maas",
|
||||
"document": DOCUMENT,
|
||||
"api_key": None,
|
||||
"api_base": None,
|
||||
"custom_llm_provider": "vertex_ai",
|
||||
"extra_headers": None,
|
||||
"optional_params": {},
|
||||
"vertex": NativeVertexOptions(project="project-1"),
|
||||
"timeout_seconds": 42.0,
|
||||
}
|
||||
|
||||
|
||||
def test_run_rust_ocr_prepares_request_and_wraps_response():
|
||||
bridge = RecordingBridge()
|
||||
logging_obj = RecordingLogging()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
||||
response = rust_bridge.attempt_ocr(
|
||||
response = ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
logging_obj=logging_obj,
|
||||
api_base="https://proxy.internal",
|
||||
|
|
@ -398,13 +520,8 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
|
|||
resolve_api_key=lambda _name: None,
|
||||
)
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert response.pages[0].markdown == "hello world"
|
||||
assert bridge.contexts[0].capabilities.execution_mode == "sync"
|
||||
assert bridge.contexts[0].capabilities.input_source_kind == "document_url"
|
||||
assert bridge.contexts[0].capabilities.native_response_format is False
|
||||
assert bridge.calls[0] == {
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": DOCUMENT,
|
||||
|
|
@ -416,6 +533,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
|
|||
"x-trace-id": "trace-1",
|
||||
},
|
||||
"optional_params": {"include_image_base64": True},
|
||||
"vertex": NativeVertexOptions(),
|
||||
"timeout_seconds": 12.5,
|
||||
}
|
||||
|
||||
|
|
@ -423,9 +541,10 @@ 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.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(api_key=None, timeout=None),
|
||||
resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None,
|
||||
)
|
||||
|
|
@ -436,12 +555,13 @@ 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.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
||||
def _resolver(name: str) -> str | None:
|
||||
raise AssertionError(f"resolver should not be called for {name}")
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
api_key="sk-explicit",
|
||||
timeout=None,
|
||||
|
|
@ -456,13 +576,14 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
|
|||
bridge = RecordingBridge()
|
||||
resolver_calls = []
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
||||
def _resolver(name):
|
||||
resolver_calls.append(name)
|
||||
return "sk-provider-env"
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"),
|
||||
model="provider-ocr-model",
|
||||
|
|
@ -479,9 +600,10 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
|
|||
def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="vertex_ai",
|
||||
model="mistral-ocr-maas",
|
||||
|
|
@ -496,17 +618,14 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
|||
resolve_api_key=lambda _name: None,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["optional_params"] == {
|
||||
"include_image_base64": True,
|
||||
"vertex_project": "project-1",
|
||||
"vertex_location": "us-central1",
|
||||
}
|
||||
assert bridge.calls[0]["optional_params"] == {"include_image_base64": True}
|
||||
assert bridge.calls[0]["vertex"] == NativeVertexOptions(project="project-1", location="us-central1")
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager():
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
||||
def _resolver(name: str) -> str | None:
|
||||
return {
|
||||
|
|
@ -514,7 +633,8 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana
|
|||
"VERTEXAI_LOCATION": "us-east5",
|
||||
}.get(name)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="vertex_ai",
|
||||
model="mistral-ocr-maas",
|
||||
|
|
@ -523,16 +643,16 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana
|
|||
resolve_api_key=_resolver,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["optional_params"]["vertex_project"] == "project-from-secret"
|
||||
assert bridge.calls[0]["optional_params"]["vertex_location"] == "us-east5"
|
||||
assert bridge.calls[0]["vertex"] == NativeVertexOptions(project="project-from-secret", location="us-east5")
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="azure_ai",
|
||||
model="pixtral-12b-2409",
|
||||
|
|
@ -548,9 +668,10 @@ 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.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="azure_ai",
|
||||
model="doc-intelligence/prebuilt-layout",
|
||||
|
|
@ -569,9 +690,10 @@ def test_run_rust_ocr_runs_pre_call_logging():
|
|||
logging_obj = RecordingLogging()
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
rust_bridge._OCR.sync.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
logging_obj=logging_obj,
|
||||
api_base="https://api.mistral.ai/v1",
|
||||
|
|
@ -633,7 +755,7 @@ def test_ocr_exception_type_uses_resolved_provider_context(
|
|||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(RaisingBridge())
|
||||
rust_bridge._OCR.sync.override(RaisingBridge())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
|
@ -679,7 +801,7 @@ async def test_aocr_exception_type_uses_resolved_provider_context(
|
|||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.rust(True)
|
||||
rust_bridge._AOCR.override(RaisingAsyncBridge())
|
||||
rust_bridge._OCR.asynchronous.override(RaisingAsyncBridge())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
|
@ -718,7 +840,7 @@ async def test_ocr_fallback_skips_native_preparation(
|
|||
def unexpected_preparation(*_args: object, **_kwargs: object) -> None:
|
||||
pytest.fail("Python fallback must not resolve native credentials or emit native pre_call")
|
||||
|
||||
monkeypatch.setattr(rust_bridge, "_prepare_rust_ocr_call", unexpected_preparation)
|
||||
monkeypatch.setattr(ocr_main, "_prepare_rust_ocr_call", unexpected_preparation)
|
||||
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fallback)
|
||||
|
||||
response: Final = (
|
||||
|
|
@ -731,25 +853,6 @@ async def test_ocr_fallback_skips_native_preparation(
|
|||
fallback.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_rejects_empty_python_fallback_response(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def fake_exception_type(**kwargs: object) -> CapturedException:
|
||||
captured.update(kwargs)
|
||||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", AsyncMock(return_value=None))
|
||||
|
||||
with pytest.raises(CapturedException, match="wrapped"):
|
||||
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
||||
original: Final = captured["original_exception"]
|
||||
assert isinstance(original, ValueError)
|
||||
assert str(original) == "Got an unexpected None response from the OCR API: None"
|
||||
|
||||
|
||||
def test_ocr_provider_configs_expose_api_key_env_vars():
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
|
|
|
|||
|
|
@ -2,14 +2,8 @@ from __future__ import annotations
|
|||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
|
||||
from litellm.rust_bridge import configuration, responses_websocket
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
NativeResponsesWebSocketRequest,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
|
||||
from litellm.rust_bridge.request import NativeRequestContext, NativeResponsesWebSocketRequest
|
||||
|
||||
|
||||
class _FakeNativeConnection:
|
||||
|
|
@ -33,55 +27,55 @@ class _ClosedNativeConnection:
|
|||
|
||||
|
||||
class _FakeNativeBridge:
|
||||
contexts: list[NativeRequestContext] = []
|
||||
|
||||
@classmethod
|
||||
async def connect(
|
||||
cls,
|
||||
request: NativeResponsesWebSocketRequest,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
) -> _FakeNativeConnection:
|
||||
cls.contexts.append(context)
|
||||
return _FakeNativeConnection()
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_responses_websocket():
|
||||
responses_websocket.set_rust_responses_websocket(connection=None)
|
||||
responses_websocket.set_rust_responses_websocket(connection=None, decline=None)
|
||||
configuration.reset_rust_configuration()
|
||||
responses_websocket.set_rust_responses_websocket(
|
||||
decline=lambda model, custom_llm_provider, *, context: (
|
||||
"unsupported feature"
|
||||
if any(getattr(context.capabilities, key) for key in ("stream", "has_agentic_hook", "has_custom_client"))
|
||||
or context.capabilities.request_format == "native"
|
||||
else None
|
||||
)
|
||||
)
|
||||
yield
|
||||
responses_websocket.set_rust_responses_websocket(connection=None)
|
||||
responses_websocket.set_rust_responses_websocket(connection=None, decline=None)
|
||||
configuration.reset_rust_configuration()
|
||||
|
||||
|
||||
def test_rust_websocket_bridge_uses_process_enablement() -> None:
|
||||
configuration.rust(False)
|
||||
assert not _rust_responses_websocket_enabled("openai")
|
||||
configuration.rust(True)
|
||||
assert _rust_responses_websocket_enabled("openai")
|
||||
assert not _rust_responses_websocket_enabled("anthropic")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
|
||||
adapter = responses_websocket.ConnectionAdapter(_ClosedNativeConnection())
|
||||
adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection())
|
||||
|
||||
with pytest.raises(responses_websocket.ConnectionClosedOK):
|
||||
await adapter.recv()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bridge_reports_unavailable(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
configuration.rust(True)
|
||||
responses_websocket._RESPONSES_WEBSOCKET.override(None)
|
||||
|
||||
assert await responses_websocket.connect(
|
||||
url="wss://example.test/responses",
|
||||
headers={},
|
||||
timeout=None,
|
||||
) == NativeSkipped(NativeSkipReason.UNAVAILABLE)
|
||||
assert (
|
||||
await responses_websocket.connect(
|
||||
url="wss://example.test/responses",
|
||||
headers={},
|
||||
timeout=None,
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -97,13 +91,10 @@ async def test_enabled_bridge_connects_and_adapts_socket(
|
|||
timeout=1.0,
|
||||
)
|
||||
|
||||
assert isinstance(connection, Handled)
|
||||
connection = connection.value
|
||||
assert connection is not None
|
||||
await connection.send("response.create")
|
||||
assert await connection.recv() == "response.completed"
|
||||
await connection.close()
|
||||
assert _FakeNativeBridge.contexts[-1].capabilities.websocket_mode == "native"
|
||||
assert _FakeNativeBridge.contexts[-1].capabilities.requires_connection is True
|
||||
|
||||
|
||||
class _FailingNativeBridge:
|
||||
|
|
@ -112,74 +103,96 @@ class _FailingNativeBridge:
|
|||
cls,
|
||||
request: NativeResponsesWebSocketRequest,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
) -> _FakeNativeConnection:
|
||||
raise RuntimeError("connection failed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_failure_is_reported_to_orchestration() -> None:
|
||||
configuration.rust(True)
|
||||
responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge)
|
||||
result = await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None)
|
||||
assert isinstance(result, NativeFailed)
|
||||
assert str(result.error) == "connection failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_connection_closes_native_socket_on_consumer_failure() -> None:
|
||||
configuration.rust(True)
|
||||
socket = _FakeNativeConnection()
|
||||
|
||||
class Bridge:
|
||||
@classmethod
|
||||
async def connect(
|
||||
cls,
|
||||
request: NativeResponsesWebSocketRequest,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
) -> _FakeNativeConnection:
|
||||
return socket
|
||||
|
||||
responses_websocket.set_rust_responses_websocket(connection=Bridge)
|
||||
result = await responses_websocket.managed_connect(url="wss://example.test/responses", headers={}, timeout=1.0)
|
||||
assert isinstance(result, Handled)
|
||||
|
||||
async def use_connection() -> None:
|
||||
async with result.value as connection:
|
||||
await connection.send("hello")
|
||||
raise ValueError("consumer failed")
|
||||
|
||||
with pytest.raises(ValueError, match="consumer failed"):
|
||||
await use_connection()
|
||||
assert socket.sent == ["hello"]
|
||||
assert socket.closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_failure_does_not_authorize_python_fallback() -> None:
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
|
||||
from litellm.rust_bridge.dispatch import anative_context, provider_errors
|
||||
|
||||
configuration.rust(True)
|
||||
responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge)
|
||||
|
||||
@anative_context(
|
||||
native=lambda: responses_websocket.managed_connect(
|
||||
url="wss://example.test/responses", headers={}, timeout=None
|
||||
),
|
||||
route="responses_websocket",
|
||||
errors=lambda: provider_errors("openai", "responses websocket"),
|
||||
)
|
||||
def execute() -> AbstractAsyncContextManager[object]:
|
||||
pytest.fail("unknown native failures must not open a Python connection")
|
||||
|
||||
async def run() -> None:
|
||||
async with execute():
|
||||
pytest.fail("connection must fail before entering its body")
|
||||
|
||||
with pytest.raises(RuntimeError, match="connection failed"):
|
||||
await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("native", [False, True])
|
||||
@pytest.mark.parametrize("session_error", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_dispatch_cleans_up_without_reconnecting(native, session_error):
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
configuration.rust(True)
|
||||
native_socket = _FakeNativeConnection()
|
||||
python_socket = _FakeNativeConnection()
|
||||
connections = []
|
||||
|
||||
class Native:
|
||||
@classmethod
|
||||
async def connect(cls, request, *, options, context):
|
||||
connections.append("native")
|
||||
assert options.custom_llm_provider == "azure"
|
||||
return native_socket
|
||||
|
||||
@asynccontextmanager
|
||||
async def python():
|
||||
connections.append("python")
|
||||
try:
|
||||
yield python_socket
|
||||
finally:
|
||||
await python_socket.close()
|
||||
|
||||
responses_websocket.set_rust_responses_websocket(connection=Native)
|
||||
if not native:
|
||||
responses_websocket.set_rust_responses_websocket(
|
||||
decline=lambda model, custom_llm_provider, **features: "declined"
|
||||
)
|
||||
|
||||
async def run():
|
||||
async with responses_websocket.open_connection(
|
||||
url="wss://example.test",
|
||||
headers={},
|
||||
timeout=1,
|
||||
model="test-model",
|
||||
provider="azure",
|
||||
fallback=python,
|
||||
):
|
||||
if session_error:
|
||||
raise RuntimeError("session failed")
|
||||
|
||||
if session_error:
|
||||
with pytest.raises(RuntimeError, match="session failed"):
|
||||
await run()
|
||||
else:
|
||||
await run()
|
||||
assert connections == ["native" if native else "python"]
|
||||
assert native_socket.closed == native
|
||||
assert python_socket.closed == (not native)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_acceptance_export_uses_python_connection_once():
|
||||
from contextlib import asynccontextmanager
|
||||
|
||||
calls = []
|
||||
socket = _FakeNativeConnection()
|
||||
|
||||
@asynccontextmanager
|
||||
async def python():
|
||||
calls.append("python")
|
||||
try:
|
||||
yield socket
|
||||
finally:
|
||||
await socket.close()
|
||||
|
||||
configuration.rust(True)
|
||||
responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge)
|
||||
responses_websocket._PREFLIGHT.override(None)
|
||||
async with responses_websocket.open_connection(
|
||||
url="wss://example.test", headers={}, timeout=1, model="model", provider="openai", fallback=python
|
||||
) as connection:
|
||||
assert connection is socket
|
||||
assert calls == ["python"]
|
||||
assert socket.closed
|
||||
|
|
|
|||
|
|
@ -12,8 +12,12 @@ import pytest
|
|||
import litellm
|
||||
from litellm.rust_bridge import bindings, configuration
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
from litellm.rust_bridge.request import NativeChatCompletionsRequest, NativeRequestContext, NativeRequestOptions
|
||||
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeBedrockOptions,
|
||||
NativeRequestCapabilities,
|
||||
NativeRequestContext,
|
||||
anthropic_options,
|
||||
)
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
RUST_RESPONSE = {
|
||||
|
|
@ -96,40 +100,16 @@ class _RecordingCall:
|
|||
self.result = result if result is not None else dict(RUST_RESPONSE)
|
||||
self.error = error
|
||||
self.calls: list[dict] = []
|
||||
self.contexts: list[NativeRequestContext] = []
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
request: NativeChatCompletionsRequest,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
):
|
||||
kwargs = {
|
||||
"model": request.model,
|
||||
"messages": request.messages,
|
||||
"optional_params": request.optional_params,
|
||||
"api_key": options.api_key,
|
||||
"api_base": options.api_base,
|
||||
"custom_llm_provider": options.custom_llm_provider,
|
||||
"extra_headers": options.extra_headers,
|
||||
"timeout_seconds": options.timeout_seconds,
|
||||
}
|
||||
self.calls.append(kwargs)
|
||||
self.contexts.append(context)
|
||||
def __call__(self, request, *, options, context):
|
||||
self.calls.append({"request": request, "options": options, "context": context})
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
return self.result
|
||||
|
||||
|
||||
class _RecordingAsyncCall(_RecordingCall):
|
||||
async def __call__(
|
||||
self,
|
||||
request: NativeChatCompletionsRequest,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
):
|
||||
async def __call__(self, request, *, options, context):
|
||||
return _RecordingCall.__call__(self, request, options=options, context=context)
|
||||
|
||||
|
||||
|
|
@ -176,25 +156,17 @@ class TestGate:
|
|||
|
||||
def test_declines_streaming_and_providers_off_the_path(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
gate = _RecordingDecline()
|
||||
gate = pytest.importorskip("litellm.rust_bridge._native").chat_completions_decline
|
||||
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 == []
|
||||
|
||||
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.
|
||||
|
||||
It does that inside the function the Rust route replaces, and the core is
|
||||
handed `optional_params` only, so accepting here would send the request
|
||||
to Anthropic with the abuse-detection attribution silently missing.
|
||||
"""
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
gate = _RecordingDecline()
|
||||
gate = pytest.importorskip("litellm.rust_bridge._native").chat_completions_decline
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
assert _accepts(litellm_params={"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"
|
||||
|
||||
# Bedrock's Converse transform reads no `user_id`, and an Anthropic request
|
||||
# whose metadata carries none is one Python would not attribute either.
|
||||
|
|
@ -202,6 +174,7 @@ class TestGate:
|
|||
_accepts(
|
||||
custom_llm_provider="bedrock",
|
||||
model="bedrock/us-east-1/anthropic.claude-v2",
|
||||
optional_params={"maxTokens": 16},
|
||||
litellm_params={"metadata": {"user_id": "u-123"}},
|
||||
)
|
||||
is True
|
||||
|
|
@ -209,6 +182,19 @@ class TestGate:
|
|||
assert _accepts(litellm_params={"metadata": {"trace_id": "t-1"}}) is True
|
||||
assert _accepts(litellm_params={"metadata": {"user_id": None}}) is True
|
||||
assert _accepts(litellm_params={"metadata": None}) is True
|
||||
assert _accepts(litellm_params={"litellm_metadata": {"user_id": "u-123"}}) is True
|
||||
assert _accepts(litellm_params={"metadata": "invalid"}) is True
|
||||
assert _accepts(litellm_params={"metadata": {"trace": object()}}) is True
|
||||
assert _accepts(litellm_params={"metadata": {"user_id": object()}}) is False
|
||||
assert (
|
||||
_accepts(
|
||||
custom_llm_provider="bedrock",
|
||||
model="bedrock/us-east-1/anthropic.claude-v2",
|
||||
optional_params={"maxTokens": 16},
|
||||
litellm_params={"metadata": {"user_id": object()}},
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
def test_declines_a_bedrock_request_while_the_proxy_owns_request_metadata(self, monkeypatch):
|
||||
"""`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the
|
||||
|
|
@ -217,16 +203,16 @@ class TestGate:
|
|||
who armed `bedrock_request_metadata_fields` keeps the Python path.
|
||||
"""
|
||||
monkeypatch.setenv("LITELLM_RUST", "1")
|
||||
gate = _RecordingDecline()
|
||||
gate = pytest.importorskip("litellm.rust_bridge._native").chat_completions_decline
|
||||
bridge.set_rust_chat_completions(decline=gate)
|
||||
bedrock = {
|
||||
"custom_llm_provider": "bedrock",
|
||||
"model": "bedrock/us-east-1/anthropic.claude-v2",
|
||||
"optional_params": {"maxTokens": 16},
|
||||
}
|
||||
|
||||
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"
|
||||
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", None)
|
||||
|
|
@ -276,8 +262,7 @@ class TestSyncCall:
|
|||
|
||||
result = bridge.chat_completions(**_call_kwargs(model_response))
|
||||
|
||||
assert isinstance(result, Handled)
|
||||
result = result.value
|
||||
assert result is not None
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
assert result.model == "claude-sonnet-4-5-20260101"
|
||||
|
|
@ -291,28 +276,16 @@ class TestSyncCall:
|
|||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert native.calls[0]["timeout_seconds"] == 30.0
|
||||
assert native.calls[0]["options"].timeout_seconds == 30.0
|
||||
|
||||
def test_preserves_execution_and_client_capabilities(self):
|
||||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
bridge.chat_completions(
|
||||
**_call_kwargs(ModelResponse()),
|
||||
stream=True,
|
||||
has_custom_client=True,
|
||||
)
|
||||
assert native.contexts[0].capabilities.execution_mode == "sync"
|
||||
assert native.contexts[0].capabilities.stream is True
|
||||
assert native.contexts[0].capabilities.has_custom_client is True
|
||||
|
||||
def test_reports_unavailable_bridge(self, monkeypatch):
|
||||
def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeSkipped)
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
def test_reports_native_decline_to_orchestration(self, monkeypatch):
|
||||
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 isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeFailed)
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
|
||||
class TestAsyncCall:
|
||||
|
|
@ -320,18 +293,290 @@ class TestAsyncCall:
|
|||
async def test_builds_a_model_response(self):
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
|
||||
result = await bridge.achat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert isinstance(result, Handled)
|
||||
result = result.value
|
||||
assert result is not None
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reports_unavailable_bridge(self, monkeypatch):
|
||||
async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeSkipped)
|
||||
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reports_native_decline_to_orchestration(self, monkeypatch):
|
||||
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 isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeFailed)
|
||||
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
|
||||
class TestAsyncFallbackWrapper:
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_the_rust_response_without_running_the_fallback(self):
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
|
||||
ran = []
|
||||
|
||||
async def fallback():
|
||||
ran.append(True)
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert ran == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
|
||||
|
||||
async def fallback():
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result == "python"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
|
||||
async def fallback():
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result == "python"
|
||||
|
||||
|
||||
class TestFailureClassification:
|
||||
"""A failure the provider already saw must not be retried on the Python
|
||||
path: it would bill the customer for the same work twice."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _native_exceptions(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
|
||||
def test_a_decline_falls_back_because_nothing_was_sent(self):
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
def test_an_upstream_failure_is_surfaced_with_its_status(self):
|
||||
from litellm.exceptions import RateLimitError
|
||||
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited")))
|
||||
with pytest.raises(RateLimitError) as raised:
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert raised.value.status_code == 429
|
||||
assert "rate limited" in str(raised.value)
|
||||
|
||||
def test_a_transport_failure_with_no_response_surfaces_as_a_500(self):
|
||||
from litellm.exceptions import APIError
|
||||
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset")))
|
||||
with pytest.raises(APIError) as raised:
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert raised.value.status_code == 500
|
||||
|
||||
def test_an_unrecognized_error_is_not_swallowed(self):
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else")))
|
||||
with pytest.raises(RuntimeError):
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self):
|
||||
from litellm.exceptions import InternalServerError
|
||||
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom")))
|
||||
ran = []
|
||||
|
||||
async def fallback():
|
||||
ran.append(True)
|
||||
return "python"
|
||||
|
||||
with pytest.raises(InternalServerError):
|
||||
await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert ran == [], "a request the provider already served must not be re-issued"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_wrapper_falls_back_on_a_decline(self):
|
||||
bridge.set_rust_chat_completions(
|
||||
achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text"))
|
||||
)
|
||||
|
||||
async def fallback():
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result == "python"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_native_exception_types_does_not_authorize_python_fallback(monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(
|
||||
chat_completions=_RecordingCall(error=RuntimeError("connection failed")),
|
||||
achat_completions=_RecordingAsyncCall(error=RuntimeError("connection failed")),
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="connection failed"):
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
|
||||
async def fallback():
|
||||
pytest.fail("unknown failure must not retry through Python")
|
||||
|
||||
with pytest.raises(RuntimeError, match="connection failed"):
|
||||
await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
|
||||
|
||||
def test_provider_credentials_are_separate_from_chat_body_params():
|
||||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
configuration.rust(True)
|
||||
kwargs = _call_kwargs(ModelResponse())
|
||||
kwargs["optional_params"] = {
|
||||
"max_tokens": 32,
|
||||
}
|
||||
kwargs["bedrock"] = NativeBedrockOptions(
|
||||
aws_access_key_id="test-access-key",
|
||||
aws_secret_access_key="test-secret-key",
|
||||
)
|
||||
bridge.chat_completions(**kwargs)
|
||||
request = native.calls[0]["request"]
|
||||
options = native.calls[0]["options"]
|
||||
assert request.optional_params == {"max_tokens": 32}
|
||||
assert options.bedrock.aws_access_key_id == "test-access-key"
|
||||
assert options.bedrock.aws_secret_access_key == "test-secret-key"
|
||||
|
||||
|
||||
def test_provider_payload_extensions_cross_the_boundary_without_partitioning():
|
||||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
configuration.rust(True)
|
||||
extensions = {
|
||||
"vendor_object": {"nested": None},
|
||||
"vendor_array": [1, "two", False],
|
||||
"vendor_scalar": 0.25,
|
||||
"extra_body": {"temperature": 0.2, "config": {"replacement": True}},
|
||||
}
|
||||
|
||||
kwargs = _call_kwargs(ModelResponse())
|
||||
kwargs["optional_params"] = extensions
|
||||
bridge.chat_completions(**kwargs)
|
||||
|
||||
assert native.calls[0]["request"].optional_params == extensions
|
||||
|
||||
|
||||
def test_typed_capability_and_provider_metadata_facts_are_isolated():
|
||||
context = NativeRequestContext(
|
||||
capabilities=NativeRequestCapabilities(
|
||||
stream=True,
|
||||
has_agentic_hook=True,
|
||||
has_custom_client=True,
|
||||
request_format="native",
|
||||
)
|
||||
)
|
||||
anthropic = anthropic_options({"metadata": {"user_id": "user-123", "ignored": object()}})
|
||||
|
||||
assert context.capabilities.request_format == "native"
|
||||
assert context.capabilities.has_agentic_hook is True
|
||||
assert anthropic.user_id == "user-123"
|
||||
assert anthropic.has_user_id is True
|
||||
assert anthropic_options({"metadata": {"user_id": object()}}).has_user_id is True
|
||||
assert anthropic_options({"metadata": {"user_id": None}}).has_user_id is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider", ["anthropic", "bedrock", "openai"])
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_public_completion_discovers_any_provider(provider, asynchronous):
|
||||
native = _RecordingCall()
|
||||
anative = _RecordingAsyncCall()
|
||||
gate = _RecordingDecline()
|
||||
bridge.set_rust_chat_completions(chat_completions=native, achat_completions=anative, decline=gate)
|
||||
kwargs = {
|
||||
"model": f"{provider}/test-model",
|
||||
"messages": MESSAGES,
|
||||
"api_key": "key",
|
||||
"max_tokens": 16,
|
||||
"num_retries": 0,
|
||||
}
|
||||
response = await litellm.acompletion(**kwargs) if asynchronous else litellm.completion(**kwargs)
|
||||
assert response.choices[0].message.content == "hello from rust"
|
||||
calls = anative.calls if asynchronous else native.calls
|
||||
assert len(calls) == 1
|
||||
assert calls[0]["options"].custom_llm_provider == provider
|
||||
assert calls[0]["request"].messages == MESSAGES
|
||||
assert gate.calls[0]["custom_llm_provider"] == provider
|
||||
assert len(native.calls) + len(anative.calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
"failure", ["preflight", "missing_preflight", "decline", "unavailable", "error", "malformed", "cancelled"]
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_public_completion_fallback_contract(monkeypatch, asynchronous, failure):
|
||||
import asyncio
|
||||
import importlib
|
||||
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class Recorder(CustomLogger):
|
||||
def __init__(self):
|
||||
self.pre = 0
|
||||
self.post = 0
|
||||
|
||||
def log_pre_api_call(self, model, messages, kwargs):
|
||||
self.pre += 1
|
||||
|
||||
def log_post_api_call(self, kwargs, response_obj, start_time, end_time):
|
||||
self.post += 1
|
||||
|
||||
recorder = Recorder()
|
||||
monkeypatch.setattr(litellm, "input_callback", [recorder])
|
||||
_fake_native_bridge(monkeypatch)
|
||||
python_calls = []
|
||||
|
||||
def python(ctx):
|
||||
python_calls.append(ctx)
|
||||
|
||||
def finish():
|
||||
ctx.logging.pre_call(input=MESSAGES, api_key="key", additional_args={})
|
||||
ctx.logging.post_call(input=MESSAGES, api_key="key", original_response="python")
|
||||
return ModelResponse(choices=[{"message": {"role": "assistant", "content": "python"}}])
|
||||
|
||||
async def afinish():
|
||||
return finish()
|
||||
|
||||
return afinish() if ctx.acompletion else finish()
|
||||
|
||||
monkeypatch.setattr(importlib.import_module("litellm.main"), "_complete_python", python)
|
||||
native_error = {
|
||||
"decline": _FakeDeclined("request unsupported"),
|
||||
"error": RuntimeError("execution failed"),
|
||||
"cancelled": asyncio.CancelledError(),
|
||||
}.get(failure)
|
||||
native = _RecordingCall(result={} if failure == "malformed" else None, error=native_error)
|
||||
anative = _RecordingAsyncCall(result={} if failure == "malformed" else None, error=native_error)
|
||||
bridge.set_rust_chat_completions(
|
||||
chat_completions=native,
|
||||
achat_completions=anative,
|
||||
decline=_RecordingDecline("unsupported" if failure == "preflight" else None),
|
||||
)
|
||||
if failure == "missing_preflight":
|
||||
bridge._CHAT_PREFLIGHT.override(None)
|
||||
if failure == "unavailable":
|
||||
bridge._CHAT.sync.override(None)
|
||||
bridge._CHAT.asynchronous.override(None)
|
||||
|
||||
async def run():
|
||||
kwargs = {"model": "openai/test-model", "messages": MESSAGES, "api_key": "key", "num_retries": 0}
|
||||
return await litellm.acompletion(**kwargs) if asynchronous else litellm.completion(**kwargs)
|
||||
|
||||
if failure in {"error", "malformed", "cancelled"}:
|
||||
with pytest.raises(asyncio.CancelledError if failure == "cancelled" else Exception):
|
||||
await run()
|
||||
assert python_calls == []
|
||||
else:
|
||||
result = await run()
|
||||
assert result.choices[0].message.content == "python"
|
||||
assert len(python_calls) == 1
|
||||
assert recorder.pre == 1
|
||||
assert recorder.post == 1
|
||||
|
|
|
|||
|
|
@ -1,117 +1,460 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from types import SimpleNamespace
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.rust_bridge import runtime
|
||||
from litellm.exceptions import APIError, AuthenticationError, InternalServerError, RateLimitError
|
||||
from litellm.rust_bridge import bindings, runtime
|
||||
|
||||
|
||||
class RustBridgeDeclined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class RustUpstreamError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def native_exceptions(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(
|
||||
bindings,
|
||||
"get_native_bridge",
|
||||
lambda: SimpleNamespace(
|
||||
RustBridgeDeclined=RustBridgeDeclined,
|
||||
RustUpstreamError=RustUpstreamError,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def context() -> runtime.BridgeErrorContext:
|
||||
return runtime.BridgeErrorContext(provider="anthropic", model="model")
|
||||
|
||||
|
||||
def enabled() -> bool:
|
||||
return True
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class FallbackCase:
|
||||
process_enabled: bool | None = None
|
||||
eligible: bool = True
|
||||
binding_available: bool = True
|
||||
declined: bool = False
|
||||
expected_events: tuple[str, ...] = ()
|
||||
|
||||
|
||||
FALLBACK_CASES: Final = (
|
||||
pytest.param(
|
||||
FallbackCase(process_enabled=False, expected_events=("python",)),
|
||||
id="process-disabled",
|
||||
),
|
||||
pytest.param(
|
||||
FallbackCase(eligible=False, expected_events=("python",)),
|
||||
id="request-ineligible",
|
||||
),
|
||||
pytest.param(
|
||||
FallbackCase(binding_available=False, expected_events=("load", "python")),
|
||||
id="bridge-unavailable",
|
||||
),
|
||||
pytest.param(
|
||||
FallbackCase(declined=True, expected_events=("load", "prepare", "rust", "python")),
|
||||
id="bridge-declined",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("case", FALLBACK_CASES)
|
||||
def test_invoke_falls_back_only_before_provider_success(case: FallbackCase) -> None:
|
||||
events: list[str] = []
|
||||
|
||||
def load() -> object | None:
|
||||
events.append("load")
|
||||
return object() if case.binding_available else None
|
||||
|
||||
def call(_binding: object, _request: object) -> int:
|
||||
events.append("rust")
|
||||
if case.declined:
|
||||
raise RustBridgeDeclined("unsupported")
|
||||
return 3
|
||||
|
||||
bridge: Final = runtime.EndpointBinding(
|
||||
route="messages", load=load, enabled=lambda: case.process_enabled is not False
|
||||
)
|
||||
result: Final = bridge.invoke(
|
||||
prepare=lambda: events.append("prepare"),
|
||||
call=call,
|
||||
fallback=lambda: events.append("python") or "fallback",
|
||||
adapt=str,
|
||||
error_context=context(),
|
||||
eligible=case.eligible,
|
||||
)
|
||||
|
||||
assert result == "fallback"
|
||||
assert tuple(events) == case.expected_events
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("case", FALLBACK_CASES)
|
||||
async def test_ainvoke_matches_sync_fallback_contract(case: FallbackCase) -> None:
|
||||
events: list[str] = []
|
||||
|
||||
def load() -> object | None:
|
||||
events.append("load")
|
||||
return object() if case.binding_available else None
|
||||
|
||||
async def call(_binding: object, _request: object) -> int:
|
||||
events.append("rust")
|
||||
if case.declined:
|
||||
raise RustBridgeDeclined("unsupported")
|
||||
return 3
|
||||
|
||||
async def fallback() -> str:
|
||||
events.append("python")
|
||||
return "fallback"
|
||||
|
||||
bridge: Final = runtime.EndpointBinding(
|
||||
route="messages", load=load, enabled=lambda: case.process_enabled is not False
|
||||
)
|
||||
result: Final = await bridge.ainvoke(
|
||||
prepare=lambda: events.append("prepare"),
|
||||
call=call,
|
||||
fallback=fallback,
|
||||
adapt=str,
|
||||
error_context=context(),
|
||||
eligible=case.eligible,
|
||||
)
|
||||
|
||||
assert result == "fallback"
|
||||
assert tuple(events) == case.expected_events
|
||||
|
||||
|
||||
def test_invoke_adapts_native_success_without_fallback() -> None:
|
||||
bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled)
|
||||
|
||||
result: Final = bridge.invoke(
|
||||
prepare=lambda: 3,
|
||||
call=lambda _binding, request: request * 2,
|
||||
fallback=lambda: pytest.fail("fallback must not run"),
|
||||
adapt=lambda value: f"adapted-{value}",
|
||||
error_context=context(),
|
||||
)
|
||||
|
||||
assert result == "adapted-6"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ainvoke_adapts_native_success_without_fallback() -> None:
|
||||
async def call(_binding: object, request: int) -> int:
|
||||
return request * 2
|
||||
|
||||
async def fallback() -> str:
|
||||
pytest.fail("fallback must not run")
|
||||
|
||||
bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled)
|
||||
result: Final = await bridge.ainvoke(
|
||||
prepare=lambda: 3,
|
||||
call=call,
|
||||
fallback=fallback,
|
||||
adapt=lambda value: f"adapted-{value}",
|
||||
error_context=context(),
|
||||
)
|
||||
|
||||
assert result == "adapted-6"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("error", "expected_type", "expected_status", "expected_message"),
|
||||
(
|
||||
pytest.param(RustUpstreamError(401, "unauthorized"), AuthenticationError, 401, "unauthorized", id="auth"),
|
||||
pytest.param(RustUpstreamError(429, "rate limited"), RateLimitError, 429, "rate limited", id="rate-limit"),
|
||||
pytest.param(RustUpstreamError(500, "failed"), InternalServerError, 500, "failed", id="server-error"),
|
||||
pytest.param(RustUpstreamError(0, "connection reset"), APIError, 500, "connection reset", id="transport"),
|
||||
pytest.param(RustUpstreamError(403, "forbidden"), APIError, 403, "forbidden", id="other-status"),
|
||||
),
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("state", ("disabled", "ineligible", "unavailable", "handled"))
|
||||
async def test_attempt_only_prepares_selected_requests(asynchronous: bool, state: str) -> None:
|
||||
events: Final[list[str]] = []
|
||||
async def test_upstream_failure_maps_to_api_error_without_fallback(
|
||||
asynchronous: bool,
|
||||
error: RustUpstreamError,
|
||||
expected_type: type[BaseException],
|
||||
expected_status: int,
|
||||
expected_message: str,
|
||||
) -> None:
|
||||
def fail(_binding: object, _request: object) -> object:
|
||||
raise error
|
||||
|
||||
async def afail(binding: object, request: object) -> object:
|
||||
return fail(binding, request)
|
||||
|
||||
async def fallback() -> str:
|
||||
pytest.fail("fallback must not run")
|
||||
|
||||
bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled)
|
||||
|
||||
async def invoke() -> None:
|
||||
if asynchronous:
|
||||
await bridge.ainvoke(
|
||||
prepare=lambda: None, call=afail, fallback=fallback, adapt=str, error_context=context()
|
||||
)
|
||||
else:
|
||||
bridge.invoke(
|
||||
prepare=lambda: None,
|
||||
call=fail,
|
||||
fallback=lambda: pytest.fail("fallback must not run"),
|
||||
adapt=str,
|
||||
error_context=context(),
|
||||
)
|
||||
|
||||
with pytest.raises(expected_type, match=expected_message) as caught:
|
||||
await invoke()
|
||||
|
||||
assert type(caught.value) is expected_type
|
||||
assert caught.value.status_code == expected_status
|
||||
assert caught.value.llm_provider == "anthropic"
|
||||
assert caught.value.model == "model"
|
||||
assert caught.value.__cause__ is error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_upstream_failure_maps_to_api_error_without_fallback() -> None:
|
||||
async def fail(_binding: object, _request: object) -> object:
|
||||
raise RustUpstreamError(503, "overloaded")
|
||||
|
||||
async def fallback() -> object:
|
||||
pytest.fail("fallback must not run")
|
||||
|
||||
bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled)
|
||||
|
||||
with pytest.raises(APIError, match="overloaded") as caught:
|
||||
await bridge.ainvoke(prepare=lambda: None, call=fail, fallback=fallback, adapt=str, error_context=context())
|
||||
|
||||
assert caught.value.status_code == 503
|
||||
|
||||
|
||||
def test_unknown_failure_is_preserved_without_fallback() -> None:
|
||||
error: Final = RuntimeError("unknown")
|
||||
bridge: Final = runtime.EndpointBinding(route="messages", load=object, enabled=enabled)
|
||||
|
||||
with pytest.raises(RuntimeError, match="unknown") as caught:
|
||||
bridge.invoke(
|
||||
prepare=lambda: None,
|
||||
call=lambda _binding, _request: (_ for _ in ()).throw(error),
|
||||
fallback=lambda: pytest.fail("fallback must not run"),
|
||||
adapt=str,
|
||||
error_context=context(),
|
||||
)
|
||||
|
||||
assert caught.value is error
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("process_enabled", "binding_available", "declined", "expected_message"),
|
||||
(
|
||||
pytest.param(False, True, False, "native messages endpoint is disabled", id="disabled"),
|
||||
pytest.param(None, False, False, "native messages endpoint is unavailable", id="unavailable"),
|
||||
pytest.param(
|
||||
None,
|
||||
True,
|
||||
True,
|
||||
"native messages endpoint declined the request: unsupported",
|
||||
id="declined",
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_require_explains_why_rust_did_not_handle_request(
|
||||
process_enabled: bool | None,
|
||||
binding_available: bool,
|
||||
declined: bool,
|
||||
expected_message: str,
|
||||
) -> None:
|
||||
def call(_binding: object, _request: object) -> object:
|
||||
if declined:
|
||||
raise RustBridgeDeclined("unsupported")
|
||||
return object()
|
||||
|
||||
bridge: Final = runtime.EndpointBinding(
|
||||
route="messages",
|
||||
load=object if binding_available else lambda: None,
|
||||
enabled=lambda: process_enabled is not False,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match=f"^{expected_message}$"):
|
||||
bridge.require(
|
||||
prepare=lambda: None,
|
||||
call=call,
|
||||
adapt=str,
|
||||
error_context=context(),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("state", "expected", "expected_events"),
|
||||
(
|
||||
pytest.param("disabled", False, (), id="disabled"),
|
||||
pytest.param("ineligible", False, (), id="ineligible"),
|
||||
pytest.param("unavailable", False, ("load",), id="unavailable"),
|
||||
pytest.param("available", True, ("load",), id="available"),
|
||||
),
|
||||
)
|
||||
def test_can_attempt_only_enabled_available_requests(
|
||||
state: str,
|
||||
expected: bool,
|
||||
expected_events: tuple[str, ...],
|
||||
) -> None:
|
||||
events: list[str] = []
|
||||
|
||||
def load() -> object | None:
|
||||
events.append("load")
|
||||
return None if state == "unavailable" else object()
|
||||
|
||||
def prepare() -> int:
|
||||
events.append("prepare")
|
||||
return 3
|
||||
bridge: Final = runtime.EndpointBinding(route="messages", load=load, enabled=lambda: state != "disabled")
|
||||
|
||||
def call(_binding: object, request: int) -> int:
|
||||
events.append("call")
|
||||
return request * 2
|
||||
|
||||
async def acall(binding: object, request: int) -> int:
|
||||
return call(binding, request)
|
||||
|
||||
def adapt(value: int) -> str:
|
||||
events.append("adapt")
|
||||
return str(value)
|
||||
|
||||
result: Final = (
|
||||
await runtime.aattempt(
|
||||
load=load,
|
||||
enabled=state != "disabled",
|
||||
assert (
|
||||
bridge.can_attempt(
|
||||
eligible=state != "ineligible",
|
||||
prepare=prepare,
|
||||
call=acall,
|
||||
adapt=adapt,
|
||||
)
|
||||
if asynchronous
|
||||
else runtime.attempt(
|
||||
load=load,
|
||||
enabled=state != "disabled",
|
||||
eligible=state != "ineligible",
|
||||
prepare=prepare,
|
||||
call=call,
|
||||
adapt=adapt,
|
||||
)
|
||||
is expected
|
||||
)
|
||||
if state == "handled":
|
||||
assert result == runtime.Handled("6")
|
||||
assert events == ["load", "prepare", "call", "adapt"]
|
||||
else:
|
||||
assert result == runtime.NativeSkipped(runtime.NativeSkipReason(state))
|
||||
assert events == (["load"] if state == "unavailable" else [])
|
||||
assert tuple(events) == expected_events
|
||||
|
||||
|
||||
def test_native_endpoint_applies_partial_overrides_and_reset(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def native_sync() -> str:
|
||||
return "native"
|
||||
|
||||
async def native_async() -> str:
|
||||
return "native async"
|
||||
|
||||
def replacement_sync() -> str:
|
||||
return "replacement"
|
||||
|
||||
monkeypatch.setattr(
|
||||
bindings,
|
||||
"get_native_bridge",
|
||||
lambda: SimpleNamespace(chat_completions=native_sync, achat_completions=native_async),
|
||||
)
|
||||
endpoint: Final[runtime.EndpointDispatch[object, object]] = runtime.EndpointDispatch.native(
|
||||
route="test",
|
||||
sync=lambda native: native.chat_completions,
|
||||
asynchronous=lambda native: native.achat_completions,
|
||||
enabled=enabled,
|
||||
)
|
||||
|
||||
assert endpoint.sync.load() is native_sync
|
||||
assert endpoint.asynchronous.load() is native_async
|
||||
endpoint.override(sync=replacement_sync)
|
||||
assert endpoint.sync.load() is replacement_sync
|
||||
assert endpoint.asynchronous.load() is native_async
|
||||
endpoint.override(asynchronous=None)
|
||||
assert endpoint.sync.load() is replacement_sync
|
||||
assert endpoint.asynchronous.load() is None
|
||||
endpoint.reset()
|
||||
assert endpoint.sync.load() is native_sync
|
||||
assert endpoint.asynchronous.load() is native_async
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("phase", ("prepare", "call"))
|
||||
async def test_attempt_reports_failure_without_deciding_retry(asynchronous: bool, phase: str) -> None:
|
||||
error: Final = RuntimeError("native failure")
|
||||
async def test_response_adaptation_failure_never_authorizes_fallback(asynchronous: bool) -> None:
|
||||
def adapt(value: str) -> str:
|
||||
assert value == "provider response"
|
||||
raise RustBridgeDeclined("adapter failed after provider response")
|
||||
|
||||
def prepare() -> int:
|
||||
if phase == "prepare":
|
||||
raise error
|
||||
return 3
|
||||
async def native(binding: object, request: object) -> str:
|
||||
return "provider response"
|
||||
|
||||
def call(_binding: object, request: int) -> int:
|
||||
raise error
|
||||
async def fallback() -> str:
|
||||
pytest.fail("a received response must not be retried")
|
||||
|
||||
async def acall(binding: object, request: int) -> int:
|
||||
return call(binding, request)
|
||||
bridge = runtime.EndpointBinding(route="messages", load=object, enabled=enabled)
|
||||
|
||||
def adapt(value: int) -> str:
|
||||
pytest.fail("failed attempts cannot be adapted")
|
||||
|
||||
result: Final = (
|
||||
await runtime.aattempt(load=object, enabled=True, eligible=True, prepare=prepare, call=acall, adapt=adapt)
|
||||
if asynchronous
|
||||
else runtime.attempt(load=object, enabled=True, eligible=True, prepare=prepare, call=call, adapt=adapt)
|
||||
)
|
||||
assert isinstance(result, runtime.NativeFailed)
|
||||
assert result.error is error
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
async def test_adaptation_failure_remains_distinct_from_native_failure(asynchronous: bool) -> None:
|
||||
error: Final = ValueError("invalid response")
|
||||
|
||||
async def acall(_binding: object, request: int) -> int:
|
||||
return request
|
||||
|
||||
def adapt(value: int) -> str:
|
||||
raise error
|
||||
|
||||
async def run() -> None:
|
||||
async def invoke() -> None:
|
||||
if asynchronous:
|
||||
await runtime.aattempt(load=object, enabled=True, eligible=True, prepare=lambda: 3, call=acall, adapt=adapt)
|
||||
await bridge.ainvoke(
|
||||
prepare=lambda: None, call=native, fallback=fallback, adapt=adapt, error_context=context()
|
||||
)
|
||||
else:
|
||||
runtime.attempt(
|
||||
load=object,
|
||||
enabled=True,
|
||||
eligible=True,
|
||||
prepare=lambda: 3,
|
||||
call=lambda binding, request: request,
|
||||
bridge.invoke(
|
||||
prepare=lambda: None,
|
||||
call=lambda binding, request: "provider response",
|
||||
fallback=lambda: pytest.fail("a received response must not be retried"),
|
||||
adapt=adapt,
|
||||
error_context=context(),
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="invalid response") as caught:
|
||||
await run()
|
||||
assert caught.value is error
|
||||
with pytest.raises(RustBridgeDeclined, match="adapter failed"):
|
||||
await invoke()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", (False, True))
|
||||
@pytest.mark.parametrize("available, accepted", ((False, False), (True, False), (True, True)))
|
||||
async def test_preflight_runs_after_binding_selection_before_preparation(
|
||||
asynchronous: bool, available: bool, accepted: bool
|
||||
) -> None:
|
||||
events: list[str] = []
|
||||
|
||||
def load() -> object | None:
|
||||
events.append("load")
|
||||
return object() if available else None
|
||||
|
||||
def preflight() -> runtime.PythonFallback | None:
|
||||
events.append("preflight")
|
||||
return None if accepted else runtime.PythonFallback(runtime.PythonFallbackReason.NATIVE_DECLINED)
|
||||
|
||||
def prepare() -> int:
|
||||
events.append("prepare")
|
||||
return 7
|
||||
|
||||
def call(binding: object, request: int) -> int:
|
||||
events.append("native")
|
||||
return request
|
||||
|
||||
async def acall(binding: object, request: int) -> int:
|
||||
return call(binding, request)
|
||||
|
||||
def fallback() -> str:
|
||||
events.append("python")
|
||||
return "3"
|
||||
|
||||
async def afallback() -> str:
|
||||
return fallback()
|
||||
|
||||
endpoint: Final = runtime.EndpointBinding(route="ocr", load=load, enabled=enabled)
|
||||
result: Final = (
|
||||
await endpoint.ainvoke(
|
||||
prepare=prepare, call=acall, fallback=afallback, adapt=str, error_context=context(), preflight=preflight
|
||||
)
|
||||
if asynchronous
|
||||
else endpoint.invoke(
|
||||
prepare=prepare, call=call, fallback=fallback, adapt=str, error_context=context(), preflight=preflight
|
||||
)
|
||||
)
|
||||
assert result == ("7" if available and accepted else "3")
|
||||
assert events == (
|
||||
["load", "preflight", "prepare", "native"]
|
||||
if available and accepted
|
||||
else ["load", "preflight", "python"] if available else ["load", "python"]
|
||||
)
|
||||
|
||||
|
||||
def test_preflight_failure_is_not_a_native_decline() -> None:
|
||||
endpoint: Final = runtime.EndpointBinding(route="ocr", load=object, enabled=enabled)
|
||||
|
||||
def preflight() -> runtime.PythonFallback | None:
|
||||
raise ValueError("invalid acceptance contract")
|
||||
|
||||
with pytest.raises(ValueError, match="invalid acceptance contract"):
|
||||
endpoint.invoke(
|
||||
prepare=lambda: pytest.fail("must not prepare"),
|
||||
call=lambda binding, request: pytest.fail("must not invoke"),
|
||||
fallback=lambda: pytest.fail("must not fall back"),
|
||||
adapt=str,
|
||||
error_context=context(),
|
||||
preflight=preflight,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -4,16 +4,29 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
|
||||
from litellm.rust_bridge.request import NativeRequestContext, NativeRequestOptions, NativeTranscriptionRequest
|
||||
from litellm.rust_bridge.runtime import Handled
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
NativeTranscriptionRequest,
|
||||
)
|
||||
|
||||
rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_rust_transcription() -> None:
|
||||
rust_bridge.configure_rust_transcription(
|
||||
transcription=None,
|
||||
atranscription=None,
|
||||
decline=lambda model, custom_llm_provider, *, context: None,
|
||||
)
|
||||
yield
|
||||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=None, decline=None)
|
||||
|
||||
|
||||
class SyncBridge:
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
self.contexts: list[NativeRequestContext] = []
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
|
|
@ -23,9 +36,13 @@ class SyncBridge:
|
|||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{"model": request.model, "audio": request.audio, "optional_params": request.optional_params}
|
||||
{
|
||||
"model": request.model,
|
||||
"audio": request.audio,
|
||||
"optional_params": request.optional_params,
|
||||
"bedrock": options.bedrock,
|
||||
}
|
||||
)
|
||||
self.contexts.append(context)
|
||||
return {"text": "hello"}
|
||||
|
||||
|
||||
|
|
@ -52,17 +69,9 @@ def test_enabled_sync_bridge_receives_audio() -> None:
|
|||
extra_headers=None,
|
||||
optional_params={"temperature": 0},
|
||||
timeout=5.0,
|
||||
stream=True,
|
||||
has_custom_client=True,
|
||||
input_source_kind="file",
|
||||
)
|
||||
assert isinstance(result, Handled)
|
||||
assert result.value == {"text": "hello"}
|
||||
assert result == {"text": "hello"}
|
||||
assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"}
|
||||
assert bridge.contexts[0].capabilities.execution_mode == "sync"
|
||||
assert bridge.contexts[0].capabilities.stream is True
|
||||
assert bridge.contexts[0].capabilities.has_custom_client is True
|
||||
assert bridge.contexts[0].capabilities.input_source_kind == "file"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -78,7 +87,7 @@ async def test_enabled_async_bridge() -> None:
|
|||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
assert result == Handled({"text": "async"})
|
||||
assert result == {"text": "async"}
|
||||
|
||||
|
||||
def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
|
|
@ -89,8 +98,7 @@ def test_loader_returns_none_without_native_extension(monkeypatch: pytest.Monkey
|
|||
|
||||
|
||||
def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
rust_bridge.configure_rust_transcription(transcription=None)
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
|
||||
monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None)
|
||||
|
||||
with pytest.raises(RuntimeError, match="bridge is unavailable"):
|
||||
BedrockAudioTranscriptionRustDispatch().audio_transcriptions(
|
||||
|
|
@ -107,8 +115,10 @@ def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) ->
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
rust_bridge.configure_rust_transcription(atranscription=None)
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
|
||||
async def unavailable(**_: object) -> None:
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(rust_bridge, "atranscription", unavailable)
|
||||
|
||||
with pytest.raises(RuntimeError, match="bridge is unavailable"):
|
||||
await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions(
|
||||
|
|
@ -125,7 +135,7 @@ async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPat
|
|||
|
||||
def test_bedrock_transcription_uses_rust_only_path() -> None:
|
||||
rust_bridge.configure_rust_transcription(
|
||||
transcription=lambda *_args, **_: {"text": "rust"},
|
||||
transcription=lambda request, *, options, context: {"text": "rust"},
|
||||
atranscription=None,
|
||||
)
|
||||
try:
|
||||
|
|
@ -141,7 +151,9 @@ def test_bedrock_transcription_uses_rust_only_path() -> None:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_atranscription_uses_rust_only_path() -> None:
|
||||
async def rust_response(*_args: object, **_: object) -> dict[str, object]:
|
||||
async def rust_response(
|
||||
request: NativeTranscriptionRequest, *, options: object, context: NativeRequestContext
|
||||
) -> dict[str, object]:
|
||||
return {"text": "rust"}
|
||||
|
||||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue