mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
5e23db8e03
commit
b544f2244b
30 changed files with 1689 additions and 64 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
));
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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}")]
|
||||
|
|
|
|||
|
|
@ -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, ¶ms)?;
|
||||
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())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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())
|
||||
49
litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs
Normal file
49
litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs
Normal 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(())
|
||||
}
|
||||
|
|
@ -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;
|
||||
}
|
||||
};
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
@ -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")),
|
||||
}
|
||||
}
|
||||
|
|
@ -1 +1,2 @@
|
|||
pub(crate) mod document_intelligence;
|
||||
pub(crate) mod mistral;
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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)),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
));
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue