refactor(native): separate request data from execution context

This commit is contained in:
Yujong Lee 2026-09-05 12:59:50 -07:00
parent f0d47db104
commit 0f886a6c30
61 changed files with 1841 additions and 1391 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

@ -31,6 +31,15 @@ from litellm.rust_bridge.protocols import (
RustChatCompletions,
RustChatCompletionsDecline,
)
from litellm.rust_bridge.request import (
NativeChatCompletionsRequest,
NativeRequestContext,
NativeRequestOptions,
PreparedNativeCall,
call_native,
provider_connection_params,
provider_request_params,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointBinding,
@ -235,17 +244,23 @@ def chat_completions(
return _build_model_response(rust_response, model_response)
return _CHAT.invoke(
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_chat_completions, timeout_seconds: rust_chat_completions(
model=model,
messages=messages,
optional_params=optional_params,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_seconds,
prepare=lambda: PreparedNativeCall(
NativeChatCompletionsRequest(
model=model,
messages=messages,
optional_params=provider_request_params(optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=lambda: None,
adapt=adapt,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
@ -270,17 +285,23 @@ async def achat_completions(
return _build_model_response(rust_response, model_response)
return await _CHAT.ainvoke(
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_achat_completions, timeout_seconds: rust_achat_completions(
model=model,
messages=messages,
optional_params=optional_params,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_seconds,
prepare=lambda: PreparedNativeCall(
NativeChatCompletionsRequest(
model=model,
messages=messages,
optional_params=provider_request_params(optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=async_none,
adapt=adapt,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
@ -315,17 +336,23 @@ async def achat_completions_or_fallback(
return _build_model_response(rust_response, model_response)
return await _CHAT.ainvoke(
prepare=lambda: timeout_to_seconds(timeout),
call=lambda rust_achat_completions, timeout_seconds: rust_achat_completions(
model=model,
messages=messages,
optional_params=optional_params,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_seconds,
prepare=lambda: PreparedNativeCall(
NativeChatCompletionsRequest(
model=model,
messages=messages,
optional_params=provider_request_params(optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=python_fallback,
adapt=adapt,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),

View file

@ -8,6 +8,13 @@ import httpx
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.protocols import RustAmessages, RustMessages
from litellm.rust_bridge.request import (
NativeMessagesRequest,
NativeRequestContext,
NativeRequestOptions,
PreparedNativeCall,
call_native,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointDispatch,
@ -61,16 +68,21 @@ def messages(
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
return _MESSAGES.invoke(
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(
NativeMessagesRequest(
model=model,
body=body,
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=lambda: None,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
@ -88,16 +100,21 @@ async def amessages(
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
return await _MESSAGES.ainvoke(
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(
NativeMessagesRequest(
model=model,
body=body,
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=async_none,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),

View file

@ -9,6 +9,15 @@ import httpx
from litellm.rust_bridge import configuration as _configuration
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.protocols import RustAocr, RustOcr
from litellm.rust_bridge.request import (
NativeOCRRequest,
NativeRequestContext,
NativeRequestOptions,
PreparedNativeCall,
call_native,
provider_connection_params,
provider_request_params,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointDispatch,
@ -67,17 +76,23 @@ def ocr(
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
return _OCR.invoke(
prepare=lambda: _timeout_to_seconds(timeout),
call=lambda rust_ocr, timeout_seconds: rust_ocr(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_seconds,
prepare=lambda: PreparedNativeCall(
NativeOCRRequest(
model=model,
document=document,
optional_params=provider_request_params(optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=_timeout_to_seconds(timeout),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=lambda: None,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
@ -96,17 +111,23 @@ async def aocr(
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
return await _OCR.ainvoke(
prepare=lambda: _timeout_to_seconds(timeout),
call=lambda rust_aocr, timeout_seconds: rust_aocr(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_seconds,
prepare=lambda: PreparedNativeCall(
NativeOCRRequest(
model=model,
document=document,
optional_params=provider_request_params(optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=_timeout_to_seconds(timeout),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=async_none,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),

View file

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

View file

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

View file

@ -13,6 +13,13 @@ from litellm.rust_bridge.protocols import (
RustResponsesWebSocket,
RustResponsesWebSocketConnection,
)
from litellm.rust_bridge.request import (
NativeRequestContext,
NativeRequestOptions,
NativeResponsesWebSocketRequest,
PreparedNativeCall,
call_native,
)
from litellm.rust_bridge.runtime import (
AsyncEndpointDispatch,
BridgeErrorContext,
@ -34,9 +41,9 @@ def set_rust_responses_websocket(
) -> None:
if not isinstance(connection, Unchanged):
if connection is None:
_RESPONSES_WEBSOCKET.reset()
_RESPONSES_WEBSOCKET.asynchronous.reset()
else:
_RESPONSES_WEBSOCKET.override(connection)
_RESPONSES_WEBSOCKET.asynchronous.override(connection)
class _ConnectionAdapter:
@ -63,12 +70,14 @@ async def connect(
timeout: float | httpx.Timeout | None,
) -> _ConnectionAdapter | None:
connection: Final = await _RESPONSES_WEBSOCKET.ainvoke(
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(
NativeResponsesWebSocketRequest(
url=url,
options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)),
),
context=NativeRequestContext(),
),
call=lambda connection_type, request: call_native(connection_type.connect, request),
fallback=async_none,
adapt=identity,
error_context=BridgeErrorContext(provider="openai", model="responses websocket"),

View file

@ -6,6 +6,15 @@ import httpx
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription
from litellm.rust_bridge.request import (
NativeRequestContext,
NativeRequestOptions,
NativeTranscriptionRequest,
PreparedNativeCall,
call_native,
provider_connection_params,
provider_request_params,
)
from litellm.rust_bridge.runtime import (
BridgeErrorContext,
EndpointDispatch,
@ -61,17 +70,23 @@ def transcription(
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
return _TRANSCRIPTION.invoke(
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(
NativeTranscriptionRequest(
model=model,
audio=audio,
optional_params=provider_request_params(optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=lambda: None,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
@ -90,17 +105,23 @@ async def atranscription(
timeout: float | httpx.Timeout | None,
) -> dict[str, object] | None:
return await _TRANSCRIPTION.ainvoke(
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(
NativeTranscriptionRequest(
model=model,
audio=audio,
optional_params=provider_request_params(optional_params),
options=NativeRequestOptions(
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
timeout_seconds=timeout_to_seconds(timeout),
provider_connection=provider_connection_params(optional_params),
),
),
context=NativeRequestContext(),
),
call=call_native,
fallback=async_none,
adapt=identity,
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),

View file

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

View file

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

View file

@ -9,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
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
@ -40,23 +41,19 @@ class RecordingMessages:
def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
request: NativeMessagesRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"body": body,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"timeout_seconds": timeout_seconds,
"model": request.model,
"body": request.body,
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"timeout_seconds": request.options.timeout_seconds,
}
)
return dict(FAKE_MESSAGES_RESPONSE)
@ -68,23 +65,19 @@ class RecordingAsyncMessages:
async def __call__(
self,
model: str,
body: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
timeout_seconds: float | None,
request: NativeMessagesRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"body": body,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"timeout_seconds": timeout_seconds,
"model": request.model,
"body": request.body,
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"timeout_seconds": request.options.timeout_seconds,
}
)
return dict(FAKE_MESSAGES_RESPONSE)
@ -94,7 +87,7 @@ class ExplodingAsyncMessages:
def __init__(self) -> None:
self.calls = 0
async def __call__(self, **kwargs: object) -> dict[str, object]:
async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]:
self.calls += 1
raise AssertionError("bridge must not be called")
@ -103,7 +96,7 @@ class RaisingAsyncMessages:
def __init__(self) -> None:
self.calls = 0
async def __call__(self, **kwargs: object) -> dict[str, object]:
async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]:
self.calls += 1
raise RuntimeError("upstream request failed with status 400: bad request")

View file

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

View file

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

View file

@ -11,6 +11,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
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
# function onto `litellm.ocr` and shadows the submodule, so import the modules
@ -46,25 +47,20 @@ class RecordingBridge:
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeOCRRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": timeout_seconds,
"model": request.model,
"document": request.document,
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
"timeout_seconds": request.options.timeout_seconds,
}
)
return dict(FAKE_OCR_RESPONSE)
@ -78,25 +74,20 @@ class RecordingAsyncBridge:
async def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeOCRRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"document": document,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": timeout_seconds,
"model": request.model,
"document": request.document,
"api_key": request.options.api_key,
"api_base": request.options.api_base,
"custom_llm_provider": request.options.custom_llm_provider,
"extra_headers": request.options.extra_headers,
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
"timeout_seconds": request.options.timeout_seconds,
}
)
return dict(FAKE_OCR_RESPONSE)
@ -105,14 +96,9 @@ class RecordingAsyncBridge:
class RaisingBridge:
def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeOCRRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
raise RuntimeError("bridge failed")
@ -120,14 +106,9 @@ class RaisingBridge:
class RaisingAsyncBridge:
async def __call__(
self,
model: str,
document: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeOCRRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
raise RuntimeError("bridge failed")

View file

@ -4,6 +4,7 @@ import pytest
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
from litellm.rust_bridge import configuration, responses_websocket
from litellm.rust_bridge.request import NativeRequestContext, NativeResponsesWebSocketRequest
class _FakeNativeConnection:
@ -30,10 +31,9 @@ class _FakeNativeBridge:
@classmethod
async def connect(
cls,
request: NativeResponsesWebSocketRequest,
*,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
context: NativeRequestContext,
) -> _FakeNativeConnection:
return _FakeNativeConnection()
@ -101,10 +101,9 @@ class _FailingNativeBridge:
@classmethod
async def connect(
cls,
request: NativeResponsesWebSocketRequest,
*,
url: str,
headers: dict[str, str],
timeout_seconds: float | None,
context: NativeRequestContext,
) -> _FakeNativeConnection:
raise RuntimeError("connection failed")

View file

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

View file

@ -95,16 +95,16 @@ class _RecordingCall:
self.error = error
self.calls: list[dict] = []
def __call__(self, **kwargs):
self.calls.append(kwargs)
def __call__(self, request, *, context):
self.calls.append({"request": request, "context": context})
if self.error is not None:
raise self.error
return self.result
class _RecordingAsyncCall(_RecordingCall):
async def __call__(self, **kwargs):
return _RecordingCall.__call__(self, **kwargs)
async def __call__(self, request, *, context):
return _RecordingCall.__call__(self, request, context=context)
def _accepts(**overrides) -> bool:
@ -271,7 +271,7 @@ class TestSyncCall:
native = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native)
bridge.chat_completions(**_call_kwargs(ModelResponse()))
assert native.calls[0]["timeout_seconds"] == 30.0
assert native.calls[0]["request"].options.timeout_seconds == 30.0
def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
_hide_native_bridge(monkeypatch)
@ -418,3 +418,22 @@ async def test_missing_native_exception_types_does_not_authorize_python_fallback
with pytest.raises(RuntimeError, match="connection failed"):
await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
def test_provider_credentials_are_separate_from_chat_body_params():
native = _RecordingCall()
bridge.set_rust_chat_completions(chat_completions=native)
configuration.rust(True)
kwargs = _call_kwargs(ModelResponse())
kwargs["optional_params"] = {
"max_tokens": 32,
"aws_access_key_id": "test-access-key",
"aws_secret_access_key": "test-secret-key",
}
bridge.chat_completions(**kwargs)
request = native.calls[0]["request"]
assert request.optional_params == {"max_tokens": 32}
assert request.options.provider_connection == {
"aws_access_key_id": "test-access-key",
"aws_secret_access_key": "test-secret-key",
}

View file

@ -4,6 +4,7 @@ import pytest
import litellm
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
from litellm.rust_bridge.request import NativeRequestContext, NativeTranscriptionRequest
rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
@ -14,30 +15,20 @@ class SyncBridge:
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeTranscriptionRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
self.calls.append({"model": model, "audio": audio, "optional_params": optional_params})
self.calls.append({"model": request.model, "audio": request.audio, "optional_params": {**request.optional_params, **(request.options.provider_connection or {})}})
return {"text": "hello"}
class AsyncBridge:
async def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
request: NativeTranscriptionRequest,
*,
context: NativeRequestContext,
) -> dict[str, object]:
return {"text": "async"}
@ -120,7 +111,7 @@ async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPat
def test_bedrock_transcription_uses_rust_only_path() -> None:
rust_bridge.configure_rust_transcription(
transcription=lambda **_: {"text": "rust"},
transcription=lambda request, *, context: {"text": "rust"},
atranscription=None,
)
try:
@ -136,7 +127,7 @@ def test_bedrock_transcription_uses_rust_only_path() -> None:
@pytest.mark.asyncio
async def test_bedrock_atranscription_uses_rust_only_path() -> None:
async def rust_response(**_: object) -> dict[str, object]:
async def rust_response(request: NativeTranscriptionRequest, *, context: NativeRequestContext) -> dict[str, object]:
return {"text": "rust"}
rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response)

View file

@ -1,6 +1,6 @@
{
"LIT001": {
"limit": 22181
"limit": 22165
},
"LIT002": {
"limit": 26745