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