From 0f886a6c30643d97a642454ecd05c69ba203c858 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 12:59:50 -0700 Subject: [PATCH] refactor(native): separate request data from execution context --- litellm-rust/README.md | 32 +- .../src/audio_transcription/hooks.rs | 28 +- .../ai-gateway/src/audio_transcription/mod.rs | 22 +- .../src/audio_transcription/prepare.rs | 45 ++- .../src/audio_transcription/tests.rs | 45 ++- .../src/audio_transcription/types.rs | 30 +- .../ai-gateway/src/integrations/types.rs | 10 +- .../crates/ai-gateway/src/io/responses_ws.rs | 21 +- .../crates/ai-gateway/src/ocr/hooks.rs | 19 +- litellm-rust/crates/ai-gateway/src/ocr/mod.rs | 56 +++- .../crates/ai-gateway/src/ocr/prepare.rs | 106 +++--- .../crates/ai-gateway/src/ocr/types.rs | 30 +- .../ai-gateway/src/realtime/streaming.rs | 17 +- .../ai-gateway/src/routes/messages/service.rs | 33 +- .../ai-gateway/src/routes/realtime/mod.rs | 7 +- .../ai-gateway/src/routes/responses/mod.rs | 7 +- .../src/routes/responses/service.rs | 4 +- .../crates/ai-gateway/tests/ocr_lifecycle.rs | 280 ++++++++++------ .../core/src/audio_transcription/handler.rs | 2 +- .../core/src/audio_transcription/mod.rs | 6 +- .../core/src/audio_transcription/prepare.rs | 47 +-- .../core/src/audio_transcription/tests.rs | 33 +- .../core/src/audio_transcription/types.rs | 9 +- .../core/src/chat_completions/handler.rs | 6 +- .../crates/core/src/chat_completions/mod.rs | 4 +- .../core/src/chat_completions/prepare.rs | 32 +- .../crates/core/src/chat_completions/tests.rs | 169 ++++++---- .../crates/core/src/chat_completions/types.rs | 16 +- litellm-rust/crates/core/src/lib.rs | 3 + litellm-rust/crates/core/src/messages/mod.rs | 11 +- .../crates/core/src/messages/prepare.rs | 45 ++- .../crates/core/src/messages/tests.rs | 238 ++++++++----- .../crates/core/src/messages/types.rs | 7 +- .../crates/core/src/request_context.rs | 18 + .../crates/core/src/request_options.rs | 14 + .../crates/core/src/responses/types.rs | 6 + litellm-rust/crates/python-bridge/src/lib.rs | 49 ++- .../crates/python-bridge/src/marshal.rs | 174 ++++++---- .../src/routes/audio_transcription.rs | 84 ++--- .../src/routes/chat_completions.rs | 96 +++--- .../python-bridge/src/routes/definition.rs | 314 ++++++------------ .../python-bridge/src/routes/messages.rs | 69 ++-- .../crates/python-bridge/src/routes/ocr.rs | 91 ++--- litellm/rust_bridge/chat_completions.py | 87 +++-- litellm/rust_bridge/messages.py | 53 ++- litellm/rust_bridge/ocr.py | 61 ++-- litellm/rust_bridge/protocols.py | 131 ++------ litellm/rust_bridge/request.py | 122 +++++++ litellm/rust_bridge/responses_websocket.py | 23 +- litellm/rust_bridge/transcription.py | 61 ++-- .../strategies/trace_parity/sdk/execution.py | 56 +++- .../trace_parity/sdk/transcription/case.py | 7 +- .../test_rust_bridge_messages.py | 53 ++- .../chat/test_anthropic_chat_handler.py | 23 +- .../chat/test_bedrock_converse_handler.py | 27 +- tests/test_litellm/ocr/test_rust_bridge.py | 77 ++--- .../responses/test_rust_bridge_websocket.py | 11 +- .../rust_bridge/native_route_wheel_test.py | 45 ++- .../rust_bridge/test_chat_completions.py | 29 +- .../test_audio_transcription_rust_bridge.py | 29 +- type-discipline-budget.json | 2 +- 61 files changed, 1841 insertions(+), 1391 deletions(-) create mode 100644 litellm-rust/crates/core/src/request_context.rs create mode 100644 litellm-rust/crates/core/src/request_options.rs create mode 100644 litellm/rust_bridge/request.py diff --git a/litellm-rust/README.md b/litellm-rust/README.md index 650d38753e7..5720a463455 100644 --- a/litellm-rust/README.md +++ b/litellm-rust/README.md @@ -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 | diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs index dbe2d3a325b..a8bb1e53348 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs @@ -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> + 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 GuardrailContext { +fn guardrail_context(metadata: &RequestAttribution) -> GuardrailContext { GuardrailContext { call_type: CallType::Other("audio_transcription".to_string()), selected_guardrails: Vec::new(), diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs index 03d621b8414..de385cb667d 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs @@ -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 { - let PreparedAudioTranscriptionCall { request, hooks } = - prepare_audio_transcription_call(request); +pub async fn audio_transcription( + request: AudioTranscriptionRequest<'_>, + context: &LiteLlmRequestContext, + hooks: RequestHooks, +) -> Result { + 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 } diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs index a475d58635f..cb4884cceb6 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs @@ -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(), ), } } diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/tests.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/tests.rs index 5df04708b7d..d54399b9f96 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/tests.rs @@ -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"})); diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs index b470638264e..8656b839c09 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs @@ -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>, pub optional_params: Map, - pub timeout: Option, - pub callbacks: Vec>, - pub guardrails: Vec>, - 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, pub(crate) api_key: Option, pub(crate) api_base: Option, pub(crate) extra_headers: Option>, pub(crate) optional_params: Map, pub(crate) timeout: Option, } - -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(), - ) - } -} diff --git a/litellm-rust/crates/ai-gateway/src/integrations/types.rs b/litellm-rust/crates/ai-gateway/src/integrations/types.rs index 34dce93d8e0..51163467adc 100644 --- a/litellm-rust/crates/ai-gateway/src/integrations/types.rs +++ b/litellm-rust/crates/ai-gateway/src/integrations/types.rs @@ -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, - pub user_api_key_user_id: Option, - pub user_api_key_team_id: Option, +#[derive(Default)] +pub struct RequestHooks { + pub callbacks: Vec>, + pub guardrails: Vec>, } /// The self-describing payload. Field names are the EXACT JSON keys the Python diff --git a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs index 0b01747b1a5..9d22ec76e5e 100644 --- a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs +++ b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs @@ -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, - timeout: Option, + pub async fn connect( + input: ResponsesWebSocketRequest, + _context: &LiteLlmRequestContext, ) -> Result { - 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::() .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()) })?, diff --git a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs index 446b323db3a..1588aee12bf 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs @@ -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> + 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 for OcrLi } } -fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext { +fn guardrail_context(metadata: &RequestAttribution) -> GuardrailContext { GuardrailContext { call_type: CallType::Ocr, selected_guardrails: Vec::new(), diff --git a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs index 2acdd232c80..520fd88c8a5 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs @@ -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 { - let PreparedOcrCall { request, hooks } = prepare_ocr_call(request); +pub async fn ocr( + request: OcrRequest<'_>, + context: &LiteLlmRequestContext, + hooks: RequestHooks, +) -> Result { + 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 { #[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!( diff --git a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs b/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs index fa9ca1a193e..cd0d2899bc7 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs @@ -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`")) ); diff --git a/litellm-rust/crates/ai-gateway/src/ocr/types.rs b/litellm-rust/crates/ai-gateway/src/ocr/types.rs index 75a8e61ddbf..d6be6ba0d78 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/types.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/types.rs @@ -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>, pub optional_params: Map, - pub timeout: Option, - pub callbacks: Vec>, - pub guardrails: Vec>, - 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, pub(crate) api_key: Option, pub(crate) api_base: Option, pub(crate) extra_headers: Option>, @@ -37,17 +24,6 @@ pub(crate) struct PreparedOcrRequest { pub(crate) timeout: Option, } -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, diff --git a/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs b/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs index c0d72e90b77..588de11a4cb 100644 --- a/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs +++ b/litellm-rust/crates/ai-gateway/src/realtime/streaming.rs @@ -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>, 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); diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs index 5434719987b..b9dfe8dc4f8 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -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| { diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs index f9144ad1fdb..17aa1ba6137 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs @@ -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` diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs index a94853e106d..08c4ed68c43 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs @@ -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 { diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs index e8f840c0c8e..8643f861294 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/service.rs @@ -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( @@ -23,7 +23,7 @@ pub async fn run( idle_timeout: Option, loggers: Arc>>, call_id: String, - metadata: RequestMetadata, + metadata: RequestAttribution, client_in: In, client_out: Out, ) -> Result<(), Error> diff --git a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs b/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs index c3a89f4394d..abbd1424525 100644 --- a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs +++ b/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs @@ -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"); diff --git a/litellm-rust/crates/core/src/audio_transcription/handler.rs b/litellm-rust/crates/core/src/audio_transcription/handler.rs index 9a96b9d1140..e9d91d7cd3b 100644 --- a/litellm-rust/crates/core/src/audio_transcription/handler.rs +++ b/litellm-rust/crates/core/src/audio_transcription/handler.rs @@ -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?; diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs index 31b6de4b3e4..c7331a6ada0 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -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 { +pub async fn audio_transcription( + request: AudioTranscriptionRequest<'_>, + _context: &LiteLlmRequestContext, +) -> Result { execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?) .await } diff --git a/litellm-rust/crates/core/src/audio_transcription/prepare.rs b/litellm-rust/crates/core/src/audio_transcription/prepare.rs index bbef97341a9..28d975bd178 100644 --- a/litellm-rust/crates/core/src/audio_transcription/prepare.rs +++ b/litellm-rust/crates/core/src/audio_transcription/prepare.rs @@ -21,29 +21,34 @@ fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProv pub fn prepare_audio_transcription_provider_call( request: AudioTranscriptionRequest<'_>, ) -> Result { - 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, }) } diff --git a/litellm-rust/crates/core/src/audio_transcription/tests.rs b/litellm-rust/crates/core/src/audio_transcription/tests.rs index 263d63337b0..224141776c2 100644 --- a/litellm-rust/crates/core/src/audio_transcription/tests.rs +++ b/litellm-rust/crates/core/src/audio_transcription/tests.rs @@ -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"})); diff --git a/litellm-rust/crates/core/src/audio_transcription/types.rs b/litellm-rust/crates/core/src/audio_transcription/types.rs index 559d7837027..f988202cd14 100644 --- a/litellm-rust/crates/core/src/audio_transcription/types.rs +++ b/litellm-rust/crates/core/src/audio_transcription/types.rs @@ -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>, pub optional_params: Map, - pub timeout: Option, + 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, + pub(super) provider_connection: Map, pub(super) timeout: Option, } diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index 96d001e2892..bf68650258e 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -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 { 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 = 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? diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 32dea17d202..8cb9ffe3246 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -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 { - 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 diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index 3be2ba21de4..7dc68e88cd8 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -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, Error> pub(super) fn resolve_request( request: ChatCompletionsRequest<'_>, -) -> Result, Error> { - let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?; + _context: &LiteLlmRequestContext, +) -> Result { + 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 { 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, }) } diff --git a/litellm-rust/crates/core/src/chat_completions/tests.rs b/litellm-rust/crates/core/src/chat_completions/tests.rs index f8594dee447..b5718aa2815 100644 --- a/litellm-rust/crates/core/src/chat_completions/tests.rs +++ b/litellm-rust/crates/core/src/chat_completions/tests.rs @@ -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 { - 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!( diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs index 3238d09b6b5..b096df4664d 100644 --- a/litellm-rust/crates/core/src/chat_completions/types.rs +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -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, - pub api_key: Option<&'a str>, - pub api_base: Option<&'a str>, - pub custom_llm_provider: Option<&'a str>, - pub extra_headers: Option>, - pub timeout: Option, + 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, pub(super) optional_params: Map, - pub(super) api_key: Option<&'a str>, - pub(super) api_base: Option<&'a str>, - pub(super) extra_headers: Option>, - pub(super) timeout: Option, + 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, + pub(super) provider_connection: Map, pub(super) timeout: Option, } diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index b93e084f57e..952dd94ba19 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -16,3 +16,6 @@ pub mod router; pub mod routing_utils; pub use error::Error; + +pub mod request_context; +pub mod request_options; diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index cfa8bda1104..10bcb7ab451 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -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 { +pub async fn messages( + request: MessagesRequest<'_>, + _context: &LiteLlmRequestContext, +) -> Result { execute_messages_provider_call(request).await } -pub async fn messages_stream(request: MessagesRequest<'_>) -> Result { +pub async fn messages_stream( + request: MessagesRequest<'_>, + _context: &LiteLlmRequestContext, +) -> Result { execute_messages_provider_stream(request).await } diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index ec83d03f535..27ea0f14f87 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -9,20 +9,25 @@ use serde_json::{Map, Value}; pub(super) fn prepare_provider_request( request: MessagesRequest<'_>, ) -> Result { - 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, }) } diff --git a/litellm-rust/crates/core/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs index df9f7051011..fd7d8f8c453 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -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"); diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index b9f807c29fd..7ad0a36cc3e 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -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>, - pub timeout: Option, + pub options: RequestOptions, } pub(super) struct ProviderMessagesRequest { diff --git a/litellm-rust/crates/core/src/request_context.rs b/litellm-rust/crates/core/src/request_context.rs new file mode 100644 index 00000000000..ea8e6096a4f --- /dev/null +++ b/litellm-rust/crates/core/src/request_context.rs @@ -0,0 +1,18 @@ +use serde_json::{Map, Value}; + +#[derive(Clone, Debug, Default, PartialEq)] +pub struct RequestAttribution { + pub user_api_key_hash: Option, + pub user_api_key_user_id: Option, + pub user_api_key_team_id: Option, +} + +#[derive(Clone, Debug, Default, PartialEq)] +pub struct LiteLlmRequestContext { + pub metadata: Option>, + pub litellm_metadata: Option>, + pub request_metadata_fields: Vec, + pub litellm_call_id: Option, + pub request_model: Option, + pub attribution: RequestAttribution, +} diff --git a/litellm-rust/crates/core/src/request_options.rs b/litellm-rust/crates/core/src/request_options.rs new file mode 100644 index 00000000000..9ab0a81c998 --- /dev/null +++ b/litellm-rust/crates/core/src/request_options.rs @@ -0,0 +1,14 @@ +use std::time::Duration; + +use serde_json::{Map, Value}; + +#[derive(Clone, Debug, Default)] +pub struct RequestOptions { + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, + pub extra_headers: Option>, + pub extra_query: Option>, + pub timeout: Option, + pub provider_connection: Map, +} diff --git a/litellm-rust/crates/core/src/responses/types.rs b/litellm-rust/crates/core/src/responses/types.rs index 4942309992e..009310e3734 100644 --- a/litellm-rust/crates/core/src/responses/types.rs +++ b/litellm-rust/crates/core/src/responses/types.rs @@ -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, diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 384f0be5a1b..0ff434b4107 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -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, - timeout_seconds: Option, + request: WebSocketConnectRequest, + context: NativeRequestContext, ) -> PyResult> { - 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" diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index a14e4b55d82..8e4543b7997 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -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, - pub(crate) api_base: Option, - pub(crate) custom_llm_provider: Option, - pub(crate) extra_headers: Option>, - pub(crate) timeout: Option, +#[derive(FromPyObject)] +pub(crate) struct NativeRequestOptions { + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + #[pyo3(from_py_with = litellm_python_interop::from_py)] + extra_headers: Option>, + #[pyo3(from_py_with = litellm_python_interop::from_py)] + extra_query: Option>, + timeout_seconds: Option, + #[pyo3(from_py_with = litellm_python_interop::from_py)] + provider_connection: Option>, } -pub(crate) struct RouteOptionsInputs { - pub(crate) model: String, - pub(crate) api_key: Option, - pub(crate) api_base: Option, - pub(crate) custom_llm_provider: Option, - pub(crate) extra_headers: Option, - pub(crate) timeout_seconds: Option, +impl From 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 { - 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, + user_api_key_user_id: Option, + user_api_key_team_id: Option, +} + +#[derive(FromPyObject)] +pub(crate) struct NativeRequestContext { + #[pyo3(from_py_with = litellm_python_interop::from_py)] + metadata: Option>, + #[pyo3(from_py_with = litellm_python_interop::from_py)] + litellm_metadata: Option>, + request_metadata_fields: Vec, + litellm_call_id: Option, + request_model: Option, + attribution: NativeRequestAttribution, +} + +impl From 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, -) -> PyResult> { - match value { - Some(value) => object(name, value), - None => Ok(Map::new()), - } -} - -fn optional_object( - name: &'static str, - value: Option, -) -> PyResult>> { - value.map(|value| object(name, value)).transpose() -} - -fn object(name: &'static str, value: Value) -> PyResult> { - 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) -> Option { timeout_seconds.and_then(|secs| { if secs.is_finite() && secs > 0.0 { @@ -84,21 +92,55 @@ pub(crate) fn optional_timeout(timeout_seconds: Option) -> Option }) } -pub(crate) fn marshal_headers(headers: Option) -> PyResult> { - 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 } diff --git a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs index af60515b0e2..0bb93bbf132 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -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, + options: NativeRequestOptions, +} fn prepare_transcription( - inputs: AudioTranscriptionInputs, + input: AudioTranscriptionInputs, + context: NativeRequestContext, ) -> PyResult> + 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, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - extra_headers: Option, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - optional_params: Option, - timeout_seconds: Option, - }, + request = AudioTranscriptionInputs, prepare = prepare_transcription, errors = core_error_to_pyerr, } diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 08ab476005c..96dad4c6328 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -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, + options: NativeRequestOptions, +} fn prepare_chat_completions( - inputs: ChatCompletionsInputs, + input: ChatCompletionsInputs, + context: NativeRequestContext, ) -> PyResult> + 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, custom_llm_provider: Option, ) -> PyResult> { - 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, - api_key: Option, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - extra_headers: Option, - timeout_seconds: Option, - }, + request = ChatCompletionsInputs, prepare = prepare_chat_completions, errors = chat_completions_error_to_pyerr, extra = [chat_completions_decline], diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index 3285da14d5f..3957fb9976f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -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> { - 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> { - 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> { - 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> { - 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> + 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::(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}], diff --git a/litellm-rust/crates/python-bridge/src/routes/messages.rs b/litellm-rust/crates/python-bridge/src/routes/messages.rs index f69b5e9251d..2ecb3690c15 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages.rs @@ -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> + 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, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - extra_headers: Option, - timeout_seconds: Option, - }, + request = MessagesInputs, prepare = prepare_messages, errors = core_error_to_pyerr, } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index cc2f8e43cea..f4a3c8ddc6e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -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, + options: NativeRequestOptions, +} fn prepare_ocr( - inputs: OcrInputs, + input: OcrInputs, + context: NativeRequestContext, ) -> PyResult> + 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, - api_base: Option, - custom_llm_provider: Option, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - extra_headers: Option, - #[pyo3(from_py_with = litellm_python_interop::from_py)] - optional_params: Option, - timeout_seconds: Option, - }, + request = OcrInputs, prepare = prepare_ocr, errors = ocr_error_to_pyerr, } diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index ded533c5ad0..8cb5c80f2fd 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -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), diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index d21e40021af..d9ea26333bd 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -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), diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index b23f6ada45e..7bd7cdc717c 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -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), diff --git a/litellm/rust_bridge/protocols.py b/litellm/rust_bridge/protocols.py index b08fcbf81aa..76b76504a6c 100644 --- a/litellm/rust_bridge/protocols.py +++ b/litellm/rust_bridge/protocols.py @@ -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]]: ... diff --git a/litellm/rust_bridge/request.py b/litellm/rust_bridge/request.py new file mode 100644 index 00000000000..ba650f22372 --- /dev/null +++ b/litellm/rust_bridge/request.py @@ -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 diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 2616b7777dc..f46d5789bc1 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -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"), diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index c70cc72903a..7c89ae449c9 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -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), diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py b/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py index f8d7c55d4e2..e3a032c8d94 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py @@ -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}") diff --git a/tests/rust-python-harness/strategies/trace_parity/sdk/transcription/case.py b/tests/rust-python-harness/strategies/trace_parity/sdk/transcription/case.py index 2fb1054d10a..a80eb8f07b5 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/transcription/case.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/transcription/case.py @@ -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( { diff --git a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py index d844faa5d7a..5f83e9c9950 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -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") diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py index da70f422f44..6751cf7f06f 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py @@ -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( diff --git a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py index becd6ecb832..c372b295d64 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py +++ b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py @@ -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() diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 6b91cc7b75e..9fdb98baa75 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -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") diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 9ff4107f359..cbf183dea91 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -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") diff --git a/tests/test_litellm/rust_bridge/native_route_wheel_test.py b/tests/test_litellm/rust_bridge/native_route_wheel_test.py index a7f50a82a99..5bc754bb724 100644 --- a/tests/test_litellm/rust_bridge/native_route_wheel_test.py +++ b/tests/test_litellm/rust_bridge/native_route_wheel_test.py @@ -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") diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index b5d4410ab5c..ab1c4306ab9 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -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", + } diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 116b8c42ac8..72c60463fc5 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -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) diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 4087bf3373b..57fbd6875d6 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22181 + "limit": 22165 }, "LIT002": { "limit": 26745