diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 1d3bf4bcbb6..7e3d25e9c5d 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -1950,12 +1950,15 @@ dependencies = [ "base64 0.22.1", "bytes", "data-url", + "futures-util", "gcp_auth", "mime_guess", "moka", "rand 0.8.7", "reqwest 0.12.28", "rstest", + "rustls 0.23.42", + "rustls-native-certs", "serde", "serde_json", "serde_path_to_error", @@ -1964,6 +1967,7 @@ dependencies = [ "subtle", "thiserror 2.0.19", "tokio", + "tokio-tungstenite", "tracing", "tracing-subscriber", "url", @@ -1976,7 +1980,6 @@ version = "0.1.0" dependencies = [ "criterion", "futures-util", - "litellm-ai-gateway", "litellm-core", "litellm-python-interop", "litellm-token-counter", diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 74cf66e88a2..dfa61226d4e 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -13,6 +13,11 @@ name = "litellm-ai-gateway" path = "src/main.rs" required-features = ["server"] +[[bin]] +name = "trace-parity-gateway" +path = "src/bin/trace_parity_gateway.rs" +required-features = ["trace-parity"] + [dependencies] tracing.workspace = true litellm-core = { workspace = true, features = ["bedrock-auth"] } diff --git a/litellm-rust/crates/ai-gateway/benchmarks/realtime/README.md b/litellm-rust/crates/ai-gateway/benchmarks/realtime/README.md deleted file mode 100644 index 84e926af243..00000000000 --- a/litellm-rust/crates/ai-gateway/benchmarks/realtime/README.md +++ /dev/null @@ -1,55 +0,0 @@ -# Realtime gateway benchmark — pool on/off - -Measures what the gateway adds over talking to OpenAI's realtime WebSocket -directly, and what the pre-warmed connection pool removes. See -`../../src/routes/realtime/README.md` for how the pool works. - -## Results - -5000 calls / 500 concurrency, gateway at 10 instances, pool ON -(`REALTIME_POOL_SIZE=64`), upstream OpenAI `gpt-realtime`. Each leg run twice. -Times in **ms**. Phases per connection: **dial** = TCP+TLS+WS upgrade, -**session** = upgrade → `session.created` (the phase the pool removes), -**1st-audio** = `response.create` → first audio delta (OpenAI inference), -**total** = full wall-clock. - -| metric | Direct OpenAI | Gateway (pool ON) | Overhead (ms) | vs OpenAI | -| ------------------ | ------------- | ----------------- | ------------- | ---------- | -| success rate (%) | 99.8 | 99.8 | — | — | -| dial p50 (ms) | 276 | 158 | −118 | **faster** | -| session p50 (ms) | 7 | 0 | −7 | **faster** | -| 1st-audio p50 (ms) | 440 | 664 | +224 | slower¹ | -| total p50 (ms) | 816 | 1010 | +194 | slower¹ | -| total p95 (ms) | 2152 | 1970 | −182 | **faster** | -| total p99 (ms) | 2692 | 2610 | −82 | **faster** | - -The gateway is **faster than direct on 4 of 6 metrics**. The warm pool makes the -**session phase sub-millisecond** at the median — ~76% of connects hit the pool, -~70% had session < 1 ms. ¹ The two "slower" rows are not gateway overhead: -`1st-audio` is OpenAI's own inference time (the gateway only relays it), which ran -slower during the gateway legs and drags `total p50` with it. - -**Pool OFF** (control, `REALTIME_POOL_SIZE=0`): session p50 was **367 ms** — the -fresh-dial overhead the pool removes. - -## Reproduce - -The load generator lives in a separate repo: -**https://github.com/ishaan-berri/litellm-realtime-bench** - -```bash -git clone https://github.com/ishaan-berri/litellm-realtime-bench -cd litellm-realtime-bench && go build -o wsbench . - -# Direct to OpenAI (baseline) -./wsbench -host api.openai.com -key "$OPENAI_API_KEY" -m gpt-realtime -n 5000 -c 500 -t 60 - -# Through the gateway — run once with pool ON, once with REALTIME_POOL_SIZE=0 -./wsbench -host -key "$LITELLM_MASTER_KEY" -m gpt-realtime -n 5000 -c 500 -t 60 -``` - -Run the gateway with the env stand-in (`OPENAI_REALTIME_MODEL=gpt-realtime`, -`OPENAI_API_KEY`, `LITELLM_MASTER_KEY`, `REALTIME_POOL_SIZE`, `HOST=0.0.0.0`). At -500 concurrency over N instances, size the pool to `≈ 500 / N` per instance (64 was -used here for 10 instances). The bench repo's README covers running 500-concurrency -legs from a hosted multi-vCPU runner. **Never commit keys — pass them via `-key`.** 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 6f48f38c9f6..b17f17de11f 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs @@ -277,7 +277,7 @@ fn core_error_kind(error: &Error) -> &'static str { Error::InvalidProvider(_) => "InvalidProvider", Error::InvalidRequest(_) => "InvalidRequest", Error::InvalidType { .. } => "InvalidType", - Error::MissingField(_) => "MissingField", + Error::MissingField(_) | Error::MissingDocumentUrl => "MissingField", Error::Http { .. } => "HttpError", Error::InvalidResponse(_) => "InvalidResponse", Error::Network(_) => "NetworkError", diff --git a/litellm-rust/crates/ai-gateway/src/bin/trace_parity_gateway.rs b/litellm-rust/crates/ai-gateway/src/bin/trace_parity_gateway.rs new file mode 100644 index 00000000000..9036deb9871 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/bin/trace_parity_gateway.rs @@ -0,0 +1,40 @@ +use std::io::Read; + +use serde::Deserialize; +use serde_json::Value; + +#[derive(Deserialize)] +struct Input { + model_alias: String, + provider_model: String, + api_base: String, + body: Value, +} + +#[tokio::main] +async fn main() { + let mut input = String::new(); + if let Err(error) = std::io::stdin().read_to_string(&mut input) { + fail(error); + } + let input: Input = match serde_json::from_str(&input) { + Ok(input) => input, + Err(error) => fail(error), + }; + let result = litellm_ai_gateway::trace_parity::traced_messages_request( + input.model_alias, + input.provider_model, + input.api_base, + input.body, + ) + .await; + match serde_json::to_string(&result) { + Ok(result) => println!("{result}"), + Err(error) => fail(error), + } +} + +fn fail(error: impl std::fmt::Display) -> ! { + eprintln!("{error}"); + std::process::exit(1) +} 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 7f3b6b0650f..f86dd778424 100644 --- a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs +++ b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs @@ -1,5 +1,3 @@ -use std::collections::HashMap; -use std::sync::Arc; use std::time::Duration; use futures_util::stream::{SplitSink, SplitStream}; @@ -10,106 +8,21 @@ use litellm_core::auth::error::MissingCredential; use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG; use litellm_core::responses::types::ResponsesWsEvent; use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig; -use tokio::net::TcpStream; -use tokio::sync::Mutex; use tokio_tungstenite::tungstenite::Message; use tokio_tungstenite::tungstenite::client::IntoClientRequest; use tokio_tungstenite::tungstenite::http::HeaderValue; -use tokio_tungstenite::tungstenite::http::header::{AUTHORIZATION, HeaderName}; -use tokio_tungstenite::{MaybeTlsStream, WebSocketStream}; +use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION; -use crate::io::tls::connect_upstream; +use litellm_core::responses::websocket::{ResponsesUpstreamWs, connect_upstream}; use crate::constants::{ DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS, }; const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY"; -pub type ResponsesUpstreamWs = WebSocketStream>; type UpstreamTx = SplitSink; type UpstreamRx = SplitStream; -#[derive(Clone)] -pub struct ResponsesWebSocketConnection { - socket: Arc>>, -} - -impl ResponsesWebSocketConnection { - pub async fn connect_url( - url: &str, - headers: &HashMap, - timeout: Option, - ) -> Result { - let mut request = url - .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) - .map_err(|error| Error::InvalidRequest(error.to_string()))?; - request.headers_mut().insert(header_name, header_value); - } - let connect = connect_upstream(request); - let result = match timeout { - Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| { - Error::Network("Responses WebSocket connection timed out".to_string()) - })?, - None => connect.await, - }; - let (socket, _) = result.map_err(|error| match *error { - tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http { - status: response.status().as_u16(), - body: String::new(), - }, - other => Error::Network(other.to_string()), - })?; - Ok(Self { - socket: Arc::new(Mutex::new(Some(socket))), - }) - } - - pub async fn send_text(&self, text: String) -> Result<(), Error> { - let mut socket = self.socket.lock().await; - let Some(socket) = socket.as_mut() else { - return Err(Error::Network("Responses WebSocket is closed".to_string())); - }; - socket - .send(Message::Text(text)) - .await - .map_err(|error| Error::Network(error.to_string())) - } - - pub async fn recv_text(&self) -> Result, Error> { - let mut socket_guard = self.socket.lock().await; - let Some(socket) = socket_guard.as_mut() else { - return Ok(None); - }; - match socket.next().await { - Some(Ok(Message::Text(text))) => Ok(Some(text)), - Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec()) - .map(Some) - .map_err(|error| Error::InvalidResponse(error.to_string())), - Some(Ok(Message::Close(_))) | None => Ok(None), - Some(Ok(_)) => Ok(None), - Some(Err(error)) => Err(Error::Network(error.to_string())), - } - } - - pub async fn close(&self) -> Result<(), Error> { - let mut socket = self.socket.lock().await; - if let Some(socket) = socket.as_mut() { - socket - .close(None) - .await - .map_err(|error| Error::Network(error.to_string()))?; - } - *socket = None; - Ok(()) - } -} - pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result { api_key .map(str::trim) diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index 39465e28e84..3334053a0a4 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -118,7 +118,8 @@ impl IntoResponse for MessagesRouteError { | Error::Connect(_) | Error::InvalidResponse(_) | Error::InvalidType { .. } - | Error::MissingField(_) => ( + | Error::MissingField(_) + | Error::MissingDocumentUrl => ( StatusCode::BAD_GATEWAY, "messages provider request failed".to_string(), ), diff --git a/litellm-rust/crates/ai-gateway/src/trace_parity.rs b/litellm-rust/crates/ai-gateway/src/trace_parity.rs index 21123df3f1c..00c9b53e691 100644 --- a/litellm-rust/crates/ai-gateway/src/trace_parity.rs +++ b/litellm-rust/crates/ai-gateway/src/trace_parity.rs @@ -10,6 +10,7 @@ use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter}; use serde::Serialize; use serde_json::Value; use tower::ServiceExt; +use tracing::instrument::WithSubscriber; use crate::io::realtime_pool::RealtimePool; use crate::routes; @@ -21,6 +22,38 @@ pub struct GatewayResponse { pub body: Value, } +#[derive(Debug, Serialize)] +pub struct TracedGatewayResponse { + pub response: Option, + pub error: Option, + pub trace: Vec, +} + +pub async fn traced_messages_request( + model_alias: String, + provider_model: String, + api_base: String, + body: Value, +) -> TracedGatewayResponse { + let trace = litellm_core::observability::FunctionTrace::default(); + let result = messages_request(model_alias, provider_model, api_base, body) + .with_subscriber(trace.dispatcher()) + .await; + let events = trace.events(); + match result { + Ok(response) => TracedGatewayResponse { + response: Some(response), + error: None, + trace: events, + }, + Err(error) => TracedGatewayResponse { + response: None, + error: Some(error.to_string()), + trace: events, + }, + } +} + pub async fn messages_request( model_alias: String, provider_model: String, diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 0730565a5f1..09c526f73cf 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -8,6 +8,7 @@ autotests = false [dependencies] bytes.workspace = true +futures-util.workspace = true base64.workspace = true azure_core.workspace = true azure_identity.workspace = true @@ -17,12 +18,15 @@ moka.workspace = true mime_guess = "2.0.5" rand.workspace = true reqwest.workspace = true +rustls.workspace = true +rustls-native-certs.workspace = true serde.workspace = true serde_json.workspace = true serde_path_to_error = "0.1" strum.workspace = true subtle.workspace = true -tokio.workspace = true +tokio = { workspace = true, features = ["sync"] } +tokio-tungstenite.workspace = true thiserror.workspace = true tracing.workspace = true tracing-subscriber = { workspace = true, optional = true } diff --git a/litellm-rust/crates/core/src/call_lifecycle/host.rs b/litellm-rust/crates/core/src/call_lifecycle/host.rs index 0b00d9bf758..ac6ddf99b9e 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/host.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/host.rs @@ -1,3 +1,30 @@ +use std::future::Future; +use std::pin::Pin; + +pub enum HostCallStep { + Host(O), + Complete(C), +} + +pub type HostCallFuture<'a, O, C> = + Pin, crate::Error>> + Send + 'a>>; + +pub trait HostCall: Send + Sync { + type Operation: Send + 'static; + type Result: Send + 'static; + type Complete: Send + 'static; + + fn resume( + &mut self, + result: Option, + ) -> HostCallFuture<'_, Self::Operation, Self::Complete>; + + fn interrupt( + &mut self, + failure: HostFailure, + ) -> HostCallFuture<'_, Self::Operation, Self::Complete>; +} + pub enum HostStep { Ready(V), Suspend(S), diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index 1d1634dae82..359ad56c336 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -9,6 +9,8 @@ pub enum Error { }, #[error("missing required field: {0}")] MissingField(&'static str), + #[error("Document URL is required")] + MissingDocumentUrl, #[error("invalid response: {0}")] InvalidResponse(String), #[error("invalid provider: {0}")] @@ -56,6 +58,8 @@ impl Error { pub const fn http_status_code(&self) -> Option { match self { Self::InvalidRequest(_) => Some(400), + Self::MissingDocumentUrl => Some(500), + Self::Http { status, .. } => Some(*status), _ => None, } } @@ -115,6 +119,7 @@ impl From for Error { fn from(error: crate::ocr::error::OcrRequestError) -> Self { match error { crate::ocr::error::OcrRequestError::MissingField(field) => Self::MissingField(field), + crate::ocr::error::OcrRequestError::MissingDocumentUrl => Self::MissingDocumentUrl, error => Self::InvalidRequest(error.to_string()), } } diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index 9f45ab78091..394ca778d2f 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -182,3 +182,29 @@ pub(crate) fn transport_error(error: reqwest::Error) -> Error { } crate::error::TransportError::from(error).into() } + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn request_timeout_has_an_http_408_status() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let _connection = listener.accept().await.unwrap(); + tokio::time::sleep(Duration::from_secs(1)).await; + }); + let error = reqwest::Client::new() + .get(format!("http://{address}")) + .timeout(Duration::from_millis(10)) + .send() + .await + .unwrap_err(); + assert!(matches!( + transport_error(error), + Error::Http { status: 408, .. } + )); + server.abort(); + } +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs index c3a341a2921..7e8ce63b379 100644 --- a/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs +++ b/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs @@ -12,7 +12,7 @@ pub(crate) fn transform_ocr_request( params: &DeepSeekOcrParams, ) -> Result { if document.source().is_empty() { - return Err(OcrRequestError::MissingField("document URL")); + return Err(OcrRequestError::MissingDocumentUrl); } let content = OcrDocument::ImageUrl { image_url: document.source().to_string(), diff --git a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs index 4b56e8d4034..f76a7c2b232 100644 --- a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs +++ b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs @@ -13,7 +13,7 @@ pub(crate) fn transform_ocr_request( ) -> Result { let source = document.source(); if source.is_empty() { - return Err(OcrRequestError::MissingField("document URL")); + return Err(OcrRequestError::MissingDocumentUrl); } Ok(if let Some(document) = InlineDocument::parse(source)? { DocumentIntelligenceRequest::Base64Source( diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index 817eefda13a..82a32ac1ab5 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -172,9 +172,11 @@ fn map_media_error(error: MediaError) -> OcrError { body: "OCR document download failed".into(), } .into(), - MediaError::Timeout => { - TransportError::Network("OCR document download timed out".into()).into() + MediaError::Timeout => TransportError::Http { + status: 408, + body: "OCR document download timed out".into(), } + .into(), MediaError::Transport(error) => error.into(), } } diff --git a/litellm-rust/crates/core/src/ocr/error.rs b/litellm-rust/crates/core/src/ocr/error.rs index e5512b08589..55ea2cbcdae 100644 --- a/litellm-rust/crates/core/src/ocr/error.rs +++ b/litellm-rust/crates/core/src/ocr/error.rs @@ -18,6 +18,8 @@ pub enum OcrRequestError { RequestField { path: String }, #[error("missing required field: {0}")] MissingField(&'static str), + #[error("Document URL is required")] + MissingDocumentUrl, #[error("invalid OCR document data URI")] InvalidDataUri, #[error( diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs index 0f272217faf..3e7507e9ed5 100644 --- a/litellm-rust/crates/core/src/ocr/hooks.rs +++ b/litellm-rust/crates/core/src/ocr/hooks.rs @@ -36,7 +36,7 @@ pub struct OcrPostCallRequest { } pub trait OcrHooks: Send + Sync { - fn has_guardrails(&self) -> bool { + fn intercepts_requests(&self) -> bool { false } fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> { @@ -91,7 +91,7 @@ impl CallLifecycleHooks Self::PreCallFuture<'a> { Box::pin(async move { - if !self.hooks.has_guardrails() { + if !self.hooks.intercepts_requests() { return Ok(request); } let changed = self diff --git a/litellm-rust/crates/core/src/ocr/lifecycle.rs b/litellm-rust/crates/core/src/ocr/lifecycle.rs index 4e23f852eea..92c9d4b717c 100644 --- a/litellm-rust/crates/core/src/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/src/ocr/lifecycle.rs @@ -13,7 +13,9 @@ use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient}; use crate::AuthError; use crate::Error; use crate::auth::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; -use crate::call_lifecycle::host::{HostFailure, HostLifecycle, HostPhase}; +use crate::call_lifecycle::host::{ + HostCall, HostCallFuture, HostCallStep, HostFailure, HostLifecycle, HostPhase, +}; use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming}; pub type NativeResult = Result, Error>; @@ -34,7 +36,6 @@ pub enum OcrDecline { pub struct OcrAdmission { pub provider_workflow: bool, pub host_operations: bool, - pub azure_ad_token_provider: bool, pub asynchronous: bool, } @@ -43,7 +44,6 @@ impl OcrAdmission { Self { provider_workflow: true, host_operations: true, - azure_ad_token_provider: false, asynchronous: false, } } @@ -71,6 +71,17 @@ pub enum OcrHostOperation { PostCall(OcrPostCallRequest), } +impl OcrHostOperation { + pub const fn phase(&self) -> Option { + match self { + Self::Lifecycle(phase) => Some(*phase), + Self::Success { .. } => Some(HostPhase::Success), + Self::Failure { .. } => Some(HostPhase::Failure), + _ => None, + } + } +} + pub enum OcrHostResult { Request(Result<(Box, bool), Error>), Lifecycle(Result<(), HostFailure>), @@ -80,10 +91,7 @@ pub enum OcrHostResult { PostCall(Result), } -pub enum OcrCallStep { - Host(OcrHostOperation), - Complete(LiteLLMOcrResponse), -} +pub type OcrCallStep = HostCallStep; pub struct OcrCall { lifecycle: HostLifecycle, @@ -105,7 +113,7 @@ impl OcrCall { } NativeOutcome::Completed(Self { lifecycle: HostLifecycle::new(admission.asynchronous), - execution: OcrExecution::new(client, admission.azure_ad_token_provider), + execution: OcrExecution::new(client), response: None, error: None, pending: false, @@ -277,6 +285,26 @@ impl OcrCall { } } +impl HostCall for OcrCall { + type Operation = OcrHostOperation; + type Result = OcrHostResult; + type Complete = LiteLLMOcrResponse; + + fn resume( + &mut self, + result: Option, + ) -> HostCallFuture<'_, Self::Operation, Self::Complete> { + Box::pin(OcrCall::resume(self, result)) + } + + fn interrupt( + &mut self, + failure: HostFailure, + ) -> HostCallFuture<'_, Self::Operation, Self::Complete> { + Box::pin(OcrCall::interrupt(self, failure)) + } +} + struct PendingOperation { operation: OcrHostOperation, result: oneshot::Sender, @@ -295,7 +323,7 @@ struct OcrExecution { } impl OcrExecution { - fn new(client: OcrClient, azure_ad_token_provider: bool) -> Self { + fn new(client: OcrClient) -> Self { let (operations_tx, operations_rx) = mpsc::unbounded_channel(); Self { client: Some(client), @@ -305,7 +333,7 @@ impl OcrExecution { pending_result: None, execution: None, completed: false, - azure_ad_token_provider, + azure_ad_token_provider: false, terminal: Arc::default(), } } @@ -360,7 +388,7 @@ impl OcrExecution { .request .take() .expect("admitted OCR call has a request"); - let has_guardrails = request.hooks.has_guardrails(); + let intercepts_requests = request.hooks.intercepts_requests(); if self.azure_ad_token_provider { request.azure_ad_token_provider = Some(TokenProviderHandle::new(Arc::new( OcrAzureAdTokenProvider { @@ -370,7 +398,7 @@ impl OcrExecution { } request.hooks = Arc::new(ProtocolHooks { operations: self.operations_tx.clone(), - has_guardrails, + intercepts_requests, terminal: self.terminal.clone(), }); self.execution = Some(tokio::spawn(async move { @@ -404,7 +432,7 @@ impl Drop for OcrExecution { struct ProtocolHooks { operations: mpsc::UnboundedSender, - has_guardrails: bool, + intercepts_requests: bool, terminal: Arc>>, } @@ -452,8 +480,8 @@ impl ProtocolHooks { } impl OcrHooks for ProtocolHooks { - fn has_guardrails(&self) -> bool { - self.has_guardrails + fn intercepts_requests(&self) -> bool { + self.intercepts_requests } fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> { diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index ee3ae2e78cb..9934a1d9a14 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -69,7 +69,7 @@ pub(crate) async fn transform_request_body( where B: Serialize + DeserializeOwned, { - let (body, headers) = if request.hooks.has_guardrails() { + let (body, headers) = if request.hooks.intercepts_requests() { let body = serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField { path: "body".into(), })?; @@ -129,7 +129,7 @@ pub(crate) async fn guardrail_document( url: &str, headers: &[(String, String)], ) -> Result<(OcrDocument, Vec<(String, String)>), OcrError> { - if !request.hooks.has_guardrails() { + if !request.hooks.intercepts_requests() { return Ok((request.document.clone(), headers.to_vec())); } let changed = request diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index ec40ffaf291..e00a6e6d3a3 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -3,7 +3,6 @@ use crate::ocr::error::OcrResponseError; use std::collections::BTreeMap; use std::time::Duration; -use super::hooks::{OcrDuringCallRequest, OcrPreCallRequest}; use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument}; use crate::Error; use crate::auth::InputSource; @@ -54,6 +53,12 @@ const VERTEX_AUTH_OPTION_FIELDS: &[&str] = &[ "vertex_ai_location", ]; +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct OptionalParamSpec { + pub name: &'static str, + pub secret: bool, +} + #[derive(Debug)] pub struct DecodedOcrResponse { pub data: T, @@ -113,11 +118,33 @@ pub fn consumed_optional_param_names( .collect()) } +pub fn consumed_optional_params( + model: &str, + custom_llm_provider: Option<&str>, +) -> Result, Error> { + consumed_optional_param_names(model, custom_llm_provider).map(|names| { + names + .into_iter() + .map(|name| OptionalParamSpec { + name, + secret: matches!( + name, + "azure_ad_token" + | "client_secret" + | "azure_federated_token_file" + | "vertex_credentials" + | "vertex_ai_credentials" + ), + }) + .collect() + }) +} + pub fn decode_request(wire: OcrWireRequest) -> Result { let api_key_source = source_for(&wire.input_sources, "api_key"); let api_base_source = source_for(&wire.input_sources, "api_base"); let extra_headers_source = source_for(&wire.input_sources, "extra_headers"); - let document = decode_request_value(wire.document, "document")?; + let document = decode_document(wire.document)?; let headers = wire .extra_headers .unwrap_or_default() @@ -182,6 +209,16 @@ pub fn decode_request(wire: OcrWireRequest) -> Result }) } +fn decode_document(value: Value) -> Result { + let kind = value.get("type").and_then(Value::as_str); + let missing_url = matches!(kind, Some("document_url")) && value.get("document_url").is_none() + || matches!(kind, Some("image_url")) && value.get("image_url").is_none(); + if missing_url { + return Err(OcrRequestError::MissingDocumentUrl); + } + decode_request_value(value, "document") +} + fn source_for(sources: &BTreeMap, name: &str) -> InputSource { sources.get(name).copied().unwrap_or_default() } @@ -233,39 +270,6 @@ pub fn decode_response( }) } -pub fn decode_pre_call_result( - original: OcrPreCallRequest, - value: Value, -) -> Result { - #[derive(Deserialize)] - struct Changed { - document: OcrDocument, - #[serde(default)] - optional_params: Map, - } - let changed: Changed = decode_request_value(value, "guardrail")?; - Ok(OcrPreCallRequest { - document: changed.document, - optional_params: Value::Object(changed.optional_params), - ..original - }) -} - -pub fn decode_during_call_result( - original: OcrDuringCallRequest, - value: Value, -) -> Result { - #[derive(Deserialize)] - struct Changed { - body: Value, - } - let changed: Changed = decode_request_value(value, "guardrail")?; - Ok(OcrDuringCallRequest { - body: changed.body, - ..original - }) -} - #[cfg(test)] mod tests { use super::*; @@ -283,4 +287,57 @@ mod tests { assert!(vertex.contains(&"vertex_credentials")); assert!(!vertex.contains(&"pages")); } + + #[test] + fn optional_param_metadata_marks_only_credentials_as_secret() { + let azure = consumed_optional_params("model", Some("azure_ai")).unwrap(); + assert!( + azure + .iter() + .any(|spec| spec.name == "client_secret" && spec.secret) + ); + assert!( + azure + .iter() + .any(|spec| spec.name == "tenant_id" && !spec.secret) + ); + let vertex = consumed_optional_params("deepseek-ocr", Some("vertex_ai")).unwrap(); + assert!( + vertex + .iter() + .any(|spec| spec.name == "vertex_credentials" && spec.secret) + ); + assert!( + vertex + .iter() + .any(|spec| spec.name == "vertex_project" && !spec.secret) + ); + } + + #[test] + fn activation_includes_migrated_providers() { + assert!(is_supported_request("model", Some("mistral"))); + assert!(is_supported_request("pixtral-12b", Some("azure_ai"))); + assert!(is_supported_request( + "documentintelligence/prebuilt-read", + Some("azure_ai") + )); + assert!(is_supported_request("parse-v3", Some("reducto"))); + assert!(is_supported_request("parse-legacy", Some("reducto"))); + assert!(is_supported_request("mistral-ocr", Some("vertex_ai"))); + assert!(is_supported_request("deepseek-ocr", Some("vertex_ai"))); + } + + #[test] + fn missing_document_source_has_a_typed_public_error() { + for document in [ + serde_json::json!({"type": "document_url"}), + serde_json::json!({"type": "image_url"}), + ] { + assert_eq!( + decode_document(document), + Err(OcrRequestError::MissingDocumentUrl) + ); + } + } } diff --git a/litellm-rust/crates/core/src/responses/websocket.rs b/litellm-rust/crates/core/src/responses/websocket.rs index 5d037e9cf1b..34213e5f6c4 100644 --- a/litellm-rust/crates/core/src/responses/websocket.rs +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -1,3 +1,21 @@ +use std::collections::HashMap; +use std::io; +use std::sync::{Arc, OnceLock}; +use std::time::Duration; + +use futures_util::{SinkExt, StreamExt}; +use rustls::{ClientConfig, RootCertStore}; +use tokio::net::TcpStream; +use tokio::sync::Mutex; +use tokio_tungstenite::tungstenite::Message; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::error::TlsError; +use tokio_tungstenite::tungstenite::handshake::client::Response; +use tokio_tungstenite::tungstenite::http::{HeaderName, HeaderValue}; +use tokio_tungstenite::{ + Connector, MaybeTlsStream, WebSocketStream, connect_async_tls_with_config, +}; + use crate::Error; use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH}; use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult}; @@ -125,6 +143,137 @@ pub fn is_terminal_event(event_type: &ResponsesWsEventType) -> bool { ) } +pub type ResponsesUpstreamWs = WebSocketStream>; + +static TLS_CONFIG: OnceLock> = OnceLock::new(); + +fn build_tls_config() -> Result> { + let native = rustls_native_certs::load_native_certs(); + let mut store = RootCertStore::empty(); + let (added, _ignored) = store.add_parsable_certificates(native.certs); + if added == 0 { + return Err(Box::new(tokio_tungstenite::tungstenite::Error::Io( + io::Error::other(format!( + "no usable native root certificates: {:?}", + native.errors + )), + ))); + } + ClientConfig::builder_with_provider(Arc::new(rustls::crypto::ring::default_provider())) + .with_safe_default_protocol_versions() + .map(|builder| builder.with_root_certificates(store).with_no_client_auth()) + .map_err(|error| { + Box::new(tokio_tungstenite::tungstenite::Error::Tls( + TlsError::Rustls(error), + )) + }) +} + +fn tls_config() -> Result, Box> { + if let Some(config) = TLS_CONFIG.get() { + return Ok(Arc::clone(config)); + } + let built = Arc::new(build_tls_config()?); + Ok(Arc::clone(TLS_CONFIG.get_or_init(|| built))) +} + +pub async fn connect_upstream( + request: R, +) -> Result<(ResponsesUpstreamWs, Response), Box> +where + R: IntoClientRequest + Unpin, +{ + let request = request.into_client_request().map_err(Box::new)?; + let connector = match request.uri().scheme_str() { + Some("wss") => Some(Connector::Rustls(tls_config()?)), + _ => None, + }; + connect_async_tls_with_config(request, None, false, connector) + .await + .map_err(Box::new) +} + +#[derive(Clone)] +pub struct ResponsesWebSocketConnection { + socket: Arc>>, +} + +impl ResponsesWebSocketConnection { + pub async fn connect_url( + url: &str, + headers: &HashMap, + timeout: Option, + ) -> Result { + let mut request = url + .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) + .map_err(|error| Error::InvalidRequest(error.to_string()))?; + request.headers_mut().insert(header_name, header_value); + } + let connect = connect_upstream(request); + let result = match timeout { + Some(timeout) => tokio::time::timeout(timeout, connect) + .await + .map_err(|_| Error::Network("Responses WebSocket connection timed out".into()))?, + None => connect.await, + }; + let (socket, _) = result.map_err(|error| match *error { + tokio_tungstenite::tungstenite::Error::Http(response) => Error::Http { + status: response.status().as_u16(), + body: String::new(), + }, + other => Error::Network(other.to_string()), + })?; + Ok(Self { + socket: Arc::new(Mutex::new(Some(socket))), + }) + } + + pub async fn send_text(&self, text: String) -> Result<(), Error> { + let mut socket = self.socket.lock().await; + let Some(socket) = socket.as_mut() else { + return Err(Error::Network("Responses WebSocket is closed".into())); + }; + socket + .send(Message::Text(text)) + .await + .map_err(|error| Error::Network(error.to_string())) + } + + pub async fn recv_text(&self) -> Result, Error> { + let mut socket = self.socket.lock().await; + let Some(socket) = socket.as_mut() else { + return Ok(None); + }; + match socket.next().await { + Some(Ok(Message::Text(text))) => Ok(Some(text)), + Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec()) + .map(Some) + .map_err(|error| Error::InvalidResponse(error.to_string())), + Some(Ok(Message::Close(_))) | None => Ok(None), + Some(Ok(_)) => Ok(None), + Some(Err(error)) => Err(Error::Network(error.to_string())), + } + } + + pub async fn close(&self) -> Result<(), Error> { + let mut socket = self.socket.lock().await; + if let Some(socket) = socket.as_mut() { + socket + .close(None) + .await + .map_err(|error| Error::Network(error.to_string()))?; + } + *socket = None; + Ok(()) + } +} + #[cfg(test)] mod tests { use super::*; diff --git a/litellm-rust/crates/core/tests/azure_ai_ocr.rs b/litellm-rust/crates/core/tests/azure_ai_ocr.rs index d7d532cfef1..b6dc8d90b93 100644 --- a/litellm-rust/crates/core/tests/azure_ai_ocr.rs +++ b/litellm-rust/crates/core/tests/azure_ai_ocr.rs @@ -70,7 +70,7 @@ async fn facade_acquires_supplied_entra_token_for_final_request() { struct ReplaceBodyDocument; impl OcrHooks for ReplaceBodyDocument { - fn has_guardrails(&self) -> bool { + fn intercepts_requests(&self) -> bool { true } diff --git a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs index 463cee994a4..3fca59033cc 100644 --- a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs +++ b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs @@ -419,7 +419,7 @@ async fn pre_call_guardrail_receives_caller_pages_before_mapping() { struct RewritePages; impl OcrHooks for RewritePages { - fn has_guardrails(&self) -> bool { + fn intercepts_requests(&self) -> bool { true } diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index cbc7f949ce7..ac64681bb01 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -132,7 +132,7 @@ struct RecordingHooks { } impl OcrHooks for RecordingHooks { - fn has_guardrails(&self) -> bool { + fn intercepts_requests(&self) -> bool { true } @@ -189,7 +189,7 @@ impl OcrHooks for RecordingHooks { struct HeaderEditHooks; impl OcrHooks for HeaderEditHooks { - fn has_guardrails(&self) -> bool { + fn intercepts_requests(&self) -> bool { true } @@ -283,7 +283,7 @@ struct AdmissionSpy { } impl OcrHooks for AdmissionSpy { - fn has_guardrails(&self) -> bool { + fn intercepts_requests(&self) -> bool { *self.effects.lock().unwrap() += 1; true } @@ -301,7 +301,6 @@ fn admission_declines_without_invoking_hooks_or_transport() { OcrAdmission { provider_workflow: false, host_operations: true, - azure_ad_token_provider: false, asynchronous: false, }, OcrDecline::ProviderWorkflow, @@ -310,7 +309,6 @@ fn admission_declines_without_invoking_hooks_or_transport() { OcrAdmission { provider_workflow: true, host_operations: false, - azure_ad_token_provider: false, asynchronous: false, }, OcrDecline::HostOperations, diff --git a/litellm-rust/crates/core/tests/reducto_ocr.rs b/litellm-rust/crates/core/tests/reducto_ocr.rs index 80d07c49867..a15e9cae5b5 100644 --- a/litellm-rust/crates/core/tests/reducto_ocr.rs +++ b/litellm-rust/crates/core/tests/reducto_ocr.rs @@ -248,7 +248,7 @@ async fn facade_omits_native_response_by_default_and_preserves_auth_priority() { struct RewriteDocument; impl OcrHooks for RewriteDocument { - fn has_guardrails(&self) -> bool { + fn intercepts_requests(&self) -> bool { true } diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index e432cab28dc..42fad740870 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -17,7 +17,6 @@ panic-test = [] trace-parity = [ "dep:tracing", "litellm-core/observability", - "litellm-ai-gateway/trace-parity", ] [dependencies] @@ -25,7 +24,6 @@ futures-util.workspace = true tracing = { workspace = true, optional = true } litellm-core = { workspace = true, features = ["bedrock-auth"] } litellm-token-counter.workspace = true -litellm-ai-gateway = { workspace = true, default-features = false } litellm-python-interop.workspace = true pyo3.workspace = true pyo3-async-runtimes.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 0ef04bbcb7b..701c6abb68c 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -22,7 +22,8 @@ pub(crate) fn core_error_to_pyerr(err: Error) -> PyErr { Error::InvalidProvider(_) | Error::InvalidRequest(_) | Error::InvalidType { .. } - | Error::MissingField(_) => PyValueError::new_err(err.to_string()), + | Error::MissingField(_) + | Error::MissingDocumentUrl => PyValueError::new_err(err.to_string()), other => PyRuntimeError::new_err(other.to_string()), } } @@ -41,6 +42,7 @@ pub(crate) fn chat_completions_error_to_pyerr(err: Error) -> PyErr { | Error::InvalidRequest(_) | Error::InvalidType { .. } | Error::MissingField(_) + | Error::MissingDocumentUrl | Error::MissingApiKey { .. } | Error::MissingAzureAiCredentials | Error::MissingAzureDocumentIntelligenceCredentials @@ -49,9 +51,7 @@ pub(crate) fn chat_completions_error_to_pyerr(err: Error) -> PyErr { // Nothing reached the provider, so serving it on Python cannot double // bill and is the only way the caller gets an answer at all. | Error::Connect(_) => RustBridgeDeclined::new_err(err.to_string()), - Error::Http { status, body } => { - RustUpstreamError::new_err((status, format!("{status}: {body}"))) - } + Error::Http { status, body } => RustUpstreamError::new_err((status, body)), Error::Network(message) | Error::InvalidResponse(message) => { RustUpstreamError::new_err((0u16, message)) } diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 5a5aec72362..12bc57a8931 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -10,7 +10,7 @@ mod marshal; mod routes; mod token_counter; -use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; +use litellm_core::responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use pyo3::prelude::*; use pyo3::types::PyAny; use serde_json::Value; @@ -66,7 +66,7 @@ impl ResponsesWebSocketConnection { } } -#[pymodule(gil_used = false)] +#[pymodule(gil_used = true)] mod _native { use pyo3::prelude::*; @@ -154,7 +154,6 @@ mod tests { "amessages", "chat_completions", "achat_completions", - "gateway_messages", ] ); } diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs index 8e5a520b3d8..65e4f1c3fb5 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs @@ -1,10 +1,12 @@ -use std::future::Future; -use std::pin::Pin; use std::sync::Arc; use std::task::Poll; use futures_util::future::{AbortHandle, Abortable}; -use litellm_core::call_lifecycle::host::{HostFailure, HostPhase, HostStep}; +#[cfg(test)] +use litellm_core::call_lifecycle::host::HostCallFuture; +use litellm_core::call_lifecycle::host::{ + HostCall as NativeCall, HostCallStep as NativeCallStep, HostFailure, HostPhase, HostStep, +}; use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; @@ -21,23 +23,6 @@ use bindings::DeploymentHooks; pub(crate) use bindings::PythonLogger; use handle::{Execution, ExecutionBody, ExecutionStep}; -pub(crate) enum NativeCallStep { - Host(O), - Complete, -} - -type NativeCallFuture<'a, O> = - Pin, litellm_core::Error>> + Send + 'a>>; - -pub(crate) trait NativeCall: Send + Sync { - type Operation: Send + 'static; - type Result: Send + 'static; - - fn resume(&mut self, result: Option) -> NativeCallFuture<'_, Self::Operation>; - - fn interrupt(&mut self, failure: HostFailure) -> NativeCallFuture<'_, Self::Operation>; -} - pub(crate) enum OperationClass { Phase(HostPhase), Route, @@ -60,12 +45,17 @@ pub(crate) trait PythonRoute: Send + Sync { fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; } -type HostResumeStep = - HostStep::Call as NativeCall>::Operation>, Py>; +type HostResumeStep = HostStep< + NativeCallStep< + <::Call as NativeCall>::Operation, + <::Call as NativeCall>::Complete, + >, + Py, +>; struct NativeCallState { call: C, - result: Option, litellm_core::Error>>, + result: Option, litellm_core::Error>>, } enum PendingOperation { @@ -151,7 +141,11 @@ impl PythonLifecycle { } } - fn take_native_result(&self) -> PyResult::Operation>> { + fn take_native_result( + &self, + ) -> PyResult< + NativeCallStep<::Operation, ::Complete>, + > { self.call .as_ref() .ok_or_else(missing_state)? @@ -214,7 +208,7 @@ impl PythonLifecycle { loop { let operation = match step { HostStep::Suspend(awaitable) => return Ok(ExecutionStep::Await(awaitable)), - HostStep::Ready(NativeCallStep::Complete) => { + HostStep::Ready(NativeCallStep::Complete(_)) => { return self .route .state_mut() @@ -683,24 +677,19 @@ mod tests { impl NativeCall for SyntheticCall { type Operation = (); type Result = (); + type Complete = (); fn resume( &mut self, result: Option, - ) -> Pin< - Box< - dyn Future, litellm_core::Error>> - + Send - + '_, - >, - > { + ) -> HostCallFuture<'_, Self::Operation, Self::Complete> { Box::pin(async move { match (self.0, result) { (false, None) => { self.0 = true; Ok(NativeCallStep::Host(())) } - (true, Some(())) => Ok(NativeCallStep::Complete), + (true, Some(())) => Ok(NativeCallStep::Complete(())), _ => Err(litellm_core::Error::InvalidRequest( "invalid synthetic lifecycle state".into(), )), @@ -711,14 +700,8 @@ mod tests { fn interrupt( &mut self, _: HostFailure, - ) -> Pin< - Box< - dyn Future, litellm_core::Error>> - + Send - + '_, - >, - > { - Box::pin(async { Ok(NativeCallStep::Complete) }) + ) -> HostCallFuture<'_, Self::Operation, Self::Complete> { + Box::pin(async { Ok(NativeCallStep::Complete(())) }) } } diff --git a/litellm-rust/crates/python-bridge/src/routes/gateway_messages.rs b/litellm-rust/crates/python-bridge/src/routes/gateway_messages.rs deleted file mode 100644 index 97ff93f299a..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/gateway_messages.rs +++ /dev/null @@ -1,29 +0,0 @@ -use pyo3::prelude::*; -use serde_json::Value; - -use crate::errors::core_error_to_pyerr; - -#[pyfunction] -fn gateway_messages<'py>( - py: Python<'py>, - model_alias: String, - provider_model: String, - api_base: String, - #[pyo3(from_py_with = litellm_python_interop::from_py)] body: Value, -) -> PyResult> { - let future = litellm_ai_gateway::trace_parity::messages_request( - model_alias, - provider_model, - api_base, - body, - ); - crate::execution::run_async( - py, - crate::function_trace::capture(future), - core_error_to_pyerr, - ) -} - -pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> { - super::definition::add_function(module, wrap_pyfunction!(gateway_messages, module)?) -} diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index d1c26e1d9cd..97c39a5d6b3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -3,9 +3,6 @@ use pyo3::prelude::*; #[macro_use] mod definition; -#[cfg(feature = "trace-parity")] -mod gateway_messages; - mod audio_transcription; mod chat_completions; mod messages; @@ -24,7 +21,6 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { audio_transcription::register_trace(&trace)?; messages::register_trace(&trace)?; chat_completions::register_trace(&trace)?; - gateway_messages::register_trace(&trace)?; module.add_submodule(&trace)?; } Ok(()) diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs index a120f6e196c..1cbe8a179e3 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -31,19 +31,21 @@ impl PythonLogger { py: Python<'_>, kwargs: &Py, pre_call: &OcrLoggingFields, + secret_fields: &[&str], url: &str, ) -> PyResult<()> { - let redact = py - .import("litellm.rust_bridge.ocr")? - .getattr("redact_logging_params")?; let update = PyDict::new(py); - update.set_item("kwargs", redact.call1((kwargs,))?.cast_into::()?)?; + update.set_item("kwargs", redact(py, kwargs.bind(py), secret_fields)?)?; update.set_item("model", &pre_call.model)?; update.set_item( "optional_params", - redact - .call1((to_py(py, &pre_call.optional_params)?,))? - .cast_into::()?, + redact( + py, + &to_py(py, &pre_call.optional_params)? + .into_bound(py) + .cast_into::()?, + secret_fields, + )?, )?; let params = PyDict::new(py); params.set_item( @@ -93,8 +95,8 @@ impl PythonLogger { &self, py: Python<'_>, original_response: &Value, - body: &Option>, - headers: &Option>, + body: Option<&Py>, + headers: Option<&Py>, ) -> PyResult<()> { let additional = PyDict::new(py); additional.set_item("complete_input_dict", body)?; @@ -118,6 +120,26 @@ impl PythonLogger { } } +fn redact( + py: Python<'_>, + params: &Bound<'_, PyDict>, + secret_fields: &[&str], +) -> PyResult> { + let redacted = PyDict::new(py); + for (name, value) in params { + let name = name.extract::()?; + if name == "proxy_server_request" { + continue; + } + if secret_fields.contains(&name.as_str()) { + redacted.set_item(name, "****")?; + } else { + redacted.set_item(name, value)?; + } + } + Ok(redacted.unbind()) +} + pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult> { py.import("litellm.rust_bridge.ocr")? .getattr("_response")? diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs index 11e9c78251f..d43c2f88775 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs @@ -4,7 +4,9 @@ use std::path::PathBuf; use pyo3::exceptions::{PyFileNotFoundError, PyTypeError, PyValueError}; use pyo3::prelude::*; use pyo3::pybacked::PyBackedBytes; -use pyo3::types::{PyBytes, PyDict, PyString}; +#[cfg(test)] +use pyo3::types::PyDict; +use pyo3::types::{PyBytes, PyString}; use litellm_core::constants::OCR_INLINE_MAX_BYTES; use litellm_core::ocr::{OcrDocument, encode_file_document, mime_type_for_name, upload_mime_type}; @@ -97,6 +99,11 @@ impl FromPyObject<'_, '_> for FileDocumentInput { fn extract(document: Borrowed<'_, '_, PyAny>) -> PyResult { let py = document.py(); + let mime_type = match document.get_item("mime_type") { + Ok(value) => Some(value.extract::()?), + Err(error) if error.is_instance_of::(py) => None, + Err(error) => return Err(error), + }; let file = document.get_item("file").map_err(|error| { if error.is_instance_of::(py) { PyValueError::new_err("document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes") @@ -110,11 +117,6 @@ impl FromPyObject<'_, '_> for FileDocumentInput { )); } let (bytes, name) = read_file_input(py, &file)?; - let mime_type = document - .cast::()? - .get_item("mime_type")? - .map(|value| value.extract::()) - .transpose()?; Ok(Self { bytes, name, @@ -203,32 +205,35 @@ mod tests { } #[test] - fn extraction_reads_mime_type_after_consuming_file_once() { + fn extraction_validates_mime_type_before_consuming_file() { Python::initialize(); Python::attach(|py| { let locals = PyDict::new(py); py.run( c"class Reader: + def __init__(self): + self.reads = 0 def read(self): - assert document['mime_type'] == 7 - document['mime_type'] = 'image/png' + self.reads += 1 return b'abc' -document = {'file': Reader(), 'mime_type': 7}", +reader = Reader() +document = {'file': reader, 'mime_type': 7}", Some(&locals), Some(&locals), ) .unwrap(); let document = locals.get_item("document").unwrap().unwrap(); - let input: FileDocumentInput = document.extract().unwrap(); - assert_eq!(input.bytes.as_ref(), b"abc"); - assert_eq!(input.mime_type.as_deref(), Some("image/png")); - let result = file_document(py, input).unwrap(); - assert_eq!( - serde_json::to_value(result).unwrap(), - serde_json::json!({ - "type": "image_url", "image_url": "data:image/png;base64,YWJj" - }) - ); + let error = document.extract::().err().unwrap(); + assert!(error.is_instance_of::(py)); + let reads: usize = locals + .get_item("reader") + .unwrap() + .unwrap() + .getattr("reads") + .unwrap() + .extract() + .unwrap(); + assert_eq!(reads, 0); }); } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs index f5fb7bead7d..66bdfb7583e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs @@ -1,54 +1,49 @@ use litellm_core::error::Error; -use pyo3::exceptions::PyValueError; use pyo3::prelude::*; use crate::errors::{RustUpstreamError, core_error_to_pyerr}; pub(super) fn to_pyerr(error: Error) -> PyErr { - match error { - Error::MissingField("document_url" | "image_url") => { - PyValueError::new_err("Document URL is required") - } - Error::Http { status, body } => upstream_error(status, body), - Error::Network(message) if message.contains("timed out") => upstream_error(408, message), - other => { - let status = other.http_status_code(); - let error = core_error_to_pyerr(other); - if let Some(status) = status { - Python::attach(|py| { - let value = error.value(py); - value.setattr("status_code", status).ok(); - value.setattr("message", value.to_string()).ok(); - }); - } - error - } - } + let status = error.http_status_code(); + let mapped = match error { + Error::Http { status, body } => RustUpstreamError::new_err((status, body)), + other => core_error_to_pyerr(other), + }; + attach_status(mapped, status) } -fn upstream_error(status: u16, message: String) -> PyErr { - let error = RustUpstreamError::new_err((status, message.clone())); - Python::attach(|py| { - let value = error.value(py); - value.setattr("status_code", status).ok(); - value.setattr("message", message).ok(); - }); +fn attach_status(error: PyErr, status: Option) -> PyErr { + if let Some(status) = status { + Python::attach(|py| { + let value = error.value(py); + value.setattr("status_code", status).ok(); + value.setattr("message", value.to_string()).ok(); + }); + } error } #[cfg(test)] mod tests { use super::*; + use pyo3::exceptions::PyValueError; #[test] fn preserves_python_validation_and_provider_details() { Python::initialize(); Python::attach(|py| { - for field in ["document_url", "image_url"] { - let mapped = to_pyerr(Error::MissingField(field)); - assert!(mapped.is_instance_of::(py)); - assert_eq!(mapped.value(py).to_string(), "Document URL is required"); - } + let mapped = to_pyerr(Error::MissingDocumentUrl); + assert!(mapped.is_instance_of::(py)); + assert_eq!(mapped.value(py).to_string(), "Document URL is required"); + assert_eq!( + mapped + .value(py) + .getattr("status_code") + .unwrap() + .extract::() + .unwrap(), + 500 + ); let mapped = to_pyerr(Error::Http { status: 429, body: r#"{"message":"rate limited"}"#.to_string(), diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs index 79029ab777b..16b8ba52e71 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs @@ -1,51 +1,54 @@ -use serde_json::{Map, Value}; -use std::sync::Arc; - use pyo3::prelude::*; use pyo3::types::{PyDict, PyTuple}; use litellm_core::auth::ResolvedCredential; use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest, OcrPreCallRequest}; -use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_param_names, decode_request}; -use litellm_core::ocr::{ - NativeOutcome, OcrAdmission, OcrCall, OcrClient, OcrHostOperation, OcrHostResult, -}; +use litellm_core::ocr::{OcrAdmission, OcrCall, OcrClient, OcrHostOperation, OcrHostResult}; use litellm_python_interop::{ from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, }; use super::callbacks; use super::errors::to_pyerr as ocr_error_to_pyerr; -use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider}; -use crate::errors::RustBridgeDeclined; +use super::project::{ProjectedOcrFields, admitted_call, project_request}; use crate::lifecycle::{ - NativeCall, NativeCallStep, OperationClass, PythonCallState, PythonRoute, missing_state, now, - run_call, + OperationClass, PythonCallState, PythonRoute, missing_state, now, run_call, }; -use crate::marshal::{project_optional_fields, python_timeout_seconds, request_input_sources}; struct PythonOcrHost { state: PythonCallState, - request: Option>, + data: OcrHostData, +} + +enum OcrHostData { + Unprojected { request: Py }, + Projected(ProjectedOcrHost), + Released, +} + +struct ProjectedOcrHost { + fields: ProjectedOcrFields, pre_call: Option, - document: Option>, - api_key: Option>, - azure_ad_token_provider: Option, - provider: String, retained_fields: Option>, body: Option>, headers: Option>, } -struct AdmittedOcrCall { - request: litellm_core::ocr::LiteLLMOcrRequest, - document: Py, - api_key: Py, - azure_ad_token_provider: Option, - provider: String, -} - impl PythonOcrHost { + fn projected(&self) -> PyResult<&ProjectedOcrHost> { + match &self.data { + OcrHostData::Projected(projected) => Ok(projected), + _ => Err(missing_state()), + } + } + + fn projected_mut(&mut self) -> PyResult<&mut ProjectedOcrHost> { + match &mut self.data { + OcrHostData::Projected(projected) => Ok(projected), + _ => Err(missing_state()), + } + } + fn pre_call( &mut self, py: Python<'_>, @@ -63,17 +66,17 @@ impl PythonOcrHost { retained_fields.set_item(name, value)?; } } - retained_fields.set_item( - "document", - self.document.as_ref().ok_or_else(missing_state)?, - )?; - self.retained_fields = Some(retained_fields.unbind()); - self.pre_call = Some((&request).into()); + retained_fields.set_item("document", &self.projected()?.fields.document)?; + let projected = self.projected_mut()?; + projected.retained_fields = Some(retained_fields.unbind()); + projected.pre_call = Some((&request).into()); Ok(request) } fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult { let provider = self + .projected()? + .fields .azure_ad_token_provider .as_ref() .ok_or_else(missing_state)?; @@ -85,11 +88,18 @@ impl PythonOcrHost { py: Python<'_>, mut request: OcrDuringCallRequest, ) -> PyResult { - let pre_call = self.pre_call.as_ref().ok_or_else(missing_state)?; - let logger = self.state.logger()?; - logger.update_ocr(py, &self.state.kwargs, pre_call, &request.url)?; - if !logger.callbacks_needed(py, "payload")? { - logger + let projected = self.projected()?; + let pre_call = projected.pre_call.as_ref().ok_or_else(missing_state)?; + self.state.logger()?.update_ocr( + py, + &self.state.kwargs, + pre_call, + &projected.fields.secret_fields, + &request.url, + )?; + if !self.state.logger()?.callbacks_needed(py, "payload")? { + self.state + .logger()? .object(py) .call_method0("record_api_call_start_time")?; return Ok(request); @@ -102,7 +112,7 @@ impl PythonOcrHost { let body = to_py(py, &request.body)? .into_bound(py) .cast_into::()?; - if let Some(retained) = &self.retained_fields { + if let Some(retained) = &self.projected()?.retained_fields { for name in &request.retained_fields { if let Some(value) = retained.bind(py).get_item(name)? { body.set_item(name, value)?; @@ -113,9 +123,13 @@ impl PythonOcrHost { for (name, value) in &request.headers { headers.set_item(name, value)?; } - self.body = Some(body.clone().unbind()); - self.headers = Some(headers.clone().unbind()); - logger.pre_ocr(py, &self.api_key, &body, &headers, &request.url)?; + let api_key = self.projected()?.fields.api_key.clone_ref(py); + let projected = self.projected_mut()?; + projected.body = Some(body.clone().unbind()); + projected.headers = Some(headers.clone().unbind()); + self.state + .logger()? + .pre_ocr(py, &Some(api_key), &body, &headers, &request.url)?; let headers = headers .iter() .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) @@ -132,59 +146,18 @@ impl PythonOcrHost { ) -> PyResult { let logger = self.state.logger()?; if logger.callbacks_needed(py, "payload")? { - logger.post_ocr(py, &request.original_response, &self.body, &self.headers)?; + let projected = self.projected()?; + logger.post_ocr( + py, + &request.original_response, + projected.body.as_ref(), + projected.headers.as_ref(), + )?; } Ok(request) } } -impl NativeCall for OcrCall { - type Operation = OcrHostOperation; - type Result = OcrHostResult; - - fn resume( - &mut self, - result: Option, - ) -> std::pin::Pin< - Box< - dyn std::future::Future< - Output = Result, litellm_core::Error>, - > + Send - + '_, - >, - > { - Box::pin(async move { - OcrCall::resume(self, result).await.map(|step| match step { - litellm_core::ocr::OcrCallStep::Host(operation) => NativeCallStep::Host(operation), - litellm_core::ocr::OcrCallStep::Complete(_) => NativeCallStep::Complete, - }) - }) - } - - fn interrupt( - &mut self, - failure: litellm_core::call_lifecycle::host::HostFailure, - ) -> std::pin::Pin< - Box< - dyn std::future::Future< - Output = Result, litellm_core::Error>, - > + Send - + '_, - >, - > { - Box::pin(async move { - OcrCall::interrupt(self, failure) - .await - .map(|step| match step { - litellm_core::ocr::OcrCallStep::Host(operation) => { - NativeCallStep::Host(operation) - } - litellm_core::ocr::OcrCallStep::Complete(_) => NativeCallStep::Complete, - }) - }) - } -} - impl PythonRoute for PythonOcrHost { type Call = OcrCall; @@ -197,16 +170,9 @@ impl PythonRoute for PythonOcrHost { } fn classify(operation: &OcrHostOperation) -> OperationClass { - match operation { - OcrHostOperation::Lifecycle(phase) => OperationClass::Phase(*phase), - OcrHostOperation::Success { .. } => { - OperationClass::Phase(litellm_core::call_lifecycle::host::HostPhase::Success) - } - OcrHostOperation::Failure { .. } => { - OperationClass::Phase(litellm_core::call_lifecycle::host::HostPhase::Failure) - } - _ => OperationClass::Route, - } + operation + .phase() + .map_or(OperationClass::Route, OperationClass::Phase) } fn lifecycle_result() -> OcrHostResult { @@ -220,19 +186,20 @@ impl PythonRoute for PythonOcrHost { fn invoke(&mut self, py: Python<'_>, operation: OcrHostOperation) -> PyResult { Ok(match operation { OcrHostOperation::ProjectRequest => { - let projected = project_request( - py, - self.request.as_ref().ok_or_else(missing_state)?.bind(py), - self.state.kwargs.bind(py), - )?; - self.document = Some(projected.document); - self.api_key = Some(projected.api_key); - self.azure_ad_token_provider = projected.azure_ad_token_provider; - self.provider = projected.provider; - OcrHostResult::Request(Ok(( - Box::new(projected.request), - self.azure_ad_token_provider.is_some(), - ))) + let OcrHostData::Unprojected { request } = &self.data else { + return Err(missing_state()); + }; + let projected = project_request(py, request.bind(py), self.state.kwargs.bind(py))?; + let has_token_provider = projected.fields.azure_ad_token_provider.is_some(); + let request = projected.request; + self.data = OcrHostData::Projected(ProjectedOcrHost { + fields: projected.fields, + pre_call: None, + retained_fields: None, + body: None, + headers: None, + }); + OcrHostResult::Request(Ok((Box::new(request), has_token_provider))) } OcrHostOperation::AcquireAzureAdToken => { OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?)) @@ -259,8 +226,15 @@ impl PythonRoute for PythonOcrHost { self.state.end = Some(now(py)?); } let error = self.state.error.as_ref().ok_or_else(missing_state)?; - let request = self.request.as_ref().ok_or_else(missing_state)?.bind(py); - let mapped = callbacks::map_failure(py, error, request, &self.provider)?; + let (request, provider) = match &self.data { + OcrHostData::Unprojected { request } => (request.bind(py), ""), + OcrHostData::Projected(projected) => ( + projected.fields.boundary_request.bind(py), + projected.fields.provider, + ), + OcrHostData::Released => return Err(missing_state()), + }; + let mapped = callbacks::map_failure(py, error, request, provider)?; self.state .retain_error(py, PyErr::from_value(mapped.into_bound(py).into_any())); OcrHostResult::Lifecycle(Ok(())) @@ -272,201 +246,31 @@ impl PythonRoute for PythonOcrHost { } fn cleanup(&mut self) { - self.request = None; - self.pre_call = None; - self.document = None; - self.api_key = None; - self.azure_ad_token_provider = None; - self.retained_fields = None; - self.body = None; - self.headers = None; + self.data = OcrHostData::Released; } fn traverse(&self, visit: &pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> { - visit.call(&self.request)?; - visit.call(&self.document)?; - visit.call(&self.api_key)?; - if let Some(provider) = &self.azure_ad_token_provider { - provider.traverse(visit)?; - } - visit.call(&self.retained_fields)?; - visit.call(&self.body)?; - visit.call(&self.headers) - } -} - -struct OcrArguments<'a, 'py> { - request: &'a Bound<'py, PyAny>, - kwargs: &'a Bound<'py, PyDict>, -} - -impl<'py> OcrArguments<'_, 'py> { - fn lookup(&self, name: &str) -> PyResult> { - match self.kwargs.get_item(name)? { - Some(value) => Ok(value), - None => self.request.getattr(name), - } - } - - fn model(&self) -> PyResult { - self.lookup("model")?.extract() - } - - fn custom_llm_provider(&self) -> PyResult> { - self.lookup("custom_llm_provider")?.extract() - } - - fn document(&self) -> PyResult> { - Ok(CapturedDocument(self.lookup("document")?)) - } - - fn api_key(&self) -> PyResult> { - Ok(CapturedApiKey(self.lookup("api_key")?)) - } - - fn api_base(&self) -> PyResult> { - self.lookup("api_base")?.extract() - } - - fn extra_headers(&self) -> PyResult>> { - self.lookup("extra_headers")? - .extract::>>()? - .map(|value| from_py(value.bind(self.request.py()))) - .transpose() - } - - fn timeout_seconds(&self) -> PyResult> { - Ok(self - .lookup("timeout")? - .extract::>>()? - .map(|value| python_timeout_seconds(self.request.py(), value)) - .transpose()? - .flatten()) - } -} - -struct CapturedDocument<'py>(Bound<'py, PyAny>); - -impl<'py> CapturedDocument<'py> { - fn as_bound(&self) -> &Bound<'py, PyAny> { - &self.0 - } -} - -struct CapturedApiKey<'py>(Bound<'py, PyAny>); - -impl CapturedApiKey<'_> { - fn value(&self) -> PyResult> { - self.0.extract() - } - - fn into_object(self) -> Py { - self.0.unbind() - } -} - -enum DocumentKind { - File, - Other, -} - -impl FromPyObject<'_, '_> for DocumentKind { - type Error = PyErr; - - fn extract(document: Borrowed<'_, '_, PyAny>) -> PyResult { - let kind: String = document.get_item("type")?.extract()?; - - Ok(match kind.as_str() { - "file" => Self::File, - _ => Self::Other, - }) - } -} - -fn project_request( - py: Python<'_>, - request: &Bound<'_, PyAny>, - kwargs: &Bound<'_, PyDict>, -) -> PyResult { - let arguments = OcrArguments { request, kwargs }; - let model = arguments.model()?; - let custom_llm_provider = arguments.custom_llm_provider()?; - let document = arguments.document()?; - let wire_document = extract_document(py, document.as_bound())?; - let retained_document = retained_document(py, document.as_bound(), &wire_document)?; - let api_key = arguments.api_key()?; - let request_kwargs = kwargs; - let consumed = consumed_optional_param_names(&model, custom_llm_provider.as_deref()) - .map_err(ocr_error_to_pyerr)?; - let optional_params = project_optional_fields(request_kwargs, &consumed)?; - let input_sources = request_input_sources( - request_kwargs, - consumed - .iter() - .copied() - .chain(["api_key", "api_base", "extra_headers"]), - )?; - let azure_ad_token_provider = request_kwargs - .get_item("azure_ad_token_provider")? - .and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER)); - let wire = OcrWireRequest { - model, - document: wire_document, - api_key: api_key.value()?, - api_base: arguments.api_base()?, - custom_llm_provider, - extra_headers: arguments.extra_headers()?, - optional_params, - input_sources, - timeout_seconds: arguments.timeout_seconds()?, - }; - let request = decode_request(wire).map_err(ocr_error_to_pyerr)?; - let provider = request.provider_name().to_string(); - let request = request.with_host_hooks(Arc::new(BridgeOcrHooks), None); - Ok(AdmittedOcrCall { - request, - document: retained_document, - api_key: api_key.into_object(), - azure_ad_token_provider, - provider, - }) -} - -fn extract_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult { - match document.extract::()? { - DocumentKind::Other => from_py(document), - DocumentKind::File => { - let input = document.extract()?; - let encoded = super::document::file_document(py, input)?; - serde_json::to_value(encoded) - .map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string())) + match &self.data { + OcrHostData::Unprojected { request } => visit.call(request), + OcrHostData::Projected(projected) => { + visit.call(&projected.fields.boundary_request)?; + visit.call(&projected.fields.document)?; + visit.call(&projected.fields.api_key)?; + if let Some(provider) = &projected.fields.azure_ad_token_provider { + provider.traverse(visit)?; + } + visit.call(&projected.retained_fields)?; + visit.call(&projected.body)?; + visit.call(&projected.headers) + } + OcrHostData::Released => Ok(()), } } } -fn retained_document( - py: Python<'_>, - document: &Bound<'_, PyAny>, - wire_document: &Value, -) -> PyResult> { - match document.extract::()? { - DocumentKind::File => to_py(py, wire_document), - DocumentKind::Other => Ok(document.clone().unbind()), - } -} - -fn admitted_call(outcome: NativeOutcome) -> PyResult { - match outcome { - NativeOutcome::Completed(call) => Ok(call), - NativeOutcome::Declined(reason) => Err(RustBridgeDeclined::new_err(format!( - "native OCR admission declined: {reason:?}" - ))), - } -} - -struct BridgeOcrHooks; +pub(super) struct BridgeOcrHooks; impl litellm_core::ocr::hooks::OcrHooks for BridgeOcrHooks { - fn has_guardrails(&self) -> bool { + fn intercepts_requests(&self) -> bool { true } } @@ -479,13 +283,6 @@ fn _ocr_lifecycle( kwargs: Bound<'_, PyDict>, asynchronous: bool, ) -> PyResult> { - if let Ok(gil_enabled) = py.import("sys")?.getattr("_is_gil_enabled") - && !gil_enabled.call0()?.is_truthy()? - { - return Err(pyo3::exceptions::PyRuntimeError::new_err( - "native OCR requires the Python GIL", - )); - } let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?; let call = admitted_call(OcrCall::admit( client, @@ -502,15 +299,9 @@ fn _ocr_lifecycle( asynchronous, if asynchronous { "aocr" } else { "ocr" }, )?, - request: Some(request.unbind()), - pre_call: None, - document: None, - api_key: None, - azure_ad_token_provider: None, - provider: String::new(), - retained_fields: None, - body: None, - headers: None, + data: OcrHostData::Unprojected { + request: request.unbind(), + }, }; run_call(py, call, host) } @@ -518,413 +309,3 @@ fn _ocr_lifecycle( pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { module.add_function(wrap_pyfunction!(_ocr_lifecycle, module)?) } - -#[cfg(test)] -mod tests { - use litellm_core::Error; - use litellm_core::ocr::OcrDecline; - use pyo3::exceptions::{PyKeyError, PyTypeError, PyValueError}; - - use super::*; - - fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> { - let locals = PyDict::new(py); - py.run(source, Some(&locals), Some(&locals)).unwrap(); - locals - } - - fn arguments<'a, 'py>( - request: &'a Bound<'py, PyAny>, - kwargs: &'a Bound<'py, PyDict>, - ) -> OcrArguments<'a, 'py> { - OcrArguments { request, kwargs } - } - - fn stub_timeout_conversion(py: Python<'_>) { - eval( - py, - c" -import sys -import types -timeouts = types.ModuleType('litellm.rust_bridge.timeouts') -timeouts.timeout_to_seconds = lambda timeout: None if timeout is None else float(timeout) -sys.modules.setdefault('litellm', types.ModuleType('litellm')) -sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bridge')) -sys.modules['litellm.rust_bridge.timeouts'] = timeouts -", - ); - } - - #[test] - fn typed_initial_decline_uses_bridge_decline_contract() { - Python::initialize(); - Python::attach(|py| { - let Err(error) = admitted_call(NativeOutcome::Declined(OcrDecline::HostOperations)) - else { - panic!("unsupported host operations should decline admission"); - }; - assert!(error.is_instance_of::(py)); - }); - } - - #[test] - fn post_admission_error_does_not_use_bridge_decline_contract() { - Python::initialize(); - Python::attach(|py| { - let error = ocr_error_to_pyerr(Error::InvalidRequest("callback result".into())); - assert!(error.is_instance_of::(py)); - assert!(!error.is_instance_of::(py)); - }); - } - - #[test] - fn kwargs_override_request_attributes_including_explicit_none() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -class Request: - def __init__(self): - self.accesses = [] - def __getattribute__(self, name): - if name != 'accesses': - object.__getattribute__(self, 'accesses').append(name) - return object.__getattribute__(self, name) -request = Request() -request.model = 'from-request' -request.custom_llm_provider = 'mistral' -kwargs = {'model': 'from-kwargs', 'custom_llm_provider': None} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let arguments = arguments(&request, &kwargs); - assert_eq!(arguments.model().unwrap(), "from-kwargs"); - assert_eq!(arguments.custom_llm_provider().unwrap(), None); - let accesses: Vec = request.getattr("accesses").unwrap().extract().unwrap(); - assert_eq!(accesses, Vec::::new()); - }); - } - - #[test] - fn missing_kwargs_read_the_request_property_once() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -class Request: - def __init__(self): - self.reads = 0 - @property - def model(self): - self.reads += 1 - return 'mistral-ocr-latest' -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - assert_eq!( - arguments(&request, &kwargs).model().unwrap(), - "mistral-ocr-latest" - ); - assert_eq!( - request.getattr("reads").unwrap().extract::().unwrap(), - 1 - ); - }); - } - - #[test] - fn request_property_exceptions_keep_their_identity() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -failure = LookupError('model failed') -class Request: - @property - def model(self): - raise failure -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let error = arguments(&request, &kwargs).model().unwrap_err(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - }); - } - - #[test] - fn unused_raising_property_is_never_inspected() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -class Request: - @property - def unused(self): - raise RuntimeError('unused') - model = 'mistral-ocr-latest' - custom_llm_provider = None -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let arguments = arguments(&request, &kwargs); - assert_eq!(arguments.model().unwrap(), "mistral-ocr-latest"); - assert_eq!(arguments.custom_llm_provider().unwrap(), None); - }); - } - - #[test] - fn document_reader_mutations_are_visible_to_later_field_reads() { - Python::initialize(); - Python::attach(|py| { - stub_timeout_conversion(py); - let locals = eval( - py, - c" -class Request: - api_base = 'original' - timeout = 1 - @property - def document(self): - return document -class Reader: - def read(self): - Request.api_base = 'mutated' - Request.timeout = 9 - return b'abc' -document = {'type': 'file', 'file': Reader()} -request = Request() -kwargs = {} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let arguments = arguments(&request, &kwargs); - let document = arguments.document().unwrap(); - extract_document(py, document.as_bound()).unwrap(); - assert_eq!(arguments.api_base().unwrap().as_deref(), Some("mutated")); - assert_eq!(arguments.timeout_seconds().unwrap(), Some(9.0)); - }); - } - - #[test] - fn captured_api_key_keeps_the_original_python_object() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -key = object() -class Request: - api_key = None -request = Request() -kwargs = {'api_key': key} -", - ); - let request = locals.get_item("request").unwrap().unwrap(); - let kwargs = locals - .get_item("kwargs") - .unwrap() - .unwrap() - .cast_into::() - .unwrap(); - let captured = arguments(&request, &kwargs).api_key().unwrap(); - assert!( - captured - .into_object() - .bind(py) - .is(locals.get_item("key").unwrap().unwrap()) - ); - }); - } - - #[test] - fn file_documents_are_encoded_and_other_documents_keep_the_python_object() { - Python::initialize(); - Python::attach(|py| { - let file = py - .eval( - c"{'type': 'file', 'file': b'%PDF-1.4', 'mime_type': 'application/pdf'}", - None, - None, - ) - .unwrap(); - assert_eq!( - extract_document(py, &file).unwrap(), - serde_json::json!({ - "type": "document_url", - "document_url": "data:application/pdf;base64,JVBERi0xLjQ=", - }) - ); - - let original = py - .eval( - c"{'type': 'document_url', 'document_url': 'https://example.com/a.pdf'}", - None, - None, - ) - .unwrap(); - let wire = extract_document(py, &original).unwrap(); - assert_eq!( - wire, - serde_json::json!({ - "type": "document_url", - "document_url": "https://example.com/a.pdf", - }) - ); - assert!( - retained_document(py, &original, &wire) - .unwrap() - .bind(py) - .is(&original) - ); - }); - } - - #[test] - fn unknown_document_types_reach_existing_downstream_validation() { - Python::initialize(); - Python::attach(|py| { - let document = py - .eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None) - .unwrap(); - let wire_document = extract_document(py, &document).unwrap(); - assert_eq!( - wire_document, - serde_json::json!({"type": "mystery", "mystery": "x"}) - ); - let error = match decode_request(OcrWireRequest { - model: "mistral/mistral-ocr-latest".into(), - document: wire_document, - api_key: None, - api_base: None, - custom_llm_provider: None, - extra_headers: None, - optional_params: Map::new(), - input_sources: Default::default(), - timeout_seconds: None, - }) { - Ok(_) => panic!("unknown discriminators belong to core validation"), - Err(error) => error, - }; - assert!(error.to_string().contains("document")); - }); - } - - #[test] - fn document_discriminator_errors_keep_their_existing_exceptions() { - Python::initialize(); - Python::attach(|py| { - let missing = py.eval(c"{}", None, None).unwrap(); - assert!( - extract_document(py, &missing) - .unwrap_err() - .is_instance_of::(py) - ); - - let non_string = py.eval(c"{'type': 1}", None, None).unwrap(); - assert!( - extract_document(py, &non_string) - .unwrap_err() - .is_instance_of::(py) - ); - - let locals = eval( - py, - c" -failure = RuntimeError('type lookup failed') -class Document: - def __getitem__(self, key): - raise failure -document = Document() -", - ); - let error = - extract_document(py, &locals.get_item("document").unwrap().unwrap()).unwrap_err(); - assert!( - error - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - }); - } - - #[test] - fn document_kind_reads_only_type_and_classification_happens_twice() { - Python::initialize(); - Python::attach(|py| { - let locals = eval( - py, - c" -class Document(dict): - def __init__(self): - super().__init__({'file': b'abc'}) - self.reads = [] - def __getitem__(self, key): - self.reads.append(key) - if key == 'type': - return 'file' if self.reads.count('type') == 1 else 'document_url' - return super().__getitem__(key) -document = Document() -", - ); - let document = locals.get_item("document").unwrap().unwrap(); - assert!(matches!( - document.extract::().unwrap(), - DocumentKind::File - )); - let reads: Vec = document.getattr("reads").unwrap().extract().unwrap(); - assert_eq!(reads, ["type"]); - - py.run(c"document.reads = []", Some(&locals), Some(&locals)) - .unwrap(); - let wire = extract_document(py, &document).unwrap(); - let retained = retained_document(py, &document, &wire).unwrap(); - assert!(retained.bind(py).is(&document)); - let reads: Vec = document.getattr("reads").unwrap().extract().unwrap(); - assert_eq!(reads, ["type", "file", "type"]); - }); - } -} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index 8c6469a4bf7..10fa40b65ea 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -2,6 +2,7 @@ mod callbacks; mod document; mod errors; mod lifecycle; +mod project; mod value; use pyo3::prelude::*; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs new file mode 100644 index 00000000000..8b6a1b02e19 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -0,0 +1,579 @@ +use std::sync::Arc; + +use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_params, decode_request}; +use litellm_core::ocr::{LiteLLMOcrRequest, NativeOutcome, OcrCall}; +use litellm_python_interop::{ + from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, +}; +use pyo3::prelude::*; +use pyo3::types::PyDict; +use serde_json::{Map, Value}; + +use super::errors::to_pyerr as ocr_error_to_pyerr; +use super::lifecycle::BridgeOcrHooks; +use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider}; +use crate::errors::RustBridgeDeclined; +use crate::marshal::{project_optional_fields, python_timeout_seconds, request_input_sources}; + +pub(super) struct ProjectedOcrFields { + pub boundary_request: Py, + pub document: Py, + pub api_key: Py, + pub azure_ad_token_provider: Option, + pub provider: &'static str, + pub secret_fields: Vec<&'static str>, +} + +pub(super) struct ProjectedOcrCall { + pub request: LiteLLMOcrRequest, + pub fields: ProjectedOcrFields, +} + +struct OcrArguments<'a, 'py> { + request: &'a Bound<'py, PyAny>, + kwargs: &'a Bound<'py, PyDict>, +} + +impl<'py> OcrArguments<'_, 'py> { + fn lookup(&self, name: &str) -> PyResult> { + match self.kwargs.get_item(name)? { + Some(value) => Ok(value), + None => self.request.getattr(name), + } + } + + fn model(&self) -> PyResult { + self.lookup("model")?.extract() + } + + fn custom_llm_provider(&self) -> PyResult> { + self.lookup("custom_llm_provider")?.extract() + } + + fn document(&self) -> PyResult> { + self.lookup("document") + } + + fn api_key(&self) -> PyResult> { + self.lookup("api_key") + } + + fn api_base(&self) -> PyResult> { + self.lookup("api_base")?.extract() + } + + fn extra_headers(&self) -> PyResult>> { + self.lookup("extra_headers")? + .extract::>>()? + .map(|value| from_py(value.bind(self.request.py()))) + .transpose() + } + + fn timeout_seconds(&self) -> PyResult> { + Ok(self + .lookup("timeout")? + .extract::>>()? + .map(|value| python_timeout_seconds(self.request.py(), value)) + .transpose()? + .flatten()) + } +} + +enum ProjectedDocument { + File { wire: Value, retained: Py }, + Other { wire: Value, retained: Py }, +} + +impl ProjectedDocument { + fn project(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult { + let kind: String = document.get_item("type")?.extract()?; + if kind != "file" { + return Ok(Self::Other { + wire: from_py(document)?, + retained: document.clone().unbind(), + }); + } + let input = document.extract()?; + let encoded = super::document::file_document(py, input)?; + let wire = serde_json::to_value(encoded) + .map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?; + Ok(Self::File { + retained: to_py(py, &wire)?, + wire, + }) + } + + fn into_parts(self) -> (Value, Py) { + match self { + Self::File { wire, retained } | Self::Other { wire, retained } => (wire, retained), + } + } +} + +pub(super) fn project_request( + py: Python<'_>, + request: &Bound<'_, PyAny>, + kwargs: &Bound<'_, PyDict>, +) -> PyResult { + let boundary_request = request.clone().unbind(); + let arguments = OcrArguments { request, kwargs }; + let model = arguments.model()?; + let custom_llm_provider = arguments.custom_llm_provider()?; + let (wire_document, retained_document) = + ProjectedDocument::project(py, &arguments.document()?)?.into_parts(); + let api_key = arguments.api_key()?; + let specs = consumed_optional_params(&model, custom_llm_provider.as_deref()) + .map_err(ocr_error_to_pyerr)?; + let names = specs.iter().map(|spec| spec.name).collect::>(); + let optional_params = project_optional_fields(kwargs, &names)?; + let input_sources = request_input_sources( + kwargs, + names + .iter() + .copied() + .chain(["api_key", "api_base", "extra_headers"]), + )?; + let azure_ad_token_provider = kwargs + .get_item("azure_ad_token_provider")? + .and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER)); + let wire = OcrWireRequest { + model, + document: wire_document, + api_key: api_key.extract()?, + api_base: arguments.api_base()?, + custom_llm_provider, + extra_headers: arguments.extra_headers()?, + optional_params, + input_sources, + timeout_seconds: arguments.timeout_seconds()?, + }; + let request = decode_request(wire).map_err(ocr_error_to_pyerr)?; + let provider = request.provider_name(); + Ok(ProjectedOcrCall { + request: request.with_host_hooks(Arc::new(BridgeOcrHooks), None), + fields: ProjectedOcrFields { + boundary_request, + document: retained_document, + api_key: api_key.unbind(), + azure_ad_token_provider, + provider, + secret_fields: specs + .into_iter() + .filter(|spec| spec.secret) + .map(|spec| spec.name) + .collect(), + }, + }) +} + +pub(super) fn admitted_call(outcome: NativeOutcome) -> PyResult { + match outcome { + NativeOutcome::Completed(call) => Ok(call), + NativeOutcome::Declined(reason) => Err(RustBridgeDeclined::new_err(format!( + "native OCR admission declined: {reason:?}" + ))), + } +} + +#[cfg(test)] +mod tests { + use litellm_core::Error; + use litellm_core::ocr::OcrDecline; + use pyo3::exceptions::{PyKeyError, PyTypeError, PyValueError}; + + use super::*; + + fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> { + let locals = PyDict::new(py); + py.run(source, Some(&locals), Some(&locals)).unwrap(); + locals + } + + fn arguments<'a, 'py>( + request: &'a Bound<'py, PyAny>, + kwargs: &'a Bound<'py, PyDict>, + ) -> OcrArguments<'a, 'py> { + OcrArguments { request, kwargs } + } + + fn project_document( + py: Python<'_>, + document: &Bound<'_, PyAny>, + ) -> PyResult<(Value, Py)> { + ProjectedDocument::project(py, document).map(ProjectedDocument::into_parts) + } + + fn stub_timeout_conversion(py: Python<'_>) { + eval( + py, + c" +import sys +import types +timeouts = types.ModuleType('litellm.rust_bridge.timeouts') +timeouts.timeout_to_seconds = lambda timeout: None if timeout is None else float(timeout) +sys.modules.setdefault('litellm', types.ModuleType('litellm')) +sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bridge')) +sys.modules['litellm.rust_bridge.timeouts'] = timeouts +", + ); + } + + #[test] + fn typed_initial_decline_uses_bridge_decline_contract() { + Python::initialize(); + Python::attach(|py| { + let Err(error) = admitted_call(NativeOutcome::Declined(OcrDecline::HostOperations)) + else { + panic!("unsupported host operations should decline admission"); + }; + assert!(error.is_instance_of::(py)); + }); + } + + #[test] + fn post_admission_error_does_not_use_bridge_decline_contract() { + Python::initialize(); + Python::attach(|py| { + let error = ocr_error_to_pyerr(Error::InvalidRequest("callback result".into())); + assert!(error.is_instance_of::(py)); + assert!(!error.is_instance_of::(py)); + }); + } + + #[test] + fn kwargs_override_request_attributes_including_explicit_none() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Request: + def __init__(self): + self.accesses = [] + def __getattribute__(self, name): + if name != 'accesses': + object.__getattribute__(self, 'accesses').append(name) + return object.__getattribute__(self, name) +request = Request() +request.model = 'from-request' +request.custom_llm_provider = 'mistral' +kwargs = {'model': 'from-kwargs', 'custom_llm_provider': None} +", + ); + let request = locals.get_item("request").unwrap().unwrap(); + let kwargs = locals + .get_item("kwargs") + .unwrap() + .unwrap() + .cast_into::() + .unwrap(); + let arguments = arguments(&request, &kwargs); + assert_eq!(arguments.model().unwrap(), "from-kwargs"); + assert_eq!(arguments.custom_llm_provider().unwrap(), None); + let accesses: Vec = request.getattr("accesses").unwrap().extract().unwrap(); + assert_eq!(accesses, Vec::::new()); + }); + } + + #[test] + fn missing_kwargs_read_the_request_property_once() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Request: + def __init__(self): + self.reads = 0 + @property + def model(self): + self.reads += 1 + return 'mistral-ocr-latest' +request = Request() +kwargs = {} +", + ); + let request = locals.get_item("request").unwrap().unwrap(); + let kwargs = locals + .get_item("kwargs") + .unwrap() + .unwrap() + .cast_into::() + .unwrap(); + assert_eq!( + arguments(&request, &kwargs).model().unwrap(), + "mistral-ocr-latest" + ); + assert_eq!( + request.getattr("reads").unwrap().extract::().unwrap(), + 1 + ); + }); + } + + #[test] + fn request_property_exceptions_keep_their_identity() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +failure = LookupError('model failed') +class Request: + @property + def model(self): + raise failure +request = Request() +kwargs = {} +", + ); + let request = locals.get_item("request").unwrap().unwrap(); + let kwargs = locals + .get_item("kwargs") + .unwrap() + .unwrap() + .cast_into::() + .unwrap(); + let error = arguments(&request, &kwargs).model().unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + }); + } + + #[test] + fn unused_raising_property_is_never_inspected() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Request: + @property + def unused(self): + raise RuntimeError('unused') + model = 'mistral-ocr-latest' + custom_llm_provider = None +request = Request() +kwargs = {} +", + ); + let request = locals.get_item("request").unwrap().unwrap(); + let kwargs = locals + .get_item("kwargs") + .unwrap() + .unwrap() + .cast_into::() + .unwrap(); + let arguments = arguments(&request, &kwargs); + assert_eq!(arguments.model().unwrap(), "mistral-ocr-latest"); + assert_eq!(arguments.custom_llm_provider().unwrap(), None); + }); + } + + #[test] + fn document_reader_mutations_are_visible_to_later_field_reads() { + Python::initialize(); + Python::attach(|py| { + stub_timeout_conversion(py); + let locals = eval( + py, + c" +class Request: + api_base = 'original' + timeout = 1 + @property + def document(self): + return document +class Reader: + def read(self): + Request.api_base = 'mutated' + Request.timeout = 9 + return b'abc' +document = {'type': 'file', 'file': Reader()} +request = Request() +kwargs = {} +", + ); + let request = locals.get_item("request").unwrap().unwrap(); + let kwargs = locals + .get_item("kwargs") + .unwrap() + .unwrap() + .cast_into::() + .unwrap(); + let arguments = arguments(&request, &kwargs); + let document = arguments.document().unwrap(); + project_document(py, &document).unwrap(); + assert_eq!(arguments.api_base().unwrap().as_deref(), Some("mutated")); + assert_eq!(arguments.timeout_seconds().unwrap(), Some(9.0)); + }); + } + + #[test] + fn captured_api_key_keeps_the_original_python_object() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +key = object() +class Request: + api_key = None +request = Request() +kwargs = {'api_key': key} +", + ); + let request = locals.get_item("request").unwrap().unwrap(); + let kwargs = locals + .get_item("kwargs") + .unwrap() + .unwrap() + .cast_into::() + .unwrap(); + let captured = arguments(&request, &kwargs).api_key().unwrap(); + assert!( + captured + .unbind() + .bind(py) + .is(locals.get_item("key").unwrap().unwrap()) + ); + }); + } + + #[test] + fn file_documents_are_encoded_and_other_documents_keep_the_python_object() { + Python::initialize(); + Python::attach(|py| { + let file = py + .eval( + c"{'type': 'file', 'file': b'%PDF-1.4', 'mime_type': 'application/pdf'}", + None, + None, + ) + .unwrap(); + assert_eq!( + project_document(py, &file).unwrap().0, + serde_json::json!({ + "type": "document_url", + "document_url": "data:application/pdf;base64,JVBERi0xLjQ=", + }) + ); + + let original = py + .eval( + c"{'type': 'document_url', 'document_url': 'https://example.com/a.pdf'}", + None, + None, + ) + .unwrap(); + let (wire, retained) = project_document(py, &original).unwrap(); + assert_eq!( + wire, + serde_json::json!({ + "type": "document_url", + "document_url": "https://example.com/a.pdf", + }) + ); + assert!(retained.bind(py).is(&original)); + }); + } + + #[test] + fn unknown_document_types_reach_existing_downstream_validation() { + Python::initialize(); + Python::attach(|py| { + let document = py + .eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None) + .unwrap(); + let wire_document = project_document(py, &document).unwrap().0; + assert_eq!( + wire_document, + serde_json::json!({"type": "mystery", "mystery": "x"}) + ); + let error = match decode_request(OcrWireRequest { + model: "mistral/mistral-ocr-latest".into(), + document: wire_document, + api_key: None, + api_base: None, + custom_llm_provider: None, + extra_headers: None, + optional_params: Map::new(), + input_sources: Default::default(), + timeout_seconds: None, + }) { + Ok(_) => panic!("unknown discriminators belong to core validation"), + Err(error) => error, + }; + assert!(error.to_string().contains("document")); + }); + } + + #[test] + fn document_discriminator_errors_keep_their_existing_exceptions() { + Python::initialize(); + Python::attach(|py| { + let missing = py.eval(c"{}", None, None).unwrap(); + assert!( + project_document(py, &missing) + .unwrap_err() + .is_instance_of::(py) + ); + + let non_string = py.eval(c"{'type': 1}", None, None).unwrap(); + assert!( + project_document(py, &non_string) + .unwrap_err() + .is_instance_of::(py) + ); + + let locals = eval( + py, + c" +failure = RuntimeError('type lookup failed') +class Document: + def __getitem__(self, key): + raise failure +document = Document() +", + ); + let error = + project_document(py, &locals.get_item("document").unwrap().unwrap()).unwrap_err(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + }); + } + + #[test] + fn document_classification_happens_once() { + Python::initialize(); + Python::attach(|py| { + let locals = eval( + py, + c" +class Document(dict): + def __init__(self): + super().__init__({'file': b'abc'}) + self.reads = [] + def __getitem__(self, key): + self.reads.append(key) + if key == 'type': + return 'file' if self.reads.count('type') == 1 else 'document_url' + return super().__getitem__(key) +document = Document() +", + ); + let document = locals.get_item("document").unwrap().unwrap(); + let (wire, retained) = project_document(py, &document).unwrap(); + assert_eq!(wire["type"], "document_url"); + assert!(!retained.bind(py).is(&document)); + let reads: Vec = document.getattr("reads").unwrap().extract().unwrap(); + assert_eq!(reads, ["type", "mime_type", "file"]); + }); + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs index 94f760e1563..051ac19d4fb 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs @@ -1,8 +1,7 @@ use litellm_core::Error; use std::future::Future; -use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr}; -use litellm_core::ocr::wire::{OcrWireRequest, decode_request, is_supported_request}; +use litellm_core::ocr::wire::{OcrWireRequest, decode_request}; use pyo3::prelude::*; use serde_json::Value; @@ -38,37 +37,20 @@ fn prepare_ocr( extra_headers, timeout, } = options; - if is_supported_request(&model, custom_llm_provider.as_deref()) { - let request = decode_request(OcrWireRequest { - model, - document, - api_key, - api_base, - custom_llm_provider, - extra_headers, - optional_params, - input_sources, - timeout_seconds: timeout.map(|value| value.as_secs_f64()), - })?; - return litellm_core::ocr::ocr(request) - .await - .map(|response| response.into_json()); - } - run_ocr(OcrRequest { - model: &model, + let request = decode_request(OcrWireRequest { + model, document, - api_key: api_key.as_deref(), - api_base: api_base.as_deref(), - custom_llm_provider: custom_llm_provider.as_deref(), + api_key, + api_base, + custom_llm_provider, extra_headers, optional_params, - timeout, - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: Default::default(), - litellm_call_id: None, - }) - .await + input_sources, + timeout_seconds: timeout.map(|value| value.as_secs_f64()), + })?; + litellm_core::ocr::ocr(request) + .await + .map(|response| response.into_json()) }) } @@ -96,22 +78,3 @@ bridge_route! { prepare = prepare_ocr, errors = ocr_error_to_pyerr, } - -#[cfg(test)] -mod tests { - use litellm_core::ocr::wire::is_supported_request; - - #[test] - fn native_activation_includes_migrated_providers() { - assert!(is_supported_request("model", Some("mistral"))); - assert!(is_supported_request("pixtral-12b", Some("azure_ai"))); - assert!(is_supported_request( - "documentintelligence/prebuilt-read", - Some("azure_ai") - )); - assert!(is_supported_request("parse-v3", Some("reducto"))); - assert!(is_supported_request("parse-legacy", Some("reducto"))); - assert!(is_supported_request("mistral-ocr", Some("vertex_ai"))); - assert!(is_supported_request("deepseek-ocr", Some("vertex_ai"))); - } -} diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index ec71c983ff2..de8a93dd8b1 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -13,18 +13,6 @@ from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KE from litellm.rust_bridge.bindings import NativeBinding from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds -_RUST_OCR_SECRET_FIELDS: Final = frozenset( - {"azure_ad_token", "client_secret", "azure_federated_token_file", "vertex_credentials", "vertex_ai_credentials"} -) - - -def redact_logging_params(params: Mapping[str, object]) -> dict[str, object]: - return { # mutable-ok: Logging.update_from_kwargs requires concrete params - name: "****" if name in _RUST_OCR_SECRET_FIELDS else value - for name, value in params.items() - if name != "proxy_server_request" - } - @dataclass(frozen=True, slots=True) class LiteLLMOcrRequest: diff --git a/tests/rust-python-harness/strategies/trace_parity/gateway/execution.py b/tests/rust-python-harness/strategies/trace_parity/gateway/execution.py index 2bd3a50f39f..860e872dd44 100644 --- a/tests/rust-python-harness/strategies/trace_parity/gateway/execution.py +++ b/tests/rust-python-harness/strategies/trace_parity/gateway/execution.py @@ -1,7 +1,8 @@ from __future__ import annotations -import asyncio -from collections.abc import Awaitable, Callable +import json +import subprocess +from functools import cache from pathlib import Path from typing import Final, Protocol, cast @@ -28,12 +29,12 @@ class _GatewayClient(Protocol): def _collect_python(fixture: RouteFixture) -> tuple[FunctionTraceEvent, ...]: - import litellm from fastapi.testclient import TestClient + import litellm + from litellm.proxy import proxy_server from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.anthropic_endpoints.endpoints import user_api_key_auth - from litellm.proxy import proxy_server provider_model: Final = cast(str, fixture.kwargs["provider_model"]) model_alias: Final = cast(str, fixture.kwargs["model_alias"]) @@ -76,24 +77,24 @@ def _collect_python(fixture: RouteFixture) -> tuple[FunctionTraceEvent, ...]: def _collect_rust(fixture: RouteFixture) -> tuple[FunctionTraceEvent, ...]: - from litellm.rust_bridge import get_native_bridge - - bridge: Final[object | None] = get_native_bridge() - trace: Final[object | None] = getattr(bridge, "_trace", None) if bridge is not None else None - gateway_messages: Final[object | None] = getattr(trace, "gateway_messages", None) - if gateway_messages is None or not callable(gateway_messages): - raise RuntimeError("native Rust trace bridge does not expose gateway_messages") - invoke_gateway: Final = cast(Callable[[str, str, str, object], Awaitable[object]], gateway_messages) - - async def invoke() -> object: - return await invoke_gateway( - cast(str, fixture.kwargs["model_alias"]), - cast(str, fixture.kwargs["provider_model"]), - cast(str, fixture.kwargs["api_base"]), - fixture.kwargs["body"], - ) - - result: Final = asyncio.run(invoke()) + payload: Final = json.dumps( + { + "model_alias": fixture.kwargs["model_alias"], + "provider_model": fixture.kwargs["provider_model"], + "api_base": fixture.kwargs["api_base"], + "body": fixture.kwargs["body"], + } + ) + completed: Final = subprocess.run( + (_gateway_trace_binary(),), + input=payload, + capture_output=True, + text=True, + check=False, + ) + if completed.returncode != 0: + raise RuntimeError(f"Rust gateway trace failed: {completed.stderr.strip()}") + result: Final = json.loads(completed.stdout) payload: Final = TraceResponsePayload.model_validate(result) response: Final = _GatewayResponsePayload.model_validate(payload.response) if response.status != 200: @@ -101,6 +102,34 @@ def _collect_rust(fixture: RouteFixture) -> tuple[FunctionTraceEvent, ...]: return native_trace_events(payload) +@cache +def _gateway_trace_binary() -> Path: + repo_root: Final = next(parent for parent in Path(__file__).resolve().parents if (parent / "litellm-rust").is_dir()) + rust_root: Final = repo_root / "litellm-rust" + completed: Final = subprocess.run( + ( + "cargo", + "build", + "--quiet", + "--package", + "litellm-ai-gateway", + "--features", + "trace-parity", + "--bin", + "trace-parity-gateway", + "--target-dir", + rust_root / "target", + ), + cwd=rust_root, + capture_output=True, + text=True, + check=False, + ) + if completed.returncode != 0: + raise RuntimeError(f"Rust gateway trace build failed: {completed.stderr.strip()}") + return rust_root / "target" / "debug" / "trace-parity-gateway" + + def _collect(scenario: TraceScenario, engine: Engine) -> tuple[FunctionTraceEvent, ...] | TraceExecutionFailure: try: with replay_server() as provider: