inlineing more errors

This commit is contained in:
Yujong Lee 2026-09-15 12:07:13 -07:00
parent 8cfb59082a
commit 03a1c4a938
50 changed files with 734 additions and 673 deletions

View file

@ -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<Value, Error> {
let body = serde_json::to_vec(&request.body)
.map_err(|error| Error::InvalidRequest(format!("invalid audio request body: {error}")))?;
) -> Result<Value, super::Error> {
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<Vec<(String, String)>, Error> {
) -> Result<Vec<(String, String)>, super::Error> {
use std::collections::BTreeMap;
use std::time::SystemTime;

View file

@ -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<ProviderAudioTranscriptionRequest, Error> {
) -> Result<ProviderAudioTranscriptionRequest, super::Error> {
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)?;

View file

@ -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<AudioTranscriptionRequestData, Error>;
) -> Result<AudioTranscriptionRequestData, super::Error>;
fn transform_transcription_response(
&self,
model: &str,
response_json: Value,
) -> Result<AudioTranscriptionResponseData, Error>;
) -> Result<AudioTranscriptionResponseData, super::Error>;
fn complete_url(
&self,
@ -40,12 +39,12 @@ pub trait AudioTranscriptionProviderConfig: Sync {
model: &str,
optional_params: &OpaqueParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error>;
) -> Result<String, super::Error>;
fn auth_strategy(
&self,
model: &str,
optional_params: &OpaqueParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<AudioTranscriptionAuth, Error>;
) -> Result<AudioTranscriptionAuth, super::Error>;
}

View file

@ -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<String, String, String> for RecordingHooks {
type Error = Error;
type PreCallFuture<'a> = BoxFuture<'a, Result<String, Error>>;
type DuringCallFuture<'a> = BoxFuture<'a, Result<String, Error>>;
type Error = crate::messages::Error;
type PreCallFuture<'a> = BoxFuture<'a, Result<String, crate::messages::Error>>;
type DuringCallFuture<'a> = BoxFuture<'a, Result<String, crate::messages::Error>>;
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<RecordingRequest, String, String> for RecordingHooks {
type Error = Error;
type PreCallFuture<'a> = BoxFuture<'a, Result<RecordingRequest, Error>>;
type DuringCallFuture<'a> = BoxFuture<'a, Result<String, Error>>;
type Error = crate::messages::Error;
type PreCallFuture<'a> = BoxFuture<'a, Result<RecordingRequest, crate::messages::Error>>;
type DuringCallFuture<'a> = BoxFuture<'a, Result<String, crate::messages::Error>>;
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::<String, Error>(Error::Transport(crate::transport::Error::Network(
"provider down".to_string(),
)))
Err::<String, crate::messages::Error>(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()
))
);

View file

@ -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<Map<String, Value>>,
) -> Result<Vec<(String, String)>, Error> {
shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(Error::from)
) -> Result<Vec<(String, String)>, super::Error> {
shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(super::Error::from)
}

View file

@ -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<ChatCompletionsResponse, Error> {
) -> Result<ChatCompletionsResponse, super::Error> {
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<Vec<(String, String)>, Error> {
) -> Result<Vec<(String, String)>, 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",
));
}

View file

@ -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<Vec<ChatMessage>, 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<Vec<ChatMessage>, 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<ResolvedChatCompletionsRequest<'_>, Error> {
) -> Result<ResolvedChatCompletionsRequest<'_>, 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<ProviderChatCompletionsRequest, Error> {
) -> Result<ProviderChatCompletionsRequest, super::Error> {
let (headers, auth) = validate_environment(&request, &request.model, request.config)?;
let model = request.model;
let config = request.config;

View file

@ -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<ProviderChatCompletionsRequest, Error> {
) -> Result<ProviderChatCompletionsRequest, crate::chat_completions::Error> {
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, .. })
..
})
));
}
}

View file

@ -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<String>,
) -> Result<String, Error>;
) -> Result<String, super::Error>;
fn auth(
&self,
@ -42,7 +41,7 @@ pub trait ChatCompletionsProviderConfig: Sync {
model: &str,
optional_params: &OpaqueParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<ChatCompletionsAuth, Error>;
) -> Result<ChatCompletionsAuth, super::Error>;
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<ChatMessage>,
optional_params: OpaqueParams,
) -> Result<ProviderChatRequestData, Error>;
) -> Result<ProviderChatRequestData, super::Error>;
fn transform_response(
&self,
model: &str,
response: ProviderChatResponseData,
) -> Result<ChatCompletionsResponse, Error>;
) -> Result<ChatCompletionsResponse, super::Error>;
}
pub fn unsupported_param(optional_params: &OpaqueParams) -> Option<Unsupported> {

View file

@ -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<reqwest::Request, Error> {
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = crate::ocr::wire::decode_request_value::<CohereParams>(
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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
CohereParseConfig.transform_ocr_response(request, response)
}
}
fn complete_url(base: &str) -> Result<String, Error> {
fn complete_url(base: &str) -> Result<String, crate::ocr::Error> {
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<String, Error> {
.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(),
}
}

View file

@ -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<String> + Sync),
) -> Result<Option<Sourced<String>>, Error> {
) -> Result<Option<Sourced<String>>, crate::ocr::Error> {
static SERVICE: OnceLock<AzureAuthService> = 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

View file

@ -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<Option<f64
fn decode_input_params(
params: Map<String, Value>,
prefix: &str,
) -> Result<ParsedProviderParams<DocumentIntelligenceInputParams>, Error> {
) -> Result<ParsedProviderParams<DocumentIntelligenceInputParams>, 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<DocumentIntelligenceParams, Error> {
) -> Result<DocumentIntelligenceParams, crate::ocr::Error> {
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<Option<String>, Error> {
fn normalize_pages(pages: PagesInput) -> Result<Option<String>, crate::ocr::Error> {
let normalized = match pages {
PagesInput::ZeroBasedIndices(indices) => {
if indices.is_empty() {
@ -213,10 +214,11 @@ fn normalize_pages(pages: PagesInput) -> Result<Option<String>, 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::<Result<BTreeSet<_>, _>>()?
.into_iter()
@ -241,7 +243,7 @@ fn normalize_pages(pages: PagesInput) -> Result<Option<String>, 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<Option<String>, Error> {
fn normalize_features(features: FeaturesInput) -> Result<Option<String>, 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<Option<String>, 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<DocumentIntelligenceRequest, Error> {
fn transform_ocr_request(
document: OcrDocument,
) -> Result<DocumentIntelligenceRequest, crate::ocr::Error> {
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<DocumentIntelligenceRe
fn transform_ocr_response(
model: &str,
response: AzureDocumentIntelligenceOperation,
) -> Result<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
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<Value, Error> {
fn normalize_page(page: AzureDocumentIntelligencePage) -> Result<Value, crate::ocr::Error> {
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<Value, Error> {
}))
}
fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result<i64, Error> {
fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result<i64, crate::ocr::Error> {
let value = value * scale;
if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 {
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<dyn OcrHooks>,
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, Error> {
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, 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<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, Error> {
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, 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<reqwest::Request, Error> {
) -> Result<reqwest::Request, crate::ocr::Error> {
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, &params)?;
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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
transform_ocr_response(&request.model, response)
}
@ -533,8 +537,10 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
url: &str,
headers: &[(String, String)],
request: &LiteLLMOcrRequest,
) -> Result<crate::ocr::wire::DecodedOcrResponse<AzureDocumentIntelligenceOperation>, Error>
{
) -> Result<
crate::ocr::wire::DecodedOcrResponse<AzureDocumentIntelligenceOperation>,
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<DocumentIntelligenceParams, Error> {
fn map_ocr_params(
request: &LiteLLMOcrRequest,
) -> Result<DocumentIntelligenceParams, crate::ocr::Error> {
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<String, Error> {
) -> Result<String, crate::ocr::Error> {
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<String> + Sync),
) -> Result<Vec<(String, String)>, Error> {
) -> Result<Vec<(String, String)>, 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<DocumentIntelligenceParams, Error> {
fn map(value: Value) -> Result<DocumentIntelligenceParams, crate::ocr::Error> {
let fields = value.as_object().unwrap().clone();
normalize_ocr_params(decode_input_params(fields, "optional_params")?.known)
}

View file

@ -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<reqwest::Request, Error> {
) -> Result<reqwest::Request, crate::ocr::Error> {
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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
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<String>,
) -> Result<String, Error> {
) -> Result<String, crate::ocr::Error> {
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<String> + Sync),
) -> Result<Vec<(String, String)>, Error> {
) -> Result<Vec<(String, String)>, 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()))
}

View file

@ -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<Output = Result<reqwest::Request, Error>> + Send;
) -> impl Future<Output = Result<reqwest::Request, crate::ocr::Error>> + Send;
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, Error>;
) -> Result<LiteLLMOcrResponse, crate::ocr::Error>;
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<Output = Result<DecodedOcrResponse<Self::ProviderResponse>, Error>> + Send
) -> impl Future<Output = Result<DecodedOcrResponse<Self::ProviderResponse>, crate::ocr::Error>> + Send
{
async move {
let bytes = crate::ocr::client::read_response_bytes(

View file

@ -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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
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::<Result<Vec<_>, Error>>()?;
.collect::<Result<Vec<_>, crate::ocr::Error>>()?;
Ok(LiteLLMOcrResponse {
pages,
model: model.into(),
@ -148,7 +148,7 @@ impl CohereParseConfig {
model: &str,
document: OcrDocument,
params: CohereParams,
) -> Result<CohereRequest, Error> {
) -> Result<CohereRequest, crate::ocr::Error> {
validate_document(&document)?;
Ok(CohereRequest {
model: model.into(),
@ -169,7 +169,7 @@ impl BaseOcrConfig for CohereParseConfig {
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, Error> {
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = crate::ocr::wire::decode_request_value::<CohereParams>(
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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
transform_response(&request.model, response)
}
}
fn complete_url(base: &str) -> Result<String, Error> {
fn complete_url(base: &str) -> Result<String, crate::ocr::Error> {
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<String, Error> {
.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<String> + Sync),
) -> Result<Vec<(String, String)>, Error> {
) -> Result<Vec<(String, String)>, 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::<CohereParams>(json!({"output_format":"html"})).is_err());
@ -369,7 +369,7 @@ mod tests {
},
&|_| None,
),
Err(Error::Auth(_))
Err(crate::ocr::Error::Auth(_))
));
}
}

View file

@ -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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
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<MistralOcrRequest, Error> {
) -> Result<MistralOcrRequest, crate::ocr::Error> {
Ok(MistralOcrRequest {
model: model.to_string(),
document,
@ -89,7 +88,7 @@ impl BaseOcrConfig for MistralOCRConfig {
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, Error> {
) -> Result<reqwest::Request, crate::ocr::Error> {
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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
transform_ocr_response(&request.model, response)
}
}
pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result<String, Error> {
pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result<String, crate::ocr::Error> {
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<String, Error>
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<String, Error>
fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, Error> {
) -> Result<Vec<(String, String)>, 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,
}
))
));
}
}

View file

@ -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<ReductoV3Request, crate::ocr::Error> {
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<reqwest::Request, crate::ocr::Error> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoV3Params>(request)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document = prepare_document(client, document, &request.connection, &headers).await?;
let body = self.transform_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
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<ReductoLegacyRequest, crate::ocr::Error> {
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<reqwest::Request, crate::ocr::Error> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoLegacyParams>(request)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document = prepare_document(client, document, &request.connection, &headers).await?;
let body = self.transform_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
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<i64> {
.then(|| value.trunc() as i64)
}
#[tracing::instrument(
name = "transform_ocr_request",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
fn transform_v3_ocr_request(
_model: &str,
document: OcrDocument,
params: &ReductoV3Params,
) -> Result<ReductoV3Request, Error> {
Ok(ReductoV3Request {
input: document.source().to_string(),
params: params.clone(),
})
}
#[tracing::instrument(
name = "transform_ocr_request",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
fn transform_legacy_ocr_request(
_model: &str,
document: OcrDocument,
params: &ReductoLegacyParams,
) -> Result<ReductoLegacyRequest, Error> {
Ok(ReductoLegacyRequest {
document_url: document.source().to_string(),
options: params.enhance.as_ref().map(|_| params.clone()),
})
}
pub(crate) fn transform_ocr_response(
model: &str,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
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>) -> Value {
}
result
}
fn get_complete_url(api_base: Option<&str>, path: &str) -> Result<String, Error> {
fn get_complete_url(api_base: Option<&str>, path: &str) -> Result<String, crate::ocr::Error> {
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<String, Error>
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<String, Error>
fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, Error> {
) -> Result<Vec<(String, String)>, 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<OcrDocument, Error> {
) -> Result<OcrDocument, crate::ocr::Error> {
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<reqwest::Request, Error> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoLegacyParams>(request)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document = prepare_document(client, document, &request.connection, &headers).await?;
let body = transform_legacy_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, Error> {
transform_ocr_response(&request.model, response)
}
}
#[derive(Clone, Debug)]
pub(crate) struct ReductoParseV3Config;
impl BaseOcrConfig for ReductoParseV3Config {
type ProviderResponse = ReductoResponse;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["formatting", "retrieval", "settings"]
}
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, Error> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoV3Params>(request)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document = prepare_document(client, document, &request.connection, &headers).await?;
let body = transform_v3_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, Error> {
transform_ocr_response(&request.model, response)
}
}
#[cfg(test)]
mod tests {
use super::*;

View file

@ -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());
}

View file

@ -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<DeepSeekAi>,
document: OcrDocument,
params: &DeepSeekOcrParams,
) -> Result<DeepSeekOcrRequest, Error> {
) -> Result<DeepSeekOcrRequest, crate::ocr::Error> {
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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
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<DecodedContent, Error> {
fn decode_content(content: DeepSeekContent) -> Result<DecodedContent, crate::ocr::Error> {
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<Option<DeepSeekOcrResult>, Error> {
fn decode_json_content(text: &str) -> Result<Option<DeepSeekOcrResult>, 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<reqwest::Request, Error> {
) -> Result<reqwest::Request, crate::ocr::Error> {
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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
mapping::transform_ocr_response(&request.model, response)
}
}
pub(crate) fn provider_model(model: &str) -> Result<ProviderModel<DeepSeekAi>, Error> {
pub(crate) fn provider_model(model: &str) -> Result<ProviderModel<DeepSeekAi>, crate::ocr::Error> {
RoutedModel::new(model)
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.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<String, Error> {
) -> Result<String, crate::ocr::Error> {
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(),
})
}

View file

@ -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<reqwest::Request, Error> {
) -> Result<reqwest::Request, crate::ocr::Error> {
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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
MistralOCRConfig.transform_ocr_response(request, response)
}
}
@ -86,7 +85,7 @@ fn get_complete_url(
project: &str,
location: &str,
model: &str,
) -> Result<String, Error> {
) -> Result<String, crate::ocr::Error> {
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(),
})
}

View file

@ -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)
}
}

View file

@ -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<Map<String, Value>>,
) -> Result<Vec<(String, String)>, Error> {
shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(Error::from)
) -> Result<Vec<(String, String)>, super::Error> {
shared_string_headers(HEADER_CONTEXT, extra_headers).map_err(super::Error::from)
}

View file

@ -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<AnthropicMessagesResponse, Error> {
) -> Result<AnthropicMessagesResponse, super::Error> {
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<reqwest::Response, Error> {
) -> Result<reqwest::Response, super::Error> {
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),
}));

View file

@ -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<ProviderMessagesRequest, Error> {
) -> Result<ProviderMessagesRequest, super::Error> {
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<Map<String, Value>>,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<Vec<(String, String)>, Error> {
) -> Result<Vec<(String, String)>, super::Error> {
let mut headers = string_headers(extra_headers)?;
let auth_strategy = config.auth_strategy();

View file

@ -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"));
}

View file

@ -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<String>,
) -> Result<String, Error>;
) -> Result<String, super::Error>;
fn resolve_api_key(
&self,
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error>;
) -> Result<String, super::Error>;
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<AnthropicMessagesRequest, Error> {
) -> Result<AnthropicMessagesRequest, super::Error> {
Ok(request)
}
@ -56,7 +55,7 @@ pub trait AnthropicMessagesProviderConfig: Sync {
&self,
_model: &str,
response: AnthropicMessagesResponse,
) -> Result<AnthropicMessagesResponse, Error> {
) -> Result<AnthropicMessagesResponse, super::Error> {
Ok(response)
}
}

View file

@ -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<Self, TransportError> {
let document_fetcher = MediaFetcher::new().map_err(TransportError::from)?;
pub fn new(provider_http: reqwest::Client) -> Result<Self, crate::transport::Error> {
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<Self, Error> {
pub fn shared() -> Result<Self, crate::ocr::Error> {
shared_client()
}
pub async fn perform(&self, request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
#[tracing::instrument(
name = "ocr",
target = "litellm::function_trace",
level = "trace",
skip_all
)]
pub async fn perform(
&self,
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
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<reqwest::Client, TransportError> {
fn no_redirect_http() -> Result<reqwest::Client, crate::transport::Error> {
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<OcrClient, Error> {
static CLIENT: OnceLock<Result<OcrClient, TransportError>> = OnceLock::new();
pub(crate) fn shared_client() -> Result<OcrClient, crate::ocr::Error> {
static CLIENT: OnceLock<Result<OcrClient, crate::transport::Error>> = 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<LiteLLMOcrResponse, Error> {
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
shared_client()?.perform(request).await
}
@ -123,7 +132,7 @@ pub async fn read_json_response<T: DeserializeOwned>(
response: reqwest::Response,
native: bool,
max_response_bytes: usize,
) -> Result<DecodedOcrResponse<T>, Error> {
) -> Result<DecodedOcrResponse<T>, 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<T: DeserializeOwned>(
pub(crate) async fn read_response_bytes(
mut response: reqwest::Response,
max_response_bytes: usize,
) -> Result<Bytes, Error> {
) -> Result<Bytes, crate::ocr::Error> {
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();
}

View file

@ -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<OcrDocument, Error> {
) -> Result<OcrDocument, crate::ocr::Error> {
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<Option<Self>, Error> {
pub(crate) fn parse(source: &'a str) -> Result<Option<Self>, 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<Vec<u8>, Error> {
pub(crate) fn decode(&self, max_bytes: usize) -> Result<Vec<u8>, 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<OcrDocument, Error> {
) -> Result<OcrDocument, crate::ocr::Error> {
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)
);
}
}

View file

@ -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<LiteLLMOcrResponse, Error> {
) -> Result<LiteLLMOcrResponse, super::Error> {
request.response_format()?;
let context = CallLifecycleContext::new(
"ocr",
@ -62,7 +61,7 @@ impl PreparedOcrCall {
pub(crate) async fn prepare(
client: OcrClient,
request: LiteLLMOcrRequest,
) -> Result<Self, Error> {
) -> Result<Self, super::Error> {
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<OcrProviderResponse, Error> {
pub(crate) async fn execute(self) -> Result<OcrProviderResponse, super::Error> {
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<Vec<(String, String)>, Error> {
fn request_headers(request: &reqwest::Request) -> Result<Vec<(String, String)>, super::Error> {
request
.headers()
.iter()
@ -195,7 +194,7 @@ pub(crate) struct OcrProviderResponse {
}
impl OcrProviderResponse {
pub(crate) fn normalize(self) -> Result<LiteLLMOcrResponse, Error> {
pub(crate) fn normalize(self) -> Result<LiteLLMOcrResponse, super::Error> {
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<dyn OcrHooks>, bytes: &[u8]) -> Result<(), Error> {
pub(crate) async fn post_call(hooks: &Arc<dyn OcrHooks>, 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 })

View file

@ -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<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
pub type OcrHookFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, super::Error>> + Send + 'a>>;
pub type OcrLogFuture<'a> = Pin<Box<dyn Future<Output = ()> + 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<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse>
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<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse
fn async_log_failure_event<'a>(
&'a self,
context: &'a CallLifecycleContext,
error: &'a Error,
error: &'a super::Error,
timing: &'a CallLifecycleTiming,
) -> Self::FailureFuture<'a> {
self.hooks.failure(context, error, timing)

View file

@ -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<T> = Result<NativeOutcome<T>, Error>;
pub type NativeResult<T> = Result<NativeOutcome<T>, super::Error>;
#[derive(Debug, PartialEq, Eq)]
pub enum NativeOutcome<T> {
@ -54,7 +52,7 @@ pub enum OcrHostOperation {
ProjectRequest,
Lifecycle(HostPhase),
ConstructResponse(Arc<LiteLLMOcrResponse>),
MapFailure(Error),
MapFailure(super::Error),
Success {
context: CallLifecycleContext,
response: Arc<LiteLLMOcrResponse>,
@ -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<LiteLLMOcrRequest>, bool), Error>),
Lifecycle(Result<(), HostFailure<Error>>),
AzureAdToken(Result<ResolvedCredential, AuthError>),
PreCall(Result<OcrPreCallRequest, Error>),
DuringCall(Result<OcrDuringCallRequest, Error>),
PostCall(Result<OcrPostCallRequest, Error>),
Request(Result<(Box<LiteLLMOcrRequest>, bool), super::Error>),
Lifecycle(Result<(), HostFailure<super::Error>>),
AzureAdToken(Result<ResolvedCredential, litellm_auth::Error>),
PreCall(Result<OcrPreCallRequest, super::Error>),
DuringCall(Result<OcrDuringCallRequest, super::Error>),
PostCall(Result<OcrPostCallRequest, super::Error>),
}
pub type OcrCallStep = HostCallStep<OcrHostOperation, LiteLLMOcrResponse>;
@ -97,7 +95,7 @@ pub struct OcrCall {
lifecycle: HostLifecycle,
execution: OcrExecution,
response: Option<Arc<LiteLLMOcrResponse>>,
error: Option<Error>,
error: Option<super::Error>,
pending: bool,
completed: bool,
projecting: bool,
@ -122,14 +120,17 @@ impl OcrCall {
})
}
pub async fn resume(&mut self, result: Option<OcrHostResult>) -> Result<OcrCallStep, Error> {
pub async fn resume(
&mut self,
result: Option<OcrHostResult>,
) -> Result<OcrCallStep, super::Error> {
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<Error>>) {
fn accept(&mut self, result: Result<(), HostFailure<super::Error>>) {
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<Error>) -> Result<OcrCallStep, Error> {
pub async fn interrupt(
&mut self,
failure: HostFailure<super::Error>,
) -> Result<OcrCallStep, super::Error> {
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<Error>,
failure: HostFailure<super::Error>,
) -> HostCallFuture<'_, Self::Operation, Self::Complete, Self::Error> {
Box::pin(OcrCall::interrupt(self, failure))
}
@ -317,7 +325,7 @@ struct OcrExecution {
operations_tx: mpsc::UnboundedSender<PendingOperation>,
operations_rx: mpsc::UnboundedReceiver<PendingOperation>,
pending_result: Option<oneshot::Sender<OcrHostResult>>,
execution: Option<tokio::task::JoinHandle<Result<LiteLLMOcrResponse, Error>>>,
execution: Option<tokio::task::JoinHandle<Result<LiteLLMOcrResponse, super::Error>>>,
completed: bool,
azure_ad_token_provider: bool,
terminal: Arc<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
@ -339,25 +347,28 @@ impl OcrExecution {
}
}
pub async fn resume(&mut self, result: Option<OcrHostResult>) -> Result<OcrCallStep, Error> {
pub async fn resume(
&mut self,
result: Option<OcrHostResult>,
) -> Result<OcrCallStep, super::Error> {
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<OcrHostResult, Error> {
async fn invoke(&self, operation: OcrHostOperation) -> Result<OcrHostResult, super::Error> {
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(),
)))
}

View file

@ -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<T: DeserializeOwned>(
request: &LiteLLMOcrRequest,
) -> Result<ParsedProviderParams<T>, Error> {
) -> Result<ParsedProviderParams<T>, 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<B>(
headers: &[(String, String)],
retains_document: bool,
body: B,
validate: impl Fn(&B) -> Result<(), Error>,
) -> Result<reqwest::Request, Error>
validate: impl Fn(&B) -> Result<(), super::Error>,
) -> Result<reqwest::Request, super::Error>
where
B: Serialize + DeserializeOwned,
{
@ -36,7 +35,7 @@ where
let composed = OcrWireBody::<B>::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<B: Serialize>(
url: &str,
headers: &[(String, String)],
body: &B,
) -> Result<reqwest::Request, Error> {
) -> Result<reqwest::Request, super::Error> {
let builder = client
.provider_http()
.post(url)
@ -88,14 +87,14 @@ pub(crate) fn build_http_request<B: Serialize>(
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<B> {
}
impl<B: Serialize + DeserializeOwned> OcrWireBody<B> {
fn decode(value: Value, prefix: &str) -> Result<Self, Error> {
fn decode(value: Value, prefix: &str) -> Result<Self, super::Error> {
let body: B = super::wire::decode_request_value(value.clone(), prefix)?;
let Value::Object(fields) = value else {
return Err(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

View file

@ -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))
}

View file

@ -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<Self, Error> {
) -> Result<Self, super::Error> {
let (model, config) = resolve_provider_config(&model, custom_llm_provider)?;
Ok(Self {

View file

@ -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<Vec<&'static str>, Error> {
) -> Result<Vec<&'static str>, 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<Vec<OptionalParamSpec>, Error> {
) -> Result<Vec<OptionalParamSpec>, 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<LiteLLMOcrRequest, Error> {
pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, crate::ocr::Error> {
let api_key_source = source_for(&wire.input_sources, "api_key");
let api_base_source = source_for(&wire.input_sources, "api_base");
let extra_headers_source = source_for(&wire.input_sources, "extra_headers");
@ -121,16 +120,18 @@ pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, Error>
.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::<Result<Vec<_>, Error>>()?;
.collect::<Result<Vec<_>, 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<LiteLLMOcrRequest, Error>
.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<LiteLLMOcrRequest, Error>
})
}
fn decode_document(value: Value) -> Result<OcrDocument, Error> {
fn decode_document(value: Value) -> Result<OcrDocument, crate::ocr::Error> {
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<String>) -> Option<String> {
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
}
pub fn decode_request_value<T: DeserializeOwned>(value: Value, prefix: &str) -> Result<T, Error> {
pub fn decode_request_value<T: DeserializeOwned>(
value: Value,
prefix: &str,
) -> Result<T, crate::ocr::Error> {
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<T: DeserializeOwned>(value: Value, prefix: &str) ->
pub fn decode_response<T: DeserializeOwned>(
bytes: &[u8],
native: bool,
) -> Result<DecodedOcrResponse<T>, Error> {
) -> Result<DecodedOcrResponse<T>, 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)
));
}
}

View file

@ -1,5 +1,4 @@
use super::*;
use crate::chat_completions::Error;
use serde_json::json;
fn messages(value: Value) -> Vec<ChatMessage> {
@ -20,7 +19,9 @@ fn transform(model: &str, msgs: Value, opts: Value) -> Value {
.body
}
fn transform_response(body: Value) -> Result<ChatCompletionsResponse, Error> {
fn transform_response(
body: Value,
) -> Result<ChatCompletionsResponse, crate::chat_completions::Error> {
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")
);
}

View file

@ -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<String>,
) -> Result<String, Error> {
) -> Result<String, crate::chat_completions::Error> {
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<String>,
) -> Result<ChatCompletionsAuth, Error> {
) -> Result<ChatCompletionsAuth, crate::chat_completions::Error> {
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<ChatMessage>,
optional_params: OpaqueParams,
) -> Result<ProviderChatRequestData, Error> {
) -> Result<ProviderChatRequestData, crate::chat_completions::Error> {
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<ChatCompletionsResponse, Error> {
let body = response
.body
.as_object()
.ok_or_else(|| Error::InvalidResponse("messages response is not an object".into()))?;
) -> Result<ChatCompletionsResponse, crate::chat_completions::Error> {
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,

View file

@ -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<String>,
) -> Result<String, Error> {
) -> Result<String, crate::messages::Error> {
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<String>,
) -> Result<String, Error> {
resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from)
) -> Result<String, crate::messages::Error> {
resolve_anthropic_api_key(api_key, env_lookup).map_err(crate::messages::Error::from)
}
fn auth_strategy(&self) -> MessagesAuthStrategy {

View file

@ -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<String>,
) -> Result<String, Error> {
) -> Result<String, crate::messages::Error> {
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<String>,
) -> Result<String, Error> {
) -> Result<String, crate::messages::Error> {
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<String>,
) -> Result<String, Error> {
) -> Result<String, crate::messages::Error> {
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<String>,
) -> Result<String, Error> {
) -> Result<String, crate::messages::Error> {
resolve_azure_api_key(api_key, env_lookup)
}
@ -168,7 +167,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
fn transform_request(
&self,
request: AnthropicMessagesRequest,
) -> Result<AnthropicMessagesRequest, Error> {
) -> Result<AnthropicMessagesRequest, crate::messages::Error> {
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<AnthropicMessagesResponse, Error> {
) -> Result<AnthropicMessagesResponse, crate::messages::Error> {
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]

View file

@ -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<AudioTranscriptionRequestData, Error> {
) -> Result<AudioTranscriptionRequestData, crate::audio_transcription::Error> {
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<AudioTranscriptionResponseData, Error> {
) -> Result<AudioTranscriptionResponseData, crate::audio_transcription::Error> {
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<String>,
) -> Result<String, Error> {
) -> Result<String, crate::audio_transcription::Error> {
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<String>,
) -> Result<AudioTranscriptionAuth, Error> {
) -> Result<AudioTranscriptionAuth, crate::audio_transcription::Error> {
let (_, model_region) = bedrock_model_id_and_region(model);
Ok(AudioTranscriptionAuth::AwsSigV4 {
region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup),

View file

@ -1,5 +1,4 @@
use super::*;
use crate::chat_completions::Error;
use serde_json::json;
fn messages(value: Value) -> Vec<ChatMessage> {
@ -24,7 +23,9 @@ fn transform(msgs: Value, opts: Value) -> Value {
.body
}
fn transform_response(body: Value) -> Result<ChatCompletionsResponse, Error> {
fn transform_response(
body: Value,
) -> Result<ChatCompletionsResponse, crate::chat_completions::Error> {
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")
);
}

View file

@ -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<String>,
) -> Result<String, Error> {
) -> Result<String, crate::chat_completions::Error> {
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<String>,
) -> Result<ChatCompletionsAuth, Error> {
) -> Result<ChatCompletionsAuth, crate::chat_completions::Error> {
// 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<ChatMessage>,
optional_params: OpaqueParams,
) -> Result<ProviderChatRequestData, Error> {
) -> Result<ProviderChatRequestData, crate::chat_completions::Error> {
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<ChatCompletionsResponse, Error> {
let body = response
.body
.as_object()
.ok_or_else(|| Error::InvalidResponse("converse response is not an object".into()))?;
) -> Result<ChatCompletionsResponse, crate::chat_completions::Error> {
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"),

View file

@ -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<N> {
let value = String::deserialize(deserializer)?;
RoutedModel::new(&value)
.and_then(RoutedModel::into_provider::<N>)
.map_err(D::Error::custom)
.map_err(<D::Error as serde::de::Error>::custom)
}
}

View file

@ -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<ResponsesWsTransformResult, Error> {
) -> Result<ResponsesWsTransformResult, crate::responses::Error> {
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<ResponsesWsTransformResult, Error> {
) -> Result<ResponsesWsTransformResult, crate::responses::Error> {
Ok(ResponsesWsTransformResult::passthrough(event.clone()))
}
}

View file

@ -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<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
type LifecycleFuture<'a, T> =
Pin<Box<dyn Future<Output = Result<T, crate::responses::Error>> + 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<Box<dyn Future<Output = ()> + 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;

View file

@ -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<ResponsesWsTransformResult, Error>;
) -> Result<ResponsesWsTransformResult, super::Error>;
fn transform_ws_response(
&self,
event: &ResponsesWsEvent,
model: &str,
) -> Result<ResponsesWsTransformResult, Error>;
) -> Result<ResponsesWsTransformResult, super::Error>;
}
pub fn complete_websocket_url(
@ -203,22 +202,22 @@ impl ResponsesWebSocketConnection {
url: &str,
headers: &HashMap<String, String>,
timeout: Option<Duration>,
) -> Result<Self, Error> {
) -> Result<Self, super::Error> {
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::<HeaderName>()
.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<Option<String>, Error> {
pub async fn recv_text(&self) -> Result<Option<String>, 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;

View file

@ -28,7 +28,6 @@ impl From<reqwest::Error> 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(_)
));
}
}

View file

@ -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),

View file

@ -1,7 +1,6 @@
use crate::call_lifecycle::host::{HostFailure, HostLifecycle, HostPhase};
use crate::ocr::Error;
fn run(fail_at: Option<HostPhase>, asynchronous: bool) -> (Vec<HostPhase>, Vec<Error>) {
fn run(fail_at: Option<HostPhase>, asynchronous: bool) -> (Vec<HostPhase>, Vec<crate::ocr::Error>) {
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<HostPhase>, asynchronous: bool) -> (Vec<HostPhase>, Vec<E
let phase = lifecycle.phase();
events.push(phase);
let result = if Some(phase) == fail_at {
Err(HostFailure::Error(Error::InvalidRequest(
Err(HostFailure::Error(crate::ocr::Error::InvalidRequest(
"selected failure".into(),
)))
} else {
@ -81,14 +80,14 @@ fn only_provider_and_response_construction_failures_use_provider_mapping() {
fn failure_handler_errors_do_not_replace_selected_failure_or_suppress_async_dispatch() {
let mut lifecycle = HostLifecycle::new(true);
while lifecycle.phase() != HostPhase::Execute {
lifecycle.accept::<Error>(Ok(()));
lifecycle.accept::<crate::ocr::Error>(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::<Error>(Ok(()));
lifecycle.accept::<crate::ocr::Error>(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)

View file

@ -686,7 +686,7 @@ async fn missing_host_result_preserves_pending_operation() {
async fn read_bounded_response(
response: Vec<u8>,
limit: usize,
) -> Result<bytes::Bytes, super::Error> {
) -> Result<bytes::Bytes, crate::ocr::Error> {
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,