diff --git a/litellm-rust/crates/core/src/audio_transcription/handler.rs b/litellm-rust/crates/core/src/audio_transcription/handler.rs index f4557aa3c5f..6658e95558a 100644 --- a/litellm-rust/crates/core/src/audio_transcription/handler.rs +++ b/litellm-rust/crates/core/src/audio_transcription/handler.rs @@ -1,6 +1,5 @@ use serde_json::Value; -use crate::audio_transcription::Error; use crate::http_utils::{http_request, truncate_error_body}; use super::client::http_client; @@ -8,9 +7,10 @@ use super::types::ProviderAudioTranscriptionRequest; pub async fn execute_audio_transcription_provider_call( request: ProviderAudioTranscriptionRequest, -) -> Result { - let body = serde_json::to_vec(&request.body) - .map_err(|error| Error::InvalidRequest(format!("invalid audio request body: {error}")))?; +) -> Result { + let body = serde_json::to_vec(&request.body).map_err(|error| { + super::Error::InvalidRequest(format!("invalid audio request body: {error}")) + })?; let headers = signed_headers(&request, &body).await?; let mut request_builder = http_client().post(&request.url).body(body); for (key, value) in headers { @@ -19,22 +19,22 @@ pub async fn execute_audio_transcription_provider_call( if let Some(duration) = request.timeout { request_builder = request_builder.timeout(duration); } - let response = http_request(request_builder) - .await - .map_err(|error| Error::Transport(crate::transport::Error::Network(error.to_string())))?; + let response = http_request(request_builder).await.map_err(|error| { + super::Error::Transport(crate::transport::Error::Network(error.to_string())) + })?; let status = response.status(); - let text = response - .text() - .await - .map_err(|error| Error::Transport(crate::transport::Error::Network(error.to_string())))?; + let text = response.text().await.map_err(|error| { + super::Error::Transport(crate::transport::Error::Network(error.to_string())) + })?; if !status.is_success() { - return Err(Error::Transport(crate::transport::Error::Http { + return Err(super::Error::Transport(crate::transport::Error::Http { status: status.as_u16(), body: truncate_error_body(&text), })); } - let response_json = serde_json::from_str(&text) - .map_err(|error| Error::InvalidResponse(format!("invalid audio response JSON: {error}")))?; + let response_json = serde_json::from_str(&text).map_err(|error| { + super::Error::InvalidResponse(format!("invalid audio response JSON: {error}")) + })?; Ok(request .config .transform_transcription_response(&request.model, response_json)? @@ -44,7 +44,7 @@ pub async fn execute_audio_transcription_provider_call( async fn signed_headers( request: &ProviderAudioTranscriptionRequest, body: &[u8], -) -> Result, Error> { +) -> Result, super::Error> { use std::collections::BTreeMap; use std::time::SystemTime; diff --git a/litellm-rust/crates/core/src/audio_transcription/prepare.rs b/litellm-rust/crates/core/src/audio_transcription/prepare.rs index 40f2012e95e..2196a53e633 100644 --- a/litellm-rust/crates/core/src/audio_transcription/prepare.rs +++ b/litellm-rust/crates/core/src/audio_transcription/prepare.rs @@ -1,4 +1,3 @@ -use crate::audio_transcription::Error; use crate::http_utils::{has_header, string_headers}; use crate::providers::bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG; use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider}; @@ -15,7 +14,7 @@ fn provider_config(provider: &str) -> Option<&'static dyn AudioTranscriptionProv pub fn prepare_audio_transcription_provider_call( request: AudioTranscriptionRequest<'_>, -) -> Result { +) -> Result { let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) .or_else(|| { request @@ -26,13 +25,14 @@ pub fn prepare_audio_transcription_provider_call( }) }) .ok_or_else(|| { - Error::InvalidProvider( + super::Error::InvalidProvider( "unable to resolve custom_llm_provider for audio transcription request".to_string(), ) })?; let model = provider_info.model.to_string(); - let config = provider_config(provider_info.custom_llm_provider) - .ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?; + let config = provider_config(provider_info.custom_llm_provider).ok_or_else(|| { + super::Error::InvalidProvider(provider_info.custom_llm_provider.to_string()) + })?; let env_lookup = |key: &str| std::env::var(key).ok(); let mut headers = string_headers("audio transcription", request.extra_headers)?; let auth = config.auth_strategy(&model, &request.optional_params, &env_lookup)?; diff --git a/litellm-rust/crates/core/src/audio_transcription/transformation.rs b/litellm-rust/crates/core/src/audio_transcription/transformation.rs index a35bd0f8d4a..34406981180 100644 --- a/litellm-rust/crates/core/src/audio_transcription/transformation.rs +++ b/litellm-rust/crates/core/src/audio_transcription/transformation.rs @@ -1,4 +1,3 @@ -use crate::audio_transcription::Error; use serde_json::Value; use super::types::{AudioTranscriptionRequestData, AudioTranscriptionResponseData}; @@ -26,13 +25,13 @@ pub trait AudioTranscriptionProviderConfig: Sync { model: &str, audio: Value, optional_params: OpaqueParams, - ) -> Result; + ) -> Result; fn transform_transcription_response( &self, model: &str, response_json: Value, - ) -> Result; + ) -> Result; fn complete_url( &self, @@ -40,12 +39,12 @@ pub trait AudioTranscriptionProviderConfig: Sync { model: &str, optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; fn auth_strategy( &self, model: &str, optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; } diff --git a/litellm-rust/crates/core/src/call_lifecycle/mod.rs b/litellm-rust/crates/core/src/call_lifecycle/mod.rs index e0272ede0bb..dce240c3d2b 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/mod.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/mod.rs @@ -229,7 +229,6 @@ fn epoch_seconds() -> f64 { #[cfg(test)] mod tests { use super::*; - use crate::messages::Error; use std::pin::Pin; use std::sync::Mutex; @@ -255,9 +254,9 @@ mod tests { } impl CallLifecycleHooks for RecordingHooks { - type Error = Error; - type PreCallFuture<'a> = BoxFuture<'a, Result>; - type DuringCallFuture<'a> = BoxFuture<'a, Result>; + type Error = crate::messages::Error; + type PreCallFuture<'a> = BoxFuture<'a, Result>; + type DuringCallFuture<'a> = BoxFuture<'a, Result>; type SuccessFuture<'a> = BoxFuture<'a, ()>; type FailureFuture<'a> = BoxFuture<'a, ()>; @@ -299,7 +298,7 @@ mod tests { fn async_log_failure_event<'a>( &'a self, _context: &'a CallLifecycleContext, - _error: &'a Error, + _error: &'a crate::messages::Error, _timing: &'a CallLifecycleTiming, ) -> Self::FailureFuture<'a> { Box::pin(async move { @@ -309,9 +308,9 @@ mod tests { } impl CallLifecycleHooks for RecordingHooks { - type Error = Error; - type PreCallFuture<'a> = BoxFuture<'a, Result>; - type DuringCallFuture<'a> = BoxFuture<'a, Result>; + type Error = crate::messages::Error; + type PreCallFuture<'a> = BoxFuture<'a, Result>; + type DuringCallFuture<'a> = BoxFuture<'a, Result>; type SuccessFuture<'a> = BoxFuture<'a, ()>; type FailureFuture<'a> = BoxFuture<'a, ()>; @@ -351,7 +350,7 @@ mod tests { fn async_log_failure_event<'a>( &'a self, _context: &'a CallLifecycleContext, - _error: &'a Error, + _error: &'a crate::messages::Error, _timing: &'a CallLifecycleTiming, ) -> Self::FailureFuture<'a> { Box::pin(async move { @@ -389,9 +388,9 @@ mod tests { "request".to_string(), &hooks, |_request| async move { - Err::(Error::Transport(crate::transport::Error::Network( - "provider down".to_string(), - ))) + Err::(crate::messages::Error::Transport( + crate::transport::Error::Network("provider down".to_string()), + )) }, ) .await @@ -399,7 +398,7 @@ mod tests { assert_eq!( error, - Error::Transport(crate::transport::Error::Network( + crate::messages::Error::Transport(crate::transport::Error::Network( "provider down".to_string() )) ); diff --git a/litellm-rust/crates/core/src/chat_completions/common_utils.rs b/litellm-rust/crates/core/src/chat_completions/common_utils.rs index ee80173e58e..1f3343e7b89 100644 --- a/litellm-rust/crates/core/src/chat_completions/common_utils.rs +++ b/litellm-rust/crates/core/src/chat_completions/common_utils.rs @@ -1,4 +1,3 @@ -use crate::chat_completions::Error; use crate::http_utils::string_headers as shared_string_headers; use crate::providers::anthropic::chat_completions::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG; use serde_json::{Map, Value}; @@ -21,6 +20,6 @@ pub(super) fn chat_completions_provider_config( pub(super) fn string_headers( extra_headers: Option>, -) -> Result, Error> { - shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(Error::from) +) -> Result, super::Error> { + shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(super::Error::from) } diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs index ac98a5d11cc..842d63f3afb 100644 --- a/litellm-rust/crates/core/src/chat_completions/handler.rs +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -1,6 +1,5 @@ use serde_json::Value; -use crate::chat_completions::Error; use crate::http_utils::{http_request, truncate_error_body}; use super::client::http_client; @@ -13,10 +12,10 @@ use super::types::{ pub(super) async fn execute_chat_completions_provider_call( request: ResolvedChatCompletionsRequest<'_>, -) -> Result { +) -> Result { let request = prepare_provider_request(request)?; let body = serde_json::to_vec(&request.body).map_err(|err| { - Error::InvalidRequest(format!( + super::Error::InvalidRequest(format!( "failed to serialize chat completions request: {err}" )) })?; @@ -35,20 +34,19 @@ pub(super) async fn execute_chat_completions_provider_call( .map_err(crate::transport::Error::from_reqwest_before_dispatch)?; let status = response.status(); - let text = response - .text() - .await - .map_err(|err| Error::Transport(crate::transport::Error::Network(err.to_string())))?; + let text = response.text().await.map_err(|err| { + super::Error::Transport(crate::transport::Error::Network(err.to_string())) + })?; if !status.is_success() { - return Err(Error::Transport(crate::transport::Error::Http { + return Err(super::Error::Transport(crate::transport::Error::Http { status: status.as_u16(), body: truncate_error_body(&text), })); } let body: Value = serde_json::from_str(&text).map_err(|err| { - Error::InvalidResponse(format!("invalid chat completions response JSON: {err}")) + super::Error::InvalidResponse(format!("invalid chat completions response JSON: {err}")) })?; request .config @@ -65,19 +63,19 @@ pub(super) async fn execute_chat_completions_provider_call( /// second kind has already been billed, and a host that keeps a reference /// implementation must not retry those, so collapse them to one variant that /// can only mean the provider was already called. -pub(super) fn as_response_error(err: Error) -> Error { +pub(super) fn as_response_error(err: super::Error) -> super::Error { match err { - already @ (Error::InvalidResponse(_) - | Error::ResponseTransform(_) - | Error::Transport(crate::transport::Error::Http { .. })) => already, - other => Error::ResponseTransform(Box::new(other)), + already @ (super::Error::InvalidResponse(_) + | super::Error::ResponseTransform(_) + | super::Error::Transport(crate::transport::Error::Http { .. })) => already, + other => super::Error::ResponseTransform(Box::new(other)), } } pub(super) async fn signed_headers( request: &ProviderChatCompletionsRequest, body: &[u8], -) -> Result, Error> { +) -> Result, super::Error> { use std::collections::BTreeMap; use std::time::SystemTime; @@ -98,7 +96,7 @@ pub(super) async fn signed_headers( .iter() .any(|(name, _)| is_sigv4_computed_header(name)) { - return Err(Error::Unsupported( + return Err(super::Error::Unsupported( "request forwards a header AWS SigV4 computes", )); } diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index 19da2113aa3..489ea0ad701 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -1,6 +1,5 @@ use serde_json::Value; -use crate::chat_completions::Error; use crate::http_utils::has_header; use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider}; @@ -14,7 +13,7 @@ use super::types::{ pub(super) fn resolve_provider_config<'a>( model: &'a str, custom_llm_provider: Option<&'a str>, -) -> Result<(String, &'static dyn ChatCompletionsProviderConfig), Error> { +) -> Result<(String, &'static dyn ChatCompletionsProviderConfig), super::Error> { let provider_info = get_custom_llm_provider(model, custom_llm_provider) .or_else(|| { custom_llm_provider.map(|provider| CustomLlmProvider { @@ -23,33 +22,36 @@ pub(super) fn resolve_provider_config<'a>( }) }) .ok_or_else(|| { - Error::InvalidProvider( + super::Error::InvalidProvider( "unable to resolve custom_llm_provider for chat completions request".to_string(), ) })?; - let config = chat_completions_provider_config(provider_info.custom_llm_provider) - .ok_or_else(|| Error::InvalidProvider(provider_info.custom_llm_provider.to_string()))?; + let config = + chat_completions_provider_config(provider_info.custom_llm_provider).ok_or_else(|| { + super::Error::InvalidProvider(provider_info.custom_llm_provider.to_string()) + })?; Ok((provider_info.model.to_string(), config)) } -pub(super) fn parse_messages(messages: Value) -> Result, Error> { - serde_json::from_value(messages) - .map_err(|err| Error::InvalidRequest(format!("invalid chat completions messages: {err}"))) +pub(super) fn parse_messages(messages: Value) -> Result, super::Error> { + serde_json::from_value(messages).map_err(|err| { + super::Error::InvalidRequest(format!("invalid chat completions messages: {err}")) + }) } pub(super) fn resolve_request( request: ChatCompletionsRequest<'_>, -) -> Result, Error> { +) -> Result, super::Error> { let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?; let messages = parse_messages(request.messages)?; if messages.is_empty() { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "chat completions requires at least one message".to_string(), )); } request.optional_params.clone().into_provider_body()?; if let Some(reason) = config.unsupported_reason(&messages, &request.optional_params) { - return Err(Error::Unsupported(reason.0)); + return Err(super::Error::Unsupported(reason.0)); } Ok(ResolvedChatCompletionsRequest { model, @@ -67,7 +69,7 @@ fn validate_environment( request: &ResolvedChatCompletionsRequest<'_>, model: &str, config: &dyn ChatCompletionsProviderConfig, -) -> Result<(Vec<(String, String)>, ChatCompletionsAuth), Error> { +) -> Result<(Vec<(String, String)>, ChatCompletionsAuth), super::Error> { let env_lookup = |key: &str| std::env::var(key).ok(); let mut headers = string_headers(request.extra_headers.clone())?; let auth = config.auth( @@ -118,7 +120,7 @@ fn validate_environment( pub(super) fn prepare_provider_request( request: ResolvedChatCompletionsRequest<'_>, -) -> Result { +) -> Result { let (headers, auth) = validate_environment(&request, &request.model, request.config)?; let model = request.model; let config = request.config; diff --git a/litellm-rust/crates/core/src/chat_completions/tests.rs b/litellm-rust/crates/core/src/chat_completions/tests.rs index 71fe5d3ac82..d030c75189d 100644 --- a/litellm-rust/crates/core/src/chat_completions/tests.rs +++ b/litellm-rust/crates/core/src/chat_completions/tests.rs @@ -1,14 +1,12 @@ use serde_json::{Map, Value, json}; -use crate::chat_completions::Error; - use super::prepare::{prepare_provider_request, resolve_request}; use super::transformation::ChatCompletionsAuth; use super::types::{ChatCompletionsRequest, ProviderChatCompletionsRequest}; fn prepare_chat_completions_call( request: ChatCompletionsRequest<'_>, -) -> Result { +) -> Result { prepare_provider_request(resolve_request(request)?) } @@ -35,7 +33,7 @@ fn request<'a>( /// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers /// carry resolved credentials), so unwrap the failure case by hand. -fn decline(request: ChatCompletionsRequest<'_>) -> Error { +fn decline(request: ChatCompletionsRequest<'_>) -> crate::chat_completions::Error { match prepare_chat_completions_call(request) { Err(error) => error, Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url), @@ -202,7 +200,10 @@ fn declines_an_unsupported_request_before_resolving_credentials() { call.api_key = None; // No api_key is set and no env is consulted: the gate must run first, so the // error is the decline rather than a missing-credential error. - assert_eq!(decline(call), Error::Unsupported("streaming")); + assert_eq!( + decline(call), + crate::chat_completions::Error::Unsupported("streaming") + ); } #[test] @@ -214,7 +215,7 @@ fn rejects_an_unknown_provider() { json!([{"role": "user", "content": "hi"}]), json!({}), )), - Error::InvalidProvider("openai".to_string()) + crate::chat_completions::Error::InvalidProvider("openai".to_string()) ); } @@ -227,7 +228,7 @@ fn rejects_a_model_with_no_resolvable_provider() { json!([{"role": "user", "content": "hi"}]), json!({}), )), - Error::InvalidProvider(_) + crate::chat_completions::Error::InvalidProvider(_) )); } @@ -240,7 +241,9 @@ fn rejects_an_empty_or_malformed_message_list() { json!([]), json!({}), )), - Error::InvalidRequest("chat completions requires at least one message".to_string()) + crate::chat_completions::Error::InvalidRequest( + "chat completions requires at least one message".to_string() + ) ); assert!(matches!( decline(request( @@ -249,7 +252,7 @@ fn rejects_an_empty_or_malformed_message_list() { json!("not a list"), json!({}), )), - Error::InvalidRequest(_) + crate::chat_completions::Error::InvalidRequest(_) )); } @@ -264,7 +267,7 @@ fn rejects_non_string_extra_headers() { call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))])); assert_eq!( decline(call), - Error::Headers(crate::http_utils::HeaderError { + crate::chat_completions::Error::Headers(crate::http_utils::HeaderError { context: "chat completions", name: "x-trace".to_string(), actual: "number", @@ -379,7 +382,7 @@ async fn a_forwarded_header_the_signer_computes_declines_to_python() { .await .expect_err("{forwarded} should decline instead of being signed"); assert!( - matches!(error, Error::Unsupported(_)), + matches!(error, crate::chat_completions::Error::Unsupported(_)), "{forwarded} declined as {error:?}, which the host would not fall back on" ); } @@ -730,7 +733,11 @@ mod round_trip { .expect_err("response cannot be normalized"); handle.await.expect("server task"); assert!( - matches!(err, Error::InvalidResponse(_) | Error::ResponseTransform(_)), + matches!( + err, + crate::chat_completions::Error::InvalidResponse(_) + | crate::chat_completions::Error::ResponseTransform(_) + ), "expected a post-send error, got {err:?}" ); } @@ -748,7 +755,11 @@ mod round_trip { .expect_err("response cannot be normalized"); handle.await.expect("server task"); assert!( - matches!(err, Error::InvalidResponse(_) | Error::ResponseTransform(_)), + matches!( + err, + crate::chat_completions::Error::InvalidResponse(_) + | crate::chat_completions::Error::ResponseTransform(_) + ), "expected a post-send error, got {err:?}" ); } @@ -768,7 +779,10 @@ mod round_trip { assert!( matches!( err, - Error::Transport(crate::transport::Error::Http { status: 429, .. }) + crate::chat_completions::Error::Transport(crate::transport::Error::Http { + status: 429, + .. + }) ), "expected a 429, got {err:?}" ); @@ -793,7 +807,10 @@ mod round_trip { .await .expect_err("nothing is listening"); assert!( - matches!(err, Error::Transport(crate::transport::Error::Connect(_))), + matches!( + err, + crate::chat_completions::Error::Transport(crate::transport::Error::Connect(_)) + ), "expected a pre-send connect failure, got {err:?}" ); } @@ -803,26 +820,31 @@ mod round_trip { use crate::chat_completions::handler::as_response_error; for original in [ - Error::MissingField("usage"), - Error::Unsupported("non-text response content block"), - Error::InvalidRequest("whatever".to_string()), - Error::Auth(litellm_auth::Error::ProviderAuthentication( + crate::chat_completions::Error::MissingField("usage"), + crate::chat_completions::Error::Unsupported("non-text response content block"), + crate::chat_completions::Error::InvalidRequest("whatever".to_string()), + crate::chat_completions::Error::Auth(litellm_auth::Error::ProviderAuthentication( "whatever".to_string(), )), ] { let label = format!("{original:?}"); assert!( - matches!(as_response_error(original.clone()), Error::ResponseTransform(source) if *source == original), + matches!(as_response_error(original.clone()), crate::chat_completions::Error::ResponseTransform(source) if *source == original), "{label} must not stay retryable once the provider has answered" ); } // An upstream status is already unambiguous, so it survives intact. assert!(matches!( - as_response_error(Error::Transport(crate::transport::Error::Http { + as_response_error(crate::chat_completions::Error::Transport( + crate::transport::Error::Http { + status: 500, + body: "boom".to_string() + } + )), + crate::chat_completions::Error::Transport(crate::transport::Error::Http { status: 500, - body: "boom".to_string() - })), - Error::Transport(crate::transport::Error::Http { status: 500, .. }) + .. + }) )); } } diff --git a/litellm-rust/crates/core/src/chat_completions/transformation.rs b/litellm-rust/crates/core/src/chat_completions/transformation.rs index 99b4496a5b0..3d81096fee7 100644 --- a/litellm-rust/crates/core/src/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/chat_completions/transformation.rs @@ -1,4 +1,3 @@ -use crate::chat_completions::Error; use serde_json::Value; use super::types::{ @@ -34,7 +33,7 @@ pub trait ChatCompletionsProviderConfig: Sync { model: &str, optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; fn auth( &self, @@ -42,7 +41,7 @@ pub trait ChatCompletionsProviderConfig: Sync { model: &str, optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; fn default_headers(&self) -> &'static [(&'static str, &'static str)] { &[("content-type", "application/json")] @@ -80,13 +79,13 @@ pub trait ChatCompletionsProviderConfig: Sync { model: &str, messages: Vec, optional_params: OpaqueParams, - ) -> Result; + ) -> Result; fn transform_response( &self, model: &str, response: ProviderChatResponseData, - ) -> Result; + ) -> Result; } pub fn unsupported_param(optional_params: &OpaqueParams) -> Option { diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs index 2ac5f2f68f9..869be589ab1 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs @@ -1,7 +1,6 @@ use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; use crate::llms::cohere::ocr::transformation::CohereParseConfig; use crate::llms::cohere::ocr::{CohereParams, CohereResponse, validate_document}; -use crate::ocr::Error; use crate::ocr::OcrClient; use crate::ocr::document::{inline_remote_document, validate_inline_document}; use crate::ocr::prepare::{credential_env, transform_request_body}; @@ -25,7 +24,7 @@ impl BaseOcrConfig for AzureAICohereParseConfig { &self, request: &LiteLLMOcrRequest, client: &OcrClient, - ) -> Result { + ) -> Result { let params = crate::ocr::wire::decode_request_value::( serde_json::Value::Object(request.optional_params.clone().into()), "optional_params", @@ -36,7 +35,7 @@ impl BaseOcrConfig for AzureAICohereParseConfig { &request.optional_params, &request.input_sources, ) - .map_err(Error::from)? + .map_err(crate::ocr::Error::from)? }; let base = request .connection @@ -45,7 +44,7 @@ impl BaseOcrConfig for AzureAICohereParseConfig { .or_else(|| credential_env(AZURE_AI_API_BASE_ENV)) .filter(|base| !base.trim().is_empty()) .ok_or_else(|| { - Error::Auth(litellm_auth::Error::ProviderAuthentication( + crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication( "Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(), )) })?; @@ -84,12 +83,12 @@ impl BaseOcrConfig for AzureAICohereParseConfig { &self, request: &LiteLLMOcrRequest, response: CohereResponse, - ) -> Result { + ) -> Result { CohereParseConfig.transform_ocr_response(request, response) } } -fn complete_url(base: &str) -> Result { +fn complete_url(base: &str) -> Result { let mut url = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?; if !matches!(url.scheme(), "http" | "https") { return Err(invalid_api_base()); @@ -106,8 +105,8 @@ fn complete_url(base: &str) -> Result { .map_err(|_| invalid_api_base()) } -fn invalid_api_base() -> Error { - Error::RequestField { +fn invalid_api_base() -> crate::ocr::Error { + crate::ocr::Error::RequestField { path: "api_base".into(), } } diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/common_utils.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/common_utils.rs index 968a29a466b..c381e39eaae 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/common_utils.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/common_utils.rs @@ -1,6 +1,5 @@ use std::sync::OnceLock; -use crate::ocr::Error; use crate::ocr::types::OcrConnection; use litellm_auth::{InputSource, Sourced}; use litellm_auth_azure::{AzureAuthInputs, AzureAuthService}; @@ -8,7 +7,7 @@ use litellm_auth_azure::{AzureAuthInputs, AzureAuthService}; pub(super) async fn resolve_entra( config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), -) -> Result>, Error> { +) -> Result>, crate::ocr::Error> { static SERVICE: OnceLock = OnceLock::new(); SERVICE .get_or_init(AzureAuthService::default) @@ -25,13 +24,13 @@ pub(super) async fn resolve_entra( Sourced::new(value, source) }) }) - .map_err(Error::from) + .map_err(crate::ocr::Error::from) } pub(super) fn validate_destination( connection: &OcrConnection, credential_source: InputSource, -) -> Result<(), Error> { +) -> Result<(), crate::ocr::Error> { if connection.api_base.is_some() && connection.api_base_source == InputSource::Request && credential_source != InputSource::Request diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs index 01434dfcdb9..65775943c79 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs @@ -16,7 +16,6 @@ use crate::constants::{ AZURE_DI_SUBSCRIPTION_HEADER, OCR_POLL_RETRY_SECS, }; use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; -use crate::ocr::Error; use crate::ocr::OcrClient; use crate::ocr::client::read_json_response; use crate::ocr::document::InlineDocument; @@ -172,19 +171,21 @@ fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, prefix: &str, -) -> Result, Error> { +) -> Result, crate::ocr::Error> { if let Some(Value::Array(pages)) = params.get("pages") { if pages.iter().any(Value::is_boolean) { - return Err(Error::Pages("boolean page index".into())); + return Err(crate::ocr::Error::Pages("boolean page index".into())); } if pages .iter() .any(|page| page.is_number() && page.as_i64().is_none()) { - return Err(Error::Pages("page index is out of range".into())); + return Err(crate::ocr::Error::Pages( + "page index is out of range".into(), + )); } if !pages.iter().all(Value::is_i64) && !pages.iter().all(Value::is_string) { - return Err(Error::Pages("mixed page element types".into())); + return Err(crate::ocr::Error::Pages("mixed page element types".into())); } } crate::ocr::wire::decode_request_value(Value::Object(params), prefix) @@ -192,7 +193,7 @@ fn decode_input_params( fn normalize_ocr_params( params: DocumentIntelligenceInputParams, -) -> Result { +) -> Result { Ok(DocumentIntelligenceParams { pages: params.pages.map(normalize_pages).transpose()?.flatten(), features: params @@ -203,7 +204,7 @@ fn normalize_ocr_params( }) } -fn normalize_pages(pages: PagesInput) -> Result, Error> { +fn normalize_pages(pages: PagesInput) -> Result, crate::ocr::Error> { let normalized = match pages { PagesInput::ZeroBasedIndices(indices) => { if indices.is_empty() { @@ -213,10 +214,11 @@ fn normalize_pages(pages: PagesInput) -> Result, Error> { .into_iter() .map(|page| { if page < 0 { - return Err(Error::Pages("negative page index".into())); + return Err(crate::ocr::Error::Pages("negative page index".into())); } - page.checked_add(1) - .ok_or_else(|| Error::Pages("page index is out of range".into())) + page.checked_add(1).ok_or_else(|| { + crate::ocr::Error::Pages("page index is out of range".into()) + }) }) .collect::, _>>()? .into_iter() @@ -241,7 +243,7 @@ fn normalize_pages(pages: PagesInput) -> Result, Error> { .join(","), }; if !normalized.split(',').all(valid_page_token) { - return Err(Error::Pages("invalid native page range".into())); + return Err(crate::ocr::Error::Pages("invalid native page range".into())); } Ok(Some(normalized)) } @@ -262,7 +264,7 @@ fn valid_page_token(token: &str) -> bool { } } -fn normalize_features(features: FeaturesInput) -> Result, Error> { +fn normalize_features(features: FeaturesInput) -> Result, crate::ocr::Error> { let tokens = match features { FeaturesInput::Names(names) => names, FeaturesInput::CommaSeparated(names) => names.split(',').map(str::to_string).collect(), @@ -277,16 +279,18 @@ fn normalize_features(features: FeaturesInput) -> Result, Error> }; first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric) }) { - return Err(Error::Features); + return Err(crate::ocr::Error::Features); } Ok(Some(normalized.join(","))) } #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -fn transform_ocr_request(document: OcrDocument) -> Result { +fn transform_ocr_request( + document: OcrDocument, +) -> Result { let source = document.source(); if source.is_empty() { - return Err(Error::MissingDocumentUrl); + return Err(crate::ocr::Error::MissingDocumentUrl); } Ok(if let Some(document) = InlineDocument::parse(source)? { DocumentIntelligenceRequest::Base64Source { @@ -303,9 +307,9 @@ fn transform_ocr_request(document: OcrDocument) -> Result Result { +) -> Result { if response.status != Some(OperationStatus::Succeeded) { - return Err(Error::OperationStatus( + return Err(crate::ocr::Error::OperationStatus( response .status .map(|status| status.to_string()) @@ -334,12 +338,12 @@ fn transform_ocr_response( }) } -fn normalize_page(page: AzureDocumentIntelligencePage) -> Result { +fn normalize_page(page: AzureDocumentIntelligencePage) -> Result { let index = page .page_number .unwrap_or(1) .checked_sub(1) - .ok_or(Error::NumericRange("page.pageNumber"))?; + .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 { @@ -369,10 +373,10 @@ fn normalize_page(page: AzureDocumentIntelligencePage) -> Result { })) } -fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result { +fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result { let value = value * scale; if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 { - return Err(Error::NumericRange(field)); + return Err(crate::ocr::Error::NumericRange(field)); } Ok(value.trunc() as i64) } @@ -390,7 +394,7 @@ async fn read_operation_response( connection: &OcrConnection, native: bool, hooks: &Arc, -) -> Result, Error> { +) -> Result, crate::ocr::Error> { if response.status() != reqwest::StatusCode::ACCEPTED { let bytes = crate::ocr::client::read_response_bytes(response, connection.max_response_bytes) @@ -402,15 +406,15 @@ async fn read_operation_response( .headers() .get("operation-location") .and_then(|value| value.to_str().ok()) - .ok_or(Error::PollLocation)? + .ok_or(crate::ocr::Error::PollLocation)? .to_string(); - let original = Url::parse(original_url).map_err(|_| Error::PollOrigin)?; - let operation = Url::parse(&location).map_err(|_| Error::PollOrigin)?; + let original = Url::parse(original_url).map_err(|_| crate::ocr::Error::PollOrigin)?; + let operation = Url::parse(&location).map_err(|_| crate::ocr::Error::PollOrigin)?; if original.origin() != operation.origin() || !operation.username().is_empty() || operation.password().is_some() { - return Err(Error::PollOrigin); + return Err(crate::ocr::Error::PollOrigin); } let bytes = crate::ocr::client::read_response_bytes(response, connection.max_response_bytes).await?; @@ -424,16 +428,16 @@ async fn poll_operation( headers: &[(String, String)], connection: &OcrConnection, native: bool, -) -> Result, Error> { +) -> Result, crate::ocr::Error> { let deadline = Instant::now() .checked_add(connection.poll_timeout) - .ok_or(Error::PollTimeout)?; + .ok_or(crate::ocr::Error::PollTimeout)?; loop { let remaining = deadline .checked_duration_since(Instant::now()) .filter(|remaining| !remaining.is_zero()) - .ok_or(Error::PollTimeout)?; + .ok_or(crate::ocr::Error::PollTimeout)?; let builder = http_client .get(url.clone()) .timeout(remaining.min(connection.timeout)); @@ -444,7 +448,7 @@ async fn poll_operation( ); let response = tokio::time::timeout_at(deadline, crate::http_utils::http_request(builder)) .await - .map_err(|_| Error::PollTimeout)? + .map_err(|_| crate::ocr::Error::PollTimeout)? .map_err(crate::transport::Error::from)?; let retry = response .headers() @@ -462,16 +466,16 @@ async fn poll_operation( ), ) .await - .map_err(|_| Error::PollTimeout)??; + .map_err(|_| crate::ocr::Error::PollTimeout)??; match &decoded.data.status { Some(OperationStatus::Succeeded) => return Ok(decoded), Some(OperationStatus::Running | OperationStatus::NotStarted) => { tokio::time::timeout_at(deadline, tokio::time::sleep(Duration::from_secs(retry))) .await - .map_err(|_| Error::PollTimeout)?; + .map_err(|_| crate::ocr::Error::PollTimeout)?; } status => { - return Err(Error::OperationStatus( + return Err(crate::ocr::Error::OperationStatus( status .as_ref() .map(ToString::to_string) @@ -499,7 +503,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { &self, request: &LiteLLMOcrRequest, client: &OcrClient, - ) -> Result { + ) -> Result { let params = map_ocr_params(request)?; let config = AzureAuthInputs { azure_ad_token_provider: request.azure_ad_token_provider.clone(), @@ -507,12 +511,12 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { &request.optional_params, &request.input_sources, ) - .map_err(Error::from)? + .map_err(crate::ocr::Error::from)? }; let headers = validate_environment(&request.connection, &config, &credential_env).await?; let endpoint = nonblank(request.connection.api_base.clone()) .or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV))) - .ok_or_else(|| Error::Auth(litellm_auth::Error::ProviderAuthentication("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into())))?; + .ok_or_else(|| crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into())))?; let url = get_complete_url(&endpoint, &request.model, ¶ms)?; let body = transform_ocr_request(request.document.clone())?; transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await @@ -522,7 +526,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { &self, request: &LiteLLMOcrRequest, response: AzureDocumentIntelligenceOperation, - ) -> Result { + ) -> Result { transform_ocr_response(&request.model, response) } @@ -533,8 +537,10 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { url: &str, headers: &[(String, String)], request: &LiteLLMOcrRequest, - ) -> Result, Error> - { + ) -> Result< + crate::ocr::wire::DecodedOcrResponse, + crate::ocr::Error, + > { read_operation_response( client.polling_http(), response, @@ -549,7 +555,9 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { } #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -fn map_ocr_params(request: &LiteLLMOcrRequest) -> Result { +fn map_ocr_params( + request: &LiteLLMOcrRequest, +) -> Result { let params = decode_input_params(request.optional_params.clone().into(), "optional_params")?; let crate::ocr::prepare::ParsedProviderParams { known: params, @@ -562,7 +570,7 @@ fn get_complete_url( endpoint: &str, model: &str, params: &DocumentIntelligenceParams, -) -> Result { +) -> Result { let model = format!("{}:analyze", model_id(model)?); ApiUrl::parse(endpoint) .and_then(|url| url.complete_path(&["documentintelligence", "documentModels", &model])) @@ -580,7 +588,7 @@ fn get_complete_url( ) .into_string() }) - .map_err(|_| Error::RequestField { + .map_err(|_| crate::ocr::Error::RequestField { path: "api_base".into(), }) } @@ -589,7 +597,7 @@ async fn validate_environment( connection: &OcrConnection, config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), -) -> Result, Error> { +) -> Result, crate::ocr::Error> { if crate::http_utils::has_header(&connection.extra_headers, "authorization") || crate::http_utils::has_header(&connection.extra_headers, AZURE_DI_SUBSCRIPTION_HEADER) { @@ -615,7 +623,7 @@ async fn validate_environment( } let token = super::super::common_utils::resolve_entra(config, env_lookup) .await? - .ok_or(Error::MissingAzureDocumentIntelligenceCredentials)?; + .ok_or(crate::ocr::Error::MissingAzureDocumentIntelligenceCredentials)?; super::super::common_utils::validate_destination(connection, token.source())?; Ok( std::iter::once(("Authorization".into(), format!("Bearer {}", token.value()))) @@ -624,10 +632,10 @@ async fn validate_environment( ) } -fn model_id(model: &str) -> Result<&str, Error> { +fn model_id(model: &str) -> Result<&str, crate::ocr::Error> { let model = model.rsplit('/').next().unwrap_or(model); if matches!(model, "." | "..") { - return Err(Error::DotModel); + return Err(crate::ocr::Error::DotModel); } Ok(model) } @@ -644,7 +652,7 @@ mod tests { use rstest::rstest; use serde_json::{Value, json}; - fn map(value: Value) -> Result { + fn map(value: Value) -> Result { let fields = value.as_object().unwrap().clone(); normalize_ocr_params(decode_input_params(fields, "optional_params")?.known) } diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs index 283dd2bfcb1..9b960170a51 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs @@ -2,7 +2,6 @@ use crate::constants::AZURE_AI_OCR_PATH; use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; use crate::llms::mistral::ocr::MistralOcrResponse; use crate::llms::mistral::ocr::transformation::MistralOCRConfig; -use crate::ocr::Error; use crate::ocr::OcrClient; use crate::ocr::document::{inline_remote_document, validate_inline_document}; use crate::ocr::prepare::{credential_env, transform_request_body}; @@ -28,7 +27,7 @@ impl BaseOcrConfig for AzureAIOCRConfig { &self, request: &LiteLLMOcrRequest, client: &OcrClient, - ) -> Result { + ) -> Result { let params = self.map_ocr_params(&request.model, &request.optional_params); let config = AzureAuthInputs { azure_ad_token_provider: request.azure_ad_token_provider.clone(), @@ -36,7 +35,7 @@ impl BaseOcrConfig for AzureAIOCRConfig { &request.optional_params, &request.input_sources, ) - .map_err(Error::from)? + .map_err(crate::ocr::Error::from)? }; let url = get_complete_url(request.connection.api_base.as_deref(), &credential_env)?; let headers = validate_environment(&request.connection, &config, &credential_env).await?; @@ -65,7 +64,7 @@ impl BaseOcrConfig for AzureAIOCRConfig { &self, request: &LiteLLMOcrRequest, response: MistralOcrResponse, - ) -> Result { + ) -> Result { MistralOCRConfig.transform_ocr_response(request, response) } } @@ -73,15 +72,15 @@ impl BaseOcrConfig for AzureAIOCRConfig { fn get_complete_url( api_base: Option<&str>, env_lookup: &dyn Fn(&str) -> Option, -) -> Result { +) -> Result { let base = nonblank(api_base.map(str::to_string)) .or_else(|| nonblank(env_lookup(AZURE_AI_API_BASE_ENV))) - .ok_or_else(|| Error::Auth(litellm_auth::Error::ProviderAuthentication("Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter".into())))?; + .ok_or_else(|| crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication("Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter".into())))?; let path: Vec<&str> = AZURE_AI_OCR_PATH.trim_matches('/').split('/').collect(); ApiUrl::parse(&base) .and_then(|url| url.complete_path(&path)) .map(|url| url.into_string()) - .map_err(|_| Error::RequestField { + .map_err(|_| crate::ocr::Error::RequestField { path: "api_base".into(), }) } @@ -90,7 +89,7 @@ pub(super) async fn validate_environment( connection: &OcrConnection, config: &AzureAuthInputs, env_lookup: &(dyn Fn(&str) -> Option + Sync), -) -> Result, Error> { +) -> Result, crate::ocr::Error> { if crate::http_utils::has_header(&connection.extra_headers, "authorization") { if config.azure_ad_token_provider.is_some() { super::common_utils::resolve_entra(config, env_lookup).await?; @@ -110,7 +109,7 @@ pub(super) async fn validate_environment( } let key = super::common_utils::resolve_entra(config, env_lookup) .await? - .ok_or(Error::MissingAzureAiCredentials)?; + .ok_or(crate::ocr::Error::MissingAzureAiCredentials)?; super::common_utils::validate_destination(connection, key.source())?; Ok(bearer_headers(connection, key.value())) } diff --git a/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs index 983bbe48637..53a52e7ec42 100644 --- a/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs @@ -2,7 +2,6 @@ use std::future::Future; use serde::de::DeserializeOwned; -use crate::ocr::Error; use crate::ocr::OcrClient; use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrResponseFormat}; use crate::ocr::wire::DecodedOcrResponse; @@ -23,13 +22,13 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { &self, request: &LiteLLMOcrRequest, client: &OcrClient, - ) -> impl Future> + Send; + ) -> impl Future> + Send; fn transform_ocr_response( &self, request: &LiteLLMOcrRequest, response: Self::ProviderResponse, - ) -> Result; + ) -> Result; fn read_response( &self, @@ -38,7 +37,7 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { _url: &str, _headers: &[(String, String)], request: &LiteLLMOcrRequest, - ) -> impl Future, Error>> + Send + ) -> impl Future, crate::ocr::Error>> + Send { async move { let bytes = crate::ocr::client::read_response_bytes( diff --git a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs index 11cae2a9685..ad82d02f43b 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -3,7 +3,6 @@ use serde_json::{Map, Value, json}; use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE}; use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; -use crate::ocr::Error; use crate::ocr::OcrClient; use crate::ocr::document::InlineDocument; use crate::ocr::prepare::{credential_env, transform_request_body}; @@ -31,16 +30,16 @@ pub(crate) struct CohereRequest { pub output_format: OutputFormat, } -pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), Error> { +pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), crate::ocr::Error> { let OcrDocument::ImageUrl { image_url, .. } = document else { - return Err(Error::CohereImageOnly); + return Err(crate::ocr::Error::CohereImageOnly); }; if image_url.is_empty() { - return Err(Error::CohereImageOnly); + return Err(crate::ocr::Error::CohereImageOnly); } if let Some(inline) = InlineDocument::parse(image_url)? { if !inline.mime_type().type_.eq_ignore_ascii_case("image") { - return Err(Error::CohereImageOnly); + return Err(crate::ocr::Error::CohereImageOnly); } inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?; } @@ -81,14 +80,15 @@ struct CohereBilledUnits { pub(crate) fn transform_response( model: &str, response: CohereResponse, -) -> Result { +) -> Result { let pages_processed = response .meta .and_then(|meta| meta.billed_units) .and_then(|units| units.pages) .map(Ok) .unwrap_or_else(|| { - i64::try_from(response.pages.len()).map_err(|_| Error::NumericRange("pages")) + i64::try_from(response.pages.len()) + .map_err(|_| crate::ocr::Error::NumericRange("pages")) })?; let pages = response .pages @@ -96,7 +96,7 @@ pub(crate) fn transform_response( .enumerate() .map(|(position, page)| { let index = page.index.map(Ok).unwrap_or_else(|| { - i64::try_from(position).map_err(|_| Error::NumericRange("page index")) + i64::try_from(position).map_err(|_| crate::ocr::Error::NumericRange("page index")) })?; let (content, images) = page .markdown @@ -127,7 +127,7 @@ pub(crate) fn transform_response( } Ok(normalized) }) - .collect::, Error>>()?; + .collect::, crate::ocr::Error>>()?; Ok(LiteLLMOcrResponse { pages, model: model.into(), @@ -148,7 +148,7 @@ impl CohereParseConfig { model: &str, document: OcrDocument, params: CohereParams, - ) -> Result { + ) -> Result { validate_document(&document)?; Ok(CohereRequest { model: model.into(), @@ -169,7 +169,7 @@ impl BaseOcrConfig for CohereParseConfig { &self, request: &LiteLLMOcrRequest, client: &OcrClient, - ) -> Result { + ) -> Result { let params = crate::ocr::wire::decode_request_value::( serde_json::Value::Object(request.optional_params.clone().into()), "optional_params", @@ -193,12 +193,12 @@ impl BaseOcrConfig for CohereParseConfig { &self, request: &LiteLLMOcrRequest, response: CohereResponse, - ) -> Result { + ) -> Result { transform_response(&request.model, response) } } -fn complete_url(base: &str) -> Result { +fn complete_url(base: &str) -> Result { let parsed = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?; if !matches!(parsed.scheme(), "http" | "https") { return Err(invalid_api_base()); @@ -209,8 +209,8 @@ fn complete_url(base: &str) -> Result { .map_err(|_| invalid_api_base()) } -fn invalid_api_base() -> Error { - Error::RequestField { +fn invalid_api_base() -> crate::ocr::Error { + crate::ocr::Error::RequestField { path: "api_base".into(), } } @@ -218,7 +218,7 @@ fn invalid_api_base() -> Error { fn validate_environment( connection: &OcrConnection, env_lookup: &(dyn Fn(&str) -> Option + Sync), -) -> Result, Error> { +) -> Result, crate::ocr::Error> { if crate::http_utils::has_header(&connection.extra_headers, "authorization") { return Ok(connection.extra_headers.clone()); } @@ -230,7 +230,7 @@ fn validate_environment( .map(str::to_string) .or_else(|| env_lookup(COHERE_API_KEY_ENV).filter(|key| !key.trim().is_empty())) .ok_or_else(|| { - Error::Auth(litellm_auth::Error::ProviderAuthentication( + crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication( "Missing COHERE_API_KEY - set it in the environment or pass api_key".into(), )) })?; @@ -321,7 +321,7 @@ mod tests { ] { assert_eq!( validate_document(&serde_json::from_value(value).unwrap()), - Err(Error::CohereImageOnly) + Err(crate::ocr::Error::CohereImageOnly) ); } assert!(serde_json::from_value::(json!({"output_format":"html"})).is_err()); @@ -369,7 +369,7 @@ mod tests { }, &|_| None, ), - Err(Error::Auth(_)) + Err(crate::ocr::Error::Auth(_)) )); } } diff --git a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs index 59d7b095a52..1570c61717a 100644 --- a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs @@ -3,7 +3,6 @@ use serde_json::{Map, Value}; use crate::constants::MISTRAL_OCR_API_BASE; use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; -use crate::ocr::Error; use crate::ocr::OcrClient; use crate::ocr::prepare::{credential_env, transform_request_body}; use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument}; @@ -33,7 +32,7 @@ pub(crate) struct MistralOcrResponse { pub(crate) fn transform_ocr_response( model: &str, response: MistralOcrResponse, -) -> Result { +) -> Result { Ok(LiteLLMOcrResponse { pages: response.pages, model: response.model.unwrap_or_else(|| model.to_string()), @@ -55,7 +54,7 @@ impl MistralOCRConfig { model: &str, document: OcrDocument, params: &OpaqueParams, - ) -> Result { + ) -> Result { Ok(MistralOcrRequest { model: model.to_string(), document, @@ -89,7 +88,7 @@ impl BaseOcrConfig for MistralOCRConfig { &self, request: &LiteLLMOcrRequest, client: &OcrClient, - ) -> Result { + ) -> Result { let params = self.map_ocr_params(&request.model, &request.optional_params); let headers = validate_environment(&request.connection, &credential_env)?; let url = get_complete_url(request.connection.api_base.as_deref())?; @@ -101,12 +100,12 @@ impl BaseOcrConfig for MistralOCRConfig { &self, request: &LiteLLMOcrRequest, response: MistralOcrResponse, - ) -> Result { + ) -> Result { transform_ocr_response(&request.model, response) } } -pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result { +pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result { let base = api_base .map(str::trim) .filter(|base| !base.is_empty()) @@ -114,7 +113,7 @@ pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result ApiUrl::parse(base) .and_then(|url| url.complete_path(&["v1", "ocr"])) .map(|url| url.into_string()) - .map_err(|_| Error::RequestField { + .map_err(|_| crate::ocr::Error::RequestField { path: "api_base".into(), }) } @@ -122,7 +121,7 @@ pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result fn validate_environment( connection: &OcrConnection, env_lookup: &(dyn Fn(&str) -> Option + Sync), -) -> Result, Error> { +) -> Result, crate::ocr::Error> { if crate::http_utils::has_header(&connection.extra_headers, "authorization") { return Ok(connection.extra_headers.clone()); } @@ -430,10 +429,12 @@ mod tests { fn environment_rejects_missing_key() { assert!(matches!( validate_environment(&OcrConnection::default(), &|_| None), - Err(Error::Auth(litellm_auth::Error::MissingApiKey { - provider: "Mistral", - environment_variable: MISTRAL_API_KEY_ENV, - })) + Err(crate::ocr::Error::Auth( + litellm_auth::Error::MissingApiKey { + provider: "Mistral", + environment_variable: MISTRAL_API_KEY_ENV, + } + )) )); } } diff --git a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs index da0efbc2c9f..d4d66335dd8 100644 --- a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -5,7 +5,6 @@ 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; -use crate::ocr::Error; use crate::ocr::OcrClient; use crate::ocr::document::InlineDocument; use crate::ocr::prepare::{ @@ -15,6 +14,120 @@ use crate::ocr::prepare::{ use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument}; use crate::url_utils::ApiUrl; +#[derive(Clone, Debug)] +pub(crate) struct ReductoParseV3Config; + +impl ReductoParseV3Config { + #[tracing::instrument( + name = "transform_ocr_request", + target = "litellm::function_trace", + level = "trace", + skip_all + )] + fn transform_ocr_request( + &self, + _model: &str, + document: OcrDocument, + params: &ReductoV3Params, + ) -> Result { + Ok(ReductoV3Request { + input: document.source().to_string(), + params: params.clone(), + }) + } +} + +impl BaseOcrConfig for ReductoParseV3Config { + type ProviderResponse = ReductoResponse; + + fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] { + &["formatting", "retrieval", "settings"] + } + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let ParsedProviderParams { + known: params, + extra_params, + } = _prepare_ocr_request::(request)?; + let headers = validate_environment(&request.connection, &credential_env)?; + let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?; + let (document, headers) = guardrail_document(request, &url, &headers).await?; + let document = prepare_document(client, document, &request.connection, &headers).await?; + let body = self.transform_ocr_request(&request.model, document, ¶ms)?; + let body = merge_extra_params(&body, extra_params)?; + build_http_request(client, request, &url, &headers, &body) + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: ReductoResponse, + ) -> Result { + transform_ocr_response(&request.model, response) + } +} + +#[derive(Clone, Debug)] +pub(crate) struct ReductoParseLegacyConfig; + +impl ReductoParseLegacyConfig { + #[tracing::instrument( + name = "transform_ocr_request", + target = "litellm::function_trace", + level = "trace", + skip_all + )] + fn transform_ocr_request( + &self, + _model: &str, + document: OcrDocument, + params: &ReductoLegacyParams, + ) -> Result { + Ok(ReductoLegacyRequest { + document_url: document.source().to_string(), + options: params.enhance.as_ref().map(|_| params.clone()), + }) + } +} + +impl BaseOcrConfig for ReductoParseLegacyConfig { + type ProviderResponse = ReductoResponse; + + fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] { + &["enhance"] + } + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let ParsedProviderParams { + known: params, + extra_params, + } = _prepare_ocr_request::(request)?; + let headers = validate_environment(&request.connection, &credential_env)?; + let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?; + let (document, headers) = guardrail_document(request, &url, &headers).await?; + let document = prepare_document(client, document, &request.connection, &headers).await?; + let body = self.transform_ocr_request(&request.model, document, ¶ms)?; + let body = merge_extra_params(&body, extra_params)?; + build_http_request(client, request, &url, &headers, &body) + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: ReductoResponse, + ) -> Result { + transform_ocr_response(&request.model, response) + } +} + #[derive(Clone, Debug, Default, Serialize, Deserialize)] struct ReductoV3Params { #[serde(skip_serializing_if = "Option::is_none")] @@ -141,44 +254,10 @@ fn checked_truncated_i64(value: f64) -> Option { .then(|| value.trunc() as i64) } -#[tracing::instrument( - name = "transform_ocr_request", - target = "litellm::function_trace", - level = "trace", - skip_all -)] -fn transform_v3_ocr_request( - _model: &str, - document: OcrDocument, - params: &ReductoV3Params, -) -> Result { - Ok(ReductoV3Request { - input: document.source().to_string(), - params: params.clone(), - }) -} - -#[tracing::instrument( - name = "transform_ocr_request", - target = "litellm::function_trace", - level = "trace", - skip_all -)] -fn transform_legacy_ocr_request( - _model: &str, - document: OcrDocument, - params: &ReductoLegacyParams, -) -> Result { - Ok(ReductoLegacyRequest { - document_url: document.source().to_string(), - options: params.enhance.as_ref().map(|_| params.clone()), - }) -} - pub(crate) fn transform_ocr_response( model: &str, response: ReductoResponse, -) -> Result { +) -> Result { let result = match response.result { Some(result) => result.unwrap_or_default(), None => ReductoResult { @@ -248,7 +327,7 @@ fn page(index: i64, markdown: String, blocks: Option) -> Value { } result } -fn get_complete_url(api_base: Option<&str>, path: &str) -> Result { +fn get_complete_url(api_base: Option<&str>, path: &str) -> Result { let base = api_base .map(str::trim) .filter(|base| !base.is_empty()) @@ -256,7 +335,7 @@ fn get_complete_url(api_base: Option<&str>, path: &str) -> Result ApiUrl::parse(base) .and_then(|url| url.complete_path(&[path])) .map(|url| url.into_string()) - .map_err(|_| Error::RequestField { + .map_err(|_| crate::ocr::Error::RequestField { path: "api_base".into(), }) } @@ -264,7 +343,7 @@ fn get_complete_url(api_base: Option<&str>, path: &str) -> Result fn validate_environment( connection: &OcrConnection, env_lookup: &(dyn Fn(&str) -> Option + Sync), -) -> Result, Error> { +) -> Result, crate::ocr::Error> { if crate::http_utils::has_header(&connection.extra_headers, "authorization") { return Ok(connection.extra_headers.clone()); } @@ -279,7 +358,7 @@ fn validate_environment( .map(|key| key.trim().to_string()) .filter(|key| !key.is_empty()) }) - .ok_or(Error::MissingReductoApiKey)?; + .ok_or(crate::ocr::Error::MissingReductoApiKey)?; Ok( std::iter::once(("Authorization".into(), format!("Bearer {api_key}"))) .chain(connection.extra_headers.clone()) @@ -292,25 +371,26 @@ async fn prepare_document( document: OcrDocument, connection: &OcrConnection, headers: &[(String, String)], -) -> Result { +) -> Result { if document.source().starts_with(REDUCTO_ID_PREFIX) { if document.source()[REDUCTO_ID_PREFIX.len()..] .trim() .is_empty() { - return Err(Error::RequestField { + return Err(crate::ocr::Error::RequestField { path: "document file id".into(), }); } return Ok(document); } - let inline = InlineDocument::parse(document.source())?.ok_or(Error::ReductoSource)?; + let inline = + InlineDocument::parse(document.source())?.ok_or(crate::ocr::Error::ReductoSource)?; let mime = inline.mime_type().to_string(); let bytes = inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?; let part = reqwest::multipart::Part::bytes(bytes) .file_name("document") .mime_str(&mime) - .map_err(|_| Error::InvalidDataUri)?; + .map_err(|_| crate::ocr::Error::InvalidDataUri)?; let builder = client .provider_http() .post(get_complete_url(connection.api_base.as_deref(), "upload")?) @@ -337,86 +417,13 @@ async fn prepare_document( .map(str::trim) .filter(|id| !id.is_empty()); let Some(file_id) = file_id else { - return Err(Error::ResponseField { + return Err(crate::ocr::Error::ResponseField { path: "file_id".into(), }); }; Ok(document.with_source(file_id.to_string())) } -#[derive(Clone, Debug)] -pub(crate) struct ReductoParseLegacyConfig; - -impl BaseOcrConfig for ReductoParseLegacyConfig { - type ProviderResponse = ReductoResponse; - - fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] { - &["enhance"] - } - - async fn prepare_request( - &self, - request: &LiteLLMOcrRequest, - client: &OcrClient, - ) -> Result { - let ParsedProviderParams { - known: params, - extra_params, - } = _prepare_ocr_request::(request)?; - let headers = validate_environment(&request.connection, &credential_env)?; - let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?; - let (document, headers) = guardrail_document(request, &url, &headers).await?; - let document = prepare_document(client, document, &request.connection, &headers).await?; - let body = transform_legacy_ocr_request(&request.model, document, ¶ms)?; - let body = merge_extra_params(&body, extra_params)?; - build_http_request(client, request, &url, &headers, &body) - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: ReductoResponse, - ) -> Result { - transform_ocr_response(&request.model, response) - } -} -#[derive(Clone, Debug)] -pub(crate) struct ReductoParseV3Config; - -impl BaseOcrConfig for ReductoParseV3Config { - type ProviderResponse = ReductoResponse; - - fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] { - &["formatting", "retrieval", "settings"] - } - - async fn prepare_request( - &self, - request: &LiteLLMOcrRequest, - client: &OcrClient, - ) -> Result { - let ParsedProviderParams { - known: params, - extra_params, - } = _prepare_ocr_request::(request)?; - let headers = validate_environment(&request.connection, &credential_env)?; - let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?; - let (document, headers) = guardrail_document(request, &url, &headers).await?; - let document = prepare_document(client, document, &request.connection, &headers).await?; - let body = transform_v3_ocr_request(&request.model, document, ¶ms)?; - let body = merge_extra_params(&body, extra_params)?; - build_http_request(client, request, &url, &headers, &body) - } - - fn transform_ocr_response( - &self, - request: &LiteLLMOcrRequest, - response: ReductoResponse, - ) -> Result { - transform_ocr_response(&request.model, response) - } -} - #[cfg(test)] mod tests { use super::*; diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/common_utils.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/common_utils.rs index 6d1259a76e8..6340084ad7f 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/common_utils.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/common_utils.rs @@ -1,8 +1,7 @@ -use crate::ocr::Error; use crate::ocr::types::OcrConnection; use litellm_auth::InputSource; -pub(super) fn validate_destination(connection: &OcrConnection) -> Result<(), Error> { +pub(super) fn validate_destination(connection: &OcrConnection) -> Result<(), crate::ocr::Error> { if connection.api_base.is_some() && connection.api_base_source == InputSource::Request { return Err(litellm_auth::Error::RequestVertexCredentialDestination.into()); } diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs index 84c3d509361..a754662df54 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs @@ -110,7 +110,6 @@ mod mapping { use super::DeepSeekAi; use super::types::*; - use crate::ocr::Error; use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; use crate::providers::model::ProviderModel; @@ -119,9 +118,9 @@ mod mapping { provider_model: ProviderModel, document: OcrDocument, params: &DeepSeekOcrParams, - ) -> Result { + ) -> Result { if document.source().is_empty() { - return Err(Error::MissingDocumentUrl); + return Err(crate::ocr::Error::MissingDocumentUrl); } let content = OcrDocument::ImageUrl { image_url: document.source().to_string(), @@ -141,13 +140,13 @@ mod mapping { pub(crate) fn transform_ocr_response( model: &str, response: DeepSeekOcrResponse, - ) -> Result { + ) -> Result { let content = response .choices .into_iter() .next() .and_then(|choice| choice.message.content) - .ok_or(Error::EmptyContent)?; + .ok_or(crate::ocr::Error::EmptyContent)?; let decoded = decode_content(content)?; let pages = match decoded.result.pages { Some(pages) if !pages.is_empty() => pages @@ -176,17 +175,18 @@ mod mapping { fallback_markdown: String, } - fn decode_content(content: DeepSeekContent) -> Result { + fn decode_content(content: DeepSeekContent) -> Result { let (result, fallback_markdown) = match content { DeepSeekContent::Text(text) if text.is_empty() => { - return Err(Error::EmptyContent); + return Err(crate::ocr::Error::EmptyContent); } DeepSeekContent::Text(text) => (decode_json_content(&text)?, text), DeepSeekContent::Object(object) => { - let fallback = - serde_json::to_string(&object).map_err(|_| Error::ResponseField { + let fallback = serde_json::to_string(&object).map_err(|_| { + crate::ocr::Error::ResponseField { path: "choices[0].message.content".into(), - })?; + } + })?; (Some(object), fallback) } }; @@ -196,7 +196,7 @@ mod mapping { }) } - fn decode_json_content(text: &str) -> Result, Error> { + fn decode_json_content(text: &str) -> Result, crate::ocr::Error> { if !text.trim_start().starts_with('{') { return Ok(None); } @@ -206,7 +206,7 @@ mod mapping { }; serde_path_to_error::deserialize(value.into_deserializer()) .map(Some) - .map_err(|error| Error::ResponseField { + .map_err(|error| crate::ocr::Error::ResponseField { path: format!("choices[0].message.content.{}", error.path()), }) } @@ -217,7 +217,6 @@ pub(crate) use mapping::{transform_ocr_request, transform_ocr_response}; use super::common_utils::validate_destination; use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; -use crate::ocr::Error; use crate::ocr::OcrClient; use crate::ocr::prepare::{ _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, @@ -251,7 +250,7 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { &self, request: &LiteLLMOcrRequest, client: &OcrClient, - ) -> Result { + ) -> Result { validate_destination(&request.connection)?; let ParsedProviderParams { known: params, @@ -261,7 +260,7 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { &request.optional_params, &request.input_sources, ) - .map_err(Error::from)?; + .map_err(crate::ocr::Error::from)?; let authentication = client .vertex_auth() .validate_environment( @@ -271,7 +270,7 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { &credential_env, ) .await - .map_err(Error::from)?; + .map_err(crate::ocr::Error::from)?; let location = vertex::get_vertex_ai_location(&config, &credential_env) .unwrap_or_else(|| DEFAULT_LOCATION.to_string()); let url = get_complete_url( @@ -298,15 +297,15 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { &self, request: &LiteLLMOcrRequest, response: DeepSeekOcrResponse, - ) -> Result { + ) -> Result { mapping::transform_ocr_response(&request.model, response) } } -pub(crate) fn provider_model(model: &str) -> Result, Error> { +pub(crate) fn provider_model(model: &str) -> Result, crate::ocr::Error> { RoutedModel::new(model) .and_then(RoutedModel::into_provider::) - .map_err(|_| Error::RequestField { + .map_err(|_| crate::ocr::Error::RequestField { path: "model".into(), }) } @@ -315,7 +314,7 @@ fn get_complete_url( api_base: Option<&str>, project: &str, location: &str, -) -> Result { +) -> Result { let base = api_base .map(str::trim) .filter(|base| !base.is_empty()) @@ -335,7 +334,7 @@ fn get_complete_url( ]) }) .map(|url| url.into_string()) - .map_err(|_| Error::RequestField { + .map_err(|_| crate::ocr::Error::RequestField { path: "api_base".into(), }) } diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs index 259fd14452d..f16843e1ee7 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs @@ -2,7 +2,6 @@ use super::common_utils::validate_destination; use crate::llms::base_llm::ocr::transformation::BaseOcrConfig; use crate::llms::mistral::ocr::MistralOcrResponse; use crate::llms::mistral::ocr::transformation::MistralOCRConfig; -use crate::ocr::Error; use crate::ocr::OcrClient; use crate::ocr::document::{inline_remote_document, validate_inline_document}; use crate::ocr::prepare::{credential_env, transform_request_body}; @@ -25,14 +24,14 @@ impl BaseOcrConfig for VertexAIOCRConfig { &self, request: &LiteLLMOcrRequest, client: &OcrClient, - ) -> Result { + ) -> Result { validate_destination(&request.connection)?; let params = self.map_ocr_params(&request.model, &request.optional_params); let config = VertexConfig::from_sourced_optional_params( &request.optional_params, &request.input_sources, ) - .map_err(Error::from)?; + .map_err(crate::ocr::Error::from)?; let authentication = client .vertex_auth() .validate_environment( @@ -42,7 +41,7 @@ impl BaseOcrConfig for VertexAIOCRConfig { &credential_env, ) .await - .map_err(Error::from)?; + .map_err(crate::ocr::Error::from)?; let location = vertex::get_vertex_ai_location(&config, &credential_env) .unwrap_or_else(|| DEFAULT_LOCATION.to_string()); let url = get_complete_url( @@ -76,7 +75,7 @@ impl BaseOcrConfig for VertexAIOCRConfig { &self, request: &LiteLLMOcrRequest, response: MistralOcrResponse, - ) -> Result { + ) -> Result { MistralOCRConfig.transform_ocr_response(request, response) } } @@ -86,7 +85,7 @@ fn get_complete_url( project: &str, location: &str, model: &str, -) -> Result { +) -> Result { validate_location(location)?; let default_base = format!("https://{location}-aiplatform.googleapis.com"); let base = api_base @@ -109,12 +108,12 @@ fn get_complete_url( ]) }) .map(|url| url.into_string()) - .map_err(|_| Error::RequestField { + .map_err(|_| crate::ocr::Error::RequestField { path: "api_base".into(), }) } -fn validate_location(location: &str) -> Result<(), Error> { +fn validate_location(location: &str) -> Result<(), crate::ocr::Error> { let valid = !location.is_empty() && location .bytes() @@ -130,7 +129,7 @@ fn validate_location(location: &str) -> Result<(), Error> { if valid { return Ok(()); } - Err(Error::RequestField { + Err(crate::ocr::Error::RequestField { path: "vertex_location".into(), }) } diff --git a/litellm-rust/crates/core/src/media.rs b/litellm-rust/crates/core/src/media.rs index 6befd769e8f..ba26f431e57 100644 --- a/litellm-rust/crates/core/src/media.rs +++ b/litellm-rust/crates/core/src/media.rs @@ -9,7 +9,6 @@ use reqwest::Url; use reqwest::dns::{Addrs, Name, Resolve, Resolving}; use crate::constants::MEDIA_CONNECT_TIMEOUT_SECS; -use crate::transport::Error as TransportError; #[derive(Debug, thiserror::Error)] pub(crate) enum Error { @@ -30,7 +29,7 @@ pub(crate) enum Error { #[error("media download timed out")] Timeout, #[error("{0}")] - Transport(#[from] TransportError), + Transport(#[from] crate::transport::Error), } #[derive(Clone)] @@ -119,7 +118,7 @@ impl MediaFetcher { .get(url.clone()) .send() .await - .map_err(TransportError::from)?; + .map_err(crate::transport::Error::from)?; if response.status().is_redirection() { if redirects_followed == policy.max_redirects { return Err(Error::TooManyRedirects); @@ -147,7 +146,11 @@ impl MediaFetcher { .unwrap_or("application/octet-stream") .to_string(); let mut bytes = Vec::new(); - while let Some(chunk) = response.chunk().await.map_err(TransportError::from)? { + while let Some(chunk) = response + .chunk() + .await + .map_err(crate::transport::Error::from)? + { enforce_download_size(bytes.len() as u64 + chunk.len() as u64, policy.max_bytes)?; bytes.extend_from_slice(&chunk); } @@ -177,7 +180,7 @@ impl MediaFetcher { .address_resolver .resolve(host, port) .await - .map_err(|error| TransportError::Network(error.to_string()))?; + .map_err(|error| crate::transport::Error::Network(error.to_string()))?; validate_addresses(&addresses) } } diff --git a/litellm-rust/crates/core/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs index 3a05d0da15c..4d438321e05 100644 --- a/litellm-rust/crates/core/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -1,5 +1,4 @@ use crate::http_utils::string_headers as shared_string_headers; -use crate::messages::Error; use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG; use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG; use serde_json::{Map, Value}; @@ -22,6 +21,6 @@ pub(super) fn messages_provider_config( pub(super) fn string_headers( extra_headers: Option>, -) -> Result, Error> { - shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(Error::from) +) -> Result, super::Error> { + shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(super::Error::from) } diff --git a/litellm-rust/crates/core/src/messages/handler.rs b/litellm-rust/crates/core/src/messages/handler.rs index 8e84adf95ba..cd92db02d4a 100644 --- a/litellm-rust/crates/core/src/messages/handler.rs +++ b/litellm-rust/crates/core/src/messages/handler.rs @@ -1,6 +1,5 @@ use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; use crate::http_utils::http_request; -use crate::messages::Error; use super::client::http_client; use super::common_utils::truncate_error_body; @@ -9,7 +8,7 @@ use super::types::{AnthropicMessagesResponse, MessagesRequest}; pub(super) async fn execute_messages_provider_call( request: MessagesRequest<'_>, -) -> Result { +) -> Result { let request = prepare_provider_request(request)?; let mut request_builder = http_client().post(&request.url).json(&request.body); for (key, value) in &request.upstream_headers { @@ -19,34 +18,34 @@ pub(super) async fn execute_messages_provider_call( request_builder = request_builder.timeout(duration); } - let response = http_request(request_builder) - .await - .map_err(|err| Error::Transport(crate::transport::Error::Network(err.to_string())))?; + let response = http_request(request_builder).await.map_err(|err| { + super::Error::Transport(crate::transport::Error::Network(err.to_string())) + })?; let status = response.status(); - let text = response - .text() - .await - .map_err(|err| Error::Transport(crate::transport::Error::Network(err.to_string())))?; + let text = response.text().await.map_err(|err| { + super::Error::Transport(crate::transport::Error::Network(err.to_string())) + })?; if !status.is_success() { - return Err(Error::Transport(crate::transport::Error::Http { + return Err(super::Error::Transport(crate::transport::Error::Http { status: status.as_u16(), body: truncate_error_body(&text), })); } - let response = serde_json::from_str(&text) - .map_err(|err| Error::InvalidResponse(format!("invalid messages response JSON: {err}")))?; + let response = serde_json::from_str(&text).map_err(|err| { + super::Error::InvalidResponse(format!("invalid messages response JSON: {err}")) + })?; request.config.transform_response(&request.model, response) } pub(super) async fn execute_messages_provider_stream( request: MessagesRequest<'_>, -) -> Result { +) -> Result { let request = prepare_provider_request(request)?; if request.provider != ANTHROPIC_MESSAGES_PROVIDER { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "streaming messages is not supported for this provider".to_string(), )); } @@ -59,16 +58,15 @@ pub(super) async fn execute_messages_provider_stream( request_builder = request_builder.timeout(duration); } - let response = http_request(request_builder) - .await - .map_err(|err| Error::Transport(crate::transport::Error::Network(err.to_string())))?; + let response = http_request(request_builder).await.map_err(|err| { + super::Error::Transport(crate::transport::Error::Network(err.to_string())) + })?; let status = response.status(); if !status.is_success() { - let text = response - .text() - .await - .map_err(|err| Error::Transport(crate::transport::Error::Network(err.to_string())))?; - return Err(Error::Transport(crate::transport::Error::Http { + let text = response.text().await.map_err(|err| { + super::Error::Transport(crate::transport::Error::Network(err.to_string())) + })?; + return Err(super::Error::Transport(crate::transport::Error::Http { status: status.as_u16(), body: truncate_error_body(&text), })); diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index de5085f92ce..fa3b09d819d 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -1,4 +1,3 @@ -use crate::messages::Error; use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers}; @@ -8,7 +7,7 @@ use serde_json::{Map, Value}; pub(super) fn prepare_provider_request( request: MessagesRequest<'_>, -) -> Result { +) -> Result { let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) .or_else(|| { request @@ -19,7 +18,7 @@ pub(super) fn prepare_provider_request( }) }) .ok_or_else(|| { - Error::InvalidProvider( + super::Error::InvalidProvider( "unable to resolve custom_llm_provider for messages request".to_string(), ) })?; @@ -27,22 +26,22 @@ pub(super) fn prepare_provider_request( let provider = provider_info.custom_llm_provider; let config = messages_provider_config(provider) - .ok_or_else(|| Error::InvalidProvider(provider.to_string()))?; + .ok_or_else(|| super::Error::InvalidProvider(provider.to_string()))?; let env_lookup = |key: &str| std::env::var(key).ok(); let headers = validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?; let params: crate::params::OpaqueParams = serde_json::from_value(request.body) - .map_err(|_| Error::InvalidRequest("messages body must be an object".into()))?; + .map_err(|_| super::Error::InvalidRequest("messages body must be an object".into()))?; let mut fields = params.into_provider_body()?; fields.insert("model".into(), Value::String(model.clone())); let typed_request = serde_json::from_value(Value::Object(fields)).map_err(|err| { - Error::InvalidRequest(format!("invalid Anthropic messages request: {err}")) + super::Error::InvalidRequest(format!("invalid Anthropic messages request: {err}")) })?; let transformed = config.transform_request(typed_request)?; let body = serde_json::to_value(transformed).map_err(|err| { - Error::InvalidRequest(format!( + super::Error::InvalidRequest(format!( "failed to serialize Anthropic messages request: {err}" )) })?; @@ -65,7 +64,7 @@ fn validate_environment( extra_headers: Option>, api_key: Option<&str>, env_lookup: &dyn Fn(&str) -> Option, -) -> Result, Error> { +) -> Result, super::Error> { let mut headers = string_headers(extra_headers)?; let auth_strategy = config.auth_strategy(); diff --git a/litellm-rust/crates/core/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs index c79262d37ff..42f8dfd97cc 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -4,8 +4,6 @@ use serde_json::{Map, Value, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; -use crate::messages::Error; - use super::common_utils::{ has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body, }; @@ -79,7 +77,7 @@ fn string_headers_rejects_non_string_values() { let err = string_headers(Some(headers)).expect_err("non-string header rejected"); assert_eq!( err, - Error::Headers(crate::http_utils::HeaderError { + super::Error::Headers(crate::http_utils::HeaderError { context: "messages", name: "x-count".to_string(), actual: "number", @@ -359,7 +357,7 @@ async fn messages_requires_auth_when_no_key_and_no_header() { .await .expect_err("missing auth errors"); - assert!(matches!(err, Error::Auth(_))); + assert!(matches!(err, super::Error::Auth(_))); } #[tokio::test] @@ -440,7 +438,7 @@ async fn messages_maps_provider_error_status_to_http_error() { assert!(matches!( err, - Error::Transport(crate::transport::Error::Http { status: 401, .. }) + super::Error::Transport(crate::transport::Error::Http { status: 401, .. }) )); } @@ -458,5 +456,5 @@ async fn messages_rejects_unsupported_provider() { .await .expect_err("unsupported provider errors"); - assert!(matches!(err, Error::InvalidProvider(provider) if provider == "openai")); + assert!(matches!(err, super::Error::InvalidProvider(provider) if provider == "openai")); } diff --git a/litellm-rust/crates/core/src/messages/transformation.rs b/litellm-rust/crates/core/src/messages/transformation.rs index 609474b5380..8d29862998e 100644 --- a/litellm-rust/crates/core/src/messages/transformation.rs +++ b/litellm-rust/crates/core/src/messages/transformation.rs @@ -1,5 +1,4 @@ use super::types::{AnthropicMessagesRequest, AnthropicMessagesResponse}; -use crate::messages::Error; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum MessagesAuthStrategy { @@ -22,13 +21,13 @@ pub trait AnthropicMessagesProviderConfig: Sync { api_base: Option<&str>, model: &str, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; fn resolve_api_key( &self, api_key: Option<&str>, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; + ) -> Result; fn auth_strategy(&self) -> MessagesAuthStrategy { MessagesAuthStrategy::Header("x-api-key") @@ -48,7 +47,7 @@ pub trait AnthropicMessagesProviderConfig: Sync { fn transform_request( &self, request: AnthropicMessagesRequest, - ) -> Result { + ) -> Result { Ok(request) } @@ -56,7 +55,7 @@ pub trait AnthropicMessagesProviderConfig: Sync { &self, _model: &str, response: AnthropicMessagesResponse, - ) -> Result { + ) -> Result { Ok(response) } } diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index 96be69580d8..c7a29e6e530 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -4,12 +4,10 @@ use std::time::Duration; use bytes::{Bytes, BytesMut}; use serde::de::DeserializeOwned; -use super::Error; use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; use super::wire::{DecodedOcrResponse, decode_response}; use crate::constants::OCR_CONNECT_TIMEOUT_SECS; use crate::media::MediaFetcher; -use crate::transport::Error as TransportError; use litellm_auth_gcp::VertexAuth; #[derive(Clone)] @@ -21,8 +19,8 @@ pub struct OcrClient { } impl OcrClient { - pub fn new(provider_http: reqwest::Client) -> Result { - let document_fetcher = MediaFetcher::new().map_err(TransportError::from)?; + pub fn new(provider_http: reqwest::Client) -> Result { + let document_fetcher = MediaFetcher::new().map_err(crate::transport::Error::from)?; Ok(Self { provider_http, polling_http: no_redirect_http()?, @@ -31,11 +29,20 @@ impl OcrClient { }) } - pub fn shared() -> Result { + pub fn shared() -> Result { shared_client() } - pub async fn perform(&self, request: LiteLLMOcrRequest) -> Result { + #[tracing::instrument( + name = "ocr", + target = "litellm::function_trace", + level = "trace", + skip_all + )] + pub async fn perform( + &self, + request: LiteLLMOcrRequest, + ) -> Result { use super::{ NativeOutcome, OcrAdmission, OcrCall, OcrCallStep, OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult, @@ -45,7 +52,7 @@ impl OcrClient { let mut request = Some(request); let NativeOutcome::Completed(mut call) = OcrCall::admit(self.clone(), OcrAdmission::all()) else { - return Err(Error::InvalidRequest( + return Err(crate::ocr::Error::InvalidRequest( "native OCR host admission declined".into(), )); }; @@ -55,7 +62,9 @@ impl OcrClient { OcrCallStep::Host(OcrHostOperation::ProjectRequest) => { result = Some(OcrHostResult::Request(Ok(( Box::new(request.take().ok_or_else(|| { - Error::InvalidRequest("OCR request was already projected".into()) + crate::ocr::Error::InvalidRequest( + "OCR request was already projected".into(), + ) })?), false, )))) @@ -93,29 +102,29 @@ impl OcrClient { } } -fn no_redirect_http() -> Result { +fn no_redirect_http() -> Result { reqwest::Client::builder() .connect_timeout(Duration::from_secs(OCR_CONNECT_TIMEOUT_SECS)) .redirect(reqwest::redirect::Policy::none()) .build() - .map_err(TransportError::from) + .map_err(crate::transport::Error::from) } -pub(crate) fn shared_client() -> Result { - static CLIENT: OnceLock> = OnceLock::new(); +pub(crate) fn shared_client() -> Result { + static CLIENT: OnceLock> = OnceLock::new(); let client = CLIENT .get_or_init(|| { reqwest::Client::builder() .connect_timeout(Duration::from_secs(OCR_CONNECT_TIMEOUT_SECS)) .build() - .map_err(TransportError::from) + .map_err(crate::transport::Error::from) .and_then(OcrClient::new) }) .clone()?; Ok(client) } -pub async fn ocr(request: LiteLLMOcrRequest) -> Result { +pub async fn ocr(request: LiteLLMOcrRequest) -> Result { shared_client()?.perform(request).await } @@ -123,7 +132,7 @@ pub async fn read_json_response( response: reqwest::Response, native: bool, max_response_bytes: usize, -) -> Result, Error> { +) -> Result, crate::ocr::Error> { let bytes = read_response_bytes(response, max_response_bytes).await?; decode_response(&bytes, native) } @@ -131,7 +140,7 @@ pub async fn read_json_response( pub(crate) async fn read_response_bytes( mut response: reqwest::Response, max_response_bytes: usize, -) -> Result { +) -> Result { let status = response.status(); let limit = if status.is_success() { max_response_bytes @@ -143,13 +152,13 @@ pub(crate) async fn read_response_bytes( .content_length() .is_some_and(|length| length > limit as u64) { - return Err(Error::TooLarge { limit }); + return Err(crate::ocr::Error::TooLarge { limit }); } let mut bytes = BytesMut::new(); while let Some(chunk) = response.chunk().await.map_err(transport_error)? { let remaining = limit.saturating_sub(bytes.len()); if status.is_success() && chunk.len() > remaining { - return Err(Error::TooLarge { limit }); + return Err(crate::ocr::Error::TooLarge { limit }); } bytes.extend_from_slice(&chunk[..chunk.len().min(remaining)]); if !status.is_success() && bytes.len() == limit { @@ -166,9 +175,9 @@ pub(crate) async fn read_response_bytes( Ok(bytes.freeze()) } -pub(crate) fn transport_error(error: reqwest::Error) -> Error { +pub(crate) fn transport_error(error: reqwest::Error) -> crate::ocr::Error { if error.is_timeout() { - return Error::Transport(crate::transport::Error::Http { + return crate::ocr::Error::Transport(crate::transport::Error::Http { status: 408, body: "OCR request timed out".into(), }); @@ -196,7 +205,7 @@ mod tests { .unwrap_err(); assert!(matches!( transport_error(error), - Error::Transport(crate::transport::Error::Http { status: 408, .. }) + crate::ocr::Error::Transport(crate::transport::Error::Http { status: 408, .. }) )); server.abort(); } diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index 505bea8a524..8c4ca08c39a 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -4,28 +4,25 @@ use data_url::{DataUrl, DataUrlError, forgiving_base64::DecodeError}; use reqwest::Url; use serde_json::Map; -use super::Error; use super::types::{OcrConnection, OcrDocument}; use crate::constants::{OCR_INLINE_MAX_BYTES, OCR_MAX_FETCH_REDIRECTS}; -use crate::media::Error as MediaError; use crate::media::{DownloadPolicy, MediaFetcher}; -use crate::transport::Error as TransportError; pub fn encode_file_document( bytes: &[u8], file_name: Option<&str>, mime_type: Option<&str>, -) -> Result { +) -> Result { if bytes.is_empty() { - return Err(Error::EmptyFile); + return Err(crate::ocr::Error::EmptyFile); } if bytes.len() > OCR_INLINE_MAX_BYTES { - return Err(Error::InlineDocumentTooLarge); + return Err(crate::ocr::Error::InlineDocumentTooLarge); } if let Some(value) = mime_type && !valid_mime_type(value) { - return Err(Error::InvalidMimeType(value.into())); + return Err(crate::ocr::Error::InvalidMimeType(value.into())); } let mime_type = mime_type .map(str::to_string) @@ -90,11 +87,11 @@ pub fn upload_mime_type<'a>(file_name: Option<&str>, content_type: Option<&'a st pub(crate) struct InlineDocument<'a>(DataUrl<'a>); impl<'a> InlineDocument<'a> { - pub(crate) fn parse(source: &'a str) -> Result, Error> { + pub(crate) fn parse(source: &'a str) -> Result, crate::ocr::Error> { match DataUrl::process(source) { Ok(url) => Ok(Some(Self(url))), Err(DataUrlError::NotADataUrl) => Ok(None), - Err(DataUrlError::NoComma) => Err(Error::InvalidDataUri), + Err(DataUrlError::NoComma) => Err(crate::ocr::Error::InvalidDataUri), } } @@ -102,26 +99,27 @@ impl<'a> InlineDocument<'a> { self.0.mime_type() } - pub(crate) fn decode(&self, max_bytes: usize) -> Result, Error> { + pub(crate) fn decode(&self, max_bytes: usize) -> Result, crate::ocr::Error> { let mut body = Vec::new(); self.0 .decode(|bytes| { if bytes.len() > max_bytes.saturating_sub(body.len()) { - return Err(Error::InlineDocumentTooLarge); + return Err(crate::ocr::Error::InlineDocumentTooLarge); } body.extend_from_slice(bytes); Ok(()) }) .map_err(|error| match error { - DecodeError::InvalidBase64(_) => Error::InvalidDataUri, + DecodeError::InvalidBase64(_) => crate::ocr::Error::InvalidDataUri, DecodeError::WriteError(error) => error, })?; Ok(body) } } -pub(crate) fn validate_inline_document(document: &OcrDocument) -> Result<(), Error> { - let inline = InlineDocument::parse(document.source())?.ok_or(Error::InvalidDataUri)?; +pub(crate) fn validate_inline_document(document: &OcrDocument) -> Result<(), crate::ocr::Error> { + let inline = + InlineDocument::parse(document.source())?.ok_or(crate::ocr::Error::InvalidDataUri)?; inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?; Ok(()) } @@ -130,13 +128,13 @@ pub(crate) async fn inline_remote_document( fetcher: &MediaFetcher, document: OcrDocument, connection: &OcrConnection, -) -> Result { +) -> Result { let source = document.source(); if !source.starts_with("http://") && !source.starts_with("https://") { validate_inline_document(&document)?; return Ok(document); } - let url = Url::parse(source).map_err(|_| Error::RequestField { + let url = Url::parse(source).map_err(|_| crate::ocr::Error::RequestField { path: "document URL".into(), })?; let downloaded = fetcher @@ -159,25 +157,25 @@ pub(crate) async fn inline_remote_document( Ok(result) } -fn map_media_error(error: MediaError) -> Error { +fn map_media_error(error: crate::media::Error) -> crate::ocr::Error { match error { - MediaError::BlockedUrl => Error::BlockedDocumentUrl, - MediaError::DownloadDisabled => Error::DownloadDisabled, - MediaError::DownloadTooLarge => Error::DownloadTooLarge, - MediaError::TooManyRedirects => Error::TooManyRedirects, - MediaError::MissingRedirectLocation => Error::MissingRedirectLocation, - MediaError::InvalidRedirect => Error::InvalidRedirect, - MediaError::Http(status) => TransportError::Http { + crate::media::Error::BlockedUrl => crate::ocr::Error::BlockedDocumentUrl, + crate::media::Error::DownloadDisabled => crate::ocr::Error::DownloadDisabled, + crate::media::Error::DownloadTooLarge => crate::ocr::Error::DownloadTooLarge, + crate::media::Error::TooManyRedirects => crate::ocr::Error::TooManyRedirects, + crate::media::Error::MissingRedirectLocation => crate::ocr::Error::MissingRedirectLocation, + crate::media::Error::InvalidRedirect => crate::ocr::Error::InvalidRedirect, + crate::media::Error::Http(status) => crate::transport::Error::Http { status, body: "OCR document download failed".into(), } .into(), - MediaError::Timeout => TransportError::Http { + crate::media::Error::Timeout => crate::transport::Error::Http { status: 408, body: "OCR document download timed out".into(), } .into(), - MediaError::Transport(error) => error.into(), + crate::media::Error::Transport(error) => error.into(), } } @@ -254,7 +252,7 @@ mod tests { let bytes = vec![b'a'; OCR_INLINE_MAX_BYTES + 1]; assert_eq!( encode_file_document(&bytes, None, None), - Err(Error::InlineDocumentTooLarge) + Err(crate::ocr::Error::InlineDocumentTooLarge) ); let document = encode_file_document(&bytes[..OCR_INLINE_MAX_BYTES], None, None).unwrap(); let inline = InlineDocument::parse(document.source()).unwrap().unwrap(); @@ -288,7 +286,7 @@ mod tests { assert_eq!(inline.decode(expected.len()).unwrap(), expected); assert_eq!( inline.decode(expected.len() - 1), - Err(Error::InlineDocumentTooLarge) + Err(crate::ocr::Error::InlineDocumentTooLarge) ); } } diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index c2e34d81f50..2b14352a615 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -21,12 +21,11 @@ use crate::llms::vertex_ai::ocr::deepseek_transformation::{ DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig, }; use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig; -use crate::ocr::Error; pub(crate) async fn perform_ocr_request( client: &OcrClient, request: LiteLLMOcrRequest, -) -> Result { +) -> Result { request.response_format()?; let context = CallLifecycleContext::new( "ocr", @@ -62,7 +61,7 @@ impl PreparedOcrCall { pub(crate) async fn prepare( client: OcrClient, request: LiteLLMOcrRequest, - ) -> Result { + ) -> Result { let http = match request.config { OcrConfigKind::Cohere => CohereParseConfig.prepare_request(&request, &client).await?, OcrConfigKind::Mistral => MistralOCRConfig.prepare_request(&request, &client).await?, @@ -101,7 +100,7 @@ impl PreparedOcrCall { }) } - pub(crate) async fn execute(self) -> Result { + pub(crate) async fn execute(self) -> Result { let url = self.http.url().to_string(); let headers = request_headers(&self.http)?; let response = @@ -162,7 +161,7 @@ impl PreparedOcrCall { } } -fn request_headers(request: &reqwest::Request) -> Result, Error> { +fn request_headers(request: &reqwest::Request) -> Result, super::Error> { request .headers() .iter() @@ -195,7 +194,7 @@ pub(crate) struct OcrProviderResponse { } impl OcrProviderResponse { - pub(crate) fn normalize(self) -> Result { + pub(crate) fn normalize(self) -> Result { let (response, native) = match self.data { OcrProviderData::Cohere(decoded) => ( CohereParseConfig.transform_ocr_response(&self.request, decoded.data)?, @@ -242,7 +241,7 @@ impl OcrProviderResponse { } } -pub(crate) async fn post_call(hooks: &Arc, bytes: &[u8]) -> Result<(), Error> { +pub(crate) async fn post_call(hooks: &Arc, bytes: &[u8]) -> Result<(), super::Error> { let original_response = serde_json::Value::String(String::from_utf8_lossy(bytes).into_owned()); hooks .post_call(OcrPostCallRequest { original_response }) diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs index f3e4f1ce37d..2860470b68b 100644 --- a/litellm-rust/crates/core/src/ocr/hooks.rs +++ b/litellm-rust/crates/core/src/ocr/hooks.rs @@ -4,11 +4,10 @@ use std::sync::Arc; use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument}; use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; -use crate::ocr::Error; use serde::Serialize; use serde_json::Value; -pub type OcrHookFuture<'a, T> = Pin> + Send + 'a>>; +pub type OcrHookFuture<'a, T> = Pin> + Send + 'a>>; pub type OcrLogFuture<'a> = Pin + Send + 'a>>; #[derive(Clone, Debug, Serialize)] @@ -62,7 +61,7 @@ pub trait OcrHooks: Send + Sync { fn failure<'a>( &'a self, _context: &'a CallLifecycleContext, - _error: &'a Error, + _error: &'a super::Error, _timing: &'a CallLifecycleTiming, ) -> OcrLogFuture<'a> { Box::pin(async {}) @@ -80,7 +79,7 @@ pub(crate) struct OcrLifecycleHooks { impl CallLifecycleHooks for OcrLifecycleHooks { - type Error = Error; + type Error = super::Error; type PreCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>; type DuringCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>; type SuccessFuture<'a> = OcrLogFuture<'a>; @@ -137,7 +136,7 @@ impl CallLifecycleHooks( &'a self, context: &'a CallLifecycleContext, - error: &'a Error, + error: &'a super::Error, timing: &'a CallLifecycleTiming, ) -> Self::FailureFuture<'a> { self.hooks.failure(context, error, timing) diff --git a/litellm-rust/crates/core/src/ocr/lifecycle.rs b/litellm-rust/crates/core/src/ocr/lifecycle.rs index b45bda4cae7..b11f5688c59 100644 --- a/litellm-rust/crates/core/src/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/src/ocr/lifecycle.rs @@ -14,11 +14,9 @@ use crate::call_lifecycle::host::{ HostCall, HostCallFuture, HostCallStep, HostFailure, HostLifecycle, HostPhase, }; use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming}; -use crate::ocr::Error; -use litellm_auth::Error as AuthError; use litellm_auth::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; -pub type NativeResult = Result, Error>; +pub type NativeResult = Result, super::Error>; #[derive(Debug, PartialEq, Eq)] pub enum NativeOutcome { @@ -54,7 +52,7 @@ pub enum OcrHostOperation { ProjectRequest, Lifecycle(HostPhase), ConstructResponse(Arc), - MapFailure(Error), + MapFailure(super::Error), Success { context: CallLifecycleContext, response: Arc, @@ -62,7 +60,7 @@ pub enum OcrHostOperation { }, Failure { context: CallLifecycleContext, - error: Error, + error: super::Error, timing: CallLifecycleTiming, }, AcquireAzureAdToken, @@ -83,12 +81,12 @@ impl OcrHostOperation { } pub enum OcrHostResult { - Request(Result<(Box, bool), Error>), - Lifecycle(Result<(), HostFailure>), - AzureAdToken(Result), - PreCall(Result), - DuringCall(Result), - PostCall(Result), + Request(Result<(Box, bool), super::Error>), + Lifecycle(Result<(), HostFailure>), + AzureAdToken(Result), + PreCall(Result), + DuringCall(Result), + PostCall(Result), } pub type OcrCallStep = HostCallStep; @@ -97,7 +95,7 @@ pub struct OcrCall { lifecycle: HostLifecycle, execution: OcrExecution, response: Option>, - error: Option, + error: Option, pending: bool, completed: bool, projecting: bool, @@ -122,14 +120,17 @@ impl OcrCall { }) } - pub async fn resume(&mut self, result: Option) -> Result { + pub async fn resume( + &mut self, + result: Option, + ) -> Result { if self.completed { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "OCR call cannot be resumed after completion".into(), )); } if self.pending != result.is_some() { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "OCR host operation result does not match pending state".into(), )); } @@ -137,7 +138,7 @@ impl OcrCall { Some(OcrHostResult::Lifecycle(Ok(()))) if self.lifecycle.phase() == HostPhase::Execute => { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "OCR provider operation requires a typed result".into(), )); } @@ -145,7 +146,7 @@ impl OcrCall { if !matches!(result, OcrHostResult::Lifecycle(_)) && self.lifecycle.phase() != HostPhase::Execute => { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "unexpected OCR provider operation result".into(), )); } @@ -165,7 +166,7 @@ impl OcrCall { None } Some(OcrHostResult::Request(_)) => { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "unexpected OCR request projection".into(), )); } @@ -206,20 +207,20 @@ impl OcrCall { .map(Arc::unwrap_or_clone) .map(OcrCallStep::Complete) .ok_or_else(|| { - Error::InvalidRequest("OCR completed without a response".into()) + super::Error::InvalidRequest("OCR completed without a response".into()) }), }; } HostPhase::ConstructResponse => OcrHostOperation::ConstructResponse( self.response .as_ref() - .ok_or_else(|| Error::InvalidRequest("missing OCR response".into()))? + .ok_or_else(|| super::Error::InvalidRequest("missing OCR response".into()))? .clone(), ), HostPhase::MapFailure => OcrHostOperation::MapFailure( self.error .as_ref() - .ok_or_else(|| Error::InvalidRequest("missing OCR failure".into()))? + .ok_or_else(|| super::Error::InvalidRequest("missing OCR failure".into()))? .clone(), ), HostPhase::Success | HostPhase::Failure => { @@ -235,7 +236,9 @@ impl OcrCall { response: self .response .as_ref() - .ok_or_else(|| Error::InvalidRequest("missing OCR response".into()))? + .ok_or_else(|| { + super::Error::InvalidRequest("missing OCR response".into()) + })? .clone(), timing, }, @@ -244,7 +247,9 @@ impl OcrCall { error: self .error .as_ref() - .ok_or_else(|| Error::InvalidRequest("missing OCR failure".into()))? + .ok_or_else(|| { + super::Error::InvalidRequest("missing OCR failure".into()) + })? .clone(), timing, }, @@ -256,7 +261,7 @@ impl OcrCall { Ok(self.host_step(operation)) } - fn accept(&mut self, result: Result<(), HostFailure>) { + fn accept(&mut self, result: Result<(), HostFailure>) { let cancelled = matches!(&result, Err(HostFailure::Cancelled(_))); if let Some(error) = self.lifecycle.accept(result) { if cancelled { @@ -268,9 +273,12 @@ impl OcrCall { } } - pub async fn interrupt(&mut self, failure: HostFailure) -> Result { + pub async fn interrupt( + &mut self, + failure: HostFailure, + ) -> Result { if self.completed { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "OCR call cannot be interrupted after completion".into(), )); } @@ -286,7 +294,7 @@ impl OcrCall { } impl HostCall for OcrCall { - type Error = Error; + type Error = super::Error; type Operation = OcrHostOperation; type Result = OcrHostResult; type Complete = LiteLLMOcrResponse; @@ -300,7 +308,7 @@ impl HostCall for OcrCall { fn interrupt( &mut self, - failure: HostFailure, + failure: HostFailure, ) -> HostCallFuture<'_, Self::Operation, Self::Complete, Self::Error> { Box::pin(OcrCall::interrupt(self, failure)) } @@ -317,7 +325,7 @@ struct OcrExecution { operations_tx: mpsc::UnboundedSender, operations_rx: mpsc::UnboundedReceiver, pending_result: Option>, - execution: Option>>, + execution: Option>>, completed: bool, azure_ad_token_provider: bool, terminal: Arc>>, @@ -339,25 +347,28 @@ impl OcrExecution { } } - pub async fn resume(&mut self, result: Option) -> Result { + pub async fn resume( + &mut self, + result: Option, + ) -> Result { if self.completed { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "OCR call cannot be resumed after completion".into(), )); } match (self.pending_result.take(), result) { - (Some(sender), Some(result)) => sender - .send(result) - .map_err(|_| Error::InvalidRequest("OCR host operation was abandoned".into()))?, + (Some(sender), Some(result)) => sender.send(result).map_err(|_| { + super::Error::InvalidRequest("OCR host operation was abandoned".into()) + })?, (None, None) if self.execution.is_none() => self.start(), (Some(sender), None) => { self.pending_result = Some(sender); - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "OCR host operation result is required".into(), )); } (None, Some(_)) => { - return Err(Error::InvalidRequest( + return Err(super::Error::InvalidRequest( "unexpected OCR host operation result".into(), )); } @@ -365,11 +376,11 @@ impl OcrExecution { } let execution = self.execution.as_mut().ok_or_else(|| { - Error::InvalidRequest("OCR call cannot be resumed after completion".into()) + super::Error::InvalidRequest("OCR call cannot be resumed after completion".into()) })?; tokio::select! { operation = self.operations_rx.recv() => { - let operation = operation.ok_or_else(|| Error::InvalidRequest("OCR operation channel closed".into()))?; + let operation = operation.ok_or_else(|| super::Error::InvalidRequest("OCR operation channel closed".into()))?; self.pending_result = Some(operation.result); Ok(OcrCallStep::Host(operation.operation)) } @@ -377,7 +388,7 @@ impl OcrExecution { self.execution = None; self.completed = true; result - .map_err(|error| Error::Transport(crate::transport::Error::Network(format!("OCR execution task failed: {error}"))))? + .map_err(|error| super::Error::Transport(crate::transport::Error::Network(format!("OCR execution task failed: {error}"))))? .map(OcrCallStep::Complete) } } @@ -452,15 +463,17 @@ impl TokenProvider for OcrAzureAdTokenProvider { result, }) .map_err(|_| { - AuthError::AzureTokenAcquisition("OCR host driver was abandoned".into()) + litellm_auth::Error::AzureTokenAcquisition( + "OCR host driver was abandoned".into(), + ) })?; match receiver.await.map_err(|_| { - AuthError::AzureTokenAcquisition( + litellm_auth::Error::AzureTokenAcquisition( "OCR token provider operation was abandoned".into(), ) })? { OcrHostResult::AzureAdToken(result) => result, - _ => Err(AuthError::AzureTokenAcquisition( + _ => Err(litellm_auth::Error::AzureTokenAcquisition( "invalid OCR token provider host result".into(), )), } @@ -469,14 +482,14 @@ impl TokenProvider for OcrAzureAdTokenProvider { } impl ProtocolHooks { - async fn invoke(&self, operation: OcrHostOperation) -> Result { + async fn invoke(&self, operation: OcrHostOperation) -> Result { let (result, receiver) = oneshot::channel(); self.operations .send(PendingOperation { operation, result }) - .map_err(|_| Error::InvalidRequest("OCR host driver was abandoned".into()))?; + .map_err(|_| super::Error::InvalidRequest("OCR host driver was abandoned".into()))?; receiver .await - .map_err(|_| Error::InvalidRequest("OCR host operation was abandoned".into())) + .map_err(|_| super::Error::InvalidRequest("OCR host operation was abandoned".into())) } } @@ -489,7 +502,7 @@ impl OcrHooks for ProtocolHooks { Box::pin(async move { match self.invoke(OcrHostOperation::PreCall(request)).await? { OcrHostResult::PreCall(result) => result, - _ => Err(Error::InvalidRequest( + _ => Err(super::Error::InvalidRequest( "invalid OCR pre-call host result".into(), )), } @@ -503,7 +516,7 @@ impl OcrHooks for ProtocolHooks { Box::pin(async move { match self.invoke(OcrHostOperation::DuringCall(request)).await? { OcrHostResult::DuringCall(result) => result, - _ => Err(Error::InvalidRequest( + _ => Err(super::Error::InvalidRequest( "invalid OCR during-call host result".into(), )), } @@ -514,7 +527,7 @@ impl OcrHooks for ProtocolHooks { Box::pin(async move { match self.invoke(OcrHostOperation::PostCall(request)).await? { OcrHostResult::PostCall(result) => result, - _ => Err(Error::InvalidRequest( + _ => Err(super::Error::InvalidRequest( "invalid OCR post-call host result".into(), )), } @@ -539,7 +552,7 @@ impl OcrHooks for ProtocolHooks { fn failure<'a>( &'a self, context: &'a CallLifecycleContext, - _error: &'a Error, + _error: &'a super::Error, timing: &'a CallLifecycleTiming, ) -> OcrLogFuture<'a> { Box::pin(async move { @@ -565,7 +578,7 @@ impl OcrHost for NoopOcrHost { Box::pin(async move { match operation { OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err( - Error::InvalidRequest("OCR host has no request projection".into()), + super::Error::InvalidRequest("OCR host has no request projection".into()), )), OcrHostOperation::Lifecycle(_) | OcrHostOperation::ConstructResponse(_) @@ -573,7 +586,7 @@ impl OcrHost for NoopOcrHost { | OcrHostOperation::Success { .. } | OcrHostOperation::Failure { .. } => OcrHostResult::Lifecycle(Ok(())), OcrHostOperation::AcquireAzureAdToken => { - OcrHostResult::AzureAdToken(Err(AuthError::AzureTokenAcquisition( + OcrHostResult::AzureAdToken(Err(litellm_auth::Error::AzureTokenAcquisition( "OCR host has no Azure AD token provider".into(), ))) } @@ -600,7 +613,7 @@ impl OcrHost for OcrHookHost { Box::pin(async move { match operation { OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err( - Error::InvalidRequest("OCR hook host has no request projection".into()), + super::Error::InvalidRequest("OCR hook host has no request projection".into()), )), OcrHostOperation::Success { context, @@ -622,7 +635,7 @@ impl OcrHost for OcrHookHost { | OcrHostOperation::ConstructResponse(_) | OcrHostOperation::MapFailure(_) => OcrHostResult::Lifecycle(Ok(())), OcrHostOperation::AcquireAzureAdToken => { - OcrHostResult::AzureAdToken(Err(AuthError::AzureTokenAcquisition( + OcrHostResult::AzureAdToken(Err(litellm_auth::Error::AzureTokenAcquisition( "OCR hook host has no Azure AD token provider".into(), ))) } diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 2bc8a9c1fc4..4a14d626ebf 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -1,7 +1,6 @@ use serde::{Serialize, de::DeserializeOwned}; use serde_json::Value; -use super::Error; use super::OcrClient; use super::hooks::OcrDuringCallRequest; use super::types::{LiteLLMOcrRequest, OcrDocument}; @@ -10,7 +9,7 @@ pub(crate) use crate::params::{ParsedProviderParams, merge_extra_params}; pub(crate) fn _prepare_ocr_request( request: &LiteLLMOcrRequest, -) -> Result, Error> { +) -> Result, super::Error> { super::wire::decode_request_value( Value::Object(request.optional_params.provider_params().into()), "optional_params", @@ -24,8 +23,8 @@ pub(crate) async fn transform_request_body( headers: &[(String, String)], retains_document: bool, body: B, - validate: impl Fn(&B) -> Result<(), Error>, -) -> Result + validate: impl Fn(&B) -> Result<(), super::Error>, +) -> Result where B: Serialize + DeserializeOwned, { @@ -36,7 +35,7 @@ where let composed = OcrWireBody::::decode(composed, "body")?; validate(&composed.body)?; let (body, headers) = if request.hooks.intercepts_requests() { - let body = serde_json::to_value(composed).map_err(|_| Error::RequestField { + let body = serde_json::to_value(composed).map_err(|_| super::Error::RequestField { path: "body".into(), })?; let retained_fields = request @@ -79,7 +78,7 @@ pub(crate) fn build_http_request( url: &str, headers: &[(String, String)], body: &B, -) -> Result { +) -> Result { let builder = client .provider_http() .post(url) @@ -88,14 +87,14 @@ pub(crate) fn build_http_request( crate::http_utils::with_headers(builder, headers, crate::http_utils::HeaderPolicy::All) .build() .map_err(crate::transport::Error::from) - .map_err(Error::from) + .map_err(super::Error::from) } pub(crate) async fn guardrail_document( request: &LiteLLMOcrRequest, url: &str, headers: &[(String, String)], -) -> Result<(OcrDocument, Vec<(String, String)>), Error> { +) -> Result<(OcrDocument, Vec<(String, String)>), super::Error> { if !request.hooks.intercepts_requests() { return Ok((request.document.clone(), headers.to_vec())); } @@ -106,8 +105,10 @@ pub(crate) async fn guardrail_document( custom_llm_provider: request.config.provider().as_str().into(), url: url.into(), headers: headers.to_vec(), - body: serde_json::to_value(&request.document).map_err(|_| Error::RequestField { - path: "document".into(), + body: serde_json::to_value(&request.document).map_err(|_| { + super::Error::RequestField { + path: "document".into(), + } })?, retained_fields: Vec::new(), }) @@ -125,14 +126,14 @@ struct OcrWireBody { } impl OcrWireBody { - fn decode(value: Value, prefix: &str) -> Result { + fn decode(value: Value, prefix: &str) -> Result { let body: B = super::wire::decode_request_value(value.clone(), prefix)?; let Value::Object(fields) = value else { - return Err(Error::RequestField { + return Err(super::Error::RequestField { path: prefix.into(), }); }; - let known = serde_json::to_value(&body).map_err(|_| Error::RequestField { + let known = serde_json::to_value(&body).map_err(|_| super::Error::RequestField { path: prefix.into(), })?; let extra = fields diff --git a/litellm-rust/crates/core/src/ocr/provider_config.rs b/litellm-rust/crates/core/src/ocr/provider_config.rs index 17ff895d19b..531a5e9f346 100644 --- a/litellm-rust/crates/core/src/ocr/provider_config.rs +++ b/litellm-rust/crates/core/src/ocr/provider_config.rs @@ -7,7 +7,6 @@ use crate::llms::mistral::ocr::transformation::MistralOCRConfig; use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, ReductoParseV3Config}; use crate::llms::vertex_ai::ocr::deepseek_transformation::VertexAIDeepSeekOCRConfig; use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig; -use crate::ocr::Error; use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; #[derive(Clone, Copy, Debug, PartialEq, Eq)] @@ -77,7 +76,7 @@ impl OcrProvider { pub(crate) fn resolve_provider_config( model: &str, custom_llm_provider: Option<&str>, -) -> Result<(String, OcrConfigKind), Error> { +) -> Result<(String, OcrConfigKind), super::Error> { let provider = get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider { model, @@ -104,7 +103,7 @@ pub(crate) fn resolve_provider_config( OcrConfigKind::VertexDeepSeek } "vertex_ai" => OcrConfigKind::VertexAi, - value => return Err(Error::InvalidProvider(value.to_string())), + value => return Err(super::Error::InvalidProvider(value.to_string())), }; Ok((provider.model.to_string(), config)) } diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 543adbe70e5..cd119bd8bcf 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -8,7 +8,6 @@ use serde_json::{Map, Value}; use super::hooks::{NoopOcrHooks, OcrHooks}; use super::provider_config::{OcrConfigKind, resolve_provider_config}; use crate::constants::OCR_HTTP_TIMEOUT_SECS; -use crate::ocr::Error; use crate::params::OpaqueParams; use litellm_auth::{InputSource, TokenProviderHandle}; @@ -108,7 +107,7 @@ impl LiteLLMOcrRequest { document: OcrDocument, custom_llm_provider: Option<&str>, optional_params: OpaqueParams, - ) -> Result { + ) -> Result { let (model, config) = resolve_provider_config(&model, custom_llm_provider)?; Ok(Self { diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index 7614e931cf6..33feb7b327c 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -2,7 +2,6 @@ use std::collections::BTreeMap; use std::time::Duration; use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument}; -use crate::ocr::Error; use crate::params::OpaqueParams; use litellm_auth::InputSource; use serde::{ @@ -68,7 +67,7 @@ pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> b pub fn consumed_optional_param_names( model: &str, custom_llm_provider: Option<&str>, -) -> Result, Error> { +) -> Result, crate::ocr::Error> { use super::provider_config::OcrConfigKind; let (provider_model, config) = @@ -92,7 +91,7 @@ pub fn consumed_optional_param_names( pub fn consumed_optional_params( model: &str, custom_llm_provider: Option<&str>, -) -> Result, Error> { +) -> Result, crate::ocr::Error> { consumed_optional_param_names(model, custom_llm_provider).map(|names| { names .into_iter() @@ -111,7 +110,7 @@ pub fn consumed_optional_params( }) } -pub fn decode_request(wire: OcrWireRequest) -> Result { +pub fn decode_request(wire: OcrWireRequest) -> Result { let api_key_source = source_for(&wire.input_sources, "api_key"); let api_base_source = source_for(&wire.input_sources, "api_base"); let extra_headers_source = source_for(&wire.input_sources, "extra_headers"); @@ -121,16 +120,18 @@ pub fn decode_request(wire: OcrWireRequest) -> Result .unwrap_or_default() .into_iter() .map(|(name, value)| { - let value = value.as_str().ok_or_else(|| Error::RequestField { - path: format!("extra_headers.{name}"), - })?; + let value = value + .as_str() + .ok_or_else(|| crate::ocr::Error::RequestField { + path: format!("extra_headers.{name}"), + })?; Ok((name, value.to_string())) }) - .collect::, Error>>()?; + .collect::, crate::ocr::Error>>()?; let timeout = wire .timeout_seconds .map(|seconds| { - Duration::try_from_secs_f64(seconds).map_err(|_| Error::RequestField { + Duration::try_from_secs_f64(seconds).map_err(|_| crate::ocr::Error::RequestField { path: "timeout_seconds".into(), }) }) @@ -144,7 +145,7 @@ pub fn decode_request(wire: OcrWireRequest) -> Result .as_u64() .and_then(|value| usize::try_from(value).ok()) .filter(|value| *value > 0 && *value <= defaults.max_response_bytes) - .ok_or_else(|| Error::RequestField { + .ok_or_else(|| crate::ocr::Error::RequestField { path: "max_response_bytes".into(), }) }) @@ -178,12 +179,12 @@ pub fn decode_request(wire: OcrWireRequest) -> Result }) } -fn decode_document(value: Value) -> Result { +fn decode_document(value: Value) -> Result { let kind = value.get("type").and_then(Value::as_str); let missing_url = matches!(kind, Some("document_url")) && value.get("document_url").is_none() || matches!(kind, Some("image_url")) && value.get("image_url").is_none(); if missing_url { - return Err(Error::MissingDocumentUrl); + return Err(crate::ocr::Error::MissingDocumentUrl); } decode_request_value(value, "document") } @@ -197,9 +198,12 @@ fn nonblank(value: Option) -> Option { .map(|s| s.trim().to_string()) .filter(|s| !s.is_empty()) } -pub fn decode_request_value(value: Value, prefix: &str) -> Result { +pub fn decode_request_value( + value: Value, + prefix: &str, +) -> Result { serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| { - Error::RequestField { + crate::ocr::Error::RequestField { path: format!("{prefix}.{}", error.path()), } }) @@ -208,19 +212,21 @@ pub fn decode_request_value(value: Value, prefix: &str) -> pub fn decode_response( bytes: &[u8], native: bool, -) -> Result, Error> { +) -> Result, crate::ocr::Error> { let mut deserializer = serde_json::Deserializer::from_slice(bytes); let data = serde_path_to_error::deserialize(&mut deserializer).map_err(|error| { - Error::ResponseField { + crate::ocr::Error::ResponseField { path: error.path().to_string(), } })?; - deserializer.end().map_err(|_| Error::ResponseField { - path: "response".into(), - })?; + deserializer + .end() + .map_err(|_| crate::ocr::Error::ResponseField { + path: "response".into(), + })?; let native = if native { Some( - serde_json::from_slice(bytes).map_err(|_| Error::ResponseField { + serde_json::from_slice(bytes).map_err(|_| crate::ocr::Error::ResponseField { path: "response".into(), })?, ) @@ -304,11 +310,11 @@ mod tests { })) .unwrap(); let error = decode_request(wire).err().expect("missing document URL"); - assert_eq!(error, Error::MissingDocumentUrl); + assert_eq!(error, crate::ocr::Error::MissingDocumentUrl); let error = crate::Error::from(error); assert!(matches!( error, - crate::Error::Ocr(Error::MissingDocumentUrl) + crate::Error::Ocr(crate::ocr::Error::MissingDocumentUrl) )); } } diff --git a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs index 74bb343d537..431d18f294f 100644 --- a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs +++ b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs @@ -1,5 +1,4 @@ use super::*; -use crate::chat_completions::Error; use serde_json::json; fn messages(value: Value) -> Vec { @@ -20,7 +19,9 @@ fn transform(model: &str, msgs: Value, opts: Value) -> Value { .body } -fn transform_response(body: Value) -> Result { +fn transform_response( + body: Value, +) -> Result { ANTHROPIC_CHAT_COMPLETIONS_CONFIG .transform_response("claude-sonnet-4-5", ProviderChatResponseData { body }) } @@ -413,26 +414,31 @@ fn declines_a_response_carrying_a_non_text_block() { "usage": {"input_tokens": 1, "output_tokens": 1} })) .expect_err("non-text block"); - assert_eq!(err, Error::Unsupported("non-text response content block")); + assert_eq!( + err, + crate::chat_completions::Error::Unsupported("non-text response content block") + ); } #[test] fn errors_on_a_response_missing_required_fields() { assert_eq!( transform_response(json!("nope")).expect_err("not an object"), - Error::InvalidResponse("messages response is not an object".to_string()) + crate::chat_completions::Error::InvalidResponse( + "messages response is not an object".to_string() + ) ); assert_eq!( transform_response(json!({"model": "m", "usage": {}})).expect_err("no content"), - Error::MissingField("content") + crate::chat_completions::Error::MissingField("content") ); assert_eq!( transform_response(json!({"model": "m", "content": []})).expect_err("no usage"), - Error::MissingField("usage") + crate::chat_completions::Error::MissingField("usage") ); assert_eq!( transform_response(json!({"content": [], "usage": {}})).expect_err("no model"), - Error::MissingField("model") + crate::chat_completions::Error::MissingField("model") ); } diff --git a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs index 75ed2a4a6c8..9aaa90aa9e2 100644 --- a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs @@ -1,6 +1,5 @@ use serde_json::{Map, Value, json}; -use crate::chat_completions::Error; use crate::chat_completions::conversation::{Conversation, build_conversation}; use crate::chat_completions::transformation::{ ChatCompletionsAuth, ChatCompletionsProviderConfig, Unsupported, unsupported_message, @@ -80,7 +79,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { _model: &str, _optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { Ok(complete_anthropic_url(api_base, env_lookup)) } @@ -90,7 +89,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { _model: &str, _optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { Ok(ChatCompletionsAuth::Header { name: "x-api-key", value: resolve_anthropic_api_key(api_key, env_lookup)?, @@ -143,7 +142,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { model: &str, messages: Vec, optional_params: OpaqueParams, - ) -> Result { + ) -> Result { Ok(ProviderChatRequestData { body: crate::params::merge_extra_params( &anthropic_body(model, &build_conversation(&messages), Map::new()), @@ -156,16 +155,17 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { &self, _model: &str, response: ProviderChatResponseData, - ) -> Result { - let body = response - .body - .as_object() - .ok_or_else(|| Error::InvalidResponse("messages response is not an object".into()))?; + ) -> Result { + let body = response.body.as_object().ok_or_else(|| { + crate::chat_completions::Error::InvalidResponse( + "messages response is not an object".into(), + ) + })?; let content = body .get("content") .and_then(Value::as_array) - .ok_or(Error::MissingField("content"))?; + .ok_or(crate::chat_completions::Error::MissingField("content"))?; // The route declines tool and thinking requests, so a non-text block // means the response carries something this path never asked for. // Decline rather than silently dropping it; the host falls back. @@ -173,7 +173,9 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { .iter() .any(|block| block.get("type").and_then(Value::as_str) != Some("text")) { - return Err(Error::Unsupported("non-text response content block")); + return Err(crate::chat_completions::Error::Unsupported( + "non-text response content block", + )); } let text: String = content .iter() @@ -183,7 +185,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { let usage = body .get("usage") .and_then(Value::as_object) - .ok_or(Error::MissingField("usage"))?; + .ok_or(crate::chat_completions::Error::MissingField("usage"))?; let field = |name: &str| usage.get(name).and_then(Value::as_u64).unwrap_or(0); Ok(ChatCompletionsResponse { @@ -191,7 +193,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { model: body .get("model") .and_then(Value::as_str) - .ok_or(Error::MissingField("model"))? + .ok_or(crate::chat_completions::Error::MissingField("model"))? .to_string(), choices: vec![ChatCompletionsChoice { index: 0, diff --git a/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs b/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs index 080f11c8cac..300c122f27b 100644 --- a/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs @@ -1,4 +1,3 @@ -use crate::messages::Error; use crate::messages::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy}; const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY"; @@ -49,7 +48,7 @@ impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig { api_base: Option<&str>, _model: &str, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { Ok(complete_anthropic_url(api_base, env_lookup)) } @@ -57,8 +56,8 @@ impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig { &self, api_key: Option<&str>, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from) + ) -> Result { + resolve_anthropic_api_key(api_key, env_lookup).map_err(crate::messages::Error::from) } fn auth_strategy(&self) -> MessagesAuthStrategy { diff --git a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs index dd122dc4536..9abab9e9a82 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs @@ -1,4 +1,3 @@ -use crate::messages::Error; use crate::messages::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy}; use crate::messages::types::{ AnthropicMessage, AnthropicMessagesRequest, AnthropicMessagesResponse, ContentBlock, @@ -28,12 +27,12 @@ pub const AZURE_ANTHROPIC_MESSAGES_CONFIG: AzureAnthropicMessagesConfig = pub fn resolve_azure_api_key( api_key: Option<&str>, env_lookup: &dyn Fn(&str) -> Option, -) -> Result { +) -> Result { non_empty(api_key) .map(str::to_string) .or_else(|| env_lookup(AZURE_API_KEY_ENV).filter(|value| !value.trim().is_empty())) .ok_or_else(|| { - Error::from(litellm_auth::Error::MissingApiKey { + crate::messages::Error::from(litellm_auth::Error::MissingApiKey { provider: "Azure", environment_variable: AZURE_API_KEY_ENV, }) @@ -43,11 +42,11 @@ pub fn resolve_azure_api_key( pub fn complete_azure_anthropic_url( api_base: Option<&str>, env_lookup: &dyn Fn(&str) -> Option, -) -> Result { +) -> Result { let api_base = non_empty(api_base) .map(str::to_string) .or_else(|| env_lookup(AZURE_API_BASE_ENV).filter(|value| !value.trim().is_empty())) - .ok_or_else(|| Error::from(litellm_auth::Error::MissingAzureApiBase))?; + .ok_or_else(|| crate::messages::Error::from(litellm_auth::Error::MissingAzureApiBase))?; let api_base = api_base.trim_end_matches('/'); @@ -141,7 +140,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig { api_base: Option<&str>, _model: &str, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { complete_azure_anthropic_url(api_base, env_lookup) } @@ -149,7 +148,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig { &self, api_key: Option<&str>, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { resolve_azure_api_key(api_key, env_lookup) } @@ -168,7 +167,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig { fn transform_request( &self, request: AnthropicMessagesRequest, - ) -> Result { + ) -> Result { let mut request = fold_system_role_messages(request); if let Some(system) = request.system.as_mut() { strip_scope_from_system(system); @@ -184,7 +183,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig { &self, model: &str, response: AnthropicMessagesResponse, - ) -> Result { + ) -> Result { self.anthropic.transform_response(model, response) } } @@ -262,7 +261,7 @@ mod tests { "https://env.services.ai.azure.com/anthropic/v1/messages" ); let err = complete_azure_anthropic_url(Some(" "), &|_| None).expect_err("missing base"); - assert!(matches!(err, Error::Auth(_))); + assert!(matches!(err, crate::messages::Error::Auth(_))); } #[test] diff --git a/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs b/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs index cc485a70bb4..da88a0f84a8 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs @@ -1,6 +1,5 @@ use serde_json::{Map, Value, json}; -use crate::audio_transcription::Error; use crate::audio_transcription::transformation::{ AudioTranscriptionAuth, AudioTranscriptionProviderConfig, }; @@ -20,22 +19,29 @@ pub static BEDROCK_AUDIO_TRANSCRIPTION_CONFIG: BedrockAudioTranscriptionConfig = pub struct BedrockAudioTranscriptionConfig; -fn audio_fields(audio: Value) -> Result<(String, String), Error> { - let object = audio.as_object().ok_or_else(|| Error::InvalidType { - expected: "object", - actual: json_type_name(&audio), - })?; +fn audio_fields(audio: Value) -> Result<(String, String), crate::audio_transcription::Error> { + let object = + audio + .as_object() + .ok_or_else(|| crate::audio_transcription::Error::InvalidType { + expected: "object", + actual: json_type_name(&audio), + })?; let data = object .get("data") .and_then(Value::as_str) .filter(|value| !value.is_empty()) - .ok_or(Error::MissingField("audio.data"))?; + .ok_or(crate::audio_transcription::Error::MissingField( + "audio.data", + ))?; let format = object .get("format") .and_then(Value::as_str) .filter(|value| matches!(*value, "wav" | "mp3" | "flac" | "ogg")) .ok_or_else(|| { - Error::InvalidRequest("audio.format must be wav, mp3, flac, or ogg".to_string()) + crate::audio_transcription::Error::InvalidRequest( + "audio.format must be wav, mp3, flac, or ogg".to_string(), + ) })?; Ok((data.to_string(), format.to_string())) } @@ -57,7 +63,7 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig { _model: &str, audio: Value, optional_params: OpaqueParams, - ) -> Result { + ) -> Result { let (data, format) = audio_fields(audio)?; let mut instruction = "Transcribe the audio. Respond with only the transcript.".to_string(); if let Some(language) = optional_string(&optional_params, "language") { @@ -92,14 +98,16 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig { &self, _model: &str, response_json: Value, - ) -> Result { + ) -> Result { let content = response_json .get("output") .and_then(|value| value.get("message")) .and_then(|value| value.get("content")) .and_then(Value::as_array) .ok_or_else(|| { - Error::InvalidResponse("Bedrock response has no output content".to_string()) + crate::audio_transcription::Error::InvalidResponse( + "Bedrock response has no output content".to_string(), + ) })?; let mut text = String::new(); for block in content { @@ -116,7 +124,7 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig { model: &str, optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { let (model_id, model_region) = bedrock_model_id_and_region(model); let region = resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup); let endpoint = optional_params @@ -138,7 +146,7 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig { model: &str, optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { let (_, model_region) = bedrock_model_id_and_region(model); Ok(AudioTranscriptionAuth::AwsSigV4 { region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup), diff --git a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs index 694e758aa2e..d756808e896 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs @@ -1,5 +1,4 @@ use super::*; -use crate::chat_completions::Error; use serde_json::json; fn messages(value: Value) -> Vec { @@ -24,7 +23,9 @@ fn transform(msgs: Value, opts: Value) -> Value { .body } -fn transform_response(body: Value) -> Result { +fn transform_response( + body: Value, +) -> Result { BEDROCK_CHAT_COMPLETIONS_CONFIG.transform_response( "anthropic.claude-sonnet-4-5-v1:0", ProviderChatResponseData { body }, @@ -508,22 +509,27 @@ fn declines_a_response_carrying_a_tool_use_block() { "usage": {"inputTokens": 1, "outputTokens": 1} })) .expect_err("tool use block"); - assert_eq!(err, Error::Unsupported("non-text response content block")); + assert_eq!( + err, + crate::chat_completions::Error::Unsupported("non-text response content block") + ); } #[test] fn errors_on_a_response_missing_required_fields() { assert_eq!( transform_response(json!("nope")).expect_err("not an object"), - Error::InvalidResponse("converse response is not an object".to_string()) + crate::chat_completions::Error::InvalidResponse( + "converse response is not an object".to_string() + ) ); assert_eq!( transform_response(json!({"usage": {}})).expect_err("no output"), - Error::MissingField("output.message.content") + crate::chat_completions::Error::MissingField("output.message.content") ); assert_eq!( transform_response(json!({"output": {"message": {"content": []}}})).expect_err("no usage"), - Error::MissingField("usage") + crate::chat_completions::Error::MissingField("usage") ); } diff --git a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs index 8e2c8259fdc..4d13451c0b3 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs @@ -1,6 +1,5 @@ use serde_json::{Map, Value, json}; -use crate::chat_completions::Error; use crate::chat_completions::conversation::{Conversation, TurnRole, build_conversation}; use crate::chat_completions::response_utils::{finish_reason_for, unix_now, usage_from_parts}; use crate::chat_completions::transformation::{ @@ -112,7 +111,7 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { model: &str, optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { let (model_id, model_region) = bedrock_model_id_and_region(model); let region = resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup); let endpoint = optional_params @@ -139,7 +138,7 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { model: &str, optional_params: &OpaqueParams, env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { + ) -> Result { // Python reads `api_key` as the Bedrock bearer token and consults the // env only when the caller passed none, so a caller-supplied empty key // falls through to SigV4 without reaching for the environment. An @@ -210,7 +209,7 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { _model: &str, messages: Vec, optional_params: OpaqueParams, - ) -> Result { + ) -> Result { Ok(ProviderChatRequestData { body: crate::params::merge_extra_params( &converse_body(&build_conversation(&messages), &optional_params), @@ -229,18 +228,21 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { &self, model: &str, response: ProviderChatResponseData, - ) -> Result { - let body = response - .body - .as_object() - .ok_or_else(|| Error::InvalidResponse("converse response is not an object".into()))?; + ) -> Result { + let body = response.body.as_object().ok_or_else(|| { + crate::chat_completions::Error::InvalidResponse( + "converse response is not an object".into(), + ) + })?; let content = body .get("output") .and_then(|output| output.get("message")) .and_then(|message| message.get("content")) .and_then(Value::as_array) - .ok_or(Error::MissingField("output.message.content"))?; + .ok_or(crate::chat_completions::Error::MissingField( + "output.message.content", + ))?; // The route declines tool requests, so anything other than a text block // is something this path never asked for. Decline; the host falls back. if content.iter().any(|block| { @@ -248,7 +250,9 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { .as_object() .is_none_or(|block| block.len() != 1 || !block.contains_key("text")) }) { - return Err(Error::Unsupported("non-text response content block")); + return Err(crate::chat_completions::Error::Unsupported( + "non-text response content block", + )); } let text: String = content .iter() @@ -258,7 +262,7 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { let usage = body .get("usage") .and_then(Value::as_object) - .ok_or(Error::MissingField("usage"))?; + .ok_or(crate::chat_completions::Error::MissingField("usage"))?; let field = |name: &str| usage.get(name).and_then(Value::as_u64).unwrap_or(0); let computed = usage_from_parts( field("inputTokens"), diff --git a/litellm-rust/crates/core/src/providers/model.rs b/litellm-rust/crates/core/src/providers/model.rs index 652e379e4b7..fcedc4b023a 100644 --- a/litellm-rust/crates/core/src/providers/model.rs +++ b/litellm-rust/crates/core/src/providers/model.rs @@ -1,9 +1,8 @@ use std::marker::PhantomData; -use serde::{Deserialize, Deserializer, Serialize, Serializer, de::Error as _}; -use thiserror::Error; +use serde::{Deserialize, Deserializer, Serialize, Serializer}; -#[derive(Clone, Copy, Debug, Eq, PartialEq, Error)] +#[derive(Clone, Copy, Debug, Eq, PartialEq, thiserror::Error)] pub(crate) enum ModelNameError { #[error("model name cannot be empty")] EmptyModel, @@ -75,7 +74,7 @@ impl<'de, N: ModelNamespace> Deserialize<'de> for ProviderModel { let value = String::deserialize(deserializer)?; RoutedModel::new(&value) .and_then(RoutedModel::into_provider::) - .map_err(D::Error::custom) + .map_err(::custom) } } diff --git a/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs b/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs index 69cdd3d33b8..6843923e04b 100644 --- a/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs +++ b/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs @@ -1,4 +1,3 @@ -use crate::responses::Error; use crate::responses::types::{ResponsesWsEvent, ResponsesWsTransformResult}; use crate::responses::websocket::{ResponsesWebSocketProviderConfig, enforce_model}; @@ -15,7 +14,7 @@ impl ResponsesWebSocketProviderConfig for OpenAIResponsesWsConfig { &self, event: &ResponsesWsEvent, model: &str, - ) -> Result { + ) -> Result { let mut event = event.clone(); event.data = event.data.into_provider_body()?.into(); Ok(ResponsesWsTransformResult::passthrough(enforce_model( @@ -27,7 +26,7 @@ impl ResponsesWebSocketProviderConfig for OpenAIResponsesWsConfig { &self, event: &ResponsesWsEvent, _model: &str, - ) -> Result { + ) -> Result { Ok(ResponsesWsTransformResult::passthrough(event.clone())) } } diff --git a/litellm-rust/crates/core/src/responses/instrumentation.rs b/litellm-rust/crates/core/src/responses/instrumentation.rs index d51640e8a93..6871cf847d7 100644 --- a/litellm-rust/crates/core/src/responses/instrumentation.rs +++ b/litellm-rust/crates/core/src/responses/instrumentation.rs @@ -6,7 +6,6 @@ use std::time::{SystemTime, UNIX_EPOCH}; use serde_json::Value; use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; -use crate::responses::Error; use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType}; #[derive(Clone, Debug, Default, PartialEq, Eq)] @@ -205,10 +204,11 @@ impl ResponsesWsInstrumentation { } } -type LifecycleFuture<'a, T> = Pin> + Send + 'a>>; +type LifecycleFuture<'a, T> = + Pin> + Send + 'a>>; impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation { - type Error = Error; + type Error = crate::responses::Error; type PreCallFuture<'a> = LifecycleFuture<'a, ()>; type DuringCallFuture<'a> = LifecycleFuture<'a, ()>; type SuccessFuture<'a> = Pin + Send + 'a>>; @@ -247,7 +247,7 @@ impl CallLifecycleHooks<(), (), ()> for ResponsesWsInstrumentation { fn async_log_failure_event<'a>( &'a self, _context: &'a CallLifecycleContext, - _error: &'a Error, + _error: &'a crate::responses::Error, _timing: &'a CallLifecycleTiming, ) -> Self::FailureFuture<'a> { Box::pin(async move { @@ -343,7 +343,7 @@ mod tests { ), (), &instrumentation, - |_| async { Ok::<(), Error>(()) }, + |_| async { Ok::<(), crate::responses::Error>(()) }, ) .await; diff --git a/litellm-rust/crates/core/src/responses/websocket.rs b/litellm-rust/crates/core/src/responses/websocket.rs index 53d884f93d8..62d6b3fba42 100644 --- a/litellm-rust/crates/core/src/responses/websocket.rs +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -17,7 +17,6 @@ use tokio_tungstenite::{ }; use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH}; -use crate::responses::Error; use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult}; pub trait ResponsesWebSocketProviderConfig: Sync { @@ -37,13 +36,13 @@ pub trait ResponsesWebSocketProviderConfig: Sync { &self, event: &ResponsesWsEvent, model: &str, - ) -> Result; + ) -> Result; fn transform_ws_response( &self, event: &ResponsesWsEvent, model: &str, - ) -> Result; + ) -> Result; } pub fn complete_websocket_url( @@ -203,22 +202,22 @@ impl ResponsesWebSocketConnection { url: &str, headers: &HashMap, timeout: Option, - ) -> Result { + ) -> Result { let mut request = url.into_client_request().map_err(|error| { - Error::Transport(crate::transport::Error::Network(error.to_string())) + super::Error::Transport(crate::transport::Error::Network(error.to_string())) })?; for (name, value) in headers { let header_name = name .parse::() - .map_err(|error| Error::InvalidRequest(error.to_string()))?; + .map_err(|error| super::Error::InvalidRequest(error.to_string()))?; let header_value = HeaderValue::from_str(value) - .map_err(|error| Error::InvalidRequest(error.to_string()))?; + .map_err(|error| super::Error::InvalidRequest(error.to_string()))?; request.headers_mut().insert(header_name, header_value); } let connect = connect_upstream(request); let result = match timeout { Some(timeout) => tokio::time::timeout(timeout, connect).await.map_err(|_| { - Error::Transport(crate::transport::Error::Network( + super::Error::Transport(crate::transport::Error::Network( "Responses WebSocket connection timed out".into(), )) })?, @@ -226,32 +225,31 @@ impl ResponsesWebSocketConnection { }; let (socket, _) = result.map_err(|error| match *error { tokio_tungstenite::tungstenite::Error::Http(response) => { - Error::Transport(crate::transport::Error::Http { + super::Error::Transport(crate::transport::Error::Http { status: response.status().as_u16(), body: String::new(), }) } - other => Error::Transport(crate::transport::Error::Network(other.to_string())), + other => super::Error::Transport(crate::transport::Error::Network(other.to_string())), })?; Ok(Self { socket: Arc::new(Mutex::new(Some(socket))), }) } - pub async fn send_text(&self, text: String) -> Result<(), Error> { + pub async fn send_text(&self, text: String) -> Result<(), super::Error> { let mut socket = self.socket.lock().await; let Some(socket) = socket.as_mut() else { - return Err(Error::Transport(crate::transport::Error::Network( + return Err(super::Error::Transport(crate::transport::Error::Network( "Responses WebSocket is closed".into(), ))); }; - socket - .send(Message::Text(text)) - .await - .map_err(|error| Error::Transport(crate::transport::Error::Network(error.to_string()))) + socket.send(Message::Text(text)).await.map_err(|error| { + super::Error::Transport(crate::transport::Error::Network(error.to_string())) + }) } - pub async fn recv_text(&self) -> Result, Error> { + pub async fn recv_text(&self) -> Result, super::Error> { let mut socket = self.socket.lock().await; let Some(socket) = socket.as_mut() else { return Ok(None); @@ -260,20 +258,20 @@ impl ResponsesWebSocketConnection { Some(Ok(Message::Text(text))) => Ok(Some(text)), Some(Ok(Message::Binary(bytes))) => String::from_utf8(bytes.to_vec()) .map(Some) - .map_err(|error| Error::InvalidResponse(error.to_string())), + .map_err(|error| super::Error::InvalidResponse(error.to_string())), Some(Ok(Message::Close(_))) | None => Ok(None), Some(Ok(_)) => Ok(None), - Some(Err(error)) => Err(Error::Transport(crate::transport::Error::Network( + Some(Err(error)) => Err(super::Error::Transport(crate::transport::Error::Network( error.to_string(), ))), } } - pub async fn close(&self) -> Result<(), Error> { + pub async fn close(&self) -> Result<(), super::Error> { let mut socket = self.socket.lock().await; if let Some(socket) = socket.as_mut() { socket.close(None).await.map_err(|error| { - Error::Transport(crate::transport::Error::Network(error.to_string())) + super::Error::Transport(crate::transport::Error::Network(error.to_string())) })?; } *socket = None; diff --git a/litellm-rust/crates/core/src/transport/error.rs b/litellm-rust/crates/core/src/transport/error.rs index 534575ebcb5..eff15365ea8 100644 --- a/litellm-rust/crates/core/src/transport/error.rs +++ b/litellm-rust/crates/core/src/transport/error.rs @@ -28,7 +28,6 @@ impl From for Error { #[cfg(test)] mod tests { - use super::Error; #[tokio::test] async fn transport_errors_remove_urls_and_keep_dispatch_context() { let error = reqwest::Client::builder() @@ -39,8 +38,8 @@ mod tests { .send() .await .expect_err("invalid port"); - let error = Error::from_reqwest_before_dispatch(error); - assert!(matches!(error, Error::Connect(_))); + let error = crate::transport::Error::from_reqwest_before_dispatch(error); + assert!(matches!(error, crate::transport::Error::Connect(_))); assert!(!error.to_string().contains("secret")); assert!(!error.to_string().contains("private")); } @@ -69,8 +68,8 @@ mod tests { let error = response.expect_err("server does not respond"); assert!(error.is_timeout()); assert!(matches!( - Error::from_reqwest_before_dispatch(error), - Error::Network(_) + crate::transport::Error::from_reqwest_before_dispatch(error), + crate::transport::Error::Network(_) )); } } diff --git a/litellm-rust/crates/core/src/url_utils.rs b/litellm-rust/crates/core/src/url_utils.rs index 1150f93a5c7..b8d82b7a04a 100644 --- a/litellm-rust/crates/core/src/url_utils.rs +++ b/litellm-rust/crates/core/src/url_utils.rs @@ -1,9 +1,8 @@ use std::marker::PhantomData; -use thiserror::Error; use url::Url; -#[derive(Debug, Error)] +#[derive(Debug, thiserror::Error)] pub(crate) enum ApiUrlError { #[error("invalid URL: {0}")] Parse(#[from] url::ParseError), diff --git a/litellm-rust/crates/core/tests/host_lifecycle.rs b/litellm-rust/crates/core/tests/host_lifecycle.rs index 1e9f463c0ff..a07979222a2 100644 --- a/litellm-rust/crates/core/tests/host_lifecycle.rs +++ b/litellm-rust/crates/core/tests/host_lifecycle.rs @@ -1,7 +1,6 @@ use crate::call_lifecycle::host::{HostFailure, HostLifecycle, HostPhase}; -use crate::ocr::Error; -fn run(fail_at: Option, asynchronous: bool) -> (Vec, Vec) { +fn run(fail_at: Option, asynchronous: bool) -> (Vec, Vec) { let mut lifecycle = HostLifecycle::new(asynchronous); let mut events = Vec::new(); let mut failures = Vec::new(); @@ -10,7 +9,7 @@ fn run(fail_at: Option, asynchronous: bool) -> (Vec, Vec(Ok(())); + lifecycle.accept::(Ok(())); } - let selected = Error::InvalidRequest("provider".into()); + let selected = crate::ocr::Error::InvalidRequest("provider".into()); assert_eq!( lifecycle.accept(Err(HostFailure::Error(selected.clone()))), Some(selected) ); - lifecycle.accept::(Ok(())); + lifecycle.accept::(Ok(())); for phase in [ HostPhase::DeploymentFailure, HostPhase::Failure, @@ -96,7 +95,7 @@ fn failure_handler_errors_do_not_replace_selected_failure_or_suppress_async_disp ] { assert_eq!(lifecycle.phase(), phase); assert_eq!( - lifecycle.accept(Err(HostFailure::Error(Error::InvalidRequest( + lifecycle.accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest( "callback".into() )))), None @@ -108,7 +107,7 @@ fn failure_handler_errors_do_not_replace_selected_failure_or_suppress_async_disp #[test] fn cancellation_skips_terminal_dispatch() { let mut lifecycle = HostLifecycle::new(true); - let error = Error::InvalidRequest("cancelled".into()); + let error = crate::ocr::Error::InvalidRequest("cancelled".into()); assert_eq!( lifecycle.accept(Err(HostFailure::Cancelled(error.clone()))), Some(error) diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index a47c22a8009..209d6ff4d2d 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -686,7 +686,7 @@ async fn missing_host_result_preserves_pending_operation() { async fn read_bounded_response( response: Vec, limit: usize, -) -> Result { +) -> Result { use tokio::io::{AsyncReadExt, AsyncWriteExt}; let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); @@ -715,8 +715,6 @@ async fn read_bounded_response( #[tokio::test] async fn response_limit_accepts_exact_size_and_rejects_declared_and_chunked_overflow() { - use super::Error; - for response in [ "HTTP/1.1 200 OK\r\nContent-Length: 8\r\n\r\nabcdefgh", "HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n4\r\nefgh\r\n0\r\n\r\n", @@ -734,7 +732,7 @@ async fn response_limit_accepts_exact_size_and_rejects_declared_and_chunked_over ] { assert!(matches!( read_bounded_response(response.as_bytes().to_vec(), 8).await, - Err(Error::TooLarge { limit: 8 }) + Err(crate::ocr::Error::TooLarge { limit: 8 }) )); } } @@ -753,7 +751,7 @@ async fn oversized_error_retains_http_status_and_bounded_diagnostics_without_dra .await .unwrap_err(); match error { - super::Error::Transport(crate::transport::Error::Http { status, body }) => { + crate::ocr::Error::Transport(crate::transport::Error::Http { status, body }) => { assert_eq!(status, 429); assert_eq!( body,