mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
more fixing
This commit is contained in:
parent
3a64b12911
commit
31b48f6191
19 changed files with 526 additions and 250 deletions
|
|
@ -1,11 +1,11 @@
|
|||
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
|
||||
use crate::llms::cohere::ocr::transformation::{CohereParseConfig, CohereRequest};
|
||||
use crate::llms::cohere::ocr::{CohereParams, CohereResponse, validate_document};
|
||||
use crate::ocr::OcrArguments;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::document::{inline_remote_document, validate_inline_document};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument};
|
||||
use crate::params::OpaqueParams;
|
||||
use crate::url_utils::ApiUrl;
|
||||
use litellm_auth_azure::AzureAuthInputs;
|
||||
|
||||
|
|
@ -25,10 +25,10 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OpaqueParams,
|
||||
optional_params: &OpaqueParams,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<CohereParams, crate::ocr::Error> {
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
CohereParseConfig.map_ocr_params(non_default_params, optional_params, model)
|
||||
}
|
||||
|
||||
|
|
@ -65,11 +65,7 @@ impl AzureAICohereParseConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.map_ocr_params(
|
||||
&request.optional_params,
|
||||
&OpaqueParams::default(),
|
||||
&request.model,
|
||||
)?;
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let config = AzureAuthInputs {
|
||||
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
|
||||
..AzureAuthInputs::from_sourced_optional_params(
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ use base64::{Engine, engine::general_purpose::STANDARD};
|
|||
use reqwest::Url;
|
||||
use serde::{Deserialize, Deserializer, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
use serde_with::serde_as;
|
||||
use tokio::time::Instant;
|
||||
|
||||
use litellm_auth::{InputSource, Sourced};
|
||||
|
|
@ -18,6 +19,7 @@ use crate::constants::{
|
|||
use crate::llms::base_llm::ocr::transformation::{
|
||||
BaseOcrConfig, OcrRequestContext, OcrResponseContext,
|
||||
};
|
||||
use crate::ocr::OcrArguments;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::client::read_json_response;
|
||||
use crate::ocr::document::InlineDocument;
|
||||
|
|
@ -29,6 +31,7 @@ use crate::ocr::types::{
|
|||
};
|
||||
use crate::ocr::wire::DecodedOcrResponse;
|
||||
use crate::params::OpaqueParams;
|
||||
use crate::serde_compat::{FiniteF64, LaxI64};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
|
|
@ -127,13 +130,15 @@ struct AzureDocumentIntelligenceAnalyzeResult {
|
|||
pub key_value_pairs: Option<Vec<Map<String, Value>>>,
|
||||
}
|
||||
|
||||
#[serde_as]
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
struct AzureDocumentIntelligencePage {
|
||||
#[serde(rename = "pageNumber", default, deserialize_with = "optional_i64")]
|
||||
#[serde(rename = "pageNumber")]
|
||||
#[serde_as(deserialize_as = "Option<LaxI64>")]
|
||||
pub page_number: Option<i64>,
|
||||
#[serde(default, deserialize_with = "optional_f64")]
|
||||
#[serde_as(deserialize_as = "Option<FiniteF64>")]
|
||||
pub width: Option<f64>,
|
||||
#[serde(default, deserialize_with = "optional_f64")]
|
||||
#[serde_as(deserialize_as = "Option<FiniteF64>")]
|
||||
pub height: Option<f64>,
|
||||
pub unit: Option<String>,
|
||||
#[serde(default)]
|
||||
|
|
@ -145,39 +150,6 @@ struct AzureDocumentIntelligenceLine {
|
|||
pub content: Option<String>,
|
||||
}
|
||||
|
||||
fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<i64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_i64()
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected an integer")),
|
||||
Some(Value::String(value)) => value
|
||||
.parse::<i64>()
|
||||
.map(Some)
|
||||
.map_err(|_| serde::de::Error::custom("expected an integer")),
|
||||
Some(_) => Err(serde::de::Error::custom("expected an integer")),
|
||||
}
|
||||
}
|
||||
|
||||
fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<f64>, D::Error> {
|
||||
match Option::<Value>::deserialize(deserializer)? {
|
||||
None | Some(Value::Null) => Ok(None),
|
||||
Some(Value::Number(number)) => number
|
||||
.as_f64()
|
||||
.filter(|value| value.is_finite())
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a finite number")),
|
||||
Some(Value::String(value)) => value
|
||||
.parse::<f64>()
|
||||
.ok()
|
||||
.filter(|value| value.is_finite())
|
||||
.map(Some)
|
||||
.ok_or_else(|| serde::de::Error::custom("expected a finite number")),
|
||||
Some(_) => Err(serde::de::Error::custom("expected a number")),
|
||||
}
|
||||
}
|
||||
|
||||
fn decode_input_params(
|
||||
params: Map<String, Value>,
|
||||
prefix: &str,
|
||||
|
|
@ -330,7 +302,7 @@ fn transform_completed_response(
|
|||
let pages = result
|
||||
.pages
|
||||
.into_iter()
|
||||
.map(normalize_page)
|
||||
.map(transform_azure_page)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
let pages_processed =
|
||||
i64::try_from(pages.len()).map_err(|_| crate::ocr::Error::NumericRange("pages"))?;
|
||||
|
|
@ -346,26 +318,16 @@ fn transform_completed_response(
|
|||
})
|
||||
}
|
||||
|
||||
fn normalize_page(page: AzureDocumentIntelligencePage) -> Result<OcrPage, crate::ocr::Error> {
|
||||
fn transform_azure_page(page: AzureDocumentIntelligencePage) -> Result<OcrPage, crate::ocr::Error> {
|
||||
let index = page
|
||||
.page_number
|
||||
.unwrap_or(1)
|
||||
.checked_sub(1)
|
||||
.ok_or(crate::ocr::Error::NumericRange("page.pageNumber"))?;
|
||||
let scale = if page.unit.as_deref().unwrap_or("inch") == "inch" {
|
||||
AZURE_DI_DEFAULT_DPI as f64
|
||||
} else {
|
||||
1.0
|
||||
};
|
||||
let width = pixel_dimension(
|
||||
let dimensions = convert_dimensions(
|
||||
page.width.unwrap_or(AZURE_DI_DEFAULT_WIDTH),
|
||||
scale,
|
||||
"page.width",
|
||||
)?;
|
||||
let height = pixel_dimension(
|
||||
page.height.unwrap_or(AZURE_DI_DEFAULT_HEIGHT),
|
||||
scale,
|
||||
"page.height",
|
||||
page.unit.as_deref().unwrap_or("inch"),
|
||||
)?;
|
||||
let markdown = page
|
||||
.lines
|
||||
|
|
@ -376,18 +338,31 @@ fn normalize_page(page: AzureDocumentIntelligencePage) -> Result<OcrPage, crate:
|
|||
Ok(OcrPage {
|
||||
index,
|
||||
markdown,
|
||||
dimensions: Some(OcrPageDimensions {
|
||||
width: Some(width),
|
||||
height: Some(height),
|
||||
dpi: Some(AZURE_DI_DEFAULT_DPI),
|
||||
}),
|
||||
dimensions: Some(dimensions),
|
||||
..Default::default()
|
||||
})
|
||||
}
|
||||
|
||||
fn convert_dimensions(
|
||||
width: f64,
|
||||
height: f64,
|
||||
unit: &str,
|
||||
) -> Result<OcrPageDimensions, crate::ocr::Error> {
|
||||
let scale = if unit == "inch" {
|
||||
AZURE_DI_DEFAULT_DPI as f64
|
||||
} else {
|
||||
1.0
|
||||
};
|
||||
Ok(OcrPageDimensions {
|
||||
width: Some(pixel_dimension(width, scale, "page.width")?),
|
||||
height: Some(pixel_dimension(height, scale, "page.height")?),
|
||||
dpi: Some(AZURE_DI_DEFAULT_DPI),
|
||||
})
|
||||
}
|
||||
|
||||
fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result<i64, crate::ocr::Error> {
|
||||
let value = value * scale;
|
||||
if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 {
|
||||
if !value.is_finite() || value < i64::MIN as f64 || value >= -(i64::MIN as f64) {
|
||||
return Err(crate::ocr::Error::NumericRange(field));
|
||||
}
|
||||
Ok(value.trunc() as i64)
|
||||
|
|
@ -510,10 +485,10 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OpaqueParams,
|
||||
optional_params: &OpaqueParams,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
_model: &str,
|
||||
) -> Result<DocumentIntelligenceParams, crate::ocr::Error> {
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
let mapped = normalize_ocr_params(decode_input_params(
|
||||
non_default_params
|
||||
.iter()
|
||||
|
|
@ -556,7 +531,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
|
|||
)
|
||||
}))
|
||||
.collect();
|
||||
crate::ocr::wire::decode_request_value(Value::Object(fields), "optional_params")
|
||||
Ok(fields)
|
||||
}
|
||||
|
||||
async fn async_transform_ocr_request(
|
||||
|
|
@ -617,11 +592,7 @@ impl AzureDocumentIntelligenceOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.map_ocr_params(
|
||||
&request.optional_params,
|
||||
&OpaqueParams::default(),
|
||||
&request.model,
|
||||
)?;
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let config = AzureAuthInputs {
|
||||
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
|
||||
..AzureAuthInputs::from_sourced_optional_params(
|
||||
|
|
@ -752,13 +723,31 @@ mod tests {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn mapping_preserves_supplied_options_when_overrides_are_empty() {
|
||||
let supplied =
|
||||
serde_json::from_value(json!({"pages":"4", "features":"languages", "extension":true}))
|
||||
.unwrap();
|
||||
fn empty_options_do_not_create_query_fields() {
|
||||
let overrides =
|
||||
serde_json::from_value(json!({"pages":[], "features":null, "req_format":"native"}))
|
||||
.unwrap();
|
||||
let mapped = AzureDocumentIntelligenceOCRConfig
|
||||
.parse_options(&overrides, "model")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(mapped).unwrap(),
|
||||
json!({
|
||||
"req_format":"native"
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mapping_preserves_supplied_options_when_overrides_are_empty() {
|
||||
let supplied = serde_json::from_value(json!({
|
||||
"pages":"4", "features":"languages", "extension":true
|
||||
}))
|
||||
.unwrap();
|
||||
let overrides = serde_json::from_value(json!({
|
||||
"pages":[], "features":null, "req_format":"native", "ignored":true
|
||||
}))
|
||||
.unwrap();
|
||||
let mapped = AzureDocumentIntelligenceOCRConfig
|
||||
.map_ocr_params(&overrides, &supplied, "model")
|
||||
.unwrap();
|
||||
|
|
@ -770,6 +759,20 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_numbers_follow_python_validation_before_dimension_conversion() {
|
||||
let response = AzureDocumentIntelligenceOCRConfig.decode_and_normalize_response(
|
||||
"model",
|
||||
br#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":2.0,"width":" 8.5 ","height":true}]}}"#,
|
||||
OcrResponseFormat::Litellm,
|
||||
).unwrap();
|
||||
assert_eq!(response.pages[0].index, 1);
|
||||
let dimensions = response.pages[0].dimensions.as_ref().unwrap();
|
||||
assert_eq!(dimensions.width, Some(816));
|
||||
assert_eq!(dimensions.height, Some(96));
|
||||
assert!(pixel_dimension(9_223_372_036_854_775_808.0, 1.0, "width").is_err());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(json!([0, 1, 2]), Some("1,2,3"))]
|
||||
#[case(json!([2, 0, 0, 1]), Some("1,2,3"))]
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ use crate::constants::AZURE_AI_OCR_PATH;
|
|||
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
|
||||
use crate::llms::mistral::ocr::MistralOcrResponse;
|
||||
use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest};
|
||||
use crate::ocr::OcrArguments;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::document::{inline_remote_document, validate_inline_document};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
|
|
@ -28,10 +29,10 @@ impl BaseOcrConfig for AzureAIOCRConfig {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OpaqueParams,
|
||||
optional_params: &OpaqueParams,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<OpaqueParams, crate::ocr::Error> {
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
MistralOCRConfig.map_ocr_params(non_default_params, optional_params, model)
|
||||
}
|
||||
|
||||
|
|
@ -67,11 +68,7 @@ impl AzureAIOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.map_ocr_params(
|
||||
&request.optional_params,
|
||||
&OpaqueParams::default(),
|
||||
&request.model,
|
||||
)?;
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let config = AzureAuthInputs {
|
||||
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
|
||||
..AzureAuthInputs::from_sourced_optional_params(
|
||||
|
|
|
|||
|
|
@ -4,14 +4,14 @@ use std::sync::Arc;
|
|||
use serde::Serialize;
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
use crate::ocr::OcrArguments;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::hooks::OcrHooks;
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrResponseFormat};
|
||||
use crate::params::OpaqueParams;
|
||||
|
||||
pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
|
||||
type OcrParams: DeserializeOwned + Send + Sync;
|
||||
type ProviderRequest: Serialize + DeserializeOwned + Send;
|
||||
type ProviderRequest: Serialize + Send;
|
||||
type ProviderResponse: DeserializeOwned + Send;
|
||||
|
||||
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
|
||||
|
|
@ -20,14 +20,20 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
_non_default_params: &OpaqueParams,
|
||||
optional_params: &OpaqueParams,
|
||||
_non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
_model: &str,
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
Ok(optional_params.clone())
|
||||
}
|
||||
|
||||
fn parse_options(
|
||||
&self,
|
||||
arguments: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<Self::OcrParams, crate::ocr::Error> {
|
||||
crate::ocr::wire::decode_request_value(
|
||||
serde_json::Value::Object(optional_params.clone().into()),
|
||||
"optional_params",
|
||||
)
|
||||
self.map_ocr_params(arguments, &OcrArguments::default(), model)?
|
||||
.parse()
|
||||
}
|
||||
|
||||
fn async_transform_ocr_request(
|
||||
|
|
@ -45,6 +51,22 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
|
|||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, crate::ocr::Error>;
|
||||
|
||||
fn decode_and_normalize_response(
|
||||
&self,
|
||||
model: &str,
|
||||
raw_response: &[u8],
|
||||
request_format: OcrResponseFormat,
|
||||
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
|
||||
let decoded = crate::ocr::wire::decode_response::<Self::ProviderResponse>(
|
||||
raw_response,
|
||||
request_format == OcrResponseFormat::Native,
|
||||
)?;
|
||||
Ok(LiteLLMOcrResponse {
|
||||
provider_native_response: decoded.native,
|
||||
..self.normalize_response(model, decoded.data)?
|
||||
})
|
||||
}
|
||||
|
||||
fn async_transform_ocr_response(
|
||||
&self,
|
||||
model: &str,
|
||||
|
|
@ -58,14 +80,7 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
|
|||
)
|
||||
.await?;
|
||||
crate::ocr::handler::post_call(context.hooks, &bytes).await?;
|
||||
let decoded = crate::ocr::wire::decode_response::<Self::ProviderResponse>(
|
||||
&bytes,
|
||||
context.request_format == OcrResponseFormat::Native,
|
||||
)?;
|
||||
Ok(LiteLLMOcrResponse {
|
||||
provider_native_response: decoded.native,
|
||||
..self.normalize_response(model, decoded.data)?
|
||||
})
|
||||
self.decode_and_normalize_response(model, &bytes, context.request_format)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ use serde_with::serde_as;
|
|||
|
||||
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
|
||||
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
|
||||
use crate::ocr::OcrArguments;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::document::InlineDocument;
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
|
|
@ -12,7 +13,6 @@ use crate::ocr::types::{
|
|||
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageImage,
|
||||
OcrUsageInfo,
|
||||
};
|
||||
use crate::params::OpaqueParams;
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
|
||||
|
|
@ -27,17 +27,13 @@ pub(crate) enum OutputFormat {
|
|||
pub(crate) struct CohereParams {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_format: Option<OutputFormat>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub req_format: Option<crate::ocr::types::OcrResponseFormat>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: OpaqueParams,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, Serialize)]
|
||||
pub(crate) struct CohereRequest {
|
||||
pub model: String,
|
||||
pub document: CohereParseDocument,
|
||||
pub output_format: OutputFormat,
|
||||
pub output_format: String,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, Serialize)]
|
||||
|
|
@ -119,25 +115,25 @@ impl BaseOcrConfig for CohereParseConfig {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OpaqueParams,
|
||||
optional_params: &OpaqueParams,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
_model: &str,
|
||||
) -> Result<CohereParams, crate::ocr::Error> {
|
||||
if let Some(value) = non_default_params
|
||||
.get("req_format")
|
||||
.filter(|value| !value.is_null())
|
||||
{
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
let overrides: OcrArguments = non_default_params
|
||||
.select(&["output_format", "req_format"])
|
||||
.into_iter()
|
||||
.filter(|(_, value)| !value.is_null())
|
||||
.collect();
|
||||
overrides.parse::<CohereParams>()?;
|
||||
if let Some(value) = overrides.get("req_format") {
|
||||
serde_json::from_value::<crate::ocr::types::OcrResponseFormat>(value.clone())
|
||||
.map_err(|_| crate::ocr::Error::RequestFormat)?;
|
||||
}
|
||||
let fields = optional_params
|
||||
Ok(optional_params
|
||||
.iter()
|
||||
.chain(non_default_params.iter().filter(|(name, value)| {
|
||||
matches!(name.as_str(), "output_format" | "req_format") && !value.is_null()
|
||||
}))
|
||||
.chain(overrides.iter())
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect();
|
||||
crate::ocr::wire::decode_request_value(Value::Object(fields), "optional_params")
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn async_transform_ocr_request(
|
||||
|
|
@ -166,11 +162,7 @@ impl CohereParseConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.map_ocr_params(
|
||||
&request.optional_params,
|
||||
&OpaqueParams::default(),
|
||||
&request.model,
|
||||
)?;
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let headers = self.validate_environment(&request.connection, &credential_env)?;
|
||||
let url = self.get_complete_url(
|
||||
request
|
||||
|
|
@ -248,7 +240,11 @@ fn build_request(model: &str, image_url: String, params: &CohereParams) -> Coher
|
|||
CohereRequest {
|
||||
model: model.into(),
|
||||
document: CohereParseDocument::ImageUrl { image_url },
|
||||
output_format: params.output_format.unwrap_or_default(),
|
||||
output_format: match params.output_format.unwrap_or_default() {
|
||||
OutputFormat::Markdown => "markdown",
|
||||
OutputFormat::Blocks => "blocks",
|
||||
}
|
||||
.into(),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -362,6 +358,45 @@ mod tests {
|
|||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn mapping_merges_non_null_supported_overrides_and_preserves_supplied_options() {
|
||||
let supplied = serde_json::from_value(json!({
|
||||
"output_format":"blocks", "req_format":"native", "extension":false
|
||||
}))
|
||||
.unwrap();
|
||||
let overrides = serde_json::from_value(json!({
|
||||
"output_format":null, "req_format":null, "ignored":true
|
||||
}))
|
||||
.unwrap();
|
||||
for config in [false, true] {
|
||||
let mapped = if config {
|
||||
crate::llms::azure_ai::ocr::cohere_parse_transformation::AzureAICohereParseConfig
|
||||
.map_ocr_params(&overrides, &supplied, "parse")
|
||||
} else {
|
||||
CohereParseConfig.map_ocr_params(&overrides, &supplied, "parse")
|
||||
}
|
||||
.unwrap();
|
||||
assert_eq!(mapped, supplied);
|
||||
}
|
||||
for overrides in [
|
||||
json!({"output_format":"html"}),
|
||||
json!({"req_format":"invalid"}),
|
||||
] {
|
||||
let overrides = serde_json::from_value(overrides).unwrap();
|
||||
assert!(
|
||||
CohereParseConfig
|
||||
.map_ocr_params(&overrides, &supplied, "parse")
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
let overrides = serde_json::from_value(json!({"output_format":"markdown"})).unwrap();
|
||||
let mapped = CohereParseConfig
|
||||
.map_ocr_params(&overrides, &supplied, "parse")
|
||||
.unwrap();
|
||||
assert_eq!(mapped["output_format"], "markdown");
|
||||
assert_eq!(mapped["extension"], false);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn billed_pages_accept_integral_doubles_and_reject_fractional_counts() {
|
||||
let response = serde_json::from_str::<CohereResponse>(
|
||||
|
|
@ -422,19 +457,17 @@ mod tests {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn parameter_mapping_merges_supplied_options_and_ignores_null_overrides() {
|
||||
let supplied =
|
||||
serde_json::from_value(json!({"output_format":"blocks", "extension":true})).unwrap();
|
||||
let overrides = serde_json::from_value(
|
||||
json!({"output_format":null,"req_format":"native","unknown":true}),
|
||||
fn provider_options_exclude_response_controls_and_extensions() {
|
||||
let arguments = serde_json::from_value(
|
||||
json!({"output_format":"blocks","req_format":"native","unknown":true}),
|
||||
)
|
||||
.unwrap();
|
||||
let params = CohereParseConfig
|
||||
.map_ocr_params(&overrides, &supplied, "parse")
|
||||
.parse_options(&arguments, "parse")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(¶ms).unwrap(),
|
||||
json!({"output_format":"blocks","req_format":"native","extension":true})
|
||||
json!({"output_format":"blocks"})
|
||||
);
|
||||
let document = serde_json::from_value(
|
||||
json!({"type":"image_url","image_url":"https://example.com/a.png","ignored":"field"}),
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ use serde_json::Value;
|
|||
|
||||
use crate::constants::MISTRAL_OCR_API_BASE;
|
||||
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
|
||||
use crate::ocr::OcrArguments;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
use crate::ocr::types::{
|
||||
|
|
@ -79,16 +80,13 @@ impl BaseOcrConfig for MistralOCRConfig {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OpaqueParams,
|
||||
_optional_params: &OpaqueParams,
|
||||
non_default_params: &OcrArguments,
|
||||
_optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<OpaqueParams, crate::ocr::Error> {
|
||||
let supported = self.get_supported_ocr_params(model);
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
Ok(non_default_params
|
||||
.iter()
|
||||
.filter(|(name, _)| supported.contains(&name.as_str()))
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect())
|
||||
.select(self.get_supported_ocr_params(model))
|
||||
.into())
|
||||
}
|
||||
|
||||
async fn async_transform_ocr_request(
|
||||
|
|
@ -117,11 +115,7 @@ impl MistralOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.map_ocr_params(
|
||||
&request.optional_params,
|
||||
&OpaqueParams::default(),
|
||||
&request.model,
|
||||
)?;
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let headers = self.validate_environment(&request.connection, &credential_env)?;
|
||||
let url = self.get_complete_url(request.connection.api_base.as_deref())?;
|
||||
let body = self
|
||||
|
|
@ -285,7 +279,7 @@ mod tests {
|
|||
let input =
|
||||
serde_json::from_value(json!({"pages":null,"extract_header":false,"unknown":true}))
|
||||
.unwrap();
|
||||
let supplied = serde_json::from_value(json!({"pages":[1],"extract_footer":true})).unwrap();
|
||||
let supplied = serde_json::from_value(json!({"pages":[9],"extension":true})).unwrap();
|
||||
let params = MistralOCRConfig
|
||||
.map_ocr_params(&input, &supplied, "model")
|
||||
.unwrap();
|
||||
|
|
@ -295,11 +289,49 @@ mod tests {
|
|||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_transform_uses_already_mapped_params_without_filtering_again() {
|
||||
let params = serde_json::from_value(json!({"extension":{"nested":null}})).unwrap();
|
||||
let body = MistralOCRConfig
|
||||
.transform_ocr_request("model", document(), ¶ms, &[])
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(body).unwrap()["extension"],
|
||||
json!({"nested":null})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn raw_response_transform_keeps_native_payload_separate_from_typed_normalization() {
|
||||
let raw = br#"{"pages":[{"index":"2","markdown":"text"}],"provider_extension":false}"#;
|
||||
let response = MistralOCRConfig
|
||||
.decode_and_normalize_response(
|
||||
"model",
|
||||
raw,
|
||||
crate::ocr::types::OcrResponseFormat::Native,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(response.pages[0].index, 2);
|
||||
let native = response.provider_native_response.unwrap();
|
||||
assert_eq!(native["pages"][0]["index"], "2");
|
||||
assert_eq!(native["provider_extension"], false);
|
||||
assert!(response.extra_fields.is_empty());
|
||||
assert!(
|
||||
MistralOCRConfig
|
||||
.decode_and_normalize_response(
|
||||
"model",
|
||||
br#"{"pages":[{"index":0}]}"#,
|
||||
crate::ocr::types::OcrResponseFormat::Litellm
|
||||
)
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
fn mapped_params(value: Value) -> Value {
|
||||
let params = serde_json::from_value::<OpaqueParams>(value).unwrap();
|
||||
let params = serde_json::from_value::<OcrArguments>(value).unwrap();
|
||||
serde_json::to_value(
|
||||
MistralOCRConfig
|
||||
.map_ocr_params(¶ms, &OpaqueParams::default(), "model")
|
||||
.map_ocr_params(¶ms, &OcrArguments::default(), "model")
|
||||
.unwrap(),
|
||||
)
|
||||
.unwrap()
|
||||
|
|
|
|||
|
|
@ -5,11 +5,10 @@ use serde_json::{Map, Value, json};
|
|||
|
||||
use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX};
|
||||
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
|
||||
use crate::ocr::OcrArguments;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::document::InlineDocument;
|
||||
use crate::ocr::prepare::{
|
||||
build_http_request, credential_env, guardrail_document, merge_extra_params,
|
||||
};
|
||||
use crate::ocr::prepare::{build_http_request, credential_env, guardrail_document};
|
||||
use crate::ocr::types::{
|
||||
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrUsageInfo,
|
||||
};
|
||||
|
|
@ -90,15 +89,15 @@ impl BaseOcrConfig for ReductoParseV3Config {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OpaqueParams,
|
||||
optional_params: &OpaqueParams,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<ReductoV3Params, crate::ocr::Error> {
|
||||
map_ocr_params(
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
Ok(map_ocr_params(
|
||||
non_default_params,
|
||||
optional_params,
|
||||
self.get_supported_ocr_params(model),
|
||||
)
|
||||
))
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
|
|
@ -137,11 +136,7 @@ impl ReductoParseV3Config {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.map_ocr_params(
|
||||
&request.optional_params,
|
||||
&OpaqueParams::default(),
|
||||
&request.model,
|
||||
)?;
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let headers = validate_environment(&request.connection, &credential_env)?;
|
||||
let url = get_complete_url(request.connection.api_base.as_deref())?;
|
||||
let (document, headers) = guardrail_document(request, &url, &headers).await?;
|
||||
|
|
@ -157,10 +152,9 @@ impl ReductoParseV3Config {
|
|||
},
|
||||
)
|
||||
.await?;
|
||||
let extra_params = request
|
||||
let body = request
|
||||
.optional_params
|
||||
.without(self.get_supported_ocr_params(&request.model));
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
.compose_body(&body, self.get_supported_ocr_params(&request.model))?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
}
|
||||
|
|
@ -179,15 +173,15 @@ impl BaseOcrConfig for ReductoParseLegacyConfig {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OpaqueParams,
|
||||
optional_params: &OpaqueParams,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<ReductoLegacyParams, crate::ocr::Error> {
|
||||
map_ocr_params(
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
Ok(map_ocr_params(
|
||||
non_default_params,
|
||||
optional_params,
|
||||
self.get_supported_ocr_params(model),
|
||||
)
|
||||
))
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
|
|
@ -223,11 +217,7 @@ impl ReductoParseLegacyConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.map_ocr_params(
|
||||
&request.optional_params,
|
||||
&OpaqueParams::default(),
|
||||
&request.model,
|
||||
)?;
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let headers = validate_environment(&request.connection, &credential_env)?;
|
||||
let url = get_complete_url(request.connection.api_base.as_deref())?;
|
||||
let (document, headers) = guardrail_document(request, &url, &headers).await?;
|
||||
|
|
@ -243,20 +233,19 @@ impl ReductoParseLegacyConfig {
|
|||
},
|
||||
)
|
||||
.await?;
|
||||
let extra_params = request
|
||||
let body = request
|
||||
.optional_params
|
||||
.without(self.get_supported_ocr_params(&request.model));
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
.compose_body(&body, self.get_supported_ocr_params(&request.model))?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
}
|
||||
|
||||
fn map_ocr_params<T: serde::de::DeserializeOwned>(
|
||||
non_default_params: &OpaqueParams,
|
||||
optional_params: &OpaqueParams,
|
||||
fn map_ocr_params(
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
supported_params: &[&str],
|
||||
) -> Result<T, crate::ocr::Error> {
|
||||
let params = optional_params
|
||||
) -> OcrArguments {
|
||||
optional_params
|
||||
.iter()
|
||||
.chain(
|
||||
non_default_params
|
||||
|
|
@ -264,8 +253,7 @@ fn map_ocr_params<T: serde::de::DeserializeOwned>(
|
|||
.filter(|(name, _)| supported_params.contains(&name.as_str())),
|
||||
)
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect();
|
||||
crate::ocr::wire::decode_request_value(Value::Object(params), "optional_params")
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
|
||||
|
|
@ -510,6 +498,36 @@ async fn upload_bytes_async(
|
|||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn mapping_merges_supported_overrides_including_null_with_supplied_options() {
|
||||
let supplied = serde_json::from_value(json!({
|
||||
"formatting":{"old":true}, "enhance":true, "extension":false
|
||||
}))
|
||||
.unwrap();
|
||||
let overrides = serde_json::from_value(json!({
|
||||
"formatting":null, "enhance":null, "ignored":true
|
||||
}))
|
||||
.unwrap();
|
||||
let v3 = ReductoParseV3Config
|
||||
.map_ocr_params(&overrides, &supplied, "parse-v3")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(v3).unwrap(),
|
||||
json!({
|
||||
"formatting":null, "enhance":true, "extension":false
|
||||
})
|
||||
);
|
||||
let legacy = ReductoParseLegacyConfig
|
||||
.map_ocr_params(&overrides, &supplied, "parse-legacy")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(legacy).unwrap(),
|
||||
json!({
|
||||
"formatting":{"old":true}, "enhance":null, "extension":false
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_uses_shared_validation_while_block_page_numbers_are_best_effort() {
|
||||
for usage in [
|
||||
|
|
@ -535,16 +553,12 @@ mod tests {
|
|||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn v3_mapping_preserves_explicit_null_and_supplied_options() {
|
||||
async fn v3_options_preserve_explicit_null() {
|
||||
let overrides =
|
||||
serde_json::from_value(json!({"formatting":null,"settings":{},"unknown":true}))
|
||||
.unwrap();
|
||||
let supplied = serde_json::from_value(
|
||||
json!({"formatting":{"old":true},"retrieval":{"mode":"page"},"extension":false}),
|
||||
)
|
||||
.unwrap();
|
||||
let params = ReductoParseV3Config
|
||||
.map_ocr_params(&overrides, &supplied, "parse-v3")
|
||||
.parse_options(&overrides, "parse-v3")
|
||||
.unwrap();
|
||||
let client = crate::ocr::test_support::ocr_client();
|
||||
let connection = OcrConnection::default();
|
||||
|
|
@ -568,15 +582,11 @@ mod tests {
|
|||
assert_eq!(
|
||||
serde_json::to_value(body).unwrap(),
|
||||
json!({
|
||||
"input":"reducto://ready.pdf", "formatting":null, "settings":{}, "retrieval":{"mode":"page"}, "extension":false
|
||||
"input":"reducto://ready.pdf", "formatting":null, "settings":{}
|
||||
})
|
||||
);
|
||||
let absent = ReductoParseV3Config
|
||||
.map_ocr_params(
|
||||
&OpaqueParams::default(),
|
||||
&OpaqueParams::default(),
|
||||
"parse-v3",
|
||||
)
|
||||
.parse_options(&OcrArguments::default(), "parse-v3")
|
||||
.unwrap();
|
||||
assert_eq!(serde_json::to_value(absent).unwrap(), json!({}));
|
||||
}
|
||||
|
|
@ -592,9 +602,8 @@ mod tests {
|
|||
] {
|
||||
let overrides =
|
||||
serde_json::from_value(json!({"enhance":value,"unknown":true})).unwrap();
|
||||
let supplied = serde_json::from_value(json!({"enhance":{"old":true}})).unwrap();
|
||||
let params = ReductoParseLegacyConfig
|
||||
.map_ocr_params(&overrides, &supplied, "parse-legacy")
|
||||
.parse_options(&overrides, "parse-legacy")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(build_legacy_body(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use litellm_auth_gcp::{self as vertex, VertexConfig};
|
||||
|
||||
use super::transformation::VertexAIOCRConfig;
|
||||
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
|
||||
use crate::ocr::OcrClient;
|
||||
|
|
@ -12,7 +14,6 @@ use crate::ocr::types::{
|
|||
use crate::params::OpaqueParams;
|
||||
use crate::routing_utils::model::{ModelNamespace, ProviderModel, RoutedModel};
|
||||
use crate::url_utils::ApiUrl;
|
||||
use litellm_auth_gcp::{self as vertex, VertexConfig};
|
||||
|
||||
const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com";
|
||||
const MODEL_NAMESPACE: &str = "deepseek-ai";
|
||||
|
|
@ -153,11 +154,7 @@ impl VertexAIDeepSeekOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.map_ocr_params(
|
||||
&OpaqueParams::default(),
|
||||
&request.optional_params,
|
||||
&request.model,
|
||||
)?;
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let config = VertexConfig::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
|
|
@ -398,6 +395,21 @@ impl VertexAIDeepSeekOCRConfig {
|
|||
mod tests {
|
||||
use super::{VertexAIDeepSeekOCRConfig, provider_model};
|
||||
|
||||
#[test]
|
||||
fn inherited_parameter_mapping_only_returns_supplied_optional_params() {
|
||||
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
|
||||
use serde_json::json;
|
||||
|
||||
let non_default = serde_json::from_value(json!({"temperature":0.5})).unwrap();
|
||||
let supplied = serde_json::from_value(json!({"max_tokens":100,"extension":null})).unwrap();
|
||||
assert_eq!(
|
||||
VertexAIDeepSeekOCRConfig
|
||||
.map_ocr_params(&non_default, &supplied, "deepseek-ocr")
|
||||
.unwrap(),
|
||||
supplied
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn config_owns_model_namespace_and_endpoint() {
|
||||
assert_eq!(
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ use super::common_utils::validate_destination;
|
|||
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
|
||||
use crate::llms::mistral::ocr::MistralOcrResponse;
|
||||
use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest};
|
||||
use crate::ocr::OcrArguments;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::document::{inline_remote_document, validate_inline_document};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
|
|
@ -25,10 +26,10 @@ impl BaseOcrConfig for VertexAIOCRConfig {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OpaqueParams,
|
||||
optional_params: &OpaqueParams,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<OpaqueParams, crate::ocr::Error> {
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
MistralOCRConfig.map_ocr_params(non_default_params, optional_params, model)
|
||||
}
|
||||
|
||||
|
|
@ -64,11 +65,7 @@ impl VertexAIOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.map_ocr_params(
|
||||
&request.optional_params,
|
||||
&OpaqueParams::default(),
|
||||
&request.model,
|
||||
)?;
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let config = VertexConfig::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
|
|
|
|||
93
litellm-rust/crates/core/src/ocr/arguments.rs
Normal file
93
litellm-rust/crates/core/src/ocr/arguments.rs
Normal file
|
|
@ -0,0 +1,93 @@
|
|||
use std::ops::Deref;
|
||||
|
||||
use serde::{Deserialize, Serialize, de::DeserializeOwned};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
#[serde(transparent)]
|
||||
pub struct OcrArguments(Map<String, Value>);
|
||||
|
||||
impl OcrArguments {
|
||||
pub(crate) fn parse<T: DeserializeOwned>(&self) -> Result<T, super::Error> {
|
||||
super::wire::decode_request_value(Value::Object(self.0.clone()), "optional_params")
|
||||
}
|
||||
|
||||
pub(crate) fn select(&self, names: &[&str]) -> Map<String, Value> {
|
||||
self.iter()
|
||||
.filter(|(name, _)| names.contains(&name.as_str()))
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub(crate) fn compose_body<B: Serialize>(
|
||||
&self,
|
||||
body: &B,
|
||||
consumed: &[&str],
|
||||
) -> Result<Value, super::Error> {
|
||||
let Value::Object(fields) =
|
||||
serde_json::to_value(body).map_err(|_| crate::params::Error::Body)?
|
||||
else {
|
||||
return Err(crate::params::Error::Body.into());
|
||||
};
|
||||
let overrides = match self.get("extra_body") {
|
||||
None | Some(Value::Null) => None,
|
||||
Some(Value::Object(fields)) => Some(fields),
|
||||
Some(_) => return Err(crate::params::Error::ExtraBody.into()),
|
||||
};
|
||||
let extensions = self.iter().filter(|(name, _)| {
|
||||
!consumed.contains(&name.as_str())
|
||||
&& name.as_str() != "extra_body"
|
||||
&& !crate::params::is_control_param(name)
|
||||
});
|
||||
Ok(Value::Object(
|
||||
fields
|
||||
.into_iter()
|
||||
.chain(
|
||||
extensions
|
||||
.chain(overrides.into_iter().flatten())
|
||||
.filter(|(name, _)| {
|
||||
name.as_str() != "model"
|
||||
&& name.as_str() != "extra_body"
|
||||
&& !crate::params::is_control_param(name)
|
||||
})
|
||||
.map(|(name, value)| (name.clone(), value.clone())),
|
||||
)
|
||||
.collect(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
impl Deref for OcrArguments {
|
||||
type Target = Map<String, Value>;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl From<Map<String, Value>> for OcrArguments {
|
||||
fn from(values: Map<String, Value>) -> Self {
|
||||
Self(values)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<OcrArguments> for Map<String, Value> {
|
||||
fn from(arguments: OcrArguments) -> Self {
|
||||
arguments.0
|
||||
}
|
||||
}
|
||||
|
||||
impl FromIterator<(String, Value)> for OcrArguments {
|
||||
fn from_iter<T: IntoIterator<Item = (String, Value)>>(iter: T) -> Self {
|
||||
Self(iter.into_iter().collect())
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoIterator for OcrArguments {
|
||||
type Item = (String, Value);
|
||||
type IntoIter = serde_json::map::IntoIter;
|
||||
|
||||
fn into_iter(self) -> Self::IntoIter {
|
||||
self.0.into_iter()
|
||||
}
|
||||
}
|
||||
|
|
@ -1,4 +1,6 @@
|
|||
mod arguments;
|
||||
mod error;
|
||||
pub use arguments::OcrArguments;
|
||||
pub use error::Error;
|
||||
pub mod client;
|
||||
pub(crate) mod document;
|
||||
|
|
|
|||
|
|
@ -5,8 +5,6 @@ use super::OcrClient;
|
|||
use super::hooks::OcrDuringCallRequest;
|
||||
use super::types::{LiteLLMOcrRequest, OcrDocument};
|
||||
|
||||
pub(crate) use crate::params::merge_extra_params;
|
||||
|
||||
pub(crate) async fn transform_request_body<B>(
|
||||
client: &OcrClient,
|
||||
request: &LiteLLMOcrRequest,
|
||||
|
|
@ -19,10 +17,10 @@ pub(crate) async fn transform_request_body<B>(
|
|||
where
|
||||
B: Serialize + DeserializeOwned,
|
||||
{
|
||||
let extras = request
|
||||
.optional_params
|
||||
.without(request.config.get_supported_ocr_params(&request.model));
|
||||
let composed = merge_extra_params(&body, extras)?;
|
||||
let composed = request.optional_params.compose_body(
|
||||
&body,
|
||||
request.config.get_supported_ocr_params(&request.model),
|
||||
)?;
|
||||
let composed = OcrWireBody::<B>::decode(composed, "body")?;
|
||||
validate(&composed.body)?;
|
||||
let (body, headers) = if request.hooks.intercepts_requests() {
|
||||
|
|
@ -110,10 +108,10 @@ pub(crate) async fn guardrail_document(
|
|||
|
||||
#[derive(Serialize)]
|
||||
struct OcrWireBody<B> {
|
||||
#[serde(flatten)]
|
||||
#[serde(skip)]
|
||||
body: B,
|
||||
#[serde(flatten)]
|
||||
extra: crate::params::OpaqueParams,
|
||||
fields: serde_json::Map<String, Value>,
|
||||
}
|
||||
|
||||
impl<B: Serialize + DeserializeOwned> OcrWireBody<B> {
|
||||
|
|
@ -124,14 +122,7 @@ impl<B: Serialize + DeserializeOwned> OcrWireBody<B> {
|
|||
path: prefix.into(),
|
||||
});
|
||||
};
|
||||
let known = serde_json::to_value(&body).map_err(|_| super::Error::RequestField {
|
||||
path: prefix.into(),
|
||||
})?;
|
||||
let extra = fields
|
||||
.into_iter()
|
||||
.filter(|(key, _)| known.get(key).is_none())
|
||||
.collect();
|
||||
Ok(Self { body, extra })
|
||||
Ok(Self { body, fields })
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -7,10 +7,10 @@ use serde::{Deserialize, Serialize};
|
|||
use serde_json::{Map, Value};
|
||||
use serde_with::serde_as;
|
||||
|
||||
use super::OcrArguments;
|
||||
use super::hooks::{NoopOcrHooks, OcrHooks};
|
||||
use super::provider_config::{OcrConfigKind, resolve_provider_config};
|
||||
use crate::constants::OCR_HTTP_TIMEOUT_SECS;
|
||||
use crate::params::OpaqueParams;
|
||||
use litellm_auth::{InputSource, TokenProviderHandle};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
|
|
@ -97,7 +97,7 @@ pub struct LiteLLMOcrRequest {
|
|||
pub connection: OcrConnection,
|
||||
pub hooks: Arc<dyn OcrHooks>,
|
||||
pub litellm_call_id: Option<String>,
|
||||
pub optional_params: OpaqueParams,
|
||||
pub optional_params: OcrArguments,
|
||||
pub input_sources: BTreeMap<String, InputSource>,
|
||||
pub azure_ad_token_provider: Option<TokenProviderHandle>,
|
||||
pub(crate) config: OcrConfigKind,
|
||||
|
|
@ -108,7 +108,7 @@ impl LiteLLMOcrRequest {
|
|||
model: String,
|
||||
document: OcrDocument,
|
||||
custom_llm_provider: Option<&str>,
|
||||
optional_params: OpaqueParams,
|
||||
optional_params: OcrArguments,
|
||||
) -> Result<Self, super::Error> {
|
||||
let (model, config) = resolve_provider_config(&model, custom_llm_provider)?;
|
||||
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::time::Duration;
|
||||
|
||||
use super::OcrArguments;
|
||||
use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument};
|
||||
use crate::params::OpaqueParams;
|
||||
use litellm_auth::InputSource;
|
||||
use serde::{
|
||||
Deserialize,
|
||||
|
|
@ -54,7 +54,7 @@ pub struct OcrWireRequest {
|
|||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
#[serde(default)]
|
||||
pub optional_params: OpaqueParams,
|
||||
pub optional_params: OcrArguments,
|
||||
#[serde(default)]
|
||||
pub input_sources: BTreeMap<String, InputSource>,
|
||||
pub timeout_seconds: Option<f64>,
|
||||
|
|
@ -70,9 +70,9 @@ pub fn consumed_optional_param_names(
|
|||
) -> Result<Vec<&'static str>, crate::ocr::Error> {
|
||||
use super::provider_config::OcrConfigKind;
|
||||
|
||||
let (provider_model, config) =
|
||||
let (model, config) =
|
||||
super::provider_config::resolve_provider_config(model, custom_llm_provider)?;
|
||||
let provider_fields = config.get_supported_ocr_params(&provider_model);
|
||||
let provider_fields = config.get_supported_ocr_params(&model);
|
||||
let auth_fields: &[&str] = match config {
|
||||
OcrConfigKind::AzureAi
|
||||
| OcrConfigKind::AzureDocumentIntelligence
|
||||
|
|
@ -110,6 +110,17 @@ pub fn consumed_optional_params(
|
|||
})
|
||||
}
|
||||
|
||||
pub fn project_argument(
|
||||
name: &str,
|
||||
consumed: &[OptionalParamSpec],
|
||||
host_fields: &[String],
|
||||
) -> bool {
|
||||
consumed.iter().any(|field| field.name == name)
|
||||
|| (!host_fields.iter().any(|field| field == name)
|
||||
&& !crate::params::is_control_param(name)
|
||||
&& !matches!(name, "model" | "document" | "timeout" | "input_sources"))
|
||||
}
|
||||
|
||||
pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, crate::ocr::Error> {
|
||||
let api_key_source = source_for(&wire.input_sources, "api_key");
|
||||
let api_base_source = source_for(&wire.input_sources, "api_base");
|
||||
|
|
@ -255,6 +266,19 @@ pub fn decode_response<T: DeserializeOwned>(
|
|||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn core_selects_consumed_values_without_serializing_host_objects() {
|
||||
let fields = consumed_optional_params("mistral/model", None).unwrap();
|
||||
let host_fields = vec!["metadata".into(), "callbacks".into(), "id".into()];
|
||||
assert!(project_argument("future_option", &fields, &host_fields));
|
||||
assert!(project_argument("extra_body", &fields, &host_fields));
|
||||
assert!(project_argument("id", &fields, &host_fields));
|
||||
assert!(!project_argument("metadata", &fields, &host_fields));
|
||||
assert!(!project_argument("callbacks", &fields, &host_fields));
|
||||
assert!(!project_argument("api_key", &fields, &host_fields));
|
||||
assert!(!project_argument("document", &fields, &host_fields));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn option_projection_is_provider_specific_and_excludes_opaque_fields() {
|
||||
let mistral = consumed_optional_param_names("mistral/model", None).unwrap();
|
||||
|
|
|
|||
|
|
@ -247,6 +247,52 @@ async fn facade_omits_native_response_by_default_and_preserves_auth_priority() {
|
|||
|
||||
struct RewriteDocument;
|
||||
|
||||
struct RewriteHeaders;
|
||||
|
||||
impl OcrHooks for RewriteHeaders {
|
||||
fn intercepts_requests(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn during_call(
|
||||
&self,
|
||||
request: OcrDuringCallRequest,
|
||||
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
|
||||
Box::pin(async move {
|
||||
Ok(OcrDuringCallRequest {
|
||||
headers: vec![("authorization".into(), "Bearer guarded".into())],
|
||||
..request
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case("reducto/parse-v3")]
|
||||
#[case("reducto/parse-legacy")]
|
||||
#[tokio::test]
|
||||
async fn guardrail_headers_reach_upload_and_parse(#[case] model: &str) {
|
||||
let (base, seen, server) = mock_server(vec![
|
||||
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
|
||||
MockResponse::json(json!({"result":{"chunks":[]}})),
|
||||
])
|
||||
.await;
|
||||
let mut request = wire_request(model, &base, json!({}));
|
||||
request.connection.extra_headers = vec![("authorization".into(), "Bearer original".into())];
|
||||
request.hooks = Arc::new(RewriteHeaders);
|
||||
|
||||
perform_ocr(request).await.unwrap();
|
||||
server.await.unwrap();
|
||||
let requests = seen.lock().unwrap();
|
||||
assert_eq!(requests.len(), 2);
|
||||
assert!(requests[0].starts_with("POST /upload "));
|
||||
assert!(requests[1].starts_with("POST /parse "));
|
||||
for request in requests.iter() {
|
||||
assert!(request.contains("authorization: Bearer guarded"));
|
||||
assert!(!request.contains("Bearer original"));
|
||||
}
|
||||
}
|
||||
|
||||
impl OcrHooks for RewriteDocument {
|
||||
fn intercepts_requests(&self) -> bool {
|
||||
true
|
||||
|
|
|
|||
|
|
@ -90,7 +90,7 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py<PyAny>) -> PyRe
|
|||
|
||||
pub(crate) fn project_optional_fields(
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
names: &[&str],
|
||||
fields: &[litellm_core::ocr::wire::OptionalParamSpec],
|
||||
) -> PyResult<Map<String, Value>> {
|
||||
let controls: Vec<String> = kwargs
|
||||
.py()
|
||||
|
|
@ -102,13 +102,7 @@ pub(crate) fn project_optional_fields(
|
|||
.map(|(name, value)| Ok((name.extract::<String>()?, value)))
|
||||
.filter_map(|entry: PyResult<_>| match entry {
|
||||
Ok((name, value))
|
||||
if names.contains(&name.as_str())
|
||||
|| (!controls.contains(&name)
|
||||
&& !litellm_core::params::is_control_param(&name)
|
||||
&& !matches!(
|
||||
name.as_str(),
|
||||
"model" | "document" | "timeout" | "input_sources"
|
||||
)) =>
|
||||
if litellm_core::ocr::wire::project_argument(&name, fields, &controls) =>
|
||||
{
|
||||
Some(from_py(&value).map(|value| (name, value)))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -124,8 +124,7 @@ pub(super) fn project_request(
|
|||
let api_key = arguments.api_key()?;
|
||||
let specs = consumed_optional_params(&model, custom_llm_provider.as_deref())
|
||||
.map_err(ocr_error_to_pyerr)?;
|
||||
let names = specs.iter().map(|spec| spec.name).collect::<Vec<_>>();
|
||||
let optional_params = project_optional_fields(kwargs, &names)?;
|
||||
let optional_params = project_optional_fields(kwargs, &specs)?;
|
||||
let input_sources = request_input_sources(
|
||||
kwargs,
|
||||
optional_params
|
||||
|
|
|
|||
|
|
@ -67,6 +67,40 @@ async def test_public_cohere_blocks_and_usage_fallback(recording_server: Recordi
|
|||
assert response.get_provider_native_response() is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
async def test_explicit_body_preserves_provider_native_values(
|
||||
recording_server: RecordingServer, model: str, asynchronous: bool
|
||||
) -> None:
|
||||
recording_server.enqueue(ResponseSpec(body=PAYLOAD))
|
||||
document: Final = {**IMAGE, "provider_options": {"future": [None, False, 0]}}
|
||||
arguments: Final = {
|
||||
"model": model,
|
||||
"document": IMAGE,
|
||||
"api_base": recording_server.base_url,
|
||||
"api_key": "test-key",
|
||||
"output_format": "markdown",
|
||||
"future_option": {"original": True},
|
||||
"extra_body": {
|
||||
"output_format": "future-format",
|
||||
"document": document,
|
||||
"future_option": None,
|
||||
},
|
||||
}
|
||||
if asynchronous:
|
||||
await litellm.aocr(**arguments)
|
||||
else:
|
||||
litellm.ocr(**arguments)
|
||||
assert not recording_server.requests[0].headers.get("user-agent", "").startswith("python-httpx")
|
||||
assert recording_server.requests[0].body == {
|
||||
"model": model.split("/", 1)[1],
|
||||
"document": document,
|
||||
"output_format": "future-format",
|
||||
"future_option": None,
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
|
|
@ -726,7 +726,7 @@ async def test_reducto_lifecycle_retains_upload_parse_and_post_call_boundaries(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_document_intelligence_post_call_observes_submission_and_final_result(
|
||||
async def test_document_intelligence_post_call_observes_submission_before_polling(
|
||||
ocr_server: RecordingServer,
|
||||
) -> None:
|
||||
ocr_server.expected_requests = 2
|
||||
|
|
@ -757,9 +757,8 @@ async def test_document_intelligence_post_call_observes_submission_and_final_res
|
|||
response: Final = await call_aocr(
|
||||
ocr_server, model="azure_ai/doc-intelligence/prebuilt-read", litellm_logging_obj=logger
|
||||
)
|
||||
assert [methods for methods, _ in boundaries] == [("POST",), ("POST", "GET")]
|
||||
assert [methods for methods, _ in boundaries] == [("POST",)]
|
||||
assert json.loads(boundaries[0][1])["status"] == "running"
|
||||
assert json.loads(boundaries[1][1])["status"] == "succeeded"
|
||||
assert [request.method for request in ocr_server.requests] == ["POST", "GET"]
|
||||
assert ocr_server.requests[1].path == "/operations/1"
|
||||
assert response.pages == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue