refactor(ocr): separate unresolved connection state

This commit is contained in:
Yujong Lee 2026-09-15 18:42:40 -07:00
parent 056487a283
commit 2f25875ce6
15 changed files with 327 additions and 214 deletions

View file

@ -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<Self::Environment, crate::ocr::Error> {
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<String, crate::ocr::Error> {
@ -103,7 +103,7 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
impl AzureAICohereParseConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;

View file

@ -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<litellm_auth::Sourced<String>>,
api_base: Option<litellm_auth::Sourced<String>>,
dynamic_api_key: Option<litellm_auth::Sourced<String>>,
dynamic_api_base: Option<litellm_auth::Sourced<String>>,
) -> (
Option<litellm_auth::Sourced<String>>,
Option<litellm_auth::Sourced<String>>,
) {
(
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<Self::Environment, crate::ocr::Error> {
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<String, crate::ocr::Error> {
@ -603,7 +596,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
impl AzureDocumentIntelligenceOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
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

View file

@ -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<Self::Environment, crate::ocr::Error> {
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<String, crate::ocr::Error> {
@ -94,7 +94,7 @@ impl BaseOcrConfig for AzureAIOCRConfig {
impl AzureAIOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
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();

View file

@ -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<Sourced<String>>,
api_base: Option<Sourced<String>>,
dynamic_api_key: Option<Sourced<String>>,
dynamic_api_base: Option<Sourced<String>>,
) -> (Option<Sourced<String>>, Option<Sourced<String>>) {
(
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<Output = Result<Self::Environment, crate::ocr::Error>> + Send;
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
optional_params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, crate::ocr::Error>;

View file

@ -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::Environment, crate::ocr::Error> {
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<String, crate::ocr::Error> {
@ -165,7 +165,7 @@ impl BaseOcrConfig for CohereParseConfig {
impl CohereParseConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
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

View file

@ -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::Environment, crate::ocr::Error> {
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<String, crate::ocr::Error> {
@ -125,7 +125,7 @@ impl BaseOcrConfig for MistralOCRConfig {
impl MistralOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;

View file

@ -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<Self::Environment, crate::ocr::Error> {
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<String, crate::ocr::Error> {
@ -156,7 +156,7 @@ impl BaseOcrConfig for ReductoParseV3Config {
impl ReductoParseV3Config {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
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<Self::Environment, crate::ocr::Error> {
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<String, crate::ocr::Error> {
@ -256,7 +256,7 @@ impl BaseOcrConfig for ReductoParseLegacyConfig {
impl ReductoParseLegacyConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
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();

View file

@ -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<Self::Environment, crate::ocr::Error> {
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<String, crate::ocr::Error> {
@ -188,7 +188,7 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
impl VertexAIDeepSeekOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
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!(

View file

@ -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<Self::Environment, crate::ocr::Error> {
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<String, crate::ocr::Error> {
@ -101,7 +101,7 @@ impl BaseOcrConfig for VertexAIOCRConfig {
impl VertexAIOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
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

View file

@ -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<Self, super::Error> {
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?,

View file

@ -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 {

View file

@ -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<B>(
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<B: Serialize>(
client: &OcrClient,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
url: &str,
headers: &[(String, String)],
body: &B,
@ -82,7 +82,7 @@ pub(crate) fn build_http_request<B: Serialize>(
}
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<String> {
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))
}

View file

@ -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")
);
}

View file

@ -61,14 +61,16 @@ pub enum OcrResponseFormat {
Native,
}
#[derive(Clone)]
pub struct OcrConnection {
pub api_key: Option<String>,
#[derive(Clone, Default)]
pub struct OcrCredentialInputs {
pub api_key: Option<Sourced<String>>,
pub dynamic_api_key: Option<Sourced<String>>,
pub api_key_source: InputSource,
pub api_base: Option<String>,
pub api_base: Option<Sourced<String>>,
pub dynamic_api_base: Option<Sourced<String>>,
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<String>,
pub api_key_source: InputSource,
pub api_base: Option<String>,
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<Sourced<String>>,
pub api_base: Option<Sourced<String>>,
}
pub struct LiteLLMOcrRequest {
pub model: String,
pub document: OcrDocument,
pub connection: OcrConnection,
pub credentials: OcrCredentialInputs,
pub transport: OcrTransportConfig,
pub hooks: Arc<dyn OcrHooks>,
pub litellm_call_id: Option<String>,
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<dyn OcrHooks>,
pub optional_params: CallArguments,
pub input_sources: BTreeMap<String, InputSource>,
pub azure_ad_token_provider: Option<TokenProviderHandle>,
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<OcrResponseFormat, super::Error> {
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 {

View file

@ -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<LiteLLMOcrRequest, crate::
})
})
.transpose()?;
let defaults = OcrConnection::default();
let defaults = OcrTransportConfig::default();
let max_response_bytes = wire
.optional_params
.get("max_response_bytes")
@ -155,13 +155,15 @@ pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, crate::
.filter(|(name, _)| name != "max_response_bytes")
.collect(),
)?;
let connection = OcrConnection {
api_key: nonblank(wire.api_key),
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_key_source,
api_base: nonblank(wire.api_base),
api_base: nonblank(wire.api_base)
.map(|value| litellm_auth::Sourced::new(value, api_base_source)),
dynamic_api_base: None,
api_base_source,
};
let transport = OcrTransportConfig {
extra_headers: headers,
extra_headers_source,
timeout: timeout.unwrap_or(defaults.timeout),
@ -170,7 +172,8 @@ pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, crate::
poll_timeout: defaults.poll_timeout,
};
Ok(LiteLLMOcrRequest {
connection,
credentials,
transport,
input_sources: wire.input_sources,
..request
})