fixes and refactor

This commit is contained in:
Yujong Lee 2026-09-12 10:57:49 -07:00
parent d6feba712d
commit 9132a5343b
40 changed files with 1315 additions and 1167 deletions

View file

@ -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",

View file

@ -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"] }

View file

@ -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 <gateway-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`.**

View file

@ -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",

View file

@ -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)
}

View file

@ -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<MaybeTlsStream<TcpStream>>;
type UpstreamTx = SplitSink<ResponsesUpstreamWs, Message>;
type UpstreamRx = SplitStream<ResponsesUpstreamWs>;
#[derive(Clone)]
pub struct ResponsesWebSocketConnection {
socket: Arc<Mutex<Option<ResponsesUpstreamWs>>>,
}
impl ResponsesWebSocketConnection {
pub async fn connect_url(
url: &str,
headers: &HashMap<String, String>,
timeout: Option<Duration>,
) -> Result<Self, Error> {
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::<HeaderName>()
.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<Option<String>, 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<String, Error> {
api_key
.map(str::trim)

View file

@ -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(),
),

View file

@ -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<GatewayResponse>,
pub error: Option<String>,
pub trace: Vec<litellm_core::observability::FunctionTraceEvent>,
}
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,

View file

@ -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 }

View file

@ -1,3 +1,30 @@
use std::future::Future;
use std::pin::Pin;
pub enum HostCallStep<O, C> {
Host(O),
Complete(C),
}
pub type HostCallFuture<'a, O, C> =
Pin<Box<dyn Future<Output = Result<HostCallStep<O, C>, 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<Self::Result>,
) -> HostCallFuture<'_, Self::Operation, Self::Complete>;
fn interrupt(
&mut self,
failure: HostFailure,
) -> HostCallFuture<'_, Self::Operation, Self::Complete>;
}
pub enum HostStep<V, S> {
Ready(V),
Suspend(S),

View file

@ -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<u16> {
match self {
Self::InvalidRequest(_) => Some(400),
Self::MissingDocumentUrl => Some(500),
Self::Http { status, .. } => Some(*status),
_ => None,
}
}
@ -115,6 +119,7 @@ impl From<crate::ocr::error::OcrRequestError> 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()),
}
}

View file

@ -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();
}
}

View file

@ -12,7 +12,7 @@ pub(crate) fn transform_ocr_request(
params: &DeepSeekOcrParams,
) -> Result<DeepSeekOcrRequest, OcrRequestError> {
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(),

View file

@ -13,7 +13,7 @@ pub(crate) fn transform_ocr_request(
) -> Result<DocumentIntelligenceRequest, OcrRequestError> {
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(

View file

@ -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(),
}
}

View file

@ -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(

View file

@ -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<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse
request: LiteLLMOcrRequest,
) -> Self::PreCallFuture<'a> {
Box::pin(async move {
if !self.hooks.has_guardrails() {
if !self.hooks.intercepts_requests() {
return Ok(request);
}
let changed = self

View file

@ -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<T> = Result<NativeOutcome<T>, 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<HostPhase> {
match self {
Self::Lifecycle(phase) => Some(*phase),
Self::Success { .. } => Some(HostPhase::Success),
Self::Failure { .. } => Some(HostPhase::Failure),
_ => None,
}
}
}
pub enum OcrHostResult {
Request(Result<(Box<LiteLLMOcrRequest>, bool), Error>),
Lifecycle(Result<(), HostFailure>),
@ -80,10 +91,7 @@ pub enum OcrHostResult {
PostCall(Result<OcrPostCallRequest, Error>),
}
pub enum OcrCallStep {
Host(OcrHostOperation),
Complete(LiteLLMOcrResponse),
}
pub type OcrCallStep = HostCallStep<OcrHostOperation, LiteLLMOcrResponse>;
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<Self::Result>,
) -> 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<OcrHostResult>,
@ -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<PendingOperation>,
has_guardrails: bool,
intercepts_requests: bool,
terminal: Arc<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
}
@ -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> {

View file

@ -69,7 +69,7 @@ pub(crate) async fn transform_request_body<B>(
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

View file

@ -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<T> {
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<Vec<OptionalParamSpec>, 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<LiteLLMOcrRequest, Error> {
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<LiteLLMOcrRequest, Error>
})
}
fn decode_document(value: Value) -> Result<OcrDocument, OcrRequestError> {
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<String, InputSource>, name: &str) -> InputSource {
sources.get(name).copied().unwrap_or_default()
}
@ -233,39 +270,6 @@ pub fn decode_response<T: DeserializeOwned>(
})
}
pub fn decode_pre_call_result(
original: OcrPreCallRequest,
value: Value,
) -> Result<OcrPreCallRequest, OcrRequestError> {
#[derive(Deserialize)]
struct Changed {
document: OcrDocument,
#[serde(default)]
optional_params: Map<String, Value>,
}
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<OcrDuringCallRequest, OcrRequestError> {
#[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)
);
}
}
}

View file

@ -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<MaybeTlsStream<TcpStream>>;
static TLS_CONFIG: OnceLock<Arc<ClientConfig>> = OnceLock::new();
fn build_tls_config() -> Result<ClientConfig, Box<tokio_tungstenite::tungstenite::Error>> {
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<Arc<ClientConfig>, Box<tokio_tungstenite::tungstenite::Error>> {
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<R>(
request: R,
) -> Result<(ResponsesUpstreamWs, Response), Box<tokio_tungstenite::tungstenite::Error>>
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<Mutex<Option<ResponsesUpstreamWs>>>,
}
impl ResponsesWebSocketConnection {
pub async fn connect_url(
url: &str,
headers: &HashMap<String, String>,
timeout: Option<Duration>,
) -> Result<Self, Error> {
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::<HeaderName>()
.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<Option<String>, 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::*;

View file

@ -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
}

View file

@ -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
}

View file

@ -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,

View file

@ -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
}

View file

@ -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

View file

@ -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))
}

View file

@ -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",
]
);
}

View file

@ -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<O> {
Host(O),
Complete,
}
type NativeCallFuture<'a, O> =
Pin<Box<dyn Future<Output = Result<NativeCallStep<O>, 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<Self::Result>) -> 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<R> =
HostStep<NativeCallStep<<<R as PythonRoute>::Call as NativeCall>::Operation>, Py<PyAny>>;
type HostResumeStep<R> = HostStep<
NativeCallStep<
<<R as PythonRoute>::Call as NativeCall>::Operation,
<<R as PythonRoute>::Call as NativeCall>::Complete,
>,
Py<PyAny>,
>;
struct NativeCallState<C: NativeCall> {
call: C,
result: Option<Result<NativeCallStep<C::Operation>, litellm_core::Error>>,
result: Option<Result<NativeCallStep<C::Operation, C::Complete>, litellm_core::Error>>,
}
enum PendingOperation {
@ -151,7 +141,11 @@ impl<R: PythonRoute> PythonLifecycle<R> {
}
}
fn take_native_result(&self) -> PyResult<NativeCallStep<<R::Call as NativeCall>::Operation>> {
fn take_native_result(
&self,
) -> PyResult<
NativeCallStep<<R::Call as NativeCall>::Operation, <R::Call as NativeCall>::Complete>,
> {
self.call
.as_ref()
.ok_or_else(missing_state)?
@ -214,7 +208,7 @@ impl<R: PythonRoute> PythonLifecycle<R> {
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<Self::Result>,
) -> Pin<
Box<
dyn Future<Output = Result<NativeCallStep<Self::Operation>, 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<Output = Result<NativeCallStep<Self::Operation>, litellm_core::Error>>
+ Send
+ '_,
>,
> {
Box::pin(async { Ok(NativeCallStep::Complete) })
) -> HostCallFuture<'_, Self::Operation, Self::Complete> {
Box::pin(async { Ok(NativeCallStep::Complete(())) })
}
}

View file

@ -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<Bound<'py, PyAny>> {
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)?)
}

View file

@ -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(())

View file

@ -31,19 +31,21 @@ impl PythonLogger {
py: Python<'_>,
kwargs: &Py<PyDict>,
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::<PyDict>()?)?;
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::<PyDict>()?,
redact(
py,
&to_py(py, &pre_call.optional_params)?
.into_bound(py)
.cast_into::<PyDict>()?,
secret_fields,
)?,
)?;
let params = PyDict::new(py);
params.set_item(
@ -93,8 +95,8 @@ impl PythonLogger {
&self,
py: Python<'_>,
original_response: &Value,
body: &Option<Py<PyDict>>,
headers: &Option<Py<PyDict>>,
body: Option<&Py<PyDict>>,
headers: Option<&Py<PyDict>>,
) -> 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<Py<PyDict>> {
let redacted = PyDict::new(py);
for (name, value) in params {
let name = name.extract::<String>()?;
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<PyAny>> {
py.import("litellm.rust_bridge.ocr")?
.getattr("_response")?

View file

@ -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<Self> {
let py = document.py();
let mime_type = match document.get_item("mime_type") {
Ok(value) => Some(value.extract::<String>()?),
Err(error) if error.is_instance_of::<pyo3::exceptions::PyKeyError>(py) => None,
Err(error) => return Err(error),
};
let file = document.get_item("file").map_err(|error| {
if error.is_instance_of::<pyo3::exceptions::PyKeyError>(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::<PyDict>()?
.get_item("mime_type")?
.map(|value| value.extract::<String>())
.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::<FileDocumentInput>().err().unwrap();
assert!(error.is_instance_of::<PyTypeError>(py));
let reads: usize = locals
.get_item("reader")
.unwrap()
.unwrap()
.getattr("reads")
.unwrap()
.extract()
.unwrap();
assert_eq!(reads, 0);
});
}

View file

@ -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<u16>) -> 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::<PyValueError>(py));
assert_eq!(mapped.value(py).to_string(), "Document URL is required");
}
let mapped = to_pyerr(Error::MissingDocumentUrl);
assert!(mapped.is_instance_of::<pyo3::exceptions::PyValueError>(py));
assert_eq!(mapped.value(py).to_string(), "Document URL is required");
assert_eq!(
mapped
.value(py)
.getattr("status_code")
.unwrap()
.extract::<u16>()
.unwrap(),
500
);
let mapped = to_pyerr(Error::Http {
status: 429,
body: r#"{"message":"rate limited"}"#.to_string(),

View file

@ -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<Py<PyAny>>,
data: OcrHostData,
}
enum OcrHostData {
Unprojected { request: Py<PyAny> },
Projected(ProjectedOcrHost),
Released,
}
struct ProjectedOcrHost {
fields: ProjectedOcrFields,
pre_call: Option<callbacks::OcrLoggingFields>,
document: Option<Py<PyAny>>,
api_key: Option<Py<PyAny>>,
azure_ad_token_provider: Option<PythonTokenProvider>,
provider: String,
retained_fields: Option<Py<PyDict>>,
body: Option<Py<PyDict>>,
headers: Option<Py<PyDict>>,
}
struct AdmittedOcrCall {
request: litellm_core::ocr::LiteLLMOcrRequest,
document: Py<PyAny>,
api_key: Py<PyAny>,
azure_ad_token_provider: Option<PythonTokenProvider>,
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<ResolvedCredential> {
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<OcrDuringCallRequest> {
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::<PyDict>()?;
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::<String>()?, value.extract::<String>()?)))
@ -132,59 +146,18 @@ impl PythonOcrHost {
) -> PyResult<OcrPostCallRequest> {
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<Self::Result>,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<NativeCallStep<Self::Operation>, 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<NativeCallStep<Self::Operation>, 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<OcrHostResult> {
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<Bound<'py, PyAny>> {
match self.kwargs.get_item(name)? {
Some(value) => Ok(value),
None => self.request.getattr(name),
}
}
fn model(&self) -> PyResult<String> {
self.lookup("model")?.extract()
}
fn custom_llm_provider(&self) -> PyResult<Option<String>> {
self.lookup("custom_llm_provider")?.extract()
}
fn document(&self) -> PyResult<CapturedDocument<'py>> {
Ok(CapturedDocument(self.lookup("document")?))
}
fn api_key(&self) -> PyResult<CapturedApiKey<'py>> {
Ok(CapturedApiKey(self.lookup("api_key")?))
}
fn api_base(&self) -> PyResult<Option<String>> {
self.lookup("api_base")?.extract()
}
fn extra_headers(&self) -> PyResult<Option<Map<String, Value>>> {
self.lookup("extra_headers")?
.extract::<Option<Py<PyAny>>>()?
.map(|value| from_py(value.bind(self.request.py())))
.transpose()
}
fn timeout_seconds(&self) -> PyResult<Option<f64>> {
Ok(self
.lookup("timeout")?
.extract::<Option<Py<PyAny>>>()?
.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<Option<String>> {
self.0.extract()
}
fn into_object(self) -> Py<PyAny> {
self.0.unbind()
}
}
enum DocumentKind {
File,
Other,
}
impl FromPyObject<'_, '_> for DocumentKind {
type Error = PyErr;
fn extract(document: Borrowed<'_, '_, PyAny>) -> PyResult<Self> {
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<AdmittedOcrCall> {
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<Value> {
match document.extract::<DocumentKind>()? {
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<Py<PyAny>> {
match document.extract::<DocumentKind>()? {
DocumentKind::File => to_py(py, wire_document),
DocumentKind::Other => Ok(document.clone().unbind()),
}
}
fn admitted_call(outcome: NativeOutcome<OcrCall>) -> PyResult<OcrCall> {
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<Py<PyAny>> {
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::<RustBridgeDeclined>(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::<PyValueError>(py));
assert!(!error.is_instance_of::<RustBridgeDeclined>(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::<PyDict>()
.unwrap();
let arguments = arguments(&request, &kwargs);
assert_eq!(arguments.model().unwrap(), "from-kwargs");
assert_eq!(arguments.custom_llm_provider().unwrap(), None);
let accesses: Vec<String> = request.getattr("accesses").unwrap().extract().unwrap();
assert_eq!(accesses, Vec::<String>::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::<PyDict>()
.unwrap();
assert_eq!(
arguments(&request, &kwargs).model().unwrap(),
"mistral-ocr-latest"
);
assert_eq!(
request.getattr("reads").unwrap().extract::<i32>().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::<PyDict>()
.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::<PyDict>()
.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::<PyDict>()
.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::<PyDict>()
.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::<PyKeyError>(py)
);
let non_string = py.eval(c"{'type': 1}", None, None).unwrap();
assert!(
extract_document(py, &non_string)
.unwrap_err()
.is_instance_of::<PyTypeError>(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::<DocumentKind>().unwrap(),
DocumentKind::File
));
let reads: Vec<String> = 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<String> = document.getattr("reads").unwrap().extract().unwrap();
assert_eq!(reads, ["type", "file", "type"]);
});
}
}

View file

@ -2,6 +2,7 @@ mod callbacks;
mod document;
mod errors;
mod lifecycle;
mod project;
mod value;
use pyo3::prelude::*;

View file

@ -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<PyAny>,
pub document: Py<PyAny>,
pub api_key: Py<PyAny>,
pub azure_ad_token_provider: Option<PythonTokenProvider>,
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<Bound<'py, PyAny>> {
match self.kwargs.get_item(name)? {
Some(value) => Ok(value),
None => self.request.getattr(name),
}
}
fn model(&self) -> PyResult<String> {
self.lookup("model")?.extract()
}
fn custom_llm_provider(&self) -> PyResult<Option<String>> {
self.lookup("custom_llm_provider")?.extract()
}
fn document(&self) -> PyResult<Bound<'py, PyAny>> {
self.lookup("document")
}
fn api_key(&self) -> PyResult<Bound<'py, PyAny>> {
self.lookup("api_key")
}
fn api_base(&self) -> PyResult<Option<String>> {
self.lookup("api_base")?.extract()
}
fn extra_headers(&self) -> PyResult<Option<Map<String, Value>>> {
self.lookup("extra_headers")?
.extract::<Option<Py<PyAny>>>()?
.map(|value| from_py(value.bind(self.request.py())))
.transpose()
}
fn timeout_seconds(&self) -> PyResult<Option<f64>> {
Ok(self
.lookup("timeout")?
.extract::<Option<Py<PyAny>>>()?
.map(|value| python_timeout_seconds(self.request.py(), value))
.transpose()?
.flatten())
}
}
enum ProjectedDocument {
File { wire: Value, retained: Py<PyAny> },
Other { wire: Value, retained: Py<PyAny> },
}
impl ProjectedDocument {
fn project(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<Self> {
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<PyAny>) {
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<ProjectedOcrCall> {
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::<Vec<_>>();
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<OcrCall>) -> PyResult<OcrCall> {
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<PyAny>)> {
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::<RustBridgeDeclined>(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::<PyValueError>(py));
assert!(!error.is_instance_of::<RustBridgeDeclined>(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::<PyDict>()
.unwrap();
let arguments = arguments(&request, &kwargs);
assert_eq!(arguments.model().unwrap(), "from-kwargs");
assert_eq!(arguments.custom_llm_provider().unwrap(), None);
let accesses: Vec<String> = request.getattr("accesses").unwrap().extract().unwrap();
assert_eq!(accesses, Vec::<String>::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::<PyDict>()
.unwrap();
assert_eq!(
arguments(&request, &kwargs).model().unwrap(),
"mistral-ocr-latest"
);
assert_eq!(
request.getattr("reads").unwrap().extract::<i32>().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::<PyDict>()
.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::<PyDict>()
.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::<PyDict>()
.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::<PyDict>()
.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::<PyKeyError>(py)
);
let non_string = py.eval(c"{'type': 1}", None, None).unwrap();
assert!(
project_document(py, &non_string)
.unwrap_err()
.is_instance_of::<PyTypeError>(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<String> = document.getattr("reads").unwrap().extract().unwrap();
assert_eq!(reads, ["type", "mime_type", "file"]);
});
}
}

View file

@ -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")));
}
}

View file

@ -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:

View file

@ -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: