mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
refactor(native): separate request data from execution context
This commit is contained in:
parent
f0d47db104
commit
0f886a6c30
61 changed files with 1841 additions and 1391 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| {
|
||||
|
|
@ -113,11 +113,11 @@ pub(super) async fn signed_headers(
|
|||
let unsigned: BTreeMap<String, String> = request.upstream_headers.iter().cloned().collect();
|
||||
// A host with its own resolution chain hands the result down; only fall
|
||||
// back to deriving credentials here when it supplied none.
|
||||
let credentials = match host_supplied_credentials(&request.optional_params) {
|
||||
let credentials = match host_supplied_credentials(&request.provider_connection) {
|
||||
Some(credentials) => credentials,
|
||||
None => {
|
||||
resolve_credentials(
|
||||
aws_auth_config(&request.optional_params, &env_lookup),
|
||||
aws_auth_config(&request.provider_connection, &env_lookup),
|
||||
&env_lookup,
|
||||
)
|
||||
.await?
|
||||
|
|
|
|||
|
|
@ -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,8 +40,12 @@ pub(super) fn parse_messages(messages: Value) -> Result<Vec<ChatMessage>, Error>
|
|||
|
||||
pub(super) fn resolve_request(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ResolvedChatCompletionsRequest<'_>, Error> {
|
||||
let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?;
|
||||
_context: &LiteLlmRequestContext,
|
||||
) -> Result<ResolvedChatCompletionsRequest, Error> {
|
||||
let (model, config) = resolve_provider_config(
|
||||
request.model,
|
||||
request.options.custom_llm_provider.as_deref(),
|
||||
)?;
|
||||
let messages = parse_messages(request.messages)?;
|
||||
if messages.is_empty() {
|
||||
return Err(Error::InvalidRequest(
|
||||
|
|
@ -55,25 +60,22 @@ pub(super) fn resolve_request(
|
|||
config,
|
||||
messages,
|
||||
optional_params: request.optional_params,
|
||||
api_key: request.api_key,
|
||||
api_base: request.api_base,
|
||||
extra_headers: request.extra_headers,
|
||||
timeout: request.timeout,
|
||||
options: request.options,
|
||||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn validate_environment(
|
||||
request: &ResolvedChatCompletionsRequest<'_>,
|
||||
request: &ResolvedChatCompletionsRequest,
|
||||
model: &str,
|
||||
config: &dyn ChatCompletionsProviderConfig,
|
||||
) -> Result<(Vec<(String, String)>, ChatCompletionsAuth), Error> {
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let mut headers = string_headers(request.extra_headers.clone())?;
|
||||
let mut headers = string_headers(request.options.extra_headers.clone())?;
|
||||
let auth = config.auth(
|
||||
request.api_key,
|
||||
request.options.api_key.as_deref(),
|
||||
model,
|
||||
&request.optional_params,
|
||||
&request.options.provider_connection,
|
||||
&env_lookup,
|
||||
)?;
|
||||
match &auth {
|
||||
|
|
@ -117,16 +119,16 @@ fn validate_environment(
|
|||
}
|
||||
|
||||
pub(super) fn prepare_provider_request(
|
||||
request: ResolvedChatCompletionsRequest<'_>,
|
||||
request: ResolvedChatCompletionsRequest,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
let (headers, auth) = validate_environment(&request, &request.model, request.config)?;
|
||||
let model = request.model;
|
||||
let config = request.config;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
let url = config.complete_url(
|
||||
request.api_base,
|
||||
request.options.api_base.as_deref(),
|
||||
&model,
|
||||
&request.optional_params,
|
||||
&request.options.provider_connection,
|
||||
&env_lookup,
|
||||
)?;
|
||||
let transformed =
|
||||
|
|
@ -139,7 +141,7 @@ pub(super) fn prepare_provider_request(
|
|||
body: transformed.body,
|
||||
upstream_headers: headers,
|
||||
auth,
|
||||
optional_params: request.optional_params,
|
||||
timeout: request.timeout,
|
||||
provider_connection: request.options.provider_connection,
|
||||
timeout: request.options.timeout,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
use crate::request_context::LiteLlmRequestContext;
|
||||
use crate::request_options::RequestOptions;
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::error::Error;
|
||||
|
|
@ -9,7 +11,12 @@ use super::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest};
|
|||
fn prepare_chat_completions_call(
|
||||
request: ChatCompletionsRequest<'_>,
|
||||
) -> Result<ProviderChatCompletionsRequest, Error> {
|
||||
prepare_provider_request(resolve_request(request)?)
|
||||
prepare_provider_request(resolve_request(
|
||||
request,
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)?)
|
||||
}
|
||||
|
||||
fn request<'a>(
|
||||
|
|
@ -25,11 +32,15 @@ fn request<'a>(
|
|||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
},
|
||||
api_key: Some("sk-test"),
|
||||
api_base: None,
|
||||
custom_llm_provider: provider,
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
|
||||
options: RequestOptions {
|
||||
api_key: (Some("sk-test")).map(|value| value.to_string()),
|
||||
api_base: None,
|
||||
custom_llm_provider: (provider).map(|value| value.to_string()),
|
||||
extra_headers: None,
|
||||
timeout: None,
|
||||
..Default::default()
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -107,7 +118,7 @@ fn the_deployment_credential_replaces_a_caller_supplied_auth_header() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([(
|
||||
call.options.extra_headers = Some(Map::from_iter([(
|
||||
"X-Api-Key".to_string(),
|
||||
json!("sk-caller"),
|
||||
)]));
|
||||
|
|
@ -132,7 +143,7 @@ fn a_forwarded_authorization_header_suppresses_the_resolved_api_key_header() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([
|
||||
call.options.extra_headers = Some(Map::from_iter([
|
||||
(
|
||||
"Authorization".to_string(),
|
||||
json!("Bearer sk-ant-oat01-token"),
|
||||
|
|
@ -168,7 +179,7 @@ fn an_unrelated_forwarded_authorization_does_not_defer_the_resolved_key() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([
|
||||
call.options.extra_headers = Some(Map::from_iter([
|
||||
("Authorization".to_string(), json!("Bearer unrelated")),
|
||||
("X-Api-Key".to_string(), json!("sk-caller")),
|
||||
]));
|
||||
|
|
@ -199,7 +210,7 @@ fn declines_an_unsupported_request_before_resolving_credentials() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true}),
|
||||
);
|
||||
call.api_key = None;
|
||||
call.options.api_key = None;
|
||||
// No api_key is set and no env is consulted: the gate must run first, so the
|
||||
// error is the decline rather than a missing-credential error.
|
||||
assert_eq!(decline(call), Error::Unsupported("streaming"));
|
||||
|
|
@ -261,7 +272,7 @@ fn rejects_non_string_extra_headers() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
|
||||
call.options.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))]));
|
||||
assert_eq!(
|
||||
decline(call),
|
||||
Error::InvalidRequest(
|
||||
|
|
@ -279,7 +290,7 @@ fn prepares_a_bedrock_call_without_resolving_credentials() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"maxTokens": 16}),
|
||||
);
|
||||
call.api_key = None;
|
||||
call.options.api_key = None;
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
assert_eq!(
|
||||
prepared.url,
|
||||
|
|
@ -312,15 +323,18 @@ async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() {
|
|||
"bedrock/us-east-1/anthropic.claude-v2",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({
|
||||
"maxTokens": 16,
|
||||
"aws_access_key_id": "AKIDEXAMPLE",
|
||||
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
|
||||
}),
|
||||
json!({"maxTokens": 16}),
|
||||
);
|
||||
call.options.provider_connection = Map::from_iter([
|
||||
("aws_access_key_id".to_string(), json!("AKIDEXAMPLE")),
|
||||
(
|
||||
"aws_secret_access_key".to_string(),
|
||||
json!("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"),
|
||||
),
|
||||
]);
|
||||
// A key would resolve to a bearer token and never reach the signer.
|
||||
call.api_key = None;
|
||||
call.extra_headers = Some(Map::from_iter([(
|
||||
call.options.api_key = None;
|
||||
call.options.extra_headers = Some(Map::from_iter([(
|
||||
"x-request-id".to_string(),
|
||||
json!("abc-123"),
|
||||
)]));
|
||||
|
|
@ -367,14 +381,18 @@ async fn a_forwarded_header_the_signer_computes_declines_to_python() {
|
|||
"bedrock/us-east-1/anthropic.claude-v2",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({
|
||||
"maxTokens": 16,
|
||||
"aws_access_key_id": "AKIDEXAMPLE",
|
||||
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"
|
||||
}),
|
||||
json!({"maxTokens": 16}),
|
||||
);
|
||||
call.api_key = None;
|
||||
call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
|
||||
call.options.provider_connection = Map::from_iter([
|
||||
("aws_access_key_id".to_string(), json!("AKIDEXAMPLE")),
|
||||
(
|
||||
"aws_secret_access_key".to_string(),
|
||||
json!("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"),
|
||||
),
|
||||
]);
|
||||
call.options.api_key = None;
|
||||
call.options.extra_headers =
|
||||
Some(Map::from_iter([(forwarded.to_string(), json!("forged"))]));
|
||||
let prepared = prepare_chat_completions_call(call).expect("prepares");
|
||||
let error = super::handler::signed_headers(&prepared, br#"{"a":1}"#)
|
||||
.await
|
||||
|
|
@ -399,7 +417,7 @@ fn a_bedrock_deployment_bearer_outranks_a_forwarded_authorization() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"maxTokens": 16}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([(
|
||||
call.options.extra_headers = Some(Map::from_iter([(
|
||||
"Authorization".to_string(),
|
||||
json!("Bearer caller-supplied"),
|
||||
)]));
|
||||
|
|
@ -432,7 +450,7 @@ fn an_anthropic_forwarded_oauth_bearer_still_outranks_the_resolved_key() {
|
|||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
);
|
||||
call.extra_headers = Some(Map::from_iter([(
|
||||
call.options.extra_headers = Some(Map::from_iter([(
|
||||
"authorization".to_string(),
|
||||
json!("Bearer sk-ant-oat01-forwarded"),
|
||||
)]));
|
||||
|
|
@ -666,11 +684,15 @@ mod round_trip {
|
|||
Value::Object(map) => map,
|
||||
other => panic!("params must be an object, got {other}"),
|
||||
},
|
||||
api_key: Some("sk-test"),
|
||||
api_base: Some(api_base),
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: Some(std::time::Duration::from_secs(10)),
|
||||
|
||||
options: RequestOptions {
|
||||
api_key: (Some("sk-test")).map(|value| value.to_string()),
|
||||
api_base: (Some(api_base)).map(|value| value.to_string()),
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
timeout: Some(std::time::Duration::from_secs(10)),
|
||||
..Default::default()
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -679,14 +701,19 @@ mod round_trip {
|
|||
#[tokio::test]
|
||||
async fn round_trip_sends_the_translated_body_and_normalizes_the_response() {
|
||||
let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await;
|
||||
let response = chat_completions(call(
|
||||
&api_base,
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let response = chat_completions(
|
||||
call(
|
||||
&api_base,
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("call succeeds");
|
||||
|
||||
|
|
@ -724,11 +751,16 @@ mod round_trip {
|
|||
const NO_USAGE: &str =
|
||||
r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#;
|
||||
let (api_base, handle) = serve_once("200 OK", NO_USAGE).await;
|
||||
let err = chat_completions(call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let err = chat_completions(
|
||||
call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect_err("response cannot be normalized");
|
||||
handle.await.expect("server task");
|
||||
|
|
@ -742,11 +774,16 @@ mod round_trip {
|
|||
async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() {
|
||||
const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#;
|
||||
let (api_base, handle) = serve_once("200 OK", TOOL_USE).await;
|
||||
let err = chat_completions(call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let err = chat_completions(
|
||||
call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect_err("response cannot be normalized");
|
||||
handle.await.expect("server task");
|
||||
|
|
@ -760,11 +797,16 @@ mod round_trip {
|
|||
async fn an_upstream_error_status_keeps_its_code() {
|
||||
let (api_base, handle) =
|
||||
serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await;
|
||||
let err = chat_completions(call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let err = chat_completions(
|
||||
call(
|
||||
&api_base,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect_err("upstream rejects");
|
||||
handle.await.expect("server task");
|
||||
|
|
@ -785,11 +827,16 @@ mod round_trip {
|
|||
listener.local_addr().expect("has an address").port()
|
||||
// Dropped here, so the port is closed and the connect is refused.
|
||||
};
|
||||
let err = chat_completions(call(
|
||||
&format!("http://127.0.0.1:{port}/v1/messages"),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
))
|
||||
let err = chat_completions(
|
||||
call(
|
||||
&format!("http://127.0.0.1:{port}/v1/messages"),
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect_err("nothing is listening");
|
||||
assert!(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -31,6 +31,15 @@ from litellm.rust_bridge.protocols import (
|
|||
RustChatCompletions,
|
||||
RustChatCompletionsDecline,
|
||||
)
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeChatCompletionsRequest,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
provider_connection_params,
|
||||
provider_request_params,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointBinding,
|
||||
|
|
@ -235,17 +244,23 @@ def chat_completions(
|
|||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
return _CHAT.invoke(
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_chat_completions, timeout_seconds: rust_chat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeChatCompletionsRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=lambda: None,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
|
|
@ -270,17 +285,23 @@ async def achat_completions(
|
|||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
return await _CHAT.ainvoke(
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_achat_completions, timeout_seconds: rust_achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeChatCompletionsRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=async_none,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
|
|
@ -315,17 +336,23 @@ async def achat_completions_or_fallback(
|
|||
return _build_model_response(rust_response, model_response)
|
||||
|
||||
return await _CHAT.ainvoke(
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_achat_completions, timeout_seconds: rust_achat_completions(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeChatCompletionsRequest(
|
||||
model=model,
|
||||
messages=messages,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=python_fallback,
|
||||
adapt=adapt,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
|
|
|
|||
|
|
@ -8,6 +8,13 @@ import httpx
|
|||
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.protocols import RustAmessages, RustMessages
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeMessagesRequest,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointDispatch,
|
||||
|
|
@ -61,16 +68,21 @@ def messages(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
return _MESSAGES.invoke(
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_messages, timeout_seconds: rust_messages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeMessagesRequest(
|
||||
model=model,
|
||||
body=body,
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
|
|
@ -88,16 +100,21 @@ async def amessages(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
return await _MESSAGES.ainvoke(
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_amessages, timeout_seconds: rust_amessages(
|
||||
model=model,
|
||||
body=body,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeMessagesRequest(
|
||||
model=model,
|
||||
body=body,
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
|
|
|
|||
|
|
@ -9,6 +9,15 @@ import httpx
|
|||
from litellm.rust_bridge import configuration as _configuration
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.protocols import RustAocr, RustOcr
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeOCRRequest,
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
provider_connection_params,
|
||||
provider_request_params,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointDispatch,
|
||||
|
|
@ -67,17 +76,23 @@ def ocr(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
return _OCR.invoke(
|
||||
prepare=lambda: _timeout_to_seconds(timeout),
|
||||
call=lambda rust_ocr, timeout_seconds: rust_ocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeOCRRequest(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
|
|
@ -96,17 +111,23 @@ async def aocr(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
return await _OCR.ainvoke(
|
||||
prepare=lambda: _timeout_to_seconds(timeout),
|
||||
call=lambda rust_aocr, timeout_seconds: rust_aocr(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeOCRRequest(
|
||||
model=model,
|
||||
document=document,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=_timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -13,6 +13,13 @@ from litellm.rust_bridge.protocols import (
|
|||
RustResponsesWebSocket,
|
||||
RustResponsesWebSocketConnection,
|
||||
)
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
NativeResponsesWebSocketRequest,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
AsyncEndpointDispatch,
|
||||
BridgeErrorContext,
|
||||
|
|
@ -34,9 +41,9 @@ def set_rust_responses_websocket(
|
|||
) -> None:
|
||||
if not isinstance(connection, Unchanged):
|
||||
if connection is None:
|
||||
_RESPONSES_WEBSOCKET.reset()
|
||||
_RESPONSES_WEBSOCKET.asynchronous.reset()
|
||||
else:
|
||||
_RESPONSES_WEBSOCKET.override(connection)
|
||||
_RESPONSES_WEBSOCKET.asynchronous.override(connection)
|
||||
|
||||
|
||||
class _ConnectionAdapter:
|
||||
|
|
@ -63,12 +70,14 @@ async def connect(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> _ConnectionAdapter | None:
|
||||
connection: Final = await _RESPONSES_WEBSOCKET.ainvoke(
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda connection_type, timeout_seconds: connection_type.connect(
|
||||
url=url,
|
||||
headers=headers,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeResponsesWebSocketRequest(
|
||||
url=url,
|
||||
options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=lambda connection_type, request: call_native(connection_type.connect, request),
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider="openai", model="responses websocket"),
|
||||
|
|
|
|||
|
|
@ -6,6 +6,15 @@ import httpx
|
|||
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestContext,
|
||||
NativeRequestOptions,
|
||||
NativeTranscriptionRequest,
|
||||
PreparedNativeCall,
|
||||
call_native,
|
||||
provider_connection_params,
|
||||
provider_request_params,
|
||||
)
|
||||
from litellm.rust_bridge.runtime import (
|
||||
BridgeErrorContext,
|
||||
EndpointDispatch,
|
||||
|
|
@ -61,17 +70,23 @@ def transcription(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
return _TRANSCRIPTION.invoke(
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_transcription, timeout_seconds: rust_transcription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeTranscriptionRequest(
|
||||
model=model,
|
||||
audio=audio,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=lambda: None,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
|
|
@ -90,17 +105,23 @@ async def atranscription(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
) -> dict[str, object] | None:
|
||||
return await _TRANSCRIPTION.ainvoke(
|
||||
prepare=lambda: timeout_to_seconds(timeout),
|
||||
call=lambda rust_atranscription, timeout_seconds: rust_atranscription(
|
||||
model=model,
|
||||
audio=audio,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout_seconds=timeout_seconds,
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
NativeTranscriptionRequest(
|
||||
model=model,
|
||||
audio=audio,
|
||||
optional_params=provider_request_params(optional_params),
|
||||
options=NativeRequestOptions(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
timeout_seconds=timeout_to_seconds(timeout),
|
||||
provider_connection=provider_connection_params(optional_params),
|
||||
),
|
||||
),
|
||||
context=NativeRequestContext(),
|
||||
),
|
||||
call=call_native,
|
||||
fallback=async_none,
|
||||
adapt=identity,
|
||||
error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model),
|
||||
|
|
|
|||
|
|
@ -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,6 +9,7 @@ import pytest
|
|||
import litellm
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.request import NativeMessagesRequest, NativeRequestContext
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
)
|
||||
|
|
@ -40,23 +41,19 @@ class RecordingMessages:
|
|||
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
body: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
request: NativeMessagesRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"body": body,
|
||||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"model": request.model,
|
||||
"body": request.body,
|
||||
"api_key": request.options.api_key,
|
||||
"api_base": request.options.api_base,
|
||||
"custom_llm_provider": request.options.custom_llm_provider,
|
||||
"extra_headers": request.options.extra_headers,
|
||||
"timeout_seconds": request.options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
|
@ -68,23 +65,19 @@ class RecordingAsyncMessages:
|
|||
|
||||
async def __call__(
|
||||
self,
|
||||
model: str,
|
||||
body: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
timeout_seconds: float | None,
|
||||
request: NativeMessagesRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"body": body,
|
||||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"model": request.model,
|
||||
"body": request.body,
|
||||
"api_key": request.options.api_key,
|
||||
"api_base": request.options.api_base,
|
||||
"custom_llm_provider": request.options.custom_llm_provider,
|
||||
"extra_headers": request.options.extra_headers,
|
||||
"timeout_seconds": request.options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_MESSAGES_RESPONSE)
|
||||
|
|
@ -94,7 +87,7 @@ class ExplodingAsyncMessages:
|
|||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, **kwargs: object) -> dict[str, object]:
|
||||
async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise AssertionError("bridge must not be called")
|
||||
|
||||
|
|
@ -103,7 +96,7 @@ class RaisingAsyncMessages:
|
|||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
async def __call__(self, **kwargs: object) -> dict[str, object]:
|
||||
async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise RuntimeError("upstream request failed with status 400: bad request")
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -11,6 +11,7 @@ import pytest
|
|||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.request import NativeOCRRequest, NativeRequestContext
|
||||
|
||||
# `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr`
|
||||
# function onto `litellm.ocr` and shadows the submodule, so import the modules
|
||||
|
|
@ -46,25 +47,20 @@ class RecordingBridge:
|
|||
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeOCRRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"document": document,
|
||||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"optional_params": optional_params,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"model": request.model,
|
||||
"document": request.document,
|
||||
"api_key": request.options.api_key,
|
||||
"api_base": request.options.api_base,
|
||||
"custom_llm_provider": request.options.custom_llm_provider,
|
||||
"extra_headers": request.options.extra_headers,
|
||||
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
|
||||
"timeout_seconds": request.options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_OCR_RESPONSE)
|
||||
|
|
@ -78,25 +74,20 @@ class RecordingAsyncBridge:
|
|||
|
||||
async def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeOCRRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"document": document,
|
||||
"api_key": api_key,
|
||||
"api_base": api_base,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"extra_headers": extra_headers,
|
||||
"optional_params": optional_params,
|
||||
"timeout_seconds": timeout_seconds,
|
||||
"model": request.model,
|
||||
"document": request.document,
|
||||
"api_key": request.options.api_key,
|
||||
"api_base": request.options.api_base,
|
||||
"custom_llm_provider": request.options.custom_llm_provider,
|
||||
"extra_headers": request.options.extra_headers,
|
||||
"optional_params": {**request.optional_params, **(request.options.provider_connection or {})},
|
||||
"timeout_seconds": request.options.timeout_seconds,
|
||||
}
|
||||
)
|
||||
return dict(FAKE_OCR_RESPONSE)
|
||||
|
|
@ -105,14 +96,9 @@ class RecordingAsyncBridge:
|
|||
class RaisingBridge:
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeOCRRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
raise RuntimeError("bridge failed")
|
||||
|
||||
|
|
@ -120,14 +106,9 @@ class RaisingBridge:
|
|||
class RaisingAsyncBridge:
|
||||
async def __call__(
|
||||
self,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeOCRRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
raise RuntimeError("bridge failed")
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import pytest
|
|||
|
||||
from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled
|
||||
from litellm.rust_bridge import configuration, responses_websocket
|
||||
from litellm.rust_bridge.request import NativeRequestContext, NativeResponsesWebSocketRequest
|
||||
|
||||
|
||||
class _FakeNativeConnection:
|
||||
|
|
@ -30,10 +31,9 @@ class _FakeNativeBridge:
|
|||
@classmethod
|
||||
async def connect(
|
||||
cls,
|
||||
request: NativeResponsesWebSocketRequest,
|
||||
*,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout_seconds: float | None,
|
||||
context: NativeRequestContext,
|
||||
) -> _FakeNativeConnection:
|
||||
return _FakeNativeConnection()
|
||||
|
||||
|
|
@ -101,10 +101,9 @@ class _FailingNativeBridge:
|
|||
@classmethod
|
||||
async def connect(
|
||||
cls,
|
||||
request: NativeResponsesWebSocketRequest,
|
||||
*,
|
||||
url: str,
|
||||
headers: dict[str, str],
|
||||
timeout_seconds: float | None,
|
||||
context: NativeRequestContext,
|
||||
) -> _FakeNativeConnection:
|
||||
raise RuntimeError("connection failed")
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -95,16 +95,16 @@ class _RecordingCall:
|
|||
self.error = error
|
||||
self.calls: list[dict] = []
|
||||
|
||||
def __call__(self, **kwargs):
|
||||
self.calls.append(kwargs)
|
||||
def __call__(self, request, *, context):
|
||||
self.calls.append({"request": request, "context": context})
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
return self.result
|
||||
|
||||
|
||||
class _RecordingAsyncCall(_RecordingCall):
|
||||
async def __call__(self, **kwargs):
|
||||
return _RecordingCall.__call__(self, **kwargs)
|
||||
async def __call__(self, request, *, context):
|
||||
return _RecordingCall.__call__(self, request, context=context)
|
||||
|
||||
|
||||
def _accepts(**overrides) -> bool:
|
||||
|
|
@ -271,7 +271,7 @@ class TestSyncCall:
|
|||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
bridge.chat_completions(**_call_kwargs(ModelResponse()))
|
||||
assert native.calls[0]["timeout_seconds"] == 30.0
|
||||
assert native.calls[0]["request"].options.timeout_seconds == 30.0
|
||||
|
||||
def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch):
|
||||
_hide_native_bridge(monkeypatch)
|
||||
|
|
@ -418,3 +418,22 @@ async def test_missing_native_exception_types_does_not_authorize_python_fallback
|
|||
|
||||
with pytest.raises(RuntimeError, match="connection failed"):
|
||||
await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback)
|
||||
|
||||
|
||||
def test_provider_credentials_are_separate_from_chat_body_params():
|
||||
native = _RecordingCall()
|
||||
bridge.set_rust_chat_completions(chat_completions=native)
|
||||
configuration.rust(True)
|
||||
kwargs = _call_kwargs(ModelResponse())
|
||||
kwargs["optional_params"] = {
|
||||
"max_tokens": 32,
|
||||
"aws_access_key_id": "test-access-key",
|
||||
"aws_secret_access_key": "test-secret-key",
|
||||
}
|
||||
bridge.chat_completions(**kwargs)
|
||||
request = native.calls[0]["request"]
|
||||
assert request.optional_params == {"max_tokens": 32}
|
||||
assert request.options.provider_connection == {
|
||||
"aws_access_key_id": "test-access-key",
|
||||
"aws_secret_access_key": "test-secret-key",
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
|
||||
from litellm.rust_bridge.request import NativeRequestContext, NativeTranscriptionRequest
|
||||
|
||||
rust_bridge = importlib.import_module("litellm.rust_bridge.transcription")
|
||||
|
||||
|
|
@ -14,30 +15,20 @@ class SyncBridge:
|
|||
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
audio: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeTranscriptionRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append({"model": model, "audio": audio, "optional_params": optional_params})
|
||||
self.calls.append({"model": request.model, "audio": request.audio, "optional_params": {**request.optional_params, **(request.options.provider_connection or {})}})
|
||||
return {"text": "hello"}
|
||||
|
||||
|
||||
class AsyncBridge:
|
||||
async def __call__(
|
||||
self,
|
||||
model: str,
|
||||
audio: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout_seconds: float | None,
|
||||
request: NativeTranscriptionRequest,
|
||||
*,
|
||||
context: NativeRequestContext,
|
||||
) -> dict[str, object]:
|
||||
return {"text": "async"}
|
||||
|
||||
|
|
@ -120,7 +111,7 @@ async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPat
|
|||
|
||||
def test_bedrock_transcription_uses_rust_only_path() -> None:
|
||||
rust_bridge.configure_rust_transcription(
|
||||
transcription=lambda **_: {"text": "rust"},
|
||||
transcription=lambda request, *, context: {"text": "rust"},
|
||||
atranscription=None,
|
||||
)
|
||||
try:
|
||||
|
|
@ -136,7 +127,7 @@ def test_bedrock_transcription_uses_rust_only_path() -> None:
|
|||
|
||||
@pytest.mark.asyncio
|
||||
async def test_bedrock_atranscription_uses_rust_only_path() -> None:
|
||||
async def rust_response(**_: object) -> dict[str, object]:
|
||||
async def rust_response(request: NativeTranscriptionRequest, *, context: NativeRequestContext) -> dict[str, object]:
|
||||
return {"text": "rust"}
|
||||
|
||||
rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22181
|
||||
"limit": 22165
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26745
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue