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