refactor(native): separate request data from execution context

This commit is contained in:
Yujong Lee 2026-09-05 12:59:50 -07:00 committed by yujonglee
parent 7ed22070e8
commit f8d6ac7a76
62 changed files with 2530 additions and 1984 deletions

View file

@ -7,12 +7,17 @@ that makes the LLM call and hands back a typed response, the same shape as
`litellm.messages()` in Python.
```rust
let response = litellm_core::messages::messages(MessagesRequest {
model: "claude-sonnet-4-5",
body,
api_key: Some(key),
..
})
let response = litellm_core::messages::messages(
MessagesRequest {
model: "claude-sonnet-4-5",
body,
options: RequestOptions {
api_key: Some(key.to_string()),
..Default::default()
},
},
&LiteLlmRequestContext::default(),
)
.await?;
```
@ -20,6 +25,21 @@ Python continues to own configuration, retries, routing policy, logging,
callbacks, spend tracking, and customer plugins until each Rust path has parity
coverage and production evidence.
## Native request boundary
Native HTTP routes and Responses WebSocket connections accept `native(request, *, context)`
The request carries the endpoint payload and `NativeRequestOptions`: credentials,
provider routing, headers, query parameters, and timeout. `NativeRequestContext`
carries LiteLLM metadata, call identity, and attribution separately from the provider payload
Python builds the frozen request dataclasses in `litellm/rust_bridge/request.py` and
PyO3 extracts their fields before execution. Provider connection parameters, such as
AWS credentials and Vertex project/location, belong in `options.provider_connection`
rather than the request body
This boundary preserves existing Python provider preparation, preflight decisions,
fallback, and callbacks
## Crates
| Crate | Role |

View file

@ -4,6 +4,8 @@ use litellm_core::audio_transcription::{
};
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use litellm_core::error::Error;
use litellm_core::request_context::RequestAttribution;
use litellm_core::request_options::RequestOptions;
use serde_json::{Map, Value, json};
use std::future::Future;
use std::pin::Pin;
@ -15,14 +17,12 @@ use crate::integrations::custom_guardrail::{
use crate::integrations::custom_logger::{
CallType, CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
use crate::integrations::types::{
RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload,
};
use crate::integrations::types::{StandardLoggingMetadata, StandardLoggingPayload};
pub(crate) struct AudioTranscriptionLifecycleHooks {
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestMetadata,
request_metadata: RequestAttribution,
}
type AudioFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
@ -32,7 +32,7 @@ impl AudioTranscriptionLifecycleHooks {
pub(crate) fn new(
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestMetadata,
request_metadata: RequestAttribution,
) -> Self {
Self {
logger_runner,
@ -94,6 +94,7 @@ impl AudioTranscriptionLifecycleHooks {
custom_llm_provider,
audio,
api_key,
provider_connection,
api_base,
extra_headers,
optional_params,
@ -104,12 +105,17 @@ impl AudioTranscriptionLifecycleHooks {
prepare_audio_transcription_provider_call(CoreAudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: Some(&custom_llm_provider),
extra_headers,
optional_params,
timeout,
options: RequestOptions {
provider_connection,
api_key: (api_key.as_deref()).map(|value| value.to_string()),
api_base: (api_base.as_deref()).map(|value| value.to_string()),
custom_llm_provider: (Some(&custom_llm_provider))
.map(|value| value.to_string()),
extra_headers,
timeout,
..Default::default()
},
})?;
self.run_during_call_guardrails(provider_request).await
}
@ -251,7 +257,7 @@ impl CallLifecycleHooks<PreparedAudioTranscriptionRequest, ProviderAudioTranscri
}
}
fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext {
fn guardrail_context(metadata: &RequestAttribution) -> GuardrailContext {
GuardrailContext {
call_type: CallType::Other("audio_transcription".to_string()),
selected_guardrails: Vec::new(),

View file

@ -1,6 +1,8 @@
use crate::integrations::types::RequestHooks;
use litellm_core::Error;
use litellm_core::audio_transcription::execute_audio_transcription_provider_call;
use litellm_core::call_lifecycle::CallLifecycle;
use litellm_core::request_context::LiteLlmRequestContext;
use serde_json::Value;
mod hooks;
@ -11,11 +13,23 @@ pub use types::AudioTranscriptionRequest;
use prepare::{PreparedAudioTranscriptionCall, prepare_audio_transcription_call};
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
let PreparedAudioTranscriptionCall { request, hooks } =
prepare_audio_transcription_call(request);
pub async fn audio_transcription(
request: AudioTranscriptionRequest<'_>,
context: &LiteLlmRequestContext,
hooks: RequestHooks,
) -> Result<Value, Error> {
let PreparedAudioTranscriptionCall {
request,
context,
hooks,
} = prepare_audio_transcription_call(request, context, hooks);
CallLifecycle::default()
.run_request(request, &hooks, execute_audio_transcription_provider_call)
.run(
context,
request,
&hooks,
execute_audio_transcription_provider_call,
)
.await
}

View file

@ -1,3 +1,6 @@
use crate::integrations::types::RequestHooks;
use litellm_core::call_lifecycle::CallLifecycleContext;
use litellm_core::request_context::LiteLlmRequestContext;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
@ -9,38 +12,50 @@ use crate::integrations::custom_guardrail::CustomGuardrailRunner;
use crate::integrations::custom_logger::CustomLoggerRunner;
pub(crate) struct PreparedAudioTranscriptionCall {
pub(crate) context: CallLifecycleContext,
pub(crate) request: PreparedAudioTranscriptionRequest,
pub(crate) hooks: AudioTranscriptionLifecycleHooks,
}
pub(crate) fn prepare_audio_transcription_call(
request: AudioTranscriptionRequest<'_>,
context: &LiteLlmRequestContext,
hooks: RequestHooks,
) -> PreparedAudioTranscriptionCall {
let call_id = request
let call_id = context
.litellm_call_id
.map(str::to_string)
.clone()
.unwrap_or_else(new_audio_transcription_call_id);
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.unwrap_or(CustomLlmProvider {
model: request.model,
custom_llm_provider: "bedrock",
});
let provider_info = get_custom_llm_provider(
request.model,
request.options.custom_llm_provider.as_deref(),
)
.unwrap_or(CustomLlmProvider {
model: request.model,
custom_llm_provider: "bedrock",
});
PreparedAudioTranscriptionCall {
context: CallLifecycleContext::new(
"audio_transcription",
provider_info.model,
provider_info.custom_llm_provider,
call_id,
),
request: PreparedAudioTranscriptionRequest {
model: provider_info.model.to_string(),
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
litellm_call_id: call_id,
audio: request.audio,
api_key: request.api_key.map(str::to_string),
api_base: request.api_base.map(str::to_string),
extra_headers: request.extra_headers,
provider_connection: request.options.provider_connection,
api_key: request.options.api_key,
api_base: request.options.api_base,
extra_headers: request.options.extra_headers,
optional_params: request.optional_params,
timeout: request.timeout,
timeout: request.options.timeout,
},
hooks: AudioTranscriptionLifecycleHooks::new(
CustomLoggerRunner::new(request.callbacks),
CustomGuardrailRunner::new(request.guardrails),
request.request_metadata,
CustomLoggerRunner::new(hooks.callbacks),
CustomGuardrailRunner::new(hooks.guardrails),
context.attribution.clone(),
),
}
}

View file

@ -1,3 +1,6 @@
use crate::integrations::types::RequestHooks;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use std::io::{Read, Write};
use std::net::TcpListener;
use std::thread;
@ -26,26 +29,38 @@ async fn bedrock_request_is_signed_and_contains_audio() {
stream.write_all(response).expect("response");
});
let optional_params = Map::from_iter([
let provider_connection = Map::from_iter([
("aws_access_key_id".to_string(), json!("access-key")),
("aws_secret_access_key".to_string(), json!("secret-key")),
("aws_region_name".to_string(), json!("us-east-1")),
]);
let api_base = format!("http://{address}");
let response = audio_transcription(AudioTranscriptionRequest {
model: "mistral.voxtral-mini-3b-2507",
audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}),
api_key: None,
api_base: Some(&api_base),
custom_llm_provider: Some("bedrock"),
extra_headers: None,
optional_params,
timeout: None,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
})
let response = audio_transcription(
AudioTranscriptionRequest {
model: "mistral.voxtral-mini-3b-2507",
audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}),
optional_params: Map::new(),
options: RequestOptions {
provider_connection,
api_key: None,
api_base: (Some(&api_base)).map(|value| value.to_string()),
custom_llm_provider: (Some("bedrock")).map(|value| value.to_string()),
extra_headers: None,
timeout: None,
..Default::default()
},
},
&LiteLlmRequestContext {
attribution: Default::default(),
litellm_call_id: None,
..Default::default()
},
RequestHooks {
callbacks: Vec::new(),
guardrails: Vec::new(),
},
)
.await
.expect("transcription");
assert_eq!(response, json!({"text": "hello"}));

View file

@ -1,47 +1,23 @@
use std::sync::Arc;
use litellm_core::request_options::RequestOptions;
use std::time::Duration;
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest};
use serde_json::{Map, Value};
use crate::integrations::custom_guardrail::CustomGuardrail;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
pub struct AudioTranscriptionRequest<'a> {
pub model: &'a str,
pub audio: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
pub callbacks: Vec<Arc<dyn CustomLogger>>,
pub guardrails: Vec<Arc<dyn CustomGuardrail>>,
pub request_metadata: RequestMetadata,
pub litellm_call_id: Option<&'a str>,
pub options: RequestOptions,
}
pub(crate) struct PreparedAudioTranscriptionRequest {
pub(crate) model: String,
pub(crate) custom_llm_provider: String,
pub(crate) litellm_call_id: String,
pub(crate) audio: Value,
pub(crate) provider_connection: Map<String, Value>,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
pub(crate) optional_params: Map<String, Value>,
pub(crate) timeout: Option<Duration>,
}
impl CallLifecycleRequest for PreparedAudioTranscriptionRequest {
fn lifecycle_context(&self) -> CallLifecycleContext {
CallLifecycleContext::new(
"audio_transcription",
self.model.clone(),
self.custom_llm_provider.clone(),
self.litellm_call_id.clone(),
)
}
}

View file

@ -20,12 +20,10 @@ pub struct Usage {
pub total_tokens: u64,
}
/// Cost-attribution metadata threaded from the authenticated request.
#[derive(Clone, Debug, Default)]
pub struct RequestMetadata {
pub user_api_key_hash: Option<String>,
pub user_api_key_user_id: Option<String>,
pub user_api_key_team_id: Option<String>,
#[derive(Default)]
pub struct RequestHooks {
pub callbacks: Vec<std::sync::Arc<dyn super::custom_logger::CustomLogger>>,
pub guardrails: Vec<std::sync::Arc<dyn super::custom_guardrail::CustomGuardrail>>,
}
/// The self-describing payload. Field names are the EXACT JSON keys the Python

View file

@ -1,12 +1,13 @@
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use litellm_core::Error;
use litellm_core::http_utils::string_headers;
use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG;
use litellm_core::responses::types::ResponsesWsEvent;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::responses::types::{ResponsesWebSocketRequest, ResponsesWsEvent};
use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig;
use tokio::net::TcpStream;
use tokio::sync::Mutex;
@ -33,24 +34,26 @@ pub struct ResponsesWebSocketConnection {
}
impl ResponsesWebSocketConnection {
pub async fn connect_url(
url: &str,
headers: &HashMap<String, String>,
timeout: Option<Duration>,
pub async fn connect(
input: ResponsesWebSocketRequest,
_context: &LiteLlmRequestContext,
) -> Result<Self, Error> {
let mut request = url
let headers = string_headers("Responses WebSocket", input.options.extra_headers)?;
let mut request = input
.url
.as_str()
.into_client_request()
.map_err(|error| Error::Network(error.to_string()))?;
for (name, value) in headers {
let header_name = name
.parse::<HeaderName>()
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
let header_value = HeaderValue::from_str(value)
let header_value = HeaderValue::from_str(&value)
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
request.headers_mut().insert(header_name, header_value);
}
let connect = connect_async(request);
let result = match timeout {
let result = match input.options.timeout {
Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| {
Error::Network("Responses WebSocket connection timed out".to_string())
})?,

View file

@ -3,6 +3,7 @@ use litellm_core::error::Error;
use litellm_core::providers::reducto::ocr::transformation::{
build_upload_request, extract_document_source, extract_upload_file_id,
};
use litellm_core::request_context::RequestAttribution;
use serde_json::{Map, Value, json};
use std::future::Future;
use std::pin::Pin;
@ -16,14 +17,12 @@ use crate::integrations::custom_guardrail::{
use crate::integrations::custom_logger::{
CallType, CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
use crate::integrations::types::{
RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload,
};
use crate::integrations::types::{StandardLoggingMetadata, StandardLoggingPayload};
pub(crate) struct OcrLifecycleHooks {
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestMetadata,
request_metadata: RequestAttribution,
}
type OcrFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
@ -33,7 +32,7 @@ impl OcrLifecycleHooks {
pub(crate) fn new(
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestMetadata,
request_metadata: RequestAttribution,
) -> Self {
Self {
logger_runner,
@ -85,10 +84,16 @@ impl OcrLifecycleHooks {
request.api_key.as_deref(),
&env_lookup,
)?;
let url_params = request
.optional_params
.clone()
.into_iter()
.chain(request.provider_connection)
.collect();
let url = config.complete_url(
request.api_base.as_deref(),
&request.model,
&request.optional_params,
&url_params,
&env_lookup,
)?;
let model = request.model.clone();
@ -323,7 +328,7 @@ impl CallLifecycleHooks<PreparedOcrRequest, PreparedOcrRequest, Value> for OcrLi
}
}
fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext {
fn guardrail_context(metadata: &RequestAttribution) -> GuardrailContext {
GuardrailContext {
call_type: CallType::Ocr,
selected_guardrails: Vec::new(),

View file

@ -1,5 +1,7 @@
use crate::integrations::types::RequestHooks;
use litellm_core::Error;
use litellm_core::call_lifecycle::CallLifecycle;
use litellm_core::request_context::LiteLlmRequestContext;
use serde_json::Value;
mod common_utils;
@ -14,10 +16,18 @@ use handler::execute_ocr_provider_call;
use prepare::{PreparedOcrCall, prepare_ocr_call};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn ocr(request: OcrRequest<'_>) -> Result<Value, Error> {
let PreparedOcrCall { request, hooks } = prepare_ocr_call(request);
pub async fn ocr(
request: OcrRequest<'_>,
context: &LiteLlmRequestContext,
hooks: RequestHooks,
) -> Result<Value, Error> {
let PreparedOcrCall {
request,
context,
hooks,
} = prepare_ocr_call(request, context, hooks);
CallLifecycle::default()
.run_request(request, &hooks, |request| {
.run(context, request, &hooks, |request| {
execute_ocr_provider_call(request, &hooks)
})
.await
@ -25,12 +35,14 @@ pub async fn ocr(request: OcrRequest<'_>) -> Result<Value, Error> {
#[cfg(test)]
mod tests {
use crate::integrations::types::RequestHooks;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use serde_json::{Map, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use super::{OcrRequest, ocr};
use crate::integrations::types::RequestMetadata;
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
@ -72,16 +84,16 @@ mod tests {
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: None,
extra_headers: None,
optional_params: Map::new(),
timeout: None,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: RequestMetadata::default(),
litellm_call_id: None,
options: RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: None,
custom_llm_provider: None,
extra_headers: None,
timeout: None,
..Default::default()
},
}
}
@ -121,9 +133,9 @@ mod tests {
});
let api_base = format!("http://{address}");
let mut request = base_ocr_request("reducto/parse-v3");
request.api_base = Some(&api_base);
request.api_key = None;
request.extra_headers = Some(Map::from_iter([
request.options.api_base = Some(&api_base).map(|value| value.to_string());
request.options.api_key = None;
request.options.extra_headers = Some(Map::from_iter([
("Authorization".to_string(), json!("Bearer test-key")),
("x-trace-id".to_string(), json!("trace-1")),
]));
@ -140,7 +152,17 @@ mod tests {
("settings".to_string(), json!({"ocr_system": "standard"})),
]);
let response = ocr(request).await.expect("Reducto OCR succeeds");
let response = ocr(
request,
&LiteLlmRequestContext {
..Default::default()
},
RequestHooks {
..Default::default()
},
)
.await
.expect("Reducto OCR succeeds");
assert_eq!(response["pages"].as_array().map(Vec::len), Some(3));
assert_eq!(

View file

@ -1,3 +1,6 @@
use crate::integrations::types::RequestHooks;
use litellm_core::call_lifecycle::CallLifecycleContext;
use litellm_core::request_context::LiteLlmRequestContext;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
@ -11,21 +14,29 @@ use crate::integrations::custom_guardrail::CustomGuardrailRunner;
use crate::integrations::custom_logger::CustomLoggerRunner;
pub(crate) struct PreparedOcrCall {
pub(crate) context: CallLifecycleContext,
pub(crate) request: PreparedOcrRequest,
pub(crate) hooks: OcrLifecycleHooks,
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall {
let call_id = request
pub(crate) fn prepare_ocr_call(
request: OcrRequest<'_>,
context: &LiteLlmRequestContext,
hooks: RequestHooks,
) -> PreparedOcrCall {
let call_id = context
.litellm_call_id
.map(str::to_string)
.clone()
.unwrap_or_else(new_ocr_call_id);
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.unwrap_or(CustomLlmProvider {
model: request.model,
custom_llm_provider: "mistral",
});
let provider_info = get_custom_llm_provider(
request.model,
request.options.custom_llm_provider.as_deref(),
)
.unwrap_or(CustomLlmProvider {
model: request.model,
custom_llm_provider: "mistral",
});
let model = provider_info.model.to_string();
let custom_llm_provider = provider_info.custom_llm_provider.to_string();
let config = ocr_provider_config(&custom_llm_provider, &model)
@ -37,46 +48,41 @@ pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall {
let optional_params = match &config {
Ok(config) => {
let supported = config.supported_ocr_params();
let mut mapped = config.map_ocr_params(
config.map_ocr_params(
&request
.optional_params
.iter()
.filter(|(name, _)| supported.contains(&name.as_str()))
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
);
for name in [
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
] {
if let Some(value) = request.optional_params.get(name) {
mapped.insert(name.to_string(), value.clone());
}
}
mapped
)
}
Err(_) => request.optional_params,
};
PreparedOcrCall {
context: CallLifecycleContext::new(
"ocr",
model.clone(),
custom_llm_provider.clone(),
call_id,
),
request: PreparedOcrRequest {
config,
model,
custom_llm_provider,
litellm_call_id: call_id,
document: request.document,
api_key: request.api_key.map(str::to_string),
api_base: request.api_base.map(str::to_string),
extra_headers: request.extra_headers,
provider_connection: request.options.provider_connection,
api_key: request.options.api_key,
api_base: request.options.api_base,
extra_headers: request.options.extra_headers,
optional_params,
timeout: request.timeout,
timeout: request.options.timeout,
},
hooks: OcrLifecycleHooks::new(
CustomLoggerRunner::new(request.callbacks),
CustomGuardrailRunner::new(request.guardrails),
request.request_metadata,
CustomLoggerRunner::new(hooks.callbacks),
CustomGuardrailRunner::new(hooks.guardrails),
context.attribution.clone(),
),
}
}
@ -113,11 +119,13 @@ fn new_ocr_call_id() -> String {
#[cfg(test)]
mod tests {
use crate::integrations::types::RequestHooks;
use litellm_core::error::Error;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use serde_json::{Map, json};
use super::{OcrRequest, prepare_ocr_call};
use crate::integrations::types::RequestMetadata;
fn base_ocr_request(model: &str) -> OcrRequest<'_> {
OcrRequest {
@ -126,16 +134,16 @@ mod tests {
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: None,
extra_headers: None,
optional_params: Map::new(),
timeout: None,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: RequestMetadata::default(),
litellm_call_id: None,
options: RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: None,
custom_llm_provider: None,
extra_headers: None,
timeout: None,
..Default::default()
},
}
}
@ -147,7 +155,15 @@ mod tests {
#[test]
fn native_format_rejected_for_provider_without_support_as_bad_request() {
let prepared = prepare_ocr_call(request_with_format("native"));
let prepared = prepare_ocr_call(
request_with_format("native"),
&LiteLlmRequestContext {
..Default::default()
},
RequestHooks {
..Default::default()
},
);
assert!(
matches!(prepared.request.config, Err(Error::InvalidRequest(message)) if message.contains("not supported for provider"))
);
@ -155,7 +171,15 @@ mod tests {
#[test]
fn unknown_format_rejected_for_provider_without_support_as_bad_request() {
let prepared = prepare_ocr_call(request_with_format("raw"));
let prepared = prepare_ocr_call(
request_with_format("raw"),
&LiteLlmRequestContext {
..Default::default()
},
RequestHooks {
..Default::default()
},
);
assert!(
matches!(prepared.request.config, Err(Error::InvalidRequest(message)) if message.contains("Invalid `req_format`"))
);

View file

@ -1,35 +1,22 @@
use std::sync::Arc;
use litellm_core::request_options::RequestOptions;
use std::time::Duration;
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest};
use litellm_core::ocr::transformation::OcrProviderConfig;
use serde_json::{Map, Value};
use crate::integrations::custom_guardrail::CustomGuardrail;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
pub struct OcrRequest<'a> {
pub model: &'a str,
pub document: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
pub callbacks: Vec<Arc<dyn CustomLogger>>,
pub guardrails: Vec<Arc<dyn CustomGuardrail>>,
pub request_metadata: RequestMetadata,
pub litellm_call_id: Option<&'a str>,
pub options: RequestOptions,
}
pub(crate) struct PreparedOcrRequest {
pub(crate) config: Result<&'static dyn OcrProviderConfig, litellm_core::Error>,
pub(crate) model: String,
pub(crate) custom_llm_provider: String,
pub(crate) litellm_call_id: String,
pub(crate) document: Value,
pub(crate) provider_connection: Map<String, Value>,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
@ -37,17 +24,6 @@ pub(crate) struct PreparedOcrRequest {
pub(crate) timeout: Option<Duration>,
}
impl CallLifecycleRequest for PreparedOcrRequest {
fn lifecycle_context(&self) -> CallLifecycleContext {
CallLifecycleContext::new(
"ocr",
self.model.clone(),
self.custom_llm_provider.clone(),
self.litellm_call_id.clone(),
)
}
}
pub(crate) struct ProviderOcrRequest {
pub(crate) model: String,
pub(crate) config: &'static dyn OcrProviderConfig,

View file

@ -6,6 +6,7 @@
//! builds a `StandardLoggingPayload` and fans it out to every registered
//! `CustomLogger`.
use litellm_core::request_context::RequestAttribution;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
@ -16,9 +17,7 @@ use crate::constants::DEFAULT_PROVIDER;
use crate::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
use crate::integrations::types::{
RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload, Usage,
};
use crate::integrations::types::{StandardLoggingMetadata, StandardLoggingPayload, Usage};
/// Current wall-clock time as epoch seconds (float), matching the Python
/// `startTime`/`endTime` contract.
@ -54,7 +53,7 @@ pub struct RealTimeStreaming {
response_cost: f64,
start_time: f64,
end_time: f64,
metadata: RequestMetadata,
metadata: RequestAttribution,
/// Count of logging callbacks that failed to enqueue (non-fatal).
dropped: u64,
}
@ -67,7 +66,7 @@ impl RealTimeStreaming {
callbacks: Vec<Arc<dyn CustomLogger>>,
litellm_call_id: String,
model: String,
metadata: RequestMetadata,
metadata: RequestAttribution,
) -> Self {
let now = epoch_seconds();
Self {
@ -277,7 +276,7 @@ mod tests {
callbacks,
"call_abc".to_string(),
"gpt-realtime".to_string(),
RequestMetadata {
RequestAttribution {
user_api_key_hash: Some("hash123".to_string()),
user_api_key_user_id: Some("user-1".to_string()),
user_api_key_team_id: Some("team-1".to_string()),
@ -329,7 +328,7 @@ mod tests {
Vec::new(),
"call_fallback".to_string(),
"gpt-realtime".to_string(),
RequestMetadata::default(),
RequestAttribution::default(),
);
streaming.observe(&event(
@ -355,7 +354,7 @@ mod tests {
Vec::new(),
"call_xyz".to_string(),
"gpt-realtime".to_string(),
RequestMetadata::default(),
RequestAttribution::default(),
);
streaming.observe(&event(
r#"{"type":"response.done","response":{"usage":{"input_tokens":1,"output_tokens":1,"total_tokens":2}}}"#,
@ -406,7 +405,7 @@ mod tests {
callbacks,
"call_1".to_string(),
"gpt-realtime".to_string(),
RequestMetadata::default(),
RequestAttribution::default(),
);
streaming.log_messages(SessionStatus::Success).await;
assert_eq!(streaming.dropped(), 1);

View file

@ -1,3 +1,5 @@
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use std::sync::Arc;
use litellm_core::Error;
@ -52,17 +54,34 @@ pub async fn run(
let request = MessagesRequest {
model: provider_model,
body,
api_key: deployment.litellm_params.api_key.as_deref(),
api_base: deployment.litellm_params.api_base.as_deref(),
custom_llm_provider,
extra_headers,
timeout: None,
options: RequestOptions {
api_key: (deployment.litellm_params.api_key.as_deref()).map(|value| value.to_string()),
api_base: (deployment.litellm_params.api_base.as_deref())
.map(|value| value.to_string()),
custom_llm_provider: (custom_llm_provider).map(|value| value.to_string()),
extra_headers,
timeout: None,
..Default::default()
},
};
if request.body.get("stream").and_then(Value::as_bool) == Some(true) {
return messages_stream(request).await.map(MessagesResponse::Stream);
return messages_stream(
request,
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.map(MessagesResponse::Stream);
}
let response = messages(request).await?;
let response = messages(
request,
&LiteLlmRequestContext {
..Default::default()
},
)
.await?;
serde_json::to_value(response)
.map(MessagesResponse::Json)
.map_err(|err| {

View file

@ -4,6 +4,7 @@
//! socket↔events adapter. The pure logic (no axum) lives in [`service`]. Auth is
//! the `RequireMasterKey` extractor, so the handler stays thin.
use litellm_core::request_context::RequestAttribution;
mod service;
use std::sync::Arc;
@ -24,7 +25,7 @@ use serde::Deserialize;
use crate::auth::RequireMasterKey;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
use crate::realtime::streaming::{RealTimeStreaming, SessionStatus};
use crate::state::AppState;
@ -110,9 +111,9 @@ async fn bridge(
// to spend logs and every callback integration; the SHA-256 (matching the
// proxy's hash_token) keeps the plaintext master key out of all of them while
// still matching the key's hash in LiteLLM_SpendLogs.
let metadata = RequestMetadata {
let metadata = RequestAttribution {
user_api_key_hash: master_key.as_deref().map(crate::auth::hash_token),
..RequestMetadata::default()
..RequestAttribution::default()
};
// Owned by THIS task only. The splice observes it via a synchronous `&mut`

View file

@ -1,3 +1,4 @@
use litellm_core::request_context::RequestAttribution;
mod service;
use std::sync::Arc;
@ -17,7 +18,7 @@ use serde::Deserialize;
use crate::auth::RequireMasterKey;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
use crate::state::AppState;
static CALL_SEQ: AtomicU64 = AtomicU64::new(0);
@ -206,9 +207,9 @@ async fn bridge(
}
let call_id = new_call_id();
let metadata = RequestMetadata {
let metadata = RequestAttribution {
user_api_key_hash: master_key.as_deref().map(crate::auth::hash_token),
..RequestMetadata::default()
..RequestAttribution::default()
};
let client_in = Box::pin(stream.filter_map(|message| async move {
match message {

View file

@ -1,3 +1,4 @@
use litellm_core::request_context::RequestAttribution;
use std::sync::Arc;
use std::time::Duration;
@ -13,7 +14,6 @@ use litellm_core::responses::types::ResponsesWsEvent;
use crate::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
use crate::integrations::types::RequestMetadata;
#[allow(clippy::too_many_arguments)]
pub async fn run<In, Out>(
@ -23,7 +23,7 @@ pub async fn run<In, Out>(
idle_timeout: Option<Duration>,
loggers: Arc<Vec<Arc<dyn CustomLogger>>>,
call_id: String,
metadata: RequestMetadata,
metadata: RequestAttribution,
client_in: In,
client_out: Out,
) -> Result<(), Error>

View file

@ -1,3 +1,7 @@
use litellm_ai_gateway::integrations::types::RequestHooks;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_context::RequestAttribution;
use litellm_core::request_options::RequestOptions;
use std::sync::{Arc, Mutex};
use std::time::Duration;
@ -8,7 +12,7 @@ use litellm_ai_gateway::integrations::custom_guardrail::{
use litellm_ai_gateway::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails,
};
use litellm_ai_gateway::integrations::types::RequestMetadata;
use litellm_ai_gateway::ocr::{OcrRequest, ocr};
use litellm_core::error::Error;
use serde_json::{Map, Value, json};
@ -219,16 +223,16 @@ fn base_ocr_request(model: &str) -> OcrRequest<'_> {
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: None,
extra_headers: None,
optional_params: Map::new(),
timeout: None,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: RequestMetadata::default(),
litellm_call_id: None,
options: RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: None,
custom_llm_provider: None,
extra_headers: None,
timeout: None,
..Default::default()
},
}
}
@ -241,14 +245,25 @@ async fn reducto_during_call_guardrail_blocks_before_upload() {
let api_base = format!("http://{address}");
let guardrail = Arc::new(RecordingOcrGuardrail::blocking_during_call());
let mut request = base_ocr_request("reducto/parse-v3");
request.api_base = Some(&api_base);
request.options.api_base = Some(&api_base).map(|value| value.to_string());
request.document = json!({
"type": "document_url",
"document_url": "data:application/pdf;base64,JVBERi0xLjQ="
});
request.guardrails = vec![guardrail.clone()];
let hooks = RequestHooks {
guardrails: vec![guardrail.clone()],
..Default::default()
};
let error = ocr(request).await.expect_err("guardrail blocks upload");
let error = ocr(
request,
&LiteLlmRequestContext {
..Default::default()
},
hooks,
)
.await
.expect_err("guardrail blocks upload");
assert!(matches!(error, Error::InvalidRequest(_)));
assert_eq!(guardrail.events(), vec!["async_moderation_hook"]);
@ -278,13 +293,21 @@ async fn reducto_upload_error_body_is_truncated() {
});
let api_base = format!("http://{address}");
let mut request = base_ocr_request("reducto/parse-v3");
request.api_base = Some(&api_base);
request.options.api_base = Some(&api_base).map(|value| value.to_string());
request.document = json!({
"type": "document_url",
"document_url": "data:application/pdf;base64,JVBERi0xLjQ="
});
let error = ocr(request).await.expect_err("upload should fail");
let error = ocr(
request,
&LiteLlmRequestContext {
..Default::default()
},
RequestHooks::default(),
)
.await
.expect_err("upload should fail");
assert!(
matches!(error, Error::Http { status: 500, body } if body.chars().count() < 300 && body.ends_with("... (truncated)"))
@ -320,26 +343,37 @@ async fn ocr_lifecycle_runs_pre_during_and_success_hooks() {
GuardrailEventHook::PreCall,
GuardrailEventHook::DuringCall,
]));
let response = ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("mistral"),
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
callbacks: vec![logger.clone()],
guardrails: vec![guardrail.clone()],
request_metadata: RequestMetadata {
user_api_key_user_id: Some("user-1".to_string()),
let response = ocr(
OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
options: RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("mistral")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
attribution: RequestAttribution {
user_api_key_user_id: Some("user-1".to_string()),
..Default::default()
},
litellm_call_id: (Some("ocr-call-1")).map(|value| value.to_string()),
..Default::default()
},
litellm_call_id: Some("ocr-call-1"),
})
RequestHooks {
callbacks: vec![logger.clone()],
guardrails: vec![guardrail.clone()],
},
)
.await
.expect("ocr request succeeds");
@ -388,23 +422,34 @@ async fn ocr_lifecycle_runs_failure_hook_on_provider_error() {
});
let logger = Arc::new(RecordingOcrLogger::default());
let err = ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("mistral"),
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
callbacks: vec![logger.clone()],
guardrails: Vec::new(),
request_metadata: RequestMetadata::default(),
litellm_call_id: Some("ocr-call-2"),
})
let err = ocr(
OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
options: RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("mistral")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
attribution: RequestAttribution::default(),
litellm_call_id: (Some("ocr-call-2")).map(|value| value.to_string()),
..Default::default()
},
RequestHooks {
callbacks: vec![logger.clone()],
guardrails: Vec::new(),
},
)
.await
.expect_err("provider error propagates");
@ -432,23 +477,34 @@ async fn ocr_lifecycle_pre_call_block_skips_provider_socket() {
let logger = Arc::new(RecordingOcrLogger::default());
let guardrail = Arc::new(RecordingOcrGuardrail::blocking_pre_call());
let err = ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-test"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("mistral"),
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_millis(100)),
callbacks: vec![logger.clone()],
guardrails: vec![guardrail.clone()],
request_metadata: RequestMetadata::default(),
litellm_call_id: Some("ocr-call-3"),
})
let err = ocr(
OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
options: RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("mistral")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_millis(100)),
..Default::default()
},
},
&LiteLlmRequestContext {
attribution: RequestAttribution::default(),
litellm_call_id: (Some("ocr-call-3")).map(|value| value.to_string()),
..Default::default()
},
RequestHooks {
callbacks: vec![logger.clone()],
guardrails: vec![guardrail.clone()],
},
)
.await
.expect_err("guardrail blocks request");
@ -502,23 +558,34 @@ async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() {
Value::String("trace-1".to_string()),
);
let response = ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-for-rust-fallback"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("mistral"),
extra_headers: Some(headers),
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: RequestMetadata::default(),
litellm_call_id: None,
})
let response = ocr(
OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
options: RequestOptions {
api_key: (Some("sk-for-rust-fallback")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("mistral")).map(|value| value.to_string()),
extra_headers: Some(headers),
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
attribution: RequestAttribution::default(),
litellm_call_id: None,
..Default::default()
},
RequestHooks {
callbacks: Vec::new(),
guardrails: Vec::new(),
},
)
.await
.expect("ocr request succeeds");
@ -571,23 +638,34 @@ async fn document_intelligence_poll_uses_resolved_subscription_key() {
(post_request, poll_request)
});
let response = ocr(OcrRequest {
model: "doc-intelligence/prebuilt-read",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("di-key"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: RequestMetadata::default(),
litellm_call_id: None,
})
let response = ocr(
OcrRequest {
model: "doc-intelligence/prebuilt-read",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
options: RequestOptions {
api_key: (Some("di-key")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("azure_ai")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
attribution: RequestAttribution::default(),
litellm_call_id: None,
..Default::default()
},
RequestHooks {
callbacks: Vec::new(),
guardrails: Vec::new(),
},
)
.await
.expect("document intelligence request succeeds");

View file

@ -59,7 +59,7 @@ async fn signed_headers(
};
let env_lookup = |key: &str| std::env::var(key).ok();
let credentials = resolve_credentials(
aws_auth_config(&request.optional_params, &env_lookup),
aws_auth_config(&request.provider_connection, &env_lookup),
&env_lookup,
)
.await?;

View file

@ -1,4 +1,5 @@
use crate::Error;
use crate::request_context::LiteLlmRequestContext;
mod client;
mod handler;
mod prepare;
@ -12,7 +13,10 @@ pub use prepare::prepare_audio_transcription_provider_call;
pub use types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
pub async fn audio_transcription(
request: AudioTranscriptionRequest<'_>,
_context: &LiteLlmRequestContext,
) -> Result<Value, Error> {
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
.await
}

View file

@ -21,29 +21,34 @@ fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProv
pub fn prepare_audio_transcription_provider_call(
request: AudioTranscriptionRequest<'_>,
) -> Result<ProviderAudioTranscriptionRequest, Error> {
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.or_else(|| {
request
.custom_llm_provider
.map(|provider| CustomLlmProvider {
model: request.model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(
"unable to resolve custom_llm_provider for audio transcription request".to_string(),
)
})?;
let provider_info = get_custom_llm_provider(
request.model,
request.options.custom_llm_provider.as_deref(),
)
.or_else(|| {
request
.options
.custom_llm_provider
.as_deref()
.map(|provider| CustomLlmProvider {
model: request.model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(
"unable to resolve custom_llm_provider for audio transcription request".to_string(),
)
})?;
let model = provider_info.model.to_string();
let config = provider_config(provider_info.custom_llm_provider)
.ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?;
let env_lookup = |key: &str| std::env::var(key).ok();
let mut headers = string_headers("audio transcription", request.extra_headers)?;
let auth = config.auth_strategy(&model, &request.optional_params, &env_lookup)?;
let mut headers = string_headers("audio transcription", request.options.extra_headers)?;
let auth = config.auth_strategy(&model, &request.options.provider_connection, &env_lookup)?;
if matches!(auth, AudioTranscriptionAuth::Bearer)
&& !has_header(&headers, "authorization")
&& let Some(api_key) = request.api_key
&& let Some(api_key) = request.options.api_key.as_deref()
{
headers.push(("Authorization".to_string(), format!("Bearer {api_key}")));
}
@ -51,9 +56,9 @@ pub fn prepare_audio_transcription_provider_call(
headers.push(("Content-Type".to_string(), "application/json".to_string()));
}
let url = config.complete_url(
request.api_base,
request.options.api_base.as_deref(),
&model,
&request.optional_params,
&request.options.provider_connection,
&env_lookup,
)?;
let filtered_params = config.map_transcription_params(&request.optional_params);
@ -68,7 +73,7 @@ pub fn prepare_audio_transcription_provider_call(
upstream_headers: headers,
auth,
#[cfg(feature = "bedrock-auth")]
optional_params: request.optional_params,
timeout: request.timeout,
provider_connection: request.options.provider_connection,
timeout: request.options.timeout,
})
}

View file

@ -1,3 +1,5 @@
use crate::request_context::LiteLlmRequestContext;
use crate::request_options::RequestOptions;
use std::io::{Read, Write};
use std::net::TcpListener;
use std::thread;
@ -27,22 +29,31 @@ async fn bedrock_request_is_signed_and_contains_audio() {
stream.write_all(response).expect("response");
});
let optional_params = Map::from_iter([
let provider_connection = Map::from_iter([
("aws_access_key_id".to_string(), json!("access-key")),
("aws_secret_access_key".to_string(), json!("secret-key")),
("aws_region_name".to_string(), json!("us-east-1")),
]);
let api_base = format!("http://{address}");
let response = audio_transcription(AudioTranscriptionRequest {
model: "mistral.voxtral-mini-3b-2507",
audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}),
api_key: None,
api_base: Some(&api_base),
custom_llm_provider: Some("bedrock"),
extra_headers: None,
optional_params,
timeout: None,
})
let response = audio_transcription(
AudioTranscriptionRequest {
model: "mistral.voxtral-mini-3b-2507",
audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}),
optional_params: Map::new(),
options: RequestOptions {
provider_connection,
api_key: None,
api_base: (Some(&api_base)).map(|value| value.to_string()),
custom_llm_provider: (Some("bedrock")).map(|value| value.to_string()),
extra_headers: None,
timeout: None,
..Default::default()
},
},
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect("transcription");
assert_eq!(response, json!({"text": "hello"}));

View file

@ -1,3 +1,4 @@
use crate::request_options::RequestOptions;
use std::time::Duration;
use serde::{Deserialize, Serialize};
@ -8,12 +9,8 @@ use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderCo
pub struct AudioTranscriptionRequest<'a> {
pub model: &'a str,
pub audio: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
pub options: RequestOptions,
}
#[derive(Clone)]
@ -26,7 +23,7 @@ pub struct ProviderAudioTranscriptionRequest {
pub(super) upstream_headers: Vec<(String, String)>,
pub(super) auth: AudioTranscriptionAuth,
#[cfg(feature = "bedrock-auth")]
pub(super) optional_params: Map<String, Value>,
pub(super) provider_connection: Map<String, Value>,
pub(super) timeout: Option<Duration>,
}

View file

@ -13,7 +13,7 @@ use super::types::{
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) async fn execute_chat_completions_provider_call(
request: ResolvedChatCompletionsRequest<'_>,
request: ResolvedChatCompletionsRequest,
) -> Result<ChatCompletionsResponse, Error> {
let request = prepare_provider_request(request)?;
let body = serde_json::to_vec(&request.body).map_err(|err| {
@ -101,11 +101,11 @@ pub(super) async fn signed_headers(
let unsigned: BTreeMap<String, String> = request.upstream_headers.iter().cloned().collect();
// A host with its own resolution chain hands the result down; only fall
// back to deriving credentials here when it supplied none.
let credentials = match host_supplied_credentials(&request.optional_params) {
let credentials = match host_supplied_credentials(&request.provider_connection) {
Some(credentials) => credentials,
None => {
resolve_credentials(
aws_auth_config(&request.optional_params, &env_lookup),
aws_auth_config(&request.provider_connection, &env_lookup),
&env_lookup,
)
.await?

View file

@ -7,6 +7,7 @@
//! calls the provider, and returns a typed OpenAI-shaped response.
use crate::Error;
use crate::request_context::LiteLlmRequestContext;
mod client;
mod common_utils;
pub mod conversation;
@ -25,8 +26,9 @@ use types::{ChatCompletionsRequest, ChatCompletionsResponse};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn chat_completions(
request: ChatCompletionsRequest<'_>,
context: &LiteLlmRequestContext,
) -> Result<ChatCompletionsResponse, Error> {
execute_chat_completions_provider_call(resolve_request(request)?).await
execute_chat_completions_provider_call(resolve_request(request, context)?).await
}
/// Whether the core would accept this request, without resolving credentials or

View file

@ -1,3 +1,4 @@
use crate::request_context::LiteLlmRequestContext;
use serde_json::Value;
use crate::error::Error;
@ -39,9 +40,13 @@ pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error>
pub(super) fn resolve_request(
request: ChatCompletionsRequest<'_>,
) -> Result<ResolvedChatCompletionsRequest<'_>, Error> {
let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)
.map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?;
_context: &LiteLlmRequestContext,
) -> Result<ResolvedChatCompletionsRequest, Error> {
let (model, config) = resolve_provider_config(
request.model,
request.options.custom_llm_provider.as_deref(),
)
.map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?;
let messages =
parse_messages(request.messages).map_err(|_| Error::Declined("unreadable message list"))?;
if messages.is_empty() {
@ -55,25 +60,22 @@ pub(super) fn resolve_request(
config,
messages,
optional_params: request.optional_params,
api_key: request.api_key,
api_base: request.api_base,
extra_headers: request.extra_headers,
timeout: request.timeout,
options: request.options,
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
fn validate_environment(
request: &ResolvedChatCompletionsRequest<'_>,
request: &ResolvedChatCompletionsRequest,
model: &str,
config: &dyn ChatCompletionsProviderConfig,
) -> Result<(Vec<(String, String)>, ChatCompletionsAuth), Error> {
let env_lookup = |key: &str| std::env::var(key).ok();
let mut headers = string_headers(request.extra_headers.clone())?;
let mut headers = string_headers(request.options.extra_headers.clone())?;
let auth = config.auth(
request.api_key,
request.options.api_key.as_deref(),
model,
&request.optional_params,
&request.options.provider_connection,
&env_lookup,
)?;
match &auth {
@ -117,16 +119,16 @@ fn validate_environment(
}
pub(super) fn prepare_provider_request(
request: ResolvedChatCompletionsRequest<'_>,
request: ResolvedChatCompletionsRequest,
) -> Result<ProviderChatCompletionsRequest, Error> {
let (headers, auth) = validate_environment(&request, &request.model, request.config)?;
let model = request.model;
let config = request.config;
let env_lookup = |key: &str| std::env::var(key).ok();
let url = config.complete_url(
request.api_base,
request.options.api_base.as_deref(),
&model,
&request.optional_params,
&request.options.provider_connection,
&env_lookup,
)?;
let transformed =
@ -139,7 +141,7 @@ pub(super) fn prepare_provider_request(
body: transformed.body,
upstream_headers: headers,
auth,
optional_params: request.optional_params,
timeout: request.timeout,
provider_connection: request.options.provider_connection,
timeout: request.options.timeout,
})
}

View file

@ -1,3 +1,5 @@
use crate::request_context::LiteLlmRequestContext;
use crate::request_options::RequestOptions;
use serde_json::{Map, Value, json};
use crate::error::Error;
@ -9,7 +11,12 @@ use super::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest};
fn prepare_chat_completions_call(
request: ChatCompletionsRequest<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
prepare_provider_request(resolve_request(request)?)
prepare_provider_request(resolve_request(
request,
&LiteLlmRequestContext {
..Default::default()
},
)?)
}
fn request<'a>(
@ -25,11 +32,15 @@ fn request<'a>(
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: provider,
extra_headers: None,
timeout: None,
options: RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: None,
custom_llm_provider: (provider).map(|value| value.to_string()),
extra_headers: None,
timeout: None,
..Default::default()
},
}
}
@ -107,7 +118,7 @@ fn the_deployment_credential_replaces_a_caller_supplied_auth_header() {
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
call.options.extra_headers = Some(Map::from_iter([(
"X-Api-Key".to_string(),
json!("sk-caller"),
)]));
@ -132,7 +143,7 @@ fn a_forwarded_authorization_header_suppresses_the_resolved_api_key_header() {
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
call.options.extra_headers = Some(Map::from_iter([
(
"Authorization".to_string(),
json!("Bearer sk-ant-oat01-token"),
@ -168,7 +179,7 @@ fn an_unrelated_forwarded_authorization_does_not_defer_the_resolved_key() {
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([
call.options.extra_headers = Some(Map::from_iter([
("Authorization".to_string(), json!("Bearer unrelated")),
("X-Api-Key".to_string(), json!("sk-caller")),
]));
@ -199,7 +210,7 @@ fn declines_an_unsupported_request_before_resolving_credentials() {
json!([{"role": "user", "content": "hi"}]),
json!({"stream": true}),
);
call.api_key = None;
call.options.api_key = None;
// No api_key is set and no env is consulted: the gate must run first, so the
// error is the decline rather than a missing-credential error.
assert_eq!(decline(call), Error::Declined("streaming"));
@ -261,7 +272,7 @@ fn rejects_non_string_extra_headers() {
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
call.options.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
assert_eq!(
decline(call),
Error::InvalidRequest(
@ -279,7 +290,7 @@ fn prepares_a_bedrock_call_without_resolving_credentials() {
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.api_key = None;
call.options.api_key = None;
let prepared = prepare_chat_completions_call(call).expect("prepares");
assert_eq!(
prepared.url,
@ -312,15 +323,18 @@ async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() {
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
json!({"maxTokens": 16}),
);
call.options.provider_connection = Map::from_iter([
("aws_access_key_id".to_string(), json!("AKIDEXAMPLE")),
(
"aws_secret_access_key".to_string(),
json!("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"),
),
]);
// A key would resolve to a bearer token and never reach the signer.
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(
call.options.api_key = None;
call.options.extra_headers = Some(Map::from_iter([(
"x-request-id".to_string(),
json!("abc-123"),
)]));
@ -367,14 +381,18 @@ async fn a_forwarded_header_the_signer_computes_declines_to_python() {
"bedrock/us-east-1/anthropic.claude-v2",
None,
json!([{"role": "user", "content": "hi"}]),
json!({
"maxTokens": 16,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
}),
json!({"maxTokens": 16}),
);
call.api_key = None;
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
call.options.provider_connection = Map::from_iter([
("aws_access_key_id".to_string(), json!("AKIDEXAMPLE")),
(
"aws_secret_access_key".to_string(),
json!("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"),
),
]);
call.options.api_key = None;
call.options.extra_headers =
Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
let prepared = prepare_chat_completions_call(call).expect("prepares");
let error = super::handler::signed_headers(&prepared, br#"{"a":1}"#)
.await
@ -399,7 +417,7 @@ fn a_bedrock_deployment_bearer_outranks_a_forwarded_authorization() {
json!([{"role": "user", "content": "hi"}]),
json!({"maxTokens": 16}),
);
call.extra_headers = Some(Map::from_iter([(
call.options.extra_headers = Some(Map::from_iter([(
"Authorization".to_string(),
json!("Bearer caller-supplied"),
)]));
@ -432,7 +450,7 @@ fn an_anthropic_forwarded_oauth_bearer_still_outranks_the_resolved_key() {
json!([{"role": "user", "content": "hi"}]),
json!({}),
);
call.extra_headers = Some(Map::from_iter([(
call.options.extra_headers = Some(Map::from_iter([(
"authorization".to_string(),
json!("Bearer sk-ant-oat01-forwarded"),
)]));
@ -666,11 +684,15 @@ mod round_trip {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
api_key: Some("sk-test"),
api_base: Some(api_base),
custom_llm_provider: None,
extra_headers: None,
timeout: Some(std::time::Duration::from_secs(10)),
options: RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: (Some(api_base)).map(|value| value.to_string()),
custom_llm_provider: None,
extra_headers: None,
timeout: Some(std::time::Duration::from_secs(10)),
..Default::default()
},
}
}
@ -679,14 +701,19 @@ mod round_trip {
#[tokio::test]
async fn round_trip_sends_the_translated_body_and_normalizes_the_response() {
let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await;
let response = chat_completions(call(
&api_base,
json!([
{"role": "system", "content": "be terse"},
{"role": "user", "content": "hi"}
]),
json!({"max_tokens": 16}),
))
let response = chat_completions(
call(
&api_base,
json!([
{"role": "system", "content": "be terse"},
{"role": "user", "content": "hi"}
]),
json!({"max_tokens": 16}),
),
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect("call succeeds");
@ -724,11 +751,16 @@ mod round_trip {
const NO_USAGE: &str =
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#;
let (api_base, handle) = serve_once("200 OK", NO_USAGE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
let err = chat_completions(
call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
),
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
@ -742,11 +774,16 @@ mod round_trip {
async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() {
const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#;
let (api_base, handle) = serve_once("200 OK", TOOL_USE).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
let err = chat_completions(
call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
),
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect_err("response cannot be normalized");
handle.await.expect("server task");
@ -760,11 +797,16 @@ mod round_trip {
async fn an_upstream_error_status_keeps_its_code() {
let (api_base, handle) =
serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await;
let err = chat_completions(call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
let err = chat_completions(
call(
&api_base,
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
),
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect_err("upstream rejects");
handle.await.expect("server task");
@ -781,11 +823,16 @@ mod round_trip {
listener.local_addr().expect("has an address").port()
// Dropped here, so the port is closed and the connect is refused.
};
let err = chat_completions(call(
&format!("http://127.0.0.1:{port}/v1/messages"),
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
))
let err = chat_completions(
call(
&format!("http://127.0.0.1:{port}/v1/messages"),
json!([{"role": "user", "content": "hi"}]),
json!({"max_tokens": 16}),
),
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect_err("nothing is listening");
assert!(

View file

@ -1,3 +1,4 @@
use crate::request_options::RequestOptions;
use std::time::Duration;
use serde::{Deserialize, Serialize};
@ -15,22 +16,15 @@ pub struct ChatCompletionsRequest<'a> {
pub model: &'a str,
pub messages: Value,
pub optional_params: Map<String, Value>,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
pub options: RequestOptions,
}
pub(super) struct ResolvedChatCompletionsRequest<'a> {
pub(super) struct ResolvedChatCompletionsRequest {
pub(super) model: String,
pub(super) config: &'static dyn ChatCompletionsProviderConfig,
pub(super) messages: Vec<ChatMessage>,
pub(super) optional_params: Map<String, Value>,
pub(super) api_key: Option<&'a str>,
pub(super) api_base: Option<&'a str>,
pub(super) extra_headers: Option<Map<String, Value>>,
pub(super) timeout: Option<Duration>,
pub(super) options: RequestOptions,
}
pub(super) struct ProviderChatCompletionsRequest {
@ -41,7 +35,7 @@ pub(super) struct ProviderChatCompletionsRequest {
pub(super) upstream_headers: Vec<(String, String)>,
pub(super) auth: ChatCompletionsAuth,
#[cfg_attr(not(feature = "bedrock-auth"), allow(dead_code))]
pub(super) optional_params: Map<String, Value>,
pub(super) provider_connection: Map<String, Value>,
pub(super) timeout: Option<Duration>,
}

View file

@ -16,3 +16,6 @@ pub mod router;
pub mod routing_utils;
pub use error::Error;
pub mod request_context;
pub mod request_options;

View file

@ -8,6 +8,7 @@
//! can splice the event stream to its own caller.
use crate::Error;
use crate::request_context::LiteLlmRequestContext;
mod client;
mod common_utils;
mod handler;
@ -19,11 +20,17 @@ use handler::{execute_messages_provider_call, execute_messages_provider_stream};
use types::{AnthropicMessagesResponse, MessagesRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
pub async fn messages(
request: MessagesRequest<'_>,
_context: &LiteLlmRequestContext,
) -> Result<AnthropicMessagesResponse, Error> {
execute_messages_provider_call(request).await
}
pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Response, Error> {
pub async fn messages_stream(
request: MessagesRequest<'_>,
_context: &LiteLlmRequestContext,
) -> Result<reqwest::Response, Error> {
execute_messages_provider_stream(request).await
}

View file

@ -9,20 +9,25 @@ use serde_json::{Map, Value};
pub(super) fn prepare_provider_request(
request: MessagesRequest<'_>,
) -> Result<ProviderMessagesRequest, Error> {
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.or_else(|| {
request
.custom_llm_provider
.map(|provider| CustomLlmProvider {
model: request.model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(
"unable to resolve custom_llm_provider for messages request".to_string(),
)
})?;
let provider_info = get_custom_llm_provider(
request.model,
request.options.custom_llm_provider.as_deref(),
)
.or_else(|| {
request
.options
.custom_llm_provider
.as_deref()
.map(|provider| CustomLlmProvider {
model: request.model,
custom_llm_provider: provider,
})
})
.ok_or_else(|| {
Error::InvalidProvider(
"unable to resolve custom_llm_provider for messages request".to_string(),
)
})?;
let model = provider_info.model.to_string();
let provider = provider_info.custom_llm_provider;
@ -30,8 +35,12 @@ pub(super) fn prepare_provider_request(
.ok_or_else(|| Error::InvalidProvider(provider.to_string()))?;
let env_lookup = |key: &str| std::env::var(key).ok();
let headers =
validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?;
let headers = validate_environment(
config,
request.options.extra_headers,
request.options.api_key.as_deref(),
&env_lookup,
)?;
let typed_request = serde_json::from_value(request.body).map_err(|err| {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
@ -43,7 +52,7 @@ pub(super) fn prepare_provider_request(
))
})?;
let url = config.complete_url(request.api_base, &model, &env_lookup)?;
let url = config.complete_url(request.options.api_base.as_deref(), &model, &env_lookup)?;
Ok(ProviderMessagesRequest {
provider: provider.to_string(),
@ -52,7 +61,7 @@ pub(super) fn prepare_provider_request(
url,
body,
upstream_headers: headers,
timeout: request.timeout,
timeout: request.options.timeout,
})
}

View file

@ -1,3 +1,5 @@
use crate::request_context::LiteLlmRequestContext;
use crate::request_options::RequestOptions;
use std::time::Duration;
use serde_json::{Map, Value, json};
@ -131,26 +133,34 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
request
});
let response = messages(MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({
"model": "claude-sonnet-4-5",
"max_tokens": 1024,
"messages": [{
"role": "user",
"content": [{
"type": "text",
"text": "hi",
"cache_control": {"type": "ephemeral", "scope": "global"}
let response = messages(
MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({
"model": "claude-sonnet-4-5",
"max_tokens": 1024,
"messages": [{
"role": "user",
"content": [{
"type": "text",
"text": "hi",
"cache_control": {"type": "ephemeral", "scope": "global"}
}]
}]
}]
}),
api_key: Some("sk-azure"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
})
}),
options: RequestOptions {
api_key: (Some("sk-azure")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("azure_ai")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect("messages request succeeds");
@ -194,19 +204,27 @@ async fn messages_round_trip_builds_native_anthropic_request() {
request
});
let response = messages(MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({
"model": "claude-sonnet-4-5",
"max_tokens": 1024,
"messages": [{"role": "user", "content": "hi"}]
}),
api_key: Some("sk-ant"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("anthropic"),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
})
let response = messages(
MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({
"model": "claude-sonnet-4-5",
"max_tokens": 1024,
"messages": [{"role": "user", "content": "hi"}]
}),
options: RequestOptions {
api_key: (Some("sk-ant")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("anthropic")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect("messages request succeeds");
@ -251,15 +269,23 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() {
Value::String("token-efficient-tools-2025-02-19".to_string()),
);
messages(MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
api_key: Some("rust-fallback-key"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: Some(headers),
timeout: Some(Duration::from_secs(5)),
})
messages(
MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
options: RequestOptions {
api_key: (Some("rust-fallback-key")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("azure_ai")).map(|value| value.to_string()),
extra_headers: Some(headers),
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect("messages request succeeds");
@ -305,15 +331,23 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
Value::String("Bearer entra-token".to_string()),
);
messages(MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
api_key: None,
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: Some(headers),
timeout: Some(Duration::from_secs(5)),
})
messages(
MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
options: RequestOptions {
api_key: None,
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("azure_ai")).map(|value| value.to_string()),
extra_headers: Some(headers),
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect("entra id request succeeds without api key");
@ -329,15 +363,23 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() {
#[tokio::test]
async fn messages_requires_auth_when_no_key_and_no_header() {
let err = messages(MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
api_key: None,
api_base: Some("http://127.0.0.1:1"),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
timeout: Some(Duration::from_millis(50)),
})
let err = messages(
MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
options: RequestOptions {
api_key: None,
api_base: (Some("http://127.0.0.1:1")).map(|value| value.to_string()),
custom_llm_provider: (Some("azure_ai")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_millis(50)),
..Default::default()
},
},
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect_err("missing auth errors");
@ -367,15 +409,23 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() {
Value::String("Bearer ".to_string()),
);
messages(MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
api_key: Some("sk-azure"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: Some(headers),
timeout: Some(Duration::from_secs(5)),
})
messages(
MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
options: RequestOptions {
api_key: (Some("sk-azure")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("azure_ai")).map(|value| value.to_string()),
extra_headers: Some(headers),
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect("falls back to api key");
@ -408,15 +458,23 @@ async fn messages_maps_provider_error_status_to_http_error() {
.expect("writes response");
});
let err = messages(MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
api_key: Some("sk-azure"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
})
let err = messages(
MessagesRequest {
model: "claude-sonnet-4-5",
body: json!({"model": "claude-sonnet-4-5", "max_tokens": 8, "messages": []}),
options: RequestOptions {
api_key: (Some("sk-azure")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("azure_ai")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
},
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect_err("provider error propagates");
@ -425,15 +483,23 @@ async fn messages_maps_provider_error_status_to_http_error() {
#[tokio::test]
async fn messages_rejects_unsupported_provider() {
let err = messages(MessagesRequest {
model: "claude-3-5-sonnet",
body: json!({"model": "claude-3-5-sonnet", "max_tokens": 8, "messages": []}),
api_key: Some("sk"),
api_base: Some("http://127.0.0.1:1"),
custom_llm_provider: Some("openai"),
extra_headers: None,
timeout: Some(Duration::from_millis(50)),
})
let err = messages(
MessagesRequest {
model: "claude-3-5-sonnet",
body: json!({"model": "claude-3-5-sonnet", "max_tokens": 8, "messages": []}),
options: RequestOptions {
api_key: (Some("sk")).map(|value| value.to_string()),
api_base: (Some("http://127.0.0.1:1")).map(|value| value.to_string()),
custom_llm_provider: (Some("openai")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_millis(50)),
..Default::default()
},
},
&LiteLlmRequestContext {
..Default::default()
},
)
.await
.expect_err("unsupported provider errors");

View file

@ -1,3 +1,4 @@
use crate::request_options::RequestOptions;
use std::time::Duration;
use serde::{Deserialize, Serialize};
@ -8,11 +9,7 @@ use super::transformation::AnthropicMessagesProviderConfig;
pub struct MessagesRequest<'a> {
pub model: &'a str,
pub body: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
pub options: RequestOptions,
}
pub(super) struct ProviderMessagesRequest {

View file

@ -0,0 +1,18 @@
use serde_json::{Map, Value};
#[derive(Clone, Debug, Default, PartialEq)]
pub struct RequestAttribution {
pub user_api_key_hash: Option<String>,
pub user_api_key_user_id: Option<String>,
pub user_api_key_team_id: Option<String>,
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct LiteLlmRequestContext {
pub metadata: Option<Map<String, Value>>,
pub litellm_metadata: Option<Map<String, Value>>,
pub request_metadata_fields: Vec<String>,
pub litellm_call_id: Option<String>,
pub request_model: Option<String>,
pub attribution: RequestAttribution,
}

View file

@ -0,0 +1,14 @@
use std::time::Duration;
use serde_json::{Map, Value};
#[derive(Clone, Debug, Default)]
pub struct RequestOptions {
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
pub extra_query: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
pub provider_connection: Map<String, Value>,
}

View file

@ -1,6 +1,12 @@
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use serde_json::{Map, Value};
#[derive(Clone, Debug)]
pub struct ResponsesWebSocketRequest {
pub url: String,
pub options: crate::request_options::RequestOptions,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ResponsesWsEventType {
ResponseCreate,

View file

@ -7,12 +7,18 @@ mod marshal;
mod routes;
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use litellm_core::responses::types::ResponsesWebSocketRequest;
use pyo3::prelude::*;
use pyo3::types::PyAny;
use serde_json::Value;
use crate::errors::core_error_to_pyerr;
use crate::marshal::{marshal_headers, optional_timeout};
use crate::marshal::{NativeRequestContext, NativeRequestOptions};
#[derive(FromPyObject)]
struct WebSocketConnectRequest {
url: String,
options: NativeRequestOptions,
}
#[pyclass]
struct ResponsesWebSocketConnection {
@ -22,18 +28,21 @@ struct ResponsesWebSocketConnection {
#[pymethods]
impl ResponsesWebSocketConnection {
#[classmethod]
#[pyo3(signature = (url, headers=None, timeout_seconds=None))]
#[pyo3(signature = (request, *, context))]
fn connect<'py>(
_cls: &Bound<'py, pyo3::types::PyType>,
py: Python<'py>,
url: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)] headers: Option<Value>,
timeout_seconds: Option<f64>,
request: WebSocketConnectRequest,
context: NativeRequestContext,
) -> PyResult<Bound<'py, PyAny>> {
let headers = marshal_headers(headers)?;
let timeout = optional_timeout(timeout_seconds);
let options: litellm_core::request_options::RequestOptions = request.options.into();
let context: litellm_core::request_context::LiteLlmRequestContext = context.into();
let request = ResponsesWebSocketRequest {
url: request.url,
options,
};
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout)
let inner = RustResponsesWebSocketConnection::connect(request, &context)
.await
.map_err(core_error_to_pyerr)?;
Ok(ResponsesWebSocketConnection { inner })
@ -81,7 +90,6 @@ mod tests {
use std::time::Duration;
use futures_util::{SinkExt, StreamExt};
use pyo3::types::PyDict;
use tokio::net::TcpListener;
use tokio_tungstenite::{accept_async, tungstenite::Message};
@ -186,7 +194,7 @@ mod tests {
Python::attach(|py| {
let module = pyo3::wrap_pymodule!(_native)(py).into_bound(py);
let locals = PyDict::new(py);
let locals = crate::marshal::request_fixtures(py);
locals
.set_item("native", &module)
.expect("module should enter Python locals");
@ -198,7 +206,24 @@ mod tests {
import asyncio
async def exercise():
connection = await native.ResponsesWebSocketConnection.connect(url)
for request, request_context, field in (
(Request(url=123), context, 'url'),
(Request(url=url, options=Options(extra_headers=[])), context, 'extra_headers'),
(Request(url=url), replace(context, litellm_call_id=123), 'litellm_call_id'),
(Request(url=url), replace(context, attribution=Attribution(user_api_key_user_id=123)), 'user_api_key_user_id'),
):
try:
native.ResponsesWebSocketConnection.connect(request, context=request_context)
except (TypeError, ValueError) as error:
parts = []
while error is not None:
parts.append(str(error))
error = error.__cause__
assert field in ' / '.join(parts), parts
else:
raise AssertionError('invalid WebSocket input reached execution')
connection = await native.ResponsesWebSocketConnection.connect(Request(url=url), context=context)
assert type(connection) is native.ResponsesWebSocketConnection
await connection.send_text("from-python")
assert await connection.recv_text() == "from-server"

View file

@ -1,38 +1,70 @@
use std::collections::HashMap;
use std::time::Duration;
use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use serde_json::{Map, Value};
pub(crate) struct RouteOptions {
pub(crate) model: String,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) custom_llm_provider: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
pub(crate) timeout: Option<Duration>,
#[derive(FromPyObject)]
pub(crate) struct NativeRequestOptions {
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<Map<String, Value>>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_query: Option<Map<String, Value>>,
timeout_seconds: Option<f64>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
provider_connection: Option<Map<String, Value>>,
}
pub(crate) struct RouteOptionsInputs {
pub(crate) model: String,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) custom_llm_provider: Option<String>,
pub(crate) extra_headers: Option<Value>,
pub(crate) timeout_seconds: Option<f64>,
impl From<NativeRequestOptions> for litellm_core::request_options::RequestOptions {
fn from(input: NativeRequestOptions) -> Self {
Self {
api_key: input.api_key,
api_base: input.api_base,
custom_llm_provider: input.custom_llm_provider,
extra_headers: input.extra_headers,
extra_query: input.extra_query,
timeout: optional_timeout(input.timeout_seconds),
provider_connection: input.provider_connection.unwrap_or_default(),
}
}
}
impl RouteOptions {
pub(crate) fn from_python(inputs: RouteOptionsInputs) -> PyResult<Self> {
Ok(Self {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: optional_object("extra_headers", inputs.extra_headers)?,
timeout: optional_timeout(inputs.timeout_seconds),
})
#[derive(FromPyObject)]
pub(crate) struct NativeRequestAttribution {
user_api_key_hash: Option<String>,
user_api_key_user_id: Option<String>,
user_api_key_team_id: Option<String>,
}
#[derive(FromPyObject)]
pub(crate) struct NativeRequestContext {
#[pyo3(from_py_with = litellm_python_interop::from_py)]
metadata: Option<Map<String, Value>>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
litellm_metadata: Option<Map<String, Value>>,
request_metadata_fields: Vec<String>,
litellm_call_id: Option<String>,
request_model: Option<String>,
attribution: NativeRequestAttribution,
}
impl From<NativeRequestContext> for litellm_core::request_context::LiteLlmRequestContext {
fn from(input: NativeRequestContext) -> Self {
Self {
metadata: input.metadata,
litellm_metadata: input.litellm_metadata,
request_metadata_fields: input.request_metadata_fields,
litellm_call_id: input.litellm_call_id,
request_model: input.request_model,
attribution: litellm_core::request_context::RequestAttribution {
user_api_key_hash: input.attribution.user_api_key_hash,
user_api_key_user_id: input.attribution.user_api_key_user_id,
user_api_key_team_id: input.attribution.user_api_key_team_id,
},
}
}
}
@ -50,30 +82,6 @@ pub(crate) fn required_value(
)))
}
pub(crate) fn object_or_empty(
name: &'static str,
value: Option<Value>,
) -> PyResult<Map<String, Value>> {
match value {
Some(value) => object(name, value),
None => Ok(Map::new()),
}
}
fn optional_object(
name: &'static str,
value: Option<Value>,
) -> PyResult<Option<Map<String, Value>>> {
value.map(|value| object(name, value)).transpose()
}
fn object(name: &'static str, value: Value) -> PyResult<Map<String, Value>> {
match value {
Value::Object(map) => Ok(map),
_ => Err(PyValueError::new_err(format!("{name} must be a dict"))),
}
}
pub(crate) fn optional_timeout(timeout_seconds: Option<f64>) -> Option<Duration> {
timeout_seconds.and_then(|secs| {
if secs.is_finite() && secs > 0.0 {
@ -84,21 +92,55 @@ pub(crate) fn optional_timeout(timeout_seconds: Option<f64>) -> Option<Duration>
})
}
pub(crate) fn marshal_headers(headers: Option<Value>) -> PyResult<HashMap<String, String>> {
let value = match headers {
Some(headers) => headers,
None => Value::Object(Map::new()),
};
let Value::Object(headers) = value else {
return Err(PyValueError::new_err("headers must be a dict"));
};
headers
.into_iter()
.map(|(name, value)| {
value
.as_str()
.map(|value| (name, value.to_string()))
.ok_or_else(|| PyValueError::new_err("header values must be strings"))
})
.collect()
#[cfg(test)]
pub(crate) fn request_fixtures(py: Python<'_>) -> Bound<'_, pyo3::types::PyDict> {
let locals = pyo3::types::PyDict::new(py);
py.run(
c"
from dataclasses import dataclass, replace
@dataclass(frozen=True)
class Options:
api_key: object = None
api_base: object = None
custom_llm_provider: object = None
extra_headers: object = None
extra_query: object = None
timeout_seconds: object = None
provider_connection: object = None
@dataclass(frozen=True)
class Attribution:
user_api_key_hash: object = None
user_api_key_user_id: object = None
user_api_key_team_id: object = None
@dataclass(frozen=True)
class Context:
metadata: object = None
litellm_metadata: object = None
request_metadata_fields: tuple = ()
litellm_call_id: object = None
request_model: object = None
attribution: Attribution = Attribution()
@dataclass(frozen=True)
class Request:
model: str = 'model'
messages: object = None
body: object = None
audio: object = None
document: object = None
optional_params: object = None
options: Options = Options()
value: str = ''
url: str = ''
context = Context()
",
Some(&locals),
Some(&locals),
)
.expect("request dataclasses should load");
locals
}

View file

@ -1,48 +1,39 @@
use crate::errors::core_error_to_pyerr;
use crate::marshal::{NativeRequestContext, NativeRequestOptions};
use litellm_core::Error;
use litellm_core::audio_transcription::AudioTranscriptionRequest;
use litellm_core::audio_transcription::audio_transcription as run_route;
use litellm_core::request_context::LiteLlmRequestContext;
use pyo3::prelude::*;
use serde_json::{Map, Value};
use std::future::Future;
use litellm_core::audio_transcription::{
AudioTranscriptionRequest, audio_transcription as run_audio_transcription,
};
use pyo3::prelude::*;
use serde_json::Value;
use crate::errors::core_error_to_pyerr;
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
#[derive(FromPyObject)]
struct AudioTranscriptionInputs {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
audio: Value,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
optional_params: Map<String, Value>,
options: NativeRequestOptions,
}
fn prepare_transcription(
inputs: AudioTranscriptionInputs,
input: AudioTranscriptionInputs,
context: NativeRequestContext,
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
let audio = inputs.audio;
let options = RouteOptions::from_python(RouteOptionsInputs {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: inputs.extra_headers,
timeout_seconds: inputs.timeout_seconds,
})?;
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
let context: LiteLlmRequestContext = context.into();
let audio = input.audio;
Ok(async move {
let RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
} = options;
run_audio_transcription(AudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
})
run_route(
AudioTranscriptionRequest {
model: &input.model,
audio,
optional_params: input.optional_params,
options: input.options.into(),
},
&context,
)
.await
})
}
@ -50,22 +41,7 @@ fn prepare_transcription(
bridge_route! {
sync = transcription,
asynchronous = atranscription,
inputs = AudioTranscriptionInputs,
required = {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
audio: serde_json::Value,
},
optional = {
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<serde_json::Value>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
optional_params: Option<serde_json::Value>,
timeout_seconds: Option<f64>,
},
request = AudioTranscriptionInputs,
prepare = prepare_transcription,
errors = core_error_to_pyerr,
}

View file

@ -1,49 +1,40 @@
use crate::errors::chat_completions_error_to_pyerr;
use crate::marshal::{NativeRequestContext, NativeRequestOptions, required_value};
use litellm_core::Error;
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 pyo3::prelude::*;
use serde_json::{Map, Value};
use std::future::Future;
use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse};
use litellm_core::chat_completions::{
chat_completions as run_chat_completions, chat_completions_decline_reason,
};
use pyo3::prelude::*;
use serde_json::Value;
use crate::errors::chat_completions_error_to_pyerr;
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty, required_value};
#[derive(FromPyObject)]
struct ChatCompletionsInputs {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
messages: Value,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
optional_params: Map<String, Value>,
options: NativeRequestOptions,
}
fn prepare_chat_completions(
inputs: ChatCompletionsInputs,
input: ChatCompletionsInputs,
context: NativeRequestContext,
) -> PyResult<impl Future<Output = Result<ChatCompletionsResponse, Error>> + Send + 'static> {
let messages = required_value("messages", inputs.messages, Value::is_array, "list")?;
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
let options = RouteOptions::from_python(RouteOptionsInputs {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: inputs.extra_headers,
timeout_seconds: inputs.timeout_seconds,
})?;
let context: LiteLlmRequestContext = context.into();
let messages = required_value("messages", input.messages, Value::is_array, "list")?;
Ok(async move {
let RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
} = options;
run_chat_completions(ChatCompletionsRequest {
model: &model,
messages,
optional_params,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
})
run_route(
ChatCompletionsRequest {
model: &input.model,
messages,
optional_params: input.optional_params,
options: input.options.into(),
},
&context,
)
.await
})
}
@ -56,7 +47,15 @@ fn chat_completions_decline(
#[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option<Value>,
custom_llm_provider: Option<String>,
) -> PyResult<Option<String>> {
let optional_params = object_or_empty("optional_params", optional_params)?;
let optional_params = match optional_params {
None | Some(Value::Null) => Map::new(),
Some(Value::Object(params)) => params,
Some(_) => {
return Err(pyo3::exceptions::PyValueError::new_err(
"optional_params must be a dict",
));
}
};
Ok(chat_completions_decline_reason(
&model,
custom_llm_provider.as_deref(),
@ -69,22 +68,7 @@ fn chat_completions_decline(
bridge_route! {
sync = chat_completions,
asynchronous = achat_completions,
inputs = ChatCompletionsInputs,
required = {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
messages: serde_json::Value,
},
optional = {
#[pyo3(from_py_with = litellm_python_interop::from_py)]
optional_params: Option<serde_json::Value>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<serde_json::Value>,
timeout_seconds: Option<f64>,
},
request = ChatCompletionsInputs,
prepare = prepare_chat_completions,
errors = chat_completions_error_to_pyerr,
extra = [chat_completions_decline],

View file

@ -6,52 +6,34 @@ macro_rules! bridge_route {
(
sync = $sync_name:ident,
asynchronous = $async_name:ident,
inputs = $inputs:ident,
required = { $($(#[$required_attr:meta])* $required_name:ident: $required_type:ty),+ $(,)? },
optional = { $($(#[$optional_attr:meta])* $optional_name:ident: $optional_type:ty),* $(,)? },
request = $inputs:ident,
prepare = $prepare:path,
errors = $map_error:path
$(, extra = [$($extra:ident),* $(,)?])?
$(,)?
$(, extra = [$($extra:ident),* $(,)?])? $(,)?
) => {
struct $inputs {
$($required_name: $required_type,)*
$($optional_name: $optional_type),*
}
#[pyfunction]
#[pyo3(signature = ($($required_name),*, $($optional_name=None),*))]
#[allow(clippy::too_many_arguments)]
#[pyo3(signature = (request, *, context))]
fn $sync_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
request: $inputs,
context: $crate::marshal::NativeRequestContext,
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
let future = $prepare($inputs {
$($required_name,)*
$($optional_name),*
})?;
let future = $prepare(request, context)?;
$crate::execution::run_sync(py, future, $map_error)
}
#[pyfunction]
#[pyo3(signature = ($($required_name),*, $($optional_name=None),*))]
#[allow(clippy::too_many_arguments)]
#[pyo3(signature = (request, *, context))]
fn $async_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
request: $inputs,
context: $crate::marshal::NativeRequestContext,
) -> pyo3::PyResult<pyo3::Bound<'_, pyo3::PyAny>> {
let future = $prepare($inputs {
$($required_name,)*
$($optional_name),*
})?;
let future = $prepare(request, context)?;
$crate::execution::run_async(py, future, $map_error)
}
pub(super) fn register(
module: &pyo3::Bound<'_, pyo3::types::PyModule>,
) -> pyo3::PyResult<()> {
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
$($($crate::routes::definition::add_function(module, pyo3::wrap_pyfunction!($extra, module)?)?;)*)?
$crate::routes::definition::add_function(module, pyo3::wrap_pyfunction!($sync_name, module)?)?;
$crate::routes::definition::add_function(module, pyo3::wrap_pyfunction!($async_name, module)?)?;
@ -64,62 +46,36 @@ macro_rules! bridge_route {
use super::{$inputs, $map_error, $prepare};
#[pyfunction]
#[pyo3(signature = ($($required_name),*, $($optional_name=None),*))]
#[allow(clippy::too_many_arguments)]
#[pyo3(signature = (request, *, context))]
fn $sync_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
request: $inputs,
context: $crate::marshal::NativeRequestContext,
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
let future = $prepare($inputs {
$($required_name,)*
$($optional_name),*
})?;
$crate::execution::run_sync(
py,
$crate::function_trace::capture(future),
$map_error,
)
let future = $prepare(request, context)?;
$crate::execution::run_sync(py, $crate::function_trace::capture(future), $map_error)
}
#[pyfunction]
#[pyo3(signature = ($($required_name),*, $($optional_name=None),*))]
#[allow(clippy::too_many_arguments)]
#[pyo3(signature = (request, *, context))]
fn $async_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
request: $inputs,
context: $crate::marshal::NativeRequestContext,
) -> pyo3::PyResult<pyo3::Bound<'_, pyo3::PyAny>> {
let future = $prepare($inputs {
$($required_name,)*
$($optional_name),*
})?;
$crate::execution::run_async(
py,
$crate::function_trace::capture(future),
$map_error,
)
let future = $prepare(request, context)?;
$crate::execution::run_async(py, $crate::function_trace::capture(future), $map_error)
}
pub(super) fn register(
module: &pyo3::Bound<'_, pyo3::types::PyModule>,
) -> pyo3::PyResult<()> {
$crate::routes::definition::add_function(
module,
pyo3::wrap_pyfunction!($sync_name, module)?,
)?;
$crate::routes::definition::add_function(
module,
pyo3::wrap_pyfunction!($async_name, module)?,
)?;
pub(super) fn register(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
$crate::routes::definition::add_function(module, pyo3::wrap_pyfunction!($sync_name, module)?)?;
$crate::routes::definition::add_function(module, pyo3::wrap_pyfunction!($async_name, module)?)?;
Ok(())
}
}
#[cfg(feature = "trace-parity")]
pub(super) fn register_trace(
module: &pyo3::Bound<'_, pyo3::types::PyModule>,
) -> pyo3::PyResult<()> {
pub(super) fn register_trace(module: &pyo3::Bound<'_, pyo3::types::PyModule>) -> pyo3::PyResult<()> {
trace::register(module)
}
};
@ -145,7 +101,7 @@ mod tests {
use litellm_core::error::Error;
use pyo3::exceptions::PyLookupError;
use pyo3::types::{PyDict, PyList};
use pyo3::types::PyDict;
use super::*;
@ -169,12 +125,15 @@ mod tests {
FUTURE_DROPPED.load(Ordering::SeqCst)
}
#[derive(FromPyObject)]
struct EchoInputs {
value: String,
}
bridge_route! {
sync = echo,
asynchronous = aecho,
inputs = EchoInputs,
required = { value: String },
optional = {},
request = EchoInputs,
prepare = prepare_echo,
errors = map_error,
extra = [future_dropped],
@ -182,6 +141,7 @@ mod tests {
fn prepare_echo(
inputs: EchoInputs,
_context: crate::marshal::NativeRequestContext,
) -> PyResult<impl Future<Output = Result<String, Error>> + Send + 'static> {
FUTURE_DROPPED.store(false, Ordering::SeqCst);
let drop_guard = (inputs.value == "pending").then_some(DropGuard);
@ -222,25 +182,13 @@ mod tests {
let module = PyModule::new(py, "routes").expect("module should be created");
crate::routes::register(&module).expect("routes should register");
let routes = [
(
"ocr",
"aocr",
"(model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)",
),
(
"transcription",
"atranscription",
"(model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)",
),
(
"messages",
"amessages",
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)",
),
("ocr", "aocr", "(request, *, context)"),
("transcription", "atranscription", "(request, *, context)"),
("messages", "amessages", "(request, *, context)"),
(
"chat_completions",
"achat_completions",
"(model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)",
"(request, *, context)",
),
];
@ -263,129 +211,45 @@ mod tests {
}
#[test]
fn sync_and_async_routes_apply_the_same_input_validation() {
fn routes_validate_dataclass_inputs_before_execution() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "routes").expect("module should be created");
crate::routes::register(&module).expect("routes should register");
let locals = crate::marshal::request_fixtures(py);
locals.set_item("routes", module).unwrap();
py.run(c"
for names, request, expected in [
(('chat_completions', 'achat_completions'), Request(messages={}, optional_params={}), 'messages must be a list'),
(('messages', 'amessages'), Request(body=[]), 'body must be a dict'),
(('ocr', 'aocr'), Request(document={}, optional_params={}, options=Options(extra_headers=[])), 'extra_headers'),
(('transcription', 'atranscription'), Request(audio={}, optional_params={}, options=Options(timeout_seconds='bad')), 'timeout_seconds'),
(('transcription', 'atranscription'), Request(audio={}, optional_params={}, options=Options(provider_connection=[])), 'provider_connection'),
]:
errors = []
for name in names:
try:
getattr(routes, name)(request, context=context)
except (ValueError, TypeError) as error:
parts = []
while error is not None:
parts.append(str(error))
error = error.__cause__
errors.append(' / '.join(parts))
else:
raise AssertionError('invalid input reached execution')
assert errors[0] == errors[1], errors
assert expected in errors[0], (expected, errors)
let invalid_messages = PyDict::new(py);
let sync_chat_error = module
.getattr("chat_completions")
.and_then(|function| function.call1(("model", &invalid_messages)))
.expect_err("sync chat should reject a non-list messages value");
let async_chat_error = module
.getattr("achat_completions")
.and_then(|function| function.call1(("model", &invalid_messages)))
.expect_err("async chat should reject a non-list messages value");
assert_eq!(
sync_chat_error.to_string(),
"ValueError: messages must be a list"
);
assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string());
let invalid_body = PyList::empty(py);
let sync_messages_error = module
.getattr("messages")
.and_then(|function| function.call1(("model", &invalid_body)))
.expect_err("sync Messages should reject a non-dict body");
let async_messages_error = module
.getattr("amessages")
.and_then(|function| function.call1(("model", &invalid_body)))
.expect_err("async Messages should reject a non-dict body");
assert_eq!(
sync_messages_error.to_string(),
"ValueError: body must be a dict"
);
assert_eq!(
async_messages_error.to_string(),
sync_messages_error.to_string()
);
let invalid_headers = PyList::empty(py);
let kwargs = PyDict::new(py);
kwargs
.set_item("extra_headers", &invalid_headers)
.expect("kwargs should accept extra_headers");
let document = PyDict::new(py);
for (sync_name, async_name) in [("ocr", "aocr"), ("transcription", "atranscription")] {
let sync_error = module
.getattr(sync_name)
.and_then(|function| function.call(("model", &document), Some(&kwargs)))
.expect_err("sync route should reject non-dict extra_headers");
let async_error = module
.getattr(async_name)
.and_then(|function| function.call(("model", &document), Some(&kwargs)))
.expect_err("async route should reject non-dict extra_headers");
assert_eq!(
sync_error.to_string(),
"ValueError: extra_headers must be a dict"
);
assert_eq!(async_error.to_string(), sync_error.to_string());
}
});
}
#[test]
fn route_input_validation_preserves_left_to_right_order() {
Python::initialize();
Python::attach(|py| {
let module = PyModule::new(py, "routes").expect("module should be created");
crate::routes::register(&module).expect("routes should register");
let invalid = PyList::empty(py);
let chat_kwargs = PyDict::new(py);
chat_kwargs
.set_item("optional_params", &invalid)
.expect("kwargs should accept optional_params");
chat_kwargs
.set_item("extra_headers", &invalid)
.expect("kwargs should accept extra_headers");
let invalid_messages = PyDict::new(py);
let error = module
.getattr("chat_completions")
.and_then(|function| {
function.call(("model", &invalid_messages), Some(&chat_kwargs))
})
.expect_err("messages should be validated first");
assert_eq!(error.to_string(), "ValueError: messages must be a list");
let valid_messages = PyList::empty(py);
let error = module
.getattr("chat_completions")
.and_then(|function| function.call(("model", &valid_messages), Some(&chat_kwargs)))
.expect_err("optional_params should be validated before headers");
assert_eq!(
error.to_string(),
"ValueError: optional_params must be a dict"
);
let headers_kwargs = PyDict::new(py);
headers_kwargs
.set_item("extra_headers", &invalid)
.expect("kwargs should accept extra_headers");
let invalid_body = PyList::empty(py);
let error = module
.getattr("messages")
.and_then(|function| function.call(("model", &invalid_body), Some(&headers_kwargs)))
.expect_err("body should be validated before headers");
assert_eq!(error.to_string(), "ValueError: body must be a dict");
let invalid_payload =
PyModule::new(py, "invalid_payload").expect("invalid payload should be created");
for name in ["ocr", "transcription"] {
let error = module
.getattr(name)
.and_then(|function| {
function.call(("model", &invalid_payload), Some(&headers_kwargs))
})
.expect_err("payload should be validated before headers");
assert!(!error.to_string().contains("extra_headers"));
}
for field in ('metadata', 'litellm_metadata', 'request_metadata_fields'):
invalid_context = replace(context, **{field: object()})
try:
routes.chat_completions(Request(messages=[], optional_params={}), context=invalid_context)
except (ValueError, TypeError) as error:
assert field in str(error)
else:
raise AssertionError('invalid context reached execution')
", Some(&locals), Some(&locals)).expect("native input validation should match");
});
}
@ -398,14 +262,30 @@ mod tests {
let sync_value: String = module
.getattr("echo")
.and_then(|function| function.call1(("sync",)))
.and_then(|function| {
let locals = crate::marshal::request_fixtures(py);
py.eval(c"Request(value=\"sync\")", Some(&locals), Some(&locals))
.and_then(|request| {
let kwargs = PyDict::new(py);
kwargs.set_item("context", locals.get_item("context")?.unwrap())?;
function.call((request,), Some(&kwargs))
})
})
.and_then(|value| value.extract())
.expect("sync route should return its value");
assert_eq!(sync_value, "sync");
let sync_error = module
.getattr("echo")
.and_then(|function| function.call1(("error",)))
.and_then(|function| {
let locals = crate::marshal::request_fixtures(py);
py.eval(c"Request(value=\"error\")", Some(&locals), Some(&locals))
.and_then(|request| {
let kwargs = PyDict::new(py);
kwargs.set_item("context", locals.get_item("context")?.unwrap())?;
function.call((request,), Some(&kwargs))
})
})
.expect_err("sync route should map its error");
assert!(sync_error.is_instance_of::<PyLookupError>(py));
assert_eq!(
@ -413,7 +293,7 @@ mod tests {
"LookupError: invalid request: synthetic error"
);
let locals = PyDict::new(py);
let locals = crate::marshal::request_fixtures(py);
locals
.set_item("routes", &module)
.expect("module should enter Python locals");
@ -422,17 +302,17 @@ mod tests {
import asyncio
async def exercise():
assert await routes.aecho("async") == "async"
assert await routes.aecho(Request(value="async"), context=context) == "async"
try:
await routes.aecho("error")
await routes.aecho(Request(value="error"), context=context)
except LookupError as error:
assert str(error) == "invalid request: synthetic error"
else:
raise AssertionError("mapped error was not raised")
try:
await routes.aecho("panic")
await routes.aecho(Request(value="panic"), context=context)
except BaseException as error:
assert type(error).__name__ == "PanicException"
assert str(error) == "synthetic panic"
@ -440,14 +320,14 @@ async def exercise():
raise AssertionError("panic was not raised")
try:
await routes.aecho("map_panic")
await routes.aecho(Request(value="map_panic"), context=context)
except BaseException as error:
assert type(error).__name__ == "PanicException"
assert str(error) == "synthetic mapper panic"
else:
raise AssertionError("mapper panic was not raised")
task = asyncio.ensure_future(routes.aecho("pending"))
task = asyncio.ensure_future(routes.aecho(Request(value="pending"), context=context))
await asyncio.sleep(0)
task.cancel()
try:
@ -479,13 +359,13 @@ asyncio.run(exercise())
Python::attach(|py| {
let module = PyModule::new(py, "synthetic").expect("module should be created");
synthetic::register_trace(&module).expect("trace routes should register");
let locals = PyDict::new(py);
let locals = crate::marshal::request_fixtures(py);
locals
.set_item("routes", &module)
.expect("module should enter Python locals");
let code = CString::new(
r#"
result = routes.echo("traced")
result = routes.echo(Request(value="traced"), context=context)
assert result == {
"response": "traced",
"trace": [{"function": "execute_echo", "depth": 0}],

View file

@ -1,44 +1,36 @@
use crate::errors::core_error_to_pyerr;
use crate::marshal::{NativeRequestContext, NativeRequestOptions, required_value};
use litellm_core::Error;
use litellm_core::messages::messages as run_messages;
use litellm_core::messages::messages as run_route;
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
use litellm_core::request_context::LiteLlmRequestContext;
use pyo3::prelude::*;
use serde_json::Value;
use std::future::Future;
use crate::errors::core_error_to_pyerr;
use crate::marshal::{RouteOptions, RouteOptionsInputs, required_value};
#[derive(FromPyObject)]
struct MessagesInputs {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
body: Value,
options: NativeRequestOptions,
}
fn prepare_messages(
inputs: MessagesInputs,
input: MessagesInputs,
context: NativeRequestContext,
) -> PyResult<impl Future<Output = Result<AnthropicMessagesResponse, Error>> + Send + 'static> {
let body = required_value("body", inputs.body, Value::is_object, "dict")?;
let options = RouteOptions::from_python(RouteOptionsInputs {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: inputs.extra_headers,
timeout_seconds: inputs.timeout_seconds,
})?;
let context: LiteLlmRequestContext = context.into();
let body = required_value("body", input.body, Value::is_object, "dict")?;
Ok(async move {
let RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
} = options;
run_messages(MessagesRequest {
model: &model,
body,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
})
run_route(
MessagesRequest {
model: &input.model,
body,
options: input.options.into(),
},
&context,
)
.await
})
}
@ -46,20 +38,7 @@ fn prepare_messages(
bridge_route! {
sync = messages,
asynchronous = amessages,
inputs = MessagesInputs,
required = {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
body: serde_json::Value,
},
optional = {
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<serde_json::Value>,
timeout_seconds: Option<f64>,
},
request = MessagesInputs,
prepare = prepare_messages,
errors = core_error_to_pyerr,
}

View file

@ -1,50 +1,44 @@
use crate::errors::ocr_error_to_pyerr;
use crate::marshal::{NativeRequestContext, NativeRequestOptions};
use litellm_ai_gateway::integrations::types::RequestHooks;
use litellm_ai_gateway::io::ocr::OcrRequest;
use litellm_ai_gateway::io::ocr::ocr as run_route;
use litellm_core::Error;
use litellm_core::request_context::LiteLlmRequestContext;
use pyo3::prelude::*;
use serde_json::{Map, Value};
use std::future::Future;
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
use pyo3::prelude::*;
use serde_json::Value;
use crate::errors::ocr_error_to_pyerr;
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
#[derive(FromPyObject)]
struct OcrInputs {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
document: Value,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
optional_params: Map<String, Value>,
options: NativeRequestOptions,
}
fn prepare_ocr(
inputs: OcrInputs,
input: OcrInputs,
context: NativeRequestContext,
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
let document = inputs.document;
let options = RouteOptions::from_python(RouteOptionsInputs {
model: inputs.model,
api_key: inputs.api_key,
api_base: inputs.api_base,
custom_llm_provider: inputs.custom_llm_provider,
extra_headers: inputs.extra_headers,
timeout_seconds: inputs.timeout_seconds,
})?;
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
let context: LiteLlmRequestContext = context.into();
let document = input.document;
Ok(async move {
let RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
} = options;
run_ocr(OcrRequest {
model: &model,
document,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
})
run_route(
OcrRequest {
model: &input.model,
document,
optional_params: input.optional_params,
options: input.options.into(),
},
&context,
RequestHooks {
callbacks: Vec::new(),
guardrails: Vec::new(),
},
)
.await
})
}
@ -52,22 +46,7 @@ fn prepare_ocr(
bridge_route! {
sync = ocr,
asynchronous = aocr,
inputs = OcrInputs,
required = {
model: String,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
document: serde_json::Value,
},
optional = {
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
extra_headers: Option<serde_json::Value>,
#[pyo3(from_py_with = litellm_python_interop::from_py)]
optional_params: Option<serde_json::Value>,
timeout_seconds: Option<f64>,
},
request = OcrInputs,
prepare = prepare_ocr,
errors = ocr_error_to_pyerr,
}

View file

@ -8,6 +8,7 @@ import mimetypes
import os
import re
from collections.abc import Callable, Coroutine, Mapping
from dataclasses import dataclass
from io import IOBase
from typing import Any, Final, cast
@ -17,16 +18,24 @@ 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 (
NativeRequestOptions,
PreparedNativeCall,
provider_connection_params,
provider_request_params,
)
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager, client
@ -35,6 +44,28 @@ 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
_RUST_OCR_PROVIDERS: Final = {
"mistral",
"azure_ai",
"vertex_ai",
}
def _prepare_ocr_request(
model: str,
document: Mapping[str, object],
@ -44,7 +75,7 @@ def _prepare_ocr_request(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
kwargs: dict[str, object],
) -> rust_ocr_bridge.PreparedOCRRequest:
) -> _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 +172,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,
@ -156,71 +187,143 @@ def _prepare_ocr_request(
)
@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_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
return False
if not prepared_request.provider_config.supports_rust_bridge():
return False
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
def _rust_bridge_optional_params(
prepared_request: _PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
) -> dict[str, object]:
optional_params: Final = dict(prepared_request.optional_params)
if prepared_request.custom_llm_provider == "vertex_ai":
vertex_project: Final = (
prepared_request.litellm_params.get("vertex_project")
or prepared_request.litellm_params.get("vertex_ai_project")
or litellm.vertex_project
or resolve_secret("VERTEXAI_PROJECT")
)
vertex_location: Final = (
prepared_request.litellm_params.get("vertex_location")
or prepared_request.litellm_params.get("vertex_ai_location")
or litellm.vertex_location
or resolve_secret("VERTEXAI_LOCATION")
or resolve_secret("VERTEX_LOCATION")
)
if vertex_project is not None:
optional_params["vertex_project"] = vertex_project
if vertex_location is not None:
optional_params["vertex_location"] = vertex_location
return optional_params
def _rust_bridge_api_base(
prepared_request: _PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
) -> str | None:
if prepared_request.api_base is not None:
return prepared_request.api_base
if prepared_request.custom_llm_provider == "azure_ai":
if is_azure_document_intelligence_model(prepared_request.model):
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
return resolve_secret("AZURE_AI_API_BASE")
return None
def _prepare_rust_ocr_call(
prepared_request: _PreparedOCRRequest,
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
def _attempt_ocr(
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
is_async: bool,
) -> DispatchResult[OCRResponse]:
return rust_ocr_bridge.attempt_ocr(prepared_request=prepared_request, resolve_api_key=resolve_api_key)
@native_first(
native=_attempt_ocr,
route="ocr",
errors=lambda prepared_request, resolve_api_key, is_async: 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,
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
return base_llm_http_handler.ocr(
model=prepared_request.model,
document=prepared_request.document,
optional_params=prepared_request.optional_params,
timeout=prepared_request.effective_timeout,
logging_obj=prepared_request.litellm_logging_obj,
api_key=prepared_request.api_key,
resolved_complete_url: Final = provider_config.get_complete_url(
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,
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=provider_request_params(rust_optional_params),
options=NativeRequestOptions(
provider_connection=provider_connection_params(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),
),
)
)
def _run_rust_ocr(
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
fallback: Callable[[], OCRResponse | Coroutine[object, object, OCRResponse]],
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
return rust_ocr_bridge.dispatch_ocr(
prepare=lambda: _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
),
fallback=fallback,
adapt=OCRResponse.model_validate,
model=prepared_request.model,
provider=prepared_request.custom_llm_provider,
eligible=_rust_ocr_supported(prepared_request),
)
async def _run_rust_aocr(
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
fallback: Callable[[], Coroutine[object, object, OCRResponse]],
) -> OCRResponse:
return await rust_ocr_bridge.adispatch_ocr(
prepare=lambda: _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
),
fallback=fallback,
adapt=OCRResponse.model_validate,
model=prepared_request.model,
provider=prepared_request.custom_llm_provider,
eligible=_rust_ocr_supported(prepared_request),
)
@client
@ -319,7 +422,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 +687,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,

View file

@ -4,12 +4,16 @@ 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, Mapping, Sequence
from typing import TYPE_CHECKING, Final, Protocol
import httpx
@ -20,14 +24,28 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
convert_to_model_response_object,
)
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.configuration import rust_enabled
from litellm.rust_bridge.protocols import (
RustAchatCompletions,
RustChatCompletions,
RustChatCompletionsDecline,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
from litellm.rust_bridge.request import (
NativeChatCompletionsRequest,
NativeRequestContext,
NativeRequestOptions,
PreparedNativeCall,
call_native,
provider_connection_params,
provider_request_params,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointBinding,
EndpointDispatch,
async_none,
)
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.types.utils import ModelResponse
@ -85,10 +103,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,
)
@ -102,14 +126,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()
@ -178,24 +202,14 @@ def rust_chat_completions_accepts(
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
),
)
def _build_model_response(
@ -224,31 +238,32 @@ def chat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
on_response: ResponseObserver,
eligible: bool = True,
) -> 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, timeout_seconds: float | None) -> Mapping[str, object]:
return native(
model=model,
messages=messages,
optional_params=optional_params,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_seconds,
)
return attempt(
load=_CHAT.load,
enabled=rust_enabled(),
eligible=eligible,
prepare=lambda: timeout_to_seconds(timeout),
call=call,
return _CHAT.invoke(
prepare=lambda: PreparedNativeCall(
NativeChatCompletionsRequest(
model=model,
messages=messages,
optional_params=provider_request_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),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=lambda: None,
adapt=adapt,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
)
@ -264,29 +279,81 @@ async def achat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
on_response: ResponseObserver,
eligible: bool = True,
) -> 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, timeout_seconds: float | None) -> Mapping[str, object]:
return await native(
model=model,
messages=messages,
optional_params=optional_params,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_seconds,
)
return await aattempt(
load=_ACHAT.load,
enabled=rust_enabled(),
eligible=eligible,
prepare=lambda: timeout_to_seconds(timeout),
call=call,
return await _CHAT.ainvoke(
prepare=lambda: PreparedNativeCall(
NativeChatCompletionsRequest(
model=model,
messages=messages,
optional_params=provider_request_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),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
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]],
) -> 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=provider_request_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),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=python_fallback,
adapt=adapt,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
)

View file

@ -6,13 +6,30 @@ from typing import Final
import httpx
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.protocols import RustAmessages, RustMessages
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
from litellm.rust_bridge.request import (
NativeMessagesRequest,
NativeRequestContext,
NativeRequestOptions,
PreparedNativeCall,
call_native,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointDispatch,
always_enabled,
async_none,
identity,
)
from litellm.rust_bridge.timeouts import timeout_to_seconds
_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=always_enabled,
)
def set_rust_messages(
@ -22,22 +39,22 @@ def set_rust_messages(
) -> None:
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(
@ -49,22 +66,26 @@ def messages(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout: float | httpx.Timeout | None,
) -> DispatchResult[dict[str, object]]:
return attempt(
load=_MESSAGES.load,
enabled=True,
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_messages, timeout_seconds: rust_messages(
model=model,
body=body,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_seconds,
) -> dict[str, object] | None:
return _MESSAGES.invoke(
prepare=lambda: PreparedNativeCall(
NativeMessagesRequest(
model=model,
body=body,
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),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=lambda: None,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
)
@ -77,20 +98,24 @@ async def amessages(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout: float | httpx.Timeout | None,
) -> DispatchResult[dict[str, object]]:
return await aattempt(
load=_AMESSAGES.load,
enabled=True,
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_amessages, timeout_seconds: rust_amessages(
model=model,
body=body,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_seconds,
) -> dict[str, object] | None:
return await _MESSAGES.ainvoke(
prepare=lambda: PreparedNativeCall(
NativeMessagesRequest(
model=model,
body=body,
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),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=async_none,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
)

View file

@ -1,211 +1,90 @@
"""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
from . import configuration as _configuration
from .bindings import UNCHANGED, Unchanged
from .protocols import RustAocr, RustOcr
from .request import NativeOCRRequest, PreparedNativeCall, call_native
from .runtime import (
BridgeErrorContext,
EndpointDispatch,
)
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.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])
rust_ocr_enabled = _configuration.rust_ocr_enabled
rust = _configuration.rust
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
@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",
}
_OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native(
route="ocr",
sync=lambda native: native.ocr,
asynchronous=lambda native: native.aocr,
enabled=_configuration.rust_ocr_enabled,
)
def set_rust_ocr(
*,
ocr: RustOcr | None | Unchanged = UNCHANGED,
aocr: RustAocr | None | Unchanged = UNCHANGED,
) -> None:
if not isinstance(ocr, Unchanged):
if ocr is None:
_OCR.sync.reset()
else:
_OCR.sync.override(ocr)
if not isinstance(aocr, Unchanged):
if aocr is None:
_OCR.asynchronous.reset()
else:
_OCR.asynchronous.override(aocr)
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 _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,
) -> ResultT:
return _OCR.invoke(
prepare=prepare,
call=call_native,
fallback=fallback,
adapt=adapt,
error_context=BridgeErrorContext(provider=provider, model=model),
eligible=eligible,
)
def attempt_ocr(
prepared_request: PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> DispatchResult[OCRResponse]:
return attempt(
load=_OCR.load,
enabled=_configuration.rust_enabled(),
prepare=lambda: _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
),
call=lambda native, prepared: native(
model=prepared_request.model,
document=prepared_request.document,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
optional_params=prepared.optional_params,
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
),
adapt=OCRResponse.model_validate,
eligible=_rust_ocr_supported(prepared_request),
)
async def aattempt_ocr(
prepared_request: PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> DispatchResult[OCRResponse]:
return await aattempt(
load=_AOCR.load,
enabled=_configuration.rust_enabled(),
prepare=lambda: _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
),
call=lambda native, prepared: native(
model=prepared_request.model,
document=prepared_request.document,
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
optional_params=prepared.optional_params,
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
),
adapt=OCRResponse.model_validate,
eligible=_rust_ocr_supported(prepared_request),
async def adispatch_ocr(
*,
prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]],
fallback: Callable[[], Awaitable[ResultT]],
adapt: Callable[[Mapping[str, object]], ResultT],
model: str,
provider: str,
eligible: bool,
) -> ResultT:
return await _OCR.ainvoke(
prepare=prepare,
call=call_native,
fallback=fallback,
adapt=adapt,
error_context=BridgeErrorContext(provider=provider, model=model),
eligible=eligible,
)

View file

@ -3,33 +3,24 @@ from __future__ import annotations
from collections.abc import Awaitable, Mapping, Sequence
from typing import Protocol
from .request import (
NativeChatCompletionsRequest,
NativeFunction,
NativeMessagesRequest,
NativeOCRRequest,
NativeRequestContext,
NativeResponsesWebSocketRequest,
NativeTranscriptionRequest,
)
class RustChatCompletions(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
) -> Mapping[str, object]: ...
class RustAchatCompletions(Protocol):
def __call__(
self,
model: str,
messages: Sequence[object],
optional_params: Mapping[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: Mapping[str, object] | None,
timeout_seconds: float | None,
) -> Awaitable[Mapping[str, object]]: ...
RustChatCompletions = NativeFunction[NativeChatCompletionsRequest, Mapping[str, object]]
RustAchatCompletions = NativeFunction[NativeChatCompletionsRequest, Awaitable[Mapping[str, object]]]
RustMessages = NativeFunction[NativeMessagesRequest, dict[str, object]]
RustAmessages = NativeFunction[NativeMessagesRequest, Awaitable[dict[str, object]]]
RustOcr = NativeFunction[NativeOCRRequest, dict[str, object]]
RustAocr = NativeFunction[NativeOCRRequest, Awaitable[dict[str, object]]]
RustTranscription = NativeFunction[NativeTranscriptionRequest, dict[str, object]]
RustAtranscription = NativeFunction[NativeTranscriptionRequest, Awaitable[dict[str, object]]]
class RustChatCompletionsDecline(Protocol):
@ -54,9 +45,9 @@ class RustResponsesWebSocketConnection(Protocol):
@classmethod
async def connect(
cls,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
request: NativeResponsesWebSocketRequest,
*,
context: NativeRequestContext,
) -> RustResponsesWebSocket: ...
@ -96,85 +87,3 @@ class NativeModule(Protocol):
@property
def atranscription(self) -> RustAtranscription: ...
class RustMessages(Protocol):
def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
) -> dict[str, object]: ...
class RustAmessages(Protocol):
def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]: ...
class RustOcr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]: ...
class RustAocr(Protocol):
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]: ...
class RustTranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]: ...
class RustAtranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]: ...

View file

@ -0,0 +1,122 @@
from __future__ import annotations
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from typing import Final, Generic, Protocol, TypeVar
@dataclass(frozen=True, slots=True)
class NativeRequestOptions:
api_key: str | None = None
api_base: str | None = None
custom_llm_provider: str | None = None
extra_headers: Mapping[str, object] | None = None
extra_query: Mapping[str, object] | None = None
timeout_seconds: float | None = None
provider_connection: Mapping[str, object] | None = None
@dataclass(frozen=True, slots=True)
class RequestAttribution:
user_api_key_hash: str | None = None
user_api_key_user_id: str | None = None
user_api_key_team_id: str | None = None
@dataclass(frozen=True, slots=True)
class NativeRequestContext:
metadata: Mapping[str, object] | None = None
litellm_metadata: Mapping[str, object] | None = None
request_metadata_fields: tuple[str, ...] = ()
litellm_call_id: str | None = None
request_model: str | None = None
attribution: RequestAttribution = RequestAttribution()
RequestT = TypeVar("RequestT")
RequestContraT = TypeVar("RequestContraT", contravariant=True)
ResultT = TypeVar("ResultT", covariant=True)
@dataclass(frozen=True, slots=True)
class PreparedNativeCall(Generic[RequestT]):
request: RequestT
context: NativeRequestContext = NativeRequestContext()
class NativeFunction(Protocol[RequestContraT, ResultT]):
def __call__(self, request: RequestContraT, *, context: NativeRequestContext) -> ResultT: ...
def call_native(native: NativeFunction[RequestT, ResultT], prepared: PreparedNativeCall[RequestT]) -> ResultT:
return native(prepared.request, context=prepared.context)
_PROVIDER_CONNECTION_FIELDS: Final = frozenset(
(
"aws_access_key_id",
"aws_secret_access_key",
"aws_session_token",
"aws_region_name",
"aws_session_name",
"aws_profile_name",
"aws_role_name",
"aws_web_identity_token",
"aws_sts_endpoint",
"aws_external_id",
"aws_bedrock_runtime_endpoint",
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
)
)
def provider_connection_params(params: Mapping[str, object]) -> dict[str, object]:
return { # mutable-ok: PyO3 boundary payload
key: value for key, value in params.items() if key in _PROVIDER_CONNECTION_FIELDS
}
def provider_request_params(params: Mapping[str, object]) -> dict[str, object]:
return { # mutable-ok: PyO3 boundary payload
key: value for key, value in params.items() if key not in _PROVIDER_CONNECTION_FIELDS
}
@dataclass(frozen=True, slots=True)
class NativeChatCompletionsRequest:
model: str
messages: Sequence[object]
optional_params: Mapping[str, object]
options: NativeRequestOptions
@dataclass(frozen=True, slots=True)
class NativeMessagesRequest:
model: str
body: dict[str, object]
options: NativeRequestOptions
@dataclass(frozen=True, slots=True)
class NativeOCRRequest:
model: str
document: dict[str, object]
optional_params: dict[str, object]
options: NativeRequestOptions
@dataclass(frozen=True, slots=True)
class NativeTranscriptionRequest:
model: str
audio: dict[str, object]
optional_params: dict[str, object]
options: NativeRequestOptions
@dataclass(frozen=True, slots=True)
class NativeResponsesWebSocketRequest:
url: str
options: NativeRequestOptions

View file

@ -2,24 +2,36 @@
from __future__ import annotations
from collections.abc import AsyncGenerator
from contextlib import AbstractAsyncContextManager, asynccontextmanager
from typing import Final
import httpx
from websockets.exceptions import ConnectionClosedOK
from litellm.rust_bridge.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,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result
from litellm.rust_bridge.request import (
NativeRequestContext,
NativeRequestOptions,
NativeResponsesWebSocketRequest,
PreparedNativeCall,
call_native,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointBinding,
async_none,
identity,
)
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,
)
@ -34,7 +46,7 @@ def set_rust_responses_websocket(
_RESPONSES_WEBSOCKET.override(connection)
class ConnectionAdapter:
class _ConnectionAdapter:
def __init__(self, connection: RustResponsesWebSocket):
self._connection: Final[RustResponsesWebSocket] = connection
@ -56,34 +68,18 @@ async def connect(
url: str,
headers: dict[str, str],
timeout: float | httpx.Timeout | None,
) -> DispatchResult[ConnectionAdapter]:
return await aattempt(
load=_RESPONSES_WEBSOCKET.load,
enabled=rust_enabled(),
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda connection_type, timeout_seconds: connection_type.connect(
url=url,
headers=headers,
timeout_seconds=timeout_seconds,
) -> _ConnectionAdapter | None:
connection: Final = await _RESPONSES_WEBSOCKET.ainvoke(
prepare=lambda: PreparedNativeCall(
NativeResponsesWebSocketRequest(
url=url,
options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)),
),
context=NativeRequestContext(),
),
adapt=ConnectionAdapter,
call=lambda connection_type, request: call_native(connection_type.connect, request),
fallback=async_none,
adapt=identity,
error_context=BridgeErrorContext(provider="openai", model="responses websocket"),
)
@asynccontextmanager
async def _connection_context(connection: ConnectionAdapter) -> AsyncGenerator[ConnectionAdapter, None]:
try:
yield connection
finally:
await connection.close()
async def managed_connect(
*,
url: str,
headers: dict[str, str],
timeout: float | httpx.Timeout | None,
) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]:
result: Final = await connect(url=url, headers=headers, timeout=timeout)
return adapt_result(result, _connection_context)
return None if connection is None else _ConnectionAdapter(connection)

View file

@ -4,38 +4,58 @@ from typing import Final
import httpx
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
from litellm.rust_bridge.request import (
NativeRequestContext,
NativeRequestOptions,
NativeTranscriptionRequest,
PreparedNativeCall,
call_native,
provider_connection_params,
provider_request_params,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointDispatch,
always_enabled,
async_none,
identity,
)
from litellm.rust_bridge.timeouts import timeout_to_seconds
_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,
)
def configure_rust_transcription(
enabled: bool = True,
*,
transcription: RustTranscription | None | Unchanged = UNCHANGED,
atranscription: RustAtranscription | None | Unchanged = UNCHANGED,
) -> None:
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(
@ -48,23 +68,28 @@ def transcription(
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> DispatchResult[dict[str, object]]:
return attempt(
load=_TRANSCRIPTION.load,
enabled=True,
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_transcription, timeout_seconds: rust_transcription(
model=model,
audio=audio,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_seconds,
) -> dict[str, object] | None:
return _TRANSCRIPTION.invoke(
prepare=lambda: PreparedNativeCall(
NativeTranscriptionRequest(
model=model,
audio=audio,
optional_params=provider_request_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),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=lambda: None,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
)
@ -78,21 +103,26 @@ async def atranscription(
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
) -> DispatchResult[dict[str, object]]:
return await aattempt(
load=_ATRANSCRIPTION.load,
enabled=True,
eligible=True,
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_atranscription, timeout_seconds: rust_atranscription(
model=model,
audio=audio,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_seconds,
) -> dict[str, object] | None:
return await _TRANSCRIPTION.ainvoke(
prepare=lambda: PreparedNativeCall(
NativeTranscriptionRequest(
model=model,
audio=audio,
optional_params=provider_request_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),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=async_none,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
)

View file

@ -60,6 +60,55 @@ def _collect(
return tuple(profiler.events)
def _native_kwargs(route: str, kwargs: dict[str, object]) -> dict[str, object]:
from pydantic import TypeAdapter
from litellm.rust_bridge.chat_completions import NativeChatCompletionsRequest
from litellm.rust_bridge.messages import NativeMessagesRequest
from litellm.rust_bridge.ocr import NativeOCRRequest
from litellm.rust_bridge.request import (
NativeRequestContext,
NativeRequestOptions,
provider_connection_params,
provider_request_params,
)
from litellm.rust_bridge.transcription import NativeTranscriptionRequest
params: Final = TypeAdapter(dict[str, object]).validate_python(kwargs.get("optional_params", {}))
options: Final = TypeAdapter(NativeRequestOptions).validate_python(
{
**{
key: kwargs.get(key)
for key in (
"api_key",
"api_base",
"custom_llm_provider",
"extra_headers",
"extra_query",
"timeout_seconds",
)
},
"provider_connection": provider_connection_params(params),
}
)
payload: Final = {
key: value
for key, value in kwargs.items()
if key not in {"api_key", "api_base", "custom_llm_provider", "extra_headers", "timeout_seconds"}
}
request_type: Final = {
"chat_completions": NativeChatCompletionsRequest,
"messages": NativeMessagesRequest,
"ocr": NativeOCRRequest,
"transcription": NativeTranscriptionRequest,
"audio_transcription": NativeTranscriptionRequest,
}[route]
request: Final = TypeAdapter(request_type).validate_python(
{**payload, "optional_params": provider_request_params(params), "options": options}
)
return {"request": request, "context": NativeRequestContext()}
def collect_trace(
spec: RouteSpec, engine: Engine, *, asynchronous: bool
) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure:
@ -77,7 +126,12 @@ def collect_trace(
"api_base": provider.url,
**({"timeout_seconds": 5} if engine == "rust" else {"timeout": 5}),
}
events: Final = _collect(function, kwargs, engine, asynchronous=asynchronous)
events: Final = _collect(
function,
_native_kwargs(spec.route, kwargs) if engine == "rust" else kwargs,
engine,
asynchronous=asynchronous,
)
provider.take_requests(len(fixture.provider_responses))
except Exception as error:
return TraceExecutionFailure(engine, f"{type(error).__name__}: {error}")

View file

@ -61,9 +61,12 @@ def _fixture(engine: Engine, _base_url: str) -> RouteFixture:
}
audio: Final = _audio_bytes()
payload: Final = (
{"audio": {"data": base64.b64encode(audio).decode(), "format": "wav"}, "optional_params": credentials}
{
"audio": {"data": base64.b64encode(audio).decode(), "format": "wav"},
"optional_params": {**credentials, "language": "en"},
}
if engine == "rust"
else {"file": ("sample.wav", audio, "audio/wav"), **credentials}
else {"file": ("sample.wav", audio, "audio/wav"), "language": "en", **credentials}
)
response: Final = json.dumps(
{

View file

@ -9,7 +9,7 @@ import pytest
import litellm
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import configuration
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
from litellm.rust_bridge.request import NativeMessagesRequest, NativeRequestContext
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
@ -41,23 +41,19 @@ class RecordingMessages:
def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
request: NativeMessagesRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"body": body,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"timeout_seconds": timeout_seconds,
"model": request.model,
"body": request.body,
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"timeout_seconds": request.options.timeout_seconds,
}
)
return dict(FAKE_MESSAGES_RESPONSE)
@ -69,23 +65,19 @@ class RecordingAsyncMessages:
async def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
request: NativeMessagesRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"body": body,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"timeout_seconds": timeout_seconds,
"model": request.model,
"body": request.body,
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"timeout_seconds": request.options.timeout_seconds,
}
)
return dict(FAKE_MESSAGES_RESPONSE)
@ -95,7 +87,7 @@ class ExplodingAsyncMessages:
def __init__(self) -> None:
self.calls = 0
async def __call__(self, **kwargs: object) -> dict[str, object]:
async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]:
self.calls += 1
raise AssertionError("bridge must not be called")
@ -104,7 +96,7 @@ class RaisingAsyncMessages:
def __init__(self) -> None:
self.calls = 0
async def __call__(self, **kwargs: object) -> dict[str, object]:
async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]:
self.calls += 1
raise RuntimeError("upstream request failed with status 400: bad request")
@ -127,6 +119,16 @@ def test_load_rust_messages_returns_injected_impl():
assert rust_messages.load_rust_messages() is bridge
def test_bare_rust_still_toggles_ocr():
from litellm.rust_bridge.ocr import rust_ocr_enabled
litellm.rust(True)
assert rust_ocr_enabled() is True
litellm.rust(False)
assert rust_ocr_enabled() is False
def test_load_rust_amessages_returns_injected_impl():
bridge = RecordingAsyncMessages()
litellm.rust(True)
@ -134,7 +136,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",
@ -151,7 +153,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():
@ -169,7 +171,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,
@ -197,7 +199,7 @@ 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
@ -215,7 +217,7 @@ def _gate(**overrides):
"timeout": 30.0,
}
kwargs.update(overrides)
return BaseLLMHTTPHandler._attempt_rust_anthropic_messages(**kwargs)
return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs)
@pytest.mark.asyncio
@ -226,8 +228,7 @@ async def test_gate_invokes_rust_and_marks_response_header():
response = await _gate()
assert isinstance(response, Handled)
response = response.value
assert response is not None
assert response["id"] == "msg_123"
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
call = bridge.calls[0]
@ -240,13 +241,13 @@ async def test_gate_invokes_rust_and_marks_response_header():
@pytest.mark.asyncio
async def test_gate_reports_failure_to_harness():
async def test_gate_propagates_unknown_native_errors():
bridge = RaisingAsyncMessages()
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate()
assert isinstance(response, NativeFailed)
with pytest.raises(RuntimeError, match="bad request"):
await _gate()
assert bridge.calls == 1
@ -257,7 +258,7 @@ async def test_gate_skips_rust_when_flag_absent():
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
assert isinstance(response, NativeSkipped)
assert response is None
assert bridge.calls == 0
@ -269,11 +270,22 @@ async def test_gate_uses_process_enable_without_request_override():
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
assert isinstance(response, Handled)
response = response.value
assert response is not None
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
@pytest.mark.asyncio
async def test_gate_ignores_request_flag_when_process_enabled():
bridge = RecordingAsyncMessages()
litellm.rust(True)
rust_messages.set_rust_messages(amessages=bridge)
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False))
assert response is not None
assert len(bridge.calls) == 1
@pytest.mark.asyncio
async def test_gate_invokes_rust_for_native_anthropic_provider():
bridge = RecordingAsyncMessages()
@ -288,8 +300,7 @@ async def test_gate_invokes_rust_for_native_anthropic_provider():
headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"},
)
assert isinstance(response, Handled)
response = response.value
assert response is not None
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"
@ -306,8 +317,7 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch):
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
)
assert isinstance(response, Handled)
response = response.value
assert response is not None
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
@ -322,7 +332,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
)
assert isinstance(response, NativeSkipped)
assert response is None
assert bridge.calls == 0
@ -334,7 +344,7 @@ async def test_gate_skips_rust_for_unsupported_provider():
response = await _gate(custom_llm_provider="openai")
assert isinstance(response, NativeSkipped)
assert response is None
assert bridge.calls == 0
@ -346,7 +356,7 @@ async def test_gate_skips_rust_for_agentic_hook():
response = await _gate(has_agentic_hook=True)
assert isinstance(response, NativeSkipped)
assert response is None
assert bridge.calls == 0
@ -362,8 +372,7 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag():
request_body=streaming_body,
)
assert isinstance(response, Handled)
response = response.value
assert response is not None
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
@ -396,105 +405,4 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch):
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, **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 is None

View file

@ -2309,8 +2309,19 @@ class TestRustChatCompletionsHook:
seen["gate"].append(kwargs)
return decline_reason
def native(**kwargs):
seen["call"].append(kwargs)
def native(request, *, context):
seen["call"].append(
{
"model": request.model,
"messages": request.messages,
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"timeout_seconds": request.options.timeout_seconds,
}
)
if sync_error is not None:
raise sync_error
return dict(sync_result if sync_result is not None else self.RUST_RESPONSE)
@ -2464,7 +2475,7 @@ class TestRustChatCompletionsHook:
RustBridgeDeclined = _Declined
RustUpstreamError = type("_Upstream", (Exception,), {})
def declining_native(**_kwargs):
def declining_native(request, *, context):
raise _Declined("blank message text")
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
@ -2501,7 +2512,7 @@ class TestRustChatCompletionsHook:
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
async def declining_native(**_kwargs):
async def declining_native(request, *, context):
raise _Declined("blank message text")
bridge.set_rust_chat_completions(
@ -2528,7 +2539,7 @@ class TestRustChatCompletionsHook:
from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion
from litellm.rust_bridge import chat_completions as bridge
async def native(**_kwargs):
async def native(request, *, context):
return dict(self.RUST_RESPONSE)
bridge.set_rust_chat_completions(
@ -2561,7 +2572,7 @@ class TestRustChatCompletionsHook:
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
def declining_native(**_kwargs):
def declining_native(request, *, context):
raise _Declined("blank message text")
bridge.set_rust_chat_completions(

View file

@ -65,8 +65,19 @@ def _inject(*, decline_reason=None, error: Exception | None = None):
seen["gate"].append(kwargs)
return decline_reason
def native(**kwargs):
seen["call"].append(kwargs)
def native(request, *, context):
seen["call"].append(
{
"model": request.model,
"messages": request.messages,
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"timeout_seconds": request.options.timeout_seconds,
}
)
if error is not None:
raise error
return dict(RUST_RESPONSE)
@ -207,7 +218,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative())
async def declining_native(**_kwargs):
async def declining_native(request, *, context):
raise _Declined("blank message text")
bridge.set_rust_chat_completions(
@ -237,7 +248,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch):
@pytest.mark.asyncio
async def test_the_async_path_serves_the_rust_response_without_the_fallback():
async def native(**_kwargs):
async def native(request, *, context):
return dict(RUST_RESPONSE)
bridge.set_rust_chat_completions(
@ -271,7 +282,7 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines():
RustBridgeDeclined = _Declined
RustUpstreamError = type("_Upstream", (Exception,), {})
async def declining_native(**_kwargs):
async def declining_native(request, *, context):
raise _Declined("blank message text")
logging_obj = MagicMock()
@ -384,7 +395,7 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines():
RustBridgeDeclined = _Declined
RustUpstreamError = type("_Upstream", (Exception,), {})
def declining_native(**_kwargs):
def declining_native(request, *, context):
raise _Declined("blank message text")
logging_obj = MagicMock()
@ -438,7 +449,7 @@ async def test_post_call_logging_fires_on_the_async_rust_path():
cannot drift apart the way the pre_call suppression once did."""
import json
async def native(**_kwargs):
async def native(request, *, context):
return dict(RUST_RESPONSE)
bridge.set_rust_chat_completions(
@ -470,7 +481,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines():
RustBridgeDeclined = _Declined
RustUpstreamError = type("_Upstream", (Exception,), {})
def declining_native(**_kwargs):
def declining_native(request, *, context):
raise _Declined("blank message text")
logging_obj, calls = _recording_logging_obj()

View file

@ -12,7 +12,7 @@ import pytest
import litellm
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.rust_bridge import configuration
from litellm.rust_bridge.runtime import Handled
from litellm.rust_bridge.request import NativeOCRRequest, NativeRequestContext, NativeRequestOptions, PreparedNativeCall
from litellm.rust_bridge.timeouts import timeout_to_seconds
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
@ -49,25 +49,20 @@ class RecordingBridge:
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeOCRRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": timeout_seconds,
"model": request.model,
"document": request.document,
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
"timeout_seconds": request.options.timeout_seconds,
}
)
return dict(FAKE_OCR_RESPONSE)
@ -81,25 +76,20 @@ class RecordingAsyncBridge:
async def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeOCRRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": timeout_seconds,
"model": request.model,
"document": request.document,
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
"timeout_seconds": request.options.timeout_seconds,
}
)
return dict(FAKE_OCR_RESPONSE)
@ -108,14 +98,9 @@ class RecordingAsyncBridge:
class RaisingBridge:
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeOCRRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
raise RuntimeError("bridge failed")
@ -123,14 +108,9 @@ class RaisingBridge:
class RaisingAsyncBridge:
async def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeOCRRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
raise RuntimeError("bridge failed")
@ -164,9 +144,6 @@ class FakeOCRConfig:
def get_api_key_env_var(self) -> str:
return self.api_key_env_var
def supports_rust_bridge(self) -> bool:
return True
def validate_environment(
self,
*,
@ -203,7 +180,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,
@ -392,13 +369,103 @@ 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.set_rust_ocr(ocr=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]},
"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.set_rust_ocr(aocr=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",
provider_connection={"vertex_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_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)
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",
@ -409,8 +476,6 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
resolve_api_key=lambda _name: None,
)
assert isinstance(response, Handled)
response = response.value
assert isinstance(response, OCRResponse)
assert response.pages[0].markdown == "hello world"
assert bridge.calls[0] == {
@ -433,7 +498,8 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
litellm.rust(True)
rust_bridge._OCR.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,
)
@ -449,7 +515,8 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver():
def _resolver(name: str) -> str | None:
raise AssertionError(f"resolver should not be called for {name}")
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,
@ -470,7 +537,8 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
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",
@ -489,7 +557,8 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
litellm.rust(True)
rust_bridge._OCR.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",
@ -522,7 +591,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",
@ -540,7 +610,8 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
litellm.rust(True)
rust_bridge._OCR.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",
@ -558,7 +629,8 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
litellm.rust(True)
rust_bridge._OCR.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",
@ -579,7 +651,8 @@ def test_run_rust_ocr_runs_pre_call_logging():
litellm.rust(True)
rust_bridge._OCR.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",
@ -726,7 +799,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 = (
@ -739,25 +812,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,

View file

@ -4,7 +4,7 @@ import pytest
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
from litellm.rust_bridge import configuration, responses_websocket
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
from litellm.rust_bridge.request import NativeRequestContext, NativeResponsesWebSocketRequest
class _FakeNativeConnection:
@ -31,10 +31,9 @@ class _FakeNativeBridge:
@classmethod
async def connect(
cls,
request: NativeResponsesWebSocketRequest,
*,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
context: NativeRequestContext,
) -> _FakeNativeConnection:
return _FakeNativeConnection()
@ -58,22 +57,25 @@ def test_rust_websocket_bridge_uses_process_enablement() -> None:
@pytest.mark.asyncio
async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
adapter = responses_websocket.ConnectionAdapter(_ClosedNativeConnection())
adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection())
with pytest.raises(responses_websocket.ConnectionClosedOK):
await adapter.recv()
@pytest.mark.asyncio
async def test_bridge_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
@ -89,8 +91,7 @@ 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()
@ -100,72 +101,17 @@ class _FailingNativeBridge:
@classmethod
async def connect(
cls,
request: NativeResponsesWebSocketRequest,
*,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
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, *, url: str, headers: dict[str, str], timeout_seconds: float | None
) -> _FakeNativeConnection:
return socket
responses_websocket.set_rust_responses_websocket(connection=Bridge)
result = await responses_websocket.managed_connect(url="wss://example.test/responses", headers={}, timeout=1.0)
assert isinstance(result, Handled)
async def use_connection() -> None:
async with result.value as connection:
await connection.send("hello")
raise ValueError("consumer failed")
with pytest.raises(ValueError, match="consumer failed"):
await use_connection()
assert socket.sent == ["hello"]
assert socket.closed
@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 run()
await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None)

View file

@ -10,6 +10,7 @@ import sys
import tempfile
import threading
import zipfile
from dataclasses import make_dataclass
from http.client import HTTPMessage
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
@ -79,6 +80,10 @@ def assert_native_request(
raise AssertionError(f"unexpected outcome marker: {outcome!r}")
if not isinstance(body, dict):
raise TypeError(f"{route} sent {type(body).__name__}, expected a JSON object")
assert all(
marker not in json.dumps(body)
for marker in ("must-not-reach-provider", "native-wheel-call", "native-user", "native-secret-key")
)
if route == "ocr":
assert path == "/v1/ocr"
assert headers.get("authorization") == "Bearer sk-native"
@ -123,7 +128,7 @@ def load_native(native_path: Path) -> object:
return native_module
def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
def _route_inputs(route: str, api_base: str, outcome: str) -> dict[str, object]:
common: Final = {
"api_base": api_base,
"extra_headers": {"x-test-outcome": outcome, "x-test-route": route},
@ -170,6 +175,44 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
raise AssertionError(f"unknown route: {route}")
def _record(name: str, fields: dict[str, object]) -> object:
record_type: Final = make_dataclass(name, tuple(fields), frozen=True, slots=True)
return record_type(**fields)
def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
inputs: Final = _route_inputs(route, api_base, outcome)
params: Final = inputs.get("optional_params", {})
connection_fields: Final = frozenset(("aws_access_key_id", "aws_secret_access_key", "aws_region_name"))
options: Final = _record("RequestOptions", {
"api_key": inputs.get("api_key"),
"api_base": inputs.get("api_base"),
"custom_llm_provider": inputs.get("custom_llm_provider"),
"extra_headers": inputs.get("extra_headers"),
"extra_query": None,
"timeout_seconds": inputs.get("timeout_seconds"),
"provider_connection": {key: value for key, value in params.items() if key in connection_fields},
})
request: Final = _record("Request", {
**{key: value for key, value in inputs.items() if key in {"model", "document", "audio", "body", "messages"}},
**({"optional_params": {key: value for key, value in params.items() if key not in connection_fields}} if route != "messages" else {}),
"options": options,
})
context: Final = _record("RequestContext", {
"metadata": None,
"litellm_metadata": {"internal_marker": "must-not-reach-provider"},
"request_metadata_fields": (),
"litellm_call_id": "native-wheel-call",
"request_model": inputs["model"],
"attribution": _record("Attribution", {
"user_api_key_hash": None,
"user_api_key_user_id": "native-user",
"user_api_key_team_id": None,
}),
})
return {"request": request, "context": context}
def assert_success(route: str, response: object) -> None:
if not isinstance(response, dict):
raise TypeError(f"{route} returned {type(response).__name__}, expected dict")

View file

@ -12,7 +12,6 @@ import pytest
import litellm
from litellm.rust_bridge import bindings, configuration
from litellm.rust_bridge import chat_completions as bridge
from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeFailed
from litellm.types.utils import ModelResponse
RUST_RESPONSE = {
@ -96,16 +95,16 @@ class _RecordingCall:
self.error = error
self.calls: list[dict] = []
def __call__(self, **kwargs):
self.calls.append(kwargs)
def __call__(self, request, *, context):
self.calls.append({"request": request, "context": context})
if self.error is not None:
raise self.error
return self.result
class _RecordingAsyncCall(_RecordingCall):
async def __call__(self, **kwargs):
return _RecordingCall.__call__(self, **kwargs)
async def __call__(self, request, *, context):
return _RecordingCall.__call__(self, request, context=context)
def _accepts(**overrides) -> bool:
@ -251,8 +250,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"
@ -266,16 +264,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]["request"].options.timeout_seconds == 30.0
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:
@ -283,18 +281,152 @@ 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,
"aws_access_key_id": "test-access-key",
"aws_secret_access_key": "test-secret-key",
}
bridge.chat_completions(**kwargs)
request = native.calls[0]["request"]
assert request.optional_params == {"max_tokens": 32}
assert request.options.provider_connection == {
"aws_access_key_id": "test-access-key",
"aws_secret_access_key": "test-secret-key",
}

View file

@ -4,7 +4,7 @@ import pytest
import litellm
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
from litellm.rust_bridge.runtime import Handled
from litellm.rust_bridge.request import NativeRequestContext, NativeTranscriptionRequest
rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
@ -15,37 +15,27 @@ class SyncBridge:
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeTranscriptionRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append({"model": model, "audio": audio, "optional_params": optional_params})
self.calls.append({"model": request.model, "audio": request.audio, "optional_params": {**request.optional_params, **(request.options.provider_connection or {})}})
return {"text": "hello"}
class AsyncBridge:
async def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeTranscriptionRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
return {"text": "async"}
def test_enabled_sync_bridge_receives_audio() -> None:
bridge = SyncBridge()
rust_bridge.configure_rust_transcription(transcription=bridge)
rust_bridge.configure_rust_transcription(True, transcription=bridge)
result = rust_bridge.transcription(
model="mistral.voxtral-mini-3b-2507",
audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"},
@ -56,14 +46,13 @@ def test_enabled_sync_bridge_receives_audio() -> None:
optional_params={"temperature": 0},
timeout=5.0,
)
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"}
@pytest.mark.asyncio
async def test_enabled_async_bridge() -> None:
rust_bridge.configure_rust_transcription(atranscription=AsyncBridge())
rust_bridge.configure_rust_transcription(True, atranscription=AsyncBridge())
result = await rust_bridge.atranscription(
model="mistral.voxtral-mini-3b-2507",
audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"},
@ -74,7 +63,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:
@ -85,8 +74,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(
@ -103,8 +91,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(
@ -121,7 +111,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 **_: {"text": "rust"},
transcription=lambda request, *, context: {"text": "rust"},
atranscription=None,
)
try:
@ -137,7 +127,7 @@ 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(**_: object) -> dict[str, object]:
async def rust_response(request: NativeTranscriptionRequest, *, context: NativeRequestContext) -> dict[str, object]:
return {"text": "rust"}
rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response)

View file

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