From 43a19d81ab83564278357f800d30ef8cd314b61e Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 12 Sep 2026 07:40:04 -0700 Subject: [PATCH] wip --- .gitignore | 3 + litellm-rust/crates/core/AGENTS.md | 2 +- litellm-rust/crates/core/CLAUDE.md | 66 --- litellm-rust/crates/core/src/constants.rs | 3 + litellm-rust/crates/core/src/error.rs | 9 + .../core/src/ocr/adapters/azure/cohere.rs | 131 ++++++ .../azure/document_intelligence/mod.rs | 3 +- .../azure/document_intelligence/polling.rs | 8 +- .../core/src/ocr/adapters/azure/mistral.rs | 2 +- .../crates/core/src/ocr/adapters/azure/mod.rs | 7 + .../crates/core/src/ocr/adapters/cohere.rs | 123 ++++++ .../crates/core/src/ocr/adapters/mod.rs | 15 +- .../crates/core/src/ocr/codecs/cohere.rs | 235 ++++++++++ .../crates/core/src/ocr/codecs/mod.rs | 1 + litellm-rust/crates/core/src/ocr/error.rs | 4 + litellm-rust/crates/core/src/ocr/handler.rs | 16 +- litellm-rust/crates/core/src/ocr/registry.rs | 10 + litellm-rust/crates/core/src/ocr/wire.rs | 7 +- .../src/providers/azure_ai/auth/resolve.rs | 2 +- .../tests/azure_document_intelligence_ocr.rs | 13 +- .../crates/python-bridge/src/errors.rs | 25 +- .../crates/python-bridge/src/execution.rs | 16 +- litellm-rust/crates/python-bridge/src/lib.rs | 2 +- .../python-bridge/src/routes/ocr_lifecycle.rs | 26 +- litellm-rust/crates/python-interop/src/lib.rs | 4 +- .../crates/python-interop/src/marshal.rs | 32 +- litellm/litellm_core_utils/call_completion.py | 241 ----------- litellm/ocr/input.py | 16 +- litellm/proxy/common_request_processing.py | 9 +- litellm/rust_bridge/ocr_lifecycle.py | 9 - litellm/utils.py | 178 +++++--- .../test_call_completion.py | 402 ------------------ ...st_azure_ai_cohere_parse_transformation.py | 112 ----- .../ocr/test_cohere_parse_transformation.py | 196 ++------- .../test_deferred_guardrail_logging.py | 34 +- .../rust_bridge/test_ocr_lifecycle.py | 34 +- tests/test_litellm/test_utils.py | 32 ++ tests/test_litellm_rust/ocr/test_cohere.py | 141 ++++++ tests/test_litellm_rust/ocr/test_dispatch.py | 38 +- tests/test_litellm_rust/ocr/test_lifecycle.py | 11 +- 40 files changed, 1096 insertions(+), 1122 deletions(-) delete mode 100644 litellm-rust/crates/core/CLAUDE.md create mode 100644 litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs create mode 100644 litellm-rust/crates/core/src/ocr/adapters/cohere.rs create mode 100644 litellm-rust/crates/core/src/ocr/codecs/cohere.rs delete mode 100644 litellm/litellm_core_utils/call_completion.py delete mode 100644 tests/test_litellm/litellm_core_utils/test_call_completion.py create mode 100644 tests/test_litellm_rust/ocr/test_cohere.py diff --git a/.gitignore b/.gitignore index deb0acae56e..7da917ce450 100644 --- a/.gitignore +++ b/.gitignore @@ -147,3 +147,6 @@ crash.*.log ui/litellm-dashboard/out/ litellm.log + +.coverage-rust +coverage-rust.xml diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index aee8b4937ef..9ba7bfb5323 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -2,6 +2,6 @@ litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-leve A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate. -Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`. +Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or host-specific callback execution. Core owns lifecycle sequencing and callback payload construction; hosts execute the selected integrations. Env reads are limited to credential fallback in a route's `prepare.rs`. Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates. diff --git a/litellm-rust/crates/core/CLAUDE.md b/litellm-rust/crates/core/CLAUDE.md deleted file mode 100644 index 5d36305ded5..00000000000 --- a/litellm-rust/crates/core/CLAUDE.md +++ /dev/null @@ -1,66 +0,0 @@ -# CLAUDE.md - -Rules for `litellm-rust/crates/core`. - -## Responsibility - -`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level -LiteLLM call has a public entrypoint here, named after the route -(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and -calling it returns a typed non-streaming response. - -Allowed: -- The public entrypoint for a route, plus its `_stream` variant when the - route supports streaming. -- Provider resolution, auth header construction, URL building, and the provider - HTTP call (shared reused client, connect + request timeouts). -- Shared request/response structs. -- Typed errors with stable, non-sensitive messages. -- Deterministic validation helpers. -- Serialization helpers that intentionally mirror Python output shape. -- Route templates that match Python base config responsibilities, such as - `messages::transformation::AnthropicMessagesProviderConfig`. - -Not allowed: -- Serving HTTP: axum routers, extractors, and other transport concerns. -- Filesystem, database, or cache access. -- Config file reading or rollout state; the host resolves those and passes them - in. Env reads are limited to credential fallback in a route's `prepare.rs`. -- Logging callbacks, tracing spans, spend writes, or customer callbacks. -- Provider-specific branching that belongs in `providers`. -- Panics for user/provider-controlled input. - -## Typed Contracts (core rule) - -Trait and function boundaries MUST be strongly typed. No stringly-typed JSON -(`&str` / `String` / `Vec` / bare `serde_json::Value`) as a transform -input or output. Parse wire bytes into typed structs/enums at the host edge; -`core` and `providers` operate only on those types (e.g. `RealtimeEvent`, -`RealtimeTransformResult`, `OcrRequestData`). A `type`-style discriminator is a -typed field on a struct, not a raw string threaded through the API. - -## Structure - -Use route names directly under `src/`: `messages`, `ocr`, future -`chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not -invent broad names like `engine` for route contracts. - -`src/messages` is the reference shape for a route module: - -``` -mod.rs pub async fn messages(..) (+ messages_stream) -types.rs request/response types -transformation.rs the provider template trait -prepare.rs provider resolution, auth headers, URL -handler.rs the provider call -client.rs the shared reqwest client -``` - -## Parity Rules - -- Every shared type used by a provider transform needs unit tests for - serialization shape. -- If Python parity requires always emitting a `null` field instead of omitting - it, document that in code and pin it with a test. -- Error enums should preserve enough detail for Python/HTTP hosts to map errors - consistently without exposing document contents or upstream bodies. diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index 9469d379462..3910c026dfe 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -63,3 +63,6 @@ pub(crate) const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY"; pub(crate) const REDUCTO_ID_PREFIX: &str = "reducto://"; pub(crate) const AZURE_AI_OCR_PATH: &str = "/providers/mistral/azure/ocr"; pub(crate) const MISTRAL_OCR_API_BASE: &str = "https://api.mistral.ai/v1"; + +pub(crate) const COHERE_PARSE_API_BASE: &str = "https://api.cohere.com"; +pub(crate) const COHERE_API_KEY_ENV: &str = "COHERE_API_KEY"; diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index aa4cb567554..1d1634dae82 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -52,6 +52,15 @@ pub enum Error { Unsupported(&'static str), } +impl Error { + pub const fn http_status_code(&self) -> Option { + match self { + Self::InvalidRequest(_) => Some(400), + _ => None, + } + } +} + #[derive(Debug, ThisError)] pub(crate) enum MediaError { #[error("media URL rejected by network policy")] diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs b/litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs new file mode 100644 index 00000000000..4c8455a171c --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs @@ -0,0 +1,131 @@ +use super::super::OcrAdapter; +use crate::Error; +use crate::ocr::OcrClient; +use crate::ocr::codecs::cohere::{ + CohereParams, CohereResponse, transform_request, transform_response, validate_document, +}; +use crate::ocr::document::{inline_remote_document, validate_inline_document}; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::registry::OcrProvider; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; +use crate::providers::azure_ai::auth::AzureAuthInputs; +use crate::url_utils::ApiUrl; + +const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE"; + +pub(crate) struct AzureCohereAdapter; + +impl OcrAdapter for AzureCohereAdapter { + type ProviderResponse = CohereResponse; + const PROVIDER: OcrProvider = OcrProvider::AzureAi; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let params = super::super::super::wire::decode_request_value::( + serde_json::Value::Object(request.optional_params.clone()), + "optional_params", + )?; + let mut config = AzureAuthInputs::from_sourced_optional_params( + &request.optional_params, + &request.input_sources, + ) + .map_err(Error::from)?; + config.azure_ad_token_provider = request.azure_ad_token_provider.clone(); + let base = request + .connection + .api_base + .clone() + .or_else(|| credential_env(AZURE_AI_API_BASE_ENV)) + .filter(|base| !base.trim().is_empty()) + .ok_or_else(|| { + Error::Auth( + "Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(), + ) + })?; + let headers = + super::validate_ai_environment(&request.connection, &config, &credential_env).await?; + validate_document(&request.document)?; + let remote = request.document.source().starts_with("http://") + || request.document.source().starts_with("https://"); + let document = inline_remote_document( + client.document_fetcher(), + request.document.clone(), + &request.connection, + ) + .await?; + let body = transform_request(&request.model, document, params)?; + transform_request_body( + client, + request, + &complete_url(&base)?, + &headers, + !remote, + body, + |body| { + validate_document(&body.document)?; + validate_inline_document(&body.document) + }, + ) + .await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result { + transform_response(&request.model, response) + } +} + +fn complete_url(base: &str) -> Result { + let mut url = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?; + if !matches!(url.scheme(), "http" | "https") { + return Err(invalid_api_base().into()); + } + let path = url.path().trim_end_matches('/').to_string(); + if path.ends_with("/v2/parse") { + url.set_path(&path); + return Ok(url.into()); + } + url.set_path(path.strip_suffix("/models").unwrap_or(&path)); + ApiUrl::parse(url.as_str()) + .and_then(|url| url.complete_path(&["providers", "cohere", "v2", "parse"])) + .map(|url| url.into_string()) + .map_err(|_| invalid_api_base().into()) +} + +fn invalid_api_base() -> OcrRequestError { + OcrRequestError::RequestField { + path: "api_base".into(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn completes_foundry_urls_without_duplicate_paths_and_preserves_queries() { + for suffix in [ + "", + "/models", + "/providers/cohere/v2", + "/providers/cohere/v2/parse", + ] { + assert_eq!( + complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(), + "https://example.com/providers/cohere/v2/parse?tenant=a" + ); + } + assert_eq!( + complete_url("https://example.com/v2/parse?tenant=a").unwrap(), + "https://example.com/v2/parse?tenant=a" + ); + assert!(complete_url("relative/path").is_err()); + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs index 2c9327a6ce4..e90c27ba59d 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs +++ b/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs @@ -61,7 +61,7 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter { url: &str, headers: &[(String, String)], request: &LiteLLMOcrRequest, - ) -> Result, OcrError> { + ) -> Result, OcrError> { polling::read_operation_response( client.polling_http(), response, @@ -72,7 +72,6 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter { &request.hooks, ) .await - .map(|decoded| decoded.text.into_bytes()) } } diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs b/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs index 9fe026084de..b4f70130171 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs +++ b/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs @@ -44,7 +44,7 @@ pub(super) async fn read_operation_response( } let bytes = crate::ocr::client::read_response_bytes(response).await?; crate::ocr::handler::post_call(hooks, &bytes).await?; - poll_operation(http_client, operation, headers, connection, native).await + poll_operation(http_client, operation, headers, connection, native, hooks).await } async fn poll_operation( @@ -53,6 +53,7 @@ async fn poll_operation( headers: &[(String, String)], connection: &OcrConnection, native: bool, + hooks: &Arc, ) -> Result, OcrError> { let deadline = Instant::now() .checked_add(connection.poll_timeout) @@ -88,7 +89,10 @@ async fn poll_operation( .await .map_err(|_| OcrPollingError::PollTimeout)??; match &decoded.data.status { - Some(OperationStatus::Succeeded) => return Ok(decoded), + Some(OperationStatus::Succeeded) => { + crate::ocr::handler::post_call(hooks, decoded.text.as_bytes()).await?; + return Ok(decoded); + } Some(OperationStatus::Running | OperationStatus::NotStarted) => { tokio::time::timeout_at(deadline, tokio::time::sleep(Duration::from_secs(retry))) .await diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs b/litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs index ee73eb7419c..8639590b05c 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs +++ b/litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs @@ -92,7 +92,7 @@ fn get_complete_url( }) } -async fn validate_environment( +pub(in crate::ocr::adapters) async fn validate_environment( connection: &OcrConnection, config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs index 9c02a7471c9..3d30ae6d6bd 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs +++ b/litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs @@ -1,3 +1,4 @@ +mod cohere; mod document_intelligence; mod mistral; @@ -10,8 +11,10 @@ use crate::ocr::error::OcrError; use crate::ocr::types::OcrConnection; use crate::providers::azure_ai::auth::{AzureAuthInputs, AzureAuthService}; +pub(crate) use cohere::AzureCohereAdapter; pub(crate) use document_intelligence::AzureDocumentIntelligenceAdapter; pub(crate) use mistral::AzureMistralAdapter; +pub(super) use mistral::validate_environment as validate_ai_environment; async fn resolve_entra( config: &AzureAuthInputs, @@ -22,6 +25,10 @@ async fn resolve_entra( .get_or_init(AzureAuthService::default) .get_azure_ad_token(config, env_lookup) .await + .or_else(|error| match error { + crate::AuthError::EmptyAzureToken => Ok(None), + other => Err(other), + }) .map(|credential| { credential.map(|credential| { let source = credential.source(); diff --git a/litellm-rust/crates/core/src/ocr/adapters/cohere.rs b/litellm-rust/crates/core/src/ocr/adapters/cohere.rs new file mode 100644 index 00000000000..933ead7f7f7 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/cohere.rs @@ -0,0 +1,123 @@ +use super::OcrAdapter; +use crate::Error; +use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE}; +use crate::ocr::OcrClient; +use crate::ocr::codecs::cohere::{ + CohereParams, CohereResponse, transform_request, transform_response, validate_document, +}; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::registry::OcrProvider; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection}; +use crate::url_utils::ApiUrl; + +pub(crate) struct CohereAdapter; + +impl OcrAdapter for CohereAdapter { + type ProviderResponse = CohereResponse; + const PROVIDER: OcrProvider = OcrProvider::Cohere; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let params = super::super::wire::decode_request_value::( + serde_json::Value::Object(request.optional_params.clone()), + "optional_params", + )?; + let headers = validate_environment(&request.connection, &credential_env)?; + let url = complete_url( + request + .connection + .api_base + .as_deref() + .unwrap_or(COHERE_PARSE_API_BASE), + )?; + let body = transform_request(&request.model, request.document.clone(), params)?; + transform_request_body(client, request, &url, &headers, true, body, |body| { + validate_document(&body.document) + }) + .await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result { + transform_response(&request.model, response) + } +} + +fn complete_url(base: &str) -> Result { + let parsed = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?; + if !matches!(parsed.scheme(), "http" | "https") { + return Err(invalid_api_base().into()); + } + ApiUrl::parse(base) + .and_then(|url| url.complete_path(&["v2", "parse"])) + .map(|url| url.into_string()) + .map_err(|_| invalid_api_base().into()) +} + +fn invalid_api_base() -> OcrRequestError { + OcrRequestError::RequestField { + path: "api_base".into(), + } +} + +fn validate_environment( + connection: &OcrConnection, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result, OcrError> { + if crate::http_utils::has_header(&connection.extra_headers, "authorization") { + return Ok(connection.extra_headers.clone()); + } + let key = connection + .api_key + .as_deref() + .map(str::trim) + .filter(|key| !key.is_empty()) + .map(str::to_string) + .or_else(|| env_lookup(COHERE_API_KEY_ENV).filter(|key| !key.trim().is_empty())) + .ok_or_else(|| { + Error::Auth("Missing COHERE_API_KEY - set it in the environment or pass api_key".into()) + })?; + Ok( + std::iter::once(("Authorization".into(), format!("Bearer {key}"))) + .chain(connection.extra_headers.clone()) + .collect(), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn completes_provider_urls_without_duplicate_paths_and_preserves_queries() { + for suffix in ["", "/v2", "/v2/parse"] { + assert_eq!( + complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(), + "https://example.com/v2/parse?tenant=a" + ); + } + } + + #[test] + fn rejects_invalid_urls_and_blank_keys() { + assert!(complete_url("relative/path").is_err()); + assert!(complete_url("ftp://example.com").is_err()); + assert!(matches!( + validate_environment( + &OcrConnection { + api_key: Some(" ".into()), + ..Default::default() + }, + &|_| None, + ), + Err(OcrError::Public(Error::Auth(_))) + )); + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/mod.rs index 2d304f7729d..613893b8d87 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/mod.rs +++ b/litellm-rust/crates/core/src/ocr/adapters/mod.rs @@ -8,11 +8,13 @@ use super::registry::OcrProvider; use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; mod azure; +mod cohere; mod mistral; mod reducto; mod vertex; -pub(crate) use azure::{AzureDocumentIntelligenceAdapter, AzureMistralAdapter}; +pub(crate) use azure::{AzureCohereAdapter, AzureDocumentIntelligenceAdapter, AzureMistralAdapter}; +pub(crate) use cohere::CohereAdapter; pub(crate) use mistral::MistralAdapter; pub(crate) use reducto::{ReductoLegacyAdapter, ReductoV3Adapter}; pub(crate) use vertex::{VertexDeepSeekAdapter, VertexMistralAdapter}; @@ -54,11 +56,16 @@ pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static { _url: &str, _headers: &[(String, String)], request: &LiteLLMOcrRequest, - ) -> impl Future, OcrError>> + Send { + ) -> impl Future< + Output = Result, OcrError>, + > + Send { async move { let bytes = super::client::read_response_bytes(response).await?; super::handler::post_call(&request.hooks, &bytes).await?; - Ok(bytes) + Ok(super::wire::decode_response( + &bytes, + request.response_format()? == super::types::OcrResponseFormat::Native, + )?) } } } @@ -66,6 +73,8 @@ pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static { macro_rules! for_each_ocr_adapter { ($callback:ident) => { $callback! { + Cohere, $crate::ocr::adapters::CohereAdapter, $crate::ocr::adapters::CohereAdapter, Cohere; + AzureCohere, $crate::ocr::adapters::AzureCohereAdapter, $crate::ocr::adapters::AzureCohereAdapter, AzureAi; Mistral, $crate::ocr::adapters::MistralAdapter, $crate::ocr::adapters::MistralAdapter, Mistral; AzureMistral, $crate::ocr::adapters::AzureMistralAdapter, $crate::ocr::adapters::AzureMistralAdapter, AzureAi; AzureDocumentIntelligence, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, AzureAi; diff --git a/litellm-rust/crates/core/src/ocr/codecs/cohere.rs b/litellm-rust/crates/core/src/ocr/codecs/cohere.rs new file mode 100644 index 00000000000..2e5eaccc395 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/cohere.rs @@ -0,0 +1,235 @@ +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value, json}; + +use crate::ocr::document::InlineDocument; +use crate::ocr::error::{OcrRequestError, OcrResponseError}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; + +#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)] +#[serde(rename_all = "lowercase")] +pub(crate) enum OutputFormat { + #[default] + Markdown, + Blocks, +} + +#[derive(Deserialize)] +pub(crate) struct CohereParams { + #[serde(default)] + pub output_format: OutputFormat, +} + +#[derive(Deserialize, Serialize)] +pub(crate) struct CohereRequest { + pub model: String, + pub document: OcrDocument, + pub output_format: OutputFormat, +} + +pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), OcrRequestError> { + let OcrDocument::ImageUrl { image_url, .. } = document else { + return Err(OcrRequestError::CohereImageOnly); + }; + if image_url.is_empty() { + return Err(OcrRequestError::CohereImageOnly); + } + if let Some(inline) = InlineDocument::parse(image_url)? { + if !inline.mime_type().type_.eq_ignore_ascii_case("image") { + return Err(OcrRequestError::CohereImageOnly); + } + inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?; + } + Ok(()) +} + +#[derive(Deserialize)] +pub(crate) struct CohereResponse { + #[serde(default)] + pages: Vec, + meta: Option, +} + +#[derive(Deserialize)] +struct CoherePage { + index: Option, + markdown: Option, + blocks: Option>>, +} + +#[derive(Deserialize)] +struct CohereMarkdown { + #[serde(default)] + content: String, + images: Option>>, +} + +#[derive(Deserialize)] +struct CohereMeta { + billed_units: Option, +} + +#[derive(Deserialize)] +struct CohereBilledUnits { + pages: Option, +} + +pub(crate) fn transform_response( + model: &str, + response: CohereResponse, +) -> Result { + let pages_processed = response + .meta + .and_then(|meta| meta.billed_units) + .and_then(|units| units.pages) + .map(Ok) + .unwrap_or_else(|| { + i64::try_from(response.pages.len()).map_err(|_| OcrResponseError::NumericRange("pages")) + })?; + let pages = response + .pages + .into_iter() + .enumerate() + .map(|(position, page)| { + let index = page.index.map(Ok).unwrap_or_else(|| { + i64::try_from(position).map_err(|_| OcrResponseError::NumericRange("page index")) + })?; + let (content, images) = page + .markdown + .map(|markdown| { + let images = + markdown + .images + .filter(|images| !images.is_empty()) + .map(|images| { + images + .into_iter() + .map(|mut image| { + if let Some(Value::Object(bbox)) = + image.get("bounding_box").cloned() + { + image.insert("bbox".into(), Value::Object(bbox)); + } + Value::Object(image) + }) + .collect::>() + }); + (markdown.content, images) + }) + .unwrap_or_default(); + let mut normalized = json!({"index": index, "markdown": content, "images": images}); + if let Some(blocks) = page.blocks { + normalized["blocks"] = json!(blocks); + } + Ok(normalized) + }) + .collect::, OcrResponseError>>()?; + Ok(LiteLLMOcrResponse { + pages, + model: model.into(), + document_annotation: None, + usage_info: Some(json!({"pages_processed": pages_processed})), + object: "ocr".into(), + extra_fields: Map::new(), + provider_native_response: None, + }) +} + +pub(crate) fn transform_request( + model: &str, + document: OcrDocument, + params: CohereParams, +) -> Result { + validate_document(&document)?; + Ok(CohereRequest { + model: model.into(), + document, + output_format: params.output_format, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn response_normalizes_markdown_images_blocks_and_billed_pages() { + let response = serde_json::from_value(json!({ + "pages": [ + {"index": 4, "markdown": {"content": "receipt", "images": [{"id":"image", "bounding_box":{"top_left_x":1}, "description":"scan"}]}}, + {"blocks": [{"type":"text","text":"total"}]} + ], + "meta": {"billed_units":{"pages":3}} + })).unwrap(); + let normalized = transform_response("parse-v5.0", response).unwrap(); + assert_eq!(normalized.pages[0]["index"], 4); + assert_eq!(normalized.pages[0]["markdown"], "receipt"); + assert_eq!(normalized.pages[0]["images"][0]["bbox"]["top_left_x"], 1); + assert_eq!(normalized.pages[0]["images"][0]["description"], "scan"); + assert_eq!(normalized.pages[1]["index"], 1); + assert_eq!(normalized.pages[1]["markdown"], ""); + assert_eq!(normalized.pages[1]["blocks"][0]["text"], "total"); + assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 3); + } + + #[test] + fn response_defaults_and_invalid_fields() { + for value in [ + json!({}), + json!({"meta":null}), + json!({"pages":[],"meta":{"billed_units":null}}), + ] { + let normalized = + transform_response("parse", serde_json::from_value(value).unwrap()).unwrap(); + assert!(normalized.pages.is_empty()); + assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 0); + } + for value in [ + json!({"pages":null}), + json!({"pages":[{"markdown":"text"}]}), + json!({"pages":[{"index":"bad"}]}), + ] { + assert!(serde_json::from_value::(value).is_err()); + } + let normalized = transform_response( + "parse", + serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(), + ) + .unwrap(); + assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 1); + assert!(normalized.pages[0]["images"].is_null()); + } + + #[test] + fn request_requires_image_and_supported_output_format() { + for value in [ + json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), + json!({"type":"image_url","image_url":""}), + json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}), + ] { + assert_eq!( + validate_document(&serde_json::from_value(value).unwrap()), + Err(OcrRequestError::CohereImageOnly) + ); + } + assert!(serde_json::from_value::(json!({"output_format":"html"})).is_err()); + for format in ["markdown", "blocks"] { + assert!( + serde_json::from_value::(json!({"output_format":format})).is_ok() + ); + } + let request = transform_request( + "parse-v5.0", + serde_json::from_value(json!({ + "type":"image_url", + "image_url":"https://example.com/image.png" + })) + .unwrap(), + serde_json::from_value(json!({})).unwrap(), + ) + .unwrap(); + assert_eq!( + serde_json::to_value(request).unwrap()["output_format"], + "markdown" + ); + } +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/mod.rs index 7c752749901..639b985b9ae 100644 --- a/litellm-rust/crates/core/src/ocr/codecs/mod.rs +++ b/litellm-rust/crates/core/src/ocr/codecs/mod.rs @@ -1,3 +1,4 @@ +pub(crate) mod cohere; pub(crate) mod deepseek; pub(crate) mod document_intelligence; pub(crate) mod mistral; diff --git a/litellm-rust/crates/core/src/ocr/error.rs b/litellm-rust/crates/core/src/ocr/error.rs index 522d059ec48..1aa35998e2b 100644 --- a/litellm-rust/crates/core/src/ocr/error.rs +++ b/litellm-rust/crates/core/src/ocr/error.rs @@ -4,6 +4,10 @@ use crate::error::TransportError; #[derive(Debug, Clone, PartialEq, Eq, Error)] pub enum OcrRequestError { + #[error( + "Cohere Parse only accepts `image_url` documents; document_url and PDF inputs are not supported" + )] + CohereImageOnly, #[error("Invalid `req_format`. Expected 'native' or 'litellm'.")] RequestFormat, #[error("invalid OCR request field: {path}")] diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index 6b03a8ff92d..cd1d538aaa8 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -75,10 +75,10 @@ impl PreparedOcrCall { ($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => { match self.request.adapter { $( OcrAdapterKind::$variant => { - let bytes = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?; + let decoded = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?; Ok(OcrProviderResponse { request: self.request, - bytes, + data: OcrProviderData::$variant(decoded), }) }, )+ } @@ -106,12 +106,14 @@ fn request_headers(request: &reqwest::Request) -> Result, macro_rules! provider_data { ($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => { + enum OcrProviderData { + $( $variant(super::wire::DecodedOcrResponse<<$adapter as OcrAdapter>::ProviderResponse>), )+ + } + impl OcrProviderResponse { pub(crate) fn normalize(self) -> Result { - let native = self.request.response_format()? == super::types::OcrResponseFormat::Native; - match self.request.adapter { - $( OcrAdapterKind::$variant => { - let decoded = super::wire::decode_response::<<$adapter as OcrAdapter>::ProviderResponse>(&self.bytes, native)?; + match self.data { + $( OcrProviderData::$variant(decoded) => { let response = $instance.transform_ocr_response(&self.request, decoded.data)?; Ok(LiteLLMOcrResponse { provider_native_response: decoded.native, ..response }) }, )+ @@ -123,7 +125,7 @@ macro_rules! provider_data { pub(crate) struct OcrProviderResponse { request: LiteLLMOcrRequest, - bytes: Vec, + data: OcrProviderData, } pub(crate) async fn post_call(hooks: &Arc, bytes: &[u8]) -> Result<(), Error> { diff --git a/litellm-rust/crates/core/src/ocr/registry.rs b/litellm-rust/crates/core/src/ocr/registry.rs index 1b20a91143b..f47dcb5e6fc 100644 --- a/litellm-rust/crates/core/src/ocr/registry.rs +++ b/litellm-rust/crates/core/src/ocr/registry.rs @@ -23,6 +23,7 @@ super::adapters::for_each_ocr_adapter!(define_adapter_types); #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(crate) enum OcrProvider { + Cohere, Mistral, AzureAi, Reducto, @@ -32,6 +33,7 @@ pub(crate) enum OcrProvider { impl OcrProvider { pub(crate) const fn as_str(self) -> &'static str { match self { + Self::Cohere => "cohere", Self::Mistral => "mistral", Self::AzureAi => "azure_ai", Self::Reducto => "reducto", @@ -50,6 +52,7 @@ pub(crate) fn resolve_wire_adapter( custom_llm_provider: OcrProvider::Mistral.as_str(), }); let typed_provider = match provider.custom_llm_provider { + "cohere" => OcrProvider::Cohere, "mistral" => OcrProvider::Mistral, "azure_ai" => OcrProvider::AzureAi, "reducto" => OcrProvider::Reducto, @@ -57,10 +60,17 @@ pub(crate) fn resolve_wire_adapter( value => return Err(Error::InvalidProvider(value.to_string())), }; let adapter = match typed_provider { + OcrProvider::Cohere => OcrAdapterKind::Cohere, OcrProvider::Mistral => OcrAdapterKind::Mistral, OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => { OcrAdapterKind::AzureDocumentIntelligence } + OcrProvider::AzureAi + if provider.model.to_ascii_lowercase().contains("cohere") + && provider.model.to_ascii_lowercase().contains("parse") => + { + OcrAdapterKind::AzureCohere + } OcrProvider::AzureAi => OcrAdapterKind::AzureMistral, OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => { OcrAdapterKind::ReductoLegacy diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index 7d9dad04686..00868bf6574 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -89,6 +89,7 @@ pub fn consumed_optional_param_names( let (_, adapter) = super::registry::resolve_wire_adapter(model, custom_llm_provider)?; let provider_fields: &[&str] = match adapter { + OcrAdapterKind::Cohere | OcrAdapterKind::AzureCohere => &["output_format"], OcrAdapterKind::Mistral | OcrAdapterKind::AzureMistral | OcrAdapterKind::VertexMistral => { MISTRAL_OPTION_FIELDS } @@ -98,9 +99,9 @@ pub fn consumed_optional_param_names( OcrAdapterKind::VertexDeepSeek => DEEPSEEK_OPTION_FIELDS, }; let auth_fields: &[&str] = match adapter { - OcrAdapterKind::AzureMistral | OcrAdapterKind::AzureDocumentIntelligence => { - AZURE_AUTH_OPTION_FIELDS - } + OcrAdapterKind::AzureMistral + | OcrAdapterKind::AzureDocumentIntelligence + | OcrAdapterKind::AzureCohere => AZURE_AUTH_OPTION_FIELDS, OcrAdapterKind::VertexMistral | OcrAdapterKind::VertexDeepSeek => VERTEX_AUTH_OPTION_FIELDS, _ => &[], }; diff --git a/litellm-rust/crates/core/src/providers/azure_ai/auth/resolve.rs b/litellm-rust/crates/core/src/providers/azure_ai/auth/resolve.rs index 627fbdf446f..025dd4f8740 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/auth/resolve.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/auth/resolve.rs @@ -81,7 +81,7 @@ impl AzureAuthService { AzureCredentialPlan::Caller(caller) => { let credential = caller.acquire().await?; if credential.secret().expose().is_empty() { - return Ok(None); + return Err(AuthError::EmptyAzureToken); } Ok(Some(Sourced::new(credential, InputSource::Deployment))) } diff --git a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs index d96cdad06d3..7c635ebae62 100644 --- a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs +++ b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs @@ -180,8 +180,17 @@ impl super::hooks::OcrHooks for SubmissionBoundary { request: super::hooks::OcrPostCallRequest, ) -> super::hooks::OcrHookFuture<'_, super::hooks::OcrPostCallRequest> { Box::pin(async move { - assert_eq!(self.request_count.lock().unwrap().len(), 1); - assert_eq!(request.original_response, json!(r#"{"submitted":true}"#)); + match self.request_count.lock().unwrap().len() { + 1 => assert_eq!(request.original_response, json!(r#"{"submitted":true}"#)), + 2 => assert!( + request + .original_response + .as_str() + .unwrap() + .contains("succeeded") + ), + count => panic!("unexpected callback after {count} requests"), + } Ok(request) }) } diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 5501c97f9c4..74006dae117 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -73,7 +73,18 @@ pub(crate) fn ocr_error_to_pyerr(err: Error) -> PyErr { Error::Network(message) if message.contains("timed out") => { ocr_upstream_error(408, message) } - other => core_error_to_pyerr(other), + 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 + } } } @@ -111,6 +122,18 @@ mod ocr_error_tests { .and_then(|args| args.extract()) .expect("OCR failures retain status and unprefixed provider message"); assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string())); + + let mapped = ocr_error_to_pyerr(Error::InvalidRequest("invalid format".into())); + assert!(mapped.is_instance_of::(py)); + assert_eq!( + mapped + .value(py) + .getattr("status_code") + .unwrap() + .extract::() + .unwrap(), + 400 + ); }); } } diff --git a/litellm-rust/crates/python-bridge/src/execution.rs b/litellm-rust/crates/python-bridge/src/execution.rs index 02cafc1b9bc..e14e1a0fb6d 100644 --- a/litellm-rust/crates/python-bridge/src/execution.rs +++ b/litellm-rust/crates/python-bridge/src/execution.rs @@ -60,9 +60,14 @@ where E: Send + 'static, F: Future> + Send + 'static, { - let result = run_sync_value_on(py, runtime, async move { - map_core_result(future.await, map_error) - })?; + if Handle::try_current().is_ok() { + return Err(PyRuntimeError::new_err( + "synchronous native routes cannot run from a Tokio context; use the async route", + )); + } + + let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?; + let result = map_core_result(result, map_error)?; Pythonized(result).into_pyobject(py).map(Bound::unbind) } @@ -76,8 +81,9 @@ where E: Send + 'static, F: Future> + Send + 'static, { - run_async_value(py, async move { - let result = map_core_result(future.await, map_error)?; + pyo3_async_runtimes::tokio::future_into_py(py, async move { + let result = catch_future_panic(future).await?; + let result = map_core_result(result, map_error)?; Ok(Pythonized(result)) }) } diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 551872b9928..2700ff207d4 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -65,7 +65,7 @@ impl ResponsesWebSocketConnection { } } -#[pymodule(gil_used = true)] +#[pymodule(gil_used = false)] mod _native { use pyo3::prelude::*; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr_lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/ocr_lifecycle.rs index 68ea69bda13..379a67010e8 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr_lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr_lifecycle.rs @@ -10,7 +10,9 @@ use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_param_names, dec use litellm_core::ocr::{ NativeOutcome, OcrAdmission, OcrCall, OcrClient, OcrHostOperation, OcrHostResult, }; -use litellm_python_interop::{from_py, to_py}; +use litellm_python_interop::{ + from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, +}; use crate::errors::{RustBridgeDeclined, ocr_error_to_pyerr}; use crate::lifecycle::{PythonCallState, PythonRoute, missing_state, now, run_call}; @@ -82,9 +84,14 @@ impl PythonOcrHost { fn python_pre_call( &mut self, py: Python<'_>, - request: OcrDuringCallRequest, + mut request: OcrDuringCallRequest, ) -> PyResult { let pre_call = self.pre_call.as_ref().ok_or_else(missing_state)?; + if let Some(body) = request.body.as_object_mut() { + for name in &request.retained_fields { + body.remove(name); + } + } let body = to_py(py, &request.body)? .into_bound(py) .cast_into::()?; @@ -134,11 +141,9 @@ impl PythonOcrHost { .iter() .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) .collect::>>()?; - Ok(OcrDuringCallRequest { - body: from_py(&body)?, - headers, - ..request - }) + request.body = from_py(&body)?; + request.headers = headers; + Ok(request) } fn python_post_call( @@ -413,6 +418,13 @@ fn _ocr_lifecycle( kwargs: Bound<'_, PyDict>, asynchronous: bool, ) -> PyResult> { + if let Ok(gil_enabled) = py.import("sys")?.getattr("_is_gil_enabled") + && !gil_enabled.call0()?.is_truthy()? + { + return Err(pyo3::exceptions::PyRuntimeError::new_err( + "native OCR requires the Python GIL", + )); + } let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?; let call = admitted_call(OcrCall::admit( client, diff --git a/litellm-rust/crates/python-interop/src/lib.rs b/litellm-rust/crates/python-interop/src/lib.rs index 2e562bdae70..79af79e8c61 100644 --- a/litellm-rust/crates/python-interop/src/lib.rs +++ b/litellm-rust/crates/python-interop/src/lib.rs @@ -2,4 +2,6 @@ mod gil; mod marshal; pub use gil::{release_count, release_gil}; -pub use marshal::{Pythonized, from_py, panic_to_pyerr, to_py}; +pub use marshal::{ + Pythonized, from_py, from_py_preserving_errors, panic_to_pyerr, to_py, to_py_preserving_errors, +}; diff --git a/litellm-rust/crates/python-interop/src/marshal.rs b/litellm-rust/crates/python-interop/src/marshal.rs index 0c7890243af..ed4cce862c0 100644 --- a/litellm-rust/crates/python-interop/src/marshal.rs +++ b/litellm-rust/crates/python-interop/src/marshal.rs @@ -1,12 +1,20 @@ use std::any::Any; use std::panic::{AssertUnwindSafe, catch_unwind}; +use pyo3::exceptions::PyValueError; use pyo3::panic::PanicException; use pyo3::prelude::*; use serde::Serialize; use serde::de::DeserializeOwned; pub fn from_py(value: &Bound<'_, PyAny>) -> PyResult +where + T: DeserializeOwned, +{ + pythonize::depythonize(value).map_err(|error| PyValueError::new_err(error.to_string())) +} + +pub fn from_py_preserving_errors(value: &Bound<'_, PyAny>) -> PyResult where T: DeserializeOwned, { @@ -17,14 +25,18 @@ pub fn to_py(py: Python<'_>, value: &T) -> PyResult> where T: Serialize + ?Sized, { - pythonize_bound(py, value).map(Bound::unbind) + pythonize::pythonize(py, value) + .map(Bound::unbind) + .map_err(|error| PyValueError::new_err(error.to_string())) } -fn pythonize_bound<'py, T>(py: Python<'py>, value: &T) -> PyResult> +pub fn to_py_preserving_errors(py: Python<'_>, value: &T) -> PyResult> where T: Serialize + ?Sized, { - pythonize::pythonize(py, value).map_err(PyErr::from) + pythonize::pythonize(py, value) + .map(Bound::unbind) + .map_err(PyErr::from) } pub struct Pythonized(pub T); @@ -38,7 +50,9 @@ where type Error = PyErr; fn into_pyobject(self, py: Python<'py>) -> PyResult { - catch_unwind(AssertUnwindSafe(|| pythonize_bound(py, &self.0))).map_err(panic_to_pyerr)? + catch_unwind(AssertUnwindSafe(|| pythonize::pythonize(py, &self.0))) + .map_err(panic_to_pyerr)? + .map_err(|error| PyValueError::new_err(error.to_string())) } } @@ -112,7 +126,15 @@ value = Broken() Some(&locals), ) .unwrap(); - let error = from_py::(&locals.get_item("value").unwrap().unwrap()).unwrap_err(); + let value = locals.get_item("value").unwrap().unwrap(); + let legacy_error = from_py::(&value).unwrap_err(); + assert!(legacy_error.is_instance_of::(py)); + assert!( + !legacy_error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + let error = from_py_preserving_errors::(&value).unwrap_err(); assert!( error .value(py) diff --git a/litellm/litellm_core_utils/call_completion.py b/litellm/litellm_core_utils/call_completion.py deleted file mode 100644 index b4cb6e8b276..00000000000 --- a/litellm/litellm_core_utils/call_completion.py +++ /dev/null @@ -1,241 +0,0 @@ -from __future__ import annotations - -import asyncio -import contextvars -import datetime -from collections.abc import Awaitable, Callable, Coroutine -from concurrent.futures import Future -from typing import Protocol - - -class CompletionLogging(Protocol): - def success_handler( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - def failure_handler( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - def async_success_handler( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> Coroutine[object, object, None]: ... - - def async_failure_handler( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> Awaitable[None]: ... - - def handle_sync_success_callbacks_for_async_calls( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - -class CompletionExecutor(Protocol): - def submit( - self, - function: Callable[..., object], - /, - *args: object, - ) -> Future[object]: ... - - -class Completion(Protocol): - def success( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - def failure( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - async def async_failure( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - -class PythonCompletion: - def __init__( - self, - logging_obj: CompletionLogging, - executor: CompletionExecutor | None, - *, - async_call: bool, - internal_call: bool, - completion_with_fallbacks: bool, - ) -> None: - self._logging_obj = logging_obj - self._executor = executor - self._async_call = async_call - self._internal_call = internal_call - self._completion_with_fallbacks = completion_with_fallbacks - - def success( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - if not self._async_call: - assert self._executor is not None - context = contextvars.copy_context() - self._executor.submit( - context.run, - self._logging_obj.success_handler, - result, - start_time, - end_time, - ) - return - - if not self._internal_call: - if getattr(self._logging_obj, "_defer_async_logging", False): - - def enqueue_deferred_logging() -> None: - asyncio.create_task(self._dispatch_async_success(result, start_time, end_time)) - - setattr( # noqa: B010 # optional legacy logger field is absent from narrow test doubles - self._logging_obj, - "_enqueue_deferred_logging", - enqueue_deferred_logging, - ) - else: - asyncio.create_task(self._dispatch_async_success(result, start_time, end_time)) - - self._logging_obj.handle_sync_success_callbacks_for_async_calls( - result=result, - start_time=start_time, - end_time=end_time, - ) - - def failure( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - if self._async_call and self._internal_call: - return - self._logging_obj.failure_handler(exception, traceback_exception, start_time, end_time) - - async def async_failure( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - if not self._async_call or self._internal_call: - return - await self._logging_obj.async_failure_handler(exception, traceback_exception, start_time, end_time) - - async def _dispatch_async_success( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - if self._completion_with_fallbacks: - return - from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER - - GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( # pyright: ignore[reportUnknownMemberType] # legacy worker lacks generic coroutine annotations - async_coroutine=self._logging_obj.async_success_handler( - result=result, - start_time=start_time, - end_time=end_time, - ) - ) - self._logging_obj.handle_sync_success_callbacks_for_async_calls( - result=result, - start_time=start_time, - end_time=end_time, - ) - - -class CallCompletion: - def __init__(self, implementation: Completion) -> None: - self._python_implementation: Completion | None = implementation - self._implementation: Completion | None = implementation - self._attached = False - - @property - def python_implementation(self) -> Completion: - assert self._python_implementation is not None - return self._python_implementation - - def attach(self, implementation: Completion) -> bool: - if self._attached: - return False - self._implementation = implementation - self._attached = True - return True - - def release(self) -> None: - self._python_implementation = None - self._implementation = None - - def success( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - implementation = self._implementation - assert implementation is not None - implementation.success(result, start_time, end_time) - - def failure( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - implementation = self._implementation - assert implementation is not None - implementation.failure(exception, traceback_exception, start_time, end_time) - - async def async_failure( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - implementation = self._implementation - assert implementation is not None - await implementation.async_failure( - exception, - traceback_exception, - start_time, - end_time, - ) diff --git a/litellm/ocr/input.py b/litellm/ocr/input.py index 59ef4ccbebc..d8a94a6d078 100644 --- a/litellm/ocr/input.py +++ b/litellm/ocr/input.py @@ -3,7 +3,9 @@ import mimetypes import os import re from io import IOBase -from typing import Any, Final +from typing import Final, Literal, Protocol + +from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm._logging import verbose_logger @@ -31,7 +33,17 @@ def get_mime_type(file_path: str) -> str: return guessed or "application/octet-stream" -def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str, str]: +class FileReader(Protocol): + def read(self) -> bytes | str: ... + + +class FileDocument(TypedDict): + type: ReadOnly[Literal["file"]] + file: ReadOnly[bytes | os.PathLike[str] | FileReader] + mime_type: ReadOnly[NotRequired[str]] + + +def convert_file_document_to_url_document(document: FileDocument) -> dict[str, str]: file_input: Final = document.get("file") if file_input is None: raise ValueError( diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index b21999d2384..289fe086379 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -3181,10 +3181,11 @@ class ProxyBaseLLMRequestProcessing: Extracted as a static method so tests can exercise the production gating logic directly rather than reimplementing the finally block. """ - pending: Final = getattr(logging_obj, "_native_pending_logging", None) - if pending is not None: - logging_obj._native_pending_logging = None # rebind-ok: consume the native release signal once - pending.release(not exception_raised) + if getattr(logging_obj, "call_type", None) in ("ocr", "aocr"): + pending: Final = getattr(logging_obj, "_native_pending_logging", None) + if pending is not None: + logging_obj._native_pending_logging = None # rebind-ok: consume the native OCR release signal once + pending.release(not exception_raised) _enqueue_fn: Final = getattr(logging_obj, "_enqueue_deferred_logging", None) if _enqueue_fn is None: return diff --git a/litellm/rust_bridge/ocr_lifecycle.py b/litellm/rust_bridge/ocr_lifecycle.py index 1cd6c4572bc..f252ee99752 100644 --- a/litellm/rust_bridge/ocr_lifecycle.py +++ b/litellm/rust_bridge/ocr_lifecycle.py @@ -8,7 +8,6 @@ from typing import Final, Protocol, cast # noqa: TID251 # validates dynamicall import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge.bindings import NativeBinding -from litellm.rust_bridge.configuration import rust_enabled from litellm.rust_bridge.ocr import LiteLLMOcrRequest @@ -44,8 +43,6 @@ NATIVE_OCR_LIFECYCLE: Final = NativeBinding("_ocr_lifecycle", validate=_binding) def select(request: LiteLLMOcrRequest) -> NativeOcrLifecycle | None: - if not rust_enabled(): - return None if litellm.cache is not None or request.kwargs.get("caching") or request.kwargs.get("aocr"): return None return NATIVE_OCR_LIFECYCLE.load() @@ -99,12 +96,6 @@ def call_azure_ad_token_provider(provider: object) -> str: def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception: - if isinstance(error, ValueError) and "Invalid `req_format`" in str(error): - return litellm.BadRequestError( - message=str(error), - model=request.model, - llm_provider=request_provider, - ) mapper: Final = cast( # cast-ok: bounded adapter for the legacy public exception mapper ExceptionMapper, litellm.exception_type ) diff --git a/litellm/utils.py b/litellm/utils.py index f039c0a5482..1a77655a5a4 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7,6 +7,7 @@ import ast import asyncio import base64 import binascii +import contextvars import copy import datetime import hashlib @@ -80,7 +81,6 @@ from litellm.constants import ( PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO, TOOL_CHOICE_OBJECT_TOKEN_COUNT, ) -from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.litellm_core_utils.fallback_generalizations import ( match_capability_generalizations, @@ -1196,6 +1196,79 @@ def function_setup( raise e +def _dispatch_success_logging( + logging_obj: LiteLLMLoggingObject, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + is_completion_with_fallbacks: bool, + is_litellm_internal_call: bool, +) -> None: + if not is_litellm_internal_call: + if getattr(logging_obj, "_defer_async_logging", False): + + def _enqueue_deferred_logging() -> None: + asyncio.create_task( + _client_async_logging_helper( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + ) + ) + + logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging + else: + asyncio.create_task( + _client_async_logging_helper( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + ) + ) + + logging_obj.handle_sync_success_callbacks_for_async_calls( + result=result, + start_time=start_time, + end_time=end_time, + ) + + +async def _client_async_logging_helper( + logging_obj: LiteLLMLoggingObject, + result, + start_time, + end_time, + is_completion_with_fallbacks: bool, +): + if ( + is_completion_with_fallbacks is False + ): # don't log the parent event litellm.completion_with_fallbacks as a 'log_success_event', this will lead to double logging the same call - https://github.com/BerriAI/litellm/issues/7477 + print_verbose( + f"Async Wrapper: Completed Call, calling async_success_handler: {logging_obj.async_success_handler}" + ) + ################################################ + # Async Logging Worker + ################################################ + from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( + async_coroutine=logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time) + ) + + ################################################ + # Sync Logging Worker + ################################################ + logging_obj.handle_sync_success_callbacks_for_async_calls( + result=result, + start_time=start_time, + end_time=end_time, + ) + + def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tuple[int | None, dict[str, Any]]: """ Get the number of retries from the kwargs and the retry policy. @@ -1464,7 +1537,6 @@ def client(original_function): start_time: Final = datetime.datetime.now() result = None logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) - completion: CallCompletion | None = None # only set litellm_call_id if its not in kwargs if "litellm_call_id" not in kwargs: @@ -1480,17 +1552,7 @@ def client(original_function): # Type assertion: logging_obj is guaranteed to be non-None after function_setup assert logging_obj is not None, "logging_obj should not be None after function_setup" - from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor - completion = CallCompletion( - PythonCompletion( - logging_obj, - logging_executor, - async_call=False, - internal_call=False, - completion_with_fallbacks=False, - ) - ) ## LOAD CREDENTIALS load_credentials_from_list(kwargs) kwargs["litellm_logging_obj"] = logging_obj @@ -1588,12 +1650,7 @@ def client(original_function): except Exception as e: print_verbose(f"Error while checking max token limit: {e}") # MODEL CALL - invocation_kwargs: Final = ( - {**kwargs, "_litellm_call_completion": completion} - if original_function.__name__ == CallTypes.ocr.value - else kwargs - ) - result = original_function(*args, **invocation_kwargs) + result = original_function(*args, **kwargs) end_time = datetime.datetime.now() if _is_streaming_request( kwargs=kwargs, @@ -1647,8 +1704,8 @@ def client(original_function): kwargs=kwargs, ) - update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata") - update_response_metadata( + _update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata") + _update_response_metadata( result=result, logging_obj=logging_obj, model=model, @@ -1656,8 +1713,21 @@ def client(original_function): start_time=start_time, end_time=end_time, ) + + # LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated verbose_logger.info("Wrapper: Completed Call, calling success_handler") - completion.success(result, start_time, end_time) + # Copy the current context to propagate it to the background thread + # This is essential for OpenTelemetry span context propagation + ctx: Final = contextvars.copy_context() + executor: Final = getattr(sys.modules[__name__], "executor") + executor.submit( + ctx.run, + logging_obj.success_handler, + result, + start_time, + end_time, + ) + # RETURN RESULT return result except Exception as e: call_type = original_function.__name__ @@ -1731,14 +1801,11 @@ def client(original_function): end_time = datetime.datetime.now() # LOG FAILURE - handle streaming failure logging in the _next_ object, remove `handle_failure` once it's deprecated - if completion is not None: - completion.failure(e, traceback_exception, start_time, end_time) - elif logging_obj: - logging_obj.failure_handler(e, traceback_exception, start_time, end_time) + if logging_obj: + logging_obj.failure_handler( + e, traceback_exception, start_time, end_time + ) # DO NOT MAKE THREADED - router retry fallback relies on this! raise e - finally: - if completion is not None: - completion.release() @wraps(original_function) async def wrapper_async(*args, **kwargs): @@ -1747,7 +1814,6 @@ def client(original_function): result = None _update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata") logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) - completion: CallCompletion | None = None LLMCachingHandler: Final = _get_cached_llm_caching_handler() _llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler( original_function=original_function, @@ -1773,15 +1839,7 @@ def client(original_function): # Type assertion: logging_obj is guaranteed to be non-None after function_setup assert logging_obj is not None, "logging_obj should not be None after function_setup" - completion = CallCompletion( - PythonCompletion( - logging_obj, - None, - async_call=True, - internal_call=_is_litellm_internal_call, - completion_with_fallbacks=is_completion_with_fallbacks, - ) - ) + modified_kwargs: Final = await async_pre_call_deployment_hook(kwargs, call_type) if modified_kwargs is not None: kwargs = modified_kwargs @@ -1866,12 +1924,7 @@ def client(original_function): # MODEL CALL try: - invocation_kwargs: Final = ( - {**kwargs, "_litellm_call_completion": completion} - if original_function.__name__ == CallTypes.aocr.value - else kwargs - ) - result = await original_function(*args, **invocation_kwargs) + result = await original_function(*args, **kwargs) except Exception as deployment_error: _deployment_call_end_time = datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with try: @@ -1940,7 +1993,14 @@ def client(original_function): and _caching_handler_response is not None and _caching_handler_response.final_embedding_cached_response is not None ): - completion.success(result, start_time, end_time) + _dispatch_success_logging( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + is_litellm_internal_call=_is_litellm_internal_call, + ) return _llm_caching_handler._combine_cached_embedding_response_with_api_result( _caching_handler_response=_caching_handler_response, embedding_response=result, @@ -1956,7 +2016,14 @@ def client(original_function): start_time=start_time, end_time=end_time, ) - completion.success(result, start_time, end_time) + _dispatch_success_logging( + logging_obj=logging_obj, + result=result, + start_time=start_time, + end_time=end_time, + is_completion_with_fallbacks=is_completion_with_fallbacks, + is_litellm_internal_call=_is_litellm_internal_call, + ) return result except Exception as e: @@ -1964,12 +2031,17 @@ def client(original_function): # Reuse the timestamp taken right when the deployment call itself failed, before # the failure hook ran, so a slow callback doesn't inflate the reported duration. end_time = _deployment_call_end_time if _deployment_call_end_time is not None else datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with - if completion is not None: - completion.failure(e, traceback_exception, start_time, end_time) - await completion.async_failure(e, traceback_exception, start_time, end_time) - elif logging_obj and not _is_litellm_internal_call: - logging_obj.failure_handler(e, traceback_exception, start_time, end_time) - await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time) + if logging_obj and not _is_litellm_internal_call: + try: + logging_obj.failure_handler( + e, traceback_exception, start_time, end_time + ) # DO NOT MAKE THREADED - router retry fallback relies on this! + except Exception as e: + raise e + try: + await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time) + except Exception as e: + raise e call_type = original_function.__name__ num_retries, kwargs = _get_wrapper_num_retries(kwargs=kwargs, exception=e) @@ -2039,8 +2111,6 @@ def client(original_function): raise e finally: - if completion is not None: - completion.release() # Restore trace_id/session_id contextvars to their pre-call value once # this call (in this asyncio Task) is fully done - see # request_correlation_in_logs. Unlike wrapper()'s sync path, it's safe to diff --git a/tests/test_litellm/litellm_core_utils/test_call_completion.py b/tests/test_litellm/litellm_core_utils/test_call_completion.py deleted file mode 100644 index 4275728f3a1..00000000000 --- a/tests/test_litellm/litellm_core_utils/test_call_completion.py +++ /dev/null @@ -1,402 +0,0 @@ -import asyncio -import contextvars -import datetime -import weakref -from collections.abc import Callable, Coroutine -from concurrent.futures import Future, ThreadPoolExecutor -from threading import get_ident -from typing import Final -from unittest.mock import AsyncMock, MagicMock - -import pytest - -import litellm -from litellm.litellm_core_utils import thread_pool_executor -from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion -from litellm.llms.base_llm.ocr.transformation import OCRResponse -from litellm.utils import client - - -class RecordingExecutor: - def __init__(self) -> None: - self.submissions: list[tuple[Callable[..., object], tuple[object, ...]]] = [] - - def submit(self, function: Callable[..., object], *args: object) -> Future[object]: - self.submissions.append((function, args)) - future: Final[Future[object]] = Future() - future.set_result(function(*args)) - return future - - -class RecordingCompletion: - def __init__(self) -> None: - self.successes: list[object] = [] - self.failures: list[Exception] = [] - - def success( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - self.successes.append(result) - - def failure( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - self.failures.append(exception) - - async def async_failure( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - self.failures.append(exception) - - -class RecordingLogging: - def __init__(self, marker: contextvars.ContextVar[str], observed: list[tuple[object, str]]) -> None: - self._marker = marker - self._observed = observed - self._defer_async_logging = False - self._enqueue_deferred_logging: Callable[[], None] | None = None - - def success_handler( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: - self._observed.append((result, self._marker.get())) - - def failure_handler( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - async def async_success_handler( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - async def async_failure_handler( - self, - exception: Exception, - traceback_exception: str, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - def handle_sync_success_callbacks_for_async_calls( - self, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - ) -> None: ... - - -def test_python_completion_preserves_sync_context_and_response_identity() -> None: - marker: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("marker") - marker.set("request-context") - executor: Final = RecordingExecutor() - response: Final = object() - observed: Final[list[tuple[object, str]]] = [] - logging_obj: Final = RecordingLogging(marker, observed) - completion: Final = PythonCompletion( - logging_obj, - executor, - async_call=False, - internal_call=False, - completion_with_fallbacks=False, - ) - now: Final = datetime.datetime.now(datetime.timezone.utc) - - completion.success(response, now, now) - - assert observed == [(response, "request-context")] - assert len(executor.submissions) == 1 - - -def test_sync_wrapper_dispatches_with_logging_executor_and_caller_context(monkeypatch: pytest.MonkeyPatch) -> None: - marker: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("wrapper-context", default="missing") - marker.set("request-context") - caller_thread: Final = get_ident() - response: Final = object() - observed: Final[list[tuple[object, str, int]]] = [] - logging_obj: Final = MagicMock() - - def record_success(result: object, start_time: datetime.datetime, end_time: datetime.datetime) -> None: - observed.append((result, marker.get(), get_ident())) - - def ocr(**kwargs: object) -> object: - return response - - logging_obj.success_handler.side_effect = record_success - monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(logging_obj, {}))) - monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) - wrapped: Final = client(ocr) - with ThreadPoolExecutor(max_workers=1) as executor: - monkeypatch.setattr(thread_pool_executor, "executor", executor) - result: Final = wrapped() - - assert result is response - assert len(observed) == 1 - assert observed[0][0] is response - assert observed[0][1] == "request-context" - assert observed[0][2] != caller_thread - - -@pytest.mark.asyncio -async def test_call_completion_attaches_once_and_forwards_final_objects() -> None: - python_completion: Final = RecordingCompletion() - native_completion: Final = RecordingCompletion() - ignored_completion: Final = RecordingCompletion() - completion: Final = CallCompletion(python_completion) - response: Final = object() - error: Final = ValueError("mapped failure") - now: Final = datetime.datetime.now(datetime.timezone.utc) - - assert completion.python_implementation is python_completion - assert completion.attach(native_completion) - assert not completion.attach(ignored_completion) - - completion.success(response, now, now) - completion.failure(error, "traceback", now, now) - await completion.async_failure(error, "traceback", now, now) - - assert native_completion.successes == [response] - assert native_completion.failures == [error, error] - assert python_completion.successes == [] - assert ignored_completion.successes == [] - - -@pytest.mark.asyncio -async def test_python_completion_retains_deferred_success_arguments(monkeypatch: pytest.MonkeyPatch) -> None: - response: Final = object() - logging_obj: Final = MagicMock() - logging_obj._defer_async_logging = True - logging_obj.async_success_handler = AsyncMock() - completion: Final = CallCompletion( - PythonCompletion( - logging_obj, - RecordingExecutor(), - async_call=True, - internal_call=False, - completion_with_fallbacks=False, - ) - ) - worker: Final = MagicMock() - monkeypatch.setattr("litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER", worker) - scheduled: Final[list[Coroutine[object, object, None]]] = [] - monkeypatch.setattr("asyncio.create_task", scheduled.append) - now: Final = datetime.datetime.now(datetime.timezone.utc) - - completion.success(response, now, now) - completion.release() - - logging_obj._enqueue_deferred_logging() - assert len(scheduled) == 1 - await scheduled[0] - await worker.ensure_initialized_and_enqueue.call_args.kwargs["async_coroutine"] - logging_obj.async_success_handler.assert_awaited_once_with(result=response, start_time=now, end_time=now) - - -@pytest.mark.asyncio -async def test_async_ocr_wrapper_injects_completion_after_fresh_deployment_kwargs( - monkeypatch: pytest.MonkeyPatch, -) -> None: - native_completion: Final = RecordingCompletion() - original_response: Final = object() - replacement_response: Final = object() - shared_metadata: Final = {"request": "shared"} - replacement_kwargs: Final[dict[str, object]] = {"metadata": shared_metadata} - hook_input: dict[str, object] | None = None - - async def fresh_kwargs(kwargs: dict[str, object], call_type: str) -> dict[str, object]: - nonlocal hook_input - hook_input = kwargs - return replacement_kwargs - - async def aocr(**kwargs: object) -> object: - completion = kwargs.get("_litellm_call_completion") - assert isinstance(completion, CallCompletion) - assert completion.attach(native_completion) - assert kwargs["metadata"] is shared_metadata - return original_response - - async def replace_response(request_data: dict[str, object], response: object, call_type: object) -> object: - assert response is original_response - return replacement_response - - monkeypatch.setattr("litellm.utils.async_pre_call_deployment_hook", fresh_kwargs) - monkeypatch.setattr("litellm.utils.async_post_call_success_deployment_hook", replace_response) - monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {}))) - monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) - wrapped: Final = client(aocr) - - result: Final = await wrapped() - - assert result is replacement_response - assert native_completion.successes == [replacement_response] - assert hook_input is not None - assert "_litellm_call_completion" not in hook_input - assert "_litellm_call_completion" not in replacement_kwargs - assert replacement_kwargs["metadata"] is shared_metadata - - -@pytest.mark.asyncio -async def test_async_ocr_wrapper_sends_final_failure_to_attached_completion( - monkeypatch: pytest.MonkeyPatch, -) -> None: - native_completion: Final = RecordingCompletion() - mapped_error: Final = ValueError("mapped OCR failure") - - async def aocr(**kwargs: object) -> object: - completion = kwargs.get("_litellm_call_completion") - assert isinstance(completion, CallCompletion) - assert completion.attach(native_completion) - raise mapped_error - - monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {}))) - monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) - wrapped: Final = client(aocr) - - with pytest.raises(ValueError, match="mapped OCR failure") as caught: - await wrapped() - - assert caught.value is mapped_error - assert native_completion.failures == [mapped_error, mapped_error] - - -@pytest.mark.asyncio -async def test_async_ocr_wrapper_reports_metadata_failure_without_success( - monkeypatch: pytest.MonkeyPatch, -) -> None: - native_completion: Final = RecordingCompletion() - response: Final = object() - metadata_error: Final = ValueError("metadata failure") - - async def aocr(**kwargs: object) -> object: - completion = kwargs.get("_litellm_call_completion") - assert isinstance(completion, CallCompletion) - assert completion.attach(native_completion) - return response - - def fail_metadata(**kwargs: object) -> None: - raise metadata_error - - monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {}))) - monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) - monkeypatch.setattr("litellm.utils.update_response_metadata", fail_metadata) - wrapped: Final = client(aocr) - - with pytest.raises(ValueError, match="metadata failure") as caught: - await wrapped() - - assert caught.value is metadata_error - assert native_completion.successes == [] - assert native_completion.failures == [metadata_error, metadata_error] - - -@pytest.mark.parametrize("asynchronous", [False, True]) -@pytest.mark.asyncio -async def test_wrapper_completion_stays_separate_from_provider_options( - monkeypatch: pytest.MonkeyPatch, asynchronous: bool -) -> None: - native_completion: Final = RecordingCompletion() - response: Final = OCRResponse(model="mistral-ocr-latest", pages=[]) - metadata: Final = {"request": "shared"} - pages: Final = [0, 2] - - def ocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> OCRResponse: - assert kwargs["metadata"] is metadata - assert kwargs["pages"] is pages - assert "_litellm_call_completion" not in kwargs - assert _litellm_call_completion.attach(native_completion) - return response - - async def aocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> OCRResponse: - return ocr(_litellm_call_completion=_litellm_call_completion, **kwargs) - - monkeypatch.setattr( - "litellm.utils.function_setup", - MagicMock(return_value=(MagicMock(), {"metadata": metadata, "pages": pages})), - ) - monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) - wrapped: Final = client(aocr if asynchronous else ocr) - arguments: Final = { - "model": "mistral/mistral-ocr-latest", - "document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"}, - "api_key": "test-key", - "metadata": metadata, - "pages": pages, - } - - result: Final = await wrapped(**arguments) if asynchronous else wrapped(**arguments) - - assert result is response - assert native_completion.successes == [response] - assert native_completion.failures == [] - assert arguments["metadata"] is metadata - assert "_litellm_call_completion" not in arguments - - -@pytest.mark.parametrize( - ("asynchronous", "exit_path"), - [(False, "success"), (True, "success"), (False, "callback_error"), (True, "callback_error"), (True, "cancelled")], -) -@pytest.mark.asyncio -async def test_wrapper_releases_completion_resources_on_every_exit( - monkeypatch: pytest.MonkeyPatch, asynchronous: bool, exit_path: str -) -> None: - retained: Final[list[tuple[CallCompletion, weakref.ReferenceType[object]]]] = [] - callback_error: Final = RuntimeError("failure callback failed") - implementation: Final = MagicMock(spec=RecordingCompletion) - implementation.failure.side_effect = callback_error - implementation.async_failure = AsyncMock() - response: Final = object() - - def ocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> object: - retained.append((_litellm_call_completion, weakref.ref(_litellm_call_completion.python_implementation))) - assert _litellm_call_completion.attach(implementation) - if exit_path == "cancelled": - raise asyncio.CancelledError - if exit_path == "callback_error": - raise ValueError("provider failed") - return response - - async def aocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> object: - return ocr(_litellm_call_completion=_litellm_call_completion, **kwargs) - - monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {}))) - monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) - wrapped: Final = client(aocr if asynchronous else ocr) - - if exit_path == "success": - result: Final = await wrapped() if asynchronous else wrapped() - assert result is response - assert implementation.success.call_args.args[0] is response - elif exit_path == "cancelled": - with pytest.raises(asyncio.CancelledError): - await wrapped() - implementation.success.assert_not_called() - implementation.failure.assert_not_called() - else: - with pytest.raises(RuntimeError, match="failure callback failed") as caught: - await wrapped() if asynchronous else wrapped() - assert caught.value is callback_error - implementation.async_failure.assert_not_called() - - assert len(retained) == 1 - assert retained[0][1]() is None diff --git a/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py b/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py index 3f98e9b6a2d..2f457fcb25b 100644 --- a/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py +++ b/tests/test_litellm/llms/azure_ai/ocr/test_azure_ai_cohere_parse_transformation.py @@ -1,9 +1,5 @@ -import base64 -import json - import pytest -import litellm from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config from litellm.llms.azure_ai.ocr.document_intelligence.transformation import AzureDocumentIntelligenceOCRConfig @@ -12,27 +8,6 @@ from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig MODEL = "azure_ai/Cohere-parse-v5" API_BASE = "https://resource.services.ai.azure.com" PARSE_URL = f"{API_BASE}/providers/cohere/v2/parse" -IMAGE_URL = "https://example.com/receipt.png" -PNG_BYTES = base64.b64decode( - "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg==" -) -PNG_DATA_URI = f"data:image/png;base64,{base64.b64encode(PNG_BYTES).decode()}" - - -def _parse_response() -> dict: - return { - "id": "882bf973-9dfa-4d02-9d30-709247008efd", - "pages": [{"index": 0, "type": "markdown", "markdown": {"content": "# Receipt\n\nTotal Due: $4.00"}}], - "meta": {"api_version": {"version": "2"}, "billed_units": {"pages": 1}}, - } - - -@pytest.fixture() -def disable_aiohttp_transport(monkeypatch): - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - litellm.in_memory_llm_clients_cache.flush_cache() - yield - litellm.in_memory_llm_clients_cache.flush_cache() @pytest.mark.parametrize( @@ -95,90 +70,3 @@ def test_validate_environment_requires_api_base(monkeypatch) -> None: with pytest.raises(ValueError, match="AZURE_AI_API_BASE"): AzureAICohereParseConfig().validate_environment(headers={}, model="Cohere-parse-v5", api_key="key") - - -@pytest.mark.asyncio -async def test_aocr_inlines_remote_image_and_posts_to_foundry(disable_aiohttp_transport, respx_mock): - respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"}) - route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) - - response = await litellm.aocr( - model=MODEL, - document={"type": "image_url", "image_url": IMAGE_URL}, - api_base=API_BASE, - api_key="azure-key", - ) - - request = route.calls.last.request - assert request.headers["Authorization"] == "Bearer azure-key" - assert json.loads(request.content) == { - "model": "Cohere-parse-v5", - "document": {"type": "image_url", "image_url": PNG_DATA_URI}, - "output_format": "markdown", - } - assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00" - assert response.usage_info.pages_processed == 1 - - -@pytest.mark.asyncio -async def test_aocr_passes_data_uri_through_without_fetching(disable_aiohttp_transport, respx_mock): - route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) - - await litellm.aocr( - model=MODEL, - document={"type": "image_url", "image_url": PNG_DATA_URI}, - api_base=API_BASE, - api_key="azure-key", - output_format="blocks", - ) - - body = json.loads(route.calls.last.request.content) - assert body["document"]["image_url"] == PNG_DATA_URI - assert body["output_format"] == "blocks" - - -def test_ocr_sync_inlines_remote_image(respx_mock): - respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"}) - route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) - - response = litellm.ocr( - model=MODEL, - document={"type": "image_url", "image_url": IMAGE_URL}, - api_base=API_BASE, - api_key="azure-key", - ) - - assert json.loads(route.calls.last.request.content)["document"]["image_url"] == PNG_DATA_URI - assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00" - - -@pytest.mark.asyncio -async def test_aocr_rejects_pdf_before_calling_foundry(disable_aiohttp_transport, respx_mock): - route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) - - with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info: - await litellm.aocr( - model=MODEL, - document={"type": "document_url", "document_url": "https://example.com/doc.pdf"}, - api_base=API_BASE, - api_key="azure-key", - ) - - assert exc_info.value.llm_provider == "azure_ai" - assert not route.called - - -@pytest.mark.asyncio -async def test_ahealth_check_ocr_sends_an_image_to_the_foundry_cohere_parse_deployment( - disable_aiohttp_transport, respx_mock -): - route = respx_mock.post(PARSE_URL).respond(json=_parse_response()) - - result = await litellm.ahealth_check( - model_params={"model": MODEL, "api_base": API_BASE, "api_key": "test-key"}, mode="ocr" - ) - - document = json.loads(route.calls.last.request.content)["document"] - assert document["type"] == "image_url" - assert document["image_url"].startswith("data:image/png;base64,") - assert "error" not in result diff --git a/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py b/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py index cb9af56f5e0..1f120be6ffa 100644 --- a/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py +++ b/tests/test_litellm/llms/cohere/ocr/test_cohere_parse_transformation.py @@ -1,8 +1,11 @@ -import json +from typing import Final +from unittest.mock import Mock +import httpx import pytest import litellm +from litellm.llms.cohere.ocr.transformation import CohereParseConfig PARSE_URL = "https://api.cohere.com/v2/parse" MODEL = "cohere/parse-v5.0" @@ -57,173 +60,38 @@ def _blocks_response() -> dict: } -@pytest.fixture() -def disable_aiohttp_transport(monkeypatch): - monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) - litellm.in_memory_llm_clients_cache.flush_cache() - yield - litellm.in_memory_llm_clients_cache.flush_cache() - - -@pytest.mark.asyncio -async def test_aocr_sends_markdown_parse_request_and_normalizes_pages(disable_aiohttp_transport, respx_mock): - route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) - - response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key") - - request = route.calls.last.request - assert request.headers["Authorization"] == "Bearer test-key" - assert json.loads(request.content) == { - "model": "parse-v5.0", - "document": IMAGE_DOCUMENT, - "output_format": "markdown", - } - assert response.object == "ocr" - assert [page.index for page in response.pages] == [0, 1] - assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00" - assert response.pages[1].markdown == "Page two" - assert response.pages[1].images is None - image = response.pages[0].images[0] - assert image.bbox == BOUNDING_BOX - assert image.model_extra["description"] == "A parking receipt" - assert image.model_extra["bounding_box_normalized"]["bottom_right_x"] == 1 - assert response.usage_info.pages_processed == 2 - assert response.get_provider_native_response() is None - - -@pytest.mark.asyncio -async def test_aocr_usage_prefers_billed_units_over_page_count(disable_aiohttp_transport, respx_mock): - respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=3)) - - response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key") - - assert response.usage_info.pages_processed == 3 - - -@pytest.mark.asyncio -async def test_aocr_usage_falls_back_to_page_count_without_meta(disable_aiohttp_transport, respx_mock): - respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=None)) - - response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key") - - assert response.usage_info.pages_processed == 2 - - -@pytest.mark.asyncio -async def test_aocr_blocks_output_format_forwards_param_and_keeps_blocks(disable_aiohttp_transport, respx_mock): - route = respx_mock.post(PARSE_URL).respond(json=_blocks_response()) - - response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="blocks") - - assert json.loads(route.calls.last.request.content)["output_format"] == "blocks" - assert response.pages[0].markdown == "" - assert response.pages[0].model_extra["blocks"] == [{"type": "text", "text": "Total Due: $4.00"}] - assert response.usage_info.pages_processed == 1 - - -@pytest.mark.asyncio -async def test_aocr_native_format_carries_provider_payload(disable_aiohttp_transport, respx_mock): - payload = _markdown_response() - route = respx_mock.post(PARSE_URL).respond(json=payload) - - response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", req_format="native") - - assert "req_format" not in json.loads(route.calls.last.request.content) - assert response.get_provider_native_response() == payload - assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00" - - -@pytest.mark.asyncio -async def test_aocr_rejects_unknown_output_format_before_calling_provider(disable_aiohttp_transport, respx_mock): - route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) - - with pytest.raises(litellm.BadRequestError, match="Invalid `output_format`: 'html'") as exc_info: - await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="html") - - assert exc_info.value.status_code == 400 - assert not route.called - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "document", - [ - {"type": "document_url", "document_url": "https://example.com/doc.pdf"}, - {"type": "image_url", "image_url": "data:application/pdf;base64,JVBERi0="}, - {"type": "image_url", "image_url": ""}, - ], -) -async def test_aocr_rejects_non_image_documents_before_calling_provider( - disable_aiohttp_transport, respx_mock, document -): - route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) - - with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info: - await litellm.aocr(model=MODEL, document=document, api_key="test-key") - - assert exc_info.value.status_code == 400 - assert not route.called - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "api_base, expected_url", - [ - ("https://gateway.example.com", "https://gateway.example.com/v2/parse"), - ("https://gateway.example.com/cohere/", "https://gateway.example.com/cohere/v2/parse"), - ("https://gateway.example.com/v2", "https://gateway.example.com/v2/parse"), - ("https://gateway.example.com/v2/parse", "https://gateway.example.com/v2/parse"), - ], -) -async def test_aocr_posts_to_api_base_variants(disable_aiohttp_transport, respx_mock, api_base, expected_url): - route = respx_mock.post(expected_url).respond(json=_markdown_response()) - - await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", api_base=api_base) - - assert route.called - - -@pytest.mark.asyncio -async def test_aocr_surfaces_provider_error_with_its_status_and_message(disable_aiohttp_transport, respx_mock): - respx_mock.post(PARSE_URL).respond( - status_code=400, json={"id": "83b0d95e", "message": "output_format must be `blocks` or `markdown`"} +@pytest.mark.parametrize("output_format", ["markdown", "blocks"]) +def test_transform_cohere_request_filters_options(output_format: str) -> None: + config: Final = CohereParseConfig() + params: Final = config.map_ocr_params( + {"output_format": output_format, "req_format": "native", "unknown": True}, {}, "parse-v5.0" ) - - with pytest.raises(litellm.BadRequestError, match="output_format must be") as exc_info: - await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key") - - assert exc_info.value.status_code == 400 + request: Final = config.transform_ocr_request("parse-v5.0", IMAGE_DOCUMENT, params, {}) + assert request.data == {"model": "parse-v5.0", "document": IMAGE_DOCUMENT, "output_format": output_format} -@pytest.mark.asyncio -async def test_aocr_reads_api_key_from_environment(disable_aiohttp_transport, respx_mock, monkeypatch): - monkeypatch.setenv("COHERE_API_KEY", "env-key") - route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) - - await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT) - - assert route.calls.last.request.headers["Authorization"] == "Bearer env-key" +@pytest.mark.parametrize("native", [False, True]) +def test_transform_cohere_response_keeps_images_and_native_payload(native: bool) -> None: + payload: Final = _markdown_response(3) + response: Final = CohereParseConfig().transform_ocr_response( + "parse-v5.0", httpx.Response(200, json=payload), Mock(), {"req_format": "native" if native else "litellm"} + ) + assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00" + assert response.pages[0].images[0].bbox == BOUNDING_BOX + assert response.pages[0].images[0].model_extra["description"] == "A parking receipt" + assert response.pages[1].images is None + assert response.usage_info.pages_processed == 3 + assert response.get_provider_native_response() == (payload if native else None) -@pytest.mark.asyncio -async def test_aocr_without_api_key_names_the_env_var(disable_aiohttp_transport, respx_mock, monkeypatch): - monkeypatch.delenv("COHERE_API_KEY", raising=False) - monkeypatch.setattr(litellm, "cohere_key", None) - route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) - - with pytest.raises(Exception, match="Missing COHERE_API_KEY"): - await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT) - - assert not route.called +def test_transform_cohere_blocks() -> None: + response: Final = CohereParseConfig().transform_ocr_response( + "parse-v5.0", httpx.Response(200, json=_blocks_response()), Mock() + ) + assert response.pages[0].model_extra["blocks"] == [{"type": "text", "text": "Total Due: $4.00"}] + assert response.pages[0].markdown == "" -@pytest.mark.asyncio -async def test_ahealth_check_ocr_sends_an_image_cohere_parse_accepts(disable_aiohttp_transport, respx_mock): - route = respx_mock.post(PARSE_URL).respond(json=_markdown_response()) - - result = await litellm.ahealth_check(model_params={"model": MODEL, "api_key": "test-key"}, mode="ocr") - - document = json.loads(route.calls.last.request.content)["document"] - assert document["type"] == "image_url" - assert document["image_url"].startswith("data:image/png;base64,") - assert "error" not in result +def test_transform_cohere_rejects_unsupported_output_format() -> None: + with pytest.raises(litellm.UnsupportedParamsError, match="output_format"): + CohereParseConfig().map_ocr_params({"output_format": "html"}, {}, "parse-v5.0") diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index 8fde4cc9d5e..11c3d2f8b20 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -15,7 +15,7 @@ Streaming: CSW.__anext__ stores args on logging_obj at stream end. """ import asyncio -from typing import Any +from typing import Any, Final from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -297,6 +297,38 @@ async def test_no_flag_fires_create_task_normally(): # --------------------------------------------------------------------------- +@pytest.mark.parametrize("call_type", ["ocr", "aocr", "completion", "acompletion", "embedding", "responses"]) +@pytest.mark.parametrize("exception_raised", [False, True]) +def test_native_pending_logging_is_released_only_for_ocr(call_type: str, exception_raised: bool) -> None: + pending: Final = MagicMock() + enqueue: Final = MagicMock() + logger: Final = MagicMock( + call_type=call_type, + _native_pending_logging=pending, + _enqueue_deferred_logging=enqueue, + ) + + ProxyBaseLLMRequestProcessing._flush_deferred_async_logging( + logging_obj=logger, + exception_raised=exception_raised, + ) + ProxyBaseLLMRequestProcessing._flush_deferred_async_logging( + logging_obj=logger, + exception_raised=exception_raised, + ) + + if call_type in ("ocr", "aocr"): + pending.release.assert_called_once_with(not exception_raised) + assert logger._native_pending_logging is None + else: + pending.release.assert_not_called() + assert logger._native_pending_logging is pending + if exception_raised: + enqueue.assert_not_called() + else: + enqueue.assert_called_once_with() + + def test_flush_deferred_async_logging_fires_on_success(): """ Happy path: with no exception, the production flush helper invokes the diff --git a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index 8ef8f5646b6..61673b14fac 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -10,11 +10,11 @@ from litellm.rust_bridge.ocr import LiteLLMOcrRequest from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE -@pytest.mark.parametrize("enabled,available", [(True, False), (False, True), (False, False)]) -def test_public_selection_requires_supported_native_ocr(enabled: bool, available: bool) -> None: +@pytest.mark.parametrize("enabled", [True, False]) +def test_public_selection_requires_available_native_ocr(enabled: bool) -> None: native: Final = Mock(side_effect=AssertionError("must not admit")) litellm.rust(enabled) - NATIVE_OCR_LIFECYCLE.override(native if available else None) + NATIVE_OCR_LIFECYCLE.override(None) try: with pytest.raises(RuntimeError, match="Rust OCR is unavailable or does not support this request"): litellm.ocr("mistral/mistral-ocr-latest", {"type": "document_url", "document_url": "https://example.com"}) @@ -96,7 +96,7 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs() assert "timeout" not in captured[0] -@pytest.mark.parametrize("enabled", [False, True], ids=["legacy", "native"]) +@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"]) def test_public_duplicate_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None: native: Final = Mock(side_effect=AssertionError("binding errors precede admission")) document: Final = {"type": "document_url", "document_url": "https://example.com"} @@ -111,7 +111,7 @@ def test_public_duplicate_argument_error_does_not_depend_on_native_selection(ena assert native.call_count == 0 -@pytest.mark.parametrize("enabled", [False, True], ids=["legacy", "native"]) +@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"]) def test_public_missing_required_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None: native: Final = Mock(side_effect=AssertionError("binding errors precede admission")) litellm.rust(enabled) @@ -123,3 +123,27 @@ def test_public_missing_required_argument_error_does_not_depend_on_native_select NATIVE_OCR_LIFECYCLE.reset() litellm.rust(None) assert native.call_count == 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("enabled", [False, True, None]) +async def test_public_ocr_ignores_rust_flag( + monkeypatch: pytest.MonkeyPatch, asynchronous: bool, enabled: bool | None +) -> None: + from unittest.mock import AsyncMock + + monkeypatch.setenv("LITELLM_RUST", "0") + response: Final = OCRResponse(pages=[], model="mistral-ocr-latest") + native: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response) + litellm.rust(enabled) + NATIVE_OCR_LIFECYCLE.override(native) + try: + if asynchronous: + assert await litellm.aocr("mistral/mistral-ocr-latest", {}) is response + else: + assert litellm.ocr("mistral/mistral-ocr-latest", {}) is response + assert native.call_count == 1 + finally: + NATIVE_OCR_LIFECYCLE.reset() + litellm.rust(None) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index c7e46829aba..19ed31c7b22 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -1,5 +1,6 @@ import asyncio import contextlib +import contextvars import json import logging import os @@ -7,6 +8,7 @@ import queue import threading from datetime import datetime, timedelta, timezone from collections.abc import Iterator +from concurrent.futures import ThreadPoolExecutor from typing import Final from unittest.mock import AsyncMock, MagicMock, patch @@ -60,6 +62,36 @@ from litellm.utils import ( # Adds the parent directory to the system path +def test_non_ocr_wrapper_preserves_logging_executor_and_context(monkeypatch: pytest.MonkeyPatch) -> None: + marker: Final = contextvars.ContextVar("non-ocr-logging-context", default="missing") + token: Final = marker.set("caller-context") + caller_thread: Final = threading.get_ident() + response: Final = object() + logger: Final = MagicMock() + observed: Final = queue.Queue[tuple[object, str, int]]() + + def record_success(result: object, start_time: datetime, end_time: datetime) -> None: + observed.put((result, marker.get(), threading.get_ident())) + + def embedding(**kwargs: object) -> object: + return response + + logger.success_handler.side_effect = record_success + monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(logger, {}))) + try: + with ThreadPoolExecutor(max_workers=1) as executor: + monkeypatch.setattr("litellm.utils.executor", executor) + result: Final = client(embedding)() + logged_response, context, worker_thread = observed.get_nowait() + assert result is response + assert logged_response is response + assert context == "caller-context" + assert worker_thread != caller_thread + assert observed.empty() + finally: + marker.reset(token) + + def test_cloudflare_model_info_includes_rpm(local_model_cost_map: None) -> None: assert litellm.get_model_info("cloudflare/@cf/meta/llama-3.1-8b-instruct-fp8")["rpm"] == 300 assert litellm.get_model_info("cloudflare/@cf/moonshotai/kimi-k2.6")["rpm"] == 20 diff --git a/tests/test_litellm_rust/ocr/test_cohere.py b/tests/test_litellm_rust/ocr/test_cohere.py new file mode 100644 index 00000000000..2a35dc62bd1 --- /dev/null +++ b/tests/test_litellm_rust/ocr/test_cohere.py @@ -0,0 +1,141 @@ +from typing import Final + +import pytest + +import litellm +from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec + +pytestmark = pytest.mark.requires_rust_extension +MODELS: Final = ("cohere/parse-v5.0", "azure_ai/Cohere-parse-v5.0") +IMAGE: Final = {"type": "image_url", "image_url": "data:image/png;base64,YWJj"} +BOX: Final = {"top_left_x": 0, "top_left_y": 0, "bottom_right_x": 32, "bottom_right_y": 32} +PAYLOAD: Final = { + "pages": [ + { + "index": 4, + "markdown": {"content": "receipt", "images": [{"id": "image", "bounding_box": BOX, "description": "scan"}]}, + }, + {"markdown": {"content": "page two"}}, + ], + "meta": {"billed_units": {"pages": 3}}, +} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", MODELS) +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_public_cohere_request_and_normalization( + recording_server: RecordingServer, model: str, asynchronous: bool +) -> None: + recording_server.enqueue(ResponseSpec(body=PAYLOAD)) + args: Final = { + "model": model, + "document": IMAGE, + "api_base": recording_server.base_url, + "api_key": "test-key", + "req_format": "native", + "unrecognized": True, + } + response: Final = await litellm.aocr(**args) if asynchronous else litellm.ocr(**args) + request: Final = recording_server.requests[0] + assert request.path == ("/providers/cohere/v2/parse" if model.startswith("azure_ai/") else "/v2/parse") + assert request.headers["authorization"] == "Bearer test-key" + assert request.body == {"model": model.split("/", 1)[1], "document": IMAGE, "output_format": "markdown"} + assert [page.index for page in response.pages] == [4, 1] + assert response.pages[0].markdown == "receipt" + assert response.pages[0].images[0].bbox == BOX + assert response.pages[0].images[0].model_extra["description"] == "scan" + assert response.pages[1].images is None + assert response.usage_info.pages_processed == 3 + assert response.get_provider_native_response() == PAYLOAD + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", MODELS) +async def test_public_cohere_blocks_and_usage_fallback(recording_server: RecordingServer, model: str) -> None: + blocks: Final = [{"type": "text", "text": "total"}] + recording_server.enqueue(ResponseSpec(body={"pages": [{"blocks": blocks}]})) + response: Final = await litellm.aocr( + model=model, document=IMAGE, api_base=recording_server.base_url, api_key="test-key", output_format="blocks" + ) + assert recording_server.requests[0].body["output_format"] == "blocks" + assert response.pages[0].model_extra["blocks"] == blocks + assert response.pages[0].markdown == "" + assert response.usage_info.pages_processed == 1 + assert response.get_provider_native_response() is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", MODELS) +@pytest.mark.parametrize( + "document", + [ + {"type": "document_url", "document_url": "https://example.com/file.pdf"}, + {"type": "image_url", "image_url": "data:application/pdf;base64,YQ=="}, + {"type": "image_url", "image_url": ""}, + ], +) +async def test_public_cohere_rejects_non_images_before_network( + recording_server: RecordingServer, model: str, document: dict[str, str] +) -> None: + recording_server.expected_requests = 0 + with pytest.raises(litellm.BadRequestError, match="only accepts `image_url`"): + await litellm.aocr(model=model, document=document, api_base=recording_server.base_url, api_key="test-key") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", MODELS) +async def test_public_cohere_rejects_unknown_format(recording_server: RecordingServer, model: str) -> None: + recording_server.expected_requests = 0 + with pytest.raises(litellm.BadRequestError, match="output_format"): + await litellm.aocr( + model=model, document=IMAGE, api_base=recording_server.base_url, api_key="test-key", output_format="html" + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", MODELS) +async def test_public_cohere_provider_failure(recording_server: RecordingServer, model: str) -> None: + recording_server.enqueue(ResponseSpec(status=400, body={"message": "output_format must be blocks or markdown"})) + with pytest.raises(litellm.BadRequestError, match="output_format must be") as caught: + await litellm.aocr(model=model, document=IMAGE, api_base=recording_server.base_url, api_key="test-key") + assert caught.value.status_code == 400 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", MODELS) +async def test_public_cohere_health_check(recording_server: RecordingServer, model: str) -> None: + recording_server.enqueue(ResponseSpec(body=PAYLOAD)) + response: Final = await litellm.ahealth_check( + model_params={"model": model, "api_key": "test-key", "api_base": recording_server.base_url}, mode="ocr" + ) + assert "error" not in response + assert recording_server.requests[0].body["document"]["image_url"].startswith("data:image/png;base64,") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("suffix", ["", "/cohere/", "/v2", "/v2/parse"]) +async def test_public_cohere_url_variants(recording_server: RecordingServer, suffix: str) -> None: + recording_server.enqueue(ResponseSpec(body=PAYLOAD)) + await litellm.aocr(model=MODELS[0], document=IMAGE, api_base=recording_server.base_url + suffix, api_key="test-key") + assert recording_server.requests[0].path == ("/cohere/v2/parse" if suffix == "/cohere/" else "/v2/parse") + + +@pytest.mark.asyncio +async def test_public_cohere_environment_key_and_remote_url( + recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("COHERE_API_KEY", "env-key") + recording_server.enqueue(ResponseSpec(body=PAYLOAD)) + document: Final = {"type": "image_url", "image_url": "https://example.com/receipt.png"} + await litellm.aocr(model=MODELS[0], document=document, api_base=recording_server.base_url) + assert recording_server.requests[0].headers["authorization"] == "Bearer env-key" + assert recording_server.requests[0].body["document"] == document + + +@pytest.mark.asyncio +async def test_public_cohere_missing_key(recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("COHERE_API_KEY", raising=False) + recording_server.expected_requests = 0 + with pytest.raises(Exception, match="Missing COHERE_API_KEY"): + await litellm.aocr(model=MODELS[0], document=IMAGE, api_base=recording_server.base_url) diff --git a/tests/test_litellm_rust/ocr/test_dispatch.py b/tests/test_litellm_rust/ocr/test_dispatch.py index 01999b402f1..e9ccad16861 100644 --- a/tests/test_litellm_rust/ocr/test_dispatch.py +++ b/tests/test_litellm_rust/ocr/test_dispatch.py @@ -1,3 +1,4 @@ +import sys from typing import Final import pytest @@ -16,8 +17,26 @@ def ocr_server(recording_server: RecordingServer) -> RecordingServer: return recording_server -def test_public_ocr_uses_native_route_when_enabled(ocr_server: RecordingServer) -> None: - litellm.rust(True) +def test_native_ocr_rejects_disabled_gil_before_provider_call( + ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch +) -> None: + ocr_server.expected_requests = 0 + monkeypatch.setattr(sys, "_is_gil_enabled", lambda: False, raising=False) + + with pytest.raises(RuntimeError, match="native OCR requires the Python GIL"): + litellm.ocr( + model=OCR_MODEL, + document=OCR_DOCUMENT, + api_key="test-key", + api_base=ocr_server.base_url, + ) + + assert not ocr_server.requests + + +@pytest.mark.parametrize("enabled", [False, True, None]) +def test_public_ocr_uses_native_route_independently_of_flag(ocr_server: RecordingServer, enabled: bool | None) -> None: + litellm.rust(enabled) response: Final = litellm.ocr( model=OCR_MODEL, document=OCR_DOCUMENT, @@ -29,18 +48,3 @@ def test_public_ocr_uses_native_route_when_enabled(ocr_server: RecordingServer) assert response.pages[0].markdown == "native OCR response" assert len(ocr_server.requests) == 1 assert not ocr_server.requests[0].headers.get("user-agent", "").startswith("python-httpx") - - -def test_public_ocr_fails_before_network_when_native_is_disabled(ocr_server: RecordingServer) -> None: - litellm.rust(False) - ocr_server.expected_requests = 0 - - with pytest.raises(RuntimeError, match="Rust OCR is unavailable"): - litellm.ocr( - model=OCR_MODEL, - document=OCR_DOCUMENT, - api_key="test-key", - api_base=ocr_server.base_url, - ) - - assert ocr_server.requests == [] diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index bf3692a316f..ccde5e51722 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -1,6 +1,7 @@ import asyncio import datetime import gc +import json import sys import threading import weakref @@ -711,7 +712,9 @@ async def test_reducto_lifecycle_retains_upload_parse_and_post_call_boundaries( @pytest.mark.asyncio -async def test_document_intelligence_post_call_runs_before_polling(ocr_server: RecordingServer) -> None: +async def test_document_intelligence_post_call_observes_submission_and_final_result( + ocr_server: RecordingServer, +) -> None: ocr_server.expected_requests = 2 ocr_server.enqueue( ResponseSpec( @@ -725,7 +728,7 @@ async def test_document_intelligence_post_call_runs_before_polling(ocr_server: R class Observe(Logging): def post_call(self, *args, **kwargs): - boundaries.append(tuple(request.method for request in ocr_server.requests)) + boundaries.append((tuple(request.method for request in ocr_server.requests), kwargs["original_response"])) return super().post_call(*args, **kwargs) logger: Final = Observe( @@ -740,7 +743,9 @@ async def test_document_intelligence_post_call_runs_before_polling(ocr_server: R response: Final = await call_aocr( ocr_server, model="azure_ai/doc-intelligence/prebuilt-read", litellm_logging_obj=logger ) - assert boundaries == [("POST",)] + assert [methods for methods, _ in boundaries] == [("POST",), ("POST", "GET")] + assert json.loads(boundaries[0][1])["status"] == "running" + assert json.loads(boundaries[1][1])["status"] == "succeeded" assert [request.method for request in ocr_server.requests] == ["POST", "GET"] assert ocr_server.requests[1].path == "/operations/1" assert response.pages == []