feat(ocr): add Azure Document Intelligence adapter (#40534)

* feat(ocr): add Azure Document Intelligence

* fix(ocr): decline missing Document Intelligence credentials

* fix(ocr): map Document Intelligence credentials

* test(ocr): expose Azure transport to adapter tests

* fix(ocr): preserve native responses through Rust bridge

* feat(core): add URL query pair completion

* fix(ocr): declare Document Intelligence native responses

* fix(auth): preserve Azure credential provenance in OCR adapters

* refactor(ocr): use shared native response handling

* refactor(ocr): preserve Document Intelligence extra params

* refactor(ocr): adopt request preparation contract

* refactor(ocr): keep native response handling behind bridge

* fix(ocr): prevent credential-bearing polling redirects

* fix(ocr): update Azure auth imports

* fix(ocr): bound Document Intelligence polling rate

* fix(ocr): preserve proxy credential provenance

* test(ocr): assert native bridge format support
This commit is contained in:
yujonglee 2026-09-11 16:22:55 -07:00 • committed by GitHub
parent 5e23db8e03
commit b544f2244b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
30 changed files with 1689 additions and 64 deletions

View file

@ -272,7 +272,8 @@ fn core_error_kind(error: &Error) -> &'static str {
Error::Auth(_)
| Error::MissingApiKey { .. }
| Error::MissingAzureAiCredentials
| Error::MissingAzureAiCredentialsOrAdToken => "AuthError",
| Error::MissingAzureAiCredentialsOrAdToken
| Error::MissingAzureDocumentIntelligenceCredentials => "AuthError",
Error::InvalidProvider(_) => "InvalidProvider",
Error::InvalidRequest(_) => "InvalidRequest",
Error::InvalidType { .. } => "InvalidType",

View file

@ -389,7 +389,8 @@ fn core_error_kind(error: &Error) -> &'static str {
Error::Auth(_)
| Error::MissingApiKey { .. }
| Error::MissingAzureAiCredentials
| Error::MissingAzureAiCredentialsOrAdToken => "AuthError",
| Error::MissingAzureAiCredentialsOrAdToken
| Error::MissingAzureDocumentIntelligenceCredentials => "AuthError",
Error::InvalidProvider(_) => "InvalidProvider",
Error::InvalidRequest(_) => "InvalidRequest",
Error::InvalidType { .. } => "InvalidType",

View file

@ -34,10 +34,10 @@ mod tests {
use litellm_core::ocr::wire::is_supported_request;
#[test]
fn core_activation_excludes_unmigrated_azure_document_intelligence() {
fn core_activation_includes_azure_document_intelligence() {
assert!(is_supported_request("model", Some("mistral")));
assert!(is_supported_request("pixtral-12b", Some("azure_ai")));
assert!(!is_supported_request(
assert!(is_supported_request(
"doc-intelligence/prebuilt-layout",
Some("azure_ai")
));

View file

@ -117,7 +117,8 @@ impl IntoResponse for MessagesRouteError {
| Error::MissingField(_)
| Error::MissingApiKey { .. }
| Error::MissingAzureAiCredentials
| Error::MissingAzureAiCredentialsOrAdToken => (
| Error::MissingAzureAiCredentialsOrAdToken
| Error::MissingAzureDocumentIntelligenceCredentials => (
StatusCode::BAD_GATEWAY,
"messages provider request failed".to_string(),
),

View file

@ -51,5 +51,12 @@ pub(crate) const OCR_CONNECT_TIMEOUT_SECS: u64 = 10;
pub(crate) const OCR_INLINE_MAX_BYTES: usize = 50 * 1024 * 1024;
pub(crate) const OCR_DOWNLOAD_MAX_BYTES: u64 = 50 * 1024 * 1024;
pub(crate) const OCR_MAX_FETCH_REDIRECTS: usize = 10;
pub(crate) const OCR_POLL_TIMEOUT_SECS: u64 = 120;
pub(crate) const OCR_POLL_RETRY_SECS: u64 = 2;
pub(crate) const AZURE_DI_API_VERSION: &str = "2024-11-30";
pub(crate) const AZURE_DI_SUBSCRIPTION_HEADER: &str = "Ocp-Apim-Subscription-Key";
pub(crate) const AZURE_DI_DEFAULT_DPI: i64 = 96;
pub(crate) const AZURE_DI_DEFAULT_WIDTH: f64 = 8.5;
pub(crate) const AZURE_DI_DEFAULT_HEIGHT: f64 = 11.0;
pub(crate) const AZURE_AI_OCR_PATH: &str = "/providers/mistral/azure/ocr";
pub(crate) const MISTRAL_OCR_API_BASE: &str = "https://api.mistral.ai/v1";

View file

@ -27,6 +27,10 @@ pub enum Error {
MissingAzureAiCredentials,
#[error("Missing Azure AI credentials - set AZURE_AI_API_KEY or provide azure_ad_token")]
MissingAzureAiCredentialsOrAdToken,
#[error(
"invalid authentication configuration: Missing Azure Document Intelligence credentials - set AZURE_DOCUMENT_INTELLIGENCE_API_KEY or configure Entra ID"
)]
MissingAzureDocumentIntelligenceCredentials,
#[error("upstream request failed with status {status}: {body}")]
Http { status: u16, body: String },
#[error("upstream network error: {0}")]

View file

@ -0,0 +1,214 @@
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::ocr::wire::DecodedOcrResponse;
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 config = 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 = document_intelligence::transform_ocr_request(request.document.clone())?;
transform_request_body(client, request, &url, &headers, 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<DecodedOcrResponse<Self::ProviderResponse>, OcrError> {
polling::read_operation_response(
client.polling_http(),
response,
url,
headers,
&request.connection,
request.response_format()? == OcrResponseFormat::Native,
)
.await
}
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
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

@ -0,0 +1,100 @@
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::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,
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
if response.status() != reqwest::StatusCode::ACCEPTED {
return read_json_response(response, native).await;
}
let location = response
.headers()
.get("operation-location")
.and_then(|value| value.to_str().ok())
.ok_or(OcrPollingError::PollLocation)?;
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());
}
poll_operation(http_client, operation, headers, connection, native).await
}
async fn poll_operation(
http_client: &reqwest::Client,
url: Url,
headers: &[(String, String)],
connection: &OcrConnection,
native: bool,
) -> 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),
)
.await
.map_err(|_| OcrPollingError::PollTimeout)??;
match &decoded.data.status {
Some(OperationStatus::Succeeded) => 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,8 +1,5 @@
use std::sync::OnceLock;
use super::OcrAdapter;
use super::super::OcrAdapter;
use crate::Error;
use crate::auth::error::AuthConfigurationError;
use crate::auth::{InputSource, Sourced};
use crate::constants::AZURE_AI_OCR_PATH;
use crate::ocr::OcrClient;
@ -14,7 +11,7 @@ use crate::ocr::prepare::{
};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
use crate::providers::azure_ai::auth::{AzureAuthInputs, AzureAuthService};
use crate::providers::azure_ai::auth::AzureAuthInputs;
use crate::url_utils::ApiUrl;
const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY";
@ -92,7 +89,7 @@ async fn validate_environment(
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
validate_destination(connection, connection.extra_headers_source)?;
super::validate_destination(connection, connection.extra_headers_source)?;
return Ok(connection.extra_headers.clone());
}
let key = nonblank(connection.api_key.clone())
@ -102,41 +99,16 @@ async fn validate_environment(
.map(|value| Sourced::new(value, InputSource::Environment))
});
if let Some(key) = key {
validate_destination(connection, key.source())?;
super::validate_destination(connection, key.source())?;
return Ok(bearer_headers(connection, key.value()));
}
static SERVICE: OnceLock<AzureAuthService> = OnceLock::new();
let key = SERVICE
.get_or_init(AzureAuthService::default)
.get_azure_ad_token(config, env_lookup)
.await
.map_err(Error::from)?
.map(|credential| {
let source = credential.source();
let value = credential.value().secret().expose().to_string();
Sourced::new(value, source)
})
let key = super::resolve_entra(config, env_lookup)
.await?
.ok_or(Error::MissingAzureAiCredentials)?;
validate_destination(connection, key.source())?;
super::validate_destination(connection, key.source())?;
Ok(bearer_headers(connection, key.value()))
}
fn validate_destination(
connection: &OcrConnection,
credential_source: InputSource,
) -> Result<(), OcrError> {
if connection.api_base.is_some()
&& connection.api_base_source == InputSource::Request
&& credential_source != InputSource::Request
{
return Err(Error::from(crate::AuthError::Configuration(
AuthConfigurationError::RequestAzureCredentialDestination,
))
.into());
}
Ok(())
}
fn bearer_headers(connection: &OcrConnection, key: &str) -> Vec<(String, String)> {
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
.chain(connection.extra_headers.clone())
@ -177,9 +149,9 @@ mod tests {
..Default::default()
};
assert_eq!(
validate_environment(&connection, &Default::default(), &|_| Some(
"environment-key".into()
))
validate_environment(&connection, &Default::default(), &|_| {
Some("environment-key".into())
})
.await
.unwrap(),
connection.extra_headers
@ -193,9 +165,9 @@ mod tests {
..Default::default()
};
assert_eq!(
validate_environment(&connection, &Default::default(), &|_| Some(
"environment-key".into()
))
validate_environment(&connection, &Default::default(), &|_| {
Some("environment-key".into())
})
.await
.unwrap()[0],
("Authorization".into(), "Bearer request-key".into())

View file

@ -0,0 +1,49 @@
mod document_intelligence;
mod mistral;
use std::sync::OnceLock;
use crate::Error;
use crate::auth::error::AuthConfigurationError;
use crate::auth::{InputSource, Sourced};
use crate::ocr::error::OcrError;
use crate::ocr::types::OcrConnection;
use crate::providers::azure_ai::auth::{AzureAuthInputs, AzureAuthService};
pub(crate) use document_intelligence::AzureDocumentIntelligenceAdapter;
pub(crate) use mistral::AzureMistralAdapter;
async fn resolve_entra(
config: &AzureAuthInputs,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Option<Sourced<String>>, Error> {
static SERVICE: OnceLock<AzureAuthService> = OnceLock::new();
SERVICE
.get_or_init(AzureAuthService::default)
.get_azure_ad_token(config, env_lookup)
.await
.map(|credential| {
credential.map(|credential| {
let source = credential.source();
let value = credential.value().secret().expose().to_string();
Sourced::new(value, source)
})
})
.map_err(Error::from)
}
fn validate_destination(
connection: &OcrConnection,
credential_source: InputSource,
) -> Result<(), OcrError> {
if connection.api_base.is_some()
&& connection.api_base_source == InputSource::Request
&& credential_source != InputSource::Request
{
return Err(Error::from(crate::AuthError::Configuration(
AuthConfigurationError::RequestAzureCredentialDestination,
))
.into());
}
Ok(())
}

View file

@ -8,10 +8,10 @@ use super::registry::OcrProvider;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrResponseFormat};
use super::wire::DecodedOcrResponse;
mod azure_mistral;
mod azure;
mod mistral;
pub(crate) use azure_mistral::AzureMistralAdapter;
pub(crate) use azure::{AzureDocumentIntelligenceAdapter, AzureMistralAdapter};
pub(crate) use mistral::MistralAdapter;
/// Converts a complete LiteLLM OCR call to provider HTTP and normalizes its response.
@ -65,6 +65,7 @@ macro_rules! for_each_ocr_adapter {
$callback! {
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;
}
};
}

View file

@ -15,6 +15,7 @@ use crate::media::MediaFetcher;
#[derive(Clone)]
pub struct OcrClient {
provider_http: reqwest::Client,
polling_http: reqwest::Client,
document_fetcher: MediaFetcher,
}
@ -23,6 +24,7 @@ impl OcrClient {
let document_fetcher = MediaFetcher::new().map_err(TransportError::from)?;
Ok(Self {
provider_http,
polling_http: no_redirect_http()?,
document_fetcher,
})
}
@ -41,6 +43,10 @@ impl OcrClient {
&self.provider_http
}
pub(crate) fn polling_http(&self) -> &reqwest::Client {
&self.polling_http
}
pub(crate) fn document_fetcher(&self) -> &MediaFetcher {
&self.document_fetcher
}
@ -49,11 +55,20 @@ impl OcrClient {
pub(crate) fn for_test(provider_http: reqwest::Client, document_http: reqwest::Client) -> Self {
Self {
provider_http,
polling_http: no_redirect_http().expect("test polling client builds"),
document_fetcher: MediaFetcher::for_test(document_http),
}
}
}
fn no_redirect_http() -> Result<reqwest::Client, TransportError> {
reqwest::Client::builder()
.connect_timeout(Duration::from_secs(OCR_CONNECT_TIMEOUT_SECS))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(TransportError::from)
}
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
static CLIENT: OnceLock<Result<OcrClient, TransportError>> = OnceLock::new();
let client = CLIENT

View file

@ -0,0 +1,9 @@
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

@ -0,0 +1,195 @@
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!(["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

@ -0,0 +1,111 @@
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::MissingField("document URL"));
}
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(
"key_value_pairs".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

@ -0,0 +1,138 @@
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 +1,2 @@
pub(crate) mod document_intelligence;
pub(crate) mod mistral;

View file

@ -22,6 +22,12 @@ pub enum OcrRequestError {
DownloadTooLarge,
#[error("OCR document download exceeded the redirect limit")]
TooManyRedirects,
#[error("invalid OCR pages: {0}")]
Pages(String),
#[error("invalid OCR features")]
Features,
#[error("OCR model cannot be a dot segment")]
DotModel,
}
#[derive(Debug, Clone, PartialEq, Eq, Error)]
@ -32,6 +38,20 @@ pub enum OcrResponseError {
MissingRedirectLocation,
#[error("OCR document redirect location is invalid")]
InvalidRedirect,
#[error("OCR operation ended with status {0}")]
OperationStatus(String),
#[error("OCR response numeric value is out of range: {0}")]
NumericRange(&'static str),
}
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum OcrPollingError {
#[error("OCR accepted response is missing a valid operation-location")]
PollLocation,
#[error("OCR operation-location must use the submission origin without credentials")]
PollOrigin,
#[error("OCR polling timed out")]
PollTimeout,
}
#[derive(Debug, Error)]
@ -43,6 +63,8 @@ pub enum OcrError {
#[error("{0}")]
Transport(#[from] TransportError),
#[error("{0}")]
Polling(#[from] OcrPollingError),
#[error("{0}")]
Public(#[from] crate::Error),
}
@ -52,6 +74,7 @@ impl From<OcrError> for crate::Error {
OcrError::Request(error) => error.into(),
OcrError::Response(error) => error.into(),
OcrError::Transport(error) => error.into(),
OcrError::Polling(error) => crate::Error::InvalidResponse(error.to_string()),
OcrError::Public(error) => error,
}
}

View file

@ -18,6 +18,9 @@ pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocumen
#[path = "../../tests/azure_ai_ocr.rs"]
mod azure_ai_tests;
#[cfg(test)]
#[path = "../../tests/azure_document_intelligence_ocr.rs"]
mod azure_document_intelligence_tests;
#[cfg(test)]
#[path = "../../tests/ocr/support.rs"]
pub(crate) mod test_support;
#[cfg(test)]

View file

@ -52,9 +52,10 @@ pub(crate) fn resolve_wire_adapter(
};
match typed_provider {
OcrProvider::Mistral => Ok((provider.model.to_string(), OcrAdapterKind::Mistral)),
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
Err(Error::InvalidProvider("azure_ai".into()))
}
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => Ok((
provider.model.to_string(),
OcrAdapterKind::AzureDocumentIntelligence,
)),
OcrProvider::AzureAi => Ok((provider.model.to_string(), OcrAdapterKind::AzureMistral)),
}
}

View file

@ -74,6 +74,7 @@ pub struct OcrConnection {
pub extra_headers_source: InputSource,
pub timeout: Duration,
pub max_download_bytes: u64,
pub poll_timeout: Duration,
}
impl Default for OcrConnection {
@ -87,6 +88,7 @@ impl Default for OcrConnection {
extra_headers_source: InputSource::Deployment,
timeout: Duration::from_secs(OCR_HTTP_TIMEOUT_SECS),
max_download_bytes: crate::constants::OCR_DOWNLOAD_MAX_BYTES,
poll_timeout: Duration::from_secs(crate::constants::OCR_POLL_TIMEOUT_SECS),
}
}
}

View file

@ -81,6 +81,7 @@ pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, Error>
extra_headers_source,
timeout: timeout.unwrap_or(defaults.timeout),
max_download_bytes: defaults.max_download_bytes,
poll_timeout: defaults.poll_timeout,
};
Ok(LiteLLMOcrRequest {
connection,

View file

@ -60,6 +60,14 @@ impl ApiUrl<Base> {
}
impl ApiUrl<Complete> {
pub(crate) fn append_query_pairs<'a>(
mut self,
pairs: impl IntoIterator<Item = (&'a str, &'a str)>,
) -> Self {
self.url.query_pairs_mut().extend_pairs(pairs);
self
}
pub(crate) fn into_string(self) -> String {
self.url.into()
}
@ -92,4 +100,19 @@ mod tests {
.expect("url builds");
assert_eq!(actual, "https://example.test/v1/ocr?tenant=a");
}
#[test]
fn appended_query_pairs_are_encoded() {
let actual = ApiUrl::parse("https://example.test")
.and_then(|url| url.complete_path(&["analyze"]))
.map(|url| {
url.append_query_pairs([("model", "name with spaces")])
.into_string()
})
.expect("url builds");
assert_eq!(
actual,
"https://example.test/analyze?model=name+with+spaces"
);
}
}

View file

@ -0,0 +1,395 @@
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use super::wire::{OcrWireRequest, decode_request};
fn query_value(url: &str, key: &str) -> Option<String> {
url::Url::parse(url)
.unwrap()
.query_pairs()
.find_map(|(name, value)| (name == key).then(|| value.into_owned()))
}
#[tokio::test]
async fn facade_maps_pages_features_and_url_document() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":{"pages":[]}
}))])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"]}),
);
request.document = serde_json::from_value(json!({
"type":"document_url",
"document_url":"https://example.com/document.pdf"
}))
.unwrap();
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let target = request.split_whitespace().nth(1).unwrap();
let url = format!("{base}{target}");
assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3"));
assert_eq!(
query_value(&url, "features").as_deref(),
Some("keyValuePairs,languages")
);
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({"urlSource":"https://example.com/document.pdf"})
);
}
#[tokio::test]
async fn rejects_invalid_pages_features_and_format() {
for options in [
json!({"pages":[true]}),
json!({"pages":[1,"2"]}),
json!({"pages":[-1]}),
json!({"pages":"1&&features=bad"}),
json!({"features":"languages&pages=1"}),
json!({"req_format":"azure"}),
] {
let result = decode_request(OcrWireRequest {
model: "azure_ai/doc-intelligence/prebuilt-read".into(),
document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
api_key: Some("key".into()),
api_base: Some("http://127.0.0.1:1".into()),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone(),
input_sources: Default::default(),
timeout_seconds: None,
});
let rejected = match result {
Ok(request) => perform_ocr(request).await.is_err(),
Err(_) => true,
};
assert!(rejected, "accepted {options}");
}
}
#[tokio::test]
async fn inline_document_decodes_to_base64_source() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(body, json!({"base64Source":"YWJj"}));
}
#[tokio::test]
async fn immediate_response_normalizes_pages_and_preserves_native() {
let operation = json!({
"status":"succeeded",
"operationExtension":42,
"analyzeResult":{
"content":"A\n\nB",
"tables":[{"cells":[]}],
"keyValuePairs":[{"key":{"content":"A"}}],
"pages":[{
"pageNumber":"2",
"width":"8.5",
"height":11,
"unit":"inch",
"lines":[{"content":"A"},{"content":null},{"content":"B"}]
}]
}
});
let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await;
let result = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0]["index"], 1);
assert_eq!(result.pages[0]["markdown"], "A\n\nB");
assert_eq!(
result.pages[0]["dimensions"],
json!({"width":816,"height":1056,"dpi":96})
);
assert_eq!(result.usage_info, Some(json!({"pages_processed":1})));
assert_eq!(result.provider_native_response, Some(operation));
}
#[tokio::test]
async fn accepted_response_polls_to_success_with_only_credentials() {
let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}});
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "0".into())],
body: json!({"status":"running"}),
},
MockResponse::json(operation.clone()),
])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
);
request
.connection
.extra_headers
.push(("X-Trace".into(), "initial-only".into()));
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(result.provider_native_response, Some(operation));
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 3);
assert!(requests[0].to_ascii_lowercase().contains("x-trace:"));
for poll in &requests[1..] {
assert!(!poll.to_ascii_lowercase().contains("x-trace:"));
assert!(
poll.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: test-key")
);
}
}
#[tokio::test]
async fn polling_forwards_bearer_credentials() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.connection.api_key = None;
request.connection.extra_headers = vec![("Authorization".into(), "Bearer token".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert!(
requests[1]
.to_ascii_lowercase()
.contains("authorization: bearer token")
);
}
#[tokio::test]
async fn polling_does_not_follow_redirects() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 302,
headers: vec![("Location", "{base}/redirected".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
assert!(error.to_string().contains("status 302"), "{error}");
assert_eq!(seen.lock().unwrap().len(), 2);
server.abort();
}
#[tokio::test]
async fn polling_rejects_terminal_failure() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"failed"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("status failed"));
}
#[tokio::test]
async fn malformed_provider_pages_report_response_paths() {
for (analysis, path) in [
(json!({"pages":null}), "pages"),
(json!({"pages":[null]}), "pages[0]"),
(json!({"pages":[{"lines":null}]}), "lines"),
(json!({"pages":[{"width":"bad"}]}), "width"),
] {
let (base, _, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":analysis
}))])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains(path), "{error}");
}
}
#[tokio::test]
async fn rejects_missing_invalid_and_cross_origin_operation_locations() {
for headers in [
Vec::new(),
vec![("Operation-Location", "/relative".into())],
vec![("Operation-Location", "http://example.com/operation".into())],
vec![(
"Operation-Location",
"http://user:password@127.0.0.1/operation".into(),
)],
] {
let (base, _, server) = mock_server(vec![MockResponse {
status: 202,
headers,
body: json!({}),
}])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("operation-location"));
}
}
#[tokio::test]
async fn polling_deadline_bounds_retry_delay() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "9999".into())],
body: json!({"status":"notStarted"}),
},
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.connection.poll_timeout = std::time::Duration::from_millis(100);
let error = tokio::time::timeout(std::time::Duration::from_secs(1), perform_ocr(request))
.await
.unwrap()
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("timed out"));
}
#[tokio::test]
async fn model_id_is_encoded_and_dot_segments_are_rejected() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
perform_ocr(wire_request(
"azure_ai/doc-intelligence/a ?#é",
&base,
json!({}),
))
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("a%20%3F%23%C3%A9:analyze"));
for model in [
"azure_ai/doc-intelligence/.",
"azure_ai/doc-intelligence/..",
] {
let error = perform_ocr(wire_request(model, "http://127.0.0.1:1", json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains("dot segment"));
}
}
#[tokio::test]
async fn pre_call_guardrail_receives_caller_pages_before_mapping() {
use crate::ocr::hooks::{OcrHookFuture, OcrHooks, OcrPreCallRequest};
use std::sync::Arc;
struct RewritePages;
impl OcrHooks for RewritePages {
fn has_guardrails(&self) -> bool {
true
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
Box::pin(async move {
assert_eq!(request.optional_params["pages"], json!([0, 2]));
Ok(OcrPreCallRequest {
optional_params: json!({"pages": [1]}),
..request
})
})
}
}
let (base, seen, server) =
mock_server(vec![MockResponse::json(json!({"status": "succeeded"}))]).await;
let request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"pages": [0, 2]}),
)
.with_host_hooks(Arc::new(RewritePages), None);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
let target = requests[0].split_whitespace().nth(1).unwrap();
assert_eq!(
query_value(&format!("{base}{target}"), "pages").as_deref(),
Some("2")
);
assert_eq!(requests.len(), 1);
}

View file

@ -44,6 +44,7 @@ pub(crate) fn chat_completions_error_to_pyerr(err: Error) -> PyErr {
| Error::MissingApiKey { .. }
| Error::MissingAzureAiCredentials
| Error::MissingAzureAiCredentialsOrAdToken
| Error::MissingAzureDocumentIntelligenceCredentials
| Error::Routing(_)
// Nothing reached the provider, so serving it on Python cannot double
// bill and is the only way the caller gets an answer at all.

View file

@ -102,10 +102,10 @@ mod tests {
use litellm_core::ocr::wire::is_supported_request;
#[test]
fn native_activation_excludes_unmigrated_azure_document_intelligence() {
fn native_activation_includes_azure_document_intelligence() {
assert!(is_supported_request("model", Some("mistral")));
assert!(is_supported_request("pixtral-12b", Some("azure_ai")));
assert!(!is_supported_request(
assert!(is_supported_request(
"documentintelligence/prebuilt-read",
Some("azure_ai")
));

View file

@ -2,14 +2,44 @@
from __future__ import annotations
from collections.abc import Awaitable, Mapping
from collections.abc import Awaitable, Callable, Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables
import httpx
from litellm.rust_bridge.bindings import NativeBinding
import litellm
from litellm.constants import request_timeout
from litellm.llms.azure_ai.ocr.common_utils import is_azure_cohere_parse_model
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
from litellm.rust_bridge.bindings import NativeBinding, native_exception_types
from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import ProviderConfigManager
_RUST_OCR_PROVIDERS: Final = frozenset({"mistral", "azure_ai", "vertex_ai"})
_RUST_OCR_CONFIG_FIELDS: Final = frozenset(
{
"azure_ad_token",
"tenant_id",
"client_id",
"client_secret",
"azure_scope",
"azure_authority_host",
"azure_credential",
"azure_federated_token_file",
"vertex_credentials",
"vertex_ai_credentials",
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
}
)
_RUST_OCR_SECRET_FIELDS: Final = frozenset(
{"azure_ad_token", "client_secret", "azure_federated_token_file", "vertex_credentials", "vertex_ai_credentials"}
)
@dataclass(frozen=True, slots=True)
@ -57,6 +87,26 @@ class RustAocr(Protocol):
raise NotImplementedError
class _OCRLogging(Protocol):
def update_from_kwargs(
self,
*,
kwargs: dict[str, object],
model: str,
optional_params: dict[str, object],
litellm_params: dict[str, object],
custom_llm_provider: str | None,
) -> None: ...
def pre_call(
self,
*,
input: str,
api_key: str | None,
additional_args: dict[str, object],
) -> None: ...
def _as_ocr(value: object) -> RustOcr | None:
return cast(RustOcr, value) if callable(value) else None
@ -77,6 +127,255 @@ def load_rust_aocr() -> RustAocr | None:
return _AOCR.load()
def provider(request: LiteLLMOcrRequest) -> str | None:
if request.custom_llm_provider is not None:
return request.custom_llm_provider
prefix: Final = request.model.partition("/")[0]
if prefix in _RUST_OCR_PROVIDERS:
return prefix
if request.model.startswith("mistral-ocr"):
return "mistral"
return None
def supported(request: LiteLLMOcrRequest) -> bool:
request_provider: Final = provider(request)
if request_provider not in _RUST_OCR_PROVIDERS:
return False
if request_provider == "azure_ai":
return (
not is_azure_cohere_parse_model(request.model)
and not callable(request.kwargs.get("azure_ad_token_provider"))
and request.kwargs.get("azure_username") is None
and request.kwargs.get("azure_password") is None
)
return True
def _optional_params(request: LiteLLMOcrRequest, resolve_secret: Callable[[str], str | None]) -> Mapping[str, object]:
optional_params: Final = MappingProxyType(
{
name: value
for name, value in request.kwargs.items()
if (name not in GenericLiteLLMParams.model_fields or name in _RUST_OCR_CONFIG_FIELDS)
and name not in ("litellm_logging_obj", "aocr", "litellm_call_id", "proxy_server_request")
}
)
request_provider: Final = provider(request)
if request_provider == "azure_ai" and litellm.enable_azure_ad_token_refresh is True:
return MappingProxyType({**optional_params, "enable_azure_ad_token_refresh": True})
if request_provider != "vertex_ai":
return optional_params
project: Final = (
request.kwargs.get("vertex_project")
or request.kwargs.get("vertex_ai_project")
or litellm.vertex_project
or resolve_secret("VERTEXAI_PROJECT")
)
location: Final = (
request.kwargs.get("vertex_location")
or request.kwargs.get("vertex_ai_location")
or litellm.vertex_location
or resolve_secret("VERTEXAI_LOCATION")
or resolve_secret("VERTEX_LOCATION")
)
vertex_params: Final = MappingProxyType(
{
name: value
for name, value in (("vertex_project", project), ("vertex_location", location))
if value is not None
}
)
return MappingProxyType({**optional_params, **vertex_params})
def _input_sources(request: LiteLLMOcrRequest, optional_params: Mapping[str, object]) -> Mapping[str, str]:
proxy_request_value: Final = request.kwargs.get("proxy_server_request")
if not isinstance(proxy_request_value, Mapping):
return MappingProxyType({})
proxy_request: Final = cast( # cast-ok: runtime Mapping check narrows metadata with unknown key and value types
Mapping[object, object], proxy_request_value
)
credential_fields_value: Final = proxy_request.get("credential_fields", ())
credential_fields: Final = (
frozenset(name for name in credential_fields_value if isinstance(name, str))
if isinstance(credential_fields_value, (list, tuple, set, frozenset))
else frozenset()
)
request_fields_value: Final = proxy_request.get("body_fields")
request_fields: Sequence[object]
if isinstance(request_fields_value, Sequence) and not isinstance(request_fields_value, (str, bytes)):
request_fields = cast( # cast-ok: runtime Sequence check excludes scalar strings and bytes
Sequence[object], request_fields_value
)
else:
body_value: Final = proxy_request.get("body")
request_fields = (
tuple(cast(Mapping[object, object], body_value)) # cast-ok: runtime Mapping check establishes iterable keys
if isinstance(body_value, Mapping)
else ()
)
names: Final = frozenset(optional_params) | frozenset({"api_key", "api_base", "extra_headers"})
request_sources: Final = MappingProxyType(
{name: "request" for name in names if name in request_fields or name in credential_fields}
)
if litellm.enable_azure_ad_token_refresh is True and "enable_azure_ad_token_refresh" in optional_params:
return MappingProxyType({**request_sources, "enable_azure_ad_token_refresh": "deployment"})
return request_sources
def _marshal(
request: LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
convert_file_document: Callable[[dict[str, object]], dict[str, str]],
) -> LiteLLMOcrRequest:
if not isinstance(request.document, dict):
raise TypeError(f"document must be a dict with 'type' and URL/file field, got {type(request.document)}")
document: Final = (
convert_file_document(request.document) if request.document.get("type") == "file" else request.document
)
request_provider: Final = provider(request)
api_key: Final = (
request.api_key or resolve_secret("MISTRAL_API_KEY") if request_provider == "mistral" else request.api_key
)
optional_params: Final = _optional_params(request, resolve_secret)
input_sources: Final = _input_sources(request, optional_params)
logged_optional_params: Final = MappingProxyType(
{name: "****" if name in _RUST_OCR_SECRET_FIELDS else value for name, value in optional_params.items()}
)
logged_kwargs: Final = MappingProxyType(
{
name: "****" if name in _RUST_OCR_SECRET_FIELDS else value
for name, value in request.kwargs.items()
if name != "proxy_server_request"
}
)
logging_obj: Final = cast( # cast-ok: client decorator injects the logging object through untyped kwargs
_OCRLogging, request.kwargs["litellm_logging_obj"]
)
logging_obj.update_from_kwargs(
kwargs=dict(logged_kwargs), # mutable-ok: legacy logging mutates its kwargs copy
model=request.model,
optional_params=dict(logged_optional_params), # mutable-ok: legacy logging requires concrete dict params
litellm_params={ # mutable-ok: legacy logging requires a concrete params dict
"litellm_call_id": request.kwargs.get("litellm_call_id"),
"api_base": request.api_base,
},
custom_llm_provider=request_provider,
)
logging_obj.pre_call(
input="OCR document processing",
api_key=api_key,
additional_args={ # mutable-ok: pre_call mutates the additional_args dict
"complete_input_dict": { # mutable-ok: callbacks consume a JSON-serializable request dict
"model": request.model,
"document": document,
**logged_optional_params,
},
"api_base": request.api_base or "",
"headers": request.extra_headers or {}, # mutable-ok: logging callbacks consume a concrete headers dict
},
)
return LiteLLMOcrRequest(
model=request.model,
document=document,
api_key=api_key,
api_base=request.api_base,
timeout=request.timeout if request.timeout is not None else request_timeout,
custom_llm_provider=request.custom_llm_provider,
extra_headers=request.extra_headers,
kwargs=optional_params,
input_sources=input_sources,
)
def _map_error(error: Exception, request: LiteLLMOcrRequest) -> Exception:
exception_types: Final = native_exception_types()
if exception_types is None or not isinstance(error, exception_types[1]):
return error
request_provider: Final = provider(request)
if request_provider is None:
return error
provider_config: Final = ProviderConfigManager.get_provider_ocr_config(
model=request.model.removeprefix(f"{request_provider}/"), provider=litellm.LlmProviders(request_provider)
)
if provider_config is None:
return error
error_args: Final = cast( # cast-ok: BaseException.args exposes Any while native errors carry scalar args
tuple[object, ...], error.args
)
status: Final = error_args[0] if error_args and isinstance(error_args[0], int) else 500
message: Final = str(error_args[1]) if len(error_args) > 1 else str(error)
error_factory: Final = cast( # cast-ok: legacy provider error factories have untyped callable parameters
Callable[..., Exception], provider_config.get_error_class
)
return error_factory(
error_message=message,
status_code=status or 500,
headers={}, # mutable-ok: provider error factories require a concrete headers dict
)
def _response(response: Mapping[str, object]) -> OCRResponse:
provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY)
normalized: Final = OCRResponse.model_validate(
MappingProxyType({key: value for key, value in response.items() if key != PROVIDER_NATIVE_RESPONSE_KEY})
)
if isinstance(provider_native_response, Mapping):
normalized.set_provider_native_response(provider_native_response)
return normalized
def run(
request: LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
convert_file_document: Callable[[dict[str, object]], dict[str, str]],
) -> OCRResponse | None:
if load_rust_ocr() is None:
return None
marshalled: Final = _marshal(request, resolve_secret, convert_file_document)
try:
response: Final = ocr(
model=marshalled.model,
document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict
api_key=marshalled.api_key,
api_base=marshalled.api_base,
custom_llm_provider=marshalled.custom_llm_provider,
extra_headers=marshalled.extra_headers,
optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict
input_sources=marshalled.input_sources,
timeout=marshalled.timeout,
)
except Exception as error:
raise _map_error(error, request) from error
return _response(response) if response is not None else None
async def arun(
request: LiteLLMOcrRequest,
resolve_secret: Callable[[str], str | None],
convert_file_document: Callable[[dict[str, object]], dict[str, str]],
) -> OCRResponse | None:
if load_rust_aocr() is None:
return None
marshalled: Final = _marshal(request, resolve_secret, convert_file_document)
try:
response: Final = await aocr(
model=marshalled.model,
document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict
api_key=marshalled.api_key,
api_base=marshalled.api_base,
custom_llm_provider=marshalled.custom_llm_provider,
extra_headers=marshalled.extra_headers,
optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict
input_sources=marshalled.input_sources,
timeout=marshalled.timeout,
)
except Exception as error:
raise _map_error(error, request) from error
return _response(response) if response is not None else None
def ocr(
*,
model: str,

View file

@ -1,13 +1,11 @@
"""
Tests for the OCR `req_format` option in the SDK request path:
providers that don't support a native response must reject it, and the Rust
bridge (which only returns the normalized shape) must not serve native requests.
Tests for the OCR `req_format` option in the SDK request path.
"""
import pytest
import litellm
from litellm.ocr.main import _rust_ocr_supported
from litellm.rust_bridge import ocr as rust_ocr_bridge
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
@ -30,16 +28,33 @@ def _request(
@pytest.mark.parametrize("optional_params", [{}, {"req_format": "litellm"}])
def test_rust_ocr_serves_default_format(optional_params):
assert _rust_ocr_supported(_request(optional_params)) is True
assert rust_ocr_bridge.supported(_request(optional_params)) is True
def test_rust_ocr_skipped_for_native_format():
assert _rust_ocr_supported(_request({"req_format": "native"})) is False
def test_rust_ocr_serves_native_format_for_document_intelligence():
assert rust_ocr_bridge.supported(_request({"req_format": "native"})) is True
def test_rust_ocr_response_retains_provider_native_response():
provider_response = {"status": "succeeded", "analyzeResult": {"content": "native"}}
response = rust_ocr_bridge._response(
{
"pages": [],
"model": "prebuilt-layout",
"document_annotation": None,
"usage_info": {"pages_processed": 0},
"object": "ocr",
"provider_native_response": provider_response,
}
)
assert response.get_provider_native_response() == provider_response
assert response.model_dump().get("provider_native_response") is None
@pytest.mark.parametrize("model", ["cohere/cohere-parse", "azure_ai/cohere-parse"])
def test_rust_ocr_skipped_for_unsupported_models(model):
assert _rust_ocr_supported(_request({}, model)) is False
assert rust_ocr_bridge.supported(_request({}, model)) is False
@pytest.mark.asyncio

View file

@ -708,6 +708,24 @@ def test_prepare_rust_ocr_call_preserves_proxy_input_sources():
"api_key": "request",
}
marshaled = rust_bridge._marshal(
build_request(
custom_llm_provider="azure_ai",
model="pixtral-12b-2409",
api_key="request-key",
api_base="https://azure.example.com",
litellm_params={
"proxy_server_request": {
"body": {"api_base": "https://azure.example.com"},
"credential_fields": ("api_key",),
}
},
),
lambda _name: None,
lambda document: document,
)
assert marshaled.input_sources == {"api_base": "request", "api_key": "request"}
def test_rust_ocr_logging_redacts_azure_credentials():
bridge = RecordingBridge()

View file

@ -73,7 +73,7 @@ def assert_native_request(
headers: HTTPMessage,
body: object,
) -> None:
if route not in {"ocr", "azure_ocr", "transcription", "messages", "chat_completions"}:
if route not in {"ocr", "azure_ocr", "azure_di", "transcription", "messages", "chat_completions"}:
raise AssertionError(f"unexpected route marker: {route!r}")
if outcome not in {"success", "429", "hang"}:
raise AssertionError(f"unexpected outcome marker: {outcome!r}")
@ -92,6 +92,13 @@ def assert_native_request(
assert body["model"] == "mistral-ocr-2505"
assert body["document"]["document_url"] == "data:application/pdf;base64,YWJj"
return
if route == "azure_di":
assert path.startswith("/documentintelligence/documentModels/prebuilt-read:analyze?")
assert "api-version=2024-11-30" in path
assert "pages=1%2C3" in path
assert headers.get("ocp-apim-subscription-key") == "di-key"
assert body == {"base64Source": "YWJj"}
return
if route == "transcription":
assert path == "/model/mistral.voxtral-mini-3b-2507/converse"
assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ")
@ -115,6 +122,8 @@ def native_response(status: int, route: str | None) -> bytes:
return b'{"error":"native-rate-limit"}'
if route in {"ocr", "azure_ocr"}:
return b'{"pages":[{"index":0,"markdown":"native-ocr"}]}'
if route == "azure_di":
return b'{"status":"succeeded","analyzeResult":{"pages":[]}}'
if route == "transcription":
return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}'
return ANTHROPIC_RESPONSE
@ -201,6 +210,18 @@ def azure_ocr_kwargs(api_base: str) -> dict[str, object]:
}
def azure_di_kwargs(api_base: str) -> dict[str, object]:
return {
"model": "doc-intelligence/prebuilt-read",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"api_key": "di-key",
"api_base": api_base,
"custom_llm_provider": "azure_ai",
"extra_headers": {"x-test-outcome": "success", "x-test-route": "azure_di"},
"optional_params": {"req_format": "native", "pages": [0, 2]},
}
def success_value(route: str, response: dict[object, object]) -> object:
if route == "ocr":
return response["pages"][0]["markdown"]
@ -232,6 +253,8 @@ def exercise_sync(native: object, api_base: str) -> None:
else:
raise AssertionError(f"{route} accepted a 429 response")
assert_success("ocr", native.ocr(**azure_ocr_kwargs(api_base)))
di_response: Final = native.ocr(**azure_di_kwargs(api_base))
assert di_response["provider_native_response"]["status"] == "succeeded"
async def exercise_async(native: object, api_base: str) -> None:
@ -245,6 +268,8 @@ async def exercise_async(native: object, api_base: str) -> None:
else:
raise AssertionError(f"a{route} accepted a 429 response")
assert_success("ocr", await native.aocr(**azure_ocr_kwargs(api_base)))
di_response: Final = await native.aocr(**azure_di_kwargs(api_base))
assert di_response["provider_native_response"]["status"] == "succeeded"
async def exercise_async_concurrency(native: object, api_base: str) -> None: