mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
wip
This commit is contained in:
parent
31b48f6191
commit
5de196af63
21 changed files with 702 additions and 458 deletions
466
litellm-rust/crates/core/src/call_arguments.rs
Normal file
466
litellm-rust/crates/core/src/call_arguments.rs
Normal file
|
|
@ -0,0 +1,466 @@
|
|||
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 CallArguments(Map<String, Value>);
|
||||
|
||||
impl CallArguments {
|
||||
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()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
|
||||
#[error("invalid argument: {path}")]
|
||||
pub struct ArgumentError {
|
||||
pub path: String,
|
||||
}
|
||||
|
||||
pub fn parse_options<T: DeserializeOwned>(arguments: &CallArguments) -> Result<T, ArgumentError> {
|
||||
use serde::de::IntoDeserializer;
|
||||
serde_path_to_error::deserialize(Value::Object(arguments.0.clone()).into_deserializer())
|
||||
.map_err(|error| ArgumentError {
|
||||
path: error.path().to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct ArgumentSpec {
|
||||
pub name: &'static str,
|
||||
pub secret: bool,
|
||||
}
|
||||
|
||||
pub fn should_project(name: &str, consumed: &[ArgumentSpec], bound_fields: &[&str]) -> bool {
|
||||
consumed.iter().any(|field| field.name == name)
|
||||
|| (!bound_fields.contains(&name) && !is_control(name))
|
||||
}
|
||||
|
||||
pub fn is_control(name: &str) -> bool {
|
||||
crate::params::is_control_param(name) || HOST_CONTROLS.contains(&name)
|
||||
}
|
||||
|
||||
const HOST_CONTROLS: &[&str] = &[
|
||||
"_agentic_loop_api_surface",
|
||||
"_agentic_loop_depth",
|
||||
"_agentic_loop_fingerprints",
|
||||
"_code_interpreter_interception_active",
|
||||
"_code_interpreter_interception_converted_stream",
|
||||
"_code_interpreter_interception_sandbox_key",
|
||||
"_code_interpreter_interception_session_scoped",
|
||||
"_headroom_interception_converted_stream",
|
||||
"_litellm_strip_stream_usage",
|
||||
"_router_weights",
|
||||
"_websearch_interception_converted_stream",
|
||||
"_websearch_interception_emit_native_blocks",
|
||||
"acompletion",
|
||||
"adaptive_router_config",
|
||||
"adaptive_router_default_model",
|
||||
"aembedding",
|
||||
"aimg_generation",
|
||||
"allm_passthrough_route",
|
||||
"allow_client_keepalive_override",
|
||||
"allowed_model_region",
|
||||
"allowed_openai_params",
|
||||
"annotation_cost_per_page",
|
||||
"api_version",
|
||||
"arize_api_key",
|
||||
"arize_space_id",
|
||||
"arize_space_key",
|
||||
"assistant_continue_message",
|
||||
"async_call",
|
||||
"atext_completion",
|
||||
"attempted_targets",
|
||||
"auto_router_config",
|
||||
"auto_router_config_path",
|
||||
"auto_router_default_model",
|
||||
"auto_router_embedding_model",
|
||||
"auto_router_max_input_chars",
|
||||
"auto_router_model_compression",
|
||||
"auto_router_routing_compression",
|
||||
"aws_batch_role_arn",
|
||||
"azure",
|
||||
"azure_password",
|
||||
"azure_username",
|
||||
"base_model",
|
||||
"bedrock_tags",
|
||||
"bos_token",
|
||||
"budget_duration",
|
||||
"cache",
|
||||
"cache_creation_input_audio_token_cost",
|
||||
"cache_creation_input_token_cost",
|
||||
"cache_creation_input_token_cost_above_1hr",
|
||||
"cache_creation_input_token_cost_above_200k_tokens",
|
||||
"cache_creation_input_token_cost_above_272k_tokens",
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex",
|
||||
"cache_creation_input_token_cost_above_272k_tokens_priority",
|
||||
"cache_creation_input_token_cost_flex",
|
||||
"cache_creation_input_token_cost_priority",
|
||||
"cache_creation_input_token_cost_ultrafast",
|
||||
"cache_key",
|
||||
"cache_read_input_audio_token_cost",
|
||||
"cache_read_input_token_cost",
|
||||
"cache_read_input_token_cost_above_200k_tokens",
|
||||
"cache_read_input_token_cost_above_200k_tokens_priority",
|
||||
"cache_read_input_token_cost_above_272k_tokens",
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex",
|
||||
"cache_read_input_token_cost_above_272k_tokens_priority",
|
||||
"cache_read_input_token_cost_above_512k_tokens",
|
||||
"cache_read_input_token_cost_flex",
|
||||
"cache_read_input_token_cost_priority",
|
||||
"cache_read_input_token_cost_ultrafast",
|
||||
"caching",
|
||||
"caching_groups",
|
||||
"citation_cost_per_token",
|
||||
"client",
|
||||
"client_side_timeout",
|
||||
"complete_response",
|
||||
"completion_call_id",
|
||||
"complexity_router_config",
|
||||
"complexity_router_default_model",
|
||||
"configurable_clientside_auth_params",
|
||||
"context_window_fallback_dict",
|
||||
"cooldown_time",
|
||||
"cost_per_query",
|
||||
"custom_prompt_dict",
|
||||
"data_residency",
|
||||
"dd_agent_host",
|
||||
"dd_agent_port",
|
||||
"dd_api_key",
|
||||
"dd_site",
|
||||
"default_api_key_rpm_limit",
|
||||
"default_api_key_tpm_limit",
|
||||
"disable_add_transform_inline_image_block",
|
||||
"enable_json_schema_validation",
|
||||
"enable_prompt_caching",
|
||||
"enable_tag_filtering",
|
||||
"ensure_alternating_roles",
|
||||
"eos_token",
|
||||
"fallback_depth",
|
||||
"fallbacks",
|
||||
"fastest_response",
|
||||
"final_prompt_value",
|
||||
"force_timeout",
|
||||
"gcs_bucket_name",
|
||||
"gcs_path_service_account",
|
||||
"google_maps_grounding_cost_per_query",
|
||||
"headers",
|
||||
"hf_model_name",
|
||||
"humanloop_api_key",
|
||||
"id",
|
||||
"input_cost_per_audio_per_second",
|
||||
"input_cost_per_audio_per_second_above_128k_tokens",
|
||||
"input_cost_per_audio_token",
|
||||
"input_cost_per_audio_token_batches",
|
||||
"input_cost_per_character",
|
||||
"input_cost_per_character_above_128k_tokens",
|
||||
"input_cost_per_image",
|
||||
"input_cost_per_image_above_128k_tokens",
|
||||
"input_cost_per_image_token",
|
||||
"input_cost_per_image_token_batches",
|
||||
"input_cost_per_pixel",
|
||||
"input_cost_per_query",
|
||||
"input_cost_per_second",
|
||||
"input_cost_per_token",
|
||||
"input_cost_per_token_above_128k_tokens",
|
||||
"input_cost_per_token_above_200k_tokens",
|
||||
"input_cost_per_token_above_200k_tokens_priority",
|
||||
"input_cost_per_token_above_272k_tokens",
|
||||
"input_cost_per_token_above_272k_tokens_flex",
|
||||
"input_cost_per_token_above_272k_tokens_priority",
|
||||
"input_cost_per_token_above_512k_tokens",
|
||||
"input_cost_per_token_batches",
|
||||
"input_cost_per_token_cache_hit",
|
||||
"input_cost_per_token_flex",
|
||||
"input_cost_per_token_priority",
|
||||
"input_cost_per_token_ultrafast",
|
||||
"input_cost_per_video_per_second",
|
||||
"input_cost_per_video_per_second_above_128k_tokens",
|
||||
"input_cost_per_video_per_second_above_15s_interval",
|
||||
"input_cost_per_video_per_second_above_8s_interval",
|
||||
"input_cost_per_video_token",
|
||||
"input_cost_per_video_token_batches",
|
||||
"itpm",
|
||||
"keepalive_seconds",
|
||||
"langfuse_environment",
|
||||
"langfuse_host",
|
||||
"langfuse_prompt_version",
|
||||
"langfuse_public_key",
|
||||
"langfuse_secret",
|
||||
"langfuse_secret_key",
|
||||
"langsmith_api_key",
|
||||
"langsmith_base_url",
|
||||
"langsmith_project",
|
||||
"langsmith_sampling_rate",
|
||||
"langsmith_tenant_id",
|
||||
"litellm_credential_name",
|
||||
"litellm_disabled_callbacks",
|
||||
"litellm_request_debug",
|
||||
"litellm_session_id",
|
||||
"litellm_system_prompt",
|
||||
"litellm_trace_id",
|
||||
"litellm_trusted_callback_vars",
|
||||
"logger_fn",
|
||||
"max_agentic_loops",
|
||||
"max_budget",
|
||||
"max_fallbacks",
|
||||
"max_parallel_requests",
|
||||
"merge_reasoning_content_in_choices",
|
||||
"metadata",
|
||||
"mock_response",
|
||||
"mock_timeout",
|
||||
"model_alias_map",
|
||||
"model_config",
|
||||
"model_file_id_mapping",
|
||||
"model_info",
|
||||
"model_list",
|
||||
"newrelic_api_key",
|
||||
"newrelic_region",
|
||||
"no-log",
|
||||
"num_retries",
|
||||
"ocr_cost_per_credit",
|
||||
"ocr_cost_per_page",
|
||||
"order",
|
||||
"otpm",
|
||||
"output_cost_per_audio_per_second",
|
||||
"output_cost_per_audio_token",
|
||||
"output_cost_per_character",
|
||||
"output_cost_per_character_above_128k_tokens",
|
||||
"output_cost_per_image",
|
||||
"output_cost_per_image_token",
|
||||
"output_cost_per_pixel",
|
||||
"output_cost_per_reasoning_token",
|
||||
"output_cost_per_reasoning_token_flex",
|
||||
"output_cost_per_reasoning_token_priority",
|
||||
"output_cost_per_second",
|
||||
"output_cost_per_second_1080p",
|
||||
"output_cost_per_second_480p",
|
||||
"output_cost_per_second_4k",
|
||||
"output_cost_per_second_720p",
|
||||
"output_cost_per_token",
|
||||
"output_cost_per_token_above_128k_tokens",
|
||||
"output_cost_per_token_above_200k_tokens",
|
||||
"output_cost_per_token_above_200k_tokens_priority",
|
||||
"output_cost_per_token_above_272k_tokens",
|
||||
"output_cost_per_token_above_272k_tokens_flex",
|
||||
"output_cost_per_token_above_272k_tokens_priority",
|
||||
"output_cost_per_token_above_512k_tokens",
|
||||
"output_cost_per_token_batches",
|
||||
"output_cost_per_token_flex",
|
||||
"output_cost_per_token_priority",
|
||||
"output_cost_per_token_ultrafast",
|
||||
"output_cost_per_video_per_second",
|
||||
"output_cost_per_video_token",
|
||||
"output_vector_size",
|
||||
"posthog_api_key",
|
||||
"posthog_api_url",
|
||||
"preset_cache_key",
|
||||
"prompt_environment",
|
||||
"prompt_id",
|
||||
"prompt_label",
|
||||
"prompt_variables",
|
||||
"prompt_version",
|
||||
"provider_specific_header",
|
||||
"quality_router_config",
|
||||
"quality_router_default_model",
|
||||
"region_name",
|
||||
"regional_endpoint_uplift_multiplier",
|
||||
"regional_processing_uplift_multiplier_eu",
|
||||
"regional_processing_uplift_multiplier_us",
|
||||
"retry_policy",
|
||||
"retry_strategy",
|
||||
"roles",
|
||||
"routing_strategy",
|
||||
"rpm",
|
||||
"rust",
|
||||
"s3_bucket_name",
|
||||
"s3_output_bucket_name",
|
||||
"s3_region_name",
|
||||
"search_context_cost_per_query",
|
||||
"search_tool_name",
|
||||
"secret_fields",
|
||||
"self",
|
||||
"shared_session",
|
||||
"ssl_verify",
|
||||
"stream_response",
|
||||
"stream_timeout",
|
||||
"supports_system_message",
|
||||
"tags",
|
||||
"text_completion",
|
||||
"tiered_pricing",
|
||||
"tpm",
|
||||
"ttl",
|
||||
"turn_off_message_logging",
|
||||
"use_chat_completions_api",
|
||||
"use_client",
|
||||
"use_in_pass_through",
|
||||
"use_litellm_proxy",
|
||||
"use_xai_oauth",
|
||||
"user_continue_message",
|
||||
"verbose",
|
||||
"wandb_api_key",
|
||||
"weave_project_id",
|
||||
"weight",
|
||||
];
|
||||
|
||||
pub fn compose_body<B: Serialize>(
|
||||
arguments: &CallArguments,
|
||||
body: &B,
|
||||
consumed: &[&str],
|
||||
) -> Result<Value, crate::params::Error> {
|
||||
let Value::Object(fields) =
|
||||
serde_json::to_value(body).map_err(|_| crate::params::Error::Body)?
|
||||
else {
|
||||
return Err(crate::params::Error::Body);
|
||||
};
|
||||
let overrides = match arguments.get("extra_body") {
|
||||
None | Some(Value::Null) => None,
|
||||
Some(Value::Object(fields)) => Some(fields),
|
||||
Some(_) => return Err(crate::params::Error::ExtraBody),
|
||||
};
|
||||
let extensions = arguments.iter().filter(|(name, _)| {
|
||||
!consumed.contains(&name.as_str()) && name.as_str() != "extra_body" && !is_control(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 CallArguments {
|
||||
type Target = Map<String, Value>;
|
||||
|
||||
fn deref(&self) -> &Self::Target {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl From<Map<String, Value>> for CallArguments {
|
||||
fn from(values: Map<String, Value>) -> Self {
|
||||
Self(values)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<CallArguments> for Map<String, Value> {
|
||||
fn from(arguments: CallArguments) -> Self {
|
||||
arguments.0
|
||||
}
|
||||
}
|
||||
|
||||
impl FromIterator<(String, Value)> for CallArguments {
|
||||
fn from_iter<T: IntoIterator<Item = (String, Value)>>(iter: T) -> Self {
|
||||
Self(iter.into_iter().collect())
|
||||
}
|
||||
}
|
||||
|
||||
impl IntoIterator for CallArguments {
|
||||
type Item = (String, Value);
|
||||
type IntoIter = serde_json::map::IntoIter;
|
||||
|
||||
fn into_iter(self) -> Self::IntoIter {
|
||||
self.0.into_iter()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn composition_preserves_extensions_and_applies_shallow_explicit_overrides() {
|
||||
let original = json!({
|
||||
"known": false, "future": {"old": 1}, "null": null, "zero": 0,
|
||||
"metadata": {"host": true}, "shared_session": "host", "api_key": "secret",
|
||||
"extra_body": {
|
||||
"known": null, "future": {"new": [false, 0, null]},
|
||||
"metadata": {"provider": true}, "model": "ignored", "api_key": "ignored"
|
||||
}
|
||||
});
|
||||
let arguments = serde_json::from_value(original.clone()).unwrap();
|
||||
let body = compose_body(
|
||||
&arguments,
|
||||
&json!({"model":"resolved", "known":false}),
|
||||
&["known"],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
body,
|
||||
json!({
|
||||
"model":"resolved", "known":null, "future":{"new":[false,0,null]},
|
||||
"null":null, "zero":0, "metadata":{"provider":true}
|
||||
})
|
||||
);
|
||||
assert_eq!(serde_json::to_value(arguments).unwrap(), original);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn projection_prioritizes_consumed_fields_and_keeps_unknown_names() {
|
||||
let fields = [ArgumentSpec {
|
||||
name: "id",
|
||||
secret: false,
|
||||
}];
|
||||
assert!(should_project("id", &fields, &[]));
|
||||
assert!(!should_project("id", &[], &[]));
|
||||
assert!(should_project("future_option", &[], &[]));
|
||||
assert!(!should_project("document", &fields, &["document"]));
|
||||
assert!(!should_project("metadata", &fields, &[]));
|
||||
assert!(!should_project("callbacks", &fields, &[]));
|
||||
assert!(!should_project("ocr_cost_per_page", &fields, &[]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_extra_body_is_rejected_without_coercing_it_to_empty() {
|
||||
for value in [json!(false), json!(0), json!([]), json!("")] {
|
||||
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
|
||||
assert_eq!(
|
||||
compose_body(&arguments, &json!({}), &[]),
|
||||
Err(crate::params::Error::ExtraBody)
|
||||
);
|
||||
}
|
||||
let arguments = serde_json::from_value(json!({"extra_body":null})).unwrap();
|
||||
assert_eq!(
|
||||
compose_body(&arguments, &json!({}), &[]).unwrap(),
|
||||
json!({})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn typed_views_preserve_missing_and_explicit_null_in_the_source() {
|
||||
#[derive(Deserialize)]
|
||||
struct Options {
|
||||
enabled: Option<bool>,
|
||||
}
|
||||
let arguments: CallArguments =
|
||||
serde_json::from_value(json!({"enabled":null,"future":0})).unwrap();
|
||||
assert!(
|
||||
parse_options::<Options>(&arguments)
|
||||
.unwrap()
|
||||
.enabled
|
||||
.is_none()
|
||||
);
|
||||
assert_eq!(arguments.get("enabled"), Some(&Value::Null));
|
||||
assert_eq!(arguments.get("missing"), None);
|
||||
let invalid = serde_json::from_value(json!({"enabled":0})).unwrap();
|
||||
assert_eq!(
|
||||
parse_options::<Options>(&invalid).err().unwrap().path,
|
||||
"enabled"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
pub mod audio_transcription;
|
||||
pub mod call_arguments;
|
||||
pub mod call_lifecycle;
|
||||
pub mod chat_completions;
|
||||
pub mod constants;
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
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::llms::cohere::ocr::{CohereOptions, CohereResponse, validate_document};
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::document::{inline_remote_document, validate_inline_document};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
|
|
@ -15,7 +14,7 @@ const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
|
|||
pub(crate) struct AzureAICohereParseConfig;
|
||||
|
||||
impl BaseOcrConfig for AzureAICohereParseConfig {
|
||||
type OcrParams = CohereParams;
|
||||
type OcrParams = CohereOptions;
|
||||
type ProviderRequest = CohereRequest;
|
||||
type ProviderResponse = CohereResponse;
|
||||
|
||||
|
|
@ -23,20 +22,11 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
|
|||
CohereParseConfig.get_supported_ocr_params(model)
|
||||
}
|
||||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
CohereParseConfig.map_ocr_params(non_default_params, optional_params, model)
|
||||
}
|
||||
|
||||
async fn async_transform_ocr_request(
|
||||
&self,
|
||||
model: &str,
|
||||
document: OcrDocument,
|
||||
optional_params: &CohereParams,
|
||||
optional_params: &CohereOptions,
|
||||
headers: &[(String, String)],
|
||||
context: OcrRequestContext<'_>,
|
||||
) -> Result<CohereRequest, crate::ocr::Error> {
|
||||
|
|
@ -65,7 +55,7 @@ impl AzureAICohereParseConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
|
||||
let config = AzureAuthInputs {
|
||||
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
|
||||
..AzureAuthInputs::from_sourced_optional_params(
|
||||
|
|
@ -110,8 +100,9 @@ impl AzureAICohereParseConfig {
|
|||
!remote,
|
||||
body,
|
||||
|body| {
|
||||
validate_document(&body.document.as_document())?;
|
||||
validate_inline_document(&body.document.as_document())
|
||||
let document = crate::ocr::prepare::body_document(body)?;
|
||||
validate_document(&document)?;
|
||||
validate_inline_document(&document)
|
||||
},
|
||||
)
|
||||
.await
|
||||
|
|
|
|||
|
|
@ -12,6 +12,7 @@ use tokio::time::Instant;
|
|||
use litellm_auth::{InputSource, Sourced};
|
||||
use litellm_auth_azure::AzureAuthInputs;
|
||||
|
||||
use crate::call_arguments::CallArguments;
|
||||
use crate::constants::{
|
||||
AZURE_DI_API_VERSION, AZURE_DI_DEFAULT_DPI, AZURE_DI_DEFAULT_HEIGHT, AZURE_DI_DEFAULT_WIDTH,
|
||||
AZURE_DI_SUBSCRIPTION_HEADER, OCR_POLL_RETRY_SECS,
|
||||
|
|
@ -19,7 +20,6 @@ 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;
|
||||
|
|
@ -30,7 +30,6 @@ use crate::ocr::types::{
|
|||
OcrResponseFormat, OcrUsageInfo,
|
||||
};
|
||||
use crate::ocr::wire::DecodedOcrResponse;
|
||||
use crate::params::OpaqueParams;
|
||||
use crate::serde_compat::{FiniteF64, LaxI64};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
|
|
@ -61,10 +60,6 @@ pub(crate) struct DocumentIntelligenceParams {
|
|||
pub pages: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub features: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub req_format: Option<OcrResponseFormat>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: OpaqueParams,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
|
|
@ -183,8 +178,6 @@ fn normalize_ocr_params(
|
|||
.map(normalize_features)
|
||||
.transpose()?
|
||||
.flatten(),
|
||||
req_format: None,
|
||||
extra_fields: OpaqueParams::default(),
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -485,53 +478,13 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
arguments: &CallArguments,
|
||||
_model: &str,
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
let mapped = normalize_ocr_params(decode_input_params(
|
||||
non_default_params
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "pages" | "features"))
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect(),
|
||||
) -> Result<DocumentIntelligenceParams, crate::ocr::Error> {
|
||||
normalize_ocr_params(decode_input_params(
|
||||
arguments.select(&["pages", "features"]),
|
||||
"optional_params",
|
||||
)?)?;
|
||||
let request_format = non_default_params
|
||||
.get("req_format")
|
||||
.filter(|value| !value.is_null())
|
||||
.map(|value| {
|
||||
serde_json::from_value::<OcrResponseFormat>(value.clone())
|
||||
.map_err(|_| crate::ocr::Error::RequestFormat)
|
||||
})
|
||||
.transpose()?;
|
||||
let fields = optional_params
|
||||
.iter()
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.chain(
|
||||
mapped
|
||||
.pages
|
||||
.map(|value| ("pages".into(), Value::String(value))),
|
||||
)
|
||||
.chain(
|
||||
mapped
|
||||
.features
|
||||
.map(|value| ("features".into(), Value::String(value))),
|
||||
)
|
||||
.chain(request_format.map(|value| {
|
||||
(
|
||||
"req_format".into(),
|
||||
Value::String(
|
||||
match value {
|
||||
OcrResponseFormat::Litellm => "litellm",
|
||||
OcrResponseFormat::Native => "native",
|
||||
}
|
||||
.into(),
|
||||
),
|
||||
)
|
||||
}))
|
||||
.collect();
|
||||
Ok(fields)
|
||||
)?)
|
||||
}
|
||||
|
||||
async fn async_transform_ocr_request(
|
||||
|
|
@ -592,7 +545,7 @@ impl AzureDocumentIntelligenceOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
|
||||
let config = AzureAuthInputs {
|
||||
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
|
||||
..AzureAuthInputs::from_sourced_optional_params(
|
||||
|
|
@ -728,35 +681,27 @@ mod tests {
|
|||
serde_json::from_value(json!({"pages":[], "features":null, "req_format":"native"}))
|
||||
.unwrap();
|
||||
let mapped = AzureDocumentIntelligenceOCRConfig
|
||||
.parse_options(&overrides, "model")
|
||||
.map_ocr_params(&overrides, "model")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(mapped).unwrap(),
|
||||
json!({
|
||||
"req_format":"native"
|
||||
})
|
||||
);
|
||||
assert_eq!(serde_json::to_value(mapped).unwrap(), json!({}));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mapping_preserves_supplied_options_when_overrides_are_empty() {
|
||||
let supplied = serde_json::from_value(json!({
|
||||
fn options_normalize_query_fields_without_consuming_extensions() {
|
||||
let arguments = 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")
|
||||
.map_ocr_params(&arguments, "model")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(mapped).unwrap(),
|
||||
json!({
|
||||
"pages":"4", "features":"languages", "extension":true, "req_format":"native"
|
||||
"pages":"4", "features":"languages"
|
||||
})
|
||||
);
|
||||
assert_eq!(arguments["extension"], true);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ 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};
|
||||
|
|
@ -27,15 +26,6 @@ impl BaseOcrConfig for AzureAIOCRConfig {
|
|||
MistralOCRConfig.get_supported_ocr_params(model)
|
||||
}
|
||||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
MistralOCRConfig.map_ocr_params(non_default_params, optional_params, model)
|
||||
}
|
||||
|
||||
async fn async_transform_ocr_request(
|
||||
&self,
|
||||
model: &str,
|
||||
|
|
@ -68,7 +58,7 @@ impl AzureAIOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
|
||||
let config = AzureAuthInputs {
|
||||
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
|
||||
..AzureAuthInputs::from_sourced_optional_params(
|
||||
|
|
@ -102,7 +92,7 @@ impl AzureAIOCRConfig {
|
|||
&headers,
|
||||
retains_document,
|
||||
body,
|
||||
|body| validate_inline_document(&body.document),
|
||||
|body| validate_inline_document(&crate::ocr::prepare::body_document(body)?),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ use std::sync::Arc;
|
|||
use serde::Serialize;
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
use crate::ocr::OcrArguments;
|
||||
use crate::call_arguments::{CallArguments, parse_options};
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::hooks::OcrHooks;
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrResponseFormat};
|
||||
|
|
@ -20,20 +20,14 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
|
|||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
_non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
_model: &str,
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
Ok(optional_params.clone())
|
||||
}
|
||||
|
||||
fn parse_options(
|
||||
&self,
|
||||
arguments: &OcrArguments,
|
||||
arguments: &CallArguments,
|
||||
model: &str,
|
||||
) -> Result<Self::OcrParams, crate::ocr::Error> {
|
||||
self.map_ocr_params(arguments, &OcrArguments::default(), model)?
|
||||
.parse()
|
||||
Ok(parse_options(
|
||||
&arguments
|
||||
.select(self.get_supported_ocr_params(model))
|
||||
.into(),
|
||||
)?)
|
||||
}
|
||||
|
||||
fn async_transform_ocr_request(
|
||||
|
|
|
|||
|
|
@ -1,3 +1,3 @@
|
|||
pub(crate) mod transformation;
|
||||
|
||||
pub(crate) use transformation::{CohereParams, CohereResponse, validate_document};
|
||||
pub(crate) use transformation::{CohereOptions, CohereResponse, validate_document};
|
||||
|
|
|
|||
|
|
@ -5,7 +5,6 @@ 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};
|
||||
|
|
@ -24,7 +23,7 @@ pub(crate) enum OutputFormat {
|
|||
}
|
||||
|
||||
#[derive(Default, Deserialize, Serialize)]
|
||||
pub(crate) struct CohereParams {
|
||||
pub(crate) struct CohereOptions {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub output_format: Option<OutputFormat>,
|
||||
}
|
||||
|
|
@ -43,16 +42,6 @@ pub(crate) enum CohereParseDocument {
|
|||
ImageUrl { image_url: String },
|
||||
}
|
||||
|
||||
impl CohereParseDocument {
|
||||
pub(crate) fn as_document(&self) -> OcrDocument {
|
||||
let Self::ImageUrl { image_url } = self;
|
||||
OcrDocument::ImageUrl {
|
||||
image_url: image_url.clone(),
|
||||
extra_fields: Default::default(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct CohereResponse {
|
||||
#[serde(default)]
|
||||
|
|
@ -96,7 +85,7 @@ impl CohereParseConfig {
|
|||
&self,
|
||||
model: &str,
|
||||
document: OcrDocument,
|
||||
optional_params: &CohereParams,
|
||||
optional_params: &CohereOptions,
|
||||
_headers: &[(String, String)],
|
||||
) -> Result<CohereRequest, crate::ocr::Error> {
|
||||
let image_url = image_url(document)?;
|
||||
|
|
@ -105,7 +94,7 @@ impl CohereParseConfig {
|
|||
}
|
||||
|
||||
impl BaseOcrConfig for CohereParseConfig {
|
||||
type OcrParams = CohereParams;
|
||||
type OcrParams = CohereOptions;
|
||||
type ProviderRequest = CohereRequest;
|
||||
type ProviderResponse = CohereResponse;
|
||||
|
||||
|
|
@ -113,34 +102,11 @@ impl BaseOcrConfig for CohereParseConfig {
|
|||
&["output_format", "req_format"]
|
||||
}
|
||||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
_model: &str,
|
||||
) -> 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)?;
|
||||
}
|
||||
Ok(optional_params
|
||||
.iter()
|
||||
.chain(overrides.iter())
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn async_transform_ocr_request(
|
||||
&self,
|
||||
model: &str,
|
||||
document: OcrDocument,
|
||||
optional_params: &CohereParams,
|
||||
optional_params: &CohereOptions,
|
||||
headers: &[(String, String)],
|
||||
_context: OcrRequestContext<'_>,
|
||||
) -> Result<CohereRequest, crate::ocr::Error> {
|
||||
|
|
@ -162,7 +128,7 @@ impl CohereParseConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
|
||||
let headers = self.validate_environment(&request.connection, &credential_env)?;
|
||||
let url = self.get_complete_url(
|
||||
request
|
||||
|
|
@ -184,7 +150,7 @@ impl CohereParseConfig {
|
|||
)
|
||||
.await?;
|
||||
transform_request_body(client, request, &url, &headers, true, body, |body| {
|
||||
validate_document(&body.document.as_document())
|
||||
validate_document(&crate::ocr::prepare::body_document(body)?)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
|
@ -236,7 +202,7 @@ fn image_url(document: OcrDocument) -> Result<String, crate::ocr::Error> {
|
|||
Ok(image_url)
|
||||
}
|
||||
|
||||
fn build_request(model: &str, image_url: String, params: &CohereParams) -> CohereRequest {
|
||||
fn build_request(model: &str, image_url: String, params: &CohereOptions) -> CohereRequest {
|
||||
CohereRequest {
|
||||
model: model.into(),
|
||||
document: CohereParseDocument::ImageUrl { image_url },
|
||||
|
|
@ -358,43 +324,62 @@ 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
|
||||
#[tokio::test]
|
||||
async fn composed_body_preserves_native_document_fields_and_untyped_overrides() {
|
||||
let mut request = crate::ocr::test_support::wire_request(
|
||||
"cohere/parse",
|
||||
"https://example.com",
|
||||
json!({
|
||||
"output_format":"markdown", "metadata":{"host":true},
|
||||
"extra_body":{
|
||||
"output_format": {"future":true},
|
||||
"document":{"type":"image_url","image_url":"https://example.com/a.png",
|
||||
"provider_options":{"nested":[false,0,null]}}
|
||||
}
|
||||
}),
|
||||
);
|
||||
request.document = serde_json::from_value(json!({
|
||||
"type":"image_url","image_url":"https://example.com/original.png"
|
||||
}))
|
||||
.unwrap();
|
||||
let overrides = serde_json::from_value(json!({
|
||||
"output_format":null, "req_format":null, "ignored":true
|
||||
let http = CohereParseConfig
|
||||
.prepare_request(&request, &crate::ocr::test_support::ocr_client())
|
||||
.await
|
||||
.unwrap();
|
||||
let body: Value = serde_json::from_slice(http.body().unwrap().as_bytes().unwrap()).unwrap();
|
||||
assert_eq!(
|
||||
body,
|
||||
json!({
|
||||
"model":"parse", "output_format":{"future":true},
|
||||
"document":{"type":"image_url","image_url":"https://example.com/a.png",
|
||||
"provider_options":{"nested":[false,0,null]}}
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn options_read_known_fields_without_changing_arguments() {
|
||||
let arguments = serde_json::from_value(json!({
|
||||
"output_format":"blocks", "req_format":"native", "extension":false
|
||||
}))
|
||||
.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")
|
||||
.map_ocr_params(&arguments, "parse")
|
||||
} else {
|
||||
CohereParseConfig.map_ocr_params(&overrides, &supplied, "parse")
|
||||
CohereParseConfig.map_ocr_params(&arguments, "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()
|
||||
assert_eq!(
|
||||
serde_json::to_value(mapped).unwrap(),
|
||||
json!({"output_format":"blocks"})
|
||||
);
|
||||
}
|
||||
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);
|
||||
assert_eq!(arguments["req_format"], "native");
|
||||
assert_eq!(arguments["extension"], false);
|
||||
let invalid = serde_json::from_value(json!({"output_format":"html"})).unwrap();
|
||||
assert!(CohereParseConfig.map_ocr_params(&invalid, "parse").is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -463,7 +448,7 @@ mod tests {
|
|||
)
|
||||
.unwrap();
|
||||
let params = CohereParseConfig
|
||||
.parse_options(&arguments, "parse")
|
||||
.map_ocr_params(&arguments, "parse")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(¶ms).unwrap(),
|
||||
|
|
@ -668,10 +653,10 @@ mod tests {
|
|||
Err(crate::ocr::Error::CohereImageOnly)
|
||||
);
|
||||
}
|
||||
assert!(serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err());
|
||||
assert!(serde_json::from_value::<CohereOptions>(json!({"output_format":"html"})).is_err());
|
||||
for format in ["markdown", "blocks"] {
|
||||
assert!(
|
||||
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
|
||||
serde_json::from_value::<CohereOptions>(json!({"output_format":format})).is_ok()
|
||||
);
|
||||
}
|
||||
let request = CohereParseConfig
|
||||
|
|
|
|||
|
|
@ -3,7 +3,6 @@ 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::{
|
||||
|
|
@ -78,17 +77,6 @@ impl BaseOcrConfig for MistralOCRConfig {
|
|||
]
|
||||
}
|
||||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OcrArguments,
|
||||
_optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
Ok(non_default_params
|
||||
.select(self.get_supported_ocr_params(model))
|
||||
.into())
|
||||
}
|
||||
|
||||
async fn async_transform_ocr_request(
|
||||
&self,
|
||||
model: &str,
|
||||
|
|
@ -115,7 +103,7 @@ impl MistralOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let params = self.map_ocr_params(&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
|
||||
|
|
@ -275,18 +263,17 @@ mod tests {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn parameter_mapping_uses_only_supported_non_default_values() {
|
||||
fn map_ocr_params_selects_known_fields_without_changing_arguments() {
|
||||
let input =
|
||||
serde_json::from_value(json!({"pages":null,"extract_header":false,"unknown":true}))
|
||||
.unwrap();
|
||||
let supplied = serde_json::from_value(json!({"pages":[9],"extension":true})).unwrap();
|
||||
let params = MistralOCRConfig
|
||||
.map_ocr_params(&input, &supplied, "model")
|
||||
.unwrap();
|
||||
let params = MistralOCRConfig.map_ocr_params(&input, "model").unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(params).unwrap(),
|
||||
json!({"pages":null,"extract_header":false})
|
||||
);
|
||||
assert_eq!(input["unknown"], true);
|
||||
assert_eq!(input.get("pages"), Some(&Value::Null));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -328,13 +315,8 @@ mod tests {
|
|||
}
|
||||
|
||||
fn mapped_params(value: Value) -> Value {
|
||||
let params = serde_json::from_value::<OcrArguments>(value).unwrap();
|
||||
serde_json::to_value(
|
||||
MistralOCRConfig
|
||||
.map_ocr_params(¶ms, &OcrArguments::default(), "model")
|
||||
.unwrap(),
|
||||
)
|
||||
.unwrap()
|
||||
let params = serde_json::from_value(value).unwrap();
|
||||
serde_json::to_value(MistralOCRConfig.map_ocr_params(¶ms, "model").unwrap()).unwrap()
|
||||
}
|
||||
|
||||
fn document() -> OcrDocument {
|
||||
|
|
@ -402,7 +384,7 @@ mod tests {
|
|||
}
|
||||
|
||||
#[rstest]
|
||||
fn map_ocr_params_filters_unsupported_non_default_params() {
|
||||
fn map_ocr_params_excludes_extensions_from_the_provider_options() {
|
||||
let mapped = mapped_params(json!({"extract_header":true,"unsupported_param":"value"}));
|
||||
assert_eq!(mapped["extract_header"], true);
|
||||
assert!(mapped.get("unsupported_param").is_none());
|
||||
|
|
|
|||
|
|
@ -3,9 +3,9 @@ use std::collections::BTreeMap;
|
|||
use serde::{Deserialize, Deserializer, Serialize};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::call_arguments::compose_body;
|
||||
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};
|
||||
|
|
@ -87,19 +87,6 @@ impl BaseOcrConfig for ReductoParseV3Config {
|
|||
&["formatting", "retrieval", "settings"]
|
||||
}
|
||||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
Ok(map_ocr_params(
|
||||
non_default_params,
|
||||
optional_params,
|
||||
self.get_supported_ocr_params(model),
|
||||
))
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "async_transform_ocr_request",
|
||||
target = "litellm::function_trace",
|
||||
|
|
@ -136,7 +123,7 @@ impl ReductoParseV3Config {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let params = self.map_ocr_params(&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?;
|
||||
|
|
@ -152,9 +139,11 @@ impl ReductoParseV3Config {
|
|||
},
|
||||
)
|
||||
.await?;
|
||||
let body = request
|
||||
.optional_params
|
||||
.compose_body(&body, self.get_supported_ocr_params(&request.model))?;
|
||||
let body = compose_body(
|
||||
&request.optional_params,
|
||||
&body,
|
||||
self.get_supported_ocr_params(&request.model),
|
||||
)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
}
|
||||
|
|
@ -171,19 +160,6 @@ impl BaseOcrConfig for ReductoParseLegacyConfig {
|
|||
&["enhance"]
|
||||
}
|
||||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
Ok(map_ocr_params(
|
||||
non_default_params,
|
||||
optional_params,
|
||||
self.get_supported_ocr_params(model),
|
||||
))
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "async_transform_ocr_request",
|
||||
target = "litellm::function_trace",
|
||||
|
|
@ -217,7 +193,7 @@ impl ReductoParseLegacyConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let params = self.map_ocr_params(&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?;
|
||||
|
|
@ -233,29 +209,15 @@ impl ReductoParseLegacyConfig {
|
|||
},
|
||||
)
|
||||
.await?;
|
||||
let body = request
|
||||
.optional_params
|
||||
.compose_body(&body, self.get_supported_ocr_params(&request.model))?;
|
||||
let body = compose_body(
|
||||
&request.optional_params,
|
||||
&body,
|
||||
self.get_supported_ocr_params(&request.model),
|
||||
)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
}
|
||||
|
||||
fn map_ocr_params(
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
supported_params: &[&str],
|
||||
) -> OcrArguments {
|
||||
optional_params
|
||||
.iter()
|
||||
.chain(
|
||||
non_default_params
|
||||
.iter()
|
||||
.filter(|(name, _)| supported_params.contains(&name.as_str())),
|
||||
)
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
|
||||
deserializer: D,
|
||||
) -> Result<Option<Option<T>>, D::Error> {
|
||||
|
|
@ -499,31 +461,27 @@ 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();
|
||||
fn options_preserve_null_and_select_the_provider_fields() {
|
||||
let overrides = serde_json::from_value(json!({
|
||||
"formatting":null, "enhance":null, "ignored":true
|
||||
}))
|
||||
.unwrap();
|
||||
let v3 = ReductoParseV3Config
|
||||
.map_ocr_params(&overrides, &supplied, "parse-v3")
|
||||
.map_ocr_params(&overrides, "parse-v3")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(v3).unwrap(),
|
||||
json!({
|
||||
"formatting":null, "enhance":true, "extension":false
|
||||
"formatting":null
|
||||
})
|
||||
);
|
||||
let legacy = ReductoParseLegacyConfig
|
||||
.map_ocr_params(&overrides, &supplied, "parse-legacy")
|
||||
.map_ocr_params(&overrides, "parse-legacy")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(legacy).unwrap(),
|
||||
json!({
|
||||
"formatting":{"old":true}, "enhance":null, "extension":false
|
||||
"enhance":null
|
||||
})
|
||||
);
|
||||
}
|
||||
|
|
@ -558,7 +516,7 @@ mod tests {
|
|||
serde_json::from_value(json!({"formatting":null,"settings":{},"unknown":true}))
|
||||
.unwrap();
|
||||
let params = ReductoParseV3Config
|
||||
.parse_options(&overrides, "parse-v3")
|
||||
.map_ocr_params(&overrides, "parse-v3")
|
||||
.unwrap();
|
||||
let client = crate::ocr::test_support::ocr_client();
|
||||
let connection = OcrConnection::default();
|
||||
|
|
@ -586,7 +544,7 @@ mod tests {
|
|||
})
|
||||
);
|
||||
let absent = ReductoParseV3Config
|
||||
.parse_options(&OcrArguments::default(), "parse-v3")
|
||||
.map_ocr_params(&crate::call_arguments::CallArguments::default(), "parse-v3")
|
||||
.unwrap();
|
||||
assert_eq!(serde_json::to_value(absent).unwrap(), json!({}));
|
||||
}
|
||||
|
|
@ -603,7 +561,7 @@ mod tests {
|
|||
let overrides =
|
||||
serde_json::from_value(json!({"enhance":value,"unknown":true})).unwrap();
|
||||
let params = ReductoParseLegacyConfig
|
||||
.parse_options(&overrides, "parse-legacy")
|
||||
.map_ocr_params(&overrides, "parse-legacy")
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(build_legacy_body(
|
||||
|
|
|
|||
|
|
@ -154,7 +154,7 @@ impl VertexAIDeepSeekOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
|
||||
let config = VertexConfig::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
|
|
@ -396,17 +396,25 @@ mod tests {
|
|||
use super::{VertexAIDeepSeekOCRConfig, provider_model};
|
||||
|
||||
#[test]
|
||||
fn inherited_parameter_mapping_only_returns_supplied_optional_params() {
|
||||
fn unconsumed_options_remain_available_for_body_composition() {
|
||||
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();
|
||||
let arguments =
|
||||
serde_json::from_value(json!({"temperature":0.5,"extension":null})).unwrap();
|
||||
assert_eq!(
|
||||
VertexAIDeepSeekOCRConfig
|
||||
.map_ocr_params(&non_default, &supplied, "deepseek-ocr")
|
||||
serde_json::to_value(
|
||||
VertexAIDeepSeekOCRConfig
|
||||
.map_ocr_params(&arguments, "deepseek-ocr")
|
||||
.unwrap()
|
||||
)
|
||||
.unwrap(),
|
||||
json!({})
|
||||
);
|
||||
assert_eq!(
|
||||
crate::call_arguments::compose_body(&arguments, &json!({"model":"deepseek-ocr"}), &[])
|
||||
.unwrap(),
|
||||
supplied
|
||||
json!({"model":"deepseek-ocr","temperature":0.5,"extension":null})
|
||||
);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ 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};
|
||||
|
|
@ -24,15 +23,6 @@ impl BaseOcrConfig for VertexAIOCRConfig {
|
|||
MistralOCRConfig.get_supported_ocr_params(model)
|
||||
}
|
||||
|
||||
fn map_ocr_params(
|
||||
&self,
|
||||
non_default_params: &OcrArguments,
|
||||
optional_params: &OcrArguments,
|
||||
model: &str,
|
||||
) -> Result<OcrArguments, crate::ocr::Error> {
|
||||
MistralOCRConfig.map_ocr_params(non_default_params, optional_params, model)
|
||||
}
|
||||
|
||||
async fn async_transform_ocr_request(
|
||||
&self,
|
||||
model: &str,
|
||||
|
|
@ -65,7 +55,7 @@ impl VertexAIOCRConfig {
|
|||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, crate::ocr::Error> {
|
||||
let params = self.parse_options(&request.optional_params, &request.model)?;
|
||||
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
|
||||
let config = VertexConfig::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
|
|
@ -103,7 +93,7 @@ impl VertexAIOCRConfig {
|
|||
&authentication.headers,
|
||||
retains_document,
|
||||
body,
|
||||
|body| validate_inline_document(&body.document),
|
||||
|body| validate_inline_document(&crate::ocr::prepare::body_document(body)?),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,93 +0,0 @@
|
|||
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()
|
||||
}
|
||||
}
|
||||
|
|
@ -88,6 +88,14 @@ pub enum Error {
|
|||
Headers(#[from] crate::http_utils::HeaderError),
|
||||
}
|
||||
|
||||
impl From<crate::call_arguments::ArgumentError> for Error {
|
||||
fn from(error: crate::call_arguments::ArgumentError) -> Self {
|
||||
Self::RequestField {
|
||||
path: format!("optional_params.{}", error.path),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Error {
|
||||
pub fn is_request(&self) -> bool {
|
||||
matches!(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,4 @@
|
|||
mod arguments;
|
||||
mod error;
|
||||
pub use arguments::OcrArguments;
|
||||
pub use error::Error;
|
||||
pub mod client;
|
||||
pub(crate) mod document;
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use serde::{Serialize, de::DeserializeOwned};
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::OcrClient;
|
||||
|
|
@ -12,21 +12,19 @@ pub(crate) async fn transform_request_body<B>(
|
|||
headers: &[(String, String)],
|
||||
retains_document: bool,
|
||||
body: B,
|
||||
validate: impl Fn(&B) -> Result<(), super::Error>,
|
||||
validate: impl Fn(&Value) -> Result<(), super::Error>,
|
||||
) -> Result<reqwest::Request, super::Error>
|
||||
where
|
||||
B: Serialize + DeserializeOwned,
|
||||
B: Serialize,
|
||||
{
|
||||
let composed = request.optional_params.compose_body(
|
||||
let composed = crate::call_arguments::compose_body(
|
||||
&request.optional_params,
|
||||
&body,
|
||||
request.config.get_supported_ocr_params(&request.model),
|
||||
)?;
|
||||
let composed = OcrWireBody::<B>::decode(composed, "body")?;
|
||||
validate(&composed.body)?;
|
||||
validate(&composed)?;
|
||||
let (body, headers) = if request.hooks.intercepts_requests() {
|
||||
let body = serde_json::to_value(composed).map_err(|_| super::Error::RequestField {
|
||||
path: "body".into(),
|
||||
})?;
|
||||
let body = composed;
|
||||
let retained_fields = request
|
||||
.optional_params
|
||||
.keys()
|
||||
|
|
@ -52,9 +50,13 @@ where
|
|||
retained_fields,
|
||||
})
|
||||
.await?;
|
||||
let body = OcrWireBody::<B>::decode(changed.body, "guardrail.body")?;
|
||||
validate(&body.body)?;
|
||||
(body, changed.headers)
|
||||
if !changed.body.is_object() {
|
||||
return Err(super::Error::RequestField {
|
||||
path: "guardrail.body".into(),
|
||||
});
|
||||
}
|
||||
validate(&changed.body)?;
|
||||
(changed.body, changed.headers)
|
||||
} else {
|
||||
(composed, headers.to_vec())
|
||||
};
|
||||
|
|
@ -106,24 +108,19 @@ pub(crate) async fn guardrail_document(
|
|||
Ok((document, changed.headers))
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct OcrWireBody<B> {
|
||||
#[serde(skip)]
|
||||
body: B,
|
||||
#[serde(flatten)]
|
||||
fields: serde_json::Map<String, Value>,
|
||||
}
|
||||
|
||||
impl<B: Serialize + DeserializeOwned> OcrWireBody<B> {
|
||||
fn decode(value: Value, prefix: &str) -> Result<Self, super::Error> {
|
||||
let body: B = super::wire::decode_request_value(value.clone(), prefix)?;
|
||||
let Value::Object(fields) = value else {
|
||||
return Err(super::Error::RequestField {
|
||||
path: prefix.into(),
|
||||
});
|
||||
};
|
||||
Ok(Self { body, fields })
|
||||
}
|
||||
pub(crate) fn body_document(body: &Value) -> Result<OcrDocument, super::Error> {
|
||||
let document = body
|
||||
.get("document")
|
||||
.and_then(Value::as_object)
|
||||
.ok_or_else(|| super::Error::RequestField {
|
||||
path: "body.document".into(),
|
||||
})?;
|
||||
let source = document
|
||||
.iter()
|
||||
.filter(|(name, _)| matches!(name.as_str(), "type" | "image_url" | "document_url"))
|
||||
.map(|(name, value)| (name.clone(), value.clone()))
|
||||
.collect();
|
||||
super::wire::decode_request_value(Value::Object(source), "body.document")
|
||||
}
|
||||
|
||||
pub(crate) fn credential_env(name: &str) -> Option<String> {
|
||||
|
|
|
|||
|
|
@ -7,9 +7,9 @@ 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::call_arguments::CallArguments;
|
||||
use crate::constants::OCR_HTTP_TIMEOUT_SECS;
|
||||
use litellm_auth::{InputSource, TokenProviderHandle};
|
||||
|
||||
|
|
@ -97,7 +97,7 @@ pub struct LiteLLMOcrRequest {
|
|||
pub connection: OcrConnection,
|
||||
pub hooks: Arc<dyn OcrHooks>,
|
||||
pub litellm_call_id: Option<String>,
|
||||
pub optional_params: OcrArguments,
|
||||
pub optional_params: CallArguments,
|
||||
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: OcrArguments,
|
||||
optional_params: CallArguments,
|
||||
) -> 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::call_arguments::{ArgumentSpec, CallArguments};
|
||||
use litellm_auth::InputSource;
|
||||
use serde::{
|
||||
Deserialize,
|
||||
|
|
@ -11,6 +11,7 @@ use serde::{
|
|||
use serde_json::{Map, Value};
|
||||
|
||||
const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body", "max_response_bytes"];
|
||||
pub const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"];
|
||||
const AZURE_AUTH_OPTION_FIELDS: &[&str] = &[
|
||||
"azure_ad_token",
|
||||
"tenant_id",
|
||||
|
|
@ -31,12 +32,6 @@ const VERTEX_AUTH_OPTION_FIELDS: &[&str] = &[
|
|||
"vertex_ai_location",
|
||||
];
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub struct OptionalParamSpec {
|
||||
pub name: &'static str,
|
||||
pub secret: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct DecodedOcrResponse<T> {
|
||||
pub data: T,
|
||||
|
|
@ -54,7 +49,7 @@ pub struct OcrWireRequest {
|
|||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
#[serde(default)]
|
||||
pub optional_params: OcrArguments,
|
||||
pub optional_params: CallArguments,
|
||||
#[serde(default)]
|
||||
pub input_sources: BTreeMap<String, InputSource>,
|
||||
pub timeout_seconds: Option<f64>,
|
||||
|
|
@ -91,11 +86,11 @@ pub fn consumed_optional_param_names(
|
|||
pub fn consumed_optional_params(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<Vec<OptionalParamSpec>, crate::ocr::Error> {
|
||||
) -> Result<Vec<ArgumentSpec>, crate::ocr::Error> {
|
||||
consumed_optional_param_names(model, custom_llm_provider).map(|names| {
|
||||
names
|
||||
.into_iter()
|
||||
.map(|name| OptionalParamSpec {
|
||||
.map(|name| ArgumentSpec {
|
||||
name,
|
||||
secret: matches!(
|
||||
name,
|
||||
|
|
@ -110,17 +105,6 @@ 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");
|
||||
|
|
@ -269,14 +253,14 @@ mod tests {
|
|||
#[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));
|
||||
use crate::call_arguments::should_project;
|
||||
assert!(should_project("future_option", &fields, BOUND_FIELDS));
|
||||
assert!(should_project("extra_body", &fields, BOUND_FIELDS));
|
||||
assert!(should_project("id", &fields, BOUND_FIELDS));
|
||||
assert!(!should_project("metadata", &fields, BOUND_FIELDS));
|
||||
assert!(!should_project("callbacks", &fields, BOUND_FIELDS));
|
||||
assert!(!should_project("api_key", &fields, BOUND_FIELDS));
|
||||
assert!(!should_project("document", &fields, BOUND_FIELDS));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -90,19 +90,15 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py<PyAny>) -> PyRe
|
|||
|
||||
pub(crate) fn project_optional_fields(
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
fields: &[litellm_core::ocr::wire::OptionalParamSpec],
|
||||
fields: &[litellm_core::call_arguments::ArgumentSpec],
|
||||
bound_fields: &[&str],
|
||||
) -> PyResult<Map<String, Value>> {
|
||||
let controls: Vec<String> = kwargs
|
||||
.py()
|
||||
.import("litellm.types.utils")?
|
||||
.getattr("all_litellm_params")?
|
||||
.extract()?;
|
||||
kwargs
|
||||
.iter()
|
||||
.map(|(name, value)| Ok((name.extract::<String>()?, value)))
|
||||
.filter_map(|entry: PyResult<_>| match entry {
|
||||
Ok((name, value))
|
||||
if litellm_core::ocr::wire::project_argument(&name, fields, &controls) =>
|
||||
if litellm_core::call_arguments::should_project(&name, fields, bound_fields) =>
|
||||
{
|
||||
Some(from_py(&value).map(|value| (name, value)))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -29,12 +29,12 @@ pub(super) struct ProjectedOcrCall {
|
|||
pub fields: ProjectedOcrFields,
|
||||
}
|
||||
|
||||
struct OcrArguments<'a, 'py> {
|
||||
struct PythonOcrFields<'a, 'py> {
|
||||
request: &'a Bound<'py, PyAny>,
|
||||
kwargs: &'a Bound<'py, PyDict>,
|
||||
}
|
||||
|
||||
impl<'py> OcrArguments<'_, 'py> {
|
||||
impl<'py> PythonOcrFields<'_, 'py> {
|
||||
fn lookup(&self, name: &str) -> PyResult<Bound<'py, PyAny>> {
|
||||
match self.kwargs.get_item(name)? {
|
||||
Some(value) => Ok(value),
|
||||
|
|
@ -116,7 +116,7 @@ pub(super) fn project_request(
|
|||
kwargs: &Bound<'_, PyDict>,
|
||||
) -> PyResult<ProjectedOcrCall> {
|
||||
let boundary_request = request.clone().unbind();
|
||||
let arguments = OcrArguments { request, kwargs };
|
||||
let arguments = PythonOcrFields { request, kwargs };
|
||||
let model = arguments.model()?;
|
||||
let custom_llm_provider = arguments.custom_llm_provider()?;
|
||||
let (wire_document, retained_document) =
|
||||
|
|
@ -124,7 +124,8 @@ 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 optional_params = project_optional_fields(kwargs, &specs)?;
|
||||
let optional_params =
|
||||
project_optional_fields(kwargs, &specs, litellm_core::ocr::wire::BOUND_FIELDS)?;
|
||||
let input_sources = request_input_sources(
|
||||
kwargs,
|
||||
optional_params
|
||||
|
|
@ -191,8 +192,8 @@ mod tests {
|
|||
fn arguments<'a, 'py>(
|
||||
request: &'a Bound<'py, PyAny>,
|
||||
kwargs: &'a Bound<'py, PyDict>,
|
||||
) -> OcrArguments<'a, 'py> {
|
||||
OcrArguments { request, kwargs }
|
||||
) -> PythonOcrFields<'a, 'py> {
|
||||
PythonOcrFields { request, kwargs }
|
||||
}
|
||||
|
||||
fn project_document(
|
||||
|
|
|
|||
|
|
@ -93,6 +93,49 @@ async def test_request_level_custom_pricing_reaches_logging_params_and_bills_the
|
|||
assert "ocr_cost_per_page" not in ocr_server.requests[0].body
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
async def test_shared_boundary_preserves_provider_fields_and_python_objects(
|
||||
ocr_server: RecordingServer, asynchronous: bool
|
||||
) -> None:
|
||||
marker: Final = object()
|
||||
observed: Final = []
|
||||
|
||||
class Observe(Logging):
|
||||
def pre_call(self, input, api_key, additional_args):
|
||||
observed.append(self.model_call_details["litellm_params"]["metadata"]["marker"])
|
||||
|
||||
logger: Final = Observe(
|
||||
model="mistral-ocr-latest",
|
||||
messages=[],
|
||||
stream=False,
|
||||
call_type="aocr" if asynchronous else "ocr",
|
||||
start_time=datetime.datetime.now(),
|
||||
litellm_call_id="shared-boundary",
|
||||
function_id="shared-boundary",
|
||||
)
|
||||
arguments: Final = {
|
||||
"id": "provider-id",
|
||||
"future_option": {"old": True},
|
||||
"explicit_null": None,
|
||||
"metadata": {"marker": marker},
|
||||
"shared_session": marker,
|
||||
"litellm_logging_obj": logger,
|
||||
"extra_body": {"future_option": {"nested": [None, False, 0]}, "metadata": {"provider": True}},
|
||||
}
|
||||
if asynchronous:
|
||||
await call_aocr(ocr_server, **arguments)
|
||||
else:
|
||||
call_ocr(ocr_server, **arguments)
|
||||
body: Final = ocr_server.requests[0].body
|
||||
assert body["id"] == "provider-id"
|
||||
assert body["future_option"] == {"nested": [None, False, 0]}
|
||||
assert "explicit_null" in body and body["explicit_null"] is None
|
||||
assert body["metadata"] == {"provider": True}
|
||||
assert "shared_session" not in body
|
||||
assert observed == [marker] and observed[0] is marker
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_replacement_finalized_before_dispatch_in_caller_task(ocr_server: RecordingServer) -> None:
|
||||
caller: Final = asyncio.current_task()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue