From 2f25875ce6c18495a3494bfe0742e80bc9b7bfa8 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Tue, 15 Sep 2026 18:42:40 -0700 Subject: [PATCH] refactor(ocr): separate unresolved connection state --- .../ocr/cohere_parse_transformation.rs | 8 +- .../document_intelligence/transformation.rs | 43 +++--- .../src/llms/azure_ai/ocr/transformation.rs | 14 +- .../src/llms/base_llm/ocr/transformation.rs | 30 ++-- .../src/llms/cohere/ocr/transformation.rs | 14 +- .../src/llms/mistral/ocr/transformation.rs | 8 +- .../src/llms/reducto/ocr/transformation.rs | 20 +-- .../vertex_ai/ocr/deepseek_transformation.rs | 15 +- .../src/llms/vertex_ai/ocr/transformation.rs | 19 ++- litellm-rust/crates/core/src/ocr/handler.rs | 6 +- litellm-rust/crates/core/src/ocr/lifecycle.rs | 10 +- litellm-rust/crates/core/src/ocr/prepare.rs | 67 ++++----- .../crates/core/src/ocr/provider_config.rs | 131 +++++++++-------- litellm-rust/crates/core/src/ocr/types.rs | 137 ++++++++++++++++-- litellm-rust/crates/core/src/ocr/wire.rs | 19 ++- 15 files changed, 327 insertions(+), 214 deletions(-) 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 503c826da67..9f28dc8e0ff 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 @@ -4,7 +4,7 @@ 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::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument, PreparedOcrRequest}; use crate::url_utils::ApiUrl; const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE"; @@ -27,7 +27,7 @@ impl BaseOcrConfig for AzureAICohereParseConfig { async fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { BaseOcrConfig::validate_environment( @@ -40,7 +40,7 @@ impl BaseOcrConfig for AzureAICohereParseConfig { fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _params: &Self::OcrParams, _environment: &Self::Environment, ) -> Result { @@ -103,7 +103,7 @@ impl BaseOcrConfig for AzureAICohereParseConfig { impl AzureAICohereParseConfig { pub(crate) async fn prepare_request( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let params = self.map_ocr_params(&request.optional_params, &request.model)?; 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 5b2a8f7838d..8f12dbefa27 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 @@ -26,8 +26,8 @@ use crate::ocr::document::InlineDocument; use crate::ocr::hooks::OcrHooks; use crate::ocr::prepare::{credential_env, transform_request_body}; use crate::ocr::types::{ - LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageDimensions, - OcrResponseFormat, OcrUsageInfo, + LiteLLMOcrResponse, OcrConnection, OcrCredentialInputs, OcrDocument, OcrPage, + OcrPageDimensions, OcrResponseFormat, OcrUsageInfo, PreparedOcrRequest, ResolvedOcrCredentials, }; use crate::ocr::wire::DecodedOcrResponse; use crate::serde_compat::{FiniteF64, LaxI64}; @@ -476,33 +476,26 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { Some(AZURE_DI_API_KEY_ENV) } - fn resolve_connection_params( - &self, - api_key: Option>, - api_base: Option>, - dynamic_api_key: Option>, - dynamic_api_base: Option>, - ) -> ( - Option>, - Option>, - ) { - ( - api_key.and_then(|key| { - dynamic_api_key + fn resolve_connection_params(&self, inputs: OcrCredentialInputs) -> ResolvedOcrCredentials { + ResolvedOcrCredentials { + api_key: inputs.api_key.and_then(|key| { + inputs + .dynamic_api_key .filter(|value| !value.value().is_empty()) .or(Some(key)) }), - api_base.and_then(|base| { - dynamic_api_base + api_base: inputs.api_base.and_then(|base| { + inputs + .dynamic_api_base .filter(|value| !value.value().is_empty()) .or(Some(base)) }), - ) + } } async fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _client: &OcrClient, ) -> Result { let config = AzureAuthInputs { @@ -518,7 +511,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, params: &Self::OcrParams, _environment: &Self::Environment, ) -> Result { @@ -603,7 +596,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { impl AzureDocumentIntelligenceOCRConfig { pub(crate) async fn prepare_request( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let params = self.map_ocr_params(&request.optional_params, &request.model)?; @@ -1026,7 +1019,7 @@ mod tests { json!({"req_format":"native"}), ); request - .connection + .transport .extra_headers .push(("X-Trace".into(), "initial-only".into())); @@ -1110,8 +1103,8 @@ mod tests { ]) .await; let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})); - request.connection.api_key = None; - request.connection.extra_headers = vec![("Authorization".into(), "Bearer token".into())]; + request.credentials.api_key = None; + request.transport.extra_headers = vec![("Authorization".into(), "Bearer token".into())]; perform_ocr(request).await.unwrap(); server.await.unwrap(); @@ -1246,7 +1239,7 @@ mod tests { ]) .await; let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})); - request.connection.poll_timeout = std::time::Duration::from_millis(100); + request.transport.poll_timeout = std::time::Duration::from_millis(100); let error = tokio::time::timeout(std::time::Duration::from_secs(1), perform_ocr(request)) .await 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 ea1aca0b66b..e4e3a9fc570 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 @@ -4,7 +4,7 @@ use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequ 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::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, PreparedOcrRequest}; use crate::params::OpaqueParams; use crate::url_utils::ApiUrl; use litellm_auth::{InputSource, Sourced}; @@ -27,7 +27,7 @@ impl BaseOcrConfig for AzureAIOCRConfig { async fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _client: &OcrClient, ) -> Result { let config = AzureAuthInputs { @@ -43,7 +43,7 @@ impl BaseOcrConfig for AzureAIOCRConfig { fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _params: &Self::OcrParams, _environment: &Self::Environment, ) -> Result { @@ -94,7 +94,7 @@ impl BaseOcrConfig for AzureAIOCRConfig { impl AzureAIOCRConfig { pub(crate) async fn prepare_request( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let params = self.map_ocr_params(&request.optional_params, &request.model)?; @@ -308,8 +308,8 @@ mod tests { &base, json!({"include_image_base64":true}), ); - request.connection.api_key = None; - request.connection.extra_headers = vec![( + request.credentials.api_key = None; + request.transport.extra_headers = vec![( "Authorization".into(), "Bearer python-prepared-token".into(), )]; @@ -345,7 +345,7 @@ mod tests { &base, json!({"azure_ad_token":"rust-owned-token"}), ); - request.connection.api_key = None; + request.credentials.api_key = None; perform_ocr(request).await.unwrap(); server.await.unwrap(); 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 e8ccdda75ab..03ec373afea 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 @@ -1,7 +1,6 @@ use std::future::Future; use std::sync::Arc; -use litellm_auth::Sourced; use serde::Serialize; use serde::de::DeserializeOwned; @@ -9,7 +8,8 @@ use crate::call_arguments::{CallArguments, parse_options}; use crate::ocr::OcrClient; use crate::ocr::hooks::OcrHooks; use crate::ocr::types::{ - LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrResponseFormat, + LiteLLMOcrResponse, OcrConnection, OcrCredentialInputs, OcrDocument, OcrResponseFormat, + PreparedOcrRequest, ResolvedOcrCredentials, }; const HEALTH_CHECK_PDF_DATA_URI: &str = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y="; @@ -23,21 +23,17 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { None } - fn resolve_connection_params( - &self, - api_key: Option>, - api_base: Option>, - dynamic_api_key: Option>, - dynamic_api_base: Option>, - ) -> (Option>, Option>) { - ( - dynamic_api_key + fn resolve_connection_params(&self, inputs: OcrCredentialInputs) -> ResolvedOcrCredentials { + ResolvedOcrCredentials { + api_key: inputs + .dynamic_api_key .filter(|value| !value.value().is_empty()) - .or(api_key), - dynamic_api_base + .or(inputs.api_key), + api_base: inputs + .dynamic_api_base .filter(|value| !value.value().is_empty()) - .or(api_base), - ) + .or(inputs.api_base), + } } fn get_health_check_document(&self) -> OcrDocument { @@ -49,13 +45,13 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> impl Future> + Send; fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, optional_params: &Self::OcrParams, environment: &Self::Environment, ) -> Result; 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 d0b92127f81..62ac1788d97 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -2,16 +2,16 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; use serde_with::serde_as; -use crate::serde_compat::LaxI64; 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::types::{ - LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageImage, - OcrUsageInfo, + LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageImage, OcrUsageInfo, + PreparedOcrRequest, }; +use crate::serde_compat::LaxI64; use crate::url_utils::ApiUrl; const COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI: &str = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"; @@ -100,7 +100,7 @@ impl BaseOcrConfig for CohereParseConfig { async fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _client: &OcrClient, ) -> Result { self.validate_environment(&request.connection, &credential_env) @@ -108,7 +108,7 @@ impl BaseOcrConfig for CohereParseConfig { fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _params: &Self::OcrParams, _environment: &Self::Environment, ) -> Result { @@ -165,7 +165,7 @@ impl BaseOcrConfig for CohereParseConfig { impl CohereParseConfig { pub(crate) async fn prepare_request( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let params = self.map_ocr_params(&request.optional_params, &request.model)?; @@ -380,6 +380,7 @@ mod tests { "type":"image_url","image_url":"https://example.com/original.png" })) .unwrap(); + let request = crate::ocr::prepare::prepare_request(request); let http = CohereParseConfig .prepare_request(&request, &crate::ocr::test_support::ocr_client()) .await @@ -525,6 +526,7 @@ mod tests { request.response_format().unwrap(), crate::ocr::types::OcrResponseFormat::Litellm ); + let request = crate::ocr::prepare::prepare_request(request); let http = CohereParseConfig .prepare_request(&request, &crate::ocr::test_support::ocr_client()) .await 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 5401c819bc0..350359bfa7c 100644 --- a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs @@ -6,7 +6,7 @@ use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContex use crate::ocr::OcrClient; use crate::ocr::prepare::{credential_env, transform_request_body}; use crate::ocr::types::{ - LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrUsageInfo, + LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrUsageInfo, PreparedOcrRequest, }; use crate::params::OpaqueParams; use crate::url_utils::ApiUrl; @@ -48,7 +48,7 @@ impl BaseOcrConfig for MistralOCRConfig { async fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _client: &OcrClient, ) -> Result { self.validate_environment(&request.connection, &credential_env) @@ -56,7 +56,7 @@ impl BaseOcrConfig for MistralOCRConfig { fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _params: &Self::OcrParams, _environment: &Self::Environment, ) -> Result { @@ -125,7 +125,7 @@ impl BaseOcrConfig for MistralOCRConfig { impl MistralOCRConfig { pub(crate) async fn prepare_request( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let params = self.map_ocr_params(&request.optional_params, &request.model)?; 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 aa6ce333757..0863e202575 100644 --- a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -10,7 +10,7 @@ use crate::ocr::OcrClient; use crate::ocr::document::InlineDocument; use crate::ocr::prepare::{build_http_request, credential_env, guardrail_document}; use crate::ocr::types::{ - LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrUsageInfo, + LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrUsageInfo, PreparedOcrRequest, }; use crate::params::OpaqueParams; use crate::url_utils::ApiUrl; @@ -85,7 +85,7 @@ impl BaseOcrConfig for ReductoParseV3Config { async fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _client: &OcrClient, ) -> Result { validate_environment(&request.connection, &credential_env) @@ -93,7 +93,7 @@ impl BaseOcrConfig for ReductoParseV3Config { fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _params: &Self::OcrParams, _environment: &Self::Environment, ) -> Result { @@ -156,7 +156,7 @@ impl BaseOcrConfig for ReductoParseV3Config { impl ReductoParseV3Config { pub(crate) async fn prepare_request( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let params = self.map_ocr_params(&request.optional_params, &request.model)?; @@ -194,7 +194,7 @@ impl BaseOcrConfig for ReductoParseLegacyConfig { async fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { ReductoParseV3Config @@ -204,7 +204,7 @@ impl BaseOcrConfig for ReductoParseLegacyConfig { fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, params: &Self::OcrParams, environment: &Self::Environment, ) -> Result { @@ -256,7 +256,7 @@ impl BaseOcrConfig for ReductoParseLegacyConfig { impl ReductoParseLegacyConfig { pub(crate) async fn prepare_request( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let params = self.map_ocr_params(&request.optional_params, &request.model)?; @@ -766,7 +766,7 @@ mod tests { ]) .await; let mut request = wire_request(&format!("reducto/{model}"), &base, json!({})); - request.connection.extra_headers = vec![ + request.transport.extra_headers = vec![ ("Content-Type".into(), "application/json".into()), ("X-Trace".into(), "upload-test".into()), ]; @@ -921,7 +921,7 @@ mod tests { let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await; let mut request = wire_request("reducto/parse-v3", &base, json!({})); request.document = request.document.with_source("reducto://ready.pdf".into()); - request.connection.extra_headers = vec![("authorization".into(), "Bearer existing".into())]; + request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())]; let response = perform_ocr(request).await.unwrap(); server.await.unwrap(); @@ -966,7 +966,7 @@ mod tests { ]) .await; let mut request = wire_request(model, &base, json!({})); - request.connection.extra_headers = vec![("authorization".into(), "Bearer original".into())]; + request.transport.extra_headers = vec![("authorization".into(), "Bearer original".into())]; request.hooks = Arc::new(RewriteHeaders); perform_ocr(request).await.unwrap(); 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 4c12c02cb02..1e8f607449d 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 @@ -8,8 +8,8 @@ use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContex use crate::ocr::OcrClient; use crate::ocr::prepare::{credential_env, transform_request_body}; use crate::ocr::types::{ - LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument, OcrPage, OcrPageDimensions, OcrPageImage, - OcrUsageInfo, + LiteLLMOcrResponse, OcrDocument, OcrPage, OcrPageDimensions, OcrPageImage, OcrUsageInfo, + PreparedOcrRequest, }; use crate::params::OpaqueParams; use crate::routing_utils::model::{ModelNamespace, ProviderModel, RoutedModel}; @@ -108,7 +108,7 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { async fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { BaseOcrConfig::validate_environment(&VertexAIOCRConfig, request, client).await @@ -116,7 +116,7 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _params: &Self::OcrParams, environment: &Self::Environment, ) -> Result { @@ -188,7 +188,7 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { impl VertexAIDeepSeekOCRConfig { pub(crate) async fn prepare_request( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let params = self.map_ocr_params(&request.optional_params, &request.model)?; @@ -700,7 +700,10 @@ mod tests { "https://caller.example", json!({"vertex_project":"project-1"}), ); - request.connection.api_base_source = InputSource::Request; + request.credentials.api_base = Some(litellm_auth::Sourced::new( + "https://caller.example".into(), + InputSource::Request, + )); let error = perform_ocr(request).await.unwrap_err(); assert!( 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 d81ea31bc2b..03a80a384c3 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 @@ -6,7 +6,7 @@ use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequ 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::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, PreparedOcrRequest}; use crate::params::OpaqueParams; use crate::url_utils::ApiUrl; @@ -26,7 +26,7 @@ impl BaseOcrConfig for VertexAIOCRConfig { async fn validate_environment( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let config = VertexConfig::from_sourced_optional_params( @@ -39,7 +39,7 @@ impl BaseOcrConfig for VertexAIOCRConfig { fn get_complete_url( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, _params: &Self::OcrParams, environment: &Self::Environment, ) -> Result { @@ -101,7 +101,7 @@ impl BaseOcrConfig for VertexAIOCRConfig { impl VertexAIOCRConfig { pub(crate) async fn prepare_request( &self, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, client: &OcrClient, ) -> Result { let params = self.map_ocr_params(&request.optional_params, &request.model)?; @@ -286,8 +286,8 @@ mod tests { &base, json!({"vertex_project":"project-1"}), ); - request.connection.api_key = None; - request.connection.extra_headers = vec![("authorization".into(), "Bearer supplied".into())]; + request.credentials.api_key = None; + request.transport.extra_headers = vec![("authorization".into(), "Bearer supplied".into())]; perform_ocr(request).await.unwrap(); server.await.unwrap(); @@ -316,7 +316,10 @@ mod tests { "https://caller.example", json!({"vertex_project":"project-1"}), ); - request.connection.api_base_source = InputSource::Request; + request.credentials.api_base = Some(litellm_auth::Sourced::new( + "https://caller.example".into(), + InputSource::Request, + )); let error = perform_ocr(request).await.unwrap_err(); assert!( @@ -349,6 +352,8 @@ mod tests { options.clone(), ); let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options); + let direct = crate::ocr::prepare::prepare_request(direct); + let vertex = crate::ocr::prepare::prepare_request(vertex); let direct_http = MistralOCRConfig .prepare_request(&direct, &client) .await diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index 73ebd38bf2e..4d2f6dcc803 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -3,7 +3,7 @@ use std::sync::Arc; use super::OcrClient; use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest}; use super::provider_config::OcrConfigKind; -use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; +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; @@ -45,7 +45,7 @@ pub(crate) async fn perform_ocr_request( pub(crate) struct PreparedOcrCall { client: OcrClient, - request: LiteLLMOcrRequest, + request: PreparedOcrRequest, http: reqwest::Request, } @@ -54,7 +54,7 @@ impl PreparedOcrCall { client: OcrClient, request: LiteLLMOcrRequest, ) -> Result { - let request = super::prepare::resolve_connection_params(request); + 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?, diff --git a/litellm-rust/crates/core/src/ocr/lifecycle.rs b/litellm-rust/crates/core/src/ocr/lifecycle.rs index f72094e2291..2dcec88eae0 100644 --- a/litellm-rust/crates/core/src/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/src/ocr/lifecycle.rs @@ -752,10 +752,10 @@ mod tests { mock_server(vec![MockResponse::json(json!({"pages": []}))]).await; let request = wire_request("mistral/model", "https://unused.invalid", json!({})); let request = crate::ocr::LiteLLMOcrRequest { - connection: crate::ocr::OcrConnection { + credentials: crate::ocr::types::OcrCredentialInputs { dynamic_api_key: Some(Sourced::new("dynamic-key".into(), InputSource::Deployment)), dynamic_api_base: Some(Sourced::new(base, InputSource::Deployment)), - ..request.connection + ..request.credentials }, ..request }; @@ -1496,7 +1496,7 @@ mod tests { "http://localhost", json!({"max_response_bytes": 123}), ); - assert_eq!(request.connection.max_response_bytes, 123); + assert_eq!(request.transport.max_response_bytes, 123); assert!(!request.optional_params.contains_key("max_response_bytes")); for value in [ json!(0), @@ -1555,9 +1555,9 @@ mod tests { let request = wire_request("azure_ai/mistral-ocr", "https://example.invalid", json!({})); let request = crate::ocr::LiteLLMOcrRequest { - connection: crate::ocr::OcrConnection { + transport: crate::ocr::types::OcrTransportConfig { extra_headers: vec![("authorization".into(), "Bearer test-key".into())], - ..request.connection + ..request.transport }, azure_ad_token_provider: Some(litellm_auth::TokenProviderHandle::new(Arc::new( PendingToken { diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index b071a8fb4b8..089293c7303 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -3,11 +3,11 @@ use serde_json::Value; use super::OcrClient; use super::hooks::OcrDuringCallRequest; -use super::types::{LiteLLMOcrRequest, OcrDocument}; +use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest}; pub(crate) async fn transform_request_body( client: &OcrClient, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, url: &str, headers: &[(String, String)], retains_document: bool, @@ -65,7 +65,7 @@ where pub(crate) fn build_http_request( client: &OcrClient, - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, url: &str, headers: &[(String, String)], body: &B, @@ -82,7 +82,7 @@ pub(crate) fn build_http_request( } pub(crate) async fn guardrail_document( - request: &LiteLLMOcrRequest, + request: &PreparedOcrRequest, url: &str, headers: &[(String, String)], ) -> Result<(OcrDocument, Vec<(String, String)>), super::Error> { @@ -127,10 +127,10 @@ pub(crate) fn credential_env(name: &str) -> Option { std::env::var(name).ok() } -pub(crate) fn resolve_connection_params(request: LiteLLMOcrRequest) -> LiteLLMOcrRequest { +pub(crate) fn prepare_request(request: LiteLLMOcrRequest) -> PreparedOcrRequest { use litellm_auth::{InputSource, Sourced}; - let connection = request.connection; + let credentials = request.credentials.clone(); let api_base_env = match request.config.provider() { super::provider_config::OcrProvider::Mistral => Some("MISTRAL_API_BASE"), super::provider_config::OcrProvider::AzureAi => Some("AZURE_AI_API_BASE"), @@ -138,38 +138,29 @@ pub(crate) fn resolve_connection_params(request: LiteLLMOcrRequest) -> LiteLLMOc | super::provider_config::OcrProvider::Reducto | super::provider_config::OcrProvider::VertexAi => None, }; - let dynamic_api_key = connection.dynamic_api_key.or_else(|| { - connection - .api_key - .clone() - .map(|value| Sourced::new(value, connection.api_key_source)) - .or_else(|| { - request - .config - .get_api_key_env_var() - .and_then(credential_env) - .map(|value| Sourced::new(value, InputSource::Environment)) - }) + let dynamic_api_key = credentials.dynamic_api_key.or_else(|| { + credentials.api_key.clone().or_else(|| { + request + .config + .get_api_key_env_var() + .and_then(credential_env) + .map(|value| Sourced::new(value, InputSource::Environment)) + }) }); - let dynamic_api_base = connection.dynamic_api_base.or_else(|| { - connection - .api_base - .clone() - .map(|value| Sourced::new(value, connection.api_base_source)) - .or_else(|| { - api_base_env - .and_then(credential_env) - .map(|value| Sourced::new(value, InputSource::Environment)) - }) + let dynamic_api_base = credentials.dynamic_api_base.or_else(|| { + credentials.api_base.clone().or_else(|| { + api_base_env + .and_then(credential_env) + .map(|value| Sourced::new(value, InputSource::Environment)) + }) }); - LiteLLMOcrRequest { - connection: request - .config - .resolve_connection_params(super::OcrConnection { - dynamic_api_key, - dynamic_api_base, - ..connection - }), - ..request - } + let resolved = request + .config + .resolve_connection_params(super::types::OcrCredentialInputs { + dynamic_api_key, + dynamic_api_base, + ..credentials + }); + let transport = request.transport.clone(); + PreparedOcrRequest::new(request, OcrConnection::new(resolved, transport)) } diff --git a/litellm-rust/crates/core/src/ocr/provider_config.rs b/litellm-rust/crates/core/src/ocr/provider_config.rs index 08d3b810083..57fefd1b051 100644 --- a/litellm-rust/crates/core/src/ocr/provider_config.rs +++ b/litellm-rust/crates/core/src/ocr/provider_config.rs @@ -1,4 +1,4 @@ -use super::types::{OcrConnection, OcrDocument}; +use super::types::{OcrCredentialInputs, OcrDocument, 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; @@ -9,7 +9,6 @@ use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, Reduct use crate::llms::vertex_ai::ocr::deepseek_transformation::VertexAIDeepSeekOCRConfig; use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig; use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; -use litellm_auth::Sourced; use strum::{EnumString, IntoStaticStr}; macro_rules! dispatch_config { @@ -66,37 +65,11 @@ impl OcrConfigKind { dispatch_config!(self, get_health_check_document()) } - pub(crate) fn resolve_connection_params(self, connection: OcrConnection) -> OcrConnection { - let api_key = connection - .api_key - .map(|value| Sourced::new(value, connection.api_key_source)); - let api_base = connection - .api_base - .map(|value| Sourced::new(value, connection.api_base_source)); - let (api_key, api_base) = dispatch_config!( - self, - resolve_connection_params( - api_key, - api_base, - connection.dynamic_api_key, - connection.dynamic_api_base, - ) - ); - OcrConnection { - api_key_source: api_key - .as_ref() - .map(Sourced::source) - .unwrap_or(connection.api_key_source), - api_base_source: api_base - .as_ref() - .map(Sourced::source) - .unwrap_or(connection.api_base_source), - api_key: api_key.map(Sourced::into_value), - api_base: api_base.map(Sourced::into_value), - dynamic_api_key: None, - dynamic_api_base: None, - ..connection - } + pub(crate) fn resolve_connection_params( + self, + inputs: OcrCredentialInputs, + ) -> ResolvedOcrCredentials { + dispatch_config!(self, resolve_connection_params(inputs)) } pub(crate) fn get_error_class( @@ -183,7 +156,7 @@ fn is_document_intelligence_model(model: &str) -> bool { #[cfg(test)] mod tests { use super::*; - use litellm_auth::InputSource; + use litellm_auth::{InputSource, Sourced}; #[test] fn provider_names_round_trip_exactly() { @@ -261,34 +234,66 @@ mod tests { #[test] fn connection_resolution_preserves_dynamic_precedence_and_input_sources() { - let connection = OcrConfigKind::Mistral.resolve_connection_params(OcrConnection { - api_key: Some("explicit-key".into()), - api_base: Some("https://explicit.test".into()), + 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: Some(Sourced::new("dynamic-key".into(), InputSource::Environment)), dynamic_api_base: Some(Sourced::new( "https://dynamic.test".into(), InputSource::Request, )), - ..Default::default() }); - assert_eq!(connection.api_key.as_deref(), Some("dynamic-key")); - assert_eq!(connection.api_base.as_deref(), Some("https://dynamic.test")); - assert_eq!(connection.api_key_source, InputSource::Environment); - assert_eq!(connection.api_base_source, InputSource::Request); + assert_eq!( + connection + .api_key + .as_ref() + .map(|value| value.value().as_str()), + Some("dynamic-key") + ); + assert_eq!( + connection + .api_base + .as_ref() + .map(|value| value.value().as_str()), + Some("https://dynamic.test") + ); + assert_eq!( + connection.api_key.as_ref().map(Sourced::source), + Some(InputSource::Environment) + ); + assert_eq!( + 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(OcrConnection { - api_key: Some("explicit-key".into()), - api_base: Some("https://explicit.test".into()), - dynamic_api_key: dynamic.clone(), - dynamic_api_base: dynamic, - ..Default::default() - }); - assert_eq!(connection.api_key.as_deref(), Some("explicit-key")); + 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_base.as_deref(), + 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") ); } @@ -302,10 +307,12 @@ mod tests { (None, Some("base")), (Some("key"), Some("base")), ] { - let connection = - OcrConfigKind::AzureDocumentIntelligence.resolve_connection_params(OcrConnection { - api_key: explicit_key.map(str::to_string), - api_base: explicit_base.map(str::to_string), + 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, @@ -314,14 +321,20 @@ mod tests { "https://dynamic.test".into(), InputSource::Deployment, )), - ..Default::default() - }); + }, + ); assert_eq!( - connection.api_key.as_deref(), + connection + .api_key + .as_ref() + .map(|value| value.value().as_str()), explicit_key.map(|_| "dynamic-key") ); assert_eq!( - connection.api_base.as_deref(), + connection + .api_base + .as_ref() + .map(|value| value.value().as_str()), explicit_base.map(|_| "https://dynamic.test") ); } diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 9e2f99a0cf9..dc4de6fc5d0 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -61,14 +61,16 @@ pub enum OcrResponseFormat { Native, } -#[derive(Clone)] -pub struct OcrConnection { - pub api_key: Option, +#[derive(Clone, Default)] +pub struct OcrCredentialInputs { + pub api_key: Option>, pub dynamic_api_key: Option>, - pub api_key_source: InputSource, - pub api_base: Option, + pub api_base: Option>, pub dynamic_api_base: Option>, - pub api_base_source: InputSource, +} + +#[derive(Clone)] +pub struct OcrTransportConfig { pub extra_headers: Vec<(String, String)>, pub extra_headers_source: InputSource, pub timeout: Duration, @@ -77,15 +79,9 @@ pub struct OcrConnection { pub poll_timeout: Duration, } -impl Default for OcrConnection { +impl Default for OcrTransportConfig { fn default() -> Self { Self { - api_key: None, - dynamic_api_key: None, - api_key_source: InputSource::Deployment, - api_base: None, - dynamic_api_base: None, - api_base_source: InputSource::Deployment, extra_headers: Vec::new(), extra_headers_source: InputSource::Deployment, timeout: Duration::from_secs(OCR_HTTP_TIMEOUT_SECS), @@ -96,10 +92,67 @@ impl Default for OcrConnection { } } +#[derive(Clone)] +pub struct OcrConnection { + pub api_key: Option, + pub api_key_source: InputSource, + pub api_base: Option, + pub api_base_source: InputSource, + pub extra_headers: Vec<(String, String)>, + pub extra_headers_source: InputSource, + pub timeout: Duration, + pub max_download_bytes: u64, + pub max_response_bytes: usize, + pub poll_timeout: Duration, +} + +impl OcrConnection { + pub(crate) fn new(credentials: ResolvedOcrCredentials, transport: OcrTransportConfig) -> Self { + let api_key_source = credentials + .api_key + .as_ref() + .map(Sourced::source) + .unwrap_or(InputSource::Deployment); + let api_base_source = credentials + .api_base + .as_ref() + .map(Sourced::source) + .unwrap_or(InputSource::Deployment); + Self { + api_key: credentials.api_key.map(Sourced::into_value), + api_key_source, + api_base: credentials.api_base.map(Sourced::into_value), + api_base_source, + extra_headers: transport.extra_headers, + extra_headers_source: transport.extra_headers_source, + timeout: transport.timeout, + max_download_bytes: transport.max_download_bytes, + max_response_bytes: transport.max_response_bytes, + poll_timeout: transport.poll_timeout, + } + } +} + +impl Default for OcrConnection { + fn default() -> Self { + Self::new( + ResolvedOcrCredentials::default(), + OcrTransportConfig::default(), + ) + } +} + +#[derive(Clone, Default)] +pub(crate) struct ResolvedOcrCredentials { + pub api_key: Option>, + pub api_base: Option>, +} + pub struct LiteLLMOcrRequest { pub model: String, pub document: OcrDocument, - pub connection: OcrConnection, + pub credentials: OcrCredentialInputs, + pub transport: OcrTransportConfig, pub hooks: Arc, pub litellm_call_id: Option, pub optional_params: CallArguments, @@ -120,7 +173,8 @@ impl LiteLLMOcrRequest { Ok(Self { model, document, - connection: OcrConnection::default(), + credentials: OcrCredentialInputs::default(), + transport: OcrTransportConfig::default(), hooks: Arc::new(NoopOcrHooks), litellm_call_id: None, optional_params, @@ -158,6 +212,59 @@ impl LiteLLMOcrRequest { } } +pub(crate) struct PreparedOcrRequest { + pub model: String, + pub document: OcrDocument, + pub connection: OcrConnection, + pub hooks: Arc, + pub optional_params: CallArguments, + pub input_sources: BTreeMap, + pub azure_ad_token_provider: Option, + pub(crate) config: OcrConfigKind, +} + +impl PreparedOcrRequest { + pub(crate) fn new(request: LiteLLMOcrRequest, connection: OcrConnection) -> Self { + let LiteLLMOcrRequest { + model, + document, + credentials: _, + transport: _, + hooks, + litellm_call_id: _, + optional_params, + input_sources, + azure_ad_token_provider, + config, + } = request; + Self { + model, + document, + connection, + hooks, + optional_params, + input_sources, + azure_ad_token_provider, + config, + } + } + + pub(crate) fn response_format(&self) -> Result { + self.optional_params + .get("req_format") + .filter(|value| !value.is_null()) + .map(|value| { + serde_json::from_value(value.clone()).map_err(|_| super::Error::RequestFormat) + }) + .transpose() + .map(|format| format.unwrap_or_default()) + } + + pub(crate) fn provider_name(&self) -> &'static str { + self.config.provider().into() + } +} + #[serde_as] #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] pub struct OcrPageDimensions { diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index ddf0c9e7d67..73e1babdb49 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -1,7 +1,7 @@ use std::collections::BTreeMap; use std::time::Duration; -use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument}; +use super::types::{LiteLLMOcrRequest, OcrDocument, OcrTransportConfig}; use crate::call_arguments::{ArgumentSpec, CallArguments}; use litellm_auth::InputSource; use serde::{ @@ -131,7 +131,7 @@ pub fn decode_request(wire: OcrWireRequest) -> Result Result Result