mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-11 22:51:28 +00:00
refactor(native): separate request data from execution context
This commit is contained in:
parent
7ed22070e8
commit
f8d6ac7a76
62 changed files with 2530 additions and 1984 deletions
|
|
@ -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 |
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"}));
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
})?,
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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`"))
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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| {
|
||||
|
|
|
|||
|
|
@ -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`
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
||||
|
|
|
|||
|
|
@ -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?;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"}));
|
||||
|
|
|
|||
|
|
@ -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>,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -13,7 +13,7 @@ use super::types::{
|
|||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(super) async fn execute_chat_completions_provider_call(
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
request: ResolvedChatCompletionsRequest,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let request = prepare_provider_request(request)?;
|
||||
let body = serde_json::to_vec(&request.body).map_err(|err| {
|
||||
|
|
@ -101,11 +101,11 @@ pub(super) async fn signed_headers(
|
|||
let unsigned: BTreeMap<String, String> = request.upstream_headers.iter().cloned().collect();
|
||||
// A host with its own resolution chain hands the result down; only fall
|
||||
// back to deriving credentials here when it supplied none.
|
||||
let credentials = match host_supplied_credentials(&request.optional_params) {
|
||||
let credentials = match host_supplied_credentials(&request.provider_connection) {
|
||||
Some(credentials) => credentials,
|
||||
None => {
|
||||
resolve_credentials(
|
||||
aws_auth_config(&request.optional_params, &env_lookup),
|
||||
aws_auth_config(&request.provider_connection, &env_lookup),
|
||||
&env_lookup,
|
||||
)
|
||||
.await?
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use crate::request_context::LiteLlmRequestContext;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::error::Error;
|
||||
|
|
@ -39,9 +40,13 @@ pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error>
|
|||
|
||||
pub(super) fn resolve_request(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ResolvedChatCompletionsRequest<'_>, Error> {
|
||||
let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)
|
||||
.map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?;
|
||||
_context: &LiteLlmRequestContext,
|
||||
) -> Result<ResolvedChatCompletionsRequest, Error> {
|
||||
let (model, config) = resolve_provider_config(
|
||||
request.model,
|
||||
request.options.custom_llm_provider.as_deref(),
|
||||
)
|
||||
.map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?;
|
||||
let messages =
|
||||
parse_messages(request.messages).map_err(|_| Error::Declined("unreadable message list"))?;
|
||||
if messages.is_empty() {
|
||||
|
|
@ -55,25 +60,22 @@ pub(super) fn resolve_request(
|
|||
config,
|
||||
messages,
|
||||
optional_params: request.optional_params,
|
||||
api_key: request.api_key,
|
||||
api_base: request.api_base,
|
||||
extra_headers: request.extra_headers,
|
||||
timeout: request.timeout,
|
||||
options: request.options,
|
||||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn validate_environment(
|
||||
request: &ResolvedChatCompletionsRequest<'_>,
|
||||
request: &ResolvedChatCompletionsRequest,
|
||||
model: &str,
|
||||
config: &dyn ChatCompletionsProviderConfig,
|
||||
) -> Result<(Vec<(String, String)>, ChatCompletionsAuth), Error> {
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let mut headers = string_headers(request.extra_headers.clone())?;
|
||||
let mut headers = string_headers(request.options.extra_headers.clone())?;
|
||||
let auth = config.auth(
|
||||
request.api_key,
|
||||
request.options.api_key.as_deref(),
|
||||
model,
|
||||
&request.optional_params,
|
||||
&request.options.provider_connection,
|
||||
&env_lookup,
|
||||
)?;
|
||||
match &auth {
|
||||
|
|
@ -117,16 +119,16 @@ fn validate_environment(
|
|||
}
|
||||
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
request: ResolvedChatCompletionsRequest,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
let (headers, auth) = validate_environment(&request, &request.model, request.config)?;
|
||||
let model = request.model;
|
||||
let config = request.config;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let url = config.complete_url(
|
||||
request.api_base,
|
||||
request.options.api_base.as_deref(),
|
||||
&model,
|
||||
&request.optional_params,
|
||||
&request.options.provider_connection,
|
||||
&env_lookup,
|
||||
)?;
|
||||
let transformed =
|
||||
|
|
@ -139,7 +141,7 @@ pub(super) fn prepare_provider_request(
|
|||
body: transformed.body,
|
||||
upstream_headers: headers,
|
||||
auth,
|
||||
optional_params: request.optional_params,
|
||||
timeout: request.timeout,
|
||||
provider_connection: request.options.provider_connection,
|
||||
timeout: request.options.timeout,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
use crate::request_context::LiteLlmRequestContext;
|
||||
use crate::request_options::RequestOptions;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::error::Error;
|
||||
|
|
@ -9,7 +11,12 @@ use super::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest};
|
|||
fn prepare_chat_completions_call(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
prepare_provider_request(resolve_request(request)?)
|
||||
prepare_provider_request(resolve_request(
|
||||
request,
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)?)
|
||||
}
|
||||
|
||||
fn request<'a>(
|
||||
|
|
@ -25,11 +32,15 @@ fn request<'a>(
|
|||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
},
|
||||
api_key: Some("sk-test"),
|
||||
api_base: None,
|
||||
custom_llm_provider: provider,
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
|
||||
options: RequestOptions {
|
||||
api_key: (Some("sk-test")).map(|value| value.to_string()),
|
||||
api_base: None,
|
||||
custom_llm_provider: (provider).map(|value| value.to_string()),
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
..Default::default()
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -107,7 +118,7 @@ fn the_deployment_credential_replaces_a_caller_supplied_auth_header() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([(
|
||||
call.options.extra_headers = Some(Map::from_iter([(
|
||||
"X-Api-Key".to_string(),
|
||||
json!("sk-caller"),
|
||||
)]));
|
||||
|
|
@ -132,7 +143,7 @@ fn a_forwarded_authorization_header_suppresses_the_resolved_api_key_header() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([
|
||||
call.options.extra_headers = Some(Map::from_iter([
|
||||
(
|
||||
"Authorization".to_string(),
|
||||
json!("Bearer sk-ant-oat01-token"),
|
||||
|
|
@ -168,7 +179,7 @@ fn an_unrelated_forwarded_authorization_does_not_defer_the_resolved_key() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([
|
||||
call.options.extra_headers = Some(Map::from_iter([
|
||||
("Authorization".to_string(), json!("Bearer unrelated")),
|
||||
("X-Api-Key".to_string(), json!("sk-caller")),
|
||||
]));
|
||||
|
|
@ -199,7 +210,7 @@ fn declines_an_unsupported_request_before_resolving_credentials() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true}),
|
||||
);
|
||||
call.api_key = None;
|
||||
call.options.api_key = None;
|
||||
// No api_key is set and no env is consulted: the gate must run first, so the
|
||||
// error is the decline rather than a missing-credential error.
|
||||
assert_eq!(decline(call), Error::Declined("streaming"));
|
||||
|
|
@ -261,7 +272,7 @@ fn rejects_non_string_extra_headers() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
|
||||
call.options.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
|
||||
assert_eq!(
|
||||
decline(call),
|
||||
Error::InvalidRequest(
|
||||
|
|
@ -279,7 +290,7 @@ fn prepares_a_bedrock_call_without_resolving_credentials() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"maxTokens": 16}),
|
||||
);
|
||||
call.api_key = None;
|
||||
call.options.api_key = None;
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
assert_eq!(
|
||||
prepared.url,
|
||||
|
|
@ -312,15 +323,18 @@ async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() {
|
|||
"bedrock/us-east-1/anthropic.claude-v2",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({
|
||||
"maxTokens": 16,
|
||||
"aws_access_key_id": "AKIDEXAMPLE",
|
||||
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
|
||||
}),
|
||||
json!({"maxTokens": 16}),
|
||||
);
|
||||
call.options.provider_connection = Map::from_iter([
|
||||
("aws_access_key_id".to_string(), json!("AKIDEXAMPLE")),
|
||||
(
|
||||
"aws_secret_access_key".to_string(),
|
||||
json!("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"),
|
||||
),
|
||||
]);
|
||||
// A key would resolve to a bearer token and never reach the signer.
|
||||
call.api_key = None;
|
||||
call.extra_headers = Some(Map::from_iter([(
|
||||
call.options.api_key = None;
|
||||
call.options.extra_headers = Some(Map::from_iter([(
|
||||
"x-request-id".to_string(),
|
||||
json!("abc-123"),
|
||||
)]));
|
||||
|
|
@ -367,14 +381,18 @@ async fn a_forwarded_header_the_signer_computes_declines_to_python() {
|
|||
"bedrock/us-east-1/anthropic.claude-v2",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({
|
||||
"maxTokens": 16,
|
||||
"aws_access_key_id": "AKIDEXAMPLE",
|
||||
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
|
||||
}),
|
||||
json!({"maxTokens": 16}),
|
||||
);
|
||||
call.api_key = None;
|
||||
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
|
||||
call.options.provider_connection = Map::from_iter([
|
||||
("aws_access_key_id".to_string(), json!("AKIDEXAMPLE")),
|
||||
(
|
||||
"aws_secret_access_key".to_string(),
|
||||
json!("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"),
|
||||
),
|
||||
]);
|
||||
call.options.api_key = None;
|
||||
call.options.extra_headers =
|
||||
Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let error = super::handler::signed_headers(&prepared, br#"{"a":1}"#)
|
||||
.await
|
||||
|
|
@ -399,7 +417,7 @@ fn a_bedrock_deployment_bearer_outranks_a_forwarded_authorization() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"maxTokens": 16}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([(
|
||||
call.options.extra_headers = Some(Map::from_iter([(
|
||||
"Authorization".to_string(),
|
||||
json!("Bearer caller-supplied"),
|
||||
)]));
|
||||
|
|
@ -432,7 +450,7 @@ fn an_anthropic_forwarded_oauth_bearer_still_outranks_the_resolved_key() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([(
|
||||
call.options.extra_headers = Some(Map::from_iter([(
|
||||
"authorization".to_string(),
|
||||
json!("Bearer sk-ant-oat01-forwarded"),
|
||||
)]));
|
||||
|
|
@ -666,11 +684,15 @@ mod round_trip {
|
|||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
},
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some(api_base),
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: Some(std::time::Duration::from_secs(10)),
|
||||
|
||||
options: RequestOptions {
|
||||
api_key: (Some("sk-test")).map(|value| value.to_string()),
|
||||
api_base: (Some(api_base)).map(|value| value.to_string()),
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: Some(std::time::Duration::from_secs(10)),
|
||||
..Default::default()
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -679,14 +701,19 @@ mod round_trip {
|
|||
#[tokio::test]
|
||||
async fn round_trip_sends_the_translated_body_and_normalizes_the_response() {
|
||||
let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await;
|
||||
let response = chat_completions(call(
|
||||
&api_base,
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let response = chat_completions(
|
||||
call(
|
||||
&api_base,
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("call succeeds");
|
||||
|
||||
|
|
@ -724,11 +751,16 @@ mod round_trip {
|
|||
const NO_USAGE: &str =
|
||||
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#;
|
||||
let (api_base, handle) = serve_once("200 OK", NO_USAGE).await;
|
||||
let err = chat_completions(call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let err = chat_completions(
|
||||
call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect_err("response cannot be normalized");
|
||||
handle.await.expect("server task");
|
||||
|
|
@ -742,11 +774,16 @@ mod round_trip {
|
|||
async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() {
|
||||
const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#;
|
||||
let (api_base, handle) = serve_once("200 OK", TOOL_USE).await;
|
||||
let err = chat_completions(call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let err = chat_completions(
|
||||
call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect_err("response cannot be normalized");
|
||||
handle.await.expect("server task");
|
||||
|
|
@ -760,11 +797,16 @@ mod round_trip {
|
|||
async fn an_upstream_error_status_keeps_its_code() {
|
||||
let (api_base, handle) =
|
||||
serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await;
|
||||
let err = chat_completions(call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let err = chat_completions(
|
||||
call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect_err("upstream rejects");
|
||||
handle.await.expect("server task");
|
||||
|
|
@ -781,11 +823,16 @@ mod round_trip {
|
|||
listener.local_addr().expect("has an address").port()
|
||||
// Dropped here, so the port is closed and the connect is refused.
|
||||
};
|
||||
let err = chat_completions(call(
|
||||
&format!("http://127.0.0.1:{port}/v1/messages"),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let err = chat_completions(
|
||||
call(
|
||||
&format!("http://127.0.0.1:{port}/v1/messages"),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect_err("nothing is listening");
|
||||
assert!(
|
||||
|
|
|
|||
|
|
@ -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>,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -16,3 +16,6 @@ pub mod router;
|
|||
pub mod routing_utils;
|
||||
|
||||
pub use error::Error;
|
||||
|
||||
pub mod request_context;
|
||||
pub mod request_options;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
18
litellm-rust/crates/core/src/request_context.rs
Normal file
18
litellm-rust/crates/core/src/request_context.rs
Normal 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,
|
||||
}
|
||||
14
litellm-rust/crates/core/src/request_options.rs
Normal file
14
litellm-rust/crates/core/src/request_options.rs
Normal 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>,
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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}],
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,6 +8,7 @@ import mimetypes
|
|||
import os
|
||||
import re
|
||||
from collections.abc import Callable, Coroutine, Mapping
|
||||
from dataclasses import dataclass
|
||||
from io import IOBase
|
||||
from typing import Any, Final, cast
|
||||
|
||||
|
|
@ -17,16 +18,24 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import request_timeout
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model
|
||||
from litellm.llms.azure_ai.ocr.common_utils import (
|
||||
is_azure_document_intelligence_model,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import (
|
||||
OCR_REQUEST_FORMAT_PARAM,
|
||||
BaseOCRConfig,
|
||||
OCRResponse,
|
||||
parse_ocr_request_format,
|
||||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors
|
||||
from litellm.rust_bridge.runtime import DispatchResult
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
provider_connection_params,
|
||||
provider_request_params,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.utils import ProviderConfigManager, client
|
||||
|
||||
|
|
@ -35,6 +44,28 @@ base_llm_http_handler = BaseLLMHTTPHandler()
|
|||
#################################################
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PreparedOCRRequest:
|
||||
model: str
|
||||
document: dict[str, Any]
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
custom_llm_provider: str
|
||||
extra_headers: dict[str, object] | None
|
||||
provider_config: BaseOCRConfig
|
||||
optional_params: dict[str, object]
|
||||
litellm_params: dict[str, object]
|
||||
effective_timeout: float | httpx.Timeout
|
||||
litellm_logging_obj: LiteLLMLoggingObj
|
||||
|
||||
|
||||
_RUST_OCR_PROVIDERS: Final = {
|
||||
"mistral",
|
||||
"azure_ai",
|
||||
"vertex_ai",
|
||||
}
|
||||
|
||||
|
||||
def _prepare_ocr_request(
|
||||
model: str,
|
||||
document: Mapping[str, object],
|
||||
|
|
@ -44,7 +75,7 @@ def _prepare_ocr_request(
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
kwargs: dict[str, object],
|
||||
) -> rust_ocr_bridge.PreparedOCRRequest:
|
||||
) -> _PreparedOCRRequest:
|
||||
litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
|
||||
litellm_call_id: Final = cast(str | None, kwargs.get("litellm_call_id", None))
|
||||
|
||||
|
|
@ -141,7 +172,7 @@ def _prepare_ocr_request(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
||||
return rust_ocr_bridge.PreparedOCRRequest(
|
||||
return _PreparedOCRRequest(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
|
|
@ -156,71 +187,143 @@ def _prepare_ocr_request(
|
|||
)
|
||||
|
||||
|
||||
@anative_first(
|
||||
native=rust_ocr_bridge.aattempt_ocr,
|
||||
route="ocr",
|
||||
errors=lambda prepared_request, resolve_api_key: provider_errors(
|
||||
prepared_request.custom_llm_provider, prepared_request.model
|
||||
),
|
||||
)
|
||||
async def _execute_aocr(
|
||||
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
|
||||
def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
|
||||
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
|
||||
return False
|
||||
if not prepared_request.provider_config.supports_rust_bridge():
|
||||
return False
|
||||
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
|
||||
|
||||
|
||||
def _rust_bridge_optional_params(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> dict[str, object]:
|
||||
optional_params: Final = dict(prepared_request.optional_params)
|
||||
if prepared_request.custom_llm_provider == "vertex_ai":
|
||||
vertex_project: Final = (
|
||||
prepared_request.litellm_params.get("vertex_project")
|
||||
or prepared_request.litellm_params.get("vertex_ai_project")
|
||||
or litellm.vertex_project
|
||||
or resolve_secret("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_location: Final = (
|
||||
prepared_request.litellm_params.get("vertex_location")
|
||||
or prepared_request.litellm_params.get("vertex_ai_location")
|
||||
or litellm.vertex_location
|
||||
or resolve_secret("VERTEXAI_LOCATION")
|
||||
or resolve_secret("VERTEX_LOCATION")
|
||||
)
|
||||
if vertex_project is not None:
|
||||
optional_params["vertex_project"] = vertex_project
|
||||
if vertex_location is not None:
|
||||
optional_params["vertex_location"] = vertex_location
|
||||
return optional_params
|
||||
|
||||
|
||||
def _rust_bridge_api_base(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> str | None:
|
||||
if prepared_request.api_base is not None:
|
||||
return prepared_request.api_base
|
||||
if prepared_request.custom_llm_provider == "azure_ai":
|
||||
if is_azure_document_intelligence_model(prepared_request.model):
|
||||
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
return resolve_secret("AZURE_AI_API_BASE")
|
||||
return None
|
||||
|
||||
|
||||
def _prepare_rust_ocr_call(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> OCRResponse:
|
||||
pending: Final = base_llm_http_handler.ocr(
|
||||
) -> PreparedNativeCall[rust_ocr_bridge.NativeOCRRequest]:
|
||||
provider_config: Final = prepared_request.provider_config
|
||||
api_key_env_var: Final = provider_config.get_api_key_env_var()
|
||||
resolved_api_key: Final = prepared_request.api_key or (
|
||||
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
|
||||
)
|
||||
resolved_headers: Final = provider_config.validate_environment(
|
||||
headers=prepared_request.extra_headers or {},
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=prepared_request.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
logging_obj=prepared_request.litellm_logging_obj,
|
||||
api_key=prepared_request.api_key,
|
||||
api_key=resolved_api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
aocr=True,
|
||||
headers=prepared_request.extra_headers,
|
||||
provider_config=prepared_request.provider_config,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
response: Final = await pending if asyncio.iscoroutine(pending) else pending
|
||||
if response is None:
|
||||
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
|
||||
return response
|
||||
|
||||
|
||||
def _attempt_ocr(
|
||||
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
is_async: bool,
|
||||
) -> DispatchResult[OCRResponse]:
|
||||
return rust_ocr_bridge.attempt_ocr(prepared_request=prepared_request, resolve_api_key=resolve_api_key)
|
||||
|
||||
|
||||
@native_first(
|
||||
native=_attempt_ocr,
|
||||
route="ocr",
|
||||
errors=lambda prepared_request, resolve_api_key, is_async: provider_errors(
|
||||
prepared_request.custom_llm_provider, prepared_request.model
|
||||
),
|
||||
)
|
||||
def _execute_ocr(
|
||||
prepared_request: rust_ocr_bridge.PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
is_async: bool,
|
||||
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
return base_llm_http_handler.ocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=prepared_request.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
logging_obj=prepared_request.litellm_logging_obj,
|
||||
api_key=prepared_request.api_key,
|
||||
resolved_complete_url: Final = provider_config.get_complete_url(
|
||||
api_base=prepared_request.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
aocr=is_async,
|
||||
headers=prepared_request.extra_headers,
|
||||
provider_config=prepared_request.provider_config,
|
||||
model=prepared_request.model,
|
||||
optional_params=prepared_request.optional_params,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
|
||||
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
|
||||
prepared_request.litellm_logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=resolved_api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": {
|
||||
"model": prepared_request.model,
|
||||
"document": prepared_request.document,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": resolved_complete_url,
|
||||
"headers": resolved_headers,
|
||||
},
|
||||
)
|
||||
return PreparedNativeCall(
|
||||
request=rust_ocr_bridge.NativeOCRRequest(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
optional_params=provider_request_params(rust_optional_params),
|
||||
options=NativeRequestOptions(
|
||||
provider_connection=provider_connection_params(rust_optional_params),
|
||||
api_key=resolved_api_key,
|
||||
api_base=rust_api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=cast( # cast-ok: provider header normalization returns string-object pairs
|
||||
dict[str, object], resolved_headers
|
||||
),
|
||||
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _run_rust_ocr(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
fallback: Callable[[], OCRResponse | Coroutine[object, object, OCRResponse]],
|
||||
) -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
return rust_ocr_bridge.dispatch_ocr(
|
||||
prepare=lambda: _prepare_rust_ocr_call(
|
||||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
),
|
||||
fallback=fallback,
|
||||
adapt=OCRResponse.model_validate,
|
||||
model=prepared_request.model,
|
||||
provider=prepared_request.custom_llm_provider,
|
||||
eligible=_rust_ocr_supported(prepared_request),
|
||||
)
|
||||
|
||||
|
||||
async def _run_rust_aocr(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
fallback: Callable[[], Coroutine[object, object, OCRResponse]],
|
||||
) -> OCRResponse:
|
||||
return await rust_ocr_bridge.adispatch_ocr(
|
||||
prepare=lambda: _prepare_rust_ocr_call(
|
||||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
),
|
||||
fallback=fallback,
|
||||
adapt=OCRResponse.model_validate,
|
||||
model=prepared_request.model,
|
||||
provider=prepared_request.custom_llm_provider,
|
||||
eligible=_rust_ocr_supported(prepared_request),
|
||||
)
|
||||
|
||||
|
||||
@client
|
||||
|
|
@ -319,7 +422,31 @@ async def aocr(
|
|||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str)
|
||||
async def python_fallback() -> OCRResponse:
|
||||
pending: Final = base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared.effective_timeout,
|
||||
logging_obj=prepared.litellm_logging_obj,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared.custom_llm_provider,
|
||||
aocr=True,
|
||||
headers=prepared.extra_headers,
|
||||
provider_config=prepared.provider_config,
|
||||
litellm_params=prepared.litellm_params,
|
||||
)
|
||||
response: Final = await pending if asyncio.iscoroutine(pending) else pending
|
||||
if response is None:
|
||||
raise ValueError(f"Got an unexpected None response from the OCR API: {response}")
|
||||
return response
|
||||
|
||||
return await _run_rust_aocr(
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
fallback=python_fallback,
|
||||
)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
@ -560,7 +687,27 @@ def ocr(
|
|||
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
return _execute_ocr(prepared_request=prepared, resolve_api_key=get_secret_str, is_async=_is_async)
|
||||
def python_fallback() -> OCRResponse | Coroutine[object, object, OCRResponse]:
|
||||
return base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared.effective_timeout,
|
||||
logging_obj=prepared.litellm_logging_obj,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared.custom_llm_provider,
|
||||
aocr=_is_async,
|
||||
headers=prepared.extra_headers,
|
||||
provider_config=prepared.provider_config,
|
||||
litellm_params=prepared.litellm_params,
|
||||
)
|
||||
|
||||
return _run_rust_ocr(
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
fallback=python_fallback,
|
||||
)
|
||||
except Exception as e:
|
||||
raise litellm.exception_type(
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -4,12 +4,16 @@ The Rust core owns the conversation translation, the provider call, and the
|
|||
response normalization for the subset of `/chat/completions` requests it
|
||||
accepts. This module only marshals inputs and hands the normalized result to
|
||||
LiteLLM's existing `ModelResponse` builder.
|
||||
|
||||
``None`` means the provider was never called, so the caller is free to serve the
|
||||
request on the Python path. A failure after the call was issued raises instead:
|
||||
retrying it there would bill the customer for the same work twice.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Mapping, Sequence
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Final, Protocol
|
||||
|
||||
import httpx
|
||||
|
|
@ -20,14 +24,28 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo
|
|||
convert_to_model_response_object,
|
||||
)
|
||||
from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.protocols import (
|
||||
RustAchatCompletions,
|
||||
RustChatCompletions,
|
||||
RustChatCompletionsDecline,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeChatCompletionsRequest,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
provider_connection_params,
|
||||
provider_request_params,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointBinding,
|
||||
EndpointDispatch,
|
||||
async_none,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -85,10 +103,16 @@ def response_logger(
|
|||
return log
|
||||
|
||||
|
||||
_CHAT: Final[NativeBinding[RustChatCompletions]] = NativeBinding(lambda native: native.chat_completions)
|
||||
_ACHAT: Final[NativeBinding[RustAchatCompletions]] = NativeBinding(lambda native: native.achat_completions)
|
||||
_CHAT_PREFLIGHT: Final[NativeBinding[RustChatCompletionsDecline]] = NativeBinding(
|
||||
lambda native: native.chat_completions_decline
|
||||
_CHAT: Final[EndpointDispatch[RustChatCompletions, RustAchatCompletions]] = EndpointDispatch.native(
|
||||
route="chat_completions",
|
||||
sync=lambda native: native.chat_completions,
|
||||
asynchronous=lambda native: native.achat_completions,
|
||||
enabled=rust_enabled,
|
||||
)
|
||||
_CHAT_PREFLIGHT: Final[EndpointBinding[RustChatCompletionsDecline]] = EndpointBinding.native(
|
||||
route="chat_completions",
|
||||
select=lambda native: native.chat_completions_decline,
|
||||
enabled=rust_enabled,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -102,14 +126,14 @@ def set_rust_chat_completions(
|
|||
patching module attributes."""
|
||||
if not isinstance(chat_completions, Unchanged):
|
||||
if chat_completions is None:
|
||||
_CHAT.reset()
|
||||
_CHAT.sync.reset()
|
||||
else:
|
||||
_CHAT.override(chat_completions)
|
||||
_CHAT.sync.override(chat_completions)
|
||||
if not isinstance(achat_completions, Unchanged):
|
||||
if achat_completions is None:
|
||||
_ACHAT.reset()
|
||||
_CHAT.asynchronous.reset()
|
||||
else:
|
||||
_ACHAT.override(achat_completions)
|
||||
_CHAT.asynchronous.override(achat_completions)
|
||||
if not isinstance(decline, Unchanged):
|
||||
if decline is None:
|
||||
_CHAT_PREFLIGHT.reset()
|
||||
|
|
@ -178,24 +202,14 @@ def rust_chat_completions_accepts(
|
|||
if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params):
|
||||
verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path")
|
||||
return False
|
||||
if not rust_enabled():
|
||||
return False
|
||||
decline: Final = _CHAT_PREFLIGHT.load()
|
||||
if decline is None:
|
||||
return False
|
||||
try:
|
||||
reason: Final = decline(
|
||||
return _CHAT_PREFLIGHT.accepts(
|
||||
check=lambda decline: decline(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
except Exception as error: # noqa: BLE001 # capability checks perform no provider I/O
|
||||
verbose_logger.debug("Native chat acceptance check failed: %s", error)
|
||||
return False
|
||||
if reason is not None:
|
||||
verbose_logger.debug("Native chat request is ineligible: %s", reason)
|
||||
return reason is None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _build_model_response(
|
||||
|
|
@ -224,31 +238,32 @@ def chat_completions(
|
|||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
eligible: bool = True,
|
||||
) -> DispatchResult[ModelResponse]:
|
||||
) -> ModelResponse | None:
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
def call(native: RustChatCompletions, timeout_seconds: float | None) -> Mapping[str, object]:
|
||||
return native(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
|
||||
return attempt(
|
||||
load=_CHAT.load,
|
||||
enabled=rust_enabled(),
|
||||
eligible=eligible,
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=call,
|
||||
return _CHAT.invoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeChatCompletionsRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=lambda: None,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -264,29 +279,81 @@ async def achat_completions(
|
|||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
eligible: bool = True,
|
||||
) -> DispatchResult[ModelResponse]:
|
||||
) -> ModelResponse | None:
|
||||
def adapt(rust_response: Mapping[str, object]) -> ModelResponse:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
async def call(native: RustAchatCompletions, timeout_seconds: float | None) -> Mapping[str, object]:
|
||||
return await native(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
)
|
||||
|
||||
return await aattempt(
|
||||
load=_ACHAT.load,
|
||||
enabled=rust_enabled(),
|
||||
eligible=eligible,
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=call,
|
||||
return await _CHAT.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeChatCompletionsRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=async_none,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
async def achat_completions_or_fallback(
|
||||
*,
|
||||
model: str,
|
||||
messages: Sequence[object],
|
||||
optional_params: Mapping[str, object],
|
||||
model_response: ModelResponse,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: Mapping[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
on_response: ResponseObserver,
|
||||
python_fallback: Callable[[], Awaitable[object]],
|
||||
) -> object:
|
||||
"""Await the Rust path, falling back to the caller's own Python path when
|
||||
the bridge is unavailable or the call fails.
|
||||
|
||||
The caller supplies the fallback, so the bridge stays free of provider
|
||||
dispatch. This exists because a caller that dispatches asynchronously has
|
||||
already returned a coroutine by the time a Rust failure surfaces, and so
|
||||
cannot fall back on its own.
|
||||
"""
|
||||
|
||||
def adapt(rust_response: Mapping[str, object]) -> object:
|
||||
on_response(rust_response)
|
||||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
return await _CHAT.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeChatCompletionsRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -6,13 +6,30 @@ from typing import Final
|
|||
|
||||
import httpx
|
||||
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.protocols import RustAmessages, RustMessages
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeMessagesRequest,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointDispatch,
|
||||
always_enabled,
|
||||
async_none,
|
||||
identity,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
_MESSAGES: Final[NativeBinding[RustMessages]] = NativeBinding(lambda native: native.messages)
|
||||
_AMESSAGES: Final[NativeBinding[RustAmessages]] = NativeBinding(lambda native: native.amessages)
|
||||
_MESSAGES: Final[EndpointDispatch[RustMessages, RustAmessages]] = EndpointDispatch.native(
|
||||
route="messages",
|
||||
sync=lambda native: native.messages,
|
||||
asynchronous=lambda native: native.amessages,
|
||||
enabled=always_enabled,
|
||||
)
|
||||
|
||||
|
||||
def set_rust_messages(
|
||||
|
|
@ -22,22 +39,22 @@ def set_rust_messages(
|
|||
) -> None:
|
||||
if not isinstance(messages, Unchanged):
|
||||
if messages is None:
|
||||
_MESSAGES.reset()
|
||||
_MESSAGES.sync.reset()
|
||||
else:
|
||||
_MESSAGES.override(messages)
|
||||
_MESSAGES.sync.override(messages)
|
||||
if not isinstance(amessages, Unchanged):
|
||||
if amessages is None:
|
||||
_AMESSAGES.reset()
|
||||
_MESSAGES.asynchronous.reset()
|
||||
else:
|
||||
_AMESSAGES.override(amessages)
|
||||
_MESSAGES.asynchronous.override(amessages)
|
||||
|
||||
|
||||
def load_rust_messages() -> RustMessages | None:
|
||||
return _MESSAGES.load()
|
||||
return _MESSAGES.sync.load()
|
||||
|
||||
|
||||
def load_rust_amessages() -> RustAmessages | None:
|
||||
return _AMESSAGES.load()
|
||||
return _MESSAGES.asynchronous.load()
|
||||
|
||||
|
||||
def messages(
|
||||
|
|
@ -49,22 +66,26 @@ def messages(
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[dict[str, object]]:
|
||||
return attempt(
|
||||
load=_MESSAGES.load,
|
||||
enabled=True,
|
||||
eligible=True,
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_messages, timeout_seconds: rust_messages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
) -> dict[str, object] | None:
|
||||
return _MESSAGES.invoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeMessagesRequest(
|
||||
model=model,
|
||||
body=body,
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -77,20 +98,24 @@ async def amessages(
|
|||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[dict[str, object]]:
|
||||
return await aattempt(
|
||||
load=_AMESSAGES.load,
|
||||
enabled=True,
|
||||
eligible=True,
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_amessages, timeout_seconds: rust_amessages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
) -> dict[str, object] | None:
|
||||
return await _MESSAGES.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeMessagesRequest(
|
||||
model=model,
|
||||
body=body,
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,211 +1,90 @@
|
|||
"""Thin Python wrapper for the native Rust OCR bridge."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Final
|
||||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from typing import Final, TypeVar
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter
|
||||
from . import configuration as _configuration
|
||||
from .bindings import UNCHANGED, Unchanged
|
||||
from .protocols import RustAocr, RustOcr
|
||||
from .request import NativeOCRRequest, PreparedNativeCall, call_native
|
||||
from .runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointDispatch,
|
||||
)
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model
|
||||
from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse
|
||||
from litellm.rust_bridge import configuration as _configuration
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.protocols import RustAocr, RustOcr
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
_OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr)
|
||||
_AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr)
|
||||
_HEADERS: Final = TypeAdapter(dict[str, object])
|
||||
rust_ocr_enabled = _configuration.rust_ocr_enabled
|
||||
rust = _configuration.rust
|
||||
ResultT = TypeVar("ResultT")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PreparedOCRRequest:
|
||||
model: str
|
||||
document: dict[str, object]
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
custom_llm_provider: str
|
||||
extra_headers: dict[str, object] | None
|
||||
provider_config: BaseOCRConfig
|
||||
optional_params: dict[str, object]
|
||||
litellm_params: dict[str, object]
|
||||
effective_timeout: float | httpx.Timeout
|
||||
litellm_logging_obj: LiteLLMLoggingObj
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _PreparedRustOCRCall:
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
headers: dict[str, object]
|
||||
optional_params: dict[str, object]
|
||||
|
||||
|
||||
_RUST_OCR_PROVIDERS: Final = frozenset(
|
||||
{
|
||||
"mistral",
|
||||
"azure_ai",
|
||||
"vertex_ai",
|
||||
}
|
||||
_OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native(
|
||||
route="ocr",
|
||||
sync=lambda native: native.ocr,
|
||||
asynchronous=lambda native: native.aocr,
|
||||
enabled=_configuration.rust_ocr_enabled,
|
||||
)
|
||||
|
||||
|
||||
def set_rust_ocr(
|
||||
*,
|
||||
ocr: RustOcr | None | Unchanged = UNCHANGED,
|
||||
aocr: RustAocr | None | Unchanged = UNCHANGED,
|
||||
) -> None:
|
||||
if not isinstance(ocr, Unchanged):
|
||||
if ocr is None:
|
||||
_OCR.sync.reset()
|
||||
else:
|
||||
_OCR.sync.override(ocr)
|
||||
if not isinstance(aocr, Unchanged):
|
||||
if aocr is None:
|
||||
_OCR.asynchronous.reset()
|
||||
else:
|
||||
_OCR.asynchronous.override(aocr)
|
||||
|
||||
|
||||
def load_rust_ocr() -> RustOcr | None:
|
||||
return _OCR.load()
|
||||
return _OCR.sync.load()
|
||||
|
||||
|
||||
def load_rust_aocr() -> RustAocr | None:
|
||||
return _AOCR.load()
|
||||
return _OCR.asynchronous.load()
|
||||
|
||||
|
||||
def _rust_ocr_supported(prepared_request: PreparedOCRRequest) -> bool:
|
||||
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
|
||||
return False
|
||||
if not prepared_request.provider_config.supports_rust_bridge():
|
||||
return False
|
||||
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
|
||||
|
||||
|
||||
def _rust_bridge_optional_params(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> dict[str, object]:
|
||||
if prepared_request.custom_llm_provider != "vertex_ai":
|
||||
return prepared_request.optional_params
|
||||
vertex_project: Final = (
|
||||
prepared_request.litellm_params.get("vertex_project")
|
||||
or prepared_request.litellm_params.get("vertex_ai_project")
|
||||
or litellm.vertex_project
|
||||
or resolve_secret("VERTEXAI_PROJECT")
|
||||
)
|
||||
vertex_location: Final = (
|
||||
prepared_request.litellm_params.get("vertex_location")
|
||||
or prepared_request.litellm_params.get("vertex_ai_location")
|
||||
or litellm.vertex_location
|
||||
or resolve_secret("VERTEXAI_LOCATION")
|
||||
or resolve_secret("VERTEX_LOCATION")
|
||||
)
|
||||
return {
|
||||
**prepared_request.optional_params,
|
||||
**{
|
||||
name: value
|
||||
for name, value in (("vertex_project", vertex_project), ("vertex_location", vertex_location))
|
||||
if value is not None
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def _rust_bridge_api_base(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> str | None:
|
||||
if prepared_request.api_base is not None:
|
||||
return prepared_request.api_base
|
||||
if prepared_request.custom_llm_provider == "azure_ai":
|
||||
if is_azure_document_intelligence_model(prepared_request.model):
|
||||
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
return resolve_secret("AZURE_AI_API_BASE")
|
||||
return None
|
||||
|
||||
|
||||
def _prepare_rust_ocr_call(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> _PreparedRustOCRCall:
|
||||
provider_config: Final = prepared_request.provider_config
|
||||
api_key_env_var: Final = provider_config.get_api_key_env_var()
|
||||
resolved_api_key: Final = prepared_request.api_key or (
|
||||
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
|
||||
)
|
||||
resolved_headers: Final = _HEADERS.validate_python(
|
||||
provider_config.validate_environment(
|
||||
headers=prepared_request.extra_headers or {},
|
||||
model=prepared_request.model,
|
||||
api_key=resolved_api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
)
|
||||
resolved_complete_url: Final = provider_config.get_complete_url(
|
||||
api_base=prepared_request.api_base,
|
||||
model=prepared_request.model,
|
||||
optional_params=prepared_request.optional_params,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
|
||||
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
|
||||
prepared_request.litellm_logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=resolved_api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": {
|
||||
"model": prepared_request.model,
|
||||
"document": prepared_request.document,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": resolved_complete_url,
|
||||
"headers": resolved_headers,
|
||||
},
|
||||
)
|
||||
return _PreparedRustOCRCall(
|
||||
api_key=resolved_api_key,
|
||||
api_base=rust_api_base,
|
||||
headers=resolved_headers,
|
||||
optional_params=rust_optional_params,
|
||||
def dispatch_ocr(
|
||||
*,
|
||||
prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]],
|
||||
fallback: Callable[[], ResultT],
|
||||
adapt: Callable[[Mapping[str, object]], ResultT],
|
||||
model: str,
|
||||
provider: str,
|
||||
eligible: bool,
|
||||
) -> ResultT:
|
||||
return _OCR.invoke(
|
||||
prepare=prepare,
|
||||
call=call_native,
|
||||
fallback=fallback,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=provider, model=model),
|
||||
eligible=eligible,
|
||||
)
|
||||
|
||||
|
||||
def attempt_ocr(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> DispatchResult[OCRResponse]:
|
||||
return attempt(
|
||||
load=_OCR.load,
|
||||
enabled=_configuration.rust_enabled(),
|
||||
prepare=lambda: _prepare_rust_ocr_call(
|
||||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
),
|
||||
call=lambda native, prepared: native(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
|
||||
),
|
||||
adapt=OCRResponse.model_validate,
|
||||
eligible=_rust_ocr_supported(prepared_request),
|
||||
)
|
||||
|
||||
|
||||
async def aattempt_ocr(
|
||||
prepared_request: PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> DispatchResult[OCRResponse]:
|
||||
return await aattempt(
|
||||
load=_AOCR.load,
|
||||
enabled=_configuration.rust_enabled(),
|
||||
prepare=lambda: _prepare_rust_ocr_call(
|
||||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
),
|
||||
call=lambda native, prepared: native(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout),
|
||||
),
|
||||
adapt=OCRResponse.model_validate,
|
||||
eligible=_rust_ocr_supported(prepared_request),
|
||||
async def adispatch_ocr(
|
||||
*,
|
||||
prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]],
|
||||
fallback: Callable[[], Awaitable[ResultT]],
|
||||
adapt: Callable[[Mapping[str, object]], ResultT],
|
||||
model: str,
|
||||
provider: str,
|
||||
eligible: bool,
|
||||
) -> ResultT:
|
||||
return await _OCR.ainvoke(
|
||||
prepare=prepare,
|
||||
call=call_native,
|
||||
fallback=fallback,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=provider, model=model),
|
||||
eligible=eligible,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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]]: ...
|
||||
|
|
|
|||
122
litellm/rust_bridge/request.py
Normal file
122
litellm/rust_bridge/request.py
Normal 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
|
||||
|
|
@ -2,24 +2,36 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncGenerator
|
||||
from contextlib import AbstractAsyncContextManager, asynccontextmanager
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from websockets.exceptions import ConnectionClosedOK
|
||||
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.protocols import (
|
||||
RustResponsesWebSocket,
|
||||
RustResponsesWebSocketConnection,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
NativeResponsesWebSocketRequest,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointBinding,
|
||||
async_none,
|
||||
identity,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
_RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding(
|
||||
lambda native: native.ResponsesWebSocketConnection,
|
||||
_RESPONSES_WEBSOCKET: Final[EndpointBinding[RustResponsesWebSocketConnection]] = EndpointBinding.native(
|
||||
route="responses_websocket",
|
||||
select=lambda native: native.ResponsesWebSocketConnection,
|
||||
enabled=rust_enabled,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -34,7 +46,7 @@ def set_rust_responses_websocket(
|
|||
_RESPONSES_WEBSOCKET.override(connection)
|
||||
|
||||
|
||||
class ConnectionAdapter:
|
||||
class _ConnectionAdapter:
|
||||
def __init__(self, connection: RustResponsesWebSocket):
|
||||
self._connection: Final[RustResponsesWebSocket] = connection
|
||||
|
||||
|
|
@ -56,34 +68,18 @@ async def connect(
|
|||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[ConnectionAdapter]:
|
||||
return await aattempt(
|
||||
load=_RESPONSES_WEBSOCKET.load,
|
||||
enabled=rust_enabled(),
|
||||
eligible=True,
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda connection_type, timeout_seconds: connection_type.connect(
|
||||
url=url,
|
||||
headers=headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
) -> _ConnectionAdapter | None:
|
||||
connection: Final = await _RESPONSES_WEBSOCKET.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeResponsesWebSocketRequest(
|
||||
url=url,
|
||||
options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
adapt=ConnectionAdapter,
|
||||
call=lambda connection_type, request: call_native(connection_type.connect, request),
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider="openai", model="responses websocket"),
|
||||
)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def _connection_context(connection: ConnectionAdapter) -> AsyncGenerator[ConnectionAdapter, None]:
|
||||
try:
|
||||
yield connection
|
||||
finally:
|
||||
await connection.close()
|
||||
|
||||
|
||||
async def managed_connect(
|
||||
*,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]:
|
||||
result: Final = await connect(url=url, headers=headers, timeout=timeout)
|
||||
return adapt_result(result, _connection_context)
|
||||
return None if connection is None else _ConnectionAdapter(connection)
|
||||
|
|
|
|||
|
|
@ -4,38 +4,58 @@ from typing import Final
|
|||
|
||||
import httpx
|
||||
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription
|
||||
from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
NativeTranscriptionRequest,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
provider_connection_params,
|
||||
provider_request_params,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointDispatch,
|
||||
always_enabled,
|
||||
async_none,
|
||||
identity,
|
||||
)
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
_TRANSCRIPTION: Final[NativeBinding[RustTranscription]] = NativeBinding(lambda native: native.transcription)
|
||||
_ATRANSCRIPTION: Final[NativeBinding[RustAtranscription]] = NativeBinding(lambda native: native.atranscription)
|
||||
_TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] = EndpointDispatch.native(
|
||||
route="audio transcription",
|
||||
sync=lambda native: native.transcription,
|
||||
asynchronous=lambda native: native.atranscription,
|
||||
enabled=always_enabled,
|
||||
)
|
||||
|
||||
|
||||
def configure_rust_transcription(
|
||||
enabled: bool = True,
|
||||
*,
|
||||
transcription: RustTranscription | None | Unchanged = UNCHANGED,
|
||||
atranscription: RustAtranscription | None | Unchanged = UNCHANGED,
|
||||
) -> None:
|
||||
if not isinstance(transcription, Unchanged):
|
||||
if transcription is None:
|
||||
_TRANSCRIPTION.reset()
|
||||
_TRANSCRIPTION.sync.reset()
|
||||
else:
|
||||
_TRANSCRIPTION.override(transcription)
|
||||
_TRANSCRIPTION.sync.override(transcription)
|
||||
if not isinstance(atranscription, Unchanged):
|
||||
if atranscription is None:
|
||||
_ATRANSCRIPTION.reset()
|
||||
_TRANSCRIPTION.asynchronous.reset()
|
||||
else:
|
||||
_ATRANSCRIPTION.override(atranscription)
|
||||
_TRANSCRIPTION.asynchronous.override(atranscription)
|
||||
|
||||
|
||||
def load_rust_transcription() -> RustTranscription | None:
|
||||
return _TRANSCRIPTION.load()
|
||||
return _TRANSCRIPTION.sync.load()
|
||||
|
||||
|
||||
def load_rust_atranscription() -> RustAtranscription | None:
|
||||
return _ATRANSCRIPTION.load()
|
||||
return _TRANSCRIPTION.asynchronous.load()
|
||||
|
||||
|
||||
def transcription(
|
||||
|
|
@ -48,23 +68,28 @@ def transcription(
|
|||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[dict[str, object]]:
|
||||
return attempt(
|
||||
load=_TRANSCRIPTION.load,
|
||||
enabled=True,
|
||||
eligible=True,
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_transcription, timeout_seconds: rust_transcription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_seconds,
|
||||
) -> dict[str, object] | None:
|
||||
return _TRANSCRIPTION.invoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeTranscriptionRequest(
|
||||
model=model,
|
||||
audio=audio,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -78,21 +103,26 @@ async def atranscription(
|
|||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> DispatchResult[dict[str, object]]:
|
||||
return await aattempt(
|
||||
load=_ATRANSCRIPTION.load,
|
||||
enabled=True,
|
||||
eligible=True,
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_atranscription, timeout_seconds: rust_atranscription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_seconds,
|
||||
) -> dict[str, object] | None:
|
||||
return await _TRANSCRIPTION.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeTranscriptionRequest(
|
||||
model=model,
|
||||
audio=audio,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ import pytest
|
|||
import litellm
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
|
||||
from litellm.rust_bridge.request import NativeMessagesRequest, NativeRequestContext
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
|
|
@ -41,23 +41,19 @@ class RecordingMessages:
|
|||
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
body: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
request: NativeMessagesRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"body": body,
|
||||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"model": request.model,
|
||||
"body": request.body,
|
||||
"api_key": request.options.api_key,
|
||||
"api_base": request.options.api_base,
|
||||
"custom_llm_provider": request.options.custom_llm_provider,
|
||||
"extra_headers": request.options.extra_headers,
|
||||
"timeout_seconds": request.options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
|
@ -69,23 +65,19 @@ class RecordingAsyncMessages:
|
|||
|
||||
async def __call__(
|
||||
self,
|
||||
model: str,
|
||||
body: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
request: NativeMessagesRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"body": body,
|
||||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"model": request.model,
|
||||
"body": request.body,
|
||||
"api_key": request.options.api_key,
|
||||
"api_base": request.options.api_base,
|
||||
"custom_llm_provider": request.options.custom_llm_provider,
|
||||
"extra_headers": request.options.extra_headers,
|
||||
"timeout_seconds": request.options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
|
@ -95,7 +87,7 @@ class ExplodingAsyncMessages:
|
|||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, **kwargs: object) -> dict[str, object]:
|
||||
async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise AssertionError("bridge must not be called")
|
||||
|
||||
|
|
@ -104,7 +96,7 @@ class RaisingAsyncMessages:
|
|||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, **kwargs: object) -> dict[str, object]:
|
||||
async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise RuntimeError("upstream request failed with status 400: bad request")
|
||||
|
||||
|
|
@ -127,6 +119,16 @@ def test_load_rust_messages_returns_injected_impl():
|
|||
assert rust_messages.load_rust_messages() is bridge
|
||||
|
||||
|
||||
def test_bare_rust_still_toggles_ocr():
|
||||
from litellm.rust_bridge.ocr import rust_ocr_enabled
|
||||
|
||||
litellm.rust(True)
|
||||
assert rust_ocr_enabled() is True
|
||||
|
||||
litellm.rust(False)
|
||||
assert rust_ocr_enabled() is False
|
||||
|
||||
|
||||
def test_load_rust_amessages_returns_injected_impl():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
|
|
@ -134,7 +136,7 @@ def test_load_rust_amessages_returns_injected_impl():
|
|||
assert rust_messages.load_rust_amessages() is bridge
|
||||
|
||||
|
||||
def test_messages_wrapper_reports_unavailable(monkeypatch):
|
||||
def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge.bindings"),
|
||||
"get_native_bridge",
|
||||
|
|
@ -151,7 +153,7 @@ def test_messages_wrapper_reports_unavailable(monkeypatch):
|
|||
extra_headers={},
|
||||
timeout=30.0,
|
||||
)
|
||||
assert result == NativeSkipped(NativeSkipReason.UNAVAILABLE)
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_messages_wrapper_forwards_args_and_converts_timeout():
|
||||
|
|
@ -169,7 +171,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout():
|
|||
timeout=httpx.Timeout(600.0, read=42.0),
|
||||
)
|
||||
|
||||
assert response == Handled(FAKE_MESSAGES_RESPONSE)
|
||||
assert response == FAKE_MESSAGES_RESPONSE
|
||||
assert bridge.calls[0] == {
|
||||
"model": "claude-sonnet-4-5",
|
||||
"body": REQUEST_BODY,
|
||||
|
|
@ -197,7 +199,7 @@ async def test_amessages_wrapper_forwards_args():
|
|||
timeout=12.5,
|
||||
)
|
||||
|
||||
assert response == Handled(FAKE_MESSAGES_RESPONSE)
|
||||
assert response == FAKE_MESSAGES_RESPONSE
|
||||
assert bridge.calls[0]["model"] == "claude-sonnet-4-5"
|
||||
assert bridge.calls[0]["timeout_seconds"] == 12.5
|
||||
|
||||
|
|
@ -215,7 +217,7 @@ def _gate(**overrides):
|
|||
"timeout": 30.0,
|
||||
}
|
||||
kwargs.update(overrides)
|
||||
return BaseLLMHTTPHandler._attempt_rust_anthropic_messages(**kwargs)
|
||||
return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -226,8 +228,7 @@ async def test_gate_invokes_rust_and_marks_response_header():
|
|||
|
||||
response = await _gate()
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response is not None
|
||||
assert response["id"] == "msg_123"
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
call = bridge.calls[0]
|
||||
|
|
@ -240,13 +241,13 @@ async def test_gate_invokes_rust_and_marks_response_header():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_reports_failure_to_harness():
|
||||
async def test_gate_propagates_unknown_native_errors():
|
||||
bridge = RaisingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate()
|
||||
assert isinstance(response, NativeFailed)
|
||||
with pytest.raises(RuntimeError, match="bad request"):
|
||||
await _gate()
|
||||
assert bridge.calls == 1
|
||||
|
||||
|
||||
|
|
@ -257,7 +258,7 @@ async def test_gate_skips_rust_when_flag_absent():
|
|||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert response is None
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
|
|
@ -269,11 +270,22 @@ async def test_gate_uses_process_enable_without_request_override():
|
|||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure"))
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response is not None
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "azure_ai"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_ignores_request_flag_when_process_enabled():
|
||||
bridge = RecordingAsyncMessages()
|
||||
litellm.rust(True)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
|
||||
response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False))
|
||||
|
||||
assert response is not None
|
||||
assert len(bridge.calls) == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_invokes_rust_for_native_anthropic_provider():
|
||||
bridge = RecordingAsyncMessages()
|
||||
|
|
@ -288,8 +300,7 @@ async def test_gate_invokes_rust_for_native_anthropic_provider():
|
|||
headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"},
|
||||
)
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response is not None
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
|
||||
assert bridge.calls[0]["api_key"] == "sk-ant"
|
||||
|
|
@ -306,8 +317,7 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch):
|
|||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
)
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response is not None
|
||||
assert bridge.calls[0]["custom_llm_provider"] == "anthropic"
|
||||
|
||||
|
||||
|
|
@ -322,7 +332,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch):
|
|||
litellm_params=GenericLiteLLMParams(api_key="sk-ant"),
|
||||
)
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert response is None
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
|
|
@ -334,7 +344,7 @@ async def test_gate_skips_rust_for_unsupported_provider():
|
|||
|
||||
response = await _gate(custom_llm_provider="openai")
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert response is None
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
|
|
@ -346,7 +356,7 @@ async def test_gate_skips_rust_for_agentic_hook():
|
|||
|
||||
response = await _gate(has_agentic_hook=True)
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
assert response is None
|
||||
assert bridge.calls == 0
|
||||
|
||||
|
||||
|
|
@ -362,8 +372,7 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag():
|
|||
request_body=streaming_body,
|
||||
)
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert response is not None
|
||||
assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
assert "stream" not in bridge.calls[0]["body"]
|
||||
assert bridge.calls[0]["body"] == REQUEST_BODY
|
||||
|
|
@ -396,105 +405,4 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch):
|
|||
|
||||
response = await _gate()
|
||||
|
||||
assert isinstance(response, NativeSkipped)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("selection", ("native", "disabled", "failed", "declined", "upstream"))
|
||||
async def test_messages_handler_runs_selected_backend_once(selection: str, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from datetime import datetime
|
||||
from types import SimpleNamespace
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.exceptions import RateLimitError
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import AnthropicMessagesConfig
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
from litellm.rust_bridge import bindings
|
||||
|
||||
class Declined(Exception):
|
||||
pass
|
||||
|
||||
class Upstream(Exception):
|
||||
pass
|
||||
|
||||
error = (
|
||||
Upstream(429, "rate limited")
|
||||
if selection == "upstream"
|
||||
else Declined("unsupported")
|
||||
if selection == "declined"
|
||||
else RuntimeError("native failed")
|
||||
if selection == "failed"
|
||||
else None
|
||||
)
|
||||
|
||||
class Native:
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, **kwargs: object) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
if error is not None:
|
||||
raise error
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
||||
bridge = Native()
|
||||
monkeypatch.setattr(
|
||||
bindings,
|
||||
"get_native_bridge",
|
||||
lambda: SimpleNamespace(
|
||||
RustBridgeDeclined=Declined,
|
||||
RustUpstreamError=Upstream,
|
||||
),
|
||||
)
|
||||
rust_messages.set_rust_messages(amessages=bridge)
|
||||
litellm.rust(selection != "disabled")
|
||||
requests: list[httpx.Request] = []
|
||||
|
||||
def respond(request: httpx.Request) -> httpx.Response:
|
||||
requests.append(request)
|
||||
return httpx.Response(200, json=FAKE_MESSAGES_RESPONSE)
|
||||
|
||||
logging_obj = Logging(
|
||||
model=FAKE_MESSAGES_RESPONSE["model"],
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="anthropic_messages",
|
||||
start_time=datetime.now(),
|
||||
litellm_call_id="harness-test",
|
||||
function_id="harness-test",
|
||||
)
|
||||
client = AsyncHTTPHandler()
|
||||
await client.client.aclose()
|
||||
async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as transport:
|
||||
client.client = transport
|
||||
|
||||
async def run():
|
||||
return await BaseLLMHTTPHandler().async_anthropic_messages_handler(
|
||||
model=FAKE_MESSAGES_RESPONSE["model"],
|
||||
messages=[{"role": "user", "content": "hello"}],
|
||||
anthropic_messages_provider_config=AnthropicMessagesConfig(),
|
||||
anthropic_messages_optional_request_params={"max_tokens": 10},
|
||||
custom_llm_provider="anthropic",
|
||||
litellm_params=GenericLiteLLMParams(),
|
||||
logging_obj=logging_obj,
|
||||
api_key="sk-test",
|
||||
api_base="https://example.test",
|
||||
client=client,
|
||||
)
|
||||
|
||||
if selection in ("failed", "upstream"):
|
||||
with pytest.raises(RateLimitError if selection == "upstream" else RuntimeError) as caught:
|
||||
await run()
|
||||
if selection == "upstream":
|
||||
assert caught.value.__cause__ is error
|
||||
assert caught.value.llm_provider == "anthropic"
|
||||
assert caught.value.model == FAKE_MESSAGES_RESPONSE["model"]
|
||||
else:
|
||||
assert caught.value is error
|
||||
else:
|
||||
response = await run()
|
||||
assert response["id"] == FAKE_MESSAGES_RESPONSE["id"]
|
||||
assert len(requests) == (1 if selection in ("disabled", "declined") else 0)
|
||||
assert bridge.calls == (0 if selection == "disabled" else 1)
|
||||
assert response is None
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import pytest
|
|||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.runtime import Handled
|
||||
from litellm.rust_bridge.request import NativeOCRRequest, NativeRequestContext, NativeRequestOptions, PreparedNativeCall
|
||||
from litellm.rust_bridge.timeouts import timeout_to_seconds
|
||||
|
||||
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
|
||||
|
|
@ -49,25 +49,20 @@ class RecordingBridge:
|
|||
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeOCRRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"document": document,
|
||||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"optional_params": optional_params,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"model": request.model,
|
||||
"document": request.document,
|
||||
"api_key": request.options.api_key,
|
||||
"api_base": request.options.api_base,
|
||||
"custom_llm_provider": request.options.custom_llm_provider,
|
||||
"extra_headers": request.options.extra_headers,
|
||||
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
|
||||
"timeout_seconds": request.options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_OCR_RESPONSE)
|
||||
|
|
@ -81,25 +76,20 @@ class RecordingAsyncBridge:
|
|||
|
||||
async def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeOCRRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"document": document,
|
||||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"optional_params": optional_params,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"model": request.model,
|
||||
"document": request.document,
|
||||
"api_key": request.options.api_key,
|
||||
"api_base": request.options.api_base,
|
||||
"custom_llm_provider": request.options.custom_llm_provider,
|
||||
"extra_headers": request.options.extra_headers,
|
||||
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
|
||||
"timeout_seconds": request.options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_OCR_RESPONSE)
|
||||
|
|
@ -108,14 +98,9 @@ class RecordingAsyncBridge:
|
|||
class RaisingBridge:
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeOCRRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
raise RuntimeError("bridge failed")
|
||||
|
||||
|
|
@ -123,14 +108,9 @@ class RaisingBridge:
|
|||
class RaisingAsyncBridge:
|
||||
async def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeOCRRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
raise RuntimeError("bridge failed")
|
||||
|
||||
|
|
@ -164,9 +144,6 @@ class FakeOCRConfig:
|
|||
def get_api_key_env_var(self) -> str:
|
||||
return self.api_key_env_var
|
||||
|
||||
def supports_rust_bridge(self) -> bool:
|
||||
return True
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
*,
|
||||
|
|
@ -203,7 +180,7 @@ def build_prepared_request(
|
|||
litellm_params: dict[str, object] | None = None,
|
||||
timeout: float | httpx.Timeout | None = 12.5,
|
||||
) -> Any:
|
||||
return rust_bridge.PreparedOCRRequest(
|
||||
return ocr_main._PreparedOCRRequest(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
|
|
@ -392,13 +369,103 @@ def test_timeout_to_seconds_handles_float_timeout_and_none():
|
|||
assert timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0
|
||||
|
||||
|
||||
def test_bridge_wrapper_forwards_prepared_args_and_wraps_response():
|
||||
bridge = RecordingBridge()
|
||||
|
||||
litellm.rust(True)
|
||||
|
||||
rust_bridge.set_rust_ocr(ocr=bridge)
|
||||
response = rust_bridge.dispatch_ocr(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeOCRRequest(
|
||||
model="mistral-ocr-latest",
|
||||
document=DOCUMENT,
|
||||
optional_params={"include_image_base64": True, "pages": [0]},
|
||||
options=NativeRequestOptions(
|
||||
api_key="sk-test",
|
||||
api_base="https://proxy.internal",
|
||||
custom_llm_provider="mistral",
|
||||
extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"},
|
||||
timeout_seconds=12.5,
|
||||
),
|
||||
),
|
||||
),
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
adapt=dict,
|
||||
model="mistral-ocr-latest",
|
||||
provider="mistral",
|
||||
eligible=True,
|
||||
)
|
||||
|
||||
assert response == FAKE_OCR_RESPONSE
|
||||
call = bridge.calls[0]
|
||||
assert call == {
|
||||
"model": "mistral-ocr-latest",
|
||||
"document": DOCUMENT,
|
||||
"api_key": "sk-test",
|
||||
"api_base": "https://proxy.internal",
|
||||
"custom_llm_provider": "mistral",
|
||||
"extra_headers": {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
},
|
||||
"optional_params": {"include_image_base64": True, "pages": [0]},
|
||||
"timeout_seconds": 12.5,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
|
||||
bridge = RecordingAsyncBridge()
|
||||
|
||||
litellm.rust(True)
|
||||
|
||||
rust_bridge.set_rust_ocr(aocr=bridge)
|
||||
|
||||
async def unexpected_fallback():
|
||||
pytest.fail("unexpected Python fallback")
|
||||
|
||||
response = await rust_bridge.adispatch_ocr(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
request=NativeOCRRequest(
|
||||
model="mistral-ocr-maas",
|
||||
document=DOCUMENT,
|
||||
optional_params={},
|
||||
options=NativeRequestOptions(
|
||||
custom_llm_provider="vertex_ai",
|
||||
provider_connection={"vertex_project": "project-1"},
|
||||
timeout_seconds=42.0,
|
||||
),
|
||||
),
|
||||
),
|
||||
fallback=unexpected_fallback,
|
||||
adapt=dict,
|
||||
model="mistral-ocr-maas",
|
||||
provider="vertex_ai",
|
||||
eligible=True,
|
||||
)
|
||||
|
||||
assert response == FAKE_OCR_RESPONSE
|
||||
assert bridge.calls[0] == {
|
||||
"model": "mistral-ocr-maas",
|
||||
"document": DOCUMENT,
|
||||
"api_key": None,
|
||||
"api_base": None,
|
||||
"custom_llm_provider": "vertex_ai",
|
||||
"extra_headers": None,
|
||||
"optional_params": {"vertex_project": "project-1"},
|
||||
"timeout_seconds": 42.0,
|
||||
}
|
||||
|
||||
|
||||
def test_run_rust_ocr_prepares_request_and_wraps_response():
|
||||
bridge = RecordingBridge()
|
||||
logging_obj = RecordingLogging()
|
||||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
|
||||
response = rust_bridge.attempt_ocr(
|
||||
response = ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
logging_obj=logging_obj,
|
||||
api_base="https://proxy.internal",
|
||||
|
|
@ -409,8 +476,6 @@ def test_run_rust_ocr_prepares_request_and_wraps_response():
|
|||
resolve_api_key=lambda _name: None,
|
||||
)
|
||||
|
||||
assert isinstance(response, Handled)
|
||||
response = response.value
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert response.pages[0].markdown == "hello world"
|
||||
assert bridge.calls[0] == {
|
||||
|
|
@ -433,7 +498,8 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
|||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(api_key=None, timeout=None),
|
||||
resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None,
|
||||
)
|
||||
|
|
@ -449,7 +515,8 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver():
|
|||
def _resolver(name: str) -> str | None:
|
||||
raise AssertionError(f"resolver should not be called for {name}")
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
api_key="sk-explicit",
|
||||
timeout=None,
|
||||
|
|
@ -470,7 +537,8 @@ def test_run_rust_ocr_uses_provider_api_key_env_var():
|
|||
resolver_calls.append(name)
|
||||
return "sk-provider-env"
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"),
|
||||
model="provider-ocr-model",
|
||||
|
|
@ -489,7 +557,8 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
|||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="vertex_ai",
|
||||
model="mistral-ocr-maas",
|
||||
|
|
@ -522,7 +591,8 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana
|
|||
"VERTEXAI_LOCATION": "us-east5",
|
||||
}.get(name)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="vertex_ai",
|
||||
model="mistral-ocr-maas",
|
||||
|
|
@ -540,7 +610,8 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
|
|||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="azure_ai",
|
||||
model="pixtral-12b-2409",
|
||||
|
|
@ -558,7 +629,8 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
|
|||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="azure_ai",
|
||||
model="doc-intelligence/prebuilt-layout",
|
||||
|
|
@ -579,7 +651,8 @@ def test_run_rust_ocr_runs_pre_call_logging():
|
|||
litellm.rust(True)
|
||||
rust_bridge._OCR.override(bridge)
|
||||
|
||||
rust_bridge.attempt_ocr(
|
||||
ocr_main._run_rust_ocr(
|
||||
fallback=lambda: pytest.fail("unexpected Python fallback"),
|
||||
prepared_request=build_prepared_request(
|
||||
logging_obj=logging_obj,
|
||||
api_base="https://api.mistral.ai/v1",
|
||||
|
|
@ -726,7 +799,7 @@ async def test_ocr_fallback_skips_native_preparation(
|
|||
def unexpected_preparation(*_args: object, **_kwargs: object) -> None:
|
||||
pytest.fail("Python fallback must not resolve native credentials or emit native pre_call")
|
||||
|
||||
monkeypatch.setattr(rust_bridge, "_prepare_rust_ocr_call", unexpected_preparation)
|
||||
monkeypatch.setattr(ocr_main, "_prepare_rust_ocr_call", unexpected_preparation)
|
||||
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fallback)
|
||||
|
||||
response: Final = (
|
||||
|
|
@ -739,25 +812,6 @@ async def test_ocr_fallback_skips_native_preparation(
|
|||
fallback.assert_called_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_rejects_empty_python_fallback_response(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def fake_exception_type(**kwargs: object) -> CapturedException:
|
||||
captured.update(kwargs)
|
||||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", AsyncMock(return_value=None))
|
||||
|
||||
with pytest.raises(CapturedException, match="wrapped"):
|
||||
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
||||
original: Final = captured["original_exception"]
|
||||
assert isinstance(original, ValueError)
|
||||
assert str(original) == "Got an unexpected None response from the OCR API: None"
|
||||
|
||||
|
||||
def test_ocr_provider_configs_expose_api_key_env_vars():
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import pytest
|
|||
|
||||
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
|
||||
from litellm.rust_bridge import configuration, responses_websocket
|
||||
from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason
|
||||
from litellm.rust_bridge.request import NativeRequestContext, NativeResponsesWebSocketRequest
|
||||
|
||||
|
||||
class _FakeNativeConnection:
|
||||
|
|
@ -31,10 +31,9 @@ class _FakeNativeBridge:
|
|||
@classmethod
|
||||
async def connect(
|
||||
cls,
|
||||
request: NativeResponsesWebSocketRequest,
|
||||
*,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout_seconds: float | None,
|
||||
context: NativeRequestContext,
|
||||
) -> _FakeNativeConnection:
|
||||
return _FakeNativeConnection()
|
||||
|
||||
|
|
@ -58,22 +57,25 @@ def test_rust_websocket_bridge_uses_process_enablement() -> None:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None:
|
||||
adapter = responses_websocket.ConnectionAdapter(_ClosedNativeConnection())
|
||||
adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection())
|
||||
|
||||
with pytest.raises(responses_websocket.ConnectionClosedOK):
|
||||
await adapter.recv()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bridge_reports_unavailable(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
configuration.rust(True)
|
||||
responses_websocket._RESPONSES_WEBSOCKET.override(None)
|
||||
|
||||
assert await responses_websocket.connect(
|
||||
url="wss://example.test/responses",
|
||||
headers={},
|
||||
timeout=None,
|
||||
) == NativeSkipped(NativeSkipReason.UNAVAILABLE)
|
||||
assert (
|
||||
await responses_websocket.connect(
|
||||
url="wss://example.test/responses",
|
||||
headers={},
|
||||
timeout=None,
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -89,8 +91,7 @@ async def test_enabled_bridge_connects_and_adapts_socket(
|
|||
timeout=1.0,
|
||||
)
|
||||
|
||||
assert isinstance(connection, Handled)
|
||||
connection = connection.value
|
||||
assert connection is not None
|
||||
await connection.send("response.create")
|
||||
assert await connection.recv() == "response.completed"
|
||||
await connection.close()
|
||||
|
|
@ -100,72 +101,17 @@ class _FailingNativeBridge:
|
|||
@classmethod
|
||||
async def connect(
|
||||
cls,
|
||||
request: NativeResponsesWebSocketRequest,
|
||||
*,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout_seconds: float | None,
|
||||
context: NativeRequestContext,
|
||||
) -> _FakeNativeConnection:
|
||||
raise RuntimeError("connection failed")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_failure_is_reported_to_orchestration() -> None:
|
||||
configuration.rust(True)
|
||||
responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge)
|
||||
result = await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None)
|
||||
assert isinstance(result, NativeFailed)
|
||||
assert str(result.error) == "connection failed"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_managed_connection_closes_native_socket_on_consumer_failure() -> None:
|
||||
configuration.rust(True)
|
||||
socket = _FakeNativeConnection()
|
||||
|
||||
class Bridge:
|
||||
@classmethod
|
||||
async def connect(
|
||||
cls, *, url: str, headers: dict[str, str], timeout_seconds: float | None
|
||||
) -> _FakeNativeConnection:
|
||||
return socket
|
||||
|
||||
responses_websocket.set_rust_responses_websocket(connection=Bridge)
|
||||
result = await responses_websocket.managed_connect(url="wss://example.test/responses", headers={}, timeout=1.0)
|
||||
assert isinstance(result, Handled)
|
||||
|
||||
async def use_connection() -> None:
|
||||
async with result.value as connection:
|
||||
await connection.send("hello")
|
||||
raise ValueError("consumer failed")
|
||||
|
||||
with pytest.raises(ValueError, match="consumer failed"):
|
||||
await use_connection()
|
||||
assert socket.sent == ["hello"]
|
||||
assert socket.closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_failure_does_not_authorize_python_fallback() -> None:
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
|
||||
from litellm.rust_bridge.dispatch import anative_context, provider_errors
|
||||
|
||||
configuration.rust(True)
|
||||
responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge)
|
||||
|
||||
@anative_context(
|
||||
native=lambda: responses_websocket.managed_connect(
|
||||
url="wss://example.test/responses", headers={}, timeout=None
|
||||
),
|
||||
route="responses_websocket",
|
||||
errors=lambda: provider_errors("openai", "responses websocket"),
|
||||
)
|
||||
def execute() -> AbstractAsyncContextManager[object]:
|
||||
pytest.fail("unknown native failures must not open a Python connection")
|
||||
|
||||
async def run() -> None:
|
||||
async with execute():
|
||||
pytest.fail("connection must fail before entering its body")
|
||||
|
||||
with pytest.raises(RuntimeError, match="connection failed"):
|
||||
await run()
|
||||
await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -12,7 +12,6 @@ import pytest
|
|||
import litellm
|
||||
from litellm.rust_bridge import bindings, configuration
|
||||
from litellm.rust_bridge import chat_completions as bridge
|
||||
from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeFailed
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
RUST_RESPONSE = {
|
||||
|
|
@ -96,16 +95,16 @@ class _RecordingCall:
|
|||
self.error = error
|
||||
self.calls: list[dict] = []
|
||||
|
||||
def __call__(self, **kwargs):
|
||||
self.calls.append(kwargs)
|
||||
def __call__(self, request, *, context):
|
||||
self.calls.append({"request": request, "context": context})
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
return self.result
|
||||
|
||||
|
||||
class _RecordingAsyncCall(_RecordingCall):
|
||||
async def __call__(self, **kwargs):
|
||||
return _RecordingCall.__call__(self, **kwargs)
|
||||
async def __call__(self, request, *, context):
|
||||
return _RecordingCall.__call__(self, request, context=context)
|
||||
|
||||
|
||||
def _accepts(**overrides) -> bool:
|
||||
|
|
@ -251,8 +250,7 @@ class TestSyncCall:
|
|||
|
||||
result = bridge.chat_completions(**_call_kwargs(model_response))
|
||||
|
||||
assert isinstance(result, Handled)
|
||||
result = result.value
|
||||
assert result is not None
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result.choices[0].finish_reason == "stop"
|
||||
assert result.model == "claude-sonnet-4-5-20260101"
|
||||
|
|
@ -266,16 +264,16 @@ class TestSyncCall:
|
|||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert native.calls[0]["timeout_seconds"] == 30.0
|
||||
assert native.calls[0]["request"].options.timeout_seconds == 30.0
|
||||
|
||||
def test_reports_unavailable_bridge(self, monkeypatch):
|
||||
def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeSkipped)
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
def test_reports_native_decline_to_orchestration(self, monkeypatch):
|
||||
def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
|
||||
assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeFailed)
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
|
||||
class TestAsyncCall:
|
||||
|
|
@ -283,18 +281,152 @@ class TestAsyncCall:
|
|||
async def test_builds_a_model_response(self):
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
|
||||
result = await bridge.achat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert isinstance(result, Handled)
|
||||
result = result.value
|
||||
assert result is not None
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reports_unavailable_bridge(self, monkeypatch):
|
||||
async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeSkipped)
|
||||
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reports_native_decline_to_orchestration(self, monkeypatch):
|
||||
async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
|
||||
assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeFailed)
|
||||
assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
|
||||
class TestAsyncFallbackWrapper:
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_the_rust_response_without_running_the_fallback(self):
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall())
|
||||
ran = []
|
||||
|
||||
async def fallback():
|
||||
ran.append(True)
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result.choices[0].message.content == "hello from rust"
|
||||
assert ran == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")))
|
||||
|
||||
async def fallback():
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result == "python"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
|
||||
async def fallback():
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result == "python"
|
||||
|
||||
|
||||
class TestFailureClassification:
|
||||
"""A failure the provider already saw must not be retried on the Python
|
||||
path: it would bill the customer for the same work twice."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _native_exceptions(self, monkeypatch):
|
||||
_fake_native_bridge(monkeypatch)
|
||||
|
||||
def test_a_decline_falls_back_because_nothing_was_sent(self):
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming")))
|
||||
assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None
|
||||
|
||||
def test_an_upstream_failure_is_surfaced_with_its_status(self):
|
||||
from litellm.exceptions import RateLimitError
|
||||
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited")))
|
||||
with pytest.raises(RateLimitError) as raised:
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert raised.value.status_code == 429
|
||||
assert "rate limited" in str(raised.value)
|
||||
|
||||
def test_a_transport_failure_with_no_response_surfaces_as_a_500(self):
|
||||
from litellm.exceptions import APIError
|
||||
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset")))
|
||||
with pytest.raises(APIError) as raised:
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert raised.value.status_code == 500
|
||||
|
||||
def test_an_unrecognized_error_is_not_swallowed(self):
|
||||
bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else")))
|
||||
with pytest.raises(RuntimeError):
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self):
|
||||
from litellm.exceptions import InternalServerError
|
||||
|
||||
bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom")))
|
||||
ran = []
|
||||
|
||||
async def fallback():
|
||||
ran.append(True)
|
||||
return "python"
|
||||
|
||||
with pytest.raises(InternalServerError):
|
||||
await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert ran == [], "a request the provider already served must not be re-issued"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_the_async_wrapper_falls_back_on_a_decline(self):
|
||||
bridge.set_rust_chat_completions(
|
||||
achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text"))
|
||||
)
|
||||
|
||||
async def fallback():
|
||||
return "python"
|
||||
|
||||
result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
assert result == "python"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_missing_native_exception_types_does_not_authorize_python_fallback(monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
bridge.set_rust_chat_completions(
|
||||
chat_completions=_RecordingCall(error=RuntimeError("connection failed")),
|
||||
achat_completions=_RecordingAsyncCall(error=RuntimeError("connection failed")),
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="connection failed"):
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
|
||||
async def fallback():
|
||||
pytest.fail("unknown failure must not retry through Python")
|
||||
|
||||
with pytest.raises(RuntimeError, match="connection failed"):
|
||||
await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
|
||||
|
||||
def test_provider_credentials_are_separate_from_chat_body_params():
|
||||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
configuration.rust(True)
|
||||
kwargs = _call_kwargs(ModelResponse())
|
||||
kwargs["optional_params"] = {
|
||||
"max_tokens": 32,
|
||||
"aws_access_key_id": "test-access-key",
|
||||
"aws_secret_access_key": "test-secret-key",
|
||||
}
|
||||
bridge.chat_completions(**kwargs)
|
||||
request = native.calls[0]["request"]
|
||||
assert request.optional_params == {"max_tokens": 32}
|
||||
assert request.options.provider_connection == {
|
||||
"aws_access_key_id": "test-access-key",
|
||||
"aws_secret_access_key": "test-secret-key",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
|
||||
from litellm.rust_bridge.runtime import Handled
|
||||
from litellm.rust_bridge.request import NativeRequestContext, NativeTranscriptionRequest
|
||||
|
||||
rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
|
||||
|
||||
|
|
@ -15,37 +15,27 @@ class SyncBridge:
|
|||
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
audio: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeTranscriptionRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append({"model": model, "audio": audio, "optional_params": optional_params})
|
||||
self.calls.append({"model": request.model, "audio": request.audio, "optional_params": {**request.optional_params, **(request.options.provider_connection or {})}})
|
||||
return {"text": "hello"}
|
||||
|
||||
|
||||
class AsyncBridge:
|
||||
async def __call__(
|
||||
self,
|
||||
model: str,
|
||||
audio: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeTranscriptionRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
return {"text": "async"}
|
||||
|
||||
|
||||
def test_enabled_sync_bridge_receives_audio() -> None:
|
||||
bridge = SyncBridge()
|
||||
rust_bridge.configure_rust_transcription(transcription=bridge)
|
||||
rust_bridge.configure_rust_transcription(True, transcription=bridge)
|
||||
result = rust_bridge.transcription(
|
||||
model="mistral.voxtral-mini-3b-2507",
|
||||
audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"},
|
||||
|
|
@ -56,14 +46,13 @@ def test_enabled_sync_bridge_receives_audio() -> None:
|
|||
optional_params={"temperature": 0},
|
||||
timeout=5.0,
|
||||
)
|
||||
assert isinstance(result, Handled)
|
||||
assert result.value == {"text": "hello"}
|
||||
assert result == {"text": "hello"}
|
||||
assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enabled_async_bridge() -> None:
|
||||
rust_bridge.configure_rust_transcription(atranscription=AsyncBridge())
|
||||
rust_bridge.configure_rust_transcription(True, atranscription=AsyncBridge())
|
||||
result = await rust_bridge.atranscription(
|
||||
model="mistral.voxtral-mini-3b-2507",
|
||||
audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"},
|
||||
|
|
@ -74,7 +63,7 @@ async def test_enabled_async_bridge() -> None:
|
|||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
assert result == Handled({"text": "async"})
|
||||
assert result == {"text": "async"}
|
||||
|
||||
|
||||
def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
|
|
@ -85,8 +74,7 @@ def test_loader_returns_none_without_native_extension(monkeypatch: pytest.Monkey
|
|||
|
||||
|
||||
def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
rust_bridge.configure_rust_transcription(transcription=None)
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
|
||||
monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None)
|
||||
|
||||
with pytest.raises(RuntimeError, match="bridge is unavailable"):
|
||||
BedrockAudioTranscriptionRustDispatch().audio_transcriptions(
|
||||
|
|
@ -103,8 +91,10 @@ def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) ->
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
rust_bridge.configure_rust_transcription(atranscription=None)
|
||||
monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None)
|
||||
async def unavailable(**_: object) -> None:
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(rust_bridge, "atranscription", unavailable)
|
||||
|
||||
with pytest.raises(RuntimeError, match="bridge is unavailable"):
|
||||
await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions(
|
||||
|
|
@ -121,7 +111,7 @@ async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPat
|
|||
|
||||
def test_bedrock_transcription_uses_rust_only_path() -> None:
|
||||
rust_bridge.configure_rust_transcription(
|
||||
transcription=lambda **_: {"text": "rust"},
|
||||
transcription=lambda request, *, context: {"text": "rust"},
|
||||
atranscription=None,
|
||||
)
|
||||
try:
|
||||
|
|
@ -137,7 +127,7 @@ def test_bedrock_transcription_uses_rust_only_path() -> None:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_atranscription_uses_rust_only_path() -> None:
|
||||
async def rust_response(**_: object) -> dict[str, object]:
|
||||
async def rust_response(request: NativeTranscriptionRequest, *, context: NativeRequestContext) -> dict[str, object]:
|
||||
return {"text": "rust"}
|
||||
|
||||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22165
|
||||
"limit": 22155
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26729
|
||||
|
|
@ -15,7 +15,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT006": {
|
||||
"limit": 1022
|
||||
"limit": 1027
|
||||
},
|
||||
"LIT007": {
|
||||
"limit": 0
|
||||
|
|
@ -27,10 +27,10 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16419
|
||||
"limit": 16426
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5497
|
||||
"limit": 5506
|
||||
},
|
||||
"LIT012": {
|
||||
"limit": 4486
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue