mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(ocr): mirror Python provider layout
This commit is contained in:
parent
4e996400e2
commit
4d903da85c
59 changed files with 2817 additions and 2611 deletions
|
|
@ -6,6 +6,7 @@ pub mod chat_completions;
|
|||
pub mod constants;
|
||||
pub mod error;
|
||||
pub mod http_utils;
|
||||
pub(crate) mod llms;
|
||||
mod media;
|
||||
pub mod messages;
|
||||
pub mod ocr;
|
||||
|
|
|
|||
1
litellm-rust/crates/core/src/llms/azure_ai/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/azure_ai/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
|
|
@ -1,40 +1,42 @@
|
|||
use super::super::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::llms::cohere::ocr::transformation::CohereParseConfig;
|
||||
use crate::llms::cohere::ocr::{CohereParams, CohereResponse, validate_document};
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::cohere::{
|
||||
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
|
||||
};
|
||||
use crate::ocr::document::{inline_remote_document, validate_inline_document};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
use crate::providers::azure_ai::auth::AzureAuthInputs;
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
|
||||
|
||||
pub(crate) struct AzureCohereAdapter;
|
||||
#[derive(Default)]
|
||||
pub(crate) struct AzureAICohereParseConfig {
|
||||
cohere: CohereParseConfig,
|
||||
}
|
||||
|
||||
impl OcrAdapter for AzureCohereAdapter {
|
||||
impl BaseOcrConfig for AzureAICohereParseConfig {
|
||||
type ProviderResponse = CohereResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::AzureAi;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let params = super::super::super::wire::decode_request_value::<CohereParams>(
|
||||
let params = crate::ocr::wire::decode_request_value::<CohereParams>(
|
||||
serde_json::Value::Object(request.optional_params.clone()),
|
||||
"optional_params",
|
||||
)?;
|
||||
let mut config = AzureAuthInputs::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
)
|
||||
.map_err(Error::from)?;
|
||||
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
|
||||
let config = AzureAuthInputs {
|
||||
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
|
||||
..AzureAuthInputs::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
)
|
||||
.map_err(Error::from)?
|
||||
};
|
||||
let base = request
|
||||
.connection
|
||||
.api_base
|
||||
|
|
@ -46,8 +48,12 @@ impl OcrAdapter for AzureCohereAdapter {
|
|||
"Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(),
|
||||
)
|
||||
})?;
|
||||
let headers =
|
||||
super::validate_ai_environment(&request.connection, &config, &credential_env).await?;
|
||||
let headers = super::transformation::validate_environment(
|
||||
&request.connection,
|
||||
&config,
|
||||
&credential_env,
|
||||
)
|
||||
.await?;
|
||||
validate_document(&request.document)?;
|
||||
let remote = request.document.source().starts_with("http://")
|
||||
|| request.document.source().starts_with("https://");
|
||||
|
|
@ -57,7 +63,9 @@ impl OcrAdapter for AzureCohereAdapter {
|
|||
&request.connection,
|
||||
)
|
||||
.await?;
|
||||
let body = transform_request(&request.model, document, params)?;
|
||||
let body = self
|
||||
.cohere
|
||||
.transform_ocr_request(&request.model, document, params)?;
|
||||
transform_request_body(
|
||||
client,
|
||||
request,
|
||||
|
|
@ -76,9 +84,9 @@ impl OcrAdapter for AzureCohereAdapter {
|
|||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
response: CohereResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
transform_response(&request.model, response)
|
||||
self.cohere.transform_ocr_response(request, response)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -1,7 +1,3 @@
|
|||
mod cohere;
|
||||
mod document_intelligence;
|
||||
mod mistral;
|
||||
|
||||
use std::sync::OnceLock;
|
||||
|
||||
use crate::Error;
|
||||
|
|
@ -11,12 +7,7 @@ use crate::ocr::error::OcrError;
|
|||
use crate::ocr::types::OcrConnection;
|
||||
use crate::providers::azure_ai::auth::{AzureAuthInputs, AzureAuthService};
|
||||
|
||||
pub(crate) use cohere::AzureCohereAdapter;
|
||||
pub(crate) use document_intelligence::AzureDocumentIntelligenceAdapter;
|
||||
pub(crate) use mistral::AzureMistralAdapter;
|
||||
pub(super) use mistral::validate_environment as validate_ai_environment;
|
||||
|
||||
async fn resolve_entra(
|
||||
pub(super) async fn resolve_entra(
|
||||
config: &AzureAuthInputs,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Option<Sourced<String>>, Error> {
|
||||
|
|
@ -39,7 +30,7 @@ async fn resolve_entra(
|
|||
.map_err(Error::from)
|
||||
}
|
||||
|
||||
fn validate_destination(
|
||||
pub(super) fn validate_destination(
|
||||
connection: &OcrConnection,
|
||||
credential_source: InputSource,
|
||||
) -> Result<(), OcrError> {
|
||||
|
|
@ -0,0 +1,3 @@
|
|||
pub(crate) mod transformation;
|
||||
|
||||
pub(crate) use transformation::AzureDocumentIntelligenceOperation;
|
||||
|
|
@ -0,0 +1,822 @@
|
|||
mod types {
|
||||
use serde::{Deserialize, Deserializer, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum PagesInput {
|
||||
ZeroBasedIndices(Vec<i64>),
|
||||
NativeTokens(Vec<String>),
|
||||
NativeRange(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum FeaturesInput {
|
||||
Names(Vec<String>),
|
||||
CommaSeparated(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub(crate) struct DocumentIntelligenceInputParams {
|
||||
pub pages: Option<PagesInput>,
|
||||
pub features: Option<FeaturesInput>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize)]
|
||||
pub(crate) struct DocumentIntelligenceParams {
|
||||
pub pages: Option<String>,
|
||||
pub features: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) enum DocumentIntelligenceRequest {
|
||||
#[serde(rename = "urlSource")]
|
||||
UrlSource(String),
|
||||
#[serde(rename = "base64Source")]
|
||||
Base64Source(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub(crate) enum OperationStatus {
|
||||
Succeeded,
|
||||
Running,
|
||||
NotStarted,
|
||||
Failed,
|
||||
Unknown(String),
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for OperationStatus {
|
||||
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
|
||||
Ok(match String::deserialize(deserializer)?.as_str() {
|
||||
"succeeded" => Self::Succeeded,
|
||||
"running" => Self::Running,
|
||||
"notStarted" => Self::NotStarted,
|
||||
"failed" => Self::Failed,
|
||||
value => Self::Unknown(value.to_string()),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for OperationStatus {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str(match self {
|
||||
Self::Succeeded => "succeeded",
|
||||
Self::Running => "running",
|
||||
Self::NotStarted => "notStarted",
|
||||
Self::Failed => "failed",
|
||||
Self::Unknown(value) => value,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct AzureDocumentIntelligenceOperation {
|
||||
pub status: Option<OperationStatus>,
|
||||
#[serde(rename = "analyzeResult")]
|
||||
pub analyze_result: Option<AzureDocumentIntelligenceAnalyzeResult>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct AzureDocumentIntelligenceAnalyzeResult {
|
||||
pub content: Option<String>,
|
||||
#[serde(default)]
|
||||
pub pages: Vec<AzureDocumentIntelligencePage>,
|
||||
pub tables: Option<Vec<Map<String, Value>>>,
|
||||
#[serde(rename = "keyValuePairs")]
|
||||
pub key_value_pairs: Option<Vec<Map<String, Value>>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct AzureDocumentIntelligencePage {
|
||||
#[serde(rename = "pageNumber", default, deserialize_with = "optional_i64")]
|
||||
pub page_number: Option<i64>,
|
||||
#[serde(default, deserialize_with = "optional_f64")]
|
||||
pub width: Option<f64>,
|
||||
#[serde(default, deserialize_with = "optional_f64")]
|
||||
pub height: Option<f64>,
|
||||
pub unit: Option<String>,
|
||||
#[serde(default)]
|
||||
pub lines: Vec<AzureDocumentIntelligenceLine>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct AzureDocumentIntelligenceLine {
|
||||
pub content: Option<String>,
|
||||
}
|
||||
|
||||
fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<i64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_i64()
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected an integer")),
|
||||
Some(Value::String(value)) => value
|
||||
.parse::<i64>()
|
||||
.map(Some)
|
||||
.map_err(|_| serde::de::Error::custom("expected an integer")),
|
||||
Some(_) => Err(serde::de::Error::custom("expected an integer")),
|
||||
}
|
||||
}
|
||||
|
||||
fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<f64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_f64()
|
||||
.filter(|value| value.is_finite())
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a finite number")),
|
||||
Some(Value::String(value)) => value
|
||||
.parse::<f64>()
|
||||
.ok()
|
||||
.filter(|value| value.is_finite())
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a finite number")),
|
||||
Some(_) => Err(serde::de::Error::custom("expected a number")),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) use types::{
|
||||
AzureDocumentIntelligenceOperation, DocumentIntelligenceParams, OperationStatus,
|
||||
};
|
||||
|
||||
mod params {
|
||||
use std::collections::BTreeSet;
|
||||
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::types::{
|
||||
DocumentIntelligenceInputParams, DocumentIntelligenceParams, FeaturesInput, PagesInput,
|
||||
};
|
||||
use crate::ocr::error::OcrRequestError;
|
||||
use crate::ocr::prepare::ParsedProviderParams;
|
||||
|
||||
pub(crate) fn decode_input_params(
|
||||
params: Map<String, Value>,
|
||||
prefix: &str,
|
||||
) -> Result<ParsedProviderParams<DocumentIntelligenceInputParams>, OcrRequestError> {
|
||||
if let Some(Value::Array(pages)) = params.get("pages") {
|
||||
if pages.iter().any(Value::is_boolean) {
|
||||
return Err(OcrRequestError::Pages("boolean page index".into()));
|
||||
}
|
||||
if pages
|
||||
.iter()
|
||||
.any(|page| page.is_number() && page.as_i64().is_none())
|
||||
{
|
||||
return Err(OcrRequestError::Pages("page index is out of range".into()));
|
||||
}
|
||||
if !pages.iter().all(Value::is_i64) && !pages.iter().all(Value::is_string) {
|
||||
return Err(OcrRequestError::Pages("mixed page element types".into()));
|
||||
}
|
||||
}
|
||||
crate::ocr::wire::decode_request_value(Value::Object(params), prefix)
|
||||
}
|
||||
|
||||
pub(crate) fn map_ocr_params(
|
||||
params: DocumentIntelligenceInputParams,
|
||||
) -> Result<DocumentIntelligenceParams, OcrRequestError> {
|
||||
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>, OcrRequestError> {
|
||||
let normalized = match pages {
|
||||
PagesInput::ZeroBasedIndices(indices) => {
|
||||
if indices.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
indices
|
||||
.into_iter()
|
||||
.map(|page| {
|
||||
if page < 0 {
|
||||
return Err(OcrRequestError::Pages("negative page index".into()));
|
||||
}
|
||||
page.checked_add(1).ok_or_else(|| {
|
||||
OcrRequestError::Pages("page index is out of range".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
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.collect::<Vec<_>>()
|
||||
.join(","),
|
||||
};
|
||||
if !normalized.split(',').all(valid_page_token) {
|
||||
return Err(OcrRequestError::Pages("invalid native page range".into()));
|
||||
}
|
||||
Ok(Some(normalized))
|
||||
}
|
||||
|
||||
fn valid_page_token(token: &str) -> bool {
|
||||
let mut parts = token.split('-');
|
||||
let start = parts.next().unwrap_or_default();
|
||||
if start.is_empty() || !start.chars().all(|character| character.is_ascii_digit()) {
|
||||
return false;
|
||||
}
|
||||
match parts.next() {
|
||||
None => true,
|
||||
Some(end) => {
|
||||
!end.is_empty()
|
||||
&& end.chars().all(|character| character.is_ascii_digit())
|
||||
&& parts.next().is_none()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_features(features: FeaturesInput) -> Result<Option<String>, OcrRequestError> {
|
||||
let tokens = match features {
|
||||
FeaturesInput::Names(names) => names,
|
||||
FeaturesInput::CommaSeparated(names) => names.split(',').map(str::to_string).collect(),
|
||||
};
|
||||
if tokens.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
let normalized = tokens.iter().map(|token| token.trim()).collect::<Vec<_>>();
|
||||
if !normalized.iter().all(|token| {
|
||||
let Some((first, rest)) = token.as_bytes().split_first() else {
|
||||
return false;
|
||||
};
|
||||
first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric)
|
||||
}) {
|
||||
return Err(OcrRequestError::Features);
|
||||
}
|
||||
Ok(Some(normalized.join(",")))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
fn map(value: Value) -> Result<DocumentIntelligenceParams, OcrRequestError> {
|
||||
let fields = value.as_object().unwrap().clone();
|
||||
map_ocr_params(decode_input_params(fields, "optional_params")?.known)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn input_params_retain_unknown_fields() {
|
||||
let parsed = decode_input_params(
|
||||
json!({
|
||||
"pages": [0],
|
||||
"future_ocr_option": true,
|
||||
"extra_body": {"provider_option": "value"}
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
.clone(),
|
||||
"optional_params",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
parsed.known.pages,
|
||||
Some(PagesInput::ZeroBasedIndices(vec![0]))
|
||||
);
|
||||
assert_eq!(parsed.extra_params["future_ocr_option"], true);
|
||||
assert_eq!(
|
||||
parsed.extra_params["extra_body"],
|
||||
json!({"provider_option": "value"})
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::to_value(map_ocr_params(parsed.known).unwrap()).unwrap(),
|
||||
json!({"pages": "1", "features": null})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(json!([0, 1, 2]), Some("1,2,3"))]
|
||||
#[case(json!([2, 0, 0, 1]), Some("1,2,3"))]
|
||||
#[case(json!([]), None)]
|
||||
#[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"))]
|
||||
fn page_mapping_matches_python(#[case] input: Value, #[case] expected: Option<&str>) {
|
||||
assert_eq!(
|
||||
map(json!({"pages": input})).unwrap().pages.as_deref(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(json!("a,b"))]
|
||||
#[case(json!([-1]))]
|
||||
#[case(json!([true, false]))]
|
||||
#[case(json!([1, "2"]))]
|
||||
#[case(json!(5))]
|
||||
fn invalid_page_mapping_matches_python(#[case] input: Value) {
|
||||
assert!(map(json!({"pages": input})).is_err());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(json!(["keyValuePairs"]), "keyValuePairs")]
|
||||
#[case(json!(["keyValuePairs", "languages"]), "keyValuePairs,languages")]
|
||||
#[case(json!("keyValuePairs"), "keyValuePairs")]
|
||||
#[case(json!("keyValuePairs,languages"), "keyValuePairs,languages")]
|
||||
#[case(json!("keyValuePairs, languages"), "keyValuePairs,languages")]
|
||||
fn feature_mapping_matches_python(#[case] input: Value, #[case] expected: &str) {
|
||||
assert_eq!(
|
||||
map(json!({"features": input})).unwrap().features.as_deref(),
|
||||
Some(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(json!("keyValuePairs&pages=9"))]
|
||||
#[case(json!("key value pairs"))]
|
||||
#[case(json!(""))]
|
||||
#[case(json!([1, 2]))]
|
||||
#[case(json!([["keyValuePairs"]]))]
|
||||
#[case(json!({"feature":"keyValuePairs"}))]
|
||||
#[case(json!(5))]
|
||||
fn invalid_feature_mapping_matches_python(#[case] input: Value) {
|
||||
assert!(map(json!({"features": input})).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_feature_list_is_omitted() {
|
||||
assert_eq!(map(json!({"features": []})).unwrap().features, None);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mod mapping {
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::types::*;
|
||||
use crate::constants::{AZURE_DI_DEFAULT_DPI, AZURE_DI_DEFAULT_HEIGHT, AZURE_DI_DEFAULT_WIDTH};
|
||||
use crate::ocr::document::InlineDocument;
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(crate) fn transform_ocr_request(
|
||||
document: OcrDocument,
|
||||
) -> Result<DocumentIntelligenceRequest, OcrRequestError> {
|
||||
let source = document.source();
|
||||
if source.is_empty() {
|
||||
return Err(OcrRequestError::MissingDocumentUrl);
|
||||
}
|
||||
Ok(if let Some(document) = InlineDocument::parse(source)? {
|
||||
DocumentIntelligenceRequest::Base64Source(
|
||||
STANDARD.encode(document.decode(crate::constants::OCR_INLINE_MAX_BYTES)?),
|
||||
)
|
||||
} else {
|
||||
DocumentIntelligenceRequest::UrlSource(source.to_string())
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_ocr_response(
|
||||
model: &str,
|
||||
response: AzureDocumentIntelligenceOperation,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
if response.status != Some(OperationStatus::Succeeded) {
|
||||
return Err(OcrResponseError::OperationStatus(
|
||||
response
|
||||
.status
|
||||
.map(|status| status.to_string())
|
||||
.unwrap_or_else(|| "None".into()),
|
||||
));
|
||||
}
|
||||
let result = response.analyze_result.unwrap_or_default();
|
||||
let pages = result
|
||||
.pages
|
||||
.into_iter()
|
||||
.map(normalize_page)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
let pages_processed = pages.len();
|
||||
let mut extra_fields = Map::new();
|
||||
extra_fields.insert("content".into(), option_value(result.content));
|
||||
extra_fields.insert("tables".into(), option_value(result.tables));
|
||||
extra_fields.insert("keyValuePairs".into(), option_value(result.key_value_pairs));
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages,
|
||||
model: model.into(),
|
||||
document_annotation: None,
|
||||
usage_info: Some(json!({"pages_processed":pages_processed})),
|
||||
object: "ocr".into(),
|
||||
extra_fields,
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn normalize_page(page: AzureDocumentIntelligencePage) -> Result<Value, OcrResponseError> {
|
||||
let index = page
|
||||
.page_number
|
||||
.unwrap_or(1)
|
||||
.checked_sub(1)
|
||||
.ok_or(OcrResponseError::NumericRange("page.pageNumber"))?;
|
||||
let scale = if page.unit.as_deref().unwrap_or("inch") == "inch" {
|
||||
AZURE_DI_DEFAULT_DPI as f64
|
||||
} else {
|
||||
1.0
|
||||
};
|
||||
let width = pixel_dimension(
|
||||
page.width.unwrap_or(AZURE_DI_DEFAULT_WIDTH),
|
||||
scale,
|
||||
"page.width",
|
||||
)?;
|
||||
let height = pixel_dimension(
|
||||
page.height.unwrap_or(AZURE_DI_DEFAULT_HEIGHT),
|
||||
scale,
|
||||
"page.height",
|
||||
)?;
|
||||
let markdown = page
|
||||
.lines
|
||||
.iter()
|
||||
.map(|line| line.content.as_deref().unwrap_or_default())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n");
|
||||
Ok(json!({
|
||||
"index":index,
|
||||
"markdown":markdown,
|
||||
"images":null,
|
||||
"dimensions":{"width":width,"height":height,"dpi":AZURE_DI_DEFAULT_DPI}
|
||||
}))
|
||||
}
|
||||
|
||||
fn pixel_dimension(
|
||||
value: f64,
|
||||
scale: f64,
|
||||
field: &'static str,
|
||||
) -> Result<i64, OcrResponseError> {
|
||||
let value = value * scale;
|
||||
if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 {
|
||||
return Err(OcrResponseError::NumericRange(field));
|
||||
}
|
||||
Ok(value.trunc() as i64)
|
||||
}
|
||||
|
||||
fn option_value<T: serde::Serialize>(value: Option<T>) -> Value {
|
||||
value
|
||||
.and_then(|value| serde_json::to_value(value).ok())
|
||||
.unwrap_or(Value::Null)
|
||||
}
|
||||
}
|
||||
|
||||
mod polling {
|
||||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use reqwest::Url;
|
||||
use tokio::time::Instant;
|
||||
|
||||
use super::{AzureDocumentIntelligenceOperation, OperationStatus};
|
||||
use crate::constants::{AZURE_DI_SUBSCRIPTION_HEADER, OCR_POLL_RETRY_SECS};
|
||||
use crate::ocr::client::read_json_response;
|
||||
use crate::ocr::error::{OcrError, OcrPollingError, OcrResponseError};
|
||||
use crate::ocr::hooks::OcrHooks;
|
||||
use crate::ocr::types::OcrConnection;
|
||||
use crate::ocr::wire::DecodedOcrResponse;
|
||||
|
||||
pub(super) async fn read_operation_response(
|
||||
http_client: &reqwest::Client,
|
||||
response: reqwest::Response,
|
||||
original_url: &str,
|
||||
headers: &[(String, String)],
|
||||
connection: &OcrConnection,
|
||||
native: bool,
|
||||
hooks: &Arc<dyn OcrHooks>,
|
||||
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
|
||||
if response.status() != reqwest::StatusCode::ACCEPTED {
|
||||
let bytes =
|
||||
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes)
|
||||
.await?;
|
||||
crate::ocr::handler::post_call(hooks, &bytes).await?;
|
||||
return Ok(crate::ocr::wire::decode_response(&bytes, native)?);
|
||||
}
|
||||
let location = response
|
||||
.headers()
|
||||
.get("operation-location")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.ok_or(OcrPollingError::PollLocation)?
|
||||
.to_string();
|
||||
let original = Url::parse(original_url).map_err(|_| OcrPollingError::PollOrigin)?;
|
||||
let operation = Url::parse(&location).map_err(|_| OcrPollingError::PollOrigin)?;
|
||||
if original.origin() != operation.origin()
|
||||
|| !operation.username().is_empty()
|
||||
|| operation.password().is_some()
|
||||
{
|
||||
return Err(OcrPollingError::PollOrigin.into());
|
||||
}
|
||||
let bytes =
|
||||
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes)
|
||||
.await?;
|
||||
crate::ocr::handler::post_call(hooks, &bytes).await?;
|
||||
poll_operation(http_client, operation, headers, connection, native, hooks).await
|
||||
}
|
||||
|
||||
async fn poll_operation(
|
||||
http_client: &reqwest::Client,
|
||||
url: Url,
|
||||
headers: &[(String, String)],
|
||||
connection: &OcrConnection,
|
||||
native: bool,
|
||||
hooks: &Arc<dyn OcrHooks>,
|
||||
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
|
||||
let deadline = Instant::now()
|
||||
.checked_add(connection.poll_timeout)
|
||||
.ok_or(OcrPollingError::PollTimeout)?;
|
||||
loop {
|
||||
let remaining = deadline
|
||||
.checked_duration_since(Instant::now())
|
||||
.filter(|remaining| !remaining.is_zero())
|
||||
.ok_or(OcrPollingError::PollTimeout)?;
|
||||
let builder = http_client
|
||||
.get(url.clone())
|
||||
.timeout(remaining.min(connection.timeout));
|
||||
let builder = crate::http_utils::with_headers(
|
||||
builder,
|
||||
headers,
|
||||
crate::http_utils::HeaderPolicy::Only(&[
|
||||
AZURE_DI_SUBSCRIPTION_HEADER,
|
||||
"authorization",
|
||||
]),
|
||||
);
|
||||
let response =
|
||||
tokio::time::timeout_at(deadline, crate::http_utils::http_request(builder))
|
||||
.await
|
||||
.map_err(|_| OcrPollingError::PollTimeout)?
|
||||
.map_err(crate::error::TransportError::from)?;
|
||||
let retry = response
|
||||
.headers()
|
||||
.get(reqwest::header::RETRY_AFTER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| value.parse::<u64>().ok())
|
||||
.unwrap_or(OCR_POLL_RETRY_SECS)
|
||||
.max(1);
|
||||
let decoded = tokio::time::timeout_at(
|
||||
deadline,
|
||||
read_json_response::<AzureDocumentIntelligenceOperation>(
|
||||
response,
|
||||
native,
|
||||
connection.max_response_bytes,
|
||||
),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| OcrPollingError::PollTimeout)??;
|
||||
match &decoded.data.status {
|
||||
Some(OperationStatus::Succeeded) => {
|
||||
crate::ocr::handler::post_call(hooks, decoded.text.as_bytes()).await?;
|
||||
return Ok(decoded);
|
||||
}
|
||||
Some(OperationStatus::Running | OperationStatus::NotStarted) => {
|
||||
tokio::time::timeout_at(
|
||||
deadline,
|
||||
tokio::time::sleep(Duration::from_secs(retry)),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| OcrPollingError::PollTimeout)?;
|
||||
}
|
||||
status => {
|
||||
return Err(OcrResponseError::OperationStatus(
|
||||
status
|
||||
.as_ref()
|
||||
.map(ToString::to_string)
|
||||
.unwrap_or_else(|| "None".into()),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
use crate::Error;
|
||||
use crate::auth::{InputSource, Sourced};
|
||||
use crate::constants::{AZURE_DI_API_VERSION, AZURE_DI_SUBSCRIPTION_HEADER};
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrResponseFormat};
|
||||
use crate::providers::azure_ai::auth::AzureAuthInputs;
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
const AZURE_DI_API_KEY_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY";
|
||||
const AZURE_DI_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT";
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct AzureDocumentIntelligenceOCRConfig;
|
||||
|
||||
impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
|
||||
type ProviderResponse = AzureDocumentIntelligenceOperation;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let params = map_ocr_params(request)?;
|
||||
let config = AzureAuthInputs {
|
||||
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
|
||||
..AzureAuthInputs::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
)
|
||||
.map_err(Error::from)?
|
||||
};
|
||||
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
|
||||
let endpoint = nonblank(request.connection.api_base.clone())
|
||||
.or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV)))
|
||||
.ok_or_else(|| Error::Auth("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into()))?;
|
||||
let url = get_complete_url(&endpoint, &request.model, ¶ms)?;
|
||||
let body = mapping::transform_ocr_request(request.document.clone())?;
|
||||
transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: AzureDocumentIntelligenceOperation,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
mapping::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
|
||||
async fn read_response(
|
||||
&self,
|
||||
client: &OcrClient,
|
||||
response: reqwest::Response,
|
||||
url: &str,
|
||||
headers: &[(String, String)],
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> Result<crate::ocr::wire::DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError>
|
||||
{
|
||||
polling::read_operation_response(
|
||||
client.polling_http(),
|
||||
response,
|
||||
url,
|
||||
headers,
|
||||
&request.connection,
|
||||
request.response_format()? == OcrResponseFormat::Native,
|
||||
&request.hooks,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
fn map_ocr_params(
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> Result<DocumentIntelligenceParams, OcrRequestError> {
|
||||
let params = params::decode_input_params(request.optional_params.clone(), "optional_params")?;
|
||||
let crate::ocr::prepare::ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params: _extra_params,
|
||||
} = params;
|
||||
params::map_ocr_params(params)
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
endpoint: &str,
|
||||
model: &str,
|
||||
params: &DocumentIntelligenceParams,
|
||||
) -> Result<String, OcrError> {
|
||||
let model = format!("{}:analyze", model_id(model)?);
|
||||
ApiUrl::parse(endpoint)
|
||||
.and_then(|url| url.complete_path(&["documentintelligence", "documentModels", &model]))
|
||||
.map(|url| {
|
||||
url.append_query_pairs(
|
||||
[("api-version", AZURE_DI_API_VERSION)]
|
||||
.into_iter()
|
||||
.chain(params.pages.iter().map(|pages| ("pages", pages.as_str())))
|
||||
.chain(
|
||||
params
|
||||
.features
|
||||
.iter()
|
||||
.map(|features| ("features", features.as_str())),
|
||||
),
|
||||
)
|
||||
.into_string()
|
||||
})
|
||||
.map_err(|_| OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
})
|
||||
.map_err(OcrError::from)
|
||||
}
|
||||
|
||||
async fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
config: &AzureAuthInputs,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization")
|
||||
|| crate::http_utils::has_header(&connection.extra_headers, AZURE_DI_SUBSCRIPTION_HEADER)
|
||||
{
|
||||
super::super::common_utils::validate_destination(
|
||||
connection,
|
||||
connection.extra_headers_source,
|
||||
)?;
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let key = nonblank(connection.api_key.clone())
|
||||
.map(|value| Sourced::new(value, connection.api_key_source))
|
||||
.or_else(|| {
|
||||
nonblank(env_lookup(AZURE_DI_API_KEY_ENV))
|
||||
.map(|value| Sourced::new(value, InputSource::Environment))
|
||||
});
|
||||
if let Some(key) = key {
|
||||
super::super::common_utils::validate_destination(connection, key.source())?;
|
||||
return Ok(
|
||||
std::iter::once((AZURE_DI_SUBSCRIPTION_HEADER.into(), key.into_value()))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
);
|
||||
}
|
||||
let token = super::super::common_utils::resolve_entra(config, env_lookup)
|
||||
.await?
|
||||
.ok_or(Error::MissingAzureDocumentIntelligenceCredentials)?;
|
||||
super::super::common_utils::validate_destination(connection, token.source())?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {}", token.value())))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
fn model_id(model: &str) -> Result<&str, OcrRequestError> {
|
||||
let model = model.rsplit('/').next().unwrap_or(model);
|
||||
if matches!(model, "." | "..") {
|
||||
return Err(OcrRequestError::DotModel);
|
||||
}
|
||||
Ok(model)
|
||||
}
|
||||
|
||||
fn nonblank(value: Option<String>) -> Option<String> {
|
||||
value
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_endpoint_cannot_receive_environment_key() {
|
||||
let connection = OcrConnection {
|
||||
api_base: Some("https://request.example".into()),
|
||||
api_base_source: InputSource::Request,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let error = validate_environment(&connection, &Default::default(), &|name| {
|
||||
(name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into())
|
||||
})
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("request-controlled Azure endpoint")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_endpoint_accepts_request_owned_key() {
|
||||
let connection = OcrConnection {
|
||||
api_key: Some("request-key".into()),
|
||||
api_key_source: InputSource::Request,
|
||||
api_base: Some("https://request.example".into()),
|
||||
api_base_source: InputSource::Request,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let headers = validate_environment(&connection, &Default::default(), &|_| None)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
headers[0],
|
||||
(AZURE_DI_SUBSCRIPTION_HEADER.into(), "request-key".into())
|
||||
);
|
||||
}
|
||||
}
|
||||
4
litellm-rust/crates/core/src/llms/azure_ai/ocr/mod.rs
Normal file
4
litellm-rust/crates/core/src/llms/azure_ai/ocr/mod.rs
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
pub(crate) mod cohere_parse_transformation;
|
||||
pub(crate) mod common_utils;
|
||||
pub(crate) mod document_intelligence;
|
||||
pub(crate) mod transformation;
|
||||
|
|
@ -1,15 +1,15 @@
|
|||
use super::super::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::auth::{InputSource, Sourced};
|
||||
use crate::constants::AZURE_AI_OCR_PATH;
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
|
||||
use crate::llms::mistral::ocr::{MistralOcrParams, MistralOcrResponse};
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse};
|
||||
use crate::ocr::document::{inline_remote_document, validate_inline_document};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
|
||||
};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
|
||||
use crate::providers::azure_ai::auth::AzureAuthInputs;
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
|
@ -17,12 +17,13 @@ use crate::url_utils::ApiUrl;
|
|||
const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY";
|
||||
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct AzureMistralAdapter;
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub(crate) struct AzureAIOCRConfig {
|
||||
mistral: MistralOCRConfig,
|
||||
}
|
||||
|
||||
impl OcrAdapter for AzureMistralAdapter {
|
||||
impl BaseOcrConfig for AzureAIOCRConfig {
|
||||
type ProviderResponse = MistralOcrResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::AzureAi;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
|
|
@ -33,12 +34,14 @@ impl OcrAdapter for AzureMistralAdapter {
|
|||
known: params,
|
||||
extra_params: _extra_params,
|
||||
} = _prepare_ocr_request::<MistralOcrParams>(request)?;
|
||||
let mut config = AzureAuthInputs::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
)
|
||||
.map_err(Error::from)?;
|
||||
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
|
||||
let config = AzureAuthInputs {
|
||||
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
|
||||
..AzureAuthInputs::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
)
|
||||
.map_err(Error::from)?
|
||||
};
|
||||
let url = get_complete_url(request.connection.api_base.as_deref(), &credential_env)?;
|
||||
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
|
||||
let retains_document = !request.document.source().starts_with("http://")
|
||||
|
|
@ -49,7 +52,9 @@ impl OcrAdapter for AzureMistralAdapter {
|
|||
&request.connection,
|
||||
)
|
||||
.await?;
|
||||
let body = mistral::transform_ocr_request(&request.model, document, ¶ms)?;
|
||||
let body = self
|
||||
.mistral
|
||||
.transform_ocr_request(&request.model, document, ¶ms)?;
|
||||
transform_request_body(
|
||||
client,
|
||||
request,
|
||||
|
|
@ -65,9 +70,9 @@ impl OcrAdapter for AzureMistralAdapter {
|
|||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
response: MistralOcrResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
mistral::transform_ocr_response(&request.model, response)
|
||||
self.mistral.transform_ocr_response(request, response)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -92,16 +97,16 @@ fn get_complete_url(
|
|||
})
|
||||
}
|
||||
|
||||
pub(in crate::ocr::adapters) async fn validate_environment(
|
||||
pub(super) async fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
config: &AzureAuthInputs,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
|
||||
if config.azure_ad_token_provider.is_some() {
|
||||
super::resolve_entra(config, env_lookup).await?;
|
||||
super::common_utils::resolve_entra(config, env_lookup).await?;
|
||||
}
|
||||
super::validate_destination(connection, connection.extra_headers_source)?;
|
||||
super::common_utils::validate_destination(connection, connection.extra_headers_source)?;
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let key = nonblank(connection.api_key.clone())
|
||||
|
|
@ -111,13 +116,13 @@ pub(in crate::ocr::adapters) async fn validate_environment(
|
|||
.map(|value| Sourced::new(value, InputSource::Environment))
|
||||
});
|
||||
if let Some(key) = key {
|
||||
super::validate_destination(connection, key.source())?;
|
||||
super::common_utils::validate_destination(connection, key.source())?;
|
||||
return Ok(bearer_headers(connection, key.value()));
|
||||
}
|
||||
let key = super::resolve_entra(config, env_lookup)
|
||||
let key = super::common_utils::resolve_entra(config, env_lookup)
|
||||
.await?
|
||||
.ok_or(Error::MissingAzureAiCredentials)?;
|
||||
super::validate_destination(connection, key.source())?;
|
||||
super::common_utils::validate_destination(connection, key.source())?;
|
||||
Ok(bearer_headers(connection, key.value()))
|
||||
}
|
||||
|
||||
1
litellm-rust/crates/core/src/llms/base_llm/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/base_llm/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
1
litellm-rust/crates/core/src/llms/base_llm/ocr/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/base_llm/ocr/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod transformation;
|
||||
|
|
@ -0,0 +1,47 @@
|
|||
use std::future::Future;
|
||||
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::error::{OcrError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrResponseFormat};
|
||||
use crate::ocr::wire::DecodedOcrResponse;
|
||||
|
||||
pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
|
||||
type ProviderResponse: DeserializeOwned + Send;
|
||||
|
||||
fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> impl Future<Output = Result<reqwest::Request, OcrError>> + Send;
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError>;
|
||||
|
||||
fn read_response(
|
||||
&self,
|
||||
_client: &OcrClient,
|
||||
response: reqwest::Response,
|
||||
_url: &str,
|
||||
_headers: &[(String, String)],
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> impl Future<Output = Result<DecodedOcrResponse<Self::ProviderResponse>, OcrError>> + Send
|
||||
{
|
||||
async move {
|
||||
let bytes = crate::ocr::client::read_response_bytes(
|
||||
response,
|
||||
request.connection.max_response_bytes,
|
||||
)
|
||||
.await?;
|
||||
crate::ocr::handler::post_call(&request.hooks, &bytes).await?;
|
||||
Ok(crate::ocr::wire::decode_response(
|
||||
&bytes,
|
||||
request.response_format()? == OcrResponseFormat::Native,
|
||||
)?)
|
||||
}
|
||||
}
|
||||
}
|
||||
1
litellm-rust/crates/core/src/llms/cohere/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/cohere/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
3
litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs
Normal file
3
litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
pub(crate) mod transformation;
|
||||
|
||||
pub(crate) use transformation::{CohereParams, CohereResponse, validate_document};
|
||||
396
litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs
Normal file
396
litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs
Normal file
|
|
@ -0,0 +1,396 @@
|
|||
mod provider {
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::ocr::document::InlineDocument;
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub(crate) enum OutputFormat {
|
||||
#[default]
|
||||
Markdown,
|
||||
Blocks,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct CohereParams {
|
||||
#[serde(default)]
|
||||
pub output_format: OutputFormat,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, Serialize)]
|
||||
pub(crate) struct CohereRequest {
|
||||
pub model: String,
|
||||
pub document: OcrDocument,
|
||||
pub output_format: OutputFormat,
|
||||
}
|
||||
|
||||
pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), OcrRequestError> {
|
||||
let OcrDocument::ImageUrl { image_url, .. } = document else {
|
||||
return Err(OcrRequestError::CohereImageOnly);
|
||||
};
|
||||
if image_url.is_empty() {
|
||||
return Err(OcrRequestError::CohereImageOnly);
|
||||
}
|
||||
if let Some(inline) = InlineDocument::parse(image_url)? {
|
||||
if !inline.mime_type().type_.eq_ignore_ascii_case("image") {
|
||||
return Err(OcrRequestError::CohereImageOnly);
|
||||
}
|
||||
inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct CohereResponse {
|
||||
#[serde(default)]
|
||||
pages: Vec<CoherePage>,
|
||||
meta: Option<CohereMeta>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CoherePage {
|
||||
index: Option<i64>,
|
||||
markdown: Option<CohereMarkdown>,
|
||||
blocks: Option<Vec<Map<String, Value>>>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CohereMarkdown {
|
||||
#[serde(default)]
|
||||
content: String,
|
||||
images: Option<Vec<Map<String, Value>>>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CohereMeta {
|
||||
billed_units: Option<CohereBilledUnits>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CohereBilledUnits {
|
||||
pages: Option<i64>,
|
||||
}
|
||||
|
||||
pub(crate) fn transform_response(
|
||||
model: &str,
|
||||
response: CohereResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
let pages_processed = response
|
||||
.meta
|
||||
.and_then(|meta| meta.billed_units)
|
||||
.and_then(|units| units.pages)
|
||||
.map(Ok)
|
||||
.unwrap_or_else(|| {
|
||||
i64::try_from(response.pages.len())
|
||||
.map_err(|_| OcrResponseError::NumericRange("pages"))
|
||||
})?;
|
||||
let pages = response
|
||||
.pages
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(position, page)| {
|
||||
let index = page.index.map(Ok).unwrap_or_else(|| {
|
||||
i64::try_from(position)
|
||||
.map_err(|_| OcrResponseError::NumericRange("page index"))
|
||||
})?;
|
||||
let (content, images) = page
|
||||
.markdown
|
||||
.map(|markdown| {
|
||||
let images =
|
||||
markdown
|
||||
.images
|
||||
.filter(|images| !images.is_empty())
|
||||
.map(|images| {
|
||||
images
|
||||
.into_iter()
|
||||
.map(|mut image| {
|
||||
if let Some(Value::Object(bbox)) =
|
||||
image.get("bounding_box").cloned()
|
||||
{
|
||||
image.insert("bbox".into(), Value::Object(bbox));
|
||||
}
|
||||
Value::Object(image)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
});
|
||||
(markdown.content, images)
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let mut normalized = json!({"index": index, "markdown": content, "images": images});
|
||||
if let Some(blocks) = page.blocks {
|
||||
normalized["blocks"] = json!(blocks);
|
||||
}
|
||||
Ok(normalized)
|
||||
})
|
||||
.collect::<Result<Vec<_>, OcrResponseError>>()?;
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages,
|
||||
model: model.into(),
|
||||
document_annotation: None,
|
||||
usage_info: Some(json!({"pages_processed": pages_processed})),
|
||||
object: "ocr".into(),
|
||||
extra_fields: Map::new(),
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_request(
|
||||
model: &str,
|
||||
document: OcrDocument,
|
||||
params: CohereParams,
|
||||
) -> Result<CohereRequest, OcrRequestError> {
|
||||
validate_document(&document)?;
|
||||
Ok(CohereRequest {
|
||||
model: model.into(),
|
||||
document,
|
||||
output_format: params.output_format,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn response_normalizes_markdown_images_blocks_and_billed_pages() {
|
||||
let response = serde_json::from_value(json!({
|
||||
"pages": [
|
||||
{
|
||||
"type":"markdown",
|
||||
"index":4,
|
||||
"markdown":{
|
||||
"content":"receipt",
|
||||
"images":[{
|
||||
"id":"image",
|
||||
"bounding_box":{"top_left_x":1,"bottom_right_x":48},
|
||||
"bounding_box_normalized":{"top_left_x":0.04,"bottom_right_x":0.15},
|
||||
"description":"scan",
|
||||
"category":"logo"
|
||||
}]
|
||||
}
|
||||
},
|
||||
{"type":"blocks","blocks":[{"type":"text","text":{"content":"total"}}]}
|
||||
],
|
||||
"meta":{"api_version":{"version":"2"},"billed_units":{"pages":3}}
|
||||
}))
|
||||
.unwrap();
|
||||
let normalized = transform_response("parse-v5.0", response).unwrap();
|
||||
assert_eq!(normalized.pages[0]["index"], 4);
|
||||
assert_eq!(normalized.pages[0]["markdown"], "receipt");
|
||||
assert_eq!(normalized.pages[0]["images"][0]["bbox"]["top_left_x"], 1);
|
||||
assert_eq!(
|
||||
normalized.pages[0]["images"][0]["bounding_box_normalized"]["bottom_right_x"],
|
||||
0.15
|
||||
);
|
||||
assert_eq!(normalized.pages[0]["images"][0]["description"], "scan");
|
||||
assert_eq!(normalized.pages[0]["images"][0]["category"], "logo");
|
||||
assert_eq!(normalized.pages[1]["index"], 1);
|
||||
assert_eq!(normalized.pages[1]["markdown"], "");
|
||||
assert_eq!(normalized.pages[1]["blocks"][0]["text"]["content"], "total");
|
||||
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_defaults_and_invalid_fields() {
|
||||
for value in [
|
||||
json!({}),
|
||||
json!({"meta":null}),
|
||||
json!({"pages":[],"meta":{"billed_units":null}}),
|
||||
] {
|
||||
let normalized =
|
||||
transform_response("parse", serde_json::from_value(value).unwrap()).unwrap();
|
||||
assert!(normalized.pages.is_empty());
|
||||
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 0);
|
||||
}
|
||||
for value in [
|
||||
json!({"pages":null}),
|
||||
json!({"pages":[{"markdown":"text"}]}),
|
||||
json!({"pages":[{"index":"bad"}]}),
|
||||
] {
|
||||
assert!(serde_json::from_value::<CohereResponse>(value).is_err());
|
||||
}
|
||||
let normalized = transform_response(
|
||||
"parse",
|
||||
serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 1);
|
||||
assert!(normalized.pages[0]["images"].is_null());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_requires_image_and_supported_output_format() {
|
||||
for value in [
|
||||
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
|
||||
json!({"type":"image_url","image_url":""}),
|
||||
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
|
||||
] {
|
||||
assert_eq!(
|
||||
validate_document(&serde_json::from_value(value).unwrap()),
|
||||
Err(OcrRequestError::CohereImageOnly)
|
||||
);
|
||||
}
|
||||
assert!(
|
||||
serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err()
|
||||
);
|
||||
for format in ["markdown", "blocks"] {
|
||||
assert!(
|
||||
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
|
||||
);
|
||||
}
|
||||
let request = transform_request(
|
||||
"parse-v5.0",
|
||||
serde_json::from_value(json!({
|
||||
"type":"image_url",
|
||||
"image_url":"https://example.com/image.png"
|
||||
}))
|
||||
.unwrap(),
|
||||
serde_json::from_value(json!({})).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(request).unwrap()["output_format"],
|
||||
"markdown"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) use provider::{
|
||||
CohereParams, CohereRequest, CohereResponse, transform_request, transform_response,
|
||||
validate_document,
|
||||
};
|
||||
|
||||
use crate::Error;
|
||||
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
#[derive(Default)]
|
||||
pub(crate) struct CohereParseConfig;
|
||||
|
||||
impl CohereParseConfig {
|
||||
pub(crate) fn transform_ocr_request(
|
||||
&self,
|
||||
model: &str,
|
||||
document: crate::ocr::types::OcrDocument,
|
||||
params: CohereParams,
|
||||
) -> Result<CohereRequest, OcrRequestError> {
|
||||
transform_request(model, document, params)
|
||||
}
|
||||
}
|
||||
|
||||
impl BaseOcrConfig for CohereParseConfig {
|
||||
type ProviderResponse = CohereResponse;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let params = crate::ocr::wire::decode_request_value::<CohereParams>(
|
||||
serde_json::Value::Object(request.optional_params.clone()),
|
||||
"optional_params",
|
||||
)?;
|
||||
let headers = validate_environment(&request.connection, &credential_env)?;
|
||||
let url = complete_url(
|
||||
request
|
||||
.connection
|
||||
.api_base
|
||||
.as_deref()
|
||||
.unwrap_or(COHERE_PARSE_API_BASE),
|
||||
)?;
|
||||
let body = self.transform_ocr_request(&request.model, request.document.clone(), params)?;
|
||||
transform_request_body(client, request, &url, &headers, true, body, |body| {
|
||||
validate_document(&body.document)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: CohereResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
transform_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
fn complete_url(base: &str) -> Result<String, OcrError> {
|
||||
let parsed = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
|
||||
if !matches!(parsed.scheme(), "http" | "https") {
|
||||
return Err(invalid_api_base().into());
|
||||
}
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v2", "parse"]))
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| invalid_api_base().into())
|
||||
}
|
||||
|
||||
fn invalid_api_base() -> OcrRequestError {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let key = connection
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key| !key.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(COHERE_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
|
||||
.ok_or_else(|| {
|
||||
Error::Auth("Missing COHERE_API_KEY - set it in the environment or pass api_key".into())
|
||||
})?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn completes_provider_urls_without_duplicate_paths_and_preserves_queries() {
|
||||
for suffix in ["", "/v2", "/v2/parse"] {
|
||||
assert_eq!(
|
||||
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
|
||||
"https://example.com/v2/parse?tenant=a"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_urls_and_blank_keys() {
|
||||
assert!(complete_url("relative/path").is_err());
|
||||
assert!(complete_url("ftp://example.com").is_err());
|
||||
assert!(matches!(
|
||||
validate_environment(
|
||||
&OcrConnection {
|
||||
api_key: Some(" ".into()),
|
||||
..Default::default()
|
||||
},
|
||||
&|_| None,
|
||||
),
|
||||
Err(OcrError::Public(Error::Auth(_)))
|
||||
));
|
||||
}
|
||||
}
|
||||
1
litellm-rust/crates/core/src/llms/mistral/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/mistral/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
3
litellm-rust/crates/core/src/llms/mistral/ocr/mod.rs
Normal file
3
litellm-rust/crates/core/src/llms/mistral/ocr/mod.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
pub(crate) mod transformation;
|
||||
|
||||
pub(crate) use transformation::{MistralOcrParams, MistralOcrResponse};
|
||||
|
|
@ -1,7 +1,65 @@
|
|||
use super::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum MistralOcrPages {
|
||||
Range(String),
|
||||
Indices(Vec<i64>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub(crate) struct MistralOcrParams {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub pages: Option<MistralOcrPages>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub include_image_base64: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub image_limit: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub image_min_size: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub bbox_annotation_format: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub document_annotation_format: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub document_annotation_prompt: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extract_header: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extract_footer: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub table_format: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub confidence_scores_granularity: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub include_blocks: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct MistralOcrRequest {
|
||||
pub model: String,
|
||||
pub document: OcrDocument,
|
||||
#[serde(flatten)]
|
||||
pub params: MistralOcrParams,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct MistralOcrResponse {
|
||||
#[serde(default)]
|
||||
pub pages: Vec<Value>,
|
||||
pub model: Option<String>,
|
||||
pub document_annotation: Option<Value>,
|
||||
pub usage_info: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(crate) fn transform_ocr_request(
|
||||
model: &str,
|
||||
document: OcrDocument,
|
||||
|
|
@ -30,7 +88,7 @@ pub(crate) fn transform_ocr_response(
|
|||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
mod mapping_tests {
|
||||
use super::*;
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
|
@ -197,19 +255,19 @@ mod tests {
|
|||
#[rstest]
|
||||
fn transform_ocr_response_preserves_blocks_and_confidence_scores() {
|
||||
let response: MistralOcrResponse = serde_json::from_value(json!({
|
||||
"pages":[{
|
||||
"index":0,
|
||||
"markdown":"hello",
|
||||
"images":[{"id":"img-0","image_base64":"data:image/png;base64,AA=="}],
|
||||
"dimensions":{"width":612,"height":792,"dpi":72},
|
||||
"blocks":[{"type":"title","bbox":{"x":1},"confidence_scores":{"mean":0.98}}],
|
||||
"confidence_scores":{"average_page_confidence_score":0.99,"minimum_page_confidence_score":0.97}
|
||||
}],
|
||||
"model":"returned-model",
|
||||
"document_annotation":"{\"language\":\"en\"}",
|
||||
"usage_info":{"pages_processed":1}
|
||||
}))
|
||||
.unwrap();
|
||||
"pages":[{
|
||||
"index":0,
|
||||
"markdown":"hello",
|
||||
"images":[{"id":"img-0","image_base64":"data:image/png;base64,AA=="}],
|
||||
"dimensions":{"width":612,"height":792,"dpi":72},
|
||||
"blocks":[{"type":"title","bbox":{"x":1},"confidence_scores":{"mean":0.98}}],
|
||||
"confidence_scores":{"average_page_confidence_score":0.99,"minimum_page_confidence_score":0.97}
|
||||
}],
|
||||
"model":"returned-model",
|
||||
"document_annotation":"{\"language\":\"en\"}",
|
||||
"usage_info":{"pages_processed":1}
|
||||
}))
|
||||
.unwrap();
|
||||
let result = transform_ocr_response("model", response)
|
||||
.unwrap()
|
||||
.into_json();
|
||||
|
|
@ -248,3 +306,158 @@ mod tests {
|
|||
assert_eq!(result["pages"][0], page);
|
||||
}
|
||||
}
|
||||
|
||||
use crate::Error;
|
||||
use crate::constants::MISTRAL_OCR_API_BASE;
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
|
||||
};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY";
|
||||
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub(crate) struct MistralOCRConfig;
|
||||
|
||||
impl MistralOCRConfig {
|
||||
pub(crate) fn transform_ocr_request(
|
||||
&self,
|
||||
model: &str,
|
||||
document: crate::ocr::types::OcrDocument,
|
||||
params: &MistralOcrParams,
|
||||
) -> Result<MistralOcrRequest, OcrRequestError> {
|
||||
transform_ocr_request(model, document, params)
|
||||
}
|
||||
}
|
||||
|
||||
impl BaseOcrConfig for MistralOCRConfig {
|
||||
type ProviderResponse = MistralOcrResponse;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params: _extra_params,
|
||||
} = _prepare_ocr_request::<MistralOcrParams>(request)?;
|
||||
let headers = validate_environment(&request.connection, &credential_env)?;
|
||||
let url = get_complete_url(request.connection.api_base.as_deref())?;
|
||||
let body = self.transform_ocr_request(&request.model, request.document.clone(), ¶ms)?;
|
||||
transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: MistralOcrResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result<String, OcrError> {
|
||||
let base = api_base
|
||||
.map(str::trim)
|
||||
.filter(|base| !base.is_empty())
|
||||
.unwrap_or(MISTRAL_OCR_API_BASE);
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v1", "ocr"]))
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let api_key = connection
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key| !key.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
|
||||
.ok_or(Error::MissingApiKey {
|
||||
provider: "Mistral",
|
||||
})?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn complete_url_defaults_and_dedupes_v1() {
|
||||
assert_eq!(
|
||||
get_complete_url(None).unwrap(),
|
||||
"https://api.mistral.ai/v1/ocr"
|
||||
);
|
||||
assert_eq!(
|
||||
get_complete_url(Some("https://example.com/v1?tenant=a")).unwrap(),
|
||||
"https://example.com/v1/ocr?tenant=a"
|
||||
);
|
||||
assert_eq!(
|
||||
get_complete_url(Some("https://example.com/v1/ocr?tenant=a")).unwrap(),
|
||||
"https://example.com/v1/ocr?tenant=a"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_prefers_explicit_key_then_environment() {
|
||||
let explicit = OcrConnection {
|
||||
api_key: Some("explicit".into()),
|
||||
..OcrConnection::default()
|
||||
};
|
||||
assert_eq!(
|
||||
validate_environment(&explicit, &|_| Some("environment".into())).unwrap()[0],
|
||||
("Authorization".into(), "Bearer explicit".into())
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
validate_environment(&OcrConnection::default(), &|_| Some("environment".into()))
|
||||
.unwrap()[0],
|
||||
("Authorization".into(), "Bearer environment".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_preserves_forwarded_authorization() {
|
||||
let connection = OcrConnection {
|
||||
extra_headers: vec![("authorization".into(), "Bearer forwarded".into())],
|
||||
..OcrConnection::default()
|
||||
};
|
||||
assert_eq!(
|
||||
validate_environment(&connection, &|_| None).unwrap(),
|
||||
connection.extra_headers
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_rejects_missing_key() {
|
||||
assert!(matches!(
|
||||
validate_environment(&OcrConnection::default(), &|_| None),
|
||||
Err(OcrError::Public(Error::MissingApiKey {
|
||||
provider: "Mistral"
|
||||
}))
|
||||
));
|
||||
}
|
||||
}
|
||||
6
litellm-rust/crates/core/src/llms/mod.rs
Normal file
6
litellm-rust/crates/core/src/llms/mod.rs
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
pub(crate) mod azure_ai;
|
||||
pub(crate) mod base_llm;
|
||||
pub(crate) mod cohere;
|
||||
pub(crate) mod mistral;
|
||||
pub(crate) mod reducto;
|
||||
pub(crate) mod vertex_ai;
|
||||
1
litellm-rust/crates/core/src/llms/reducto/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/reducto/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
3
litellm-rust/crates/core/src/llms/reducto/ocr/mod.rs
Normal file
3
litellm-rust/crates/core/src/llms/reducto/ocr/mod.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
pub(crate) mod transformation;
|
||||
|
||||
pub(crate) use transformation::ReductoResponse;
|
||||
532
litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs
Normal file
532
litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs
Normal file
|
|
@ -0,0 +1,532 @@
|
|||
mod types {
|
||||
use serde::{Deserialize, Deserializer, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoV3Params {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub formatting: Option<Map<String, Value>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub retrieval: Option<Map<String, Value>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub settings: Option<Map<String, Value>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoLegacyParams {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub enhance: Option<Map<String, Value>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoV3Request {
|
||||
pub input: String,
|
||||
#[serde(flatten)]
|
||||
pub params: ReductoV3Params,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoLegacyRequest {
|
||||
pub document_url: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub options: Option<ReductoLegacyParams>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct ReductoUploadResponse {
|
||||
pub file_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct ReductoResponse {
|
||||
#[serde(default, deserialize_with = "present_nullable")]
|
||||
pub result: Option<Option<ReductoResult>>,
|
||||
pub usage: Option<ReductoUsage>,
|
||||
#[serde(default)]
|
||||
pub chunks: Option<Vec<ReductoChunk>>,
|
||||
}
|
||||
|
||||
fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
|
||||
deserializer: D,
|
||||
) -> Result<Option<Option<T>>, D::Error> {
|
||||
Option::<T>::deserialize(deserializer).map(Some)
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct ReductoResult {
|
||||
pub chunks: Option<Vec<ReductoChunk>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct ReductoUsage {
|
||||
#[serde(default, deserialize_with = "optional_i64")]
|
||||
pub num_pages: Option<i64>,
|
||||
#[serde(default, deserialize_with = "optional_f64")]
|
||||
pub credits: Option<f64>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct ReductoChunk {
|
||||
pub content: Option<String>,
|
||||
pub blocks: Option<Vec<ReductoBlock>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoBlock {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub content: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub bbox: Option<ReductoBoundingBox>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoBoundingBox {
|
||||
#[serde(default, deserialize_with = "optional_i64")]
|
||||
pub page: Option<i64>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
|
||||
fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<i64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_i64()
|
||||
.or_else(|| number.as_f64().and_then(checked_truncated_i64))
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected an integer")),
|
||||
Some(Value::String(value)) => value
|
||||
.trim()
|
||||
.parse::<i64>()
|
||||
.map(Some)
|
||||
.map_err(|_| serde::de::Error::custom("expected an integer")),
|
||||
Some(Value::Bool(value)) => Ok(Some(i64::from(value))),
|
||||
Some(_) => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<f64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_f64()
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a number")),
|
||||
Some(Value::String(value)) => value
|
||||
.trim()
|
||||
.parse::<f64>()
|
||||
.map(Some)
|
||||
.map_err(|_| serde::de::Error::custom("expected a number")),
|
||||
Some(_) => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn checked_truncated_i64(value: f64) -> Option<i64> {
|
||||
(value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64)
|
||||
.then(|| value.trunc() as i64)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) use mapping::transform_ocr_response;
|
||||
pub(crate) use types::ReductoResponse;
|
||||
|
||||
mod mapping {
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::types::*;
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "transform_ocr_request",
|
||||
target = "litellm::function_trace",
|
||||
level = "trace",
|
||||
skip_all
|
||||
)]
|
||||
pub(crate) fn transform_v3_ocr_request(
|
||||
_model: &str,
|
||||
document: OcrDocument,
|
||||
params: &ReductoV3Params,
|
||||
) -> Result<ReductoV3Request, OcrRequestError> {
|
||||
Ok(ReductoV3Request {
|
||||
input: document.source().to_string(),
|
||||
params: params.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "transform_ocr_request",
|
||||
target = "litellm::function_trace",
|
||||
level = "trace",
|
||||
skip_all
|
||||
)]
|
||||
pub(crate) fn transform_legacy_ocr_request(
|
||||
_model: &str,
|
||||
document: OcrDocument,
|
||||
params: &ReductoLegacyParams,
|
||||
) -> Result<ReductoLegacyRequest, OcrRequestError> {
|
||||
Ok(ReductoLegacyRequest {
|
||||
document_url: document.source().to_string(),
|
||||
options: params.enhance.as_ref().map(|_| params.clone()),
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_ocr_response(
|
||||
model: &str,
|
||||
response: ReductoResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
let result = match response.result {
|
||||
Some(result) => result.unwrap_or_default(),
|
||||
None => ReductoResult {
|
||||
chunks: response.chunks,
|
||||
},
|
||||
};
|
||||
let usage = response.usage.unwrap_or_default();
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages: build_pages(result.chunks.unwrap_or_default()),
|
||||
model: model.to_string(),
|
||||
document_annotation: None,
|
||||
usage_info: Some(json!({
|
||||
"pages_processed": usage.num_pages,
|
||||
"credits": usage.credits,
|
||||
})),
|
||||
object: "ocr".to_string(),
|
||||
extra_fields: serde_json::Map::new(),
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn build_pages(chunks: Vec<ReductoChunk>) -> Vec<Value> {
|
||||
let blocks_by_page = chunks
|
||||
.iter()
|
||||
.flat_map(|chunk| chunk.blocks.iter().flatten())
|
||||
.filter_map(|block| block.bbox.as_ref()?.page.map(|page| (page, block)))
|
||||
.fold(
|
||||
BTreeMap::<i64, Vec<&ReductoBlock>>::new(),
|
||||
|mut pages, (page, block)| {
|
||||
pages.entry(page).or_default().push(block);
|
||||
pages
|
||||
},
|
||||
);
|
||||
if blocks_by_page.is_empty() {
|
||||
let markdown = join_content(chunks.iter().map(|chunk| chunk.content.as_deref()));
|
||||
return if markdown.is_empty() {
|
||||
Vec::new()
|
||||
} else {
|
||||
vec![page(0, markdown, None)]
|
||||
};
|
||||
}
|
||||
blocks_by_page
|
||||
.into_iter()
|
||||
.map(|(index, blocks)| {
|
||||
let markdown = join_content(blocks.iter().map(|block| block.content.as_deref()));
|
||||
page(
|
||||
index.saturating_sub(1).max(0),
|
||||
markdown,
|
||||
Some(json!(blocks)),
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn join_content<'a>(content: impl Iterator<Item = Option<&'a str>>) -> String {
|
||||
content
|
||||
.flatten()
|
||||
.filter(|text| !text.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n\n")
|
||||
}
|
||||
|
||||
fn page(index: i64, markdown: String, blocks: Option<Value>) -> Value {
|
||||
let mut result = json!({"index":index,"markdown":markdown,"images":null});
|
||||
if let (Value::Object(fields), Some(blocks)) = (&mut result, blocks) {
|
||||
fields.insert("blocks".into(), blocks);
|
||||
}
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
mod common {
|
||||
use crate::Error;
|
||||
use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX};
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::ocr::document::InlineDocument;
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{OcrConnection, OcrDocument};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
pub(crate) trait BaseReductoOcrConfig: BaseOcrConfig {
|
||||
fn get_complete_url(&self, api_base: Option<&str>, path: &str) -> Result<String, OcrError> {
|
||||
get_complete_url(api_base, path)
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
&self,
|
||||
connection: &OcrConnection,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
validate_environment(connection, env_lookup)
|
||||
}
|
||||
|
||||
async fn prepare_document(
|
||||
&self,
|
||||
client: &crate::ocr::OcrClient,
|
||||
document: OcrDocument,
|
||||
connection: &OcrConnection,
|
||||
headers: &[(String, String)],
|
||||
) -> Result<OcrDocument, OcrError> {
|
||||
prepare_document(client, document, connection, headers).await
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn get_complete_url(api_base: Option<&str>, path: &str) -> Result<String, OcrError> {
|
||||
let base = api_base
|
||||
.map(str::trim)
|
||||
.filter(|base| !base.is_empty())
|
||||
.unwrap_or(REDUCTO_API_BASE);
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&[path]))
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let api_key = connection
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key| !key.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| {
|
||||
env_lookup(REDUCTO_API_KEY_ENV)
|
||||
.map(|key| key.trim().to_string())
|
||||
.filter(|key| !key.is_empty())
|
||||
})
|
||||
.ok_or(Error::MissingReductoApiKey)?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) async fn prepare_document(
|
||||
client: &crate::ocr::OcrClient,
|
||||
document: OcrDocument,
|
||||
connection: &OcrConnection,
|
||||
headers: &[(String, String)],
|
||||
) -> Result<OcrDocument, OcrError> {
|
||||
if document.source().starts_with(REDUCTO_ID_PREFIX) {
|
||||
if document.source()[REDUCTO_ID_PREFIX.len()..]
|
||||
.trim()
|
||||
.is_empty()
|
||||
{
|
||||
return Err(OcrRequestError::RequestField {
|
||||
path: "document file id".into(),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
return Ok(document);
|
||||
}
|
||||
let inline =
|
||||
InlineDocument::parse(document.source())?.ok_or(OcrRequestError::ReductoSource)?;
|
||||
let mime = inline.mime_type().to_string();
|
||||
let bytes = inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
|
||||
let part = reqwest::multipart::Part::bytes(bytes)
|
||||
.file_name("document")
|
||||
.mime_str(&mime)
|
||||
.map_err(|_| OcrRequestError::InvalidDataUri)?;
|
||||
let builder = client
|
||||
.provider_http()
|
||||
.post(get_complete_url(connection.api_base.as_deref(), "upload")?)
|
||||
.multipart(reqwest::multipart::Form::new().part("file", part))
|
||||
.timeout(connection.timeout);
|
||||
let builder = crate::http_utils::with_headers(
|
||||
builder,
|
||||
headers,
|
||||
crate::http_utils::HeaderPolicy::Except(&["content-type", "content-length"]),
|
||||
);
|
||||
let response = crate::http_utils::http_request(builder)
|
||||
.await
|
||||
.map_err(crate::error::TransportError::from)?;
|
||||
let uploaded =
|
||||
crate::ocr::client::read_json_response::<super::types::ReductoUploadResponse>(
|
||||
response,
|
||||
false,
|
||||
connection.max_response_bytes,
|
||||
)
|
||||
.await?
|
||||
.data;
|
||||
let file_id = uploaded
|
||||
.file_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|id| !id.is_empty());
|
||||
let Some(file_id) = file_id else {
|
||||
return Err(OcrResponseError::ResponseField {
|
||||
path: "file_id".into(),
|
||||
}
|
||||
.into());
|
||||
};
|
||||
Ok(document.with_source(file_id.to_string()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn explicit_key_precedes_environment_key() {
|
||||
let connection = OcrConnection {
|
||||
api_key: Some("passed-key".into()),
|
||||
..Default::default()
|
||||
};
|
||||
let headers = validate_environment(&connection, &|_| Some("env-key".into())).unwrap();
|
||||
assert_eq!(headers[0].1, "Bearer passed-key");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn blank_explicit_key_uses_environment_key() {
|
||||
let connection = OcrConnection {
|
||||
api_key: Some(" ".into()),
|
||||
..Default::default()
|
||||
};
|
||||
let headers = validate_environment(&connection, &|_| Some(" env-key ".into())).unwrap();
|
||||
assert_eq!(headers[0].1, "Bearer env-key");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn existing_authorization_skips_key_lookup() {
|
||||
let connection = OcrConnection {
|
||||
extra_headers: vec![("authorization".into(), "Bearer existing".into())],
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(
|
||||
validate_environment(&connection, &|_| None).unwrap(),
|
||||
connection.extra_headers
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mod legacy {
|
||||
use super::common::BaseReductoOcrConfig;
|
||||
use super::mapping as reducto;
|
||||
use super::types::{ReductoLegacyParams, ReductoResponse};
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::error::{OcrError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env,
|
||||
guardrail_document, merge_extra_params,
|
||||
};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct ReductoParseLegacyConfig;
|
||||
|
||||
impl BaseReductoOcrConfig for ReductoParseLegacyConfig {}
|
||||
|
||||
impl BaseOcrConfig for ReductoParseLegacyConfig {
|
||||
type ProviderResponse = ReductoResponse;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params,
|
||||
} = _prepare_ocr_request::<ReductoLegacyParams>(request)?;
|
||||
let headers = self.validate_environment(&request.connection, &credential_env)?;
|
||||
let url = self.get_complete_url(request.connection.api_base.as_deref(), "parse")?;
|
||||
let (document, headers) = guardrail_document(request, &url, &headers).await?;
|
||||
let document = self
|
||||
.prepare_document(client, document, &request.connection, &headers)
|
||||
.await?;
|
||||
let body = reducto::transform_legacy_ocr_request(&request.model, document, ¶ms)?;
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: ReductoResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
reducto::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) use legacy::ReductoParseLegacyConfig;
|
||||
|
||||
mod v3 {
|
||||
use super::common::BaseReductoOcrConfig;
|
||||
use super::mapping as reducto;
|
||||
use super::types::{ReductoResponse, ReductoV3Params};
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::error::{OcrError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env,
|
||||
guardrail_document, merge_extra_params,
|
||||
};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct ReductoParseV3Config;
|
||||
|
||||
impl BaseReductoOcrConfig for ReductoParseV3Config {}
|
||||
|
||||
impl BaseOcrConfig for ReductoParseV3Config {
|
||||
type ProviderResponse = ReductoResponse;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params,
|
||||
} = _prepare_ocr_request::<ReductoV3Params>(request)?;
|
||||
let headers = self.validate_environment(&request.connection, &credential_env)?;
|
||||
let url = self.get_complete_url(request.connection.api_base.as_deref(), "parse")?;
|
||||
let (document, headers) = guardrail_document(request, &url, &headers).await?;
|
||||
let document = self
|
||||
.prepare_document(client, document, &request.connection, &headers)
|
||||
.await?;
|
||||
let body = reducto::transform_v3_ocr_request(&request.model, document, ¶ms)?;
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: ReductoResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
reducto::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) use v3::ReductoParseV3Config;
|
||||
1
litellm-rust/crates/core/src/llms/vertex_ai/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/vertex_ai/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
|
|
@ -1,16 +1,10 @@
|
|||
mod deepseek;
|
||||
mod mistral;
|
||||
|
||||
use crate::Error;
|
||||
use crate::auth::InputSource;
|
||||
use crate::auth::error::AuthConfigurationError;
|
||||
use crate::ocr::error::OcrError;
|
||||
use crate::ocr::types::OcrConnection;
|
||||
|
||||
pub(crate) use deepseek::VertexDeepSeekAdapter;
|
||||
pub(crate) use mistral::VertexMistralAdapter;
|
||||
|
||||
fn validate_destination(connection: &OcrConnection) -> Result<(), OcrError> {
|
||||
pub(super) fn validate_destination(connection: &OcrConnection) -> Result<(), OcrError> {
|
||||
if connection.api_base.is_some() && connection.api_base_source == InputSource::Request {
|
||||
return Err(Error::from(crate::AuthError::Configuration(
|
||||
AuthConfigurationError::RequestVertexCredentialDestination,
|
||||
|
|
@ -0,0 +1,347 @@
|
|||
mod types {
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrParams {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stream: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_tokens: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub top_p: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub n: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stop: Option<StopSequences>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum StopSequences {
|
||||
One(String),
|
||||
Many(Vec<String>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrRequest {
|
||||
pub model: String,
|
||||
pub messages: Vec<DeepSeekOcrMessage>,
|
||||
#[serde(flatten)]
|
||||
pub params: DeepSeekOcrParams,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrMessage {
|
||||
pub role: UserRole,
|
||||
pub content: Vec<crate::ocr::types::OcrDocument>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub(crate) enum UserRole {
|
||||
User,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrResponse {
|
||||
#[serde(default)]
|
||||
pub choices: Vec<DeepSeekChoice>,
|
||||
pub usage: Option<Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct DeepSeekChoice {
|
||||
pub message: DeepSeekResponseMessage,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct DeepSeekResponseMessage {
|
||||
pub content: Option<DeepSeekContent>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum DeepSeekContent {
|
||||
Text(String),
|
||||
Object(DeepSeekOcrResult),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrResult {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub pages: Option<Vec<DeepSeekPage>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub model: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub usage_info: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub document_annotation: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekPage {
|
||||
#[serde(default)]
|
||||
pub index: i64,
|
||||
#[serde(default)]
|
||||
pub markdown: String,
|
||||
pub images: Option<Value>,
|
||||
pub dimensions: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) use types::DeepSeekOcrParams;
|
||||
pub(crate) use types::DeepSeekOcrResponse;
|
||||
|
||||
mod mapping {
|
||||
use serde::de::IntoDeserializer;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::types::*;
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(crate) fn transform_ocr_request(
|
||||
provider_model: &str,
|
||||
document: OcrDocument,
|
||||
params: &DeepSeekOcrParams,
|
||||
) -> Result<DeepSeekOcrRequest, OcrRequestError> {
|
||||
if document.source().is_empty() {
|
||||
return Err(OcrRequestError::MissingDocumentUrl);
|
||||
}
|
||||
let content = OcrDocument::ImageUrl {
|
||||
image_url: document.source().to_string(),
|
||||
extra_fields: serde_json::Map::new(),
|
||||
};
|
||||
Ok(DeepSeekOcrRequest {
|
||||
model: provider_model.to_string(),
|
||||
messages: vec![DeepSeekOcrMessage {
|
||||
role: UserRole::User,
|
||||
content: vec![content],
|
||||
}],
|
||||
params: params.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_ocr_response(
|
||||
model: &str,
|
||||
response: DeepSeekOcrResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
let content = response
|
||||
.choices
|
||||
.into_iter()
|
||||
.next()
|
||||
.and_then(|choice| choice.message.content)
|
||||
.ok_or(OcrResponseError::EmptyContent)?;
|
||||
let decoded = decode_content(content)?;
|
||||
let pages = match decoded.result.pages {
|
||||
Some(pages) if !pages.is_empty() => pages
|
||||
.into_iter()
|
||||
.map(|page| serde_json::to_value(page).expect("DeepSeek page serializes"))
|
||||
.collect(),
|
||||
_ => vec![json!({
|
||||
"index":0,
|
||||
"markdown":decoded.fallback_markdown,
|
||||
"images":null
|
||||
})],
|
||||
};
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages,
|
||||
model: decoded.result.model.unwrap_or_else(|| model.to_string()),
|
||||
document_annotation: decoded.result.document_annotation,
|
||||
usage_info: decoded.result.usage_info.or(response.usage),
|
||||
object: "ocr".into(),
|
||||
extra_fields: decoded.result.extra_fields,
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
struct DecodedContent {
|
||||
result: DeepSeekOcrResult,
|
||||
fallback_markdown: String,
|
||||
}
|
||||
|
||||
fn decode_content(content: DeepSeekContent) -> Result<DecodedContent, OcrResponseError> {
|
||||
let (result, fallback_markdown) = match content {
|
||||
DeepSeekContent::Text(text) if text.is_empty() => {
|
||||
return Err(OcrResponseError::EmptyContent);
|
||||
}
|
||||
DeepSeekContent::Text(text) => (decode_json_content(&text)?, text),
|
||||
DeepSeekContent::Object(object) => {
|
||||
let fallback = serde_json::to_string(&object).map_err(|_| {
|
||||
OcrResponseError::ResponseField {
|
||||
path: "choices[0].message.content".into(),
|
||||
}
|
||||
})?;
|
||||
(Some(object), fallback)
|
||||
}
|
||||
};
|
||||
Ok(DecodedContent {
|
||||
result: result.unwrap_or_default(),
|
||||
fallback_markdown,
|
||||
})
|
||||
}
|
||||
|
||||
fn decode_json_content(text: &str) -> Result<Option<DeepSeekOcrResult>, OcrResponseError> {
|
||||
if !text.trim_start().starts_with('{') {
|
||||
return Ok(None);
|
||||
}
|
||||
let value = match serde_json::from_str::<Value>(text) {
|
||||
Ok(value) => value,
|
||||
Err(_) => return Ok(None),
|
||||
};
|
||||
serde_path_to_error::deserialize(value.into_deserializer())
|
||||
.map(Some)
|
||||
.map_err(|error| OcrResponseError::ResponseField {
|
||||
path: format!("choices[0].message.content.{}", error.path()),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) use mapping::{transform_ocr_request, transform_ocr_response};
|
||||
|
||||
use super::common_utils::validate_destination;
|
||||
use crate::Error;
|
||||
use crate::auth::vertex::{self, VertexConfig};
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
|
||||
};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
use crate::url_utils::ApiUrl;
|
||||
const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com";
|
||||
const MODEL_NAMESPACE: &str = "deepseek-ai";
|
||||
const DEFAULT_LOCATION: &str = "us-central1";
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct VertexAIDeepSeekOCRConfig;
|
||||
|
||||
impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
|
||||
type ProviderResponse = DeepSeekOcrResponse;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
validate_destination(&request.connection)?;
|
||||
let ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params: _extra_params,
|
||||
} = _prepare_ocr_request::<DeepSeekOcrParams>(request)?;
|
||||
let config = VertexConfig::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
)
|
||||
.map_err(Error::from)?;
|
||||
let authentication = client
|
||||
.vertex_auth()
|
||||
.validate_environment(
|
||||
request.connection.extra_headers.clone(),
|
||||
request.connection.api_key.as_deref(),
|
||||
&config,
|
||||
&credential_env,
|
||||
)
|
||||
.await
|
||||
.map_err(Error::from)?;
|
||||
let location = vertex::get_vertex_ai_location(&config, &credential_env)
|
||||
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
|
||||
let url = get_complete_url(
|
||||
request.connection.api_base.as_deref(),
|
||||
&authentication.project_id,
|
||||
&location,
|
||||
)?;
|
||||
let document = request.document.clone();
|
||||
let body =
|
||||
mapping::transform_ocr_request(&provider_model(&request.model), document, ¶ms)?;
|
||||
transform_request_body(
|
||||
client,
|
||||
request,
|
||||
&url,
|
||||
&authentication.headers,
|
||||
false,
|
||||
body,
|
||||
|_| Ok(()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: DeepSeekOcrResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
mapping::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_model(model: &str) -> String {
|
||||
if model.starts_with(&format!("{MODEL_NAMESPACE}/")) {
|
||||
model.to_string()
|
||||
} else {
|
||||
format!("{MODEL_NAMESPACE}/{model}")
|
||||
}
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
api_base: Option<&str>,
|
||||
project: &str,
|
||||
location: &str,
|
||||
) -> Result<String, OcrError> {
|
||||
let base = api_base
|
||||
.map(str::trim)
|
||||
.filter(|base| !base.is_empty())
|
||||
.unwrap_or(DEFAULT_API_BASE);
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| {
|
||||
url.complete_path(&[
|
||||
"v1",
|
||||
"projects",
|
||||
project,
|
||||
"locations",
|
||||
location,
|
||||
"endpoints",
|
||||
"openapi",
|
||||
"chat",
|
||||
"completions",
|
||||
])
|
||||
})
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{get_complete_url, provider_model};
|
||||
|
||||
#[test]
|
||||
fn config_owns_model_namespace_and_endpoint() {
|
||||
assert_eq!(
|
||||
provider_model("deepseek-ocr-maas"),
|
||||
"deepseek-ai/deepseek-ocr-maas"
|
||||
);
|
||||
assert_eq!(
|
||||
provider_model("deepseek-ai/deepseek-ocr-maas"),
|
||||
"deepseek-ai/deepseek-ocr-maas"
|
||||
);
|
||||
assert_eq!(
|
||||
get_complete_url(None, "proj-1", "europe-west4").unwrap(),
|
||||
"https://aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/endpoints/openapi/chat/completions"
|
||||
);
|
||||
}
|
||||
}
|
||||
3
litellm-rust/crates/core/src/llms/vertex_ai/ocr/mod.rs
Normal file
3
litellm-rust/crates/core/src/llms/vertex_ai/ocr/mod.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
pub(crate) mod common_utils;
|
||||
pub(crate) mod deepseek_transformation;
|
||||
pub(crate) mod transformation;
|
||||
|
|
@ -1,25 +1,26 @@
|
|||
use super::super::OcrAdapter;
|
||||
use super::validate_destination;
|
||||
use super::common_utils::validate_destination;
|
||||
use crate::Error;
|
||||
use crate::auth::vertex::{self, VertexConfig};
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
|
||||
use crate::llms::mistral::ocr::{MistralOcrParams, MistralOcrResponse};
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse};
|
||||
use crate::ocr::document::{inline_remote_document, validate_inline_document};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
|
||||
};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
use crate::url_utils::ApiUrl;
|
||||
const DEFAULT_LOCATION: &str = "us-central1";
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct VertexMistralAdapter;
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub(crate) struct VertexAIOCRConfig {
|
||||
mistral: MistralOCRConfig,
|
||||
}
|
||||
|
||||
impl OcrAdapter for VertexMistralAdapter {
|
||||
impl BaseOcrConfig for VertexAIOCRConfig {
|
||||
type ProviderResponse = MistralOcrResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::VertexAi;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
|
|
@ -62,7 +63,9 @@ impl OcrAdapter for VertexMistralAdapter {
|
|||
&request.connection,
|
||||
)
|
||||
.await?;
|
||||
let body = mistral::transform_ocr_request(&request.model, document, ¶ms)?;
|
||||
let body = self
|
||||
.mistral
|
||||
.transform_ocr_request(&request.model, document, ¶ms)?;
|
||||
transform_request_body(
|
||||
client,
|
||||
request,
|
||||
|
|
@ -78,9 +81,9 @@ impl OcrAdapter for VertexMistralAdapter {
|
|||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
response: MistralOcrResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
mistral::transform_ocr_response(&request.model, response)
|
||||
self.mistral.transform_ocr_response(request, response)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -1,214 +0,0 @@
|
|||
use super::super::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::auth::{InputSource, Sourced};
|
||||
use crate::constants::{AZURE_DI_API_VERSION, AZURE_DI_SUBSCRIPTION_HEADER};
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::document_intelligence::{
|
||||
self, AzureDocumentIntelligenceOperation, DocumentIntelligenceParams,
|
||||
};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrResponseFormat};
|
||||
use crate::providers::azure_ai::auth::AzureAuthInputs;
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
mod polling;
|
||||
|
||||
const AZURE_DI_API_KEY_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY";
|
||||
const AZURE_DI_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT";
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct AzureDocumentIntelligenceAdapter;
|
||||
|
||||
impl OcrAdapter for AzureDocumentIntelligenceAdapter {
|
||||
type ProviderResponse = AzureDocumentIntelligenceOperation;
|
||||
const PROVIDER: OcrProvider = OcrProvider::AzureAi;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let params = map_ocr_params(request)?;
|
||||
let mut config = AzureAuthInputs::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
)
|
||||
.map_err(Error::from)?;
|
||||
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
|
||||
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
|
||||
let endpoint = nonblank(request.connection.api_base.clone())
|
||||
.or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV)))
|
||||
.ok_or_else(|| Error::Auth("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into()))?;
|
||||
let url = get_complete_url(&endpoint, &request.model, ¶ms)?;
|
||||
let body = document_intelligence::transform_ocr_request(request.document.clone())?;
|
||||
transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
document_intelligence::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
|
||||
async fn read_response(
|
||||
&self,
|
||||
client: &OcrClient,
|
||||
response: reqwest::Response,
|
||||
url: &str,
|
||||
headers: &[(String, String)],
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> Result<crate::ocr::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError> {
|
||||
polling::read_operation_response(
|
||||
client.polling_http(),
|
||||
response,
|
||||
url,
|
||||
headers,
|
||||
&request.connection,
|
||||
request.response_format()? == OcrResponseFormat::Native,
|
||||
&request.hooks,
|
||||
)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
fn map_ocr_params(
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> Result<DocumentIntelligenceParams, OcrRequestError> {
|
||||
let params = document_intelligence::decode_input_params(
|
||||
request.optional_params.clone(),
|
||||
"optional_params",
|
||||
)?;
|
||||
let crate::ocr::prepare::ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params: _extra_params,
|
||||
} = params;
|
||||
document_intelligence::map_ocr_params(params)
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
endpoint: &str,
|
||||
model: &str,
|
||||
params: &DocumentIntelligenceParams,
|
||||
) -> Result<String, OcrError> {
|
||||
let model = format!("{}:analyze", model_id(model)?);
|
||||
ApiUrl::parse(endpoint)
|
||||
.and_then(|url| url.complete_path(&["documentintelligence", "documentModels", &model]))
|
||||
.map(|url| {
|
||||
url.append_query_pairs(
|
||||
[("api-version", AZURE_DI_API_VERSION)]
|
||||
.into_iter()
|
||||
.chain(params.pages.iter().map(|pages| ("pages", pages.as_str())))
|
||||
.chain(
|
||||
params
|
||||
.features
|
||||
.iter()
|
||||
.map(|features| ("features", features.as_str())),
|
||||
),
|
||||
)
|
||||
.into_string()
|
||||
})
|
||||
.map_err(|_| OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
})
|
||||
.map_err(OcrError::from)
|
||||
}
|
||||
|
||||
async fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
config: &AzureAuthInputs,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization")
|
||||
|| crate::http_utils::has_header(&connection.extra_headers, AZURE_DI_SUBSCRIPTION_HEADER)
|
||||
{
|
||||
super::validate_destination(connection, connection.extra_headers_source)?;
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let key = nonblank(connection.api_key.clone())
|
||||
.map(|value| Sourced::new(value, connection.api_key_source))
|
||||
.or_else(|| {
|
||||
nonblank(env_lookup(AZURE_DI_API_KEY_ENV))
|
||||
.map(|value| Sourced::new(value, InputSource::Environment))
|
||||
});
|
||||
if let Some(key) = key {
|
||||
super::validate_destination(connection, key.source())?;
|
||||
return Ok(
|
||||
std::iter::once((AZURE_DI_SUBSCRIPTION_HEADER.into(), key.into_value()))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
);
|
||||
}
|
||||
let token = super::resolve_entra(config, env_lookup)
|
||||
.await?
|
||||
.ok_or(Error::MissingAzureDocumentIntelligenceCredentials)?;
|
||||
super::validate_destination(connection, token.source())?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {}", token.value())))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
fn model_id(model: &str) -> Result<&str, OcrRequestError> {
|
||||
let model = model.rsplit('/').next().unwrap_or(model);
|
||||
if matches!(model, "." | "..") {
|
||||
return Err(OcrRequestError::DotModel);
|
||||
}
|
||||
Ok(model)
|
||||
}
|
||||
|
||||
fn nonblank(value: Option<String>) -> Option<String> {
|
||||
value
|
||||
.map(|value| value.trim().to_string())
|
||||
.filter(|value| !value.is_empty())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_endpoint_cannot_receive_environment_key() {
|
||||
let connection = OcrConnection {
|
||||
api_base: Some("https://request.example".into()),
|
||||
api_base_source: InputSource::Request,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let error = validate_environment(&connection, &Default::default(), &|name| {
|
||||
(name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into())
|
||||
})
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("request-controlled Azure endpoint")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_endpoint_accepts_request_owned_key() {
|
||||
let connection = OcrConnection {
|
||||
api_key: Some("request-key".into()),
|
||||
api_key_source: InputSource::Request,
|
||||
api_base: Some("https://request.example".into()),
|
||||
api_base_source: InputSource::Request,
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let headers = validate_environment(&connection, &Default::default(), &|_| None)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
headers[0],
|
||||
(AZURE_DI_SUBSCRIPTION_HEADER.into(), "request-key".into())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,119 +0,0 @@
|
|||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use reqwest::Url;
|
||||
use tokio::time::Instant;
|
||||
|
||||
use crate::constants::{AZURE_DI_SUBSCRIPTION_HEADER, OCR_POLL_RETRY_SECS};
|
||||
use crate::ocr::client::read_json_response;
|
||||
use crate::ocr::codecs::document_intelligence::{
|
||||
AzureDocumentIntelligenceOperation, OperationStatus,
|
||||
};
|
||||
use crate::ocr::error::{OcrError, OcrPollingError, OcrResponseError};
|
||||
use crate::ocr::hooks::OcrHooks;
|
||||
use crate::ocr::types::OcrConnection;
|
||||
use crate::ocr::wire::DecodedOcrResponse;
|
||||
|
||||
pub(super) async fn read_operation_response(
|
||||
http_client: &reqwest::Client,
|
||||
response: reqwest::Response,
|
||||
original_url: &str,
|
||||
headers: &[(String, String)],
|
||||
connection: &OcrConnection,
|
||||
native: bool,
|
||||
hooks: &Arc<dyn OcrHooks>,
|
||||
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
|
||||
if response.status() != reqwest::StatusCode::ACCEPTED {
|
||||
let bytes =
|
||||
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes)
|
||||
.await?;
|
||||
crate::ocr::handler::post_call(hooks, &bytes).await?;
|
||||
return Ok(crate::ocr::wire::decode_response(&bytes, native)?);
|
||||
}
|
||||
let location = response
|
||||
.headers()
|
||||
.get("operation-location")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.ok_or(OcrPollingError::PollLocation)?
|
||||
.to_string();
|
||||
let original = Url::parse(original_url).map_err(|_| OcrPollingError::PollOrigin)?;
|
||||
let operation = Url::parse(&location).map_err(|_| OcrPollingError::PollOrigin)?;
|
||||
if original.origin() != operation.origin()
|
||||
|| !operation.username().is_empty()
|
||||
|| operation.password().is_some()
|
||||
{
|
||||
return Err(OcrPollingError::PollOrigin.into());
|
||||
}
|
||||
let bytes =
|
||||
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes).await?;
|
||||
crate::ocr::handler::post_call(hooks, &bytes).await?;
|
||||
poll_operation(http_client, operation, headers, connection, native, hooks).await
|
||||
}
|
||||
|
||||
async fn poll_operation(
|
||||
http_client: &reqwest::Client,
|
||||
url: Url,
|
||||
headers: &[(String, String)],
|
||||
connection: &OcrConnection,
|
||||
native: bool,
|
||||
hooks: &Arc<dyn OcrHooks>,
|
||||
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
|
||||
let deadline = Instant::now()
|
||||
.checked_add(connection.poll_timeout)
|
||||
.ok_or(OcrPollingError::PollTimeout)?;
|
||||
loop {
|
||||
let remaining = deadline
|
||||
.checked_duration_since(Instant::now())
|
||||
.filter(|remaining| !remaining.is_zero())
|
||||
.ok_or(OcrPollingError::PollTimeout)?;
|
||||
let builder = http_client
|
||||
.get(url.clone())
|
||||
.timeout(remaining.min(connection.timeout));
|
||||
let builder = crate::http_utils::with_headers(
|
||||
builder,
|
||||
headers,
|
||||
crate::http_utils::HeaderPolicy::Only(&[AZURE_DI_SUBSCRIPTION_HEADER, "authorization"]),
|
||||
);
|
||||
let response = tokio::time::timeout_at(deadline, crate::http_utils::http_request(builder))
|
||||
.await
|
||||
.map_err(|_| OcrPollingError::PollTimeout)?
|
||||
.map_err(crate::error::TransportError::from)?;
|
||||
let retry = response
|
||||
.headers()
|
||||
.get(reqwest::header::RETRY_AFTER)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(|value| value.parse::<u64>().ok())
|
||||
.unwrap_or(OCR_POLL_RETRY_SECS)
|
||||
.max(1);
|
||||
let decoded = tokio::time::timeout_at(
|
||||
deadline,
|
||||
read_json_response::<AzureDocumentIntelligenceOperation>(
|
||||
response,
|
||||
native,
|
||||
connection.max_response_bytes,
|
||||
),
|
||||
)
|
||||
.await
|
||||
.map_err(|_| OcrPollingError::PollTimeout)??;
|
||||
match &decoded.data.status {
|
||||
Some(OperationStatus::Succeeded) => {
|
||||
crate::ocr::handler::post_call(hooks, decoded.text.as_bytes()).await?;
|
||||
return Ok(decoded);
|
||||
}
|
||||
Some(OperationStatus::Running | OperationStatus::NotStarted) => {
|
||||
tokio::time::timeout_at(deadline, tokio::time::sleep(Duration::from_secs(retry)))
|
||||
.await
|
||||
.map_err(|_| OcrPollingError::PollTimeout)?;
|
||||
}
|
||||
status => {
|
||||
return Err(OcrResponseError::OperationStatus(
|
||||
status
|
||||
.as_ref()
|
||||
.map(ToString::to_string)
|
||||
.unwrap_or_else(|| "None".into()),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -1,123 +0,0 @@
|
|||
use super::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::cohere::{
|
||||
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
|
||||
};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
pub(crate) struct CohereAdapter;
|
||||
|
||||
impl OcrAdapter for CohereAdapter {
|
||||
type ProviderResponse = CohereResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::Cohere;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let params = super::super::wire::decode_request_value::<CohereParams>(
|
||||
serde_json::Value::Object(request.optional_params.clone()),
|
||||
"optional_params",
|
||||
)?;
|
||||
let headers = validate_environment(&request.connection, &credential_env)?;
|
||||
let url = complete_url(
|
||||
request
|
||||
.connection
|
||||
.api_base
|
||||
.as_deref()
|
||||
.unwrap_or(COHERE_PARSE_API_BASE),
|
||||
)?;
|
||||
let body = transform_request(&request.model, request.document.clone(), params)?;
|
||||
transform_request_body(client, request, &url, &headers, true, body, |body| {
|
||||
validate_document(&body.document)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
transform_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
fn complete_url(base: &str) -> Result<String, OcrError> {
|
||||
let parsed = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
|
||||
if !matches!(parsed.scheme(), "http" | "https") {
|
||||
return Err(invalid_api_base().into());
|
||||
}
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v2", "parse"]))
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| invalid_api_base().into())
|
||||
}
|
||||
|
||||
fn invalid_api_base() -> OcrRequestError {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let key = connection
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key| !key.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(COHERE_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
|
||||
.ok_or_else(|| {
|
||||
Error::Auth("Missing COHERE_API_KEY - set it in the environment or pass api_key".into())
|
||||
})?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn completes_provider_urls_without_duplicate_paths_and_preserves_queries() {
|
||||
for suffix in ["", "/v2", "/v2/parse"] {
|
||||
assert_eq!(
|
||||
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
|
||||
"https://example.com/v2/parse?tenant=a"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_urls_and_blank_keys() {
|
||||
assert!(complete_url("relative/path").is_err());
|
||||
assert!(complete_url("ftp://example.com").is_err());
|
||||
assert!(matches!(
|
||||
validate_environment(
|
||||
&OcrConnection {
|
||||
api_key: Some(" ".into()),
|
||||
..Default::default()
|
||||
},
|
||||
&|_| None,
|
||||
),
|
||||
Err(OcrError::Public(Error::Auth(_)))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
@ -1,147 +0,0 @@
|
|||
use super::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::constants::MISTRAL_OCR_API_BASE;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
|
||||
};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY";
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct MistralAdapter;
|
||||
|
||||
impl OcrAdapter for MistralAdapter {
|
||||
type ProviderResponse = MistralOcrResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::Mistral;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params: _extra_params,
|
||||
} = _prepare_ocr_request::<MistralOcrParams>(request)?;
|
||||
let headers = validate_environment(&request.connection, &credential_env)?;
|
||||
let url = get_complete_url(request.connection.api_base.as_deref())?;
|
||||
let body =
|
||||
mistral::transform_ocr_request(&request.model, request.document.clone(), ¶ms)?;
|
||||
transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
mistral::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result<String, OcrError> {
|
||||
let base = api_base
|
||||
.map(str::trim)
|
||||
.filter(|base| !base.is_empty())
|
||||
.unwrap_or(MISTRAL_OCR_API_BASE);
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v1", "ocr"]))
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let api_key = connection
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key| !key.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
|
||||
.ok_or(Error::MissingApiKey {
|
||||
provider: "Mistral",
|
||||
})?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn complete_url_defaults_and_dedupes_v1() {
|
||||
assert_eq!(
|
||||
get_complete_url(None).unwrap(),
|
||||
"https://api.mistral.ai/v1/ocr"
|
||||
);
|
||||
assert_eq!(
|
||||
get_complete_url(Some("https://example.com/v1?tenant=a")).unwrap(),
|
||||
"https://example.com/v1/ocr?tenant=a"
|
||||
);
|
||||
assert_eq!(
|
||||
get_complete_url(Some("https://example.com/v1/ocr?tenant=a")).unwrap(),
|
||||
"https://example.com/v1/ocr?tenant=a"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_prefers_explicit_key_then_environment() {
|
||||
let explicit = OcrConnection {
|
||||
api_key: Some("explicit".into()),
|
||||
..OcrConnection::default()
|
||||
};
|
||||
assert_eq!(
|
||||
validate_environment(&explicit, &|_| Some("environment".into())).unwrap()[0],
|
||||
("Authorization".into(), "Bearer explicit".into())
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
validate_environment(&OcrConnection::default(), &|_| Some("environment".into()))
|
||||
.unwrap()[0],
|
||||
("Authorization".into(), "Bearer environment".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_preserves_forwarded_authorization() {
|
||||
let connection = OcrConnection {
|
||||
extra_headers: vec![("authorization".into(), "Bearer forwarded".into())],
|
||||
..OcrConnection::default()
|
||||
};
|
||||
assert_eq!(
|
||||
validate_environment(&connection, &|_| None).unwrap(),
|
||||
connection.extra_headers
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_rejects_missing_key() {
|
||||
assert!(matches!(
|
||||
validate_environment(&OcrConnection::default(), &|_| None),
|
||||
Err(OcrError::Public(Error::MissingApiKey {
|
||||
provider: "Mistral"
|
||||
}))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
@ -1,91 +0,0 @@
|
|||
use std::future::Future;
|
||||
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
use super::OcrClient;
|
||||
use super::error::{OcrError, OcrResponseError};
|
||||
use super::registry::OcrProvider;
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
|
||||
mod azure;
|
||||
mod cohere;
|
||||
mod mistral;
|
||||
mod reducto;
|
||||
mod vertex;
|
||||
|
||||
pub(crate) use azure::{AzureCohereAdapter, AzureDocumentIntelligenceAdapter, AzureMistralAdapter};
|
||||
pub(crate) use cohere::CohereAdapter;
|
||||
pub(crate) use mistral::MistralAdapter;
|
||||
pub(crate) use reducto::{ReductoLegacyAdapter, ReductoV3Adapter};
|
||||
pub(crate) use vertex::{VertexDeepSeekAdapter, VertexMistralAdapter};
|
||||
|
||||
/// Converts a complete LiteLLM OCR call to provider HTTP and normalizes its response.
|
||||
pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
|
||||
/// Provider JSON schema; direct and Vertex Mistral share `MistralOcrResponse`.
|
||||
type ProviderResponse: DeserializeOwned + Send;
|
||||
|
||||
const PROVIDER: OcrProvider;
|
||||
|
||||
/// Prepares the complete provider HTTP request.
|
||||
/// `request` contains the model, document, connection, and unmapped caller options.
|
||||
/// `client` supplies reusable provider and document HTTP clients.
|
||||
/// Returns the complete HTTP request, whereas Python returns body data.
|
||||
fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> impl Future<Output = Result<reqwest::Request, OcrError>> + Send;
|
||||
|
||||
/// Python: `transform_ocr_response`.
|
||||
/// `request` supplies caller context, including the fallback model.
|
||||
/// `response` is the decoded provider payload; the output is the shared LiteLLM schema.
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError>;
|
||||
|
||||
/// Decodes provider HTTP; adapters may override this to poll asynchronous operations.
|
||||
/// Python performs that polling inside `async_transform_ocr_response`.
|
||||
/// `client` is reused for polling; `response` is the initial HTTP response.
|
||||
/// `url` and `headers` describe the submitted call; `request` supplies limits and format.
|
||||
fn read_response(
|
||||
&self,
|
||||
_client: &OcrClient,
|
||||
response: reqwest::Response,
|
||||
_url: &str,
|
||||
_headers: &[(String, String)],
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> impl Future<
|
||||
Output = Result<super::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError>,
|
||||
> + Send {
|
||||
async move {
|
||||
let bytes =
|
||||
super::client::read_response_bytes(response, request.connection.max_response_bytes)
|
||||
.await?;
|
||||
super::handler::post_call(&request.hooks, &bytes).await?;
|
||||
Ok(super::wire::decode_response(
|
||||
&bytes,
|
||||
request.response_format()? == super::types::OcrResponseFormat::Native,
|
||||
)?)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
macro_rules! for_each_ocr_adapter {
|
||||
($callback:ident) => {
|
||||
$callback! {
|
||||
Cohere, $crate::ocr::adapters::CohereAdapter, $crate::ocr::adapters::CohereAdapter, Cohere;
|
||||
AzureCohere, $crate::ocr::adapters::AzureCohereAdapter, $crate::ocr::adapters::AzureCohereAdapter, AzureAi;
|
||||
Mistral, $crate::ocr::adapters::MistralAdapter, $crate::ocr::adapters::MistralAdapter, Mistral;
|
||||
AzureMistral, $crate::ocr::adapters::AzureMistralAdapter, $crate::ocr::adapters::AzureMistralAdapter, AzureAi;
|
||||
AzureDocumentIntelligence, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, AzureAi;
|
||||
ReductoLegacy, $crate::ocr::adapters::ReductoLegacyAdapter, $crate::ocr::adapters::ReductoLegacyAdapter, Reducto;
|
||||
ReductoV3, $crate::ocr::adapters::ReductoV3Adapter, $crate::ocr::adapters::ReductoV3Adapter, Reducto;
|
||||
VertexMistral, $crate::ocr::adapters::VertexMistralAdapter, $crate::ocr::adapters::VertexMistralAdapter, VertexAi;
|
||||
VertexDeepSeek, $crate::ocr::adapters::VertexDeepSeekAdapter, $crate::ocr::adapters::VertexDeepSeekAdapter, VertexAi;
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
pub(crate) use for_each_ocr_adapter;
|
||||
|
|
@ -1,45 +0,0 @@
|
|||
use super::super::OcrAdapter;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::reducto::{self, ReductoLegacyParams, ReductoResponse};
|
||||
use crate::ocr::error::{OcrError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env,
|
||||
guardrail_document, merge_extra_params,
|
||||
};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct ReductoLegacyAdapter;
|
||||
|
||||
impl OcrAdapter for ReductoLegacyAdapter {
|
||||
type ProviderResponse = ReductoResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::Reducto;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params,
|
||||
} = _prepare_ocr_request::<ReductoLegacyParams>(request)?;
|
||||
let headers = super::validate_environment(&request.connection, &credential_env)?;
|
||||
let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?;
|
||||
let (document, headers) = guardrail_document(request, &url, &headers).await?;
|
||||
let document =
|
||||
super::prepare_document(client, document, &request.connection, &headers).await?;
|
||||
let body = reducto::transform_legacy_ocr_request(&request.model, document, ¶ms)?;
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
reducto::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,148 +0,0 @@
|
|||
mod legacy;
|
||||
mod v3;
|
||||
|
||||
use crate::Error;
|
||||
use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX};
|
||||
use crate::ocr::document::InlineDocument;
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{OcrConnection, OcrDocument};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
pub(crate) use legacy::ReductoLegacyAdapter;
|
||||
pub(crate) use v3::ReductoV3Adapter;
|
||||
|
||||
pub(super) fn get_complete_url(api_base: Option<&str>, path: &str) -> Result<String, OcrError> {
|
||||
let base = api_base
|
||||
.map(str::trim)
|
||||
.filter(|base| !base.is_empty())
|
||||
.unwrap_or(REDUCTO_API_BASE);
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&[path]))
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let api_key = connection
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key| !key.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| {
|
||||
env_lookup(REDUCTO_API_KEY_ENV)
|
||||
.map(|key| key.trim().to_string())
|
||||
.filter(|key| !key.is_empty())
|
||||
})
|
||||
.ok_or(Error::MissingReductoApiKey)?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) async fn prepare_document(
|
||||
client: &crate::ocr::OcrClient,
|
||||
document: OcrDocument,
|
||||
connection: &OcrConnection,
|
||||
headers: &[(String, String)],
|
||||
) -> Result<OcrDocument, OcrError> {
|
||||
if document.source().starts_with(REDUCTO_ID_PREFIX) {
|
||||
if document.source()[REDUCTO_ID_PREFIX.len()..]
|
||||
.trim()
|
||||
.is_empty()
|
||||
{
|
||||
return Err(OcrRequestError::RequestField {
|
||||
path: "document file id".into(),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
return Ok(document);
|
||||
}
|
||||
let inline = InlineDocument::parse(document.source())?.ok_or(OcrRequestError::ReductoSource)?;
|
||||
let mime = inline.mime_type().to_string();
|
||||
let bytes = inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
|
||||
let part = reqwest::multipart::Part::bytes(bytes)
|
||||
.file_name("document")
|
||||
.mime_str(&mime)
|
||||
.map_err(|_| OcrRequestError::InvalidDataUri)?;
|
||||
let builder = client
|
||||
.provider_http()
|
||||
.post(get_complete_url(connection.api_base.as_deref(), "upload")?)
|
||||
.multipart(reqwest::multipart::Form::new().part("file", part))
|
||||
.timeout(connection.timeout);
|
||||
let builder = crate::http_utils::with_headers(
|
||||
builder,
|
||||
headers,
|
||||
crate::http_utils::HeaderPolicy::Except(&["content-type", "content-length"]),
|
||||
);
|
||||
let response = crate::http_utils::http_request(builder)
|
||||
.await
|
||||
.map_err(crate::error::TransportError::from)?;
|
||||
let uploaded = crate::ocr::client::read_json_response::<
|
||||
crate::ocr::codecs::reducto::ReductoUploadResponse,
|
||||
>(response, false, connection.max_response_bytes)
|
||||
.await?
|
||||
.data;
|
||||
let file_id = uploaded
|
||||
.file_id
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|id| !id.is_empty());
|
||||
let Some(file_id) = file_id else {
|
||||
return Err(OcrResponseError::ResponseField {
|
||||
path: "file_id".into(),
|
||||
}
|
||||
.into());
|
||||
};
|
||||
Ok(document.with_source(file_id.to_string()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn explicit_key_precedes_environment_key() {
|
||||
let connection = OcrConnection {
|
||||
api_key: Some("passed-key".into()),
|
||||
..Default::default()
|
||||
};
|
||||
let headers = validate_environment(&connection, &|_| Some("env-key".into())).unwrap();
|
||||
assert_eq!(headers[0].1, "Bearer passed-key");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn blank_explicit_key_uses_environment_key() {
|
||||
let connection = OcrConnection {
|
||||
api_key: Some(" ".into()),
|
||||
..Default::default()
|
||||
};
|
||||
let headers = validate_environment(&connection, &|_| Some(" env-key ".into())).unwrap();
|
||||
assert_eq!(headers[0].1, "Bearer env-key");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn existing_authorization_skips_key_lookup() {
|
||||
let connection = OcrConnection {
|
||||
extra_headers: vec![("authorization".into(), "Bearer existing".into())],
|
||||
..Default::default()
|
||||
};
|
||||
assert_eq!(
|
||||
validate_environment(&connection, &|_| None).unwrap(),
|
||||
connection.extra_headers
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,45 +0,0 @@
|
|||
use super::super::OcrAdapter;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::reducto::{self, ReductoResponse, ReductoV3Params};
|
||||
use crate::ocr::error::{OcrError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env,
|
||||
guardrail_document, merge_extra_params,
|
||||
};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct ReductoV3Adapter;
|
||||
|
||||
impl OcrAdapter for ReductoV3Adapter {
|
||||
type ProviderResponse = ReductoResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::Reducto;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params,
|
||||
} = _prepare_ocr_request::<ReductoV3Params>(request)?;
|
||||
let headers = super::validate_environment(&request.connection, &credential_env)?;
|
||||
let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?;
|
||||
let (document, headers) = guardrail_document(request, &url, &headers).await?;
|
||||
let document =
|
||||
super::prepare_document(client, document, &request.connection, &headers).await?;
|
||||
let body = reducto::transform_v3_ocr_request(&request.model, document, ¶ms)?;
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
reducto::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,140 +0,0 @@
|
|||
use super::super::OcrAdapter;
|
||||
use super::validate_destination;
|
||||
use crate::Error;
|
||||
use crate::auth::vertex::{self, VertexConfig};
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::deepseek::{self, DeepSeekOcrParams, DeepSeekOcrResponse};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
|
||||
};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
use crate::url_utils::ApiUrl;
|
||||
const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com";
|
||||
const MODEL_NAMESPACE: &str = "deepseek-ai";
|
||||
const DEFAULT_LOCATION: &str = "us-central1";
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct VertexDeepSeekAdapter;
|
||||
|
||||
impl OcrAdapter for VertexDeepSeekAdapter {
|
||||
type ProviderResponse = DeepSeekOcrResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::VertexAi;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
validate_destination(&request.connection)?;
|
||||
let ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params: _extra_params,
|
||||
} = _prepare_ocr_request::<DeepSeekOcrParams>(request)?;
|
||||
let config = VertexConfig::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
)
|
||||
.map_err(Error::from)?;
|
||||
let authentication = client
|
||||
.vertex_auth()
|
||||
.validate_environment(
|
||||
request.connection.extra_headers.clone(),
|
||||
request.connection.api_key.as_deref(),
|
||||
&config,
|
||||
&credential_env,
|
||||
)
|
||||
.await
|
||||
.map_err(Error::from)?;
|
||||
let location = vertex::get_vertex_ai_location(&config, &credential_env)
|
||||
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
|
||||
let url = get_complete_url(
|
||||
request.connection.api_base.as_deref(),
|
||||
&authentication.project_id,
|
||||
&location,
|
||||
)?;
|
||||
let document = request.document.clone();
|
||||
let body =
|
||||
deepseek::transform_ocr_request(&provider_model(&request.model), document, ¶ms)?;
|
||||
transform_request_body(
|
||||
client,
|
||||
request,
|
||||
&url,
|
||||
&authentication.headers,
|
||||
false,
|
||||
body,
|
||||
|_| Ok(()),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
deepseek::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_model(model: &str) -> String {
|
||||
if model.starts_with(&format!("{MODEL_NAMESPACE}/")) {
|
||||
model.to_string()
|
||||
} else {
|
||||
format!("{MODEL_NAMESPACE}/{model}")
|
||||
}
|
||||
}
|
||||
|
||||
fn get_complete_url(
|
||||
api_base: Option<&str>,
|
||||
project: &str,
|
||||
location: &str,
|
||||
) -> Result<String, OcrError> {
|
||||
let base = api_base
|
||||
.map(str::trim)
|
||||
.filter(|base| !base.is_empty())
|
||||
.unwrap_or(DEFAULT_API_BASE);
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| {
|
||||
url.complete_path(&[
|
||||
"v1",
|
||||
"projects",
|
||||
project,
|
||||
"locations",
|
||||
location,
|
||||
"endpoints",
|
||||
"openapi",
|
||||
"chat",
|
||||
"completions",
|
||||
])
|
||||
})
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{get_complete_url, provider_model};
|
||||
|
||||
#[test]
|
||||
fn adapter_owns_model_namespace_and_endpoint() {
|
||||
assert_eq!(
|
||||
provider_model("deepseek-ocr-maas"),
|
||||
"deepseek-ai/deepseek-ocr-maas"
|
||||
);
|
||||
assert_eq!(
|
||||
provider_model("deepseek-ai/deepseek-ocr-maas"),
|
||||
"deepseek-ai/deepseek-ocr-maas"
|
||||
);
|
||||
assert_eq!(
|
||||
get_complete_url(None, "proj-1", "europe-west4").unwrap(),
|
||||
"https://aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/endpoints/openapi/chat/completions"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,254 +0,0 @@
|
|||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::ocr::document::InlineDocument;
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub(crate) enum OutputFormat {
|
||||
#[default]
|
||||
Markdown,
|
||||
Blocks,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct CohereParams {
|
||||
#[serde(default)]
|
||||
pub output_format: OutputFormat,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, Serialize)]
|
||||
pub(crate) struct CohereRequest {
|
||||
pub model: String,
|
||||
pub document: OcrDocument,
|
||||
pub output_format: OutputFormat,
|
||||
}
|
||||
|
||||
pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), OcrRequestError> {
|
||||
let OcrDocument::ImageUrl { image_url, .. } = document else {
|
||||
return Err(OcrRequestError::CohereImageOnly);
|
||||
};
|
||||
if image_url.is_empty() {
|
||||
return Err(OcrRequestError::CohereImageOnly);
|
||||
}
|
||||
if let Some(inline) = InlineDocument::parse(image_url)? {
|
||||
if !inline.mime_type().type_.eq_ignore_ascii_case("image") {
|
||||
return Err(OcrRequestError::CohereImageOnly);
|
||||
}
|
||||
inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct CohereResponse {
|
||||
#[serde(default)]
|
||||
pages: Vec<CoherePage>,
|
||||
meta: Option<CohereMeta>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CoherePage {
|
||||
index: Option<i64>,
|
||||
markdown: Option<CohereMarkdown>,
|
||||
blocks: Option<Vec<Map<String, Value>>>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CohereMarkdown {
|
||||
#[serde(default)]
|
||||
content: String,
|
||||
images: Option<Vec<Map<String, Value>>>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CohereMeta {
|
||||
billed_units: Option<CohereBilledUnits>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CohereBilledUnits {
|
||||
pages: Option<i64>,
|
||||
}
|
||||
|
||||
pub(crate) fn transform_response(
|
||||
model: &str,
|
||||
response: CohereResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
let pages_processed = response
|
||||
.meta
|
||||
.and_then(|meta| meta.billed_units)
|
||||
.and_then(|units| units.pages)
|
||||
.map(Ok)
|
||||
.unwrap_or_else(|| {
|
||||
i64::try_from(response.pages.len()).map_err(|_| OcrResponseError::NumericRange("pages"))
|
||||
})?;
|
||||
let pages = response
|
||||
.pages
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(position, page)| {
|
||||
let index = page.index.map(Ok).unwrap_or_else(|| {
|
||||
i64::try_from(position).map_err(|_| OcrResponseError::NumericRange("page index"))
|
||||
})?;
|
||||
let (content, images) = page
|
||||
.markdown
|
||||
.map(|markdown| {
|
||||
let images =
|
||||
markdown
|
||||
.images
|
||||
.filter(|images| !images.is_empty())
|
||||
.map(|images| {
|
||||
images
|
||||
.into_iter()
|
||||
.map(|mut image| {
|
||||
if let Some(Value::Object(bbox)) =
|
||||
image.get("bounding_box").cloned()
|
||||
{
|
||||
image.insert("bbox".into(), Value::Object(bbox));
|
||||
}
|
||||
Value::Object(image)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
});
|
||||
(markdown.content, images)
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let mut normalized = json!({"index": index, "markdown": content, "images": images});
|
||||
if let Some(blocks) = page.blocks {
|
||||
normalized["blocks"] = json!(blocks);
|
||||
}
|
||||
Ok(normalized)
|
||||
})
|
||||
.collect::<Result<Vec<_>, OcrResponseError>>()?;
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages,
|
||||
model: model.into(),
|
||||
document_annotation: None,
|
||||
usage_info: Some(json!({"pages_processed": pages_processed})),
|
||||
object: "ocr".into(),
|
||||
extra_fields: Map::new(),
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_request(
|
||||
model: &str,
|
||||
document: OcrDocument,
|
||||
params: CohereParams,
|
||||
) -> Result<CohereRequest, OcrRequestError> {
|
||||
validate_document(&document)?;
|
||||
Ok(CohereRequest {
|
||||
model: model.into(),
|
||||
document,
|
||||
output_format: params.output_format,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn response_normalizes_markdown_images_blocks_and_billed_pages() {
|
||||
let response = serde_json::from_value(json!({
|
||||
"pages": [
|
||||
{
|
||||
"type":"markdown",
|
||||
"index":4,
|
||||
"markdown":{
|
||||
"content":"receipt",
|
||||
"images":[{
|
||||
"id":"image",
|
||||
"bounding_box":{"top_left_x":1,"bottom_right_x":48},
|
||||
"bounding_box_normalized":{"top_left_x":0.04,"bottom_right_x":0.15},
|
||||
"description":"scan",
|
||||
"category":"logo"
|
||||
}]
|
||||
}
|
||||
},
|
||||
{"type":"blocks","blocks":[{"type":"text","text":{"content":"total"}}]}
|
||||
],
|
||||
"meta":{"api_version":{"version":"2"},"billed_units":{"pages":3}}
|
||||
}))
|
||||
.unwrap();
|
||||
let normalized = transform_response("parse-v5.0", response).unwrap();
|
||||
assert_eq!(normalized.pages[0]["index"], 4);
|
||||
assert_eq!(normalized.pages[0]["markdown"], "receipt");
|
||||
assert_eq!(normalized.pages[0]["images"][0]["bbox"]["top_left_x"], 1);
|
||||
assert_eq!(
|
||||
normalized.pages[0]["images"][0]["bounding_box_normalized"]["bottom_right_x"],
|
||||
0.15
|
||||
);
|
||||
assert_eq!(normalized.pages[0]["images"][0]["description"], "scan");
|
||||
assert_eq!(normalized.pages[0]["images"][0]["category"], "logo");
|
||||
assert_eq!(normalized.pages[1]["index"], 1);
|
||||
assert_eq!(normalized.pages[1]["markdown"], "");
|
||||
assert_eq!(normalized.pages[1]["blocks"][0]["text"]["content"], "total");
|
||||
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_defaults_and_invalid_fields() {
|
||||
for value in [
|
||||
json!({}),
|
||||
json!({"meta":null}),
|
||||
json!({"pages":[],"meta":{"billed_units":null}}),
|
||||
] {
|
||||
let normalized =
|
||||
transform_response("parse", serde_json::from_value(value).unwrap()).unwrap();
|
||||
assert!(normalized.pages.is_empty());
|
||||
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 0);
|
||||
}
|
||||
for value in [
|
||||
json!({"pages":null}),
|
||||
json!({"pages":[{"markdown":"text"}]}),
|
||||
json!({"pages":[{"index":"bad"}]}),
|
||||
] {
|
||||
assert!(serde_json::from_value::<CohereResponse>(value).is_err());
|
||||
}
|
||||
let normalized = transform_response(
|
||||
"parse",
|
||||
serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 1);
|
||||
assert!(normalized.pages[0]["images"].is_null());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_requires_image_and_supported_output_format() {
|
||||
for value in [
|
||||
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
|
||||
json!({"type":"image_url","image_url":""}),
|
||||
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
|
||||
] {
|
||||
assert_eq!(
|
||||
validate_document(&serde_json::from_value(value).unwrap()),
|
||||
Err(OcrRequestError::CohereImageOnly)
|
||||
);
|
||||
}
|
||||
assert!(serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err());
|
||||
for format in ["markdown", "blocks"] {
|
||||
assert!(
|
||||
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
|
||||
);
|
||||
}
|
||||
let request = transform_request(
|
||||
"parse-v5.0",
|
||||
serde_json::from_value(json!({
|
||||
"type":"image_url",
|
||||
"image_url":"https://example.com/image.png"
|
||||
}))
|
||||
.unwrap(),
|
||||
serde_json::from_value(json!({})).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(request).unwrap()["output_format"],
|
||||
"markdown"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
mod transformation;
|
||||
mod types;
|
||||
|
||||
pub(crate) use transformation::{transform_ocr_request, transform_ocr_response};
|
||||
pub(crate) use types::{DeepSeekOcrParams, DeepSeekOcrResponse};
|
||||
|
|
@ -1,101 +0,0 @@
|
|||
use serde::de::IntoDeserializer;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::types::*;
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
pub(crate) fn transform_ocr_request(
|
||||
provider_model: &str,
|
||||
document: OcrDocument,
|
||||
params: &DeepSeekOcrParams,
|
||||
) -> Result<DeepSeekOcrRequest, OcrRequestError> {
|
||||
if document.source().is_empty() {
|
||||
return Err(OcrRequestError::MissingDocumentUrl);
|
||||
}
|
||||
let content = OcrDocument::ImageUrl {
|
||||
image_url: document.source().to_string(),
|
||||
extra_fields: serde_json::Map::new(),
|
||||
};
|
||||
Ok(DeepSeekOcrRequest {
|
||||
model: provider_model.to_string(),
|
||||
messages: vec![DeepSeekOcrMessage {
|
||||
role: UserRole::User,
|
||||
content: vec![content],
|
||||
}],
|
||||
params: params.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_ocr_response(
|
||||
model: &str,
|
||||
response: DeepSeekOcrResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
let content = response
|
||||
.choices
|
||||
.into_iter()
|
||||
.next()
|
||||
.and_then(|choice| choice.message.content)
|
||||
.ok_or(OcrResponseError::EmptyContent)?;
|
||||
let decoded = decode_content(content)?;
|
||||
let pages = match decoded.result.pages {
|
||||
Some(pages) if !pages.is_empty() => pages
|
||||
.into_iter()
|
||||
.map(|page| serde_json::to_value(page).expect("DeepSeek page serializes"))
|
||||
.collect(),
|
||||
_ => vec![json!({
|
||||
"index":0,
|
||||
"markdown":decoded.fallback_markdown,
|
||||
"images":null
|
||||
})],
|
||||
};
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages,
|
||||
model: decoded.result.model.unwrap_or_else(|| model.to_string()),
|
||||
document_annotation: decoded.result.document_annotation,
|
||||
usage_info: decoded.result.usage_info.or(response.usage),
|
||||
object: "ocr".into(),
|
||||
extra_fields: decoded.result.extra_fields,
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
struct DecodedContent {
|
||||
result: DeepSeekOcrResult,
|
||||
fallback_markdown: String,
|
||||
}
|
||||
|
||||
fn decode_content(content: DeepSeekContent) -> Result<DecodedContent, OcrResponseError> {
|
||||
let (result, fallback_markdown) = match content {
|
||||
DeepSeekContent::Text(text) if text.is_empty() => {
|
||||
return Err(OcrResponseError::EmptyContent);
|
||||
}
|
||||
DeepSeekContent::Text(text) => (decode_json_content(&text)?, text),
|
||||
DeepSeekContent::Object(object) => {
|
||||
let fallback =
|
||||
serde_json::to_string(&object).map_err(|_| OcrResponseError::ResponseField {
|
||||
path: "choices[0].message.content".into(),
|
||||
})?;
|
||||
(Some(object), fallback)
|
||||
}
|
||||
};
|
||||
Ok(DecodedContent {
|
||||
result: result.unwrap_or_default(),
|
||||
fallback_markdown,
|
||||
})
|
||||
}
|
||||
|
||||
fn decode_json_content(text: &str) -> Result<Option<DeepSeekOcrResult>, OcrResponseError> {
|
||||
if !text.trim_start().starts_with('{') {
|
||||
return Ok(None);
|
||||
}
|
||||
let value = match serde_json::from_str::<Value>(text) {
|
||||
Ok(value) => value,
|
||||
Err(_) => return Ok(None),
|
||||
};
|
||||
serde_path_to_error::deserialize(value.into_deserializer())
|
||||
.map(Some)
|
||||
.map_err(|error| OcrResponseError::ResponseField {
|
||||
path: format!("choices[0].message.content.{}", error.path()),
|
||||
})
|
||||
}
|
||||
|
|
@ -1,95 +0,0 @@
|
|||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrParams {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stream: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub temperature: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub max_tokens: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub top_p: Option<f64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub n: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub stop: Option<StopSequences>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum StopSequences {
|
||||
One(String),
|
||||
Many(Vec<String>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrRequest {
|
||||
pub model: String,
|
||||
pub messages: Vec<DeepSeekOcrMessage>,
|
||||
#[serde(flatten)]
|
||||
pub params: DeepSeekOcrParams,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrMessage {
|
||||
pub role: UserRole,
|
||||
pub content: Vec<crate::ocr::types::OcrDocument>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub(crate) enum UserRole {
|
||||
User,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrResponse {
|
||||
#[serde(default)]
|
||||
pub choices: Vec<DeepSeekChoice>,
|
||||
pub usage: Option<Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct DeepSeekChoice {
|
||||
pub message: DeepSeekResponseMessage,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct DeepSeekResponseMessage {
|
||||
pub content: Option<DeepSeekContent>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum DeepSeekContent {
|
||||
Text(String),
|
||||
Object(DeepSeekOcrResult),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekOcrResult {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub pages: Option<Vec<DeepSeekPage>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub model: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub usage_info: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub document_annotation: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct DeepSeekPage {
|
||||
#[serde(default)]
|
||||
pub index: i64,
|
||||
#[serde(default)]
|
||||
pub markdown: String,
|
||||
pub images: Option<Value>,
|
||||
pub dimensions: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
|
|
@ -1,9 +0,0 @@
|
|||
mod params;
|
||||
mod transformation;
|
||||
mod types;
|
||||
|
||||
pub(crate) use params::{decode_input_params, map_ocr_params};
|
||||
pub(crate) use transformation::{transform_ocr_request, transform_ocr_response};
|
||||
pub(crate) use types::{
|
||||
AzureDocumentIntelligenceOperation, DocumentIntelligenceParams, OperationStatus,
|
||||
};
|
||||
|
|
@ -1,219 +0,0 @@
|
|||
use std::collections::BTreeSet;
|
||||
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::types::{
|
||||
DocumentIntelligenceInputParams, DocumentIntelligenceParams, FeaturesInput, PagesInput,
|
||||
};
|
||||
use crate::ocr::error::OcrRequestError;
|
||||
use crate::ocr::prepare::ParsedProviderParams;
|
||||
|
||||
pub(crate) fn decode_input_params(
|
||||
params: Map<String, Value>,
|
||||
prefix: &str,
|
||||
) -> Result<ParsedProviderParams<DocumentIntelligenceInputParams>, OcrRequestError> {
|
||||
if let Some(Value::Array(pages)) = params.get("pages") {
|
||||
if pages.iter().any(Value::is_boolean) {
|
||||
return Err(OcrRequestError::Pages("boolean page index".into()));
|
||||
}
|
||||
if pages
|
||||
.iter()
|
||||
.any(|page| page.is_number() && page.as_i64().is_none())
|
||||
{
|
||||
return Err(OcrRequestError::Pages("page index is out of range".into()));
|
||||
}
|
||||
if !pages.iter().all(Value::is_i64) && !pages.iter().all(Value::is_string) {
|
||||
return Err(OcrRequestError::Pages("mixed page element types".into()));
|
||||
}
|
||||
}
|
||||
crate::ocr::wire::decode_request_value(Value::Object(params), prefix)
|
||||
}
|
||||
|
||||
pub(crate) fn map_ocr_params(
|
||||
params: DocumentIntelligenceInputParams,
|
||||
) -> Result<DocumentIntelligenceParams, OcrRequestError> {
|
||||
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>, OcrRequestError> {
|
||||
let normalized = match pages {
|
||||
PagesInput::ZeroBasedIndices(indices) => {
|
||||
if indices.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
indices
|
||||
.into_iter()
|
||||
.map(|page| {
|
||||
if page < 0 {
|
||||
return Err(OcrRequestError::Pages("negative page index".into()));
|
||||
}
|
||||
page.checked_add(1)
|
||||
.ok_or_else(|| OcrRequestError::Pages("page index is out of range".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
|
||||
.split(',')
|
||||
.map(str::trim)
|
||||
.collect::<Vec<_>>()
|
||||
.join(","),
|
||||
};
|
||||
if !normalized.split(',').all(valid_page_token) {
|
||||
return Err(OcrRequestError::Pages("invalid native page range".into()));
|
||||
}
|
||||
Ok(Some(normalized))
|
||||
}
|
||||
|
||||
fn valid_page_token(token: &str) -> bool {
|
||||
let mut parts = token.split('-');
|
||||
let start = parts.next().unwrap_or_default();
|
||||
if start.is_empty() || !start.chars().all(|character| character.is_ascii_digit()) {
|
||||
return false;
|
||||
}
|
||||
match parts.next() {
|
||||
None => true,
|
||||
Some(end) => {
|
||||
!end.is_empty()
|
||||
&& end.chars().all(|character| character.is_ascii_digit())
|
||||
&& parts.next().is_none()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_features(features: FeaturesInput) -> Result<Option<String>, OcrRequestError> {
|
||||
let tokens = match features {
|
||||
FeaturesInput::Names(names) => names,
|
||||
FeaturesInput::CommaSeparated(names) => names.split(',').map(str::to_string).collect(),
|
||||
};
|
||||
if tokens.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
let normalized = tokens.iter().map(|token| token.trim()).collect::<Vec<_>>();
|
||||
if !normalized.iter().all(|token| {
|
||||
let Some((first, rest)) = token.as_bytes().split_first() else {
|
||||
return false;
|
||||
};
|
||||
first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric)
|
||||
}) {
|
||||
return Err(OcrRequestError::Features);
|
||||
}
|
||||
Ok(Some(normalized.join(",")))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
fn map(value: Value) -> Result<DocumentIntelligenceParams, OcrRequestError> {
|
||||
let fields = value.as_object().unwrap().clone();
|
||||
map_ocr_params(decode_input_params(fields, "optional_params")?.known)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn input_params_retain_unknown_fields() {
|
||||
let parsed = decode_input_params(
|
||||
json!({
|
||||
"pages": [0],
|
||||
"future_ocr_option": true,
|
||||
"extra_body": {"provider_option": "value"}
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
.clone(),
|
||||
"optional_params",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(
|
||||
parsed.known.pages,
|
||||
Some(PagesInput::ZeroBasedIndices(vec![0]))
|
||||
);
|
||||
assert_eq!(parsed.extra_params["future_ocr_option"], true);
|
||||
assert_eq!(
|
||||
parsed.extra_params["extra_body"],
|
||||
json!({"provider_option": "value"})
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::to_value(map_ocr_params(parsed.known).unwrap()).unwrap(),
|
||||
json!({"pages": "1", "features": null})
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(json!([0, 1, 2]), Some("1,2,3"))]
|
||||
#[case(json!([2, 0, 0, 1]), Some("1,2,3"))]
|
||||
#[case(json!([]), None)]
|
||||
#[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"))]
|
||||
fn page_mapping_matches_python(#[case] input: Value, #[case] expected: Option<&str>) {
|
||||
assert_eq!(
|
||||
map(json!({"pages": input})).unwrap().pages.as_deref(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(json!("a,b"))]
|
||||
#[case(json!([-1]))]
|
||||
#[case(json!([true, false]))]
|
||||
#[case(json!([1, "2"]))]
|
||||
#[case(json!(5))]
|
||||
fn invalid_page_mapping_matches_python(#[case] input: Value) {
|
||||
assert!(map(json!({"pages": input})).is_err());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(json!(["keyValuePairs"]), "keyValuePairs")]
|
||||
#[case(json!(["keyValuePairs", "languages"]), "keyValuePairs,languages")]
|
||||
#[case(json!("keyValuePairs"), "keyValuePairs")]
|
||||
#[case(json!("keyValuePairs,languages"), "keyValuePairs,languages")]
|
||||
#[case(json!("keyValuePairs, languages"), "keyValuePairs,languages")]
|
||||
fn feature_mapping_matches_python(#[case] input: Value, #[case] expected: &str) {
|
||||
assert_eq!(
|
||||
map(json!({"features": input})).unwrap().features.as_deref(),
|
||||
Some(expected)
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(json!("keyValuePairs&pages=9"))]
|
||||
#[case(json!("key value pairs"))]
|
||||
#[case(json!(""))]
|
||||
#[case(json!([1, 2]))]
|
||||
#[case(json!([["keyValuePairs"]]))]
|
||||
#[case(json!({"feature":"keyValuePairs"}))]
|
||||
#[case(json!(5))]
|
||||
fn invalid_feature_mapping_matches_python(#[case] input: Value) {
|
||||
assert!(map(json!({"features": input})).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_feature_list_is_omitted() {
|
||||
assert_eq!(map(json!({"features": []})).unwrap().features, None);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,107 +0,0 @@
|
|||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use super::types::*;
|
||||
use crate::constants::{AZURE_DI_DEFAULT_DPI, AZURE_DI_DEFAULT_HEIGHT, AZURE_DI_DEFAULT_WIDTH};
|
||||
use crate::ocr::document::InlineDocument;
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
pub(crate) fn transform_ocr_request(
|
||||
document: OcrDocument,
|
||||
) -> Result<DocumentIntelligenceRequest, OcrRequestError> {
|
||||
let source = document.source();
|
||||
if source.is_empty() {
|
||||
return Err(OcrRequestError::MissingDocumentUrl);
|
||||
}
|
||||
Ok(if let Some(document) = InlineDocument::parse(source)? {
|
||||
DocumentIntelligenceRequest::Base64Source(
|
||||
STANDARD.encode(document.decode(crate::constants::OCR_INLINE_MAX_BYTES)?),
|
||||
)
|
||||
} else {
|
||||
DocumentIntelligenceRequest::UrlSource(source.to_string())
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_ocr_response(
|
||||
model: &str,
|
||||
response: AzureDocumentIntelligenceOperation,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
if response.status != Some(OperationStatus::Succeeded) {
|
||||
return Err(OcrResponseError::OperationStatus(
|
||||
response
|
||||
.status
|
||||
.map(|status| status.to_string())
|
||||
.unwrap_or_else(|| "None".into()),
|
||||
));
|
||||
}
|
||||
let result = response.analyze_result.unwrap_or_default();
|
||||
let pages = result
|
||||
.pages
|
||||
.into_iter()
|
||||
.map(normalize_page)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
let pages_processed = pages.len();
|
||||
let mut extra_fields = Map::new();
|
||||
extra_fields.insert("content".into(), option_value(result.content));
|
||||
extra_fields.insert("tables".into(), option_value(result.tables));
|
||||
extra_fields.insert("keyValuePairs".into(), option_value(result.key_value_pairs));
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages,
|
||||
model: model.into(),
|
||||
document_annotation: None,
|
||||
usage_info: Some(json!({"pages_processed":pages_processed})),
|
||||
object: "ocr".into(),
|
||||
extra_fields,
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn normalize_page(page: AzureDocumentIntelligencePage) -> Result<Value, OcrResponseError> {
|
||||
let index = page
|
||||
.page_number
|
||||
.unwrap_or(1)
|
||||
.checked_sub(1)
|
||||
.ok_or(OcrResponseError::NumericRange("page.pageNumber"))?;
|
||||
let scale = if page.unit.as_deref().unwrap_or("inch") == "inch" {
|
||||
AZURE_DI_DEFAULT_DPI as f64
|
||||
} else {
|
||||
1.0
|
||||
};
|
||||
let width = pixel_dimension(
|
||||
page.width.unwrap_or(AZURE_DI_DEFAULT_WIDTH),
|
||||
scale,
|
||||
"page.width",
|
||||
)?;
|
||||
let height = pixel_dimension(
|
||||
page.height.unwrap_or(AZURE_DI_DEFAULT_HEIGHT),
|
||||
scale,
|
||||
"page.height",
|
||||
)?;
|
||||
let markdown = page
|
||||
.lines
|
||||
.iter()
|
||||
.map(|line| line.content.as_deref().unwrap_or_default())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n");
|
||||
Ok(json!({
|
||||
"index":index,
|
||||
"markdown":markdown,
|
||||
"images":null,
|
||||
"dimensions":{"width":width,"height":height,"dpi":AZURE_DI_DEFAULT_DPI}
|
||||
}))
|
||||
}
|
||||
|
||||
fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result<i64, OcrResponseError> {
|
||||
let value = value * scale;
|
||||
if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 {
|
||||
return Err(OcrResponseError::NumericRange(field));
|
||||
}
|
||||
Ok(value.trunc() as i64)
|
||||
}
|
||||
|
||||
fn option_value<T: serde::Serialize>(value: Option<T>) -> Value {
|
||||
value
|
||||
.and_then(|value| serde_json::to_value(value).ok())
|
||||
.unwrap_or(Value::Null)
|
||||
}
|
||||
|
|
@ -1,138 +0,0 @@
|
|||
use serde::{Deserialize, Deserializer, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum PagesInput {
|
||||
ZeroBasedIndices(Vec<i64>),
|
||||
NativeTokens(Vec<String>),
|
||||
NativeRange(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum FeaturesInput {
|
||||
Names(Vec<String>),
|
||||
CommaSeparated(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub(crate) struct DocumentIntelligenceInputParams {
|
||||
pub pages: Option<PagesInput>,
|
||||
pub features: Option<FeaturesInput>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize)]
|
||||
pub(crate) struct DocumentIntelligenceParams {
|
||||
pub pages: Option<String>,
|
||||
pub features: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) enum DocumentIntelligenceRequest {
|
||||
#[serde(rename = "urlSource")]
|
||||
UrlSource(String),
|
||||
#[serde(rename = "base64Source")]
|
||||
Base64Source(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub(crate) enum OperationStatus {
|
||||
Succeeded,
|
||||
Running,
|
||||
NotStarted,
|
||||
Failed,
|
||||
Unknown(String),
|
||||
}
|
||||
|
||||
impl<'de> Deserialize<'de> for OperationStatus {
|
||||
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
|
||||
Ok(match String::deserialize(deserializer)?.as_str() {
|
||||
"succeeded" => Self::Succeeded,
|
||||
"running" => Self::Running,
|
||||
"notStarted" => Self::NotStarted,
|
||||
"failed" => Self::Failed,
|
||||
value => Self::Unknown(value.to_string()),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for OperationStatus {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter.write_str(match self {
|
||||
Self::Succeeded => "succeeded",
|
||||
Self::Running => "running",
|
||||
Self::NotStarted => "notStarted",
|
||||
Self::Failed => "failed",
|
||||
Self::Unknown(value) => value,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct AzureDocumentIntelligenceOperation {
|
||||
pub status: Option<OperationStatus>,
|
||||
#[serde(rename = "analyzeResult")]
|
||||
pub analyze_result: Option<AzureDocumentIntelligenceAnalyzeResult>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct AzureDocumentIntelligenceAnalyzeResult {
|
||||
pub content: Option<String>,
|
||||
#[serde(default)]
|
||||
pub pages: Vec<AzureDocumentIntelligencePage>,
|
||||
pub tables: Option<Vec<Map<String, Value>>>,
|
||||
#[serde(rename = "keyValuePairs")]
|
||||
pub key_value_pairs: Option<Vec<Map<String, Value>>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct AzureDocumentIntelligencePage {
|
||||
#[serde(rename = "pageNumber", default, deserialize_with = "optional_i64")]
|
||||
pub page_number: Option<i64>,
|
||||
#[serde(default, deserialize_with = "optional_f64")]
|
||||
pub width: Option<f64>,
|
||||
#[serde(default, deserialize_with = "optional_f64")]
|
||||
pub height: Option<f64>,
|
||||
pub unit: Option<String>,
|
||||
#[serde(default)]
|
||||
pub lines: Vec<AzureDocumentIntelligenceLine>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct AzureDocumentIntelligenceLine {
|
||||
pub content: Option<String>,
|
||||
}
|
||||
|
||||
fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<i64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_i64()
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected an integer")),
|
||||
Some(Value::String(value)) => value
|
||||
.parse::<i64>()
|
||||
.map(Some)
|
||||
.map_err(|_| serde::de::Error::custom("expected an integer")),
|
||||
Some(_) => Err(serde::de::Error::custom("expected an integer")),
|
||||
}
|
||||
}
|
||||
|
||||
fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<f64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_f64()
|
||||
.filter(|value| value.is_finite())
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a finite number")),
|
||||
Some(Value::String(value)) => value
|
||||
.parse::<f64>()
|
||||
.ok()
|
||||
.filter(|value| value.is_finite())
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a finite number")),
|
||||
Some(_) => Err(serde::de::Error::custom("expected a number")),
|
||||
}
|
||||
}
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
mod transformation;
|
||||
mod types;
|
||||
|
||||
pub(crate) use transformation::{transform_ocr_request, transform_ocr_response};
|
||||
pub(crate) use types::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse};
|
||||
|
|
@ -1,60 +0,0 @@
|
|||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::ocr::types::OcrDocument;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(untagged)]
|
||||
pub(crate) enum MistralOcrPages {
|
||||
Range(String),
|
||||
Indices(Vec<i64>),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub(crate) struct MistralOcrParams {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub pages: Option<MistralOcrPages>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub include_image_base64: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub image_limit: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub image_min_size: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub bbox_annotation_format: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub document_annotation_format: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub document_annotation_prompt: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extract_header: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extract_footer: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub table_format: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub confidence_scores_granularity: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub include_blocks: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct MistralOcrRequest {
|
||||
pub model: String,
|
||||
pub document: OcrDocument,
|
||||
#[serde(flatten)]
|
||||
pub params: MistralOcrParams,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct MistralOcrResponse {
|
||||
#[serde(default)]
|
||||
pub pages: Vec<Value>,
|
||||
pub model: Option<String>,
|
||||
pub document_annotation: Option<Value>,
|
||||
pub usage_info: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
pub(crate) mod cohere;
|
||||
pub(crate) mod deepseek;
|
||||
pub(crate) mod document_intelligence;
|
||||
pub(crate) mod mistral;
|
||||
pub(crate) mod reducto;
|
||||
|
|
@ -1,9 +0,0 @@
|
|||
mod transformation;
|
||||
mod types;
|
||||
|
||||
pub(crate) use transformation::{
|
||||
transform_legacy_ocr_request, transform_ocr_response, transform_v3_ocr_request,
|
||||
};
|
||||
pub(crate) use types::{
|
||||
ReductoLegacyParams, ReductoResponse, ReductoUploadResponse, ReductoV3Params,
|
||||
};
|
||||
|
|
@ -1,103 +0,0 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::types::*;
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
pub(crate) fn transform_v3_ocr_request(
|
||||
_model: &str,
|
||||
document: OcrDocument,
|
||||
params: &ReductoV3Params,
|
||||
) -> Result<ReductoV3Request, OcrRequestError> {
|
||||
Ok(ReductoV3Request {
|
||||
input: document.source().to_string(),
|
||||
params: params.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_legacy_ocr_request(
|
||||
_model: &str,
|
||||
document: OcrDocument,
|
||||
params: &ReductoLegacyParams,
|
||||
) -> Result<ReductoLegacyRequest, OcrRequestError> {
|
||||
Ok(ReductoLegacyRequest {
|
||||
document_url: document.source().to_string(),
|
||||
options: params.enhance.as_ref().map(|_| params.clone()),
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_ocr_response(
|
||||
model: &str,
|
||||
response: ReductoResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
let result = match response.result {
|
||||
Some(result) => result.unwrap_or_default(),
|
||||
None => ReductoResult {
|
||||
chunks: response.chunks,
|
||||
},
|
||||
};
|
||||
let usage = response.usage.unwrap_or_default();
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages: build_pages(result.chunks.unwrap_or_default()),
|
||||
model: model.to_string(),
|
||||
document_annotation: None,
|
||||
usage_info: Some(json!({
|
||||
"pages_processed": usage.num_pages,
|
||||
"credits": usage.credits,
|
||||
})),
|
||||
object: "ocr".to_string(),
|
||||
extra_fields: serde_json::Map::new(),
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn build_pages(chunks: Vec<ReductoChunk>) -> Vec<Value> {
|
||||
let blocks_by_page = chunks
|
||||
.iter()
|
||||
.flat_map(|chunk| chunk.blocks.iter().flatten())
|
||||
.filter_map(|block| block.bbox.as_ref()?.page.map(|page| (page, block)))
|
||||
.fold(
|
||||
BTreeMap::<i64, Vec<&ReductoBlock>>::new(),
|
||||
|mut pages, (page, block)| {
|
||||
pages.entry(page).or_default().push(block);
|
||||
pages
|
||||
},
|
||||
);
|
||||
if blocks_by_page.is_empty() {
|
||||
let markdown = join_content(chunks.iter().map(|chunk| chunk.content.as_deref()));
|
||||
return if markdown.is_empty() {
|
||||
Vec::new()
|
||||
} else {
|
||||
vec![page(0, markdown, None)]
|
||||
};
|
||||
}
|
||||
blocks_by_page
|
||||
.into_iter()
|
||||
.map(|(index, blocks)| {
|
||||
let markdown = join_content(blocks.iter().map(|block| block.content.as_deref()));
|
||||
page(
|
||||
index.saturating_sub(1).max(0),
|
||||
markdown,
|
||||
Some(json!(blocks)),
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn join_content<'a>(content: impl Iterator<Item = Option<&'a str>>) -> String {
|
||||
content
|
||||
.flatten()
|
||||
.filter(|text| !text.is_empty())
|
||||
.collect::<Vec<_>>()
|
||||
.join("\n\n")
|
||||
}
|
||||
|
||||
fn page(index: i64, markdown: String, blocks: Option<Value>) -> Value {
|
||||
let mut result = json!({"index":index,"markdown":markdown,"images":null});
|
||||
if let (Value::Object(fields), Some(blocks)) = (&mut result, blocks) {
|
||||
fields.insert("blocks".into(), blocks);
|
||||
}
|
||||
result
|
||||
}
|
||||
|
|
@ -1,128 +0,0 @@
|
|||
use serde::{Deserialize, Deserializer, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoV3Params {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub formatting: Option<Map<String, Value>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub retrieval: Option<Map<String, Value>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub settings: Option<Map<String, Value>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoLegacyParams {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub enhance: Option<Map<String, Value>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoV3Request {
|
||||
pub input: String,
|
||||
#[serde(flatten)]
|
||||
pub params: ReductoV3Params,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoLegacyRequest {
|
||||
pub document_url: String,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub options: Option<ReductoLegacyParams>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct ReductoUploadResponse {
|
||||
pub file_id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct ReductoResponse {
|
||||
#[serde(default, deserialize_with = "present_nullable")]
|
||||
pub result: Option<Option<ReductoResult>>,
|
||||
pub usage: Option<ReductoUsage>,
|
||||
#[serde(default)]
|
||||
pub chunks: Option<Vec<ReductoChunk>>,
|
||||
}
|
||||
|
||||
fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
|
||||
deserializer: D,
|
||||
) -> Result<Option<Option<T>>, D::Error> {
|
||||
Option::<T>::deserialize(deserializer).map(Some)
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct ReductoResult {
|
||||
pub chunks: Option<Vec<ReductoChunk>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct ReductoUsage {
|
||||
#[serde(default, deserialize_with = "optional_i64")]
|
||||
pub num_pages: Option<i64>,
|
||||
#[serde(default, deserialize_with = "optional_f64")]
|
||||
pub credits: Option<f64>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub(crate) struct ReductoChunk {
|
||||
pub content: Option<String>,
|
||||
pub blocks: Option<Vec<ReductoBlock>>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoBlock {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub content: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub bbox: Option<ReductoBoundingBox>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct ReductoBoundingBox {
|
||||
#[serde(default, deserialize_with = "optional_i64")]
|
||||
pub page: Option<i64>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
|
||||
fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<i64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_i64()
|
||||
.or_else(|| number.as_f64().and_then(checked_truncated_i64))
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected an integer")),
|
||||
Some(Value::String(value)) => value
|
||||
.trim()
|
||||
.parse::<i64>()
|
||||
.map(Some)
|
||||
.map_err(|_| serde::de::Error::custom("expected an integer")),
|
||||
Some(Value::Bool(value)) => Ok(Some(i64::from(value))),
|
||||
Some(_) => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<f64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_f64()
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a number")),
|
||||
Some(Value::String(value)) => value
|
||||
.trim()
|
||||
.parse::<f64>()
|
||||
.map(Some)
|
||||
.map_err(|_| serde::de::Error::custom("expected a number")),
|
||||
Some(_) => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
fn checked_truncated_i64(value: f64) -> Option<i64> {
|
||||
(value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64)
|
||||
.then(|| value.trunc() as i64)
|
||||
}
|
||||
|
|
@ -1,11 +1,27 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use super::OcrClient;
|
||||
use super::adapters::OcrAdapter;
|
||||
use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest};
|
||||
use super::registry::OcrAdapterKind;
|
||||
use super::provider_config::OcrConfigKind;
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
use super::wire::DecodedOcrResponse;
|
||||
use crate::Error;
|
||||
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
|
||||
use std::sync::Arc;
|
||||
use crate::llms::azure_ai::ocr::cohere_parse_transformation::AzureAICohereParseConfig;
|
||||
use crate::llms::azure_ai::ocr::document_intelligence::AzureDocumentIntelligenceOperation;
|
||||
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::cohere::ocr::CohereResponse;
|
||||
use crate::llms::cohere::ocr::transformation::CohereParseConfig;
|
||||
use crate::llms::mistral::ocr::MistralOcrResponse;
|
||||
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
|
||||
use crate::llms::reducto::ocr::ReductoResponse;
|
||||
use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, ReductoParseV3Config};
|
||||
use crate::llms::vertex_ai::ocr::deepseek_transformation::{
|
||||
DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig,
|
||||
};
|
||||
use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig;
|
||||
|
||||
pub(crate) async fn perform_ocr_request(
|
||||
client: &OcrClient,
|
||||
|
|
@ -15,7 +31,7 @@ pub(crate) async fn perform_ocr_request(
|
|||
let context = CallLifecycleContext::new(
|
||||
"ocr",
|
||||
request.model.clone(),
|
||||
request.adapter.provider().as_str(),
|
||||
request.config.provider().as_str(),
|
||||
request
|
||||
.litellm_call_id
|
||||
.clone()
|
||||
|
|
@ -47,14 +63,45 @@ impl PreparedOcrCall {
|
|||
client: OcrClient,
|
||||
request: LiteLLMOcrRequest,
|
||||
) -> Result<Self, Error> {
|
||||
macro_rules! prepare_adapter {
|
||||
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
|
||||
match request.adapter {
|
||||
$( OcrAdapterKind::$variant => $instance.prepare_request(&request, &client).await?, )+
|
||||
}
|
||||
};
|
||||
}
|
||||
let http = super::adapters::for_each_ocr_adapter!(prepare_adapter);
|
||||
let http = match request.config {
|
||||
OcrConfigKind::Cohere => CohereParseConfig.prepare_request(&request, &client).await?,
|
||||
OcrConfigKind::Mistral => MistralOCRConfig.prepare_request(&request, &client).await?,
|
||||
OcrConfigKind::AzureAi => {
|
||||
AzureAIOCRConfig::default()
|
||||
.prepare_request(&request, &client)
|
||||
.await?
|
||||
}
|
||||
OcrConfigKind::AzureCohere => {
|
||||
AzureAICohereParseConfig::default()
|
||||
.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::default()
|
||||
.prepare_request(&request, &client)
|
||||
.await?
|
||||
}
|
||||
OcrConfigKind::VertexDeepSeek => {
|
||||
VertexAIDeepSeekOCRConfig
|
||||
.prepare_request(&request, &client)
|
||||
.await?
|
||||
}
|
||||
};
|
||||
Ok(Self {
|
||||
client,
|
||||
request,
|
||||
|
|
@ -71,20 +118,57 @@ impl PreparedOcrCall {
|
|||
))
|
||||
.await
|
||||
.map_err(super::client::transport_error)?;
|
||||
macro_rules! read_adapter {
|
||||
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
|
||||
match self.request.adapter {
|
||||
$( OcrAdapterKind::$variant => {
|
||||
let decoded = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?;
|
||||
Ok(OcrProviderResponse {
|
||||
request: self.request,
|
||||
data: OcrProviderData::$variant(decoded),
|
||||
})
|
||||
}, )+
|
||||
}
|
||||
};
|
||||
}
|
||||
super::adapters::for_each_ocr_adapter!(read_adapter)
|
||||
let data = match self.request.config {
|
||||
OcrConfigKind::Cohere => OcrProviderData::Cohere(
|
||||
CohereParseConfig
|
||||
.read_response(&self.client, response, &url, &headers, &self.request)
|
||||
.await?,
|
||||
),
|
||||
OcrConfigKind::Mistral => OcrProviderData::Mistral(
|
||||
MistralOCRConfig
|
||||
.read_response(&self.client, response, &url, &headers, &self.request)
|
||||
.await?,
|
||||
),
|
||||
OcrConfigKind::AzureAi => OcrProviderData::AzureAi(
|
||||
AzureAIOCRConfig::default()
|
||||
.read_response(&self.client, response, &url, &headers, &self.request)
|
||||
.await?,
|
||||
),
|
||||
OcrConfigKind::AzureCohere => OcrProviderData::AzureCohere(
|
||||
AzureAICohereParseConfig::default()
|
||||
.read_response(&self.client, response, &url, &headers, &self.request)
|
||||
.await?,
|
||||
),
|
||||
OcrConfigKind::AzureDocumentIntelligence => OcrProviderData::AzureDocumentIntelligence(
|
||||
AzureDocumentIntelligenceOCRConfig
|
||||
.read_response(&self.client, response, &url, &headers, &self.request)
|
||||
.await?,
|
||||
),
|
||||
OcrConfigKind::ReductoLegacy => OcrProviderData::ReductoLegacy(
|
||||
ReductoParseLegacyConfig
|
||||
.read_response(&self.client, response, &url, &headers, &self.request)
|
||||
.await?,
|
||||
),
|
||||
OcrConfigKind::ReductoV3 => OcrProviderData::ReductoV3(
|
||||
ReductoParseV3Config
|
||||
.read_response(&self.client, response, &url, &headers, &self.request)
|
||||
.await?,
|
||||
),
|
||||
OcrConfigKind::VertexAi => OcrProviderData::VertexAi(
|
||||
VertexAIOCRConfig::default()
|
||||
.read_response(&self.client, response, &url, &headers, &self.request)
|
||||
.await?,
|
||||
),
|
||||
OcrConfigKind::VertexDeepSeek => OcrProviderData::VertexDeepSeek(
|
||||
VertexAIDeepSeekOCRConfig
|
||||
.read_response(&self.client, response, &url, &headers, &self.request)
|
||||
.await?,
|
||||
),
|
||||
};
|
||||
Ok(OcrProviderResponse {
|
||||
request: self.request,
|
||||
data,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -104,23 +188,16 @@ fn request_headers(request: &reqwest::Request) -> Result<Vec<(String, String)>,
|
|||
.collect()
|
||||
}
|
||||
|
||||
macro_rules! provider_data {
|
||||
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
|
||||
enum OcrProviderData {
|
||||
$( $variant(super::wire::DecodedOcrResponse<<$adapter as OcrAdapter>::ProviderResponse>), )+
|
||||
}
|
||||
|
||||
impl OcrProviderResponse {
|
||||
pub(crate) fn normalize(self) -> Result<LiteLLMOcrResponse, Error> {
|
||||
match self.data {
|
||||
$( OcrProviderData::$variant(decoded) => {
|
||||
let response = $instance.transform_ocr_response(&self.request, decoded.data)?;
|
||||
Ok(LiteLLMOcrResponse { provider_native_response: decoded.native, ..response })
|
||||
}, )+
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
enum OcrProviderData {
|
||||
Cohere(DecodedOcrResponse<CohereResponse>),
|
||||
Mistral(DecodedOcrResponse<MistralOcrResponse>),
|
||||
AzureAi(DecodedOcrResponse<MistralOcrResponse>),
|
||||
AzureCohere(DecodedOcrResponse<CohereResponse>),
|
||||
AzureDocumentIntelligence(DecodedOcrResponse<AzureDocumentIntelligenceOperation>),
|
||||
ReductoLegacy(DecodedOcrResponse<ReductoResponse>),
|
||||
ReductoV3(DecodedOcrResponse<ReductoResponse>),
|
||||
VertexAi(DecodedOcrResponse<MistralOcrResponse>),
|
||||
VertexDeepSeek(DecodedOcrResponse<DeepSeekOcrResponse>),
|
||||
}
|
||||
|
||||
pub(crate) struct OcrProviderResponse {
|
||||
|
|
@ -128,6 +205,55 @@ pub(crate) struct OcrProviderResponse {
|
|||
data: OcrProviderData,
|
||||
}
|
||||
|
||||
impl OcrProviderResponse {
|
||||
pub(crate) fn normalize(self) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let (response, native) = match self.data {
|
||||
OcrProviderData::Cohere(decoded) => (
|
||||
CohereParseConfig.transform_ocr_response(&self.request, decoded.data)?,
|
||||
decoded.native,
|
||||
),
|
||||
OcrProviderData::Mistral(decoded) => (
|
||||
MistralOCRConfig.transform_ocr_response(&self.request, decoded.data)?,
|
||||
decoded.native,
|
||||
),
|
||||
OcrProviderData::AzureAi(decoded) => (
|
||||
AzureAIOCRConfig::default().transform_ocr_response(&self.request, decoded.data)?,
|
||||
decoded.native,
|
||||
),
|
||||
OcrProviderData::AzureCohere(decoded) => (
|
||||
AzureAICohereParseConfig::default()
|
||||
.transform_ocr_response(&self.request, decoded.data)?,
|
||||
decoded.native,
|
||||
),
|
||||
OcrProviderData::AzureDocumentIntelligence(decoded) => (
|
||||
AzureDocumentIntelligenceOCRConfig
|
||||
.transform_ocr_response(&self.request, decoded.data)?,
|
||||
decoded.native,
|
||||
),
|
||||
OcrProviderData::ReductoLegacy(decoded) => (
|
||||
ReductoParseLegacyConfig.transform_ocr_response(&self.request, decoded.data)?,
|
||||
decoded.native,
|
||||
),
|
||||
OcrProviderData::ReductoV3(decoded) => (
|
||||
ReductoParseV3Config.transform_ocr_response(&self.request, decoded.data)?,
|
||||
decoded.native,
|
||||
),
|
||||
OcrProviderData::VertexAi(decoded) => (
|
||||
VertexAIOCRConfig::default().transform_ocr_response(&self.request, decoded.data)?,
|
||||
decoded.native,
|
||||
),
|
||||
OcrProviderData::VertexDeepSeek(decoded) => (
|
||||
VertexAIDeepSeekOCRConfig.transform_ocr_response(&self.request, decoded.data)?,
|
||||
decoded.native,
|
||||
),
|
||||
};
|
||||
Ok(LiteLLMOcrResponse {
|
||||
provider_native_response: native,
|
||||
..response
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn post_call(hooks: &Arc<dyn OcrHooks>, bytes: &[u8]) -> Result<(), Error> {
|
||||
let original_response = serde_json::Value::String(String::from_utf8_lossy(bytes).into_owned());
|
||||
hooks
|
||||
|
|
@ -135,5 +261,3 @@ pub(crate) async fn post_call(hooks: &Arc<dyn OcrHooks>, bytes: &[u8]) -> Result
|
|||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
super::adapters::for_each_ocr_adapter!(provider_data);
|
||||
|
|
|
|||
|
|
@ -1,13 +1,11 @@
|
|||
mod adapters;
|
||||
pub mod client;
|
||||
mod codecs;
|
||||
mod document;
|
||||
pub(crate) mod document;
|
||||
pub mod error;
|
||||
mod handler;
|
||||
pub(crate) mod handler;
|
||||
pub mod hooks;
|
||||
mod lifecycle;
|
||||
mod prepare;
|
||||
mod registry;
|
||||
pub(crate) mod prepare;
|
||||
mod provider_config;
|
||||
pub mod types;
|
||||
pub mod wire;
|
||||
|
||||
|
|
|
|||
|
|
@ -83,7 +83,7 @@ where
|
|||
.hooks
|
||||
.during_call(OcrDuringCallRequest {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: request.adapter.provider().as_str().into(),
|
||||
custom_llm_provider: request.config.provider().as_str().into(),
|
||||
url: url.into(),
|
||||
headers: headers.to_vec(),
|
||||
body,
|
||||
|
|
@ -135,7 +135,7 @@ pub(crate) async fn guardrail_document(
|
|||
.hooks
|
||||
.during_call(OcrDuringCallRequest {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: request.adapter.provider().as_str().into(),
|
||||
custom_llm_provider: request.config.provider().as_str().into(),
|
||||
url: url.into(),
|
||||
headers: headers.to_vec(),
|
||||
body: serde_json::to_value(&request.document).map_err(|_| {
|
||||
|
|
|
|||
130
litellm-rust/crates/core/src/ocr/provider_config.rs
Normal file
130
litellm-rust/crates/core/src/ocr/provider_config.rs
Normal file
|
|
@ -0,0 +1,130 @@
|
|||
use crate::Error;
|
||||
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum OcrConfigKind {
|
||||
Cohere,
|
||||
Mistral,
|
||||
AzureAi,
|
||||
AzureCohere,
|
||||
AzureDocumentIntelligence,
|
||||
ReductoLegacy,
|
||||
ReductoV3,
|
||||
VertexAi,
|
||||
VertexDeepSeek,
|
||||
}
|
||||
|
||||
impl OcrConfigKind {
|
||||
pub(crate) const fn provider(self) -> OcrProvider {
|
||||
match self {
|
||||
Self::Cohere => OcrProvider::Cohere,
|
||||
Self::Mistral => OcrProvider::Mistral,
|
||||
Self::AzureAi | Self::AzureCohere | Self::AzureDocumentIntelligence => {
|
||||
OcrProvider::AzureAi
|
||||
}
|
||||
Self::ReductoLegacy | Self::ReductoV3 => OcrProvider::Reducto,
|
||||
Self::VertexAi | Self::VertexDeepSeek => OcrProvider::VertexAi,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum OcrProvider {
|
||||
Cohere,
|
||||
Mistral,
|
||||
AzureAi,
|
||||
Reducto,
|
||||
VertexAi,
|
||||
}
|
||||
|
||||
impl OcrProvider {
|
||||
pub(crate) const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Cohere => "cohere",
|
||||
Self::Mistral => "mistral",
|
||||
Self::AzureAi => "azure_ai",
|
||||
Self::Reducto => "reducto",
|
||||
Self::VertexAi => "vertex_ai",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_provider_config(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<(String, OcrConfigKind), Error> {
|
||||
let provider =
|
||||
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: OcrProvider::Mistral.as_str(),
|
||||
});
|
||||
let config = match provider.custom_llm_provider {
|
||||
"cohere" => OcrConfigKind::Cohere,
|
||||
"mistral" => OcrConfigKind::Mistral,
|
||||
"azure_ai" if is_document_intelligence_model(provider.model) => {
|
||||
OcrConfigKind::AzureDocumentIntelligence
|
||||
}
|
||||
"azure_ai"
|
||||
if provider.model.to_ascii_lowercase().contains("cohere")
|
||||
&& provider.model.to_ascii_lowercase().contains("parse") =>
|
||||
{
|
||||
OcrConfigKind::AzureCohere
|
||||
}
|
||||
"azure_ai" => OcrConfigKind::AzureAi,
|
||||
"reducto" if provider.model.eq_ignore_ascii_case("parse-legacy") => {
|
||||
OcrConfigKind::ReductoLegacy
|
||||
}
|
||||
"reducto" => OcrConfigKind::ReductoV3,
|
||||
"vertex_ai" if provider.model.to_ascii_lowercase().contains("deepseek") => {
|
||||
OcrConfigKind::VertexDeepSeek
|
||||
}
|
||||
"vertex_ai" => OcrConfigKind::VertexAi,
|
||||
value => return Err(Error::InvalidProvider(value.to_string())),
|
||||
};
|
||||
Ok((provider.model.to_string(), config))
|
||||
}
|
||||
|
||||
fn is_document_intelligence_model(model: &str) -> bool {
|
||||
let model = model.to_ascii_lowercase();
|
||||
model.contains("doc-intelligence") || model.contains("documentintelligence")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[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() {
|
||||
assert_eq!(
|
||||
resolve_provider_config("reducto/parse-legacy", None)
|
||||
.unwrap()
|
||||
.1,
|
||||
OcrConfigKind::ReductoLegacy
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_provider_config("reducto/future-parse-model", None)
|
||||
.unwrap()
|
||||
.1,
|
||||
OcrConfigKind::ReductoV3
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_provider_config("azure_ai/doc-intelligence/prebuilt-layout", None)
|
||||
.unwrap()
|
||||
.1,
|
||||
OcrConfigKind::AzureDocumentIntelligence
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,132 +0,0 @@
|
|||
use super::adapters::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
macro_rules! define_adapter_types {
|
||||
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum OcrAdapterKind {
|
||||
$( $variant, )+
|
||||
}
|
||||
|
||||
impl OcrAdapterKind {
|
||||
pub(crate) const fn provider(self) -> OcrProvider {
|
||||
match self {
|
||||
$( Self::$variant => <$adapter>::PROVIDER, )+
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
super::adapters::for_each_ocr_adapter!(define_adapter_types);
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum OcrProvider {
|
||||
Cohere,
|
||||
Mistral,
|
||||
AzureAi,
|
||||
Reducto,
|
||||
VertexAi,
|
||||
}
|
||||
|
||||
impl OcrProvider {
|
||||
pub(crate) const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Cohere => "cohere",
|
||||
Self::Mistral => "mistral",
|
||||
Self::AzureAi => "azure_ai",
|
||||
Self::Reducto => "reducto",
|
||||
Self::VertexAi => "vertex_ai",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_wire_adapter(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<(String, OcrAdapterKind), Error> {
|
||||
let provider =
|
||||
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: OcrProvider::Mistral.as_str(),
|
||||
});
|
||||
let typed_provider = match provider.custom_llm_provider {
|
||||
"cohere" => OcrProvider::Cohere,
|
||||
"mistral" => OcrProvider::Mistral,
|
||||
"azure_ai" => OcrProvider::AzureAi,
|
||||
"reducto" => OcrProvider::Reducto,
|
||||
"vertex_ai" => OcrProvider::VertexAi,
|
||||
value => return Err(Error::InvalidProvider(value.to_string())),
|
||||
};
|
||||
let adapter = match typed_provider {
|
||||
OcrProvider::Cohere => OcrAdapterKind::Cohere,
|
||||
OcrProvider::Mistral => OcrAdapterKind::Mistral,
|
||||
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
|
||||
OcrAdapterKind::AzureDocumentIntelligence
|
||||
}
|
||||
OcrProvider::AzureAi
|
||||
if provider.model.to_ascii_lowercase().contains("cohere")
|
||||
&& provider.model.to_ascii_lowercase().contains("parse") =>
|
||||
{
|
||||
OcrAdapterKind::AzureCohere
|
||||
}
|
||||
OcrProvider::AzureAi => OcrAdapterKind::AzureMistral,
|
||||
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
|
||||
OcrAdapterKind::ReductoLegacy
|
||||
}
|
||||
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-v3") => {
|
||||
OcrAdapterKind::ReductoV3
|
||||
}
|
||||
OcrProvider::Reducto => OcrAdapterKind::ReductoV3,
|
||||
OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
|
||||
OcrAdapterKind::VertexDeepSeek
|
||||
}
|
||||
OcrProvider::VertexAi => OcrAdapterKind::VertexMistral,
|
||||
};
|
||||
Ok((provider.model.to_string(), adapter))
|
||||
}
|
||||
|
||||
fn is_document_intelligence_model(model: &str) -> bool {
|
||||
let model = model.to_ascii_lowercase();
|
||||
model.contains("doc-intelligence") || model.contains("documentintelligence")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn provider_models_are_preserved_without_a_local_allowlist() {
|
||||
let cases = [
|
||||
("mistral/future-ocr-model", OcrAdapterKind::Mistral),
|
||||
("azure_ai/future-ocr-model", OcrAdapterKind::AzureMistral),
|
||||
];
|
||||
|
||||
for (qualified_model, expected_adapter) in cases {
|
||||
let expected_model = qualified_model.split_once('/').unwrap().1;
|
||||
let (model, adapter) = resolve_wire_adapter(qualified_model, None).unwrap();
|
||||
assert_eq!(model, expected_model);
|
||||
assert_eq!(adapter, expected_adapter);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_reducto_models_use_the_current_protocol() {
|
||||
let (model, adapter) = resolve_wire_adapter("reducto/future-parse-model", None).unwrap();
|
||||
assert_eq!(model, "future-parse-model");
|
||||
assert_eq!(adapter, OcrAdapterKind::ReductoV3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn known_protocol_models_still_select_specialized_adapters() {
|
||||
let (model, adapter) = resolve_wire_adapter("reducto/parse-legacy", None).unwrap();
|
||||
assert_eq!(model, "parse-legacy");
|
||||
assert_eq!(adapter, OcrAdapterKind::ReductoLegacy);
|
||||
|
||||
let (model, adapter) =
|
||||
resolve_wire_adapter("azure_ai/doc-intelligence/prebuilt-layout", None).unwrap();
|
||||
assert_eq!(model, "doc-intelligence/prebuilt-layout");
|
||||
assert_eq!(adapter, OcrAdapterKind::AzureDocumentIntelligence);
|
||||
}
|
||||
}
|
||||
|
|
@ -6,7 +6,7 @@ use serde::{Deserialize, Serialize};
|
|||
use serde_json::{Map, Value};
|
||||
|
||||
use super::hooks::{NoopOcrHooks, OcrHooks};
|
||||
use super::registry::{OcrAdapterKind, resolve_wire_adapter};
|
||||
use super::provider_config::{OcrConfigKind, resolve_provider_config};
|
||||
use crate::Error;
|
||||
use crate::auth::{InputSource, TokenProviderHandle};
|
||||
use crate::constants::OCR_HTTP_TIMEOUT_SECS;
|
||||
|
|
@ -98,7 +98,7 @@ pub struct LiteLLMOcrRequest {
|
|||
pub optional_params: Map<String, Value>,
|
||||
pub input_sources: BTreeMap<String, InputSource>,
|
||||
pub azure_ad_token_provider: Option<TokenProviderHandle>,
|
||||
pub(crate) adapter: OcrAdapterKind,
|
||||
pub(crate) config: OcrConfigKind,
|
||||
}
|
||||
|
||||
impl LiteLLMOcrRequest {
|
||||
|
|
@ -108,7 +108,7 @@ impl LiteLLMOcrRequest {
|
|||
custom_llm_provider: Option<&str>,
|
||||
optional_params: Map<String, Value>,
|
||||
) -> Result<Self, Error> {
|
||||
let (model, adapter_kind) = resolve_wire_adapter(&model, custom_llm_provider)?;
|
||||
let (model, config) = resolve_provider_config(&model, custom_llm_provider)?;
|
||||
|
||||
Ok(Self {
|
||||
model,
|
||||
|
|
@ -119,7 +119,7 @@ impl LiteLLMOcrRequest {
|
|||
optional_params,
|
||||
input_sources: BTreeMap::new(),
|
||||
azure_ad_token_provider: None,
|
||||
adapter: adapter_kind,
|
||||
config,
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -137,7 +137,7 @@ impl LiteLLMOcrRequest {
|
|||
}
|
||||
|
||||
pub fn provider_name(&self) -> &'static str {
|
||||
self.adapter.provider().as_str()
|
||||
self.config.provider().as_str()
|
||||
}
|
||||
|
||||
pub fn with_host_hooks(
|
||||
|
|
|
|||
|
|
@ -83,31 +83,31 @@ pub struct OcrWireRequest {
|
|||
}
|
||||
|
||||
pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> bool {
|
||||
super::registry::resolve_wire_adapter(model, custom_llm_provider).is_ok()
|
||||
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>, Error> {
|
||||
use super::registry::OcrAdapterKind;
|
||||
use super::provider_config::OcrConfigKind;
|
||||
|
||||
let (_, adapter) = super::registry::resolve_wire_adapter(model, custom_llm_provider)?;
|
||||
let provider_fields: &[&str] = match adapter {
|
||||
OcrAdapterKind::Cohere | OcrAdapterKind::AzureCohere => &["output_format"],
|
||||
OcrAdapterKind::Mistral | OcrAdapterKind::AzureMistral | OcrAdapterKind::VertexMistral => {
|
||||
let (_, config) = super::provider_config::resolve_provider_config(model, custom_llm_provider)?;
|
||||
let provider_fields: &[&str] = match config {
|
||||
OcrConfigKind::Cohere | OcrConfigKind::AzureCohere => &["output_format"],
|
||||
OcrConfigKind::Mistral | OcrConfigKind::AzureAi | OcrConfigKind::VertexAi => {
|
||||
MISTRAL_OPTION_FIELDS
|
||||
}
|
||||
OcrAdapterKind::AzureDocumentIntelligence => DOCUMENT_INTELLIGENCE_OPTION_FIELDS,
|
||||
OcrAdapterKind::ReductoV3 => REDUCTO_V3_OPTION_FIELDS,
|
||||
OcrAdapterKind::ReductoLegacy => REDUCTO_LEGACY_OPTION_FIELDS,
|
||||
OcrAdapterKind::VertexDeepSeek => DEEPSEEK_OPTION_FIELDS,
|
||||
OcrConfigKind::AzureDocumentIntelligence => DOCUMENT_INTELLIGENCE_OPTION_FIELDS,
|
||||
OcrConfigKind::ReductoV3 => REDUCTO_V3_OPTION_FIELDS,
|
||||
OcrConfigKind::ReductoLegacy => REDUCTO_LEGACY_OPTION_FIELDS,
|
||||
OcrConfigKind::VertexDeepSeek => DEEPSEEK_OPTION_FIELDS,
|
||||
};
|
||||
let auth_fields: &[&str] = match adapter {
|
||||
OcrAdapterKind::AzureMistral
|
||||
| OcrAdapterKind::AzureDocumentIntelligence
|
||||
| OcrAdapterKind::AzureCohere => AZURE_AUTH_OPTION_FIELDS,
|
||||
OcrAdapterKind::VertexMistral | OcrAdapterKind::VertexDeepSeek => VERTEX_AUTH_OPTION_FIELDS,
|
||||
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
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::ocr::codecs::deepseek::{
|
||||
use crate::llms::vertex_ai::ocr::deepseek_transformation::{
|
||||
DeepSeekOcrParams, DeepSeekOcrResponse, transform_ocr_request, transform_ocr_response,
|
||||
};
|
||||
use crate::ocr::types::OcrDocument;
|
||||
|
|
@ -65,7 +65,10 @@ fn request_maps_both_document_types_to_image_content(#[case] document: Value) {
|
|||
#[case(json!("[]"), "[]")]
|
||||
#[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")]
|
||||
#[case(json!({"pages":[{"markdown":"object"}]}), "object")]
|
||||
fn response_codec_handles_text_json_and_objects(#[case] content: Value, #[case] expected: &str) {
|
||||
fn response_transform_handles_text_json_and_objects(
|
||||
#[case] content: Value,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
let response: DeepSeekOcrResponse = serde_json::from_value(
|
||||
json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}),
|
||||
)
|
||||
|
|
@ -102,7 +105,7 @@ fn structured_result_maps_pages_usage_model_and_annotation() {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn response_codec_rejects_missing_empty_and_malformed_content() {
|
||||
fn response_transform_rejects_missing_empty_and_malformed_content() {
|
||||
for value in [
|
||||
json!({"choices":[]}),
|
||||
json!({"choices":[{"message":{"content":""}}]}),
|
||||
|
|
|
|||
|
|
@ -182,7 +182,7 @@ async fn rejects_invalid_document_sources_before_network(#[case] source: &str) {
|
|||
|
||||
#[test]
|
||||
fn response_normalization_groups_blocks_and_distinguishes_null_result() {
|
||||
use crate::ocr::codecs::reducto::{ReductoResponse, transform_ocr_response};
|
||||
use crate::llms::reducto::ocr::transformation::{ReductoResponse, transform_ocr_response};
|
||||
|
||||
let raw = json!({"usage":{"num_pages":"2","credits":"3"},"result":{"type":"full","chunks":[
|
||||
{"blocks":[{
|
||||
|
|
|
|||
|
|
@ -96,10 +96,12 @@ async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
|
|||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn adapters_build_complete_requests_and_share_mistral_normalization() {
|
||||
async fn configs_build_complete_requests_and_share_mistral_normalization() {
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::ocr::adapters::{MistralAdapter, OcrAdapter, VertexMistralAdapter};
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
|
||||
use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig;
|
||||
use crate::ocr::test_support::ocr_client;
|
||||
|
||||
let client = ocr_client();
|
||||
|
|
@ -116,11 +118,11 @@ async fn adapters_build_complete_requests_and_share_mistral_normalization() {
|
|||
options.clone(),
|
||||
);
|
||||
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
|
||||
let direct_http = MistralAdapter
|
||||
let direct_http = MistralOCRConfig
|
||||
.prepare_request(&direct, &client)
|
||||
.await
|
||||
.unwrap();
|
||||
let vertex_http = VertexMistralAdapter
|
||||
let vertex_http = VertexAIOCRConfig::default()
|
||||
.prepare_request(&vertex, &client)
|
||||
.await
|
||||
.unwrap();
|
||||
|
|
@ -146,11 +148,11 @@ async fn adapters_build_complete_requests_and_share_mistral_normalization() {
|
|||
);
|
||||
}
|
||||
let payload = json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"});
|
||||
let direct_response = MistralAdapter
|
||||
let direct_response = MistralOCRConfig
|
||||
.transform_ocr_response(&direct, serde_json::from_value(payload.clone()).unwrap())
|
||||
.unwrap()
|
||||
.into_json();
|
||||
let vertex_response = VertexMistralAdapter
|
||||
let vertex_response = VertexAIOCRConfig::default()
|
||||
.transform_ocr_response(&vertex, serde_json::from_value(payload).unwrap())
|
||||
.unwrap()
|
||||
.into_json();
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue