From 1298be50323630e53b5f2fdc230b04555104db7a Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 12:59:50 -0700 Subject: [PATCH 1/4] 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 | 34 +- .../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/ocr/main.py | 275 +++++++++++---- litellm/rust_bridge/chat_completions.py | 197 +++++++---- litellm/rust_bridge/messages.py | 101 +++--- litellm/rust_bridge/ocr.py | 263 ++++----------- litellm/rust_bridge/protocols.py | 131 ++------ litellm/rust_bridge/request.py | 122 +++++++ litellm/rust_bridge/responses_websocket.py | 68 ++-- litellm/rust_bridge/transcription.py | 110 +++--- .../strategies/trace_parity/sdk/execution.py | 56 +++- .../trace_parity/sdk/transcription/case.py | 7 +- .../test_rust_bridge_messages.py | 218 ++++-------- .../chat/test_anthropic_chat_handler.py | 23 +- .../chat/test_bedrock_converse_handler.py | 27 +- tests/test_litellm/ocr/test_rust_bridge.py | 222 ++++++++----- .../responses/test_rust_bridge_websocket.py | 88 +---- .../rust_bridge/native_route_wheel_test.py | 45 ++- .../rust_bridge/test_chat_completions.py | 168 +++++++++- .../test_audio_transcription_rust_bridge.py | 48 ++- type-discipline-budget.json | 8 +- 62 files changed, 2530 insertions(+), 1984 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 c1c7b911a08..ee4a99ca094 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 9103b942684..0a018b19a0e 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 a70d1b78a3a..efa7351abcb 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| { @@ -101,11 +101,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 dee150e0971..f4e74b8563b 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,9 +40,13 @@ 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) - .map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?; + _context: &LiteLlmRequestContext, +) -> Result { + let (model, config) = resolve_provider_config( + request.model, + request.options.custom_llm_provider.as_deref(), + ) + .map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?; let messages = parse_messages(request.messages).map_err(|_| Error::Declined("unreadable message list"))?; if messages.is_empty() { @@ -55,25 +60,22 @@ pub(super) fn resolve_request( config, messages, optional_params: request.optional_params, - api_key: request.api_key, - api_base: request.api_base, - extra_headers: request.extra_headers, - timeout: request.timeout, + options: request.options, }) } #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] fn validate_environment( - request: &ResolvedChatCompletionsRequest<'_>, + request: &ResolvedChatCompletionsRequest, model: &str, config: &dyn ChatCompletionsProviderConfig, ) -> Result<(Vec<(String, String)>, ChatCompletionsAuth), Error> { let env_lookup = |key: &str| std::env::var(key).ok(); - let mut headers = string_headers(request.extra_headers.clone())?; + let mut headers = string_headers(request.options.extra_headers.clone())?; let auth = config.auth( - request.api_key, + request.options.api_key.as_deref(), model, - &request.optional_params, + &request.options.provider_connection, &env_lookup, )?; match &auth { @@ -117,16 +119,16 @@ fn validate_environment( } pub(super) fn prepare_provider_request( - request: ResolvedChatCompletionsRequest<'_>, + request: ResolvedChatCompletionsRequest, ) -> Result { 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 0104729650d..d25d45b9de0 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::Declined("streaming")); @@ -261,7 +272,7 @@ fn rejects_non_string_extra_headers() { json!([{"role": "user", "content": "hi"}]), json!({}), ); - call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))])); + call.options.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))])); assert_eq!( decline(call), Error::InvalidRequest( @@ -279,7 +290,7 @@ fn prepares_a_bedrock_call_without_resolving_credentials() { json!([{"role": "user", "content": "hi"}]), json!({"maxTokens": 16}), ); - call.api_key = None; + call.options.api_key = None; let prepared = prepare_chat_completions_call(call).expect("prepares"); assert_eq!( prepared.url, @@ -312,15 +323,18 @@ async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() { "bedrock/us-east-1/anthropic.claude-v2", None, json!([{"role": "user", "content": "hi"}]), - json!({ - "maxTokens": 16, - "aws_access_key_id": "AKIDEXAMPLE", - "aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY" - }), + json!({"maxTokens": 16}), ); + call.options.provider_connection = Map::from_iter([ + ("aws_access_key_id".to_string(), json!("AKIDEXAMPLE")), + ( + "aws_secret_access_key".to_string(), + json!("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"), + ), + ]); // A key would resolve to a bearer token and never reach the signer. - call.api_key = None; - call.extra_headers = Some(Map::from_iter([( + call.options.api_key = None; + call.options.extra_headers = Some(Map::from_iter([( "x-request-id".to_string(), json!("abc-123"), )])); @@ -367,14 +381,18 @@ async fn a_forwarded_header_the_signer_computes_declines_to_python() { "bedrock/us-east-1/anthropic.claude-v2", None, json!([{"role": "user", "content": "hi"}]), - json!({ - "maxTokens": 16, - "aws_access_key_id": "AKIDEXAMPLE", - "aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY" - }), + json!({"maxTokens": 16}), ); - call.api_key = None; - call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])); + call.options.provider_connection = Map::from_iter([ + ("aws_access_key_id".to_string(), json!("AKIDEXAMPLE")), + ( + "aws_secret_access_key".to_string(), + json!("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY"), + ), + ]); + call.options.api_key = None; + call.options.extra_headers = + Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])); let prepared = prepare_chat_completions_call(call).expect("prepares"); let error = super::handler::signed_headers(&prepared, br#"{"a":1}"#) .await @@ -399,7 +417,7 @@ fn a_bedrock_deployment_bearer_outranks_a_forwarded_authorization() { json!([{"role": "user", "content": "hi"}]), json!({"maxTokens": 16}), ); - call.extra_headers = Some(Map::from_iter([( + call.options.extra_headers = Some(Map::from_iter([( "Authorization".to_string(), json!("Bearer caller-supplied"), )])); @@ -432,7 +450,7 @@ fn an_anthropic_forwarded_oauth_bearer_still_outranks_the_resolved_key() { json!([{"role": "user", "content": "hi"}]), json!({}), ); - call.extra_headers = Some(Map::from_iter([( + call.options.extra_headers = Some(Map::from_iter([( "authorization".to_string(), json!("Bearer sk-ant-oat01-forwarded"), )])); @@ -666,11 +684,15 @@ mod round_trip { Value::Object(map) => map, other => panic!("params must be an object, got {other}"), }, - api_key: Some("sk-test"), - api_base: Some(api_base), - custom_llm_provider: None, - extra_headers: None, - timeout: Some(std::time::Duration::from_secs(10)), + + options: RequestOptions { + api_key: (Some("sk-test")).map(|value| value.to_string()), + api_base: (Some(api_base)).map(|value| value.to_string()), + custom_llm_provider: None, + extra_headers: None, + timeout: Some(std::time::Duration::from_secs(10)), + ..Default::default() + }, } } @@ -679,14 +701,19 @@ mod round_trip { #[tokio::test] async fn round_trip_sends_the_translated_body_and_normalizes_the_response() { let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await; - let response = chat_completions(call( - &api_base, - json!([ - {"role": "system", "content": "be terse"}, - {"role": "user", "content": "hi"} - ]), - json!({"max_tokens": 16}), - )) + let response = chat_completions( + call( + &api_base, + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"} + ]), + json!({"max_tokens": 16}), + ), + &LiteLlmRequestContext { + ..Default::default() + }, + ) .await .expect("call succeeds"); @@ -724,11 +751,16 @@ mod round_trip { const NO_USAGE: &str = r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#; let (api_base, handle) = serve_once("200 OK", NO_USAGE).await; - let err = chat_completions(call( - &api_base, - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 16}), - )) + let err = chat_completions( + call( + &api_base, + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + ), + &LiteLlmRequestContext { + ..Default::default() + }, + ) .await .expect_err("response cannot be normalized"); handle.await.expect("server task"); @@ -742,11 +774,16 @@ mod round_trip { async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() { const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#; let (api_base, handle) = serve_once("200 OK", TOOL_USE).await; - let err = chat_completions(call( - &api_base, - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 16}), - )) + let err = chat_completions( + call( + &api_base, + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + ), + &LiteLlmRequestContext { + ..Default::default() + }, + ) .await .expect_err("response cannot be normalized"); handle.await.expect("server task"); @@ -760,11 +797,16 @@ mod round_trip { async fn an_upstream_error_status_keeps_its_code() { let (api_base, handle) = serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await; - let err = chat_completions(call( - &api_base, - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 16}), - )) + let err = chat_completions( + call( + &api_base, + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + ), + &LiteLlmRequestContext { + ..Default::default() + }, + ) .await .expect_err("upstream rejects"); handle.await.expect("server task"); @@ -781,11 +823,16 @@ mod round_trip { listener.local_addr().expect("has an address").port() // Dropped here, so the port is closed and the connect is refused. }; - let err = chat_completions(call( - &format!("http://127.0.0.1:{port}/v1/messages"), - json!([{"role": "user", "content": "hi"}]), - json!({"max_tokens": 16}), - )) + let err = chat_completions( + call( + &format!("http://127.0.0.1:{port}/v1/messages"), + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + ), + &LiteLlmRequestContext { + ..Default::default() + }, + ) .await .expect_err("nothing is listening"); assert!( 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/ocr/main.py b/litellm/ocr/main.py index 7fea86e10b8..9c77fc39909 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -8,6 +8,7 @@ import mimetypes import os import re from collections.abc import Callable, Coroutine, Mapping +from dataclasses import dataclass from io import IOBase from typing import Any, Final, cast @@ -17,16 +18,24 @@ import litellm from litellm._logging import verbose_logger from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model +from litellm.llms.azure_ai.ocr.common_utils import ( + is_azure_document_intelligence_model, +) from litellm.llms.base_llm.ocr.transformation import ( OCR_REQUEST_FORMAT_PARAM, + BaseOCRConfig, OCRResponse, parse_ocr_request_format, ) from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import ocr as rust_ocr_bridge -from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors -from litellm.rust_bridge.runtime import DispatchResult +from litellm.rust_bridge.request import ( + NativeRequestOptions, + PreparedNativeCall, + provider_connection_params, + provider_request_params, +) +from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -35,6 +44,28 @@ base_llm_http_handler = BaseLLMHTTPHandler() ################################################# +@dataclass +class _PreparedOCRRequest: + model: str + document: dict[str, Any] + api_key: str | None + api_base: str | None + custom_llm_provider: str + extra_headers: dict[str, object] | None + provider_config: BaseOCRConfig + optional_params: dict[str, object] + litellm_params: dict[str, object] + effective_timeout: float | httpx.Timeout + litellm_logging_obj: LiteLLMLoggingObj + + +_RUST_OCR_PROVIDERS: Final = { + "mistral", + "azure_ai", + "vertex_ai", +} + + def _prepare_ocr_request( model: str, document: Mapping[str, object], @@ -44,7 +75,7 @@ def _prepare_ocr_request( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, kwargs: dict[str, object], -) -> rust_ocr_bridge.PreparedOCRRequest: +) -> _PreparedOCRRequest: litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj")) litellm_call_id: Final = cast(str | None, kwargs.get("litellm_call_id", None)) @@ -141,7 +172,7 @@ def _prepare_ocr_request( custom_llm_provider=custom_llm_provider, ) - return rust_ocr_bridge.PreparedOCRRequest( + return _PreparedOCRRequest( model=model, document=document, api_key=api_key, @@ -156,71 +187,143 @@ def _prepare_ocr_request( ) -@anative_first( - native=rust_ocr_bridge.aattempt_ocr, - route="ocr", - errors=lambda prepared_request, resolve_api_key: provider_errors( - prepared_request.custom_llm_provider, prepared_request.model - ), -) -async def _execute_aocr( - prepared_request: rust_ocr_bridge.PreparedOCRRequest, +def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool: + if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": + return False + if not prepared_request.provider_config.supports_rust_bridge(): + return False + return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS + + +def _rust_bridge_optional_params( + prepared_request: _PreparedOCRRequest, + resolve_secret: Callable[[str], str | None], +) -> dict[str, object]: + optional_params: Final = dict(prepared_request.optional_params) + if prepared_request.custom_llm_provider == "vertex_ai": + vertex_project: Final = ( + prepared_request.litellm_params.get("vertex_project") + or prepared_request.litellm_params.get("vertex_ai_project") + or litellm.vertex_project + or resolve_secret("VERTEXAI_PROJECT") + ) + vertex_location: Final = ( + prepared_request.litellm_params.get("vertex_location") + or prepared_request.litellm_params.get("vertex_ai_location") + or litellm.vertex_location + or resolve_secret("VERTEXAI_LOCATION") + or resolve_secret("VERTEX_LOCATION") + ) + if vertex_project is not None: + optional_params["vertex_project"] = vertex_project + if vertex_location is not None: + optional_params["vertex_location"] = vertex_location + return optional_params + + +def _rust_bridge_api_base( + prepared_request: _PreparedOCRRequest, + resolve_secret: Callable[[str], str | None], +) -> str | None: + if prepared_request.api_base is not None: + return prepared_request.api_base + if prepared_request.custom_llm_provider == "azure_ai": + if is_azure_document_intelligence_model(prepared_request.model): + return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") + return resolve_secret("AZURE_AI_API_BASE") + return None + + +def _prepare_rust_ocr_call( + prepared_request: _PreparedOCRRequest, resolve_api_key: Callable[[str], str | None], -) -> OCRResponse: - pending: Final = base_llm_http_handler.ocr( +) -> PreparedNativeCall[rust_ocr_bridge.NativeOCRRequest]: + provider_config: Final = prepared_request.provider_config + api_key_env_var: Final = provider_config.get_api_key_env_var() + resolved_api_key: Final = prepared_request.api_key or ( + resolve_api_key(api_key_env_var) if api_key_env_var is not None else None + ) + resolved_headers: Final = provider_config.validate_environment( + headers=prepared_request.extra_headers or {}, model=prepared_request.model, - document=prepared_request.document, - optional_params=prepared_request.optional_params, - timeout=prepared_request.effective_timeout, - logging_obj=prepared_request.litellm_logging_obj, - api_key=prepared_request.api_key, + api_key=resolved_api_key, api_base=prepared_request.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - aocr=True, - headers=prepared_request.extra_headers, - provider_config=prepared_request.provider_config, litellm_params=prepared_request.litellm_params, ) - response: Final = await pending if asyncio.iscoroutine(pending) else pending - if response is None: - raise ValueError(f"Got an unexpected None response from the OCR API: {response}") - return response - - -def _attempt_ocr( - prepared_request: rust_ocr_bridge.PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], - is_async: bool, -) -> DispatchResult[OCRResponse]: - return rust_ocr_bridge.attempt_ocr(prepared_request=prepared_request, resolve_api_key=resolve_api_key) - - -@native_first( - native=_attempt_ocr, - route="ocr", - errors=lambda prepared_request, resolve_api_key, is_async: provider_errors( - prepared_request.custom_llm_provider, prepared_request.model - ), -) -def _execute_ocr( - prepared_request: rust_ocr_bridge.PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], - is_async: bool, -) -> OCRResponse | Coroutine[object, object, OCRResponse]: - return base_llm_http_handler.ocr( - model=prepared_request.model, - document=prepared_request.document, - optional_params=prepared_request.optional_params, - timeout=prepared_request.effective_timeout, - logging_obj=prepared_request.litellm_logging_obj, - api_key=prepared_request.api_key, + resolved_complete_url: Final = provider_config.get_complete_url( api_base=prepared_request.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - aocr=is_async, - headers=prepared_request.extra_headers, - provider_config=prepared_request.provider_config, + model=prepared_request.model, + optional_params=prepared_request.optional_params, litellm_params=prepared_request.litellm_params, ) + rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key) + rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key) + prepared_request.litellm_logging_obj.pre_call( + input="OCR document processing", + api_key=resolved_api_key, + additional_args={ + "complete_input_dict": { + "model": prepared_request.model, + "document": prepared_request.document, + **rust_optional_params, + }, + "api_base": resolved_complete_url, + "headers": resolved_headers, + }, + ) + return PreparedNativeCall( + request=rust_ocr_bridge.NativeOCRRequest( + model=prepared_request.model, + document=prepared_request.document, + optional_params=provider_request_params(rust_optional_params), + options=NativeRequestOptions( + provider_connection=provider_connection_params(rust_optional_params), + api_key=resolved_api_key, + api_base=rust_api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=cast( # cast-ok: provider header normalization returns string-object pairs + dict[str, object], resolved_headers + ), + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + ), + ) + ) + + +def _run_rust_ocr( + prepared_request: _PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], + fallback: Callable[[], OCRResponse | Coroutine[object, object, OCRResponse]], +) -> OCRResponse | Coroutine[object, object, OCRResponse]: + return rust_ocr_bridge.dispatch_ocr( + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + fallback=fallback, + adapt=OCRResponse.model_validate, + model=prepared_request.model, + provider=prepared_request.custom_llm_provider, + eligible=_rust_ocr_supported(prepared_request), + ) + + +async def _run_rust_aocr( + prepared_request: _PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], + fallback: Callable[[], Coroutine[object, object, OCRResponse]], +) -> OCRResponse: + return await rust_ocr_bridge.adispatch_ocr( + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + fallback=fallback, + adapt=OCRResponse.model_validate, + model=prepared_request.model, + provider=prepared_request.custom_llm_provider, + eligible=_rust_ocr_supported(prepared_request), + ) @client @@ -319,7 +422,31 @@ async def aocr( from litellm.secret_managers.main import get_secret_str - return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str) + async def python_fallback() -> OCRResponse: + pending: Final = base_llm_http_handler.ocr( + model=prepared.model, + document=prepared.document, + optional_params=prepared.optional_params, + timeout=prepared.effective_timeout, + logging_obj=prepared.litellm_logging_obj, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared.custom_llm_provider, + aocr=True, + headers=prepared.extra_headers, + provider_config=prepared.provider_config, + litellm_params=prepared.litellm_params, + ) + response: Final = await pending if asyncio.iscoroutine(pending) else pending + if response is None: + raise ValueError(f"Got an unexpected None response from the OCR API: {response}") + return response + + return await _run_rust_aocr( + prepared_request=prepared, + resolve_api_key=get_secret_str, + fallback=python_fallback, + ) except Exception as e: raise litellm.exception_type( model=model, @@ -560,7 +687,27 @@ def ocr( from litellm.secret_managers.main import get_secret_str - return _execute_ocr(prepared_request=prepared, resolve_api_key=get_secret_str, is_async=_is_async) + def python_fallback() -> OCRResponse | Coroutine[object, object, OCRResponse]: + return base_llm_http_handler.ocr( + model=prepared.model, + document=prepared.document, + optional_params=prepared.optional_params, + timeout=prepared.effective_timeout, + logging_obj=prepared.litellm_logging_obj, + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared.custom_llm_provider, + aocr=_is_async, + headers=prepared.extra_headers, + provider_config=prepared.provider_config, + litellm_params=prepared.litellm_params, + ) + + return _run_rust_ocr( + prepared_request=prepared, + resolve_api_key=get_secret_str, + fallback=python_fallback, + ) except Exception as e: raise litellm.exception_type( model=model, diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index c463a82c431..8cb5c80f2fd 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -4,12 +4,16 @@ The Rust core owns the conversation translation, the provider call, and the response normalization for the subset of `/chat/completions` requests it accepts. This module only marshals inputs and hands the normalized result to LiteLLM's existing `ModelResponse` builder. + +``None`` means the provider was never called, so the caller is free to serve the +request on the Python path. A failure after the call was issued raises instead: +retrying it there would bill the customer for the same work twice. """ from __future__ import annotations import json -from collections.abc import Mapping, Sequence +from collections.abc import Awaitable, Callable, Mapping, Sequence from typing import TYPE_CHECKING, Final, Protocol import httpx @@ -20,14 +24,28 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo convert_to_model_response_object, ) from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned -from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, Unchanged from litellm.rust_bridge.configuration import rust_enabled from litellm.rust_bridge.protocols import ( RustAchatCompletions, RustChatCompletions, RustChatCompletionsDecline, ) -from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt +from litellm.rust_bridge.request import ( + NativeChatCompletionsRequest, + NativeRequestContext, + NativeRequestOptions, + PreparedNativeCall, + call_native, + provider_connection_params, + provider_request_params, +) +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + EndpointBinding, + EndpointDispatch, + async_none, +) from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.utils import ModelResponse @@ -85,10 +103,16 @@ def response_logger( return log -_CHAT: Final[NativeBinding[RustChatCompletions]] = NativeBinding(lambda native: native.chat_completions) -_ACHAT: Final[NativeBinding[RustAchatCompletions]] = NativeBinding(lambda native: native.achat_completions) -_CHAT_PREFLIGHT: Final[NativeBinding[RustChatCompletionsDecline]] = NativeBinding( - lambda native: native.chat_completions_decline +_CHAT: Final[EndpointDispatch[RustChatCompletions, RustAchatCompletions]] = EndpointDispatch.native( + route="chat_completions", + sync=lambda native: native.chat_completions, + asynchronous=lambda native: native.achat_completions, + enabled=rust_enabled, +) +_CHAT_PREFLIGHT: Final[EndpointBinding[RustChatCompletionsDecline]] = EndpointBinding.native( + route="chat_completions", + select=lambda native: native.chat_completions_decline, + enabled=rust_enabled, ) @@ -102,14 +126,14 @@ def set_rust_chat_completions( patching module attributes.""" if not isinstance(chat_completions, Unchanged): if chat_completions is None: - _CHAT.reset() + _CHAT.sync.reset() else: - _CHAT.override(chat_completions) + _CHAT.sync.override(chat_completions) if not isinstance(achat_completions, Unchanged): if achat_completions is None: - _ACHAT.reset() + _CHAT.asynchronous.reset() else: - _ACHAT.override(achat_completions) + _CHAT.asynchronous.override(achat_completions) if not isinstance(decline, Unchanged): if decline is None: _CHAT_PREFLIGHT.reset() @@ -178,24 +202,14 @@ def rust_chat_completions_accepts( if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params): verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path") return False - if not rust_enabled(): - return False - decline: Final = _CHAT_PREFLIGHT.load() - if decline is None: - return False - try: - reason: Final = decline( + return _CHAT_PREFLIGHT.accepts( + check=lambda decline: decline( model=model, messages=messages, optional_params=optional_params, custom_llm_provider=custom_llm_provider, - ) - except Exception as error: # noqa: BLE001 # capability checks perform no provider I/O - verbose_logger.debug("Native chat acceptance check failed: %s", error) - return False - if reason is not None: - verbose_logger.debug("Native chat request is ineligible: %s", reason) - return reason is None + ), + ) def _build_model_response( @@ -224,31 +238,32 @@ def chat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, - eligible: bool = True, -) -> DispatchResult[ModelResponse]: +) -> ModelResponse | None: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) return _build_model_response(rust_response, model_response) - def call(native: RustChatCompletions, timeout_seconds: float | None) -> Mapping[str, object]: - return native( - model=model, - messages=messages, - optional_params=optional_params, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_seconds, - ) - - return attempt( - load=_CHAT.load, - enabled=rust_enabled(), - eligible=eligible, - prepare=lambda: timeout_to_seconds(timeout), - call=call, + return _CHAT.invoke( + prepare=lambda: PreparedNativeCall( + NativeChatCompletionsRequest( + model=model, + messages=messages, + optional_params=provider_request_params(optional_params), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + provider_connection=provider_connection_params(optional_params), + ), + ), + context=NativeRequestContext(), + ), + call=call_native, + fallback=lambda: None, adapt=adapt, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -264,29 +279,81 @@ async def achat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, - eligible: bool = True, -) -> DispatchResult[ModelResponse]: +) -> ModelResponse | None: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) return _build_model_response(rust_response, model_response) - async def call(native: RustAchatCompletions, timeout_seconds: float | None) -> Mapping[str, object]: - return await native( - model=model, - messages=messages, - optional_params=optional_params, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_seconds, - ) - - return await aattempt( - load=_ACHAT.load, - enabled=rust_enabled(), - eligible=eligible, - prepare=lambda: timeout_to_seconds(timeout), - call=call, + return await _CHAT.ainvoke( + prepare=lambda: PreparedNativeCall( + NativeChatCompletionsRequest( + model=model, + messages=messages, + optional_params=provider_request_params(optional_params), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + provider_connection=provider_connection_params(optional_params), + ), + ), + context=NativeRequestContext(), + ), + call=call_native, + fallback=async_none, adapt=adapt, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), + ) + + +async def achat_completions_or_fallback( + *, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object], + model_response: ModelResponse, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout: float | httpx.Timeout | None, + on_response: ResponseObserver, + python_fallback: Callable[[], Awaitable[object]], +) -> object: + """Await the Rust path, falling back to the caller's own Python path when + the bridge is unavailable or the call fails. + + The caller supplies the fallback, so the bridge stays free of provider + dispatch. This exists because a caller that dispatches asynchronously has + already returned a coroutine by the time a Rust failure surfaces, and so + cannot fall back on its own. + """ + + def adapt(rust_response: Mapping[str, object]) -> object: + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + return await _CHAT.ainvoke( + prepare=lambda: PreparedNativeCall( + NativeChatCompletionsRequest( + model=model, + messages=messages, + optional_params=provider_request_params(optional_params), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + provider_connection=provider_connection_params(optional_params), + ), + ), + context=NativeRequestContext(), + ), + call=call_native, + fallback=python_fallback, + adapt=adapt, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 160e6f0f743..d9ea26333bd 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -6,13 +6,30 @@ from typing import Final import httpx -from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, Unchanged from litellm.rust_bridge.protocols import RustAmessages, RustMessages -from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity +from litellm.rust_bridge.request import ( + NativeMessagesRequest, + NativeRequestContext, + NativeRequestOptions, + PreparedNativeCall, + call_native, +) +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + EndpointDispatch, + always_enabled, + async_none, + identity, +) from litellm.rust_bridge.timeouts import timeout_to_seconds -_MESSAGES: Final[NativeBinding[RustMessages]] = NativeBinding(lambda native: native.messages) -_AMESSAGES: Final[NativeBinding[RustAmessages]] = NativeBinding(lambda native: native.amessages) +_MESSAGES: Final[EndpointDispatch[RustMessages, RustAmessages]] = EndpointDispatch.native( + route="messages", + sync=lambda native: native.messages, + asynchronous=lambda native: native.amessages, + enabled=always_enabled, +) def set_rust_messages( @@ -22,22 +39,22 @@ def set_rust_messages( ) -> None: if not isinstance(messages, Unchanged): if messages is None: - _MESSAGES.reset() + _MESSAGES.sync.reset() else: - _MESSAGES.override(messages) + _MESSAGES.sync.override(messages) if not isinstance(amessages, Unchanged): if amessages is None: - _AMESSAGES.reset() + _MESSAGES.asynchronous.reset() else: - _AMESSAGES.override(amessages) + _MESSAGES.asynchronous.override(amessages) def load_rust_messages() -> RustMessages | None: - return _MESSAGES.load() + return _MESSAGES.sync.load() def load_rust_amessages() -> RustAmessages | None: - return _AMESSAGES.load() + return _MESSAGES.asynchronous.load() def messages( @@ -49,22 +66,26 @@ def messages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: float | httpx.Timeout | None, -) -> DispatchResult[dict[str, object]]: - return attempt( - load=_MESSAGES.load, - enabled=True, - eligible=True, - prepare=lambda: timeout_to_seconds(timeout), - call=lambda rust_messages, timeout_seconds: rust_messages( - model=model, - body=body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_seconds, +) -> dict[str, object] | None: + return _MESSAGES.invoke( + prepare=lambda: PreparedNativeCall( + NativeMessagesRequest( + model=model, + body=body, + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ), + ), + context=NativeRequestContext(), ), + call=call_native, + fallback=lambda: None, adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -77,20 +98,24 @@ async def amessages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: float | httpx.Timeout | None, -) -> DispatchResult[dict[str, object]]: - return await aattempt( - load=_AMESSAGES.load, - enabled=True, - eligible=True, - prepare=lambda: timeout_to_seconds(timeout), - call=lambda rust_amessages, timeout_seconds: rust_amessages( - model=model, - body=body, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_seconds, +) -> dict[str, object] | None: + return await _MESSAGES.ainvoke( + prepare=lambda: PreparedNativeCall( + NativeMessagesRequest( + model=model, + body=body, + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ), + ), + context=NativeRequestContext(), ), + call=call_native, + fallback=async_none, adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index c04645cb29d..320f4b6718f 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -1,211 +1,90 @@ +"""Thin Python wrapper for the native Rust OCR bridge.""" + from __future__ import annotations -from collections.abc import Callable -from dataclasses import dataclass -from typing import Final +from collections.abc import Awaitable, Callable, Mapping +from typing import Final, TypeVar -import httpx -from pydantic import TypeAdapter +from . import configuration as _configuration +from .bindings import UNCHANGED, Unchanged +from .protocols import RustAocr, RustOcr +from .request import NativeOCRRequest, PreparedNativeCall, call_native +from .runtime import ( + BridgeErrorContext, + EndpointDispatch, +) -import litellm -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model -from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse -from litellm.rust_bridge import configuration as _configuration -from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.protocols import RustAocr, RustOcr -from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt -from litellm.rust_bridge.timeouts import timeout_to_seconds - -_OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr) -_AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr) -_HEADERS: Final = TypeAdapter(dict[str, object]) +rust_ocr_enabled = _configuration.rust_ocr_enabled +rust = _configuration.rust +ResultT = TypeVar("ResultT") -@dataclass(frozen=True, slots=True) -class PreparedOCRRequest: - model: str - document: dict[str, object] - api_key: str | None - api_base: str | None - custom_llm_provider: str - extra_headers: dict[str, object] | None - provider_config: BaseOCRConfig - optional_params: dict[str, object] - litellm_params: dict[str, object] - effective_timeout: float | httpx.Timeout - litellm_logging_obj: LiteLLMLoggingObj - - -@dataclass(frozen=True, slots=True) -class _PreparedRustOCRCall: - api_key: str | None - api_base: str | None - headers: dict[str, object] - optional_params: dict[str, object] - - -_RUST_OCR_PROVIDERS: Final = frozenset( - { - "mistral", - "azure_ai", - "vertex_ai", - } +_OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native( + route="ocr", + sync=lambda native: native.ocr, + asynchronous=lambda native: native.aocr, + enabled=_configuration.rust_ocr_enabled, ) +def set_rust_ocr( + *, + ocr: RustOcr | None | Unchanged = UNCHANGED, + aocr: RustAocr | None | Unchanged = UNCHANGED, +) -> None: + if not isinstance(ocr, Unchanged): + if ocr is None: + _OCR.sync.reset() + else: + _OCR.sync.override(ocr) + if not isinstance(aocr, Unchanged): + if aocr is None: + _OCR.asynchronous.reset() + else: + _OCR.asynchronous.override(aocr) + + def load_rust_ocr() -> RustOcr | None: - return _OCR.load() + return _OCR.sync.load() def load_rust_aocr() -> RustAocr | None: - return _AOCR.load() + return _OCR.asynchronous.load() -def _rust_ocr_supported(prepared_request: PreparedOCRRequest) -> bool: - if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": - return False - if not prepared_request.provider_config.supports_rust_bridge(): - return False - return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS - - -def _rust_bridge_optional_params( - prepared_request: PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> dict[str, object]: - if prepared_request.custom_llm_provider != "vertex_ai": - return prepared_request.optional_params - vertex_project: Final = ( - prepared_request.litellm_params.get("vertex_project") - or prepared_request.litellm_params.get("vertex_ai_project") - or litellm.vertex_project - or resolve_secret("VERTEXAI_PROJECT") - ) - vertex_location: Final = ( - prepared_request.litellm_params.get("vertex_location") - or prepared_request.litellm_params.get("vertex_ai_location") - or litellm.vertex_location - or resolve_secret("VERTEXAI_LOCATION") - or resolve_secret("VERTEX_LOCATION") - ) - return { - **prepared_request.optional_params, - **{ - name: value - for name, value in (("vertex_project", vertex_project), ("vertex_location", vertex_location)) - if value is not None - }, - } - - -def _rust_bridge_api_base( - prepared_request: PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> str | None: - if prepared_request.api_base is not None: - return prepared_request.api_base - if prepared_request.custom_llm_provider == "azure_ai": - if is_azure_document_intelligence_model(prepared_request.model): - return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") - return resolve_secret("AZURE_AI_API_BASE") - return None - - -def _prepare_rust_ocr_call( - prepared_request: PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], -) -> _PreparedRustOCRCall: - provider_config: Final = prepared_request.provider_config - api_key_env_var: Final = provider_config.get_api_key_env_var() - resolved_api_key: Final = prepared_request.api_key or ( - resolve_api_key(api_key_env_var) if api_key_env_var is not None else None - ) - resolved_headers: Final = _HEADERS.validate_python( - provider_config.validate_environment( - headers=prepared_request.extra_headers or {}, - model=prepared_request.model, - api_key=resolved_api_key, - api_base=prepared_request.api_base, - litellm_params=prepared_request.litellm_params, - ) - ) - resolved_complete_url: Final = provider_config.get_complete_url( - api_base=prepared_request.api_base, - model=prepared_request.model, - optional_params=prepared_request.optional_params, - litellm_params=prepared_request.litellm_params, - ) - rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key) - rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key) - prepared_request.litellm_logging_obj.pre_call( - input="OCR document processing", - api_key=resolved_api_key, - additional_args={ - "complete_input_dict": { - "model": prepared_request.model, - "document": prepared_request.document, - **rust_optional_params, - }, - "api_base": resolved_complete_url, - "headers": resolved_headers, - }, - ) - return _PreparedRustOCRCall( - api_key=resolved_api_key, - api_base=rust_api_base, - headers=resolved_headers, - optional_params=rust_optional_params, +def dispatch_ocr( + *, + prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]], + fallback: Callable[[], ResultT], + adapt: Callable[[Mapping[str, object]], ResultT], + model: str, + provider: str, + eligible: bool, +) -> ResultT: + return _OCR.invoke( + prepare=prepare, + call=call_native, + fallback=fallback, + adapt=adapt, + error_context=BridgeErrorContext(provider=provider, model=model), + eligible=eligible, ) -def attempt_ocr( - prepared_request: PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], -) -> DispatchResult[OCRResponse]: - return attempt( - load=_OCR.load, - enabled=_configuration.rust_enabled(), - prepare=lambda: _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ), - call=lambda native, prepared: native( - model=prepared_request.model, - document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, - timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), - ), - adapt=OCRResponse.model_validate, - eligible=_rust_ocr_supported(prepared_request), - ) - - -async def aattempt_ocr( - prepared_request: PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], -) -> DispatchResult[OCRResponse]: - return await aattempt( - load=_AOCR.load, - enabled=_configuration.rust_enabled(), - prepare=lambda: _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ), - call=lambda native, prepared: native( - model=prepared_request.model, - document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, - timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), - ), - adapt=OCRResponse.model_validate, - eligible=_rust_ocr_supported(prepared_request), +async def adispatch_ocr( + *, + prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]], + fallback: Callable[[], Awaitable[ResultT]], + adapt: Callable[[Mapping[str, object]], ResultT], + model: str, + provider: str, + eligible: bool, +) -> ResultT: + return await _OCR.ainvoke( + prepare=prepare, + call=call_native, + fallback=fallback, + adapt=adapt, + error_context=BridgeErrorContext(provider=provider, model=model), + eligible=eligible, ) 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 d673ab8431c..a48405a7884 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -2,24 +2,36 @@ from __future__ import annotations -from collections.abc import AsyncGenerator -from contextlib import AbstractAsyncContextManager, asynccontextmanager from typing import Final import httpx from websockets.exceptions import ConnectionClosedOK -from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, Unchanged from litellm.rust_bridge.configuration import rust_enabled from litellm.rust_bridge.protocols import ( RustResponsesWebSocket, RustResponsesWebSocketConnection, ) -from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result +from litellm.rust_bridge.request import ( + NativeRequestContext, + NativeRequestOptions, + NativeResponsesWebSocketRequest, + PreparedNativeCall, + call_native, +) +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + EndpointBinding, + async_none, + identity, +) from litellm.rust_bridge.timeouts import timeout_to_seconds -_RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding( - lambda native: native.ResponsesWebSocketConnection, +_RESPONSES_WEBSOCKET: Final[EndpointBinding[RustResponsesWebSocketConnection]] = EndpointBinding.native( + route="responses_websocket", + select=lambda native: native.ResponsesWebSocketConnection, + enabled=rust_enabled, ) @@ -34,7 +46,7 @@ def set_rust_responses_websocket( _RESPONSES_WEBSOCKET.override(connection) -class ConnectionAdapter: +class _ConnectionAdapter: def __init__(self, connection: RustResponsesWebSocket): self._connection: Final[RustResponsesWebSocket] = connection @@ -56,34 +68,18 @@ async def connect( url: str, headers: dict[str, str], timeout: float | httpx.Timeout | None, -) -> DispatchResult[ConnectionAdapter]: - return await aattempt( - load=_RESPONSES_WEBSOCKET.load, - enabled=rust_enabled(), - eligible=True, - prepare=lambda: timeout_to_seconds(timeout), - call=lambda connection_type, timeout_seconds: connection_type.connect( - url=url, - headers=headers, - timeout_seconds=timeout_seconds, +) -> _ConnectionAdapter | None: + connection: Final = await _RESPONSES_WEBSOCKET.ainvoke( + prepare=lambda: PreparedNativeCall( + NativeResponsesWebSocketRequest( + url=url, + options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)), + ), + context=NativeRequestContext(), ), - adapt=ConnectionAdapter, + call=lambda connection_type, request: call_native(connection_type.connect, request), + fallback=async_none, + adapt=identity, + error_context=BridgeErrorContext(provider="openai", model="responses websocket"), ) - - -@asynccontextmanager -async def _connection_context(connection: ConnectionAdapter) -> AsyncGenerator[ConnectionAdapter, None]: - try: - yield connection - finally: - await connection.close() - - -async def managed_connect( - *, - url: str, - headers: dict[str, str], - timeout: float | httpx.Timeout | None, -) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]: - result: Final = await connect(url=url, headers=headers, timeout=timeout) - return adapt_result(result, _connection_context) + return None if connection is None else _ConnectionAdapter(connection) diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 4970b86b18d..7c89ae449c9 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -4,38 +4,58 @@ from typing import Final import httpx -from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, Unchanged from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription -from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity +from litellm.rust_bridge.request import ( + NativeRequestContext, + NativeRequestOptions, + NativeTranscriptionRequest, + PreparedNativeCall, + call_native, + provider_connection_params, + provider_request_params, +) +from litellm.rust_bridge.runtime import ( + BridgeErrorContext, + EndpointDispatch, + always_enabled, + async_none, + identity, +) from litellm.rust_bridge.timeouts import timeout_to_seconds -_TRANSCRIPTION: Final[NativeBinding[RustTranscription]] = NativeBinding(lambda native: native.transcription) -_ATRANSCRIPTION: Final[NativeBinding[RustAtranscription]] = NativeBinding(lambda native: native.atranscription) +_TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] = EndpointDispatch.native( + route="audio transcription", + sync=lambda native: native.transcription, + asynchronous=lambda native: native.atranscription, + enabled=always_enabled, +) def configure_rust_transcription( + enabled: bool = True, *, transcription: RustTranscription | None | Unchanged = UNCHANGED, atranscription: RustAtranscription | None | Unchanged = UNCHANGED, ) -> None: if not isinstance(transcription, Unchanged): if transcription is None: - _TRANSCRIPTION.reset() + _TRANSCRIPTION.sync.reset() else: - _TRANSCRIPTION.override(transcription) + _TRANSCRIPTION.sync.override(transcription) if not isinstance(atranscription, Unchanged): if atranscription is None: - _ATRANSCRIPTION.reset() + _TRANSCRIPTION.asynchronous.reset() else: - _ATRANSCRIPTION.override(atranscription) + _TRANSCRIPTION.asynchronous.override(atranscription) def load_rust_transcription() -> RustTranscription | None: - return _TRANSCRIPTION.load() + return _TRANSCRIPTION.sync.load() def load_rust_atranscription() -> RustAtranscription | None: - return _ATRANSCRIPTION.load() + return _TRANSCRIPTION.asynchronous.load() def transcription( @@ -48,23 +68,28 @@ def transcription( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, -) -> DispatchResult[dict[str, object]]: - return attempt( - load=_TRANSCRIPTION.load, - enabled=True, - eligible=True, - prepare=lambda: timeout_to_seconds(timeout), - call=lambda rust_transcription, timeout_seconds: rust_transcription( - model=model, - audio=audio, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_seconds, +) -> dict[str, object] | None: + return _TRANSCRIPTION.invoke( + prepare=lambda: PreparedNativeCall( + NativeTranscriptionRequest( + model=model, + audio=audio, + optional_params=provider_request_params(optional_params), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + provider_connection=provider_connection_params(optional_params), + ), + ), + context=NativeRequestContext(), ), + call=call_native, + fallback=lambda: None, adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -78,21 +103,26 @@ async def atranscription( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, -) -> DispatchResult[dict[str, object]]: - return await aattempt( - load=_ATRANSCRIPTION.load, - enabled=True, - eligible=True, - prepare=lambda: timeout_to_seconds(timeout), - call=lambda rust_atranscription, timeout_seconds: rust_atranscription( - model=model, - audio=audio, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - optional_params=optional_params, - timeout_seconds=timeout_seconds, +) -> dict[str, object] | None: + return await _TRANSCRIPTION.ainvoke( + prepare=lambda: PreparedNativeCall( + NativeTranscriptionRequest( + model=model, + audio=audio, + optional_params=provider_request_params(optional_params), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + provider_connection=provider_connection_params(optional_params), + ), + ), + context=NativeRequestContext(), ), + call=call_native, + fallback=async_none, adapt=identity, + error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) 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 ec198939806..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,7 +9,7 @@ import pytest import litellm from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import configuration -from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason +from litellm.rust_bridge.request import NativeMessagesRequest, NativeRequestContext from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) @@ -41,23 +41,19 @@ class RecordingMessages: def __call__( self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, + request: NativeMessagesRequest, + *, + context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( { - "model": model, - "body": body, - "api_key": api_key, - "api_base": api_base, - "custom_llm_provider": custom_llm_provider, - "extra_headers": extra_headers, - "timeout_seconds": timeout_seconds, + "model": request.model, + "body": request.body, + "api_key": request.options.api_key, + "api_base": request.options.api_base, + "custom_llm_provider": request.options.custom_llm_provider, + "extra_headers": request.options.extra_headers, + "timeout_seconds": request.options.timeout_seconds, } ) return dict(FAKE_MESSAGES_RESPONSE) @@ -69,23 +65,19 @@ class RecordingAsyncMessages: async def __call__( self, - model: str, - body: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - timeout_seconds: float | None, + request: NativeMessagesRequest, + *, + context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( { - "model": model, - "body": body, - "api_key": api_key, - "api_base": api_base, - "custom_llm_provider": custom_llm_provider, - "extra_headers": extra_headers, - "timeout_seconds": timeout_seconds, + "model": request.model, + "body": request.body, + "api_key": request.options.api_key, + "api_base": request.options.api_base, + "custom_llm_provider": request.options.custom_llm_provider, + "extra_headers": request.options.extra_headers, + "timeout_seconds": request.options.timeout_seconds, } ) return dict(FAKE_MESSAGES_RESPONSE) @@ -95,7 +87,7 @@ class ExplodingAsyncMessages: def __init__(self) -> None: self.calls = 0 - async def __call__(self, **kwargs: object) -> dict[str, object]: + async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]: self.calls += 1 raise AssertionError("bridge must not be called") @@ -104,7 +96,7 @@ class RaisingAsyncMessages: def __init__(self) -> None: self.calls = 0 - async def __call__(self, **kwargs: object) -> dict[str, object]: + async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]: self.calls += 1 raise RuntimeError("upstream request failed with status 400: bad request") @@ -127,6 +119,16 @@ def test_load_rust_messages_returns_injected_impl(): assert rust_messages.load_rust_messages() is bridge +def test_bare_rust_still_toggles_ocr(): + from litellm.rust_bridge.ocr import rust_ocr_enabled + + litellm.rust(True) + assert rust_ocr_enabled() is True + + litellm.rust(False) + assert rust_ocr_enabled() is False + + def test_load_rust_amessages_returns_injected_impl(): bridge = RecordingAsyncMessages() litellm.rust(True) @@ -134,7 +136,7 @@ def test_load_rust_amessages_returns_injected_impl(): assert rust_messages.load_rust_amessages() is bridge -def test_messages_wrapper_reports_unavailable(monkeypatch): +def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): monkeypatch.setattr( importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", @@ -151,7 +153,7 @@ def test_messages_wrapper_reports_unavailable(monkeypatch): extra_headers={}, timeout=30.0, ) - assert result == NativeSkipped(NativeSkipReason.UNAVAILABLE) + assert result is None def test_messages_wrapper_forwards_args_and_converts_timeout(): @@ -169,7 +171,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout(): timeout=httpx.Timeout(600.0, read=42.0), ) - assert response == Handled(FAKE_MESSAGES_RESPONSE) + assert response == FAKE_MESSAGES_RESPONSE assert bridge.calls[0] == { "model": "claude-sonnet-4-5", "body": REQUEST_BODY, @@ -197,7 +199,7 @@ async def test_amessages_wrapper_forwards_args(): timeout=12.5, ) - assert response == Handled(FAKE_MESSAGES_RESPONSE) + assert response == FAKE_MESSAGES_RESPONSE assert bridge.calls[0]["model"] == "claude-sonnet-4-5" assert bridge.calls[0]["timeout_seconds"] == 12.5 @@ -215,7 +217,7 @@ def _gate(**overrides): "timeout": 30.0, } kwargs.update(overrides) - return BaseLLMHTTPHandler._attempt_rust_anthropic_messages(**kwargs) + return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs) @pytest.mark.asyncio @@ -226,8 +228,7 @@ async def test_gate_invokes_rust_and_marks_response_header(): response = await _gate() - assert isinstance(response, Handled) - response = response.value + assert response is not None assert response["id"] == "msg_123" assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} call = bridge.calls[0] @@ -240,13 +241,13 @@ async def test_gate_invokes_rust_and_marks_response_header(): @pytest.mark.asyncio -async def test_gate_reports_failure_to_harness(): +async def test_gate_propagates_unknown_native_errors(): bridge = RaisingAsyncMessages() litellm.rust(True) rust_messages.set_rust_messages(amessages=bridge) - response = await _gate() - assert isinstance(response, NativeFailed) + with pytest.raises(RuntimeError, match="bad request"): + await _gate() assert bridge.calls == 1 @@ -257,7 +258,7 @@ async def test_gate_skips_rust_when_flag_absent(): response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure")) - assert isinstance(response, NativeSkipped) + assert response is None assert bridge.calls == 0 @@ -269,11 +270,22 @@ async def test_gate_uses_process_enable_without_request_override(): response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure")) - assert isinstance(response, Handled) - response = response.value + assert response is not None assert bridge.calls[0]["custom_llm_provider"] == "azure_ai" +@pytest.mark.asyncio +async def test_gate_ignores_request_flag_when_process_enabled(): + bridge = RecordingAsyncMessages() + litellm.rust(True) + rust_messages.set_rust_messages(amessages=bridge) + + response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False)) + + assert response is not None + assert len(bridge.calls) == 1 + + @pytest.mark.asyncio async def test_gate_invokes_rust_for_native_anthropic_provider(): bridge = RecordingAsyncMessages() @@ -288,8 +300,7 @@ async def test_gate_invokes_rust_for_native_anthropic_provider(): headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"}, ) - assert isinstance(response, Handled) - response = response.value + assert response is not None assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} assert bridge.calls[0]["custom_llm_provider"] == "anthropic" assert bridge.calls[0]["api_key"] == "sk-ant" @@ -306,8 +317,7 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch): litellm_params=GenericLiteLLMParams(api_key="sk-ant"), ) - assert isinstance(response, Handled) - response = response.value + assert response is not None assert bridge.calls[0]["custom_llm_provider"] == "anthropic" @@ -322,7 +332,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch): litellm_params=GenericLiteLLMParams(api_key="sk-ant"), ) - assert isinstance(response, NativeSkipped) + assert response is None assert bridge.calls == 0 @@ -334,7 +344,7 @@ async def test_gate_skips_rust_for_unsupported_provider(): response = await _gate(custom_llm_provider="openai") - assert isinstance(response, NativeSkipped) + assert response is None assert bridge.calls == 0 @@ -346,7 +356,7 @@ async def test_gate_skips_rust_for_agentic_hook(): response = await _gate(has_agentic_hook=True) - assert isinstance(response, NativeSkipped) + assert response is None assert bridge.calls == 0 @@ -362,8 +372,7 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag(): request_body=streaming_body, ) - assert isinstance(response, Handled) - response = response.value + assert response is not None assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} assert "stream" not in bridge.calls[0]["body"] assert bridge.calls[0]["body"] == REQUEST_BODY @@ -396,105 +405,4 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch): response = await _gate() - assert isinstance(response, NativeSkipped) - - -@pytest.mark.asyncio -@pytest.mark.parametrize("selection", ("native", "disabled", "failed", "declined", "upstream")) -async def test_messages_handler_runs_selected_backend_once(selection: str, monkeypatch: pytest.MonkeyPatch) -> None: - from datetime import datetime - from types import SimpleNamespace - - import httpx - - from litellm.exceptions import RateLimitError - from litellm.litellm_core_utils.litellm_logging import Logging - from litellm.llms.anthropic.experimental_pass_through.messages.transformation import AnthropicMessagesConfig - from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - from litellm.rust_bridge import bindings - - class Declined(Exception): - pass - - class Upstream(Exception): - pass - - error = ( - Upstream(429, "rate limited") - if selection == "upstream" - else Declined("unsupported") - if selection == "declined" - else RuntimeError("native failed") - if selection == "failed" - else None - ) - - class Native: - def __init__(self) -> None: - self.calls = 0 - - async def __call__(self, **kwargs: object) -> dict[str, object]: - self.calls += 1 - if error is not None: - raise error - return dict(FAKE_MESSAGES_RESPONSE) - - bridge = Native() - monkeypatch.setattr( - bindings, - "get_native_bridge", - lambda: SimpleNamespace( - RustBridgeDeclined=Declined, - RustUpstreamError=Upstream, - ), - ) - rust_messages.set_rust_messages(amessages=bridge) - litellm.rust(selection != "disabled") - requests: list[httpx.Request] = [] - - def respond(request: httpx.Request) -> httpx.Response: - requests.append(request) - return httpx.Response(200, json=FAKE_MESSAGES_RESPONSE) - - logging_obj = Logging( - model=FAKE_MESSAGES_RESPONSE["model"], - messages=[], - stream=False, - call_type="anthropic_messages", - start_time=datetime.now(), - litellm_call_id="harness-test", - function_id="harness-test", - ) - client = AsyncHTTPHandler() - await client.client.aclose() - async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as transport: - client.client = transport - - async def run(): - return await BaseLLMHTTPHandler().async_anthropic_messages_handler( - model=FAKE_MESSAGES_RESPONSE["model"], - messages=[{"role": "user", "content": "hello"}], - anthropic_messages_provider_config=AnthropicMessagesConfig(), - anthropic_messages_optional_request_params={"max_tokens": 10}, - custom_llm_provider="anthropic", - litellm_params=GenericLiteLLMParams(), - logging_obj=logging_obj, - api_key="sk-test", - api_base="https://example.test", - client=client, - ) - - if selection in ("failed", "upstream"): - with pytest.raises(RateLimitError if selection == "upstream" else RuntimeError) as caught: - await run() - if selection == "upstream": - assert caught.value.__cause__ is error - assert caught.value.llm_provider == "anthropic" - assert caught.value.model == FAKE_MESSAGES_RESPONSE["model"] - else: - assert caught.value is error - else: - response = await run() - assert response["id"] == FAKE_MESSAGES_RESPONSE["id"] - assert len(requests) == (1 if selection in ("disabled", "declined") else 0) - assert bridge.calls == (0 if selection == "disabled" else 1) + assert response is None 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 7e04a4f0f4b..8720e73e77d 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -12,7 +12,7 @@ import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge import configuration -from litellm.rust_bridge.runtime import Handled +from litellm.rust_bridge.request import NativeOCRRequest, NativeRequestContext, NativeRequestOptions, PreparedNativeCall from litellm.rust_bridge.timeouts import timeout_to_seconds # `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` @@ -49,25 +49,20 @@ class RecordingBridge: def __call__( self, - model: str, - document: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, + request: NativeOCRRequest, + *, + context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( { - "model": model, - "document": document, - "api_key": api_key, - "api_base": api_base, - "custom_llm_provider": custom_llm_provider, - "extra_headers": extra_headers, - "optional_params": optional_params, - "timeout_seconds": timeout_seconds, + "model": request.model, + "document": request.document, + "api_key": request.options.api_key, + "api_base": request.options.api_base, + "custom_llm_provider": request.options.custom_llm_provider, + "extra_headers": request.options.extra_headers, + "optional_params": {**request.optional_params, **(request.options.provider_connection or {})}, + "timeout_seconds": request.options.timeout_seconds, } ) return dict(FAKE_OCR_RESPONSE) @@ -81,25 +76,20 @@ class RecordingAsyncBridge: async def __call__( self, - model: str, - document: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, + request: NativeOCRRequest, + *, + context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( { - "model": model, - "document": document, - "api_key": api_key, - "api_base": api_base, - "custom_llm_provider": custom_llm_provider, - "extra_headers": extra_headers, - "optional_params": optional_params, - "timeout_seconds": timeout_seconds, + "model": request.model, + "document": request.document, + "api_key": request.options.api_key, + "api_base": request.options.api_base, + "custom_llm_provider": request.options.custom_llm_provider, + "extra_headers": request.options.extra_headers, + "optional_params": {**request.optional_params, **(request.options.provider_connection or {})}, + "timeout_seconds": request.options.timeout_seconds, } ) return dict(FAKE_OCR_RESPONSE) @@ -108,14 +98,9 @@ class RecordingAsyncBridge: class RaisingBridge: def __call__( self, - model: str, - document: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, + request: NativeOCRRequest, + *, + context: NativeRequestContext, ) -> dict[str, object]: raise RuntimeError("bridge failed") @@ -123,14 +108,9 @@ class RaisingBridge: class RaisingAsyncBridge: async def __call__( self, - model: str, - document: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, + request: NativeOCRRequest, + *, + context: NativeRequestContext, ) -> dict[str, object]: raise RuntimeError("bridge failed") @@ -164,9 +144,6 @@ class FakeOCRConfig: def get_api_key_env_var(self) -> str: return self.api_key_env_var - def supports_rust_bridge(self) -> bool: - return True - def validate_environment( self, *, @@ -203,7 +180,7 @@ def build_prepared_request( litellm_params: dict[str, object] | None = None, timeout: float | httpx.Timeout | None = 12.5, ) -> Any: - return rust_bridge.PreparedOCRRequest( + return ocr_main._PreparedOCRRequest( model=model, document=document, api_key=api_key, @@ -392,13 +369,103 @@ def test_timeout_to_seconds_handles_float_timeout_and_none(): assert timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0 +def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): + bridge = RecordingBridge() + + litellm.rust(True) + + rust_bridge.set_rust_ocr(ocr=bridge) + response = rust_bridge.dispatch_ocr( + prepare=lambda: PreparedNativeCall( + request=NativeOCRRequest( + model="mistral-ocr-latest", + document=DOCUMENT, + optional_params={"include_image_base64": True, "pages": [0]}, + options=NativeRequestOptions( + api_key="sk-test", + api_base="https://proxy.internal", + custom_llm_provider="mistral", + extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"}, + timeout_seconds=12.5, + ), + ), + ), + fallback=lambda: pytest.fail("unexpected Python fallback"), + adapt=dict, + model="mistral-ocr-latest", + provider="mistral", + eligible=True, + ) + + assert response == FAKE_OCR_RESPONSE + call = bridge.calls[0] + assert call == { + "model": "mistral-ocr-latest", + "document": DOCUMENT, + "api_key": "sk-test", + "api_base": "https://proxy.internal", + "custom_llm_provider": "mistral", + "extra_headers": { + "Authorization": "Bearer sk-test", + "x-trace-id": "trace-1", + }, + "optional_params": {"include_image_base64": True, "pages": [0]}, + "timeout_seconds": 12.5, + } + + +@pytest.mark.asyncio +async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): + bridge = RecordingAsyncBridge() + + litellm.rust(True) + + rust_bridge.set_rust_ocr(aocr=bridge) + + async def unexpected_fallback(): + pytest.fail("unexpected Python fallback") + + response = await rust_bridge.adispatch_ocr( + prepare=lambda: PreparedNativeCall( + request=NativeOCRRequest( + model="mistral-ocr-maas", + document=DOCUMENT, + optional_params={}, + options=NativeRequestOptions( + custom_llm_provider="vertex_ai", + provider_connection={"vertex_project": "project-1"}, + timeout_seconds=42.0, + ), + ), + ), + fallback=unexpected_fallback, + adapt=dict, + model="mistral-ocr-maas", + provider="vertex_ai", + eligible=True, + ) + + assert response == FAKE_OCR_RESPONSE + assert bridge.calls[0] == { + "model": "mistral-ocr-maas", + "document": DOCUMENT, + "api_key": None, + "api_base": None, + "custom_llm_provider": "vertex_ai", + "extra_headers": None, + "optional_params": {"vertex_project": "project-1"}, + "timeout_seconds": 42.0, + } + + def test_run_rust_ocr_prepares_request_and_wraps_response(): bridge = RecordingBridge() logging_obj = RecordingLogging() litellm.rust(True) rust_bridge._OCR.override(bridge) - response = rust_bridge.attempt_ocr( + response = ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://proxy.internal", @@ -409,8 +476,6 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): resolve_api_key=lambda _name: None, ) - assert isinstance(response, Handled) - response = response.value assert isinstance(response, OCRResponse) assert response.pages[0].markdown == "hello world" assert bridge.calls[0] == { @@ -433,7 +498,8 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): litellm.rust(True) rust_bridge._OCR.override(bridge) - rust_bridge.attempt_ocr( + ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request(api_key=None, timeout=None), resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None, ) @@ -449,7 +515,8 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver(): def _resolver(name: str) -> str | None: raise AssertionError(f"resolver should not be called for {name}") - rust_bridge.attempt_ocr( + ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( api_key="sk-explicit", timeout=None, @@ -470,7 +537,8 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): resolver_calls.append(name) return "sk-provider-env" - rust_bridge.attempt_ocr( + ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"), model="provider-ocr-model", @@ -489,7 +557,8 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): litellm.rust(True) rust_bridge._OCR.override(bridge) - rust_bridge.attempt_ocr( + ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -522,7 +591,8 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana "VERTEXAI_LOCATION": "us-east5", }.get(name) - rust_bridge.attempt_ocr( + ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -540,7 +610,8 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): litellm.rust(True) rust_bridge._OCR.override(bridge) - rust_bridge.attempt_ocr( + ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", @@ -558,7 +629,8 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): litellm.rust(True) rust_bridge._OCR.override(bridge) - rust_bridge.attempt_ocr( + ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="doc-intelligence/prebuilt-layout", @@ -579,7 +651,8 @@ def test_run_rust_ocr_runs_pre_call_logging(): litellm.rust(True) rust_bridge._OCR.override(bridge) - rust_bridge.attempt_ocr( + ocr_main._run_rust_ocr( + fallback=lambda: pytest.fail("unexpected Python fallback"), prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://api.mistral.ai/v1", @@ -726,7 +799,7 @@ async def test_ocr_fallback_skips_native_preparation( def unexpected_preparation(*_args: object, **_kwargs: object) -> None: pytest.fail("Python fallback must not resolve native credentials or emit native pre_call") - monkeypatch.setattr(rust_bridge, "_prepare_rust_ocr_call", unexpected_preparation) + monkeypatch.setattr(ocr_main, "_prepare_rust_ocr_call", unexpected_preparation) monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fallback) response: Final = ( @@ -739,25 +812,6 @@ async def test_ocr_fallback_skips_native_preparation( fallback.assert_called_once() -@pytest.mark.asyncio -async def test_aocr_rejects_empty_python_fallback_response(monkeypatch: pytest.MonkeyPatch) -> None: - captured: dict[str, object] = {} - - def fake_exception_type(**kwargs: object) -> CapturedException: - captured.update(kwargs) - return CapturedException("wrapped") - - monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) - monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", AsyncMock(return_value=None)) - - with pytest.raises(CapturedException, match="wrapped"): - await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") - - original: Final = captured["original_exception"] - assert isinstance(original, ValueError) - assert str(original) == "Got an unexpected None response from the OCR API: None" - - def test_ocr_provider_configs_expose_api_key_env_vars(): from litellm.llms.azure_ai.ocr.document_intelligence.transformation import ( AzureDocumentIntelligenceOCRConfig, diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index 9752f95198f..cbf183dea91 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -4,7 +4,7 @@ import pytest from litellm.llms.custom_httpx.llm_http_handler import _rust_responses_websocket_enabled from litellm.rust_bridge import configuration, responses_websocket -from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason +from litellm.rust_bridge.request import NativeRequestContext, NativeResponsesWebSocketRequest class _FakeNativeConnection: @@ -31,10 +31,9 @@ class _FakeNativeBridge: @classmethod async def connect( cls, + request: NativeResponsesWebSocketRequest, *, - url: str, - headers: dict[str, str], - timeout_seconds: float | None, + context: NativeRequestContext, ) -> _FakeNativeConnection: return _FakeNativeConnection() @@ -58,22 +57,25 @@ def test_rust_websocket_bridge_uses_process_enablement() -> None: @pytest.mark.asyncio async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None: - adapter = responses_websocket.ConnectionAdapter(_ClosedNativeConnection()) + adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection()) with pytest.raises(responses_websocket.ConnectionClosedOK): await adapter.recv() @pytest.mark.asyncio -async def test_bridge_reports_unavailable(monkeypatch: pytest.MonkeyPatch) -> None: +async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None: configuration.rust(True) responses_websocket._RESPONSES_WEBSOCKET.override(None) - assert await responses_websocket.connect( - url="wss://example.test/responses", - headers={}, - timeout=None, - ) == NativeSkipped(NativeSkipReason.UNAVAILABLE) + assert ( + await responses_websocket.connect( + url="wss://example.test/responses", + headers={}, + timeout=None, + ) + is None + ) @pytest.mark.asyncio @@ -89,8 +91,7 @@ async def test_enabled_bridge_connects_and_adapts_socket( timeout=1.0, ) - assert isinstance(connection, Handled) - connection = connection.value + assert connection is not None await connection.send("response.create") assert await connection.recv() == "response.completed" await connection.close() @@ -100,72 +101,17 @@ class _FailingNativeBridge: @classmethod async def connect( cls, + request: NativeResponsesWebSocketRequest, *, - url: str, - headers: dict[str, str], - timeout_seconds: float | None, + context: NativeRequestContext, ) -> _FakeNativeConnection: raise RuntimeError("connection failed") -@pytest.mark.asyncio -async def test_connection_failure_is_reported_to_orchestration() -> None: - configuration.rust(True) - responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge) - result = await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None) - assert isinstance(result, NativeFailed) - assert str(result.error) == "connection failed" - - -@pytest.mark.asyncio -async def test_managed_connection_closes_native_socket_on_consumer_failure() -> None: - configuration.rust(True) - socket = _FakeNativeConnection() - - class Bridge: - @classmethod - async def connect( - cls, *, url: str, headers: dict[str, str], timeout_seconds: float | None - ) -> _FakeNativeConnection: - return socket - - responses_websocket.set_rust_responses_websocket(connection=Bridge) - result = await responses_websocket.managed_connect(url="wss://example.test/responses", headers={}, timeout=1.0) - assert isinstance(result, Handled) - - async def use_connection() -> None: - async with result.value as connection: - await connection.send("hello") - raise ValueError("consumer failed") - - with pytest.raises(ValueError, match="consumer failed"): - await use_connection() - assert socket.sent == ["hello"] - assert socket.closed - - @pytest.mark.asyncio async def test_connection_failure_does_not_authorize_python_fallback() -> None: - from contextlib import AbstractAsyncContextManager - - from litellm.rust_bridge.dispatch import anative_context, provider_errors - configuration.rust(True) responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge) - @anative_context( - native=lambda: responses_websocket.managed_connect( - url="wss://example.test/responses", headers={}, timeout=None - ), - route="responses_websocket", - errors=lambda: provider_errors("openai", "responses websocket"), - ) - def execute() -> AbstractAsyncContextManager[object]: - pytest.fail("unknown native failures must not open a Python connection") - - async def run() -> None: - async with execute(): - pytest.fail("connection must fail before entering its body") - with pytest.raises(RuntimeError, match="connection failed"): - await run() + await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None) 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 f3f2ac651bf..4b07099cc78 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -12,7 +12,6 @@ import pytest import litellm from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge import chat_completions as bridge -from litellm.rust_bridge.runtime import Handled, NativeSkipped, NativeFailed from litellm.types.utils import ModelResponse RUST_RESPONSE = { @@ -96,16 +95,16 @@ class _RecordingCall: self.error = error self.calls: list[dict] = [] - def __call__(self, **kwargs): - self.calls.append(kwargs) + def __call__(self, request, *, context): + self.calls.append({"request": request, "context": context}) if self.error is not None: raise self.error return self.result class _RecordingAsyncCall(_RecordingCall): - async def __call__(self, **kwargs): - return _RecordingCall.__call__(self, **kwargs) + async def __call__(self, request, *, context): + return _RecordingCall.__call__(self, request, context=context) def _accepts(**overrides) -> bool: @@ -251,8 +250,7 @@ class TestSyncCall: result = bridge.chat_completions(**_call_kwargs(model_response)) - assert isinstance(result, Handled) - result = result.value + assert result is not None assert result.choices[0].message.content == "hello from rust" assert result.choices[0].finish_reason == "stop" assert result.model == "claude-sonnet-4-5-20260101" @@ -266,16 +264,16 @@ class TestSyncCall: native = _RecordingCall() bridge.set_rust_chat_completions(chat_completions=native) bridge.chat_completions(**_call_kwargs(ModelResponse())) - assert native.calls[0]["timeout_seconds"] == 30.0 + assert native.calls[0]["request"].options.timeout_seconds == 30.0 - def test_reports_unavailable_bridge(self, monkeypatch): + def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeSkipped) + assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None - def test_reports_native_decline_to_orchestration(self, monkeypatch): + def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) - assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeFailed) + assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None class TestAsyncCall: @@ -283,18 +281,152 @@ class TestAsyncCall: async def test_builds_a_model_response(self): bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) result = await bridge.achat_completions(**_call_kwargs(ModelResponse())) - assert isinstance(result, Handled) - result = result.value + assert result is not None assert result.choices[0].message.content == "hello from rust" assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} @pytest.mark.asyncio - async def test_reports_unavailable_bridge(self, monkeypatch): + async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeSkipped) + assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None @pytest.mark.asyncio - async def test_reports_native_decline_to_orchestration(self, monkeypatch): + async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))) - assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeFailed) + assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None + + +class TestAsyncFallbackWrapper: + @pytest.mark.asyncio + async def test_returns_the_rust_response_without_running_the_fallback(self): + bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) + ran = [] + + async def fallback(): + ran.append(True) + return "python" + + result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) + assert result.choices[0].message.content == "hello from rust" + assert ran == [] + + @pytest.mark.asyncio + async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch): + _fake_native_bridge(monkeypatch) + bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))) + + async def fallback(): + return "python" + + result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) + assert result == "python" + + @pytest.mark.asyncio + async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch): + _hide_native_bridge(monkeypatch) + + async def fallback(): + return "python" + + result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) + assert result == "python" + + +class TestFailureClassification: + """A failure the provider already saw must not be retried on the Python + path: it would bill the customer for the same work twice.""" + + @pytest.fixture(autouse=True) + def _native_exceptions(self, monkeypatch): + _fake_native_bridge(monkeypatch) + + def test_a_decline_falls_back_because_nothing_was_sent(self): + bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) + assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + + def test_an_upstream_failure_is_surfaced_with_its_status(self): + from litellm.exceptions import RateLimitError + + bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited"))) + with pytest.raises(RateLimitError) as raised: + bridge.chat_completions(**_call_kwargs(ModelResponse())) + assert raised.value.status_code == 429 + assert "rate limited" in str(raised.value) + + def test_a_transport_failure_with_no_response_surfaces_as_a_500(self): + from litellm.exceptions import APIError + + bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset"))) + with pytest.raises(APIError) as raised: + bridge.chat_completions(**_call_kwargs(ModelResponse())) + assert raised.value.status_code == 500 + + def test_an_unrecognized_error_is_not_swallowed(self): + bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else"))) + with pytest.raises(RuntimeError): + bridge.chat_completions(**_call_kwargs(ModelResponse())) + + @pytest.mark.asyncio + async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self): + from litellm.exceptions import InternalServerError + + bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom"))) + ran = [] + + async def fallback(): + ran.append(True) + return "python" + + with pytest.raises(InternalServerError): + await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) + assert ran == [], "a request the provider already served must not be re-issued" + + @pytest.mark.asyncio + async def test_the_async_wrapper_falls_back_on_a_decline(self): + bridge.set_rust_chat_completions( + achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text")) + ) + + async def fallback(): + return "python" + + result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) + assert result == "python" + + +@pytest.mark.asyncio +async def test_missing_native_exception_types_does_not_authorize_python_fallback(monkeypatch): + _hide_native_bridge(monkeypatch) + bridge.set_rust_chat_completions( + chat_completions=_RecordingCall(error=RuntimeError("connection failed")), + achat_completions=_RecordingAsyncCall(error=RuntimeError("connection failed")), + ) + + with pytest.raises(RuntimeError, match="connection failed"): + bridge.chat_completions(**_call_kwargs(ModelResponse())) + + async def fallback(): + pytest.fail("unknown failure must not retry through Python") + + with pytest.raises(RuntimeError, match="connection failed"): + await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) + + +def test_provider_credentials_are_separate_from_chat_body_params(): + native = _RecordingCall() + bridge.set_rust_chat_completions(chat_completions=native) + configuration.rust(True) + kwargs = _call_kwargs(ModelResponse()) + kwargs["optional_params"] = { + "max_tokens": 32, + "aws_access_key_id": "test-access-key", + "aws_secret_access_key": "test-secret-key", + } + bridge.chat_completions(**kwargs) + request = native.calls[0]["request"] + assert request.optional_params == {"max_tokens": 32} + assert request.options.provider_connection == { + "aws_access_key_id": "test-access-key", + "aws_secret_access_key": "test-secret-key", + } diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 0cbe0bf5277..72c60463fc5 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -4,7 +4,7 @@ import pytest import litellm from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch -from litellm.rust_bridge.runtime import Handled +from litellm.rust_bridge.request import NativeRequestContext, NativeTranscriptionRequest rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") @@ -15,37 +15,27 @@ class SyncBridge: def __call__( self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, + request: NativeTranscriptionRequest, + *, + context: NativeRequestContext, ) -> dict[str, object]: - self.calls.append({"model": model, "audio": audio, "optional_params": optional_params}) + self.calls.append({"model": request.model, "audio": request.audio, "optional_params": {**request.optional_params, **(request.options.provider_connection or {})}}) return {"text": "hello"} class AsyncBridge: async def __call__( self, - model: str, - audio: dict[str, object], - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: dict[str, object] | None, - optional_params: dict[str, object], - timeout_seconds: float | None, + request: NativeTranscriptionRequest, + *, + context: NativeRequestContext, ) -> dict[str, object]: return {"text": "async"} def test_enabled_sync_bridge_receives_audio() -> None: bridge = SyncBridge() - rust_bridge.configure_rust_transcription(transcription=bridge) + rust_bridge.configure_rust_transcription(True, transcription=bridge) result = rust_bridge.transcription( model="mistral.voxtral-mini-3b-2507", audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"}, @@ -56,14 +46,13 @@ def test_enabled_sync_bridge_receives_audio() -> None: optional_params={"temperature": 0}, timeout=5.0, ) - assert isinstance(result, Handled) - assert result.value == {"text": "hello"} + assert result == {"text": "hello"} assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"} @pytest.mark.asyncio async def test_enabled_async_bridge() -> None: - rust_bridge.configure_rust_transcription(atranscription=AsyncBridge()) + rust_bridge.configure_rust_transcription(True, atranscription=AsyncBridge()) result = await rust_bridge.atranscription( model="mistral.voxtral-mini-3b-2507", audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"}, @@ -74,7 +63,7 @@ async def test_enabled_async_bridge() -> None: optional_params={}, timeout=None, ) - assert result == Handled({"text": "async"}) + assert result == {"text": "async"} def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None: @@ -85,8 +74,7 @@ def test_loader_returns_none_without_native_extension(monkeypatch: pytest.Monkey def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - rust_bridge.configure_rust_transcription(transcription=None) - monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) + monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None) with pytest.raises(RuntimeError, match="bridge is unavailable"): BedrockAudioTranscriptionRustDispatch().audio_transcriptions( @@ -103,8 +91,10 @@ def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> @pytest.mark.asyncio async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - rust_bridge.configure_rust_transcription(atranscription=None) - monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) + async def unavailable(**_: object) -> None: + return None + + monkeypatch.setattr(rust_bridge, "atranscription", unavailable) with pytest.raises(RuntimeError, match="bridge is unavailable"): await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions( @@ -121,7 +111,7 @@ async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPat def test_bedrock_transcription_uses_rust_only_path() -> None: rust_bridge.configure_rust_transcription( - transcription=lambda **_: {"text": "rust"}, + transcription=lambda request, *, context: {"text": "rust"}, atranscription=None, ) try: @@ -137,7 +127,7 @@ def test_bedrock_transcription_uses_rust_only_path() -> None: @pytest.mark.asyncio async def test_bedrock_atranscription_uses_rust_only_path() -> None: - async def rust_response(**_: object) -> dict[str, object]: + async def rust_response(request: NativeTranscriptionRequest, *, context: NativeRequestContext) -> dict[str, object]: return {"text": "rust"} rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response) diff --git a/type-discipline-budget.json b/type-discipline-budget.json index 78b28d2d775..5035894f3f9 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -1,6 +1,6 @@ { "LIT001": { - "limit": 22165 + "limit": 22155 }, "LIT002": { "limit": 26729 @@ -15,7 +15,7 @@ "limit": 0 }, "LIT006": { - "limit": 1022 + "limit": 1027 }, "LIT007": { "limit": 0 @@ -27,10 +27,10 @@ "limit": 0 }, "LIT010": { - "limit": 16419 + "limit": 16426 }, "LIT011": { - "limit": 5497 + "limit": 5506 }, "LIT012": { "limit": 4486 From c03d3f8ce036e2147c8e1406086572680e3b0aca Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 16:59:57 -0700 Subject: [PATCH 2/4] refactor(native): separate request options and context --- litellm-rust/README.md | 14 +- .../src/audio_transcription/hooks.rs | 28 +- .../ai-gateway/src/audio_transcription/mod.rs | 4 +- .../src/audio_transcription/prepare.rs | 27 +- .../src/audio_transcription/tests.rs | 32 +- .../src/audio_transcription/types.rs | 5 +- .../crates/ai-gateway/src/io/responses_ws.rs | 6 +- .../crates/ai-gateway/src/ocr/hooks.rs | 2 +- litellm-rust/crates/ai-gateway/src/ocr/mod.rs | 36 +- .../crates/ai-gateway/src/ocr/prepare.rs | 38 +- .../crates/ai-gateway/src/ocr/types.rs | 5 +- .../crates/ai-gateway/tests/ocr_lifecycle.rs | 124 +++-- .../core/src/audio_transcription/handler.rs | 2 +- .../core/src/audio_transcription/mod.rs | 9 +- .../core/src/audio_transcription/prepare.rs | 52 +- .../core/src/audio_transcription/tests.rs | 31 +- .../core/src/audio_transcription/types.rs | 6 +- .../core/src/chat_completions/handler.rs | 11 +- .../crates/core/src/chat_completions/mod.rs | 5 +- .../core/src/chat_completions/prepare.rs | 26 +- .../crates/core/src/chat_completions/tests.rs | 91 ++-- .../crates/core/src/chat_completions/types.rs | 5 +- .../crates/core/src/messages/handler.rs | 6 +- litellm-rust/crates/core/src/messages/mod.rs | 7 +- .../crates/core/src/messages/prepare.rs | 43 +- .../crates/core/src/messages/tests.rs | 128 ++--- .../crates/core/src/messages/types.rs | 2 - .../crates/core/src/request_context.rs | 15 +- .../crates/core/src/request_options.rs | 71 ++- .../crates/core/src/responses/types.rs | 1 - litellm-rust/crates/python-bridge/src/lib.rs | 27 +- .../crates/python-bridge/src/marshal.rs | 138 ++++- .../src/routes/audio_transcription.rs | 4 +- .../src/routes/chat_completions.rs | 4 +- .../python-bridge/src/routes/definition.rs | 65 ++- .../python-bridge/src/routes/messages.rs | 4 +- .../crates/python-bridge/src/routes/ocr.rs | 4 +- litellm/llms/anthropic/chat/handler.py | 197 +++---- litellm/llms/bedrock/chat/converse_handler.py | 315 ++++++------ litellm/llms/bedrock/request_metadata.py | 6 +- litellm/ocr/main.py | 25 +- litellm/rust_bridge/chat_completions.py | 67 +-- litellm/rust_bridge/messages.py | 28 +- litellm/rust_bridge/protocols.py | 2 + litellm/rust_bridge/request.py | 131 +++-- litellm/rust_bridge/responses_websocket.py | 2 +- litellm/rust_bridge/transcription.py | 39 +- .../strategies/trace_parity/sdk/execution.py | 49 +- .../test_rust_bridge_messages.py | 30 +- .../chat/test_anthropic_chat_handler.py | 485 +++++------------- .../chat/test_bedrock_converse_handler.py | 123 ++--- tests/test_litellm/ocr/test_rust_bridge.py | 77 +-- .../responses/test_rust_bridge_websocket.py | 2 + .../rust_bridge/native_route_wheel_test.py | 103 ++-- .../rust_bridge/test_chat_completions.py | 61 ++- .../test_audio_transcription_rust_bridge.py | 23 +- 56 files changed, 1508 insertions(+), 1335 deletions(-) diff --git a/litellm-rust/README.md b/litellm-rust/README.md index 5720a463455..c76b33ef5a2 100644 --- a/litellm-rust/README.md +++ b/litellm-rust/README.md @@ -27,15 +27,15 @@ 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 +Native HTTP routes and Responses WebSocket connections accept +`native(request, *, options, context)`. The request carries only endpoint payload. +`NativeRequestOptions` carries credentials, typed provider configuration, routing, +headers, query parameters, and timeout. `NativeRequestContext` carries call identity, +attribution, and typed capability facts separately from the provider payload. Python builds the frozen request dataclasses in `litellm/rust_bridge/request.py` and -PyO3 extracts their fields before execution. Provider connection parameters, such as -AWS credentials and Vertex project/location, belong in `options.provider_connection` -rather than the request body +PyO3 extracts their fields before execution. AWS credentials and metadata policy +belong in `options.bedrock`; Vertex project/location belongs in `options.vertex`. This boundary preserves existing Python provider preparation, preflight decisions, fallback, and callbacks 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 ee4a99ca094..3b1201b1666 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs @@ -94,29 +94,29 @@ impl AudioTranscriptionLifecycleHooks { custom_llm_provider, audio, api_key, - provider_connection, + bedrock, api_base, extra_headers, optional_params, timeout, .. } = request; - let provider_request = - prepare_audio_transcription_provider_call(CoreAudioTranscriptionRequest { + let provider_request = prepare_audio_transcription_provider_call( + CoreAudioTranscriptionRequest { model: &model, audio, optional_params, - 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() - }, - })?; + }, + RequestOptions { + bedrock: Some(bedrock), + api_key: (api_key.as_deref()).map(|value| value.to_string()), + api_base: (api_base.as_deref()).map(|value| value.to_string()), + custom_llm_provider: (Some(&custom_llm_provider)).map(|value| value.to_string()), + extra_headers, + timeout, + ..Default::default() + }, + )?; self.run_during_call_guardrails(provider_request).await } 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 de385cb667d..642e2c2b168 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs @@ -3,6 +3,7 @@ use litellm_core::Error; use litellm_core::audio_transcription::execute_audio_transcription_provider_call; use litellm_core::call_lifecycle::CallLifecycle; use litellm_core::request_context::LiteLlmRequestContext; +use litellm_core::request_options::RequestOptions; use serde_json::Value; mod hooks; @@ -15,6 +16,7 @@ use prepare::{PreparedAudioTranscriptionCall, prepare_audio_transcription_call}; pub async fn audio_transcription( request: AudioTranscriptionRequest<'_>, + options: &RequestOptions, context: &LiteLlmRequestContext, hooks: RequestHooks, ) -> Result { @@ -22,7 +24,7 @@ pub async fn audio_transcription( request, context, hooks, - } = prepare_audio_transcription_call(request, context, hooks); + } = prepare_audio_transcription_call(request, options.clone(), context, hooks); CallLifecycle::default() .run( context, 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 cb4884cceb6..f551b632793 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs @@ -1,6 +1,7 @@ use crate::integrations::types::RequestHooks; use litellm_core::call_lifecycle::CallLifecycleContext; use litellm_core::request_context::LiteLlmRequestContext; +use litellm_core::request_options::RequestOptions; use std::sync::atomic::{AtomicU64, Ordering}; use std::time::{SystemTime, UNIX_EPOCH}; @@ -19,6 +20,7 @@ pub(crate) struct PreparedAudioTranscriptionCall { pub(crate) fn prepare_audio_transcription_call( request: AudioTranscriptionRequest<'_>, + options: RequestOptions, context: &LiteLlmRequestContext, hooks: RequestHooks, ) -> PreparedAudioTranscriptionCall { @@ -26,14 +28,13 @@ pub(crate) fn prepare_audio_transcription_call( .litellm_call_id .clone() .unwrap_or_else(new_audio_transcription_call_id); - 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", - }); + let provider_info = + get_custom_llm_provider(request.model, options.custom_llm_provider.as_deref()).unwrap_or( + CustomLlmProvider { + model: request.model, + custom_llm_provider: "bedrock", + }, + ); PreparedAudioTranscriptionCall { context: CallLifecycleContext::new( "audio_transcription", @@ -45,12 +46,12 @@ pub(crate) fn prepare_audio_transcription_call( model: provider_info.model.to_string(), custom_llm_provider: provider_info.custom_llm_provider.to_string(), audio: request.audio, - provider_connection: request.options.provider_connection, - api_key: request.options.api_key, - api_base: request.options.api_base, - extra_headers: request.options.extra_headers, + bedrock: options.bedrock.unwrap_or_default(), + api_key: options.api_key, + api_base: options.api_base, + extra_headers: options.extra_headers, optional_params: request.optional_params, - timeout: request.options.timeout, + timeout: options.timeout, }, hooks: AudioTranscriptionLifecycleHooks::new( CustomLoggerRunner::new(hooks.callbacks), 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 d54399b9f96..ea7c6fb0074 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/tests.rs @@ -1,6 +1,6 @@ use crate::integrations::types::RequestHooks; use litellm_core::request_context::LiteLlmRequestContext; -use litellm_core::request_options::RequestOptions; +use litellm_core::request_options::{BedrockOptions, RequestOptions}; use std::io::{Read, Write}; use std::net::TcpListener; use std::thread; @@ -29,27 +29,27 @@ async fn bedrock_request_is_signed_and_contains_audio() { stream.write_all(response).expect("response"); }); - 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 bedrock = BedrockOptions { + aws_access_key_id: Some("access-key".to_string()), + aws_secret_access_key: Some("secret-key".to_string()), + aws_region_name: Some("us-east-1".to_string()), + ..Default::default() + }; let api_base = format!("http://{address}"); let response = audio_transcription( AudioTranscriptionRequest { model: "mistral.voxtral-mini-3b-2507", audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}), 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() - }, + }, + &RequestOptions { + bedrock: Some(bedrock), + api_key: None, + api_base: (Some(&api_base)).map(|value| value.to_string()), + custom_llm_provider: (Some("bedrock")).map(|value| value.to_string()), + extra_headers: None, + timeout: None, + ..Default::default() }, &LiteLlmRequestContext { attribution: Default::default(), 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 8656b839c09..aee24cf4bb0 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs @@ -1,4 +1,4 @@ -use litellm_core::request_options::RequestOptions; +use litellm_core::request_options::BedrockOptions; use std::time::Duration; use serde_json::{Map, Value}; @@ -7,14 +7,13 @@ pub struct AudioTranscriptionRequest<'a> { pub model: &'a str, pub audio: Value, pub optional_params: Map, - pub options: RequestOptions, } pub(crate) struct PreparedAudioTranscriptionRequest { pub(crate) model: String, pub(crate) custom_llm_provider: String, pub(crate) audio: Value, - pub(crate) provider_connection: Map, + pub(crate) bedrock: BedrockOptions, pub(crate) api_key: Option, pub(crate) api_base: Option, pub(crate) extra_headers: Option>, 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 9d22ec76e5e..8409cdcbaec 100644 --- a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs +++ b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs @@ -7,6 +7,7 @@ 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::request_context::LiteLlmRequestContext; +use litellm_core::request_options::RequestOptions; use litellm_core::responses::types::{ResponsesWebSocketRequest, ResponsesWsEvent}; use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig; use tokio::net::TcpStream; @@ -36,9 +37,10 @@ pub struct ResponsesWebSocketConnection { impl ResponsesWebSocketConnection { pub async fn connect( input: ResponsesWebSocketRequest, + options: &RequestOptions, _context: &LiteLlmRequestContext, ) -> Result { - let headers = string_headers("Responses WebSocket", input.options.extra_headers)?; + let headers = string_headers("Responses WebSocket", options.extra_headers.clone())?; let mut request = input .url .as_str() @@ -53,7 +55,7 @@ impl ResponsesWebSocketConnection { request.headers_mut().insert(header_name, header_value); } let connect = connect_async(request); - let result = match input.options.timeout { + let result = match options.timeout { Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| { Error::Network("Responses WebSocket connection timed out".to_string()) })?, diff --git a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs index 0a018b19a0e..ff9b47b4ccc 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs @@ -88,7 +88,7 @@ impl OcrLifecycleHooks { .optional_params .clone() .into_iter() - .chain(request.provider_connection) + .chain(request.vertex.into_map()) .collect(); let url = config.complete_url( request.api_base.as_deref(), diff --git a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs index 520fd88c8a5..328b42be352 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs @@ -2,6 +2,7 @@ use crate::integrations::types::RequestHooks; use litellm_core::Error; use litellm_core::call_lifecycle::CallLifecycle; use litellm_core::request_context::LiteLlmRequestContext; +use litellm_core::request_options::RequestOptions; use serde_json::Value; mod common_utils; @@ -18,6 +19,7 @@ use prepare::{PreparedOcrCall, prepare_ocr_call}; #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn ocr( request: OcrRequest<'_>, + options: &RequestOptions, context: &LiteLlmRequestContext, hooks: RequestHooks, ) -> Result { @@ -25,7 +27,7 @@ pub async fn ocr( request, context, hooks, - } = prepare_ocr_call(request, context, hooks); + } = prepare_ocr_call(request, options.clone(), context, hooks); CallLifecycle::default() .run(context, request, &hooks, |request| { execute_ocr_provider_call(request, &hooks) @@ -77,16 +79,17 @@ mod tests { String::from_utf8(request).expect("request is utf8") } - fn base_ocr_request(model: &str) -> OcrRequest<'_> { - OcrRequest { - model, - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - optional_params: Map::new(), - - options: RequestOptions { + fn base_ocr_request(model: &str) -> (OcrRequest<'_>, RequestOptions) { + ( + OcrRequest { + model, + document: json!({ + "type": "document_url", + "document_url": "https://example.com/doc.pdf" + }), + optional_params: Map::new(), + }, + RequestOptions { api_key: (Some("sk-test")).map(|value| value.to_string()), api_base: None, custom_llm_provider: None, @@ -94,7 +97,7 @@ mod tests { timeout: None, ..Default::default() }, - } + ) } #[tokio::test] @@ -132,10 +135,10 @@ mod tests { (upload_request, parse_request) }); let api_base = format!("http://{address}"); - let mut request = base_ocr_request("reducto/parse-v3"); - request.options.api_base = Some(&api_base).map(|value| value.to_string()); - request.options.api_key = None; - request.options.extra_headers = Some(Map::from_iter([ + let (mut request, mut options) = base_ocr_request("reducto/parse-v3"); + options.api_base = Some(&api_base).map(|value| value.to_string()); + options.api_key = None; + options.extra_headers = Some(Map::from_iter([ ("Authorization".to_string(), json!("Bearer test-key")), ("x-trace-id".to_string(), json!("trace-1")), ])); @@ -154,6 +157,7 @@ mod tests { let response = ocr( request, + &options, &LiteLlmRequestContext { ..Default::default() }, diff --git a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs b/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs index cd0d2899bc7..1fe884489ad 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs @@ -1,6 +1,7 @@ use crate::integrations::types::RequestHooks; use litellm_core::call_lifecycle::CallLifecycleContext; use litellm_core::request_context::LiteLlmRequestContext; +use litellm_core::request_options::RequestOptions; use std::sync::atomic::{AtomicU64, Ordering}; use std::time::{SystemTime, UNIX_EPOCH}; @@ -22,6 +23,7 @@ pub(crate) struct PreparedOcrCall { #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub(crate) fn prepare_ocr_call( request: OcrRequest<'_>, + options: RequestOptions, context: &LiteLlmRequestContext, hooks: RequestHooks, ) -> PreparedOcrCall { @@ -29,14 +31,13 @@ pub(crate) fn prepare_ocr_call( .litellm_call_id .clone() .unwrap_or_else(new_ocr_call_id); - 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 provider_info = + get_custom_llm_provider(request.model, options.custom_llm_provider.as_deref()).unwrap_or( + CustomLlmProvider { + model: request.model, + custom_llm_provider: "mistral", + }, + ); let model = provider_info.model.to_string(); let custom_llm_provider = provider_info.custom_llm_provider.to_string(); let config = ocr_provider_config(&custom_llm_provider, &model) @@ -72,12 +73,12 @@ pub(crate) fn prepare_ocr_call( model, custom_llm_provider, document: request.document, - provider_connection: request.options.provider_connection, - api_key: request.options.api_key, - api_base: request.options.api_base, - extra_headers: request.options.extra_headers, + vertex: options.vertex.unwrap_or_default(), + api_key: options.api_key, + api_base: options.api_base, + extra_headers: options.extra_headers, optional_params, - timeout: request.options.timeout, + timeout: options.timeout, }, hooks: OcrLifecycleHooks::new( CustomLoggerRunner::new(hooks.callbacks), @@ -135,15 +136,6 @@ mod tests { "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: None, - custom_llm_provider: None, - extra_headers: None, - timeout: None, - ..Default::default() - }, } } @@ -157,6 +149,7 @@ mod tests { fn native_format_rejected_for_provider_without_support_as_bad_request() { let prepared = prepare_ocr_call( request_with_format("native"), + RequestOptions::default(), &LiteLlmRequestContext { ..Default::default() }, @@ -173,6 +166,7 @@ mod tests { fn unknown_format_rejected_for_provider_without_support_as_bad_request() { let prepared = prepare_ocr_call( request_with_format("raw"), + RequestOptions::default(), &LiteLlmRequestContext { ..Default::default() }, diff --git a/litellm-rust/crates/ai-gateway/src/ocr/types.rs b/litellm-rust/crates/ai-gateway/src/ocr/types.rs index d6be6ba0d78..9be628a2c81 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/types.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/types.rs @@ -1,14 +1,13 @@ -use litellm_core::request_options::RequestOptions; use std::time::Duration; use litellm_core::ocr::transformation::OcrProviderConfig; +use litellm_core::request_options::VertexOptions; use serde_json::{Map, Value}; pub struct OcrRequest<'a> { pub model: &'a str, pub document: Value, pub optional_params: Map, - pub options: RequestOptions, } pub(crate) struct PreparedOcrRequest { @@ -16,7 +15,7 @@ pub(crate) struct PreparedOcrRequest { pub(crate) model: String, pub(crate) custom_llm_provider: String, pub(crate) document: Value, - pub(crate) provider_connection: Map, + pub(crate) vertex: VertexOptions, pub(crate) api_key: Option, pub(crate) api_base: Option, pub(crate) extra_headers: Option>, diff --git a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs b/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs index abbd1424525..f68ac32333a 100644 --- a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs +++ b/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs @@ -216,24 +216,21 @@ impl CustomGuardrail for RecordingOcrGuardrail { } } -fn base_ocr_request(model: &str) -> OcrRequest<'_> { - OcrRequest { - model, - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - optional_params: Map::new(), - - options: RequestOptions { - api_key: (Some("sk-test")).map(|value| value.to_string()), - api_base: None, - custom_llm_provider: None, - extra_headers: None, - timeout: None, +fn base_ocr_request(model: &str) -> (OcrRequest<'_>, RequestOptions) { + ( + OcrRequest { + model, + document: json!({ + "type": "document_url", + "document_url": "https://example.com/doc.pdf" + }), + optional_params: Map::new(), + }, + RequestOptions { + api_key: Some("sk-test".to_string()), ..Default::default() }, - } + ) } #[tokio::test] @@ -244,8 +241,8 @@ async fn reducto_during_call_guardrail_blocks_before_upload() { let address = listener.local_addr().expect("listener has local address"); let api_base = format!("http://{address}"); let guardrail = Arc::new(RecordingOcrGuardrail::blocking_during_call()); - let mut request = base_ocr_request("reducto/parse-v3"); - request.options.api_base = Some(&api_base).map(|value| value.to_string()); + let (mut request, mut options) = base_ocr_request("reducto/parse-v3"); + options.api_base = Some(&api_base).map(|value| value.to_string()); request.document = json!({ "type": "document_url", "document_url": "data:application/pdf;base64,JVBERi0xLjQ=" @@ -257,6 +254,7 @@ async fn reducto_during_call_guardrail_blocks_before_upload() { let error = ocr( request, + &options, &LiteLlmRequestContext { ..Default::default() }, @@ -292,8 +290,8 @@ async fn reducto_upload_error_body_is_truncated() { .expect("writes upload response"); }); let api_base = format!("http://{address}"); - let mut request = base_ocr_request("reducto/parse-v3"); - request.options.api_base = Some(&api_base).map(|value| value.to_string()); + let (mut request, mut options) = base_ocr_request("reducto/parse-v3"); + options.api_base = Some(&api_base).map(|value| value.to_string()); request.document = json!({ "type": "document_url", "document_url": "data:application/pdf;base64,JVBERi0xLjQ=" @@ -301,6 +299,7 @@ async fn reducto_upload_error_body_is_truncated() { let error = ocr( request, + &options, &LiteLlmRequestContext { ..Default::default() }, @@ -351,15 +350,14 @@ async fn ocr_lifecycle_runs_pre_during_and_success_hooks() { "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() - }, + }, + &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 { @@ -430,15 +428,14 @@ async fn ocr_lifecycle_runs_failure_hook_on_provider_error() { "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() - }, + }, + &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(), @@ -485,15 +482,14 @@ async fn ocr_lifecycle_pre_call_block_skips_provider_socket() { "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() - }, + }, + &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(), @@ -566,15 +562,14 @@ async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() { "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() - }, + }, + &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(), @@ -646,15 +641,14 @@ async fn document_intelligence_poll_uses_resolved_subscription_key() { "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() - }, + }, + &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(), diff --git a/litellm-rust/crates/core/src/audio_transcription/handler.rs b/litellm-rust/crates/core/src/audio_transcription/handler.rs index e9d91d7cd3b..ca29c40d4f5 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.provider_connection, &env_lookup), + aws_auth_config(&request.bedrock.into_map(), &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 c7331a6ada0..798be28803a 100644 --- a/litellm-rust/crates/core/src/audio_transcription/mod.rs +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -1,5 +1,6 @@ use crate::Error; use crate::request_context::LiteLlmRequestContext; +use crate::request_options::RequestOptions; mod client; mod handler; mod prepare; @@ -15,10 +16,14 @@ pub use types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest}; #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn audio_transcription( request: AudioTranscriptionRequest<'_>, + options: &RequestOptions, _context: &LiteLlmRequestContext, ) -> Result { - execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call(request)?) - .await + execute_audio_transcription_provider_call(prepare_audio_transcription_provider_call( + request, + options.clone(), + )?) + .await } #[cfg(test)] diff --git a/litellm-rust/crates/core/src/audio_transcription/prepare.rs b/litellm-rust/crates/core/src/audio_transcription/prepare.rs index 28d975bd178..3c382315864 100644 --- a/litellm-rust/crates/core/src/audio_transcription/prepare.rs +++ b/litellm-rust/crates/core/src/audio_transcription/prepare.rs @@ -2,6 +2,7 @@ use crate::error::Error; use crate::http_utils::{has_header, string_headers}; #[cfg(feature = "bedrock-auth")] use crate::providers::bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG; +use crate::request_options::RequestOptions; use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderConfig}; @@ -20,35 +21,36 @@ fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProv #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub fn prepare_audio_transcription_provider_call( request: AudioTranscriptionRequest<'_>, + options: RequestOptions, ) -> Result { - 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, + let provider_info = + get_custom_llm_provider(request.model, options.custom_llm_provider.as_deref()) + .or_else(|| { + options + .custom_llm_provider + .as_deref() + .map(|provider| CustomLlmProvider { + model: request.model, + custom_llm_provider: provider, + }) }) - }) - .ok_or_else(|| { - Error::InvalidProvider( - "unable to resolve custom_llm_provider for audio transcription request".to_string(), - ) - })?; + .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.options.extra_headers)?; - let auth = config.auth_strategy(&model, &request.options.provider_connection, &env_lookup)?; + let mut headers = string_headers("audio transcription", options.extra_headers)?; + let bedrock = options.bedrock.unwrap_or_default(); + let bedrock_options = bedrock.clone().into_map(); + let auth = config.auth_strategy(&model, &bedrock_options, &env_lookup)?; if matches!(auth, AudioTranscriptionAuth::Bearer) && !has_header(&headers, "authorization") - && let Some(api_key) = request.options.api_key.as_deref() + && let Some(api_key) = options.api_key.as_deref() { headers.push(("Authorization".to_string(), format!("Bearer {api_key}"))); } @@ -56,9 +58,9 @@ pub fn prepare_audio_transcription_provider_call( headers.push(("Content-Type".to_string(), "application/json".to_string())); } let url = config.complete_url( - request.options.api_base.as_deref(), + options.api_base.as_deref(), &model, - &request.options.provider_connection, + &bedrock_options, &env_lookup, )?; let filtered_params = config.map_transcription_params(&request.optional_params); @@ -73,7 +75,7 @@ pub fn prepare_audio_transcription_provider_call( upstream_headers: headers, auth, #[cfg(feature = "bedrock-auth")] - provider_connection: request.options.provider_connection, - timeout: request.options.timeout, + bedrock, + timeout: 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 224141776c2..98cb93da5df 100644 --- a/litellm-rust/crates/core/src/audio_transcription/tests.rs +++ b/litellm-rust/crates/core/src/audio_transcription/tests.rs @@ -1,5 +1,5 @@ use crate::request_context::LiteLlmRequestContext; -use crate::request_options::RequestOptions; +use crate::request_options::{BedrockOptions, RequestOptions}; use std::io::{Read, Write}; use std::net::TcpListener; use std::thread; @@ -29,26 +29,27 @@ async fn bedrock_request_is_signed_and_contains_audio() { stream.write_all(response).expect("response"); }); - 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 bedrock = BedrockOptions { + aws_access_key_id: Some("access-key".to_string()), + aws_secret_access_key: Some("secret-key".to_string()), + aws_region_name: Some("us-east-1".to_string()), + ..Default::default() + }; let api_base = format!("http://{address}"); let response = audio_transcription( AudioTranscriptionRequest { model: "mistral.voxtral-mini-3b-2507", audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}), 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() - }, + }, + &RequestOptions { + bedrock: Some(bedrock), + api_key: None, + api_base: (Some(&api_base)).map(|value| value.to_string()), + custom_llm_provider: (Some("bedrock")).map(|value| value.to_string()), + extra_headers: None, + timeout: None, + ..Default::default() }, &LiteLlmRequestContext { ..Default::default() diff --git a/litellm-rust/crates/core/src/audio_transcription/types.rs b/litellm-rust/crates/core/src/audio_transcription/types.rs index f988202cd14..5ecec7df642 100644 --- a/litellm-rust/crates/core/src/audio_transcription/types.rs +++ b/litellm-rust/crates/core/src/audio_transcription/types.rs @@ -1,16 +1,16 @@ -use crate::request_options::RequestOptions; use std::time::Duration; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use crate::request_options::BedrockOptions; + use super::transformation::{AudioTranscriptionAuth, AudioTranscriptionProviderConfig}; pub struct AudioTranscriptionRequest<'a> { pub model: &'a str, pub audio: Value, pub optional_params: Map, - pub options: RequestOptions, } #[derive(Clone)] @@ -23,7 +23,7 @@ pub struct ProviderAudioTranscriptionRequest { pub(super) upstream_headers: Vec<(String, String)>, pub(super) auth: AudioTranscriptionAuth, #[cfg(feature = "bedrock-auth")] - pub(super) provider_connection: Map, + pub(super) bedrock: BedrockOptions, 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 efa7351abcb..ee35d26443d 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -101,15 +101,10 @@ 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.provider_connection) { + let bedrock = request.bedrock.into_map(); + let credentials = match host_supplied_credentials(&bedrock) { Some(credentials) => credentials, - None => { - resolve_credentials( - aws_auth_config(&request.provider_connection, &env_lookup), - &env_lookup, - ) - .await? - } + None => resolve_credentials(aws_auth_config(&bedrock, &env_lookup), &env_lookup).await?, }; let signature = sign_bedrock_post( &request.url, diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs index 8cb9ffe3246..b6139a57b70 100644 --- a/litellm-rust/crates/core/src/chat_completions/mod.rs +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -8,6 +8,7 @@ use crate::Error; use crate::request_context::LiteLlmRequestContext; +use crate::request_options::RequestOptions; mod client; mod common_utils; pub mod conversation; @@ -26,9 +27,11 @@ use types::{ChatCompletionsRequest, ChatCompletionsResponse}; #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn chat_completions( request: ChatCompletionsRequest<'_>, + options: &RequestOptions, context: &LiteLlmRequestContext, ) -> Result { - execute_chat_completions_provider_call(resolve_request(request, context)?).await + execute_chat_completions_provider_call(resolve_request(request, options.clone(), 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 f4e74b8563b..d22a66ed6c0 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -1,4 +1,5 @@ use crate::request_context::LiteLlmRequestContext; +use crate::request_options::RequestOptions; use serde_json::Value; use crate::error::Error; @@ -40,12 +41,11 @@ pub(super) fn parse_messages(messages: Value) -> Result, Error> pub(super) fn resolve_request( request: ChatCompletionsRequest<'_>, + options: RequestOptions, _context: &LiteLlmRequestContext, ) -> Result { - let (model, config) = resolve_provider_config( - request.model, - request.options.custom_llm_provider.as_deref(), - ) + let (model, config) = + resolve_provider_config(request.model, options.custom_llm_provider.as_deref()) .map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?; let messages = parse_messages(request.messages).map_err(|_| Error::Declined("unreadable message list"))?; @@ -60,7 +60,7 @@ pub(super) fn resolve_request( config, messages, optional_params: request.optional_params, - options: request.options, + options, }) } @@ -75,7 +75,12 @@ fn validate_environment( let auth = config.auth( request.options.api_key.as_deref(), model, - &request.options.provider_connection, + &request + .options + .bedrock + .clone() + .unwrap_or_default() + .into_map(), &env_lookup, )?; match &auth { @@ -128,7 +133,12 @@ pub(super) fn prepare_provider_request( let url = config.complete_url( request.options.api_base.as_deref(), &model, - &request.options.provider_connection, + &request + .options + .bedrock + .clone() + .unwrap_or_default() + .into_map(), &env_lookup, )?; let transformed = @@ -141,7 +151,7 @@ pub(super) fn prepare_provider_request( body: transformed.body, upstream_headers: headers, auth, - provider_connection: request.options.provider_connection, + bedrock: request.options.bedrock.unwrap_or_default(), 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 d25d45b9de0..a02904db716 100644 --- a/litellm-rust/crates/core/src/chat_completions/tests.rs +++ b/litellm-rust/crates/core/src/chat_completions/tests.rs @@ -1,5 +1,5 @@ use crate::request_context::LiteLlmRequestContext; -use crate::request_options::RequestOptions; +use crate::request_options::{BedrockOptions, RequestOptions}; use serde_json::{Map, Value, json}; use crate::error::Error; @@ -8,11 +8,17 @@ use super::prepare::{prepare_provider_request, resolve_request}; use super::transformation::ChatCompletionsAuth; use super::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest}; +struct TestChatCompletionsCall<'a> { + request: ChatCompletionsRequest<'a>, + options: RequestOptions, +} + fn prepare_chat_completions_call( - request: ChatCompletionsRequest<'_>, + call: TestChatCompletionsCall<'_>, ) -> Result { prepare_provider_request(resolve_request( - request, + call.request, + call.options, &LiteLlmRequestContext { ..Default::default() }, @@ -24,15 +30,16 @@ fn request<'a>( provider: Option<&'a str>, messages: Value, optional_params: Value, -) -> ChatCompletionsRequest<'a> { - ChatCompletionsRequest { - model, - messages, - optional_params: match optional_params { - Value::Object(map) => map, - other => panic!("params must be an object, got {other}"), +) -> TestChatCompletionsCall<'a> { + TestChatCompletionsCall { + request: ChatCompletionsRequest { + model, + messages, + optional_params: match optional_params { + Value::Object(map) => map, + other => panic!("params must be an object, got {other}"), + }, }, - options: RequestOptions { api_key: (Some("sk-test")).map(|value| value.to_string()), api_base: None, @@ -46,7 +53,7 @@ fn request<'a>( /// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers /// carry resolved credentials), so unwrap the failure case by hand. -fn decline(request: ChatCompletionsRequest<'_>) -> Error { +fn decline(request: TestChatCompletionsCall<'_>) -> Error { match prepare_chat_completions_call(request) { Err(error) => error, Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url), @@ -325,13 +332,11 @@ async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() { json!([{"role": "user", "content": "hi"}]), 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"), - ), - ]); + call.options.bedrock = Some(BedrockOptions { + aws_access_key_id: Some("AKIDEXAMPLE".to_string()), + aws_secret_access_key: Some("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".to_string()), + ..Default::default() + }); // A key would resolve to a bearer token and never reach the signer. call.options.api_key = None; call.options.extra_headers = Some(Map::from_iter([( @@ -383,13 +388,11 @@ async fn a_forwarded_header_the_signer_computes_declines_to_python() { json!([{"role": "user", "content": "hi"}]), 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"), - ), - ]); + call.options.bedrock = Some(BedrockOptions { + aws_access_key_id: Some("AKIDEXAMPLE".to_string()), + aws_secret_access_key: Some("wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY".to_string()), + ..Default::default() + }); call.options.api_key = None; call.options.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])); @@ -613,7 +616,7 @@ mod round_trip { use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; - use crate::chat_completions::chat_completions; + use crate::chat_completions::chat_completions as run_chat_completions; async fn read_http_request(socket: &mut TcpStream) -> String { let mut request = Vec::new(); @@ -676,15 +679,16 @@ mod round_trip { (format!("http://127.0.0.1:{port}/v1/messages"), handle) } - fn call(api_base: &str, messages: Value, params: Value) -> ChatCompletionsRequest<'_> { - ChatCompletionsRequest { - model: "anthropic/claude-sonnet-4-5", - messages, - optional_params: match params { - Value::Object(map) => map, - other => panic!("params must be an object, got {other}"), + fn call(api_base: &str, messages: Value, params: Value) -> TestChatCompletionsCall<'_> { + TestChatCompletionsCall { + request: ChatCompletionsRequest { + model: "anthropic/claude-sonnet-4-5", + messages, + optional_params: match params { + Value::Object(map) => map, + other => panic!("params must be an object, got {other}"), + }, }, - options: RequestOptions { api_key: (Some("sk-test")).map(|value| value.to_string()), api_base: (Some(api_base)).map(|value| value.to_string()), @@ -696,12 +700,19 @@ mod round_trip { } } + async fn execute( + call: TestChatCompletionsCall<'_>, + context: &LiteLlmRequestContext, + ) -> Result { + run_chat_completions(call.request, &call.options, context).await + } + const GOOD_BODY: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#; #[tokio::test] async fn round_trip_sends_the_translated_body_and_normalizes_the_response() { let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await; - let response = chat_completions( + let response = execute( call( &api_base, json!([ @@ -751,7 +762,7 @@ 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( + let err = execute( call( &api_base, json!([{"role": "user", "content": "hi"}]), @@ -774,7 +785,7 @@ 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( + let err = execute( call( &api_base, json!([{"role": "user", "content": "hi"}]), @@ -797,7 +808,7 @@ 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( + let err = execute( call( &api_base, json!([{"role": "user", "content": "hi"}]), @@ -823,7 +834,7 @@ 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( + let err = execute( call( &format!("http://127.0.0.1:{port}/v1/messages"), json!([{"role": "user", "content": "hi"}]), diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs index b096df4664d..e9c459b6108 100644 --- a/litellm-rust/crates/core/src/chat_completions/types.rs +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -1,4 +1,4 @@ -use crate::request_options::RequestOptions; +use crate::request_options::{BedrockOptions, RequestOptions}; use std::time::Duration; use serde::{Deserialize, Serialize}; @@ -16,7 +16,6 @@ pub struct ChatCompletionsRequest<'a> { pub model: &'a str, pub messages: Value, pub optional_params: Map, - pub options: RequestOptions, } pub(super) struct ResolvedChatCompletionsRequest { @@ -35,7 +34,7 @@ pub(super) struct ProviderChatCompletionsRequest { pub(super) upstream_headers: Vec<(String, String)>, pub(super) auth: ChatCompletionsAuth, #[cfg_attr(not(feature = "bedrock-auth"), allow(dead_code))] - pub(super) provider_connection: Map, + pub(super) bedrock: BedrockOptions, pub(super) timeout: Option, } diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index 61ff81bcdc8..1897d550b39 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -10,8 +10,9 @@ use super::types::{AnthropicMessagesResponse, MessagesRequest}; #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub(super) async fn execute_messages_provider_call( request: MessagesRequest<'_>, + options: crate::request_options::RequestOptions, ) -> Result { - let request = prepare_provider_request(request)?; + let request = prepare_provider_request(request, options)?; let mut request_builder = http_client().post(&request.url).json(&request.body); for (key, value) in &request.upstream_headers { request_builder = request_builder.header(key, value); @@ -44,8 +45,9 @@ pub(super) async fn execute_messages_provider_call( pub(super) async fn execute_messages_provider_stream( request: MessagesRequest<'_>, + options: crate::request_options::RequestOptions, ) -> Result { - let request = prepare_provider_request(request)?; + let request = prepare_provider_request(request, options)?; if request.provider != ANTHROPIC_MESSAGES_PROVIDER { return Err(Error::InvalidRequest( "streaming messages is not supported for this provider".to_string(), diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index 10bcb7ab451..61360bb17be 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -9,6 +9,7 @@ use crate::Error; use crate::request_context::LiteLlmRequestContext; +use crate::request_options::RequestOptions; mod client; mod common_utils; mod handler; @@ -22,16 +23,18 @@ use types::{AnthropicMessagesResponse, MessagesRequest}; #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn messages( request: MessagesRequest<'_>, + options: &RequestOptions, _context: &LiteLlmRequestContext, ) -> Result { - execute_messages_provider_call(request).await + execute_messages_provider_call(request, options.clone()).await } pub async fn messages_stream( request: MessagesRequest<'_>, + options: &RequestOptions, _context: &LiteLlmRequestContext, ) -> Result { - execute_messages_provider_stream(request).await + execute_messages_provider_stream(request, options.clone()).await } #[cfg(test)] diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index 27ea0f14f87..2a75dfc30cf 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -1,4 +1,5 @@ use crate::error::Error; +use crate::request_options::RequestOptions; use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers}; @@ -8,26 +9,24 @@ use serde_json::{Map, Value}; pub(super) fn prepare_provider_request( request: MessagesRequest<'_>, + options: RequestOptions, ) -> Result { - 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, + let provider_info = + get_custom_llm_provider(request.model, options.custom_llm_provider.as_deref()) + .or_else(|| { + options + .custom_llm_provider + .as_deref() + .map(|provider| CustomLlmProvider { + model: request.model, + custom_llm_provider: provider, + }) }) - }) - .ok_or_else(|| { - Error::InvalidProvider( - "unable to resolve custom_llm_provider for messages request".to_string(), - ) - })?; + .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; @@ -37,8 +36,8 @@ pub(super) fn prepare_provider_request( let headers = validate_environment( config, - request.options.extra_headers, - request.options.api_key.as_deref(), + options.extra_headers, + options.api_key.as_deref(), &env_lookup, )?; @@ -52,7 +51,7 @@ pub(super) fn prepare_provider_request( )) })?; - let url = config.complete_url(request.options.api_base.as_deref(), &model, &env_lookup)?; + let url = config.complete_url(options.api_base.as_deref(), &model, &env_lookup)?; Ok(ProviderMessagesRequest { provider: provider.to_string(), @@ -61,7 +60,7 @@ pub(super) fn prepare_provider_request( url, body, upstream_headers: headers, - timeout: request.options.timeout, + timeout: options.timeout, }) } diff --git a/litellm-rust/crates/core/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs index fd7d8f8c453..9c74bbb7034 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -148,14 +148,14 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through() }] }] }), - 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() - }, + }, + &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() @@ -212,14 +212,14 @@ async fn messages_round_trip_builds_native_anthropic_request() { "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() - }, + }, + &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() @@ -273,14 +273,14 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() { 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() - }, + }, + &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() @@ -335,14 +335,14 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() { 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() - }, + }, + &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() @@ -367,14 +367,14 @@ async fn messages_requires_auth_when_no_key_and_no_header() { 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() - }, + }, + &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() @@ -413,14 +413,14 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() { 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() - }, + }, + &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() @@ -462,14 +462,14 @@ async fn messages_maps_provider_error_status_to_http_error() { 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() - }, + }, + &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() @@ -487,14 +487,14 @@ async fn messages_rejects_unsupported_provider() { 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() - }, + }, + &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() diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index 7ad0a36cc3e..608e3f77fff 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -1,4 +1,3 @@ -use crate::request_options::RequestOptions; use std::time::Duration; use serde::{Deserialize, Serialize}; @@ -9,7 +8,6 @@ use super::transformation::AnthropicMessagesProviderConfig; pub struct MessagesRequest<'a> { pub model: &'a str, pub body: Value, - 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 index ea8e6096a4f..2796edf9993 100644 --- a/litellm-rust/crates/core/src/request_context.rs +++ b/litellm-rust/crates/core/src/request_context.rs @@ -1,5 +1,3 @@ -use serde_json::{Map, Value}; - #[derive(Clone, Debug, Default, PartialEq)] pub struct RequestAttribution { pub user_api_key_hash: Option, @@ -7,12 +5,19 @@ pub struct RequestAttribution { pub user_api_key_team_id: Option, } +#[derive(Clone, Debug, Default, PartialEq)] +pub struct RequestCapabilities { + pub stream: bool, + pub has_agentic_hook: bool, + pub has_custom_client: bool, + pub request_format: 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 trace_id: Option, pub request_model: Option, pub attribution: RequestAttribution, + pub capabilities: RequestCapabilities, } diff --git a/litellm-rust/crates/core/src/request_options.rs b/litellm-rust/crates/core/src/request_options.rs index 9ab0a81c998..66c7ed89366 100644 --- a/litellm-rust/crates/core/src/request_options.rs +++ b/litellm-rust/crates/core/src/request_options.rs @@ -2,6 +2,73 @@ use std::time::Duration; use serde_json::{Map, Value}; +#[derive(Clone, Debug, Default)] +pub struct BedrockOptions { + pub aws_access_key_id: Option, + pub aws_secret_access_key: Option, + pub aws_session_token: Option, + pub aws_region_name: Option, + pub aws_session_name: Option, + pub aws_profile_name: Option, + pub aws_role_name: Option, + pub aws_web_identity_token: Option, + pub aws_sts_endpoint: Option, + pub aws_external_id: Option, + pub aws_bedrock_runtime_endpoint: Option, + pub request_metadata_fields: Vec, + pub request_metadata: Option>, +} + +impl BedrockOptions { + pub fn into_map(&self) -> Map { + [ + ("aws_access_key_id", self.aws_access_key_id.clone()), + ("aws_secret_access_key", self.aws_secret_access_key.clone()), + ("aws_session_token", self.aws_session_token.clone()), + ("aws_region_name", self.aws_region_name.clone()), + ("aws_session_name", self.aws_session_name.clone()), + ("aws_profile_name", self.aws_profile_name.clone()), + ("aws_role_name", self.aws_role_name.clone()), + ( + "aws_web_identity_token", + self.aws_web_identity_token.clone(), + ), + ("aws_sts_endpoint", self.aws_sts_endpoint.clone()), + ("aws_external_id", self.aws_external_id.clone()), + ( + "aws_bedrock_runtime_endpoint", + self.aws_bedrock_runtime_endpoint.clone(), + ), + ] + .into_iter() + .filter_map(|(name, value)| value.map(|value| (name.to_string(), Value::String(value)))) + .collect() + } +} + +#[derive(Clone, Debug, Default)] +pub struct AnthropicOptions { + pub user_id: Option, +} + +#[derive(Clone, Debug, Default)] +pub struct VertexOptions { + pub project: Option, + pub location: Option, +} + +impl VertexOptions { + pub fn into_map(&self) -> Map { + [ + ("vertex_project", self.project.clone()), + ("vertex_location", self.location.clone()), + ] + .into_iter() + .filter_map(|(name, value)| value.map(|value| (name.to_string(), Value::String(value)))) + .collect() + } +} + #[derive(Clone, Debug, Default)] pub struct RequestOptions { pub api_key: Option, @@ -10,5 +77,7 @@ pub struct RequestOptions { pub extra_headers: Option>, pub extra_query: Option>, pub timeout: Option, - pub provider_connection: Map, + pub bedrock: Option, + pub anthropic: Option, + pub vertex: Option, } diff --git a/litellm-rust/crates/core/src/responses/types.rs b/litellm-rust/crates/core/src/responses/types.rs index 009310e3734..65a6951df0f 100644 --- a/litellm-rust/crates/core/src/responses/types.rs +++ b/litellm-rust/crates/core/src/responses/types.rs @@ -4,7 +4,6 @@ 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)] diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 0ff434b4107..5dbbb100ec9 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -17,7 +17,6 @@ use crate::marshal::{NativeRequestContext, NativeRequestOptions}; #[derive(FromPyObject)] struct WebSocketConnectRequest { url: String, - options: NativeRequestOptions, } #[pyclass] @@ -28,21 +27,19 @@ struct ResponsesWebSocketConnection { #[pymethods] impl ResponsesWebSocketConnection { #[classmethod] - #[pyo3(signature = (request, *, context))] + #[pyo3(signature = (request, *, options, context))] fn connect<'py>( _cls: &Bound<'py, pyo3::types::PyType>, py: Python<'py>, request: WebSocketConnectRequest, + options: NativeRequestOptions, context: NativeRequestContext, ) -> PyResult> { - let options: litellm_core::request_options::RequestOptions = request.options.into(); + let options: litellm_core::request_options::RequestOptions = options.into(); let context: litellm_core::request_context::LiteLlmRequestContext = context.into(); - let request = ResponsesWebSocketRequest { - url: request.url, - options, - }; + let request = ResponsesWebSocketRequest { url: request.url }; pyo3_async_runtimes::tokio::future_into_py(py, async move { - let inner = RustResponsesWebSocketConnection::connect(request, &context) + let inner = RustResponsesWebSocketConnection::connect(request, &options, &context) .await .map_err(core_error_to_pyerr)?; Ok(ResponsesWebSocketConnection { inner }) @@ -206,14 +203,14 @@ mod tests { import asyncio async def exercise(): - 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'), + for request, request_options, request_context, field in ( + (Request(url=123), options, context, 'url'), + (Request(url=url), Options(extra_headers=[]), context, 'extra_headers'), + (Request(url=url), options, replace(context, litellm_call_id=123), 'litellm_call_id'), + (Request(url=url), options, replace(context, attribution=Attribution(user_api_key_user_id=123)), 'user_api_key_user_id'), ): try: - native.ResponsesWebSocketConnection.connect(request, context=request_context) + native.ResponsesWebSocketConnection.connect(request, options=request_options, context=request_context) except (TypeError, ValueError) as error: parts = [] while error is not None: @@ -223,7 +220,7 @@ async def exercise(): else: raise AssertionError('invalid WebSocket input reached execution') - connection = await native.ResponsesWebSocketConnection.connect(Request(url=url), context=context) + connection = await native.ResponsesWebSocketConnection.connect(Request(url=url), options=options, context=context) assert type(connection) is native.ResponsesWebSocketConnection await connection.send_text("from-python") assert await connection.recv_text() == "from-server" diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 8e4543b7997..1b6b37f04de 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -4,6 +4,71 @@ use pyo3::exceptions::PyValueError; use pyo3::prelude::*; use serde_json::{Map, Value}; +#[derive(FromPyObject)] +struct NativeBedrockOptions { + aws_access_key_id: Option, + aws_secret_access_key: Option, + aws_session_token: Option, + aws_region_name: Option, + aws_session_name: Option, + aws_profile_name: Option, + aws_role_name: Option, + aws_web_identity_token: Option, + aws_sts_endpoint: Option, + aws_external_id: Option, + aws_bedrock_runtime_endpoint: Option, + request_metadata_fields: Vec, + request_metadata: Option>, +} + +impl From for litellm_core::request_options::BedrockOptions { + fn from(input: NativeBedrockOptions) -> Self { + Self { + aws_access_key_id: input.aws_access_key_id, + aws_secret_access_key: input.aws_secret_access_key, + aws_session_token: input.aws_session_token, + aws_region_name: input.aws_region_name, + aws_session_name: input.aws_session_name, + aws_profile_name: input.aws_profile_name, + aws_role_name: input.aws_role_name, + aws_web_identity_token: input.aws_web_identity_token, + aws_sts_endpoint: input.aws_sts_endpoint, + aws_external_id: input.aws_external_id, + aws_bedrock_runtime_endpoint: input.aws_bedrock_runtime_endpoint, + request_metadata_fields: input.request_metadata_fields, + request_metadata: input.request_metadata, + } + } +} + +#[derive(FromPyObject)] +struct NativeAnthropicOptions { + user_id: Option, +} + +impl From for litellm_core::request_options::AnthropicOptions { + fn from(input: NativeAnthropicOptions) -> Self { + Self { + user_id: input.user_id, + } + } +} + +#[derive(FromPyObject)] +struct NativeVertexOptions { + project: Option, + location: Option, +} + +impl From for litellm_core::request_options::VertexOptions { + fn from(input: NativeVertexOptions) -> Self { + Self { + project: input.project, + location: input.location, + } + } +} + #[derive(FromPyObject)] pub(crate) struct NativeRequestOptions { api_key: Option, @@ -14,8 +79,9 @@ pub(crate) struct NativeRequestOptions { #[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>, + bedrock: Option, + anthropic: Option, + vertex: Option, } impl From for litellm_core::request_options::RequestOptions { @@ -27,7 +93,9 @@ impl From for litellm_core::request_options::RequestOption extra_headers: input.extra_headers, extra_query: input.extra_query, timeout: optional_timeout(input.timeout_seconds), - provider_connection: input.provider_connection.unwrap_or_default(), + bedrock: input.bedrock.map(Into::into), + anthropic: input.anthropic.map(Into::into), + vertex: input.vertex.map(Into::into), } } } @@ -41,29 +109,38 @@ pub(crate) struct NativeRequestAttribution { #[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, + trace_id: Option, request_model: Option, attribution: NativeRequestAttribution, + capabilities: NativeRequestCapabilities, +} + +#[derive(FromPyObject)] +struct NativeRequestCapabilities { + stream: bool, + has_agentic_hook: bool, + has_custom_client: bool, + request_format: Option, } 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, + trace_id: input.trace_id, request_model: input.request_model, attribution: litellm_core::request_context::RequestAttribution { user_api_key_hash: input.attribution.user_api_key_hash, user_api_key_user_id: input.attribution.user_api_key_user_id, user_api_key_team_id: input.attribution.user_api_key_team_id, }, + capabilities: litellm_core::request_context::RequestCapabilities { + stream: input.capabilities.stream, + has_agentic_hook: input.capabilities.has_agentic_hook, + has_custom_client: input.capabilities.has_custom_client, + request_format: input.capabilities.request_format, + }, } } } @@ -107,7 +184,37 @@ class Options: extra_headers: object = None extra_query: object = None timeout_seconds: object = None - provider_connection: object = None + bedrock: object = None + anthropic: object = None + vertex: object = None + +@dataclass(frozen=True) +class BedrockOptions: + aws_access_key_id: object = None + aws_secret_access_key: object = None + aws_session_token: object = None + aws_region_name: object = None + aws_session_name: object = None + aws_profile_name: object = None + aws_role_name: object = None + aws_web_identity_token: object = None + aws_sts_endpoint: object = None + aws_external_id: object = None + aws_bedrock_runtime_endpoint: object = None + request_metadata_fields: object = () + request_metadata: object = None + +@dataclass(frozen=True) +class Capabilities: + stream: object = False + has_agentic_hook: object = False + has_custom_client: object = False + request_format: object = None + +@dataclass(frozen=True) +class VertexOptions: + project: object = None + location: object = None @dataclass(frozen=True) class Attribution: @@ -117,12 +224,11 @@ class Attribution: @dataclass(frozen=True) class Context: - metadata: object = None - litellm_metadata: object = None - request_metadata_fields: tuple = () litellm_call_id: object = None + trace_id: object = None request_model: object = None attribution: Attribution = Attribution() + capabilities: Capabilities = Capabilities() @dataclass(frozen=True) class Request: @@ -132,11 +238,11 @@ class Request: audio: object = None document: object = None optional_params: object = None - options: Options = Options() value: str = '' url: str = '' context = Context() +options = Options() ", Some(&locals), Some(&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 0bb93bbf132..18d351b3791 100644 --- a/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs +++ b/litellm-rust/crates/python-bridge/src/routes/audio_transcription.rs @@ -15,11 +15,11 @@ struct AudioTranscriptionInputs { audio: Value, #[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Map, - options: NativeRequestOptions, } fn prepare_transcription( input: AudioTranscriptionInputs, + options: NativeRequestOptions, context: NativeRequestContext, ) -> PyResult> + Send + 'static> { let context: LiteLlmRequestContext = context.into(); @@ -30,8 +30,8 @@ fn prepare_transcription( model: &input.model, audio, optional_params: input.optional_params, - options: input.options.into(), }, + &options.into(), &context, ) .await 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 96dad4c6328..aad6c3c737f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -16,11 +16,11 @@ struct ChatCompletionsInputs { messages: Value, #[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Map, - options: NativeRequestOptions, } fn prepare_chat_completions( input: ChatCompletionsInputs, + options: NativeRequestOptions, context: NativeRequestContext, ) -> PyResult> + Send + 'static> { let context: LiteLlmRequestContext = context.into(); @@ -31,8 +31,8 @@ fn prepare_chat_completions( model: &input.model, messages, optional_params: input.optional_params, - options: input.options.into(), }, + &options.into(), &context, ) .await diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index 3957fb9976f..cf38fef9d55 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -12,24 +12,26 @@ macro_rules! bridge_route { $(, extra = [$($extra:ident),* $(,)?])? $(,)? ) => { #[pyfunction] - #[pyo3(signature = (request, *, context))] + #[pyo3(signature = (request, *, options, context))] fn $sync_name( py: pyo3::Python<'_>, request: $inputs, + options: $crate::marshal::NativeRequestOptions, context: $crate::marshal::NativeRequestContext, ) -> pyo3::PyResult> { - let future = $prepare(request, context)?; + let future = $prepare(request, options, context)?; $crate::execution::run_sync(py, future, $map_error) } #[pyfunction] - #[pyo3(signature = (request, *, context))] + #[pyo3(signature = (request, *, options, context))] fn $async_name( py: pyo3::Python<'_>, request: $inputs, + options: $crate::marshal::NativeRequestOptions, context: $crate::marshal::NativeRequestContext, ) -> pyo3::PyResult> { - let future = $prepare(request, context)?; + let future = $prepare(request, options, context)?; $crate::execution::run_async(py, future, $map_error) } @@ -46,24 +48,26 @@ macro_rules! bridge_route { use super::{$inputs, $map_error, $prepare}; #[pyfunction] - #[pyo3(signature = (request, *, context))] + #[pyo3(signature = (request, *, options, context))] fn $sync_name( py: pyo3::Python<'_>, request: $inputs, + options: $crate::marshal::NativeRequestOptions, context: $crate::marshal::NativeRequestContext, ) -> pyo3::PyResult> { - let future = $prepare(request, context)?; + let future = $prepare(request, options, context)?; $crate::execution::run_sync(py, $crate::function_trace::capture(future), $map_error) } #[pyfunction] - #[pyo3(signature = (request, *, context))] + #[pyo3(signature = (request, *, options, context))] fn $async_name( py: pyo3::Python<'_>, request: $inputs, + options: $crate::marshal::NativeRequestOptions, context: $crate::marshal::NativeRequestContext, ) -> pyo3::PyResult> { - let future = $prepare(request, context)?; + let future = $prepare(request, options, context)?; $crate::execution::run_async(py, $crate::function_trace::capture(future), $map_error) } @@ -141,6 +145,7 @@ mod tests { fn prepare_echo( inputs: EchoInputs, + _options: crate::marshal::NativeRequestOptions, _context: crate::marshal::NativeRequestContext, ) -> PyResult> + Send + 'static> { FUTURE_DROPPED.store(false, Ordering::SeqCst); @@ -182,13 +187,17 @@ mod tests { let module = PyModule::new(py, "routes").expect("module should be created"); crate::routes::register(&module).expect("routes should register"); let routes = [ - ("ocr", "aocr", "(request, *, context)"), - ("transcription", "atranscription", "(request, *, context)"), - ("messages", "amessages", "(request, *, context)"), + ("ocr", "aocr", "(request, *, options, context)"), + ( + "transcription", + "atranscription", + "(request, *, options, context)", + ), + ("messages", "amessages", "(request, *, options, context)"), ( "chat_completions", "achat_completions", - "(request, *, context)", + "(request, *, options, context)", ), ]; @@ -219,17 +228,17 @@ mod tests { 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'), +for names, request, request_options, expected in [ + (('chat_completions', 'achat_completions'), Request(messages={}, optional_params={}), options, 'messages must be a list'), + (('messages', 'amessages'), Request(body=[]), options, 'body must be a dict'), + (('ocr', 'aocr'), Request(document={}, optional_params={}), Options(extra_headers=[]), 'extra_headers'), + (('transcription', 'atranscription'), Request(audio={}, optional_params={}), Options(timeout_seconds='bad'), 'timeout_seconds'), + (('transcription', 'atranscription'), Request(audio={}, optional_params={}), Options(bedrock=[]), 'bedrock'), ]: errors = [] for name in names: try: - getattr(routes, name)(request, context=context) + getattr(routes, name)(request, options=request_options, context=context) except (ValueError, TypeError) as error: parts = [] while error is not None: @@ -241,10 +250,10 @@ for names, request, expected in [ assert errors[0] == errors[1], errors assert expected in errors[0], (expected, errors) -for field in ('metadata', 'litellm_metadata', 'request_metadata_fields'): +for field in ('litellm_call_id', 'trace_id', 'request_model'): invalid_context = replace(context, **{field: object()}) try: - routes.chat_completions(Request(messages=[], optional_params={}), context=invalid_context) + routes.chat_completions(Request(messages=[], optional_params={}), options=options, context=invalid_context) except (ValueError, TypeError) as error: assert field in str(error) else: @@ -267,6 +276,7 @@ for field in ('metadata', 'litellm_metadata', 'request_metadata_fields'): py.eval(c"Request(value=\"sync\")", Some(&locals), Some(&locals)) .and_then(|request| { let kwargs = PyDict::new(py); + kwargs.set_item("options", locals.get_item("options")?.unwrap())?; kwargs.set_item("context", locals.get_item("context")?.unwrap())?; function.call((request,), Some(&kwargs)) }) @@ -282,6 +292,7 @@ for field in ('metadata', 'litellm_metadata', 'request_metadata_fields'): py.eval(c"Request(value=\"error\")", Some(&locals), Some(&locals)) .and_then(|request| { let kwargs = PyDict::new(py); + kwargs.set_item("options", locals.get_item("options")?.unwrap())?; kwargs.set_item("context", locals.get_item("context")?.unwrap())?; function.call((request,), Some(&kwargs)) }) @@ -302,17 +313,17 @@ for field in ('metadata', 'litellm_metadata', 'request_metadata_fields'): import asyncio async def exercise(): - assert await routes.aecho(Request(value="async"), context=context) == "async" + assert await routes.aecho(Request(value="async"), options=options, context=context) == "async" try: - await routes.aecho(Request(value="error"), context=context) + await routes.aecho(Request(value="error"), options=options, context=context) except LookupError as error: assert str(error) == "invalid request: synthetic error" else: raise AssertionError("mapped error was not raised") try: - await routes.aecho(Request(value="panic"), context=context) + await routes.aecho(Request(value="panic"), options=options, context=context) except BaseException as error: assert type(error).__name__ == "PanicException" assert str(error) == "synthetic panic" @@ -320,14 +331,14 @@ async def exercise(): raise AssertionError("panic was not raised") try: - await routes.aecho(Request(value="map_panic"), context=context) + await routes.aecho(Request(value="map_panic"), options=options, context=context) except BaseException as error: assert type(error).__name__ == "PanicException" assert str(error) == "synthetic mapper panic" else: raise AssertionError("mapper panic was not raised") - task = asyncio.ensure_future(routes.aecho(Request(value="pending"), context=context)) + task = asyncio.ensure_future(routes.aecho(Request(value="pending"), options=options, context=context)) await asyncio.sleep(0) task.cancel() try: @@ -365,7 +376,7 @@ asyncio.run(exercise()) .expect("module should enter Python locals"); let code = CString::new( r#" -result = routes.echo(Request(value="traced"), context=context) +result = routes.echo(Request(value="traced"), options=options, 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 2ecb3690c15..2ee126b9807 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages.rs @@ -13,11 +13,11 @@ struct MessagesInputs { model: String, #[pyo3(from_py_with = litellm_python_interop::from_py)] body: Value, - options: NativeRequestOptions, } fn prepare_messages( input: MessagesInputs, + options: NativeRequestOptions, context: NativeRequestContext, ) -> PyResult> + Send + 'static> { let context: LiteLlmRequestContext = context.into(); @@ -27,8 +27,8 @@ fn prepare_messages( MessagesRequest { model: &input.model, body, - options: input.options.into(), }, + &options.into(), &context, ) .await diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index f4a3c8ddc6e..fceb4b1b865 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -16,11 +16,11 @@ struct OcrInputs { document: Value, #[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Map, - options: NativeRequestOptions, } fn prepare_ocr( input: OcrInputs, + options: NativeRequestOptions, context: NativeRequestContext, ) -> PyResult> + Send + 'static> { let context: LiteLlmRequestContext = context.into(); @@ -31,8 +31,8 @@ fn prepare_ocr( model: &input.model, document, optional_params: input.optional_params, - options: input.options.into(), }, + &options.into(), &context, RequestHooks { callbacks: Vec::new(), diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index bcc2235e11e..cc9430f6bf4 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -27,8 +27,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts -from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors -from litellm.rust_bridge.runtime import DispatchResult +from litellm.rust_bridge.request import anthropic_options from litellm.types.llms.anthropic import ( ContentBlockDelta, ContentBlockStart, @@ -370,7 +369,15 @@ class AnthropicChatCompletion(BaseLLM): if config is None: raise ValueError(f"Provider config not found for model: {model} and provider: {custom_llm_provider}") - def prepare_python() -> tuple[dict[str, str], dict[str, object]]: # mutable-ok: stream mutates data + def build_request() -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream + """Translate the request the Python way, returning `(headers, data)`. + + The pair stays mutable because the streaming path rewrites it in + place (`data["stream"] = True`) before sending. + + Shared by the normal path and by the Rust path's fallback, which + builds it only when the Rust call did not serve the request. + """ request_data: Final = config.transform_request( model=model, messages=messages, @@ -378,29 +385,12 @@ class AnthropicChatCompletion(BaseLLM): litellm_params=litellm_params, headers=headers, ) - python_headers, data = update_request_with_filtered_beta( + return update_request_with_filtered_beta( headers=headers, request_data=request_data, provider=custom_llm_provider, ) - ## LOGGING - # Reaching here with `serves_via_rust` set means the Rust attempt - # declined at call time, before the provider was called, and already - # logged this request. That is the same attempt continuing. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, - api_key=api_key, - additional_args={ - "complete_input_dict": data, - "api_base": api_base, - "headers": python_headers, - }, - ) - print_verbose(f"_is_function_call: {_is_function_call}") - return python_headers, data - # The Rust core owns the whole call for the subset it accepts, so ask # before transforming: whichever path runs emits pre_call exactly once. # `get_config` merges the class-level defaults (Anthropic's required @@ -417,26 +407,68 @@ class AnthropicChatCompletion(BaseLLM): litellm_params=litellm_params, stream=stream, ) - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "model": model, - "messages": messages, - **rust_optional_params, - }, - "api_base": api_base, - "headers": headers, - } if serves_via_rust: + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "model": model, + "messages": messages, + **rust_optional_params, + }, + "api_base": api_base, + "headers": headers, + } logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args) - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, - messages=messages, - api_key=api_key, - additional_args=rust_logging_args, - ) + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key=api_key, + additional_args=rust_logging_args, + ) + if acompletion is True: - def native_completion() -> DispatchResult[ModelResponse]: - return rust_chat_completions_bridge.chat_completions( + async def python_fallback() -> "ModelResponse | CustomStreamWrapper": + # pre_call already fired for this request above. The Rust + # path only declines before the provider is called, so this + # is the same attempt continuing, not a second one. + fallback_headers, fallback_data = build_request() + return await self.acompletion_function( + model=model, + messages=messages, + data=fallback_data, + api_base=api_base, + custom_prompt_dict=custom_prompt_dict, + model_response=model_response, + print_verbose=print_verbose, + encoding=encoding, + api_key=api_key, + provider_config=config, + logging_obj=logging_obj, + optional_params=optional_params, + stream=stream, + _is_function_call=_is_function_call, + litellm_params=litellm_params, + logger_fn=logger_fn, + headers=fallback_headers, + client=client, + json_mode=json_mode, + timeout=timeout, + ) + + return rust_chat_completions_bridge.achat_completions_or_fallback( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + python_fallback=python_fallback, + anthropic=anthropic_options(litellm_params), + ) + rust_response: Final = rust_chat_completions_bridge.chat_completions( model=model, messages=messages, optional_params=rust_optional_params, @@ -447,37 +479,35 @@ class AnthropicChatCompletion(BaseLLM): extra_headers=headers, timeout=timeout, on_response=log_rust_post_call, - eligible=serves_via_rust, + anthropic=anthropic_options(litellm_params), ) + if rust_response is not None: + return rust_response - async def native_acompletion() -> DispatchResult[ModelResponse]: - return await rust_chat_completions_bridge.achat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, + headers, data = build_request() + + ## LOGGING + # Reaching here with `serves_via_rust` set means the Rust attempt + # declined at call time, before the provider was called, and already + # logged this request. That is the same attempt continuing. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - eligible=serves_via_rust, + additional_args={ + "complete_input_dict": data, + "api_base": api_base, + "headers": headers, + }, ) - - @anative_first( - native=native_acompletion, - route="chat_completions", - errors=lambda: provider_errors(custom_llm_provider or "", model), - ) - async def execute_async() -> ModelResponse | CustomStreamWrapper: - headers, data = prepare_python() + print_verbose(f"_is_function_call: {_is_function_call}") + if acompletion is True: if ( stream is True ): # if function call - fake the streaming (need complete blocks for output parsing in openai format) print_verbose("makes async anthropic streaming POST request") data["stream"] = stream - return await self.acompletion_stream_function( + return self.acompletion_stream_function( model=model, messages=messages, data=data, @@ -499,7 +529,7 @@ class AnthropicChatCompletion(BaseLLM): client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None), ) else: - return await self.acompletion_function( + return self.acompletion_function( model=model, messages=messages, data=data, @@ -521,14 +551,7 @@ class AnthropicChatCompletion(BaseLLM): json_mode=json_mode, timeout=timeout, ) - - @native_first( - native=native_completion, - route="chat_completions", - errors=lambda: provider_errors(custom_llm_provider or "", model), - ) - def execute_sync() -> ModelResponse | CustomStreamWrapper: - headers, data = prepare_python() + else: ## COMPLETION CALL if ( stream is True @@ -560,12 +583,13 @@ class AnthropicChatCompletion(BaseLLM): ) else: - python_client: Final = ( - client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout}) - ) + if client is None or not isinstance(client, HTTPHandler): + client = _get_httpx_client(params={"timeout": timeout}) + else: + client = client try: - response: Final = python_client.post( + response: Final = client.post( api_base, headers=headers, data=json.dumps(data), @@ -586,21 +610,20 @@ class AnthropicChatCompletion(BaseLLM): status_code=status_code, headers=error_headers, ) - return config.transform_response( - model=model, - raw_response=response, - model_response=model_response, - logging_obj=logging_obj, - api_key=api_key, - request_data=data, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - encoding=encoding, - json_mode=json_mode, - ) - return execute_async() if acompletion else execute_sync() + return config.transform_response( + model=model, + raw_response=response, + model_response=model_response, + logging_obj=logging_obj, + api_key=api_key, + request_data=data, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + encoding=encoding, + json_mode=json_mode, + ) def embedding(self): # logic for parsing in - calling - parsing out model embedding calls diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 8d6cf23f47e..19ffb695af1 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -10,6 +10,10 @@ from litellm.anthropic_beta_headers_manager import ( update_headers_with_filtered_beta, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject +from litellm.llms.bedrock.request_metadata import ( + get_bedrock_request_metadata_fields, + resolve_bedrock_request_metadata, +) from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, @@ -18,8 +22,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts -from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors -from litellm.rust_bridge.runtime import DispatchResult +from litellm.rust_bridge.request import NativeBedrockOptions from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -395,11 +398,15 @@ class BedrockConverseLLM(BaseAWSLLM): # resolved so both paths sign as the same principal. Bearer-token auth # resolves no SigV4 principal at all, and each path reads that token # itself. - rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy - **optional_params, - **_sigv4_principal(credentials), - "aws_region_name": aws_region_name, - } + rust_optional_params: Final = optional_params + rust_bedrock_options: Final = NativeBedrockOptions( + aws_access_key_id=None if credentials is None else credentials.access_key, + aws_secret_access_key=None if credentials is None else credentials.secret_key, + aws_session_token=None if credentials is None else credentials.token, + aws_region_name=aws_region_name, + request_metadata_fields=get_bedrock_request_metadata_fields(), + request_metadata=resolve_bedrock_request_metadata(litellm_params, optional_params.get("requestMetadata")), + ) serves_via_rust: Final = rust_chat_completions_accepts( model=model, messages=messages, @@ -408,25 +415,55 @@ class BedrockConverseLLM(BaseAWSLLM): litellm_params=litellm_params, stream=stream, ) - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "messages": messages, - **optional_params, - }, - "api_base": proxy_endpoint_url, - "headers": headers, - } if serves_via_rust: + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "messages": messages, + **optional_params, + }, + "api_base": proxy_endpoint_url, + "headers": headers, + } logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args) - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, - messages=messages, - api_key="", - additional_args=rust_logging_args, - ) - - def native_completion() -> DispatchResult[ModelResponse]: - return rust_chat_completions_bridge.chat_completions( + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key="", + additional_args=rust_logging_args, + ) + if acompletion: + return rust_chat_completions_bridge.achat_completions_or_fallback( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=proxy_endpoint_url, + custom_llm_provider="bedrock", + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + bedrock=rust_bedrock_options, + python_fallback=lambda: self.async_completion( + model=model, + messages=messages, + api_base=proxy_endpoint_url, + model_response=model_response, + encoding=encoding, + logging_obj=logging_obj, + optional_params=optional_params, + stream=stream, + litellm_params=litellm_params, + logger_fn=logger_fn, + headers=headers, + timeout=timeout, + client=client, + credentials=credentials, + api_key=api_key, + skip_pre_call_logging=True, + ), + ) + rust_response: Final = rust_chat_completions_bridge.chat_completions( model=model, messages=messages, optional_params=rust_optional_params, @@ -437,33 +474,17 @@ class BedrockConverseLLM(BaseAWSLLM): extra_headers=headers, timeout=timeout, on_response=log_rust_post_call, - eligible=serves_via_rust, + bedrock=rust_bedrock_options, ) + if rust_response is not None: + return rust_response - async def native_acompletion() -> DispatchResult[ModelResponse]: - return await rust_chat_completions_bridge.achat_completions( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=proxy_endpoint_url, - custom_llm_provider="bedrock", - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - eligible=serves_via_rust, - ) - - @anative_first( - native=native_acompletion, - route="chat_completions", - errors=lambda: provider_errors("bedrock", model), - ) - async def execute_async() -> ModelResponse | CustomStreamWrapper: - python_client: Final = None if isinstance(client, HTTPHandler) else client + ### ROUTING (ASYNC, STREAMING, SYNC) + if acompletion: + if isinstance(client, HTTPHandler): + client = None if stream is True: - return await self.async_streaming( + return self.async_streaming( model=model, messages=messages, api_base=proxy_endpoint_url, @@ -476,7 +497,7 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=logger_fn, headers=headers, timeout=timeout, - client=python_client, + client=client, json_mode=json_mode, fake_stream=fake_stream, credentials=credentials, @@ -484,7 +505,7 @@ class BedrockConverseLLM(BaseAWSLLM): stream_chunk_size=stream_chunk_size, ) ### ASYNC COMPLETION - return await self.async_completion( + return self.async_completion( model=model, messages=messages, api_base=proxy_endpoint_url, @@ -497,112 +518,108 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=logger_fn, headers=headers, timeout=timeout, - client=python_client, + client=client, credentials=credentials, api_key=api_key, - skip_pre_call_logging=serves_via_rust, ) - @native_first( - native=native_completion, - route="chat_completions", - errors=lambda: provider_errors("bedrock", model), + ## TRANSFORMATION ## + + _data: Final = litellm.AmazonConverseConfig()._transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=extra_headers, ) - def execute_sync() -> ModelResponse | CustomStreamWrapper: - ## TRANSFORMATION ## + data: Final = json.dumps(_data) - _data: Final = litellm.AmazonConverseConfig()._transform_request( - model=model, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - headers=extra_headers, - ) - data: Final = json.dumps(_data) + prepped: Final = self.get_request_headers( + credentials=credentials, + aws_region_name=aws_region_name, + extra_headers=extra_headers, + endpoint_url=proxy_endpoint_url, + data=data, + headers=headers, + api_key=api_key, + ) - prepped: Final = self.get_request_headers( - credentials=credentials, - aws_region_name=aws_region_name, - extra_headers=extra_headers, - endpoint_url=proxy_endpoint_url, - data=data, - headers=headers, - api_key=api_key, - ) - - ## LOGGING - # Reaching here with `serves_via_rust` set means the synchronous Rust - # attempt declined at call time, before the provider was called, and - # already logged this request. That is the same attempt continuing. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) - resolved_timeout: Final = httpx.Timeout(timeout) if isinstance(timeout, (float, int)) else timeout - python_client: Final = ( - _get_httpx_client({"timeout": resolved_timeout} if resolved_timeout is not None else None) - if client is None or isinstance(client, AsyncHTTPHandler) - else client - ) - - if stream is not None and stream is True: - completion_stream, response_headers = make_sync_call( - client=python_client, - api_base=proxy_endpoint_url, - headers=prepped.headers, - data=data, - model=model, - messages=messages, - logging_obj=logging_obj, - json_mode=json_mode, - fake_stream=fake_stream, - stream_chunk_size=stream_chunk_size, - ) - streaming_response: Final = CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="bedrock", - logging_obj=logging_obj, - _response_headers=response_headers, - ) - - return streaming_response - - ### COMPLETION - - try: - response: Final = python_client.post( - url=proxy_endpoint_url, - headers=prepped.headers, - data=data, - logging_obj=logging_obj, - ) - response.raise_for_status() - except httpx.HTTPStatusError as err: - error_code: Final = err.response.status_code - raise BedrockError(status_code=error_code, message=err.response.text) - except httpx.TimeoutException: - raise BedrockError(status_code=408, message="Timeout error occurred.") - - sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response( - model=model, - response=response, - model_response=model_response, - stream=stream if isinstance(stream, bool) else False, - logging_obj=logging_obj, + ## LOGGING + # Reaching here with `serves_via_rust` set means the synchronous Rust + # attempt declined at call time, before the provider was called, and + # already logged this request. That is the same attempt continuing. + # The asynchronous branch above returns before this point, and hands + # its own fallback `skip_pre_call_logging=True` for the same reason. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, api_key="", - data=data, - messages=messages, - optional_params=optional_params, - encoding=encoding, + additional_args={ + "complete_input_dict": data, + "api_base": proxy_endpoint_url, + "headers": prepped.headers, + }, ) - sync_transformed_response.set_provider_response_headers(response.headers) - return sync_transformed_response + if client is None or isinstance(client, AsyncHTTPHandler): + _params: Final = {} + if timeout is not None: + if isinstance(timeout, float) or isinstance(timeout, int): + timeout = httpx.Timeout(timeout) + _params["timeout"] = timeout + client = _get_httpx_client(_params) + else: + client = client - return execute_async() if acompletion else execute_sync() + if stream is not None and stream is True: + completion_stream, response_headers = make_sync_call( + client=(client if client is not None and isinstance(client, HTTPHandler) else None), + api_base=proxy_endpoint_url, + headers=prepped.headers, + data=data, + model=model, + messages=messages, + logging_obj=logging_obj, + json_mode=json_mode, + fake_stream=fake_stream, + stream_chunk_size=stream_chunk_size, + ) + streaming_response: Final = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="bedrock", + logging_obj=logging_obj, + _response_headers=response_headers, + ) + + return streaming_response + + ### COMPLETION + + try: + response: Final = client.post( + url=proxy_endpoint_url, + headers=prepped.headers, + data=data, + logging_obj=logging_obj, + ) + response.raise_for_status() + except httpx.HTTPStatusError as err: + error_code: Final = err.response.status_code + raise BedrockError(status_code=error_code, message=err.response.text) + except httpx.TimeoutException: + raise BedrockError(status_code=408, message="Timeout error occurred.") + + sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response( + model=model, + response=response, + model_response=model_response, + stream=stream if isinstance(stream, bool) else False, + logging_obj=logging_obj, + api_key="", + data=data, + messages=messages, + optional_params=optional_params, + encoding=encoding, + ) + sync_transformed_response.set_provider_response_headers(response.headers) + return sync_transformed_response diff --git a/litellm/llms/bedrock/request_metadata.py b/litellm/llms/bedrock/request_metadata.py index 1f4e5886508..599afec42e8 100644 --- a/litellm/llms/bedrock/request_metadata.py +++ b/litellm/llms/bedrock/request_metadata.py @@ -43,7 +43,7 @@ def _text_pairs(source: object) -> tuple[tuple[str, str], ...]: return tuple((key, value) for key, value in source.items() if isinstance(key, str) and isinstance(value, str)) -def _allowed_fields() -> tuple[str, ...]: +def get_bedrock_request_metadata_fields() -> tuple[str, ...]: """ The operator allow-list, deduplicated so a field repeated in config cannot consume a second reserved slot and shrink the client budget for nothing. First occurrence wins, which keeps @@ -121,7 +121,7 @@ def resolve_bedrock_request_metadata( been validated (and rejected with a 400) by the Converse transformation, so it is only filtered here for the reserved identity prefix and the remaining slot budget. """ - allowed_fields: Final = _allowed_fields() + allowed_fields: Final = get_bedrock_request_metadata_fields() if not allowed_fields: return None sources: Final = _metadata_sources(litellm_params) @@ -146,7 +146,7 @@ def bedrock_request_metadata_is_owned() -> bool: "fall back to whatever the caller supplied", or the reserved-prefix guarantee is bypassable by anyone who can make the resolver produce nothing. """ - return bool(_allowed_fields()) + return bool(get_bedrock_request_metadata_fields()) def bedrock_request_metadata_headers( diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 9c77fc39909..bebb4b44589 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -32,8 +32,7 @@ from litellm.rust_bridge import ocr as rust_ocr_bridge from litellm.rust_bridge.request import ( NativeRequestOptions, PreparedNativeCall, - provider_connection_params, - provider_request_params, + vertex_options, ) from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.router import GenericLiteLLMParams @@ -275,18 +274,18 @@ def _prepare_rust_ocr_call( request=rust_ocr_bridge.NativeOCRRequest( model=prepared_request.model, document=prepared_request.document, - optional_params=provider_request_params(rust_optional_params), - options=NativeRequestOptions( - provider_connection=provider_connection_params(rust_optional_params), - api_key=resolved_api_key, - api_base=rust_api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=cast( # cast-ok: provider header normalization returns string-object pairs - dict[str, object], resolved_headers - ), - timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + optional_params=prepared_request.optional_params, + ), + options=NativeRequestOptions( + vertex=vertex_options(rust_optional_params), + api_key=resolved_api_key, + api_base=rust_api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=cast( # cast-ok: provider header normalization returns string-object pairs + dict[str, object], resolved_headers ), - ) + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + ), ) diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index 8cb5c80f2fd..a46a1f541bd 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -32,13 +32,13 @@ from litellm.rust_bridge.protocols import ( RustChatCompletionsDecline, ) from litellm.rust_bridge.request import ( + NativeAnthropicOptions, + NativeBedrockOptions, NativeChatCompletionsRequest, NativeRequestContext, NativeRequestOptions, PreparedNativeCall, call_native, - provider_connection_params, - provider_request_params, ) from litellm.rust_bridge.runtime import ( BridgeErrorContext, @@ -238,6 +238,8 @@ def chat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, + bedrock: NativeBedrockOptions | None = None, + anthropic: NativeAnthropicOptions | None = None, ) -> ModelResponse | None: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) @@ -248,15 +250,16 @@ def chat_completions( 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), - ), + optional_params=optional_params, + ), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + bedrock=bedrock, + anthropic=anthropic, ), context=NativeRequestContext(), ), @@ -279,6 +282,8 @@ async def achat_completions( extra_headers: Mapping[str, object] | None, timeout: float | httpx.Timeout | None, on_response: ResponseObserver, + bedrock: NativeBedrockOptions | None = None, + anthropic: NativeAnthropicOptions | None = None, ) -> ModelResponse | None: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) @@ -289,15 +294,16 @@ async def achat_completions( 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), - ), + optional_params=optional_params, + ), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + bedrock=bedrock, + anthropic=anthropic, ), context=NativeRequestContext(), ), @@ -321,6 +327,8 @@ async def achat_completions_or_fallback( timeout: float | httpx.Timeout | None, on_response: ResponseObserver, python_fallback: Callable[[], Awaitable[object]], + bedrock: NativeBedrockOptions | None = None, + anthropic: NativeAnthropicOptions | None = None, ) -> object: """Await the Rust path, falling back to the caller's own Python path when the bridge is unavailable or the call fails. @@ -340,15 +348,16 @@ async def achat_completions_or_fallback( 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), - ), + optional_params=optional_params, + ), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + bedrock=bedrock, + anthropic=anthropic, ), context=NativeRequestContext(), ), diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index d9ea26333bd..f41cff140f8 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -72,13 +72,13 @@ def messages( 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), - ), + ), + 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(), ), @@ -104,13 +104,13 @@ async def amessages( 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), - ), + ), + 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(), ), diff --git a/litellm/rust_bridge/protocols.py b/litellm/rust_bridge/protocols.py index 76b76504a6c..dc779b346ab 100644 --- a/litellm/rust_bridge/protocols.py +++ b/litellm/rust_bridge/protocols.py @@ -9,6 +9,7 @@ from .request import ( NativeMessagesRequest, NativeOCRRequest, NativeRequestContext, + NativeRequestOptions, NativeResponsesWebSocketRequest, NativeTranscriptionRequest, ) @@ -47,6 +48,7 @@ class RustResponsesWebSocketConnection(Protocol): cls, request: NativeResponsesWebSocketRequest, *, + options: NativeRequestOptions, context: NativeRequestContext, ) -> RustResponsesWebSocket: ... diff --git a/litellm/rust_bridge/request.py b/litellm/rust_bridge/request.py index ba650f22372..4239b9763a1 100644 --- a/litellm/rust_bridge/request.py +++ b/litellm/rust_bridge/request.py @@ -2,7 +2,70 @@ from __future__ import annotations from collections.abc import Mapping, Sequence from dataclasses import dataclass -from typing import Final, Generic, Protocol, TypeVar +from typing import Generic, Protocol, TypeVar + + +@dataclass(frozen=True, slots=True) +class NativeBedrockOptions: + aws_access_key_id: str | None = None + aws_secret_access_key: str | None = None + aws_session_token: str | None = None + aws_region_name: str | None = None + aws_session_name: str | None = None + aws_profile_name: str | None = None + aws_role_name: str | None = None + aws_web_identity_token: str | None = None + aws_sts_endpoint: str | None = None + aws_external_id: str | None = None + aws_bedrock_runtime_endpoint: str | None = None + request_metadata_fields: tuple[str, ...] = () + request_metadata: Mapping[str, str] | None = None + + +@dataclass(frozen=True, slots=True) +class NativeAnthropicOptions: + user_id: str | None = None + + +@dataclass(frozen=True, slots=True) +class NativeVertexOptions: + project: str | None = None + location: str | None = None + + +def bedrock_options(params: Mapping[str, object]) -> NativeBedrockOptions: + def string(name: str) -> str | None: + value = params.get(name) + return value if isinstance(value, str) else None + + return NativeBedrockOptions( + aws_access_key_id=string("aws_access_key_id"), + aws_secret_access_key=string("aws_secret_access_key"), + aws_session_token=string("aws_session_token"), + aws_region_name=string("aws_region_name"), + aws_session_name=string("aws_session_name"), + aws_profile_name=string("aws_profile_name"), + aws_role_name=string("aws_role_name"), + aws_web_identity_token=string("aws_web_identity_token"), + aws_sts_endpoint=string("aws_sts_endpoint"), + aws_external_id=string("aws_external_id"), + aws_bedrock_runtime_endpoint=string("aws_bedrock_runtime_endpoint"), + ) + + +def anthropic_options(litellm_params: Mapping[str, object] | None) -> NativeAnthropicOptions: + metadata = None if litellm_params is None else litellm_params.get("metadata") + user_id = metadata.get("user_id") if isinstance(metadata, Mapping) else None + return NativeAnthropicOptions(user_id=user_id if isinstance(user_id, str) else None) + + +def vertex_options(params: Mapping[str, object]) -> NativeVertexOptions: + project = params.get("vertex_project") or params.get("vertex_ai_project") + location = params.get("vertex_location") or params.get("vertex_ai_location") + return NativeVertexOptions( + project=project if isinstance(project, str) else None, + location=location if isinstance(location, str) else None, + ) @dataclass(frozen=True, slots=True) @@ -13,7 +76,9 @@ class NativeRequestOptions: 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 + bedrock: NativeBedrockOptions | None = None + anthropic: NativeAnthropicOptions | None = None + vertex: NativeVertexOptions | None = None @dataclass(frozen=True, slots=True) @@ -23,14 +88,21 @@ class RequestAttribution: user_api_key_team_id: str | None = None +@dataclass(frozen=True, slots=True) +class NativeRequestCapabilities: + stream: bool = False + has_agentic_hook: bool = False + has_custom_client: bool = False + request_format: 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 + trace_id: str | None = None request_model: str | None = None attribution: RequestAttribution = RequestAttribution() + capabilities: NativeRequestCapabilities = NativeRequestCapabilities() RequestT = TypeVar("RequestT") @@ -41,48 +113,22 @@ ResultT = TypeVar("ResultT", covariant=True) @dataclass(frozen=True, slots=True) class PreparedNativeCall(Generic[RequestT]): request: RequestT + options: NativeRequestOptions = NativeRequestOptions() context: NativeRequestContext = NativeRequestContext() class NativeFunction(Protocol[RequestContraT, ResultT]): - def __call__(self, request: RequestContraT, *, context: NativeRequestContext) -> ResultT: ... + def __call__( + self, + request: RequestContraT, + *, + options: NativeRequestOptions, + context: NativeRequestContext, + ) -> ResultT: ... def call_native(native: NativeFunction[RequestT, ResultT], prepared: PreparedNativeCall[RequestT]) -> ResultT: - return native(prepared.request, 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 - } + return native(prepared.request, options=prepared.options, context=prepared.context) @dataclass(frozen=True, slots=True) @@ -90,14 +136,12 @@ 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) @@ -105,7 +149,6 @@ class NativeOCRRequest: model: str document: dict[str, object] optional_params: dict[str, object] - options: NativeRequestOptions @dataclass(frozen=True, slots=True) @@ -113,10 +156,8 @@ 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 a48405a7884..b7098aaddb4 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -73,8 +73,8 @@ async def connect( prepare=lambda: PreparedNativeCall( NativeResponsesWebSocketRequest( url=url, - options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)), ), + options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)), context=NativeRequestContext(), ), call=lambda connection_type, request: call_native(connection_type.connect, request), diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 7c89ae449c9..0300cd84883 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -11,9 +11,8 @@ from litellm.rust_bridge.request import ( NativeRequestOptions, NativeTranscriptionRequest, PreparedNativeCall, + bedrock_options, call_native, - provider_connection_params, - provider_request_params, ) from litellm.rust_bridge.runtime import ( BridgeErrorContext, @@ -74,15 +73,15 @@ def transcription( 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), - ), + optional_params=optional_params, + ), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + bedrock=bedrock_options(optional_params), ), context=NativeRequestContext(), ), @@ -109,15 +108,15 @@ async def atranscription( 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), - ), + optional_params=optional_params, + ), + options=NativeRequestOptions( + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + bedrock=bedrock_options(optional_params), ), context=NativeRequestContext(), ), 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 e3a032c8d94..5bd0358eeb2 100644 --- a/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py +++ b/tests/rust-python-harness/strategies/trace_parity/sdk/execution.py @@ -69,8 +69,8 @@ def _native_kwargs(route: str, kwargs: dict[str, object]) -> dict[str, object]: from litellm.rust_bridge.request import ( NativeRequestContext, NativeRequestOptions, - provider_connection_params, - provider_request_params, + bedrock_options, + vertex_options, ) from litellm.rust_bridge.transcription import NativeTranscriptionRequest @@ -88,25 +88,36 @@ def _native_kwargs(route: str, kwargs: dict[str, object]) -> dict[str, object]: "timeout_seconds", ) }, - "provider_connection": provider_connection_params(params), + "bedrock": bedrock_options(params), + "vertex": vertex_options(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()} + if route == "chat_completions": + request: Final = NativeChatCompletionsRequest( + model=TypeAdapter(str).validate_python(kwargs.get("model")), + messages=TypeAdapter(list[object]).validate_python(kwargs.get("messages")), + optional_params=params, + ) + elif route == "messages": + request = NativeMessagesRequest( + model=TypeAdapter(str).validate_python(kwargs.get("model")), + body=TypeAdapter(dict[str, object]).validate_python(kwargs.get("body")), + ) + elif route == "ocr": + request = NativeOCRRequest( + model=TypeAdapter(str).validate_python(kwargs.get("model")), + document=TypeAdapter(dict[str, object]).validate_python(kwargs.get("document")), + optional_params=params, + ) + elif route in {"transcription", "audio_transcription"}: + request = NativeTranscriptionRequest( + model=TypeAdapter(str).validate_python(kwargs.get("model")), + audio=TypeAdapter(dict[str, object]).validate_python(kwargs.get("audio")), + optional_params=params, + ) + else: + raise ValueError(f"unsupported native trace route: {route}") + return {"request": request, "options": options, "context": NativeRequestContext()} def collect_trace( 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 5f83e9c9950..59787822ab6 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -43,17 +43,18 @@ class RecordingMessages: self, request: NativeMessagesRequest, *, + options: object, context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( { "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, + "api_key": options.api_key, + "api_base": options.api_base, + "custom_llm_provider": options.custom_llm_provider, + "extra_headers": options.extra_headers, + "timeout_seconds": options.timeout_seconds, } ) return dict(FAKE_MESSAGES_RESPONSE) @@ -67,17 +68,18 @@ class RecordingAsyncMessages: self, request: NativeMessagesRequest, *, + options: object, context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( { "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, + "api_key": options.api_key, + "api_base": options.api_base, + "custom_llm_provider": options.custom_llm_provider, + "extra_headers": options.extra_headers, + "timeout_seconds": options.timeout_seconds, } ) return dict(FAKE_MESSAGES_RESPONSE) @@ -87,7 +89,9 @@ class ExplodingAsyncMessages: def __init__(self) -> None: self.calls = 0 - async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]: + async def __call__( + self, request: NativeMessagesRequest, *, options: object, context: NativeRequestContext + ) -> dict[str, object]: self.calls += 1 raise AssertionError("bridge must not be called") @@ -96,7 +100,9 @@ class RaisingAsyncMessages: def __init__(self) -> None: self.calls = 0 - async def __call__(self, request: NativeMessagesRequest, *, context: NativeRequestContext) -> dict[str, object]: + async def __call__( + self, request: NativeMessagesRequest, *, options: object, 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 6751cf7f06f..019022ed919 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 @@ -23,9 +23,7 @@ async def test_make_call_passes_logging_obj_to_client_post(): mock_client = AsyncMock() mock_response = MagicMock() mock_response.aiter_lines = MagicMock( - return_value=iter( - [b'data: {"type":"message_start"}\n', b'data: {"type":"message_delta"}\n'] - ) + return_value=iter([b'data: {"type":"message_start"}\n', b'data: {"type":"message_delta"}\n']) ) mock_client.post.return_value = mock_response @@ -94,9 +92,7 @@ def test_redacted_thinking_content_block_delta(): "data": "EuoBCoYBGAIiQJ/SxkPAgqxhKok29YrpJHRUJ0OT8ahCHKAwyhmRuUhtdmDX9+mn4gDzKNv3fVpQdB01zEPMzNY3QuTCd+1bdtEqQK6JuKHqdndbwpr81oVWb4wxd1GqF/7Jkw74IlQa27oobX+KuRkopr9Dllt/RDe7Se0sI1IkU7tJIAQCoP46OAwSDF51P09q67xhHlQ3ihoM2aOVlkghq/X0w8NlIjBMNvXYNbjhyrOcIg6kPFn2ed/KK7Cm5prYAtXCwkb4Wr5tUSoSHu9T5hKdJRbr6WsqEc7Lle7FULqMLZGkhqXyc3BA", }, } - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) model_response = model_response_iterator.chunk_parser(chunk=chunk) print(f"\n\nmodel_response: {model_response}\n\n") assert model_response.choices[0].delta.thinking_blocks is not None @@ -104,19 +100,14 @@ def test_redacted_thinking_content_block_delta(): print( f"\n\nmodel_response.choices[0].delta.thinking_blocks[0]: {model_response.choices[0].delta.thinking_blocks[0]}\n\n" ) - assert ( - model_response.choices[0].delta.thinking_blocks[0]["type"] - == "redacted_thinking" - ) + assert model_response.choices[0].delta.thinking_blocks[0]["type"] == "redacted_thinking" assert model_response.choices[0].delta.provider_specific_fields is not None assert "thinking_blocks" in model_response.choices[0].delta.provider_specific_fields def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) chunks = [ { "type": "content_block_start", @@ -140,17 +131,12 @@ def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): }, ] - parsed_chunks = [ - model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks - ] + parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" - for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks ) thinking_blocks = tuple( - block - for chunk in parsed_chunks - for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) expected_delta_blocks = ( {"type": "thinking", "thinking": "Step 1. "}, @@ -164,18 +150,12 @@ def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): assert reasoning_content == "Step 1. Step 2." assert thinking_blocks == (*expected_delta_blocks, expected_thinking_block) - assert parsed_chunks[1].choices[0].delta.provider_specific_fields == { - "thinking_blocks": [expected_delta_blocks[0]] - } - assert parsed_chunks[-1].choices[0].delta.provider_specific_fields == { - "thinking_blocks": [expected_thinking_block] - } + assert parsed_chunks[1].choices[0].delta.provider_specific_fields == {"thinking_blocks": [expected_delta_blocks[0]]} + assert parsed_chunks[-1].choices[0].delta.provider_specific_fields == {"thinking_blocks": [expected_thinking_block]} def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) chunks = [ { "type": "content_block_start", @@ -195,17 +175,12 @@ def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): {"type": "content_block_stop", "index": 0}, ] - parsed_chunks = [ - model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks - ] + parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" - for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks ) thinking_blocks = tuple( - block - for chunk in parsed_chunks - for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) assert reasoning_content == "Step 1. Step 2." @@ -216,9 +191,7 @@ def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) chunks = [ { "type": "content_block_start", @@ -237,17 +210,12 @@ def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): }, ] - parsed_chunks = [ - model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks - ] + parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" - for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks ) thinking_blocks = tuple( - block - for chunk in parsed_chunks - for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) assert reasoning_content == "Step 1. Step 2." @@ -258,9 +226,7 @@ def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): def test_handle_json_mode_chunk_response_format_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) response_format_tool = ChatCompletionToolCallChunk( id="tool_123", type="function", @@ -271,9 +237,7 @@ def test_handle_json_mode_chunk_response_format_tool(): index=0, ) - text, tool_use = model_response_iterator._handle_json_mode_chunk( - "", response_format_tool - ) + text, tool_use = model_response_iterator._handle_json_mode_chunk("", response_format_tool) print(f"\n\nresponse_format_tool text: {text}\n\n") print(f"\n\nresponse_format_tool tool_use: {tool_use}\n\n") @@ -282,15 +246,11 @@ def test_handle_json_mode_chunk_response_format_tool(): def test_handle_json_mode_chunk_regular_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) regular_tool = ChatCompletionToolCallChunk( id="tool_456", type="function", - function=ChatCompletionToolCallFunctionChunk( - name="get_weather", arguments='{"location": "San Francisco, CA"}' - ), + function=ChatCompletionToolCallFunctionChunk(name="get_weather", arguments='{"location": "San Francisco, CA"}'), index=0, ) @@ -304,17 +264,13 @@ def test_handle_json_mode_chunk_regular_tool(): def test_handle_json_mode_chunk_streaming_response_format_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: response_format tool with id and name, but no arguments first_chunk = ChatCompletionToolCallChunk( id="tool_123", type="function", - function=ChatCompletionToolCallFunctionChunk( - name=RESPONSE_FORMAT_TOOL_NAME, arguments="" - ), + function=ChatCompletionToolCallFunctionChunk(name=RESPONSE_FORMAT_TOOL_NAME, arguments=""), index=0, ) @@ -322,9 +278,7 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): second_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments='{"question": "What is the weather?"' - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments='{"question": "What is the weather?"'), index=0, ) @@ -332,9 +286,7 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): third_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments=', "answer": "It is sunny"}' - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments=', "answer": "It is sunny"}'), index=0, ) @@ -365,9 +317,7 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): def test_handle_json_mode_chunk_streaming_regular_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: regular tool with id and name, but no arguments first_chunk = ChatCompletionToolCallChunk( @@ -381,9 +331,7 @@ def test_handle_json_mode_chunk_streaming_regular_tool(): second_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments='{"location": "San Francisco, CA"}' - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments='{"location": "San Francisco, CA"}'), index=0, ) @@ -408,27 +356,19 @@ def test_handle_json_mode_chunk_streaming_regular_tool(): def test_response_format_tool_finish_reason(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: response_format tool response_format_tool = ChatCompletionToolCallChunk( id="tool_123", type="function", - function=ChatCompletionToolCallFunctionChunk( - name=RESPONSE_FORMAT_TOOL_NAME, arguments='{"answer": "test"}' - ), + function=ChatCompletionToolCallFunctionChunk(name=RESPONSE_FORMAT_TOOL_NAME, arguments='{"answer": "test"}'), index=0, ) # Process the tool call (should set converted_response_format_tool flag) - text, tool_use = model_response_iterator._handle_json_mode_chunk( - "", response_format_tool - ) - print( - f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n" - ) + text, tool_use = model_response_iterator._handle_json_mode_chunk("", response_format_tool) + print(f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n") # Simulate message_delta chunk with tool_use stop_reason message_delta_chunk = { @@ -447,25 +387,19 @@ def test_response_format_tool_finish_reason(): def test_regular_tool_finish_reason(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: regular tool (not response_format) regular_tool = ChatCompletionToolCallChunk( id="tool_456", type="function", - function=ChatCompletionToolCallFunctionChunk( - name="get_weather", arguments='{"location": "San Francisco, CA"}' - ), + function=ChatCompletionToolCallFunctionChunk(name="get_weather", arguments='{"location": "San Francisco, CA"}'), index=0, ) # Process the tool call (should NOT set converted_response_format_tool flag) text, tool_use = model_response_iterator._handle_json_mode_chunk("", regular_tool) - print( - f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n" - ) + print(f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n") # Simulate message_delta chunk with tool_use stop_reason message_delta_chunk = { @@ -525,9 +459,7 @@ def test_text_only_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert ( - parsed.choices[0].index == 0 - ), f"Expected index=0, got {parsed.choices[0].index}" + assert parsed.choices[0].index == 0, f"Expected index=0, got {parsed.choices[0].index}" def test_streaming_thinking_deltas_count_reasoning_tokens_in_usage(): @@ -704,9 +636,7 @@ def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinkin ] self._write_response( content_type="text/event-stream", - body="".join( - f"data: {json.dumps(event)}\n\n" for event in events - ).encode("utf-8"), + body="".join(f"data: {json.dumps(event)}\n\n" for event in events).encode("utf-8"), ) return @@ -787,13 +717,9 @@ def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinkin assert content_chunks == [answer_text] assert stream_usage is not None stream_completion_details = stream_usage["completion_tokens_details"] - assert ( - stream_completion_details["reasoning_tokens"] - == non_stream_details.reasoning_tokens - ) + assert stream_completion_details["reasoning_tokens"] == non_stream_details.reasoning_tokens assert stream_completion_details["text_tokens"] == ( - stream_usage["completion_tokens"] - - stream_completion_details["reasoning_tokens"] + stream_usage["completion_tokens"] - stream_completion_details["reasoning_tokens"] ) assert requests_seen == [ { @@ -885,9 +811,9 @@ def test_text_and_tool_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert ( - parsed.choices[0].index == 0 - ), f"Expected index=0 for chunk type {chunk.get('type')}, got {parsed.choices[0].index}" + assert parsed.choices[0].index == 0, ( + f"Expected index=0 for chunk type {chunk.get('type')}, got {parsed.choices[0].index}" + ) def test_multiple_tools_streaming_has_index_zero(): @@ -940,15 +866,11 @@ def test_multiple_tools_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert ( - parsed.choices[0].index == 0 - ), f"Expected index=0, got {parsed.choices[0].index}" + assert parsed.choices[0].index == 0, f"Expected index=0, got {parsed.choices[0].index}" def test_streaming_chunks_have_stable_ids(): - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) first_chunk = { "type": "content_block_delta", "index": 0, @@ -973,9 +895,7 @@ def test_partial_json_chunk_accumulation(): This tests the fix for https://github.com/BerriAI/litellm/issues/17473 where network fragmentation can cause SSE data to arrive in partial chunks. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) partial_chunk_1 = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hel' partial_chunk_2 = 'lo"}}' @@ -983,31 +903,21 @@ def test_partial_json_chunk_accumulation(): # First partial chunk should return None (still accumulating) result1 = iterator._parse_sse_data(f"data:{partial_chunk_1}") assert result1 is None, "First partial chunk should return None while accumulating" - assert ( - iterator.chunk_type == "accumulated_json" - ), "Should switch to accumulated_json mode" - assert ( - iterator.accumulated_json == partial_chunk_1 - ), "Should have accumulated first part" + assert iterator.chunk_type == "accumulated_json", "Should switch to accumulated_json mode" + assert iterator.accumulated_json == partial_chunk_1, "Should have accumulated first part" # Second partial chunk should complete the JSON and return a parsed result result2 = iterator._parse_sse_data(f"data:{partial_chunk_2}") assert result2 is not None, "Second chunk should return parsed result" - assert ( - iterator.accumulated_json == "" - ), "Buffer should be cleared after successful parse" - assert ( - result2.choices[0].delta.content == "Hello" - ), f"Expected 'Hello', got '{result2.choices[0].delta.content}'" + assert iterator.accumulated_json == "", "Buffer should be cleared after successful parse" + assert result2.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result2.choices[0].delta.content}'" def test_complete_json_chunk_no_accumulation(): """ Test that complete JSON chunks are parsed immediately without accumulation. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) complete_chunk = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}' @@ -1015,18 +925,14 @@ def test_complete_json_chunk_no_accumulation(): assert result is not None, "Complete chunk should return parsed result immediately" assert iterator.chunk_type == "valid_json", "Should remain in valid_json mode" assert iterator.accumulated_json == "", "Buffer should remain empty" - assert ( - result.choices[0].delta.content == "Hello" - ), f"Expected 'Hello', got '{result.choices[0].delta.content}'" + assert result.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result.choices[0].delta.content}'" def test_multiple_partial_chunks_accumulation(): """ Test that multiple partial chunks can be accumulated across several iterations. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Split a JSON chunk into three parts part1 = '{"type":"content_block_del' @@ -1054,17 +960,11 @@ def test_accumulated_json_partial_fragment_returns_none_without_parsing(): unlike Vertex which already deferred parsing until the buffer could close. A fragment that can't close a JSON value must not trigger a decode attempt. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" - with patch.object( - json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode - ) as spy: - result = iterator._handle_accumulated_json_chunk( - '{"type":"content_block_delta","index":0,"delta":' - ) + with patch.object(json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode) as spy: + result = iterator._handle_accumulated_json_chunk('{"type":"content_block_delta","index":0,"delta":') assert result is None assert spy.call_count == 0, "incomplete buffer should not be parsed" @@ -1076,21 +976,15 @@ def test_accumulated_json_does_not_reparse_every_fragment(): fragment. """ text = "x" * 200_000 - blob = json.dumps( - {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}} - ) + blob = json.dumps({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}}) fragments = [blob[i : i + 4096] for i in range(0, len(blob), 4096)] assert len(fragments) > 10, "need a multi-fragment payload to exercise the bug" - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" parsed = None - with patch.object( - json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode - ) as spy: + with patch.object(json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode) as spy: for fragment in fragments: out = iterator._handle_accumulated_json_chunk(fragment) if out is not None: @@ -1114,9 +1008,7 @@ def test_accumulated_json_concatenated_envelopes_do_not_wedge(): and keeps the remainder, so both values surface across two calls. """ obj = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"a"}}' - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" first = iterator._handle_accumulated_json_chunk(obj + obj) @@ -1137,9 +1029,7 @@ def test_accumulated_json_heuristic_passes_but_value_still_incomplete(): heuristic must let the parse attempt through, and pop_next_value finding nothing must propagate as None rather than raising. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" result = iterator._handle_accumulated_json_chunk('{"type": {"nested": 1}') @@ -1153,9 +1043,7 @@ def test_accumulated_json_setter_and_sync_end_of_stream_drain(): underlying stream ends, instead of being silently dropped. """ obj = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"a"}}' - iterator = ModelResponseIterator( - streaming_response=iter([]), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=iter([]), sync_stream=True, json_mode=False) iterator.chunk_type = "accumulated_json" iterator.accumulated_json = obj # exercises the setter @@ -1170,9 +1058,7 @@ def test_accumulated_json_async_end_of_stream_drain(): import asyncio obj = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"a"}}' - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) iterator.chunk_type = "accumulated_json" iterator.accumulated_json = obj mock_async_iterator = MagicMock() @@ -1194,9 +1080,7 @@ def test_web_search_tool_result_no_extra_tool_calls(): The issue was that web_search_tool_result blocks have input_json_delta events with {} that were incorrectly being converted to tool calls. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate the streaming sequence: # 1. server_tool_use block starts (web_search) @@ -1271,9 +1155,7 @@ def test_web_search_tool_result_no_extra_tool_calls(): # Should have exactly 2 tool calls: # 1. From content_block_start (server_tool_use) with id and name # 2. From content_block_delta with the actual query - assert ( - len(tool_calls_emitted) == 2 - ), f"Expected 2 tool calls, got {len(tool_calls_emitted)}" + assert len(tool_calls_emitted) == 2, f"Expected 2 tool calls, got {len(tool_calls_emitted)}" # First tool call should have the id and name assert tool_calls_emitted[0]["id"] == "srvtoolu_01ABC123" @@ -1289,9 +1171,7 @@ def test_current_content_block_type_tracking(): """ Test that current_content_block_type is properly tracked and reset. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Initially should be None assert iterator.current_content_block_type is None @@ -1344,9 +1224,7 @@ def test_web_search_tool_result_captured_in_provider_specific_fields(): The web_search_tool_result content comes ALL AT ONCE in content_block_start, not in deltas, so we need to capture it there. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate the streaming sequence with web_search_tool_result chunks = [ @@ -1417,23 +1295,15 @@ def test_web_search_tool_result_captured_in_provider_specific_fields(): and parsed.choices[0].delta.provider_specific_fields and "web_search_results" in parsed.choices[0].delta.provider_specific_fields ): - web_search_results = parsed.choices[0].delta.provider_specific_fields[ - "web_search_results" - ] + web_search_results = parsed.choices[0].delta.provider_specific_fields["web_search_results"] # Verify web_search_results was captured assert web_search_results is not None, "web_search_results should be captured" assert len(web_search_results) == 1, "Should have 1 web_search_tool_result block" - assert ( - web_search_results[0]["type"] == "web_search_tool_result" - ), "Block type should be web_search_tool_result" - assert ( - web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123" - ), "tool_use_id should match" + assert web_search_results[0]["type"] == "web_search_tool_result", "Block type should be web_search_tool_result" + assert web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123", "tool_use_id should match" assert len(web_search_results[0]["content"]) == 2, "Should have 2 search results" - assert ( - web_search_results[0]["content"][0]["title"] == "Fun Otter Facts" - ), "First result title should match" + assert web_search_results[0]["content"][0]["title"] == "Fun Otter Facts", "First result title should match" def test_web_fetch_tool_result_captured_in_provider_specific_fields(): @@ -1447,9 +1317,7 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields(): The web_fetch_tool_result content comes ALL AT ONCE in content_block_start, not in deltas, so we need to capture it there. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate the streaming sequence with web_fetch_tool_result chunks = [ @@ -1520,25 +1388,15 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields(): and parsed.choices[0].delta.provider_specific_fields and "web_search_results" in parsed.choices[0].delta.provider_specific_fields ): - web_search_results = parsed.choices[0].delta.provider_specific_fields[ - "web_search_results" - ] + web_search_results = parsed.choices[0].delta.provider_specific_fields["web_search_results"] # Verify web_fetch_tool_result was captured (stored in web_search_results list) assert web_search_results is not None, "web_search_results should be captured" assert len(web_search_results) == 1, "Should have 1 web_fetch_tool_result block" - assert ( - web_search_results[0]["type"] == "web_fetch_tool_result" - ), "Block type should be web_fetch_tool_result" - assert ( - web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123" - ), "tool_use_id should match" - assert ( - web_search_results[0]["content"]["url"] == "https://example.com" - ), "URL should match" - assert ( - web_search_results[0]["content"]["content"]["title"] == "Example Page" - ), "Title should match" + assert web_search_results[0]["type"] == "web_fetch_tool_result", "Block type should be web_fetch_tool_result" + assert web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123", "tool_use_id should match" + assert web_search_results[0]["content"]["url"] == "https://example.com", "URL should match" + assert web_search_results[0]["content"]["content"]["title"] == "Example Page", "Title should match" def test_web_fetch_tool_result_no_extra_tool_calls(): @@ -1551,9 +1409,7 @@ def test_web_fetch_tool_result_no_extra_tool_calls(): The issue was that web_fetch_tool_result blocks have input_json_delta events with {} that were incorrectly being converted to tool calls. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # to verify it doesn't emit tool calls chunks = [ @@ -1597,9 +1453,9 @@ def test_web_fetch_tool_result_no_extra_tool_calls(): tool_call_count += 1 # Should have 0 tool calls - web_fetch_tool_result should not emit tool calls - assert ( - tool_call_count == 0 - ), f"Expected 0 tool calls, got {tool_call_count}. web_fetch_tool_result should not emit tool calls" + assert tool_call_count == 0, ( + f"Expected 0 tool calls, got {tool_call_count}. web_fetch_tool_result should not emit tool calls" + ) def test_container_in_provider_specific_fields_streaming(): @@ -1609,9 +1465,7 @@ def test_container_in_provider_specific_fields_streaming(): When container with skills is used, the container field should be present in the provider_specific_fields of the message_delta chunk. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate streaming chunks chunks = [ @@ -1679,20 +1533,12 @@ def test_container_in_provider_specific_fields_streaming(): and parsed.choices[0].delta.provider_specific_fields and "container" in parsed.choices[0].delta.provider_specific_fields ): - container_field = parsed.choices[0].delta.provider_specific_fields[ - "container" - ] + container_field = parsed.choices[0].delta.provider_specific_fields["container"] # Verify container was captured - assert ( - container_field is not None - ), "container should be captured in provider_specific_fields" - assert ( - container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p" - ), "container id should match" - assert ( - container_field["expires_at"] == "2025-12-16T04:57:16.913181Z" - ), "expires_at should match" + assert container_field is not None, "container should be captured in provider_specific_fields" + assert container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p", "container id should match" + assert container_field["expires_at"] == "2025-12-16T04:57:16.913181Z", "expires_at should match" assert len(container_field["skills"]) == 1, "Should have 1 skill" assert container_field["skills"][0]["skill_id"] == "pptx", "skill_id should be pptx" assert container_field["skills"][0]["version"] == "20251013", "version should match" @@ -1705,9 +1551,7 @@ def test_container_in_provider_specific_fields_non_streaming(): When container with skills is used in non-streaming, the container field should be present in the provider_specific_fields of the response. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) # Simulate a message_delta chunk with container (as it would appear in non-streaming) message_delta_chunk = { @@ -1743,21 +1587,13 @@ def test_container_in_provider_specific_fields_non_streaming(): # Verify container is in provider_specific_fields assert model_response.choices[0].delta.provider_specific_fields is not None assert "container" in model_response.choices[0].delta.provider_specific_fields - container_field = model_response.choices[0].delta.provider_specific_fields[ - "container" - ] + container_field = model_response.choices[0].delta.provider_specific_fields["container"] assert container_field["id"] == "container_abc123xyz", "container id should match" - assert ( - container_field["expires_at"] == "2025-12-20T10:30:00.000000Z" - ), "expires_at should match" + assert container_field["expires_at"] == "2025-12-20T10:30:00.000000Z", "expires_at should match" assert len(container_field["skills"]) == 2, "Should have 2 skills" - assert ( - container_field["skills"][0]["skill_id"] == "code_execution" - ), "First skill_id should be code_execution" - assert ( - container_field["skills"][1]["skill_id"] == "pptx" - ), "Second skill_id should be pptx" + assert container_field["skills"][0]["skill_id"] == "code_execution", "First skill_id should be code_execution" + assert container_field["skills"][1]["skill_id"] == "pptx", "Second skill_id should be pptx" def test_container_absent_when_not_provided(): @@ -1766,9 +1602,7 @@ def test_container_absent_when_not_provided(): This ensures we don't add empty or None container fields. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) # message_delta without container message_delta_chunk = { @@ -1787,9 +1621,9 @@ def test_container_absent_when_not_provided(): # Verify container is NOT in provider_specific_fields when not provided if model_response.choices[0].delta.provider_specific_fields: - assert ( - "container" not in model_response.choices[0].delta.provider_specific_fields - ), "container should not be present when not provided in delta" + assert "container" not in model_response.choices[0].delta.provider_specific_fields, ( + "container should not be present when not provided in delta" + ) def test_streaming_code_execution_produces_code_interpreter_results(): @@ -1985,8 +1819,7 @@ def test_streaming_multiple_code_executions_no_duplicates(): # Second (final) emission: cumulative list with BOTH results # This is what stream_chunk_builder will pick as "last value wins" assert len(emissions[1]) == 2, ( - f"Expected final emission to have 2 results, got {len(emissions[1])}. " - f"IDs: {[r.id for r in emissions[1]]}" + f"Expected final emission to have 2 results, got {len(emissions[1])}. IDs: {[r.id for r in emissions[1]]}" ) assert emissions[1][0].id == "srvtoolu_01AAA" assert emissions[1][0].code == "echo first" @@ -2150,9 +1983,7 @@ def test_empty_output_produces_null_outputs(): assert code_results is not None, "No code_interpreter_results emitted" assert len(code_results) == 1 assert code_results[0].id == "srvtoolu_01AAA" - assert ( - code_results[0].outputs is None - ), f"Expected outputs=None for empty execution, got {code_results[0].outputs}" + assert code_results[0].outputs is None, f"Expected outputs=None for empty execution, got {code_results[0].outputs}" def test_non_bash_tool_result_skipped(): @@ -2215,12 +2046,10 @@ def test_non_bash_tool_result_skipped(): code_results = psf["code_interpreter_results"] # code_interpreter_results should be emitted but empty (no bash results) - assert ( - code_results is not None - ), "Expected code_interpreter_results key to be emitted" - assert ( - len(code_results) == 0 - ), f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}" + assert code_results is not None, "Expected code_interpreter_results key to be emitted" + assert len(code_results) == 0, ( + f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}" + ) class TestRustChatCompletionsHook: @@ -2257,13 +2086,9 @@ class TestRustChatCompletionsHook: from litellm.rust_bridge import chat_completions as bridge monkeypatch.setenv("LITELLM_RUST", "1") - bridge.set_rust_chat_completions( - chat_completions=None, achat_completions=None, decline=None - ) + bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) yield - bridge.set_rust_chat_completions( - chat_completions=None, achat_completions=None, decline=None - ) + bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) @staticmethod def _completion_kwargs(**overrides): @@ -2309,17 +2134,17 @@ class TestRustChatCompletionsHook: seen["gate"].append(kwargs) return decline_reason - def native(request, *, context): + def native(request, *, options, 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, + "optional_params": request.optional_params, + "api_key": options.api_key, + "api_base": options.api_base, + "custom_llm_provider": options.custom_llm_provider, + "extra_headers": options.extra_headers, + "timeout_seconds": options.timeout_seconds, } ) if sync_error is not None: @@ -2372,9 +2197,7 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion seen = self._inject() - AnthropicChatCompletion().completion( - **self._completion_kwargs(optional_params={"max_tokens": 7}) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(optional_params={"max_tokens": 7})) assert seen["call"][0]["optional_params"]["max_tokens"] == 7 def test_without_the_opt_in_the_core_is_never_consulted(self, monkeypatch): @@ -2383,15 +2206,14 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.transformation import AnthropicConfig seen = self._inject() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ) as transform, patch.object( - AnthropicChatCompletion, "acompletion_function" + with ( + patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ) as transform, + patch.object(AnthropicChatCompletion, "acompletion_function"), ): try: - AnthropicChatCompletion().completion( - **self._completion_kwargs(litellm_params={}) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(litellm_params={})) except Exception: # The Python path goes on to make an HTTP call; reaching it is # the assertion, so the network failure below is expected. @@ -2405,9 +2227,7 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.transformation import AnthropicConfig seen = self._inject(decline_reason="unrecognized request parameter") - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: AnthropicChatCompletion().completion(**self._completion_kwargs()) except Exception: @@ -2420,9 +2240,7 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.transformation import AnthropicConfig seen = self._inject() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: AnthropicChatCompletion().completion( **self._completion_kwargs(optional_params={"max_tokens": 16, "stream": True}) @@ -2436,9 +2254,7 @@ class TestRustChatCompletionsHook: seen = self._inject() logging_obj = MagicMock() - AnthropicChatCompletion().completion( - **self._completion_kwargs(logging_obj=logging_obj) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) assert logging_obj.pre_call.call_count == 1 assert len(seen["call"]) == 1 @@ -2452,9 +2268,7 @@ class TestRustChatCompletionsHook: self._inject() logging_obj = MagicMock() - AnthropicChatCompletion().completion( - **self._completion_kwargs(logging_obj=logging_obj) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) assert logging_obj.post_call.call_count == 1 logged = logging_obj.post_call.call_args.kwargs["original_response"] @@ -2475,22 +2289,16 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(request, *, context): + def declining_native(request, *, options, context): raise _Declined("blank message text") monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, chat_completions=declining_native - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, chat_completions=declining_native) logging_obj, calls = self._recording_logging_obj() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: - AnthropicChatCompletion().completion( - **self._completion_kwargs(logging_obj=logging_obj) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) except Exception: # The Python path goes on to make an HTTP call; the log count is # the assertion, so a failure past this point is expected. @@ -2512,24 +2320,18 @@ class TestRustChatCompletionsHook: monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - async def declining_native(request, *, context): + async def declining_native(request, *, options, context): raise _Declined("blank message text") - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, achat_completions=declining_native - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=declining_native) sentinel = object() async def python_path(**_kwargs): return sentinel - with patch.object( - AnthropicChatCompletion, "acompletion_function", side_effect=python_path - ) as python_call: - result = await AnthropicChatCompletion().completion( - **self._completion_kwargs(acompletion=True) - ) + with patch.object(AnthropicChatCompletion, "acompletion_function", side_effect=python_path) as python_call: + result = await AnthropicChatCompletion().completion(**self._completion_kwargs(acompletion=True)) assert result is sentinel assert python_call.called, "a failing rust call must re-enter the python path" @@ -2539,23 +2341,18 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion from litellm.rust_bridge import chat_completions as bridge - async def native(request, *, context): + async def native(request, *, options, context): return dict(self.RUST_RESPONSE) - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, achat_completions=native - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=native) with patch.object(AnthropicChatCompletion, "acompletion_function") as python_call: - result = await AnthropicChatCompletion().completion( - **self._completion_kwargs(acompletion=True) - ) + result = await AnthropicChatCompletion().completion(**self._completion_kwargs(acompletion=True)) assert result.choices[0].message.content == "hello from rust" assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} assert not python_call.called - def test_pre_call_logging_fires_once_when_the_sync_rust_call_declines(self, monkeypatch): """One request, one pre_call, on the synchronous path too. Without the suppression the Python path logs a second time for the same attempt.""" @@ -2572,30 +2369,22 @@ class TestRustChatCompletionsHook: monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - def declining_native(request, *, context): + def declining_native(request, *, options, context): raise _Declined("blank message text") - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, chat_completions=declining_native - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, chat_completions=declining_native) logging_obj, calls = self._recording_logging_obj() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: - AnthropicChatCompletion().completion( - **self._completion_kwargs(logging_obj=logging_obj) - ) + AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) except Exception: # The Python path goes on to make an HTTP call; the log count is # the assertion, so a failure past this point is expected. pass assert len(calls["pre_call"]) == 1 - assert calls["pre_call"][0]["additional_args"]["complete_input_dict"]["model"] == ( - "claude-sonnet-4-5" - ) + assert calls["pre_call"][0]["additional_args"]["complete_input_dict"]["model"] == ("claude-sonnet-4-5") def test_pre_call_logging_still_fires_when_rust_is_not_involved(self, monkeypatch): """The suppression must not swallow the log on the ordinary path.""" @@ -2605,9 +2394,7 @@ class TestRustChatCompletionsHook: self._inject() logging_obj, calls = self._recording_logging_obj() - with patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ): + with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): try: AnthropicChatCompletion().completion( **self._completion_kwargs(litellm_params={}, logging_obj=logging_obj) 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 c372b295d64..c614e040d4c 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 @@ -10,8 +10,8 @@ from unittest.mock import MagicMock, patch import httpx import pytest - from botocore.credentials import Credentials + from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.rust_bridge import chat_completions as bridge @@ -49,13 +49,9 @@ RESOLVED_CREDENTIALS = Credentials( @pytest.fixture(autouse=True) def reset_bridge(monkeypatch): monkeypatch.setenv("LITELLM_RUST", "1") - bridge.set_rust_chat_completions( - chat_completions=None, achat_completions=None, decline=None - ) + bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) yield - bridge.set_rust_chat_completions( - chat_completions=None, achat_completions=None, decline=None - ) + bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) def _inject(*, decline_reason=None, error: Exception | None = None): @@ -65,17 +61,18 @@ def _inject(*, decline_reason=None, error: Exception | None = None): seen["gate"].append(kwargs) return decline_reason - def native(request, *, context): + def native(request, *, options, 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, + "optional_params": request.optional_params, + "bedrock": options.bedrock, + "api_key": options.api_key, + "api_base": options.api_base, + "custom_llm_provider": options.custom_llm_provider, + "extra_headers": options.extra_headers, + "timeout_seconds": options.timeout_seconds, } ) if error is not None: @@ -137,20 +134,18 @@ def test_the_core_receives_the_credentials_this_handler_already_resolved(): seen = _inject() _run() - params = seen["call"][0]["optional_params"] - assert params["aws_access_key_id"] == "AKIARESOLVED" - assert params["aws_secret_access_key"] == "resolved-secret" - assert params["aws_session_token"] == "resolved-token" - assert params["aws_region_name"] == "us-east-1" + bedrock = seen["call"][0]["bedrock"] + assert bedrock.aws_access_key_id == "AKIARESOLVED" + assert bedrock.aws_secret_access_key == "resolved-secret" + assert bedrock.aws_session_token == "resolved-token" + assert bedrock.aws_region_name == "us-east-1" def test_the_core_receives_the_converse_url_this_handler_already_built(): seen = _inject() _run() - assert seen["call"][0]["api_base"].endswith( - "/model/anthropic.claude-sonnet-4-5-v1%3A0/converse" - ) + assert seen["call"][0]["api_base"].endswith("/model/anthropic.claude-sonnet-4-5-v1%3A0/converse") assert "bedrock-runtime.us-east-1.amazonaws.com" in seen["call"][0]["api_base"] @@ -218,12 +213,10 @@ 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(request, *, context): + async def declining_native(request, *, options, context): raise _Declined("blank message text") - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, achat_completions=declining_native - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=declining_native) sentinel = object() @@ -231,16 +224,10 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): return sentinel with ( - patch.object( - BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS - ), - patch.object( - BedrockConverseLLM, "async_completion", side_effect=python_path - ) as python_call, + patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS), + patch.object(BedrockConverseLLM, "async_completion", side_effect=python_path) as python_call, ): - result = await BedrockConverseLLM().completion( - **_completion_kwargs(acompletion=True) - ) + result = await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True)) assert result is sentinel assert python_call.called, "a failing rust call must re-enter the python path" @@ -248,22 +235,16 @@ 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(request, *, context): + async def native(request, *, options, context): return dict(RUST_RESPONSE) - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, achat_completions=native - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=native) with ( - patch.object( - BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS - ), + patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS), patch.object(BedrockConverseLLM, "async_completion") as python_call, ): - result = await BedrockConverseLLM().completion( - **_completion_kwargs(acompletion=True) - ) + result = await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True)) assert result.choices[0].message.content == "hello from rust" assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} @@ -282,7 +263,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(request, *, context): + async def declining_native(request, *, options, context): raise _Declined("blank message text") logging_obj = MagicMock() @@ -294,19 +275,11 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines(): with ( patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()), - patch.object( - BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS - ), - patch.object( - BedrockConverseLLM, "async_completion", side_effect=python_path - ), + patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS), + patch.object(BedrockConverseLLM, "async_completion", side_effect=python_path), ): - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, achat_completions=declining_native - ) - await BedrockConverseLLM().completion( - **_completion_kwargs(acompletion=True, logging_obj=logging_obj) - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=declining_native) + await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True, logging_obj=logging_obj)) assert logging_obj.pre_call.call_count == 1 assert served and served[0]["skip_pre_call_logging"] is True @@ -395,15 +368,13 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(request, *, context): + def declining_native(request, *, options, context): raise _Declined("blank message text") logging_obj = MagicMock() with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()): - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, chat_completions=declining_native - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, chat_completions=declining_native) response = _run( logging_obj=logging_obj, client=_sync_client_returning_converse_response(), @@ -449,20 +420,14 @@ 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(request, *, context): + async def native(request, *, options, context): return dict(RUST_RESPONSE) - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, achat_completions=native - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=native) logging_obj = MagicMock() - with patch.object( - BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS - ): - await BedrockConverseLLM().completion( - **_completion_kwargs(acompletion=True, logging_obj=logging_obj) - ) + with patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS): + await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True, logging_obj=logging_obj)) assert logging_obj.post_call.call_count == 1 logged = logging_obj.post_call.call_args.kwargs["original_response"] @@ -481,15 +446,13 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(request, *, context): + def declining_native(request, *, options, context): raise _Declined("blank message text") logging_obj, calls = _recording_logging_obj() with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()): - bridge.set_rust_chat_completions( - decline=lambda **_kwargs: None, chat_completions=declining_native - ) + bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, chat_completions=declining_native) response = _run( logging_obj=logging_obj, client=_sync_client_returning_converse_response(), @@ -523,9 +486,11 @@ def test_the_rust_opt_in_needs_no_sigv4_principal(): response = _run(credentials=None, api_key="bedrock-bearer-token") assert response.choices[0].message.content == "hello from rust" - params = seen["call"][0]["optional_params"] - assert not {"aws_access_key_id", "aws_secret_access_key", "aws_session_token"} & params.keys() - assert params["aws_region_name"] == "us-east-1" + bedrock = seen["call"][0]["bedrock"] + assert bedrock.aws_access_key_id is None + assert bedrock.aws_secret_access_key is None + assert bedrock.aws_session_token is None + assert bedrock.aws_region_name == "us-east-1" assert seen["call"][0]["api_key"] == "bedrock-bearer-token" diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 8720e73e77d..2a9b314ac85 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -12,7 +12,13 @@ import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge import configuration -from litellm.rust_bridge.request import NativeOCRRequest, NativeRequestContext, NativeRequestOptions, PreparedNativeCall +from litellm.rust_bridge.request import ( + NativeOCRRequest, + NativeRequestContext, + NativeRequestOptions, + NativeVertexOptions, + PreparedNativeCall, +) from litellm.rust_bridge.timeouts import timeout_to_seconds # `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` @@ -51,18 +57,20 @@ class RecordingBridge: self, request: NativeOCRRequest, *, + options: NativeRequestOptions, context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( { "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, + "api_key": options.api_key, + "api_base": options.api_base, + "custom_llm_provider": options.custom_llm_provider, + "extra_headers": options.extra_headers, + "optional_params": request.optional_params, + "vertex": options.vertex, + "timeout_seconds": options.timeout_seconds, } ) return dict(FAKE_OCR_RESPONSE) @@ -78,18 +86,20 @@ class RecordingAsyncBridge: self, request: NativeOCRRequest, *, + options: NativeRequestOptions, context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( { "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, + "api_key": options.api_key, + "api_base": options.api_base, + "custom_llm_provider": options.custom_llm_provider, + "extra_headers": options.extra_headers, + "optional_params": request.optional_params, + "vertex": options.vertex, + "timeout_seconds": options.timeout_seconds, } ) return dict(FAKE_OCR_RESPONSE) @@ -100,6 +110,7 @@ class RaisingBridge: self, request: NativeOCRRequest, *, + options: NativeRequestOptions, context: NativeRequestContext, ) -> dict[str, object]: raise RuntimeError("bridge failed") @@ -110,6 +121,7 @@ class RaisingAsyncBridge: self, request: NativeOCRRequest, *, + options: NativeRequestOptions, context: NativeRequestContext, ) -> dict[str, object]: raise RuntimeError("bridge failed") @@ -381,13 +393,13 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): model="mistral-ocr-latest", document=DOCUMENT, optional_params={"include_image_base64": True, "pages": [0]}, - options=NativeRequestOptions( - api_key="sk-test", - api_base="https://proxy.internal", - custom_llm_provider="mistral", - extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"}, - timeout_seconds=12.5, - ), + ), + options=NativeRequestOptions( + api_key="sk-test", + api_base="https://proxy.internal", + custom_llm_provider="mistral", + extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"}, + timeout_seconds=12.5, ), ), fallback=lambda: pytest.fail("unexpected Python fallback"), @@ -410,6 +422,7 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): "x-trace-id": "trace-1", }, "optional_params": {"include_image_base64": True, "pages": [0]}, + "vertex": None, "timeout_seconds": 12.5, } @@ -431,11 +444,11 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): model="mistral-ocr-maas", document=DOCUMENT, optional_params={}, - options=NativeRequestOptions( - custom_llm_provider="vertex_ai", - provider_connection={"vertex_project": "project-1"}, - timeout_seconds=42.0, - ), + ), + options=NativeRequestOptions( + custom_llm_provider="vertex_ai", + vertex=NativeVertexOptions(project="project-1"), + timeout_seconds=42.0, ), ), fallback=unexpected_fallback, @@ -453,7 +466,8 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): "api_base": None, "custom_llm_provider": "vertex_ai", "extra_headers": None, - "optional_params": {"vertex_project": "project-1"}, + "optional_params": {}, + "vertex": NativeVertexOptions(project="project-1"), "timeout_seconds": 42.0, } @@ -489,6 +503,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): "x-trace-id": "trace-1", }, "optional_params": {"include_image_base64": True}, + "vertex": NativeVertexOptions(), "timeout_seconds": 12.5, } @@ -573,11 +588,8 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): resolve_api_key=lambda _name: None, ) - assert bridge.calls[0]["optional_params"] == { - "include_image_base64": True, - "vertex_project": "project-1", - "vertex_location": "us-central1", - } + assert bridge.calls[0]["optional_params"] == {"include_image_base64": True} + assert bridge.calls[0]["vertex"] == NativeVertexOptions(project="project-1", location="us-central1") def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager(): @@ -601,8 +613,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana resolve_api_key=_resolver, ) - assert bridge.calls[0]["optional_params"]["vertex_project"] == "project-from-secret" - assert bridge.calls[0]["optional_params"]["vertex_location"] == "us-east5" + assert bridge.calls[0]["vertex"] == NativeVertexOptions(project="project-from-secret", location="us-east5") def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index cbf183dea91..aea5137db66 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -33,6 +33,7 @@ class _FakeNativeBridge: cls, request: NativeResponsesWebSocketRequest, *, + options: object, context: NativeRequestContext, ) -> _FakeNativeConnection: return _FakeNativeConnection() @@ -103,6 +104,7 @@ class _FailingNativeBridge: cls, request: NativeResponsesWebSocketRequest, *, + options: object, 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 5bc754bb724..a24d5ba1ae9 100644 --- a/tests/test_litellm/rust_bridge/native_route_wheel_test.py +++ b/tests/test_litellm/rust_bridge/native_route_wheel_test.py @@ -183,34 +183,74 @@ def _record(name: str, fields: dict[str, object]) -> object: 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} + bedrock: Final = _record( + "BedrockOptions", + { + "aws_access_key_id": params.get("aws_access_key_id"), + "aws_secret_access_key": params.get("aws_secret_access_key"), + "aws_session_token": None, + "aws_region_name": params.get("aws_region_name"), + "aws_session_name": None, + "aws_profile_name": None, + "aws_role_name": None, + "aws_web_identity_token": None, + "aws_sts_endpoint": None, + "aws_external_id": None, + "aws_bedrock_runtime_endpoint": None, + "request_metadata_fields": (), + "request_metadata": None, + }, + ) + options: Final = _record( + "RequestOptions", + { + "api_key": inputs.get("api_key"), + "api_base": inputs.get("api_base"), + "custom_llm_provider": inputs.get("custom_llm_provider"), + "extra_headers": inputs.get("extra_headers"), + "extra_query": None, + "timeout_seconds": inputs.get("timeout_seconds"), + "bedrock": bedrock, + "anthropic": None, + "vertex": None, + }, + ) + request_params: Final = {"language": params.get("language")} if route == "transcription" else params + request: Final = _record( + "Request", + { + **{ + key: value for key, value in inputs.items() if key in {"model", "document", "audio", "body", "messages"} + }, + **({"optional_params": request_params} if route != "messages" else {}), + }, + ) + context: Final = _record( + "RequestContext", + { + "litellm_call_id": "native-wheel-call", + "trace_id": None, + "request_model": inputs["model"], + "attribution": _record( + "Attribution", + { + "user_api_key_hash": None, + "user_api_key_user_id": "native-user", + "user_api_key_team_id": None, + }, + ), + "capabilities": _record( + "Capabilities", + { + "stream": False, + "has_agentic_hook": False, + "has_custom_client": False, + "request_format": None, + }, + ), + }, + ) + return {"request": request, "options": options, "context": context} def assert_success(route: str, response: object) -> None: @@ -270,12 +310,7 @@ async def exercise_async(native: object, api_base: str) -> None: async def exercise_async_concurrency(native: object, api_base: str) -> None: responses: Final = await asyncio.wait_for( - asyncio.gather( - *( - native.amessages(**route_kwargs("messages", api_base, "success")) - for _ in range(32) - ) - ), + asyncio.gather(*(native.amessages(**route_kwargs("messages", api_base, "success")) for _ in range(32))), timeout=15, ) for response in responses: diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index 4b07099cc78..f81b7050134 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -12,6 +12,12 @@ import pytest import litellm from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge import chat_completions as bridge +from litellm.rust_bridge.request import ( + NativeBedrockOptions, + NativeRequestCapabilities, + NativeRequestContext, + anthropic_options, +) from litellm.types.utils import ModelResponse RUST_RESPONSE = { @@ -95,16 +101,16 @@ class _RecordingCall: self.error = error self.calls: list[dict] = [] - def __call__(self, request, *, context): - self.calls.append({"request": request, "context": context}) + def __call__(self, request, *, options, context): + self.calls.append({"request": request, "options": options, "context": context}) if self.error is not None: raise self.error return self.result class _RecordingAsyncCall(_RecordingCall): - async def __call__(self, request, *, context): - return _RecordingCall.__call__(self, request, context=context) + async def __call__(self, request, *, options, context): + return _RecordingCall.__call__(self, request, options=options, context=context) def _accepts(**overrides) -> bool: @@ -264,7 +270,7 @@ class TestSyncCall: native = _RecordingCall() bridge.set_rust_chat_completions(chat_completions=native) bridge.chat_completions(**_call_kwargs(ModelResponse())) - assert native.calls[0]["request"].options.timeout_seconds == 30.0 + assert native.calls[0]["options"].timeout_seconds == 30.0 def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): _hide_native_bridge(monkeypatch) @@ -420,13 +426,48 @@ def test_provider_credentials_are_separate_from_chat_body_params(): kwargs = _call_kwargs(ModelResponse()) kwargs["optional_params"] = { "max_tokens": 32, - "aws_access_key_id": "test-access-key", - "aws_secret_access_key": "test-secret-key", } + kwargs["bedrock"] = NativeBedrockOptions( + aws_access_key_id="test-access-key", + aws_secret_access_key="test-secret-key", + ) bridge.chat_completions(**kwargs) request = native.calls[0]["request"] + options = native.calls[0]["options"] 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", + assert options.bedrock.aws_access_key_id == "test-access-key" + assert options.bedrock.aws_secret_access_key == "test-secret-key" + + +def test_provider_payload_extensions_cross_the_boundary_without_partitioning(): + native = _RecordingCall() + bridge.set_rust_chat_completions(chat_completions=native) + configuration.rust(True) + extensions = { + "vendor_object": {"nested": None}, + "vendor_array": [1, "two", False], + "vendor_scalar": 0.25, + "extra_body": {"temperature": 0.2, "config": {"replacement": True}}, } + + kwargs = _call_kwargs(ModelResponse()) + kwargs["optional_params"] = extensions + bridge.chat_completions(**kwargs) + + assert native.calls[0]["request"].optional_params == extensions + + +def test_typed_capability_and_provider_metadata_facts_are_isolated(): + context = NativeRequestContext( + capabilities=NativeRequestCapabilities( + stream=True, + has_agentic_hook=True, + has_custom_client=True, + request_format="native", + ) + ) + anthropic = anthropic_options({"metadata": {"user_id": "user-123", "ignored": object()}}) + + assert context.capabilities.request_format == "native" + assert context.capabilities.has_agentic_hook is True + assert anthropic.user_id == "user-123" diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 72c60463fc5..57c936c170b 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -4,7 +4,11 @@ import pytest import litellm from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch -from litellm.rust_bridge.request import NativeRequestContext, NativeTranscriptionRequest +from litellm.rust_bridge.request import ( + NativeRequestContext, + NativeRequestOptions, + NativeTranscriptionRequest, +) rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") @@ -17,9 +21,17 @@ class SyncBridge: self, request: NativeTranscriptionRequest, *, + options: NativeRequestOptions, context: NativeRequestContext, ) -> dict[str, object]: - self.calls.append({"model": request.model, "audio": request.audio, "optional_params": {**request.optional_params, **(request.options.provider_connection or {})}}) + self.calls.append( + { + "model": request.model, + "audio": request.audio, + "optional_params": request.optional_params, + "bedrock": options.bedrock, + } + ) return {"text": "hello"} @@ -28,6 +40,7 @@ class AsyncBridge: self, request: NativeTranscriptionRequest, *, + options: NativeRequestOptions, context: NativeRequestContext, ) -> dict[str, object]: return {"text": "async"} @@ -111,7 +124,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 request, *, context: {"text": "rust"}, + transcription=lambda request, *, options, context: {"text": "rust"}, atranscription=None, ) try: @@ -127,7 +140,9 @@ 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(request: NativeTranscriptionRequest, *, context: NativeRequestContext) -> dict[str, object]: + async def rust_response( + request: NativeTranscriptionRequest, *, options: object, context: NativeRequestContext + ) -> dict[str, object]: return {"text": "rust"} rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response) From 986296b3a1f5f270f8c9d961669570dda9a2a479 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 21:35:34 -0700 Subject: [PATCH 3/4] fix(native): preserve typed capability context --- .../crates/core/src/request_context.rs | 5 + .../crates/python-bridge/src/marshal.rs | 15 + litellm/llms/anthropic/chat/handler.py | 200 ++++--- litellm/llms/bedrock/chat/converse_handler.py | 318 ++++++------ litellm/llms/custom_httpx/llm_http_handler.py | 7 + litellm/ocr/main.py | 282 +++------- litellm/rust_bridge/chat_completions.py | 173 +++---- litellm/rust_bridge/messages.py | 83 +-- litellm/rust_bridge/ocr.py | 308 +++++++++-- litellm/rust_bridge/request.py | 5 + litellm/rust_bridge/responses_websocket.py | 73 ++- litellm/rust_bridge/transcription.py | 86 ++-- .../test_rust_bridge_messages.py | 191 ++++++- .../chat/test_anthropic_chat_handler.py | 486 +++++++++++++----- .../chat/test_bedrock_converse_handler.py | 124 +++-- tests/test_litellm/ocr/test_rust_bridge.py | 175 ++----- .../responses/test_rust_bridge_websocket.py | 96 +++- .../rust_bridge/test_chat_completions.py | 244 ++------- .../test_audio_transcription_rust_bridge.py | 47 +- 19 files changed, 1595 insertions(+), 1323 deletions(-) diff --git a/litellm-rust/crates/core/src/request_context.rs b/litellm-rust/crates/core/src/request_context.rs index 2796edf9993..8f02af57f59 100644 --- a/litellm-rust/crates/core/src/request_context.rs +++ b/litellm-rust/crates/core/src/request_context.rs @@ -7,10 +7,15 @@ pub struct RequestAttribution { #[derive(Clone, Debug, Default, PartialEq)] pub struct RequestCapabilities { + pub execution_mode: Option, pub stream: bool, pub has_agentic_hook: bool, pub has_custom_client: bool, pub request_format: Option, + pub input_source_kind: Option, + pub native_response_format: bool, + pub websocket_mode: Option, + pub requires_connection: bool, } #[derive(Clone, Debug, Default, PartialEq)] diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 1b6b37f04de..122a4b6f9e1 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -118,10 +118,15 @@ pub(crate) struct NativeRequestContext { #[derive(FromPyObject)] struct NativeRequestCapabilities { + execution_mode: Option, stream: bool, has_agentic_hook: bool, has_custom_client: bool, request_format: Option, + input_source_kind: Option, + native_response_format: bool, + websocket_mode: Option, + requires_connection: bool, } impl From for litellm_core::request_context::LiteLlmRequestContext { @@ -136,10 +141,15 @@ impl From for litellm_core::request_context::LiteLlmReques user_api_key_team_id: input.attribution.user_api_key_team_id, }, capabilities: litellm_core::request_context::RequestCapabilities { + execution_mode: input.capabilities.execution_mode, stream: input.capabilities.stream, has_agentic_hook: input.capabilities.has_agentic_hook, has_custom_client: input.capabilities.has_custom_client, request_format: input.capabilities.request_format, + input_source_kind: input.capabilities.input_source_kind, + native_response_format: input.capabilities.native_response_format, + websocket_mode: input.capabilities.websocket_mode, + requires_connection: input.capabilities.requires_connection, }, } } @@ -206,10 +216,15 @@ class BedrockOptions: @dataclass(frozen=True) class Capabilities: + execution_mode: object = None stream: object = False has_agentic_hook: object = False has_custom_client: object = False request_format: object = None + input_source_kind: object = None + native_response_format: object = False + websocket_mode: object = None + requires_connection: object = False @dataclass(frozen=True) class VertexOptions: diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index cc9430f6bf4..a0c84615224 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -27,7 +27,9 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts +from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors from litellm.rust_bridge.request import anthropic_options +from litellm.rust_bridge.runtime import DispatchResult from litellm.types.llms.anthropic import ( ContentBlockDelta, ContentBlockStart, @@ -369,15 +371,7 @@ class AnthropicChatCompletion(BaseLLM): if config is None: raise ValueError(f"Provider config not found for model: {model} and provider: {custom_llm_provider}") - def build_request() -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream - """Translate the request the Python way, returning `(headers, data)`. - - The pair stays mutable because the streaming path rewrites it in - place (`data["stream"] = True`) before sending. - - Shared by the normal path and by the Rust path's fallback, which - builds it only when the Rust call did not serve the request. - """ + def prepare_python() -> tuple[dict[str, str], dict[str, object]]: # mutable-ok: stream mutates data request_data: Final = config.transform_request( model=model, messages=messages, @@ -385,12 +379,29 @@ class AnthropicChatCompletion(BaseLLM): litellm_params=litellm_params, headers=headers, ) - return update_request_with_filtered_beta( + python_headers, data = update_request_with_filtered_beta( headers=headers, request_data=request_data, provider=custom_llm_provider, ) + ## LOGGING + # Reaching here with `serves_via_rust` set means the Rust attempt + # declined at call time, before the provider was called, and already + # logged this request. That is the same attempt continuing. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, + api_key=api_key, + additional_args={ + "complete_input_dict": data, + "api_base": api_base, + "headers": python_headers, + }, + ) + print_verbose(f"_is_function_call: {_is_function_call}") + return python_headers, data + # The Rust core owns the whole call for the subset it accepts, so ask # before transforming: whichever path runs emits pre_call exactly once. # `get_config` merges the class-level defaults (Anthropic's required @@ -407,68 +418,26 @@ class AnthropicChatCompletion(BaseLLM): litellm_params=litellm_params, stream=stream, ) + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "model": model, + "messages": messages, + **rust_optional_params, + }, + "api_base": api_base, + "headers": headers, + } if serves_via_rust: - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "model": model, - "messages": messages, - **rust_optional_params, - }, - "api_base": api_base, - "headers": headers, - } logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args) - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, - messages=messages, - api_key=api_key, - additional_args=rust_logging_args, - ) - if acompletion is True: + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key=api_key, + additional_args=rust_logging_args, + ) - async def python_fallback() -> "ModelResponse | CustomStreamWrapper": - # pre_call already fired for this request above. The Rust - # path only declines before the provider is called, so this - # is the same attempt continuing, not a second one. - fallback_headers, fallback_data = build_request() - return await self.acompletion_function( - model=model, - messages=messages, - data=fallback_data, - api_base=api_base, - custom_prompt_dict=custom_prompt_dict, - model_response=model_response, - print_verbose=print_verbose, - encoding=encoding, - api_key=api_key, - provider_config=config, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - _is_function_call=_is_function_call, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=fallback_headers, - client=client, - json_mode=json_mode, - timeout=timeout, - ) - - return rust_chat_completions_bridge.achat_completions_or_fallback( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - python_fallback=python_fallback, - anthropic=anthropic_options(litellm_params), - ) - rust_response: Final = rust_chat_completions_bridge.chat_completions( + def native_completion() -> DispatchResult[ModelResponse]: + return rust_chat_completions_bridge.chat_completions( model=model, messages=messages, optional_params=rust_optional_params, @@ -480,34 +449,42 @@ class AnthropicChatCompletion(BaseLLM): timeout=timeout, on_response=log_rust_post_call, anthropic=anthropic_options(litellm_params), + stream=bool(stream), + has_custom_client=client is not None, + eligible=serves_via_rust, ) - if rust_response is not None: - return rust_response - headers, data = build_request() - - ## LOGGING - # Reaching here with `serves_via_rust` set means the Rust attempt - # declined at call time, before the provider was called, and already - # logged this request. That is the same attempt continuing. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, + async def native_acompletion() -> DispatchResult[ModelResponse]: + return await rust_chat_completions_bridge.achat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, api_key=api_key, - additional_args={ - "complete_input_dict": data, - "api_base": api_base, - "headers": headers, - }, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + anthropic=anthropic_options(litellm_params), + stream=bool(stream), + has_custom_client=client is not None, + eligible=serves_via_rust, ) - print_verbose(f"_is_function_call: {_is_function_call}") - if acompletion is True: + + @anative_first( + native=native_acompletion, + route="chat_completions", + errors=lambda: provider_errors(custom_llm_provider or "", model), + ) + async def execute_async() -> ModelResponse | CustomStreamWrapper: + headers, data = prepare_python() if ( stream is True ): # if function call - fake the streaming (need complete blocks for output parsing in openai format) print_verbose("makes async anthropic streaming POST request") data["stream"] = stream - return self.acompletion_stream_function( + return await self.acompletion_stream_function( model=model, messages=messages, data=data, @@ -529,7 +506,7 @@ class AnthropicChatCompletion(BaseLLM): client=(client if client is not None and isinstance(client, AsyncHTTPHandler) else None), ) else: - return self.acompletion_function( + return await self.acompletion_function( model=model, messages=messages, data=data, @@ -551,7 +528,14 @@ class AnthropicChatCompletion(BaseLLM): json_mode=json_mode, timeout=timeout, ) - else: + + @native_first( + native=native_completion, + route="chat_completions", + errors=lambda: provider_errors(custom_llm_provider or "", model), + ) + def execute_sync() -> ModelResponse | CustomStreamWrapper: + headers, data = prepare_python() ## COMPLETION CALL if ( stream is True @@ -583,13 +567,12 @@ class AnthropicChatCompletion(BaseLLM): ) else: - if client is None or not isinstance(client, HTTPHandler): - client = _get_httpx_client(params={"timeout": timeout}) - else: - client = client + python_client: Final = ( + client if isinstance(client, HTTPHandler) else _get_httpx_client(params={"timeout": timeout}) + ) try: - response: Final = client.post( + response: Final = python_client.post( api_base, headers=headers, data=json.dumps(data), @@ -610,20 +593,21 @@ class AnthropicChatCompletion(BaseLLM): status_code=status_code, headers=error_headers, ) + return config.transform_response( + model=model, + raw_response=response, + model_response=model_response, + logging_obj=logging_obj, + api_key=api_key, + request_data=data, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + encoding=encoding, + json_mode=json_mode, + ) - return config.transform_response( - model=model, - raw_response=response, - model_response=model_response, - logging_obj=logging_obj, - api_key=api_key, - request_data=data, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - encoding=encoding, - json_mode=json_mode, - ) + return execute_async() if acompletion else execute_sync() def embedding(self): # logic for parsing in - calling - parsing out model embedding calls diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 19ffb695af1..982af97a6c4 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -10,10 +10,6 @@ from litellm.anthropic_beta_headers_manager import ( update_headers_with_filtered_beta, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObject -from litellm.llms.bedrock.request_metadata import ( - get_bedrock_request_metadata_fields, - resolve_bedrock_request_metadata, -) from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, HTTPHandler, @@ -22,7 +18,9 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts -from litellm.rust_bridge.request import NativeBedrockOptions +from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors +from litellm.rust_bridge.request import bedrock_options +from litellm.rust_bridge.runtime import DispatchResult from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -398,15 +396,11 @@ class BedrockConverseLLM(BaseAWSLLM): # resolved so both paths sign as the same principal. Bearer-token auth # resolves no SigV4 principal at all, and each path reads that token # itself. - rust_optional_params: Final = optional_params - rust_bedrock_options: Final = NativeBedrockOptions( - aws_access_key_id=None if credentials is None else credentials.access_key, - aws_secret_access_key=None if credentials is None else credentials.secret_key, - aws_session_token=None if credentials is None else credentials.token, - aws_region_name=aws_region_name, - request_metadata_fields=get_bedrock_request_metadata_fields(), - request_metadata=resolve_bedrock_request_metadata(litellm_params, optional_params.get("requestMetadata")), - ) + rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy + **optional_params, + **_sigv4_principal(credentials), + "aws_region_name": aws_region_name, + } serves_via_rust: Final = rust_chat_completions_accepts( model=model, messages=messages, @@ -415,55 +409,25 @@ class BedrockConverseLLM(BaseAWSLLM): litellm_params=litellm_params, stream=stream, ) + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "messages": messages, + **optional_params, + }, + "api_base": proxy_endpoint_url, + "headers": headers, + } if serves_via_rust: - rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict - "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent - "messages": messages, - **optional_params, - }, - "api_base": proxy_endpoint_url, - "headers": headers, - } logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args) - log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( - logging_obj=logging_obj, - messages=messages, - api_key="", - additional_args=rust_logging_args, - ) - if acompletion: - return rust_chat_completions_bridge.achat_completions_or_fallback( - model=model, - messages=messages, - optional_params=rust_optional_params, - model_response=model_response, - api_key=api_key, - api_base=proxy_endpoint_url, - custom_llm_provider="bedrock", - extra_headers=headers, - timeout=timeout, - on_response=log_rust_post_call, - bedrock=rust_bedrock_options, - python_fallback=lambda: self.async_completion( - model=model, - messages=messages, - api_base=proxy_endpoint_url, - model_response=model_response, - encoding=encoding, - logging_obj=logging_obj, - optional_params=optional_params, - stream=stream, - litellm_params=litellm_params, - logger_fn=logger_fn, - headers=headers, - timeout=timeout, - client=client, - credentials=credentials, - api_key=api_key, - skip_pre_call_logging=True, - ), - ) - rust_response: Final = rust_chat_completions_bridge.chat_completions( + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key="", + additional_args=rust_logging_args, + ) + + def native_completion() -> DispatchResult[ModelResponse]: + return rust_chat_completions_bridge.chat_completions( model=model, messages=messages, optional_params=rust_optional_params, @@ -474,17 +438,39 @@ class BedrockConverseLLM(BaseAWSLLM): extra_headers=headers, timeout=timeout, on_response=log_rust_post_call, - bedrock=rust_bedrock_options, + bedrock=bedrock_options(rust_optional_params), + stream=bool(stream), + has_custom_client=client is not None, + eligible=serves_via_rust, ) - if rust_response is not None: - return rust_response - ### ROUTING (ASYNC, STREAMING, SYNC) - if acompletion: - if isinstance(client, HTTPHandler): - client = None + async def native_acompletion() -> DispatchResult[ModelResponse]: + return await rust_chat_completions_bridge.achat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=proxy_endpoint_url, + custom_llm_provider="bedrock", + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + bedrock=bedrock_options(rust_optional_params), + stream=bool(stream), + has_custom_client=client is not None, + eligible=serves_via_rust, + ) + + @anative_first( + native=native_acompletion, + route="chat_completions", + errors=lambda: provider_errors("bedrock", model), + ) + async def execute_async() -> ModelResponse | CustomStreamWrapper: + python_client: Final = None if isinstance(client, HTTPHandler) else client if stream is True: - return self.async_streaming( + return await self.async_streaming( model=model, messages=messages, api_base=proxy_endpoint_url, @@ -497,7 +483,7 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=logger_fn, headers=headers, timeout=timeout, - client=client, + client=python_client, json_mode=json_mode, fake_stream=fake_stream, credentials=credentials, @@ -505,7 +491,7 @@ class BedrockConverseLLM(BaseAWSLLM): stream_chunk_size=stream_chunk_size, ) ### ASYNC COMPLETION - return self.async_completion( + return await self.async_completion( model=model, messages=messages, api_base=proxy_endpoint_url, @@ -518,108 +504,112 @@ class BedrockConverseLLM(BaseAWSLLM): logger_fn=logger_fn, headers=headers, timeout=timeout, - client=client, + client=python_client, credentials=credentials, api_key=api_key, + skip_pre_call_logging=serves_via_rust, + ) + + @native_first( + native=native_completion, + route="chat_completions", + errors=lambda: provider_errors("bedrock", model), + ) + def execute_sync() -> ModelResponse | CustomStreamWrapper: + ## TRANSFORMATION ## + + _data: Final = litellm.AmazonConverseConfig()._transform_request( + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=extra_headers, + ) + data: Final = json.dumps(_data) + + prepped: Final = self.get_request_headers( + credentials=credentials, + aws_region_name=aws_region_name, + extra_headers=extra_headers, + endpoint_url=proxy_endpoint_url, + data=data, + headers=headers, + api_key=api_key, ) - ## TRANSFORMATION ## + ## LOGGING + # Reaching here with `serves_via_rust` set means the synchronous Rust + # attempt declined at call time, before the provider was called, and + # already logged this request. That is the same attempt continuing. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": proxy_endpoint_url, + "headers": prepped.headers, + }, + ) + resolved_timeout: Final = httpx.Timeout(timeout) if isinstance(timeout, (float, int)) else timeout + python_client: Final = ( + _get_httpx_client({"timeout": resolved_timeout} if resolved_timeout is not None else None) + if client is None or isinstance(client, AsyncHTTPHandler) + else client + ) - _data: Final = litellm.AmazonConverseConfig()._transform_request( - model=model, - messages=messages, - optional_params=optional_params, - litellm_params=litellm_params, - headers=extra_headers, - ) - data: Final = json.dumps(_data) + if stream is not None and stream is True: + completion_stream, response_headers = make_sync_call( + client=python_client, + api_base=proxy_endpoint_url, + headers=prepped.headers, + data=data, + model=model, + messages=messages, + logging_obj=logging_obj, + json_mode=json_mode, + fake_stream=fake_stream, + stream_chunk_size=stream_chunk_size, + ) + streaming_response: Final = CustomStreamWrapper( + completion_stream=completion_stream, + model=model, + custom_llm_provider="bedrock", + logging_obj=logging_obj, + _response_headers=response_headers, + ) - prepped: Final = self.get_request_headers( - credentials=credentials, - aws_region_name=aws_region_name, - extra_headers=extra_headers, - endpoint_url=proxy_endpoint_url, - data=data, - headers=headers, - api_key=api_key, - ) + return streaming_response - ## LOGGING - # Reaching here with `serves_via_rust` set means the synchronous Rust - # attempt declined at call time, before the provider was called, and - # already logged this request. That is the same attempt continuing. - # The asynchronous branch above returns before this point, and hands - # its own fallback `skip_pre_call_logging=True` for the same reason. - if not serves_via_rust: - logging_obj.pre_call( - input=messages, + ### COMPLETION + + try: + response: Final = python_client.post( + url=proxy_endpoint_url, + headers=prepped.headers, + data=data, + logging_obj=logging_obj, + ) + response.raise_for_status() + except httpx.HTTPStatusError as err: + error_code: Final = err.response.status_code + raise BedrockError(status_code=error_code, message=err.response.text) + except httpx.TimeoutException: + raise BedrockError(status_code=408, message="Timeout error occurred.") + + sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response( + model=model, + response=response, + model_response=model_response, + stream=stream if isinstance(stream, bool) else False, + logging_obj=logging_obj, api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) - if client is None or isinstance(client, AsyncHTTPHandler): - _params: Final = {} - if timeout is not None: - if isinstance(timeout, float) or isinstance(timeout, int): - timeout = httpx.Timeout(timeout) - _params["timeout"] = timeout - client = _get_httpx_client(_params) - else: - client = client - - if stream is not None and stream is True: - completion_stream, response_headers = make_sync_call( - client=(client if client is not None and isinstance(client, HTTPHandler) else None), - api_base=proxy_endpoint_url, - headers=prepped.headers, data=data, - model=model, messages=messages, - logging_obj=logging_obj, - json_mode=json_mode, - fake_stream=fake_stream, - stream_chunk_size=stream_chunk_size, - ) - streaming_response: Final = CustomStreamWrapper( - completion_stream=completion_stream, - model=model, - custom_llm_provider="bedrock", - logging_obj=logging_obj, - _response_headers=response_headers, + optional_params=optional_params, + encoding=encoding, ) + sync_transformed_response.set_provider_response_headers(response.headers) + return sync_transformed_response - return streaming_response - - ### COMPLETION - - try: - response: Final = client.post( - url=proxy_endpoint_url, - headers=prepped.headers, - data=data, - logging_obj=logging_obj, - ) - response.raise_for_status() - except httpx.HTTPStatusError as err: - error_code: Final = err.response.status_code - raise BedrockError(status_code=error_code, message=err.response.text) - except httpx.TimeoutException: - raise BedrockError(status_code=408, message="Timeout error occurred.") - - sync_transformed_response: Final = litellm.AmazonConverseConfig()._transform_response( - model=model, - response=response, - model_response=model_response, - stream=stream if isinstance(stream, bool) else False, - logging_obj=logging_obj, - api_key="", - data=data, - messages=messages, - optional_params=optional_params, - encoding=encoding, - ) - sync_transformed_response.set_provider_response_headers(response.headers) - return sync_transformed_response + return execute_async() if acompletion else execute_sync() diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index ee5a49374af..5aa580ae47a 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2232,6 +2232,8 @@ class BaseLLMHTTPHandler: custom_llm_provider=custom_llm_provider, litellm_params=litellm_params, has_agentic_hook=self._has_agentic_completion_hook(logging_obj), + stream=bool(stream), + has_custom_client=client is not None, model=model, api_key=api_key, api_base=api_base, @@ -2388,6 +2390,8 @@ class BaseLLMHTTPHandler: custom_llm_provider: str, litellm_params: GenericLiteLLMParams, has_agentic_hook: bool, + stream: bool, + has_custom_client: bool, model: str, api_key: str | None, api_base: str | None, @@ -2415,6 +2419,9 @@ class BaseLLMHTTPHandler: custom_llm_provider=custom_llm_provider, extra_headers=headers, timeout=timeout, + stream=stream, + has_custom_client=has_custom_client, + has_agentic_hook=has_agentic_hook, ) def adapt(rust_response: dict[str, object]) -> AnthropicMessagesResponse: diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index bebb4b44589..7fea86e10b8 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -8,7 +8,6 @@ import mimetypes import os import re from collections.abc import Callable, Coroutine, Mapping -from dataclasses import dataclass from io import IOBase from typing import Any, Final, cast @@ -18,23 +17,16 @@ import litellm from litellm._logging import verbose_logger from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.azure_ai.ocr.common_utils import ( - is_azure_document_intelligence_model, -) +from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model from litellm.llms.base_llm.ocr.transformation import ( OCR_REQUEST_FORMAT_PARAM, - BaseOCRConfig, OCRResponse, parse_ocr_request_format, ) from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.rust_bridge import ocr as rust_ocr_bridge -from litellm.rust_bridge.request import ( - NativeRequestOptions, - PreparedNativeCall, - vertex_options, -) -from litellm.rust_bridge.timeouts import timeout_to_seconds +from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors +from litellm.rust_bridge.runtime import DispatchResult from litellm.types.router import GenericLiteLLMParams from litellm.utils import ProviderConfigManager, client @@ -43,28 +35,6 @@ base_llm_http_handler = BaseLLMHTTPHandler() ################################################# -@dataclass -class _PreparedOCRRequest: - model: str - document: dict[str, Any] - api_key: str | None - api_base: str | None - custom_llm_provider: str - extra_headers: dict[str, object] | None - provider_config: BaseOCRConfig - optional_params: dict[str, object] - litellm_params: dict[str, object] - effective_timeout: float | httpx.Timeout - litellm_logging_obj: LiteLLMLoggingObj - - -_RUST_OCR_PROVIDERS: Final = { - "mistral", - "azure_ai", - "vertex_ai", -} - - def _prepare_ocr_request( model: str, document: Mapping[str, object], @@ -74,7 +44,7 @@ def _prepare_ocr_request( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, kwargs: dict[str, object], -) -> _PreparedOCRRequest: +) -> rust_ocr_bridge.PreparedOCRRequest: litellm_logging_obj: Final = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj")) litellm_call_id: Final = cast(str | None, kwargs.get("litellm_call_id", None)) @@ -171,7 +141,7 @@ def _prepare_ocr_request( custom_llm_provider=custom_llm_provider, ) - return _PreparedOCRRequest( + return rust_ocr_bridge.PreparedOCRRequest( model=model, document=document, api_key=api_key, @@ -186,142 +156,70 @@ def _prepare_ocr_request( ) -def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool: - if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": - return False - if not prepared_request.provider_config.supports_rust_bridge(): - return False - return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS - - -def _rust_bridge_optional_params( - prepared_request: _PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> dict[str, object]: - optional_params: Final = dict(prepared_request.optional_params) - if prepared_request.custom_llm_provider == "vertex_ai": - vertex_project: Final = ( - prepared_request.litellm_params.get("vertex_project") - or prepared_request.litellm_params.get("vertex_ai_project") - or litellm.vertex_project - or resolve_secret("VERTEXAI_PROJECT") - ) - vertex_location: Final = ( - prepared_request.litellm_params.get("vertex_location") - or prepared_request.litellm_params.get("vertex_ai_location") - or litellm.vertex_location - or resolve_secret("VERTEXAI_LOCATION") - or resolve_secret("VERTEX_LOCATION") - ) - if vertex_project is not None: - optional_params["vertex_project"] = vertex_project - if vertex_location is not None: - optional_params["vertex_location"] = vertex_location - return optional_params - - -def _rust_bridge_api_base( - prepared_request: _PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> str | None: - if prepared_request.api_base is not None: - return prepared_request.api_base - if prepared_request.custom_llm_provider == "azure_ai": - if is_azure_document_intelligence_model(prepared_request.model): - return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") - return resolve_secret("AZURE_AI_API_BASE") - return None - - -def _prepare_rust_ocr_call( - prepared_request: _PreparedOCRRequest, +@anative_first( + native=rust_ocr_bridge.aattempt_ocr, + route="ocr", + errors=lambda prepared_request, resolve_api_key: provider_errors( + prepared_request.custom_llm_provider, prepared_request.model + ), +) +async def _execute_aocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, resolve_api_key: Callable[[str], str | None], -) -> PreparedNativeCall[rust_ocr_bridge.NativeOCRRequest]: - provider_config: Final = prepared_request.provider_config - api_key_env_var: Final = provider_config.get_api_key_env_var() - resolved_api_key: Final = prepared_request.api_key or ( - resolve_api_key(api_key_env_var) if api_key_env_var is not None else None - ) - resolved_headers: Final = provider_config.validate_environment( - headers=prepared_request.extra_headers or {}, - model=prepared_request.model, - api_key=resolved_api_key, - api_base=prepared_request.api_base, - litellm_params=prepared_request.litellm_params, - ) - resolved_complete_url: Final = provider_config.get_complete_url( - api_base=prepared_request.api_base, - model=prepared_request.model, - optional_params=prepared_request.optional_params, - litellm_params=prepared_request.litellm_params, - ) - rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key) - rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key) - prepared_request.litellm_logging_obj.pre_call( - input="OCR document processing", - api_key=resolved_api_key, - additional_args={ - "complete_input_dict": { - "model": prepared_request.model, - "document": prepared_request.document, - **rust_optional_params, - }, - "api_base": resolved_complete_url, - "headers": resolved_headers, - }, - ) - return PreparedNativeCall( - request=rust_ocr_bridge.NativeOCRRequest( - model=prepared_request.model, - document=prepared_request.document, - optional_params=prepared_request.optional_params, - ), - options=NativeRequestOptions( - vertex=vertex_options(rust_optional_params), - api_key=resolved_api_key, - api_base=rust_api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=cast( # cast-ok: provider header normalization returns string-object pairs - dict[str, object], resolved_headers - ), - timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), - ), - ) - - -def _run_rust_ocr( - prepared_request: _PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], - fallback: Callable[[], OCRResponse | Coroutine[object, object, OCRResponse]], -) -> OCRResponse | Coroutine[object, object, OCRResponse]: - return rust_ocr_bridge.dispatch_ocr( - prepare=lambda: _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ), - fallback=fallback, - adapt=OCRResponse.model_validate, - model=prepared_request.model, - provider=prepared_request.custom_llm_provider, - eligible=_rust_ocr_supported(prepared_request), - ) - - -async def _run_rust_aocr( - prepared_request: _PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], - fallback: Callable[[], Coroutine[object, object, OCRResponse]], ) -> OCRResponse: - return await rust_ocr_bridge.adispatch_ocr( - prepare=lambda: _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ), - fallback=fallback, - adapt=OCRResponse.model_validate, + pending: Final = base_llm_http_handler.ocr( model=prepared_request.model, - provider=prepared_request.custom_llm_provider, - eligible=_rust_ocr_supported(prepared_request), + document=prepared_request.document, + optional_params=prepared_request.optional_params, + timeout=prepared_request.effective_timeout, + logging_obj=prepared_request.litellm_logging_obj, + api_key=prepared_request.api_key, + api_base=prepared_request.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + aocr=True, + headers=prepared_request.extra_headers, + provider_config=prepared_request.provider_config, + litellm_params=prepared_request.litellm_params, + ) + response: Final = await pending if asyncio.iscoroutine(pending) else pending + if response is None: + raise ValueError(f"Got an unexpected None response from the OCR API: {response}") + return response + + +def _attempt_ocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], + is_async: bool, +) -> DispatchResult[OCRResponse]: + return rust_ocr_bridge.attempt_ocr(prepared_request=prepared_request, resolve_api_key=resolve_api_key) + + +@native_first( + native=_attempt_ocr, + route="ocr", + errors=lambda prepared_request, resolve_api_key, is_async: provider_errors( + prepared_request.custom_llm_provider, prepared_request.model + ), +) +def _execute_ocr( + prepared_request: rust_ocr_bridge.PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], + is_async: bool, +) -> OCRResponse | Coroutine[object, object, OCRResponse]: + return base_llm_http_handler.ocr( + model=prepared_request.model, + document=prepared_request.document, + optional_params=prepared_request.optional_params, + timeout=prepared_request.effective_timeout, + logging_obj=prepared_request.litellm_logging_obj, + api_key=prepared_request.api_key, + api_base=prepared_request.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + aocr=is_async, + headers=prepared_request.extra_headers, + provider_config=prepared_request.provider_config, + litellm_params=prepared_request.litellm_params, ) @@ -421,31 +319,7 @@ async def aocr( from litellm.secret_managers.main import get_secret_str - async def python_fallback() -> OCRResponse: - pending: Final = base_llm_http_handler.ocr( - model=prepared.model, - document=prepared.document, - optional_params=prepared.optional_params, - timeout=prepared.effective_timeout, - logging_obj=prepared.litellm_logging_obj, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared.custom_llm_provider, - aocr=True, - headers=prepared.extra_headers, - provider_config=prepared.provider_config, - litellm_params=prepared.litellm_params, - ) - response: Final = await pending if asyncio.iscoroutine(pending) else pending - if response is None: - raise ValueError(f"Got an unexpected None response from the OCR API: {response}") - return response - - return await _run_rust_aocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, - fallback=python_fallback, - ) + return await _execute_aocr(prepared_request=prepared, resolve_api_key=get_secret_str) except Exception as e: raise litellm.exception_type( model=model, @@ -686,27 +560,7 @@ def ocr( from litellm.secret_managers.main import get_secret_str - def python_fallback() -> OCRResponse | Coroutine[object, object, OCRResponse]: - return base_llm_http_handler.ocr( - model=prepared.model, - document=prepared.document, - optional_params=prepared.optional_params, - timeout=prepared.effective_timeout, - logging_obj=prepared.litellm_logging_obj, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared.custom_llm_provider, - aocr=_is_async, - headers=prepared.extra_headers, - provider_config=prepared.provider_config, - litellm_params=prepared.litellm_params, - ) - - return _run_rust_ocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, - fallback=python_fallback, - ) + return _execute_ocr(prepared_request=prepared, resolve_api_key=get_secret_str, is_async=_is_async) except Exception as e: raise litellm.exception_type( model=model, diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index a46a1f541bd..6af8471e2b3 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -4,16 +4,12 @@ The Rust core owns the conversation translation, the provider call, and the response normalization for the subset of `/chat/completions` requests it accepts. This module only marshals inputs and hands the normalized result to LiteLLM's existing `ModelResponse` builder. - -``None`` means the provider was never called, so the caller is free to serve the -request on the Python path. A failure after the call was issued raises instead: -retrying it there would bill the customer for the same work twice. """ from __future__ import annotations import json -from collections.abc import Awaitable, Callable, Mapping, Sequence +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Final, Protocol import httpx @@ -24,7 +20,7 @@ from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response impo convert_to_model_response_object, ) from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned -from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged from litellm.rust_bridge.configuration import rust_enabled from litellm.rust_bridge.protocols import ( RustAchatCompletions, @@ -35,17 +31,13 @@ from litellm.rust_bridge.request import ( NativeAnthropicOptions, NativeBedrockOptions, NativeChatCompletionsRequest, + NativeRequestCapabilities, NativeRequestContext, NativeRequestOptions, PreparedNativeCall, call_native, ) -from litellm.rust_bridge.runtime import ( - BridgeErrorContext, - EndpointBinding, - EndpointDispatch, - async_none, -) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt from litellm.rust_bridge.timeouts import timeout_to_seconds from litellm.types.utils import ModelResponse @@ -103,16 +95,10 @@ def response_logger( return log -_CHAT: Final[EndpointDispatch[RustChatCompletions, RustAchatCompletions]] = EndpointDispatch.native( - route="chat_completions", - sync=lambda native: native.chat_completions, - asynchronous=lambda native: native.achat_completions, - enabled=rust_enabled, -) -_CHAT_PREFLIGHT: Final[EndpointBinding[RustChatCompletionsDecline]] = EndpointBinding.native( - route="chat_completions", - select=lambda native: native.chat_completions_decline, - enabled=rust_enabled, +_CHAT: Final[NativeBinding[RustChatCompletions]] = NativeBinding(lambda native: native.chat_completions) +_ACHAT: Final[NativeBinding[RustAchatCompletions]] = NativeBinding(lambda native: native.achat_completions) +_CHAT_PREFLIGHT: Final[NativeBinding[RustChatCompletionsDecline]] = NativeBinding( + lambda native: native.chat_completions_decline ) @@ -126,14 +112,14 @@ def set_rust_chat_completions( patching module attributes.""" if not isinstance(chat_completions, Unchanged): if chat_completions is None: - _CHAT.sync.reset() + _CHAT.reset() else: - _CHAT.sync.override(chat_completions) + _CHAT.override(chat_completions) if not isinstance(achat_completions, Unchanged): if achat_completions is None: - _CHAT.asynchronous.reset() + _ACHAT.reset() else: - _CHAT.asynchronous.override(achat_completions) + _ACHAT.override(achat_completions) if not isinstance(decline, Unchanged): if decline is None: _CHAT_PREFLIGHT.reset() @@ -202,14 +188,24 @@ def rust_chat_completions_accepts( if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params): verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path") return False - return _CHAT_PREFLIGHT.accepts( - check=lambda decline: decline( + if not rust_enabled(): + return False + decline: Final = _CHAT_PREFLIGHT.load() + if decline is None: + return False + try: + reason: Final = decline( model=model, messages=messages, optional_params=optional_params, custom_llm_provider=custom_llm_provider, - ), - ) + ) + except Exception as error: # noqa: BLE001 # capability checks perform no provider I/O + verbose_logger.debug("Native chat acceptance check failed: %s", error) + return False + if reason is not None: + verbose_logger.debug("Native chat request is ineligible: %s", reason) + return reason is None def _build_model_response( @@ -240,18 +236,23 @@ def chat_completions( on_response: ResponseObserver, bedrock: NativeBedrockOptions | None = None, anthropic: NativeAnthropicOptions | None = None, -) -> ModelResponse | None: + stream: bool = False, + has_custom_client: bool = False, + eligible: bool = True, +) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) return _build_model_response(rust_response, model_response) - return _CHAT.invoke( + def call(native: RustChatCompletions, prepared: PreparedNativeCall[NativeChatCompletionsRequest]) -> Mapping[str, object]: + return call_native(native, prepared) + + return attempt( + load=_CHAT.load, + enabled=rust_enabled(), + eligible=eligible, prepare=lambda: PreparedNativeCall( - NativeChatCompletionsRequest( - model=model, - messages=messages, - optional_params=optional_params, - ), + request=NativeChatCompletionsRequest(model=model, messages=messages, optional_params=optional_params), options=NativeRequestOptions( api_key=api_key, api_base=api_base, @@ -261,12 +262,16 @@ def chat_completions( bedrock=bedrock, anthropic=anthropic, ), - context=NativeRequestContext(), + context=NativeRequestContext( + capabilities=NativeRequestCapabilities( + execution_mode="sync", + stream=stream, + has_custom_client=has_custom_client, + ) + ), ), - call=call_native, - fallback=lambda: None, + call=call, adapt=adapt, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -284,18 +289,26 @@ async def achat_completions( on_response: ResponseObserver, bedrock: NativeBedrockOptions | None = None, anthropic: NativeAnthropicOptions | None = None, -) -> ModelResponse | None: + stream: bool = False, + has_custom_client: bool = False, + eligible: bool = True, +) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) return _build_model_response(rust_response, model_response) - return await _CHAT.ainvoke( + async def call( + native: RustAchatCompletions, + prepared: PreparedNativeCall[NativeChatCompletionsRequest], + ) -> Mapping[str, object]: + return await call_native(native, prepared) + + return await aattempt( + load=_ACHAT.load, + enabled=rust_enabled(), + eligible=eligible, prepare=lambda: PreparedNativeCall( - NativeChatCompletionsRequest( - model=model, - messages=messages, - optional_params=optional_params, - ), + request=NativeChatCompletionsRequest(model=model, messages=messages, optional_params=optional_params), options=NativeRequestOptions( api_key=api_key, api_base=api_base, @@ -305,64 +318,14 @@ async def achat_completions( bedrock=bedrock, anthropic=anthropic, ), - context=NativeRequestContext(), - ), - call=call_native, - fallback=async_none, - adapt=adapt, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), - ) - - -async def achat_completions_or_fallback( - *, - model: str, - messages: Sequence[object], - optional_params: Mapping[str, object], - model_response: ModelResponse, - api_key: str | None, - api_base: str | None, - custom_llm_provider: str | None, - extra_headers: Mapping[str, object] | None, - timeout: float | httpx.Timeout | None, - on_response: ResponseObserver, - python_fallback: Callable[[], Awaitable[object]], - bedrock: NativeBedrockOptions | None = None, - anthropic: NativeAnthropicOptions | None = None, -) -> object: - """Await the Rust path, falling back to the caller's own Python path when - the bridge is unavailable or the call fails. - - The caller supplies the fallback, so the bridge stays free of provider - dispatch. This exists because a caller that dispatches asynchronously has - already returned a coroutine by the time a Rust failure surfaces, and so - cannot fall back on its own. - """ - - def adapt(rust_response: Mapping[str, object]) -> object: - on_response(rust_response) - return _build_model_response(rust_response, model_response) - - return await _CHAT.ainvoke( - prepare=lambda: PreparedNativeCall( - NativeChatCompletionsRequest( - model=model, - messages=messages, - optional_params=optional_params, + context=NativeRequestContext( + capabilities=NativeRequestCapabilities( + execution_mode="async", + stream=stream, + has_custom_client=has_custom_client, + ) ), - options=NativeRequestOptions( - api_key=api_key, - api_base=api_base, - custom_llm_provider=custom_llm_provider, - extra_headers=extra_headers, - timeout_seconds=timeout_to_seconds(timeout), - bedrock=bedrock, - anthropic=anthropic, - ), - context=NativeRequestContext(), ), - call=call_native, - fallback=python_fallback, + call=call, 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 f41cff140f8..445009e4cfb 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -6,30 +6,21 @@ from typing import Final import httpx -from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged from litellm.rust_bridge.protocols import RustAmessages, RustMessages from litellm.rust_bridge.request import ( NativeMessagesRequest, + NativeRequestCapabilities, NativeRequestContext, NativeRequestOptions, PreparedNativeCall, call_native, ) -from litellm.rust_bridge.runtime import ( - BridgeErrorContext, - EndpointDispatch, - always_enabled, - async_none, - identity, -) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds -_MESSAGES: Final[EndpointDispatch[RustMessages, RustAmessages]] = EndpointDispatch.native( - route="messages", - sync=lambda native: native.messages, - asynchronous=lambda native: native.amessages, - enabled=always_enabled, -) +_MESSAGES: Final[NativeBinding[RustMessages]] = NativeBinding(lambda native: native.messages) +_AMESSAGES: Final[NativeBinding[RustAmessages]] = NativeBinding(lambda native: native.amessages) def set_rust_messages( @@ -39,22 +30,22 @@ def set_rust_messages( ) -> None: if not isinstance(messages, Unchanged): if messages is None: - _MESSAGES.sync.reset() + _MESSAGES.reset() else: - _MESSAGES.sync.override(messages) + _MESSAGES.override(messages) if not isinstance(amessages, Unchanged): if amessages is None: - _MESSAGES.asynchronous.reset() + _AMESSAGES.reset() else: - _MESSAGES.asynchronous.override(amessages) + _AMESSAGES.override(amessages) def load_rust_messages() -> RustMessages | None: - return _MESSAGES.sync.load() + return _MESSAGES.load() def load_rust_amessages() -> RustAmessages | None: - return _MESSAGES.asynchronous.load() + return _AMESSAGES.load() def messages( @@ -66,13 +57,16 @@ def messages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - return _MESSAGES.invoke( + stream: bool = False, + has_custom_client: bool = False, + has_agentic_hook: bool = False, +) -> DispatchResult[dict[str, object]]: + return attempt( + load=_MESSAGES.load, + enabled=True, + eligible=True, prepare=lambda: PreparedNativeCall( - NativeMessagesRequest( - model=model, - body=body, - ), + request=NativeMessagesRequest(model=model, body=body), options=NativeRequestOptions( api_key=api_key, api_base=api_base, @@ -80,12 +74,17 @@ def messages( extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), ), - context=NativeRequestContext(), + context=NativeRequestContext( + capabilities=NativeRequestCapabilities( + execution_mode="sync", + stream=stream, + has_custom_client=has_custom_client, + has_agentic_hook=has_agentic_hook, + ) + ), ), call=call_native, - fallback=lambda: None, adapt=identity, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -98,13 +97,16 @@ async def amessages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - return await _MESSAGES.ainvoke( + stream: bool = False, + has_custom_client: bool = False, + has_agentic_hook: bool = False, +) -> DispatchResult[dict[str, object]]: + return await aattempt( + load=_AMESSAGES.load, + enabled=True, + eligible=True, prepare=lambda: PreparedNativeCall( - NativeMessagesRequest( - model=model, - body=body, - ), + request=NativeMessagesRequest(model=model, body=body), options=NativeRequestOptions( api_key=api_key, api_base=api_base, @@ -112,10 +114,15 @@ async def amessages( extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), ), - context=NativeRequestContext(), + context=NativeRequestContext( + capabilities=NativeRequestCapabilities( + execution_mode="async", + stream=stream, + has_custom_client=has_custom_client, + has_agentic_hook=has_agentic_hook, + ) + ), ), 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 320f4b6718f..e2236317853 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -1,29 +1,68 @@ -"""Thin Python wrapper for the native Rust OCR bridge.""" - from __future__ import annotations -from collections.abc import Awaitable, Callable, Mapping -from typing import Final, TypeVar +from collections.abc import Callable +from dataclasses import dataclass +from typing import Final -from . import configuration as _configuration -from .bindings import UNCHANGED, Unchanged -from .protocols import RustAocr, RustOcr -from .request import NativeOCRRequest, PreparedNativeCall, call_native -from .runtime import ( - BridgeErrorContext, - EndpointDispatch, +import httpx +from pydantic import TypeAdapter + +import litellm +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model +from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse +from litellm.rust_bridge import configuration as _configuration +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.protocols import RustAocr, RustOcr +from litellm.rust_bridge.request import ( + NativeOCRRequest, + NativeRequestCapabilities, + NativeRequestContext, + NativeRequestOptions, + PreparedNativeCall, + call_native, + vertex_options, ) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt +from litellm.rust_bridge.timeouts import timeout_to_seconds -rust_ocr_enabled = _configuration.rust_ocr_enabled -rust = _configuration.rust -ResultT = TypeVar("ResultT") +rust: Final = _configuration.rust +rust_ocr_enabled: Final = _configuration.rust_ocr_enabled + +_OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr) +_AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr) +_HEADERS: Final = TypeAdapter(dict[str, object]) -_OCR: Final[EndpointDispatch[RustOcr, RustAocr]] = EndpointDispatch.native( - route="ocr", - sync=lambda native: native.ocr, - asynchronous=lambda native: native.aocr, - enabled=_configuration.rust_ocr_enabled, +@dataclass(frozen=True, slots=True) +class PreparedOCRRequest: + model: str + document: dict[str, object] + api_key: str | None + api_base: str | None + custom_llm_provider: str + extra_headers: dict[str, object] | None + provider_config: BaseOCRConfig + optional_params: dict[str, object] + litellm_params: dict[str, object] + effective_timeout: float | httpx.Timeout + litellm_logging_obj: LiteLLMLoggingObj + + +@dataclass(frozen=True, slots=True) +class _PreparedRustOCRCall: + api_key: str | None + api_base: str | None + headers: dict[str, object] + optional_params: dict[str, object] + + +_RUST_OCR_PROVIDERS: Final = frozenset( + { + "mistral", + "azure_ai", + "vertex_ai", + } ) @@ -34,57 +73,212 @@ def set_rust_ocr( ) -> None: if not isinstance(ocr, Unchanged): if ocr is None: - _OCR.sync.reset() + _OCR.reset() else: - _OCR.sync.override(ocr) + _OCR.override(ocr) if not isinstance(aocr, Unchanged): if aocr is None: - _OCR.asynchronous.reset() + _AOCR.reset() else: - _OCR.asynchronous.override(aocr) + _AOCR.override(aocr) def load_rust_ocr() -> RustOcr | None: - return _OCR.sync.load() + return _OCR.load() def load_rust_aocr() -> RustAocr | None: - return _OCR.asynchronous.load() + return _AOCR.load() -def dispatch_ocr( - *, - prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]], - fallback: Callable[[], ResultT], - adapt: Callable[[Mapping[str, object]], ResultT], - model: str, - provider: str, - eligible: bool, -) -> ResultT: - return _OCR.invoke( - prepare=prepare, - call=call_native, - fallback=fallback, - adapt=adapt, - error_context=BridgeErrorContext(provider=provider, model=model), - eligible=eligible, +def _rust_ocr_supported(prepared_request: PreparedOCRRequest) -> bool: + if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": + return False + if not prepared_request.provider_config.supports_rust_bridge(): + return False + return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS + + +def _ocr_input_source_kind(document: dict[str, object]) -> str: + if "document_url" in document: + return "document_url" + if "image_url" in document: + return "image_url" + if "file" in document: + return "file" + return "inline" + + +def _rust_bridge_optional_params( + prepared_request: PreparedOCRRequest, + resolve_secret: Callable[[str], str | None], +) -> dict[str, object]: + if prepared_request.custom_llm_provider != "vertex_ai": + return prepared_request.optional_params + vertex_project: Final = ( + prepared_request.litellm_params.get("vertex_project") + or prepared_request.litellm_params.get("vertex_ai_project") + or litellm.vertex_project + or resolve_secret("VERTEXAI_PROJECT") + ) + vertex_location: Final = ( + prepared_request.litellm_params.get("vertex_location") + or prepared_request.litellm_params.get("vertex_ai_location") + or litellm.vertex_location + or resolve_secret("VERTEXAI_LOCATION") + or resolve_secret("VERTEX_LOCATION") + ) + return { + **prepared_request.optional_params, + **{ + name: value + for name, value in (("vertex_project", vertex_project), ("vertex_location", vertex_location)) + if value is not None + }, + } + + +def _rust_bridge_api_base( + prepared_request: PreparedOCRRequest, + resolve_secret: Callable[[str], str | None], +) -> str | None: + if prepared_request.api_base is not None: + return prepared_request.api_base + if prepared_request.custom_llm_provider == "azure_ai": + if is_azure_document_intelligence_model(prepared_request.model): + return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") + return resolve_secret("AZURE_AI_API_BASE") + return None + + +def _prepare_rust_ocr_call( + prepared_request: PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> _PreparedRustOCRCall: + provider_config: Final = prepared_request.provider_config + api_key_env_var: Final = provider_config.get_api_key_env_var() + resolved_api_key: Final = prepared_request.api_key or ( + resolve_api_key(api_key_env_var) if api_key_env_var is not None else None + ) + resolved_headers: Final = _HEADERS.validate_python( + provider_config.validate_environment( + headers=prepared_request.extra_headers or {}, + model=prepared_request.model, + api_key=resolved_api_key, + api_base=prepared_request.api_base, + litellm_params=prepared_request.litellm_params, + ) + ) + resolved_complete_url: Final = provider_config.get_complete_url( + api_base=prepared_request.api_base, + model=prepared_request.model, + optional_params=prepared_request.optional_params, + litellm_params=prepared_request.litellm_params, + ) + rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key) + rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key) + prepared_request.litellm_logging_obj.pre_call( + input="OCR document processing", + api_key=resolved_api_key, + additional_args={ + "complete_input_dict": { + "model": prepared_request.model, + "document": prepared_request.document, + **rust_optional_params, + }, + "api_base": resolved_complete_url, + "headers": resolved_headers, + }, + ) + return _PreparedRustOCRCall( + api_key=resolved_api_key, + api_base=rust_api_base, + headers=resolved_headers, + optional_params=rust_optional_params, ) -async def adispatch_ocr( - *, - prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]], - fallback: Callable[[], Awaitable[ResultT]], - adapt: Callable[[Mapping[str, object]], ResultT], - model: str, - provider: str, - eligible: bool, -) -> ResultT: - return await _OCR.ainvoke( - prepare=prepare, - call=call_native, - fallback=fallback, - adapt=adapt, - error_context=BridgeErrorContext(provider=provider, model=model), - eligible=eligible, +def attempt_ocr( + prepared_request: PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> DispatchResult[OCRResponse]: + return attempt( + load=_OCR.load, + enabled=rust_ocr_enabled(), + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + call=lambda native, prepared: call_native( + native, + PreparedNativeCall( + request=NativeOCRRequest( + model=prepared_request.model, + document=prepared_request.document, + optional_params=prepared.optional_params, + ), + options=NativeRequestOptions( + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=prepared.headers, + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + vertex=vertex_options(prepared.optional_params), + ), + context=NativeRequestContext( + capabilities=NativeRequestCapabilities( + execution_mode="sync", + input_source_kind=_ocr_input_source_kind(prepared_request.document), + native_response_format=( + prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" + ), + ) + ), + ), + ), + adapt=OCRResponse.model_validate, + eligible=_rust_ocr_supported(prepared_request), + ) + + +async def aattempt_ocr( + prepared_request: PreparedOCRRequest, + resolve_api_key: Callable[[str], str | None], +) -> DispatchResult[OCRResponse]: + return await aattempt( + load=_AOCR.load, + enabled=rust_ocr_enabled(), + prepare=lambda: _prepare_rust_ocr_call( + prepared_request=prepared_request, + resolve_api_key=resolve_api_key, + ), + call=lambda native, prepared: call_native( + native, + PreparedNativeCall( + request=NativeOCRRequest( + model=prepared_request.model, + document=prepared_request.document, + optional_params=prepared.optional_params, + ), + options=NativeRequestOptions( + api_key=prepared.api_key, + api_base=prepared.api_base, + custom_llm_provider=prepared_request.custom_llm_provider, + extra_headers=prepared.headers, + timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), + vertex=vertex_options(prepared.optional_params), + ), + context=NativeRequestContext( + capabilities=NativeRequestCapabilities( + execution_mode="async", + input_source_kind=_ocr_input_source_kind(prepared_request.document), + native_response_format=( + prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" + ), + ) + ), + ), + ), + adapt=OCRResponse.model_validate, + eligible=_rust_ocr_supported(prepared_request), ) diff --git a/litellm/rust_bridge/request.py b/litellm/rust_bridge/request.py index 4239b9763a1..e81d10f587f 100644 --- a/litellm/rust_bridge/request.py +++ b/litellm/rust_bridge/request.py @@ -90,10 +90,15 @@ class RequestAttribution: @dataclass(frozen=True, slots=True) class NativeRequestCapabilities: + execution_mode: str | None = None stream: bool = False has_agentic_hook: bool = False has_custom_client: bool = False request_format: str | None = None + input_source_kind: str | None = None + native_response_format: bool = False + websocket_mode: str | None = None + requires_connection: bool = False @dataclass(frozen=True, slots=True) diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index b7098aaddb4..2dc3e18bd2c 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -2,36 +2,32 @@ from __future__ import annotations +from collections.abc import AsyncGenerator +from contextlib import AbstractAsyncContextManager, asynccontextmanager from typing import Final import httpx from websockets.exceptions import ConnectionClosedOK -from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged from litellm.rust_bridge.configuration import rust_enabled from litellm.rust_bridge.protocols import ( RustResponsesWebSocket, RustResponsesWebSocketConnection, ) from litellm.rust_bridge.request import ( + NativeRequestCapabilities, NativeRequestContext, NativeRequestOptions, NativeResponsesWebSocketRequest, PreparedNativeCall, call_native, ) -from litellm.rust_bridge.runtime import ( - BridgeErrorContext, - EndpointBinding, - async_none, - identity, -) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result from litellm.rust_bridge.timeouts import timeout_to_seconds -_RESPONSES_WEBSOCKET: Final[EndpointBinding[RustResponsesWebSocketConnection]] = EndpointBinding.native( - route="responses_websocket", - select=lambda native: native.ResponsesWebSocketConnection, - enabled=rust_enabled, +_RESPONSES_WEBSOCKET: Final[NativeBinding[RustResponsesWebSocketConnection]] = NativeBinding( + lambda native: native.ResponsesWebSocketConnection, ) @@ -46,7 +42,7 @@ def set_rust_responses_websocket( _RESPONSES_WEBSOCKET.override(connection) -class _ConnectionAdapter: +class ConnectionAdapter: def __init__(self, connection: RustResponsesWebSocket): self._connection: Final[RustResponsesWebSocket] = connection @@ -68,18 +64,49 @@ async def connect( url: str, headers: dict[str, str], timeout: float | httpx.Timeout | None, -) -> _ConnectionAdapter | None: - connection: Final = await _RESPONSES_WEBSOCKET.ainvoke( + websocket_mode: str = "native", + requires_connection: bool = True, +) -> DispatchResult[ConnectionAdapter]: + return await aattempt( + load=_RESPONSES_WEBSOCKET.load, + enabled=rust_enabled(), + eligible=True, prepare=lambda: PreparedNativeCall( - NativeResponsesWebSocketRequest( - url=url, - ), + request=NativeResponsesWebSocketRequest(url=url), options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)), - context=NativeRequestContext(), + context=NativeRequestContext( + capabilities=NativeRequestCapabilities( + websocket_mode=websocket_mode, + requires_connection=requires_connection, + ) + ), ), - call=lambda connection_type, request: call_native(connection_type.connect, request), - fallback=async_none, - adapt=identity, - error_context=BridgeErrorContext(provider="openai", model="responses websocket"), + call=lambda connection_type, prepared: call_native(connection_type.connect, prepared), + adapt=ConnectionAdapter, ) - return None if connection is None else _ConnectionAdapter(connection) + + +@asynccontextmanager +async def _connection_context(connection: ConnectionAdapter) -> AsyncGenerator[ConnectionAdapter, None]: + try: + yield connection + finally: + await connection.close() + + +async def managed_connect( + *, + url: str, + headers: dict[str, str], + timeout: float | httpx.Timeout | None, + websocket_mode: str = "managed", + requires_connection: bool = True, +) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]: + result: Final = await connect( + url=url, + headers=headers, + timeout=timeout, + websocket_mode=websocket_mode, + requires_connection=requires_connection, + ) + return adapt_result(result, _connection_context) diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 0300cd84883..97f4eb485d1 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -4,9 +4,10 @@ from typing import Final import httpx -from litellm.rust_bridge.bindings import UNCHANGED, Unchanged +from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged from litellm.rust_bridge.protocols import RustAtranscription, RustTranscription from litellm.rust_bridge.request import ( + NativeRequestCapabilities, NativeRequestContext, NativeRequestOptions, NativeTranscriptionRequest, @@ -14,47 +15,36 @@ from litellm.rust_bridge.request import ( bedrock_options, call_native, ) -from litellm.rust_bridge.runtime import ( - BridgeErrorContext, - EndpointDispatch, - always_enabled, - async_none, - identity, -) +from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds -_TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] = EndpointDispatch.native( - route="audio transcription", - sync=lambda native: native.transcription, - asynchronous=lambda native: native.atranscription, - enabled=always_enabled, -) +_TRANSCRIPTION: Final[NativeBinding[RustTranscription]] = NativeBinding(lambda native: native.transcription) +_ATRANSCRIPTION: Final[NativeBinding[RustAtranscription]] = NativeBinding(lambda native: native.atranscription) def configure_rust_transcription( - enabled: bool = True, *, transcription: RustTranscription | None | Unchanged = UNCHANGED, atranscription: RustAtranscription | None | Unchanged = UNCHANGED, ) -> None: if not isinstance(transcription, Unchanged): if transcription is None: - _TRANSCRIPTION.sync.reset() + _TRANSCRIPTION.reset() else: - _TRANSCRIPTION.sync.override(transcription) + _TRANSCRIPTION.override(transcription) if not isinstance(atranscription, Unchanged): if atranscription is None: - _TRANSCRIPTION.asynchronous.reset() + _ATRANSCRIPTION.reset() else: - _TRANSCRIPTION.asynchronous.override(atranscription) + _ATRANSCRIPTION.override(atranscription) def load_rust_transcription() -> RustTranscription | None: - return _TRANSCRIPTION.sync.load() + return _TRANSCRIPTION.load() def load_rust_atranscription() -> RustAtranscription | None: - return _TRANSCRIPTION.asynchronous.load() + return _ATRANSCRIPTION.load() def transcription( @@ -67,14 +57,16 @@ def transcription( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - return _TRANSCRIPTION.invoke( + stream: bool = False, + has_custom_client: bool = False, + input_source_kind: str | None = None, +) -> DispatchResult[dict[str, object]]: + return attempt( + load=_TRANSCRIPTION.load, + enabled=True, + eligible=True, prepare=lambda: PreparedNativeCall( - NativeTranscriptionRequest( - model=model, - audio=audio, - optional_params=optional_params, - ), + request=NativeTranscriptionRequest(model=model, audio=audio, optional_params=optional_params), options=NativeRequestOptions( api_key=api_key, api_base=api_base, @@ -83,12 +75,17 @@ def transcription( timeout_seconds=timeout_to_seconds(timeout), bedrock=bedrock_options(optional_params), ), - context=NativeRequestContext(), + context=NativeRequestContext( + capabilities=NativeRequestCapabilities( + execution_mode="sync", + stream=stream, + has_custom_client=has_custom_client, + input_source_kind=input_source_kind, + ) + ), ), call=call_native, - fallback=lambda: None, adapt=identity, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) @@ -102,14 +99,16 @@ async def atranscription( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, -) -> dict[str, object] | None: - return await _TRANSCRIPTION.ainvoke( + stream: bool = False, + has_custom_client: bool = False, + input_source_kind: str | None = None, +) -> DispatchResult[dict[str, object]]: + return await aattempt( + load=_ATRANSCRIPTION.load, + enabled=True, + eligible=True, prepare=lambda: PreparedNativeCall( - NativeTranscriptionRequest( - model=model, - audio=audio, - optional_params=optional_params, - ), + request=NativeTranscriptionRequest(model=model, audio=audio, optional_params=optional_params), options=NativeRequestOptions( api_key=api_key, api_base=api_base, @@ -118,10 +117,15 @@ async def atranscription( timeout_seconds=timeout_to_seconds(timeout), bedrock=bedrock_options(optional_params), ), - context=NativeRequestContext(), + context=NativeRequestContext( + capabilities=NativeRequestCapabilities( + execution_mode="async", + stream=stream, + has_custom_client=has_custom_client, + input_source_kind=input_source_kind, + ) + ), ), call=call_native, - fallback=async_none, adapt=identity, - error_context=BridgeErrorContext(provider=custom_llm_provider or "", model=model), ) 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 59787822ab6..7af9a080da6 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -9,7 +9,8 @@ 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.rust_bridge.request import NativeMessagesRequest, NativeRequestContext, NativeRequestOptions +from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason from litellm.types.llms.anthropic_messages.anthropic_response import ( AnthropicMessagesResponse, ) @@ -38,12 +39,13 @@ REQUEST_BODY: dict[str, object] = { class RecordingMessages: def __init__(self) -> None: self.calls: list[dict[str, object]] = [] + self.contexts: list[NativeRequestContext] = [] def __call__( self, request: NativeMessagesRequest, *, - options: object, + options: NativeRequestOptions, context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( @@ -57,18 +59,20 @@ class RecordingMessages: "timeout_seconds": options.timeout_seconds, } ) + self.contexts.append(context) return dict(FAKE_MESSAGES_RESPONSE) class RecordingAsyncMessages: def __init__(self) -> None: self.calls: list[dict[str, object]] = [] + self.contexts: list[NativeRequestContext] = [] async def __call__( self, request: NativeMessagesRequest, *, - options: object, + options: NativeRequestOptions, context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( @@ -82,6 +86,7 @@ class RecordingAsyncMessages: "timeout_seconds": options.timeout_seconds, } ) + self.contexts.append(context) return dict(FAKE_MESSAGES_RESPONSE) @@ -89,9 +94,7 @@ class ExplodingAsyncMessages: def __init__(self) -> None: self.calls = 0 - async def __call__( - self, request: NativeMessagesRequest, *, options: object, context: NativeRequestContext - ) -> dict[str, object]: + async def __call__(self, *args: object, **kwargs: object) -> dict[str, object]: self.calls += 1 raise AssertionError("bridge must not be called") @@ -100,9 +103,7 @@ class RaisingAsyncMessages: def __init__(self) -> None: self.calls = 0 - async def __call__( - self, request: NativeMessagesRequest, *, options: object, context: NativeRequestContext - ) -> dict[str, object]: + async def __call__(self, *args: object, **kwargs: object) -> dict[str, object]: self.calls += 1 raise RuntimeError("upstream request failed with status 400: bad request") @@ -142,7 +143,7 @@ def test_load_rust_amessages_returns_injected_impl(): assert rust_messages.load_rust_amessages() is bridge -def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): +def test_messages_wrapper_reports_unavailable(monkeypatch): monkeypatch.setattr( importlib.import_module("litellm.rust_bridge.bindings"), "get_native_bridge", @@ -159,7 +160,7 @@ def test_messages_wrapper_returns_none_when_bridge_absent(monkeypatch): extra_headers={}, timeout=30.0, ) - assert result is None + assert result == NativeSkipped(NativeSkipReason.UNAVAILABLE) def test_messages_wrapper_forwards_args_and_converts_timeout(): @@ -177,7 +178,7 @@ def test_messages_wrapper_forwards_args_and_converts_timeout(): timeout=httpx.Timeout(600.0, read=42.0), ) - assert response == FAKE_MESSAGES_RESPONSE + assert response == Handled(FAKE_MESSAGES_RESPONSE) assert bridge.calls[0] == { "model": "claude-sonnet-4-5", "body": REQUEST_BODY, @@ -205,16 +206,43 @@ async def test_amessages_wrapper_forwards_args(): timeout=12.5, ) - assert response == FAKE_MESSAGES_RESPONSE + assert response == Handled(FAKE_MESSAGES_RESPONSE) assert bridge.calls[0]["model"] == "claude-sonnet-4-5" assert bridge.calls[0]["timeout_seconds"] == 12.5 +@pytest.mark.asyncio +async def test_amessages_wrapper_preserves_capability_facts(): + bridge = RecordingAsyncMessages() + rust_messages.set_rust_messages(amessages=bridge) + + await rust_messages.amessages( + model="claude-sonnet-4-5", + body=REQUEST_BODY, + api_key=None, + api_base=None, + custom_llm_provider="anthropic", + extra_headers=None, + timeout=None, + stream=True, + has_custom_client=True, + has_agentic_hook=True, + ) + + capabilities = bridge.contexts[0].capabilities + assert capabilities.execution_mode == "async" + assert capabilities.stream is True + assert capabilities.has_custom_client is True + assert capabilities.has_agentic_hook is True + + def _gate(**overrides): kwargs = { "custom_llm_provider": "azure_ai", "litellm_params": GenericLiteLLMParams(api_key="sk-azure"), "has_agentic_hook": False, + "stream": False, + "has_custom_client": False, "model": "claude-sonnet-4-5", "api_key": "sk-azure", "api_base": "https://resource.services.ai.azure.com/anthropic", @@ -223,7 +251,7 @@ def _gate(**overrides): "timeout": 30.0, } kwargs.update(overrides) - return BaseLLMHTTPHandler._maybe_rust_anthropic_messages(**kwargs) + return BaseLLMHTTPHandler._attempt_rust_anthropic_messages(**kwargs) @pytest.mark.asyncio @@ -234,7 +262,8 @@ async def test_gate_invokes_rust_and_marks_response_header(): response = await _gate() - assert response is not None + assert isinstance(response, Handled) + response = response.value assert response["id"] == "msg_123" assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} call = bridge.calls[0] @@ -247,13 +276,13 @@ async def test_gate_invokes_rust_and_marks_response_header(): @pytest.mark.asyncio -async def test_gate_propagates_unknown_native_errors(): +async def test_gate_reports_failure_to_harness(): bridge = RaisingAsyncMessages() litellm.rust(True) rust_messages.set_rust_messages(amessages=bridge) - with pytest.raises(RuntimeError, match="bad request"): - await _gate() + response = await _gate() + assert isinstance(response, NativeFailed) assert bridge.calls == 1 @@ -264,7 +293,7 @@ async def test_gate_skips_rust_when_flag_absent(): response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure")) - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -276,7 +305,8 @@ async def test_gate_uses_process_enable_without_request_override(): response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure")) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert bridge.calls[0]["custom_llm_provider"] == "azure_ai" @@ -288,7 +318,8 @@ async def test_gate_ignores_request_flag_when_process_enabled(): response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False)) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert len(bridge.calls) == 1 @@ -306,7 +337,8 @@ async def test_gate_invokes_rust_for_native_anthropic_provider(): headers={"x-api-key": "sk-ant", "anthropic-version": "2023-06-01"}, ) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} assert bridge.calls[0]["custom_llm_provider"] == "anthropic" assert bridge.calls[0]["api_key"] == "sk-ant" @@ -323,7 +355,8 @@ async def test_gate_invokes_rust_when_env_var_set(monkeypatch): litellm_params=GenericLiteLLMParams(api_key="sk-ant"), ) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert bridge.calls[0]["custom_llm_provider"] == "anthropic" @@ -338,7 +371,7 @@ async def test_gate_env_var_falsey_does_not_enable(monkeypatch): litellm_params=GenericLiteLLMParams(api_key="sk-ant"), ) - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -350,7 +383,7 @@ async def test_gate_skips_rust_for_unsupported_provider(): response = await _gate(custom_llm_provider="openai") - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -362,7 +395,7 @@ async def test_gate_skips_rust_for_agentic_hook(): response = await _gate(has_agentic_hook=True) - assert response is None + assert isinstance(response, NativeSkipped) assert bridge.calls == 0 @@ -378,7 +411,8 @@ async def test_gate_streams_through_rust_when_eligible_and_strips_stream_flag(): request_body=streaming_body, ) - assert response is not None + assert isinstance(response, Handled) + response = response.value assert response["_hidden_params"]["additional_headers"] == {"x-litellm-rust": "true"} assert "stream" not in bridge.calls[0]["body"] assert bridge.calls[0]["body"] == REQUEST_BODY @@ -411,4 +445,105 @@ async def test_gate_falls_back_when_bridge_unavailable(monkeypatch): response = await _gate() - assert response is None + assert isinstance(response, NativeSkipped) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("selection", ("native", "disabled", "failed", "declined", "upstream")) +async def test_messages_handler_runs_selected_backend_once(selection: str, monkeypatch: pytest.MonkeyPatch) -> None: + from datetime import datetime + from types import SimpleNamespace + + import httpx + + from litellm.exceptions import RateLimitError + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.llms.anthropic.experimental_pass_through.messages.transformation import AnthropicMessagesConfig + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.rust_bridge import bindings + + class Declined(Exception): + pass + + class Upstream(Exception): + pass + + error = ( + Upstream(429, "rate limited") + if selection == "upstream" + else Declined("unsupported") + if selection == "declined" + else RuntimeError("native failed") + if selection == "failed" + else None + ) + + class Native: + def __init__(self) -> None: + self.calls = 0 + + async def __call__(self, *args: object, **kwargs: object) -> dict[str, object]: + self.calls += 1 + if error is not None: + raise error + return dict(FAKE_MESSAGES_RESPONSE) + + bridge = Native() + monkeypatch.setattr( + bindings, + "get_native_bridge", + lambda: SimpleNamespace( + RustBridgeDeclined=Declined, + RustUpstreamError=Upstream, + ), + ) + rust_messages.set_rust_messages(amessages=bridge) + litellm.rust(selection != "disabled") + requests: list[httpx.Request] = [] + + def respond(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(200, json=FAKE_MESSAGES_RESPONSE) + + logging_obj = Logging( + model=FAKE_MESSAGES_RESPONSE["model"], + messages=[], + stream=False, + call_type="anthropic_messages", + start_time=datetime.now(), + litellm_call_id="harness-test", + function_id="harness-test", + ) + client = AsyncHTTPHandler() + await client.client.aclose() + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as transport: + client.client = transport + + async def run(): + return await BaseLLMHTTPHandler().async_anthropic_messages_handler( + model=FAKE_MESSAGES_RESPONSE["model"], + messages=[{"role": "user", "content": "hello"}], + anthropic_messages_provider_config=AnthropicMessagesConfig(), + anthropic_messages_optional_request_params={"max_tokens": 10}, + custom_llm_provider="anthropic", + litellm_params=GenericLiteLLMParams(), + logging_obj=logging_obj, + api_key="sk-test", + api_base="https://example.test", + client=client, + ) + + if selection in ("failed", "upstream"): + with pytest.raises(RateLimitError if selection == "upstream" else RuntimeError) as caught: + await run() + if selection == "upstream": + assert caught.value.__cause__ is error + assert caught.value.llm_provider == "anthropic" + assert caught.value.model == FAKE_MESSAGES_RESPONSE["model"] + else: + assert caught.value is error + else: + response = await run() + assert response["id"] == FAKE_MESSAGES_RESPONSE["id"] + assert len(requests) == (1 if selection in ("disabled", "declined") else 0) + assert bridge.calls == (0 if selection == "disabled" else 1) 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 019022ed919..da70f422f44 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 @@ -23,7 +23,9 @@ async def test_make_call_passes_logging_obj_to_client_post(): mock_client = AsyncMock() mock_response = MagicMock() mock_response.aiter_lines = MagicMock( - return_value=iter([b'data: {"type":"message_start"}\n', b'data: {"type":"message_delta"}\n']) + return_value=iter( + [b'data: {"type":"message_start"}\n', b'data: {"type":"message_delta"}\n'] + ) ) mock_client.post.return_value = mock_response @@ -92,7 +94,9 @@ def test_redacted_thinking_content_block_delta(): "data": "EuoBCoYBGAIiQJ/SxkPAgqxhKok29YrpJHRUJ0OT8ahCHKAwyhmRuUhtdmDX9+mn4gDzKNv3fVpQdB01zEPMzNY3QuTCd+1bdtEqQK6JuKHqdndbwpr81oVWb4wxd1GqF/7Jkw74IlQa27oobX+KuRkopr9Dllt/RDe7Se0sI1IkU7tJIAQCoP46OAwSDF51P09q67xhHlQ3ihoM2aOVlkghq/X0w8NlIjBMNvXYNbjhyrOcIg6kPFn2ed/KK7Cm5prYAtXCwkb4Wr5tUSoSHu9T5hKdJRbr6WsqEc7Lle7FULqMLZGkhqXyc3BA", }, } - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=False, json_mode=False + ) model_response = model_response_iterator.chunk_parser(chunk=chunk) print(f"\n\nmodel_response: {model_response}\n\n") assert model_response.choices[0].delta.thinking_blocks is not None @@ -100,14 +104,19 @@ def test_redacted_thinking_content_block_delta(): print( f"\n\nmodel_response.choices[0].delta.thinking_blocks[0]: {model_response.choices[0].delta.thinking_blocks[0]}\n\n" ) - assert model_response.choices[0].delta.thinking_blocks[0]["type"] == "redacted_thinking" + assert ( + model_response.choices[0].delta.thinking_blocks[0]["type"] + == "redacted_thinking" + ) assert model_response.choices[0].delta.provider_specific_fields is not None assert "thinking_blocks" in model_response.choices[0].delta.provider_specific_fields def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) chunks = [ { "type": "content_block_start", @@ -131,12 +140,17 @@ def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): }, ] - parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] + parsed_chunks = [ + model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks + ] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" + for chunk in parsed_chunks ) thinking_blocks = tuple( - block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block + for chunk in parsed_chunks + for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) expected_delta_blocks = ( {"type": "thinking", "thinking": "Step 1. "}, @@ -150,12 +164,18 @@ def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): assert reasoning_content == "Step 1. Step 2." assert thinking_blocks == (*expected_delta_blocks, expected_thinking_block) - assert parsed_chunks[1].choices[0].delta.provider_specific_fields == {"thinking_blocks": [expected_delta_blocks[0]]} - assert parsed_chunks[-1].choices[0].delta.provider_specific_fields == {"thinking_blocks": [expected_thinking_block]} + assert parsed_chunks[1].choices[0].delta.provider_specific_fields == { + "thinking_blocks": [expected_delta_blocks[0]] + } + assert parsed_chunks[-1].choices[0].delta.provider_specific_fields == { + "thinking_blocks": [expected_thinking_block] + } def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) chunks = [ { "type": "content_block_start", @@ -175,12 +195,17 @@ def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): {"type": "content_block_stop", "index": 0}, ] - parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] + parsed_chunks = [ + model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks + ] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" + for chunk in parsed_chunks ) thinking_blocks = tuple( - block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block + for chunk in parsed_chunks + for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) assert reasoning_content == "Step 1. Step 2." @@ -191,7 +216,9 @@ def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) chunks = [ { "type": "content_block_start", @@ -210,12 +237,17 @@ def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): }, ] - parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] + parsed_chunks = [ + model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks + ] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" + for chunk in parsed_chunks ) thinking_blocks = tuple( - block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block + for chunk in parsed_chunks + for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) assert reasoning_content == "Step 1. Step 2." @@ -226,7 +258,9 @@ def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): def test_handle_json_mode_chunk_response_format_tool(): - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=True + ) response_format_tool = ChatCompletionToolCallChunk( id="tool_123", type="function", @@ -237,7 +271,9 @@ def test_handle_json_mode_chunk_response_format_tool(): index=0, ) - text, tool_use = model_response_iterator._handle_json_mode_chunk("", response_format_tool) + text, tool_use = model_response_iterator._handle_json_mode_chunk( + "", response_format_tool + ) print(f"\n\nresponse_format_tool text: {text}\n\n") print(f"\n\nresponse_format_tool tool_use: {tool_use}\n\n") @@ -246,11 +282,15 @@ def test_handle_json_mode_chunk_response_format_tool(): def test_handle_json_mode_chunk_regular_tool(): - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=True + ) regular_tool = ChatCompletionToolCallChunk( id="tool_456", type="function", - function=ChatCompletionToolCallFunctionChunk(name="get_weather", arguments='{"location": "San Francisco, CA"}'), + function=ChatCompletionToolCallFunctionChunk( + name="get_weather", arguments='{"location": "San Francisco, CA"}' + ), index=0, ) @@ -264,13 +304,17 @@ def test_handle_json_mode_chunk_regular_tool(): def test_handle_json_mode_chunk_streaming_response_format_tool(): - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=True + ) # First chunk: response_format tool with id and name, but no arguments first_chunk = ChatCompletionToolCallChunk( id="tool_123", type="function", - function=ChatCompletionToolCallFunctionChunk(name=RESPONSE_FORMAT_TOOL_NAME, arguments=""), + function=ChatCompletionToolCallFunctionChunk( + name=RESPONSE_FORMAT_TOOL_NAME, arguments="" + ), index=0, ) @@ -278,7 +322,9 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): second_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk(name=None, arguments='{"question": "What is the weather?"'), + function=ChatCompletionToolCallFunctionChunk( + name=None, arguments='{"question": "What is the weather?"' + ), index=0, ) @@ -286,7 +332,9 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): third_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk(name=None, arguments=', "answer": "It is sunny"}'), + function=ChatCompletionToolCallFunctionChunk( + name=None, arguments=', "answer": "It is sunny"}' + ), index=0, ) @@ -317,7 +365,9 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): def test_handle_json_mode_chunk_streaming_regular_tool(): - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=True + ) # First chunk: regular tool with id and name, but no arguments first_chunk = ChatCompletionToolCallChunk( @@ -331,7 +381,9 @@ def test_handle_json_mode_chunk_streaming_regular_tool(): second_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk(name=None, arguments='{"location": "San Francisco, CA"}'), + function=ChatCompletionToolCallFunctionChunk( + name=None, arguments='{"location": "San Francisco, CA"}' + ), index=0, ) @@ -356,19 +408,27 @@ def test_handle_json_mode_chunk_streaming_regular_tool(): def test_response_format_tool_finish_reason(): - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=True + ) # First chunk: response_format tool response_format_tool = ChatCompletionToolCallChunk( id="tool_123", type="function", - function=ChatCompletionToolCallFunctionChunk(name=RESPONSE_FORMAT_TOOL_NAME, arguments='{"answer": "test"}'), + function=ChatCompletionToolCallFunctionChunk( + name=RESPONSE_FORMAT_TOOL_NAME, arguments='{"answer": "test"}' + ), index=0, ) # Process the tool call (should set converted_response_format_tool flag) - text, tool_use = model_response_iterator._handle_json_mode_chunk("", response_format_tool) - print(f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n") + text, tool_use = model_response_iterator._handle_json_mode_chunk( + "", response_format_tool + ) + print( + f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n" + ) # Simulate message_delta chunk with tool_use stop_reason message_delta_chunk = { @@ -387,19 +447,25 @@ def test_response_format_tool_finish_reason(): def test_regular_tool_finish_reason(): - model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) + model_response_iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=True + ) # First chunk: regular tool (not response_format) regular_tool = ChatCompletionToolCallChunk( id="tool_456", type="function", - function=ChatCompletionToolCallFunctionChunk(name="get_weather", arguments='{"location": "San Francisco, CA"}'), + function=ChatCompletionToolCallFunctionChunk( + name="get_weather", arguments='{"location": "San Francisco, CA"}' + ), index=0, ) # Process the tool call (should NOT set converted_response_format_tool flag) text, tool_use = model_response_iterator._handle_json_mode_chunk("", regular_tool) - print(f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n") + print( + f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n" + ) # Simulate message_delta chunk with tool_use stop_reason message_delta_chunk = { @@ -459,7 +525,9 @@ def test_text_only_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert parsed.choices[0].index == 0, f"Expected index=0, got {parsed.choices[0].index}" + assert ( + parsed.choices[0].index == 0 + ), f"Expected index=0, got {parsed.choices[0].index}" def test_streaming_thinking_deltas_count_reasoning_tokens_in_usage(): @@ -636,7 +704,9 @@ def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinkin ] self._write_response( content_type="text/event-stream", - body="".join(f"data: {json.dumps(event)}\n\n" for event in events).encode("utf-8"), + body="".join( + f"data: {json.dumps(event)}\n\n" for event in events + ).encode("utf-8"), ) return @@ -717,9 +787,13 @@ def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinkin assert content_chunks == [answer_text] assert stream_usage is not None stream_completion_details = stream_usage["completion_tokens_details"] - assert stream_completion_details["reasoning_tokens"] == non_stream_details.reasoning_tokens + assert ( + stream_completion_details["reasoning_tokens"] + == non_stream_details.reasoning_tokens + ) assert stream_completion_details["text_tokens"] == ( - stream_usage["completion_tokens"] - stream_completion_details["reasoning_tokens"] + stream_usage["completion_tokens"] + - stream_completion_details["reasoning_tokens"] ) assert requests_seen == [ { @@ -811,9 +885,9 @@ def test_text_and_tool_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert parsed.choices[0].index == 0, ( - f"Expected index=0 for chunk type {chunk.get('type')}, got {parsed.choices[0].index}" - ) + assert ( + parsed.choices[0].index == 0 + ), f"Expected index=0 for chunk type {chunk.get('type')}, got {parsed.choices[0].index}" def test_multiple_tools_streaming_has_index_zero(): @@ -866,11 +940,15 @@ def test_multiple_tools_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert parsed.choices[0].index == 0, f"Expected index=0, got {parsed.choices[0].index}" + assert ( + parsed.choices[0].index == 0 + ), f"Expected index=0, got {parsed.choices[0].index}" def test_streaming_chunks_have_stable_ids(): - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=False, json_mode=False + ) first_chunk = { "type": "content_block_delta", "index": 0, @@ -895,7 +973,9 @@ def test_partial_json_chunk_accumulation(): This tests the fix for https://github.com/BerriAI/litellm/issues/17473 where network fragmentation can cause SSE data to arrive in partial chunks. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) partial_chunk_1 = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hel' partial_chunk_2 = 'lo"}}' @@ -903,21 +983,31 @@ def test_partial_json_chunk_accumulation(): # First partial chunk should return None (still accumulating) result1 = iterator._parse_sse_data(f"data:{partial_chunk_1}") assert result1 is None, "First partial chunk should return None while accumulating" - assert iterator.chunk_type == "accumulated_json", "Should switch to accumulated_json mode" - assert iterator.accumulated_json == partial_chunk_1, "Should have accumulated first part" + assert ( + iterator.chunk_type == "accumulated_json" + ), "Should switch to accumulated_json mode" + assert ( + iterator.accumulated_json == partial_chunk_1 + ), "Should have accumulated first part" # Second partial chunk should complete the JSON and return a parsed result result2 = iterator._parse_sse_data(f"data:{partial_chunk_2}") assert result2 is not None, "Second chunk should return parsed result" - assert iterator.accumulated_json == "", "Buffer should be cleared after successful parse" - assert result2.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result2.choices[0].delta.content}'" + assert ( + iterator.accumulated_json == "" + ), "Buffer should be cleared after successful parse" + assert ( + result2.choices[0].delta.content == "Hello" + ), f"Expected 'Hello', got '{result2.choices[0].delta.content}'" def test_complete_json_chunk_no_accumulation(): """ Test that complete JSON chunks are parsed immediately without accumulation. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) complete_chunk = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}' @@ -925,14 +1015,18 @@ def test_complete_json_chunk_no_accumulation(): assert result is not None, "Complete chunk should return parsed result immediately" assert iterator.chunk_type == "valid_json", "Should remain in valid_json mode" assert iterator.accumulated_json == "", "Buffer should remain empty" - assert result.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result.choices[0].delta.content}'" + assert ( + result.choices[0].delta.content == "Hello" + ), f"Expected 'Hello', got '{result.choices[0].delta.content}'" def test_multiple_partial_chunks_accumulation(): """ Test that multiple partial chunks can be accumulated across several iterations. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) # Split a JSON chunk into three parts part1 = '{"type":"content_block_del' @@ -960,11 +1054,17 @@ def test_accumulated_json_partial_fragment_returns_none_without_parsing(): unlike Vertex which already deferred parsing until the buffer could close. A fragment that can't close a JSON value must not trigger a decode attempt. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) iterator.chunk_type = "accumulated_json" - with patch.object(json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode) as spy: - result = iterator._handle_accumulated_json_chunk('{"type":"content_block_delta","index":0,"delta":') + with patch.object( + json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode + ) as spy: + result = iterator._handle_accumulated_json_chunk( + '{"type":"content_block_delta","index":0,"delta":' + ) assert result is None assert spy.call_count == 0, "incomplete buffer should not be parsed" @@ -976,15 +1076,21 @@ def test_accumulated_json_does_not_reparse_every_fragment(): fragment. """ text = "x" * 200_000 - blob = json.dumps({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}}) + blob = json.dumps( + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}} + ) fragments = [blob[i : i + 4096] for i in range(0, len(blob), 4096)] assert len(fragments) > 10, "need a multi-fragment payload to exercise the bug" - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) iterator.chunk_type = "accumulated_json" parsed = None - with patch.object(json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode) as spy: + with patch.object( + json.JSONDecoder, "raw_decode", autospec=True, side_effect=json.JSONDecoder.raw_decode + ) as spy: for fragment in fragments: out = iterator._handle_accumulated_json_chunk(fragment) if out is not None: @@ -1008,7 +1114,9 @@ def test_accumulated_json_concatenated_envelopes_do_not_wedge(): and keeps the remainder, so both values surface across two calls. """ obj = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"a"}}' - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) iterator.chunk_type = "accumulated_json" first = iterator._handle_accumulated_json_chunk(obj + obj) @@ -1029,7 +1137,9 @@ def test_accumulated_json_heuristic_passes_but_value_still_incomplete(): heuristic must let the parse attempt through, and pop_next_value finding nothing must propagate as None rather than raising. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) iterator.chunk_type = "accumulated_json" result = iterator._handle_accumulated_json_chunk('{"type": {"nested": 1}') @@ -1043,7 +1153,9 @@ def test_accumulated_json_setter_and_sync_end_of_stream_drain(): underlying stream ends, instead of being silently dropped. """ obj = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"a"}}' - iterator = ModelResponseIterator(streaming_response=iter([]), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=iter([]), sync_stream=True, json_mode=False + ) iterator.chunk_type = "accumulated_json" iterator.accumulated_json = obj # exercises the setter @@ -1058,7 +1170,9 @@ def test_accumulated_json_async_end_of_stream_drain(): import asyncio obj = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"a"}}' - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=False, json_mode=False + ) iterator.chunk_type = "accumulated_json" iterator.accumulated_json = obj mock_async_iterator = MagicMock() @@ -1080,7 +1194,9 @@ def test_web_search_tool_result_no_extra_tool_calls(): The issue was that web_search_tool_result blocks have input_json_delta events with {} that were incorrectly being converted to tool calls. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) # Simulate the streaming sequence: # 1. server_tool_use block starts (web_search) @@ -1155,7 +1271,9 @@ def test_web_search_tool_result_no_extra_tool_calls(): # Should have exactly 2 tool calls: # 1. From content_block_start (server_tool_use) with id and name # 2. From content_block_delta with the actual query - assert len(tool_calls_emitted) == 2, f"Expected 2 tool calls, got {len(tool_calls_emitted)}" + assert ( + len(tool_calls_emitted) == 2 + ), f"Expected 2 tool calls, got {len(tool_calls_emitted)}" # First tool call should have the id and name assert tool_calls_emitted[0]["id"] == "srvtoolu_01ABC123" @@ -1171,7 +1289,9 @@ def test_current_content_block_type_tracking(): """ Test that current_content_block_type is properly tracked and reset. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) # Initially should be None assert iterator.current_content_block_type is None @@ -1224,7 +1344,9 @@ def test_web_search_tool_result_captured_in_provider_specific_fields(): The web_search_tool_result content comes ALL AT ONCE in content_block_start, not in deltas, so we need to capture it there. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) # Simulate the streaming sequence with web_search_tool_result chunks = [ @@ -1295,15 +1417,23 @@ def test_web_search_tool_result_captured_in_provider_specific_fields(): and parsed.choices[0].delta.provider_specific_fields and "web_search_results" in parsed.choices[0].delta.provider_specific_fields ): - web_search_results = parsed.choices[0].delta.provider_specific_fields["web_search_results"] + web_search_results = parsed.choices[0].delta.provider_specific_fields[ + "web_search_results" + ] # Verify web_search_results was captured assert web_search_results is not None, "web_search_results should be captured" assert len(web_search_results) == 1, "Should have 1 web_search_tool_result block" - assert web_search_results[0]["type"] == "web_search_tool_result", "Block type should be web_search_tool_result" - assert web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123", "tool_use_id should match" + assert ( + web_search_results[0]["type"] == "web_search_tool_result" + ), "Block type should be web_search_tool_result" + assert ( + web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123" + ), "tool_use_id should match" assert len(web_search_results[0]["content"]) == 2, "Should have 2 search results" - assert web_search_results[0]["content"][0]["title"] == "Fun Otter Facts", "First result title should match" + assert ( + web_search_results[0]["content"][0]["title"] == "Fun Otter Facts" + ), "First result title should match" def test_web_fetch_tool_result_captured_in_provider_specific_fields(): @@ -1317,7 +1447,9 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields(): The web_fetch_tool_result content comes ALL AT ONCE in content_block_start, not in deltas, so we need to capture it there. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) # Simulate the streaming sequence with web_fetch_tool_result chunks = [ @@ -1388,15 +1520,25 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields(): and parsed.choices[0].delta.provider_specific_fields and "web_search_results" in parsed.choices[0].delta.provider_specific_fields ): - web_search_results = parsed.choices[0].delta.provider_specific_fields["web_search_results"] + web_search_results = parsed.choices[0].delta.provider_specific_fields[ + "web_search_results" + ] # Verify web_fetch_tool_result was captured (stored in web_search_results list) assert web_search_results is not None, "web_search_results should be captured" assert len(web_search_results) == 1, "Should have 1 web_fetch_tool_result block" - assert web_search_results[0]["type"] == "web_fetch_tool_result", "Block type should be web_fetch_tool_result" - assert web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123", "tool_use_id should match" - assert web_search_results[0]["content"]["url"] == "https://example.com", "URL should match" - assert web_search_results[0]["content"]["content"]["title"] == "Example Page", "Title should match" + assert ( + web_search_results[0]["type"] == "web_fetch_tool_result" + ), "Block type should be web_fetch_tool_result" + assert ( + web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123" + ), "tool_use_id should match" + assert ( + web_search_results[0]["content"]["url"] == "https://example.com" + ), "URL should match" + assert ( + web_search_results[0]["content"]["content"]["title"] == "Example Page" + ), "Title should match" def test_web_fetch_tool_result_no_extra_tool_calls(): @@ -1409,7 +1551,9 @@ def test_web_fetch_tool_result_no_extra_tool_calls(): The issue was that web_fetch_tool_result blocks have input_json_delta events with {} that were incorrectly being converted to tool calls. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) # to verify it doesn't emit tool calls chunks = [ @@ -1453,9 +1597,9 @@ def test_web_fetch_tool_result_no_extra_tool_calls(): tool_call_count += 1 # Should have 0 tool calls - web_fetch_tool_result should not emit tool calls - assert tool_call_count == 0, ( - f"Expected 0 tool calls, got {tool_call_count}. web_fetch_tool_result should not emit tool calls" - ) + assert ( + tool_call_count == 0 + ), f"Expected 0 tool calls, got {tool_call_count}. web_fetch_tool_result should not emit tool calls" def test_container_in_provider_specific_fields_streaming(): @@ -1465,7 +1609,9 @@ def test_container_in_provider_specific_fields_streaming(): When container with skills is used, the container field should be present in the provider_specific_fields of the message_delta chunk. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=True, json_mode=False + ) # Simulate streaming chunks chunks = [ @@ -1533,12 +1679,20 @@ def test_container_in_provider_specific_fields_streaming(): and parsed.choices[0].delta.provider_specific_fields and "container" in parsed.choices[0].delta.provider_specific_fields ): - container_field = parsed.choices[0].delta.provider_specific_fields["container"] + container_field = parsed.choices[0].delta.provider_specific_fields[ + "container" + ] # Verify container was captured - assert container_field is not None, "container should be captured in provider_specific_fields" - assert container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p", "container id should match" - assert container_field["expires_at"] == "2025-12-16T04:57:16.913181Z", "expires_at should match" + assert ( + container_field is not None + ), "container should be captured in provider_specific_fields" + assert ( + container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p" + ), "container id should match" + assert ( + container_field["expires_at"] == "2025-12-16T04:57:16.913181Z" + ), "expires_at should match" assert len(container_field["skills"]) == 1, "Should have 1 skill" assert container_field["skills"][0]["skill_id"] == "pptx", "skill_id should be pptx" assert container_field["skills"][0]["version"] == "20251013", "version should match" @@ -1551,7 +1705,9 @@ def test_container_in_provider_specific_fields_non_streaming(): When container with skills is used in non-streaming, the container field should be present in the provider_specific_fields of the response. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=False, json_mode=False + ) # Simulate a message_delta chunk with container (as it would appear in non-streaming) message_delta_chunk = { @@ -1587,13 +1743,21 @@ def test_container_in_provider_specific_fields_non_streaming(): # Verify container is in provider_specific_fields assert model_response.choices[0].delta.provider_specific_fields is not None assert "container" in model_response.choices[0].delta.provider_specific_fields - container_field = model_response.choices[0].delta.provider_specific_fields["container"] + container_field = model_response.choices[0].delta.provider_specific_fields[ + "container" + ] assert container_field["id"] == "container_abc123xyz", "container id should match" - assert container_field["expires_at"] == "2025-12-20T10:30:00.000000Z", "expires_at should match" + assert ( + container_field["expires_at"] == "2025-12-20T10:30:00.000000Z" + ), "expires_at should match" assert len(container_field["skills"]) == 2, "Should have 2 skills" - assert container_field["skills"][0]["skill_id"] == "code_execution", "First skill_id should be code_execution" - assert container_field["skills"][1]["skill_id"] == "pptx", "Second skill_id should be pptx" + assert ( + container_field["skills"][0]["skill_id"] == "code_execution" + ), "First skill_id should be code_execution" + assert ( + container_field["skills"][1]["skill_id"] == "pptx" + ), "Second skill_id should be pptx" def test_container_absent_when_not_provided(): @@ -1602,7 +1766,9 @@ def test_container_absent_when_not_provided(): This ensures we don't add empty or None container fields. """ - iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) + iterator = ModelResponseIterator( + streaming_response=MagicMock(), sync_stream=False, json_mode=False + ) # message_delta without container message_delta_chunk = { @@ -1621,9 +1787,9 @@ def test_container_absent_when_not_provided(): # Verify container is NOT in provider_specific_fields when not provided if model_response.choices[0].delta.provider_specific_fields: - assert "container" not in model_response.choices[0].delta.provider_specific_fields, ( - "container should not be present when not provided in delta" - ) + assert ( + "container" not in model_response.choices[0].delta.provider_specific_fields + ), "container should not be present when not provided in delta" def test_streaming_code_execution_produces_code_interpreter_results(): @@ -1819,7 +1985,8 @@ def test_streaming_multiple_code_executions_no_duplicates(): # Second (final) emission: cumulative list with BOTH results # This is what stream_chunk_builder will pick as "last value wins" assert len(emissions[1]) == 2, ( - f"Expected final emission to have 2 results, got {len(emissions[1])}. IDs: {[r.id for r in emissions[1]]}" + f"Expected final emission to have 2 results, got {len(emissions[1])}. " + f"IDs: {[r.id for r in emissions[1]]}" ) assert emissions[1][0].id == "srvtoolu_01AAA" assert emissions[1][0].code == "echo first" @@ -1983,7 +2150,9 @@ def test_empty_output_produces_null_outputs(): assert code_results is not None, "No code_interpreter_results emitted" assert len(code_results) == 1 assert code_results[0].id == "srvtoolu_01AAA" - assert code_results[0].outputs is None, f"Expected outputs=None for empty execution, got {code_results[0].outputs}" + assert ( + code_results[0].outputs is None + ), f"Expected outputs=None for empty execution, got {code_results[0].outputs}" def test_non_bash_tool_result_skipped(): @@ -2046,10 +2215,12 @@ def test_non_bash_tool_result_skipped(): code_results = psf["code_interpreter_results"] # code_interpreter_results should be emitted but empty (no bash results) - assert code_results is not None, "Expected code_interpreter_results key to be emitted" - assert len(code_results) == 0, ( - f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}" - ) + assert ( + code_results is not None + ), "Expected code_interpreter_results key to be emitted" + assert ( + len(code_results) == 0 + ), f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}" class TestRustChatCompletionsHook: @@ -2086,9 +2257,13 @@ class TestRustChatCompletionsHook: from litellm.rust_bridge import chat_completions as bridge monkeypatch.setenv("LITELLM_RUST", "1") - bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) yield - bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) @staticmethod def _completion_kwargs(**overrides): @@ -2134,19 +2309,8 @@ class TestRustChatCompletionsHook: seen["gate"].append(kwargs) return decline_reason - def native(request, *, options, context): - seen["call"].append( - { - "model": request.model, - "messages": request.messages, - "optional_params": request.optional_params, - "api_key": options.api_key, - "api_base": options.api_base, - "custom_llm_provider": options.custom_llm_provider, - "extra_headers": options.extra_headers, - "timeout_seconds": options.timeout_seconds, - } - ) + def native(**kwargs): + seen["call"].append(kwargs) if sync_error is not None: raise sync_error return dict(sync_result if sync_result is not None else self.RUST_RESPONSE) @@ -2197,7 +2361,9 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion seen = self._inject() - AnthropicChatCompletion().completion(**self._completion_kwargs(optional_params={"max_tokens": 7})) + AnthropicChatCompletion().completion( + **self._completion_kwargs(optional_params={"max_tokens": 7}) + ) assert seen["call"][0]["optional_params"]["max_tokens"] == 7 def test_without_the_opt_in_the_core_is_never_consulted(self, monkeypatch): @@ -2206,14 +2372,15 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.transformation import AnthropicConfig seen = self._inject() - with ( - patch.object( - AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} - ) as transform, - patch.object(AnthropicChatCompletion, "acompletion_function"), + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ) as transform, patch.object( + AnthropicChatCompletion, "acompletion_function" ): try: - AnthropicChatCompletion().completion(**self._completion_kwargs(litellm_params={})) + AnthropicChatCompletion().completion( + **self._completion_kwargs(litellm_params={}) + ) except Exception: # The Python path goes on to make an HTTP call; reaching it is # the assertion, so the network failure below is expected. @@ -2227,7 +2394,9 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.transformation import AnthropicConfig seen = self._inject(decline_reason="unrecognized request parameter") - with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): try: AnthropicChatCompletion().completion(**self._completion_kwargs()) except Exception: @@ -2240,7 +2409,9 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.transformation import AnthropicConfig seen = self._inject() - with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): try: AnthropicChatCompletion().completion( **self._completion_kwargs(optional_params={"max_tokens": 16, "stream": True}) @@ -2254,7 +2425,9 @@ class TestRustChatCompletionsHook: seen = self._inject() logging_obj = MagicMock() - AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) + AnthropicChatCompletion().completion( + **self._completion_kwargs(logging_obj=logging_obj) + ) assert logging_obj.pre_call.call_count == 1 assert len(seen["call"]) == 1 @@ -2268,7 +2441,9 @@ class TestRustChatCompletionsHook: self._inject() logging_obj = MagicMock() - AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) + AnthropicChatCompletion().completion( + **self._completion_kwargs(logging_obj=logging_obj) + ) assert logging_obj.post_call.call_count == 1 logged = logging_obj.post_call.call_args.kwargs["original_response"] @@ -2289,16 +2464,22 @@ class TestRustChatCompletionsHook: RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(request, *, options, context): + def declining_native(**_kwargs): raise _Declined("blank message text") monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, chat_completions=declining_native) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, chat_completions=declining_native + ) logging_obj, calls = self._recording_logging_obj() - with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): try: - AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) + AnthropicChatCompletion().completion( + **self._completion_kwargs(logging_obj=logging_obj) + ) except Exception: # The Python path goes on to make an HTTP call; the log count is # the assertion, so a failure past this point is expected. @@ -2320,18 +2501,24 @@ class TestRustChatCompletionsHook: monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - async def declining_native(request, *, options, context): + async def declining_native(**_kwargs): raise _Declined("blank message text") - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=declining_native) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=declining_native + ) sentinel = object() async def python_path(**_kwargs): return sentinel - with patch.object(AnthropicChatCompletion, "acompletion_function", side_effect=python_path) as python_call: - result = await AnthropicChatCompletion().completion(**self._completion_kwargs(acompletion=True)) + with patch.object( + AnthropicChatCompletion, "acompletion_function", side_effect=python_path + ) as python_call: + result = await AnthropicChatCompletion().completion( + **self._completion_kwargs(acompletion=True) + ) assert result is sentinel assert python_call.called, "a failing rust call must re-enter the python path" @@ -2341,18 +2528,23 @@ class TestRustChatCompletionsHook: from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion from litellm.rust_bridge import chat_completions as bridge - async def native(request, *, options, context): + async def native(**_kwargs): return dict(self.RUST_RESPONSE) - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=native) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=native + ) with patch.object(AnthropicChatCompletion, "acompletion_function") as python_call: - result = await AnthropicChatCompletion().completion(**self._completion_kwargs(acompletion=True)) + result = await AnthropicChatCompletion().completion( + **self._completion_kwargs(acompletion=True) + ) assert result.choices[0].message.content == "hello from rust" assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} assert not python_call.called + def test_pre_call_logging_fires_once_when_the_sync_rust_call_declines(self, monkeypatch): """One request, one pre_call, on the synchronous path too. Without the suppression the Python path logs a second time for the same attempt.""" @@ -2369,22 +2561,30 @@ class TestRustChatCompletionsHook: monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - def declining_native(request, *, options, context): + def declining_native(**_kwargs): raise _Declined("blank message text") - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, chat_completions=declining_native) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, chat_completions=declining_native + ) logging_obj, calls = self._recording_logging_obj() - with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): try: - AnthropicChatCompletion().completion(**self._completion_kwargs(logging_obj=logging_obj)) + AnthropicChatCompletion().completion( + **self._completion_kwargs(logging_obj=logging_obj) + ) except Exception: # The Python path goes on to make an HTTP call; the log count is # the assertion, so a failure past this point is expected. pass assert len(calls["pre_call"]) == 1 - assert calls["pre_call"][0]["additional_args"]["complete_input_dict"]["model"] == ("claude-sonnet-4-5") + assert calls["pre_call"][0]["additional_args"]["complete_input_dict"]["model"] == ( + "claude-sonnet-4-5" + ) def test_pre_call_logging_still_fires_when_rust_is_not_involved(self, monkeypatch): """The suppression must not swallow the log on the ordinary path.""" @@ -2394,7 +2594,9 @@ class TestRustChatCompletionsHook: self._inject() logging_obj, calls = self._recording_logging_obj() - with patch.object(AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []}): + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): try: AnthropicChatCompletion().completion( **self._completion_kwargs(litellm_params={}, logging_obj=logging_obj) 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 c614e040d4c..becd6ecb832 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 @@ -10,8 +10,8 @@ from unittest.mock import MagicMock, patch import httpx import pytest -from botocore.credentials import Credentials +from botocore.credentials import Credentials from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.rust_bridge import chat_completions as bridge @@ -49,9 +49,13 @@ RESOLVED_CREDENTIALS = Credentials( @pytest.fixture(autouse=True) def reset_bridge(monkeypatch): monkeypatch.setenv("LITELLM_RUST", "1") - bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) yield - bridge.set_rust_chat_completions(chat_completions=None, achat_completions=None, decline=None) + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) def _inject(*, decline_reason=None, error: Exception | None = None): @@ -61,20 +65,8 @@ def _inject(*, decline_reason=None, error: Exception | None = None): seen["gate"].append(kwargs) return decline_reason - def native(request, *, options, context): - seen["call"].append( - { - "model": request.model, - "messages": request.messages, - "optional_params": request.optional_params, - "bedrock": options.bedrock, - "api_key": options.api_key, - "api_base": options.api_base, - "custom_llm_provider": options.custom_llm_provider, - "extra_headers": options.extra_headers, - "timeout_seconds": options.timeout_seconds, - } - ) + def native(**kwargs): + seen["call"].append(kwargs) if error is not None: raise error return dict(RUST_RESPONSE) @@ -134,18 +126,20 @@ def test_the_core_receives_the_credentials_this_handler_already_resolved(): seen = _inject() _run() - bedrock = seen["call"][0]["bedrock"] - assert bedrock.aws_access_key_id == "AKIARESOLVED" - assert bedrock.aws_secret_access_key == "resolved-secret" - assert bedrock.aws_session_token == "resolved-token" - assert bedrock.aws_region_name == "us-east-1" + params = seen["call"][0]["optional_params"] + assert params["aws_access_key_id"] == "AKIARESOLVED" + assert params["aws_secret_access_key"] == "resolved-secret" + assert params["aws_session_token"] == "resolved-token" + assert params["aws_region_name"] == "us-east-1" def test_the_core_receives_the_converse_url_this_handler_already_built(): seen = _inject() _run() - assert seen["call"][0]["api_base"].endswith("/model/anthropic.claude-sonnet-4-5-v1%3A0/converse") + assert seen["call"][0]["api_base"].endswith( + "/model/anthropic.claude-sonnet-4-5-v1%3A0/converse" + ) assert "bedrock-runtime.us-east-1.amazonaws.com" in seen["call"][0]["api_base"] @@ -213,10 +207,12 @@ 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(request, *, options, context): + async def declining_native(**_kwargs): raise _Declined("blank message text") - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=declining_native) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=declining_native + ) sentinel = object() @@ -224,10 +220,16 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): return sentinel with ( - patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS), - patch.object(BedrockConverseLLM, "async_completion", side_effect=python_path) as python_call, + patch.object( + BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS + ), + patch.object( + BedrockConverseLLM, "async_completion", side_effect=python_path + ) as python_call, ): - result = await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True)) + result = await BedrockConverseLLM().completion( + **_completion_kwargs(acompletion=True) + ) assert result is sentinel assert python_call.called, "a failing rust call must re-enter the python path" @@ -235,16 +237,22 @@ 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(request, *, options, context): + async def native(**_kwargs): return dict(RUST_RESPONSE) - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=native) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=native + ) with ( - patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS), + patch.object( + BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS + ), patch.object(BedrockConverseLLM, "async_completion") as python_call, ): - result = await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True)) + result = await BedrockConverseLLM().completion( + **_completion_kwargs(acompletion=True) + ) assert result.choices[0].message.content == "hello from rust" assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} @@ -263,7 +271,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(request, *, options, context): + async def declining_native(**_kwargs): raise _Declined("blank message text") logging_obj = MagicMock() @@ -275,11 +283,19 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines(): with ( patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()), - patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS), - patch.object(BedrockConverseLLM, "async_completion", side_effect=python_path), + patch.object( + BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS + ), + patch.object( + BedrockConverseLLM, "async_completion", side_effect=python_path + ), ): - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=declining_native) - await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True, logging_obj=logging_obj)) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=declining_native + ) + await BedrockConverseLLM().completion( + **_completion_kwargs(acompletion=True, logging_obj=logging_obj) + ) assert logging_obj.pre_call.call_count == 1 assert served and served[0]["skip_pre_call_logging"] is True @@ -368,13 +384,15 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(request, *, options, context): + def declining_native(**_kwargs): raise _Declined("blank message text") logging_obj = MagicMock() with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()): - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, chat_completions=declining_native) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, chat_completions=declining_native + ) response = _run( logging_obj=logging_obj, client=_sync_client_returning_converse_response(), @@ -420,14 +438,20 @@ 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(request, *, options, context): + async def native(**_kwargs): return dict(RUST_RESPONSE) - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, achat_completions=native) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=native + ) logging_obj = MagicMock() - with patch.object(BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS): - await BedrockConverseLLM().completion(**_completion_kwargs(acompletion=True, logging_obj=logging_obj)) + with patch.object( + BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS + ): + await BedrockConverseLLM().completion( + **_completion_kwargs(acompletion=True, logging_obj=logging_obj) + ) assert logging_obj.post_call.call_count == 1 logged = logging_obj.post_call.call_args.kwargs["original_response"] @@ -446,13 +470,15 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(request, *, options, context): + def declining_native(**_kwargs): raise _Declined("blank message text") logging_obj, calls = _recording_logging_obj() with patch("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()): - bridge.set_rust_chat_completions(decline=lambda **_kwargs: None, chat_completions=declining_native) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, chat_completions=declining_native + ) response = _run( logging_obj=logging_obj, client=_sync_client_returning_converse_response(), @@ -486,11 +512,9 @@ def test_the_rust_opt_in_needs_no_sigv4_principal(): response = _run(credentials=None, api_key="bedrock-bearer-token") assert response.choices[0].message.content == "hello from rust" - bedrock = seen["call"][0]["bedrock"] - assert bedrock.aws_access_key_id is None - assert bedrock.aws_secret_access_key is None - assert bedrock.aws_session_token is None - assert bedrock.aws_region_name == "us-east-1" + params = seen["call"][0]["optional_params"] + assert not {"aws_access_key_id", "aws_secret_access_key", "aws_session_token"} & params.keys() + assert params["aws_region_name"] == "us-east-1" assert seen["call"][0]["api_key"] == "bedrock-bearer-token" diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 2a9b314ac85..cd30a0b698f 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -12,13 +12,8 @@ import pytest import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge import configuration -from litellm.rust_bridge.request import ( - NativeOCRRequest, - NativeRequestContext, - NativeRequestOptions, - NativeVertexOptions, - PreparedNativeCall, -) +from litellm.rust_bridge.request import NativeOCRRequest, NativeRequestContext, NativeRequestOptions +from litellm.rust_bridge.runtime import Handled from litellm.rust_bridge.timeouts import timeout_to_seconds # `litellm/__init__.py` does `from .ocr.main import *`, which binds the `ocr` @@ -52,6 +47,7 @@ class RecordingBridge: def __init__(self) -> None: self.calls: list[dict[str, object]] = [] + self.contexts: list[NativeRequestContext] = [] def __call__( self, @@ -69,10 +65,10 @@ class RecordingBridge: "custom_llm_provider": options.custom_llm_provider, "extra_headers": options.extra_headers, "optional_params": request.optional_params, - "vertex": options.vertex, "timeout_seconds": options.timeout_seconds, } ) + self.contexts.append(context) return dict(FAKE_OCR_RESPONSE) @@ -81,6 +77,7 @@ class RecordingAsyncBridge: def __init__(self) -> None: self.calls: list[dict[str, object]] = [] + self.contexts: list[NativeRequestContext] = [] async def __call__( self, @@ -98,10 +95,10 @@ class RecordingAsyncBridge: "custom_llm_provider": options.custom_llm_provider, "extra_headers": options.extra_headers, "optional_params": request.optional_params, - "vertex": options.vertex, "timeout_seconds": options.timeout_seconds, } ) + self.contexts.append(context) return dict(FAKE_OCR_RESPONSE) @@ -156,6 +153,9 @@ class FakeOCRConfig: def get_api_key_env_var(self) -> str: return self.api_key_env_var + def supports_rust_bridge(self) -> bool: + return True + def validate_environment( self, *, @@ -192,7 +192,7 @@ def build_prepared_request( litellm_params: dict[str, object] | None = None, timeout: float | httpx.Timeout | None = 12.5, ) -> Any: - return ocr_main._PreparedOCRRequest( + return rust_bridge.PreparedOCRRequest( model=model, document=document, api_key=api_key, @@ -381,105 +381,13 @@ def test_timeout_to_seconds_handles_float_timeout_and_none(): assert timeout_to_seconds(httpx.Timeout(30.0, read=42.0)) == 42.0 -def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): - bridge = RecordingBridge() - - litellm.rust(True) - - rust_bridge.set_rust_ocr(ocr=bridge) - response = rust_bridge.dispatch_ocr( - prepare=lambda: PreparedNativeCall( - request=NativeOCRRequest( - model="mistral-ocr-latest", - document=DOCUMENT, - optional_params={"include_image_base64": True, "pages": [0]}, - ), - options=NativeRequestOptions( - api_key="sk-test", - api_base="https://proxy.internal", - custom_llm_provider="mistral", - extra_headers={"Authorization": "Bearer sk-test", "x-trace-id": "trace-1"}, - timeout_seconds=12.5, - ), - ), - fallback=lambda: pytest.fail("unexpected Python fallback"), - adapt=dict, - model="mistral-ocr-latest", - provider="mistral", - eligible=True, - ) - - assert response == FAKE_OCR_RESPONSE - call = bridge.calls[0] - assert call == { - "model": "mistral-ocr-latest", - "document": DOCUMENT, - "api_key": "sk-test", - "api_base": "https://proxy.internal", - "custom_llm_provider": "mistral", - "extra_headers": { - "Authorization": "Bearer sk-test", - "x-trace-id": "trace-1", - }, - "optional_params": {"include_image_base64": True, "pages": [0]}, - "vertex": None, - "timeout_seconds": 12.5, - } - - -@pytest.mark.asyncio -async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): - bridge = RecordingAsyncBridge() - - litellm.rust(True) - - rust_bridge.set_rust_ocr(aocr=bridge) - - async def unexpected_fallback(): - pytest.fail("unexpected Python fallback") - - response = await rust_bridge.adispatch_ocr( - prepare=lambda: PreparedNativeCall( - request=NativeOCRRequest( - model="mistral-ocr-maas", - document=DOCUMENT, - optional_params={}, - ), - options=NativeRequestOptions( - custom_llm_provider="vertex_ai", - vertex=NativeVertexOptions(project="project-1"), - timeout_seconds=42.0, - ), - ), - fallback=unexpected_fallback, - adapt=dict, - model="mistral-ocr-maas", - provider="vertex_ai", - eligible=True, - ) - - assert response == FAKE_OCR_RESPONSE - assert bridge.calls[0] == { - "model": "mistral-ocr-maas", - "document": DOCUMENT, - "api_key": None, - "api_base": None, - "custom_llm_provider": "vertex_ai", - "extra_headers": None, - "optional_params": {}, - "vertex": NativeVertexOptions(project="project-1"), - "timeout_seconds": 42.0, - } - - def test_run_rust_ocr_prepares_request_and_wraps_response(): bridge = RecordingBridge() logging_obj = RecordingLogging() litellm.rust(True) rust_bridge._OCR.override(bridge) - response = ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + response = rust_bridge.attempt_ocr( prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://proxy.internal", @@ -490,8 +398,13 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): resolve_api_key=lambda _name: None, ) + assert isinstance(response, Handled) + response = response.value assert isinstance(response, OCRResponse) assert response.pages[0].markdown == "hello world" + assert bridge.contexts[0].capabilities.execution_mode == "sync" + assert bridge.contexts[0].capabilities.input_source_kind == "document_url" + assert bridge.contexts[0].capabilities.native_response_format is False assert bridge.calls[0] == { "model": "mistral-ocr-latest", "document": DOCUMENT, @@ -503,7 +416,6 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): "x-trace-id": "trace-1", }, "optional_params": {"include_image_base64": True}, - "vertex": NativeVertexOptions(), "timeout_seconds": 12.5, } @@ -513,8 +425,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request(api_key=None, timeout=None), resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None, ) @@ -530,8 +441,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver(): def _resolver(name: str) -> str | None: raise AssertionError(f"resolver should not be called for {name}") - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( api_key="sk-explicit", timeout=None, @@ -552,8 +462,7 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): resolver_calls.append(name) return "sk-provider-env" - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"), model="provider-ocr-model", @@ -572,8 +481,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -588,8 +496,11 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): resolve_api_key=lambda _name: None, ) - assert bridge.calls[0]["optional_params"] == {"include_image_base64": True} - assert bridge.calls[0]["vertex"] == NativeVertexOptions(project="project-1", location="us-central1") + assert bridge.calls[0]["optional_params"] == { + "include_image_base64": True, + "vertex_project": "project-1", + "vertex_location": "us-central1", + } def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager(): @@ -603,8 +514,7 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana "VERTEXAI_LOCATION": "us-east5", }.get(name) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", @@ -613,7 +523,8 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana resolve_api_key=_resolver, ) - assert bridge.calls[0]["vertex"] == NativeVertexOptions(project="project-from-secret", location="us-east5") + assert bridge.calls[0]["optional_params"]["vertex_project"] == "project-from-secret" + assert bridge.calls[0]["optional_params"]["vertex_location"] == "us-east5" def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): @@ -621,8 +532,7 @@ def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", @@ -640,8 +550,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( custom_llm_provider="azure_ai", model="doc-intelligence/prebuilt-layout", @@ -662,8 +571,7 @@ def test_run_rust_ocr_runs_pre_call_logging(): litellm.rust(True) rust_bridge._OCR.override(bridge) - ocr_main._run_rust_ocr( - fallback=lambda: pytest.fail("unexpected Python fallback"), + rust_bridge.attempt_ocr( prepared_request=build_prepared_request( logging_obj=logging_obj, api_base="https://api.mistral.ai/v1", @@ -810,7 +718,7 @@ async def test_ocr_fallback_skips_native_preparation( def unexpected_preparation(*_args: object, **_kwargs: object) -> None: pytest.fail("Python fallback must not resolve native credentials or emit native pre_call") - monkeypatch.setattr(ocr_main, "_prepare_rust_ocr_call", unexpected_preparation) + monkeypatch.setattr(rust_bridge, "_prepare_rust_ocr_call", unexpected_preparation) monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", fallback) response: Final = ( @@ -823,6 +731,25 @@ async def test_ocr_fallback_skips_native_preparation( fallback.assert_called_once() +@pytest.mark.asyncio +async def test_aocr_rejects_empty_python_fallback_response(monkeypatch: pytest.MonkeyPatch) -> None: + captured: dict[str, object] = {} + + def fake_exception_type(**kwargs: object) -> CapturedException: + captured.update(kwargs) + return CapturedException("wrapped") + + monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type) + monkeypatch.setattr(ocr_main.base_llm_http_handler, "ocr", AsyncMock(return_value=None)) + + with pytest.raises(CapturedException, match="wrapped"): + await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test") + + original: Final = captured["original_exception"] + assert isinstance(original, ValueError) + assert str(original) == "Got an unexpected None response from the OCR API: None" + + def test_ocr_provider_configs_expose_api_key_env_vars(): from litellm.llms.azure_ai.ocr.document_intelligence.transformation import ( AzureDocumentIntelligenceOCRConfig, diff --git a/tests/test_litellm/responses/test_rust_bridge_websocket.py b/tests/test_litellm/responses/test_rust_bridge_websocket.py index aea5137db66..de8f91b6580 100644 --- a/tests/test_litellm/responses/test_rust_bridge_websocket.py +++ b/tests/test_litellm/responses/test_rust_bridge_websocket.py @@ -4,7 +4,12 @@ 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 +from litellm.rust_bridge.request import ( + NativeRequestContext, + NativeRequestOptions, + NativeResponsesWebSocketRequest, +) +from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped, NativeSkipReason class _FakeNativeConnection: @@ -28,14 +33,17 @@ class _ClosedNativeConnection: class _FakeNativeBridge: + contexts: list[NativeRequestContext] = [] + @classmethod async def connect( cls, request: NativeResponsesWebSocketRequest, *, - options: object, + options: NativeRequestOptions, context: NativeRequestContext, ) -> _FakeNativeConnection: + cls.contexts.append(context) return _FakeNativeConnection() @@ -58,25 +66,22 @@ def test_rust_websocket_bridge_uses_process_enablement() -> None: @pytest.mark.asyncio async def test_adapter_raises_clean_close_when_rust_connection_ends() -> None: - adapter = responses_websocket._ConnectionAdapter(_ClosedNativeConnection()) + adapter = responses_websocket.ConnectionAdapter(_ClosedNativeConnection()) with pytest.raises(responses_websocket.ConnectionClosedOK): await adapter.recv() @pytest.mark.asyncio -async def test_bridge_unavailable_returns_none(monkeypatch: pytest.MonkeyPatch) -> None: +async def test_bridge_reports_unavailable(monkeypatch: pytest.MonkeyPatch) -> None: configuration.rust(True) responses_websocket._RESPONSES_WEBSOCKET.override(None) - assert ( - await responses_websocket.connect( - url="wss://example.test/responses", - headers={}, - timeout=None, - ) - is None - ) + assert await responses_websocket.connect( + url="wss://example.test/responses", + headers={}, + timeout=None, + ) == NativeSkipped(NativeSkipReason.UNAVAILABLE) @pytest.mark.asyncio @@ -92,10 +97,13 @@ async def test_enabled_bridge_connects_and_adapts_socket( timeout=1.0, ) - assert connection is not None + assert isinstance(connection, Handled) + connection = connection.value await connection.send("response.create") assert await connection.recv() == "response.completed" await connection.close() + assert _FakeNativeBridge.contexts[-1].capabilities.websocket_mode == "native" + assert _FakeNativeBridge.contexts[-1].capabilities.requires_connection is True class _FailingNativeBridge: @@ -104,16 +112,74 @@ class _FailingNativeBridge: cls, request: NativeResponsesWebSocketRequest, *, - options: object, + options: NativeRequestOptions, context: NativeRequestContext, ) -> _FakeNativeConnection: raise RuntimeError("connection failed") +@pytest.mark.asyncio +async def test_connection_failure_is_reported_to_orchestration() -> None: + configuration.rust(True) + responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge) + result = await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None) + assert isinstance(result, NativeFailed) + assert str(result.error) == "connection failed" + + +@pytest.mark.asyncio +async def test_managed_connection_closes_native_socket_on_consumer_failure() -> None: + configuration.rust(True) + socket = _FakeNativeConnection() + + class Bridge: + @classmethod + async def connect( + cls, + request: NativeResponsesWebSocketRequest, + *, + options: NativeRequestOptions, + context: NativeRequestContext, + ) -> _FakeNativeConnection: + return socket + + responses_websocket.set_rust_responses_websocket(connection=Bridge) + result = await responses_websocket.managed_connect(url="wss://example.test/responses", headers={}, timeout=1.0) + assert isinstance(result, Handled) + + async def use_connection() -> None: + async with result.value as connection: + await connection.send("hello") + raise ValueError("consumer failed") + + with pytest.raises(ValueError, match="consumer failed"): + await use_connection() + assert socket.sent == ["hello"] + assert socket.closed + + @pytest.mark.asyncio async def test_connection_failure_does_not_authorize_python_fallback() -> None: + from contextlib import AbstractAsyncContextManager + + from litellm.rust_bridge.dispatch import anative_context, provider_errors + configuration.rust(True) responses_websocket.set_rust_responses_websocket(connection=_FailingNativeBridge) + @anative_context( + native=lambda: responses_websocket.managed_connect( + url="wss://example.test/responses", headers={}, timeout=None + ), + route="responses_websocket", + errors=lambda: provider_errors("openai", "responses websocket"), + ) + def execute() -> AbstractAsyncContextManager[object]: + pytest.fail("unknown native failures must not open a Python connection") + + async def run() -> None: + async with execute(): + pytest.fail("connection must fail before entering its body") + with pytest.raises(RuntimeError, match="connection failed"): - await responses_websocket.connect(url="wss://example.test/responses", headers={}, timeout=None) + await run() diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py index f81b7050134..31163ebaa05 100644 --- a/tests/test_litellm/rust_bridge/test_chat_completions.py +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -12,12 +12,8 @@ import pytest import litellm from litellm.rust_bridge import bindings, configuration from litellm.rust_bridge import chat_completions as bridge -from litellm.rust_bridge.request import ( - NativeBedrockOptions, - NativeRequestCapabilities, - NativeRequestContext, - anthropic_options, -) +from litellm.rust_bridge.request import NativeChatCompletionsRequest, NativeRequestContext, NativeRequestOptions +from litellm.rust_bridge.runtime import Handled, NativeFailed, NativeSkipped from litellm.types.utils import ModelResponse RUST_RESPONSE = { @@ -100,16 +96,40 @@ class _RecordingCall: self.result = result if result is not None else dict(RUST_RESPONSE) self.error = error self.calls: list[dict] = [] + self.contexts: list[NativeRequestContext] = [] - def __call__(self, request, *, options, context): - self.calls.append({"request": request, "options": options, "context": context}) + def __call__( + self, + request: NativeChatCompletionsRequest, + *, + options: NativeRequestOptions, + context: NativeRequestContext, + ): + kwargs = { + "model": request.model, + "messages": request.messages, + "optional_params": request.optional_params, + "api_key": options.api_key, + "api_base": options.api_base, + "custom_llm_provider": options.custom_llm_provider, + "extra_headers": options.extra_headers, + "timeout_seconds": options.timeout_seconds, + } + self.calls.append(kwargs) + self.contexts.append(context) if self.error is not None: raise self.error return self.result class _RecordingAsyncCall(_RecordingCall): - async def __call__(self, request, *, options, context): + async def __call__( + self, + request: NativeChatCompletionsRequest, + *, + options: NativeRequestOptions, + context: NativeRequestContext, + ): return _RecordingCall.__call__(self, request, options=options, context=context) @@ -256,7 +276,8 @@ class TestSyncCall: result = bridge.chat_completions(**_call_kwargs(model_response)) - assert result is not None + assert isinstance(result, Handled) + result = result.value assert result.choices[0].message.content == "hello from rust" assert result.choices[0].finish_reason == "stop" assert result.model == "claude-sonnet-4-5-20260101" @@ -270,16 +291,28 @@ class TestSyncCall: native = _RecordingCall() bridge.set_rust_chat_completions(chat_completions=native) bridge.chat_completions(**_call_kwargs(ModelResponse())) - assert native.calls[0]["options"].timeout_seconds == 30.0 + assert native.calls[0]["timeout_seconds"] == 30.0 - def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): + def test_preserves_execution_and_client_capabilities(self): + native = _RecordingCall() + bridge.set_rust_chat_completions(chat_completions=native) + bridge.chat_completions( + **_call_kwargs(ModelResponse()), + stream=True, + has_custom_client=True, + ) + assert native.contexts[0].capabilities.execution_mode == "sync" + assert native.contexts[0].capabilities.stream is True + assert native.contexts[0].capabilities.has_custom_client is True + + def test_reports_unavailable_bridge(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeSkipped) - def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): + def test_reports_native_decline_to_orchestration(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + assert isinstance(bridge.chat_completions(**_call_kwargs(ModelResponse())), NativeFailed) class TestAsyncCall: @@ -287,187 +320,18 @@ class TestAsyncCall: async def test_builds_a_model_response(self): bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) result = await bridge.achat_completions(**_call_kwargs(ModelResponse())) - assert result is not None + assert isinstance(result, Handled) + result = result.value assert result.choices[0].message.content == "hello from rust" assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} @pytest.mark.asyncio - async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): + async def test_reports_unavailable_bridge(self, monkeypatch): _hide_native_bridge(monkeypatch) - assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None + assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeSkipped) @pytest.mark.asyncio - async def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): + async def test_reports_native_decline_to_orchestration(self, monkeypatch): _fake_native_bridge(monkeypatch) bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))) - assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None - - -class TestAsyncFallbackWrapper: - @pytest.mark.asyncio - async def test_returns_the_rust_response_without_running_the_fallback(self): - bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) - ran = [] - - async def fallback(): - ran.append(True) - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result.choices[0].message.content == "hello from rust" - assert ran == [] - - @pytest.mark.asyncio - async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch): - _fake_native_bridge(monkeypatch) - bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming"))) - - async def fallback(): - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result == "python" - - @pytest.mark.asyncio - async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch): - _hide_native_bridge(monkeypatch) - - async def fallback(): - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result == "python" - - -class TestFailureClassification: - """A failure the provider already saw must not be retried on the Python - path: it would bill the customer for the same work twice.""" - - @pytest.fixture(autouse=True) - def _native_exceptions(self, monkeypatch): - _fake_native_bridge(monkeypatch) - - def test_a_decline_falls_back_because_nothing_was_sent(self): - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeDeclined("streaming"))) - assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None - - def test_an_upstream_failure_is_surfaced_with_its_status(self): - from litellm.exceptions import RateLimitError - - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited"))) - with pytest.raises(RateLimitError) as raised: - bridge.chat_completions(**_call_kwargs(ModelResponse())) - assert raised.value.status_code == 429 - assert "rate limited" in str(raised.value) - - def test_a_transport_failure_with_no_response_surfaces_as_a_500(self): - from litellm.exceptions import APIError - - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset"))) - with pytest.raises(APIError) as raised: - bridge.chat_completions(**_call_kwargs(ModelResponse())) - assert raised.value.status_code == 500 - - def test_an_unrecognized_error_is_not_swallowed(self): - bridge.set_rust_chat_completions(chat_completions=_RecordingCall(error=RuntimeError("something else"))) - with pytest.raises(RuntimeError): - bridge.chat_completions(**_call_kwargs(ModelResponse())) - - @pytest.mark.asyncio - async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self): - from litellm.exceptions import InternalServerError - - bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom"))) - ran = [] - - async def fallback(): - ran.append(True) - return "python" - - with pytest.raises(InternalServerError): - await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert ran == [], "a request the provider already served must not be re-issued" - - @pytest.mark.asyncio - async def test_the_async_wrapper_falls_back_on_a_decline(self): - bridge.set_rust_chat_completions( - achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text")) - ) - - async def fallback(): - return "python" - - result = await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - assert result == "python" - - -@pytest.mark.asyncio -async def test_missing_native_exception_types_does_not_authorize_python_fallback(monkeypatch): - _hide_native_bridge(monkeypatch) - bridge.set_rust_chat_completions( - chat_completions=_RecordingCall(error=RuntimeError("connection failed")), - achat_completions=_RecordingAsyncCall(error=RuntimeError("connection failed")), - ) - - with pytest.raises(RuntimeError, match="connection failed"): - bridge.chat_completions(**_call_kwargs(ModelResponse())) - - async def fallback(): - pytest.fail("unknown failure must not retry through Python") - - with pytest.raises(RuntimeError, match="connection failed"): - await bridge.achat_completions_or_fallback(**_call_kwargs(ModelResponse()), python_fallback=fallback) - - -def test_provider_credentials_are_separate_from_chat_body_params(): - native = _RecordingCall() - bridge.set_rust_chat_completions(chat_completions=native) - configuration.rust(True) - kwargs = _call_kwargs(ModelResponse()) - kwargs["optional_params"] = { - "max_tokens": 32, - } - kwargs["bedrock"] = NativeBedrockOptions( - aws_access_key_id="test-access-key", - aws_secret_access_key="test-secret-key", - ) - bridge.chat_completions(**kwargs) - request = native.calls[0]["request"] - options = native.calls[0]["options"] - assert request.optional_params == {"max_tokens": 32} - assert options.bedrock.aws_access_key_id == "test-access-key" - assert options.bedrock.aws_secret_access_key == "test-secret-key" - - -def test_provider_payload_extensions_cross_the_boundary_without_partitioning(): - native = _RecordingCall() - bridge.set_rust_chat_completions(chat_completions=native) - configuration.rust(True) - extensions = { - "vendor_object": {"nested": None}, - "vendor_array": [1, "two", False], - "vendor_scalar": 0.25, - "extra_body": {"temperature": 0.2, "config": {"replacement": True}}, - } - - kwargs = _call_kwargs(ModelResponse()) - kwargs["optional_params"] = extensions - bridge.chat_completions(**kwargs) - - assert native.calls[0]["request"].optional_params == extensions - - -def test_typed_capability_and_provider_metadata_facts_are_isolated(): - context = NativeRequestContext( - capabilities=NativeRequestCapabilities( - stream=True, - has_agentic_hook=True, - has_custom_client=True, - request_format="native", - ) - ) - anthropic = anthropic_options({"metadata": {"user_id": "user-123", "ignored": object()}}) - - assert context.capabilities.request_format == "native" - assert context.capabilities.has_agentic_hook is True - assert anthropic.user_id == "user-123" + assert isinstance(await bridge.achat_completions(**_call_kwargs(ModelResponse())), NativeFailed) diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py index 57c936c170b..6b58859155a 100644 --- a/tests/test_litellm/test_audio_transcription_rust_bridge.py +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -4,11 +4,8 @@ import pytest import litellm from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch -from litellm.rust_bridge.request import ( - NativeRequestContext, - NativeRequestOptions, - NativeTranscriptionRequest, -) +from litellm.rust_bridge.request import NativeRequestContext, NativeRequestOptions, NativeTranscriptionRequest +from litellm.rust_bridge.runtime import Handled rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") @@ -16,6 +13,7 @@ rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") class SyncBridge: def __init__(self) -> None: self.calls: list[dict[str, object]] = [] + self.contexts: list[NativeRequestContext] = [] def __call__( self, @@ -25,13 +23,9 @@ class SyncBridge: context: NativeRequestContext, ) -> dict[str, object]: self.calls.append( - { - "model": request.model, - "audio": request.audio, - "optional_params": request.optional_params, - "bedrock": options.bedrock, - } + {"model": request.model, "audio": request.audio, "optional_params": request.optional_params} ) + self.contexts.append(context) return {"text": "hello"} @@ -48,7 +42,7 @@ class AsyncBridge: def test_enabled_sync_bridge_receives_audio() -> None: bridge = SyncBridge() - rust_bridge.configure_rust_transcription(True, transcription=bridge) + rust_bridge.configure_rust_transcription(transcription=bridge) result = rust_bridge.transcription( model="mistral.voxtral-mini-3b-2507", audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"}, @@ -58,14 +52,22 @@ def test_enabled_sync_bridge_receives_audio() -> None: extra_headers=None, optional_params={"temperature": 0}, timeout=5.0, + stream=True, + has_custom_client=True, + input_source_kind="file", ) - assert result == {"text": "hello"} + assert isinstance(result, Handled) + assert result.value == {"text": "hello"} assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"} + assert bridge.contexts[0].capabilities.execution_mode == "sync" + assert bridge.contexts[0].capabilities.stream is True + assert bridge.contexts[0].capabilities.has_custom_client is True + assert bridge.contexts[0].capabilities.input_source_kind == "file" @pytest.mark.asyncio async def test_enabled_async_bridge() -> None: - rust_bridge.configure_rust_transcription(True, atranscription=AsyncBridge()) + rust_bridge.configure_rust_transcription(atranscription=AsyncBridge()) result = await rust_bridge.atranscription( model="mistral.voxtral-mini-3b-2507", audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"}, @@ -76,7 +78,7 @@ async def test_enabled_async_bridge() -> None: optional_params={}, timeout=None, ) - assert result == {"text": "async"} + assert result == Handled({"text": "async"}) def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None: @@ -87,7 +89,8 @@ def test_loader_returns_none_without_native_extension(monkeypatch: pytest.Monkey def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None) + rust_bridge.configure_rust_transcription(transcription=None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) with pytest.raises(RuntimeError, match="bridge is unavailable"): BedrockAudioTranscriptionRustDispatch().audio_transcriptions( @@ -104,10 +107,8 @@ def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> @pytest.mark.asyncio async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: - async def unavailable(**_: object) -> None: - return None - - monkeypatch.setattr(rust_bridge, "atranscription", unavailable) + rust_bridge.configure_rust_transcription(atranscription=None) + monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: None) with pytest.raises(RuntimeError, match="bridge is unavailable"): await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions( @@ -124,7 +125,7 @@ async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPat def test_bedrock_transcription_uses_rust_only_path() -> None: rust_bridge.configure_rust_transcription( - transcription=lambda request, *, options, context: {"text": "rust"}, + transcription=lambda *_args, **_: {"text": "rust"}, atranscription=None, ) try: @@ -140,9 +141,7 @@ def test_bedrock_transcription_uses_rust_only_path() -> None: @pytest.mark.asyncio async def test_bedrock_atranscription_uses_rust_only_path() -> None: - async def rust_response( - request: NativeTranscriptionRequest, *, options: object, context: NativeRequestContext - ) -> dict[str, object]: + async def rust_response(*_args: object, **_: object) -> dict[str, object]: return {"text": "rust"} rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response) From 81832eee4897d3661511c89b3d2b73f78797d2e3 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 5 Sep 2026 23:04:57 -0700 Subject: [PATCH 4/4] feat(native): preserve request identity and capabilities --- .../core/src/chat_completions/prepare.rs | 2 +- litellm/llms/anthropic/chat/handler.py | 9 +++- .../bedrock/audio_transcription/__init__.py | 34 ++++++++++++- litellm/llms/bedrock/chat/converse_handler.py | 9 +++- litellm/llms/custom_httpx/llm_http_handler.py | 14 ++++++ litellm/main.py | 2 + litellm/rust_bridge/chat_completions.py | 21 +++++--- litellm/rust_bridge/messages.py | 17 ++++--- litellm/rust_bridge/ocr.py | 49 ++++++++----------- litellm/rust_bridge/request.py | 44 +++++++++++++++-- litellm/rust_bridge/responses_websocket.py | 12 +++-- litellm/rust_bridge/transcription.py | 21 +++++--- .../test_rust_bridge_messages.py | 23 --------- .../chat/test_bedrock_converse_handler.py | 25 +++++++--- .../rust_bridge/test_request_context.py | 40 +++++++++++++++ 15 files changed, 231 insertions(+), 91 deletions(-) create mode 100644 tests/test_litellm/rust_bridge/test_request_context.py diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index d22a66ed6c0..92012e9cf4a 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -46,7 +46,7 @@ pub(super) fn resolve_request( ) -> Result { let (model, config) = resolve_provider_config(request.model, options.custom_llm_provider.as_deref()) - .map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?; + .map_err(|_| Error::Declined("provider is not on the rust chat completions path"))?; let messages = parse_messages(request.messages).map_err(|_| Error::Declined("unreadable message list"))?; if messages.is_empty() { diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index a0c84615224..4bb9d0ba338 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -28,7 +28,7 @@ from litellm.llms.custom_httpx.http_handler import ( from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors -from litellm.rust_bridge.request import anthropic_options +from litellm.rust_bridge.request import anthropic_options, request_context from litellm.rust_bridge.runtime import DispatchResult from litellm.types.llms.anthropic import ( ContentBlockDelta, @@ -435,6 +435,11 @@ class AnthropicChatCompletion(BaseLLM): api_key=api_key, additional_args=rust_logging_args, ) + rust_context: Final = request_context( + logging_obj=logging_obj, + request_model=logging_obj.model, + litellm_params=litellm_params, + ) def native_completion() -> DispatchResult[ModelResponse]: return rust_chat_completions_bridge.chat_completions( @@ -452,6 +457,7 @@ class AnthropicChatCompletion(BaseLLM): stream=bool(stream), has_custom_client=client is not None, eligible=serves_via_rust, + context=rust_context, ) async def native_acompletion() -> DispatchResult[ModelResponse]: @@ -470,6 +476,7 @@ class AnthropicChatCompletion(BaseLLM): stream=bool(stream), has_custom_client=client is not None, eligible=serves_via_rust, + context=rust_context, ) @anative_first( diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py index 6fc1a82d2c6..87847953a8f 100644 --- a/litellm/llms/bedrock/audio_transcription/__init__.py +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -1,11 +1,14 @@ import base64 +from io import IOBase from typing import Final, NoReturn import httpx from litellm.litellm_core_utils.audio_utils.utils import process_audio_file +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.rust_bridge import transcription as rust_transcription_bridge from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors +from litellm.rust_bridge.request import request_context from litellm.rust_bridge.runtime import DispatchResult, adapt_result from litellm.types.utils import FileTypes, TranscriptionResponse @@ -19,6 +22,17 @@ async def _aunavailable() -> NoReturn: class BedrockAudioTranscriptionRustDispatch: + @staticmethod + def _input_source_kind(audio_file: FileTypes) -> str: + content: Final = audio_file[1] if isinstance(audio_file, tuple) else audio_file + if isinstance(content, (bytes, bytearray, memoryview)): + return "bytes" + if isinstance(content, IOBase): + return "file" + if isinstance(content, str): + return "path" + return "opaque" + @staticmethod def _audio_payload(audio_file: FileTypes) -> dict[str, object]: processed_audio: Final = process_audio_file(audio_file) @@ -52,6 +66,7 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + logging_obj: Logging | None = None, ) -> DispatchResult[TranscriptionResponse]: result: Final = rust_transcription_bridge.transcription( model=model, @@ -62,13 +77,19 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers=extra_headers, optional_params=optional_params, timeout=timeout, + input_source_kind=self._input_source_kind(audio_file), + context=request_context( + logging_obj=logging_obj, + request_model=logging_obj.model if logging_obj is not None else model, + litellm_params=logging_obj.litellm_params if logging_obj is not None else None, + ), ) return adapt_result(result, lambda response: TranscriptionResponse(**response)) @native_first( native=_attempt_audio_transcriptions, route="audio transcription", - errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: ( + errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout, logging_obj=None: ( provider_errors(custom_llm_provider, model) ), ) @@ -83,6 +104,7 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + logging_obj: Logging | None = None, ) -> TranscriptionResponse: _unavailable() @@ -97,6 +119,7 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + logging_obj: Logging | None = None, ) -> DispatchResult[TranscriptionResponse]: result: Final = await rust_transcription_bridge.atranscription( model=model, @@ -107,13 +130,19 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers=extra_headers, optional_params=optional_params, timeout=timeout, + input_source_kind=self._input_source_kind(audio_file), + context=request_context( + logging_obj=logging_obj, + request_model=logging_obj.model if logging_obj is not None else model, + litellm_params=logging_obj.litellm_params if logging_obj is not None else None, + ), ) return adapt_result(result, lambda response: TranscriptionResponse(**response)) @anative_first( native=_attempt_async_audio_transcriptions, route="audio transcription", - errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout: ( + errors=lambda self, model, audio_file, api_key, api_base, custom_llm_provider, extra_headers, optional_params, timeout, logging_obj=None: ( provider_errors(custom_llm_provider, model) ), ) @@ -128,5 +157,6 @@ class BedrockAudioTranscriptionRustDispatch: extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + logging_obj: Logging | None = None, ) -> TranscriptionResponse: await _aunavailable() diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 982af97a6c4..2d2db6c17a5 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -19,7 +19,7 @@ from litellm.llms.custom_httpx.http_handler import ( from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts from litellm.rust_bridge.dispatch import anative_first, native_first, provider_errors -from litellm.rust_bridge.request import bedrock_options +from litellm.rust_bridge.request import bedrock_options, request_context from litellm.rust_bridge.runtime import DispatchResult from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -425,6 +425,11 @@ class BedrockConverseLLM(BaseAWSLLM): api_key="", additional_args=rust_logging_args, ) + rust_context: Final = request_context( + logging_obj=logging_obj, + request_model=logging_obj.model, + litellm_params=litellm_params, + ) def native_completion() -> DispatchResult[ModelResponse]: return rust_chat_completions_bridge.chat_completions( @@ -442,6 +447,7 @@ class BedrockConverseLLM(BaseAWSLLM): stream=bool(stream), has_custom_client=client is not None, eligible=serves_via_rust, + context=rust_context, ) async def native_acompletion() -> DispatchResult[ModelResponse]: @@ -460,6 +466,7 @@ class BedrockConverseLLM(BaseAWSLLM): stream=bool(stream), has_custom_client=client is not None, eligible=serves_via_rust, + context=rust_context, ) @anative_first( diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 5aa580ae47a..4e7cc94ae8c 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2244,6 +2244,7 @@ class BaseLLMHTTPHandler: stream=stream or False, custom_llm_provider=custom_llm_provider, ), + logging_obj=logging_obj, ) return adapt_result(result, self._rust_anthropic_messages_fake_stream) if stream else result @@ -2398,6 +2399,7 @@ class BaseLLMHTTPHandler: headers: dict, request_body: dict, timeout: float | httpx.Timeout | None, + logging_obj: LiteLLMLoggingObj | None = None, ) -> DispatchResult[AnthropicMessagesResponse]: if custom_llm_provider not in ("azure_ai", "anthropic"): return NativeSkipped(NativeSkipReason.INELIGIBLE) @@ -2409,6 +2411,7 @@ class BaseLLMHTTPHandler: return NativeSkipped(NativeSkipReason.INELIGIBLE) from litellm.rust_bridge import messages as rust_messages_bridge + from litellm.rust_bridge.request import request_context upstream_body: Final = {key: value for key, value in request_body.items() if key != "stream"} result: Final = await rust_messages_bridge.amessages( @@ -2422,6 +2425,11 @@ class BaseLLMHTTPHandler: stream=stream, has_custom_client=has_custom_client, has_agentic_hook=has_agentic_hook, + context=request_context( + logging_obj=logging_obj, + request_model=logging_obj.model if logging_obj is not None else model, + litellm_params=litellm_params.model_dump(), + ), ) def adapt(rust_response: dict[str, object]) -> AnthropicMessagesResponse: @@ -6509,6 +6517,7 @@ class BaseLLMHTTPHandler: ) from litellm.rust_bridge import responses_websocket as rust_responses_websocket + from litellm.rust_bridge.request import request_context async def attempt_connection() -> DispatchResult[ AbstractAsyncContextManager[rust_responses_websocket.ConnectionAdapter] @@ -6519,6 +6528,11 @@ class BaseLLMHTTPHandler: url=ws_url, headers={str(key): str(value) for key, value in headers.items()}, timeout=timeout, + context=request_context( + logging_obj=logging_obj, + request_model=logging_obj.model, + litellm_params=litellm_params.model_dump(), + ), ) @anative_context( diff --git a/litellm/main.py b/litellm/main.py index 2929790f2bd..bf1f669240b 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -7895,6 +7895,7 @@ def transcription( extra_headers=extra_headers, optional_params=optional_params, timeout=timeout, + logging_obj=litellm_logging_obj, ) else: response = dispatch.audio_transcriptions( @@ -7906,6 +7907,7 @@ def transcription( extra_headers=extra_headers, optional_params=optional_params, timeout=timeout, + logging_obj=litellm_logging_obj, ) elif provider_config is not None: response = base_llm_http_handler.audio_transcriptions( diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py index 6af8471e2b3..553b8b60e6a 100644 --- a/litellm/rust_bridge/chat_completions.py +++ b/litellm/rust_bridge/chat_completions.py @@ -36,6 +36,7 @@ from litellm.rust_bridge.request import ( NativeRequestOptions, PreparedNativeCall, call_native, + with_capabilities, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -239,12 +240,15 @@ def chat_completions( stream: bool = False, has_custom_client: bool = False, eligible: bool = True, + context: NativeRequestContext | None = None, ) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) return _build_model_response(rust_response, model_response) - def call(native: RustChatCompletions, prepared: PreparedNativeCall[NativeChatCompletionsRequest]) -> Mapping[str, object]: + def call( + native: RustChatCompletions, prepared: PreparedNativeCall[NativeChatCompletionsRequest] + ) -> Mapping[str, object]: return call_native(native, prepared) return attempt( @@ -262,12 +266,13 @@ def chat_completions( bedrock=bedrock, anthropic=anthropic, ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="sync", stream=stream, has_custom_client=has_custom_client, - ) + ), ), ), call=call, @@ -292,6 +297,7 @@ async def achat_completions( stream: bool = False, has_custom_client: bool = False, eligible: bool = True, + context: NativeRequestContext | None = None, ) -> DispatchResult[ModelResponse]: def adapt(rust_response: Mapping[str, object]) -> ModelResponse: on_response(rust_response) @@ -318,12 +324,13 @@ async def achat_completions( bedrock=bedrock, anthropic=anthropic, ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="async", stream=stream, has_custom_client=has_custom_client, - ) + ), ), ), call=call, diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 445009e4cfb..97ea14f2aa8 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -15,6 +15,7 @@ from litellm.rust_bridge.request import ( NativeRequestOptions, PreparedNativeCall, call_native, + with_capabilities, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -60,6 +61,7 @@ def messages( stream: bool = False, has_custom_client: bool = False, has_agentic_hook: bool = False, + context: NativeRequestContext | None = None, ) -> DispatchResult[dict[str, object]]: return attempt( load=_MESSAGES.load, @@ -74,13 +76,14 @@ def messages( extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="sync", stream=stream, has_custom_client=has_custom_client, has_agentic_hook=has_agentic_hook, - ) + ), ), ), call=call_native, @@ -100,6 +103,7 @@ async def amessages( stream: bool = False, has_custom_client: bool = False, has_agentic_hook: bool = False, + context: NativeRequestContext | None = None, ) -> DispatchResult[dict[str, object]]: return await aattempt( load=_AMESSAGES.load, @@ -114,13 +118,14 @@ async def amessages( extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="async", stream=stream, has_custom_client=has_custom_client, has_agentic_hook=has_agentic_hook, - ) + ), ), ), call=call_native, diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index e2236317853..516a79e385b 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -12,23 +12,20 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.llms.azure_ai.ocr.common_utils import is_azure_document_intelligence_model from litellm.llms.base_llm.ocr.transformation import OCR_REQUEST_FORMAT_PARAM, BaseOCRConfig, OCRResponse from litellm.rust_bridge import configuration as _configuration -from litellm.rust_bridge.bindings import UNCHANGED, NativeBinding, Unchanged +from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.protocols import RustAocr, RustOcr from litellm.rust_bridge.request import ( NativeOCRRequest, NativeRequestCapabilities, - NativeRequestContext, NativeRequestOptions, PreparedNativeCall, call_native, + request_context, vertex_options, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt from litellm.rust_bridge.timeouts import timeout_to_seconds -rust: Final = _configuration.rust -rust_ocr_enabled: Final = _configuration.rust_ocr_enabled - _OCR: Final[NativeBinding[RustOcr]] = NativeBinding(lambda native: native.ocr) _AOCR: Final[NativeBinding[RustAocr]] = NativeBinding(lambda native: native.aocr) _HEADERS: Final = TypeAdapter(dict[str, object]) @@ -66,23 +63,6 @@ _RUST_OCR_PROVIDERS: Final = frozenset( ) -def set_rust_ocr( - *, - ocr: RustOcr | None | Unchanged = UNCHANGED, - aocr: RustAocr | None | Unchanged = UNCHANGED, -) -> None: - if not isinstance(ocr, Unchanged): - if ocr is None: - _OCR.reset() - else: - _OCR.override(ocr) - if not isinstance(aocr, Unchanged): - if aocr is None: - _AOCR.reset() - else: - _AOCR.override(aocr) - - def load_rust_ocr() -> RustOcr | None: return _OCR.load() @@ -109,6 +89,11 @@ def _ocr_input_source_kind(document: dict[str, object]) -> str: return "inline" +def _ocr_request_format(optional_params: dict[str, object]) -> str | None: + value = optional_params.get(OCR_REQUEST_FORMAT_PARAM) + return value if isinstance(value, str) else None + + def _rust_bridge_optional_params( prepared_request: PreparedOCRRequest, resolve_secret: Callable[[str], str | None], @@ -204,7 +189,7 @@ def attempt_ocr( ) -> DispatchResult[OCRResponse]: return attempt( load=_OCR.load, - enabled=rust_ocr_enabled(), + enabled=_configuration.rust_enabled(), prepare=lambda: _prepare_rust_ocr_call( prepared_request=prepared_request, resolve_api_key=resolve_api_key, @@ -225,14 +210,18 @@ def attempt_ocr( timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), vertex=vertex_options(prepared.optional_params), ), - context=NativeRequestContext( + context=request_context( + logging_obj=prepared_request.litellm_logging_obj, + request_model=prepared_request.model, + litellm_params=prepared_request.litellm_params, capabilities=NativeRequestCapabilities( execution_mode="sync", input_source_kind=_ocr_input_source_kind(prepared_request.document), + request_format=_ocr_request_format(prepared_request.optional_params), native_response_format=( prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" ), - ) + ), ), ), ), @@ -247,7 +236,7 @@ async def aattempt_ocr( ) -> DispatchResult[OCRResponse]: return await aattempt( load=_AOCR.load, - enabled=rust_ocr_enabled(), + enabled=_configuration.rust_enabled(), prepare=lambda: _prepare_rust_ocr_call( prepared_request=prepared_request, resolve_api_key=resolve_api_key, @@ -268,14 +257,18 @@ async def aattempt_ocr( timeout_seconds=timeout_to_seconds(prepared_request.effective_timeout), vertex=vertex_options(prepared.optional_params), ), - context=NativeRequestContext( + context=request_context( + logging_obj=prepared_request.litellm_logging_obj, + request_model=prepared_request.model, + litellm_params=prepared_request.litellm_params, capabilities=NativeRequestCapabilities( execution_mode="async", input_source_kind=_ocr_input_source_kind(prepared_request.document), + request_format=_ocr_request_format(prepared_request.optional_params), native_response_format=( prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" ), - ) + ), ), ), ), diff --git a/litellm/rust_bridge/request.py b/litellm/rust_bridge/request.py index e81d10f587f..562ab1f4950 100644 --- a/litellm/rust_bridge/request.py +++ b/litellm/rust_bridge/request.py @@ -1,7 +1,8 @@ from __future__ import annotations from collections.abc import Mapping, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, replace +from types import MappingProxyType from typing import Generic, Protocol, TypeVar @@ -110,6 +111,43 @@ class NativeRequestContext: capabilities: NativeRequestCapabilities = NativeRequestCapabilities() +def request_context( + *, + logging_obj: object | None, + request_model: str, + litellm_params: Mapping[str, object] | None = None, + capabilities: NativeRequestCapabilities | None = None, +) -> NativeRequestContext: + params = litellm_params if litellm_params is not None else MappingProxyType({}) + metadata_value = params.get("metadata") or params.get("litellm_metadata") + metadata = metadata_value if isinstance(metadata_value, Mapping) else MappingProxyType({}) + + def string(name: str) -> str | None: + value = params.get(name, metadata.get(name)) + return value if isinstance(value, str) else None + + call_id = getattr(logging_obj, "litellm_call_id", None) + trace_id = getattr(logging_obj, "litellm_trace_id", None) + return NativeRequestContext( + litellm_call_id=call_id if isinstance(call_id, str) else None, + trace_id=trace_id if isinstance(trace_id, str) else None, + request_model=request_model, + attribution=RequestAttribution( + user_api_key_hash=string("user_api_key_hash"), + user_api_key_user_id=string("user_api_key_user_id"), + user_api_key_team_id=string("user_api_key_team_id"), + ), + capabilities=capabilities or NativeRequestCapabilities(), + ) + + +def with_capabilities( + context: NativeRequestContext, + capabilities: NativeRequestCapabilities, +) -> NativeRequestContext: + return replace(context, capabilities=capabilities) + + RequestT = TypeVar("RequestT") RequestContraT = TypeVar("RequestContraT", contravariant=True) ResultT = TypeVar("ResultT", covariant=True) @@ -152,14 +190,14 @@ class NativeMessagesRequest: @dataclass(frozen=True, slots=True) class NativeOCRRequest: model: str - document: dict[str, object] + document: object optional_params: dict[str, object] @dataclass(frozen=True, slots=True) class NativeTranscriptionRequest: model: str - audio: dict[str, object] + audio: object optional_params: dict[str, object] diff --git a/litellm/rust_bridge/responses_websocket.py b/litellm/rust_bridge/responses_websocket.py index 2dc3e18bd2c..a1ffe107cda 100644 --- a/litellm/rust_bridge/responses_websocket.py +++ b/litellm/rust_bridge/responses_websocket.py @@ -22,6 +22,7 @@ from litellm.rust_bridge.request import ( NativeResponsesWebSocketRequest, PreparedNativeCall, call_native, + with_capabilities, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, adapt_result from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -66,6 +67,7 @@ async def connect( timeout: float | httpx.Timeout | None, websocket_mode: str = "native", requires_connection: bool = True, + context: NativeRequestContext | None = None, ) -> DispatchResult[ConnectionAdapter]: return await aattempt( load=_RESPONSES_WEBSOCKET.load, @@ -74,11 +76,13 @@ async def connect( prepare=lambda: PreparedNativeCall( request=NativeResponsesWebSocketRequest(url=url), options=NativeRequestOptions(extra_headers=headers, timeout_seconds=timeout_to_seconds(timeout)), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( + execution_mode="async", websocket_mode=websocket_mode, requires_connection=requires_connection, - ) + ), ), ), call=lambda connection_type, prepared: call_native(connection_type.connect, prepared), @@ -101,6 +105,7 @@ async def managed_connect( timeout: float | httpx.Timeout | None, websocket_mode: str = "managed", requires_connection: bool = True, + context: NativeRequestContext | None = None, ) -> DispatchResult[AbstractAsyncContextManager[ConnectionAdapter]]: result: Final = await connect( url=url, @@ -108,5 +113,6 @@ async def managed_connect( timeout=timeout, websocket_mode=websocket_mode, requires_connection=requires_connection, + context=context, ) return adapt_result(result, _connection_context) diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py index 97f4eb485d1..8d1997e565c 100644 --- a/litellm/rust_bridge/transcription.py +++ b/litellm/rust_bridge/transcription.py @@ -14,6 +14,7 @@ from litellm.rust_bridge.request import ( PreparedNativeCall, bedrock_options, call_native, + with_capabilities, ) from litellm.rust_bridge.runtime import DispatchResult, aattempt, attempt, identity from litellm.rust_bridge.timeouts import timeout_to_seconds @@ -50,7 +51,7 @@ def load_rust_atranscription() -> RustAtranscription | None: def transcription( *, model: str, - audio: dict[str, object], + audio: object, api_key: str | None, api_base: str | None, custom_llm_provider: str | None, @@ -60,6 +61,7 @@ def transcription( stream: bool = False, has_custom_client: bool = False, input_source_kind: str | None = None, + context: NativeRequestContext | None = None, ) -> DispatchResult[dict[str, object]]: return attempt( load=_TRANSCRIPTION.load, @@ -75,13 +77,14 @@ def transcription( timeout_seconds=timeout_to_seconds(timeout), bedrock=bedrock_options(optional_params), ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="sync", stream=stream, has_custom_client=has_custom_client, input_source_kind=input_source_kind, - ) + ), ), ), call=call_native, @@ -92,7 +95,7 @@ def transcription( async def atranscription( *, model: str, - audio: dict[str, object], + audio: object, api_key: str | None, api_base: str | None, custom_llm_provider: str | None, @@ -102,6 +105,7 @@ async def atranscription( stream: bool = False, has_custom_client: bool = False, input_source_kind: str | None = None, + context: NativeRequestContext | None = None, ) -> DispatchResult[dict[str, object]]: return await aattempt( load=_ATRANSCRIPTION.load, @@ -117,13 +121,14 @@ async def atranscription( timeout_seconds=timeout_to_seconds(timeout), bedrock=bedrock_options(optional_params), ), - context=NativeRequestContext( - capabilities=NativeRequestCapabilities( + context=with_capabilities( + context or NativeRequestContext(), + NativeRequestCapabilities( execution_mode="async", stream=stream, has_custom_client=has_custom_client, input_source_kind=input_source_kind, - ) + ), ), ), call=call_native, 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 7af9a080da6..6327d3d914a 100644 --- a/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py +++ b/tests/test_litellm/anthropic_interface/test_rust_bridge_messages.py @@ -126,16 +126,6 @@ def test_load_rust_messages_returns_injected_impl(): assert rust_messages.load_rust_messages() is bridge -def test_bare_rust_still_toggles_ocr(): - from litellm.rust_bridge.ocr import rust_ocr_enabled - - litellm.rust(True) - assert rust_ocr_enabled() is True - - litellm.rust(False) - assert rust_ocr_enabled() is False - - def test_load_rust_amessages_returns_injected_impl(): bridge = RecordingAsyncMessages() litellm.rust(True) @@ -310,19 +300,6 @@ async def test_gate_uses_process_enable_without_request_override(): assert bridge.calls[0]["custom_llm_provider"] == "azure_ai" -@pytest.mark.asyncio -async def test_gate_ignores_request_flag_when_process_enabled(): - bridge = RecordingAsyncMessages() - litellm.rust(True) - rust_messages.set_rust_messages(amessages=bridge) - - response = await _gate(litellm_params=GenericLiteLLMParams(api_key="sk-azure", rust=False)) - - assert isinstance(response, Handled) - response = response.value - assert len(bridge.calls) == 1 - - @pytest.mark.asyncio async def test_gate_invokes_rust_for_native_anthropic_provider(): bridge = RecordingAsyncMessages() 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..bd1d93e7a3a 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,17 @@ def _inject(*, decline_reason=None, error: Exception | None = None): seen["gate"].append(kwargs) return decline_reason - def native(**kwargs): - seen["call"].append(kwargs) + def native(request, *, options, context): + seen["call"].append( + { + "model": request.model, + "messages": request.messages, + "optional_params": request.optional_params, + "api_key": options.api_key, + "api_base": options.api_base, + "context": context, + } + ) if error is not None: raise error return dict(RUST_RESPONSE) @@ -207,7 +216,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): monkeypatch.setattr("litellm.rust_bridge.bindings.get_native_bridge", lambda: _FakeNative()) - async def declining_native(**_kwargs): + async def declining_native(_request, **_kwargs): raise _Declined("blank message text") bridge.set_rust_chat_completions( @@ -237,7 +246,7 @@ async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): @pytest.mark.asyncio async def test_the_async_path_serves_the_rust_response_without_the_fallback(): - async def native(**_kwargs): + async def native(_request, **_kwargs): return dict(RUST_RESPONSE) bridge.set_rust_chat_completions( @@ -271,7 +280,7 @@ async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - async def declining_native(**_kwargs): + async def declining_native(_request, **_kwargs): raise _Declined("blank message text") logging_obj = MagicMock() @@ -384,7 +393,7 @@ def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(**_kwargs): + def declining_native(_request, **_kwargs): raise _Declined("blank message text") logging_obj = MagicMock() @@ -438,7 +447,7 @@ async def test_post_call_logging_fires_on_the_async_rust_path(): cannot drift apart the way the pre_call suppression once did.""" import json - async def native(**_kwargs): + async def native(_request, **_kwargs): return dict(RUST_RESPONSE) bridge.set_rust_chat_completions( @@ -470,7 +479,7 @@ def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(): RustBridgeDeclined = _Declined RustUpstreamError = type("_Upstream", (Exception,), {}) - def declining_native(**_kwargs): + def declining_native(_request, **_kwargs): raise _Declined("blank message text") logging_obj, calls = _recording_logging_obj() diff --git a/tests/test_litellm/rust_bridge/test_request_context.py b/tests/test_litellm/rust_bridge/test_request_context.py new file mode 100644 index 00000000000..1717fac035c --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_request_context.py @@ -0,0 +1,40 @@ +from types import SimpleNamespace + +from litellm.rust_bridge.request import NativeRequestCapabilities, request_context + + +def test_request_context_preserves_identity_attribution_and_capabilities() -> None: + capabilities = NativeRequestCapabilities(execution_mode="async", stream=True) + + context = request_context( + logging_obj=SimpleNamespace(litellm_call_id="call-1", litellm_trace_id="trace-1"), + request_model="router-alias", + litellm_params={ + "metadata": { + "user_api_key_hash": "hash-1", + "user_api_key_user_id": "user-1", + "user_api_key_team_id": "team-1", + } + }, + capabilities=capabilities, + ) + + assert context.litellm_call_id == "call-1" + assert context.trace_id == "trace-1" + assert context.request_model == "router-alias" + assert context.attribution.user_api_key_hash == "hash-1" + assert context.attribution.user_api_key_user_id == "user-1" + assert context.attribution.user_api_key_team_id == "team-1" + assert context.capabilities is capabilities + + +def test_request_context_ignores_untyped_identity_values() -> None: + context = request_context( + logging_obj=SimpleNamespace(litellm_call_id=1, litellm_trace_id=[]), + request_model="model", + litellm_params={"user_api_key_user_id": 42}, + ) + + assert context.litellm_call_id is None + assert context.trace_id is None + assert context.attribution.user_api_key_user_id is None