This commit is contained in:
yujonglee 2026-09-06 07:52:31 +00:00 • committed by GitHub
commit e2dd92d18e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
68 changed files with 2581 additions and 1450 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, *, options, context)`. The request carries only endpoint payload.
`NativeRequestOptions` carries credentials, typed provider configuration, routing,
headers, query parameters, and timeout. `NativeRequestContext` carries call identity,
attribution, and typed capability facts separately from the provider payload.
Python builds the frozen request dataclasses in `litellm/rust_bridge/request.py` and
PyO3 extracts their fields before execution. AWS credentials and metadata policy
belong in `options.bedrock`; Vertex project/location belongs in `options.vertex`.
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,23 +94,29 @@ impl AudioTranscriptionLifecycleHooks {
custom_llm_provider,
audio,
api_key,
bedrock,
api_base,
extra_headers,
optional_params,
timeout,
..
} = request;
let provider_request =
prepare_audio_transcription_provider_call(CoreAudioTranscriptionRequest {
let provider_request = 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,
},
RequestOptions {
bedrock: Some(bedrock),
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,9 @@
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 litellm_core::request_options::RequestOptions;
use serde_json::Value;
mod hooks;
@ -11,11 +14,24 @@ 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<'_>,
options: &RequestOptions,
context: &LiteLlmRequestContext,
hooks: RequestHooks,
) -> Result<Value, Error> {
let PreparedAudioTranscriptionCall {
request,
context,
hooks,
} = prepare_audio_transcription_call(request, options.clone(), 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,7 @@
use crate::integrations::types::RequestHooks;
use litellm_core::call_lifecycle::CallLifecycleContext;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
@ -9,38 +13,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<'_>,
options: RequestOptions,
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, 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,
bedrock: options.bedrock.unwrap_or_default(),
api_key: options.api_key,
api_base: options.api_base,
extra_headers: options.extra_headers,
optional_params: request.optional_params,
timeout: request.timeout,
timeout: 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::{BedrockOptions, 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([
("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 bedrock = BedrockOptions {
aws_access_key_id: Some("access-key".to_string()),
aws_secret_access_key: Some("secret-key".to_string()),
aws_region_name: Some("us-east-1".to_string()),
..Default::default()
};
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(),
},
&RequestOptions {
bedrock: Some(bedrock),
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,22 @@
use std::sync::Arc;
use litellm_core::request_options::BedrockOptions;
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(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) bedrock: BedrockOptions,
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,14 @@
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::request_options::RequestOptions;
use litellm_core::responses::types::{ResponsesWebSocketRequest, ResponsesWsEvent};
use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig;
use tokio::net::TcpStream;
use tokio::sync::Mutex;
@ -33,24 +35,27 @@ 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,
options: &RequestOptions,
_context: &LiteLlmRequestContext,
) -> Result<Self, Error> {
let mut request = url
let headers = string_headers("Responses WebSocket", options.extra_headers.clone())?;
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 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.vertex.into_map())
.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,8 @@
use crate::integrations::types::RequestHooks;
use litellm_core::Error;
use litellm_core::call_lifecycle::CallLifecycle;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use serde_json::Value;
mod common_utils;
@ -14,10 +17,19 @@ 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<'_>,
options: &RequestOptions,
context: &LiteLlmRequestContext,
hooks: RequestHooks,
) -> Result<Value, Error> {
let PreparedOcrCall {
request,
context,
hooks,
} = prepare_ocr_call(request, options.clone(), context, hooks);
CallLifecycle::default()
.run_request(request, &hooks, |request| {
.run(context, request, &hooks, |request| {
execute_ocr_provider_call(request, &hooks)
})
.await
@ -25,12 +37,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();
@ -65,24 +79,25 @@ mod tests {
String::from_utf8(request).expect("request is utf8")
}
fn base_ocr_request(model: &str) -> OcrRequest<'_> {
OcrRequest {
model,
document: json!({
"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,
}
fn base_ocr_request(model: &str) -> (OcrRequest<'_>, RequestOptions) {
(
OcrRequest {
model,
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
},
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()
},
)
}
#[tokio::test]
@ -120,10 +135,10 @@ mod tests {
(upload_request, parse_request)
});
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([
let (mut request, mut options) = base_ocr_request("reducto/parse-v3");
options.api_base = Some(&api_base).map(|value| value.to_string());
options.api_key = None;
options.extra_headers = Some(Map::from_iter([
("Authorization".to_string(), json!("Bearer test-key")),
("x-trace-id".to_string(), json!("trace-1")),
]));
@ -140,7 +155,18 @@ mod tests {
("settings".to_string(), json!({"ocr_system": "standard"})),
]);
let response = ocr(request).await.expect("Reducto OCR succeeds");
let response = ocr(
request,
&options,
&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,7 @@
use crate::integrations::types::RequestHooks;
use litellm_core::call_lifecycle::CallLifecycleContext;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
@ -11,21 +15,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<'_>,
options: RequestOptions,
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, 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 +49,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,
vertex: options.vertex.unwrap_or_default(),
api_key: options.api_key,
api_base: options.api_base,
extra_headers: options.extra_headers,
optional_params,
timeout: request.timeout,
timeout: 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 +120,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 +135,7 @@ 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,
}
}
@ -147,7 +147,16 @@ 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"),
RequestOptions::default(),
&LiteLlmRequestContext {
..Default::default()
},
RequestHooks {
..Default::default()
},
);
assert!(
matches!(prepared.request.config, Err(Error::InvalidRequest(message)) if message.contains("not supported for provider"))
);
@ -155,7 +164,16 @@ 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"),
RequestOptions::default(),
&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,21 @@
use std::sync::Arc;
use std::time::Duration;
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest};
use litellm_core::ocr::transformation::OcrProviderConfig;
use litellm_core::request_options::VertexOptions;
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(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) vertex: VertexOptions,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
@ -37,17 +23,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};
@ -212,24 +216,21 @@ impl CustomGuardrail for RecordingOcrGuardrail {
}
}
fn base_ocr_request(model: &str) -> OcrRequest<'_> {
OcrRequest {
model,
document: json!({
"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,
}
fn base_ocr_request(model: &str) -> (OcrRequest<'_>, RequestOptions) {
(
OcrRequest {
model,
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
},
RequestOptions {
api_key: Some("sk-test".to_string()),
..Default::default()
},
)
}
#[tokio::test]
@ -240,15 +241,27 @@ async fn reducto_during_call_guardrail_blocks_before_upload() {
let address = listener.local_addr().expect("listener has local address");
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);
let (mut request, mut options) = base_ocr_request("reducto/parse-v3");
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,
&options,
&LiteLlmRequestContext {
..Default::default()
},
hooks,
)
.await
.expect_err("guardrail blocks upload");
assert!(matches!(error, Error::InvalidRequest(_)));
assert_eq!(guardrail.events(), vec!["async_moderation_hook"]);
@ -277,14 +290,23 @@ async fn reducto_upload_error_body_is_truncated() {
.expect("writes upload response");
});
let api_base = format!("http://{address}");
let mut request = base_ocr_request("reducto/parse-v3");
request.api_base = Some(&api_base);
let (mut request, mut options) = base_ocr_request("reducto/parse-v3");
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,
&options,
&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 +342,36 @@ 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(),
},
&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()
},
litellm_call_id: Some("ocr-call-1"),
})
&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()
},
RequestHooks {
callbacks: vec![logger.clone()],
guardrails: vec![guardrail.clone()],
},
)
.await
.expect("ocr request succeeds");
@ -388,23 +420,33 @@ 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(),
},
&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 +474,33 @@ 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(),
},
&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 +554,33 @@ 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(),
},
&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 +633,33 @@ 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(),
},
&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.bedrock.into_map(), &env_lookup),
&env_lookup,
)
.await?;

View file

@ -1,4 +1,6 @@
use crate::Error;
use crate::request_context::LiteLlmRequestContext;
use crate::request_options::RequestOptions;
mod client;
mod handler;
mod prepare;
@ -12,9 +14,16 @@ 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> {
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?)
.await
pub async fn audio_transcription(
request: AudioTranscriptionRequest<'_>,
options: &RequestOptions,
_context: &LiteLlmRequestContext,
) -> Result<Value, Error> {
execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(
request,
options.clone(),
)?)
.await
}
#[cfg(test)]

View file

@ -2,6 +2,7 @@ use crate::error::Error;
use crate::http_utils::{has_header, string_headers};
#[cfg(feature = "bedrock-auth")]
use crate::providers::bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG;
use crate::request_options::RequestOptions;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderConfig};
@ -20,30 +21,36 @@ fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProv
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub fn prepare_audio_transcription_provider_call(
request: AudioTranscriptionRequest<'_>,
options: RequestOptions,
) -> 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, options.custom_llm_provider.as_deref())
.or_else(|| {
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", options.extra_headers)?;
let bedrock = options.bedrock.unwrap_or_default();
let bedrock_options = bedrock.clone().into_map();
let auth = config.auth_strategy(&model, &bedrock_options, &env_lookup)?;
if matches!(auth, AudioTranscriptionAuth::Bearer)
&& !has_header(&headers, "authorization")
&& let Some(api_key) = request.api_key
&& let Some(api_key) = options.api_key.as_deref()
{
headers.push(("Authorization".to_string(), format!("Bearer {api_key}")));
}
@ -51,9 +58,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,
options.api_base.as_deref(),
&model,
&request.optional_params,
&bedrock_options,
&env_lookup,
)?;
let filtered_params = config.map_transcription_params(&request.optional_params);
@ -68,7 +75,7 @@ pub fn prepare_audio_transcription_provider_call(
upstream_headers: headers,
auth,
#[cfg(feature = "bedrock-auth")]
optional_params: request.optional_params,
timeout: request.timeout,
bedrock,
timeout: options.timeout,
})
}

View file

@ -1,3 +1,5 @@
use crate::request_context::LiteLlmRequestContext;
use crate::request_options::{BedrockOptions, RequestOptions};
use std::io::{Read, Write};
use std::net::TcpListener;
use std::thread;
@ -27,22 +29,32 @@ async fn bedrock_request_is_signed_and_contains_audio() {
stream.write_all(response).expect("response");
});
let optional_params = 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 bedrock = BedrockOptions {
aws_access_key_id: Some("access-key".to_string()),
aws_secret_access_key: Some("secret-key".to_string()),
aws_region_name: Some("us-east-1".to_string()),
..Default::default()
};
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(),
},
&RequestOptions {
bedrock: Some(bedrock),
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

@ -3,17 +3,14 @@ use std::time::Duration;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use crate::request_options::BedrockOptions;
use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderConfig};
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>,
}
#[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) bedrock: BedrockOptions,
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,15 +101,10 @@ 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 bedrock = request.bedrock.into_map();
let credentials = match host_supplied_credentials(&bedrock) {
Some(credentials) => credentials,
None => {
resolve_credentials(
aws_auth_config(&request.optional_params, &env_lookup),
&env_lookup,
)
.await?
}
None => resolve_credentials(aws_auth_config(&bedrock, &env_lookup), &env_lookup).await?,
};
let signature = sign_bedrock_post(
&request.url,

View file

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

View file

@ -1,3 +1,5 @@
use crate::request_context::LiteLlmRequestContext;
use crate::request_options::RequestOptions;
use serde_json::Value;
use crate::error::Error;
@ -39,9 +41,12 @@ 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"))?;
options: RequestOptions,
_context: &LiteLlmRequestContext,
) -> Result<ResolvedChatCompletionsRequest, Error> {
let (model, config) =
resolve_provider_config(request.model, options.custom_llm_provider.as_deref())
.map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?;
let messages =
parse_messages(request.messages).map_err(|_| Error::Declined("unreadable message list"))?;
if messages.is_empty() {
@ -55,25 +60,27 @@ 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,
})
}
#[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
.bedrock
.clone()
.unwrap_or_default()
.into_map(),
&env_lookup,
)?;
match &auth {
@ -117,16 +124,21 @@ 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
.bedrock
.clone()
.unwrap_or_default()
.into_map(),
&env_lookup,
)?;
let transformed =
@ -139,7 +151,7 @@ pub(super) fn prepare_provider_request(
body: transformed.body,
upstream_headers: headers,
auth,
optional_params: request.optional_params,
timeout: request.timeout,
bedrock: request.options.bedrock.unwrap_or_default(),
timeout: request.options.timeout,
})
}

View file

@ -1,3 +1,5 @@
use crate::request_context::LiteLlmRequestContext;
use crate::request_options::{BedrockOptions, RequestOptions};
use serde_json::{Map, Value, json};
use crate::error::Error;
@ -6,10 +8,21 @@ use super::prepare::{prepare_provider_request, resolve_request};
use super::transformation::ChatCompletionsAuth;
use super::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest};
struct TestChatCompletionsCall<'a> {
request: ChatCompletionsRequest<'a>,
options: RequestOptions,
}
fn prepare_chat_completions_call(
request: ChatCompletionsRequest<'_>,
call: TestChatCompletionsCall<'_>,
) -> Result<ProviderChatCompletionsRequest, Error> {
prepare_provider_request(resolve_request(request)?)
prepare_provider_request(resolve_request(
call.request,
call.options,
&LiteLlmRequestContext {
..Default::default()
},
)?)
}
fn request<'a>(
@ -17,25 +30,30 @@ fn request<'a>(
provider: Option<&'a str>,
messages: Value,
optional_params: Value,
) -> ChatCompletionsRequest<'a> {
ChatCompletionsRequest {
model,
messages,
optional_params: match optional_params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
) -> TestChatCompletionsCall<'a> {
TestChatCompletionsCall {
request: ChatCompletionsRequest {
model,
messages,
optional_params: match optional_params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
},
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()
},
api_key: Some("sk-test"),
api_base: None,
custom_llm_provider: provider,
extra_headers: None,
timeout: None,
}
}
/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers
/// carry resolved credentials), so unwrap the failure case by hand.
fn decline(request: ChatCompletionsRequest<'_>) -> Error {
fn decline(request: TestChatCompletionsCall<'_>) -> Error {
match prepare_chat_completions_call(request) {
Err(error) => error,
Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url),
@ -107,7 +125,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 +150,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 +186,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 +217,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 +279,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 +297,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 +330,16 @@ 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.bedrock = Some(BedrockOptions {
aws_access_key_id: Some("AKIDEXAMPLE".to_string()),
aws_secret_access_key: Some("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".to_string()),
..Default::default()
});
// 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 +386,16 @@ 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.bedrock = Some(BedrockOptions {
aws_access_key_id: Some("AKIDEXAMPLE".to_string()),
aws_secret_access_key: Some("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".to_string()),
..Default::default()
});
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 +420,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 +453,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"),
)]));
@ -595,7 +616,7 @@ mod round_trip {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use crate::chat_completions::chat_completions;
use crate::chat_completions::chat_completions as run_chat_completions;
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
@ -658,35 +679,52 @@ mod round_trip {
(format!("http://127.0.0.1:{port}/v1/messages"), handle)
}
fn call(api_base: &str, messages: Value, params: Value) -> ChatCompletionsRequest<'_> {
ChatCompletionsRequest {
model: "anthropic/claude-sonnet-4-5",
messages,
optional_params: match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
fn call(api_base: &str, messages: Value, params: Value) -> TestChatCompletionsCall<'_> {
TestChatCompletionsCall {
request: ChatCompletionsRequest {
model: "anthropic/claude-sonnet-4-5",
messages,
optional_params: match params {
Value::Object(map) => map,
other => panic!("params must be an object, got {other}"),
},
},
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()
},
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)),
}
}
async fn execute(
call: TestChatCompletionsCall<'_>,
context: &LiteLlmRequestContext,
) -> Result<super::super::types::ChatCompletionsResponse, Error> {
run_chat_completions(call.request, &call.options, context).await
}
const GOOD_BODY: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#;
#[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 = execute(
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 +762,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 = execute(
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 +785,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 = execute(
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 +808,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 = execute(
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 +834,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 = execute(
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::{BedrockOptions, RequestOptions};
use std::time::Duration;
use serde::{Deserialize, Serialize};
@ -15,22 +16,14 @@ 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(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 +34,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) bedrock: BedrockOptions,
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

@ -10,8 +10,9 @@ use super::types::{AnthropicMessagesResponse, MessagesRequest};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) async fn execute_messages_provider_call(
request: MessagesRequest<'_>,
options: crate::request_options::RequestOptions,
) -> Result<AnthropicMessagesResponse, Error> {
let request = prepare_provider_request(request)?;
let request = prepare_provider_request(request, options)?;
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
@ -44,8 +45,9 @@ pub(super) async fn execute_messages_provider_call(
pub(super) async fn execute_messages_provider_stream(
request: MessagesRequest<'_>,
options: crate::request_options::RequestOptions,
) -> Result<reqwest::Response, Error> {
let request = prepare_provider_request(request)?;
let request = prepare_provider_request(request, options)?;
if request.provider != ANTHROPIC_MESSAGES_PROVIDER {
return Err(Error::InvalidRequest(
"streaming messages is not supported for this provider".to_string(),

View file

@ -8,6 +8,8 @@
//! can splice the event stream to its own caller.
use crate::Error;
use crate::request_context::LiteLlmRequestContext;
use crate::request_options::RequestOptions;
mod client;
mod common_utils;
mod handler;
@ -19,12 +21,20 @@ 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> {
execute_messages_provider_call(request).await
pub async fn messages(
request: MessagesRequest<'_>,
options: &RequestOptions,
_context: &LiteLlmRequestContext,
) -> Result<AnthropicMessagesResponse, Error> {
execute_messages_provider_call(request, options.clone()).await
}
pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Response, Error> {
execute_messages_provider_stream(request).await
pub async fn messages_stream(
request: MessagesRequest<'_>,
options: &RequestOptions,
_context: &LiteLlmRequestContext,
) -> Result<reqwest::Response, Error> {
execute_messages_provider_stream(request, options.clone()).await
}
#[cfg(test)]

View file

@ -1,4 +1,5 @@
use crate::error::Error;
use crate::request_options::RequestOptions;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers};
@ -8,21 +9,24 @@ use serde_json::{Map, Value};
pub(super) fn prepare_provider_request(
request: MessagesRequest<'_>,
options: RequestOptions,
) -> 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, options.custom_llm_provider.as_deref())
.or_else(|| {
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 +34,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,
options.extra_headers,
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 +51,7 @@ pub(super) fn prepare_provider_request(
))
})?;
let url = config.complete_url(request.api_base, &model, &env_lookup)?;
let url = config.complete_url(options.api_base.as_deref(), &model, &env_lookup)?;
Ok(ProviderMessagesRequest {
provider: provider.to_string(),
@ -52,7 +60,7 @@ pub(super) fn prepare_provider_request(
url,
body,
upstream_headers: headers,
timeout: request.timeout,
timeout: 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)),
})
}),
},
&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"}]
}),
},
&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": []}),
},
&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": []}),
},
&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": []}),
},
&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": []}),
},
&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": []}),
},
&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": []}),
},
&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

@ -8,11 +8,6 @@ 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(super) struct ProviderMessagesRequest {

View file

@ -0,0 +1,28 @@
#[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 RequestCapabilities {
pub execution_mode: Option<String>,
pub stream: bool,
pub has_agentic_hook: bool,
pub has_custom_client: bool,
pub request_format: Option<String>,
pub input_source_kind: Option<String>,
pub native_response_format: bool,
pub websocket_mode: Option<String>,
pub requires_connection: bool,
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct LiteLlmRequestContext {
pub litellm_call_id: Option<String>,
pub trace_id: Option<String>,
pub request_model: Option<String>,
pub attribution: RequestAttribution,
pub capabilities: RequestCapabilities,
}

View file

@ -0,0 +1,83 @@
use std::time::Duration;
use serde_json::{Map, Value};
#[derive(Clone, Debug, Default)]
pub struct BedrockOptions {
pub aws_access_key_id: Option<String>,
pub aws_secret_access_key: Option<String>,
pub aws_session_token: Option<String>,
pub aws_region_name: Option<String>,
pub aws_session_name: Option<String>,
pub aws_profile_name: Option<String>,
pub aws_role_name: Option<String>,
pub aws_web_identity_token: Option<String>,
pub aws_sts_endpoint: Option<String>,
pub aws_external_id: Option<String>,
pub aws_bedrock_runtime_endpoint: Option<String>,
pub request_metadata_fields: Vec<String>,
pub request_metadata: Option<std::collections::BTreeMap<String, String>>,
}
impl BedrockOptions {
pub fn into_map(&self) -> Map<String, Value> {
[
("aws_access_key_id", self.aws_access_key_id.clone()),
("aws_secret_access_key", self.aws_secret_access_key.clone()),
("aws_session_token", self.aws_session_token.clone()),
("aws_region_name", self.aws_region_name.clone()),
("aws_session_name", self.aws_session_name.clone()),
("aws_profile_name", self.aws_profile_name.clone()),
("aws_role_name", self.aws_role_name.clone()),
(
"aws_web_identity_token",
self.aws_web_identity_token.clone(),
),
("aws_sts_endpoint", self.aws_sts_endpoint.clone()),
("aws_external_id", self.aws_external_id.clone()),
(
"aws_bedrock_runtime_endpoint",
self.aws_bedrock_runtime_endpoint.clone(),
),
]
.into_iter()
.filter_map(|(name, value)| value.map(|value| (name.to_string(), Value::String(value))))
.collect()
}
}
#[derive(Clone, Debug, Default)]
pub struct AnthropicOptions {
pub user_id: Option<String>,
}
#[derive(Clone, Debug, Default)]
pub struct VertexOptions {
pub project: Option<String>,
pub location: Option<String>,
}
impl VertexOptions {
pub fn into_map(&self) -> Map<String, Value> {
[
("vertex_project", self.project.clone()),
("vertex_location", self.location.clone()),
]
.into_iter()
.filter_map(|(name, value)| value.map(|value| (name.to_string(), Value::String(value))))
.collect()
}
}
#[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 bedrock: Option<BedrockOptions>,
pub anthropic: Option<AnthropicOptions>,
pub vertex: Option<VertexOptions>,
}

View file

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

View file

@ -7,12 +7,17 @@ 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,
}
#[pyclass]
struct ResponsesWebSocketConnection {
@ -22,18 +27,19 @@ struct ResponsesWebSocketConnection {
#[pymethods]
impl ResponsesWebSocketConnection {
#[classmethod]
#[pyo3(signature = (url, headers=None, timeout_seconds=None))]
#[pyo3(signature = (request, *, options, 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,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> PyResult<Bound<'py, PyAny>> {
let headers = marshal_headers(headers)?;
let timeout = optional_timeout(timeout_seconds);
let options: litellm_core::request_options::RequestOptions = options.into();
let context: litellm_core::request_context::LiteLlmRequestContext = context.into();
let request = ResponsesWebSocketRequest { url: request.url };
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout)
let inner = RustResponsesWebSocketConnection::connect(request, &options, &context)
.await
.map_err(core_error_to_pyerr)?;
Ok(ResponsesWebSocketConnection { inner })
@ -81,7 +87,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 +191,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 +203,24 @@ mod tests {
import asyncio
async def exercise():
connection = await native.ResponsesWebSocketConnection.connect(url)
for request, request_options, request_context, field in (
(Request(url=123), options, context, 'url'),
(Request(url=url), Options(extra_headers=[]), context, 'extra_headers'),
(Request(url=url), options, replace(context, litellm_call_id=123), 'litellm_call_id'),
(Request(url=url), options, replace(context, attribution=Attribution(user_api_key_user_id=123)), 'user_api_key_user_id'),
):
try:
native.ResponsesWebSocketConnection.connect(request, options=request_options, 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), options=options, 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,157 @@
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)]
struct NativeBedrockOptions {
aws_access_key_id: Option<String>,
aws_secret_access_key: Option<String>,
aws_session_token: Option<String>,
aws_region_name: Option<String>,
aws_session_name: Option<String>,
aws_profile_name: Option<String>,
aws_role_name: Option<String>,
aws_web_identity_token: Option<String>,
aws_sts_endpoint: Option<String>,
aws_external_id: Option<String>,
aws_bedrock_runtime_endpoint: Option<String>,
request_metadata_fields: Vec<String>,
request_metadata: Option<std::collections::BTreeMap<String, String>>,
}
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<NativeBedrockOptions> for litellm_core::request_options::BedrockOptions {
fn from(input: NativeBedrockOptions) -> Self {
Self {
aws_access_key_id: input.aws_access_key_id,
aws_secret_access_key: input.aws_secret_access_key,
aws_session_token: input.aws_session_token,
aws_region_name: input.aws_region_name,
aws_session_name: input.aws_session_name,
aws_profile_name: input.aws_profile_name,
aws_role_name: input.aws_role_name,
aws_web_identity_token: input.aws_web_identity_token,
aws_sts_endpoint: input.aws_sts_endpoint,
aws_external_id: input.aws_external_id,
aws_bedrock_runtime_endpoint: input.aws_bedrock_runtime_endpoint,
request_metadata_fields: input.request_metadata_fields,
request_metadata: input.request_metadata,
}
}
}
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)]
struct NativeAnthropicOptions {
user_id: Option<String>,
}
impl From<NativeAnthropicOptions> for litellm_core::request_options::AnthropicOptions {
fn from(input: NativeAnthropicOptions) -> Self {
Self {
user_id: input.user_id,
}
}
}
#[derive(FromPyObject)]
struct NativeVertexOptions {
project: Option<String>,
location: Option<String>,
}
impl From<NativeVertexOptions> for litellm_core::request_options::VertexOptions {
fn from(input: NativeVertexOptions) -> Self {
Self {
project: input.project,
location: input.location,
}
}
}
#[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>,
bedrock: Option<NativeBedrockOptions>,
anthropic: Option<NativeAnthropicOptions>,
vertex: Option<NativeVertexOptions>,
}
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),
bedrock: input.bedrock.map(Into::into),
anthropic: input.anthropic.map(Into::into),
vertex: input.vertex.map(Into::into),
}
}
}
#[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 {
litellm_call_id: Option<String>,
trace_id: Option<String>,
request_model: Option<String>,
attribution: NativeRequestAttribution,
capabilities: NativeRequestCapabilities,
}
#[derive(FromPyObject)]
struct NativeRequestCapabilities {
execution_mode: Option<String>,
stream: bool,
has_agentic_hook: bool,
has_custom_client: bool,
request_format: Option<String>,
input_source_kind: Option<String>,
native_response_format: bool,
websocket_mode: Option<String>,
requires_connection: bool,
}
impl From<NativeRequestContext> for litellm_core::request_context::LiteLlmRequestContext {
fn from(input: NativeRequestContext) -> Self {
Self {
litellm_call_id: input.litellm_call_id,
trace_id: input.trace_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,
},
capabilities: litellm_core::request_context::RequestCapabilities {
execution_mode: input.capabilities.execution_mode,
stream: input.capabilities.stream,
has_agentic_hook: input.capabilities.has_agentic_hook,
has_custom_client: input.capabilities.has_custom_client,
request_format: input.capabilities.request_format,
input_source_kind: input.capabilities.input_source_kind,
native_response_format: input.capabilities.native_response_format,
websocket_mode: input.capabilities.websocket_mode,
requires_connection: input.capabilities.requires_connection,
},
}
}
}
@ -50,30 +169,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 +179,89 @@ 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
bedrock: object = None
anthropic: object = None
vertex: object = None
@dataclass(frozen=True)
class BedrockOptions:
aws_access_key_id: object = None
aws_secret_access_key: object = None
aws_session_token: object = None
aws_region_name: object = None
aws_session_name: object = None
aws_profile_name: object = None
aws_role_name: object = None
aws_web_identity_token: object = None
aws_sts_endpoint: object = None
aws_external_id: object = None
aws_bedrock_runtime_endpoint: object = None
request_metadata_fields: object = ()
request_metadata: object = None
@dataclass(frozen=True)
class Capabilities:
execution_mode: object = None
stream: object = False
has_agentic_hook: object = False
has_custom_client: object = False
request_format: object = None
input_source_kind: object = None
native_response_format: object = False
websocket_mode: object = None
requires_connection: object = False
@dataclass(frozen=True)
class VertexOptions:
project: object = None
location: 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:
litellm_call_id: object = None
trace_id: object = None
request_model: object = None
attribution: Attribution = Attribution()
capabilities: Capabilities = Capabilities()
@dataclass(frozen=True)
class Request:
model: str = 'model'
messages: object = None
body: object = None
audio: object = None
document: object = None
optional_params: object = None
value: str = ''
url: str = ''
context = Context()
options = Options()
",
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>,
}
fn prepare_transcription(
inputs: AudioTranscriptionInputs,
input: AudioTranscriptionInputs,
options: NativeRequestOptions,
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.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>,
}
fn prepare_chat_completions(
inputs: ChatCompletionsInputs,
input: ChatCompletionsInputs,
options: NativeRequestOptions,
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.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,36 @@ 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, *, options, context))]
fn $sync_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
request: $inputs,
options: $crate::marshal::NativeRequestOptions,
context: $crate::marshal::NativeRequestContext,
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
let future = $prepare($inputs {
$($required_name,)*
$($optional_name),*
})?;
let future = $prepare(request, options, 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, *, options, context))]
fn $async_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
request: $inputs,
options: $crate::marshal::NativeRequestOptions,
context: $crate::marshal::NativeRequestContext,
) -> pyo3::PyResult<pyo3::Bound<'_, pyo3::PyAny>> {
let future = $prepare($inputs {
$($required_name,)*
$($optional_name),*
})?;
let future = $prepare(request, options, 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 +48,38 @@ 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, *, options, context))]
fn $sync_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
request: $inputs,
options: $crate::marshal::NativeRequestOptions,
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, options, 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, *, options, context))]
fn $async_name(
py: pyo3::Python<'_>,
$($(#[$required_attr])* $required_name: $required_type,)*
$($(#[$optional_attr])* $optional_name: $optional_type,)*
request: $inputs,
options: $crate::marshal::NativeRequestOptions,
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, options, 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 +105,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 +129,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 +145,8 @@ mod tests {
fn prepare_echo(
inputs: EchoInputs,
_options: crate::marshal::NativeRequestOptions,
_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 +187,17 @@ 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)",
),
("ocr", "aocr", "(request, *, options, context)"),
(
"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)",
"(request, *, options, context)",
),
("messages", "amessages", "(request, *, options, 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, *, options, context)",
),
];
@ -263,129 +220,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, request_options, expected in [
(('chat_completions', 'achat_completions'), Request(messages={}, optional_params={}), options, 'messages must be a list'),
(('messages', 'amessages'), Request(body=[]), options, 'body must be a dict'),
(('ocr', 'aocr'), Request(document={}, optional_params={}), Options(extra_headers=[]), 'extra_headers'),
(('transcription', 'atranscription'), Request(audio={}, optional_params={}), Options(timeout_seconds='bad'), 'timeout_seconds'),
(('transcription', 'atranscription'), Request(audio={}, optional_params={}), Options(bedrock=[]), 'bedrock'),
]:
errors = []
for name in names:
try:
getattr(routes, name)(request, options=request_options, 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 ('litellm_call_id', 'trace_id', 'request_model'):
invalid_context = replace(context, **{field: object()})
try:
routes.chat_completions(Request(messages=[], optional_params={}), options=options, 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 +271,32 @@ 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("options", locals.get_item("options")?.unwrap())?;
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("options", locals.get_item("options")?.unwrap())?;
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 +304,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 +313,17 @@ mod tests {
import asyncio
async def exercise():
assert await routes.aecho("async") == "async"
assert await routes.aecho(Request(value="async"), options=options, context=context) == "async"
try:
await routes.aecho("error")
await routes.aecho(Request(value="error"), options=options, 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"), options=options, context=context)
except BaseException as error:
assert type(error).__name__ == "PanicException"
assert str(error) == "synthetic panic"
@ -440,14 +331,14 @@ async def exercise():
raise AssertionError("panic was not raised")
try:
await routes.aecho("map_panic")
await routes.aecho(Request(value="map_panic"), options=options, 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"), options=options, context=context))
await asyncio.sleep(0)
task.cancel()
try:
@ -479,13 +370,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"), options=options, 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,
}
fn prepare_messages(
inputs: MessagesInputs,
input: MessagesInputs,
options: NativeRequestOptions,
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.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>,
}
fn prepare_ocr(
inputs: OcrInputs,
input: OcrInputs,
options: NativeRequestOptions,
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.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

@ -28,6 +28,7 @@ from litellm.llms.custom_httpx.http_handler import (
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors
from litellm.rust_bridge.request import anthropic_options, request_context
from litellm.rust_bridge.runtime import DispatchResult
from litellm.types.llms.anthropic import (
ContentBlockDelta,
@ -434,6 +435,11 @@ class AnthropicChatCompletion(BaseLLM):
api_key=api_key,
additional_args=rust_logging_args,
)
rust_context: Final = request_context(
logging_obj=logging_obj,
request_model=logging_obj.model,
litellm_params=litellm_params,
)
def native_completion() -> DispatchResult[ModelResponse]:
return rust_chat_completions_bridge.chat_completions(
@ -447,7 +453,11 @@ class AnthropicChatCompletion(BaseLLM):
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
anthropic=anthropic_options(litellm_params),
stream=bool(stream),
has_custom_client=client is not None,
eligible=serves_via_rust,
context=rust_context,
)
async def native_acompletion() -> DispatchResult[ModelResponse]:
@ -462,7 +472,11 @@ class AnthropicChatCompletion(BaseLLM):
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
anthropic=anthropic_options(litellm_params),
stream=bool(stream),
has_custom_client=client is not None,
eligible=serves_via_rust,
context=rust_context,
)
@anative_first(

View file

@ -1,11 +1,14 @@
import base64
from io import IOBase
from typing import Final, NoReturn
import httpx
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.rust_bridge import transcription as rust_transcription_bridge
from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors
from litellm.rust_bridge.request import request_context
from litellm.rust_bridge.runtime import DispatchResult, adapt_result
from litellm.types.utils import FileTypes, TranscriptionResponse
@ -19,6 +22,17 @@ async def _aunavailable() -> NoReturn:
class BedrockAudioTranscriptionRustDispatch:
@staticmethod
def _input_source_kind(audio_file: FileTypes) -> str:
content: Final = audio_file[1] if isinstance(audio_file, tuple) else audio_file
if isinstance(content, (bytes, bytearray, memoryview)):
return "bytes"
if isinstance(content, IOBase):
return "file"
if isinstance(content, str):
return "path"
return "opaque"
@staticmethod
def _audio_payload(audio_file: FileTypes) -> dict[str, object]:
processed_audio: Final = process_audio_file(audio_file)
@ -52,6 +66,7 @@ class BedrockAudioTranscriptionRustDispatch:
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
logging_obj: Logging | None = None,
) -> DispatchResult[TranscriptionResponse]:
result: Final = rust_transcription_bridge.transcription(
model=model,
@ -62,13 +77,19 @@ class BedrockAudioTranscriptionRustDispatch:
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
input_source_kind=self._input_source_kind(audio_file),
context=request_context(
logging_obj=logging_obj,
request_model=logging_obj.model if logging_obj is not None else model,
litellm_params=logging_obj.litellm_params if logging_obj is not None else None,
),
)
return adapt_result(result, lambda response: TranscriptionResponse(**response))
@native_first(
native=_attempt_audio_transcriptions,
route="audio transcription",
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: (
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout, logging_obj=None: (
provider_errors(custom_llm_provider, model)
),
)
@ -83,6 +104,7 @@ class BedrockAudioTranscriptionRustDispatch:
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
logging_obj: Logging | None = None,
) -> TranscriptionResponse:
_unavailable()
@ -97,6 +119,7 @@ class BedrockAudioTranscriptionRustDispatch:
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
logging_obj: Logging | None = None,
) -> DispatchResult[TranscriptionResponse]:
result: Final = await rust_transcription_bridge.atranscription(
model=model,
@ -107,13 +130,19 @@ class BedrockAudioTranscriptionRustDispatch:
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
input_source_kind=self._input_source_kind(audio_file),
context=request_context(
logging_obj=logging_obj,
request_model=logging_obj.model if logging_obj is not None else model,
litellm_params=logging_obj.litellm_params if logging_obj is not None else None,
),
)
return adapt_result(result, lambda response: TranscriptionResponse(**response))
@anative_first(
native=_attempt_async_audio_transcriptions,
route="audio transcription",
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: (
errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout, logging_obj=None: (
provider_errors(custom_llm_provider, model)
),
)
@ -128,5 +157,6 @@ class BedrockAudioTranscriptionRustDispatch:
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
logging_obj: Logging | None = None,
) -> TranscriptionResponse:
await _aunavailable()

View file

@ -19,6 +19,7 @@ from litellm.llms.custom_httpx.http_handler import (
from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge
from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors
from litellm.rust_bridge.request import bedrock_options, request_context
from litellm.rust_bridge.runtime import DispatchResult
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
@ -424,6 +425,11 @@ class BedrockConverseLLM(BaseAWSLLM):
api_key="",
additional_args=rust_logging_args,
)
rust_context: Final = request_context(
logging_obj=logging_obj,
request_model=logging_obj.model,
litellm_params=litellm_params,
)
def native_completion() -> DispatchResult[ModelResponse]:
return rust_chat_completions_bridge.chat_completions(
@ -437,7 +443,11 @@ class BedrockConverseLLM(BaseAWSLLM):
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
bedrock=bedrock_options(rust_optional_params),
stream=bool(stream),
has_custom_client=client is not None,
eligible=serves_via_rust,
context=rust_context,
)
async def native_acompletion() -> DispatchResult[ModelResponse]:
@ -452,7 +462,11 @@ class BedrockConverseLLM(BaseAWSLLM):
extra_headers=headers,
timeout=timeout,
on_response=log_rust_post_call,
bedrock=bedrock_options(rust_optional_params),
stream=bool(stream),
has_custom_client=client is not None,
eligible=serves_via_rust,
context=rust_context,
)
@anative_first(

View file

@ -43,7 +43,7 @@ def _text_pairs(source: object) -> tuple[tuple[str, str], ...]:
return tuple((key, value) for key, value in source.items() if isinstance(key, str) and isinstance(value, str))
def _allowed_fields() -> tuple[str, ...]:
def get_bedrock_request_metadata_fields() -> tuple[str, ...]:
"""
The operator allow-list, deduplicated so a field repeated in config cannot consume a second
reserved slot and shrink the client budget for nothing. First occurrence wins, which keeps
@ -121,7 +121,7 @@ def resolve_bedrock_request_metadata(
been validated (and rejected with a 400) by the Converse transformation, so it is only
filtered here for the reserved identity prefix and the remaining slot budget.
"""
allowed_fields: Final = _allowed_fields()
allowed_fields: Final = get_bedrock_request_metadata_fields()
if not allowed_fields:
return None
sources: Final = _metadata_sources(litellm_params)
@ -146,7 +146,7 @@ def bedrock_request_metadata_is_owned() -> bool:
"fall back to whatever the caller supplied", or the reserved-prefix guarantee is bypassable
by anyone who can make the resolver produce nothing.
"""
return bool(_allowed_fields())
return bool(get_bedrock_request_metadata_fields())
def bedrock_request_metadata_headers(

View file

@ -2232,6 +2232,8 @@ class BaseLLMHTTPHandler:
custom_llm_provider=custom_llm_provider,
litellm_params=litellm_params,
has_agentic_hook=self._has_agentic_completion_hook(logging_obj),
stream=bool(stream),
has_custom_client=client is not None,
model=model,
api_key=api_key,
api_base=api_base,
@ -2242,6 +2244,7 @@ class BaseLLMHTTPHandler:
stream=stream or False,
custom_llm_provider=custom_llm_provider,
),
logging_obj=logging_obj,
)
return adapt_result(result, self._rust_anthropic_messages_fake_stream) if stream else result
@ -2388,12 +2391,15 @@ class BaseLLMHTTPHandler:
custom_llm_provider: str,
litellm_params: GenericLiteLLMParams,
has_agentic_hook: bool,
stream: bool,
has_custom_client: bool,
model: str,
api_key: str | None,
api_base: str | None,
headers: dict,
request_body: dict,
timeout: float | httpx.Timeout | None,
logging_obj: LiteLLMLoggingObj | None = None,
) -> DispatchResult[AnthropicMessagesResponse]:
if custom_llm_provider not in ("azure_ai", "anthropic"):
return NativeSkipped(NativeSkipReason.INELIGIBLE)
@ -2405,6 +2411,7 @@ class BaseLLMHTTPHandler:
return NativeSkipped(NativeSkipReason.INELIGIBLE)
from litellm.rust_bridge import messages as rust_messages_bridge
from litellm.rust_bridge.request import request_context
upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"}
result: Final = await rust_messages_bridge.amessages(
@ -2415,6 +2422,14 @@ class BaseLLMHTTPHandler:
custom_llm_provider=custom_llm_provider,
extra_headers=headers,
timeout=timeout,
stream=stream,
has_custom_client=has_custom_client,
has_agentic_hook=has_agentic_hook,
context=request_context(
logging_obj=logging_obj,
request_model=logging_obj.model if logging_obj is not None else model,
litellm_params=litellm_params.model_dump(),
),
)
def adapt(rust_response: dict[str, object]) -> AnthropicMessagesResponse:
@ -6502,6 +6517,7 @@ class BaseLLMHTTPHandler:
)
from litellm.rust_bridge import responses_websocket as rust_responses_websocket
from litellm.rust_bridge.request import request_context
async def attempt_connection() -> DispatchResult[
AbstractAsyncContextManager[rust_responses_websocket.ConnectionAdapter]
@ -6512,6 +6528,11 @@ class BaseLLMHTTPHandler:
url=ws_url,
headers={str(key): str(value) for key, value in headers.items()},
timeout=timeout,
context=request_context(
logging_obj=logging_obj,
request_model=logging_obj.model,
litellm_params=litellm_params.model_dump(),
),
)
@anative_context(

View file

@ -7895,6 +7895,7 @@ def transcription(
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
logging_obj=litellm_logging_obj,
)
else:
response = dispatch.audio_transcriptions(
@ -7906,6 +7907,7 @@ def transcription(
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
logging_obj=litellm_logging_obj,
)
elif provider_config is not None:
response = base_llm_http_handler.audio_transcriptions(

View file

@ -27,6 +27,17 @@ from litellm.rust_bridge.protocols import (
RustChatCompletions,
RustChatCompletionsDecline,
)
from litellm.rust_bridge.request import (
NativeAnthropicOptions,
NativeBedrockOptions,
NativeChatCompletionsRequest,
NativeRequestCapabilities,
NativeRequestContext,
NativeRequestOptions,
PreparedNativeCall,
call_native,
with_capabilities,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
from litellm.rust_bridge.timeouts import timeout_to_seconds
from litellm.types.utils import ModelResponse
@ -224,29 +235,46 @@ def chat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
on_response: ResponseObserver,
bedrock: NativeBedrockOptions | None = None,
anthropic: NativeAnthropicOptions | None = None,
stream: bool = False,
has_custom_client: bool = False,
eligible: bool = True,
context: NativeRequestContext | None = None,
) -> DispatchResult[ModelResponse]:
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,
)
def call(
native: RustChatCompletions, prepared: PreparedNativeCall[NativeChatCompletionsRequest]
) -> Mapping[str, object]:
return call_native(native, prepared)
return attempt(
load=_CHAT.load,
enabled=rust_enabled(),
eligible=eligible,
prepare=lambda: timeout_to_seconds(timeout),
prepare=lambda: PreparedNativeCall(
request=NativeChatCompletionsRequest(model=model, messages=messages, optional_params=optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
bedrock=bedrock,
anthropic=anthropic,
),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="sync",
stream=stream,
has_custom_client=has_custom_client,
),
),
),
call=call,
adapt=adapt,
)
@ -264,29 +292,47 @@ async def achat_completions(
extra_headers: Mapping[str, object] | None,
timeout: float | httpx.Timeout | None,
on_response: ResponseObserver,
bedrock: NativeBedrockOptions | None = None,
anthropic: NativeAnthropicOptions | None = None,
stream: bool = False,
has_custom_client: bool = False,
eligible: bool = True,
context: NativeRequestContext | None = None,
) -> DispatchResult[ModelResponse]:
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,
)
async def call(
native: RustAchatCompletions,
prepared: PreparedNativeCall[NativeChatCompletionsRequest],
) -> Mapping[str, object]:
return await call_native(native, prepared)
return await aattempt(
load=_ACHAT.load,
enabled=rust_enabled(),
eligible=eligible,
prepare=lambda: timeout_to_seconds(timeout),
prepare=lambda: PreparedNativeCall(
request=NativeChatCompletionsRequest(model=model, messages=messages, optional_params=optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
bedrock=bedrock,
anthropic=anthropic,
),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="async",
stream=stream,
has_custom_client=has_custom_client,
),
),
),
call=call,
adapt=adapt,
)

View file

@ -8,6 +8,15 @@ import httpx
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.protocols import RustAmessages, RustMessages
from litellm.rust_bridge.request import (
NativeMessagesRequest,
NativeRequestCapabilities,
NativeRequestContext,
NativeRequestOptions,
PreparedNativeCall,
call_native,
with_capabilities,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -49,21 +58,35 @@ def messages(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout: float | httpx.Timeout | None,
stream: bool = False,
has_custom_client: bool = False,
has_agentic_hook: bool = False,
context: NativeRequestContext | None = None,
) -> DispatchResult[dict[str, object]]:
return attempt(
load=_MESSAGES.load,
enabled=True,
eligible=True,
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,
prepare=lambda: PreparedNativeCall(
request=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=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="sync",
stream=stream,
has_custom_client=has_custom_client,
has_agentic_hook=has_agentic_hook,
),
),
),
call=call_native,
adapt=identity,
)
@ -77,20 +100,34 @@ async def amessages(
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout: float | httpx.Timeout | None,
stream: bool = False,
has_custom_client: bool = False,
has_agentic_hook: bool = False,
context: NativeRequestContext | None = None,
) -> DispatchResult[dict[str, object]]:
return await aattempt(
load=_AMESSAGES.load,
enabled=True,
eligible=True,
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,
prepare=lambda: PreparedNativeCall(
request=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=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="async",
stream=stream,
has_custom_client=has_custom_client,
has_agentic_hook=has_agentic_hook,
),
),
),
call=call_native,
adapt=identity,
)

View file

@ -14,6 +14,15 @@ from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, B
from litellm.rust_bridge import configuration as _configuration
from litellm.rust_bridge.bindings import NativeBinding
from litellm.rust_bridge.protocols import RustAocr, RustOcr
from litellm.rust_bridge.request import (
NativeOCRRequest,
NativeRequestCapabilities,
NativeRequestOptions,
PreparedNativeCall,
call_native,
request_context,
vertex_options,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -70,6 +79,21 @@ def _rust_ocr_supported(prepared_request: PreparedOCRRequest) -> bool:
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
def _ocr_input_source_kind(document: dict[str, object]) -> str:
if "document_url" in document:
return "document_url"
if "image_url" in document:
return "image_url"
if "file" in document:
return "file"
return "inline"
def _ocr_request_format(optional_params: dict[str, object]) -> str | None:
value = optional_params.get(OCR_REQUEST_FORMAT_PARAM)
return value if isinstance(value, str) else None
def _rust_bridge_optional_params(
prepared_request: PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
@ -170,15 +194,36 @@ def attempt_ocr(
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),
call=lambda native, prepared: call_native(
native,
PreparedNativeCall(
request=NativeOCRRequest(
model=prepared_request.model,
document=prepared_request.document,
optional_params=prepared.optional_params,
),
options=NativeRequestOptions(
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
vertex=vertex_options(prepared.optional_params),
),
context=request_context(
logging_obj=prepared_request.litellm_logging_obj,
request_model=prepared_request.model,
litellm_params=prepared_request.litellm_params,
capabilities=NativeRequestCapabilities(
execution_mode="sync",
input_source_kind=_ocr_input_source_kind(prepared_request.document),
request_format=_ocr_request_format(prepared_request.optional_params),
native_response_format=(
prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"
),
),
),
),
),
adapt=OCRResponse.model_validate,
eligible=_rust_ocr_supported(prepared_request),
@ -196,15 +241,36 @@ async def aattempt_ocr(
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),
call=lambda native, prepared: call_native(
native,
PreparedNativeCall(
request=NativeOCRRequest(
model=prepared_request.model,
document=prepared_request.document,
optional_params=prepared.optional_params,
),
options=NativeRequestOptions(
api_key=prepared.api_key,
api_base=prepared.api_base,
custom_llm_provider=prepared_request.custom_llm_provider,
extra_headers=prepared.headers,
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
vertex=vertex_options(prepared.optional_params),
),
context=request_context(
logging_obj=prepared_request.litellm_logging_obj,
request_model=prepared_request.model,
litellm_params=prepared_request.litellm_params,
capabilities=NativeRequestCapabilities(
execution_mode="async",
input_source_kind=_ocr_input_source_kind(prepared_request.document),
request_format=_ocr_request_format(prepared_request.optional_params),
native_response_format=(
prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"
),
),
),
),
),
adapt=OCRResponse.model_validate,
eligible=_rust_ocr_supported(prepared_request),

View file

@ -3,33 +3,25 @@ from __future__ import annotations
from collections.abc import Awaitable, Mapping, Sequence
from typing import Protocol
from .request import (
NativeChatCompletionsRequest,
NativeFunction,
NativeMessagesRequest,
NativeOCRRequest,
NativeRequestContext,
NativeRequestOptions,
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 +46,10 @@ class RustResponsesWebSocketConnection(Protocol):
@classmethod
async def connect(
cls,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
request: NativeResponsesWebSocketRequest,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> RustResponsesWebSocket: ...
@ -96,85 +89,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,206 @@
from __future__ import annotations
from collections.abc import Mapping, Sequence
from dataclasses import dataclass, replace
from types import MappingProxyType
from typing import Generic, Protocol, TypeVar
@dataclass(frozen=True, slots=True)
class NativeBedrockOptions:
aws_access_key_id: str | None = None
aws_secret_access_key: str | None = None
aws_session_token: str | None = None
aws_region_name: str | None = None
aws_session_name: str | None = None
aws_profile_name: str | None = None
aws_role_name: str | None = None
aws_web_identity_token: str | None = None
aws_sts_endpoint: str | None = None
aws_external_id: str | None = None
aws_bedrock_runtime_endpoint: str | None = None
request_metadata_fields: tuple[str, ...] = ()
request_metadata: Mapping[str, str] | None = None
@dataclass(frozen=True, slots=True)
class NativeAnthropicOptions:
user_id: str | None = None
@dataclass(frozen=True, slots=True)
class NativeVertexOptions:
project: str | None = None
location: str | None = None
def bedrock_options(params: Mapping[str, object]) -> NativeBedrockOptions:
def string(name: str) -> str | None:
value = params.get(name)
return value if isinstance(value, str) else None
return NativeBedrockOptions(
aws_access_key_id=string("aws_access_key_id"),
aws_secret_access_key=string("aws_secret_access_key"),
aws_session_token=string("aws_session_token"),
aws_region_name=string("aws_region_name"),
aws_session_name=string("aws_session_name"),
aws_profile_name=string("aws_profile_name"),
aws_role_name=string("aws_role_name"),
aws_web_identity_token=string("aws_web_identity_token"),
aws_sts_endpoint=string("aws_sts_endpoint"),
aws_external_id=string("aws_external_id"),
aws_bedrock_runtime_endpoint=string("aws_bedrock_runtime_endpoint"),
)
def anthropic_options(litellm_params: Mapping[str, object] | None) -> NativeAnthropicOptions:
metadata = None if litellm_params is None else litellm_params.get("metadata")
user_id = metadata.get("user_id") if isinstance(metadata, Mapping) else None
return NativeAnthropicOptions(user_id=user_id if isinstance(user_id, str) else None)
def vertex_options(params: Mapping[str, object]) -> NativeVertexOptions:
project = params.get("vertex_project") or params.get("vertex_ai_project")
location = params.get("vertex_location") or params.get("vertex_ai_location")
return NativeVertexOptions(
project=project if isinstance(project, str) else None,
location=location if isinstance(location, str) else None,
)
@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
bedrock: NativeBedrockOptions | None = None
anthropic: NativeAnthropicOptions | None = None
vertex: NativeVertexOptions | 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 NativeRequestCapabilities:
execution_mode: str | None = None
stream: bool = False
has_agentic_hook: bool = False
has_custom_client: bool = False
request_format: str | None = None
input_source_kind: str | None = None
native_response_format: bool = False
websocket_mode: str | None = None
requires_connection: bool = False
@dataclass(frozen=True, slots=True)
class NativeRequestContext:
litellm_call_id: str | None = None
trace_id: str | None = None
request_model: str | None = None
attribution: RequestAttribution = RequestAttribution()
capabilities: NativeRequestCapabilities = NativeRequestCapabilities()
def request_context(
*,
logging_obj: object | None,
request_model: str,
litellm_params: Mapping[str, object] | None = None,
capabilities: NativeRequestCapabilities | None = None,
) -> NativeRequestContext:
params = litellm_params if litellm_params is not None else MappingProxyType({})
metadata_value = params.get("metadata") or params.get("litellm_metadata")
metadata = metadata_value if isinstance(metadata_value, Mapping) else MappingProxyType({})
def string(name: str) -> str | None:
value = params.get(name, metadata.get(name))
return value if isinstance(value, str) else None
call_id = getattr(logging_obj, "litellm_call_id", None)
trace_id = getattr(logging_obj, "litellm_trace_id", None)
return NativeRequestContext(
litellm_call_id=call_id if isinstance(call_id, str) else None,
trace_id=trace_id if isinstance(trace_id, str) else None,
request_model=request_model,
attribution=RequestAttribution(
user_api_key_hash=string("user_api_key_hash"),
user_api_key_user_id=string("user_api_key_user_id"),
user_api_key_team_id=string("user_api_key_team_id"),
),
capabilities=capabilities or NativeRequestCapabilities(),
)
def with_capabilities(
context: NativeRequestContext,
capabilities: NativeRequestCapabilities,
) -> NativeRequestContext:
return replace(context, capabilities=capabilities)
RequestT = TypeVar("RequestT")
RequestContraT = TypeVar("RequestContraT", contravariant=True)
ResultT = TypeVar("ResultT", covariant=True)
@dataclass(frozen=True, slots=True)
class PreparedNativeCall(Generic[RequestT]):
request: RequestT
options: NativeRequestOptions = NativeRequestOptions()
context: NativeRequestContext = NativeRequestContext()
class NativeFunction(Protocol[RequestContraT, ResultT]):
def __call__(
self,
request: RequestContraT,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> ResultT: ...
def call_native(native: NativeFunction[RequestT, ResultT], prepared: PreparedNativeCall[RequestT]) -> ResultT:
return native(prepared.request, options=prepared.options, context=prepared.context)
@dataclass(frozen=True, slots=True)
class NativeChatCompletionsRequest:
model: str
messages: Sequence[object]
optional_params: Mapping[str, object]
@dataclass(frozen=True, slots=True)
class NativeMessagesRequest:
model: str
body: dict[str, object]
@dataclass(frozen=True, slots=True)
class NativeOCRRequest:
model: str
document: object
optional_params: dict[str, object]
@dataclass(frozen=True, slots=True)
class NativeTranscriptionRequest:
model: str
audio: object
optional_params: dict[str, object]
@dataclass(frozen=True, slots=True)
class NativeResponsesWebSocketRequest:
url: str

View file

@ -15,6 +15,15 @@ from litellm.rust_bridge.protocols import (
RustResponsesWebSocket,
RustResponsesWebSocketConnection,
)
from litellm.rust_bridge.request import (
NativeRequestCapabilities,
NativeRequestContext,
NativeRequestOptions,
NativeResponsesWebSocketRequest,
PreparedNativeCall,
call_native,
with_capabilities,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -56,17 +65,27 @@ async def connect(
url: str,
headers: dict[str, str],
timeout: float | httpx.Timeout | None,
websocket_mode: str = "native",
requires_connection: bool = True,
context: NativeRequestContext | None = 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,
prepare=lambda: PreparedNativeCall(
request=NativeResponsesWebSocketRequest(url=url),
options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="async",
websocket_mode=websocket_mode,
requires_connection=requires_connection,
),
),
),
call=lambda connection_type, prepared: call_native(connection_type.connect, prepared),
adapt=ConnectionAdapter,
)
@ -84,6 +103,16 @@ async def managed_connect(
url: str,
headers: dict[str, str],
timeout: float | httpx.Timeout | None,
websocket_mode: str = "managed",
requires_connection: bool = True,
context: NativeRequestContext | None = None,
) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]:
result: Final = await connect(url=url, headers=headers, timeout=timeout)
result: Final = await connect(
url=url,
headers=headers,
timeout=timeout,
websocket_mode=websocket_mode,
requires_connection=requires_connection,
context=context,
)
return adapt_result(result, _connection_context)

View file

@ -6,6 +6,16 @@ import httpx
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription
from litellm.rust_bridge.request import (
NativeRequestCapabilities,
NativeRequestContext,
NativeRequestOptions,
NativeTranscriptionRequest,
PreparedNativeCall,
bedrock_options,
call_native,
with_capabilities,
)
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -41,29 +51,43 @@ def load_rust_atranscription() -> RustAtranscription | None:
def transcription(
*,
model: str,
audio: dict[str, object],
audio: object,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
stream: bool = False,
has_custom_client: bool = False,
input_source_kind: str | None = None,
context: NativeRequestContext | None = None,
) -> DispatchResult[dict[str, object]]:
return attempt(
load=_TRANSCRIPTION.load,
enabled=True,
eligible=True,
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,
prepare=lambda: PreparedNativeCall(
request=NativeTranscriptionRequest(model=model, audio=audio, optional_params=optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
bedrock=bedrock_options(optional_params),
),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="sync",
stream=stream,
has_custom_client=has_custom_client,
input_source_kind=input_source_kind,
),
),
),
call=call_native,
adapt=identity,
)
@ -71,28 +95,42 @@ def transcription(
async def atranscription(
*,
model: str,
audio: dict[str, object],
audio: object,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: float | httpx.Timeout | None,
stream: bool = False,
has_custom_client: bool = False,
input_source_kind: str | None = None,
context: NativeRequestContext | None = None,
) -> DispatchResult[dict[str, object]]:
return await aattempt(
load=_ATRANSCRIPTION.load,
enabled=True,
eligible=True,
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,
prepare=lambda: PreparedNativeCall(
request=NativeTranscriptionRequest(model=model, audio=audio, optional_params=optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
bedrock=bedrock_options(optional_params),
),
context=with_capabilities(
context or NativeRequestContext(),
NativeRequestCapabilities(
execution_mode="async",
stream=stream,
has_custom_client=has_custom_client,
input_source_kind=input_source_kind,
),
),
),
call=call_native,
adapt=identity,
)

View file

@ -60,6 +60,66 @@ 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,
bedrock_options,
vertex_options,
)
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",
)
},
"bedrock": bedrock_options(params),
"vertex": vertex_options(params),
}
)
if route == "chat_completions":
request: Final = NativeChatCompletionsRequest(
model=TypeAdapter(str).validate_python(kwargs.get("model")),
messages=TypeAdapter(list[object]).validate_python(kwargs.get("messages")),
optional_params=params,
)
elif route == "messages":
request = NativeMessagesRequest(
model=TypeAdapter(str).validate_python(kwargs.get("model")),
body=TypeAdapter(dict[str, object]).validate_python(kwargs.get("body")),
)
elif route == "ocr":
request = NativeOCRRequest(
model=TypeAdapter(str).validate_python(kwargs.get("model")),
document=TypeAdapter(dict[str, object]).validate_python(kwargs.get("document")),
optional_params=params,
)
elif route in {"transcription", "audio_transcription"}:
request = NativeTranscriptionRequest(
model=TypeAdapter(str).validate_python(kwargs.get("model")),
audio=TypeAdapter(dict[str, object]).validate_python(kwargs.get("audio")),
optional_params=params,
)
else:
raise ValueError(f"unsupported native trace route: {route}")
return {"request": request, "options": options, "context": NativeRequestContext()}
def collect_trace(
spec: RouteSpec, engine: Engine, *, asynchronous: bool
) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure:
@ -77,7 +137,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,6 +9,7 @@ import pytest
import litellm
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.rust_bridge import configuration
from litellm.rust_bridge.request import NativeMessagesRequest, NativeRequestContext, NativeRequestOptions
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
@ -38,56 +39,54 @@ REQUEST_BODY: dict[str, object] = {
class RecordingMessages:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
self.contexts: list[NativeRequestContext] = []
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,
*,
options: NativeRequestOptions,
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": options.api_key,
"api_base": options.api_base,
"custom_llm_provider": options.custom_llm_provider,
"extra_headers": options.extra_headers,
"timeout_seconds": options.timeout_seconds,
}
)
self.contexts.append(context)
return dict(FAKE_MESSAGES_RESPONSE)
class RecordingAsyncMessages:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
self.contexts: list[NativeRequestContext] = []
async def __call__(
self,
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,
*,
options: NativeRequestOptions,
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": options.api_key,
"api_base": options.api_base,
"custom_llm_provider": options.custom_llm_provider,
"extra_headers": options.extra_headers,
"timeout_seconds": options.timeout_seconds,
}
)
self.contexts.append(context)
return dict(FAKE_MESSAGES_RESPONSE)
@ -95,7 +94,7 @@ class ExplodingAsyncMessages:
def __init__(self) -> None:
self.calls = 0
async def __call__(self, **kwargs: object) -> dict[str, object]:
async def __call__(self, *args: object, **kwargs: object) -> dict[str, object]:
self.calls += 1
raise AssertionError("bridge must not be called")
@ -104,7 +103,7 @@ class RaisingAsyncMessages:
def __init__(self) -> None:
self.calls = 0
async def __call__(self, **kwargs: object) -> dict[str, object]:
async def __call__(self, *args: object, **kwargs: object) -> dict[str, object]:
self.calls += 1
raise RuntimeError("upstream request failed with status 400: bad request")
@ -202,11 +201,38 @@ async def test_amessages_wrapper_forwards_args():
assert bridge.calls[0]["timeout_seconds"] == 12.5
@pytest.mark.asyncio
async def test_amessages_wrapper_preserves_capability_facts():
bridge = RecordingAsyncMessages()
rust_messages.set_rust_messages(amessages=bridge)
await rust_messages.amessages(
model="claude-sonnet-4-5",
body=REQUEST_BODY,
api_key=None,
api_base=None,
custom_llm_provider="anthropic",
extra_headers=None,
timeout=None,
stream=True,
has_custom_client=True,
has_agentic_hook=True,
)
capabilities = bridge.contexts[0].capabilities
assert capabilities.execution_mode == "async"
assert capabilities.stream is True
assert capabilities.has_custom_client is True
assert capabilities.has_agentic_hook is True
def _gate(**overrides):
kwargs = {
"custom_llm_provider": "azure_ai",
"litellm_params": GenericLiteLLMParams(api_key="sk-azure"),
"has_agentic_hook": False,
"stream": False,
"has_custom_client": False,
"model": "claude-sonnet-4-5",
"api_key": "sk-azure",
"api_base": "https://resource.services.ai.azure.com/anthropic",
@ -433,7 +459,7 @@ async def test_messages_handler_runs_selected_backend_once(selection: str, monke
def __init__(self) -> None:
self.calls = 0
async def __call__(self, **kwargs: object) -> dict[str, object]:
async def __call__(self, *args: object, **kwargs: object) -> dict[str, object]:
self.calls += 1
if error is not None:
raise error

View file

@ -65,8 +65,17 @@ 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, *, options, context):
seen["call"].append(
{
"model": request.model,
"messages": request.messages,
"optional_params": request.optional_params,
"api_key": options.api_key,
"api_base": options.api_base,
"context": context,
}
)
if error is not None:
raise error
return dict(RUST_RESPONSE)
@ -207,7 +216,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, **_kwargs):
raise _Declined("blank message text")
bridge.set_rust_chat_completions(
@ -237,7 +246,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, **_kwargs):
return dict(RUST_RESPONSE)
bridge.set_rust_chat_completions(
@ -271,7 +280,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, **_kwargs):
raise _Declined("blank message text")
logging_obj = MagicMock()
@ -384,7 +393,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, **_kwargs):
raise _Declined("blank message text")
logging_obj = MagicMock()
@ -438,7 +447,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, **_kwargs):
return dict(RUST_RESPONSE)
bridge.set_rust_chat_completions(
@ -470,7 +479,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, **_kwargs):
raise _Declined("blank message text")
logging_obj, calls = _recording_logging_obj()

View file

@ -12,6 +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.request import NativeOCRRequest, NativeRequestContext, NativeRequestOptions
from litellm.rust_bridge.runtime import Handled
from litellm.rust_bridge.timeouts import timeout_to_seconds
@ -46,30 +47,28 @@ class RecordingBridge:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
self.contexts: list[NativeRequestContext] = []
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,
*,
options: NativeRequestOptions,
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": options.api_key,
"api_base": options.api_base,
"custom_llm_provider": options.custom_llm_provider,
"extra_headers": options.extra_headers,
"optional_params": request.optional_params,
"timeout_seconds": options.timeout_seconds,
}
)
self.contexts.append(context)
return dict(FAKE_OCR_RESPONSE)
@ -78,44 +77,38 @@ class RecordingAsyncBridge:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
self.contexts: list[NativeRequestContext] = []
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,
*,
options: NativeRequestOptions,
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": options.api_key,
"api_base": options.api_base,
"custom_llm_provider": options.custom_llm_provider,
"extra_headers": options.extra_headers,
"optional_params": request.optional_params,
"timeout_seconds": options.timeout_seconds,
}
)
self.contexts.append(context)
return dict(FAKE_OCR_RESPONSE)
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,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> dict[str, object]:
raise RuntimeError("bridge failed")
@ -123,14 +116,10 @@ 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,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> dict[str, object]:
raise RuntimeError("bridge failed")
@ -413,6 +402,9 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
response = response.value
assert isinstance(response, OCRResponse)
assert response.pages[0].markdown == "hello world"
assert bridge.contexts[0].capabilities.execution_mode == "sync"
assert bridge.contexts[0].capabilities.input_source_kind == "document_url"
assert bridge.contexts[0].capabilities.native_response_format is False
assert bridge.calls[0] == {
"model": "mistral-ocr-latest",
"document": DOCUMENT,

View file

@ -4,6 +4,11 @@ import pytest
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
from litellm.rust_bridge import configuration, responses_websocket
from litellm.rust_bridge.request import (
NativeRequestContext,
NativeRequestOptions,
NativeResponsesWebSocketRequest,
)
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
@ -28,14 +33,17 @@ class _ClosedNativeConnection:
class _FakeNativeBridge:
contexts: list[NativeRequestContext] = []
@classmethod
async def connect(
cls,
request: NativeResponsesWebSocketRequest,
*,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> _FakeNativeConnection:
cls.contexts.append(context)
return _FakeNativeConnection()
@ -94,16 +102,18 @@ async def test_enabled_bridge_connects_and_adapts_socket(
await connection.send("response.create")
assert await connection.recv() == "response.completed"
await connection.close()
assert _FakeNativeBridge.contexts[-1].capabilities.websocket_mode == "native"
assert _FakeNativeBridge.contexts[-1].capabilities.requires_connection is True
class _FailingNativeBridge:
@classmethod
async def connect(
cls,
request: NativeResponsesWebSocketRequest,
*,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> _FakeNativeConnection:
raise RuntimeError("connection failed")
@ -125,7 +135,11 @@ async def test_managed_connection_closes_native_socket_on_consumer_failure() ->
class Bridge:
@classmethod
async def connect(
cls, *, url: str, headers: dict[str, str], timeout_seconds: float | None
cls,
request: NativeResponsesWebSocketRequest,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> _FakeNativeConnection:
return socket

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,84 @@ 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", {})
bedrock: Final = _record(
"BedrockOptions",
{
"aws_access_key_id": params.get("aws_access_key_id"),
"aws_secret_access_key": params.get("aws_secret_access_key"),
"aws_session_token": None,
"aws_region_name": params.get("aws_region_name"),
"aws_session_name": None,
"aws_profile_name": None,
"aws_role_name": None,
"aws_web_identity_token": None,
"aws_sts_endpoint": None,
"aws_external_id": None,
"aws_bedrock_runtime_endpoint": None,
"request_metadata_fields": (),
"request_metadata": None,
},
)
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"),
"bedrock": bedrock,
"anthropic": None,
"vertex": None,
},
)
request_params: Final = {"language": params.get("language")} if route == "transcription" else params
request: Final = _record(
"Request",
{
**{
key: value for key, value in inputs.items() if key in {"model", "document", "audio", "body", "messages"}
},
**({"optional_params": request_params} if route != "messages" else {}),
},
)
context: Final = _record(
"RequestContext",
{
"litellm_call_id": "native-wheel-call",
"trace_id": None,
"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,
},
),
"capabilities": _record(
"Capabilities",
{
"stream": False,
"has_agentic_hook": False,
"has_custom_client": False,
"request_format": None,
},
),
},
)
return {"request": request, "options": options, "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")
@ -227,12 +310,7 @@ async def exercise_async(native: object, api_base: str) -> None:
async def exercise_async_concurrency(native: object, api_base: str) -> None:
responses: Final = await asyncio.wait_for(
asyncio.gather(
*(
native.amessages(**route_kwargs("messages", api_base, "success"))
for _ in range(32)
)
),
asyncio.gather(*(native.amessages(**route_kwargs("messages", api_base, "success")) for _ in range(32))),
timeout=15,
)
for response in responses:

View file

@ -12,7 +12,8 @@ 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.rust_bridge.request import NativeChatCompletionsRequest, NativeRequestContext, NativeRequestOptions
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped
from litellm.types.utils import ModelResponse
RUST_RESPONSE = {
@ -95,17 +96,41 @@ class _RecordingCall:
self.result = result if result is not None else dict(RUST_RESPONSE)
self.error = error
self.calls: list[dict] = []
self.contexts: list[NativeRequestContext] = []
def __call__(self, **kwargs):
def __call__(
self,
request: NativeChatCompletionsRequest,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
):
kwargs = {
"model": request.model,
"messages": request.messages,
"optional_params": request.optional_params,
"api_key": options.api_key,
"api_base": options.api_base,
"custom_llm_provider": options.custom_llm_provider,
"extra_headers": options.extra_headers,
"timeout_seconds": options.timeout_seconds,
}
self.calls.append(kwargs)
self.contexts.append(context)
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: NativeChatCompletionsRequest,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
):
return _RecordingCall.__call__(self, request, options=options, context=context)
def _accepts(**overrides) -> bool:
@ -268,6 +293,18 @@ class TestSyncCall:
bridge.chat_completions(**_call_kwargs(ModelResponse()))
assert native.calls[0]["timeout_seconds"] == 30.0
def test_preserves_execution_and_client_capabilities(self):
native = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native)
bridge.chat_completions(
**_call_kwargs(ModelResponse()),
stream=True,
has_custom_client=True,
)
assert native.contexts[0].capabilities.execution_mode == "sync"
assert native.contexts[0].capabilities.stream is True
assert native.contexts[0].capabilities.has_custom_client is True
def test_reports_unavailable_bridge(self, monkeypatch):
_hide_native_bridge(monkeypatch)
assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeSkipped)

View file

@ -0,0 +1,40 @@
from types import SimpleNamespace
from litellm.rust_bridge.request import NativeRequestCapabilities, request_context
def test_request_context_preserves_identity_attribution_and_capabilities() -> None:
capabilities = NativeRequestCapabilities(execution_mode="async", stream=True)
context = request_context(
logging_obj=SimpleNamespace(litellm_call_id="call-1", litellm_trace_id="trace-1"),
request_model="router-alias",
litellm_params={
"metadata": {
"user_api_key_hash": "hash-1",
"user_api_key_user_id": "user-1",
"user_api_key_team_id": "team-1",
}
},
capabilities=capabilities,
)
assert context.litellm_call_id == "call-1"
assert context.trace_id == "trace-1"
assert context.request_model == "router-alias"
assert context.attribution.user_api_key_hash == "hash-1"
assert context.attribution.user_api_key_user_id == "user-1"
assert context.attribution.user_api_key_team_id == "team-1"
assert context.capabilities is capabilities
def test_request_context_ignores_untyped_identity_values() -> None:
context = request_context(
logging_obj=SimpleNamespace(litellm_call_id=1, litellm_trace_id=[]),
request_model="model",
litellm_params={"user_api_key_user_id": 42},
)
assert context.litellm_call_id is None
assert context.trace_id is None
assert context.attribution.user_api_key_user_id is None

View file

@ -4,6 +4,7 @@ import pytest
import litellm
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
from litellm.rust_bridge.request import NativeRequestContext, NativeRequestOptions, NativeTranscriptionRequest
from litellm.rust_bridge.runtime import Handled
rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
@ -12,33 +13,29 @@ rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
class SyncBridge:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
self.contexts: list[NativeRequestContext] = []
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,
*,
options: NativeRequestOptions,
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}
)
self.contexts.append(context)
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,
*,
options: NativeRequestOptions,
context: NativeRequestContext,
) -> dict[str, object]:
return {"text": "async"}
@ -55,10 +52,17 @@ def test_enabled_sync_bridge_receives_audio() -> None:
extra_headers=None,
optional_params={"temperature": 0},
timeout=5.0,
stream=True,
has_custom_client=True,
input_source_kind="file",
)
assert isinstance(result, Handled)
assert result.value == {"text": "hello"}
assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"}
assert bridge.contexts[0].capabilities.execution_mode == "sync"
assert bridge.contexts[0].capabilities.stream is True
assert bridge.contexts[0].capabilities.has_custom_client is True
assert bridge.contexts[0].capabilities.input_source_kind == "file"
@pytest.mark.asyncio
@ -121,7 +125,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 *_args, **_: {"text": "rust"},
atranscription=None,
)
try:
@ -137,7 +141,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(*_args: object, **_: object) -> 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