mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
refactor(ocr): separate unresolved connection state
This commit is contained in:
parent
056487a283
commit
2f25875ce6
15 changed files with 327 additions and 214 deletions
|
|
@ -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)?;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)?;
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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?,
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
})
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue