mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
inlineing more errors
This commit is contained in:
parent
8cfb59082a
commit
03a1c4a938
50 changed files with 734 additions and 673 deletions
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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)?;
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
))
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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, .. })
|
||||
..
|
||||
})
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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, ¶ms)?;
|
||||
let body = transform_ocr_request(request.document.clone())?;
|
||||
transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await
|
||||
|
|
@ -522,7 +526,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
|
|||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: AzureDocumentIntelligenceOperation,
|
||||
) -> Result<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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(_))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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, ¶ms)?;
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: ReductoResponse,
|
||||
) -> Result<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, ¶ms)?;
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: ReductoResponse,
|
||||
) -> Result<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, ¶ms)?;
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: ReductoResponse,
|
||||
) -> Result<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, ¶ms)?;
|
||||
let body = merge_extra_params(&body, extra_params)?;
|
||||
build_http_request(client, request, &url, &headers, &body)
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: ReductoResponse,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}));
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 })
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
)))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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(_)
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue