From 4d903da85c26d0520c075b6730d9e1b446092d70 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Tue, 15 Sep 2026 08:31:17 -0700 Subject: [PATCH] refactor(ocr): mirror Python provider layout --- litellm-rust/crates/core/src/lib.rs | 1 + .../crates/core/src/llms/azure_ai/mod.rs | 1 + .../ocr/cohere_parse_transformation.rs} | 48 +- .../azure_ai/ocr/common_utils.rs} | 13 +- .../azure_ai/ocr/document_intelligence/mod.rs | 3 + .../document_intelligence/transformation.rs | 822 ++++++++++++++++++ .../crates/core/src/llms/azure_ai/ocr/mod.rs | 4 + .../azure_ai/ocr/transformation.rs} | 49 +- .../crates/core/src/llms/base_llm/mod.rs | 1 + .../crates/core/src/llms/base_llm/ocr/mod.rs | 1 + .../src/llms/base_llm/ocr/transformation.rs | 47 + .../crates/core/src/llms/cohere/mod.rs | 1 + .../crates/core/src/llms/cohere/ocr/mod.rs | 3 + .../src/llms/cohere/ocr/transformation.rs | 396 +++++++++ .../crates/core/src/llms/mistral/mod.rs | 1 + .../crates/core/src/llms/mistral/ocr/mod.rs | 3 + .../mistral/ocr}/transformation.rs | 243 +++++- litellm-rust/crates/core/src/llms/mod.rs | 6 + .../crates/core/src/llms/reducto/mod.rs | 1 + .../crates/core/src/llms/reducto/ocr/mod.rs | 3 + .../src/llms/reducto/ocr/transformation.rs | 532 ++++++++++++ .../crates/core/src/llms/vertex_ai/mod.rs | 1 + .../vertex_ai/ocr/common_utils.rs} | 8 +- .../vertex_ai/ocr/deepseek_transformation.rs | 347 ++++++++ .../crates/core/src/llms/vertex_ai/ocr/mod.rs | 3 + .../vertex_ai/ocr/transformation.rs} | 25 +- .../azure/document_intelligence/mod.rs | 214 ----- .../azure/document_intelligence/polling.rs | 119 --- .../crates/core/src/ocr/adapters/cohere.rs | 123 --- .../crates/core/src/ocr/adapters/mistral.rs | 147 ---- .../crates/core/src/ocr/adapters/mod.rs | 91 -- .../core/src/ocr/adapters/reducto/legacy.rs | 45 - .../core/src/ocr/adapters/reducto/mod.rs | 148 ---- .../core/src/ocr/adapters/reducto/v3.rs | 45 - .../core/src/ocr/adapters/vertex/deepseek.rs | 140 --- .../crates/core/src/ocr/codecs/cohere.rs | 254 ------ .../core/src/ocr/codecs/deepseek/mod.rs | 5 - .../src/ocr/codecs/deepseek/transformation.rs | 101 --- .../core/src/ocr/codecs/deepseek/types.rs | 95 -- .../ocr/codecs/document_intelligence/mod.rs | 9 - .../codecs/document_intelligence/params.rs | 219 ----- .../document_intelligence/transformation.rs | 107 --- .../ocr/codecs/document_intelligence/types.rs | 138 --- .../crates/core/src/ocr/codecs/mistral/mod.rs | 5 - .../core/src/ocr/codecs/mistral/types.rs | 60 -- .../crates/core/src/ocr/codecs/mod.rs | 5 - .../crates/core/src/ocr/codecs/reducto/mod.rs | 9 - .../src/ocr/codecs/reducto/transformation.rs | 103 --- .../core/src/ocr/codecs/reducto/types.rs | 128 --- litellm-rust/crates/core/src/ocr/handler.rs | 214 ++++- litellm-rust/crates/core/src/ocr/mod.rs | 10 +- litellm-rust/crates/core/src/ocr/prepare.rs | 4 +- .../crates/core/src/ocr/provider_config.rs | 130 +++ litellm-rust/crates/core/src/ocr/registry.rs | 132 --- litellm-rust/crates/core/src/ocr/types.rs | 10 +- litellm-rust/crates/core/src/ocr/wire.rs | 30 +- .../crates/core/tests/deepseek_ocr.rs | 9 +- litellm-rust/crates/core/tests/reducto_ocr.rs | 2 +- .../crates/core/tests/vertex_ai_ocr.rs | 14 +- 59 files changed, 2817 insertions(+), 2611 deletions(-) create mode 100644 litellm-rust/crates/core/src/llms/azure_ai/mod.rs rename litellm-rust/crates/core/src/{ocr/adapters/azure/cohere.rs => llms/azure_ai/ocr/cohere_parse_transformation.rs} (75%) rename litellm-rust/crates/core/src/{ocr/adapters/azure/mod.rs => llms/azure_ai/ocr/common_utils.rs} (81%) create mode 100644 litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/mod.rs create mode 100644 litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs create mode 100644 litellm-rust/crates/core/src/llms/azure_ai/ocr/mod.rs rename litellm-rust/crates/core/src/{ocr/adapters/azure/mistral.rs => llms/azure_ai/ocr/transformation.rs} (83%) create mode 100644 litellm-rust/crates/core/src/llms/base_llm/mod.rs create mode 100644 litellm-rust/crates/core/src/llms/base_llm/ocr/mod.rs create mode 100644 litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs create mode 100644 litellm-rust/crates/core/src/llms/cohere/mod.rs create mode 100644 litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs create mode 100644 litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs create mode 100644 litellm-rust/crates/core/src/llms/mistral/mod.rs create mode 100644 litellm-rust/crates/core/src/llms/mistral/ocr/mod.rs rename litellm-rust/crates/core/src/{ocr/codecs/mistral => llms/mistral/ocr}/transformation.rs (51%) create mode 100644 litellm-rust/crates/core/src/llms/mod.rs create mode 100644 litellm-rust/crates/core/src/llms/reducto/mod.rs create mode 100644 litellm-rust/crates/core/src/llms/reducto/ocr/mod.rs create mode 100644 litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs create mode 100644 litellm-rust/crates/core/src/llms/vertex_ai/mod.rs rename litellm-rust/crates/core/src/{ocr/adapters/vertex/mod.rs => llms/vertex_ai/ocr/common_utils.rs} (68%) create mode 100644 litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs create mode 100644 litellm-rust/crates/core/src/llms/vertex_ai/ocr/mod.rs rename litellm-rust/crates/core/src/{ocr/adapters/vertex/mistral.rs => llms/vertex_ai/ocr/transformation.rs} (88%) delete mode 100644 litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs delete mode 100644 litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs delete mode 100644 litellm-rust/crates/core/src/ocr/adapters/cohere.rs delete mode 100644 litellm-rust/crates/core/src/ocr/adapters/mistral.rs delete mode 100644 litellm-rust/crates/core/src/ocr/adapters/mod.rs delete mode 100644 litellm-rust/crates/core/src/ocr/adapters/reducto/legacy.rs delete mode 100644 litellm-rust/crates/core/src/ocr/adapters/reducto/mod.rs delete mode 100644 litellm-rust/crates/core/src/ocr/adapters/reducto/v3.rs delete mode 100644 litellm-rust/crates/core/src/ocr/adapters/vertex/deepseek.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/cohere.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/deepseek/mod.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/deepseek/types.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/document_intelligence/mod.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/document_intelligence/params.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/document_intelligence/types.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/mod.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/reducto/mod.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/reducto/transformation.rs delete mode 100644 litellm-rust/crates/core/src/ocr/codecs/reducto/types.rs create mode 100644 litellm-rust/crates/core/src/ocr/provider_config.rs delete mode 100644 litellm-rust/crates/core/src/ocr/registry.rs diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index f1d888ec12c..5f4bc84b4a8 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -6,6 +6,7 @@ pub mod chat_completions; pub mod constants; pub mod error; pub mod http_utils; +pub(crate) mod llms; mod media; pub mod messages; pub mod ocr; diff --git a/litellm-rust/crates/core/src/llms/azure_ai/mod.rs b/litellm-rust/crates/core/src/llms/azure_ai/mod.rs new file mode 100644 index 00000000000..079e0c41eae --- /dev/null +++ b/litellm-rust/crates/core/src/llms/azure_ai/mod.rs @@ -0,0 +1 @@ +pub(crate) mod ocr; diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs similarity index 75% rename from litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs rename to litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs index 4c8455a171c..0c82e2e3ee4 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs @@ -1,40 +1,42 @@ -use super::super::OcrAdapter; use crate::Error; +use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; +use crate::llms::cohere::ocr::transformation::CohereParseConfig; +use crate::llms::cohere::ocr::{CohereParams, CohereResponse, validate_document}; 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; +#[derive(Default)] +pub(crate) struct AzureAICohereParseConfig { + cohere: CohereParseConfig, +} -impl OcrAdapter for AzureCohereAdapter { +impl BaseOcrConfig for AzureAICohereParseConfig { 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::( + let params = crate::ocr::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 config = AzureAuthInputs { + azure_ad_token_provider: request.azure_ad_token_provider.clone(), + ..AzureAuthInputs::from_sourced_optional_params( + &request.optional_params, + &request.input_sources, + ) + .map_err(Error::from)? + }; let base = request .connection .api_base @@ -46,8 +48,12 @@ impl OcrAdapter for AzureCohereAdapter { "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?; + let headers = super::transformation::validate_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://"); @@ -57,7 +63,9 @@ impl OcrAdapter for AzureCohereAdapter { &request.connection, ) .await?; - let body = transform_request(&request.model, document, params)?; + let body = self + .cohere + .transform_ocr_request(&request.model, document, params)?; transform_request_body( client, request, @@ -76,9 +84,9 @@ impl OcrAdapter for AzureCohereAdapter { fn transform_ocr_response( &self, request: &LiteLLMOcrRequest, - response: Self::ProviderResponse, + response: CohereResponse, ) -> Result { - transform_response(&request.model, response) + self.cohere.transform_ocr_response(request, response) } } diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/common_utils.rs similarity index 81% rename from litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs rename to litellm-rust/crates/core/src/llms/azure_ai/ocr/common_utils.rs index 3d30ae6d6bd..ef11d35da7f 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/common_utils.rs @@ -1,7 +1,3 @@ -mod cohere; -mod document_intelligence; -mod mistral; - use std::sync::OnceLock; use crate::Error; @@ -11,12 +7,7 @@ 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( +pub(super) async fn resolve_entra( config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), ) -> Result>, Error> { @@ -39,7 +30,7 @@ async fn resolve_entra( .map_err(Error::from) } -fn validate_destination( +pub(super) fn validate_destination( connection: &OcrConnection, credential_source: InputSource, ) -> Result<(), OcrError> { diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/mod.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/mod.rs new file mode 100644 index 00000000000..453049a4f2c --- /dev/null +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/mod.rs @@ -0,0 +1,3 @@ +pub(crate) mod transformation; + +pub(crate) use transformation::AzureDocumentIntelligenceOperation; diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs new file mode 100644 index 00000000000..1be4a3934a4 --- /dev/null +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs @@ -0,0 +1,822 @@ +mod types { + use serde::{Deserialize, Deserializer, Serialize}; + use serde_json::{Map, Value}; + + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] + #[serde(untagged)] + pub(crate) enum PagesInput { + ZeroBasedIndices(Vec), + NativeTokens(Vec), + NativeRange(String), + } + + #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] + #[serde(untagged)] + pub(crate) enum FeaturesInput { + Names(Vec), + CommaSeparated(String), + } + + #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] + pub(crate) struct DocumentIntelligenceInputParams { + pub pages: Option, + pub features: Option, + } + + #[derive(Clone, Debug, PartialEq, Serialize)] + pub(crate) struct DocumentIntelligenceParams { + pub pages: Option, + pub features: Option, + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + pub(crate) enum DocumentIntelligenceRequest { + #[serde(rename = "urlSource")] + UrlSource(String), + #[serde(rename = "base64Source")] + Base64Source(String), + } + + #[derive(Clone, Debug, PartialEq)] + pub(crate) enum OperationStatus { + Succeeded, + Running, + NotStarted, + Failed, + Unknown(String), + } + + impl<'de> Deserialize<'de> for OperationStatus { + fn deserialize>(deserializer: D) -> Result { + Ok(match String::deserialize(deserializer)?.as_str() { + "succeeded" => Self::Succeeded, + "running" => Self::Running, + "notStarted" => Self::NotStarted, + "failed" => Self::Failed, + value => Self::Unknown(value.to_string()), + }) + } + } + + impl std::fmt::Display for OperationStatus { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(match self { + Self::Succeeded => "succeeded", + Self::Running => "running", + Self::NotStarted => "notStarted", + Self::Failed => "failed", + Self::Unknown(value) => value, + }) + } + } + + #[derive(Clone, Debug, Deserialize)] + pub(crate) struct AzureDocumentIntelligenceOperation { + pub status: Option, + #[serde(rename = "analyzeResult")] + pub analyze_result: Option, + } + + #[derive(Clone, Debug, Default, Deserialize)] + pub(crate) struct AzureDocumentIntelligenceAnalyzeResult { + pub content: Option, + #[serde(default)] + pub pages: Vec, + pub tables: Option>>, + #[serde(rename = "keyValuePairs")] + pub key_value_pairs: Option>>, + } + + #[derive(Clone, Debug, Deserialize)] + pub(crate) struct AzureDocumentIntelligencePage { + #[serde(rename = "pageNumber", default, deserialize_with = "optional_i64")] + pub page_number: Option, + #[serde(default, deserialize_with = "optional_f64")] + pub width: Option, + #[serde(default, deserialize_with = "optional_f64")] + pub height: Option, + pub unit: Option, + #[serde(default)] + pub lines: Vec, + } + + #[derive(Clone, Debug, Deserialize)] + pub(crate) struct AzureDocumentIntelligenceLine { + pub content: Option, + } + + fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::deserialize(deserializer)? { + None | Some(Value::Null) => Ok(None), + Some(Value::Number(number)) => number + .as_i64() + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected an integer")), + Some(Value::String(value)) => value + .parse::() + .map(Some) + .map_err(|_| serde::de::Error::custom("expected an integer")), + Some(_) => Err(serde::de::Error::custom("expected an integer")), + } + } + + fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::deserialize(deserializer)? { + None | Some(Value::Null) => Ok(None), + Some(Value::Number(number)) => number + .as_f64() + .filter(|value| value.is_finite()) + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected a finite number")), + Some(Value::String(value)) => value + .parse::() + .ok() + .filter(|value| value.is_finite()) + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected a finite number")), + Some(_) => Err(serde::de::Error::custom("expected a number")), + } + } +} + +pub(crate) use types::{ + AzureDocumentIntelligenceOperation, DocumentIntelligenceParams, OperationStatus, +}; + +mod params { + use std::collections::BTreeSet; + + use serde_json::{Map, Value}; + + use super::types::{ + DocumentIntelligenceInputParams, DocumentIntelligenceParams, FeaturesInput, PagesInput, + }; + use crate::ocr::error::OcrRequestError; + use crate::ocr::prepare::ParsedProviderParams; + + pub(crate) fn decode_input_params( + params: Map, + prefix: &str, + ) -> Result, OcrRequestError> { + if let Some(Value::Array(pages)) = params.get("pages") { + if pages.iter().any(Value::is_boolean) { + return Err(OcrRequestError::Pages("boolean page index".into())); + } + if pages + .iter() + .any(|page| page.is_number() && page.as_i64().is_none()) + { + return Err(OcrRequestError::Pages("page index is out of range".into())); + } + if !pages.iter().all(Value::is_i64) && !pages.iter().all(Value::is_string) { + return Err(OcrRequestError::Pages("mixed page element types".into())); + } + } + crate::ocr::wire::decode_request_value(Value::Object(params), prefix) + } + + pub(crate) fn map_ocr_params( + params: DocumentIntelligenceInputParams, + ) -> Result { + Ok(DocumentIntelligenceParams { + pages: params.pages.map(normalize_pages).transpose()?.flatten(), + features: params + .features + .map(normalize_features) + .transpose()? + .flatten(), + }) + } + + fn normalize_pages(pages: PagesInput) -> Result, OcrRequestError> { + let normalized = match pages { + PagesInput::ZeroBasedIndices(indices) => { + if indices.is_empty() { + return Ok(None); + } + indices + .into_iter() + .map(|page| { + if page < 0 { + return Err(OcrRequestError::Pages("negative page index".into())); + } + page.checked_add(1).ok_or_else(|| { + OcrRequestError::Pages("page index is out of range".into()) + }) + }) + .collect::, _>>()? + .into_iter() + .map(|page| page.to_string()) + .collect::>() + .join(",") + } + PagesInput::NativeTokens(tokens) => { + if tokens.is_empty() { + return Ok(None); + } + tokens + .iter() + .map(|token| token.trim()) + .collect::>() + .join(",") + } + PagesInput::NativeRange(range) => range + .split(',') + .map(str::trim) + .collect::>() + .join(","), + }; + if !normalized.split(',').all(valid_page_token) { + return Err(OcrRequestError::Pages("invalid native page range".into())); + } + Ok(Some(normalized)) + } + + fn valid_page_token(token: &str) -> bool { + let mut parts = token.split('-'); + let start = parts.next().unwrap_or_default(); + if start.is_empty() || !start.chars().all(|character| character.is_ascii_digit()) { + return false; + } + match parts.next() { + None => true, + Some(end) => { + !end.is_empty() + && end.chars().all(|character| character.is_ascii_digit()) + && parts.next().is_none() + } + } + } + + fn normalize_features(features: FeaturesInput) -> Result, OcrRequestError> { + let tokens = match features { + FeaturesInput::Names(names) => names, + FeaturesInput::CommaSeparated(names) => names.split(',').map(str::to_string).collect(), + }; + if tokens.is_empty() { + return Ok(None); + } + let normalized = tokens.iter().map(|token| token.trim()).collect::>(); + if !normalized.iter().all(|token| { + let Some((first, rest)) = token.as_bytes().split_first() else { + return false; + }; + first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric) + }) { + return Err(OcrRequestError::Features); + } + Ok(Some(normalized.join(","))) + } + + #[cfg(test)] + mod tests { + use rstest::rstest; + use serde_json::{Value, json}; + + use super::*; + + fn map(value: Value) -> Result { + let fields = value.as_object().unwrap().clone(); + map_ocr_params(decode_input_params(fields, "optional_params")?.known) + } + + #[test] + fn input_params_retain_unknown_fields() { + let parsed = decode_input_params( + json!({ + "pages": [0], + "future_ocr_option": true, + "extra_body": {"provider_option": "value"} + }) + .as_object() + .unwrap() + .clone(), + "optional_params", + ) + .unwrap(); + + assert_eq!( + parsed.known.pages, + Some(PagesInput::ZeroBasedIndices(vec![0])) + ); + assert_eq!(parsed.extra_params["future_ocr_option"], true); + assert_eq!( + parsed.extra_params["extra_body"], + json!({"provider_option": "value"}) + ); + assert_eq!( + serde_json::to_value(map_ocr_params(parsed.known).unwrap()).unwrap(), + json!({"pages": "1", "features": null}) + ); + } + + #[rstest] + #[case(json!([0, 1, 2]), Some("1,2,3"))] + #[case(json!([2, 0, 0, 1]), Some("1,2,3"))] + #[case(json!([]), None)] + #[case(json!("3-9"), Some("3-9"))] + #[case(json!("1-3, 5"), Some("1-3,5"))] + #[case(json!(["1", "3-5"]), Some("1,3-5"))] + fn page_mapping_matches_python(#[case] input: Value, #[case] expected: Option<&str>) { + assert_eq!( + map(json!({"pages": input})).unwrap().pages.as_deref(), + expected + ); + } + + #[rstest] + #[case(json!("a,b"))] + #[case(json!([-1]))] + #[case(json!([true, false]))] + #[case(json!([1, "2"]))] + #[case(json!(5))] + fn invalid_page_mapping_matches_python(#[case] input: Value) { + assert!(map(json!({"pages": input})).is_err()); + } + + #[rstest] + #[case(json!(["keyValuePairs"]), "keyValuePairs")] + #[case(json!(["keyValuePairs", "languages"]), "keyValuePairs,languages")] + #[case(json!("keyValuePairs"), "keyValuePairs")] + #[case(json!("keyValuePairs,languages"), "keyValuePairs,languages")] + #[case(json!("keyValuePairs, languages"), "keyValuePairs,languages")] + fn feature_mapping_matches_python(#[case] input: Value, #[case] expected: &str) { + assert_eq!( + map(json!({"features": input})).unwrap().features.as_deref(), + Some(expected) + ); + } + + #[rstest] + #[case(json!("keyValuePairs&pages=9"))] + #[case(json!("key value pairs"))] + #[case(json!(""))] + #[case(json!([1, 2]))] + #[case(json!([["keyValuePairs"]]))] + #[case(json!({"feature":"keyValuePairs"}))] + #[case(json!(5))] + fn invalid_feature_mapping_matches_python(#[case] input: Value) { + assert!(map(json!({"features": input})).is_err()); + } + + #[test] + fn empty_feature_list_is_omitted() { + assert_eq!(map(json!({"features": []})).unwrap().features, None); + } + } +} + +mod mapping { + use base64::{Engine, engine::general_purpose::STANDARD}; + use serde_json::{Map, Value, json}; + + use super::types::*; + use crate::constants::{AZURE_DI_DEFAULT_DPI, AZURE_DI_DEFAULT_HEIGHT, AZURE_DI_DEFAULT_WIDTH}; + use crate::ocr::document::InlineDocument; + use crate::ocr::error::{OcrRequestError, OcrResponseError}; + use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; + + #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] + pub(crate) fn transform_ocr_request( + document: OcrDocument, + ) -> Result { + let source = document.source(); + if source.is_empty() { + return Err(OcrRequestError::MissingDocumentUrl); + } + Ok(if let Some(document) = InlineDocument::parse(source)? { + DocumentIntelligenceRequest::Base64Source( + STANDARD.encode(document.decode(crate::constants::OCR_INLINE_MAX_BYTES)?), + ) + } else { + DocumentIntelligenceRequest::UrlSource(source.to_string()) + }) + } + + pub(crate) fn transform_ocr_response( + model: &str, + response: AzureDocumentIntelligenceOperation, + ) -> Result { + if response.status != Some(OperationStatus::Succeeded) { + return Err(OcrResponseError::OperationStatus( + response + .status + .map(|status| status.to_string()) + .unwrap_or_else(|| "None".into()), + )); + } + let result = response.analyze_result.unwrap_or_default(); + let pages = result + .pages + .into_iter() + .map(normalize_page) + .collect::, _>>()?; + let pages_processed = pages.len(); + let mut extra_fields = Map::new(); + extra_fields.insert("content".into(), option_value(result.content)); + extra_fields.insert("tables".into(), option_value(result.tables)); + extra_fields.insert("keyValuePairs".into(), option_value(result.key_value_pairs)); + Ok(LiteLLMOcrResponse { + pages, + model: model.into(), + document_annotation: None, + usage_info: Some(json!({"pages_processed":pages_processed})), + object: "ocr".into(), + extra_fields, + provider_native_response: None, + }) + } + + fn normalize_page(page: AzureDocumentIntelligencePage) -> Result { + let index = page + .page_number + .unwrap_or(1) + .checked_sub(1) + .ok_or(OcrResponseError::NumericRange("page.pageNumber"))?; + let scale = if page.unit.as_deref().unwrap_or("inch") == "inch" { + AZURE_DI_DEFAULT_DPI as f64 + } else { + 1.0 + }; + let width = pixel_dimension( + page.width.unwrap_or(AZURE_DI_DEFAULT_WIDTH), + scale, + "page.width", + )?; + let height = pixel_dimension( + page.height.unwrap_or(AZURE_DI_DEFAULT_HEIGHT), + scale, + "page.height", + )?; + let markdown = page + .lines + .iter() + .map(|line| line.content.as_deref().unwrap_or_default()) + .collect::>() + .join("\n"); + Ok(json!({ + "index":index, + "markdown":markdown, + "images":null, + "dimensions":{"width":width,"height":height,"dpi":AZURE_DI_DEFAULT_DPI} + })) + } + + fn pixel_dimension( + value: f64, + scale: f64, + field: &'static str, + ) -> Result { + let value = value * scale; + if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 { + return Err(OcrResponseError::NumericRange(field)); + } + Ok(value.trunc() as i64) + } + + fn option_value(value: Option) -> Value { + value + .and_then(|value| serde_json::to_value(value).ok()) + .unwrap_or(Value::Null) + } +} + +mod polling { + use std::sync::Arc; + use std::time::Duration; + + use reqwest::Url; + use tokio::time::Instant; + + use super::{AzureDocumentIntelligenceOperation, OperationStatus}; + use crate::constants::{AZURE_DI_SUBSCRIPTION_HEADER, OCR_POLL_RETRY_SECS}; + use crate::ocr::client::read_json_response; + use crate::ocr::error::{OcrError, OcrPollingError, OcrResponseError}; + use crate::ocr::hooks::OcrHooks; + use crate::ocr::types::OcrConnection; + use crate::ocr::wire::DecodedOcrResponse; + + pub(super) async fn read_operation_response( + http_client: &reqwest::Client, + response: reqwest::Response, + original_url: &str, + headers: &[(String, String)], + connection: &OcrConnection, + native: bool, + hooks: &Arc, + ) -> Result, OcrError> { + if response.status() != reqwest::StatusCode::ACCEPTED { + let bytes = + crate::ocr::client::read_response_bytes(response, connection.max_response_bytes) + .await?; + crate::ocr::handler::post_call(hooks, &bytes).await?; + return Ok(crate::ocr::wire::decode_response(&bytes, native)?); + } + let location = response + .headers() + .get("operation-location") + .and_then(|value| value.to_str().ok()) + .ok_or(OcrPollingError::PollLocation)? + .to_string(); + let original = Url::parse(original_url).map_err(|_| OcrPollingError::PollOrigin)?; + let operation = Url::parse(&location).map_err(|_| OcrPollingError::PollOrigin)?; + if original.origin() != operation.origin() + || !operation.username().is_empty() + || operation.password().is_some() + { + return Err(OcrPollingError::PollOrigin.into()); + } + let bytes = + crate::ocr::client::read_response_bytes(response, connection.max_response_bytes) + .await?; + crate::ocr::handler::post_call(hooks, &bytes).await?; + poll_operation(http_client, operation, headers, connection, native, hooks).await + } + + async fn poll_operation( + http_client: &reqwest::Client, + url: Url, + headers: &[(String, String)], + connection: &OcrConnection, + native: bool, + hooks: &Arc, + ) -> Result, OcrError> { + let deadline = Instant::now() + .checked_add(connection.poll_timeout) + .ok_or(OcrPollingError::PollTimeout)?; + loop { + let remaining = deadline + .checked_duration_since(Instant::now()) + .filter(|remaining| !remaining.is_zero()) + .ok_or(OcrPollingError::PollTimeout)?; + let builder = http_client + .get(url.clone()) + .timeout(remaining.min(connection.timeout)); + let builder = crate::http_utils::with_headers( + builder, + headers, + crate::http_utils::HeaderPolicy::Only(&[ + AZURE_DI_SUBSCRIPTION_HEADER, + "authorization", + ]), + ); + let response = + tokio::time::timeout_at(deadline, crate::http_utils::http_request(builder)) + .await + .map_err(|_| OcrPollingError::PollTimeout)? + .map_err(crate::error::TransportError::from)?; + let retry = response + .headers() + .get(reqwest::header::RETRY_AFTER) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.parse::().ok()) + .unwrap_or(OCR_POLL_RETRY_SECS) + .max(1); + let decoded = tokio::time::timeout_at( + deadline, + read_json_response::( + response, + native, + connection.max_response_bytes, + ), + ) + .await + .map_err(|_| OcrPollingError::PollTimeout)??; + match &decoded.data.status { + 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 + .map_err(|_| OcrPollingError::PollTimeout)?; + } + status => { + return Err(OcrResponseError::OperationStatus( + status + .as_ref() + .map(ToString::to_string) + .unwrap_or_else(|| "None".into()), + ) + .into()); + } + } + } + } +} + +use crate::Error; +use crate::auth::{InputSource, Sourced}; +use crate::constants::{AZURE_DI_API_VERSION, AZURE_DI_SUBSCRIPTION_HEADER}; +use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; +use crate::ocr::OcrClient; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrResponseFormat}; +use crate::providers::azure_ai::auth::AzureAuthInputs; +use crate::url_utils::ApiUrl; + +const AZURE_DI_API_KEY_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"; +const AZURE_DI_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT"; + +#[derive(Clone, Debug)] +pub(crate) struct AzureDocumentIntelligenceOCRConfig; + +impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { + type ProviderResponse = AzureDocumentIntelligenceOperation; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let params = map_ocr_params(request)?; + let config = AzureAuthInputs { + azure_ad_token_provider: request.azure_ad_token_provider.clone(), + ..AzureAuthInputs::from_sourced_optional_params( + &request.optional_params, + &request.input_sources, + ) + .map_err(Error::from)? + }; + let headers = validate_environment(&request.connection, &config, &credential_env).await?; + let endpoint = nonblank(request.connection.api_base.clone()) + .or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV))) + .ok_or_else(|| Error::Auth("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into()))?; + let url = get_complete_url(&endpoint, &request.model, ¶ms)?; + let body = mapping::transform_ocr_request(request.document.clone())?; + transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: AzureDocumentIntelligenceOperation, + ) -> Result { + mapping::transform_ocr_response(&request.model, response) + } + + async fn read_response( + &self, + client: &OcrClient, + response: reqwest::Response, + url: &str, + headers: &[(String, String)], + request: &LiteLLMOcrRequest, + ) -> Result, OcrError> + { + polling::read_operation_response( + client.polling_http(), + response, + url, + headers, + &request.connection, + request.response_format()? == OcrResponseFormat::Native, + &request.hooks, + ) + .await + } +} + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +fn map_ocr_params( + request: &LiteLLMOcrRequest, +) -> Result { + let params = params::decode_input_params(request.optional_params.clone(), "optional_params")?; + let crate::ocr::prepare::ParsedProviderParams { + known: params, + extra_params: _extra_params, + } = params; + params::map_ocr_params(params) +} + +fn get_complete_url( + endpoint: &str, + model: &str, + params: &DocumentIntelligenceParams, +) -> Result { + let model = format!("{}:analyze", model_id(model)?); + ApiUrl::parse(endpoint) + .and_then(|url| url.complete_path(&["documentintelligence", "documentModels", &model])) + .map(|url| { + url.append_query_pairs( + [("api-version", AZURE_DI_API_VERSION)] + .into_iter() + .chain(params.pages.iter().map(|pages| ("pages", pages.as_str()))) + .chain( + params + .features + .iter() + .map(|features| ("features", features.as_str())), + ), + ) + .into_string() + }) + .map_err(|_| OcrRequestError::RequestField { + path: "api_base".into(), + }) + .map_err(OcrError::from) +} + +async fn validate_environment( + connection: &OcrConnection, + config: &AzureAuthInputs, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result, OcrError> { + if crate::http_utils::has_header(&connection.extra_headers, "authorization") + || crate::http_utils::has_header(&connection.extra_headers, AZURE_DI_SUBSCRIPTION_HEADER) + { + super::super::common_utils::validate_destination( + connection, + connection.extra_headers_source, + )?; + return Ok(connection.extra_headers.clone()); + } + let key = nonblank(connection.api_key.clone()) + .map(|value| Sourced::new(value, connection.api_key_source)) + .or_else(|| { + nonblank(env_lookup(AZURE_DI_API_KEY_ENV)) + .map(|value| Sourced::new(value, InputSource::Environment)) + }); + if let Some(key) = key { + super::super::common_utils::validate_destination(connection, key.source())?; + return Ok( + std::iter::once((AZURE_DI_SUBSCRIPTION_HEADER.into(), key.into_value())) + .chain(connection.extra_headers.clone()) + .collect(), + ); + } + let token = super::super::common_utils::resolve_entra(config, env_lookup) + .await? + .ok_or(Error::MissingAzureDocumentIntelligenceCredentials)?; + super::super::common_utils::validate_destination(connection, token.source())?; + Ok( + std::iter::once(("Authorization".into(), format!("Bearer {}", token.value()))) + .chain(connection.extra_headers.clone()) + .collect(), + ) +} + +fn model_id(model: &str) -> Result<&str, OcrRequestError> { + let model = model.rsplit('/').next().unwrap_or(model); + if matches!(model, "." | "..") { + return Err(OcrRequestError::DotModel); + } + Ok(model) +} + +fn nonblank(value: Option) -> Option { + value + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn request_endpoint_cannot_receive_environment_key() { + let connection = OcrConnection { + api_base: Some("https://request.example".into()), + api_base_source: InputSource::Request, + ..Default::default() + }; + + let error = validate_environment(&connection, &Default::default(), &|name| { + (name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into()) + }) + .await + .unwrap_err(); + + assert!( + error + .to_string() + .contains("request-controlled Azure endpoint") + ); + } + + #[tokio::test] + async fn request_endpoint_accepts_request_owned_key() { + let connection = OcrConnection { + api_key: Some("request-key".into()), + api_key_source: InputSource::Request, + api_base: Some("https://request.example".into()), + api_base_source: InputSource::Request, + ..Default::default() + }; + + let headers = validate_environment(&connection, &Default::default(), &|_| None) + .await + .unwrap(); + + assert_eq!( + headers[0], + (AZURE_DI_SUBSCRIPTION_HEADER.into(), "request-key".into()) + ); + } +} diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/mod.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/mod.rs new file mode 100644 index 00000000000..e106f50b0a7 --- /dev/null +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/mod.rs @@ -0,0 +1,4 @@ +pub(crate) mod cohere_parse_transformation; +pub(crate) mod common_utils; +pub(crate) mod document_intelligence; +pub(crate) mod transformation; diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs similarity index 83% rename from litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs rename to litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs index 8639590b05c..b87b5a6fe66 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs @@ -1,15 +1,15 @@ -use super::super::OcrAdapter; use crate::Error; use crate::auth::{InputSource, Sourced}; use crate::constants::AZURE_AI_OCR_PATH; +use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; +use crate::llms::mistral::ocr::transformation::MistralOCRConfig; +use crate::llms::mistral::ocr::{MistralOcrParams, MistralOcrResponse}; use crate::ocr::OcrClient; -use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse}; use crate::ocr::document::{inline_remote_document, validate_inline_document}; use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; use crate::ocr::prepare::{ _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, }; -use crate::ocr::registry::OcrProvider; use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection}; use crate::providers::azure_ai::auth::AzureAuthInputs; use crate::url_utils::ApiUrl; @@ -17,12 +17,13 @@ use crate::url_utils::ApiUrl; const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY"; const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE"; -#[derive(Clone, Debug)] -pub(crate) struct AzureMistralAdapter; +#[derive(Clone, Debug, Default)] +pub(crate) struct AzureAIOCRConfig { + mistral: MistralOCRConfig, +} -impl OcrAdapter for AzureMistralAdapter { +impl BaseOcrConfig for AzureAIOCRConfig { type ProviderResponse = MistralOcrResponse; - const PROVIDER: OcrProvider = OcrProvider::AzureAi; async fn prepare_request( &self, @@ -33,12 +34,14 @@ impl OcrAdapter for AzureMistralAdapter { known: params, extra_params: _extra_params, } = _prepare_ocr_request::(request)?; - 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 config = AzureAuthInputs { + azure_ad_token_provider: request.azure_ad_token_provider.clone(), + ..AzureAuthInputs::from_sourced_optional_params( + &request.optional_params, + &request.input_sources, + ) + .map_err(Error::from)? + }; let url = get_complete_url(request.connection.api_base.as_deref(), &credential_env)?; let headers = validate_environment(&request.connection, &config, &credential_env).await?; let retains_document = !request.document.source().starts_with("http://") @@ -49,7 +52,9 @@ impl OcrAdapter for AzureMistralAdapter { &request.connection, ) .await?; - let body = mistral::transform_ocr_request(&request.model, document, ¶ms)?; + let body = self + .mistral + .transform_ocr_request(&request.model, document, ¶ms)?; transform_request_body( client, request, @@ -65,9 +70,9 @@ impl OcrAdapter for AzureMistralAdapter { fn transform_ocr_response( &self, request: &LiteLLMOcrRequest, - response: Self::ProviderResponse, + response: MistralOcrResponse, ) -> Result { - mistral::transform_ocr_response(&request.model, response) + self.mistral.transform_ocr_response(request, response) } } @@ -92,16 +97,16 @@ fn get_complete_url( }) } -pub(in crate::ocr::adapters) async fn validate_environment( +pub(super) async fn validate_environment( connection: &OcrConnection, config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), ) -> Result, OcrError> { if crate::http_utils::has_header(&connection.extra_headers, "authorization") { if config.azure_ad_token_provider.is_some() { - super::resolve_entra(config, env_lookup).await?; + super::common_utils::resolve_entra(config, env_lookup).await?; } - super::validate_destination(connection, connection.extra_headers_source)?; + super::common_utils::validate_destination(connection, connection.extra_headers_source)?; return Ok(connection.extra_headers.clone()); } let key = nonblank(connection.api_key.clone()) @@ -111,13 +116,13 @@ pub(in crate::ocr::adapters) async fn validate_environment( .map(|value| Sourced::new(value, InputSource::Environment)) }); if let Some(key) = key { - super::validate_destination(connection, key.source())?; + super::common_utils::validate_destination(connection, key.source())?; return Ok(bearer_headers(connection, key.value())); } - let key = super::resolve_entra(config, env_lookup) + let key = super::common_utils::resolve_entra(config, env_lookup) .await? .ok_or(Error::MissingAzureAiCredentials)?; - super::validate_destination(connection, key.source())?; + super::common_utils::validate_destination(connection, key.source())?; Ok(bearer_headers(connection, key.value())) } diff --git a/litellm-rust/crates/core/src/llms/base_llm/mod.rs b/litellm-rust/crates/core/src/llms/base_llm/mod.rs new file mode 100644 index 00000000000..079e0c41eae --- /dev/null +++ b/litellm-rust/crates/core/src/llms/base_llm/mod.rs @@ -0,0 +1 @@ +pub(crate) mod ocr; diff --git a/litellm-rust/crates/core/src/llms/base_llm/ocr/mod.rs b/litellm-rust/crates/core/src/llms/base_llm/ocr/mod.rs new file mode 100644 index 00000000000..080f0a1183f --- /dev/null +++ b/litellm-rust/crates/core/src/llms/base_llm/ocr/mod.rs @@ -0,0 +1 @@ +pub(crate) mod transformation; diff --git a/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs new file mode 100644 index 00000000000..66c5c5133ae --- /dev/null +++ b/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs @@ -0,0 +1,47 @@ +use std::future::Future; + +use serde::de::DeserializeOwned; + +use crate::ocr::OcrClient; +use crate::ocr::error::{OcrError, OcrResponseError}; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrResponseFormat}; +use crate::ocr::wire::DecodedOcrResponse; + +pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { + type ProviderResponse: DeserializeOwned + Send; + + fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> impl Future> + Send; + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result; + + fn read_response( + &self, + _client: &OcrClient, + response: reqwest::Response, + _url: &str, + _headers: &[(String, String)], + request: &LiteLLMOcrRequest, + ) -> impl Future, OcrError>> + Send + { + async move { + let bytes = crate::ocr::client::read_response_bytes( + response, + request.connection.max_response_bytes, + ) + .await?; + crate::ocr::handler::post_call(&request.hooks, &bytes).await?; + Ok(crate::ocr::wire::decode_response( + &bytes, + request.response_format()? == OcrResponseFormat::Native, + )?) + } + } +} diff --git a/litellm-rust/crates/core/src/llms/cohere/mod.rs b/litellm-rust/crates/core/src/llms/cohere/mod.rs new file mode 100644 index 00000000000..079e0c41eae --- /dev/null +++ b/litellm-rust/crates/core/src/llms/cohere/mod.rs @@ -0,0 +1 @@ +pub(crate) mod ocr; diff --git a/litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs b/litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs new file mode 100644 index 00000000000..477e9afd141 --- /dev/null +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs @@ -0,0 +1,3 @@ +pub(crate) mod transformation; + +pub(crate) use transformation::{CohereParams, CohereResponse, validate_document}; diff --git a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs new file mode 100644 index 00000000000..29a50de0d9d --- /dev/null +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -0,0 +1,396 @@ +mod provider { + 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": [ + { + "type":"markdown", + "index":4, + "markdown":{ + "content":"receipt", + "images":[{ + "id":"image", + "bounding_box":{"top_left_x":1,"bottom_right_x":48}, + "bounding_box_normalized":{"top_left_x":0.04,"bottom_right_x":0.15}, + "description":"scan", + "category":"logo" + }] + } + }, + {"type":"blocks","blocks":[{"type":"text","text":{"content":"total"}}]} + ], + "meta":{"api_version":{"version":"2"},"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]["bounding_box_normalized"]["bottom_right_x"], + 0.15 + ); + assert_eq!(normalized.pages[0]["images"][0]["description"], "scan"); + assert_eq!(normalized.pages[0]["images"][0]["category"], "logo"); + assert_eq!(normalized.pages[1]["index"], 1); + assert_eq!(normalized.pages[1]["markdown"], ""); + assert_eq!(normalized.pages[1]["blocks"][0]["text"]["content"], "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" + ); + } + } +} + +pub(crate) use provider::{ + CohereParams, CohereRequest, CohereResponse, transform_request, transform_response, + validate_document, +}; + +use crate::Error; +use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE}; +use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; +use crate::ocr::OcrClient; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection}; +use crate::url_utils::ApiUrl; + +#[derive(Default)] +pub(crate) struct CohereParseConfig; + +impl CohereParseConfig { + pub(crate) fn transform_ocr_request( + &self, + model: &str, + document: crate::ocr::types::OcrDocument, + params: CohereParams, + ) -> Result { + transform_request(model, document, params) + } +} + +impl BaseOcrConfig for CohereParseConfig { + type ProviderResponse = CohereResponse; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let params = crate::ocr::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 = self.transform_ocr_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: CohereResponse, + ) -> 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/llms/mistral/mod.rs b/litellm-rust/crates/core/src/llms/mistral/mod.rs new file mode 100644 index 00000000000..079e0c41eae --- /dev/null +++ b/litellm-rust/crates/core/src/llms/mistral/mod.rs @@ -0,0 +1 @@ +pub(crate) mod ocr; diff --git a/litellm-rust/crates/core/src/llms/mistral/ocr/mod.rs b/litellm-rust/crates/core/src/llms/mistral/ocr/mod.rs new file mode 100644 index 00000000000..546ac9eb50a --- /dev/null +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/mod.rs @@ -0,0 +1,3 @@ +pub(crate) mod transformation; + +pub(crate) use transformation::{MistralOcrParams, MistralOcrResponse}; diff --git a/litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs similarity index 51% rename from litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs rename to litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs index e8073905548..bfe70588b2b 100644 --- a/litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs @@ -1,7 +1,65 @@ -use super::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse}; +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; + use crate::ocr::error::{OcrRequestError, OcrResponseError}; use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub(crate) enum MistralOcrPages { + Range(String), + Indices(Vec), +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub(crate) struct MistralOcrParams { + #[serde(skip_serializing_if = "Option::is_none")] + pub pages: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub include_image_base64: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub image_limit: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub image_min_size: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub bbox_annotation_format: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub document_annotation_format: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub document_annotation_prompt: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub extract_header: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub extract_footer: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub table_format: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub confidence_scores_granularity: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub include_blocks: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub id: Option, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct MistralOcrRequest { + pub model: String, + pub document: OcrDocument, + #[serde(flatten)] + pub params: MistralOcrParams, +} + +#[derive(Clone, Debug, Default, Deserialize)] +pub(crate) struct MistralOcrResponse { + #[serde(default)] + pub pages: Vec, + pub model: Option, + pub document_annotation: Option, + pub usage_info: Option, + #[serde(flatten)] + pub extra_fields: Map, +} +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub(crate) fn transform_ocr_request( model: &str, document: OcrDocument, @@ -30,7 +88,7 @@ pub(crate) fn transform_ocr_response( } #[cfg(test)] -mod tests { +mod mapping_tests { use super::*; use rstest::rstest; use serde_json::{Value, json}; @@ -197,19 +255,19 @@ mod tests { #[rstest] fn transform_ocr_response_preserves_blocks_and_confidence_scores() { let response: MistralOcrResponse = serde_json::from_value(json!({ - "pages":[{ - "index":0, - "markdown":"hello", - "images":[{"id":"img-0","image_base64":"data:image/png;base64,AA=="}], - "dimensions":{"width":612,"height":792,"dpi":72}, - "blocks":[{"type":"title","bbox":{"x":1},"confidence_scores":{"mean":0.98}}], - "confidence_scores":{"average_page_confidence_score":0.99,"minimum_page_confidence_score":0.97} - }], - "model":"returned-model", - "document_annotation":"{\"language\":\"en\"}", - "usage_info":{"pages_processed":1} - })) - .unwrap(); + "pages":[{ + "index":0, + "markdown":"hello", + "images":[{"id":"img-0","image_base64":"data:image/png;base64,AA=="}], + "dimensions":{"width":612,"height":792,"dpi":72}, + "blocks":[{"type":"title","bbox":{"x":1},"confidence_scores":{"mean":0.98}}], + "confidence_scores":{"average_page_confidence_score":0.99,"minimum_page_confidence_score":0.97} + }], + "model":"returned-model", + "document_annotation":"{\"language\":\"en\"}", + "usage_info":{"pages_processed":1} + })) + .unwrap(); let result = transform_ocr_response("model", response) .unwrap() .into_json(); @@ -248,3 +306,158 @@ mod tests { assert_eq!(result["pages"][0], page); } } + +use crate::Error; +use crate::constants::MISTRAL_OCR_API_BASE; +use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; +use crate::ocr::OcrClient; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, +}; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection}; +use crate::url_utils::ApiUrl; + +const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY"; + +#[derive(Clone, Debug, Default)] +pub(crate) struct MistralOCRConfig; + +impl MistralOCRConfig { + pub(crate) fn transform_ocr_request( + &self, + model: &str, + document: crate::ocr::types::OcrDocument, + params: &MistralOcrParams, + ) -> Result { + transform_ocr_request(model, document, params) + } +} + +impl BaseOcrConfig for MistralOCRConfig { + type ProviderResponse = MistralOcrResponse; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let ParsedProviderParams { + known: params, + extra_params: _extra_params, + } = _prepare_ocr_request::(request)?; + let headers = validate_environment(&request.connection, &credential_env)?; + let url = get_complete_url(request.connection.api_base.as_deref())?; + let body = self.transform_ocr_request(&request.model, request.document.clone(), ¶ms)?; + transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: MistralOcrResponse, + ) -> Result { + transform_ocr_response(&request.model, response) + } +} + +pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result { + let base = api_base + .map(str::trim) + .filter(|base| !base.is_empty()) + .unwrap_or(MISTRAL_OCR_API_BASE); + ApiUrl::parse(base) + .and_then(|url| url.complete_path(&["v1", "ocr"])) + .map(|url| url.into_string()) + .map_err(|_| { + OcrRequestError::RequestField { + path: "api_base".into(), + } + .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 api_key = connection + .api_key + .as_deref() + .map(str::trim) + .filter(|key| !key.is_empty()) + .map(str::to_string) + .or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty())) + .ok_or(Error::MissingApiKey { + provider: "Mistral", + })?; + Ok( + std::iter::once(("Authorization".into(), format!("Bearer {api_key}"))) + .chain(connection.extra_headers.clone()) + .collect(), + ) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn complete_url_defaults_and_dedupes_v1() { + assert_eq!( + get_complete_url(None).unwrap(), + "https://api.mistral.ai/v1/ocr" + ); + assert_eq!( + get_complete_url(Some("https://example.com/v1?tenant=a")).unwrap(), + "https://example.com/v1/ocr?tenant=a" + ); + assert_eq!( + get_complete_url(Some("https://example.com/v1/ocr?tenant=a")).unwrap(), + "https://example.com/v1/ocr?tenant=a" + ); + } + + #[test] + fn environment_prefers_explicit_key_then_environment() { + let explicit = OcrConnection { + api_key: Some("explicit".into()), + ..OcrConnection::default() + }; + assert_eq!( + validate_environment(&explicit, &|_| Some("environment".into())).unwrap()[0], + ("Authorization".into(), "Bearer explicit".into()) + ); + + assert_eq!( + validate_environment(&OcrConnection::default(), &|_| Some("environment".into())) + .unwrap()[0], + ("Authorization".into(), "Bearer environment".into()) + ); + } + + #[test] + fn environment_preserves_forwarded_authorization() { + let connection = OcrConnection { + extra_headers: vec![("authorization".into(), "Bearer forwarded".into())], + ..OcrConnection::default() + }; + assert_eq!( + validate_environment(&connection, &|_| None).unwrap(), + connection.extra_headers + ); + } + + #[test] + fn environment_rejects_missing_key() { + assert!(matches!( + validate_environment(&OcrConnection::default(), &|_| None), + Err(OcrError::Public(Error::MissingApiKey { + provider: "Mistral" + })) + )); + } +} diff --git a/litellm-rust/crates/core/src/llms/mod.rs b/litellm-rust/crates/core/src/llms/mod.rs new file mode 100644 index 00000000000..3dad380f833 --- /dev/null +++ b/litellm-rust/crates/core/src/llms/mod.rs @@ -0,0 +1,6 @@ +pub(crate) mod azure_ai; +pub(crate) mod base_llm; +pub(crate) mod cohere; +pub(crate) mod mistral; +pub(crate) mod reducto; +pub(crate) mod vertex_ai; diff --git a/litellm-rust/crates/core/src/llms/reducto/mod.rs b/litellm-rust/crates/core/src/llms/reducto/mod.rs new file mode 100644 index 00000000000..079e0c41eae --- /dev/null +++ b/litellm-rust/crates/core/src/llms/reducto/mod.rs @@ -0,0 +1 @@ +pub(crate) mod ocr; diff --git a/litellm-rust/crates/core/src/llms/reducto/ocr/mod.rs b/litellm-rust/crates/core/src/llms/reducto/ocr/mod.rs new file mode 100644 index 00000000000..6405500b44a --- /dev/null +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/mod.rs @@ -0,0 +1,3 @@ +pub(crate) mod transformation; + +pub(crate) use transformation::ReductoResponse; diff --git a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs new file mode 100644 index 00000000000..2053e9e9ee3 --- /dev/null +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -0,0 +1,532 @@ +mod types { + use serde::{Deserialize, Deserializer, Serialize}; + use serde_json::{Map, Value}; + + #[derive(Clone, Debug, Default, Serialize, Deserialize)] + pub(crate) struct ReductoV3Params { + #[serde(skip_serializing_if = "Option::is_none")] + pub formatting: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub retrieval: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub settings: Option>, + } + + #[derive(Clone, Debug, Default, Serialize, Deserialize)] + pub(crate) struct ReductoLegacyParams { + #[serde(skip_serializing_if = "Option::is_none")] + pub enhance: Option>, + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + pub(crate) struct ReductoV3Request { + pub input: String, + #[serde(flatten)] + pub params: ReductoV3Params, + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + pub(crate) struct ReductoLegacyRequest { + pub document_url: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub options: Option, + } + + #[derive(Deserialize)] + pub(crate) struct ReductoUploadResponse { + pub file_id: Option, + } + + #[derive(Clone, Debug, Deserialize)] + pub(crate) struct ReductoResponse { + #[serde(default, deserialize_with = "present_nullable")] + pub result: Option>, + pub usage: Option, + #[serde(default)] + pub chunks: Option>, + } + + fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>( + deserializer: D, + ) -> Result>, D::Error> { + Option::::deserialize(deserializer).map(Some) + } + + #[derive(Clone, Debug, Default, Deserialize)] + pub(crate) struct ReductoResult { + pub chunks: Option>, + } + + #[derive(Clone, Debug, Default, Deserialize)] + pub(crate) struct ReductoUsage { + #[serde(default, deserialize_with = "optional_i64")] + pub num_pages: Option, + #[serde(default, deserialize_with = "optional_f64")] + pub credits: Option, + } + + #[derive(Clone, Debug, Deserialize)] + pub(crate) struct ReductoChunk { + pub content: Option, + pub blocks: Option>, + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + pub(crate) struct ReductoBlock { + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub bbox: Option, + #[serde(flatten)] + pub extra_fields: Map, + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + pub(crate) struct ReductoBoundingBox { + #[serde(default, deserialize_with = "optional_i64")] + pub page: Option, + #[serde(flatten)] + pub extra_fields: Map, + } + + fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::deserialize(deserializer)? { + None | Some(Value::Null) => Ok(None), + Some(Value::Number(number)) => number + .as_i64() + .or_else(|| number.as_f64().and_then(checked_truncated_i64)) + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected an integer")), + Some(Value::String(value)) => value + .trim() + .parse::() + .map(Some) + .map_err(|_| serde::de::Error::custom("expected an integer")), + Some(Value::Bool(value)) => Ok(Some(i64::from(value))), + Some(_) => Ok(None), + } + } + + fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::deserialize(deserializer)? { + None | Some(Value::Null) => Ok(None), + Some(Value::Number(number)) => number + .as_f64() + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected a number")), + Some(Value::String(value)) => value + .trim() + .parse::() + .map(Some) + .map_err(|_| serde::de::Error::custom("expected a number")), + Some(_) => Ok(None), + } + } + + fn checked_truncated_i64(value: f64) -> Option { + (value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64) + .then(|| value.trunc() as i64) + } +} + +#[cfg(test)] +pub(crate) use mapping::transform_ocr_response; +pub(crate) use types::ReductoResponse; + +mod mapping { + use std::collections::BTreeMap; + + use serde_json::{Value, json}; + + use super::types::*; + use crate::ocr::error::{OcrRequestError, OcrResponseError}; + use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; + + #[tracing::instrument( + name = "transform_ocr_request", + target = "litellm::function_trace", + level = "trace", + skip_all + )] + pub(crate) fn transform_v3_ocr_request( + _model: &str, + document: OcrDocument, + params: &ReductoV3Params, + ) -> Result { + Ok(ReductoV3Request { + input: document.source().to_string(), + params: params.clone(), + }) + } + + #[tracing::instrument( + name = "transform_ocr_request", + target = "litellm::function_trace", + level = "trace", + skip_all + )] + pub(crate) fn transform_legacy_ocr_request( + _model: &str, + document: OcrDocument, + params: &ReductoLegacyParams, + ) -> Result { + Ok(ReductoLegacyRequest { + document_url: document.source().to_string(), + options: params.enhance.as_ref().map(|_| params.clone()), + }) + } + + pub(crate) fn transform_ocr_response( + model: &str, + response: ReductoResponse, + ) -> Result { + let result = match response.result { + Some(result) => result.unwrap_or_default(), + None => ReductoResult { + chunks: response.chunks, + }, + }; + let usage = response.usage.unwrap_or_default(); + Ok(LiteLLMOcrResponse { + pages: build_pages(result.chunks.unwrap_or_default()), + model: model.to_string(), + document_annotation: None, + usage_info: Some(json!({ + "pages_processed": usage.num_pages, + "credits": usage.credits, + })), + object: "ocr".to_string(), + extra_fields: serde_json::Map::new(), + provider_native_response: None, + }) + } + + fn build_pages(chunks: Vec) -> Vec { + let blocks_by_page = chunks + .iter() + .flat_map(|chunk| chunk.blocks.iter().flatten()) + .filter_map(|block| block.bbox.as_ref()?.page.map(|page| (page, block))) + .fold( + BTreeMap::>::new(), + |mut pages, (page, block)| { + pages.entry(page).or_default().push(block); + pages + }, + ); + if blocks_by_page.is_empty() { + let markdown = join_content(chunks.iter().map(|chunk| chunk.content.as_deref())); + return if markdown.is_empty() { + Vec::new() + } else { + vec![page(0, markdown, None)] + }; + } + blocks_by_page + .into_iter() + .map(|(index, blocks)| { + let markdown = join_content(blocks.iter().map(|block| block.content.as_deref())); + page( + index.saturating_sub(1).max(0), + markdown, + Some(json!(blocks)), + ) + }) + .collect() + } + + fn join_content<'a>(content: impl Iterator>) -> String { + content + .flatten() + .filter(|text| !text.is_empty()) + .collect::>() + .join("\n\n") + } + + fn page(index: i64, markdown: String, blocks: Option) -> Value { + let mut result = json!({"index":index,"markdown":markdown,"images":null}); + if let (Value::Object(fields), Some(blocks)) = (&mut result, blocks) { + fields.insert("blocks".into(), blocks); + } + result + } +} + +mod common { + use crate::Error; + use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX}; + use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; + use crate::ocr::document::InlineDocument; + use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; + use crate::ocr::types::{OcrConnection, OcrDocument}; + use crate::url_utils::ApiUrl; + + pub(crate) trait BaseReductoOcrConfig: BaseOcrConfig { + fn get_complete_url(&self, api_base: Option<&str>, path: &str) -> Result { + get_complete_url(api_base, path) + } + + fn validate_environment( + &self, + connection: &OcrConnection, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result, OcrError> { + validate_environment(connection, env_lookup) + } + + async fn prepare_document( + &self, + client: &crate::ocr::OcrClient, + document: OcrDocument, + connection: &OcrConnection, + headers: &[(String, String)], + ) -> Result { + prepare_document(client, document, connection, headers).await + } + } + + pub(super) fn get_complete_url(api_base: Option<&str>, path: &str) -> Result { + let base = api_base + .map(str::trim) + .filter(|base| !base.is_empty()) + .unwrap_or(REDUCTO_API_BASE); + ApiUrl::parse(base) + .and_then(|url| url.complete_path(&[path])) + .map(|url| url.into_string()) + .map_err(|_| { + OcrRequestError::RequestField { + path: "api_base".into(), + } + .into() + }) + } + + pub(super) 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 api_key = connection + .api_key + .as_deref() + .map(str::trim) + .filter(|key| !key.is_empty()) + .map(str::to_string) + .or_else(|| { + env_lookup(REDUCTO_API_KEY_ENV) + .map(|key| key.trim().to_string()) + .filter(|key| !key.is_empty()) + }) + .ok_or(Error::MissingReductoApiKey)?; + Ok( + std::iter::once(("Authorization".into(), format!("Bearer {api_key}"))) + .chain(connection.extra_headers.clone()) + .collect(), + ) + } + + pub(super) async fn prepare_document( + client: &crate::ocr::OcrClient, + document: OcrDocument, + connection: &OcrConnection, + headers: &[(String, String)], + ) -> Result { + if document.source().starts_with(REDUCTO_ID_PREFIX) { + if document.source()[REDUCTO_ID_PREFIX.len()..] + .trim() + .is_empty() + { + return Err(OcrRequestError::RequestField { + path: "document file id".into(), + } + .into()); + } + return Ok(document); + } + let inline = + InlineDocument::parse(document.source())?.ok_or(OcrRequestError::ReductoSource)?; + let mime = inline.mime_type().to_string(); + let bytes = inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?; + let part = reqwest::multipart::Part::bytes(bytes) + .file_name("document") + .mime_str(&mime) + .map_err(|_| OcrRequestError::InvalidDataUri)?; + let builder = client + .provider_http() + .post(get_complete_url(connection.api_base.as_deref(), "upload")?) + .multipart(reqwest::multipart::Form::new().part("file", part)) + .timeout(connection.timeout); + let builder = crate::http_utils::with_headers( + builder, + headers, + crate::http_utils::HeaderPolicy::Except(&["content-type", "content-length"]), + ); + let response = crate::http_utils::http_request(builder) + .await + .map_err(crate::error::TransportError::from)?; + let uploaded = + crate::ocr::client::read_json_response::( + response, + false, + connection.max_response_bytes, + ) + .await? + .data; + let file_id = uploaded + .file_id + .as_deref() + .map(str::trim) + .filter(|id| !id.is_empty()); + let Some(file_id) = file_id else { + return Err(OcrResponseError::ResponseField { + path: "file_id".into(), + } + .into()); + }; + Ok(document.with_source(file_id.to_string())) + } + + #[cfg(test)] + mod tests { + use super::*; + + #[test] + fn explicit_key_precedes_environment_key() { + let connection = OcrConnection { + api_key: Some("passed-key".into()), + ..Default::default() + }; + let headers = validate_environment(&connection, &|_| Some("env-key".into())).unwrap(); + assert_eq!(headers[0].1, "Bearer passed-key"); + } + + #[test] + fn blank_explicit_key_uses_environment_key() { + let connection = OcrConnection { + api_key: Some(" ".into()), + ..Default::default() + }; + let headers = validate_environment(&connection, &|_| Some(" env-key ".into())).unwrap(); + assert_eq!(headers[0].1, "Bearer env-key"); + } + + #[test] + fn existing_authorization_skips_key_lookup() { + let connection = OcrConnection { + extra_headers: vec![("authorization".into(), "Bearer existing".into())], + ..Default::default() + }; + assert_eq!( + validate_environment(&connection, &|_| None).unwrap(), + connection.extra_headers + ); + } + } +} + +mod legacy { + use super::common::BaseReductoOcrConfig; + use super::mapping as reducto; + use super::types::{ReductoLegacyParams, ReductoResponse}; + use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; + use crate::ocr::OcrClient; + use crate::ocr::error::{OcrError, OcrResponseError}; + use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env, + guardrail_document, merge_extra_params, + }; + use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; + + #[derive(Clone, Debug)] + pub(crate) struct ReductoParseLegacyConfig; + + impl BaseReductoOcrConfig for ReductoParseLegacyConfig {} + + impl BaseOcrConfig for ReductoParseLegacyConfig { + type ProviderResponse = ReductoResponse; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let ParsedProviderParams { + known: params, + extra_params, + } = _prepare_ocr_request::(request)?; + let headers = self.validate_environment(&request.connection, &credential_env)?; + let url = self.get_complete_url(request.connection.api_base.as_deref(), "parse")?; + let (document, headers) = guardrail_document(request, &url, &headers).await?; + let document = self + .prepare_document(client, document, &request.connection, &headers) + .await?; + let body = reducto::transform_legacy_ocr_request(&request.model, document, ¶ms)?; + let body = merge_extra_params(&body, extra_params)?; + build_http_request(client, request, &url, &headers, &body) + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: ReductoResponse, + ) -> Result { + reducto::transform_ocr_response(&request.model, response) + } + } +} + +pub(crate) use legacy::ReductoParseLegacyConfig; + +mod v3 { + use super::common::BaseReductoOcrConfig; + use super::mapping as reducto; + use super::types::{ReductoResponse, ReductoV3Params}; + use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; + use crate::ocr::OcrClient; + use crate::ocr::error::{OcrError, OcrResponseError}; + use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env, + guardrail_document, merge_extra_params, + }; + use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; + + #[derive(Clone, Debug)] + pub(crate) struct ReductoParseV3Config; + + impl BaseReductoOcrConfig for ReductoParseV3Config {} + + impl BaseOcrConfig for ReductoParseV3Config { + type ProviderResponse = ReductoResponse; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let ParsedProviderParams { + known: params, + extra_params, + } = _prepare_ocr_request::(request)?; + let headers = self.validate_environment(&request.connection, &credential_env)?; + let url = self.get_complete_url(request.connection.api_base.as_deref(), "parse")?; + let (document, headers) = guardrail_document(request, &url, &headers).await?; + let document = self + .prepare_document(client, document, &request.connection, &headers) + .await?; + let body = reducto::transform_v3_ocr_request(&request.model, document, ¶ms)?; + let body = merge_extra_params(&body, extra_params)?; + build_http_request(client, request, &url, &headers, &body) + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: ReductoResponse, + ) -> Result { + reducto::transform_ocr_response(&request.model, response) + } + } +} + +pub(crate) use v3::ReductoParseV3Config; diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/mod.rs b/litellm-rust/crates/core/src/llms/vertex_ai/mod.rs new file mode 100644 index 00000000000..079e0c41eae --- /dev/null +++ b/litellm-rust/crates/core/src/llms/vertex_ai/mod.rs @@ -0,0 +1 @@ +pub(crate) mod ocr; diff --git a/litellm-rust/crates/core/src/ocr/adapters/vertex/mod.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/common_utils.rs similarity index 68% rename from litellm-rust/crates/core/src/ocr/adapters/vertex/mod.rs rename to litellm-rust/crates/core/src/llms/vertex_ai/ocr/common_utils.rs index 270c41e647d..9daa28a1e5e 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/vertex/mod.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/common_utils.rs @@ -1,16 +1,10 @@ -mod deepseek; -mod mistral; - use crate::Error; use crate::auth::InputSource; use crate::auth::error::AuthConfigurationError; use crate::ocr::error::OcrError; use crate::ocr::types::OcrConnection; -pub(crate) use deepseek::VertexDeepSeekAdapter; -pub(crate) use mistral::VertexMistralAdapter; - -fn validate_destination(connection: &OcrConnection) -> Result<(), OcrError> { +pub(super) fn validate_destination(connection: &OcrConnection) -> Result<(), OcrError> { if connection.api_base.is_some() && connection.api_base_source == InputSource::Request { return Err(Error::from(crate::AuthError::Configuration( AuthConfigurationError::RequestVertexCredentialDestination, diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs new file mode 100644 index 00000000000..566b9780b9c --- /dev/null +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs @@ -0,0 +1,347 @@ +mod types { + use serde::{Deserialize, Serialize}; + use serde_json::{Map, Value}; + + #[derive(Clone, Debug, Default, Serialize, Deserialize)] + pub(crate) struct DeepSeekOcrParams { + #[serde(skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub top_p: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub n: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub stop: Option, + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + #[serde(untagged)] + pub(crate) enum StopSequences { + One(String), + Many(Vec), + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + pub(crate) struct DeepSeekOcrRequest { + pub model: String, + pub messages: Vec, + #[serde(flatten)] + pub params: DeepSeekOcrParams, + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + pub(crate) struct DeepSeekOcrMessage { + pub role: UserRole, + pub content: Vec, + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + #[serde(rename_all = "lowercase")] + pub(crate) enum UserRole { + User, + } + + #[derive(Clone, Debug, Deserialize)] + pub(crate) struct DeepSeekOcrResponse { + #[serde(default)] + pub choices: Vec, + pub usage: Option, + } + + #[derive(Clone, Debug, Deserialize)] + pub(crate) struct DeepSeekChoice { + pub message: DeepSeekResponseMessage, + } + + #[derive(Clone, Debug, Deserialize)] + pub(crate) struct DeepSeekResponseMessage { + pub content: Option, + } + + #[derive(Clone, Debug, Deserialize)] + #[serde(untagged)] + pub(crate) enum DeepSeekContent { + Text(String), + Object(DeepSeekOcrResult), + } + + #[derive(Clone, Debug, Default, Serialize, Deserialize)] + pub(crate) struct DeepSeekOcrResult { + #[serde(skip_serializing_if = "Option::is_none")] + pub pages: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub model: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub usage_info: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub document_annotation: Option, + #[serde(flatten)] + pub extra_fields: Map, + } + + #[derive(Clone, Debug, Serialize, Deserialize)] + pub(crate) struct DeepSeekPage { + #[serde(default)] + pub index: i64, + #[serde(default)] + pub markdown: String, + pub images: Option, + pub dimensions: Option, + #[serde(flatten)] + pub extra_fields: Map, + } +} + +pub(crate) use types::DeepSeekOcrParams; +pub(crate) use types::DeepSeekOcrResponse; + +mod mapping { + use serde::de::IntoDeserializer; + use serde_json::{Value, json}; + + use super::types::*; + use crate::ocr::error::{OcrRequestError, OcrResponseError}; + use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; + + #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] + pub(crate) fn transform_ocr_request( + provider_model: &str, + document: OcrDocument, + params: &DeepSeekOcrParams, + ) -> Result { + if document.source().is_empty() { + return Err(OcrRequestError::MissingDocumentUrl); + } + let content = OcrDocument::ImageUrl { + image_url: document.source().to_string(), + extra_fields: serde_json::Map::new(), + }; + Ok(DeepSeekOcrRequest { + model: provider_model.to_string(), + messages: vec![DeepSeekOcrMessage { + role: UserRole::User, + content: vec![content], + }], + params: params.clone(), + }) + } + + pub(crate) fn transform_ocr_response( + model: &str, + response: DeepSeekOcrResponse, + ) -> Result { + let content = response + .choices + .into_iter() + .next() + .and_then(|choice| choice.message.content) + .ok_or(OcrResponseError::EmptyContent)?; + let decoded = decode_content(content)?; + let pages = match decoded.result.pages { + Some(pages) if !pages.is_empty() => pages + .into_iter() + .map(|page| serde_json::to_value(page).expect("DeepSeek page serializes")) + .collect(), + _ => vec![json!({ + "index":0, + "markdown":decoded.fallback_markdown, + "images":null + })], + }; + Ok(LiteLLMOcrResponse { + pages, + model: decoded.result.model.unwrap_or_else(|| model.to_string()), + document_annotation: decoded.result.document_annotation, + usage_info: decoded.result.usage_info.or(response.usage), + object: "ocr".into(), + extra_fields: decoded.result.extra_fields, + provider_native_response: None, + }) + } + + struct DecodedContent { + result: DeepSeekOcrResult, + fallback_markdown: String, + } + + fn decode_content(content: DeepSeekContent) -> Result { + let (result, fallback_markdown) = match content { + DeepSeekContent::Text(text) if text.is_empty() => { + return Err(OcrResponseError::EmptyContent); + } + DeepSeekContent::Text(text) => (decode_json_content(&text)?, text), + DeepSeekContent::Object(object) => { + let fallback = serde_json::to_string(&object).map_err(|_| { + OcrResponseError::ResponseField { + path: "choices[0].message.content".into(), + } + })?; + (Some(object), fallback) + } + }; + Ok(DecodedContent { + result: result.unwrap_or_default(), + fallback_markdown, + }) + } + + fn decode_json_content(text: &str) -> Result, OcrResponseError> { + if !text.trim_start().starts_with('{') { + return Ok(None); + } + let value = match serde_json::from_str::(text) { + Ok(value) => value, + Err(_) => return Ok(None), + }; + serde_path_to_error::deserialize(value.into_deserializer()) + .map(Some) + .map_err(|error| OcrResponseError::ResponseField { + path: format!("choices[0].message.content.{}", error.path()), + }) + } +} + +#[cfg(test)] +pub(crate) use mapping::{transform_ocr_request, transform_ocr_response}; + +use super::common_utils::validate_destination; +use crate::Error; +use crate::auth::vertex::{self, VertexConfig}; +use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; +use crate::ocr::OcrClient; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, +}; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; +use crate::url_utils::ApiUrl; +const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com"; +const MODEL_NAMESPACE: &str = "deepseek-ai"; +const DEFAULT_LOCATION: &str = "us-central1"; + +#[derive(Clone, Debug)] +pub(crate) struct VertexAIDeepSeekOCRConfig; + +impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { + type ProviderResponse = DeepSeekOcrResponse; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + validate_destination(&request.connection)?; + let ParsedProviderParams { + known: params, + extra_params: _extra_params, + } = _prepare_ocr_request::(request)?; + let config = VertexConfig::from_sourced_optional_params( + &request.optional_params, + &request.input_sources, + ) + .map_err(Error::from)?; + let authentication = client + .vertex_auth() + .validate_environment( + request.connection.extra_headers.clone(), + request.connection.api_key.as_deref(), + &config, + &credential_env, + ) + .await + .map_err(Error::from)?; + let location = vertex::get_vertex_ai_location(&config, &credential_env) + .unwrap_or_else(|| DEFAULT_LOCATION.to_string()); + let url = get_complete_url( + request.connection.api_base.as_deref(), + &authentication.project_id, + &location, + )?; + let document = request.document.clone(); + let body = + mapping::transform_ocr_request(&provider_model(&request.model), document, ¶ms)?; + transform_request_body( + client, + request, + &url, + &authentication.headers, + false, + body, + |_| Ok(()), + ) + .await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: DeepSeekOcrResponse, + ) -> Result { + mapping::transform_ocr_response(&request.model, response) + } +} + +fn provider_model(model: &str) -> String { + if model.starts_with(&format!("{MODEL_NAMESPACE}/")) { + model.to_string() + } else { + format!("{MODEL_NAMESPACE}/{model}") + } +} + +fn get_complete_url( + api_base: Option<&str>, + project: &str, + location: &str, +) -> Result { + let base = api_base + .map(str::trim) + .filter(|base| !base.is_empty()) + .unwrap_or(DEFAULT_API_BASE); + ApiUrl::parse(base) + .and_then(|url| { + url.complete_path(&[ + "v1", + "projects", + project, + "locations", + location, + "endpoints", + "openapi", + "chat", + "completions", + ]) + }) + .map(|url| url.into_string()) + .map_err(|_| { + OcrRequestError::RequestField { + path: "api_base".into(), + } + .into() + }) +} + +#[cfg(test)] +mod tests { + use super::{get_complete_url, provider_model}; + + #[test] + fn config_owns_model_namespace_and_endpoint() { + assert_eq!( + provider_model("deepseek-ocr-maas"), + "deepseek-ai/deepseek-ocr-maas" + ); + assert_eq!( + provider_model("deepseek-ai/deepseek-ocr-maas"), + "deepseek-ai/deepseek-ocr-maas" + ); + assert_eq!( + get_complete_url(None, "proj-1", "europe-west4").unwrap(), + "https://aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/endpoints/openapi/chat/completions" + ); + } +} diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/mod.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/mod.rs new file mode 100644 index 00000000000..f894ec145f8 --- /dev/null +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/mod.rs @@ -0,0 +1,3 @@ +pub(crate) mod common_utils; +pub(crate) mod deepseek_transformation; +pub(crate) mod transformation; diff --git a/litellm-rust/crates/core/src/ocr/adapters/vertex/mistral.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs similarity index 88% rename from litellm-rust/crates/core/src/ocr/adapters/vertex/mistral.rs rename to litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs index 88c61725cee..eb03577d19b 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/vertex/mistral.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs @@ -1,25 +1,26 @@ -use super::super::OcrAdapter; -use super::validate_destination; +use super::common_utils::validate_destination; use crate::Error; use crate::auth::vertex::{self, VertexConfig}; +use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; +use crate::llms::mistral::ocr::transformation::MistralOCRConfig; +use crate::llms::mistral::ocr::{MistralOcrParams, MistralOcrResponse}; use crate::ocr::OcrClient; -use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse}; use crate::ocr::document::{inline_remote_document, validate_inline_document}; use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; use crate::ocr::prepare::{ _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, }; -use crate::ocr::registry::OcrProvider; use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; use crate::url_utils::ApiUrl; const DEFAULT_LOCATION: &str = "us-central1"; -#[derive(Clone, Debug)] -pub(crate) struct VertexMistralAdapter; +#[derive(Clone, Debug, Default)] +pub(crate) struct VertexAIOCRConfig { + mistral: MistralOCRConfig, +} -impl OcrAdapter for VertexMistralAdapter { +impl BaseOcrConfig for VertexAIOCRConfig { type ProviderResponse = MistralOcrResponse; - const PROVIDER: OcrProvider = OcrProvider::VertexAi; async fn prepare_request( &self, @@ -62,7 +63,9 @@ impl OcrAdapter for VertexMistralAdapter { &request.connection, ) .await?; - let body = mistral::transform_ocr_request(&request.model, document, ¶ms)?; + let body = self + .mistral + .transform_ocr_request(&request.model, document, ¶ms)?; transform_request_body( client, request, @@ -78,9 +81,9 @@ impl OcrAdapter for VertexMistralAdapter { fn transform_ocr_response( &self, request: &LiteLLMOcrRequest, - response: Self::ProviderResponse, + response: MistralOcrResponse, ) -> Result { - mistral::transform_ocr_response(&request.model, response) + self.mistral.transform_ocr_response(request, response) } } 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 deleted file mode 100644 index 22034ddc3f6..00000000000 --- a/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs +++ /dev/null @@ -1,214 +0,0 @@ -use super::super::OcrAdapter; -use crate::Error; -use crate::auth::{InputSource, Sourced}; -use crate::constants::{AZURE_DI_API_VERSION, AZURE_DI_SUBSCRIPTION_HEADER}; -use crate::ocr::OcrClient; -use crate::ocr::codecs::document_intelligence::{ - self, AzureDocumentIntelligenceOperation, DocumentIntelligenceParams, -}; -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, OcrResponseFormat}; -use crate::providers::azure_ai::auth::AzureAuthInputs; -use crate::url_utils::ApiUrl; - -mod polling; - -const AZURE_DI_API_KEY_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"; -const AZURE_DI_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT"; - -#[derive(Clone, Debug)] -pub(crate) struct AzureDocumentIntelligenceAdapter; - -impl OcrAdapter for AzureDocumentIntelligenceAdapter { - type ProviderResponse = AzureDocumentIntelligenceOperation; - const PROVIDER: OcrProvider = OcrProvider::AzureAi; - - async fn prepare_request( - &self, - request: &LiteLLMOcrRequest, - client: &OcrClient, - ) -> Result { - let params = map_ocr_params(request)?; - 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 headers = validate_environment(&request.connection, &config, &credential_env).await?; - let endpoint = nonblank(request.connection.api_base.clone()) - .or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV))) - .ok_or_else(|| Error::Auth("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into()))?; - let url = get_complete_url(&endpoint, &request.model, ¶ms)?; - let body = document_intelligence::transform_ocr_request(request.document.clone())?; - transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: Self::ProviderResponse, - ) -> Result { - document_intelligence::transform_ocr_response(&request.model, response) - } - - async fn read_response( - &self, - client: &OcrClient, - response: reqwest::Response, - url: &str, - headers: &[(String, String)], - request: &LiteLLMOcrRequest, - ) -> Result, OcrError> { - polling::read_operation_response( - client.polling_http(), - response, - url, - headers, - &request.connection, - request.response_format()? == OcrResponseFormat::Native, - &request.hooks, - ) - .await - } -} - -fn map_ocr_params( - request: &LiteLLMOcrRequest, -) -> Result { - let params = document_intelligence::decode_input_params( - request.optional_params.clone(), - "optional_params", - )?; - let crate::ocr::prepare::ParsedProviderParams { - known: params, - extra_params: _extra_params, - } = params; - document_intelligence::map_ocr_params(params) -} - -fn get_complete_url( - endpoint: &str, - model: &str, - params: &DocumentIntelligenceParams, -) -> Result { - let model = format!("{}:analyze", model_id(model)?); - ApiUrl::parse(endpoint) - .and_then(|url| url.complete_path(&["documentintelligence", "documentModels", &model])) - .map(|url| { - url.append_query_pairs( - [("api-version", AZURE_DI_API_VERSION)] - .into_iter() - .chain(params.pages.iter().map(|pages| ("pages", pages.as_str()))) - .chain( - params - .features - .iter() - .map(|features| ("features", features.as_str())), - ), - ) - .into_string() - }) - .map_err(|_| OcrRequestError::RequestField { - path: "api_base".into(), - }) - .map_err(OcrError::from) -} - -async fn validate_environment( - connection: &OcrConnection, - config: &AzureAuthInputs, - env_lookup: &(dyn Fn(&str) -> Option + Sync), -) -> Result, OcrError> { - if crate::http_utils::has_header(&connection.extra_headers, "authorization") - || crate::http_utils::has_header(&connection.extra_headers, AZURE_DI_SUBSCRIPTION_HEADER) - { - super::validate_destination(connection, connection.extra_headers_source)?; - return Ok(connection.extra_headers.clone()); - } - let key = nonblank(connection.api_key.clone()) - .map(|value| Sourced::new(value, connection.api_key_source)) - .or_else(|| { - nonblank(env_lookup(AZURE_DI_API_KEY_ENV)) - .map(|value| Sourced::new(value, InputSource::Environment)) - }); - if let Some(key) = key { - super::validate_destination(connection, key.source())?; - return Ok( - std::iter::once((AZURE_DI_SUBSCRIPTION_HEADER.into(), key.into_value())) - .chain(connection.extra_headers.clone()) - .collect(), - ); - } - let token = super::resolve_entra(config, env_lookup) - .await? - .ok_or(Error::MissingAzureDocumentIntelligenceCredentials)?; - super::validate_destination(connection, token.source())?; - Ok( - std::iter::once(("Authorization".into(), format!("Bearer {}", token.value()))) - .chain(connection.extra_headers.clone()) - .collect(), - ) -} - -fn model_id(model: &str) -> Result<&str, OcrRequestError> { - let model = model.rsplit('/').next().unwrap_or(model); - if matches!(model, "." | "..") { - return Err(OcrRequestError::DotModel); - } - Ok(model) -} - -fn nonblank(value: Option) -> Option { - value - .map(|value| value.trim().to_string()) - .filter(|value| !value.is_empty()) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[tokio::test] - async fn request_endpoint_cannot_receive_environment_key() { - let connection = OcrConnection { - api_base: Some("https://request.example".into()), - api_base_source: InputSource::Request, - ..Default::default() - }; - - let error = validate_environment(&connection, &Default::default(), &|name| { - (name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into()) - }) - .await - .unwrap_err(); - - assert!( - error - .to_string() - .contains("request-controlled Azure endpoint") - ); - } - - #[tokio::test] - async fn request_endpoint_accepts_request_owned_key() { - let connection = OcrConnection { - api_key: Some("request-key".into()), - api_key_source: InputSource::Request, - api_base: Some("https://request.example".into()), - api_base_source: InputSource::Request, - ..Default::default() - }; - - let headers = validate_environment(&connection, &Default::default(), &|_| None) - .await - .unwrap(); - - assert_eq!( - headers[0], - (AZURE_DI_SUBSCRIPTION_HEADER.into(), "request-key".into()) - ); - } -} 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 deleted file mode 100644 index 6ed1e4441d4..00000000000 --- a/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs +++ /dev/null @@ -1,119 +0,0 @@ -use std::sync::Arc; -use std::time::Duration; - -use reqwest::Url; -use tokio::time::Instant; - -use crate::constants::{AZURE_DI_SUBSCRIPTION_HEADER, OCR_POLL_RETRY_SECS}; -use crate::ocr::client::read_json_response; -use crate::ocr::codecs::document_intelligence::{ - AzureDocumentIntelligenceOperation, OperationStatus, -}; -use crate::ocr::error::{OcrError, OcrPollingError, OcrResponseError}; -use crate::ocr::hooks::OcrHooks; -use crate::ocr::types::OcrConnection; -use crate::ocr::wire::DecodedOcrResponse; - -pub(super) async fn read_operation_response( - http_client: &reqwest::Client, - response: reqwest::Response, - original_url: &str, - headers: &[(String, String)], - connection: &OcrConnection, - native: bool, - hooks: &Arc, -) -> Result, OcrError> { - if response.status() != reqwest::StatusCode::ACCEPTED { - let bytes = - crate::ocr::client::read_response_bytes(response, connection.max_response_bytes) - .await?; - crate::ocr::handler::post_call(hooks, &bytes).await?; - return Ok(crate::ocr::wire::decode_response(&bytes, native)?); - } - let location = response - .headers() - .get("operation-location") - .and_then(|value| value.to_str().ok()) - .ok_or(OcrPollingError::PollLocation)? - .to_string(); - let original = Url::parse(original_url).map_err(|_| OcrPollingError::PollOrigin)?; - let operation = Url::parse(&location).map_err(|_| OcrPollingError::PollOrigin)?; - if original.origin() != operation.origin() - || !operation.username().is_empty() - || operation.password().is_some() - { - return Err(OcrPollingError::PollOrigin.into()); - } - let bytes = - crate::ocr::client::read_response_bytes(response, connection.max_response_bytes).await?; - crate::ocr::handler::post_call(hooks, &bytes).await?; - poll_operation(http_client, operation, headers, connection, native, hooks).await -} - -async fn poll_operation( - http_client: &reqwest::Client, - url: Url, - headers: &[(String, String)], - connection: &OcrConnection, - native: bool, - hooks: &Arc, -) -> Result, OcrError> { - let deadline = Instant::now() - .checked_add(connection.poll_timeout) - .ok_or(OcrPollingError::PollTimeout)?; - loop { - let remaining = deadline - .checked_duration_since(Instant::now()) - .filter(|remaining| !remaining.is_zero()) - .ok_or(OcrPollingError::PollTimeout)?; - let builder = http_client - .get(url.clone()) - .timeout(remaining.min(connection.timeout)); - let builder = crate::http_utils::with_headers( - builder, - headers, - crate::http_utils::HeaderPolicy::Only(&[AZURE_DI_SUBSCRIPTION_HEADER, "authorization"]), - ); - let response = tokio::time::timeout_at(deadline, crate::http_utils::http_request(builder)) - .await - .map_err(|_| OcrPollingError::PollTimeout)? - .map_err(crate::error::TransportError::from)?; - let retry = response - .headers() - .get(reqwest::header::RETRY_AFTER) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.parse::().ok()) - .unwrap_or(OCR_POLL_RETRY_SECS) - .max(1); - let decoded = tokio::time::timeout_at( - deadline, - read_json_response::( - response, - native, - connection.max_response_bytes, - ), - ) - .await - .map_err(|_| OcrPollingError::PollTimeout)??; - match &decoded.data.status { - 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 - .map_err(|_| OcrPollingError::PollTimeout)?; - } - status => { - return Err(OcrResponseError::OperationStatus( - status - .as_ref() - .map(ToString::to_string) - .unwrap_or_else(|| "None".into()), - ) - .into()); - } - } - } -} diff --git a/litellm-rust/crates/core/src/ocr/adapters/cohere.rs b/litellm-rust/crates/core/src/ocr/adapters/cohere.rs deleted file mode 100644 index 933ead7f7f7..00000000000 --- a/litellm-rust/crates/core/src/ocr/adapters/cohere.rs +++ /dev/null @@ -1,123 +0,0 @@ -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/mistral.rs b/litellm-rust/crates/core/src/ocr/adapters/mistral.rs deleted file mode 100644 index cdbc2c3effc..00000000000 --- a/litellm-rust/crates/core/src/ocr/adapters/mistral.rs +++ /dev/null @@ -1,147 +0,0 @@ -use super::OcrAdapter; -use crate::Error; -use crate::constants::MISTRAL_OCR_API_BASE; -use crate::ocr::OcrClient; -use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse}; -use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; -use crate::ocr::prepare::{ - _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, -}; -use crate::ocr::registry::OcrProvider; -use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection}; -use crate::url_utils::ApiUrl; - -const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY"; - -#[derive(Clone, Debug)] -pub(crate) struct MistralAdapter; - -impl OcrAdapter for MistralAdapter { - type ProviderResponse = MistralOcrResponse; - const PROVIDER: OcrProvider = OcrProvider::Mistral; - - async fn prepare_request( - &self, - request: &LiteLLMOcrRequest, - client: &OcrClient, - ) -> Result { - let ParsedProviderParams { - known: params, - extra_params: _extra_params, - } = _prepare_ocr_request::(request)?; - let headers = validate_environment(&request.connection, &credential_env)?; - let url = get_complete_url(request.connection.api_base.as_deref())?; - let body = - mistral::transform_ocr_request(&request.model, request.document.clone(), ¶ms)?; - transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: Self::ProviderResponse, - ) -> Result { - mistral::transform_ocr_response(&request.model, response) - } -} - -pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result { - let base = api_base - .map(str::trim) - .filter(|base| !base.is_empty()) - .unwrap_or(MISTRAL_OCR_API_BASE); - ApiUrl::parse(base) - .and_then(|url| url.complete_path(&["v1", "ocr"])) - .map(|url| url.into_string()) - .map_err(|_| { - OcrRequestError::RequestField { - path: "api_base".into(), - } - .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 api_key = connection - .api_key - .as_deref() - .map(str::trim) - .filter(|key| !key.is_empty()) - .map(str::to_string) - .or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty())) - .ok_or(Error::MissingApiKey { - provider: "Mistral", - })?; - Ok( - std::iter::once(("Authorization".into(), format!("Bearer {api_key}"))) - .chain(connection.extra_headers.clone()) - .collect(), - ) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn complete_url_defaults_and_dedupes_v1() { - assert_eq!( - get_complete_url(None).unwrap(), - "https://api.mistral.ai/v1/ocr" - ); - assert_eq!( - get_complete_url(Some("https://example.com/v1?tenant=a")).unwrap(), - "https://example.com/v1/ocr?tenant=a" - ); - assert_eq!( - get_complete_url(Some("https://example.com/v1/ocr?tenant=a")).unwrap(), - "https://example.com/v1/ocr?tenant=a" - ); - } - - #[test] - fn environment_prefers_explicit_key_then_environment() { - let explicit = OcrConnection { - api_key: Some("explicit".into()), - ..OcrConnection::default() - }; - assert_eq!( - validate_environment(&explicit, &|_| Some("environment".into())).unwrap()[0], - ("Authorization".into(), "Bearer explicit".into()) - ); - - assert_eq!( - validate_environment(&OcrConnection::default(), &|_| Some("environment".into())) - .unwrap()[0], - ("Authorization".into(), "Bearer environment".into()) - ); - } - - #[test] - fn environment_preserves_forwarded_authorization() { - let connection = OcrConnection { - extra_headers: vec![("authorization".into(), "Bearer forwarded".into())], - ..OcrConnection::default() - }; - assert_eq!( - validate_environment(&connection, &|_| None).unwrap(), - connection.extra_headers - ); - } - - #[test] - fn environment_rejects_missing_key() { - assert!(matches!( - validate_environment(&OcrConnection::default(), &|_| None), - Err(OcrError::Public(Error::MissingApiKey { - provider: "Mistral" - })) - )); - } -} diff --git a/litellm-rust/crates/core/src/ocr/adapters/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/mod.rs deleted file mode 100644 index d473fcad280..00000000000 --- a/litellm-rust/crates/core/src/ocr/adapters/mod.rs +++ /dev/null @@ -1,91 +0,0 @@ -use std::future::Future; - -use serde::de::DeserializeOwned; - -use super::OcrClient; -use super::error::{OcrError, OcrResponseError}; -use super::registry::OcrProvider; -use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; - -mod azure; -mod cohere; -mod mistral; -mod reducto; -mod vertex; - -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}; - -/// Converts a complete LiteLLM OCR call to provider HTTP and normalizes its response. -pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static { - /// Provider JSON schema; direct and Vertex Mistral share `MistralOcrResponse`. - type ProviderResponse: DeserializeOwned + Send; - - const PROVIDER: OcrProvider; - - /// Prepares the complete provider HTTP request. - /// `request` contains the model, document, connection, and unmapped caller options. - /// `client` supplies reusable provider and document HTTP clients. - /// Returns the complete HTTP request, whereas Python returns body data. - fn prepare_request( - &self, - request: &LiteLLMOcrRequest, - client: &OcrClient, - ) -> impl Future> + Send; - - /// Python: `transform_ocr_response`. - /// `request` supplies caller context, including the fallback model. - /// `response` is the decoded provider payload; the output is the shared LiteLLM schema. - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: Self::ProviderResponse, - ) -> Result; - - /// Decodes provider HTTP; adapters may override this to poll asynchronous operations. - /// Python performs that polling inside `async_transform_ocr_response`. - /// `client` is reused for polling; `response` is the initial HTTP response. - /// `url` and `headers` describe the submitted call; `request` supplies limits and format. - fn read_response( - &self, - _client: &OcrClient, - response: reqwest::Response, - _url: &str, - _headers: &[(String, String)], - request: &LiteLLMOcrRequest, - ) -> impl Future< - Output = Result, OcrError>, - > + Send { - async move { - let bytes = - super::client::read_response_bytes(response, request.connection.max_response_bytes) - .await?; - super::handler::post_call(&request.hooks, &bytes).await?; - Ok(super::wire::decode_response( - &bytes, - request.response_format()? == super::types::OcrResponseFormat::Native, - )?) - } - } -} - -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; - ReductoLegacy, $crate::ocr::adapters::ReductoLegacyAdapter, $crate::ocr::adapters::ReductoLegacyAdapter, Reducto; - ReductoV3, $crate::ocr::adapters::ReductoV3Adapter, $crate::ocr::adapters::ReductoV3Adapter, Reducto; - VertexMistral, $crate::ocr::adapters::VertexMistralAdapter, $crate::ocr::adapters::VertexMistralAdapter, VertexAi; - VertexDeepSeek, $crate::ocr::adapters::VertexDeepSeekAdapter, $crate::ocr::adapters::VertexDeepSeekAdapter, VertexAi; - } - }; -} - -pub(crate) use for_each_ocr_adapter; diff --git a/litellm-rust/crates/core/src/ocr/adapters/reducto/legacy.rs b/litellm-rust/crates/core/src/ocr/adapters/reducto/legacy.rs deleted file mode 100644 index 8889bcd1b45..00000000000 --- a/litellm-rust/crates/core/src/ocr/adapters/reducto/legacy.rs +++ /dev/null @@ -1,45 +0,0 @@ -use super::super::OcrAdapter; -use crate::ocr::OcrClient; -use crate::ocr::codecs::reducto::{self, ReductoLegacyParams, ReductoResponse}; -use crate::ocr::error::{OcrError, OcrResponseError}; -use crate::ocr::prepare::{ - _prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env, - guardrail_document, merge_extra_params, -}; -use crate::ocr::registry::OcrProvider; -use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; - -#[derive(Clone, Debug)] -pub(crate) struct ReductoLegacyAdapter; - -impl OcrAdapter for ReductoLegacyAdapter { - type ProviderResponse = ReductoResponse; - const PROVIDER: OcrProvider = OcrProvider::Reducto; - - async fn prepare_request( - &self, - request: &LiteLLMOcrRequest, - client: &OcrClient, - ) -> Result { - let ParsedProviderParams { - known: params, - extra_params, - } = _prepare_ocr_request::(request)?; - let headers = super::validate_environment(&request.connection, &credential_env)?; - let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?; - let (document, headers) = guardrail_document(request, &url, &headers).await?; - let document = - super::prepare_document(client, document, &request.connection, &headers).await?; - let body = reducto::transform_legacy_ocr_request(&request.model, document, ¶ms)?; - let body = merge_extra_params(&body, extra_params)?; - build_http_request(client, request, &url, &headers, &body) - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: Self::ProviderResponse, - ) -> Result { - reducto::transform_ocr_response(&request.model, response) - } -} diff --git a/litellm-rust/crates/core/src/ocr/adapters/reducto/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/reducto/mod.rs deleted file mode 100644 index 2dafe291674..00000000000 --- a/litellm-rust/crates/core/src/ocr/adapters/reducto/mod.rs +++ /dev/null @@ -1,148 +0,0 @@ -mod legacy; -mod v3; - -use crate::Error; -use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX}; -use crate::ocr::document::InlineDocument; -use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; -use crate::ocr::types::{OcrConnection, OcrDocument}; -use crate::url_utils::ApiUrl; - -pub(crate) use legacy::ReductoLegacyAdapter; -pub(crate) use v3::ReductoV3Adapter; - -pub(super) fn get_complete_url(api_base: Option<&str>, path: &str) -> Result { - let base = api_base - .map(str::trim) - .filter(|base| !base.is_empty()) - .unwrap_or(REDUCTO_API_BASE); - ApiUrl::parse(base) - .and_then(|url| url.complete_path(&[path])) - .map(|url| url.into_string()) - .map_err(|_| { - OcrRequestError::RequestField { - path: "api_base".into(), - } - .into() - }) -} - -pub(super) 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 api_key = connection - .api_key - .as_deref() - .map(str::trim) - .filter(|key| !key.is_empty()) - .map(str::to_string) - .or_else(|| { - env_lookup(REDUCTO_API_KEY_ENV) - .map(|key| key.trim().to_string()) - .filter(|key| !key.is_empty()) - }) - .ok_or(Error::MissingReductoApiKey)?; - Ok( - std::iter::once(("Authorization".into(), format!("Bearer {api_key}"))) - .chain(connection.extra_headers.clone()) - .collect(), - ) -} - -pub(super) async fn prepare_document( - client: &crate::ocr::OcrClient, - document: OcrDocument, - connection: &OcrConnection, - headers: &[(String, String)], -) -> Result { - if document.source().starts_with(REDUCTO_ID_PREFIX) { - if document.source()[REDUCTO_ID_PREFIX.len()..] - .trim() - .is_empty() - { - return Err(OcrRequestError::RequestField { - path: "document file id".into(), - } - .into()); - } - return Ok(document); - } - let inline = InlineDocument::parse(document.source())?.ok_or(OcrRequestError::ReductoSource)?; - let mime = inline.mime_type().to_string(); - let bytes = inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?; - let part = reqwest::multipart::Part::bytes(bytes) - .file_name("document") - .mime_str(&mime) - .map_err(|_| OcrRequestError::InvalidDataUri)?; - let builder = client - .provider_http() - .post(get_complete_url(connection.api_base.as_deref(), "upload")?) - .multipart(reqwest::multipart::Form::new().part("file", part)) - .timeout(connection.timeout); - let builder = crate::http_utils::with_headers( - builder, - headers, - crate::http_utils::HeaderPolicy::Except(&["content-type", "content-length"]), - ); - let response = crate::http_utils::http_request(builder) - .await - .map_err(crate::error::TransportError::from)?; - let uploaded = crate::ocr::client::read_json_response::< - crate::ocr::codecs::reducto::ReductoUploadResponse, - >(response, false, connection.max_response_bytes) - .await? - .data; - let file_id = uploaded - .file_id - .as_deref() - .map(str::trim) - .filter(|id| !id.is_empty()); - let Some(file_id) = file_id else { - return Err(OcrResponseError::ResponseField { - path: "file_id".into(), - } - .into()); - }; - Ok(document.with_source(file_id.to_string())) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn explicit_key_precedes_environment_key() { - let connection = OcrConnection { - api_key: Some("passed-key".into()), - ..Default::default() - }; - let headers = validate_environment(&connection, &|_| Some("env-key".into())).unwrap(); - assert_eq!(headers[0].1, "Bearer passed-key"); - } - - #[test] - fn blank_explicit_key_uses_environment_key() { - let connection = OcrConnection { - api_key: Some(" ".into()), - ..Default::default() - }; - let headers = validate_environment(&connection, &|_| Some(" env-key ".into())).unwrap(); - assert_eq!(headers[0].1, "Bearer env-key"); - } - - #[test] - fn existing_authorization_skips_key_lookup() { - let connection = OcrConnection { - extra_headers: vec![("authorization".into(), "Bearer existing".into())], - ..Default::default() - }; - assert_eq!( - validate_environment(&connection, &|_| None).unwrap(), - connection.extra_headers - ); - } -} diff --git a/litellm-rust/crates/core/src/ocr/adapters/reducto/v3.rs b/litellm-rust/crates/core/src/ocr/adapters/reducto/v3.rs deleted file mode 100644 index c272d31b67e..00000000000 --- a/litellm-rust/crates/core/src/ocr/adapters/reducto/v3.rs +++ /dev/null @@ -1,45 +0,0 @@ -use super::super::OcrAdapter; -use crate::ocr::OcrClient; -use crate::ocr::codecs::reducto::{self, ReductoResponse, ReductoV3Params}; -use crate::ocr::error::{OcrError, OcrResponseError}; -use crate::ocr::prepare::{ - _prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env, - guardrail_document, merge_extra_params, -}; -use crate::ocr::registry::OcrProvider; -use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; - -#[derive(Clone, Debug)] -pub(crate) struct ReductoV3Adapter; - -impl OcrAdapter for ReductoV3Adapter { - type ProviderResponse = ReductoResponse; - const PROVIDER: OcrProvider = OcrProvider::Reducto; - - async fn prepare_request( - &self, - request: &LiteLLMOcrRequest, - client: &OcrClient, - ) -> Result { - let ParsedProviderParams { - known: params, - extra_params, - } = _prepare_ocr_request::(request)?; - let headers = super::validate_environment(&request.connection, &credential_env)?; - let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?; - let (document, headers) = guardrail_document(request, &url, &headers).await?; - let document = - super::prepare_document(client, document, &request.connection, &headers).await?; - let body = reducto::transform_v3_ocr_request(&request.model, document, ¶ms)?; - let body = merge_extra_params(&body, extra_params)?; - build_http_request(client, request, &url, &headers, &body) - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: Self::ProviderResponse, - ) -> Result { - reducto::transform_ocr_response(&request.model, response) - } -} diff --git a/litellm-rust/crates/core/src/ocr/adapters/vertex/deepseek.rs b/litellm-rust/crates/core/src/ocr/adapters/vertex/deepseek.rs deleted file mode 100644 index d16b3e7f386..00000000000 --- a/litellm-rust/crates/core/src/ocr/adapters/vertex/deepseek.rs +++ /dev/null @@ -1,140 +0,0 @@ -use super::super::OcrAdapter; -use super::validate_destination; -use crate::Error; -use crate::auth::vertex::{self, VertexConfig}; -use crate::ocr::OcrClient; -use crate::ocr::codecs::deepseek::{self, DeepSeekOcrParams, DeepSeekOcrResponse}; -use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; -use crate::ocr::prepare::{ - _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, -}; -use crate::ocr::registry::OcrProvider; -use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; -use crate::url_utils::ApiUrl; -const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com"; -const MODEL_NAMESPACE: &str = "deepseek-ai"; -const DEFAULT_LOCATION: &str = "us-central1"; - -#[derive(Clone, Debug)] -pub(crate) struct VertexDeepSeekAdapter; - -impl OcrAdapter for VertexDeepSeekAdapter { - type ProviderResponse = DeepSeekOcrResponse; - const PROVIDER: OcrProvider = OcrProvider::VertexAi; - - async fn prepare_request( - &self, - request: &LiteLLMOcrRequest, - client: &OcrClient, - ) -> Result { - validate_destination(&request.connection)?; - let ParsedProviderParams { - known: params, - extra_params: _extra_params, - } = _prepare_ocr_request::(request)?; - let config = VertexConfig::from_sourced_optional_params( - &request.optional_params, - &request.input_sources, - ) - .map_err(Error::from)?; - let authentication = client - .vertex_auth() - .validate_environment( - request.connection.extra_headers.clone(), - request.connection.api_key.as_deref(), - &config, - &credential_env, - ) - .await - .map_err(Error::from)?; - let location = vertex::get_vertex_ai_location(&config, &credential_env) - .unwrap_or_else(|| DEFAULT_LOCATION.to_string()); - let url = get_complete_url( - request.connection.api_base.as_deref(), - &authentication.project_id, - &location, - )?; - let document = request.document.clone(); - let body = - deepseek::transform_ocr_request(&provider_model(&request.model), document, ¶ms)?; - transform_request_body( - client, - request, - &url, - &authentication.headers, - false, - body, - |_| Ok(()), - ) - .await - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: Self::ProviderResponse, - ) -> Result { - deepseek::transform_ocr_response(&request.model, response) - } -} - -fn provider_model(model: &str) -> String { - if model.starts_with(&format!("{MODEL_NAMESPACE}/")) { - model.to_string() - } else { - format!("{MODEL_NAMESPACE}/{model}") - } -} - -fn get_complete_url( - api_base: Option<&str>, - project: &str, - location: &str, -) -> Result { - let base = api_base - .map(str::trim) - .filter(|base| !base.is_empty()) - .unwrap_or(DEFAULT_API_BASE); - ApiUrl::parse(base) - .and_then(|url| { - url.complete_path(&[ - "v1", - "projects", - project, - "locations", - location, - "endpoints", - "openapi", - "chat", - "completions", - ]) - }) - .map(|url| url.into_string()) - .map_err(|_| { - OcrRequestError::RequestField { - path: "api_base".into(), - } - .into() - }) -} - -#[cfg(test)] -mod tests { - use super::{get_complete_url, provider_model}; - - #[test] - fn adapter_owns_model_namespace_and_endpoint() { - assert_eq!( - provider_model("deepseek-ocr-maas"), - "deepseek-ai/deepseek-ocr-maas" - ); - assert_eq!( - provider_model("deepseek-ai/deepseek-ocr-maas"), - "deepseek-ai/deepseek-ocr-maas" - ); - assert_eq!( - get_complete_url(None, "proj-1", "europe-west4").unwrap(), - "https://aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/endpoints/openapi/chat/completions" - ); - } -} diff --git a/litellm-rust/crates/core/src/ocr/codecs/cohere.rs b/litellm-rust/crates/core/src/ocr/codecs/cohere.rs deleted file mode 100644 index 649432f39d3..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/cohere.rs +++ /dev/null @@ -1,254 +0,0 @@ -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": [ - { - "type":"markdown", - "index":4, - "markdown":{ - "content":"receipt", - "images":[{ - "id":"image", - "bounding_box":{"top_left_x":1,"bottom_right_x":48}, - "bounding_box_normalized":{"top_left_x":0.04,"bottom_right_x":0.15}, - "description":"scan", - "category":"logo" - }] - } - }, - {"type":"blocks","blocks":[{"type":"text","text":{"content":"total"}}]} - ], - "meta":{"api_version":{"version":"2"},"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]["bounding_box_normalized"]["bottom_right_x"], - 0.15 - ); - assert_eq!(normalized.pages[0]["images"][0]["description"], "scan"); - assert_eq!(normalized.pages[0]["images"][0]["category"], "logo"); - assert_eq!(normalized.pages[1]["index"], 1); - assert_eq!(normalized.pages[1]["markdown"], ""); - assert_eq!(normalized.pages[1]["blocks"][0]["text"]["content"], "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/deepseek/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/deepseek/mod.rs deleted file mode 100644 index 682b3addde7..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/deepseek/mod.rs +++ /dev/null @@ -1,5 +0,0 @@ -mod transformation; -mod types; - -pub(crate) use transformation::{transform_ocr_request, transform_ocr_response}; -pub(crate) use types::{DeepSeekOcrParams, DeepSeekOcrResponse}; diff --git a/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs deleted file mode 100644 index 999ac6cf032..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs +++ /dev/null @@ -1,101 +0,0 @@ -use serde::de::IntoDeserializer; -use serde_json::{Value, json}; - -use super::types::*; -use crate::ocr::error::{OcrRequestError, OcrResponseError}; -use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; - -pub(crate) fn transform_ocr_request( - provider_model: &str, - document: OcrDocument, - params: &DeepSeekOcrParams, -) -> Result { - if document.source().is_empty() { - return Err(OcrRequestError::MissingDocumentUrl); - } - let content = OcrDocument::ImageUrl { - image_url: document.source().to_string(), - extra_fields: serde_json::Map::new(), - }; - Ok(DeepSeekOcrRequest { - model: provider_model.to_string(), - messages: vec![DeepSeekOcrMessage { - role: UserRole::User, - content: vec![content], - }], - params: params.clone(), - }) -} - -pub(crate) fn transform_ocr_response( - model: &str, - response: DeepSeekOcrResponse, -) -> Result { - let content = response - .choices - .into_iter() - .next() - .and_then(|choice| choice.message.content) - .ok_or(OcrResponseError::EmptyContent)?; - let decoded = decode_content(content)?; - let pages = match decoded.result.pages { - Some(pages) if !pages.is_empty() => pages - .into_iter() - .map(|page| serde_json::to_value(page).expect("DeepSeek page serializes")) - .collect(), - _ => vec![json!({ - "index":0, - "markdown":decoded.fallback_markdown, - "images":null - })], - }; - Ok(LiteLLMOcrResponse { - pages, - model: decoded.result.model.unwrap_or_else(|| model.to_string()), - document_annotation: decoded.result.document_annotation, - usage_info: decoded.result.usage_info.or(response.usage), - object: "ocr".into(), - extra_fields: decoded.result.extra_fields, - provider_native_response: None, - }) -} - -struct DecodedContent { - result: DeepSeekOcrResult, - fallback_markdown: String, -} - -fn decode_content(content: DeepSeekContent) -> Result { - let (result, fallback_markdown) = match content { - DeepSeekContent::Text(text) if text.is_empty() => { - return Err(OcrResponseError::EmptyContent); - } - DeepSeekContent::Text(text) => (decode_json_content(&text)?, text), - DeepSeekContent::Object(object) => { - let fallback = - serde_json::to_string(&object).map_err(|_| OcrResponseError::ResponseField { - path: "choices[0].message.content".into(), - })?; - (Some(object), fallback) - } - }; - Ok(DecodedContent { - result: result.unwrap_or_default(), - fallback_markdown, - }) -} - -fn decode_json_content(text: &str) -> Result, OcrResponseError> { - if !text.trim_start().starts_with('{') { - return Ok(None); - } - let value = match serde_json::from_str::(text) { - Ok(value) => value, - Err(_) => return Ok(None), - }; - serde_path_to_error::deserialize(value.into_deserializer()) - .map(Some) - .map_err(|error| OcrResponseError::ResponseField { - path: format!("choices[0].message.content.{}", error.path()), - }) -} diff --git a/litellm-rust/crates/core/src/ocr/codecs/deepseek/types.rs b/litellm-rust/crates/core/src/ocr/codecs/deepseek/types.rs deleted file mode 100644 index 0ce2d9913f7..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/deepseek/types.rs +++ /dev/null @@ -1,95 +0,0 @@ -use serde::{Deserialize, Serialize}; -use serde_json::{Map, Value}; - -#[derive(Clone, Debug, Default, Serialize, Deserialize)] -pub(crate) struct DeepSeekOcrParams { - #[serde(skip_serializing_if = "Option::is_none")] - pub stream: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub temperature: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub max_tokens: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub top_p: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub n: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub stop: Option, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -#[serde(untagged)] -pub(crate) enum StopSequences { - One(String), - Many(Vec), -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub(crate) struct DeepSeekOcrRequest { - pub model: String, - pub messages: Vec, - #[serde(flatten)] - pub params: DeepSeekOcrParams, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub(crate) struct DeepSeekOcrMessage { - pub role: UserRole, - pub content: Vec, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -#[serde(rename_all = "lowercase")] -pub(crate) enum UserRole { - User, -} - -#[derive(Clone, Debug, Deserialize)] -pub(crate) struct DeepSeekOcrResponse { - #[serde(default)] - pub choices: Vec, - pub usage: Option, -} - -#[derive(Clone, Debug, Deserialize)] -pub(crate) struct DeepSeekChoice { - pub message: DeepSeekResponseMessage, -} - -#[derive(Clone, Debug, Deserialize)] -pub(crate) struct DeepSeekResponseMessage { - pub content: Option, -} - -#[derive(Clone, Debug, Deserialize)] -#[serde(untagged)] -pub(crate) enum DeepSeekContent { - Text(String), - Object(DeepSeekOcrResult), -} - -#[derive(Clone, Debug, Default, Serialize, Deserialize)] -pub(crate) struct DeepSeekOcrResult { - #[serde(skip_serializing_if = "Option::is_none")] - pub pages: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub model: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub usage_info: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub document_annotation: Option, - #[serde(flatten)] - pub extra_fields: Map, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub(crate) struct DeepSeekPage { - #[serde(default)] - pub index: i64, - #[serde(default)] - pub markdown: String, - pub images: Option, - pub dimensions: Option, - #[serde(flatten)] - pub extra_fields: Map, -} diff --git a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/mod.rs deleted file mode 100644 index 8031f2124a3..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/mod.rs +++ /dev/null @@ -1,9 +0,0 @@ -mod params; -mod transformation; -mod types; - -pub(crate) use params::{decode_input_params, map_ocr_params}; -pub(crate) use transformation::{transform_ocr_request, transform_ocr_response}; -pub(crate) use types::{ - AzureDocumentIntelligenceOperation, DocumentIntelligenceParams, OperationStatus, -}; diff --git a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/params.rs b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/params.rs deleted file mode 100644 index 9389f93b8e3..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/params.rs +++ /dev/null @@ -1,219 +0,0 @@ -use std::collections::BTreeSet; - -use serde_json::{Map, Value}; - -use super::types::{ - DocumentIntelligenceInputParams, DocumentIntelligenceParams, FeaturesInput, PagesInput, -}; -use crate::ocr::error::OcrRequestError; -use crate::ocr::prepare::ParsedProviderParams; - -pub(crate) fn decode_input_params( - params: Map, - prefix: &str, -) -> Result, OcrRequestError> { - if let Some(Value::Array(pages)) = params.get("pages") { - if pages.iter().any(Value::is_boolean) { - return Err(OcrRequestError::Pages("boolean page index".into())); - } - if pages - .iter() - .any(|page| page.is_number() && page.as_i64().is_none()) - { - return Err(OcrRequestError::Pages("page index is out of range".into())); - } - if !pages.iter().all(Value::is_i64) && !pages.iter().all(Value::is_string) { - return Err(OcrRequestError::Pages("mixed page element types".into())); - } - } - crate::ocr::wire::decode_request_value(Value::Object(params), prefix) -} - -pub(crate) fn map_ocr_params( - params: DocumentIntelligenceInputParams, -) -> Result { - Ok(DocumentIntelligenceParams { - pages: params.pages.map(normalize_pages).transpose()?.flatten(), - features: params - .features - .map(normalize_features) - .transpose()? - .flatten(), - }) -} - -fn normalize_pages(pages: PagesInput) -> Result, OcrRequestError> { - let normalized = match pages { - PagesInput::ZeroBasedIndices(indices) => { - if indices.is_empty() { - return Ok(None); - } - indices - .into_iter() - .map(|page| { - if page < 0 { - return Err(OcrRequestError::Pages("negative page index".into())); - } - page.checked_add(1) - .ok_or_else(|| OcrRequestError::Pages("page index is out of range".into())) - }) - .collect::, _>>()? - .into_iter() - .map(|page| page.to_string()) - .collect::>() - .join(",") - } - PagesInput::NativeTokens(tokens) => { - if tokens.is_empty() { - return Ok(None); - } - tokens - .iter() - .map(|token| token.trim()) - .collect::>() - .join(",") - } - PagesInput::NativeRange(range) => range - .split(',') - .map(str::trim) - .collect::>() - .join(","), - }; - if !normalized.split(',').all(valid_page_token) { - return Err(OcrRequestError::Pages("invalid native page range".into())); - } - Ok(Some(normalized)) -} - -fn valid_page_token(token: &str) -> bool { - let mut parts = token.split('-'); - let start = parts.next().unwrap_or_default(); - if start.is_empty() || !start.chars().all(|character| character.is_ascii_digit()) { - return false; - } - match parts.next() { - None => true, - Some(end) => { - !end.is_empty() - && end.chars().all(|character| character.is_ascii_digit()) - && parts.next().is_none() - } - } -} - -fn normalize_features(features: FeaturesInput) -> Result, OcrRequestError> { - let tokens = match features { - FeaturesInput::Names(names) => names, - FeaturesInput::CommaSeparated(names) => names.split(',').map(str::to_string).collect(), - }; - if tokens.is_empty() { - return Ok(None); - } - let normalized = tokens.iter().map(|token| token.trim()).collect::>(); - if !normalized.iter().all(|token| { - let Some((first, rest)) = token.as_bytes().split_first() else { - return false; - }; - first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric) - }) { - return Err(OcrRequestError::Features); - } - Ok(Some(normalized.join(","))) -} - -#[cfg(test)] -mod tests { - use rstest::rstest; - use serde_json::{Value, json}; - - use super::*; - - fn map(value: Value) -> Result { - let fields = value.as_object().unwrap().clone(); - map_ocr_params(decode_input_params(fields, "optional_params")?.known) - } - - #[test] - fn input_params_retain_unknown_fields() { - let parsed = decode_input_params( - json!({ - "pages": [0], - "future_ocr_option": true, - "extra_body": {"provider_option": "value"} - }) - .as_object() - .unwrap() - .clone(), - "optional_params", - ) - .unwrap(); - - assert_eq!( - parsed.known.pages, - Some(PagesInput::ZeroBasedIndices(vec![0])) - ); - assert_eq!(parsed.extra_params["future_ocr_option"], true); - assert_eq!( - parsed.extra_params["extra_body"], - json!({"provider_option": "value"}) - ); - assert_eq!( - serde_json::to_value(map_ocr_params(parsed.known).unwrap()).unwrap(), - json!({"pages": "1", "features": null}) - ); - } - - #[rstest] - #[case(json!([0, 1, 2]), Some("1,2,3"))] - #[case(json!([2, 0, 0, 1]), Some("1,2,3"))] - #[case(json!([]), None)] - #[case(json!("3-9"), Some("3-9"))] - #[case(json!("1-3, 5"), Some("1-3,5"))] - #[case(json!(["1", "3-5"]), Some("1,3-5"))] - fn page_mapping_matches_python(#[case] input: Value, #[case] expected: Option<&str>) { - assert_eq!( - map(json!({"pages": input})).unwrap().pages.as_deref(), - expected - ); - } - - #[rstest] - #[case(json!("a,b"))] - #[case(json!([-1]))] - #[case(json!([true, false]))] - #[case(json!([1, "2"]))] - #[case(json!(5))] - fn invalid_page_mapping_matches_python(#[case] input: Value) { - assert!(map(json!({"pages": input})).is_err()); - } - - #[rstest] - #[case(json!(["keyValuePairs"]), "keyValuePairs")] - #[case(json!(["keyValuePairs", "languages"]), "keyValuePairs,languages")] - #[case(json!("keyValuePairs"), "keyValuePairs")] - #[case(json!("keyValuePairs,languages"), "keyValuePairs,languages")] - #[case(json!("keyValuePairs, languages"), "keyValuePairs,languages")] - fn feature_mapping_matches_python(#[case] input: Value, #[case] expected: &str) { - assert_eq!( - map(json!({"features": input})).unwrap().features.as_deref(), - Some(expected) - ); - } - - #[rstest] - #[case(json!("keyValuePairs&pages=9"))] - #[case(json!("key value pairs"))] - #[case(json!(""))] - #[case(json!([1, 2]))] - #[case(json!([["keyValuePairs"]]))] - #[case(json!({"feature":"keyValuePairs"}))] - #[case(json!(5))] - fn invalid_feature_mapping_matches_python(#[case] input: Value) { - assert!(map(json!({"features": input})).is_err()); - } - - #[test] - fn empty_feature_list_is_omitted() { - assert_eq!(map(json!({"features": []})).unwrap().features, None); - } -} diff --git a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs deleted file mode 100644 index 018d7eb9c65..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs +++ /dev/null @@ -1,107 +0,0 @@ -use base64::{Engine, engine::general_purpose::STANDARD}; -use serde_json::{Map, Value, json}; - -use super::types::*; -use crate::constants::{AZURE_DI_DEFAULT_DPI, AZURE_DI_DEFAULT_HEIGHT, AZURE_DI_DEFAULT_WIDTH}; -use crate::ocr::document::InlineDocument; -use crate::ocr::error::{OcrRequestError, OcrResponseError}; -use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; - -pub(crate) fn transform_ocr_request( - document: OcrDocument, -) -> Result { - let source = document.source(); - if source.is_empty() { - return Err(OcrRequestError::MissingDocumentUrl); - } - Ok(if let Some(document) = InlineDocument::parse(source)? { - DocumentIntelligenceRequest::Base64Source( - STANDARD.encode(document.decode(crate::constants::OCR_INLINE_MAX_BYTES)?), - ) - } else { - DocumentIntelligenceRequest::UrlSource(source.to_string()) - }) -} - -pub(crate) fn transform_ocr_response( - model: &str, - response: AzureDocumentIntelligenceOperation, -) -> Result { - if response.status != Some(OperationStatus::Succeeded) { - return Err(OcrResponseError::OperationStatus( - response - .status - .map(|status| status.to_string()) - .unwrap_or_else(|| "None".into()), - )); - } - let result = response.analyze_result.unwrap_or_default(); - let pages = result - .pages - .into_iter() - .map(normalize_page) - .collect::, _>>()?; - let pages_processed = pages.len(); - let mut extra_fields = Map::new(); - extra_fields.insert("content".into(), option_value(result.content)); - extra_fields.insert("tables".into(), option_value(result.tables)); - extra_fields.insert("keyValuePairs".into(), option_value(result.key_value_pairs)); - Ok(LiteLLMOcrResponse { - pages, - model: model.into(), - document_annotation: None, - usage_info: Some(json!({"pages_processed":pages_processed})), - object: "ocr".into(), - extra_fields, - provider_native_response: None, - }) -} - -fn normalize_page(page: AzureDocumentIntelligencePage) -> Result { - let index = page - .page_number - .unwrap_or(1) - .checked_sub(1) - .ok_or(OcrResponseError::NumericRange("page.pageNumber"))?; - let scale = if page.unit.as_deref().unwrap_or("inch") == "inch" { - AZURE_DI_DEFAULT_DPI as f64 - } else { - 1.0 - }; - let width = pixel_dimension( - page.width.unwrap_or(AZURE_DI_DEFAULT_WIDTH), - scale, - "page.width", - )?; - let height = pixel_dimension( - page.height.unwrap_or(AZURE_DI_DEFAULT_HEIGHT), - scale, - "page.height", - )?; - let markdown = page - .lines - .iter() - .map(|line| line.content.as_deref().unwrap_or_default()) - .collect::>() - .join("\n"); - Ok(json!({ - "index":index, - "markdown":markdown, - "images":null, - "dimensions":{"width":width,"height":height,"dpi":AZURE_DI_DEFAULT_DPI} - })) -} - -fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result { - let value = value * scale; - if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 { - return Err(OcrResponseError::NumericRange(field)); - } - Ok(value.trunc() as i64) -} - -fn option_value(value: Option) -> Value { - value - .and_then(|value| serde_json::to_value(value).ok()) - .unwrap_or(Value::Null) -} diff --git a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/types.rs b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/types.rs deleted file mode 100644 index 793f4547e99..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/types.rs +++ /dev/null @@ -1,138 +0,0 @@ -use serde::{Deserialize, Deserializer, Serialize}; -use serde_json::{Map, Value}; - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(untagged)] -pub(crate) enum PagesInput { - ZeroBasedIndices(Vec), - NativeTokens(Vec), - NativeRange(String), -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(untagged)] -pub(crate) enum FeaturesInput { - Names(Vec), - CommaSeparated(String), -} - -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub(crate) struct DocumentIntelligenceInputParams { - pub pages: Option, - pub features: Option, -} - -#[derive(Clone, Debug, PartialEq, Serialize)] -pub(crate) struct DocumentIntelligenceParams { - pub pages: Option, - pub features: Option, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub(crate) enum DocumentIntelligenceRequest { - #[serde(rename = "urlSource")] - UrlSource(String), - #[serde(rename = "base64Source")] - Base64Source(String), -} - -#[derive(Clone, Debug, PartialEq)] -pub(crate) enum OperationStatus { - Succeeded, - Running, - NotStarted, - Failed, - Unknown(String), -} - -impl<'de> Deserialize<'de> for OperationStatus { - fn deserialize>(deserializer: D) -> Result { - Ok(match String::deserialize(deserializer)?.as_str() { - "succeeded" => Self::Succeeded, - "running" => Self::Running, - "notStarted" => Self::NotStarted, - "failed" => Self::Failed, - value => Self::Unknown(value.to_string()), - }) - } -} - -impl std::fmt::Display for OperationStatus { - fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - formatter.write_str(match self { - Self::Succeeded => "succeeded", - Self::Running => "running", - Self::NotStarted => "notStarted", - Self::Failed => "failed", - Self::Unknown(value) => value, - }) - } -} - -#[derive(Clone, Debug, Deserialize)] -pub(crate) struct AzureDocumentIntelligenceOperation { - pub status: Option, - #[serde(rename = "analyzeResult")] - pub analyze_result: Option, -} - -#[derive(Clone, Debug, Default, Deserialize)] -pub(crate) struct AzureDocumentIntelligenceAnalyzeResult { - pub content: Option, - #[serde(default)] - pub pages: Vec, - pub tables: Option>>, - #[serde(rename = "keyValuePairs")] - pub key_value_pairs: Option>>, -} - -#[derive(Clone, Debug, Deserialize)] -pub(crate) struct AzureDocumentIntelligencePage { - #[serde(rename = "pageNumber", default, deserialize_with = "optional_i64")] - pub page_number: Option, - #[serde(default, deserialize_with = "optional_f64")] - pub width: Option, - #[serde(default, deserialize_with = "optional_f64")] - pub height: Option, - pub unit: Option, - #[serde(default)] - pub lines: Vec, -} - -#[derive(Clone, Debug, Deserialize)] -pub(crate) struct AzureDocumentIntelligenceLine { - pub content: Option, -} - -fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { - match Option::::deserialize(deserializer)? { - None | Some(Value::Null) => Ok(None), - Some(Value::Number(number)) => number - .as_i64() - .map(Some) - .ok_or_else(|| serde::de::Error::custom("expected an integer")), - Some(Value::String(value)) => value - .parse::() - .map(Some) - .map_err(|_| serde::de::Error::custom("expected an integer")), - Some(_) => Err(serde::de::Error::custom("expected an integer")), - } -} - -fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { - match Option::::deserialize(deserializer)? { - None | Some(Value::Null) => Ok(None), - Some(Value::Number(number)) => number - .as_f64() - .filter(|value| value.is_finite()) - .map(Some) - .ok_or_else(|| serde::de::Error::custom("expected a finite number")), - Some(Value::String(value)) => value - .parse::() - .ok() - .filter(|value| value.is_finite()) - .map(Some) - .ok_or_else(|| serde::de::Error::custom("expected a finite number")), - Some(_) => Err(serde::de::Error::custom("expected a number")), - } -} diff --git a/litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs deleted file mode 100644 index eea4254779e..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs +++ /dev/null @@ -1,5 +0,0 @@ -mod transformation; -mod types; - -pub(crate) use transformation::{transform_ocr_request, transform_ocr_response}; -pub(crate) use types::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse}; diff --git a/litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs b/litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs deleted file mode 100644 index e0bc8a267d2..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs +++ /dev/null @@ -1,60 +0,0 @@ -use serde::{Deserialize, Serialize}; -use serde_json::{Map, Value}; - -use crate::ocr::types::OcrDocument; - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(untagged)] -pub(crate) enum MistralOcrPages { - Range(String), - Indices(Vec), -} - -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -pub(crate) struct MistralOcrParams { - #[serde(skip_serializing_if = "Option::is_none")] - pub pages: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub include_image_base64: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub image_limit: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub image_min_size: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub bbox_annotation_format: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub document_annotation_format: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub document_annotation_prompt: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub extract_header: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub extract_footer: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub table_format: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub confidence_scores_granularity: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub include_blocks: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub id: Option, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub(crate) struct MistralOcrRequest { - pub model: String, - pub document: OcrDocument, - #[serde(flatten)] - pub params: MistralOcrParams, -} - -#[derive(Clone, Debug, Default, Deserialize)] -pub(crate) struct MistralOcrResponse { - #[serde(default)] - pub pages: Vec, - pub model: Option, - pub document_annotation: Option, - pub usage_info: Option, - #[serde(flatten)] - pub extra_fields: Map, -} diff --git a/litellm-rust/crates/core/src/ocr/codecs/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/mod.rs deleted file mode 100644 index 639b985b9ae..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/mod.rs +++ /dev/null @@ -1,5 +0,0 @@ -pub(crate) mod cohere; -pub(crate) mod deepseek; -pub(crate) mod document_intelligence; -pub(crate) mod mistral; -pub(crate) mod reducto; diff --git a/litellm-rust/crates/core/src/ocr/codecs/reducto/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/reducto/mod.rs deleted file mode 100644 index 3fff40451c6..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/reducto/mod.rs +++ /dev/null @@ -1,9 +0,0 @@ -mod transformation; -mod types; - -pub(crate) use transformation::{ - transform_legacy_ocr_request, transform_ocr_response, transform_v3_ocr_request, -}; -pub(crate) use types::{ - ReductoLegacyParams, ReductoResponse, ReductoUploadResponse, ReductoV3Params, -}; diff --git a/litellm-rust/crates/core/src/ocr/codecs/reducto/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/reducto/transformation.rs deleted file mode 100644 index f4c8338c134..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/reducto/transformation.rs +++ /dev/null @@ -1,103 +0,0 @@ -use std::collections::BTreeMap; - -use serde_json::{Value, json}; - -use super::types::*; -use crate::ocr::error::{OcrRequestError, OcrResponseError}; -use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; - -pub(crate) fn transform_v3_ocr_request( - _model: &str, - document: OcrDocument, - params: &ReductoV3Params, -) -> Result { - Ok(ReductoV3Request { - input: document.source().to_string(), - params: params.clone(), - }) -} - -pub(crate) fn transform_legacy_ocr_request( - _model: &str, - document: OcrDocument, - params: &ReductoLegacyParams, -) -> Result { - Ok(ReductoLegacyRequest { - document_url: document.source().to_string(), - options: params.enhance.as_ref().map(|_| params.clone()), - }) -} - -pub(crate) fn transform_ocr_response( - model: &str, - response: ReductoResponse, -) -> Result { - let result = match response.result { - Some(result) => result.unwrap_or_default(), - None => ReductoResult { - chunks: response.chunks, - }, - }; - let usage = response.usage.unwrap_or_default(); - Ok(LiteLLMOcrResponse { - pages: build_pages(result.chunks.unwrap_or_default()), - model: model.to_string(), - document_annotation: None, - usage_info: Some(json!({ - "pages_processed": usage.num_pages, - "credits": usage.credits, - })), - object: "ocr".to_string(), - extra_fields: serde_json::Map::new(), - provider_native_response: None, - }) -} - -fn build_pages(chunks: Vec) -> Vec { - let blocks_by_page = chunks - .iter() - .flat_map(|chunk| chunk.blocks.iter().flatten()) - .filter_map(|block| block.bbox.as_ref()?.page.map(|page| (page, block))) - .fold( - BTreeMap::>::new(), - |mut pages, (page, block)| { - pages.entry(page).or_default().push(block); - pages - }, - ); - if blocks_by_page.is_empty() { - let markdown = join_content(chunks.iter().map(|chunk| chunk.content.as_deref())); - return if markdown.is_empty() { - Vec::new() - } else { - vec![page(0, markdown, None)] - }; - } - blocks_by_page - .into_iter() - .map(|(index, blocks)| { - let markdown = join_content(blocks.iter().map(|block| block.content.as_deref())); - page( - index.saturating_sub(1).max(0), - markdown, - Some(json!(blocks)), - ) - }) - .collect() -} - -fn join_content<'a>(content: impl Iterator>) -> String { - content - .flatten() - .filter(|text| !text.is_empty()) - .collect::>() - .join("\n\n") -} - -fn page(index: i64, markdown: String, blocks: Option) -> Value { - let mut result = json!({"index":index,"markdown":markdown,"images":null}); - if let (Value::Object(fields), Some(blocks)) = (&mut result, blocks) { - fields.insert("blocks".into(), blocks); - } - result -} diff --git a/litellm-rust/crates/core/src/ocr/codecs/reducto/types.rs b/litellm-rust/crates/core/src/ocr/codecs/reducto/types.rs deleted file mode 100644 index c03720cc8ae..00000000000 --- a/litellm-rust/crates/core/src/ocr/codecs/reducto/types.rs +++ /dev/null @@ -1,128 +0,0 @@ -use serde::{Deserialize, Deserializer, Serialize}; -use serde_json::{Map, Value}; - -#[derive(Clone, Debug, Default, Serialize, Deserialize)] -pub(crate) struct ReductoV3Params { - #[serde(skip_serializing_if = "Option::is_none")] - pub formatting: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub retrieval: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub settings: Option>, -} - -#[derive(Clone, Debug, Default, Serialize, Deserialize)] -pub(crate) struct ReductoLegacyParams { - #[serde(skip_serializing_if = "Option::is_none")] - pub enhance: Option>, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub(crate) struct ReductoV3Request { - pub input: String, - #[serde(flatten)] - pub params: ReductoV3Params, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub(crate) struct ReductoLegacyRequest { - pub document_url: String, - #[serde(skip_serializing_if = "Option::is_none")] - pub options: Option, -} - -#[derive(Deserialize)] -pub(crate) struct ReductoUploadResponse { - pub file_id: Option, -} - -#[derive(Clone, Debug, Deserialize)] -pub(crate) struct ReductoResponse { - #[serde(default, deserialize_with = "present_nullable")] - pub result: Option>, - pub usage: Option, - #[serde(default)] - pub chunks: Option>, -} - -fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>( - deserializer: D, -) -> Result>, D::Error> { - Option::::deserialize(deserializer).map(Some) -} - -#[derive(Clone, Debug, Default, Deserialize)] -pub(crate) struct ReductoResult { - pub chunks: Option>, -} - -#[derive(Clone, Debug, Default, Deserialize)] -pub(crate) struct ReductoUsage { - #[serde(default, deserialize_with = "optional_i64")] - pub num_pages: Option, - #[serde(default, deserialize_with = "optional_f64")] - pub credits: Option, -} - -#[derive(Clone, Debug, Deserialize)] -pub(crate) struct ReductoChunk { - pub content: Option, - pub blocks: Option>, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub(crate) struct ReductoBlock { - #[serde(skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub bbox: Option, - #[serde(flatten)] - pub extra_fields: Map, -} - -#[derive(Clone, Debug, Serialize, Deserialize)] -pub(crate) struct ReductoBoundingBox { - #[serde(default, deserialize_with = "optional_i64")] - pub page: Option, - #[serde(flatten)] - pub extra_fields: Map, -} - -fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { - match Option::::deserialize(deserializer)? { - None | Some(Value::Null) => Ok(None), - Some(Value::Number(number)) => number - .as_i64() - .or_else(|| number.as_f64().and_then(checked_truncated_i64)) - .map(Some) - .ok_or_else(|| serde::de::Error::custom("expected an integer")), - Some(Value::String(value)) => value - .trim() - .parse::() - .map(Some) - .map_err(|_| serde::de::Error::custom("expected an integer")), - Some(Value::Bool(value)) => Ok(Some(i64::from(value))), - Some(_) => Ok(None), - } -} - -fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { - match Option::::deserialize(deserializer)? { - None | Some(Value::Null) => Ok(None), - Some(Value::Number(number)) => number - .as_f64() - .map(Some) - .ok_or_else(|| serde::de::Error::custom("expected a number")), - Some(Value::String(value)) => value - .trim() - .parse::() - .map(Some) - .map_err(|_| serde::de::Error::custom("expected a number")), - Some(_) => Ok(None), - } -} - -fn checked_truncated_i64(value: f64) -> Option { - (value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64) - .then(|| value.trunc() as i64) -} diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index cd1d538aaa8..96dfa19d7cc 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -1,11 +1,27 @@ +use std::sync::Arc; + use super::OcrClient; -use super::adapters::OcrAdapter; use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest}; -use super::registry::OcrAdapterKind; +use super::provider_config::OcrConfigKind; use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; +use super::wire::DecodedOcrResponse; use crate::Error; use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext}; -use std::sync::Arc; +use crate::llms::azure_ai::ocr::cohere_parse_transformation::AzureAICohereParseConfig; +use crate::llms::azure_ai::ocr::document_intelligence::AzureDocumentIntelligenceOperation; +use crate::llms::azure_ai::ocr::document_intelligence::transformation::AzureDocumentIntelligenceOCRConfig; +use crate::llms::azure_ai::ocr::transformation::AzureAIOCRConfig; +use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; +use crate::llms::cohere::ocr::CohereResponse; +use crate::llms::cohere::ocr::transformation::CohereParseConfig; +use crate::llms::mistral::ocr::MistralOcrResponse; +use crate::llms::mistral::ocr::transformation::MistralOCRConfig; +use crate::llms::reducto::ocr::ReductoResponse; +use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, ReductoParseV3Config}; +use crate::llms::vertex_ai::ocr::deepseek_transformation::{ + DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig, +}; +use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig; pub(crate) async fn perform_ocr_request( client: &OcrClient, @@ -15,7 +31,7 @@ pub(crate) async fn perform_ocr_request( let context = CallLifecycleContext::new( "ocr", request.model.clone(), - request.adapter.provider().as_str(), + request.config.provider().as_str(), request .litellm_call_id .clone() @@ -47,14 +63,45 @@ impl PreparedOcrCall { client: OcrClient, request: LiteLLMOcrRequest, ) -> Result { - macro_rules! prepare_adapter { - ($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => { - match request.adapter { - $( OcrAdapterKind::$variant => $instance.prepare_request(&request, &client).await?, )+ - } - }; - } - let http = super::adapters::for_each_ocr_adapter!(prepare_adapter); + let http = match request.config { + OcrConfigKind::Cohere => CohereParseConfig.prepare_request(&request, &client).await?, + OcrConfigKind::Mistral => MistralOCRConfig.prepare_request(&request, &client).await?, + OcrConfigKind::AzureAi => { + AzureAIOCRConfig::default() + .prepare_request(&request, &client) + .await? + } + OcrConfigKind::AzureCohere => { + AzureAICohereParseConfig::default() + .prepare_request(&request, &client) + .await? + } + OcrConfigKind::AzureDocumentIntelligence => { + AzureDocumentIntelligenceOCRConfig + .prepare_request(&request, &client) + .await? + } + OcrConfigKind::ReductoLegacy => { + ReductoParseLegacyConfig + .prepare_request(&request, &client) + .await? + } + OcrConfigKind::ReductoV3 => { + ReductoParseV3Config + .prepare_request(&request, &client) + .await? + } + OcrConfigKind::VertexAi => { + VertexAIOCRConfig::default() + .prepare_request(&request, &client) + .await? + } + OcrConfigKind::VertexDeepSeek => { + VertexAIDeepSeekOCRConfig + .prepare_request(&request, &client) + .await? + } + }; Ok(Self { client, request, @@ -71,20 +118,57 @@ impl PreparedOcrCall { )) .await .map_err(super::client::transport_error)?; - macro_rules! read_adapter { - ($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => { - match self.request.adapter { - $( OcrAdapterKind::$variant => { - let decoded = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?; - Ok(OcrProviderResponse { - request: self.request, - data: OcrProviderData::$variant(decoded), - }) - }, )+ - } - }; - } - super::adapters::for_each_ocr_adapter!(read_adapter) + let data = match self.request.config { + OcrConfigKind::Cohere => OcrProviderData::Cohere( + CohereParseConfig + .read_response(&self.client, response, &url, &headers, &self.request) + .await?, + ), + OcrConfigKind::Mistral => OcrProviderData::Mistral( + MistralOCRConfig + .read_response(&self.client, response, &url, &headers, &self.request) + .await?, + ), + OcrConfigKind::AzureAi => OcrProviderData::AzureAi( + AzureAIOCRConfig::default() + .read_response(&self.client, response, &url, &headers, &self.request) + .await?, + ), + OcrConfigKind::AzureCohere => OcrProviderData::AzureCohere( + AzureAICohereParseConfig::default() + .read_response(&self.client, response, &url, &headers, &self.request) + .await?, + ), + OcrConfigKind::AzureDocumentIntelligence => OcrProviderData::AzureDocumentIntelligence( + AzureDocumentIntelligenceOCRConfig + .read_response(&self.client, response, &url, &headers, &self.request) + .await?, + ), + OcrConfigKind::ReductoLegacy => OcrProviderData::ReductoLegacy( + ReductoParseLegacyConfig + .read_response(&self.client, response, &url, &headers, &self.request) + .await?, + ), + OcrConfigKind::ReductoV3 => OcrProviderData::ReductoV3( + ReductoParseV3Config + .read_response(&self.client, response, &url, &headers, &self.request) + .await?, + ), + OcrConfigKind::VertexAi => OcrProviderData::VertexAi( + VertexAIOCRConfig::default() + .read_response(&self.client, response, &url, &headers, &self.request) + .await?, + ), + OcrConfigKind::VertexDeepSeek => OcrProviderData::VertexDeepSeek( + VertexAIDeepSeekOCRConfig + .read_response(&self.client, response, &url, &headers, &self.request) + .await?, + ), + }; + Ok(OcrProviderResponse { + request: self.request, + data, + }) } } @@ -104,23 +188,16 @@ fn request_headers(request: &reqwest::Request) -> Result, .collect() } -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 { - match self.data { - $( OcrProviderData::$variant(decoded) => { - let response = $instance.transform_ocr_response(&self.request, decoded.data)?; - Ok(LiteLLMOcrResponse { provider_native_response: decoded.native, ..response }) - }, )+ - } - } - } - }; +enum OcrProviderData { + Cohere(DecodedOcrResponse), + Mistral(DecodedOcrResponse), + AzureAi(DecodedOcrResponse), + AzureCohere(DecodedOcrResponse), + AzureDocumentIntelligence(DecodedOcrResponse), + ReductoLegacy(DecodedOcrResponse), + ReductoV3(DecodedOcrResponse), + VertexAi(DecodedOcrResponse), + VertexDeepSeek(DecodedOcrResponse), } pub(crate) struct OcrProviderResponse { @@ -128,6 +205,55 @@ pub(crate) struct OcrProviderResponse { data: OcrProviderData, } +impl OcrProviderResponse { + pub(crate) fn normalize(self) -> Result { + let (response, native) = match self.data { + OcrProviderData::Cohere(decoded) => ( + CohereParseConfig.transform_ocr_response(&self.request, decoded.data)?, + decoded.native, + ), + OcrProviderData::Mistral(decoded) => ( + MistralOCRConfig.transform_ocr_response(&self.request, decoded.data)?, + decoded.native, + ), + OcrProviderData::AzureAi(decoded) => ( + AzureAIOCRConfig::default().transform_ocr_response(&self.request, decoded.data)?, + decoded.native, + ), + OcrProviderData::AzureCohere(decoded) => ( + AzureAICohereParseConfig::default() + .transform_ocr_response(&self.request, decoded.data)?, + decoded.native, + ), + OcrProviderData::AzureDocumentIntelligence(decoded) => ( + AzureDocumentIntelligenceOCRConfig + .transform_ocr_response(&self.request, decoded.data)?, + decoded.native, + ), + OcrProviderData::ReductoLegacy(decoded) => ( + ReductoParseLegacyConfig.transform_ocr_response(&self.request, decoded.data)?, + decoded.native, + ), + OcrProviderData::ReductoV3(decoded) => ( + ReductoParseV3Config.transform_ocr_response(&self.request, decoded.data)?, + decoded.native, + ), + OcrProviderData::VertexAi(decoded) => ( + VertexAIOCRConfig::default().transform_ocr_response(&self.request, decoded.data)?, + decoded.native, + ), + OcrProviderData::VertexDeepSeek(decoded) => ( + VertexAIDeepSeekOCRConfig.transform_ocr_response(&self.request, decoded.data)?, + decoded.native, + ), + }; + Ok(LiteLLMOcrResponse { + provider_native_response: native, + ..response + }) + } +} + pub(crate) async fn post_call(hooks: &Arc, bytes: &[u8]) -> Result<(), Error> { let original_response = serde_json::Value::String(String::from_utf8_lossy(bytes).into_owned()); hooks @@ -135,5 +261,3 @@ pub(crate) async fn post_call(hooks: &Arc, bytes: &[u8]) -> Result .await?; Ok(()) } - -super::adapters::for_each_ocr_adapter!(provider_data); diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index e29fd6ac572..a66eafc7aae 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,13 +1,11 @@ -mod adapters; pub mod client; -mod codecs; -mod document; +pub(crate) mod document; pub mod error; -mod handler; +pub(crate) mod handler; pub mod hooks; mod lifecycle; -mod prepare; -mod registry; +pub(crate) mod prepare; +mod provider_config; pub mod types; pub mod wire; diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 241cb09529a..af342e519e6 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -83,7 +83,7 @@ where .hooks .during_call(OcrDuringCallRequest { model: request.model.clone(), - custom_llm_provider: request.adapter.provider().as_str().into(), + custom_llm_provider: request.config.provider().as_str().into(), url: url.into(), headers: headers.to_vec(), body, @@ -135,7 +135,7 @@ pub(crate) async fn guardrail_document( .hooks .during_call(OcrDuringCallRequest { model: request.model.clone(), - custom_llm_provider: request.adapter.provider().as_str().into(), + custom_llm_provider: request.config.provider().as_str().into(), url: url.into(), headers: headers.to_vec(), body: serde_json::to_value(&request.document).map_err(|_| { diff --git a/litellm-rust/crates/core/src/ocr/provider_config.rs b/litellm-rust/crates/core/src/ocr/provider_config.rs new file mode 100644 index 00000000000..f361ca49b0c --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/provider_config.rs @@ -0,0 +1,130 @@ +use crate::Error; +use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum OcrConfigKind { + Cohere, + Mistral, + AzureAi, + AzureCohere, + AzureDocumentIntelligence, + ReductoLegacy, + ReductoV3, + VertexAi, + VertexDeepSeek, +} + +impl OcrConfigKind { + pub(crate) const fn provider(self) -> OcrProvider { + match self { + Self::Cohere => OcrProvider::Cohere, + Self::Mistral => OcrProvider::Mistral, + Self::AzureAi | Self::AzureCohere | Self::AzureDocumentIntelligence => { + OcrProvider::AzureAi + } + Self::ReductoLegacy | Self::ReductoV3 => OcrProvider::Reducto, + Self::VertexAi | Self::VertexDeepSeek => OcrProvider::VertexAi, + } + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) enum OcrProvider { + Cohere, + Mistral, + AzureAi, + Reducto, + VertexAi, +} + +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", + Self::VertexAi => "vertex_ai", + } + } +} + +pub(crate) fn resolve_provider_config( + model: &str, + custom_llm_provider: Option<&str>, +) -> Result<(String, OcrConfigKind), Error> { + let provider = + get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider { + model, + custom_llm_provider: OcrProvider::Mistral.as_str(), + }); + let config = match provider.custom_llm_provider { + "cohere" => OcrConfigKind::Cohere, + "mistral" => OcrConfigKind::Mistral, + "azure_ai" if is_document_intelligence_model(provider.model) => { + OcrConfigKind::AzureDocumentIntelligence + } + "azure_ai" + if provider.model.to_ascii_lowercase().contains("cohere") + && provider.model.to_ascii_lowercase().contains("parse") => + { + OcrConfigKind::AzureCohere + } + "azure_ai" => OcrConfigKind::AzureAi, + "reducto" if provider.model.eq_ignore_ascii_case("parse-legacy") => { + OcrConfigKind::ReductoLegacy + } + "reducto" => OcrConfigKind::ReductoV3, + "vertex_ai" if provider.model.to_ascii_lowercase().contains("deepseek") => { + OcrConfigKind::VertexDeepSeek + } + "vertex_ai" => OcrConfigKind::VertexAi, + value => return Err(Error::InvalidProvider(value.to_string())), + }; + Ok((provider.model.to_string(), config)) +} + +fn is_document_intelligence_model(model: &str) -> bool { + let model = model.to_ascii_lowercase(); + model.contains("doc-intelligence") || model.contains("documentintelligence") +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn provider_models_are_preserved_without_a_local_allowlist() { + for (qualified_model, expected_config) in [ + ("mistral/future-ocr-model", OcrConfigKind::Mistral), + ("azure_ai/future-ocr-model", OcrConfigKind::AzureAi), + ] { + let expected_model = qualified_model.split_once('/').unwrap().1; + let (model, config) = resolve_provider_config(qualified_model, None).unwrap(); + assert_eq!(model, expected_model); + assert_eq!(config, expected_config); + } + } + + #[test] + fn provider_specific_models_select_their_config() { + assert_eq!( + resolve_provider_config("reducto/parse-legacy", None) + .unwrap() + .1, + OcrConfigKind::ReductoLegacy + ); + assert_eq!( + resolve_provider_config("reducto/future-parse-model", None) + .unwrap() + .1, + OcrConfigKind::ReductoV3 + ); + assert_eq!( + resolve_provider_config("azure_ai/doc-intelligence/prebuilt-layout", None) + .unwrap() + .1, + OcrConfigKind::AzureDocumentIntelligence + ); + } +} diff --git a/litellm-rust/crates/core/src/ocr/registry.rs b/litellm-rust/crates/core/src/ocr/registry.rs deleted file mode 100644 index 54f57cc3ca2..00000000000 --- a/litellm-rust/crates/core/src/ocr/registry.rs +++ /dev/null @@ -1,132 +0,0 @@ -use super::adapters::OcrAdapter; -use crate::Error; -use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider}; - -macro_rules! define_adapter_types { - ($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => { - #[derive(Clone, Copy, Debug, PartialEq, Eq)] - pub(crate) enum OcrAdapterKind { - $( $variant, )+ - } - - impl OcrAdapterKind { - pub(crate) const fn provider(self) -> OcrProvider { - match self { - $( Self::$variant => <$adapter>::PROVIDER, )+ - } - } - } - }; -} - -super::adapters::for_each_ocr_adapter!(define_adapter_types); - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub(crate) enum OcrProvider { - Cohere, - Mistral, - AzureAi, - Reducto, - VertexAi, -} - -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", - Self::VertexAi => "vertex_ai", - } - } -} - -pub(crate) fn resolve_wire_adapter( - model: &str, - custom_llm_provider: Option<&str>, -) -> Result<(String, OcrAdapterKind), Error> { - let provider = - get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider { - model, - 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, - "vertex_ai" => OcrProvider::VertexAi, - 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 - } - OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-v3") => { - OcrAdapterKind::ReductoV3 - } - OcrProvider::Reducto => OcrAdapterKind::ReductoV3, - OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => { - OcrAdapterKind::VertexDeepSeek - } - OcrProvider::VertexAi => OcrAdapterKind::VertexMistral, - }; - Ok((provider.model.to_string(), adapter)) -} - -fn is_document_intelligence_model(model: &str) -> bool { - let model = model.to_ascii_lowercase(); - model.contains("doc-intelligence") || model.contains("documentintelligence") -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn provider_models_are_preserved_without_a_local_allowlist() { - let cases = [ - ("mistral/future-ocr-model", OcrAdapterKind::Mistral), - ("azure_ai/future-ocr-model", OcrAdapterKind::AzureMistral), - ]; - - for (qualified_model, expected_adapter) in cases { - let expected_model = qualified_model.split_once('/').unwrap().1; - let (model, adapter) = resolve_wire_adapter(qualified_model, None).unwrap(); - assert_eq!(model, expected_model); - assert_eq!(adapter, expected_adapter); - } - } - - #[test] - fn unknown_reducto_models_use_the_current_protocol() { - let (model, adapter) = resolve_wire_adapter("reducto/future-parse-model", None).unwrap(); - assert_eq!(model, "future-parse-model"); - assert_eq!(adapter, OcrAdapterKind::ReductoV3); - } - - #[test] - fn known_protocol_models_still_select_specialized_adapters() { - let (model, adapter) = resolve_wire_adapter("reducto/parse-legacy", None).unwrap(); - assert_eq!(model, "parse-legacy"); - assert_eq!(adapter, OcrAdapterKind::ReductoLegacy); - - let (model, adapter) = - resolve_wire_adapter("azure_ai/doc-intelligence/prebuilt-layout", None).unwrap(); - assert_eq!(model, "doc-intelligence/prebuilt-layout"); - assert_eq!(adapter, OcrAdapterKind::AzureDocumentIntelligence); - } -} diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 76df8b42806..544f46fb797 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -6,7 +6,7 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; use super::hooks::{NoopOcrHooks, OcrHooks}; -use super::registry::{OcrAdapterKind, resolve_wire_adapter}; +use super::provider_config::{OcrConfigKind, resolve_provider_config}; use crate::Error; use crate::auth::{InputSource, TokenProviderHandle}; use crate::constants::OCR_HTTP_TIMEOUT_SECS; @@ -98,7 +98,7 @@ pub struct LiteLLMOcrRequest { pub optional_params: Map, pub input_sources: BTreeMap, pub azure_ad_token_provider: Option, - pub(crate) adapter: OcrAdapterKind, + pub(crate) config: OcrConfigKind, } impl LiteLLMOcrRequest { @@ -108,7 +108,7 @@ impl LiteLLMOcrRequest { custom_llm_provider: Option<&str>, optional_params: Map, ) -> Result { - let (model, adapter_kind) = resolve_wire_adapter(&model, custom_llm_provider)?; + let (model, config) = resolve_provider_config(&model, custom_llm_provider)?; Ok(Self { model, @@ -119,7 +119,7 @@ impl LiteLLMOcrRequest { optional_params, input_sources: BTreeMap::new(), azure_ad_token_provider: None, - adapter: adapter_kind, + config, }) } @@ -137,7 +137,7 @@ impl LiteLLMOcrRequest { } pub fn provider_name(&self) -> &'static str { - self.adapter.provider().as_str() + self.config.provider().as_str() } pub fn with_host_hooks( diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index 6dc6b34b73d..dc006f44252 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -83,31 +83,31 @@ pub struct OcrWireRequest { } pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> bool { - super::registry::resolve_wire_adapter(model, custom_llm_provider).is_ok() + super::provider_config::resolve_provider_config(model, custom_llm_provider).is_ok() } pub fn consumed_optional_param_names( model: &str, custom_llm_provider: Option<&str>, ) -> Result, Error> { - use super::registry::OcrAdapterKind; + use super::provider_config::OcrConfigKind; - 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 => { + let (_, config) = super::provider_config::resolve_provider_config(model, custom_llm_provider)?; + let provider_fields: &[&str] = match config { + OcrConfigKind::Cohere | OcrConfigKind::AzureCohere => &["output_format"], + OcrConfigKind::Mistral | OcrConfigKind::AzureAi | OcrConfigKind::VertexAi => { MISTRAL_OPTION_FIELDS } - OcrAdapterKind::AzureDocumentIntelligence => DOCUMENT_INTELLIGENCE_OPTION_FIELDS, - OcrAdapterKind::ReductoV3 => REDUCTO_V3_OPTION_FIELDS, - OcrAdapterKind::ReductoLegacy => REDUCTO_LEGACY_OPTION_FIELDS, - OcrAdapterKind::VertexDeepSeek => DEEPSEEK_OPTION_FIELDS, + OcrConfigKind::AzureDocumentIntelligence => DOCUMENT_INTELLIGENCE_OPTION_FIELDS, + OcrConfigKind::ReductoV3 => REDUCTO_V3_OPTION_FIELDS, + OcrConfigKind::ReductoLegacy => REDUCTO_LEGACY_OPTION_FIELDS, + OcrConfigKind::VertexDeepSeek => DEEPSEEK_OPTION_FIELDS, }; - let auth_fields: &[&str] = match adapter { - OcrAdapterKind::AzureMistral - | OcrAdapterKind::AzureDocumentIntelligence - | OcrAdapterKind::AzureCohere => AZURE_AUTH_OPTION_FIELDS, - OcrAdapterKind::VertexMistral | OcrAdapterKind::VertexDeepSeek => VERTEX_AUTH_OPTION_FIELDS, + let auth_fields: &[&str] = match config { + OcrConfigKind::AzureAi + | OcrConfigKind::AzureDocumentIntelligence + | OcrConfigKind::AzureCohere => AZURE_AUTH_OPTION_FIELDS, + OcrConfigKind::VertexAi | OcrConfigKind::VertexDeepSeek => VERTEX_AUTH_OPTION_FIELDS, _ => &[], }; Ok(COMMON_OPTION_FIELDS diff --git a/litellm-rust/crates/core/tests/deepseek_ocr.rs b/litellm-rust/crates/core/tests/deepseek_ocr.rs index 4ba39561dcd..ad37846eda1 100644 --- a/litellm-rust/crates/core/tests/deepseek_ocr.rs +++ b/litellm-rust/crates/core/tests/deepseek_ocr.rs @@ -1,7 +1,7 @@ use rstest::rstest; use serde_json::{Value, json}; -use crate::ocr::codecs::deepseek::{ +use crate::llms::vertex_ai::ocr::deepseek_transformation::{ DeepSeekOcrParams, DeepSeekOcrResponse, transform_ocr_request, transform_ocr_response, }; use crate::ocr::types::OcrDocument; @@ -65,7 +65,10 @@ fn request_maps_both_document_types_to_image_content(#[case] document: Value) { #[case(json!("[]"), "[]")] #[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")] #[case(json!({"pages":[{"markdown":"object"}]}), "object")] -fn response_codec_handles_text_json_and_objects(#[case] content: Value, #[case] expected: &str) { +fn response_transform_handles_text_json_and_objects( + #[case] content: Value, + #[case] expected: &str, +) { let response: DeepSeekOcrResponse = serde_json::from_value( json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}), ) @@ -102,7 +105,7 @@ fn structured_result_maps_pages_usage_model_and_annotation() { } #[test] -fn response_codec_rejects_missing_empty_and_malformed_content() { +fn response_transform_rejects_missing_empty_and_malformed_content() { for value in [ json!({"choices":[]}), json!({"choices":[{"message":{"content":""}}]}), diff --git a/litellm-rust/crates/core/tests/reducto_ocr.rs b/litellm-rust/crates/core/tests/reducto_ocr.rs index a15e9cae5b5..dde2a67add7 100644 --- a/litellm-rust/crates/core/tests/reducto_ocr.rs +++ b/litellm-rust/crates/core/tests/reducto_ocr.rs @@ -182,7 +182,7 @@ async fn rejects_invalid_document_sources_before_network(#[case] source: &str) { #[test] fn response_normalization_groups_blocks_and_distinguishes_null_result() { - use crate::ocr::codecs::reducto::{ReductoResponse, transform_ocr_response}; + use crate::llms::reducto::ocr::transformation::{ReductoResponse, transform_ocr_response}; let raw = json!({"usage":{"num_pages":"2","credits":"3"},"result":{"type":"full","chunks":[ {"blocks":[{ diff --git a/litellm-rust/crates/core/tests/vertex_ai_ocr.rs b/litellm-rust/crates/core/tests/vertex_ai_ocr.rs index 96a19dd62b4..adbf6de7434 100644 --- a/litellm-rust/crates/core/tests/vertex_ai_ocr.rs +++ b/litellm-rust/crates/core/tests/vertex_ai_ocr.rs @@ -96,10 +96,12 @@ async fn request_controlled_api_base_is_rejected_before_vertex_auth() { } #[tokio::test] -async fn adapters_build_complete_requests_and_share_mistral_normalization() { +async fn configs_build_complete_requests_and_share_mistral_normalization() { use std::time::Duration; - use crate::ocr::adapters::{MistralAdapter, OcrAdapter, VertexMistralAdapter}; + use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; + use crate::llms::mistral::ocr::transformation::MistralOCRConfig; + use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig; use crate::ocr::test_support::ocr_client; let client = ocr_client(); @@ -116,11 +118,11 @@ async fn adapters_build_complete_requests_and_share_mistral_normalization() { options.clone(), ); let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options); - let direct_http = MistralAdapter + let direct_http = MistralOCRConfig .prepare_request(&direct, &client) .await .unwrap(); - let vertex_http = VertexMistralAdapter + let vertex_http = VertexAIOCRConfig::default() .prepare_request(&vertex, &client) .await .unwrap(); @@ -146,11 +148,11 @@ async fn adapters_build_complete_requests_and_share_mistral_normalization() { ); } let payload = json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"}); - let direct_response = MistralAdapter + let direct_response = MistralOCRConfig .transform_ocr_response(&direct, serde_json::from_value(payload.clone()).unwrap()) .unwrap() .into_json(); - let vertex_response = VertexMistralAdapter + let vertex_response = VertexAIOCRConfig::default() .transform_ocr_response(&vertex, serde_json::from_value(payload).unwrap()) .unwrap() .into_json();