This commit is contained in:
Yujong Lee 2026-09-15 10:49:40 -07:00
parent 57e7bce5ed
commit 0d63ecd931
8 changed files with 866 additions and 938 deletions

View file

@ -81,9 +81,12 @@ pub enum Error {
#[error("credential acquisition failed: Azure OIDC reference did not resolve to a value")]
UnresolvedOidcReference,
#[error(
"Missing {provider} API Key - A call is being made to {provider} but no key is set either in the environment variables or via params"
"Missing {provider} API Key - Set `api_key` or the {environment_variable} environment variable"
)]
MissingApiKey { provider: &'static str },
MissingApiKey {
provider: &'static str,
environment_variable: &'static str,
},
#[error(
"Missing {provider} API Base - Set {environment_variable} environment variable or pass api_base parameter"
)]
@ -91,24 +94,27 @@ pub enum Error {
provider: &'static str,
environment_variable: &'static str,
},
#[error(
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
)]
MissingAnthropicApiKey,
#[error("Missing Azure API Key - Set `api_key` or the AZURE_API_KEY environment variable")]
MissingAzureApiKey,
#[error(
"Missing Azure API Base - Set `api_base` or the AZURE_API_BASE environment variable. Expected format: https://<resource-name>.services.ai.azure.com/anthropic"
)]
MissingAzureApiBase,
#[error(
"Missing OpenAI API Key - a realtime call is being made but no key was passed via params or the OPENAI_API_KEY environment variable"
)]
MissingOpenAiRealtimeApiKey,
#[error(
"Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable"
)]
MissingOpenAiResponsesApiKey,
#[error("invalid authentication header")]
InvalidHeader,
}
#[cfg(test)]
mod tests {
use super::Error;
#[test]
fn missing_api_key_names_provider_and_environment_variable() {
assert_eq!(
Error::MissingApiKey {
provider: "Anthropic",
environment_variable: "ANTHROPIC_API_KEY",
}
.to_string(),
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
);
}
}

View file

@ -5,3 +5,5 @@ A route module owns everything the call needs: types, the provider template trai
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or host-specific callback execution. Core owns lifecycle sequencing and callback payload construction; hosts execute the selected integrations. Env reads are limited to credential fallback in a route's `prepare.rs`.
Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates.
Provider ports mirror the Python path under `litellm/llms/`, but they use Rust structure rather than copying Python inheritance. Keep provider transformation files flat by default: imports, constants and wire types, concrete configs and trait implementations, private helpers, then one test module. Add a production submodule only when it creates a real privacy, conditional compilation, or reuse boundary. Share behavior with private functions or explicit delegation; add a provider-specific base trait only when callers need that interface.

View file

@ -20,9 +20,12 @@ pub enum Error {
#[error("{0}")]
Auth(String),
#[error(
"Missing {provider} API Key - A call is being made to {provider} but no key is set either in the environment variables or via params"
"Missing {provider} API Key - Set `api_key` or the {environment_variable} environment variable"
)]
MissingApiKey { provider: &'static str },
MissingApiKey {
provider: &'static str,
environment_variable: &'static str,
},
#[error(
"invalid authentication configuration: Missing Azure AI credentials - set AZURE_AI_API_KEY or configure Entra ID"
)]
@ -144,7 +147,13 @@ impl From<TransportError> for Error {
impl From<crate::AuthError> for Error {
fn from(error: crate::AuthError) -> Self {
match error {
crate::AuthError::MissingApiKey { provider } => Self::MissingApiKey { provider },
crate::AuthError::MissingApiKey {
provider,
environment_variable,
} => Self::MissingApiKey {
provider,
environment_variable,
},
error => Self::Auth(error.to_string()),
}
}
@ -172,10 +181,21 @@ mod transport_tests {
use super::*;
#[test]
fn missing_auth_key_preserves_provider_in_public_error() {
fn missing_auth_key_preserves_guidance_in_public_error() {
let error = Error::from(crate::AuthError::MissingApiKey {
provider: "Vertex",
environment_variable: "GOOGLE_APPLICATION_CREDENTIALS",
});
assert_eq!(
Error::from(crate::AuthError::MissingApiKey { provider: "Vertex" }),
Error::MissingApiKey { provider: "Vertex" }
error,
Error::MissingApiKey {
provider: "Vertex",
environment_variable: "GOOGLE_APPLICATION_CREDENTIALS",
}
);
assert_eq!(
error.to_string(),
"Missing Vertex API Key - Set `api_key` or the GOOGLE_APPLICATION_CREDENTIALS environment variable"
);
}

View file

@ -1,143 +1,151 @@
mod provider {
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value, json};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value, json};
use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
use crate::Error;
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::OcrClient;
use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
use crate::url_utils::ApiUrl;
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum OutputFormat {
#[default]
Markdown,
Blocks,
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum OutputFormat {
#[default]
Markdown,
Blocks,
}
#[derive(Deserialize)]
pub(crate) struct CohereParams {
#[serde(default)]
pub output_format: OutputFormat,
}
#[derive(Deserialize, Serialize)]
pub(crate) struct CohereRequest {
pub model: String,
pub document: OcrDocument,
pub output_format: OutputFormat,
}
pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), OcrRequestError> {
let OcrDocument::ImageUrl { image_url, .. } = document else {
return Err(OcrRequestError::CohereImageOnly);
};
if image_url.is_empty() {
return Err(OcrRequestError::CohereImageOnly);
}
#[derive(Deserialize)]
pub(crate) struct CohereParams {
#[serde(default)]
pub output_format: OutputFormat,
}
#[derive(Deserialize, Serialize)]
pub(crate) struct CohereRequest {
pub model: String,
pub document: OcrDocument,
pub output_format: OutputFormat,
}
pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), OcrRequestError> {
let OcrDocument::ImageUrl { image_url, .. } = document else {
return Err(OcrRequestError::CohereImageOnly);
};
if image_url.is_empty() {
if let Some(inline) = InlineDocument::parse(image_url)? {
if !inline.mime_type().type_.eq_ignore_ascii_case("image") {
return Err(OcrRequestError::CohereImageOnly);
}
if let Some(inline) = InlineDocument::parse(image_url)? {
if !inline.mime_type().type_.eq_ignore_ascii_case("image") {
return Err(OcrRequestError::CohereImageOnly);
}
inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
}
Ok(())
inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
}
Ok(())
}
#[derive(Deserialize)]
pub(crate) struct CohereResponse {
#[serde(default)]
pages: Vec<CoherePage>,
meta: Option<CohereMeta>,
}
#[derive(Deserialize)]
pub(crate) struct CohereResponse {
#[serde(default)]
pages: Vec<CoherePage>,
meta: Option<CohereMeta>,
}
#[derive(Deserialize)]
struct CoherePage {
index: Option<i64>,
markdown: Option<CohereMarkdown>,
blocks: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CoherePage {
index: Option<i64>,
markdown: Option<CohereMarkdown>,
blocks: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMarkdown {
#[serde(default)]
content: String,
images: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMarkdown {
#[serde(default)]
content: String,
images: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMeta {
billed_units: Option<CohereBilledUnits>,
}
#[derive(Deserialize)]
struct CohereMeta {
billed_units: Option<CohereBilledUnits>,
}
#[derive(Deserialize)]
struct CohereBilledUnits {
pages: Option<i64>,
}
#[derive(Deserialize)]
struct CohereBilledUnits {
pages: Option<i64>,
}
pub(crate) fn transform_response(
model: &str,
response: CohereResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
let pages_processed = response
.meta
.and_then(|meta| meta.billed_units)
.and_then(|units| units.pages)
.map(Ok)
.unwrap_or_else(|| {
i64::try_from(response.pages.len())
.map_err(|_| OcrResponseError::NumericRange("pages"))
pub(crate) fn transform_response(
model: &str,
response: CohereResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
let pages_processed = response
.meta
.and_then(|meta| meta.billed_units)
.and_then(|units| units.pages)
.map(Ok)
.unwrap_or_else(|| {
i64::try_from(response.pages.len()).map_err(|_| OcrResponseError::NumericRange("pages"))
})?;
let pages = response
.pages
.into_iter()
.enumerate()
.map(|(position, page)| {
let index = page.index.map(Ok).unwrap_or_else(|| {
i64::try_from(position).map_err(|_| OcrResponseError::NumericRange("page index"))
})?;
let pages = response
.pages
.into_iter()
.enumerate()
.map(|(position, page)| {
let index = page.index.map(Ok).unwrap_or_else(|| {
i64::try_from(position)
.map_err(|_| OcrResponseError::NumericRange("page index"))
})?;
let (content, images) = page
.markdown
.map(|markdown| {
let images =
markdown
.images
.filter(|images| !images.is_empty())
.map(|images| {
images
.into_iter()
.map(|mut image| {
if let Some(Value::Object(bbox)) =
image.get("bounding_box").cloned()
{
image.insert("bbox".into(), Value::Object(bbox));
}
Value::Object(image)
})
.collect::<Vec<_>>()
});
(markdown.content, images)
})
.unwrap_or_default();
let mut normalized = json!({"index": index, "markdown": content, "images": images});
if let Some(blocks) = page.blocks {
normalized["blocks"] = json!(blocks);
}
Ok(normalized)
})
.collect::<Result<Vec<_>, OcrResponseError>>()?;
Ok(LiteLLMOcrResponse {
pages,
model: model.into(),
document_annotation: None,
usage_info: Some(json!({"pages_processed": pages_processed})),
object: "ocr".into(),
extra_fields: Map::new(),
provider_native_response: None,
let (content, images) = page
.markdown
.map(|markdown| {
let images =
markdown
.images
.filter(|images| !images.is_empty())
.map(|images| {
images
.into_iter()
.map(|mut image| {
if let Some(Value::Object(bbox)) =
image.get("bounding_box").cloned()
{
image.insert("bbox".into(), Value::Object(bbox));
}
Value::Object(image)
})
.collect::<Vec<_>>()
});
(markdown.content, images)
})
.unwrap_or_default();
let mut normalized = json!({"index": index, "markdown": content, "images": images});
if let Some(blocks) = page.blocks {
normalized["blocks"] = json!(blocks);
}
Ok(normalized)
})
}
.collect::<Result<Vec<_>, OcrResponseError>>()?;
Ok(LiteLLMOcrResponse {
pages,
model: model.into(),
document_annotation: None,
usage_info: Some(json!({"pages_processed": pages_processed})),
object: "ocr".into(),
extra_fields: Map::new(),
provider_native_response: None,
})
}
pub(crate) fn transform_request(
#[derive(Default)]
pub(crate) struct CohereParseConfig;
impl CohereParseConfig {
pub(crate) fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
params: CohereParams,
@ -149,142 +157,6 @@ mod provider {
output_format: params.output_format,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn response_normalizes_markdown_images_blocks_and_billed_pages() {
let response = serde_json::from_value(json!({
"pages": [
{
"type":"markdown",
"index":4,
"markdown":{
"content":"receipt",
"images":[{
"id":"image",
"bounding_box":{"top_left_x":1,"bottom_right_x":48},
"bounding_box_normalized":{"top_left_x":0.04,"bottom_right_x":0.15},
"description":"scan",
"category":"logo"
}]
}
},
{"type":"blocks","blocks":[{"type":"text","text":{"content":"total"}}]}
],
"meta":{"api_version":{"version":"2"},"billed_units":{"pages":3}}
}))
.unwrap();
let normalized = transform_response("parse-v5.0", response).unwrap();
assert_eq!(normalized.pages[0]["index"], 4);
assert_eq!(normalized.pages[0]["markdown"], "receipt");
assert_eq!(normalized.pages[0]["images"][0]["bbox"]["top_left_x"], 1);
assert_eq!(
normalized.pages[0]["images"][0]["bounding_box_normalized"]["bottom_right_x"],
0.15
);
assert_eq!(normalized.pages[0]["images"][0]["description"], "scan");
assert_eq!(normalized.pages[0]["images"][0]["category"], "logo");
assert_eq!(normalized.pages[1]["index"], 1);
assert_eq!(normalized.pages[1]["markdown"], "");
assert_eq!(normalized.pages[1]["blocks"][0]["text"]["content"], "total");
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 3);
}
#[test]
fn response_defaults_and_invalid_fields() {
for value in [
json!({}),
json!({"meta":null}),
json!({"pages":[],"meta":{"billed_units":null}}),
] {
let normalized =
transform_response("parse", serde_json::from_value(value).unwrap()).unwrap();
assert!(normalized.pages.is_empty());
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 0);
}
for value in [
json!({"pages":null}),
json!({"pages":[{"markdown":"text"}]}),
json!({"pages":[{"index":"bad"}]}),
] {
assert!(serde_json::from_value::<CohereResponse>(value).is_err());
}
let normalized = transform_response(
"parse",
serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(),
)
.unwrap();
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 1);
assert!(normalized.pages[0]["images"].is_null());
}
#[test]
fn request_requires_image_and_supported_output_format() {
for value in [
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
json!({"type":"image_url","image_url":""}),
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
] {
assert_eq!(
validate_document(&serde_json::from_value(value).unwrap()),
Err(OcrRequestError::CohereImageOnly)
);
}
assert!(
serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err()
);
for format in ["markdown", "blocks"] {
assert!(
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
);
}
let request = transform_request(
"parse-v5.0",
serde_json::from_value(json!({
"type":"image_url",
"image_url":"https://example.com/image.png"
}))
.unwrap(),
serde_json::from_value(json!({})).unwrap(),
)
.unwrap();
assert_eq!(
serde_json::to_value(request).unwrap()["output_format"],
"markdown"
);
}
}
}
pub(crate) use provider::{
CohereParams, CohereRequest, CohereResponse, transform_request, transform_response,
validate_document,
};
use crate::Error;
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::OcrClient;
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
use crate::url_utils::ApiUrl;
#[derive(Default)]
pub(crate) struct CohereParseConfig;
impl CohereParseConfig {
pub(crate) fn transform_ocr_request(
&self,
model: &str,
document: crate::ocr::types::OcrDocument,
params: CohereParams,
) -> Result<CohereRequest, OcrRequestError> {
transform_request(model, document, params)
}
}
impl BaseOcrConfig for CohereParseConfig {
@ -372,6 +244,108 @@ fn validate_environment(
mod tests {
use super::*;
#[test]
fn response_normalizes_markdown_images_blocks_and_billed_pages() {
let response = serde_json::from_value(json!({
"pages": [
{
"type":"markdown",
"index":4,
"markdown":{
"content":"receipt",
"images":[{
"id":"image",
"bounding_box":{"top_left_x":1,"bottom_right_x":48},
"bounding_box_normalized":{"top_left_x":0.04,"bottom_right_x":0.15},
"description":"scan",
"category":"logo"
}]
}
},
{"type":"blocks","blocks":[{"type":"text","text":{"content":"total"}}]}
],
"meta":{"api_version":{"version":"2"},"billed_units":{"pages":3}}
}))
.unwrap();
let normalized = transform_response("parse-v5.0", response).unwrap();
assert_eq!(normalized.pages[0]["index"], 4);
assert_eq!(normalized.pages[0]["markdown"], "receipt");
assert_eq!(normalized.pages[0]["images"][0]["bbox"]["top_left_x"], 1);
assert_eq!(
normalized.pages[0]["images"][0]["bounding_box_normalized"]["bottom_right_x"],
0.15
);
assert_eq!(normalized.pages[0]["images"][0]["description"], "scan");
assert_eq!(normalized.pages[0]["images"][0]["category"], "logo");
assert_eq!(normalized.pages[1]["index"], 1);
assert_eq!(normalized.pages[1]["markdown"], "");
assert_eq!(normalized.pages[1]["blocks"][0]["text"]["content"], "total");
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 3);
}
#[test]
fn response_defaults_and_invalid_fields() {
for value in [
json!({}),
json!({"meta":null}),
json!({"pages":[],"meta":{"billed_units":null}}),
] {
let normalized =
transform_response("parse", serde_json::from_value(value).unwrap()).unwrap();
assert!(normalized.pages.is_empty());
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 0);
}
for value in [
json!({"pages":null}),
json!({"pages":[{"markdown":"text"}]}),
json!({"pages":[{"index":"bad"}]}),
] {
assert!(serde_json::from_value::<CohereResponse>(value).is_err());
}
let normalized = transform_response(
"parse",
serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(),
)
.unwrap();
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 1);
assert!(normalized.pages[0]["images"].is_null());
}
#[test]
fn request_requires_image_and_supported_output_format() {
for value in [
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
json!({"type":"image_url","image_url":""}),
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
] {
assert_eq!(
validate_document(&serde_json::from_value(value).unwrap()),
Err(OcrRequestError::CohereImageOnly)
);
}
assert!(serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err());
for format in ["markdown", "blocks"] {
assert!(
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
);
}
let request = CohereParseConfig
.transform_ocr_request(
"parse-v5.0",
serde_json::from_value(json!({
"type":"image_url",
"image_url":"https://example.com/image.png"
}))
.unwrap(),
serde_json::from_value(json!({})).unwrap(),
)
.unwrap();
assert_eq!(
serde_json::to_value(request).unwrap()["output_format"],
"markdown"
);
}
#[test]
fn completes_provider_urls_without_duplicate_paths_and_preserves_queries() {
for suffix in ["", "/v2", "/v2/parse"] {

View file

@ -1,9 +1,17 @@
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
use crate::Error;
use crate::constants::MISTRAL_OCR_API_BASE;
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::OcrClient;
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
use crate::params::OpaqueParams;
use crate::url_utils::ApiUrl;
const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY";
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct MistralOcrRequest {
@ -23,19 +31,6 @@ pub(crate) struct MistralOcrResponse {
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn transform_ocr_request(
model: &str,
document: OcrDocument,
params: &OpaqueParams,
) -> Result<MistralOcrRequest, OcrRequestError> {
Ok(MistralOcrRequest {
model: model.to_string(),
document,
params: params.clone(),
})
}
pub(crate) fn transform_ocr_response(
model: &str,
response: MistralOcrResponse,
@ -51,8 +46,109 @@ pub(crate) fn transform_ocr_response(
})
}
#[derive(Clone, Debug, Default)]
pub(crate) struct MistralOCRConfig;
impl MistralOCRConfig {
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
params: &OpaqueParams,
) -> Result<MistralOcrRequest, OcrRequestError> {
Ok(MistralOcrRequest {
model: model.to_string(),
document,
params: params.clone(),
})
}
}
impl BaseOcrConfig for MistralOCRConfig {
type ProviderResponse = MistralOcrResponse;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&[
"pages",
"include_image_base64",
"image_limit",
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
"document_annotation_prompt",
"extract_header",
"extract_footer",
"table_format",
"confidence_scores_granularity",
"include_blocks",
"id",
]
}
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = self.map_ocr_params(&request.model, &request.optional_params);
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref())?;
let body = self.transform_ocr_request(&request.model, request.document.clone(), &params)?;
transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: MistralOcrResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_ocr_response(&request.model, response)
}
}
pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result<String, OcrError> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(MISTRAL_OCR_API_BASE);
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v1", "ocr"]))
.map(|url| url.into_string())
.map_err(|_| {
OcrRequestError::RequestField {
path: "api_base".into(),
}
.into()
})
}
fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let api_key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.ok_or(Error::MissingApiKey {
provider: "Mistral",
})?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
#[cfg(test)]
mod mapping_tests {
mod tests {
use super::*;
use rstest::rstest;
use serde_json::{Value, json};
@ -181,9 +277,12 @@ mod mapping_tests {
#[case("id", json!("req-123"))]
fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: OpaqueParams = serde_json::from_value(json!({name: value.clone()})).unwrap();
let result =
serde_json::to_value(transform_ocr_request("model", document(), &params).unwrap())
.unwrap();
let result = serde_json::to_value(
MistralOCRConfig
.transform_ocr_request("model", document(), &params)
.unwrap(),
)
.unwrap();
assert_eq!(result["model"], "model");
assert_eq!(result[name], value);
}
@ -202,7 +301,9 @@ mod mapping_tests {
) {
let params: OpaqueParams = serde_json::from_value(json!({name:value.clone()})).unwrap();
let result = serde_json::to_value(
transform_ocr_request("mistral-ocr-latest", document(), &params).unwrap(),
MistralOCRConfig
.transform_ocr_request("mistral-ocr-latest", document(), &params)
.unwrap(),
)
.unwrap();
assert_eq!(result[name], value);
@ -218,7 +319,9 @@ mod mapping_tests {
}))
.unwrap();
let result = serde_json::to_value(
transform_ocr_request("mistral-ocr-latest", document(), &params).unwrap(),
MistralOCRConfig
.transform_ocr_request("mistral-ocr-latest", document(), &params)
.unwrap(),
)
.unwrap();
assert_eq!(result["table_format"], "html");
@ -279,118 +382,6 @@ mod mapping_tests {
.into_json();
assert_eq!(result["pages"][0], page);
}
}
use crate::Error;
use crate::constants::MISTRAL_OCR_API_BASE;
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::OcrClient;
use crate::ocr::error::OcrError;
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::types::{LiteLLMOcrRequest, OcrConnection};
use crate::url_utils::ApiUrl;
const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY";
#[derive(Clone, Debug, Default)]
pub(crate) struct MistralOCRConfig;
impl MistralOCRConfig {
pub(crate) fn transform_ocr_request(
&self,
model: &str,
document: crate::ocr::types::OcrDocument,
params: &OpaqueParams,
) -> Result<MistralOcrRequest, OcrRequestError> {
transform_ocr_request(model, document, params)
}
}
impl BaseOcrConfig for MistralOCRConfig {
type ProviderResponse = MistralOcrResponse;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&[
"pages",
"include_image_base64",
"image_limit",
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
"document_annotation_prompt",
"extract_header",
"extract_footer",
"table_format",
"confidence_scores_granularity",
"include_blocks",
"id",
]
}
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = self.map_ocr_params(&request.model, &request.optional_params);
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref())?;
let body = self.transform_ocr_request(&request.model, request.document.clone(), &params)?;
transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: MistralOcrResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_ocr_response(&request.model, response)
}
}
pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result<String, OcrError> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(MISTRAL_OCR_API_BASE);
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v1", "ocr"]))
.map(|url| url.into_string())
.map_err(|_| {
OcrRequestError::RequestField {
path: "api_base".into(),
}
.into()
})
}
fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let api_key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.ok_or(Error::MissingApiKey {
provider: "Mistral",
})?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn complete_url_defaults_and_dedupes_v1() {
@ -443,7 +434,7 @@ mod tests {
assert!(matches!(
validate_environment(&OcrConnection::default(), &|_| None),
Err(OcrError::Public(Error::MissingApiKey {
provider: "Mistral"
provider: "Mistral",
}))
));
}

View file

@ -1,540 +1,461 @@
mod types {
use serde::{Deserialize, Deserializer, Serialize};
use serde_json::{Map, Value};
use std::collections::BTreeMap;
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub(crate) struct ReductoV3Params {
#[serde(skip_serializing_if = "Option::is_none")]
pub formatting: Option<Map<String, Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub retrieval: Option<Map<String, Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub settings: Option<Map<String, Value>>,
use serde::{Deserialize, Deserializer, Serialize};
use serde_json::{Map, Value, json};
use crate::Error;
use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX};
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::OcrClient;
use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env,
guardrail_document, merge_extra_params,
};
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
use crate::url_utils::ApiUrl;
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
struct ReductoV3Params {
#[serde(skip_serializing_if = "Option::is_none")]
pub formatting: Option<Map<String, Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub retrieval: Option<Map<String, Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub settings: Option<Map<String, Value>>,
}
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
struct ReductoLegacyParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub enhance: Option<Map<String, Value>>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
struct ReductoV3Request {
pub input: String,
#[serde(flatten)]
pub params: ReductoV3Params,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
struct ReductoLegacyRequest {
pub document_url: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub options: Option<ReductoLegacyParams>,
}
#[derive(Deserialize)]
struct ReductoUploadResponse {
pub file_id: Option<String>,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct ReductoResponse {
#[serde(default, deserialize_with = "present_nullable")]
result: Option<Option<ReductoResult>>,
usage: Option<ReductoUsage>,
#[serde(default)]
chunks: Option<Vec<ReductoChunk>>,
}
fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
deserializer: D,
) -> Result<Option<Option<T>>, D::Error> {
Option::<T>::deserialize(deserializer).map(Some)
}
#[derive(Clone, Debug, Default, Deserialize)]
struct ReductoResult {
pub chunks: Option<Vec<ReductoChunk>>,
}
#[derive(Clone, Debug, Default, Deserialize)]
struct ReductoUsage {
#[serde(default, deserialize_with = "optional_i64")]
pub num_pages: Option<i64>,
#[serde(default, deserialize_with = "optional_f64")]
pub credits: Option<f64>,
}
#[derive(Clone, Debug, Deserialize)]
struct ReductoChunk {
pub content: Option<String>,
pub blocks: Option<Vec<ReductoBlock>>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
struct ReductoBlock {
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub bbox: Option<ReductoBoundingBox>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
struct ReductoBoundingBox {
#[serde(default, deserialize_with = "optional_i64")]
pub page: Option<i64>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<i64>, D::Error> {
match Option::<Value>::deserialize(deserializer)? {
None | Some(Value::Null) => Ok(None),
Some(Value::Number(number)) => number
.as_i64()
.or_else(|| number.as_f64().and_then(checked_truncated_i64))
.map(Some)
.ok_or_else(|| serde::de::Error::custom("expected an integer")),
Some(Value::String(value)) => value
.trim()
.parse::<i64>()
.map(Some)
.map_err(|_| serde::de::Error::custom("expected an integer")),
Some(Value::Bool(value)) => Ok(Some(i64::from(value))),
Some(_) => Ok(None),
}
}
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub(crate) struct ReductoLegacyParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub enhance: Option<Map<String, Value>>,
fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<f64>, D::Error> {
match Option::<Value>::deserialize(deserializer)? {
None | Some(Value::Null) => Ok(None),
Some(Value::Number(number)) => number
.as_f64()
.map(Some)
.ok_or_else(|| serde::de::Error::custom("expected a number")),
Some(Value::String(value)) => value
.trim()
.parse::<f64>()
.map(Some)
.map_err(|_| serde::de::Error::custom("expected a number")),
Some(_) => Ok(None),
}
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct ReductoV3Request {
pub input: String,
#[serde(flatten)]
pub params: ReductoV3Params,
fn checked_truncated_i64(value: f64) -> Option<i64> {
(value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64)
.then(|| value.trunc() as i64)
}
#[tracing::instrument(
name = "transform_ocr_request",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
fn transform_v3_ocr_request(
_model: &str,
document: OcrDocument,
params: &ReductoV3Params,
) -> Result<ReductoV3Request, OcrRequestError> {
Ok(ReductoV3Request {
input: document.source().to_string(),
params: params.clone(),
})
}
#[tracing::instrument(
name = "transform_ocr_request",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
fn transform_legacy_ocr_request(
_model: &str,
document: OcrDocument,
params: &ReductoLegacyParams,
) -> Result<ReductoLegacyRequest, OcrRequestError> {
Ok(ReductoLegacyRequest {
document_url: document.source().to_string(),
options: params.enhance.as_ref().map(|_| params.clone()),
})
}
pub(crate) fn transform_ocr_response(
model: &str,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
let result = match response.result {
Some(result) => result.unwrap_or_default(),
None => ReductoResult {
chunks: response.chunks,
},
};
let usage = response.usage.unwrap_or_default();
Ok(LiteLLMOcrResponse {
pages: build_pages(result.chunks.unwrap_or_default()),
model: model.to_string(),
document_annotation: None,
usage_info: Some(json!({
"pages_processed": usage.num_pages,
"credits": usage.credits,
})),
object: "ocr".to_string(),
extra_fields: serde_json::Map::new(),
provider_native_response: None,
})
}
fn build_pages(chunks: Vec<ReductoChunk>) -> Vec<Value> {
let blocks_by_page = chunks
.iter()
.flat_map(|chunk| chunk.blocks.iter().flatten())
.filter_map(|block| block.bbox.as_ref()?.page.map(|page| (page, block)))
.fold(
BTreeMap::<i64, Vec<&ReductoBlock>>::new(),
|mut pages, (page, block)| {
pages.entry(page).or_default().push(block);
pages
},
);
if blocks_by_page.is_empty() {
let markdown = join_content(chunks.iter().map(|chunk| chunk.content.as_deref()));
return if markdown.is_empty() {
Vec::new()
} else {
vec![page(0, markdown, None)]
};
}
blocks_by_page
.into_iter()
.map(|(index, blocks)| {
let markdown = join_content(blocks.iter().map(|block| block.content.as_deref()));
page(
index.saturating_sub(1).max(0),
markdown,
Some(json!(blocks)),
)
})
.collect()
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct ReductoLegacyRequest {
pub document_url: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub options: Option<ReductoLegacyParams>,
fn join_content<'a>(content: impl Iterator<Item = Option<&'a str>>) -> String {
content
.flatten()
.filter(|text| !text.is_empty())
.collect::<Vec<_>>()
.join("\n\n")
}
fn page(index: i64, markdown: String, blocks: Option<Value>) -> Value {
let mut result = json!({"index":index,"markdown":markdown,"images":null});
if let (Value::Object(fields), Some(blocks)) = (&mut result, blocks) {
fields.insert("blocks".into(), blocks);
}
result
}
fn get_complete_url(api_base: Option<&str>, path: &str) -> Result<String, OcrError> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(REDUCTO_API_BASE);
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&[path]))
.map(|url| url.into_string())
.map_err(|_| {
OcrRequestError::RequestField {
path: "api_base".into(),
}
.into()
})
}
#[derive(Deserialize)]
pub(crate) struct ReductoUploadResponse {
pub file_id: Option<String>,
fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let api_key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| {
env_lookup(REDUCTO_API_KEY_ENV)
.map(|key| key.trim().to_string())
.filter(|key| !key.is_empty())
})
.ok_or(Error::MissingReductoApiKey)?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct ReductoResponse {
#[serde(default, deserialize_with = "present_nullable")]
pub result: Option<Option<ReductoResult>>,
pub usage: Option<ReductoUsage>,
#[serde(default)]
pub chunks: Option<Vec<ReductoChunk>>,
}
fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
deserializer: D,
) -> Result<Option<Option<T>>, D::Error> {
Option::<T>::deserialize(deserializer).map(Some)
}
#[derive(Clone, Debug, Default, Deserialize)]
pub(crate) struct ReductoResult {
pub chunks: Option<Vec<ReductoChunk>>,
}
#[derive(Clone, Debug, Default, Deserialize)]
pub(crate) struct ReductoUsage {
#[serde(default, deserialize_with = "optional_i64")]
pub num_pages: Option<i64>,
#[serde(default, deserialize_with = "optional_f64")]
pub credits: Option<f64>,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct ReductoChunk {
pub content: Option<String>,
pub blocks: Option<Vec<ReductoBlock>>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct ReductoBlock {
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub bbox: Option<ReductoBoundingBox>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct ReductoBoundingBox {
#[serde(default, deserialize_with = "optional_i64")]
pub page: Option<i64>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<i64>, D::Error> {
match Option::<Value>::deserialize(deserializer)? {
None | Some(Value::Null) => Ok(None),
Some(Value::Number(number)) => number
.as_i64()
.or_else(|| number.as_f64().and_then(checked_truncated_i64))
.map(Some)
.ok_or_else(|| serde::de::Error::custom("expected an integer")),
Some(Value::String(value)) => value
.trim()
.parse::<i64>()
.map(Some)
.map_err(|_| serde::de::Error::custom("expected an integer")),
Some(Value::Bool(value)) => Ok(Some(i64::from(value))),
Some(_) => Ok(None),
async fn prepare_document(
client: &crate::ocr::OcrClient,
document: OcrDocument,
connection: &OcrConnection,
headers: &[(String, String)],
) -> Result<OcrDocument, OcrError> {
if document.source().starts_with(REDUCTO_ID_PREFIX) {
if document.source()[REDUCTO_ID_PREFIX.len()..]
.trim()
.is_empty()
{
return Err(OcrRequestError::RequestField {
path: "document file id".into(),
}
.into());
}
return Ok(document);
}
fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<f64>, D::Error> {
match Option::<Value>::deserialize(deserializer)? {
None | Some(Value::Null) => Ok(None),
Some(Value::Number(number)) => number
.as_f64()
.map(Some)
.ok_or_else(|| serde::de::Error::custom("expected a number")),
Some(Value::String(value)) => value
.trim()
.parse::<f64>()
.map(Some)
.map_err(|_| serde::de::Error::custom("expected a number")),
Some(_) => Ok(None),
let inline = InlineDocument::parse(document.source())?.ok_or(OcrRequestError::ReductoSource)?;
let mime = inline.mime_type().to_string();
let bytes = inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
let part = reqwest::multipart::Part::bytes(bytes)
.file_name("document")
.mime_str(&mime)
.map_err(|_| OcrRequestError::InvalidDataUri)?;
let builder = client
.provider_http()
.post(get_complete_url(connection.api_base.as_deref(), "upload")?)
.multipart(reqwest::multipart::Form::new().part("file", part))
.timeout(connection.timeout);
let builder = crate::http_utils::with_headers(
builder,
headers,
crate::http_utils::HeaderPolicy::Except(&["content-type", "content-length"]),
);
let response = crate::http_utils::http_request(builder)
.await
.map_err(crate::error::TransportError::from)?;
let uploaded = crate::ocr::client::read_json_response::<ReductoUploadResponse>(
response,
false,
connection.max_response_bytes,
)
.await?
.data;
let file_id = uploaded
.file_id
.as_deref()
.map(str::trim)
.filter(|id| !id.is_empty());
let Some(file_id) = file_id else {
return Err(OcrResponseError::ResponseField {
path: "file_id".into(),
}
.into());
};
Ok(document.with_source(file_id.to_string()))
}
#[derive(Clone, Debug)]
pub(crate) struct ReductoParseLegacyConfig;
impl BaseOcrConfig for ReductoParseLegacyConfig {
type ProviderResponse = ReductoResponse;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["enhance"]
}
fn checked_truncated_i64(value: f64) -> Option<i64> {
(value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64)
.then(|| value.trunc() as i64)
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoLegacyParams>(request)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document = prepare_document(client, document, &request.connection, &headers).await?;
let body = transform_legacy_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_ocr_response(&request.model, response)
}
}
#[derive(Clone, Debug)]
pub(crate) struct ReductoParseV3Config;
impl BaseOcrConfig for ReductoParseV3Config {
type ProviderResponse = ReductoResponse;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["formatting", "retrieval", "settings"]
}
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoV3Params>(request)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document = prepare_document(client, document, &request.connection, &headers).await?;
let body = transform_v3_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_ocr_response(&request.model, response)
}
}
#[cfg(test)]
pub(crate) use mapping::transform_ocr_response;
pub(crate) use types::ReductoResponse;
mod tests {
use super::*;
mod mapping {
use std::collections::BTreeMap;
use serde_json::{Value, json};
use super::types::*;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
#[tracing::instrument(
name = "transform_ocr_request",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
pub(crate) fn transform_v3_ocr_request(
_model: &str,
document: OcrDocument,
params: &ReductoV3Params,
) -> Result<ReductoV3Request, OcrRequestError> {
Ok(ReductoV3Request {
input: document.source().to_string(),
params: params.clone(),
})
}
#[tracing::instrument(
name = "transform_ocr_request",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
pub(crate) fn transform_legacy_ocr_request(
_model: &str,
document: OcrDocument,
params: &ReductoLegacyParams,
) -> Result<ReductoLegacyRequest, OcrRequestError> {
Ok(ReductoLegacyRequest {
document_url: document.source().to_string(),
options: params.enhance.as_ref().map(|_| params.clone()),
})
}
pub(crate) fn transform_ocr_response(
model: &str,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
let result = match response.result {
Some(result) => result.unwrap_or_default(),
None => ReductoResult {
chunks: response.chunks,
},
#[test]
fn explicit_key_precedes_environment_key() {
let connection = OcrConnection {
api_key: Some("passed-key".into()),
..Default::default()
};
let usage = response.usage.unwrap_or_default();
Ok(LiteLLMOcrResponse {
pages: build_pages(result.chunks.unwrap_or_default()),
model: model.to_string(),
document_annotation: None,
usage_info: Some(json!({
"pages_processed": usage.num_pages,
"credits": usage.credits,
})),
object: "ocr".to_string(),
extra_fields: serde_json::Map::new(),
provider_native_response: None,
})
let headers = validate_environment(&connection, &|_| Some("env-key".into())).unwrap();
assert_eq!(headers[0].1, "Bearer passed-key");
}
fn build_pages(chunks: Vec<ReductoChunk>) -> Vec<Value> {
let blocks_by_page = chunks
.iter()
.flat_map(|chunk| chunk.blocks.iter().flatten())
.filter_map(|block| block.bbox.as_ref()?.page.map(|page| (page, block)))
.fold(
BTreeMap::<i64, Vec<&ReductoBlock>>::new(),
|mut pages, (page, block)| {
pages.entry(page).or_default().push(block);
pages
},
);
if blocks_by_page.is_empty() {
let markdown = join_content(chunks.iter().map(|chunk| chunk.content.as_deref()));
return if markdown.is_empty() {
Vec::new()
} else {
vec![page(0, markdown, None)]
};
}
blocks_by_page
.into_iter()
.map(|(index, blocks)| {
let markdown = join_content(blocks.iter().map(|block| block.content.as_deref()));
page(
index.saturating_sub(1).max(0),
markdown,
Some(json!(blocks)),
)
})
.collect()
#[test]
fn blank_explicit_key_uses_environment_key() {
let connection = OcrConnection {
api_key: Some(" ".into()),
..Default::default()
};
let headers = validate_environment(&connection, &|_| Some(" env-key ".into())).unwrap();
assert_eq!(headers[0].1, "Bearer env-key");
}
fn join_content<'a>(content: impl Iterator<Item = Option<&'a str>>) -> String {
content
.flatten()
.filter(|text| !text.is_empty())
.collect::<Vec<_>>()
.join("\n\n")
}
fn page(index: i64, markdown: String, blocks: Option<Value>) -> Value {
let mut result = json!({"index":index,"markdown":markdown,"images":null});
if let (Value::Object(fields), Some(blocks)) = (&mut result, blocks) {
fields.insert("blocks".into(), blocks);
}
result
}
}
mod common {
use crate::Error;
use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX};
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::types::{OcrConnection, OcrDocument};
use crate::url_utils::ApiUrl;
pub(crate) trait BaseReductoOcrConfig: BaseOcrConfig {
fn get_complete_url(&self, api_base: Option<&str>, path: &str) -> Result<String, OcrError> {
get_complete_url(api_base, path)
}
fn validate_environment(
&self,
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
validate_environment(connection, env_lookup)
}
async fn prepare_document(
&self,
client: &crate::ocr::OcrClient,
document: OcrDocument,
connection: &OcrConnection,
headers: &[(String, String)],
) -> Result<OcrDocument, OcrError> {
prepare_document(client, document, connection, headers).await
}
}
pub(super) fn get_complete_url(api_base: Option<&str>, path: &str) -> Result<String, OcrError> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(REDUCTO_API_BASE);
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&[path]))
.map(|url| url.into_string())
.map_err(|_| {
OcrRequestError::RequestField {
path: "api_base".into(),
}
.into()
})
}
pub(super) fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let api_key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| {
env_lookup(REDUCTO_API_KEY_ENV)
.map(|key| key.trim().to_string())
.filter(|key| !key.is_empty())
})
.ok_or(Error::MissingReductoApiKey)?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
pub(super) async fn prepare_document(
client: &crate::ocr::OcrClient,
document: OcrDocument,
connection: &OcrConnection,
headers: &[(String, String)],
) -> Result<OcrDocument, OcrError> {
if document.source().starts_with(REDUCTO_ID_PREFIX) {
if document.source()[REDUCTO_ID_PREFIX.len()..]
.trim()
.is_empty()
{
return Err(OcrRequestError::RequestField {
path: "document file id".into(),
}
.into());
}
return Ok(document);
}
let inline =
InlineDocument::parse(document.source())?.ok_or(OcrRequestError::ReductoSource)?;
let mime = inline.mime_type().to_string();
let bytes = inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
let part = reqwest::multipart::Part::bytes(bytes)
.file_name("document")
.mime_str(&mime)
.map_err(|_| OcrRequestError::InvalidDataUri)?;
let builder = client
.provider_http()
.post(get_complete_url(connection.api_base.as_deref(), "upload")?)
.multipart(reqwest::multipart::Form::new().part("file", part))
.timeout(connection.timeout);
let builder = crate::http_utils::with_headers(
builder,
headers,
crate::http_utils::HeaderPolicy::Except(&["content-type", "content-length"]),
#[test]
fn existing_authorization_skips_key_lookup() {
let connection = OcrConnection {
extra_headers: vec![("authorization".into(), "Bearer existing".into())],
..Default::default()
};
assert_eq!(
validate_environment(&connection, &|_| None).unwrap(),
connection.extra_headers
);
let response = crate::http_utils::http_request(builder)
.await
.map_err(crate::error::TransportError::from)?;
let uploaded =
crate::ocr::client::read_json_response::<super::types::ReductoUploadResponse>(
response,
false,
connection.max_response_bytes,
)
.await?
.data;
let file_id = uploaded
.file_id
.as_deref()
.map(str::trim)
.filter(|id| !id.is_empty());
let Some(file_id) = file_id else {
return Err(OcrResponseError::ResponseField {
path: "file_id".into(),
}
.into());
};
Ok(document.with_source(file_id.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn explicit_key_precedes_environment_key() {
let connection = OcrConnection {
api_key: Some("passed-key".into()),
..Default::default()
};
let headers = validate_environment(&connection, &|_| Some("env-key".into())).unwrap();
assert_eq!(headers[0].1, "Bearer passed-key");
}
#[test]
fn blank_explicit_key_uses_environment_key() {
let connection = OcrConnection {
api_key: Some(" ".into()),
..Default::default()
};
let headers = validate_environment(&connection, &|_| Some(" env-key ".into())).unwrap();
assert_eq!(headers[0].1, "Bearer env-key");
}
#[test]
fn existing_authorization_skips_key_lookup() {
let connection = OcrConnection {
extra_headers: vec![("authorization".into(), "Bearer existing".into())],
..Default::default()
};
assert_eq!(
validate_environment(&connection, &|_| None).unwrap(),
connection.extra_headers
);
}
}
}
mod legacy {
use super::common::BaseReductoOcrConfig;
use super::mapping as reducto;
use super::types::{ReductoLegacyParams, ReductoResponse};
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::OcrClient;
use crate::ocr::error::{OcrError, OcrResponseError};
use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env,
guardrail_document, merge_extra_params,
};
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
#[derive(Clone, Debug)]
pub(crate) struct ReductoParseLegacyConfig;
impl BaseReductoOcrConfig for ReductoParseLegacyConfig {}
impl BaseOcrConfig for ReductoParseLegacyConfig {
type ProviderResponse = ReductoResponse;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["enhance"]
}
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoLegacyParams>(request)?;
let headers = self.validate_environment(&request.connection, &credential_env)?;
let url = self.get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document = self
.prepare_document(client, document, &request.connection, &headers)
.await?;
let body = reducto::transform_legacy_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
reducto::transform_ocr_response(&request.model, response)
}
}
}
pub(crate) use legacy::ReductoParseLegacyConfig;
mod v3 {
use super::common::BaseReductoOcrConfig;
use super::mapping as reducto;
use super::types::{ReductoResponse, ReductoV3Params};
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::OcrClient;
use crate::ocr::error::{OcrError, OcrResponseError};
use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env,
guardrail_document, merge_extra_params,
};
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
#[derive(Clone, Debug)]
pub(crate) struct ReductoParseV3Config;
impl BaseReductoOcrConfig for ReductoParseV3Config {}
impl BaseOcrConfig for ReductoParseV3Config {
type ProviderResponse = ReductoResponse;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["formatting", "retrieval", "settings"]
}
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoV3Params>(request)?;
let headers = self.validate_environment(&request.connection, &credential_env)?;
let url = self.get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document = self
.prepare_document(client, document, &request.connection, &headers)
.await?;
let body = reducto::transform_v3_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
reducto::transform_ocr_response(&request.model, response)
}
}
}
pub(crate) use v3::ReductoParseV3Config;

View file

@ -21,7 +21,12 @@ pub fn resolve_anthropic_api_key(
non_empty(api_key)
.map(str::to_string)
.or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty()))
.ok_or_else(|| Error::from(crate::AuthError::MissingAnthropicApiKey))
.ok_or_else(|| {
Error::from(crate::AuthError::MissingApiKey {
provider: "Anthropic",
environment_variable: ANTHROPIC_API_KEY_ENV,
})
})
}
pub fn complete_anthropic_url(
@ -113,10 +118,12 @@ mod tests {
resolve_anthropic_api_key(Some(" "), &with_env).unwrap(),
"sk-env"
);
assert!(matches!(
resolve_anthropic_api_key(None, &|_| None).expect_err("missing key"),
Error::Auth(_)
));
assert_eq!(
resolve_anthropic_api_key(None, &|_| None)
.expect_err("missing key")
.to_string(),
"Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable"
);
}
#[test]

View file

@ -32,7 +32,12 @@ pub fn resolve_azure_api_key(
non_empty(api_key)
.map(str::to_string)
.or_else(|| env_lookup(AZURE_API_KEY_ENV).filter(|value| !value.trim().is_empty()))
.ok_or_else(|| Error::from(crate::AuthError::MissingAzureApiKey))
.ok_or_else(|| {
Error::from(crate::AuthError::MissingApiKey {
provider: "Azure",
environment_variable: AZURE_API_KEY_ENV,
})
})
}
pub fn complete_azure_anthropic_url(
@ -271,10 +276,12 @@ mod tests {
resolve_azure_api_key(Some(" "), &with_env).unwrap(),
"sk-env"
);
assert!(matches!(
resolve_azure_api_key(None, &|_| None).expect_err("missing key"),
Error::Auth(_)
));
assert_eq!(
resolve_azure_api_key(None, &|_| None)
.expect_err("missing key")
.to_string(),
"Missing Azure API Key - Set `api_key` or the AZURE_API_KEY environment variable"
);
}
#[test]