refactor(rust): drop the ai-gateway dependency from python-bridge

Move the Responses WebSocket transport (connection type, dial, splice)
from ai-gateway io/responses_ws into core responses::connection, with
send/recv locked on separate halves and a Drop that closes the upstream
socket when the last clone is abandoned.

Move the OCR pipeline (prepare, provider call, document fetch, Azure DI
polling) into core ocr, mirroring the audio_transcription split: core
owns prepare/execute plus the ocr()/ocr_with_observer() entrypoints, the
gateway keeps only its lifecycle hooks (custom loggers/guardrails) and
request type, re-exporting through io/ocr unchanged.

Delete litellm-ai-gateway from python-bridge's dependencies and guard
the direction with manifest tests: python-bridge depends only on core +
python-interop, and core carries no pyo3/pythonize/pyo3-async-runtimes.
This commit is contained in:
Yujong Lee 2026-09-03 14:09:59 -07:00
parent d9ad7ae2e6
commit a4bb1c8066
28 changed files with 1747 additions and 2099 deletions

View file

@ -1446,6 +1446,8 @@ dependencies = [
"aws-smithy-runtime-api",
"aws-types",
"base64",
"futures-channel",
"futures-util",
"rand 0.8.7",
"reqwest",
"rstest",
@ -1454,6 +1456,7 @@ dependencies = [
"sha2 0.10.9",
"thiserror 2.0.19",
"tokio",
"tokio-tungstenite",
"tracing",
"tracing-subscriber",
]
@ -1464,7 +1467,6 @@ version = "0.1.0"
dependencies = [
"criterion",
"futures-util",
"litellm-ai-gateway",
"litellm-core",
"litellm-python-interop",
"pyo3",

View file

@ -1,14 +0,0 @@
use std::sync::OnceLock;
use std::time::Duration;
const HTTP_CLIENT_TIMEOUT_SECS: u64 = 600;
pub(crate) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(HTTP_CLIENT_TIMEOUT_SECS))
.build()
.expect("failed to build reqwest client")
})
}

View file

@ -29,9 +29,6 @@ pub(crate) const DEFAULT_FLUSH_INTERVAL_MS: u64 = 500;
#[cfg(feature = "server")]
pub(crate) const DEFAULT_PROVIDER: &str = "openai";
pub(crate) const DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
/// HTTP path for the non-streaming Anthropic Messages route.
#[cfg(feature = "server")]
pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages";

View file

@ -1,557 +1,9 @@
use std::sync::Arc;
use std::time::Duration;
//! Compatibility re-exports: the Responses WebSocket transport moved to
//! `litellm-core` (`litellm_core::responses::connection`) so the python bridge
//! can use it without depending on this crate. Re-exported here so existing
//! gateway imports keep working.
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use litellm_core::Error;
use litellm_core::http_utils::string_headers;
use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use litellm_core::responses::types::{ResponsesWebSocketRequest, ResponsesWsEvent};
use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig;
use tokio::net::TcpStream;
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, connect_async};
use crate::constants::{
DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS,
pub use litellm_core::responses::connection::{
ResponsesUpstreamWs, ResponsesWebSocketConnection, ResponsesWebSocketStreaming,
async_responses_websocket, responses_ws,
};
const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable";
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(
input: ResponsesWebSocketRequest,
options: &RequestOptions,
_context: &LiteLlmRequestContext,
) -> Result<Self, Error> {
if !litellm_core::responses::websocket::native_websocket_supported(
options.custom_llm_provider.as_deref().unwrap_or("openai"),
) {
return Err(Error::Unsupported("unsupported native WebSocket provider"));
}
let headers = string_headers("Responses WebSocket", options.extra_headers.clone())?;
let mut request = input
.url
.as_str()
.into_client_request()
.map_err(|error| Error::Network(error.to_string()))?;
for (name, value) in headers {
let header_name = name
.parse::<HeaderName>()
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
let header_value = HeaderValue::from_str(&value)
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
request.headers_mut().insert(header_name, header_value);
}
let connect = connect_async(request);
let result = match options.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)
.filter(|value| !value.is_empty())
.map(str::to_string)
.or_else(|| {
std::env::var(OPENAI_API_KEY_ENV)
.ok()
.filter(|value| !value.trim().is_empty())
})
.ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string()))
}
async fn dial_upstream(
model: &str,
api_key: &str,
api_base: Option<&str>,
) -> Result<ResponsesUpstreamWs, Error> {
let url = OPENAI_RESPONSES_WS_CONFIG.complete_websocket_url(api_base, model);
let mut request = url
.as_str()
.into_client_request()
.map_err(|error| Error::Network(error.to_string()))?;
request.headers_mut().insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {api_key}"))
.map_err(|error| Error::Auth(error.to_string()))?,
);
let result = tokio::time::timeout(
Duration::from_secs(DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS),
connect_async(request),
)
.await
.map_err(|_| Error::Network("Responses WebSocket connection timed out".to_string()))?;
result
.map(|(socket, _)| socket)
.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()),
})
}
pub struct ResponsesWebSocketStreaming;
impl ResponsesWebSocketStreaming {
pub async fn bidirectional_forward<In, Out>(
model: &str,
upstream_tx: UpstreamTx,
upstream_rx: UpstreamRx,
idle_timeout: Option<Duration>,
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
splice(
model,
upstream_tx,
upstream_rx,
idle_timeout,
observe,
client_in,
client_out,
)
.await
}
}
pub(crate) async fn splice<In, Out>(
model: &str,
mut upstream_tx: UpstreamTx,
mut upstream_rx: UpstreamRx,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
mut client_in: In,
mut client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
let idle =
idle_timeout.unwrap_or_else(|| Duration::from_secs(DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS));
loop {
tokio::select! {
event = client_in.next() => {
let Some(event) = event else { break };
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_request(&event, model)?
.events
{
let payload = serde_json::to_string(&outbound)
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
upstream_tx.send(Message::Text(payload))
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
}
message = upstream_rx.next() => {
let Some(message) = message else { break };
match message.map_err(|error| Error::Network(error.to_string()))? {
Message::Text(text) => {
let event = serde_json::from_str::<ResponsesWsEvent>(&text)
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
observe(&event);
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_response(&event, model)?
.events
{
client_out.send(outbound)
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
}
Message::Close(_) => break,
_ => {}
}
}
_ = tokio::time::sleep(idle) => break,
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub async fn async_responses_websocket<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
let key = resolve_api_key(api_key)?;
let upstream = dial_upstream(model, &key, api_base).await?;
let (mut upstream_tx, upstream_rx) = upstream.split();
if let Some(first_frame) = first_frame {
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_request(&first_frame, model)?
.events
{
let payload = serde_json::to_string(&outbound)
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
upstream_tx
.send(Message::Text(payload))
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
}
ResponsesWebSocketStreaming::bidirectional_forward(
model,
upstream_tx,
upstream_rx,
idle_timeout,
&mut observe,
client_in,
client_out,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn responses_ws<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Option<Duration>,
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
async_responses_websocket(
model,
api_key,
api_base,
first_frame,
idle_timeout,
observe,
client_in,
client_out,
)
.await
}
#[cfg(test)]
mod tests {
use super::*;
use futures_channel::mpsc;
use futures_util::{SinkExt, StreamExt};
use litellm_core::responses::types::ResponsesWsEventType;
use serde_json::json;
use tokio::io::AsyncWriteExt;
use tokio::net::TcpListener;
use tokio_tungstenite::accept_async;
async fn websocket_base() -> (String, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("local address");
let task = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let mut socket = accept_async(stream).await.expect("websocket handshake");
while let Some(Ok(Message::Text(text))) = socket.next().await {
let request: serde_json::Value = serde_json::from_str(&text).expect("request json");
let model = request
.get("model")
.and_then(serde_json::Value::as_str)
.or_else(|| {
request
.get("response")
.and_then(serde_json::Value::as_object)
.and_then(|response| {
response.get("model").and_then(serde_json::Value::as_str)
})
})
.expect("enforced model");
socket
.send(Message::Text(
json!({
"type": "response.created",
"response": {
"id": format!("resp-{model}"),
"model": model,
"extra": "preserved"
}
})
.to_string(),
))
.await
.expect("created event");
socket
.send(Message::Text(
json!({
"type": "response.completed",
"response": {
"id": format!("resp-{model}"),
"model": model,
"usage": {
"input_tokens": 1,
"output_tokens": 2,
"total_tokens": 3
}
}
})
.to_string(),
))
.await
.expect("completed event");
}
});
(format!("http://{address}"), task)
}
fn event(value: serde_json::Value) -> ResponsesWsEvent {
serde_json::from_value(value).expect("event")
}
#[test]
fn explicit_nonblank_key_wins() {
assert_eq!(
resolve_api_key(Some(" explicit ")).expect("key"),
"explicit"
);
}
#[test]
fn blank_key_is_not_accepted_without_environment_key() {
if std::env::var(OPENAI_API_KEY_ENV).is_err() {
assert!(resolve_api_key(Some(" ")).is_err());
}
}
#[tokio::test]
async fn forwards_events_sequentially_and_enforces_model() {
let (api_base, server) = websocket_base().await;
let (client_tx, client_rx) = mpsc::unbounded();
let (output_tx, mut output_rx) = mpsc::unbounded();
let (observed_tx, observed_rx) = mpsc::unbounded();
client_tx
.unbounded_send(event(json!({
"type": "response.create",
"model": "wrong"
})))
.expect("first request");
client_tx
.unbounded_send(event(json!({
"type": "response.create",
"response": {"model": "also-wrong"}
})))
.expect("second request");
let task = tokio::spawn(async move {
responses_ws(
"authorized-model",
Some("test-key"),
Some(&api_base),
None,
Some(Duration::from_secs(1)),
move |event| {
observed_tx
.unbounded_send(event.clone())
.expect("observe event");
},
client_rx,
output_tx,
)
.await
});
let first = output_rx.next().await.expect("first output");
let second = output_rx.next().await.expect("second output");
let third = output_rx.next().await.expect("third output");
let fourth = output_rx.next().await.expect("fourth output");
drop(client_tx);
task.await.expect("splice task").expect("successful splice");
server.await.expect("server task");
assert_eq!(first.event_type, ResponsesWsEventType::ResponseCreated);
assert_eq!(first.model(), Some("authorized-model"));
assert_eq!(first.data["response"]["extra"], "preserved");
assert_eq!(second.event_type, ResponsesWsEventType::ResponseCompleted);
assert_eq!(third.event_type, ResponsesWsEventType::ResponseCreated);
assert_eq!(fourth.event_type, ResponsesWsEventType::ResponseCompleted);
let observed: Vec<_> = observed_rx.collect().await;
assert_eq!(observed.len(), 4);
assert!(
observed
.iter()
.all(|event| event.event_type != ResponsesWsEventType::ResponseCreate)
);
}
#[tokio::test]
async fn idle_timeout_ends_without_upstream_events() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let _socket = accept_async(stream).await.expect("handshake");
tokio::time::sleep(Duration::from_secs(1)).await;
});
let (_client_tx, client_rx) = mpsc::unbounded::<ResponsesWsEvent>();
let (output_tx, mut output_rx) = mpsc::unbounded();
let result = responses_ws(
"model",
Some("key"),
Some(&format!("http://{address}")),
None,
Some(Duration::from_millis(20)),
|_| {},
client_rx,
output_tx,
)
.await;
assert!(result.is_ok());
assert!(output_rx.next().await.is_none());
server.abort();
}
#[tokio::test]
async fn dial_http_status_is_preserved() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("accept");
stream
.write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\n\r\n")
.await
.expect("response");
});
let (_client_tx, client_rx) = mpsc::unbounded::<ResponsesWsEvent>();
let (output_tx, _output_rx) = mpsc::unbounded();
let error = responses_ws(
"model",
Some("key"),
Some(&format!("http://{address}")),
None,
Some(Duration::from_millis(20)),
|_| {},
client_rx,
output_tx,
)
.await
.expect_err("status error");
assert!(matches!(error, Error::Http { status: 401, .. }));
server.await.expect("server task");
}
#[tokio::test]
async fn dial_http_500_status_is_preserved() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("accept");
stream
.write_all(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n")
.await
.expect("response");
});
let (_client_tx, client_rx) = mpsc::unbounded::<ResponsesWsEvent>();
let (output_tx, _output_rx) = mpsc::unbounded();
let error = responses_ws(
"model",
Some("key"),
Some(&format!("http://{address}")),
None,
Some(Duration::from_millis(20)),
|_| {},
client_rx,
output_tx,
)
.await
.expect_err("status error");
assert!(matches!(error, Error::Http { status: 500, .. }));
server.await.expect("server task");
}
}

View file

@ -13,7 +13,6 @@
//! binary turns on.
pub mod audio_transcription;
mod client;
pub mod io;
pub mod ocr;

View file

@ -1,99 +1,16 @@
use litellm_core::call_lifecycle::CallLifecycleContext;
use litellm_core::error::Error;
use litellm_core::ocr::transformation::OcrResponseHandling;
use litellm_core::provider_callbacks::ProviderAttemptObserver;
use litellm_core::provider_callbacks::handler::{
ProviderAttemptContext, ProviderRequest, send_provider_request,
};
use litellm_core::Error;
use litellm_core::ocr::observers::OcrObserver;
use litellm_core::ocr::{PreparedOcrRequest, execute_ocr_provider_call as core_execute};
use serde_json::Value;
use super::common_utils::poll_document_intelligence;
use super::hooks::OcrLifecycleHooks;
use super::types::PreparedOcrRequest;
use crate::client::http_client;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) async fn execute_ocr_provider_call<Observer>(
pub(crate) async fn execute_ocr_provider_call(
request: PreparedOcrRequest,
context: &CallLifecycleContext,
hooks: &OcrLifecycleHooks,
observer: &mut Observer,
) -> Result<Value, Error>
where
Observer: ProviderAttemptObserver,
Observer::Error: std::fmt::Display,
{
observer: &mut impl OcrObserver,
) -> Result<Value, Error> {
let request = hooks.prepare_provider_request(request).await?;
let provider_request = ProviderRequest {
provider: request.custom_llm_provider.clone(),
model: request.model.clone(),
body: serde_json::from_value(request.body).map_err(|error| {
Error::InvalidRequest(format!("OCR provider request must be an object: {error}"))
})?,
api_base: request.url.clone(),
headers: request.upstream_headers.iter().cloned().collect(),
};
let mut request_builder = http_client().post(&request.url);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let response = send_provider_request(
request_builder,
provider_request,
ProviderAttemptContext {
call_id: context.litellm_call_id.clone(),
trace_id: None,
attempt: 1,
},
observer,
)
.await?;
let status = response.status;
if request.config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll
&& status.as_u16() == 202
{
let operation_url = response
.headers
.get("operation-location")
.and_then(|value| value.to_str().ok())
.map(str::to_string)
.ok_or_else(|| {
Error::InvalidResponse(
"Azure Document Intelligence returned 202 but no Operation-Location header found"
.to_string(),
)
})?;
let response_json = poll_document_intelligence(
&operation_url,
&request.url,
&request.upstream_headers,
request.timeout,
)
.await?;
return Ok(request
.config
.transform_ocr_response_with_params(
&request.model,
response_json,
&request.optional_params,
)?
.into_json());
}
let response_json: Value = serde_json::from_str(&response.body)
.map_err(|err| Error::InvalidResponse(format!("invalid OCR response JSON: {err}")))?;
Ok(request
.config
.transform_ocr_response_with_params(
&request.model,
response_json,
&request.optional_params,
)?
.into_json())
core_execute(request, observer).await
}

View file

@ -1,28 +1,24 @@
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use litellm_core::error::Error;
use litellm_core::providers::reducto::ocr::transformation::{
build_upload_request, extract_document_source, extract_upload_file_id,
};
use litellm_core::request_context::RequestAttribution;
use litellm_core::ocr::{PreparedOcrRequest, ProviderOcrRequest, prepare_ocr_provider_call};
use serde_json::{Map, Value, json};
use std::future::Future;
use std::pin::Pin;
use super::common_utils::{convert_document_url_to_data_uri, string_headers, truncate_error_body};
use super::types::{PreparedOcrRequest, ProviderOcrRequest};
use crate::client::http_client;
use crate::integrations::custom_guardrail::{
CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest,
};
use crate::integrations::custom_logger::{
CallType, CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
use crate::integrations::types::{StandardLoggingMetadata, StandardLoggingPayload};
use crate::integrations::types::{
RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload,
};
pub(crate) struct OcrLifecycleHooks {
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestAttribution,
request_metadata: RequestMetadata,
}
type OcrFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
@ -32,7 +28,7 @@ impl OcrLifecycleHooks {
pub(crate) fn new(
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestAttribution,
request_metadata: RequestMetadata,
) -> Self {
Self {
logger_runner,
@ -77,83 +73,24 @@ impl OcrLifecycleHooks {
&self,
request: PreparedOcrRequest,
) -> Result<ProviderOcrRequest, Error> {
let config = request.config?;
let env_lookup = |key: &str| std::env::var(key).ok();
let upstream_headers = config.validate_environment(
string_headers(request.extra_headers)?,
request.api_key.as_deref(),
&env_lookup,
)?;
let url_params = request
.optional_params
.clone()
.into_iter()
.chain(request.vertex.into_map())
.collect();
let url = config.complete_url(
request.api_base.as_deref(),
&request.model,
&url_params,
&env_lookup,
)?;
let model = request.model.clone();
let custom_llm_provider = request.custom_llm_provider.clone();
let is_reducto = custom_llm_provider == "reducto";
let document = if is_reducto {
let guarded_document = self
.run_during_call_guardrails(&model, &custom_llm_provider, &url, request.document)
.await?;
upload_reducto_document(
&guarded_document,
request.api_base.as_deref(),
request.timeout,
&upstream_headers,
)
.await?
} else if config.requires_data_uri_document() {
convert_document_url_to_data_uri(request.document).await?
} else {
request.document
};
let optional_params = request.optional_params;
let body = config
.transform_ocr_request(&request.model, document, optional_params.clone())?
.data;
let body = if is_reducto {
body
} else {
self.run_during_call_guardrails(&model, &custom_llm_provider, &url, body)
.await?
};
Ok(ProviderOcrRequest {
model,
custom_llm_provider,
config,
url,
body,
optional_params,
upstream_headers,
timeout: request.timeout,
})
let provider_request = prepare_ocr_provider_call(request).await?;
self.run_during_call_guardrails(provider_request).await
}
async fn run_during_call_guardrails(
&self,
model: &str,
custom_llm_provider: &str,
url: &str,
body: Value,
) -> Result<Value, Error> {
request: ProviderOcrRequest,
) -> Result<ProviderOcrRequest, Error> {
if self.guardrail_runner.is_empty() {
return Ok(body);
return Ok(request);
}
let context = guardrail_context(&self.request_metadata);
let guardrail_request = GuardrailRequest::new(json!({
"model": model,
"custom_llm_provider": custom_llm_provider,
"url": url,
"body": body,
"model": request.model(),
"custom_llm_provider": request.custom_llm_provider(),
"url": request.url(),
"body": request.body(),
}));
let (guardrail_request, _) = self
.guardrail_runner
@ -161,6 +98,7 @@ impl OcrLifecycleHooks {
.await
.map_err(guardrail_error_to_core_error)?;
parse_ocr_during_call_guardrail_request(guardrail_request)
.map(|body| request.with_body(body))
}
fn standard_logging_payload(
@ -192,63 +130,6 @@ impl OcrLifecycleHooks {
}
}
async fn upload_reducto_document(
document: &Value,
api_base: Option<&str>,
timeout: Option<std::time::Duration>,
upstream_headers: &[(String, String)],
) -> Result<Value, Error> {
let source = extract_document_source(document)?;
let Some(authorization) = upstream_headers
.iter()
.find(|(name, _)| name.eq_ignore_ascii_case("authorization"))
.map(|(_, value)| value.as_str())
else {
return Err(Error::Auth(
"Reducto upload requires an Authorization header".to_string(),
));
};
let Some(upload) = build_upload_request(source, authorization, api_base) else {
return Ok(document.clone());
};
let part = reqwest::multipart::Part::bytes(upload.bytes)
.file_name(upload.file_name)
.mime_str(&upload.mime_type)
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
let form = reqwest::multipart::Form::new().part("file", part);
let mut request_builder = http_client().post(upload.url).multipart(form);
for (name, value) in upstream_headers {
if !name.eq_ignore_ascii_case("content-type")
&& !name.eq_ignore_ascii_case("content-length")
{
request_builder = request_builder.header(name, value);
}
}
if let Some(timeout) = timeout {
request_builder = request_builder.timeout(timeout);
}
let response = request_builder
.send()
.await
.map_err(|error| Error::Network(error.to_string()))?;
let status = response.status();
let body = response
.text()
.await
.map_err(|error| Error::Network(error.to_string()))?;
if !status.is_success() {
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&body),
});
}
let response_json: Value = serde_json::from_str(&body).map_err(|error| {
Error::InvalidResponse(format!("invalid Reducto upload response JSON: {error}"))
})?;
let file_id = extract_upload_file_id(&response_json)?;
Ok(json!({"type": "document_url", "document_url": file_id}))
}
impl CallLifecycleHooks<PreparedOcrRequest, PreparedOcrRequest, Value> for OcrLifecycleHooks {
type PreCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>;
type DuringCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>;
@ -329,7 +210,7 @@ impl CallLifecycleHooks<PreparedOcrRequest, PreparedOcrRequest, Value> for OcrLi
}
}
fn guardrail_context(metadata: &RequestAttribution) -> GuardrailContext {
fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext {
GuardrailContext {
call_type: CallType::Ocr,
selected_guardrails: Vec::new(),

View file

@ -6,7 +6,6 @@ use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use serde_json::Value;
mod common_utils;
mod handler;
mod hooks;
mod prepare;

View file

@ -1,181 +1,46 @@
use crate::integrations::types::RequestHooks;
use litellm_core::call_lifecycle::CallLifecycleContext;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use litellm_core::ocr::{OcrRequest as CoreOcrRequest, PreparedOcrRequest};
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use serde_json::{Map, Value};
use super::common_utils::ocr_provider_config;
use super::hooks::OcrLifecycleHooks;
use super::types::{OcrRequest, PreparedOcrRequest};
use super::types::OcrRequest;
use crate::integrations::custom_guardrail::CustomGuardrailRunner;
use crate::integrations::custom_logger::CustomLoggerRunner;
pub(crate) struct PreparedOcrCall {
pub(crate) context: CallLifecycleContext,
pub(crate) request: PreparedOcrRequest,
pub(crate) hooks: OcrLifecycleHooks,
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn prepare_ocr_call(
request: OcrRequest<'_>,
options: RequestOptions,
context: &LiteLlmRequestContext,
hooks: RequestHooks,
) -> PreparedOcrCall {
let call_id = context
.litellm_call_id
.clone()
.unwrap_or_else(new_ocr_call_id);
let provider_info =
get_custom_llm_provider(request.model, options.custom_llm_provider.as_deref()).unwrap_or(
CustomLlmProvider {
model: request.model,
custom_llm_provider: "mistral",
},
);
let model = provider_info.model.to_string();
let custom_llm_provider = provider_info.custom_llm_provider.to_string();
let config = ocr_provider_config(&custom_llm_provider, &model)
.ok_or_else(|| litellm_core::Error::InvalidProvider(custom_llm_provider.clone()))
.and_then(|config| {
validate_request_format(config, &request.optional_params, &custom_llm_provider)?;
Ok(config)
});
let optional_params = match &config {
Ok(config) => {
let supported = config.supported_ocr_params();
config.map_ocr_params(
&request
.optional_params
.iter()
.filter(|(name, _)| supported.contains(&name.as_str()))
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
)
}
Err(_) => request.optional_params,
};
pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall {
let OcrRequest {
model,
document,
api_key,
api_base,
custom_llm_provider,
extra_headers,
optional_params,
timeout,
callbacks,
guardrails,
request_metadata,
litellm_call_id,
} = request;
PreparedOcrCall {
context: CallLifecycleContext::new(
"ocr",
model.clone(),
custom_llm_provider.clone(),
call_id,
),
request: PreparedOcrRequest {
config,
request: litellm_core::ocr::prepare_ocr_call(CoreOcrRequest {
model,
document,
api_key,
api_base,
custom_llm_provider,
document: request.document,
vertex: options.vertex.unwrap_or_default(),
api_key: options.api_key,
api_base: options.api_base,
extra_headers: options.extra_headers,
extra_headers,
optional_params,
timeout: options.timeout,
},
timeout,
litellm_call_id,
}),
hooks: OcrLifecycleHooks::new(
CustomLoggerRunner::new(hooks.callbacks),
CustomGuardrailRunner::new(hooks.guardrails),
context.attribution.clone(),
CustomLoggerRunner::new(callbacks),
CustomGuardrailRunner::new(guardrails),
request_metadata,
),
}
}
fn validate_request_format(
config: &'static dyn litellm_core::ocr::transformation::OcrProviderConfig,
optional_params: &Map<String, Value>,
provider: &str,
) -> Result<(), litellm_core::Error> {
let Some(format) = optional_params.get("req_format") else {
return Ok(());
};
match format.as_str() {
Some("litellm") => Ok(()),
Some("native") if config.supported_ocr_params().contains(&"req_format") => Ok(()),
Some("native") => Err(litellm_core::Error::InvalidRequest(format!(
"`req_format=native` is not supported for provider {provider}"
))),
_ => Err(litellm_core::Error::InvalidRequest(format!(
"Invalid `req_format`: {format}. Expected `litellm` or `native`"
))),
}
}
fn new_ocr_call_id() -> String {
static COUNTER: AtomicU64 = AtomicU64::new(1);
let sequence = COUNTER.fetch_add(1, Ordering::Relaxed);
let timestamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_nanos())
.unwrap_or(0);
format!("ocr-{timestamp}-{sequence}")
}
#[cfg(test)]
mod tests {
use crate::integrations::types::RequestHooks;
use litellm_core::error::Error;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use serde_json::{Map, json};
use super::{OcrRequest, prepare_ocr_call};
fn base_ocr_request(model: &str) -> OcrRequest<'_> {
OcrRequest {
model,
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
}
}
fn request_with_format(format: &str) -> OcrRequest<'_> {
let mut request = base_ocr_request("mistral/mistral-ocr-latest");
request.optional_params = Map::from_iter([("req_format".to_string(), json!(format))]);
request
}
#[test]
fn native_format_rejected_for_provider_without_support_as_bad_request() {
let prepared = prepare_ocr_call(
request_with_format("native"),
RequestOptions::default(),
&LiteLlmRequestContext {
..Default::default()
},
RequestHooks {
..Default::default()
},
);
assert!(
matches!(prepared.request.config, Err(Error::InvalidRequest(message)) if message.contains("not supported for provider"))
);
}
#[test]
fn unknown_format_rejected_for_provider_without_support_as_bad_request() {
let prepared = prepare_ocr_call(
request_with_format("raw"),
RequestOptions::default(),
&LiteLlmRequestContext {
..Default::default()
},
RequestHooks {
..Default::default()
},
);
assert!(
matches!(prepared.request.config, Err(Error::InvalidRequest(message)) if message.contains("Invalid `req_format`"))
);
}
}

View file

@ -1,35 +1,23 @@
use std::sync::Arc;
use std::time::Duration;
use litellm_core::ocr::transformation::OcrProviderConfig;
use litellm_core::request_options::VertexOptions;
use serde_json::{Map, Value};
use crate::integrations::custom_guardrail::CustomGuardrail;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
pub struct OcrRequest<'a> {
pub model: &'a str,
pub document: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
}
pub(crate) struct PreparedOcrRequest {
pub(crate) config: Result<&'static dyn OcrProviderConfig, litellm_core::Error>,
pub(crate) model: String,
pub(crate) custom_llm_provider: String,
pub(crate) document: Value,
pub(crate) vertex: VertexOptions,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
pub(crate) optional_params: Map<String, Value>,
pub(crate) timeout: Option<Duration>,
}
pub(crate) struct ProviderOcrRequest {
pub(crate) model: String,
pub(crate) custom_llm_provider: String,
pub(crate) config: &'static dyn OcrProviderConfig,
pub(crate) url: String,
pub(crate) body: Value,
pub(crate) optional_params: Map<String, Value>,
pub(crate) upstream_headers: Vec<(String, String)>,
pub(crate) timeout: Option<Duration>,
pub timeout: Option<Duration>,
pub callbacks: Vec<Arc<dyn CustomLogger>>,
pub guardrails: Vec<Arc<dyn CustomGuardrail>>,
pub request_metadata: RequestMetadata,
pub litellm_call_id: Option<&'a str>,
}

View file

@ -1,928 +0,0 @@
use litellm_ai_gateway::integrations::types::RequestHooks;
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_context::RequestAttribution;
use litellm_core::request_options::RequestOptions;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use litellm_ai_gateway::integrations::custom_guardrail::{
CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook,
GuardrailFuture, GuardrailRequest,
};
use litellm_ai_gateway::integrations::custom_logger::{
CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails,
};
use litellm_ai_gateway::ocr::{OcrRequest, ocr, ocr_with_observer};
use litellm_core::error::Error;
use litellm_core::provider_callbacks::{
CallbackDecision, ProviderAttemptObserver, ProviderError, ProviderPostCall, ProviderPreCall,
};
use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
struct ProviderObserver {
events: Arc<Mutex<Vec<&'static str>>>,
raw_response: Option<String>,
rejected_callback: Option<&'static str>,
decision: Option<&'static str>,
}
impl ProviderAttemptObserver for ProviderObserver {
type Error = &'static str;
async fn pre_call(&mut self, input: &ProviderPreCall) -> Result<CallbackDecision, Self::Error> {
assert_eq!(input.model, "mistral-ocr-4-1");
assert_eq!(input.call_id, "observer-test");
assert_eq!(
input.request["document"]["document_url"],
"https://example.com/document.pdf"
);
assert!(input.api_base.ends_with("/v1/ocr"));
assert!(
input
.headers
.values()
.any(|value| value == "Bearer test-key")
);
self.events.lock().unwrap().push("pre");
match (self.rejected_callback, self.decision) {
(Some("pre"), _) => Err("observer failure"),
(_, Some("replace_pre")) => Ok(CallbackDecision::Replace {
payload: Value::Object(
input
.request
.iter()
.map(|(key, value)| (key.clone(), value.clone()))
.chain(std::iter::once((
"callback_replaced".to_string(),
json!(true),
)))
.collect(),
),
}),
(_, Some("reject_pre")) => Ok(CallbackDecision::Reject {
message: "callback rejected request".to_string(),
status_code: Some(400),
}),
_ => Ok(CallbackDecision::Unchanged),
}
}
async fn post_call(
&mut self,
input: &ProviderPostCall,
) -> Result<CallbackDecision, Self::Error> {
self.events.lock().unwrap().push("post");
self.raw_response = input.response.as_str().map(str::to_string);
match (self.rejected_callback, self.decision) {
(Some("post"), _) => Err("observer failure"),
(_, Some("replace_post")) => Ok(CallbackDecision::Replace {
payload: json!({"pages":[{"index":0,"markdown":"masked"}]}),
}),
_ => Ok(CallbackDecision::Unchanged),
}
}
async fn error(&mut self, input: &ProviderError) -> Result<(), Self::Error> {
assert!(input.committed);
assert!(!input.message.is_empty());
self.events.lock().unwrap().push("error");
if self.rejected_callback == Some("error") {
Err("observer failure")
} else {
Ok(())
}
}
}
fn observer_request() -> OcrRequest<'static> {
OcrRequest {
model: "mistral/mistral-ocr-4-1",
document: json!({"type":"document_url","document_url":"https://example.com/document.pdf"}),
optional_params: Map::new(),
}
}
fn observer_options(api_base: &str) -> RequestOptions {
RequestOptions {
api_key: Some("test-key".into()),
api_base: Some(api_base.into()),
custom_llm_provider: Some("mistral".into()),
timeout: Some(Duration::from_secs(2)),
..Default::default()
}
}
fn observer_context() -> LiteLlmRequestContext {
LiteLlmRequestContext {
litellm_call_id: Some("observer-test".into()),
..Default::default()
}
}
async fn observer_case(status: u16, body: &'static str, decision: Option<&'static str>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
let events = Arc::new(Mutex::new(Vec::new()));
let provider_events = Arc::clone(&events);
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let request = read_http_request(&mut socket).await;
assert!(request.starts_with("POST /v1/ocr "));
assert_eq!(
request.contains(r#""callback_replaced":true"#),
decision == Some("replace_pre")
);
provider_events.lock().unwrap().push("http");
let response = format!(
"HTTP/1.1 {status} Test\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
body.len()
);
socket.write_all(response.as_bytes()).await.unwrap();
});
let mut observer = ProviderObserver {
events: Arc::clone(&events),
raw_response: None,
rejected_callback: None,
decision,
};
let result = ocr_with_observer(
observer_request(),
&observer_options(&url),
&observer_context(),
RequestHooks::default(),
&mut observer,
)
.await;
tokio::time::timeout(Duration::from_secs(2), server)
.await
.unwrap()
.unwrap();
if status != 200 {
assert!(matches!(result, Err(Error::Http { status: actual, .. }) if actual == status));
assert_eq!(*events.lock().unwrap(), ["pre", "http", "error"]);
assert_eq!(observer.raw_response, None);
} else {
assert_eq!(*events.lock().unwrap(), ["pre", "http", "post"]);
assert_eq!(observer.raw_response.as_deref(), Some(body));
if body == "invalid-json" {
assert!(matches!(result, Err(Error::InvalidResponse(_))));
} else {
assert_eq!(
result.unwrap()["pages"][0]["markdown"],
if decision == Some("replace_post") {
"masked"
} else {
"ok"
}
);
}
}
}
#[tokio::test]
async fn provider_observers_surround_http_and_can_replace_request_or_response() {
observer_case(200, r#"{"pages":[{"index":0,"markdown":"ok"}]}"#, None).await;
observer_case(200, "invalid-json", None).await;
observer_case(401, r#"{"error":"rejected"}"#, None).await;
observer_case(
200,
r#"{"pages":[{"index":0,"markdown":"ok"}]}"#,
Some("replace_pre"),
)
.await;
observer_case(
200,
r#"{"pages":[{"index":0,"markdown":"ok"}]}"#,
Some("replace_post"),
)
.await;
}
#[tokio::test]
async fn provider_callback_rejection_stops_before_http() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
let events = Arc::new(Mutex::new(Vec::new()));
let mut observer = ProviderObserver {
events: Arc::clone(&events),
raw_response: None,
rejected_callback: None,
decision: Some("reject_pre"),
};
let result = ocr_with_observer(
observer_request(),
&observer_options(&url),
&observer_context(),
RequestHooks::default(),
&mut observer,
)
.await;
assert!(
matches!(result, Err(Error::InvalidRequest(message)) if message == "callback rejected request")
);
assert_eq!(*events.lock().unwrap(), ["pre"]);
assert!(
tokio::time::timeout(Duration::from_millis(50), listener.accept())
.await
.is_err()
);
}
#[tokio::test]
async fn invalid_ocr_preparation_does_not_call_observers_or_provider() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
let events = Arc::new(Mutex::new(Vec::new()));
let mut observer = ProviderObserver {
events: Arc::clone(&events),
raw_response: None,
rejected_callback: None,
decision: None,
};
let request = OcrRequest {
document: json!(42),
..observer_request()
};
assert!(
ocr_with_observer(
request,
&observer_options(&url),
&observer_context(),
RequestHooks::default(),
&mut observer
)
.await
.is_err()
);
assert!(events.lock().unwrap().is_empty());
assert!(
tokio::time::timeout(Duration::from_millis(50), listener.accept())
.await
.is_err()
);
}
async fn read_http_headers(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
if request.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
String::from_utf8(request).expect("request is utf8")
}
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
let header_end = loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break request.len();
}
request.extend_from_slice(&buffer[..n]);
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") {
break position + 4;
}
};
let headers = String::from_utf8_lossy(&request[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
while request.len().saturating_sub(header_end) < content_length {
let n = socket.read(&mut buffer).await.expect("reads body");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
}
String::from_utf8(request).expect("request is utf8")
}
#[derive(Clone, Debug, PartialEq)]
struct RecordedLogEvent {
hook: &'static str,
model: String,
call_type: String,
user_id: Option<String>,
response_object: Option<String>,
error_kind: Option<String>,
}
#[derive(Default)]
struct RecordingOcrLogger {
events: Mutex<Vec<RecordedLogEvent>>,
}
impl RecordingOcrLogger {
fn events(&self) -> Vec<RecordedLogEvent> {
self.events.lock().unwrap().clone()
}
}
impl CustomLogger for RecordingOcrLogger {
fn async_log_success_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
response_obj: &'a CallbackValue,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push(RecordedLogEvent {
hook: "async_log_success_event",
model: model_call_details.model.clone(),
call_type: model_call_details.call_type.to_string(),
user_id: model_call_details.metadata.user_api_key_user_id.clone(),
response_object: Some(response_obj.object.clone()),
error_kind: None,
});
Ok(())
})
}
fn async_log_failure_event<'a>(
&'a self,
model_call_details: &'a ModelCallDetails,
response_obj: Option<&'a CallbackValue>,
_timing: CallbackTiming,
) -> LogFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push(RecordedLogEvent {
hook: "async_log_failure_event",
model: model_call_details.model.clone(),
call_type: model_call_details.call_type.to_string(),
user_id: model_call_details.metadata.user_api_key_user_id.clone(),
response_object: response_obj.map(|value| value.object.clone()),
error_kind: model_call_details
.failure_error
.as_ref()
.map(|error| error.kind.clone()),
});
Ok(())
})
}
}
struct RecordingOcrGuardrail {
hooks: Vec<GuardrailEventHook>,
events: Mutex<Vec<&'static str>>,
block_pre_call: bool,
block_during_call: bool,
}
impl RecordingOcrGuardrail {
fn new(hooks: Vec<GuardrailEventHook>) -> Self {
Self {
hooks,
events: Mutex::new(Vec::new()),
block_pre_call: false,
block_during_call: false,
}
}
fn blocking_pre_call() -> Self {
Self {
hooks: vec![GuardrailEventHook::PreCall],
events: Mutex::new(Vec::new()),
block_pre_call: true,
block_during_call: false,
}
}
fn blocking_during_call() -> Self {
Self {
hooks: vec![GuardrailEventHook::DuringCall],
events: Mutex::new(Vec::new()),
block_pre_call: false,
block_during_call: true,
}
}
fn events(&self) -> Vec<&'static str> {
self.events.lock().unwrap().clone()
}
}
impl CustomGuardrail for RecordingOcrGuardrail {
fn guardrail_name(&self) -> &str {
"recording-ocr-guardrail"
}
fn supported_event_hooks(&self) -> &[GuardrailEventHook] {
&self.hooks
}
fn async_pre_call_hook<'a>(
&'a self,
_context: &'a GuardrailContext,
mut request: GuardrailRequest,
) -> GuardrailFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push("async_pre_call_hook");
if self.block_pre_call {
return Ok(GuardrailDecision::Block(GuardrailError::blocked(
"blocked before provider",
)));
}
request.data["document"]["guarded_pre"] = json!(true);
Ok(GuardrailDecision::Mask(request))
})
}
fn async_moderation_hook<'a>(
&'a self,
_context: &'a GuardrailContext,
mut request: GuardrailRequest,
) -> GuardrailFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push("async_moderation_hook");
if self.block_during_call {
return Ok(GuardrailDecision::Block(GuardrailError::blocked(
"blocked before provider",
)));
}
request.data["body"]["guarded_during"] = json!(true);
Ok(GuardrailDecision::Mask(request))
})
}
}
fn base_ocr_request(model: &str) -> (OcrRequest<'_>, RequestOptions) {
(
OcrRequest {
model,
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
},
RequestOptions {
api_key: Some("sk-test".to_string()),
..Default::default()
},
)
}
#[tokio::test]
async fn reducto_during_call_guardrail_blocks_before_upload() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let address = listener.local_addr().expect("listener has local address");
let api_base = format!("http://{address}");
let guardrail = Arc::new(RecordingOcrGuardrail::blocking_during_call());
let (mut request, mut options) = base_ocr_request("reducto/parse-v3");
options.api_base = Some(&api_base).map(|value| value.to_string());
request.document = json!({
"type": "document_url",
"document_url": "data:application/pdf;base64,JVBERi0xLjQ="
});
let hooks = RequestHooks {
guardrails: vec![guardrail.clone()],
..Default::default()
};
let error = ocr(
request,
&options,
&LiteLlmRequestContext {
..Default::default()
},
hooks,
)
.await
.expect_err("guardrail blocks upload");
assert!(matches!(error, Error::InvalidRequest(_)));
assert_eq!(guardrail.events(), vec!["async_moderation_hook"]);
let accepted = tokio::time::timeout(Duration::from_millis(100), listener.accept()).await;
assert!(accepted.is_err(), "upload socket should not be touched");
}
#[tokio::test]
async fn reducto_upload_error_body_is_truncated() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let address = listener.local_addr().expect("listener has local address");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts upload request");
let _request = read_http_request(&mut socket).await;
let body = "x".repeat(300);
let response = format!(
"HTTP/1.1 500 Internal Server Error\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
body.len(),
body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes upload response");
});
let api_base = format!("http://{address}");
let (mut request, mut options) = base_ocr_request("reducto/parse-v3");
options.api_base = Some(&api_base).map(|value| value.to_string());
request.document = json!({
"type": "document_url",
"document_url": "data:application/pdf;base64,JVBERi0xLjQ="
});
let error = ocr(
request,
&options,
&LiteLlmRequestContext {
..Default::default()
},
RequestHooks::default(),
)
.await
.expect_err("upload should fail");
assert!(
matches!(error, Error::Http { status: 500, body } if body.chars().count() < 300 && body.ends_with("... (truncated)"))
);
server.await.expect("server task completes");
}
#[tokio::test]
async fn ocr_lifecycle_runs_pre_during_and_success_hooks() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let request = read_http_request(&mut socket).await;
let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
request
});
let logger = Arc::new(RecordingOcrLogger::default());
let guardrail = Arc::new(RecordingOcrGuardrail::new(vec![
GuardrailEventHook::PreCall,
GuardrailEventHook::DuringCall,
]));
let response = ocr(
OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
},
&RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("mistral")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
&LiteLlmRequestContext {
attribution: RequestAttribution {
user_api_key_user_id: Some("user-1".to_string()),
..Default::default()
},
litellm_call_id: (Some("ocr-call-1")).map(|value| value.to_string()),
..Default::default()
},
RequestHooks {
callbacks: vec![logger.clone()],
guardrails: vec![guardrail.clone()],
},
)
.await
.expect("ocr request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
assert_eq!(
guardrail.events(),
vec!["async_pre_call_hook", "async_moderation_hook"]
);
assert_eq!(
logger.events(),
vec![RecordedLogEvent {
hook: "async_log_success_event",
model: "mistral-ocr-latest".to_string(),
call_type: "ocr".to_string(),
user_id: Some("user-1".to_string()),
response_object: Some("ocr".to_string()),
error_kind: None,
}]
);
let request = server.await.expect("server task completes");
assert!(request.contains(r#""guarded_pre":true"#), "{request}");
assert!(request.contains(r#""guarded_during":true"#), "{request}");
}
#[tokio::test]
async fn ocr_lifecycle_runs_failure_hook_on_provider_error() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let _request = read_http_request(&mut socket).await;
let response_body = "provider failed";
let response = format!(
"HTTP/1.1 500 Internal Server Error\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
});
let logger = Arc::new(RecordingOcrLogger::default());
let err = ocr(
OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
},
&RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("mistral")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
&LiteLlmRequestContext {
attribution: RequestAttribution::default(),
litellm_call_id: (Some("ocr-call-2")).map(|value| value.to_string()),
..Default::default()
},
RequestHooks {
callbacks: vec![logger.clone()],
guardrails: Vec::new(),
},
)
.await
.expect_err("provider error propagates");
assert!(matches!(err, Error::Http { status: 500, .. }));
server.await.expect("server task completes");
assert_eq!(
logger.events(),
vec![RecordedLogEvent {
hook: "async_log_failure_event",
model: "mistral-ocr-latest".to_string(),
call_type: "ocr".to_string(),
user_id: None,
response_object: Some("error".to_string()),
error_kind: Some("HttpError".to_string()),
}]
);
}
#[tokio::test]
async fn ocr_lifecycle_pre_call_block_skips_provider_socket() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let logger = Arc::new(RecordingOcrLogger::default());
let guardrail = Arc::new(RecordingOcrGuardrail::blocking_pre_call());
let err = ocr(
OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
},
&RequestOptions {
api_key: (Some("sk-test")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("mistral")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_millis(100)),
..Default::default()
},
&LiteLlmRequestContext {
attribution: RequestAttribution::default(),
litellm_call_id: (Some("ocr-call-3")).map(|value| value.to_string()),
..Default::default()
},
RequestHooks {
callbacks: vec![logger.clone()],
guardrails: vec![guardrail.clone()],
},
)
.await
.expect_err("guardrail blocks request");
assert!(matches!(err, Error::InvalidRequest(_)));
assert_eq!(guardrail.events(), vec!["async_pre_call_hook"]);
assert_eq!(
logger.events(),
vec![RecordedLogEvent {
hook: "async_log_failure_event",
model: "mistral-ocr-latest".to_string(),
call_type: "ocr".to_string(),
user_id: None,
response_object: Some("error".to_string()),
error_kind: Some("InvalidRequest".to_string()),
}]
);
let accepted = tokio::time::timeout(Duration::from_millis(100), listener.accept()).await;
assert!(accepted.is_err(), "provider socket should not be touched");
}
#[tokio::test]
async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let request = read_http_headers(&mut socket).await;
let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
request
});
let mut headers = Map::new();
headers.insert(
"Authorization".to_string(),
Value::String("Bearer sk-from-python".to_string()),
);
headers.insert(
"x-trace-id".to_string(),
Value::String("trace-1".to_string()),
);
let response = ocr(
OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
},
&RequestOptions {
api_key: (Some("sk-for-rust-fallback")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("mistral")).map(|value| value.to_string()),
extra_headers: Some(headers),
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
&LiteLlmRequestContext {
attribution: RequestAttribution::default(),
litellm_call_id: None,
..Default::default()
},
RequestHooks {
callbacks: Vec::new(),
guardrails: Vec::new(),
},
)
.await
.expect("ocr request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
let request = server.await.expect("server task completes");
let authorization_count = request
.lines()
.filter(|line| line.to_ascii_lowercase().starts_with("authorization:"))
.count();
assert_eq!(authorization_count, 1, "{request}");
assert!(
request.contains("authorization: Bearer sk-from-python")
|| request.contains("Authorization: Bearer sk-from-python"),
"{request}"
);
}
#[tokio::test]
async fn document_intelligence_poll_uses_resolved_subscription_key() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let operation_url = format!("http://{addr}/operations/1");
let server = tokio::spawn(async move {
let (mut post_socket, _) = listener.accept().await.expect("accepts post request");
let post_request = read_http_headers(&mut post_socket).await;
let post_response = format!(
"HTTP/1.1 202 Accepted\r\noperation-location: {operation_url}\r\ncontent-length: 0\r\nconnection: close\r\n\r\n"
);
post_socket
.write_all(post_response.as_bytes())
.await
.expect("writes post response");
let (mut poll_socket, _) = listener.accept().await.expect("accepts poll request");
let poll_request = read_http_headers(&mut poll_socket).await;
let response_body = r#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"ok"}]}]}}"#;
let poll_response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
poll_socket
.write_all(poll_response.as_bytes())
.await
.expect("writes poll response");
(post_request, poll_request)
});
let response = ocr(
OcrRequest {
model: "doc-intelligence/prebuilt-read",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
optional_params: Map::new(),
},
&RequestOptions {
api_key: (Some("di-key")).map(|value| value.to_string()),
api_base: (Some(&format!("http://{addr}"))).map(|value| value.to_string()),
custom_llm_provider: (Some("azure_ai")).map(|value| value.to_string()),
extra_headers: None,
timeout: Some(Duration::from_secs(5)),
..Default::default()
},
&LiteLlmRequestContext {
attribution: RequestAttribution::default(),
litellm_call_id: None,
..Default::default()
},
RequestHooks {
callbacks: Vec::new(),
guardrails: Vec::new(),
},
)
.await
.expect("document intelligence request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
let (post_request, poll_request) = server.await.expect("server task completes");
assert!(
post_request
.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: di-key"),
"{post_request}"
);
assert!(
poll_request
.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: di-key"),
"{poll_request}"
);
}

View file

@ -7,11 +7,14 @@ repository.workspace = true
[dependencies]
base64.workspace = true
futures-util.workspace = true
rand.workspace = true
reqwest.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
tokio = { workspace = true, features = ["sync"] }
tokio-tungstenite.workspace = true
tracing.workspace = true
tracing-subscriber = { workspace = true, optional = true }
sha2.workspace = true
@ -35,6 +38,7 @@ bedrock-auth = [
observability = ["dep:tracing-subscriber"]
[dev-dependencies]
futures-channel = "0.3"
rstest.workspace = true
tokio = { workspace = true, features = ["macros", "rt-multi-thread"] }
tracing-subscriber.workspace = true

View file

@ -32,6 +32,18 @@ pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600;
/// Full-request timeout ceiling for OCR provider calls, in seconds. The
/// per-request timeout from the caller still overrides this on the request
/// builder.
pub(crate) const OCR_TIMEOUT_SECS: u64 = 600;
/// Connect timeout for the Responses WebSocket upstream dial, in seconds.
pub(crate) const RESPONSES_WS_CONNECT_TIMEOUT_SECS: u64 = 10;
/// Idle timeout that ends a Responses WebSocket splice when neither side
/// produces an event, in seconds.
pub(crate) const RESPONSES_WS_IDLE_TIMEOUT_SECS: u64 = 300;
/// `object` field every non-streaming chat completion response carries.
pub const CHAT_COMPLETION_OBJECT: &str = "chat.completion";

View file

@ -0,0 +1,14 @@
use std::sync::OnceLock;
use std::time::Duration;
use crate::constants::OCR_TIMEOUT_SECS;
pub(super) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(OCR_TIMEOUT_SECS))
.build()
.unwrap_or_else(|_| reqwest::Client::new())
})
}

View file

@ -3,78 +3,16 @@ use std::time::{Duration, Instant};
use base64::Engine;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use litellm_core::error::Error;
use litellm_core::ocr::transformation::OcrProviderConfig;
use reqwest::Url;
use serde_json::{Map, Value};
use serde_json::Value;
use litellm_core::providers::azure_ai::ocr::transformation::{
AZURE_AI_OCR_CONFIG, AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG,
};
use litellm_core::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG;
use litellm_core::providers::reducto::ocr::transformation as reducto;
use litellm_core::providers::vertex_ai::ocr::transformation as vertex_ai;
use litellm_core::providers::vertex_ai::ocr::transformation::{
VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG,
};
use super::client::http_client;
use crate::Error;
use crate::http_utils::truncate_error_body;
use crate::client::http_client;
const ERROR_BODY_MAX_CHARS: usize = 256;
const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120;
const DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB: f64 = 50.0;
const MAX_SAFE_FETCH_REDIRECTS: usize = 10;
pub(super) fn truncate_error_body(body: &str) -> String {
if body.chars().count() <= ERROR_BODY_MAX_CHARS {
return body.to_string();
}
let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect();
format!("{truncated}... (truncated)")
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) fn ocr_provider_config(
provider: &str,
model: &str,
) -> Option<&'static dyn OcrProviderConfig> {
match provider {
"mistral" => Some(&MISTRAL_OCR_CONFIG),
"reducto" => reducto::config_for_model(model),
"azure_ai" if is_azure_document_intelligence_model(model) => {
Some(&AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG)
}
"azure_ai" => Some(&AZURE_AI_OCR_CONFIG),
"vertex_ai" if vertex_ai::is_deepseek_model(model) => Some(&VERTEX_AI_DEEPSEEK_OCR_CONFIG),
"vertex_ai" => Some(&VERTEX_AI_OCR_CONFIG),
_ => None,
}
}
fn is_azure_document_intelligence_model(model: &str) -> bool {
let model = model.to_ascii_lowercase();
model.contains("doc-intelligence") || model.contains("documentintelligence")
}
pub(super) fn string_headers(
extra_headers: Option<Map<String, Value>>,
) -> Result<Vec<(String, String)>, Error> {
extra_headers
.unwrap_or_default()
.into_iter()
.map(|(key, value)| {
value
.as_str()
.map(|value| (key.clone(), value.to_string()))
.ok_or_else(|| {
Error::InvalidRequest(format!(
"OCR extra_headers.{key} must be a string, got {}",
litellm_core::error::json_type_name(&value)
))
})
})
.collect()
}
const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120;
fn document_url_field(document: &Value) -> Result<Option<(&str, &str)>, Error> {
let Some(object) = document.as_object() else {
@ -336,7 +274,6 @@ fn operation_status(response_json: &Value) -> Result<&str, Error> {
}
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(super) async fn poll_document_intelligence(
operation_url: &str,
original_url: &str,
@ -395,10 +332,8 @@ pub(super) async fn poll_document_intelligence(
#[cfg(test)]
mod tests {
use litellm_core::ocr::transformation::OcrResponseHandling;
use serde_json::json;
use super::*;
use serde_json::json;
#[test]
fn blocks_private_and_metadata_ips() {
@ -443,87 +378,4 @@ mod tests {
assert_eq!(transformed, document);
}
#[test]
fn truncate_error_body_passes_short_strings_through() {
let body = "Unauthorized";
assert_eq!(truncate_error_body(body), "Unauthorized");
}
#[test]
fn truncate_error_body_caps_long_payloads() {
let body = "x".repeat(306);
let truncated = truncate_error_body(&body);
assert!(truncated.ends_with("... (truncated)"));
let prefix_chars = truncated
.strip_suffix("... (truncated)")
.expect("truncated marker present")
.chars()
.count();
assert_eq!(prefix_chars, 256);
}
#[test]
fn truncate_error_body_does_not_split_multibyte_chars() {
let body = "é".repeat(266);
let truncated = truncate_error_body(&body);
assert!(truncated.is_char_boundary(truncated.len()));
}
#[test]
fn ocr_dispatch_supports_migrated_providers() {
assert!(ocr_provider_config("mistral", "mistral-ocr-latest").is_some());
assert!(
ocr_provider_config("azure_ai", "pixtral-12b-2409")
.expect("azure ai config resolves")
.requires_data_uri_document()
);
assert_eq!(
ocr_provider_config("azure_ai", "doc-intelligence/prebuilt-read")
.expect("document intelligence config resolves")
.response_handling(),
OcrResponseHandling::AzureDocumentIntelligencePoll
);
assert!(
ocr_provider_config("vertex_ai", "deepseek-ocr-maas")
.expect("vertex deepseek config resolves")
.supported_ocr_params()
.contains(&"temperature")
);
assert!(ocr_provider_config("openai", "gpt-4o").is_none());
}
#[test]
fn string_headers_accepts_string_values() {
let headers = json!({
"x-trace-id": "trace-1"
})
.as_object()
.unwrap()
.clone();
assert_eq!(
string_headers(Some(headers)).expect("string headers accepted"),
vec![("x-trace-id".to_string(), "trace-1".to_string())]
);
}
#[test]
fn string_headers_rejects_non_string_values() {
let headers = json!({
"x-retry-count": 3
})
.as_object()
.unwrap()
.clone();
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert_eq!(
err,
Error::InvalidRequest(
"OCR extra_headers.x-retry-count must be a string, got number".to_string()
)
);
}
}

View file

@ -0,0 +1,117 @@
use reqwest::{RequestBuilder, StatusCode, header::HeaderMap};
use serde_json::Value;
use super::client::http_client;
use super::common_utils::poll_document_intelligence;
use super::observers::{OcrObserver, OcrPostCall, OcrPreCall};
use super::transformation::OcrResponseHandling;
use super::types::{OcrRequestData, ProviderOcrRequest};
use crate::Error;
use crate::http_utils::{http_request, truncate_error_body};
pub struct OcrHttpResponse {
pub status: StatusCode,
pub headers: HeaderMap,
pub body: String,
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn send_ocr_request(
request: RequestBuilder,
event: &OcrPreCall,
observer: &mut impl OcrObserver,
) -> Result<OcrHttpResponse, Error> {
if observer.pre_call(event).await.is_err() {
tracing::warn!("OCR pre-call observer failed");
}
let response = http_request(request).await.map_err(transport_error)?;
let status = response.status();
let headers = response.headers().clone();
let body = response.text().await.map_err(transport_error)?;
if !status.is_success() {
return Err(Error::Http {
status: status.as_u16(),
body: truncate_error_body(&body),
});
}
let event = OcrPostCall {
original_response: body,
};
if observer.post_call(&event).await.is_err() {
tracing::warn!("OCR post-call observer failed");
}
Ok(OcrHttpResponse {
status,
headers,
body: event.original_response,
})
}
fn transport_error(error: reqwest::Error) -> Error {
Error::Network(if error.is_timeout() {
"Request timed out".into()
} else {
error.to_string()
})
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn execute_ocr_provider_call(
request: ProviderOcrRequest,
observer: &mut impl OcrObserver,
) -> Result<Value, Error> {
let mut request_builder = http_client().post(request.url()).json(request.body());
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let event = OcrPreCall {
model: request.model().to_string(),
request: OcrRequestData {
data: request.body().clone(),
files: None,
},
api_base: request.url().to_string(),
headers: request.upstream_headers.iter().cloned().collect(),
};
let response = send_ocr_request(request_builder, &event, observer).await?;
let status = response.status;
if request.config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll
&& status.as_u16() == 202
{
let operation_url = response
.headers
.get("operation-location")
.and_then(|value| value.to_str().ok())
.map(str::to_string)
.ok_or_else(|| {
Error::InvalidResponse(
"Azure Document Intelligence returned 202 but no Operation-Location header found"
.to_string(),
)
})?;
let response_json = poll_document_intelligence(
&operation_url,
request.url(),
&request.upstream_headers,
request.timeout,
)
.await?;
return Ok(request
.config
.transform_ocr_response(request.model(), response_json)?
.into_json());
}
let response_json: Value = serde_json::from_str(&response.body)
.map_err(|err| Error::InvalidResponse(format!("invalid OCR response JSON: {err}")))?;
Ok(request
.config
.transform_ocr_response(request.model(), response_json)?
.into_json())
}

View file

@ -1,2 +1,45 @@
mod client;
mod common_utils;
pub mod handler;
pub mod observers;
pub mod prepare;
pub mod transformation;
pub mod types;
pub use handler::execute_ocr_provider_call;
pub use prepare::{prepare_ocr_call, prepare_ocr_provider_call};
pub use types::{OcrRequest, PreparedOcrRequest, ProviderOcrRequest};
use serde_json::Value;
use crate::Error;
use crate::ocr::observers::{NoopOcrObserver, OcrObserver};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn ocr(request: OcrRequest<'_>) -> Result<Value, Error> {
ocr_with_observer(request, &mut NoopOcrObserver).await
}
#[tracing::instrument(
name = "ocr",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
pub async fn ocr_with_observer(
request: OcrRequest<'_>,
observer: &mut impl OcrObserver,
) -> Result<Value, Error> {
let prepared = prepare_ocr_call(request);
let provider_request = prepare_ocr_provider_call(prepared).await?;
execute_ocr_provider_call(provider_request, observer).await
}
pub fn ocr_admitted(model: &str, provider: &str, request_format: Option<&str>) -> bool {
common_utils::ocr_provider_config(provider, model).is_some_and(|config| {
request_format != Some("native") || config.supported_ocr_params().contains(&"req_format")
})
}
#[cfg(test)]
mod tests;

View file

@ -0,0 +1,48 @@
use std::collections::BTreeMap;
use std::convert::Infallible;
use serde::Serialize;
use super::types::OcrRequestData;
#[derive(Serialize)]
pub struct OcrPreCall {
pub model: String,
pub request: OcrRequestData,
pub api_base: String,
pub headers: BTreeMap<String, String>,
}
#[derive(Serialize)]
pub struct OcrPostCall {
pub original_response: String,
}
#[macro_export]
macro_rules! ocr_observer_catalog {
($consumer:path, $($options:tt)*) => {
$consumer! {
$($options)*
{
pre_call: PreCall($crate::ocr::observers::OcrPreCall) -> () = direct;
post_call: PostCall($crate::ocr::observers::OcrPostCall) -> () = direct;
}
}
};
}
ocr_observer_catalog!(crate::define_hooks, pub trait OcrObserver;);
pub struct NoopOcrObserver;
impl OcrObserver for NoopOcrObserver {
type Error = Infallible;
async fn pre_call(&mut self, _input: &OcrPreCall) -> Result<(), Infallible> {
Ok(())
}
async fn post_call(&mut self, _input: &OcrPostCall) -> Result<(), Infallible> {
Ok(())
}
}

View file

@ -0,0 +1,130 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use crate::Error;
use crate::http_utils::string_headers;
use crate::ocr::common_utils::convert_document_url_to_data_uri;
use crate::ocr::transformation::OcrProviderConfig;
use crate::ocr::types::{OcrRequest, PreparedOcrRequest, ProviderOcrRequest};
use crate::providers::azure_ai::ocr::transformation::{
AZURE_AI_OCR_CONFIG, AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG,
};
use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG;
use crate::providers::vertex_ai::ocr::transformation as vertex_ai;
use crate::providers::vertex_ai::ocr::transformation::{
VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG,
};
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn ocr_provider_config(
provider: &str,
model: &str,
) -> Option<&'static dyn OcrProviderConfig> {
match provider {
"mistral" => Some(&MISTRAL_OCR_CONFIG),
"azure_ai" if is_azure_document_intelligence_model(model) => {
Some(&AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG)
}
"azure_ai" => Some(&AZURE_AI_OCR_CONFIG),
"vertex_ai" if vertex_ai::is_deepseek_model(model) => Some(&VERTEX_AI_DEEPSEEK_OCR_CONFIG),
"vertex_ai" => Some(&VERTEX_AI_OCR_CONFIG),
_ => None,
}
}
fn is_azure_document_intelligence_model(model: &str) -> bool {
let model = model.to_ascii_lowercase();
model.contains("doc-intelligence") || model.contains("documentintelligence")
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrRequest {
let call_id = request
.litellm_call_id
.map(str::to_string)
.unwrap_or_else(new_ocr_call_id);
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.unwrap_or(CustomLlmProvider {
model: request.model,
custom_llm_provider: "mistral",
});
let model = provider_info.model.to_string();
let custom_llm_provider = provider_info.custom_llm_provider.to_string();
let config = ocr_provider_config(&custom_llm_provider, &model)
.ok_or_else(|| Error::InvalidProvider(custom_llm_provider.clone()));
let optional_params = match &config {
Ok(config) => {
let supported = config.supported_ocr_params();
config.map_ocr_params(
&request
.optional_params
.into_iter()
.filter(|(name, _)| supported.contains(&name.as_str()))
.collect(),
)
}
Err(_) => request.optional_params,
};
PreparedOcrRequest {
config,
model,
custom_llm_provider,
litellm_call_id: call_id,
document: request.document,
api_key: request.api_key.map(str::to_string),
api_base: request.api_base.map(str::to_string),
extra_headers: request.extra_headers,
optional_params,
timeout: request.timeout,
}
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub async fn prepare_ocr_provider_call(
request: PreparedOcrRequest,
) -> Result<ProviderOcrRequest, Error> {
let config = request.config?;
let env_lookup = |key: &str| std::env::var(key).ok();
let upstream_headers = config.validate_environment(
string_headers("OCR", request.extra_headers)?,
request.api_key.as_deref(),
&env_lookup,
)?;
let url = config.complete_url(
request.api_base.as_deref(),
&request.model,
&request.optional_params,
&env_lookup,
)?;
let model = request.model.clone();
let custom_llm_provider = request.custom_llm_provider.clone();
let document = if config.requires_data_uri_document() {
convert_document_url_to_data_uri(request.document).await?
} else {
request.document
};
let body = config
.transform_ocr_request(&request.model, document, request.optional_params)?
.data;
Ok(ProviderOcrRequest {
model,
custom_llm_provider,
config,
url,
body,
upstream_headers,
timeout: request.timeout,
})
}
fn new_ocr_call_id() -> String {
static COUNTER: AtomicU64 = AtomicU64::new(1);
let sequence = COUNTER.fetch_add(1, Ordering::Relaxed);
let timestamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_nanos())
.unwrap_or(0);
format!("ocr-{timestamp}-{sequence}")
}

View file

@ -0,0 +1,362 @@
use std::sync::{Arc, Mutex};
use std::time::Duration;
use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use crate::Error;
use crate::http_utils::has_header;
use crate::ocr::observers::{OcrObserver, OcrPostCall, OcrPreCall};
use crate::ocr::prepare::ocr_provider_config;
use crate::ocr::transformation::OcrResponseHandling;
use crate::ocr::{OcrRequest, ocr, ocr_with_observer};
struct ProviderObserver {
events: Arc<Mutex<Vec<&'static str>>>,
raw_response: Option<String>,
reject: bool,
}
impl OcrObserver for ProviderObserver {
type Error = &'static str;
async fn pre_call(&mut self, input: &OcrPreCall) -> Result<(), Self::Error> {
assert_eq!(input.model, "mistral-ocr-4-1");
assert_eq!(
input.request.data["document"]["document_url"],
"https://example.com/document.pdf"
);
assert!(input.api_base.ends_with("/v1/ocr"));
assert!(
input
.headers
.values()
.any(|value| value == "Bearer test-key")
);
self.events.lock().unwrap().push("pre");
if self.reject {
Err("observer failure")
} else {
Ok(())
}
}
async fn post_call(&mut self, input: &OcrPostCall) -> Result<(), Self::Error> {
self.events.lock().unwrap().push("post");
self.raw_response = Some(input.original_response.clone());
if self.reject {
Err("observer failure")
} else {
Ok(())
}
}
}
fn observer_request(api_base: &str) -> OcrRequest<'_> {
OcrRequest {
model: "mistral/mistral-ocr-4-1",
document: json!({"type":"document_url","document_url":"https://example.com/document.pdf"}),
api_key: Some("test-key"),
api_base: Some(api_base),
custom_llm_provider: Some("mistral"),
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(2)),
litellm_call_id: Some("observer-test"),
}
}
async fn observer_case(status: u16, body: &'static str, reject: bool) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
let events = Arc::new(Mutex::new(Vec::new()));
let provider_events = Arc::clone(&events);
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let request = read_http_request(&mut socket).await;
assert!(request.starts_with("POST /v1/ocr "));
provider_events.lock().unwrap().push("http");
let response = format!(
"HTTP/1.1 {status} Test\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
body.len()
);
socket.write_all(response.as_bytes()).await.unwrap();
});
let mut observer = ProviderObserver {
events: Arc::clone(&events),
raw_response: None,
reject,
};
let result = ocr_with_observer(observer_request(&url), &mut observer).await;
tokio::time::timeout(Duration::from_secs(2), server)
.await
.unwrap()
.unwrap();
if status != 200 {
assert!(matches!(result, Err(Error::Http { status: actual, .. }) if actual == status));
assert_eq!(*events.lock().unwrap(), ["pre", "http"]);
assert_eq!(observer.raw_response, None);
} else {
assert_eq!(*events.lock().unwrap(), ["pre", "http", "post"]);
assert_eq!(observer.raw_response.as_deref(), Some(body));
if body == "invalid-json" {
assert!(matches!(result, Err(Error::InvalidResponse(_))));
} else {
assert_eq!(result.unwrap()["pages"][0]["markdown"], "ok");
}
}
}
#[tokio::test]
async fn provider_observers_surround_http_and_cannot_replace_its_outcome() {
for reject in [false, true] {
observer_case(200, r#"{"pages":[{"index":0,"markdown":"ok"}]}"#, reject).await;
observer_case(200, "invalid-json", reject).await;
observer_case(401, r#"{"error":"rejected"}"#, reject).await;
}
}
#[tokio::test]
async fn invalid_ocr_preparation_does_not_call_observers_or_provider() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let url = format!("http://{}/v1", listener.local_addr().unwrap());
let events = Arc::new(Mutex::new(Vec::new()));
let mut observer = ProviderObserver {
events: Arc::clone(&events),
raw_response: None,
reject: false,
};
let request = OcrRequest {
document: json!(42),
..observer_request(&url)
};
assert!(ocr_with_observer(request, &mut observer).await.is_err());
assert!(events.lock().unwrap().is_empty());
assert!(
tokio::time::timeout(Duration::from_millis(50), listener.accept())
.await
.is_err()
);
}
async fn read_http_headers(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
if request.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
String::from_utf8(request).expect("request is utf8")
}
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 1024];
let header_end = loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break request.len();
}
request.extend_from_slice(&buffer[..n]);
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") {
break position + 4;
}
};
let headers = String::from_utf8_lossy(&request[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
while request.len().saturating_sub(header_end) < content_length {
let n = socket.read(&mut buffer).await.expect("reads body");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
}
String::from_utf8(request).expect("request is utf8")
}
#[test]
fn ocr_dispatch_supports_migrated_providers() {
assert!(ocr_provider_config("mistral", "mistral-ocr-latest").is_some());
assert!(
ocr_provider_config("azure_ai", "pixtral-12b-2409")
.expect("azure ai config resolves")
.requires_data_uri_document()
);
assert_eq!(
ocr_provider_config("azure_ai", "doc-intelligence/prebuilt-read")
.expect("document intelligence config resolves")
.response_handling(),
OcrResponseHandling::AzureDocumentIntelligencePoll
);
assert!(
ocr_provider_config("vertex_ai", "deepseek-ocr-maas")
.expect("vertex deepseek config resolves")
.supported_ocr_params()
.contains(&"temperature")
);
assert!(ocr_provider_config("openai", "gpt-4o").is_none());
}
#[test]
fn auth_header_detection_is_case_insensitive() {
let headers = vec![
("x-trace-id".to_string(), "trace-1".to_string()),
("authorization".to_string(), "Bearer sk-test".to_string()),
];
assert!(has_header(&headers, "authorization"));
let headers = vec![("Authorization".to_string(), "Bearer sk-test".to_string())];
assert!(has_header(&headers, "authorization"));
let headers = vec![("x-trace-id".to_string(), "trace-1".to_string())];
assert!(!has_header(&headers, "authorization"));
}
#[tokio::test]
async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let request = read_http_headers(&mut socket).await;
let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
request
});
let mut headers = Map::new();
headers.insert(
"Authorization".to_string(),
Value::String("Bearer sk-from-python".to_string()),
);
headers.insert(
"x-trace-id".to_string(),
Value::String("trace-1".to_string()),
);
let response = ocr(OcrRequest {
model: "mistral-ocr-latest",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("sk-for-rust-fallback"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("mistral"),
extra_headers: Some(headers),
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
litellm_call_id: None,
})
.await
.expect("ocr request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
let request = server.await.expect("server task completes");
let authorization_count = request
.lines()
.filter(|line| line.to_ascii_lowercase().starts_with("authorization:"))
.count();
assert_eq!(authorization_count, 1, "{request}");
assert!(
request.contains("authorization: Bearer sk-from-python")
|| request.contains("Authorization: Bearer sk-from-python"),
"{request}"
);
}
#[tokio::test]
async fn document_intelligence_poll_uses_resolved_subscription_key() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let operation_url = format!("http://{addr}/operations/1");
let server = tokio::spawn(async move {
let (mut post_socket, _) = listener.accept().await.expect("accepts post request");
let post_request = read_http_headers(&mut post_socket).await;
let post_response = format!(
"HTTP/1.1 202 Accepted\r\noperation-location: {operation_url}\r\ncontent-length: 0\r\nconnection: close\r\n\r\n"
);
post_socket
.write_all(post_response.as_bytes())
.await
.expect("writes post response");
let (mut poll_socket, _) = listener.accept().await.expect("accepts poll request");
let poll_request = read_http_headers(&mut poll_socket).await;
let response_body = r#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"ok"}]}]}}"#;
let poll_response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
poll_socket
.write_all(poll_response.as_bytes())
.await
.expect("writes poll response");
(post_request, poll_request)
});
let response = ocr(OcrRequest {
model: "doc-intelligence/prebuilt-read",
document: json!({
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}),
api_key: Some("di-key"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("azure_ai"),
extra_headers: None,
optional_params: Map::new(),
timeout: Some(Duration::from_secs(5)),
litellm_call_id: None,
})
.await
.expect("document intelligence request succeeds");
assert_eq!(response["pages"][0]["markdown"], "ok");
let (post_request, poll_request) = server.await.expect("server task completes");
assert!(
post_request
.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: di-key"),
"{post_request}"
);
assert!(
poll_request
.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: di-key"),
"{poll_request}"
);
}

View file

@ -1,6 +1,80 @@
use std::time::Duration;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use crate::Error;
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest};
use crate::ocr::transformation::OcrProviderConfig;
pub struct OcrRequest<'a> {
pub model: &'a str,
pub document: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
pub litellm_call_id: Option<&'a str>,
}
pub struct PreparedOcrRequest {
pub config: Result<&'static dyn OcrProviderConfig, Error>,
pub model: String,
pub custom_llm_provider: String,
pub litellm_call_id: String,
pub document: Value,
pub api_key: Option<String>,
pub api_base: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
}
impl CallLifecycleRequest for PreparedOcrRequest {
fn lifecycle_context(&self) -> CallLifecycleContext {
CallLifecycleContext::new(
"ocr",
self.model.clone(),
self.custom_llm_provider.clone(),
self.litellm_call_id.clone(),
)
}
}
pub struct ProviderOcrRequest {
pub(super) model: String,
pub(super) custom_llm_provider: String,
pub(super) config: &'static dyn OcrProviderConfig,
pub(super) url: String,
pub(super) body: Value,
pub(super) upstream_headers: Vec<(String, String)>,
pub(super) timeout: Option<Duration>,
}
impl ProviderOcrRequest {
pub fn model(&self) -> &str {
&self.model
}
pub fn custom_llm_provider(&self) -> &str {
&self.custom_llm_provider
}
pub fn url(&self) -> &str {
&self.url
}
pub fn body(&self) -> &Value {
&self.body
}
pub fn with_body(self, body: Value) -> Self {
Self { body, ..self }
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct OcrRequestData {
pub data: Value,

View file

@ -0,0 +1,677 @@
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use tokio::net::TcpStream;
use tokio::runtime::Handle;
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, connect_async};
use crate::Error;
use crate::constants::{RESPONSES_WS_CONNECT_TIMEOUT_SECS, RESPONSES_WS_IDLE_TIMEOUT_SECS};
use crate::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG;
use crate::responses::types::ResponsesWsEvent;
use crate::responses::websocket::ResponsesWebSocketProviderConfig;
const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable";
pub type ResponsesUpstreamWs = WebSocketStream<MaybeTlsStream<TcpStream>>;
type UpstreamTx = SplitSink<ResponsesUpstreamWs, Message>;
type UpstreamRx = SplitStream<ResponsesUpstreamWs>;
/// A connected Responses WebSocket upstream.
///
/// The send and receive halves are locked separately so a pending
/// [`recv_text`](Self::recv_text) never blocks
/// [`send_text`](Self::send_text). Dropping the last clone closes the
/// upstream socket: the runtime handle captured at connect time is used to
/// flush a close frame without blocking the dropping thread.
#[derive(Clone)]
pub struct ResponsesWebSocketConnection {
tx: Arc<Mutex<Option<UpstreamTx>>>,
rx: Arc<Mutex<Option<UpstreamRx>>>,
runtime: Handle,
}
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 runtime = Handle::try_current().map_err(|error| Error::Network(error.to_string()))?;
let connect = connect_async(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()),
})?;
let (tx, rx) = socket.split();
Ok(Self {
tx: Arc::new(Mutex::new(Some(tx))),
rx: Arc::new(Mutex::new(Some(rx))),
runtime,
})
}
pub async fn send_text(&self, text: String) -> Result<(), Error> {
let mut sender = self.tx.lock().await;
let Some(sender) = sender.as_mut() else {
return Err(Error::Network("Responses WebSocket is closed".to_string()));
};
sender
.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 receiver = self.rx.lock().await;
let Some(receiver) = receiver.as_mut() else {
return Ok(None);
};
match receiver.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 sender = self.tx.lock().await;
if let Some(mut sender) = sender.take() {
sender
.send(Message::Close(None))
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
let mut receiver = self.rx.lock().await;
*receiver = None;
Ok(())
}
}
impl Drop for ResponsesWebSocketConnection {
fn drop(&mut self) {
if Arc::strong_count(&self.tx) != 1 || Arc::strong_count(&self.rx) != 1 {
return;
}
let tx = Arc::clone(&self.tx);
let rx = Arc::clone(&self.rx);
self.runtime.spawn(async move {
let mut sender = tx.lock().await;
if let Some(mut sender) = sender.take() {
let _ = sender.send(Message::Close(None)).await;
}
let mut receiver = rx.lock().await;
*receiver = None;
});
}
}
fn resolve_api_key(api_key: Option<&str>) -> Result<String, Error> {
api_key
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
.or_else(|| {
std::env::var(OPENAI_API_KEY_ENV)
.ok()
.filter(|value| !value.trim().is_empty())
})
.ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string()))
}
async fn dial_upstream(
model: &str,
api_key: &str,
api_base: Option<&str>,
) -> Result<ResponsesUpstreamWs, Error> {
let url = OPENAI_RESPONSES_WS_CONFIG.complete_websocket_url(api_base, model);
let mut request = url
.as_str()
.into_client_request()
.map_err(|error| Error::Network(error.to_string()))?;
request.headers_mut().insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {api_key}"))
.map_err(|error| Error::Auth(error.to_string()))?,
);
let result = tokio::time::timeout(
Duration::from_secs(RESPONSES_WS_CONNECT_TIMEOUT_SECS),
connect_async(request),
)
.await
.map_err(|_| Error::Network("Responses WebSocket connection timed out".to_string()))?;
result
.map(|(socket, _)| socket)
.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()),
})
}
pub struct ResponsesWebSocketStreaming;
impl ResponsesWebSocketStreaming {
pub async fn bidirectional_forward<In, Out>(
model: &str,
upstream_tx: UpstreamTx,
upstream_rx: UpstreamRx,
idle_timeout: Option<Duration>,
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
splice(
model,
upstream_tx,
upstream_rx,
idle_timeout,
observe,
client_in,
client_out,
)
.await
}
}
async fn splice<In, Out>(
model: &str,
mut upstream_tx: UpstreamTx,
mut upstream_rx: UpstreamRx,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
mut client_in: In,
mut client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
let idle = idle_timeout.unwrap_or_else(|| Duration::from_secs(RESPONSES_WS_IDLE_TIMEOUT_SECS));
loop {
tokio::select! {
event = client_in.next() => {
let Some(event) = event else { break };
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_request(&event, model)?
.events
{
let payload = serde_json::to_string(&outbound)
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
upstream_tx.send(Message::Text(payload))
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
}
message = upstream_rx.next() => {
let Some(message) = message else { break };
match message.map_err(|error| Error::Network(error.to_string()))? {
Message::Text(text) => {
let event = serde_json::from_str::<ResponsesWsEvent>(&text)
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
observe(&event);
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_response(&event, model)?
.events
{
client_out.send(outbound)
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
}
Message::Close(_) => break,
_ => {}
}
}
_ = tokio::time::sleep(idle) => break,
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub async fn async_responses_websocket<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Option<Duration>,
mut observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
let key = resolve_api_key(api_key)?;
let upstream = dial_upstream(model, &key, api_base).await?;
let (mut upstream_tx, upstream_rx) = upstream.split();
if let Some(first_frame) = first_frame {
for outbound in OPENAI_RESPONSES_WS_CONFIG
.transform_ws_request(&first_frame, model)?
.events
{
let payload = serde_json::to_string(&outbound)
.map_err(|error| Error::InvalidResponse(error.to_string()))?;
upstream_tx
.send(Message::Text(payload))
.await
.map_err(|error| Error::Network(error.to_string()))?;
}
}
ResponsesWebSocketStreaming::bidirectional_forward(
model,
upstream_tx,
upstream_rx,
idle_timeout,
&mut observe,
client_in,
client_out,
)
.await
}
#[allow(clippy::too_many_arguments)]
pub async fn responses_ws<In, Out>(
model: &str,
api_key: Option<&str>,
api_base: Option<&str>,
first_frame: Option<ResponsesWsEvent>,
idle_timeout: Option<Duration>,
observe: impl FnMut(&ResponsesWsEvent) + Send,
client_in: In,
client_out: Out,
) -> Result<(), Error>
where
In: Stream<Item = ResponsesWsEvent> + Unpin + Send,
Out: Sink<ResponsesWsEvent> + Unpin + Send,
Out::Error: std::fmt::Display,
{
async_responses_websocket(
model,
api_key,
api_base,
first_frame,
idle_timeout,
observe,
client_in,
client_out,
)
.await
}
#[cfg(test)]
mod tests {
use super::*;
use crate::responses::types::ResponsesWsEventType;
use futures_channel::mpsc;
use futures_util::{SinkExt, StreamExt};
use serde_json::json;
use tokio::io::AsyncWriteExt;
use tokio::net::TcpListener;
use tokio_tungstenite::accept_async;
async fn websocket_base() -> (String, tokio::task::JoinHandle<()>) {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("local address");
let task = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let mut socket = accept_async(stream).await.expect("websocket handshake");
while let Some(Ok(Message::Text(text))) = socket.next().await {
let request: serde_json::Value = serde_json::from_str(&text).expect("request json");
let model = request
.get("model")
.and_then(serde_json::Value::as_str)
.or_else(|| {
request
.get("response")
.and_then(serde_json::Value::as_object)
.and_then(|response| {
response.get("model").and_then(serde_json::Value::as_str)
})
})
.expect("enforced model");
socket
.send(Message::Text(
json!({
"type": "response.created",
"response": {
"id": format!("resp-{model}"),
"model": model,
"extra": "preserved"
}
})
.to_string(),
))
.await
.expect("created event");
socket
.send(Message::Text(
json!({
"type": "response.completed",
"response": {
"id": format!("resp-{model}"),
"model": model,
"usage": {
"input_tokens": 1,
"output_tokens": 2,
"total_tokens": 3
}
}
})
.to_string(),
))
.await
.expect("completed event");
}
});
(format!("http://{address}"), task)
}
fn event(value: serde_json::Value) -> ResponsesWsEvent {
serde_json::from_value(value).expect("event")
}
#[test]
fn explicit_nonblank_key_wins() {
assert_eq!(
resolve_api_key(Some(" explicit ")).expect("key"),
"explicit"
);
}
#[test]
fn blank_key_is_not_accepted_without_environment_key() {
if std::env::var(OPENAI_API_KEY_ENV).is_err() {
assert!(resolve_api_key(Some(" ")).is_err());
}
}
#[tokio::test]
async fn forwards_events_sequentially_and_enforces_model() {
let (api_base, server) = websocket_base().await;
let (client_tx, client_rx) = mpsc::unbounded();
let (output_tx, mut output_rx) = mpsc::unbounded();
let (observed_tx, observed_rx) = mpsc::unbounded();
client_tx
.unbounded_send(event(json!({
"type": "response.create",
"model": "wrong"
})))
.expect("first request");
client_tx
.unbounded_send(event(json!({
"type": "response.create",
"response": {"model": "also-wrong"}
})))
.expect("second request");
let task = tokio::spawn(async move {
responses_ws(
"authorized-model",
Some("test-key"),
Some(&api_base),
None,
Some(Duration::from_secs(1)),
move |event| {
observed_tx
.unbounded_send(event.clone())
.expect("observe event");
},
client_rx,
output_tx,
)
.await
});
let first = output_rx.next().await.expect("first output");
let second = output_rx.next().await.expect("second output");
let third = output_rx.next().await.expect("third output");
let fourth = output_rx.next().await.expect("fourth output");
drop(client_tx);
task.await.expect("splice task").expect("successful splice");
server.await.expect("server task");
assert_eq!(first.event_type, ResponsesWsEventType::ResponseCreated);
assert_eq!(first.model(), Some("authorized-model"));
assert_eq!(first.data["response"]["extra"], "preserved");
assert_eq!(second.event_type, ResponsesWsEventType::ResponseCompleted);
assert_eq!(third.event_type, ResponsesWsEventType::ResponseCreated);
assert_eq!(fourth.event_type, ResponsesWsEventType::ResponseCompleted);
let observed: Vec<_> = observed_rx.collect().await;
assert_eq!(observed.len(), 4);
assert!(
observed
.iter()
.all(|event| event.event_type != ResponsesWsEventType::ResponseCreate)
);
}
#[tokio::test]
async fn idle_timeout_ends_without_upstream_events() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let _socket = accept_async(stream).await.expect("handshake");
tokio::time::sleep(Duration::from_secs(1)).await;
});
let (_client_tx, client_rx) = mpsc::unbounded::<ResponsesWsEvent>();
let (output_tx, mut output_rx) = mpsc::unbounded();
let result = responses_ws(
"model",
Some("key"),
Some(&format!("http://{address}")),
None,
Some(Duration::from_millis(20)),
|_| {},
client_rx,
output_tx,
)
.await;
assert!(result.is_ok());
assert!(output_rx.next().await.is_none());
server.abort();
}
#[tokio::test]
async fn dial_http_status_is_preserved() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("accept");
stream
.write_all(b"HTTP/1.1 401 Unauthorized\r\nContent-Length: 0\r\n\r\n")
.await
.expect("response");
});
let (_client_tx, client_rx) = mpsc::unbounded::<ResponsesWsEvent>();
let (output_tx, _output_rx) = mpsc::unbounded();
let error = responses_ws(
"model",
Some("key"),
Some(&format!("http://{address}")),
None,
Some(Duration::from_millis(20)),
|_| {},
client_rx,
output_tx,
)
.await
.expect_err("status error");
assert!(matches!(error, Error::Http { status: 401, .. }));
server.await.expect("server task");
}
#[tokio::test]
async fn dial_http_500_status_is_preserved() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.expect("accept");
stream
.write_all(b"HTTP/1.1 500 Internal Server Error\r\nContent-Length: 0\r\n\r\n")
.await
.expect("response");
});
let (_client_tx, client_rx) = mpsc::unbounded::<ResponsesWsEvent>();
let (output_tx, _output_rx) = mpsc::unbounded();
let error = responses_ws(
"model",
Some("key"),
Some(&format!("http://{address}")),
None,
Some(Duration::from_millis(20)),
|_| {},
client_rx,
output_tx,
)
.await
.expect_err("status error");
assert!(matches!(error, Error::Http { status: 500, .. }));
server.await.expect("server task");
}
#[tokio::test]
async fn dropping_the_last_connection_closes_the_upstream_socket() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let mut socket = accept_async(stream).await.expect("handshake");
let message = socket.next().await.expect("close frame").expect("frame");
assert!(matches!(message, Message::Close(_)));
});
let connection = ResponsesWebSocketConnection::connect_url(
&format!("ws://{address}"),
&HashMap::new(),
None,
)
.await
.expect("connect");
drop(connection);
tokio::time::timeout(Duration::from_secs(2), server)
.await
.expect("close frame after drop")
.expect("server task");
}
#[tokio::test]
async fn dropping_one_clone_leaves_the_socket_open_until_the_last_drop() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let mut socket = accept_async(stream).await.expect("handshake");
let early = tokio::time::timeout(Duration::from_millis(150), socket.next()).await;
assert!(early.is_err(), "no close frame while a clone is alive");
let message = socket
.next()
.await
.expect("close frame after last drop")
.expect("frame");
assert!(matches!(message, Message::Close(_)));
});
let connection = ResponsesWebSocketConnection::connect_url(
&format!("ws://{address}"),
&HashMap::new(),
None,
)
.await
.expect("connect");
let clone = connection.clone();
drop(connection);
tokio::time::sleep(Duration::from_millis(250)).await;
drop(clone);
tokio::time::timeout(Duration::from_secs(2), server)
.await
.expect("server completes")
.expect("server task");
}
#[tokio::test]
async fn send_text_is_not_blocked_by_a_pending_recv_text() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("bind");
let address = listener.local_addr().expect("address");
let server = tokio::spawn(async move {
let (stream, _) = listener.accept().await.expect("accept");
let mut socket = accept_async(stream).await.expect("handshake");
let message = socket.next().await.expect("send frame").expect("frame");
assert_eq!(message, Message::Text("ping".into()));
socket
.send(Message::Text("pong".into()))
.await
.expect("reply");
assert!(matches!(socket.next().await, Some(Ok(Message::Close(_)))));
});
let connection = ResponsesWebSocketConnection::connect_url(
&format!("ws://{address}"),
&HashMap::new(),
None,
)
.await
.expect("connect");
let pending = tokio::spawn({
let connection = connection.clone();
async move { connection.recv_text().await }
});
tokio::time::sleep(Duration::from_millis(100)).await;
tokio::time::timeout(Duration::from_secs(1), connection.send_text("ping".into()))
.await
.expect("send completes while recv is pending")
.expect("send succeeds");
let received = tokio::time::timeout(Duration::from_secs(1), pending)
.await
.expect("recv completes")
.expect("recv task");
assert_eq!(received.expect("recv succeeds"), Some("pong".to_string()));
connection.close().await.expect("close");
tokio::time::timeout(Duration::from_secs(2), server)
.await
.expect("server completes")
.expect("server task");
}
}

View file

@ -1,3 +1,4 @@
pub mod connection;
pub mod instrumentation;
pub mod types;
pub mod websocket;

View file

@ -110,3 +110,73 @@ fn crates_directory_matches_allowlist() {
let expected: BTreeSet<String> = EXPECTED_CRATE_DIRS.iter().map(|s| s.to_string()).collect();
assert_eq!(actual, expected, "{MISMATCH}");
}
/// Parse the dependency names out of a crate manifest's `[dependencies]`
/// table.
///
/// Same hand-rolled approach as [`parse_members`]: take the lines after the
/// `[dependencies]` header up to the next table header (a line starting with
/// `[`), then keep the text before ` = ` on each non-comment line, trimmed of
/// the `.workspace`-style shorthand suffix.
fn parse_dependencies(manifest: &str) -> BTreeSet<String> {
let Some((_, after_table)) = manifest.split_once("[dependencies]") else {
return BTreeSet::new();
};
after_table
.lines()
.map(str::trim)
.take_while(|line| !line.starts_with('['))
.filter_map(|line| {
let (name, _) = line.split_once(" = ")?;
let name = name.split('.').next().unwrap_or(name);
(!name.is_empty() && !name.starts_with('#')).then(|| name.to_string())
})
.collect()
}
fn crate_manifest(root: &Path, crate_dir: &str) -> String {
fs::read_to_string(root.join("crates").join(crate_dir).join("Cargo.toml"))
.unwrap_or_else(|error| panic!("{crate_dir}/Cargo.toml should be readable: {error}"))
}
/// The python bridge depends on the domain layers, never on the gateway: the
/// bridge and the axum server are alternative hosts over `litellm-core`, so a
/// bridge -> gateway edge would drag the server crate into every cdylib build
/// and let provider I/O creep back out of core.
#[test]
fn python_bridge_dependencies_stay_on_core_and_interop() {
let manifest = crate_manifest(&workspace_root(), "python-bridge");
let dependencies = parse_dependencies(&manifest);
assert!(
dependencies.contains("litellm-core"),
"python-bridge must depend on litellm-core, got {dependencies:?}"
);
assert!(
dependencies.contains("litellm-python-interop"),
"python-bridge must depend on litellm-python-interop, got {dependencies:?}"
);
assert!(
!dependencies.contains("litellm-ai-gateway"),
"python-bridge must not depend on litellm-ai-gateway (it belongs behind the \
gateway's own host surface, not the cdylib), got {dependencies:?}"
);
}
/// `litellm-core` stays a pure Rust SDK: Python bindings (pyo3, pythonize,
/// pyo3-async-runtimes) live in `litellm-python-interop` /
/// `litellm-python-bridge`, never in core.
#[test]
fn core_dependencies_stay_python_free() {
let manifest = crate_manifest(&workspace_root(), "core");
let dependencies = parse_dependencies(&manifest);
for banned in ["pyo3", "pyo3-async-runtimes", "pythonize"] {
assert!(
!dependencies.contains(banned),
"litellm-core must not depend on {banned}; Python binding crates own it, \
got {dependencies:?}"
);
}
}

View file

@ -14,17 +14,12 @@ default = ["abi3"]
abi3 = ["pyo3/abi3-py310"]
extension-module = ["pyo3/extension-module"]
panic-test = []
trace-parity = [
"dep:tracing",
"litellm-core/observability",
"litellm-ai-gateway/trace-parity",
]
trace-parity = ["litellm-core/observability"]
[dependencies]
futures-util.workspace = true
tracing = { workspace = true, optional = true }
tracing.workspace = true
litellm-core = { workspace = true, features = ["bedrock-auth"] }
litellm-ai-gateway = { workspace = true, default-features = false }
litellm-python-interop.workspace = true
pyo3.workspace = true
pyo3-async-runtimes.workspace = true

View file

@ -9,14 +9,14 @@ mod execution;
#[cfg(feature = "trace-parity")]
mod function_trace;
mod marshal;
mod ocr_callbacks;
mod python_hook_bindings;
mod routes;
use std::sync::atomic::{AtomicU64, Ordering};
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use litellm_core::provider_callbacks::{CallbackDecision, SessionEvent, SessionObserver};
use litellm_core::responses::types::ResponsesWebSocketRequest;
use litellm_core::responses::connection::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use pyo3::prelude::*;
use pyo3::types::PyAny;
@ -65,7 +65,15 @@ impl ResponsesWebSocketConnection {
let mut observer = callback_adapter
.map(|adapter| crate::callback_bindings::python_async_session(adapter, py))
.transpose()?;
let request = ResponsesWebSocketRequest { url: request.url };
let headers = litellm_core::http_utils::string_headers(
"Responses WebSocket",
options.extra_headers.clone(),
)
.map_err(core_error_to_pyerr)?
.into_iter()
.collect();
let url = request.url;
let timeout = options.timeout;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
if let Some(observer) = observer.as_mut() {
let decision = observer
@ -83,9 +91,7 @@ impl ResponsesWebSocketConnection {
}
}
}
let inner = match RustResponsesWebSocketConnection::connect(request, &options, &context)
.await
{
let inner = match RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout).await {
Ok(inner) => inner,
Err(error) => {
if let Some(observer) = observer.as_mut() {
@ -153,6 +159,7 @@ mod _native {
litellm_python_interop::callback_runtime::register(module)?;
super::callback_bindings::register(module)?;
super::ocr_callbacks::register(module)?;
super::errors::register(module)?;
let ready_endpoints = PyDict::new(module.py());
module.add("ready_endpoints", ready_endpoints)?;

View file

@ -0,0 +1,82 @@
use std::num::NonZeroUsize;
use litellm_core::ocr::observers::{OcrObserver, OcrPostCall, OcrPreCall};
use litellm_python_interop::callback_runtime::{AsyncContext, CallbackRuntime, SyncContext};
use pyo3::prelude::*;
use crate::constants::OCR_CALLBACK_CAPACITY;
use crate::execution::PythonCallContext;
litellm_core::ocr_observer_catalog!(crate::bind_python_hooks,
pub(crate) struct PythonOcrSession;
trait OcrObserver;
);
#[pyclass(frozen)]
struct OcrRuntime(CallbackRuntime);
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
let capacity = NonZeroUsize::new(OCR_CALLBACK_CAPACITY)
.expect("OCR callback capacity is a positive constant");
module.add(
"__ocr_callback_runtime__",
OcrRuntime(CallbackRuntime::new(module, capacity)?),
)
}
pub(crate) enum PythonOcrObserver {
Disabled,
Sync(PythonOcrSession<SyncContext>),
Async(PythonOcrSession<AsyncContext>),
}
impl PythonOcrObserver {
pub(crate) fn new(
adapter: Option<Py<PyAny>>,
context: PythonCallContext<'_>,
) -> PyResult<Self> {
let Some(adapter) = adapter else {
return Ok(Self::Disabled);
};
let py = context.py;
let module = py.import("litellm.rust_bridge._native")?;
let runtime = module
.getattr("__ocr_callback_runtime__")?
.extract::<PyRef<'_, OcrRuntime>>()?
.0
.clone();
if context.asynchronous {
Ok(Self::Async(PythonOcrSession::new(
adapter.bind(py),
runtime.async_context(py)?,
)?))
} else {
Ok(Self::Sync(PythonOcrSession::new(
adapter.bind(py),
runtime.sync_context(py)?,
)?))
}
}
}
impl OcrObserver for PythonOcrObserver {
type Error = PyErr;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn pre_call(&mut self, input: &OcrPreCall) -> PyResult<()> {
match self {
Self::Disabled => Ok(()),
Self::Sync(session) => session.pre_call(input).await,
Self::Async(session) => session.pre_call(input).await,
}
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
async fn post_call(&mut self, input: &OcrPostCall) -> PyResult<()> {
match self {
Self::Disabled => Ok(()),
Self::Sync(session) => session.post_call(input).await,
Self::Async(session) => session.post_call(input).await,
}
}
}

View file

@ -1,11 +1,10 @@
use crate::callback_bindings::PythonProviderObserver;
use crate::errors::ocr_error_to_pyerr;
use crate::marshal::{NativeRequestContext, NativeRequestOptions};
use litellm_ai_gateway::integrations::types::RequestHooks;
use litellm_ai_gateway::io::ocr::OcrRequest;
use litellm_ai_gateway::io::ocr::ocr_with_observer as run_route;
use crate::marshal::{NativeRequestContext, NativeRequestOptions, required_value};
use crate::ocr_callbacks::PythonOcrObserver;
use litellm_core::Error;
use litellm_core::ocr::{OcrRequest, ocr_with_observer};
use litellm_core::request_context::LiteLlmRequestContext;
use litellm_core::request_options::RequestOptions;
use pyo3::prelude::*;
use serde_json::{Map, Value};
use std::future::Future;
@ -26,7 +25,7 @@ fn prepare_ocr(
python_context: crate::execution::PythonCallContext<'_>,
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
let context: LiteLlmRequestContext = context.into();
let provider_admitted = litellm_ai_gateway::io::ocr::ocr_admitted(
let provider_admitted = litellm_core::ocr::ocr_admitted(
&input.model,
options.provider("mistral"),
context.capabilities.request_format.as_deref(),
@ -52,19 +51,22 @@ fn prepare_ocr(
input.document
};
let document: Value = litellm_python_interop::from_py(document.bind(py))?;
let mut observer = PythonProviderObserver::new(callback_adapter, python_context)?;
let document = required_value("document", document, Value::is_object, "dict")?;
let call_id = context.litellm_call_id.clone();
let options: RequestOptions = options.into();
let mut observer = PythonOcrObserver::new(callback_adapter, python_context)?;
Ok(async move {
run_route(
ocr_with_observer(
OcrRequest {
model: &input.model,
document,
api_key: options.api_key.as_deref(),
api_base: options.api_base.as_deref(),
custom_llm_provider: options.custom_llm_provider.as_deref(),
extra_headers: options.extra_headers,
optional_params: input.optional_params,
},
&options.into(),
&context,
RequestHooks {
callbacks: Vec::new(),
guardrails: Vec::new(),
timeout: options.timeout,
litellm_call_id: call_id.as_deref(),
},
&mut observer,
)