refactor(ocr): mirror Python provider layout

This commit is contained in:
Yujong Lee 2026-09-15 08:31:17 -07:00
parent 4e996400e2
commit 4d903da85c
59 changed files with 2817 additions and 2611 deletions

View file

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

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

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

View file

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

View file

@ -0,0 +1,3 @@
pub(crate) mod transformation;
pub(crate) use transformation::AzureDocumentIntelligenceOperation;

View file

@ -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, &params)?;
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())
);
}
}

View 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;

View file

@ -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, &params)?;
let body = self
.mistral
.transform_ocr_request(&request.model, document, &params)?;
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()))
}

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

@ -0,0 +1 @@
pub(crate) mod transformation;

View file

@ -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,
)?)
}
}
}

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

@ -0,0 +1,3 @@
pub(crate) mod transformation;
pub(crate) use transformation::{CohereParams, CohereResponse, validate_document};

View 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(_)))
));
}
}

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

@ -0,0 +1,3 @@
pub(crate) mod transformation;
pub(crate) use transformation::{MistralOcrParams, MistralOcrResponse};

View file

@ -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(), &params)?;
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"
}))
));
}
}

View 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;

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

@ -0,0 +1,3 @@
pub(crate) mod transformation;
pub(crate) use transformation::ReductoResponse;

View 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, &params)?;
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, &params)?;
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;

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

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

View file

@ -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, &params)?;
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"
);
}
}

View file

@ -0,0 +1,3 @@
pub(crate) mod common_utils;
pub(crate) mod deepseek_transformation;
pub(crate) mod transformation;

View file

@ -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, &params)?;
let body = self
.mistral
.transform_ocr_request(&request.model, document, &params)?;
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)
}
}

View file

@ -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, &params)?;
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())
);
}
}

View file

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

View file

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

View file

@ -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(), &params)?;
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"
}))
));
}
}

View file

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

View file

@ -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, &params)?;
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)
}
}

View file

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

View file

@ -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, &params)?;
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)
}
}

View file

@ -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, &params)?;
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"
);
}
}

View file

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

View file

@ -1,5 +0,0 @@
mod transformation;
mod types;
pub(crate) use transformation::{transform_ocr_request, transform_ocr_response};
pub(crate) use types::{DeepSeekOcrParams, DeepSeekOcrResponse};

View file

@ -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()),
})
}

View file

@ -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>,
}

View file

@ -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,
};

View file

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

View file

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

View file

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

View file

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

View file

@ -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>,
}

View file

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

View file

@ -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,
};

View file

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

View file

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

View file

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

View file

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

View file

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

View 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
);
}
}

View file

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

View file

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

View file

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

View file

@ -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":""}}]}),

View file

@ -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":[{

View file

@ -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();