diff --git a/litellm-rust/crates/core/src/call_arguments.rs b/litellm-rust/crates/core/src/call_arguments.rs index 6bd7fb4f8dc..67852cef27d 100644 --- a/litellm-rust/crates/core/src/call_arguments.rs +++ b/litellm-rust/crates/core/src/call_arguments.rs @@ -23,11 +23,12 @@ pub struct ArgumentError { } pub fn parse_options(arguments: &CallArguments) -> Result { - use serde::de::IntoDeserializer; - serde_path_to_error::deserialize(Value::Object(arguments.0.clone()).into_deserializer()) - .map_err(|error| ArgumentError { - path: error.path().to_string(), - }) + let deserializer = serde::de::value::MapDeserializer::new( + arguments.iter().map(|(name, value)| (name.as_str(), value)), + ); + serde_path_to_error::deserialize(deserializer).map_err(|error| ArgumentError { + path: error.path().to_string(), + }) } #[derive(Clone, Copy, Debug, PartialEq, Eq)] diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs index 9f28dc8e0ff..4af36cd9be9 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs @@ -1,13 +1,12 @@ +use crate::call_arguments::CallArguments; use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext}; use crate::llms::cohere::ocr::transformation::{CohereParseConfig, CohereRequest}; use crate::llms::cohere::ocr::{CohereOptions, validate_document}; use crate::ocr::OcrClient; use crate::ocr::document::{inline_remote_document, validate_inline_document}; -use crate::ocr::prepare::{credential_env, transform_request_body}; use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument, PreparedOcrRequest}; use crate::url_utils::ApiUrl; - -const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE"; +use serde_json::Value; #[derive(Default)] pub(crate) struct AzureAICohereParseConfig; @@ -44,17 +43,10 @@ impl BaseOcrConfig for AzureAICohereParseConfig { _params: &Self::OcrParams, _environment: &Self::Environment, ) -> Result { - let base = request - .connection - .api_base - .clone() - .or_else(|| credential_env(AZURE_AI_API_BASE_ENV)) - .filter(|base| !base.trim().is_empty()) - .ok_or_else(|| { - crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication( - "Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(), - )) - })?; + let base = super::transformation::AzureAIOCRConfig::resolve_api_base( + request.connection.api_base.as_deref(), + &crate::ocr::prepare::credential_env, + )?; self.get_complete_url(&base) } @@ -72,6 +64,14 @@ impl BaseOcrConfig for AzureAICohereParseConfig { CohereParseConfig.get_supported_ocr_params(model) } + fn map_ocr_params( + &self, + arguments: &CallArguments, + model: &str, + ) -> Result { + CohereParseConfig.map_ocr_params(arguments, model) + } + async fn async_transform_ocr_request( &self, model: &str, @@ -98,37 +98,15 @@ impl BaseOcrConfig for AzureAICohereParseConfig { ) -> Result { CohereParseConfig.transform_ocr_response(model, raw_response, request_format) } -} -impl AzureAICohereParseConfig { - pub(crate) async fn prepare_request( - &self, - request: &PreparedOcrRequest, - client: &OcrClient, - ) -> Result { - let params = self.map_ocr_params(&request.optional_params, &request.model)?; - let url = BaseOcrConfig::get_complete_url(self, request, ¶ms, &Vec::new())?; - let headers = self.validate_environment(request, client).await?; - let remote = request.document.source().starts_with("http://") - || request.document.source().starts_with("https://"); - let body = self - .async_transform_ocr_request( - &request.model, - request.document.clone(), - ¶ms, - &headers, - OcrRequestContext { - client, - connection: &request.connection, - }, - ) - .await?; - transform_request_body(client, request, &url, &headers, !remote, body, |body| { - let document = crate::ocr::prepare::body_document(body)?; - validate_document(&document)?; - validate_inline_document(&document) - }) - .await + fn retains_document(&self, document: &OcrDocument) -> bool { + !document.is_remote() + } + + fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> { + let document = crate::ocr::prepare::body_document(body)?; + validate_document(&document)?; + validate_inline_document(&document) } } 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 index 8f12dbefa27..742f032e790 100644 --- 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 @@ -24,37 +24,16 @@ use crate::ocr::OcrClient; use crate::ocr::client::read_json_response; use crate::ocr::document::InlineDocument; use crate::ocr::hooks::OcrHooks; -use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::json::DecodedOcrResponse; +use crate::ocr::prepare::credential_env; use crate::ocr::types::{ LiteLLMOcrResponse, OcrConnection, OcrCredentialInputs, OcrDocument, OcrPage, OcrPageDimensions, OcrResponseFormat, OcrUsageInfo, PreparedOcrRequest, ResolvedOcrCredentials, }; -use crate::ocr::wire::DecodedOcrResponse; use crate::serde_compat::{FiniteF64, LaxI64}; use crate::url_utils::ApiUrl; -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(untagged)] -enum PagesInput { - ZeroBasedIndices(Vec), - NativeTokens(Vec), - NativeRange(String), -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -#[serde(untagged)] -enum FeaturesInput { - Names(Vec), - CommaSeparated(String), -} - -#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] -struct DocumentIntelligenceInputParams { - pub pages: Option, - pub features: Option, -} - -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[derive(Clone, Debug, PartialEq, Serialize)] pub(crate) struct DocumentIntelligenceParams { #[serde(skip_serializing_if = "Option::is_none")] pub pages: Option, @@ -145,79 +124,46 @@ struct AzureDocumentIntelligenceLine { pub content: Option, } -fn decode_input_params( - params: Map, - prefix: &str, -) -> Result { - if let Some(Value::Array(pages)) = params.get("pages") { - if pages.iter().any(Value::is_boolean) { - return Err(crate::ocr::Error::Pages("boolean page index".into())); - } - if pages - .iter() - .any(|page| page.is_number() && page.as_i64().is_none()) - { - return Err(crate::ocr::Error::Pages( - "page index is out of range".into(), - )); - } - if !pages.iter().all(Value::is_i64) && !pages.iter().all(Value::is_string) { - return Err(crate::ocr::Error::Pages("mixed page element types".into())); - } - } - crate::ocr::wire::decode_request_value(Value::Object(params), prefix) -} - -fn normalize_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, crate::ocr::Error> { +fn normalize_pages(pages: Option<&Value>) -> Result, crate::ocr::Error> { let normalized = match pages { - PagesInput::ZeroBasedIndices(indices) => { - if indices.is_empty() { - return Ok(None); - } - indices - .into_iter() - .map(|page| { - if page < 0 { - return Err(crate::ocr::Error::Pages("negative page index".into())); - } - page.checked_add(1).ok_or_else(|| { - crate::ocr::Error::Pages("page index is out of range".into()) - }) + None | Some(Value::Null) => return Ok(None), + Some(Value::Array(pages)) if pages.is_empty() => return Ok(None), + Some(Value::Array(pages)) if pages.iter().all(Value::is_number) => pages + .iter() + .map(|page| { + let page = page + .as_i64() + .ok_or_else(|| crate::ocr::Error::Pages("page index is out of range".into()))?; + if page < 0 { + return Err(crate::ocr::Error::Pages("negative page index".into())); + } + page.checked_add(1) + .ok_or_else(|| crate::ocr::Error::Pages("page index is out of range".into())) + }) + .collect::, _>>()? + .into_iter() + .map(|page| page.to_string()) + .collect::>() + .join(","), + Some(Value::Array(tokens)) => tokens + .iter() + .map(|token| { + token.as_str().map(str::trim).ok_or_else(|| { + crate::ocr::Error::Pages("expected only integers or only strings".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 + }) + .collect::, _>>()? + .join(","), + Some(Value::String(range)) => range .split(',') .map(str::trim) .collect::>() .join(","), + Some(_) => { + return Err(crate::ocr::Error::Pages( + "expected an array of integers or strings, or a native page range".into(), + )); + } }; if !normalized.split(',').all(valid_page_token) { return Err(crate::ocr::Error::Pages("invalid native page range".into())); @@ -241,10 +187,15 @@ fn valid_page_token(token: &str) -> bool { } } -fn normalize_features(features: FeaturesInput) -> Result, crate::ocr::Error> { +fn normalize_features(features: Option<&Value>) -> Result, crate::ocr::Error> { let tokens = match features { - FeaturesInput::Names(names) => names, - FeaturesInput::CommaSeparated(names) => names.split(',').map(str::to_string).collect(), + None | Some(Value::Null) => return Ok(None), + Some(Value::Array(names)) => names + .iter() + .map(|name| name.as_str().ok_or(crate::ocr::Error::Features)) + .collect::, _>>()?, + Some(Value::String(names)) => names.split(',').collect(), + Some(_) => return Err(crate::ocr::Error::Features), }; if tokens.is_empty() { return Ok(None); @@ -375,7 +326,7 @@ async fn read_operation_response( crate::ocr::client::read_response_bytes(response, connection.max_response_bytes) .await?; crate::ocr::handler::post_call(hooks, &bytes).await?; - return crate::ocr::wire::decode_response(&bytes, native); + return crate::ocr::json::decode_response(&bytes, native); } let location = response .headers() @@ -530,10 +481,10 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { arguments: &CallArguments, _model: &str, ) -> Result { - normalize_ocr_params(decode_input_params( - arguments.select(&["pages", "features"]), - "optional_params", - )?) + Ok(DocumentIntelligenceParams { + pages: normalize_pages(arguments.get("pages"))?, + features: normalize_features(arguments.get("features"))?, + }) } async fn async_transform_ocr_request( @@ -591,30 +542,10 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { ) -> Result { build_request(document) } -} -impl AzureDocumentIntelligenceOCRConfig { - pub(crate) async fn prepare_request( - &self, - request: &PreparedOcrRequest, - client: &OcrClient, - ) -> Result { - let params = self.map_ocr_params(&request.optional_params, &request.model)?; - let headers = BaseOcrConfig::validate_environment(self, request, client).await?; - let url = BaseOcrConfig::get_complete_url(self, request, ¶ms, &headers)?; - let body = self - .async_transform_ocr_request( - &request.model, - request.document.clone(), - ¶ms, - &headers, - OcrRequestContext { - client, - connection: &request.connection, - }, - ) - .await?; - transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await + /// The body is `urlSource`/`base64Source`, not a `document` field. + fn retains_document(&self, _document: &OcrDocument) -> bool { + false } } @@ -712,8 +643,8 @@ mod tests { use serde_json::{Value, json}; fn map(value: Value) -> Result { - let fields = value.as_object().unwrap().clone(); - normalize_ocr_params(decode_input_params(fields, "optional_params")?) + let arguments = serde_json::from_value(value).unwrap(); + AzureDocumentIntelligenceOCRConfig.map_ocr_params(&arguments, "model") } #[test] @@ -763,6 +694,8 @@ mod tests { #[case(json!([0, 1, 2]), Some("1,2,3"))] #[case(json!([2, 0, 0, 1]), Some("1,2,3"))] #[case(json!([]), None)] + #[case(Value::Null, None)] + #[case(json!([i64::MAX - 1]), Some("9223372036854775807"))] #[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"))] @@ -778,8 +711,14 @@ mod tests { #[case(json!([-1]))] #[case(json!([true, false]))] #[case(json!([1, "2"]))] + #[case(json!(["1", 2]))] + #[case(json!([1.0]))] + #[case(json!([i64::MAX]))] + #[case(json!([u64::MAX]))] + #[case(json!([null]))] + #[case(json!([[1]]))] #[case(json!(5))] - fn invalid_page_mapping_matches_python(#[case] input: Value) { + fn page_mapping_rejects_invalid_shapes_and_overflow(#[case] input: Value) { assert!(map(json!({"pages": input})).is_err()); } @@ -859,7 +798,6 @@ mod tests { use std::sync::{Arc, Mutex}; use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; - use crate::ocr::wire::{OcrWireRequest, decode_request}; fn query_value(url: &str, key: &str) -> Option { url::Url::parse(url) @@ -913,21 +851,12 @@ mod tests { json!({"features":"languages&pages=1"}), json!({"req_format":"azure"}), ] { - let result = decode_request(OcrWireRequest { - model: "azure_ai/doc-intelligence/prebuilt-read".into(), - document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), - api_key: Some("key".into()), - api_base: Some("http://127.0.0.1:1".into()), - custom_llm_provider: None, - extra_headers: None, - optional_params: options.as_object().unwrap().clone().into(), - input_sources: Default::default(), - timeout_seconds: None, - }); - let rejected = match result { - Ok(request) => perform_ocr(request).await.is_err(), - Err(_) => true, - }; + let request = wire_request( + "azure_ai/doc-intelligence/prebuilt-read", + "http://127.0.0.1:1", + options.clone(), + ); + let rejected = perform_ocr(request).await.is_err(); assert!(rejected, "accepted {options}"); } } diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs index e4e3a9fc570..a99a549cf84 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs @@ -1,14 +1,16 @@ +use crate::call_arguments::CallArguments; use crate::constants::AZURE_AI_OCR_PATH; use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext}; use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest}; use crate::ocr::OcrClient; use crate::ocr::document::{inline_remote_document, validate_inline_document}; -use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::prepare::credential_env; use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, PreparedOcrRequest}; use crate::params::OpaqueParams; use crate::url_utils::ApiUrl; use litellm_auth::{InputSource, Sourced}; use litellm_auth_azure::AzureAuthInputs; +use serde_json::Value; const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY"; const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE"; @@ -64,6 +66,14 @@ impl BaseOcrConfig for AzureAIOCRConfig { MistralOCRConfig.get_supported_ocr_params(model) } + fn map_ocr_params( + &self, + arguments: &CallArguments, + model: &str, + ) -> Result { + MistralOCRConfig.map_ocr_params(arguments, model) + } + async fn async_transform_ocr_request( &self, model: &str, @@ -89,53 +99,40 @@ impl BaseOcrConfig for AzureAIOCRConfig { ) -> Result { MistralOCRConfig.transform_ocr_response(model, raw_response, request_format) } -} -impl AzureAIOCRConfig { - pub(crate) async fn prepare_request( - &self, - request: &PreparedOcrRequest, - client: &OcrClient, - ) -> Result { - let params = self.map_ocr_params(&request.optional_params, &request.model)?; - let url = BaseOcrConfig::get_complete_url(self, request, ¶ms, &Vec::new())?; - let headers = BaseOcrConfig::validate_environment(self, request, client).await?; - let retains_document = !request.document.source().starts_with("http://") - && !request.document.source().starts_with("https://"); - let body = self - .async_transform_ocr_request( - &request.model, - request.document.clone(), - ¶ms, - &headers, - OcrRequestContext { - client, - connection: &request.connection, - }, - ) - .await?; - transform_request_body( - client, - request, - &url, - &headers, - retains_document, - body, - |body| validate_inline_document(&crate::ocr::prepare::body_document(body)?), - ) - .await + fn retains_document(&self, document: &OcrDocument) -> bool { + !document.is_remote() + } + + fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> { + validate_inline_document(&crate::ocr::prepare::body_document(body)?) } } impl AzureAIOCRConfig { + /// Python `AzureAIOCRConfig.validate_environment` requires the endpoint + /// before it resolves credentials; keep that order so a missing base is + /// reported without invoking any token provider. + pub(super) fn resolve_api_base( + api_base: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Result { + nonblank(api_base.map(str::to_string)) + .or_else(|| nonblank(env_lookup(AZURE_AI_API_BASE_ENV))) + .ok_or(crate::ocr::Error::Auth( + litellm_auth::Error::MissingApiBase { + provider: "Azure AI", + environment_variable: AZURE_AI_API_BASE_ENV, + }, + )) + } + fn get_complete_url( &self, api_base: Option<&str>, env_lookup: &dyn Fn(&str) -> Option, ) -> Result { - let base = nonblank(api_base.map(str::to_string)) - .or_else(|| nonblank(env_lookup(AZURE_AI_API_BASE_ENV))) - .ok_or_else(|| crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication("Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter".into())))?; + let base = Self::resolve_api_base(api_base, env_lookup)?; let path: Vec<&str> = AZURE_AI_OCR_PATH.trim_matches('/').split('/').collect(); ApiUrl::parse(&base) .and_then(|url| url.complete_path(&path)) @@ -151,6 +148,7 @@ impl AzureAIOCRConfig { config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), ) -> Result, crate::ocr::Error> { + Self::resolve_api_base(connection.api_base.as_deref(), env_lookup)?; if crate::http_utils::has_header(&connection.extra_headers, "authorization") { if config.azure_ad_token_provider.is_some() { super::common_utils::resolve_entra(config, env_lookup).await?; @@ -211,10 +209,24 @@ mod tests { ); } + #[test] + fn missing_api_base_is_structured() { + assert!(matches!( + AzureAIOCRConfig::resolve_api_base(None, &|_| None), + Err(crate::ocr::Error::Auth( + litellm_auth::Error::MissingApiBase { + provider: "Azure AI", + environment_variable: AZURE_AI_API_BASE_ENV, + } + )) + )); + } + #[tokio::test] async fn supplied_authorization_precedes_keys() { let connection = OcrConnection { api_key: Some("request-key".into()), + api_base: Some("https://example.com".into()), extra_headers: vec![("authorization".into(), "Bearer prepared".into())], ..Default::default() }; @@ -233,6 +245,7 @@ mod tests { async fn request_key_precedes_environment_key() { let connection = OcrConnection { api_key: Some("request-key".into()), + api_base: Some("https://example.com".into()), ..Default::default() }; assert_eq!( @@ -291,7 +304,7 @@ mod tests { use std::sync::Arc; - use serde_json::{Value, json}; + use serde_json::json; use crate::ocr::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks}; use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; 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 index 03ec373afea..33f3523f220 100644 --- a/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs @@ -3,8 +3,9 @@ use std::sync::Arc; use serde::Serialize; use serde::de::DeserializeOwned; +use serde_json::Value; -use crate::call_arguments::{CallArguments, parse_options}; +use crate::call_arguments::CallArguments; use crate::ocr::OcrClient; use crate::ocr::hooks::OcrHooks; use crate::ocr::types::{ @@ -12,12 +13,24 @@ use crate::ocr::types::{ PreparedOcrRequest, ResolvedOcrCredentials, }; +/// Output of `validate_environment`: whatever a provider resolves up front +/// (headers at minimum; Vertex also carries the project id). +pub(crate) trait OcrEnvironment: Send + Sync { + fn headers(&self) -> &[(String, String)]; +} + +impl OcrEnvironment for Vec<(String, String)> { + fn headers(&self) -> &[(String, String)] { + self + } +} + const HEALTH_CHECK_PDF_DATA_URI: &str = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y="; pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { - type OcrParams: DeserializeOwned + Send + Sync; + type OcrParams: Send + Sync; type ProviderRequest: Serialize + Send; - type Environment: Send + Sync; + type Environment: OcrEnvironment; fn get_api_key_env_var(&self) -> Option<&'static str> { None @@ -64,13 +77,7 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { &self, arguments: &CallArguments, model: &str, - ) -> Result { - Ok(parse_options( - &arguments - .select(self.get_supported_ocr_params(model)) - .into(), - )?) - } + ) -> Result; fn transform_ocr_request( &self, @@ -127,6 +134,57 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { headers, } } + + /// Whether the `document` field in the outgoing body is owned by the + /// provider transform and must survive guardrail body rewrites. + /// Providers that inline remote URLs return `false` for remote documents + /// so a hook may still replace the fetched payload. + fn retains_document(&self, _document: &OcrDocument) -> bool { + true + } + + /// Provider-specific check applied to the composed body, both before and + /// after guardrail hooks. Defaults to accepting any body. + fn validate_request_body(&self, _body: &Value) -> Result<(), crate::ocr::Error> { + Ok(()) + } + + /// Rust counterpart of `BaseLLMHTTPHandler._async_prepare_ocr_request`: + /// map params, validate environment, build URL, transform, compose body. + fn prepare_request( + &self, + request: &PreparedOcrRequest, + client: &OcrClient, + ) -> impl Future> + Send { + async move { + let params = self.map_ocr_params(&request.optional_params, &request.model)?; + let environment = self.validate_environment(request, client).await?; + let url = self.get_complete_url(request, ¶ms, &environment)?; + let headers = environment.headers(); + let body = self + .async_transform_ocr_request( + &request.model, + request.document.clone(), + ¶ms, + headers, + OcrRequestContext { + client, + connection: &request.connection, + }, + ) + .await?; + crate::ocr::prepare::transform_request_body( + client, + request, + &url, + headers, + self.retains_document(&request.document), + body, + |body| self.validate_request_body(body), + ) + .await + } + } } pub(crate) fn decode_and_normalize_response( @@ -135,7 +193,7 @@ pub(crate) fn decode_and_normalize_response( request_format: OcrResponseFormat, normalize: impl FnOnce(&str, T) -> Result, ) -> Result { - let decoded = crate::ocr::wire::decode_response( + let decoded = crate::ocr::json::decode_response( raw_response, request_format == OcrResponseFormat::Native, )?; diff --git a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs index 62ac1788d97..af593f465b4 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -2,11 +2,12 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; use serde_with::serde_as; +use crate::call_arguments::{CallArguments, parse_options}; use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE}; use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext}; use crate::ocr::OcrClient; use crate::ocr::document::InlineDocument; -use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::prepare::credential_env; use crate::ocr::types::{ LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageImage, OcrUsageInfo, PreparedOcrRequest, @@ -136,6 +137,14 @@ impl BaseOcrConfig for CohereParseConfig { &["output_format", "req_format"] } + fn map_ocr_params( + &self, + arguments: &CallArguments, + _model: &str, + ) -> Result { + Ok(parse_options(arguments)?) + } + async fn async_transform_ocr_request( &self, model: &str, @@ -160,33 +169,9 @@ impl BaseOcrConfig for CohereParseConfig { normalize_response, ) } -} -impl CohereParseConfig { - pub(crate) async fn prepare_request( - &self, - request: &PreparedOcrRequest, - client: &OcrClient, - ) -> Result { - let params = self.map_ocr_params(&request.optional_params, &request.model)?; - let headers = BaseOcrConfig::validate_environment(self, request, client).await?; - let url = BaseOcrConfig::get_complete_url(self, request, ¶ms, &headers)?; - let body = self - .async_transform_ocr_request( - &request.model, - request.document.clone(), - ¶ms, - &headers, - OcrRequestContext { - client, - connection: &request.connection, - }, - ) - .await?; - transform_request_body(client, request, &url, &headers, true, body, |body| { - validate_document(&crate::ocr::prepare::body_document(body)?) - }) - .await + fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> { + validate_document(&crate::ocr::prepare::body_document(body)?) } } @@ -255,7 +240,7 @@ fn page_image( if let Some(Value::Object(bbox)) = image.get("bounding_box") { image.insert("bbox".into(), Value::Object(bbox.clone())); } - crate::ocr::wire::decode_response_value(Value::Object(image), path) + crate::ocr::json::decode_response_value(Value::Object(image), path) } fn normalize_page(page: CoherePage, position: usize) -> Result { @@ -418,7 +403,11 @@ mod tests { assert_eq!(arguments["req_format"], "native"); assert_eq!(arguments["extension"], false); let invalid = serde_json::from_value(json!({"output_format":"html"})).unwrap(); - assert!(CohereParseConfig.map_ocr_params(&invalid, "parse").is_err()); + assert!(matches!( + CohereParseConfig.map_ocr_params(&invalid, "parse"), + Err(crate::ocr::Error::RequestField { path }) + if path == "optional_params.output_format" + )); } #[test] diff --git a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs index 350359bfa7c..fa6b0c3e098 100644 --- a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs @@ -1,10 +1,11 @@ use serde::{Deserialize, Serialize}; use serde_json::Value; +use crate::call_arguments::CallArguments; use crate::constants::MISTRAL_OCR_API_BASE; use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext}; use crate::ocr::OcrClient; -use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::prepare::credential_env; use crate::ocr::types::{ LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrUsageInfo, PreparedOcrRequest, }; @@ -96,6 +97,16 @@ impl BaseOcrConfig for MistralOCRConfig { ] } + fn map_ocr_params( + &self, + arguments: &CallArguments, + model: &str, + ) -> Result { + Ok(arguments + .select(self.get_supported_ocr_params(model)) + .into()) + } + async fn async_transform_ocr_request( &self, model: &str, @@ -122,31 +133,6 @@ impl BaseOcrConfig for MistralOCRConfig { } } -impl MistralOCRConfig { - pub(crate) async fn prepare_request( - &self, - request: &PreparedOcrRequest, - client: &OcrClient, - ) -> Result { - let params = self.map_ocr_params(&request.optional_params, &request.model)?; - let headers = BaseOcrConfig::validate_environment(self, request, client).await?; - let url = BaseOcrConfig::get_complete_url(self, request, ¶ms, &headers)?; - let body = self - .async_transform_ocr_request( - &request.model, - request.document.clone(), - ¶ms, - &headers, - OcrRequestContext { - client, - connection: &request.connection, - }, - ) - .await?; - transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await - } -} - pub(crate) fn normalize_response( model: &str, response: MistralOcrResponse, @@ -251,7 +237,7 @@ mod tests { "usage_info.pages_processed", ), ] { - let error = crate::ocr::wire::decode_response::( + let error = crate::ocr::json::decode_response::( &serde_json::to_vec(&payload).unwrap(), false, ) diff --git a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs index 0863e202575..7959d97dec5 100644 --- a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -3,7 +3,7 @@ use std::collections::BTreeMap; use serde::{Deserialize, Deserializer, Serialize}; use serde_json::{Map, Value, json}; -use crate::call_arguments::compose_body; +use crate::call_arguments::{CallArguments, compose_body}; use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX}; use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext}; use crate::ocr::OcrClient; @@ -117,6 +117,16 @@ impl BaseOcrConfig for ReductoParseV3Config { &["formatting", "retrieval", "settings"] } + fn map_ocr_params( + &self, + arguments: &CallArguments, + model: &str, + ) -> Result { + Ok(arguments + .select(self.get_supported_ocr_params(model)) + .into()) + } + #[tracing::instrument( name = "async_transform_ocr_request", target = "litellm::function_trace", @@ -151,36 +161,13 @@ impl BaseOcrConfig for ReductoParseV3Config { normalize_response, ) } -} -impl ReductoParseV3Config { - pub(crate) async fn prepare_request( + async fn prepare_request( &self, request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params(&request.optional_params, &request.model)?; - let headers = self.validate_environment(request, client).await?; - let url = self.get_complete_url(request, ¶ms, &headers)?; - let (document, headers) = guardrail_document(request, &url, &headers).await?; - let body = self - .async_transform_ocr_request( - &request.model, - document, - ¶ms, - &headers, - OcrRequestContext { - client, - connection: &request.connection, - }, - ) - .await?; - let body = compose_body( - &request.optional_params, - &body, - self.get_supported_ocr_params(&request.model), - )?; - build_http_request(client, request, &url, &headers, &body) + prepare_upload_request(self, request, client).await } } @@ -225,6 +212,16 @@ impl BaseOcrConfig for ReductoParseLegacyConfig { &["enhance"] } + fn map_ocr_params( + &self, + arguments: &CallArguments, + model: &str, + ) -> Result { + Ok(arguments + .select(self.get_supported_ocr_params(model)) + .into()) + } + #[tracing::instrument( name = "async_transform_ocr_request", target = "litellm::function_trace", @@ -251,39 +248,48 @@ impl BaseOcrConfig for ReductoParseLegacyConfig { ) -> Result { ReductoParseV3Config.transform_ocr_response(model, raw_response, request_format) } -} -impl ReductoParseLegacyConfig { - pub(crate) async fn prepare_request( + async fn prepare_request( &self, request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params(&request.optional_params, &request.model)?; - let headers = self.validate_environment(request, client).await?; - let url = self.get_complete_url(request, ¶ms, &headers)?; - let (document, headers) = guardrail_document(request, &url, &headers).await?; - let body = self - .async_transform_ocr_request( - &request.model, - document, - ¶ms, - &headers, - OcrRequestContext { - client, - connection: &request.connection, - }, - ) - .await?; - let body = compose_body( - &request.optional_params, - &body, - self.get_supported_ocr_params(&request.model), - )?; - build_http_request(client, request, &url, &headers, &body) + prepare_upload_request(self, request, client).await } } +/// Reducto differs from the shared `BaseOcrConfig::prepare_request` flow: +/// guardrails see the *source* document before it is uploaded, because the +/// final body only carries the opaque Reducto file id. +async fn prepare_upload_request>>( + config: &C, + request: &PreparedOcrRequest, + client: &OcrClient, +) -> Result { + let params = config.map_ocr_params(&request.optional_params, &request.model)?; + let headers = config.validate_environment(request, client).await?; + let url = config.get_complete_url(request, ¶ms, &headers)?; + let (document, headers) = guardrail_document(request, &url, &headers).await?; + let body = config + .async_transform_ocr_request( + &request.model, + document, + ¶ms, + &headers, + OcrRequestContext { + client, + connection: &request.connection, + }, + ) + .await?; + let body = compose_body( + &request.optional_params, + &body, + config.get_supported_ocr_params(&request.model), + )?; + build_http_request(client, request, &url, &headers, &body) +} + fn uploaded_file_id(document: OcrDocument) -> Result { if !document.source().starts_with(REDUCTO_ID_PREFIX) { return Err(crate::ocr::Error::ReductoSource); 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 index 1e8f607449d..80f4c0cd700 100644 --- 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 @@ -4,9 +4,10 @@ use serde_json::{Map, Value}; use litellm_auth_gcp::{self as vertex, VertexConfig}; use super::transformation::VertexAIOCRConfig; +use crate::call_arguments::CallArguments; use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext}; use crate::ocr::OcrClient; -use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::prepare::credential_env; use crate::ocr::types::{ LiteLLMOcrResponse, OcrDocument, OcrPage, OcrPageDimensions, OcrPageImage, OcrUsageInfo, PreparedOcrRequest, @@ -106,6 +107,14 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { VertexAIOCRConfig.get_api_key_env_var() } + fn map_ocr_params( + &self, + _arguments: &CallArguments, + _model: &str, + ) -> Result { + Ok(DeepSeekOcrParams::default()) + } + async fn validate_environment( &self, request: &PreparedOcrRequest, @@ -183,40 +192,11 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { .collect(), }) } -} -impl VertexAIDeepSeekOCRConfig { - pub(crate) async fn prepare_request( - &self, - request: &PreparedOcrRequest, - client: &OcrClient, - ) -> Result { - let params = self.map_ocr_params(&request.optional_params, &request.model)?; - let authentication = self.validate_environment(request, client).await?; - let url = BaseOcrConfig::get_complete_url(self, request, ¶ms, &authentication)?; - - let body = self - .async_transform_ocr_request( - &request.model, - request.document.clone(), - ¶ms, - &authentication.headers, - OcrRequestContext { - client, - connection: &request.connection, - }, - ) - .await?; - transform_request_body( - client, - request, - &url, - &authentication.headers, - false, - body, - |_| Ok(()), - ) - .await + /// The body carries the document inside `messages`, not a top-level + /// `document` field, so there is nothing for guardrails to retain. + fn retains_document(&self, _document: &OcrDocument) -> bool { + false } } @@ -267,7 +247,7 @@ pub(crate) fn normalize_response( .enumerate() .filter(|(_, page)| page.is_object()) .map(|(position, page)| { - let page: DeepSeekPage = crate::ocr::wire::decode_response_value( + let page: DeepSeekPage = crate::ocr::json::decode_response_value( page.clone(), &format!("choices[0].message.content.pages[{position}]"), )?; @@ -288,7 +268,7 @@ pub(crate) fn normalize_response( .or_else(|| (!has_pages).then_some(&response.usage)); let usage_info: Option = usage .filter(|usage| usage.is_object()) - .map(|usage| crate::ocr::wire::decode_response_value(usage.clone(), "usage_info")) + .map(|usage| crate::ocr::json::decode_response_value(usage.clone(), "usage_info")) .transpose()?; let model = match ocr_data.get("model") { Some(Value::String(model)) => model.clone(), @@ -683,11 +663,11 @@ mod tests { #[test] fn host_registration_selects_deepseek_without_affecting_mistral() { - assert!(crate::ocr::wire::is_supported_request( + assert!(crate::ocr::is_supported_request( "deepseek-ocr-maas", Some("vertex_ai") )); - assert!(crate::ocr::wire::is_supported_request( + assert!(crate::ocr::is_supported_request( "mistral-ocr-maas", Some("vertex_ai") )); diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs index 03a80a384c3..9d7774e9b65 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs @@ -1,11 +1,15 @@ use litellm_auth_gcp::{self as vertex, VertexConfig}; +use serde_json::Value; use super::common_utils::validate_destination; -use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext}; +use crate::call_arguments::CallArguments; +use crate::llms::base_llm::ocr::transformation::{ + BaseOcrConfig, OcrEnvironment, OcrRequestContext, +}; use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest}; use crate::ocr::OcrClient; use crate::ocr::document::{inline_remote_document, validate_inline_document}; -use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::prepare::credential_env; use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, PreparedOcrRequest}; use crate::params::OpaqueParams; use crate::url_utils::ApiUrl; @@ -71,6 +75,14 @@ impl BaseOcrConfig for VertexAIOCRConfig { MistralOCRConfig.get_supported_ocr_params(model) } + fn map_ocr_params( + &self, + arguments: &CallArguments, + model: &str, + ) -> Result { + MistralOCRConfig.map_ocr_params(arguments, model) + } + async fn async_transform_ocr_request( &self, model: &str, @@ -96,41 +108,19 @@ impl BaseOcrConfig for VertexAIOCRConfig { ) -> Result { MistralOCRConfig.transform_ocr_response(model, raw_response, request_format) } + + fn retains_document(&self, document: &OcrDocument) -> bool { + !document.is_remote() + } + + fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> { + validate_inline_document(&crate::ocr::prepare::body_document(body)?) + } } -impl VertexAIOCRConfig { - pub(crate) async fn prepare_request( - &self, - request: &PreparedOcrRequest, - client: &OcrClient, - ) -> Result { - let params = self.map_ocr_params(&request.optional_params, &request.model)?; - let authentication = BaseOcrConfig::validate_environment(self, request, client).await?; - let url = BaseOcrConfig::get_complete_url(self, request, ¶ms, &authentication)?; - let retains_document = !request.document.source().starts_with("http://") - && !request.document.source().starts_with("https://"); - let body = self - .async_transform_ocr_request( - &request.model, - request.document.clone(), - ¶ms, - &authentication.headers, - OcrRequestContext { - client, - connection: &request.connection, - }, - ) - .await?; - transform_request_body( - client, - request, - &url, - &authentication.headers, - retains_document, - body, - |body| validate_inline_document(&crate::ocr::prepare::body_document(body)?), - ) - .await +impl OcrEnvironment for vertex::VertexEnvironment { + fn headers(&self) -> &[(String, String)] { + &self.headers } } diff --git a/litellm-rust/crates/core/src/ocr/arguments.rs b/litellm-rust/crates/core/src/ocr/arguments.rs new file mode 100644 index 00000000000..293931e8bbb --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/arguments.rs @@ -0,0 +1,101 @@ +use crate::call_arguments::ArgumentSpec; + +use super::provider_config::{OcrConfigKind, resolve_provider_config}; + +const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body", "max_response_bytes"]; +const AZURE_AUTH_OPTION_FIELDS: &[&str] = &[ + "azure_ad_token", + "tenant_id", + "client_id", + "client_secret", + "azure_scope", + "azure_authority_host", + "azure_credential", + "azure_federated_token_file", + "enable_azure_ad_token_refresh", +]; +const VERTEX_AUTH_OPTION_FIELDS: &[&str] = &[ + "vertex_credentials", + "vertex_ai_credentials", + "vertex_project", + "vertex_ai_project", + "vertex_location", + "vertex_ai_location", +]; + +pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> bool { + resolve_provider_config(model, custom_llm_provider).is_ok() +} + +pub fn consumed_optional_param_names( + model: &str, + custom_llm_provider: Option<&str>, +) -> Result, super::Error> { + let (model, config) = resolve_provider_config(model, custom_llm_provider)?; + let provider_fields = config.get_supported_ocr_params(&model); + 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 + .iter() + .chain(provider_fields) + .chain(auth_fields) + .copied() + .collect()) +} + +pub fn consumed_optional_params( + model: &str, + custom_llm_provider: Option<&str>, +) -> Result, super::Error> { + consumed_optional_param_names(model, custom_llm_provider).map(|names| { + names + .into_iter() + .map(|name| ArgumentSpec { + name, + secret: matches!( + name, + "azure_ad_token" + | "client_secret" + | "azure_federated_token_file" + | "vertex_credentials" + | "vertex_ai_credentials" + ), + }) + .collect() + }) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn consumed_params_include_provider_options_and_mark_credentials() { + let mistral = consumed_optional_param_names("mistral/model", None).unwrap(); + assert!(mistral.contains(&"pages")); + assert!(mistral.contains(&"req_format")); + assert!(!mistral.contains(&"vertex_project")); + + let vertex = consumed_optional_param_names("vertex_ai/deepseek-ocr", None).unwrap(); + assert!(!vertex.contains(&"temperature")); + assert!(vertex.contains(&"vertex_credentials")); + assert!(!vertex.contains(&"pages")); + + let azure = consumed_optional_params("model", Some("azure_ai")).unwrap(); + assert!( + azure + .iter() + .any(|spec| spec.name == "client_secret" && spec.secret) + ); + assert!( + azure + .iter() + .any(|spec| spec.name == "tenant_id" && !spec.secret) + ); + } +} diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index c7a29e6e530..1c5b2dd68ad 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -4,8 +4,8 @@ use std::time::Duration; use bytes::{Bytes, BytesMut}; use serde::de::DeserializeOwned; +use super::json::{DecodedOcrResponse, decode_response}; use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; -use super::wire::{DecodedOcrResponse, decode_response}; use crate::constants::OCR_CONNECT_TIMEOUT_SECS; use crate::media::MediaFetcher; use litellm_auth_gcp::VertexAuth; diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index 14c8d10d3d2..827a1c47365 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -128,11 +128,11 @@ pub(crate) async fn inline_remote_document( document: OcrDocument, connection: &OcrConnection, ) -> Result { - let source = document.source(); - if !source.starts_with("http://") && !source.starts_with("https://") { + if !document.is_remote() { validate_inline_document(&document)?; return Ok(document); } + let source = document.source(); let url = Url::parse(source).map_err(|_| crate::ocr::Error::RequestField { path: "document URL".into(), })?; diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index 4d2f6dcc803..65ffaa38477 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -2,18 +2,9 @@ use std::sync::Arc; use super::OcrClient; use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest}; -use super::provider_config::OcrConfigKind; use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest}; use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext}; -use crate::llms::azure_ai::ocr::cohere_parse_transformation::AzureAICohereParseConfig; -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, OcrResponseContext}; -use crate::llms::cohere::ocr::transformation::CohereParseConfig; -use crate::llms::mistral::ocr::transformation::MistralOCRConfig; -use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, ReductoParseV3Config}; -use crate::llms::vertex_ai::ocr::deepseek_transformation::VertexAIDeepSeekOCRConfig; -use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig; +use crate::llms::base_llm::ocr::transformation::OcrResponseContext; pub(crate) async fn perform_ocr_request( client: &OcrClient, @@ -55,37 +46,7 @@ impl PreparedOcrCall { request: LiteLLMOcrRequest, ) -> Result { let request = super::prepare::prepare_request(request); - let http = match request.config { - OcrConfigKind::Cohere => CohereParseConfig.prepare_request(&request, &client).await?, - OcrConfigKind::Mistral => MistralOCRConfig.prepare_request(&request, &client).await?, - OcrConfigKind::AzureAi => AzureAIOCRConfig.prepare_request(&request, &client).await?, - OcrConfigKind::AzureCohere => { - AzureAICohereParseConfig - .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.prepare_request(&request, &client).await?, - OcrConfigKind::VertexDeepSeek => { - VertexAIDeepSeekOCRConfig - .prepare_request(&request, &client) - .await? - } - }; + let http = request.config.prepare_request(&request, &client).await?; Ok(Self { client, request, @@ -133,53 +94,10 @@ impl PreparedOcrCall { url: &url, headers: &headers, }; - match self.request.config { - OcrConfigKind::Cohere => { - CohereParseConfig - .async_transform_ocr_response(model, response, context) - .await - } - OcrConfigKind::Mistral => { - MistralOCRConfig - .async_transform_ocr_response(model, response, context) - .await - } - OcrConfigKind::AzureAi => { - AzureAIOCRConfig - .async_transform_ocr_response(model, response, context) - .await - } - OcrConfigKind::AzureCohere => { - AzureAICohereParseConfig - .async_transform_ocr_response(model, response, context) - .await - } - OcrConfigKind::AzureDocumentIntelligence => { - AzureDocumentIntelligenceOCRConfig - .async_transform_ocr_response(model, response, context) - .await - } - OcrConfigKind::ReductoLegacy => { - ReductoParseLegacyConfig - .async_transform_ocr_response(model, response, context) - .await - } - OcrConfigKind::ReductoV3 => { - ReductoParseV3Config - .async_transform_ocr_response(model, response, context) - .await - } - OcrConfigKind::VertexAi => { - VertexAIOCRConfig - .async_transform_ocr_response(model, response, context) - .await - } - OcrConfigKind::VertexDeepSeek => { - VertexAIDeepSeekOCRConfig - .async_transform_ocr_response(model, response, context) - .await - } - } + self.request + .config + .async_transform_ocr_response(model, response, context) + .await } } diff --git a/litellm-rust/crates/core/src/ocr/json.rs b/litellm-rust/crates/core/src/ocr/json.rs new file mode 100644 index 00000000000..d4651838a2d --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/json.rs @@ -0,0 +1,62 @@ +use serde::de::{DeserializeOwned, IntoDeserializer}; +use serde_json::{Map, Value}; + +#[derive(Debug)] +pub struct DecodedOcrResponse { + pub data: T, + pub native: Option>, + pub text: String, +} + +pub(crate) fn decode_request_value( + value: Value, + prefix: &str, +) -> Result { + serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| { + crate::ocr::Error::RequestField { + path: format!("{prefix}.{}", error.path()), + } + }) +} + +pub(crate) fn decode_response_value( + value: Value, + prefix: &str, +) -> Result { + serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| { + crate::ocr::Error::ResponseField { + path: format!("{prefix}.{}", error.path()), + } + }) +} + +pub(crate) fn decode_response( + bytes: &[u8], + native: bool, +) -> Result, crate::ocr::Error> { + let mut deserializer = serde_json::Deserializer::from_slice(bytes); + let data = serde_path_to_error::deserialize(&mut deserializer).map_err(|error| { + crate::ocr::Error::ResponseField { + path: error.path().to_string(), + } + })?; + deserializer + .end() + .map_err(|_| crate::ocr::Error::ResponseField { + path: "response".into(), + })?; + let native = if native { + Some( + serde_json::from_slice(bytes).map_err(|_| crate::ocr::Error::ResponseField { + path: "response".into(), + })?, + ) + } else { + None + }; + Ok(DecodedOcrResponse { + data, + native, + text: String::from_utf8_lossy(bytes).into_owned(), + }) +} diff --git a/litellm-rust/crates/core/src/ocr/lifecycle.rs b/litellm-rust/crates/core/src/ocr/lifecycle.rs index 2dcec88eae0..35289eaebde 100644 --- a/litellm-rust/crates/core/src/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/src/ocr/lifecycle.rs @@ -666,42 +666,37 @@ mod tests { OcrPreCallRequest, }; use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; - use crate::ocr::wire::{OcrWireRequest, decode_request}; use crate::ocr::{ - NativeOutcome, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline, OcrHost, - OcrHostOperation, OcrHostResult, + LiteLLMOcrRequest, NativeOutcome, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, + OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult, }; #[test] fn request_boundary_selects_mistral_and_rejects_unknown_providers() { - let request = OcrWireRequest { - model: "mistral/model".into(), - document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}), - api_key: Some("key".into()), - api_base: None, - custom_llm_provider: None, - extra_headers: None, - optional_params: json!({"extract_header":true,"unknown":42}) - .as_object() - .unwrap() - .clone() - .into(), - input_sources: Default::default(), - timeout_seconds: None, - }; - assert!(decode_request(request).is_ok()); + let document = OcrDocument::try_from( + json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}), + ) + .unwrap(); assert!( - decode_request(OcrWireRequest { - model: "model".into(), - document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}), - api_key: Some("key".into()), - api_base: None, - custom_llm_provider: Some("unknown".into()), - extra_headers: None, - optional_params: Default::default(), - input_sources: Default::default(), - timeout_seconds: None, - }) + LiteLLMOcrRequest::new( + "mistral/model".into(), + document.clone(), + None, + json!({"extract_header":true,"unknown":42}) + .as_object() + .unwrap() + .clone() + .into(), + ) + .is_ok() + ); + assert!( + LiteLLMOcrRequest::new( + "model".into(), + document, + Some("unknown"), + Default::default() + ) .is_err() ); } @@ -1507,11 +1502,20 @@ mod tests { json!(crate::constants::OCR_RESPONSE_MAX_BYTES + 1), Value::Null, ] { - let wire = serde_json::from_value(json!({ - "model": "mistral/model", "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, - "optional_params": {"max_response_bytes": value} - })).unwrap(); - let Err(error) = decode_request(wire) else { + let Err(error) = LiteLLMOcrRequest::new( + "mistral/model".into(), + OcrDocument::try_from(json!({ + "type": "document_url", + "document_url": "data:application/pdf;base64,YWJj" + })) + .unwrap(), + None, + json!({"max_response_bytes": value}) + .as_object() + .unwrap() + .clone() + .into(), + ) else { panic!("invalid response limit accepted") }; assert!(error.to_string().contains("max_response_bytes")); diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index e5d31a5fe59..bf9eecd5c79 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,15 +1,19 @@ mod error; pub use error::Error; +mod arguments; pub mod client; pub(crate) mod document; pub(crate) mod handler; pub mod hooks; +pub(crate) mod json; mod lifecycle; pub(crate) mod prepare; mod provider_config; pub mod types; -pub mod wire; +pub use arguments::{ + consumed_optional_param_names, consumed_optional_params, is_supported_request, +}; pub use client::{OcrClient, ocr}; pub use document::{encode_file_document, mime_type_for_name, upload_mime_type}; pub use lifecycle::{ @@ -18,8 +22,8 @@ pub use lifecycle::{ }; pub use provider_config::{get_api_key_env_var, get_health_check_document}; pub use types::{ - LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageDimensions, - OcrPageImage, OcrUsageInfo, + LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrCredentialInputs, OcrDocument, + OcrPage, OcrPageDimensions, OcrPageImage, OcrTransportConfig, OcrUsageInfo, }; #[cfg(test)] diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 089293c7303..ba173585cf2 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -104,7 +104,7 @@ pub(crate) async fn guardrail_document( retained_fields: Vec::new(), }) .await?; - let document = super::wire::decode_request_value(changed.body, "guardrail.document")?; + let document = super::json::decode_request_value(changed.body, "guardrail.document")?; Ok((document, changed.headers)) } @@ -120,7 +120,7 @@ pub(crate) fn body_document(body: &Value) -> Result { .filter(|(name, _)| matches!(name.as_str(), "type" | "image_url" | "document_url")) .map(|(name, value)| (name.clone(), value.clone())) .collect(); - super::wire::decode_request_value(Value::Object(source), "body.document") + super::json::decode_request_value(Value::Object(source), "body.document") } pub(crate) fn credential_env(name: &str) -> Option { diff --git a/litellm-rust/crates/core/src/ocr/provider_config.rs b/litellm-rust/crates/core/src/ocr/provider_config.rs index 57fefd1b051..7c8f2692894 100644 --- a/litellm-rust/crates/core/src/ocr/provider_config.rs +++ b/litellm-rust/crates/core/src/ocr/provider_config.rs @@ -1,8 +1,12 @@ -use super::types::{OcrCredentialInputs, OcrDocument, ResolvedOcrCredentials}; +use super::OcrClient; +use super::types::{ + LiteLLMOcrResponse, OcrCredentialInputs, OcrDocument, PreparedOcrRequest, + ResolvedOcrCredentials, +}; use crate::llms::azure_ai::ocr::cohere_parse_transformation::AzureAICohereParseConfig; 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::base_llm::ocr::transformation::{BaseOcrConfig, OcrResponseContext}; use crate::llms::cohere::ocr::transformation::CohereParseConfig; use crate::llms::mistral::ocr::transformation::MistralOCRConfig; use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, ReductoParseV3Config}; @@ -13,16 +17,22 @@ use strum::{EnumString, IntoStaticStr}; macro_rules! dispatch_config { ($config:expr, $method:ident($($argument:expr),* $(,)?)) => { + dispatch_config!(@arms $config, $method($($argument),*), ) + }; + ($config:expr, $method:ident($($argument:expr),* $(,)?).await) => { + dispatch_config!(@arms $config, $method($($argument),*), .await) + }; + (@arms $config:expr, $method:ident($($argument:expr),*), $($suffix:tt)*) => { match $config { - OcrConfigKind::Cohere => CohereParseConfig.$method($($argument),*), - OcrConfigKind::Mistral => MistralOCRConfig.$method($($argument),*), - OcrConfigKind::AzureAi => AzureAIOCRConfig.$method($($argument),*), - OcrConfigKind::AzureCohere => AzureAICohereParseConfig.$method($($argument),*), - OcrConfigKind::AzureDocumentIntelligence => AzureDocumentIntelligenceOCRConfig.$method($($argument),*), - OcrConfigKind::ReductoLegacy => ReductoParseLegacyConfig.$method($($argument),*), - OcrConfigKind::ReductoV3 => ReductoParseV3Config.$method($($argument),*), - OcrConfigKind::VertexAi => VertexAIOCRConfig.$method($($argument),*), - OcrConfigKind::VertexDeepSeek => VertexAIDeepSeekOCRConfig.$method($($argument),*), + OcrConfigKind::Cohere => CohereParseConfig.$method($($argument),*)$($suffix)*, + OcrConfigKind::Mistral => MistralOCRConfig.$method($($argument),*)$($suffix)*, + OcrConfigKind::AzureAi => AzureAIOCRConfig.$method($($argument),*)$($suffix)*, + OcrConfigKind::AzureCohere => AzureAICohereParseConfig.$method($($argument),*)$($suffix)*, + OcrConfigKind::AzureDocumentIntelligence => AzureDocumentIntelligenceOCRConfig.$method($($argument),*)$($suffix)*, + OcrConfigKind::ReductoLegacy => ReductoParseLegacyConfig.$method($($argument),*)$($suffix)*, + OcrConfigKind::ReductoV3 => ReductoParseV3Config.$method($($argument),*)$($suffix)*, + OcrConfigKind::VertexAi => VertexAIOCRConfig.$method($($argument),*)$($suffix)*, + OcrConfigKind::VertexDeepSeek => VertexAIDeepSeekOCRConfig.$method($($argument),*)$($suffix)*, } }; } @@ -80,6 +90,26 @@ impl OcrConfigKind { ) -> super::Error { dispatch_config!(self, get_error_class(message, status, headers)) } + + pub(crate) async fn prepare_request( + self, + request: &PreparedOcrRequest, + client: &OcrClient, + ) -> Result { + dispatch_config!(self, prepare_request(request, client).await) + } + + pub(crate) async fn async_transform_ocr_response( + self, + model: &str, + raw_response: reqwest::Response, + context: OcrResponseContext<'_>, + ) -> Result { + dispatch_config!( + self, + async_transform_ocr_response(model, raw_response, context).await + ) + } } pub fn get_api_key_env_var( @@ -157,79 +187,83 @@ fn is_document_intelligence_model(model: &str) -> bool { mod tests { use super::*; use litellm_auth::{InputSource, Sourced}; + use rstest::rstest; - #[test] - fn provider_names_round_trip_exactly() { - for provider in ["cohere", "mistral", "azure_ai", "reducto", "vertex_ai"] { - let (_, config) = resolve_provider_config("model", Some(provider)).unwrap(); - let resolved: &'static str = config.provider().into(); - assert_eq!(resolved, provider); - } - - for provider in ["Mistral", "unknown"] { - assert_eq!( - resolve_provider_config("model", Some(provider)), - Err(crate::ocr::Error::InvalidProvider(provider.into())) - ); - } + #[rstest] + #[case("cohere")] + #[case("mistral")] + #[case("azure_ai")] + #[case("reducto")] + #[case("vertex_ai")] + fn provider_names_round_trip_exactly(#[case] provider: &str) { + let (_, config) = resolve_provider_config("model", Some(provider)).unwrap(); + let resolved: &'static str = config.provider().into(); + assert_eq!(resolved, provider); } - #[test] - fn health_check_documents_are_valid_for_each_provider() { - for model in [ - "mistral/ocr", - "azure_ai/ocr", - "azure_ai/doc-intelligence/prebuilt-layout", - "reducto/parse-v3", - "vertex_ai/mistral-ocr", - "vertex_ai/deepseek-ocr", - ] { - let document = get_health_check_document(model, None).unwrap(); - assert!(matches!(document, OcrDocument::DocumentUrl { .. })); - let inline = crate::ocr::document::InlineDocument::parse(document.source()) - .unwrap() - .unwrap(); - assert_eq!(inline.mime_type().to_string(), "application/pdf"); - assert!(inline.decode(4096).unwrap().starts_with(b"%PDF-")); - } - for model in ["cohere/parse", "azure_ai/cohere-parse"] { - let document = get_health_check_document(model, None).unwrap(); - crate::llms::cohere::ocr::validate_document(&document).unwrap(); - let inline = crate::ocr::document::InlineDocument::parse(document.source()) - .unwrap() - .unwrap(); - assert_eq!(inline.mime_type().to_string(), "image/png"); - assert!( - inline - .decode(4096) - .unwrap() - .starts_with(b"\x89PNG\r\n\x1a\n") - ); - } + #[rstest] + #[case("Mistral")] + #[case("unknown")] + fn invalid_provider_names_are_rejected(#[case] provider: &str) { + assert_eq!( + resolve_provider_config("model", Some(provider)), + Err(crate::ocr::Error::InvalidProvider(provider.into())) + ); } - #[test] - fn api_key_metadata_follows_provider_overrides_and_python_defaults() { - for (model, expected) in [ - ("mistral/ocr", Some("MISTRAL_API_KEY")), - ("cohere/parse", Some("COHERE_API_KEY")), - ("azure_ai/ocr", Some("AZURE_AI_API_KEY")), - ("azure_ai/cohere-parse", Some("AZURE_AI_API_KEY")), - ( - "azure_ai/doc-intelligence/prebuilt-layout", - Some("AZURE_DOCUMENT_INTELLIGENCE_API_KEY"), - ), - ("vertex_ai/mistral-ocr", Some("VERTEX_AI_API_KEY")), - ("vertex_ai/deepseek-ocr", Some("VERTEX_AI_API_KEY")), - ("reducto/parse-v3", None), - ("reducto/parse-legacy", None), - ] { - assert_eq!( - get_api_key_env_var(model, None).unwrap(), - expected, - "{model}" - ); - } + #[rstest] + #[case("mistral/ocr")] + #[case("azure_ai/ocr")] + #[case("azure_ai/doc-intelligence/prebuilt-layout")] + #[case("reducto/parse-v3")] + #[case("vertex_ai/mistral-ocr")] + #[case("vertex_ai/deepseek-ocr")] + fn pdf_health_check_documents_are_valid(#[case] model: &str) { + let document = get_health_check_document(model, None).unwrap(); + assert!(matches!(document, OcrDocument::DocumentUrl { .. })); + let inline = crate::ocr::document::InlineDocument::parse(document.source()) + .unwrap() + .unwrap(); + assert_eq!(inline.mime_type().to_string(), "application/pdf"); + assert!(inline.decode(4096).unwrap().starts_with(b"%PDF-")); + } + + #[rstest] + #[case("cohere/parse")] + #[case("azure_ai/cohere-parse")] + fn png_health_check_documents_are_valid(#[case] model: &str) { + let document = get_health_check_document(model, None).unwrap(); + crate::llms::cohere::ocr::validate_document(&document).unwrap(); + let inline = crate::ocr::document::InlineDocument::parse(document.source()) + .unwrap() + .unwrap(); + assert_eq!(inline.mime_type().to_string(), "image/png"); + assert!( + inline + .decode(4096) + .unwrap() + .starts_with(b"\x89PNG\r\n\x1a\n") + ); + } + + #[rstest] + #[case("mistral/ocr", Some("MISTRAL_API_KEY"))] + #[case("cohere/parse", Some("COHERE_API_KEY"))] + #[case("azure_ai/ocr", Some("AZURE_AI_API_KEY"))] + #[case("azure_ai/cohere-parse", Some("AZURE_AI_API_KEY"))] + #[case( + "azure_ai/doc-intelligence/prebuilt-layout", + Some("AZURE_DOCUMENT_INTELLIGENCE_API_KEY") + )] + #[case("vertex_ai/mistral-ocr", Some("VERTEX_AI_API_KEY"))] + #[case("vertex_ai/deepseek-ocr", Some("VERTEX_AI_API_KEY"))] + #[case("reducto/parse-v3", None)] + #[case("reducto/parse-legacy", None)] + fn api_key_metadata_follows_provider_overrides_and_python_defaults( + #[case] model: &str, + #[case] expected: Option<&str>, + ) { + assert_eq!(get_api_key_env_var(model, None).unwrap(), expected); } #[test] @@ -268,110 +302,106 @@ mod tests { connection.api_base.as_ref().map(Sourced::source), Some(InputSource::Request) ); - for dynamic in [ - None, - Some(Sourced::new(String::new(), InputSource::Environment)), - ] { - let connection = - OcrConfigKind::Mistral.resolve_connection_params(OcrCredentialInputs { - api_key: Some(Sourced::new("explicit-key".into(), InputSource::Deployment)), - api_base: Some(Sourced::new( - "https://explicit.test".into(), - InputSource::Deployment, - )), - dynamic_api_key: dynamic.clone(), - dynamic_api_base: dynamic, - }); - assert_eq!( - connection - .api_key - .as_ref() - .map(|value| value.value().as_str()), - Some("explicit-key") - ); - assert_eq!( - connection - .api_base - .as_ref() - .map(|value| value.value().as_str()), - Some("https://explicit.test") - ); - } } - #[test] - fn document_intelligence_only_accepts_dynamic_values_for_explicit_fields() { - for (explicit_key, explicit_base) in [ - (None, None), - (Some("key"), None), - (None, Some("base")), - (Some("key"), Some("base")), - ] { - let connection = OcrConfigKind::AzureDocumentIntelligence.resolve_connection_params( - OcrCredentialInputs { - api_key: explicit_key - .map(|value| Sourced::new(value.to_string(), InputSource::Deployment)), - api_base: explicit_base - .map(|value| Sourced::new(value.to_string(), InputSource::Deployment)), - dynamic_api_key: Some(Sourced::new( - "dynamic-key".into(), - InputSource::Environment, - )), - dynamic_api_base: Some(Sourced::new( - "https://dynamic.test".into(), - InputSource::Deployment, - )), - }, - ); - assert_eq!( - connection - .api_key - .as_ref() - .map(|value| value.value().as_str()), - explicit_key.map(|_| "dynamic-key") - ); - assert_eq!( - connection - .api_base - .as_ref() - .map(|value| value.value().as_str()), - explicit_base.map(|_| "https://dynamic.test") - ); - } - } - - #[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() { + #[rstest] + #[case(None)] + #[case(Some(""))] + fn empty_or_missing_dynamic_credentials_preserve_explicit_values( + #[case] dynamic_value: Option<&str>, + ) { + let dynamic = + dynamic_value.map(|value| Sourced::new(value.into(), InputSource::Environment)); + let connection = OcrConfigKind::Mistral.resolve_connection_params(OcrCredentialInputs { + api_key: Some(Sourced::new("explicit-key".into(), InputSource::Deployment)), + api_base: Some(Sourced::new( + "https://explicit.test".into(), + InputSource::Deployment, + )), + dynamic_api_key: dynamic.clone(), + dynamic_api_base: dynamic, + }); assert_eq!( - resolve_provider_config("reducto/parse-legacy", None) - .unwrap() - .1, - OcrConfigKind::ReductoLegacy + connection + .api_key + .as_ref() + .map(|value| value.value().as_str()), + Some("explicit-key") ); assert_eq!( - resolve_provider_config("reducto/future-parse-model", None) - .unwrap() - .1, - OcrConfigKind::ReductoV3 + connection + .api_base + .as_ref() + .map(|value| value.value().as_str()), + Some("https://explicit.test") + ); + } + + #[rstest] + #[case(None, None)] + #[case(Some("key"), None)] + #[case(None, Some("base"))] + #[case(Some("key"), Some("base"))] + fn document_intelligence_only_accepts_dynamic_values_for_explicit_fields( + #[case] explicit_key: Option<&str>, + #[case] explicit_base: Option<&str>, + ) { + let connection = OcrConfigKind::AzureDocumentIntelligence.resolve_connection_params( + OcrCredentialInputs { + api_key: explicit_key + .map(|value| Sourced::new(value.into(), InputSource::Deployment)), + api_base: explicit_base + .map(|value| Sourced::new(value.into(), InputSource::Deployment)), + dynamic_api_key: Some(Sourced::new("dynamic-key".into(), InputSource::Environment)), + dynamic_api_base: Some(Sourced::new( + "https://dynamic.test".into(), + InputSource::Deployment, + )), + }, ); assert_eq!( - resolve_provider_config("azure_ai/doc-intelligence/prebuilt-layout", None) - .unwrap() - .1, - OcrConfigKind::AzureDocumentIntelligence + connection + .api_key + .as_ref() + .map(|value| value.value().as_str()), + explicit_key.map(|_| "dynamic-key") + ); + assert_eq!( + connection + .api_base + .as_ref() + .map(|value| value.value().as_str()), + explicit_base.map(|_| "https://dynamic.test") + ); + } + + #[rstest] + #[case("mistral/future-ocr-model", OcrConfigKind::Mistral)] + #[case("azure_ai/future-ocr-model", OcrConfigKind::AzureAi)] + fn provider_models_are_preserved_without_a_local_allowlist( + #[case] qualified_model: &str, + #[case] expected_config: OcrConfigKind, + ) { + 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); + } + + #[rstest] + #[case("reducto/parse-legacy", OcrConfigKind::ReductoLegacy)] + #[case("reducto/future-parse-model", OcrConfigKind::ReductoV3)] + #[case( + "azure_ai/doc-intelligence/prebuilt-layout", + OcrConfigKind::AzureDocumentIntelligence + )] + fn provider_specific_models_select_their_config( + #[case] model: &str, + #[case] expected_config: OcrConfigKind, + ) { + assert_eq!( + resolve_provider_config(model, None).unwrap().1, + expected_config ); } } diff --git a/litellm-rust/crates/core/src/ocr/test_support.rs b/litellm-rust/crates/core/src/ocr/test_support.rs index 59f4acb8a5b..784bc051b8f 100644 --- a/litellm-rust/crates/core/src/ocr/test_support.rs +++ b/litellm-rust/crates/core/src/ocr/test_support.rs @@ -4,8 +4,9 @@ use serde_json::{Value, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::TcpListener; -use crate::ocr::wire::{OcrWireRequest, decode_request}; -use crate::ocr::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient}; +use crate::ocr::{ + LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrCredentialInputs, OcrDocument, +}; pub(crate) fn ocr_client() -> OcrClient { let document_http = reqwest::Client::builder() @@ -22,18 +23,31 @@ pub(crate) async fn perform_ocr( } pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest { - decode_request(OcrWireRequest { - model: model.into(), - document: json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}), - api_key: Some("test-key".into()), - api_base: Some(base.into()), - custom_llm_provider: None, - extra_headers: None, - optional_params: options.as_object().unwrap().clone().into(), - input_sources: Default::default(), - timeout_seconds: Some(2.0), - }) - .unwrap() + let request = LiteLLMOcrRequest::new( + model.into(), + OcrDocument::try_from( + json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}), + ) + .unwrap(), + None, + options.as_object().unwrap().clone().into(), + ) + .unwrap(); + let transport = request.transport.clone().with_overrides( + Vec::new(), + Default::default(), + Some(std::time::Duration::from_secs(2)), + ); + request.with_connection_inputs( + OcrCredentialInputs::new( + Some("test-key".into()), + Default::default(), + Some(base.into()), + Default::default(), + ), + transport, + Default::default(), + ) } pub(crate) struct MockResponse { diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index dc4de6fc5d0..958d1b11e97 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -39,6 +39,11 @@ impl OcrDocument { } } + pub(crate) fn is_remote(&self) -> bool { + let source = self.source(); + source.starts_with("http://") || source.starts_with("https://") + } + pub(crate) fn with_source(self, source: String) -> Self { match self { Self::DocumentUrl { extra_fields, .. } => Self::DocumentUrl { @@ -53,6 +58,14 @@ impl OcrDocument { } } +impl TryFrom for OcrDocument { + type Error = super::Error; + + fn try_from(value: Value) -> Result { + super::json::decode_request_value(value, "document") + } +} + #[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "lowercase")] pub enum OcrResponseFormat { @@ -69,6 +82,22 @@ pub struct OcrCredentialInputs { pub dynamic_api_base: Option>, } +impl OcrCredentialInputs { + pub fn new( + api_key: Option, + api_key_source: InputSource, + api_base: Option, + api_base_source: InputSource, + ) -> Self { + Self { + api_key: nonblank(api_key).map(|value| Sourced::new(value, api_key_source)), + dynamic_api_key: None, + api_base: nonblank(api_base).map(|value| Sourced::new(value, api_base_source)), + dynamic_api_base: None, + } + } +} + #[derive(Clone)] pub struct OcrTransportConfig { pub extra_headers: Vec<(String, String)>, @@ -92,6 +121,28 @@ impl Default for OcrTransportConfig { } } +impl OcrTransportConfig { + pub fn with_overrides( + self, + extra_headers: Vec<(String, String)>, + extra_headers_source: InputSource, + timeout: Option, + ) -> Self { + Self { + extra_headers, + extra_headers_source, + timeout: timeout.unwrap_or(self.timeout), + ..self + } + } +} + +fn nonblank(value: Option) -> Option { + value + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) +} + #[derive(Clone)] pub struct OcrConnection { pub api_key: Option, @@ -169,12 +220,34 @@ impl LiteLLMOcrRequest { optional_params: CallArguments, ) -> Result { let (model, config) = resolve_provider_config(&model, custom_llm_provider)?; + let default_transport = OcrTransportConfig::default(); + let max_response_bytes = optional_params + .get("max_response_bytes") + .map(|value| { + value + .as_u64() + .and_then(|value| usize::try_from(value).ok()) + .filter(|value| *value > 0 && *value <= default_transport.max_response_bytes) + .ok_or_else(|| super::Error::RequestField { + path: "max_response_bytes".into(), + }) + }) + .transpose()? + .unwrap_or(default_transport.max_response_bytes); + let transport = OcrTransportConfig { + max_response_bytes, + ..default_transport + }; + let optional_params = optional_params + .into_iter() + .filter(|(name, _)| name != "max_response_bytes") + .collect(); Ok(Self { model, document, credentials: OcrCredentialInputs::default(), - transport: OcrTransportConfig::default(), + transport, hooks: Arc::new(NoopOcrHooks), litellm_call_id: None, optional_params, @@ -210,6 +283,20 @@ impl LiteLLMOcrRequest { ..self } } + + pub fn with_connection_inputs( + self, + credentials: OcrCredentialInputs, + transport: OcrTransportConfig, + input_sources: BTreeMap, + ) -> Self { + Self { + credentials, + transport, + input_sources, + ..self + } + } } pub(crate) struct PreparedOcrRequest { diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs deleted file mode 100644 index 73e1babdb49..00000000000 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ /dev/null @@ -1,345 +0,0 @@ -use std::collections::BTreeMap; -use std::time::Duration; - -use super::types::{LiteLLMOcrRequest, OcrDocument, OcrTransportConfig}; -use crate::call_arguments::{ArgumentSpec, CallArguments}; -use litellm_auth::InputSource; -use serde::{ - Deserialize, - de::{DeserializeOwned, IntoDeserializer}, -}; -use serde_json::{Map, Value}; - -const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body", "max_response_bytes"]; -pub const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"]; -const AZURE_AUTH_OPTION_FIELDS: &[&str] = &[ - "azure_ad_token", - "tenant_id", - "client_id", - "client_secret", - "azure_scope", - "azure_authority_host", - "azure_credential", - "azure_federated_token_file", - "enable_azure_ad_token_refresh", -]; -const VERTEX_AUTH_OPTION_FIELDS: &[&str] = &[ - "vertex_credentials", - "vertex_ai_credentials", - "vertex_project", - "vertex_ai_project", - "vertex_location", - "vertex_ai_location", -]; - -#[derive(Debug)] -pub struct DecodedOcrResponse { - pub data: T, - pub native: Option>, - pub text: String, -} - -#[derive(Deserialize)] -#[serde(deny_unknown_fields)] -pub struct OcrWireRequest { - pub model: String, - pub document: Value, - pub api_key: Option, - pub api_base: Option, - pub custom_llm_provider: Option, - pub extra_headers: Option>, - #[serde(default)] - pub optional_params: CallArguments, - #[serde(default)] - pub input_sources: BTreeMap, - pub timeout_seconds: Option, -} - -pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> bool { - 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, crate::ocr::Error> { - use super::provider_config::OcrConfigKind; - - let (model, config) = - super::provider_config::resolve_provider_config(model, custom_llm_provider)?; - let provider_fields = config.get_supported_ocr_params(&model); - 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 - .iter() - .chain(provider_fields) - .chain(auth_fields) - .copied() - .collect()) -} - -pub fn consumed_optional_params( - model: &str, - custom_llm_provider: Option<&str>, -) -> Result, crate::ocr::Error> { - consumed_optional_param_names(model, custom_llm_provider).map(|names| { - names - .into_iter() - .map(|name| ArgumentSpec { - name, - secret: matches!( - name, - "azure_ad_token" - | "client_secret" - | "azure_federated_token_file" - | "vertex_credentials" - | "vertex_ai_credentials" - ), - }) - .collect() - }) -} - -pub fn decode_request(wire: OcrWireRequest) -> Result { - let api_key_source = source_for(&wire.input_sources, "api_key"); - let api_base_source = source_for(&wire.input_sources, "api_base"); - let extra_headers_source = source_for(&wire.input_sources, "extra_headers"); - let document = decode_document(wire.document)?; - let headers = wire - .extra_headers - .unwrap_or_default() - .into_iter() - .map(|(name, value)| { - let value = value - .as_str() - .ok_or_else(|| crate::ocr::Error::RequestField { - path: format!("extra_headers.{name}"), - })?; - Ok((name, value.to_string())) - }) - .collect::, crate::ocr::Error>>()?; - let timeout = wire - .timeout_seconds - .map(|seconds| { - Duration::try_from_secs_f64(seconds).map_err(|_| crate::ocr::Error::RequestField { - path: "timeout_seconds".into(), - }) - }) - .transpose()?; - let defaults = OcrTransportConfig::default(); - let max_response_bytes = wire - .optional_params - .get("max_response_bytes") - .map(|value| { - value - .as_u64() - .and_then(|value| usize::try_from(value).ok()) - .filter(|value| *value > 0 && *value <= defaults.max_response_bytes) - .ok_or_else(|| crate::ocr::Error::RequestField { - path: "max_response_bytes".into(), - }) - }) - .transpose()? - .unwrap_or(defaults.max_response_bytes); - let request = LiteLLMOcrRequest::new( - wire.model, - document, - wire.custom_llm_provider.as_deref(), - wire.optional_params - .into_iter() - .filter(|(name, _)| name != "max_response_bytes") - .collect(), - )?; - let credentials = super::types::OcrCredentialInputs { - api_key: nonblank(wire.api_key) - .map(|value| litellm_auth::Sourced::new(value, api_key_source)), - dynamic_api_key: None, - api_base: nonblank(wire.api_base) - .map(|value| litellm_auth::Sourced::new(value, api_base_source)), - dynamic_api_base: None, - }; - let transport = OcrTransportConfig { - extra_headers: headers, - extra_headers_source, - timeout: timeout.unwrap_or(defaults.timeout), - max_download_bytes: defaults.max_download_bytes, - max_response_bytes, - poll_timeout: defaults.poll_timeout, - }; - Ok(LiteLLMOcrRequest { - credentials, - transport, - input_sources: wire.input_sources, - ..request - }) -} - -fn decode_document(value: Value) -> Result { - let kind = value.get("type").and_then(Value::as_str); - let missing_url = matches!(kind, Some("document_url")) && value.get("document_url").is_none() - || matches!(kind, Some("image_url")) && value.get("image_url").is_none(); - if missing_url { - return Err(crate::ocr::Error::MissingDocumentUrl); - } - decode_request_value(value, "document") -} - -fn source_for(sources: &BTreeMap, name: &str) -> InputSource { - sources.get(name).copied().unwrap_or_default() -} - -fn nonblank(value: Option) -> Option { - value - .map(|s| s.trim().to_string()) - .filter(|s| !s.is_empty()) -} -pub fn decode_request_value( - value: Value, - prefix: &str, -) -> Result { - serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| { - crate::ocr::Error::RequestField { - path: format!("{prefix}.{}", error.path()), - } - }) -} - -pub(crate) fn decode_response_value( - value: Value, - prefix: &str, -) -> Result { - serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| { - crate::ocr::Error::ResponseField { - path: format!("{prefix}.{}", error.path()), - } - }) -} - -pub fn decode_response( - bytes: &[u8], - native: bool, -) -> Result, crate::ocr::Error> { - let mut deserializer = serde_json::Deserializer::from_slice(bytes); - let data = serde_path_to_error::deserialize(&mut deserializer).map_err(|error| { - crate::ocr::Error::ResponseField { - path: error.path().to_string(), - } - })?; - deserializer - .end() - .map_err(|_| crate::ocr::Error::ResponseField { - path: "response".into(), - })?; - let native = if native { - Some( - serde_json::from_slice(bytes).map_err(|_| crate::ocr::Error::ResponseField { - path: "response".into(), - })?, - ) - } else { - None - }; - Ok(DecodedOcrResponse { - data, - native, - text: String::from_utf8_lossy(bytes).into_owned(), - }) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn core_selects_consumed_values_without_serializing_host_objects() { - let fields = consumed_optional_params("mistral/model", None).unwrap(); - use crate::call_arguments::should_project; - assert!(should_project("future_option", &fields, BOUND_FIELDS)); - assert!(should_project("extra_body", &fields, BOUND_FIELDS)); - assert!(should_project("id", &fields, BOUND_FIELDS)); - assert!(!should_project("metadata", &fields, BOUND_FIELDS)); - assert!(!should_project("callbacks", &fields, BOUND_FIELDS)); - assert!(!should_project("api_key", &fields, BOUND_FIELDS)); - assert!(!should_project("document", &fields, BOUND_FIELDS)); - } - - #[test] - fn option_projection_is_provider_specific_and_excludes_opaque_fields() { - let mistral = consumed_optional_param_names("mistral/model", None).unwrap(); - assert!(mistral.contains(&"pages")); - assert!(mistral.contains(&"req_format")); - assert!(!mistral.contains(&"vertex_project")); - assert!(!mistral.contains(&"opaque_extension")); - - let vertex = consumed_optional_param_names("vertex_ai/deepseek-ocr", None).unwrap(); - assert!(!vertex.contains(&"temperature")); - assert!(vertex.contains(&"vertex_credentials")); - assert!(!vertex.contains(&"pages")); - } - - #[test] - fn optional_param_metadata_marks_only_credentials_as_secret() { - let azure = consumed_optional_params("model", Some("azure_ai")).unwrap(); - assert!( - azure - .iter() - .any(|spec| spec.name == "client_secret" && spec.secret) - ); - assert!( - azure - .iter() - .any(|spec| spec.name == "tenant_id" && !spec.secret) - ); - let vertex = consumed_optional_params("deepseek-ocr", Some("vertex_ai")).unwrap(); - assert!( - vertex - .iter() - .any(|spec| spec.name == "vertex_credentials" && spec.secret) - ); - assert!( - vertex - .iter() - .any(|spec| spec.name == "vertex_project" && !spec.secret) - ); - } - - #[test] - fn activation_includes_migrated_providers() { - assert!(is_supported_request("model", Some("mistral"))); - assert!(is_supported_request("pixtral-12b", Some("azure_ai"))); - assert!(is_supported_request( - "documentintelligence/prebuilt-read", - Some("azure_ai") - )); - assert!(is_supported_request("parse-v3", Some("reducto"))); - assert!(is_supported_request("parse-legacy", Some("reducto"))); - assert!(is_supported_request("mistral-ocr", Some("vertex_ai"))); - assert!(is_supported_request("deepseek-ocr", Some("vertex_ai"))); - } - - #[test] - fn missing_document_source_has_a_typed_public_error() { - for document in [ - serde_json::json!({"type": "document_url"}), - serde_json::json!({"type": "image_url"}), - ] { - let wire = serde_json::from_value(serde_json::json!({ - "model": "mistral/model", - "document": document, - })) - .unwrap(); - let error = decode_request(wire).err().expect("missing document URL"); - assert_eq!(error, crate::ocr::Error::MissingDocumentUrl); - let error = crate::Error::from(error); - assert!(matches!( - error, - crate::Error::Ocr(crate::ocr::Error::MissingDocumentUrl) - )); - } - } -} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs index d43c2f88775..85c54823abe 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs @@ -9,8 +9,7 @@ use pyo3::types::PyDict; use pyo3::types::{PyBytes, PyString}; use litellm_core::constants::OCR_INLINE_MAX_BYTES; -use litellm_core::ocr::{OcrDocument, encode_file_document, mime_type_for_name, upload_mime_type}; -use litellm_python_interop::to_py_preserving_errors; +use litellm_core::ocr::{OcrDocument, encode_file_document}; enum FileBytes { Python(PyBackedBytes), @@ -136,44 +135,6 @@ pub(super) fn file_document(py: Python<'_>, document: FileDocumentInput) -> PyRe .map_err(|error| PyValueError::new_err(error.to_string())) } -#[pyfunction] -fn _ocr_file_document(py: Python<'_>, document: Bound<'_, PyAny>) -> PyResult> { - to_py_preserving_errors(py, &file_document(py, document.extract()?)?) -} - -#[pyfunction] -fn _ocr_mime_type(file_name: &str) -> String { - mime_type_for_name(file_name).into() -} - -#[pyfunction] -#[pyo3(signature = (file_content, file_name=None, content_type=None))] -fn _ocr_upload_document( - py: Python<'_>, - file_content: &Bound<'_, PyBytes>, - file_name: Option<&str>, - content_type: Option<&str>, -) -> PyResult> { - let bytes: PyBackedBytes = file_content.extract()?; - let document = py - .detach(|| { - encode_file_document( - &bytes, - None, - Some(upload_mime_type(file_name, content_type)), - ) - }) - .map_err(|error| PyValueError::new_err(error.to_string()))?; - to_py_preserving_errors(py, &document) -} - -pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - module.add("_OCR_MAX_FILE_BYTES", OCR_INLINE_MAX_BYTES)?; - module.add_function(wrap_pyfunction!(_ocr_upload_document, module)?)?; - module.add_function(wrap_pyfunction!(_ocr_file_document, module)?)?; - module.add_function(wrap_pyfunction!(_ocr_mime_type, module)?) -} - #[cfg(test)] mod tests { use super::*; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs index fe755f5982b..d41497c8291 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs @@ -10,7 +10,7 @@ use litellm_python_interop::{ use super::callbacks; use super::errors::to_pyerr as ocr_error_to_pyerr; -use super::project::{ProjectedOcrFields, admitted_call, project_request}; +use super::project::{ProjectedOcrCall, ProjectedOcrFields, PythonOcrInput, admitted_call}; use crate::lifecycle::{ OperationClass, PythonCallState, PythonRoute, missing_state, now, run_call, }; @@ -193,7 +193,10 @@ impl PythonRoute for PythonOcrHost { let OcrHostData::Unprojected { request } = &self.data else { return Err(missing_state()); }; - let projected = project_request(py, request.bind(py), self.state.kwargs.bind(py))?; + let projected = ProjectedOcrCall::try_from(PythonOcrInput { + request: request.bind(py), + kwargs: self.state.kwargs.bind(py), + })?; let has_token_provider = projected.fields.azure_ad_token_provider.is_some(); let request = projected.request; self.data = OcrHostData::Projected(Box::new(ProjectedOcrHost { diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index f17bf249b7f..83d12e3163b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -3,12 +3,12 @@ mod document; mod errors; mod lifecycle; mod project; +mod request; mod value; use pyo3::prelude::*; pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { value::register(module)?; - document::register(module)?; lifecycle::register(module) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 6b07fa02068..a5b94fb6da9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -1,7 +1,6 @@ use std::sync::Arc; -use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_params, decode_request}; -use litellm_core::ocr::{LiteLLMOcrRequest, NativeOutcome, OcrCall}; +use litellm_core::ocr::{LiteLLMOcrRequest, NativeOutcome, OcrCall, consumed_optional_params}; use litellm_python_interop::{ from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, }; @@ -11,10 +10,13 @@ use serde_json::{Map, Value}; use super::errors::to_pyerr as ocr_error_to_pyerr; use super::lifecycle::BridgeOcrHooks; +use super::request::BridgeOcrRequest; use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider}; use crate::errors::RustBridgeDeclined; use crate::marshal::{project_optional_fields, python_timeout_seconds, request_input_sources}; +const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"]; + pub(super) struct ProjectedOcrFields { pub boundary_request: Py, pub document: Py, @@ -29,6 +31,11 @@ pub(super) struct ProjectedOcrCall { pub fields: ProjectedOcrFields, } +pub(super) struct PythonOcrInput<'a, 'py> { + pub request: &'a Bound<'py, PyAny>, + pub kwargs: &'a Bound<'py, PyDict>, +} + struct PythonOcrFields<'a, 'py> { request: &'a Bound<'py, PyAny>, kwargs: &'a Bound<'py, PyDict>, @@ -110,60 +117,63 @@ impl ProjectedDocument { } } -pub(super) fn project_request( - py: Python<'_>, - request: &Bound<'_, PyAny>, - kwargs: &Bound<'_, PyDict>, -) -> PyResult { - let boundary_request = request.clone().unbind(); - let arguments = PythonOcrFields { request, kwargs }; - let model = arguments.model()?; - let custom_llm_provider = arguments.custom_llm_provider()?; - let (wire_document, retained_document) = - ProjectedDocument::project(py, &arguments.document()?)?.into_parts(); - let api_key = arguments.api_key()?; - let specs = consumed_optional_params(&model, custom_llm_provider.as_deref()) +impl TryFrom> for ProjectedOcrCall { + type Error = PyErr; + + fn try_from(input: PythonOcrInput<'_, '_>) -> PyResult { + let request = input.request; + let kwargs = input.kwargs; + let py = request.py(); + let boundary_request = request.clone().unbind(); + let arguments = PythonOcrFields { request, kwargs }; + let model = arguments.model()?; + let custom_llm_provider = arguments.custom_llm_provider()?; + let (wire_document, retained_document) = + ProjectedDocument::project(py, &arguments.document()?)?.into_parts(); + let api_key = arguments.api_key()?; + let specs = consumed_optional_params(&model, custom_llm_provider.as_deref()) + .map_err(ocr_error_to_pyerr)?; + let optional_params = project_optional_fields(kwargs, &specs, BOUND_FIELDS)?; + let input_sources = request_input_sources( + kwargs, + optional_params.keys().map(String::as_str).chain([ + "api_key", + "api_base", + "extra_headers", + ]), + )?; + let azure_ad_token_provider = kwargs + .get_item("azure_ad_token_provider")? + .and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER)); + let request = LiteLLMOcrRequest::try_from(BridgeOcrRequest { + model, + document: wire_document, + api_key: api_key.extract()?, + api_base: arguments.api_base()?, + custom_llm_provider, + extra_headers: arguments.extra_headers()?, + optional_params: optional_params.into(), + input_sources, + timeout_seconds: arguments.timeout_seconds()?, + }) .map_err(ocr_error_to_pyerr)?; - let optional_params = - project_optional_fields(kwargs, &specs, litellm_core::ocr::wire::BOUND_FIELDS)?; - let input_sources = request_input_sources( - kwargs, - optional_params - .keys() - .map(String::as_str) - .chain(["api_key", "api_base", "extra_headers"]), - )?; - let azure_ad_token_provider = kwargs - .get_item("azure_ad_token_provider")? - .and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER)); - let wire = OcrWireRequest { - model, - document: wire_document, - api_key: api_key.extract()?, - api_base: arguments.api_base()?, - custom_llm_provider, - extra_headers: arguments.extra_headers()?, - optional_params: optional_params.into(), - input_sources, - timeout_seconds: arguments.timeout_seconds()?, - }; - let request = decode_request(wire).map_err(ocr_error_to_pyerr)?; - let provider = request.provider_name(); - Ok(ProjectedOcrCall { - request: request.with_host_hooks(Arc::new(BridgeOcrHooks), None), - fields: ProjectedOcrFields { - boundary_request, - document: retained_document, - api_key: api_key.unbind(), - azure_ad_token_provider, - provider, - secret_fields: specs - .into_iter() - .filter(|spec| spec.secret) - .map(|spec| spec.name) - .collect(), - }, - }) + let provider = request.provider_name(); + Ok(Self { + request: request.with_host_hooks(Arc::new(BridgeOcrHooks), None), + fields: ProjectedOcrFields { + boundary_request, + document: retained_document, + api_key: api_key.unbind(), + azure_ad_token_provider, + provider, + secret_fields: specs + .into_iter() + .filter(|spec| spec.secret) + .map(|spec| spec.name) + .collect(), + }, + }) + } } pub(super) fn admitted_call(outcome: NativeOutcome) -> PyResult { @@ -183,6 +193,19 @@ mod tests { use super::*; + #[test] + fn projection_selects_consumed_values_without_serializing_host_objects() { + let fields = consumed_optional_params("mistral/model", None).unwrap(); + use litellm_core::call_arguments::should_project; + assert!(should_project("future_option", &fields, BOUND_FIELDS)); + assert!(should_project("extra_body", &fields, BOUND_FIELDS)); + assert!(should_project("id", &fields, BOUND_FIELDS)); + assert!(!should_project("metadata", &fields, BOUND_FIELDS)); + assert!(!should_project("callbacks", &fields, BOUND_FIELDS)); + assert!(!should_project("api_key", &fields, BOUND_FIELDS)); + assert!(!should_project("document", &fields, BOUND_FIELDS)); + } + fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> { let locals = PyDict::new(py); py.run(source, Some(&locals), Some(&locals)).unwrap(); @@ -493,7 +516,7 @@ kwargs = {'api_key': key} wire_document, serde_json::json!({"type": "mystery", "mystery": "x"}) ); - let error = match decode_request(OcrWireRequest { + let error = match LiteLLMOcrRequest::try_from(BridgeOcrRequest { model: "mistral/mistral-ocr-latest".into(), document: wire_document, api_key: None, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/request.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/request.rs new file mode 100644 index 00000000000..051540bdee3 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/request.rs @@ -0,0 +1,75 @@ +use std::collections::BTreeMap; +use std::time::Duration; + +use litellm_auth::InputSource; +use litellm_core::call_arguments::CallArguments; +use litellm_core::ocr::{LiteLLMOcrRequest, OcrCredentialInputs, OcrDocument}; +use serde_json::{Map, Value}; + +pub(super) struct BridgeOcrRequest { + pub model: String, + pub document: Value, + pub api_key: Option, + pub api_base: Option, + pub custom_llm_provider: Option, + pub extra_headers: Option>, + pub optional_params: CallArguments, + pub input_sources: BTreeMap, + pub timeout_seconds: Option, +} + +impl TryFrom for LiteLLMOcrRequest { + type Error = litellm_core::ocr::Error; + + fn try_from(request: BridgeOcrRequest) -> Result { + let api_key_source = source_for(&request.input_sources, "api_key"); + let api_base_source = source_for(&request.input_sources, "api_base"); + let extra_headers_source = source_for(&request.input_sources, "extra_headers"); + let timeout = request + .timeout_seconds + .map(|seconds| { + Duration::try_from_secs_f64(seconds).map_err(|_| Self::Error::RequestField { + path: "timeout_seconds".into(), + }) + }) + .transpose()?; + let headers = request + .extra_headers + .unwrap_or_default() + .into_iter() + .map(|(name, value)| { + value + .as_str() + .map(|value| (name.clone(), value.to_string())) + .ok_or_else(|| Self::Error::RequestField { + path: format!("extra_headers.{name}"), + }) + }) + .collect::, _>>()?; + let core_request = LiteLLMOcrRequest::new( + request.model, + OcrDocument::try_from(request.document)?, + request.custom_llm_provider.as_deref(), + request.optional_params, + )?; + let transport = + core_request + .transport + .clone() + .with_overrides(headers, extra_headers_source, timeout); + Ok(core_request.with_connection_inputs( + OcrCredentialInputs::new( + request.api_key, + api_key_source, + request.api_base, + api_base_source, + ), + transport, + request.input_sources, + )) + } +} + +fn source_for(sources: &BTreeMap, name: &str) -> InputSource { + sources.get(name).copied().unwrap_or_default() +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs index c7de98dee24..f2a068ead27 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs @@ -1,11 +1,12 @@ use litellm_core::ocr::Error; use std::future::Future; -use litellm_core::ocr::wire::{OcrWireRequest, decode_request}; +use litellm_core::ocr::LiteLLMOcrRequest; use pyo3::prelude::*; use serde_json::Value; use super::errors::to_pyerr as ocr_error_to_pyerr; +use super::request::BridgeOcrRequest; use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty}; fn prepare_ocr( @@ -37,7 +38,7 @@ fn prepare_ocr( extra_headers, timeout, } = options; - let request = decode_request(OcrWireRequest { + let request = LiteLLMOcrRequest::try_from(BridgeOcrRequest { model, document, api_key, diff --git a/tests/test_litellm/ocr/test_ocr_file_input.py b/tests/test_litellm/ocr/test_ocr_file_input.py index 3526d8c00d6..53b13fd36b3 100644 --- a/tests/test_litellm/ocr/test_ocr_file_input.py +++ b/tests/test_litellm/ocr/test_ocr_file_input.py @@ -12,32 +12,16 @@ Tests that: import base64 import os import tempfile -from collections.abc import Generator from io import BytesIO from pathlib import Path from typing import Final -from unittest.mock import AsyncMock, MagicMock, Mock +from unittest.mock import AsyncMock, MagicMock import orjson import pytest from starlette.datastructures import FormData -from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type - - -@pytest.fixture(autouse=True, params=["native", "disabled", "unavailable"]) -def document_runtime(request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch) -> Generator[None]: - from litellm.rust_bridge import bindings, configuration - - configuration.reset_rust_configuration() - monkeypatch.delenv("LITELLM_RUST", raising=False) - if request.param == "disabled": - monkeypatch.setenv("LITELLM_RUST", "0") - monkeypatch.setattr(bindings, "get_native_bridge", Mock(side_effect=AssertionError("Rust is disabled"))) - elif request.param == "unavailable": - monkeypatch.setattr(bindings, "get_native_bridge", lambda: None) - yield - configuration.reset_rust_configuration() +from litellm.ocr.main import convert_file_document_to_url_document, get_mime_type class TestGetMimeType: @@ -503,10 +487,9 @@ class TestProxySecurityGuard: async def test_proxy_upload_stops_reading_at_size_limit() -> None: from starlette.datastructures import UploadFile - from litellm.ocr.input import get_max_file_bytes from litellm.proxy.ocr_endpoints.endpoints import _parse_multipart_form - limit: Final = get_max_file_bytes() + limit: Final = 50 * 1024 * 1024 with tempfile.TemporaryFile() as stream: stream.truncate(limit * 2) upload: Final = UploadFile(file=stream, filename="large.pdf") diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index c600b80845a..5bcc65d75af 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -548,111 +548,3 @@ async def test_native_ocr_inherits_named_credentials_without_overwriting_argumen assert response.pages[0].markdown == "native OCR response" assert ocr_server.requests[0].headers["authorization"] == f"Bearer {expected_key}" assert ocr_server.requests[0].body["pages"] == [0, 2] - - - -@pytest.mark.parametrize("source", ["sdk", "proxy"]) -@pytest.mark.parametrize( - "filename,mime", [("scan.PNG", "image/png"), ("document.pdf", "application/pdf"), ("note.txt", "text/plain")] -) -def test_ocr_file_helpers_use_native_document_preparation(source: str, filename: str, mime: str) -> None: - from io import BytesIO - - from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type - from litellm.proxy.ocr_endpoints.endpoints import _build_document_from_upload - - file: Final = BytesIO(b"abc") - file.name = filename - document: Final = ( - convert_file_document_to_url_document({"type": "file", "file": file}) - if source == "sdk" - else _build_document_from_upload(b"abc", filename, "application/octet-stream; charset=utf-8") - ) - field: Final = "image_url" if mime.startswith("image/") else "document_url" - assert get_mime_type(filename) == mime - assert document == {"type": field, field: f"data:{mime};base64,YWJj"} - - -@pytest.mark.parametrize("attribute", ["read", "name"]) -def test_native_file_preparation_preserves_property_errors(attribute: str) -> None: - from litellm.ocr.input import convert_file_document_to_url_document - - failure: Final = LookupError("file property failed") - - class File: - def __getattribute__(self, name: str): - if name == attribute: - raise failure - return super().__getattribute__(name) - - def read(self): - return b"abc" - - with pytest.raises(LookupError) as caught: - convert_file_document_to_url_document({"type": "file", "file": File()}) - assert caught.value is failure - - -@pytest.mark.parametrize("kind", ["bytes", "path", "reader"]) -def test_native_file_preparation_rejects_oversized_input(kind: str, tmp_path: Path) -> None: - from litellm.ocr.input import FileDocument, convert_file_document_to_url_document, get_max_file_bytes - - limit: Final = get_max_file_bytes() - path: Final = tmp_path / "large.pdf" - with path.open("wb") as stream: - stream.truncate(limit + 1) - - class Reader: - def read(self) -> bytes: - return b"a" * (limit + 1) - - document: Final[FileDocument] = { - "type": "file", - "file": path if kind == "path" else Reader() if kind == "reader" else b"a" * (limit + 1), - } - with pytest.raises(ValueError, match="exceeds the size limit"): - convert_file_document_to_url_document(document) - - -@pytest.mark.parametrize("kind", ["str", "path", "reader"]) -def test_native_upload_binding_rejects_filesystem_inputs(kind: str, tmp_path: Path) -> None: - from io import BytesIO - from typing import cast # noqa: TID251 # deliberately invalid inputs exercise the native runtime boundary - - from litellm.ocr.input import convert_upload_to_url_document - - path: Final = tmp_path / "secret.pdf" - path.write_bytes(b"server secret") - source: Final = str(path) if kind == "str" else path if kind == "path" else BytesIO(b"abc") - with pytest.raises(TypeError): - convert_upload_to_url_document(cast(bytes, source), "document.pdf", None) - - -@pytest.mark.parametrize("extra_bytes", [0, 1]) -def test_native_upload_enforces_file_size_limit(extra_bytes: int) -> None: - import base64 - - from litellm.ocr.input import convert_upload_to_url_document, get_max_file_bytes - - content: Final = b"a" * (get_max_file_bytes() + extra_bytes) - if extra_bytes: - with pytest.raises(ValueError, match="exceeds the size limit"): - convert_upload_to_url_document(content, "scan.pdf", None) - return - document: Final = convert_upload_to_url_document(content, "scan.pdf", None) - assert document["type"] == "document_url" - assert base64.b64decode(document["document_url"].split(",", 1)[1]) == content - - -def test_native_file_preparation_preserves_reader_exception() -> None: - from litellm.ocr.input import convert_file_document_to_url_document - - failure: Final = RuntimeError("reader failed") - - class Reader: - def read(self) -> bytes: - raise failure - - with pytest.raises(RuntimeError) as caught: - convert_file_document_to_url_document({"type": "file", "file": Reader()}) - assert caught.value is failure