diff --git a/litellm-rust/crates/auth/src/error.rs b/litellm-rust/crates/auth/src/error.rs index 4bbd707b16e..914265ffb32 100644 --- a/litellm-rust/crates/auth/src/error.rs +++ b/litellm-rust/crates/auth/src/error.rs @@ -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://.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" + ); + } +} diff --git a/litellm-rust/crates/core/AGENTS.md b/litellm-rust/crates/core/AGENTS.md index 9ba7bfb5323..7a06da63291 100644 --- a/litellm-rust/crates/core/AGENTS.md +++ b/litellm-rust/crates/core/AGENTS.md @@ -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. diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index 74291a03547..641eab8f043 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -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 for Error { impl From 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" ); } diff --git a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs index 1e822c216e2..f30abfffbf2 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -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, - meta: Option, - } +#[derive(Deserialize)] +pub(crate) struct CohereResponse { + #[serde(default)] + pages: Vec, + meta: Option, +} - #[derive(Deserialize)] - struct CoherePage { - index: Option, - markdown: Option, - blocks: Option>>, - } +#[derive(Deserialize)] +struct CoherePage { + index: Option, + markdown: Option, + blocks: Option>>, +} - #[derive(Deserialize)] - struct CohereMarkdown { - #[serde(default)] - content: String, - images: Option>>, - } +#[derive(Deserialize)] +struct CohereMarkdown { + #[serde(default)] + content: String, + images: Option>>, +} - #[derive(Deserialize)] - struct CohereMeta { - billed_units: Option, - } +#[derive(Deserialize)] +struct CohereMeta { + billed_units: Option, +} - #[derive(Deserialize)] - struct CohereBilledUnits { - pages: Option, - } +#[derive(Deserialize)] +struct CohereBilledUnits { + pages: Option, +} - pub(crate) fn transform_response( - model: &str, - response: CohereResponse, - ) -> Result { - 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 { + 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::>() - }); - (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::, 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::>() + }); + (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::, 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::(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::(json!({"output_format":"html"})).is_err() - ); - for format in ["markdown", "blocks"] { - assert!( - serde_json::from_value::(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 { - 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::(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::(json!({"output_format":"html"})).is_err()); + for format in ["markdown", "blocks"] { + assert!( + serde_json::from_value::(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"] { diff --git a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs index 46f002b92be..6efe72cf19b 100644 --- a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs @@ -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, } -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub(crate) fn transform_ocr_request( - model: &str, - document: OcrDocument, - params: &OpaqueParams, -) -> Result { - 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 { + 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 { + 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(), ¶ms)?; + transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: MistralOcrResponse, + ) -> Result { + transform_ocr_response(&request.model, response) + } +} + +pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result { + 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 + Sync), +) -> Result, 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(), ¶ms).unwrap()) - .unwrap(); + let result = serde_json::to_value( + MistralOCRConfig + .transform_ocr_request("model", document(), ¶ms) + .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(), ¶ms).unwrap(), + MistralOCRConfig + .transform_ocr_request("mistral-ocr-latest", document(), ¶ms) + .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(), ¶ms).unwrap(), + MistralOCRConfig + .transform_ocr_request("mistral-ocr-latest", document(), ¶ms) + .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 { - 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 { - 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(), ¶ms)?; - transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: MistralOcrResponse, - ) -> Result { - transform_ocr_response(&request.model, response) - } -} - -pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result { - 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 + Sync), -) -> Result, 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", })) )); } diff --git a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs index daa9a61e781..95199b60b2d 100644 --- a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -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>, - #[serde(skip_serializing_if = "Option::is_none")] - pub retrieval: Option>, - #[serde(skip_serializing_if = "Option::is_none")] - pub settings: Option>, +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>, + #[serde(skip_serializing_if = "Option::is_none")] + pub retrieval: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub settings: Option>, +} + +#[derive(Clone, Debug, Default, Serialize, Deserialize)] +struct ReductoLegacyParams { + #[serde(skip_serializing_if = "Option::is_none")] + pub enhance: Option>, +} + +#[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, +} + +#[derive(Deserialize)] +struct ReductoUploadResponse { + pub file_id: Option, +} + +#[derive(Clone, Debug, Deserialize)] +pub(crate) struct ReductoResponse { + #[serde(default, deserialize_with = "present_nullable")] + result: Option>, + usage: Option, + #[serde(default)] + chunks: Option>, +} + +fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>( + deserializer: D, +) -> Result>, D::Error> { + Option::::deserialize(deserializer).map(Some) +} + +#[derive(Clone, Debug, Default, Deserialize)] +struct ReductoResult { + pub chunks: Option>, +} + +#[derive(Clone, Debug, Default, Deserialize)] +struct ReductoUsage { + #[serde(default, deserialize_with = "optional_i64")] + pub num_pages: Option, + #[serde(default, deserialize_with = "optional_f64")] + pub credits: Option, +} + +#[derive(Clone, Debug, Deserialize)] +struct ReductoChunk { + pub content: Option, + pub blocks: Option>, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +struct ReductoBlock { + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub bbox: Option, + #[serde(flatten)] + pub extra_fields: Map, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +struct ReductoBoundingBox { + #[serde(default, deserialize_with = "optional_i64")] + pub page: Option, + #[serde(flatten)] + pub extra_fields: Map, +} + +fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::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::() + .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>, +fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::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::() + .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 { + (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 { + 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 { + 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 { + 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) -> Vec { + 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::>::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, +fn join_content<'a>(content: impl Iterator>) -> String { + content + .flatten() + .filter(|text| !text.is_empty()) + .collect::>() + .join("\n\n") +} + +fn page(index: i64, markdown: String, blocks: Option) -> 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 { + 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, +fn validate_environment( + connection: &OcrConnection, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result, 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>, - pub usage: Option, - #[serde(default)] - pub chunks: Option>, - } - - fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>( - deserializer: D, - ) -> Result>, D::Error> { - Option::::deserialize(deserializer).map(Some) - } - - #[derive(Clone, Debug, Default, Deserialize)] - pub(crate) struct ReductoResult { - pub chunks: Option>, - } - - #[derive(Clone, Debug, Default, Deserialize)] - pub(crate) struct ReductoUsage { - #[serde(default, deserialize_with = "optional_i64")] - pub num_pages: Option, - #[serde(default, deserialize_with = "optional_f64")] - pub credits: Option, - } - - #[derive(Clone, Debug, Deserialize)] - pub(crate) struct ReductoChunk { - pub content: Option, - pub blocks: Option>, - } - - #[derive(Clone, Debug, Serialize, Deserialize)] - pub(crate) struct ReductoBlock { - #[serde(skip_serializing_if = "Option::is_none")] - pub content: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub bbox: Option, - #[serde(flatten)] - pub extra_fields: Map, - } - - #[derive(Clone, Debug, Serialize, Deserialize)] - pub(crate) struct ReductoBoundingBox { - #[serde(default, deserialize_with = "optional_i64")] - pub page: Option, - #[serde(flatten)] - pub extra_fields: Map, - } - - fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { - match Option::::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::() - .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 { + 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, D::Error> { - match Option::::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::() - .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::( + 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 { - (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 { + let ParsedProviderParams { + known: params, + extra_params, + } = _prepare_ocr_request::(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, ¶ms)?; + let body = merge_extra_params(&body, extra_params)?; + build_http_request(client, request, &url, &headers, &body) + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: ReductoResponse, + ) -> Result { + 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 { + let ParsedProviderParams { + known: params, + extra_params, + } = _prepare_ocr_request::(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, ¶ms)?; + let body = merge_extra_params(&body, extra_params)?; + build_http_request(client, request, &url, &headers, &body) + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: ReductoResponse, + ) -> Result { + 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 { - 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 { - 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 { - 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) -> Vec { - 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::>::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>) -> String { - content - .flatten() - .filter(|text| !text.is_empty()) - .collect::>() - .join("\n\n") - } - - fn page(index: i64, markdown: String, blocks: Option) -> 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 { - get_complete_url(api_base, path) - } - - fn validate_environment( - &self, - connection: &OcrConnection, - env_lookup: &(dyn Fn(&str) -> Option + Sync), - ) -> Result, OcrError> { - validate_environment(connection, env_lookup) - } - - async fn prepare_document( - &self, - client: &crate::ocr::OcrClient, - document: OcrDocument, - connection: &OcrConnection, - headers: &[(String, String)], - ) -> Result { - prepare_document(client, document, connection, headers).await - } - } - - pub(super) fn get_complete_url(api_base: Option<&str>, path: &str) -> Result { - 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 + Sync), - ) -> Result, 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 { - 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::( - 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 { - let ParsedProviderParams { - known: params, - extra_params, - } = _prepare_ocr_request::(request)?; - let headers = self.validate_environment(&request.connection, &credential_env)?; - let url = self.get_complete_url(request.connection.api_base.as_deref(), "parse")?; - let (document, headers) = guardrail_document(request, &url, &headers).await?; - let document = self - .prepare_document(client, document, &request.connection, &headers) - .await?; - let body = reducto::transform_legacy_ocr_request(&request.model, document, ¶ms)?; - let body = merge_extra_params(&body, extra_params)?; - build_http_request(client, request, &url, &headers, &body) - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: ReductoResponse, - ) -> Result { - 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 { - let ParsedProviderParams { - known: params, - extra_params, - } = _prepare_ocr_request::(request)?; - let headers = self.validate_environment(&request.connection, &credential_env)?; - let url = self.get_complete_url(request.connection.api_base.as_deref(), "parse")?; - let (document, headers) = guardrail_document(request, &url, &headers).await?; - let document = self - .prepare_document(client, document, &request.connection, &headers) - .await?; - let body = reducto::transform_v3_ocr_request(&request.model, document, ¶ms)?; - let body = merge_extra_params(&body, extra_params)?; - build_http_request(client, request, &url, &headers, &body) - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: ReductoResponse, - ) -> Result { - reducto::transform_ocr_response(&request.model, response) - } - } -} - -pub(crate) use v3::ReductoParseV3Config; diff --git a/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs b/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs index 9e8cfd3df5c..f9171e21cf9 100644 --- a/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs @@ -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] diff --git a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs index c534e20d9f7..1bdfe4b3b64 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs @@ -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]