mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat(ocr): add core foundation and Mistral adapter (#40530)
* feat(ocr): add core foundation and transport primitives * fix(ocr): decline missing Mistral credentials * fix(rust): compile trace parity on Rust 1.98 * refactor(ocr): define native response capability * refactor(auth): generalize missing API key errors * refactor(core): keep URL helpers usage scoped * refactor(ocr): support native responses across adapters * refactor(ocr): preserve unmapped provider params * refactor(ocr): distinguish request preparation from payload transforms * refactor(ocr): trace payload transformation at codec boundary
This commit is contained in:
parent
692a311efb
commit
b0d66a15b8
35 changed files with 1904 additions and 50 deletions
2
litellm-rust/Cargo.lock
generated
2
litellm-rust/Cargo.lock
generated
|
|
@ -1453,11 +1453,13 @@ dependencies = [
|
|||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"serde_path_to_error",
|
||||
"sha2 0.10.9",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"]
|
|||
tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
|
||||
futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
|
||||
base64 = "0.22"
|
||||
url = "2.5.8"
|
||||
|
||||
[profile.release]
|
||||
opt-level = 3
|
||||
|
|
|
|||
|
|
@ -269,7 +269,7 @@ fn guardrail_error_to_core_error(error: GuardrailError) -> Error {
|
|||
|
||||
fn core_error_kind(error: &Error) -> &'static str {
|
||||
match error {
|
||||
Error::Auth(_) => "AuthError",
|
||||
Error::Auth(_) | Error::MissingApiKey { .. } => "AuthError",
|
||||
Error::InvalidProvider(_) => "InvalidProvider",
|
||||
Error::InvalidRequest(_) => "InvalidRequest",
|
||||
Error::InvalidType { .. } => "InvalidType",
|
||||
|
|
|
|||
|
|
@ -386,7 +386,7 @@ fn guardrail_error_to_core_error(error: GuardrailError) -> Error {
|
|||
|
||||
fn core_error_kind(error: &Error) -> &'static str {
|
||||
match error {
|
||||
Error::Auth(_) => "AuthError",
|
||||
Error::Auth(_) | Error::MissingApiKey { .. } => "AuthError",
|
||||
Error::InvalidProvider(_) => "InvalidProvider",
|
||||
Error::InvalidRequest(_) => "InvalidRequest",
|
||||
Error::InvalidType { .. } => "InvalidType",
|
||||
|
|
|
|||
|
|
@ -114,7 +114,8 @@ impl IntoResponse for MessagesRouteError {
|
|||
| Error::Connect(_)
|
||||
| Error::InvalidResponse(_)
|
||||
| Error::InvalidType { .. }
|
||||
| Error::MissingField(_) => (
|
||||
| Error::MissingField(_)
|
||||
| Error::MissingApiKey { .. } => (
|
||||
StatusCode::BAD_GATEWAY,
|
||||
"messages provider request failed".to_string(),
|
||||
),
|
||||
|
|
|
|||
|
|
@ -47,10 +47,10 @@ pub async fn messages_request(
|
|||
.header(CONTENT_TYPE, "application/json")
|
||||
.body(Body::from(body.to_string()))
|
||||
.map_err(|error| Error::InvalidRequest(error.to_string()))?;
|
||||
let response = routes::app(state)
|
||||
.oneshot(request)
|
||||
.await
|
||||
.map_err(|error| match error {})?;
|
||||
let response = match routes::app(state).oneshot(request).await {
|
||||
Ok(response) => response,
|
||||
Err(error) => match error {},
|
||||
};
|
||||
let status: StatusCode = response.status();
|
||||
let bytes = to_bytes(response.into_body(), usize::MAX)
|
||||
.await
|
||||
|
|
|
|||
|
|
@ -4,6 +4,11 @@ version = "0.1.0"
|
|||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
autotests = false
|
||||
|
||||
[[test]]
|
||||
name = "workspace_crate_allowlist"
|
||||
path = "tests/workspace_crate_allowlist.rs"
|
||||
|
||||
[dependencies]
|
||||
base64.workspace = true
|
||||
|
|
@ -11,10 +16,13 @@ rand.workspace = true
|
|||
reqwest.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
serde_path_to_error = "0.1"
|
||||
tokio.workspace = true
|
||||
thiserror.workspace = true
|
||||
tracing.workspace = true
|
||||
tracing-subscriber = { workspace = true, optional = true }
|
||||
sha2.workspace = true
|
||||
url.workspace = true
|
||||
aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true }
|
||||
aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"], optional = true }
|
||||
aws-sdk-sts = { version = "1.108.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true }
|
||||
|
|
|
|||
|
|
@ -43,3 +43,6 @@ pub const EMPTY_TEXT_PLACEHOLDER: &str =
|
|||
"[System: Empty message content sanitised to satisfy protocol]";
|
||||
|
||||
pub const FUNCTION_TRACE_TARGET: &str = "litellm::function_trace";
|
||||
pub(crate) const OCR_HTTP_TIMEOUT_SECS: u64 = 600;
|
||||
pub(crate) const OCR_CONNECT_TIMEOUT_SECS: u64 = 10;
|
||||
pub(crate) const MISTRAL_OCR_API_BASE: &str = "https://api.mistral.ai/v1";
|
||||
|
|
|
|||
|
|
@ -17,6 +17,10 @@ pub enum Error {
|
|||
InvalidRequest(String),
|
||||
#[error("{0}")]
|
||||
Auth(String),
|
||||
#[error(
|
||||
"Missing {provider} API Key - A call is being made to {provider} but no key is set either in the environment variables or via params"
|
||||
)]
|
||||
MissingApiKey { provider: &'static str },
|
||||
#[error("upstream request failed with status {status}: {body}")]
|
||||
Http { status: u16, body: String },
|
||||
#[error("upstream network error: {0}")]
|
||||
|
|
@ -36,6 +40,59 @@ pub enum Error {
|
|||
Unsupported(&'static str),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, ThisError, PartialEq, Eq)]
|
||||
pub enum TransportError {
|
||||
#[error("upstream request failed with status {status}: {body}")]
|
||||
Http { status: u16, body: String },
|
||||
#[error("upstream network error: {0}")]
|
||||
Network(String),
|
||||
#[error("could not reach the provider: {0}")]
|
||||
Connect(String),
|
||||
}
|
||||
|
||||
impl TransportError {
|
||||
pub fn from_reqwest_before_dispatch(error: reqwest::Error) -> Self {
|
||||
let before_dispatch = !error.is_timeout() && (error.is_connect() || error.is_builder());
|
||||
let message = error.without_url().to_string();
|
||||
if before_dispatch {
|
||||
Self::Connect(message)
|
||||
} else {
|
||||
Self::Network(message)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<reqwest::Error> for TransportError {
|
||||
fn from(error: reqwest::Error) -> Self {
|
||||
Self::Network(error.without_url().to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl From<crate::ocr::error::OcrRequestError> for Error {
|
||||
fn from(error: crate::ocr::error::OcrRequestError) -> Self {
|
||||
match error {
|
||||
crate::ocr::error::OcrRequestError::MissingField(field) => Self::MissingField(field),
|
||||
error => Self::InvalidRequest(error.to_string()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<crate::ocr::error::OcrResponseError> for Error {
|
||||
fn from(error: crate::ocr::error::OcrResponseError) -> Self {
|
||||
Self::InvalidResponse(error.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl From<TransportError> for Error {
|
||||
fn from(error: TransportError) -> Self {
|
||||
match error {
|
||||
TransportError::Http { status, body } => Self::Http { status, body },
|
||||
TransportError::Network(message) => Self::Network(message),
|
||||
TransportError::Connect(message) => Self::Connect(message),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn json_type_name(value: &serde_json::Value) -> &'static str {
|
||||
match value {
|
||||
serde_json::Value::Null => "null",
|
||||
|
|
@ -46,3 +103,53 @@ pub fn json_type_name(value: &serde_json::Value) -> &'static str {
|
|||
serde_json::Value::Object(_) => "object",
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod transport_tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn transport_errors_remove_urls_and_keep_dispatch_context() {
|
||||
let error = reqwest::Client::builder()
|
||||
.no_proxy()
|
||||
.build()
|
||||
.expect("client")
|
||||
.get("http://localhost:invalid/private?api_key=secret")
|
||||
.send()
|
||||
.await
|
||||
.expect_err("invalid port");
|
||||
let error = TransportError::from_reqwest_before_dispatch(error);
|
||||
assert!(matches!(error, TransportError::Connect(_)));
|
||||
assert!(!error.to_string().contains("secret"));
|
||||
assert!(!error.to_string().contains("private"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn request_timeout_is_not_safe_to_retry_as_a_connect_failure() {
|
||||
use std::time::Duration;
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
|
||||
.await
|
||||
.expect("bind");
|
||||
let address = listener.local_addr().expect("address");
|
||||
let request = reqwest::Client::builder()
|
||||
.no_proxy()
|
||||
.build()
|
||||
.expect("client")
|
||||
.get(format!("http://{address}"))
|
||||
.timeout(Duration::from_millis(200))
|
||||
.send();
|
||||
let (response, accepted) = tokio::join!(
|
||||
request,
|
||||
tokio::time::timeout(Duration::from_secs(2), listener.accept())
|
||||
);
|
||||
let _connection = accepted
|
||||
.expect("accept deadline")
|
||||
.expect("accepted connection");
|
||||
let error = response.expect_err("server does not respond");
|
||||
assert!(error.is_timeout());
|
||||
assert!(matches!(
|
||||
TransportError::from_reqwest_before_dispatch(error),
|
||||
TransportError::Network(_)
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,10 +1,43 @@
|
|||
//! Header and upstream-body helpers shared by every route module.
|
||||
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS;
|
||||
use crate::error::{Error, json_type_name};
|
||||
|
||||
#[allow(
|
||||
dead_code,
|
||||
reason = "used by the OCR architecture in the next stacked PR"
|
||||
)]
|
||||
pub(crate) enum HeaderPolicy<'a> {
|
||||
All,
|
||||
Only(&'a [&'a str]),
|
||||
Except(&'a [&'a str]),
|
||||
}
|
||||
|
||||
#[allow(
|
||||
dead_code,
|
||||
reason = "used by the OCR architecture in the next stacked PR"
|
||||
)]
|
||||
pub(crate) fn with_headers(
|
||||
builder: reqwest::RequestBuilder,
|
||||
headers: &[(String, String)],
|
||||
policy: HeaderPolicy<'_>,
|
||||
) -> reqwest::RequestBuilder {
|
||||
headers
|
||||
.iter()
|
||||
.filter(|(name, _)| match policy {
|
||||
HeaderPolicy::All => true,
|
||||
HeaderPolicy::Only(names) => names
|
||||
.iter()
|
||||
.any(|allowed| name.eq_ignore_ascii_case(allowed)),
|
||||
HeaderPolicy::Except(names) => !names
|
||||
.iter()
|
||||
.any(|excluded| name.eq_ignore_ascii_case(excluded)),
|
||||
})
|
||||
.fold(builder, |builder, (name, value)| {
|
||||
builder.header(name, value)
|
||||
})
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub async fn http_request(
|
||||
request: reqwest::RequestBuilder,
|
||||
|
|
@ -12,8 +45,6 @@ pub async fn http_request(
|
|||
request.send().await
|
||||
}
|
||||
|
||||
/// Bound an upstream error body before it crosses a host boundary, so provider
|
||||
/// bodies stay data-minimized.
|
||||
pub fn truncate_error_body(body: &str) -> String {
|
||||
if body.chars().count() <= UPSTREAM_ERROR_BODY_MAX_CHARS {
|
||||
return body.to_string();
|
||||
|
|
@ -61,11 +92,77 @@ pub fn has_bearer_auth(headers: &[(String, String)]) -> bool {
|
|||
})
|
||||
}
|
||||
|
||||
#[allow(
|
||||
dead_code,
|
||||
reason = "used by the OCR architecture in the next stacked PR"
|
||||
)]
|
||||
pub(crate) fn deserialize_optional_param<'de, D, T>(
|
||||
deserializer: D,
|
||||
) -> Result<Option<Option<T>>, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
T: serde::Deserialize<'de>,
|
||||
{
|
||||
<Option<T> as serde::Deserialize>::deserialize(deserializer).map(Some)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case(HeaderPolicy::All, true, true)]
|
||||
#[case(HeaderPolicy::Only(&["authorization"]), true, false)]
|
||||
#[case(HeaderPolicy::Except(&["authorization"]), false, true)]
|
||||
fn forwarding_policy_preserves_matching_headers_and_duplicates(
|
||||
#[case] policy: HeaderPolicy<'_>,
|
||||
#[case] auth: bool,
|
||||
#[case] trace: bool,
|
||||
) {
|
||||
let request = with_headers(
|
||||
reqwest::Client::new().get("https://example.com"),
|
||||
&[
|
||||
("AuThOrIzAtIoN".into(), "Bearer token".into()),
|
||||
("X-Trace".into(), "first".into()),
|
||||
("x-trace".into(), "second".into()),
|
||||
],
|
||||
policy,
|
||||
)
|
||||
.build()
|
||||
.unwrap();
|
||||
assert_eq!(request.headers().contains_key("authorization"), auth);
|
||||
let traces: Vec<_> = request.headers().get_all("x-trace").iter().collect();
|
||||
if trace {
|
||||
assert_eq!(traces, ["first", "second"]);
|
||||
} else {
|
||||
assert!(traces.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn multipart_policy_leaves_content_headers_to_reqwest() {
|
||||
let request = with_headers(
|
||||
reqwest::Client::new()
|
||||
.post("https://example.com")
|
||||
.multipart(reqwest::multipart::Form::new().text("file", "abc")),
|
||||
&[
|
||||
("Content-Type".into(), "application/json".into()),
|
||||
("CONTENT-LENGTH".into(), "0".into()),
|
||||
],
|
||||
HeaderPolicy::Except(&["content-type", "content-length"]),
|
||||
)
|
||||
.build()
|
||||
.unwrap();
|
||||
assert!(
|
||||
request.headers()["content-type"]
|
||||
.to_str()
|
||||
.unwrap()
|
||||
.starts_with("multipart/form-data; boundary=")
|
||||
);
|
||||
assert_ne!(request.headers()["content-length"], "0");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn truncate_leaves_short_bodies_untouched() {
|
||||
assert_eq!(truncate_error_body("short"), "short");
|
||||
|
|
|
|||
|
|
@ -14,5 +14,6 @@ pub mod realtime;
|
|||
pub mod responses;
|
||||
pub mod router;
|
||||
pub mod routing_utils;
|
||||
mod url_utils;
|
||||
|
||||
pub use error::Error;
|
||||
|
|
|
|||
147
litellm-rust/crates/core/src/ocr/adapters/mistral.rs
Normal file
147
litellm-rust/crates/core/src/ocr/adapters/mistral.rs
Normal file
|
|
@ -0,0 +1,147 @@
|
|||
use super::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::constants::MISTRAL_OCR_API_BASE;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{
|
||||
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
|
||||
};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY";
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(crate) struct MistralAdapter;
|
||||
|
||||
impl OcrAdapter for MistralAdapter {
|
||||
type ProviderResponse = MistralOcrResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::Mistral;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let ParsedProviderParams {
|
||||
known: params,
|
||||
extra_params: _extra_params,
|
||||
} = _prepare_ocr_request::<MistralOcrParams>(request)?;
|
||||
let headers = validate_environment(&request.connection, &credential_env)?;
|
||||
let url = get_complete_url(request.connection.api_base.as_deref())?;
|
||||
let body =
|
||||
mistral::transform_ocr_request(&request.model, request.document.clone(), ¶ms)?;
|
||||
transform_request_body(client, request, &url, &headers, body, |_| Ok(())).await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
mistral::transform_ocr_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result<String, OcrError> {
|
||||
let base = api_base
|
||||
.map(str::trim)
|
||||
.filter(|base| !base.is_empty())
|
||||
.unwrap_or(MISTRAL_OCR_API_BASE);
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v1", "ocr"]))
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
.into()
|
||||
})
|
||||
}
|
||||
|
||||
fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
) -> Result<Vec<(String, String)>, OcrError> {
|
||||
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
|
||||
return Ok(connection.extra_headers.clone());
|
||||
}
|
||||
let api_key = connection
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key| !key.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
|
||||
.ok_or(Error::MissingApiKey {
|
||||
provider: "Mistral",
|
||||
})?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn complete_url_defaults_and_dedupes_v1() {
|
||||
assert_eq!(
|
||||
get_complete_url(None).unwrap(),
|
||||
"https://api.mistral.ai/v1/ocr"
|
||||
);
|
||||
assert_eq!(
|
||||
get_complete_url(Some("https://example.com/v1?tenant=a")).unwrap(),
|
||||
"https://example.com/v1/ocr?tenant=a"
|
||||
);
|
||||
assert_eq!(
|
||||
get_complete_url(Some("https://example.com/v1/ocr?tenant=a")).unwrap(),
|
||||
"https://example.com/v1/ocr?tenant=a"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_prefers_explicit_key_then_environment() {
|
||||
let explicit = OcrConnection {
|
||||
api_key: Some("explicit".into()),
|
||||
..OcrConnection::default()
|
||||
};
|
||||
assert_eq!(
|
||||
validate_environment(&explicit, &|_| Some("environment".into())).unwrap()[0],
|
||||
("Authorization".into(), "Bearer explicit".into())
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
validate_environment(&OcrConnection::default(), &|_| Some("environment".into()))
|
||||
.unwrap()[0],
|
||||
("Authorization".into(), "Bearer environment".into())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_preserves_forwarded_authorization() {
|
||||
let connection = OcrConnection {
|
||||
extra_headers: vec![("authorization".into(), "Bearer forwarded".into())],
|
||||
..OcrConnection::default()
|
||||
};
|
||||
assert_eq!(
|
||||
validate_environment(&connection, &|_| None).unwrap(),
|
||||
connection.extra_headers
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn environment_rejects_missing_key() {
|
||||
assert!(matches!(
|
||||
validate_environment(&OcrConnection::default(), &|_| None),
|
||||
Err(OcrError::Public(Error::MissingApiKey {
|
||||
provider: "Mistral"
|
||||
}))
|
||||
));
|
||||
}
|
||||
}
|
||||
69
litellm-rust/crates/core/src/ocr/adapters/mod.rs
Normal file
69
litellm-rust/crates/core/src/ocr/adapters/mod.rs
Normal file
|
|
@ -0,0 +1,69 @@
|
|||
use std::future::Future;
|
||||
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
use super::OcrClient;
|
||||
use super::error::{OcrError, OcrResponseError};
|
||||
use super::registry::OcrProvider;
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrResponseFormat};
|
||||
use super::wire::DecodedOcrResponse;
|
||||
|
||||
mod mistral;
|
||||
|
||||
pub(crate) use mistral::MistralAdapter;
|
||||
|
||||
/// Converts a complete LiteLLM OCR call to provider HTTP and normalizes its response.
|
||||
pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
|
||||
/// Provider JSON schema; direct and Vertex Mistral share `MistralOcrResponse`.
|
||||
type ProviderResponse: DeserializeOwned + Send;
|
||||
|
||||
const PROVIDER: OcrProvider;
|
||||
|
||||
/// Prepares the complete provider HTTP request.
|
||||
/// `request` contains the model, document, connection, and unmapped caller options.
|
||||
/// `client` supplies reusable provider and document HTTP clients.
|
||||
/// Returns the complete HTTP request, whereas Python returns body data.
|
||||
fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> impl Future<Output = Result<reqwest::Request, OcrError>> + Send;
|
||||
|
||||
/// Python: `transform_ocr_response`.
|
||||
/// `request` supplies caller context, including the fallback model.
|
||||
/// `response` is the decoded provider payload; the output is the shared LiteLLM schema.
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError>;
|
||||
|
||||
/// Decodes provider HTTP; adapters may override this to poll asynchronous operations.
|
||||
/// Python performs that polling inside `async_transform_ocr_response`.
|
||||
/// `client` is reused for polling; `response` is the initial HTTP response.
|
||||
/// `url` and `headers` describe the submitted call; `request` supplies limits and format.
|
||||
fn read_response(
|
||||
&self,
|
||||
_client: &OcrClient,
|
||||
response: reqwest::Response,
|
||||
_url: &str,
|
||||
_headers: &[(String, String)],
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> impl Future<Output = Result<DecodedOcrResponse<Self::ProviderResponse>, OcrError>> + Send
|
||||
{
|
||||
let retain_native = request
|
||||
.response_format()
|
||||
.map(|format| format == OcrResponseFormat::Native);
|
||||
async move { super::client::read_json_response(response, retain_native?).await }
|
||||
}
|
||||
}
|
||||
|
||||
macro_rules! for_each_ocr_adapter {
|
||||
($callback:ident) => {
|
||||
$callback! {
|
||||
Mistral, $crate::ocr::adapters::MistralAdapter, $crate::ocr::adapters::MistralAdapter, Mistral;
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
pub(crate) use for_each_ocr_adapter;
|
||||
75
litellm-rust/crates/core/src/ocr/client.rs
Normal file
75
litellm-rust/crates/core/src/ocr/client.rs
Normal file
|
|
@ -0,0 +1,75 @@
|
|||
use std::sync::OnceLock;
|
||||
use std::time::Duration;
|
||||
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
use super::error::OcrError;
|
||||
use super::handler::perform_ocr_request;
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
use super::wire::{DecodedOcrResponse, decode_response};
|
||||
use crate::Error;
|
||||
use crate::constants::OCR_CONNECT_TIMEOUT_SECS;
|
||||
use crate::error::TransportError;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct OcrClient {
|
||||
provider_http: reqwest::Client,
|
||||
}
|
||||
|
||||
impl OcrClient {
|
||||
pub fn new(provider_http: reqwest::Client) -> Result<Self, TransportError> {
|
||||
Ok(Self { provider_http })
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "ocr",
|
||||
target = "litellm::function_trace",
|
||||
level = "trace",
|
||||
skip_all
|
||||
)]
|
||||
pub async fn perform(&self, request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
|
||||
perform_ocr_request(self, request).await
|
||||
}
|
||||
|
||||
pub(crate) fn provider_http(&self) -> &reqwest::Client {
|
||||
&self.provider_http
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
pub(crate) fn for_test(provider_http: reqwest::Client) -> Self {
|
||||
Self { provider_http }
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
|
||||
static CLIENT: OnceLock<Result<OcrClient, TransportError>> = 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)
|
||||
.and_then(OcrClient::new)
|
||||
})
|
||||
.clone()?;
|
||||
client.perform(request).await
|
||||
}
|
||||
|
||||
pub async fn read_json_response<T: DeserializeOwned>(
|
||||
response: reqwest::Response,
|
||||
native: bool,
|
||||
) -> Result<DecodedOcrResponse<T>, OcrError> {
|
||||
let status = response.status();
|
||||
let bytes = response
|
||||
.bytes()
|
||||
.await
|
||||
.map_err(crate::error::TransportError::from)?;
|
||||
if !status.is_success() {
|
||||
return Err(crate::error::TransportError::Http {
|
||||
status: status.as_u16(),
|
||||
body: crate::http_utils::truncate_error_body(&String::from_utf8_lossy(&bytes)),
|
||||
}
|
||||
.into());
|
||||
}
|
||||
Ok(decode_response(&bytes, native)?)
|
||||
}
|
||||
5
litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs
Normal file
5
litellm-rust/crates/core/src/ocr/codecs/mistral/mod.rs
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
mod transformation;
|
||||
mod types;
|
||||
|
||||
pub(crate) use transformation::{transform_ocr_request, transform_ocr_response};
|
||||
pub(crate) use types::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse};
|
||||
|
|
@ -0,0 +1,102 @@
|
|||
use super::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse};
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(crate) fn transform_ocr_request(
|
||||
model: &str,
|
||||
document: OcrDocument,
|
||||
params: &MistralOcrParams,
|
||||
) -> Result<MistralOcrRequest, OcrRequestError> {
|
||||
Ok(MistralOcrRequest {
|
||||
model: model.to_string(),
|
||||
document,
|
||||
params: params.clone(),
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_ocr_response(
|
||||
model: &str,
|
||||
response: MistralOcrResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages: response.pages,
|
||||
model: response.model.unwrap_or_else(|| model.to_string()),
|
||||
document_annotation: response.document_annotation,
|
||||
usage_info: response.usage_info,
|
||||
object: "ocr".to_string(),
|
||||
extra_fields: response.extra_fields,
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
#[rstest]
|
||||
#[case("pages", json!([0, 2]))]
|
||||
#[case("include_image_base64", json!(true))]
|
||||
#[case("image_limit", json!(2))]
|
||||
#[case("image_min_size", json!(100))]
|
||||
#[case("bbox_annotation_format", json!({"type":"json_schema"}))]
|
||||
#[case("document_annotation_format", json!({"type":"json_schema"}))]
|
||||
#[case("document_annotation_prompt", json!("extract"))]
|
||||
#[case("extract_header", json!(true))]
|
||||
#[case("extract_footer", json!(false))]
|
||||
#[case("table_format", json!("html"))]
|
||||
#[case("confidence_scores_granularity", json!("word"))]
|
||||
#[case("include_blocks", json!(true))]
|
||||
#[case("id", json!("req-123"))]
|
||||
fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
|
||||
let params: MistralOcrParams =
|
||||
serde_json::from_value(json!({name: value.clone()})).unwrap();
|
||||
let document: OcrDocument = serde_json::from_value(
|
||||
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
|
||||
)
|
||||
.unwrap();
|
||||
let result =
|
||||
serde_json::to_value(transform_ocr_request("model", document, ¶ms).unwrap())
|
||||
.unwrap();
|
||||
assert_eq!(result["model"], "model");
|
||||
assert_eq!(result[name], value);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_mapping_filters_unknown_fields() {
|
||||
let params: MistralOcrParams = serde_json::from_value(json!({"unknown": true})).unwrap();
|
||||
let document: OcrDocument = serde_json::from_value(
|
||||
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
|
||||
)
|
||||
.unwrap();
|
||||
let result =
|
||||
serde_json::to_value(transform_ocr_request("model", document, ¶ms).unwrap())
|
||||
.unwrap();
|
||||
assert!(result.get("unknown").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_preserves_provider_fields() {
|
||||
let response: MistralOcrResponse = serde_json::from_value(json!({
|
||||
"pages":[{"index":0,"markdown":"hello","header":"head","confidence_scores":{"mean":0.99}}],
|
||||
"model":"returned-model",
|
||||
"usage_info":{"pages_processed":1,"future_counter":5},
|
||||
"future_response_field":"kept"
|
||||
}))
|
||||
.unwrap();
|
||||
let result = transform_ocr_response("model", response)
|
||||
.unwrap()
|
||||
.into_json();
|
||||
assert_eq!(result["pages"][0]["header"], "head");
|
||||
assert_eq!(result["usage_info"]["future_counter"], 5);
|
||||
assert_eq!(result["future_response_field"], "kept");
|
||||
assert_eq!(result["model"], "returned-model");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_rejects_null_pages() {
|
||||
assert!(serde_json::from_value::<MistralOcrResponse>(json!({"pages":null})).is_err());
|
||||
}
|
||||
}
|
||||
53
litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs
Normal file
53
litellm-rust/crates/core/src/ocr/codecs/mistral/types.rs
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::ocr::types::OcrDocument;
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
|
||||
pub(crate) struct MistralOcrParams {
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub pages: Option<Vec<i64>>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub include_image_base64: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub image_limit: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub image_min_size: Option<i64>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub bbox_annotation_format: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub document_annotation_format: Option<Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub document_annotation_prompt: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extract_header: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub extract_footer: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub table_format: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub confidence_scores_granularity: Option<String>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub include_blocks: Option<bool>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub id: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize)]
|
||||
pub(crate) struct MistralOcrRequest {
|
||||
pub model: String,
|
||||
pub document: OcrDocument,
|
||||
#[serde(flatten)]
|
||||
pub params: MistralOcrParams,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Deserialize)]
|
||||
pub(crate) struct MistralOcrResponse {
|
||||
#[serde(default)]
|
||||
pub pages: Vec<Value>,
|
||||
pub model: Option<String>,
|
||||
pub document_annotation: Option<Value>,
|
||||
pub usage_info: Option<Value>,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
}
|
||||
1
litellm-rust/crates/core/src/ocr/codecs/mod.rs
Normal file
1
litellm-rust/crates/core/src/ocr/codecs/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod mistral;
|
||||
42
litellm-rust/crates/core/src/ocr/error.rs
Normal file
42
litellm-rust/crates/core/src/ocr/error.rs
Normal file
|
|
@ -0,0 +1,42 @@
|
|||
use thiserror::Error;
|
||||
|
||||
use crate::error::TransportError;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Error)]
|
||||
pub enum OcrRequestError {
|
||||
#[error("Invalid `req_format`. Expected 'native' or 'litellm'.")]
|
||||
RequestFormat,
|
||||
#[error("invalid OCR request field: {path}")]
|
||||
RequestField { path: String },
|
||||
#[error("missing required field: {0}")]
|
||||
MissingField(&'static str),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Error)]
|
||||
pub enum OcrResponseError {
|
||||
#[error("invalid OCR response field: {path}")]
|
||||
ResponseField { path: String },
|
||||
}
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub enum OcrError {
|
||||
#[error("{0}")]
|
||||
Request(#[from] OcrRequestError),
|
||||
#[error("{0}")]
|
||||
Response(#[from] OcrResponseError),
|
||||
#[error("{0}")]
|
||||
Transport(#[from] TransportError),
|
||||
#[error("{0}")]
|
||||
Public(#[from] crate::Error),
|
||||
}
|
||||
|
||||
impl From<OcrError> for crate::Error {
|
||||
fn from(error: OcrError) -> Self {
|
||||
match error {
|
||||
OcrError::Request(error) => error.into(),
|
||||
OcrError::Response(error) => error.into(),
|
||||
OcrError::Transport(error) => error.into(),
|
||||
OcrError::Public(error) => error,
|
||||
}
|
||||
}
|
||||
}
|
||||
72
litellm-rust/crates/core/src/ocr/handler.rs
Normal file
72
litellm-rust/crates/core/src/ocr/handler.rs
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
use super::OcrClient;
|
||||
use super::adapters::OcrAdapter;
|
||||
use super::hooks::OcrLifecycleHooks;
|
||||
use super::registry::OcrAdapterKind;
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
use crate::Error;
|
||||
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
|
||||
|
||||
pub(crate) async fn perform_ocr_request(
|
||||
client: &OcrClient,
|
||||
request: LiteLLMOcrRequest,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let context = CallLifecycleContext::new(
|
||||
"ocr",
|
||||
request.model.clone(),
|
||||
request.adapter.provider().as_str(),
|
||||
request
|
||||
.litellm_call_id
|
||||
.clone()
|
||||
.unwrap_or_else(|| format!("ocr-{:032x}", rand::random::<u128>())),
|
||||
);
|
||||
let hooks = OcrLifecycleHooks {
|
||||
hooks: request.hooks.clone(),
|
||||
provider_name: context.custom_llm_provider.clone(),
|
||||
};
|
||||
CallLifecycle::default().run(context, request, &hooks, |request| async move {
|
||||
macro_rules! execute_selected_adapter {
|
||||
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
|
||||
match request.adapter {
|
||||
$( OcrAdapterKind::$variant => execute_ocr_provider_call(client, &$instance, request).await, )+
|
||||
}
|
||||
};
|
||||
}
|
||||
super::adapters::for_each_ocr_adapter!(execute_selected_adapter)
|
||||
}).await
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
async fn execute_ocr_provider_call<A: OcrAdapter>(
|
||||
client: &OcrClient,
|
||||
adapter: &A,
|
||||
request: LiteLLMOcrRequest,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let provider_request = adapter.prepare_request(&request, client).await?;
|
||||
let url = provider_request.url().to_string();
|
||||
let headers = provider_request
|
||||
.headers()
|
||||
.iter()
|
||||
.map(|(name, value)| {
|
||||
value
|
||||
.to_str()
|
||||
.map(|value| (name.to_string(), value.to_string()))
|
||||
.map_err(|_| super::error::OcrRequestError::RequestField {
|
||||
path: "headers".into(),
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
let response = crate::http_utils::http_request(reqwest::RequestBuilder::from_parts(
|
||||
client.provider_http().clone(),
|
||||
provider_request,
|
||||
))
|
||||
.await
|
||||
.map_err(crate::error::TransportError::from)?;
|
||||
let decoded = adapter
|
||||
.read_response(client, response, &url, &headers, &request)
|
||||
.await?;
|
||||
let response = adapter.transform_ocr_response(&request, decoded.data)?;
|
||||
Ok(LiteLLMOcrResponse {
|
||||
provider_native_response: decoded.native,
|
||||
..response
|
||||
})
|
||||
}
|
||||
146
litellm-rust/crates/core/src/ocr/hooks.rs
Normal file
146
litellm-rust/crates/core/src/ocr/hooks.rs
Normal file
|
|
@ -0,0 +1,146 @@
|
|||
use std::future::Future;
|
||||
use std::pin::Pin;
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument};
|
||||
use crate::Error;
|
||||
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
|
||||
use serde::Serialize;
|
||||
use serde_json::Value;
|
||||
|
||||
pub type OcrHookFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T, Error>> + Send + 'a>>;
|
||||
pub type OcrLogFuture<'a> = Pin<Box<dyn Future<Output = ()> + Send + 'a>>;
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct OcrPreCallRequest {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
pub document: OcrDocument,
|
||||
pub optional_params: Value,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct OcrDuringCallRequest {
|
||||
pub model: String,
|
||||
pub custom_llm_provider: String,
|
||||
pub url: String,
|
||||
pub body: Value,
|
||||
}
|
||||
|
||||
pub trait OcrHooks: Send + Sync {
|
||||
fn has_guardrails(&self) -> bool {
|
||||
false
|
||||
}
|
||||
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
|
||||
Box::pin(async move { Ok(request) })
|
||||
}
|
||||
fn during_call(
|
||||
&self,
|
||||
request: OcrDuringCallRequest,
|
||||
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
|
||||
Box::pin(async move { Ok(request) })
|
||||
}
|
||||
fn success<'a>(
|
||||
&'a self,
|
||||
_context: &'a CallLifecycleContext,
|
||||
_response: &'a LiteLLMOcrResponse,
|
||||
_timing: &'a CallLifecycleTiming,
|
||||
) -> OcrLogFuture<'a> {
|
||||
Box::pin(async {})
|
||||
}
|
||||
fn failure<'a>(
|
||||
&'a self,
|
||||
_context: &'a CallLifecycleContext,
|
||||
_error: &'a Error,
|
||||
_timing: &'a CallLifecycleTiming,
|
||||
) -> OcrLogFuture<'a> {
|
||||
Box::pin(async {})
|
||||
}
|
||||
}
|
||||
|
||||
pub struct NoopOcrHooks;
|
||||
impl OcrHooks for NoopOcrHooks {}
|
||||
|
||||
pub(crate) struct OcrLifecycleHooks {
|
||||
pub hooks: Arc<dyn OcrHooks>,
|
||||
pub provider_name: String,
|
||||
}
|
||||
|
||||
impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse>
|
||||
for OcrLifecycleHooks
|
||||
{
|
||||
type PreCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>;
|
||||
type DuringCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>;
|
||||
type SuccessFuture<'a> = OcrLogFuture<'a>;
|
||||
type FailureFuture<'a> = OcrLogFuture<'a>;
|
||||
|
||||
fn async_pre_call_hook<'a>(
|
||||
&'a self,
|
||||
_context: &'a CallLifecycleContext,
|
||||
request: LiteLLMOcrRequest,
|
||||
) -> Self::PreCallFuture<'a> {
|
||||
Box::pin(async move {
|
||||
if !self.hooks.has_guardrails() {
|
||||
return Ok(request);
|
||||
}
|
||||
let changed = self
|
||||
.hooks
|
||||
.pre_call(OcrPreCallRequest {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: self.provider_name.clone(),
|
||||
document: request.document,
|
||||
optional_params: Value::Object(request.optional_params),
|
||||
})
|
||||
.await?;
|
||||
let Value::Object(optional_params) = changed.optional_params else {
|
||||
return Err(super::error::OcrRequestError::RequestField {
|
||||
path: "guardrail.optional_params".into(),
|
||||
}
|
||||
.into());
|
||||
};
|
||||
Ok(LiteLLMOcrRequest {
|
||||
document: changed.document,
|
||||
optional_params,
|
||||
..request
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn async_during_call_hook<'a>(
|
||||
&'a self,
|
||||
_context: &'a CallLifecycleContext,
|
||||
request: LiteLLMOcrRequest,
|
||||
) -> Self::DuringCallFuture<'a> {
|
||||
Box::pin(async move { Ok(request) })
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "success_callback",
|
||||
target = "litellm::function_trace",
|
||||
level = "trace",
|
||||
skip_all
|
||||
)]
|
||||
fn async_log_success_event<'a>(
|
||||
&'a self,
|
||||
context: &'a CallLifecycleContext,
|
||||
response: &'a LiteLLMOcrResponse,
|
||||
timing: &'a CallLifecycleTiming,
|
||||
) -> Self::SuccessFuture<'a> {
|
||||
self.hooks.success(context, response, timing)
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "failure_callback",
|
||||
target = "litellm::function_trace",
|
||||
level = "trace",
|
||||
skip_all
|
||||
)]
|
||||
fn async_log_failure_event<'a>(
|
||||
&'a self,
|
||||
context: &'a CallLifecycleContext,
|
||||
error: &'a Error,
|
||||
timing: &'a CallLifecycleTiming,
|
||||
) -> Self::FailureFuture<'a> {
|
||||
self.hooks.failure(context, error, timing)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,2 +1,21 @@
|
|||
mod adapters;
|
||||
pub mod client;
|
||||
mod codecs;
|
||||
pub mod error;
|
||||
mod handler;
|
||||
pub mod hooks;
|
||||
mod prepare;
|
||||
mod registry;
|
||||
pub mod transformation;
|
||||
pub mod types;
|
||||
pub mod wire;
|
||||
|
||||
pub use client::{OcrClient, ocr};
|
||||
pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "../../tests/ocr/support.rs"]
|
||||
pub(crate) mod test_support;
|
||||
#[cfg(test)]
|
||||
#[path = "../../tests/ocr.rs"]
|
||||
pub(crate) mod tests;
|
||||
|
|
|
|||
142
litellm-rust/crates/core/src/ocr/prepare.rs
Normal file
142
litellm-rust/crates/core/src/ocr/prepare.rs
Normal file
|
|
@ -0,0 +1,142 @@
|
|||
use serde::{Deserialize, Serialize, de::DeserializeOwned};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::OcrClient;
|
||||
use super::error::{OcrError, OcrRequestError};
|
||||
use super::hooks::OcrDuringCallRequest;
|
||||
use super::types::LiteLLMOcrRequest;
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub(crate) struct ParsedProviderParams<T> {
|
||||
#[serde(flatten)]
|
||||
pub known: T,
|
||||
#[serde(default, flatten)]
|
||||
pub extra_params: Map<String, Value>,
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(crate) fn _prepare_ocr_request<T: DeserializeOwned>(
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> Result<ParsedProviderParams<T>, OcrRequestError> {
|
||||
super::wire::decode_request_value(
|
||||
Value::Object(request.optional_params.clone()),
|
||||
"optional_params",
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) async fn transform_request_body<B>(
|
||||
client: &OcrClient,
|
||||
request: &LiteLLMOcrRequest,
|
||||
url: &str,
|
||||
headers: &[(String, String)],
|
||||
body: B,
|
||||
validate: impl FnOnce(&B) -> Result<(), OcrRequestError>,
|
||||
) -> Result<reqwest::Request, OcrError>
|
||||
where
|
||||
B: Serialize + DeserializeOwned,
|
||||
{
|
||||
let body = if request.hooks.has_guardrails() {
|
||||
let changed = request
|
||||
.hooks
|
||||
.during_call(OcrDuringCallRequest {
|
||||
model: request.model.clone(),
|
||||
custom_llm_provider: request.adapter.provider().as_str().into(),
|
||||
url: url.into(),
|
||||
body: serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField {
|
||||
path: "body".into(),
|
||||
})?,
|
||||
})
|
||||
.await?;
|
||||
let body = OcrWireBody::<B>::decode(changed.body)?;
|
||||
validate(&body.body)?;
|
||||
body
|
||||
} else {
|
||||
OcrWireBody {
|
||||
body,
|
||||
extra: Map::new(),
|
||||
}
|
||||
};
|
||||
build_http_request(client, request, url, headers, &body)
|
||||
}
|
||||
|
||||
pub(crate) fn build_http_request<B: Serialize>(
|
||||
client: &OcrClient,
|
||||
request: &LiteLLMOcrRequest,
|
||||
url: &str,
|
||||
headers: &[(String, String)],
|
||||
body: &B,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let builder = client
|
||||
.provider_http()
|
||||
.post(url)
|
||||
.json(body)
|
||||
.timeout(request.connection.timeout);
|
||||
crate::http_utils::with_headers(builder, headers, crate::http_utils::HeaderPolicy::All)
|
||||
.build()
|
||||
.map_err(crate::error::TransportError::from)
|
||||
.map_err(OcrError::from)
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
struct OcrWireBody<B> {
|
||||
#[serde(flatten)]
|
||||
body: B,
|
||||
#[serde(flatten)]
|
||||
extra: Map<String, Value>,
|
||||
}
|
||||
|
||||
impl<B: Serialize + DeserializeOwned> OcrWireBody<B> {
|
||||
fn decode(value: Value) -> Result<Self, OcrRequestError> {
|
||||
let body: B = super::wire::decode_request_value(value.clone(), "guardrail.body")?;
|
||||
let Value::Object(fields) = value else {
|
||||
return Err(OcrRequestError::RequestField {
|
||||
path: "guardrail.body".into(),
|
||||
});
|
||||
};
|
||||
let known = serde_json::to_value(&body).map_err(|_| OcrRequestError::RequestField {
|
||||
path: "guardrail.body".into(),
|
||||
})?;
|
||||
let extra = fields
|
||||
.into_iter()
|
||||
.filter(|(key, _)| known.get(key).is_none())
|
||||
.collect();
|
||||
Ok(Self { body, extra })
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn credential_env(name: &str) -> Option<String> {
|
||||
std::env::var(name).ok()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[derive(Debug, Deserialize, PartialEq)]
|
||||
struct KnownParams {
|
||||
pages: Option<Vec<i64>>,
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parsed_provider_params_separates_known_and_extra_params() {
|
||||
let parsed: ParsedProviderParams<KnownParams> = super::super::wire::decode_request_value(
|
||||
json!({
|
||||
"pages": [0, 2],
|
||||
"future_ocr_option": true,
|
||||
"extra_body": {"provider_option": "value"}
|
||||
}),
|
||||
"optional_params",
|
||||
)
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(parsed.known.pages, Some(vec![0, 2]));
|
||||
assert_eq!(parsed.extra_params["future_ocr_option"], true);
|
||||
assert_eq!(
|
||||
parsed.extra_params["extra_body"],
|
||||
json!({"provider_option": "value"})
|
||||
);
|
||||
assert_eq!(parsed.extra_params.len(), 2);
|
||||
}
|
||||
}
|
||||
53
litellm-rust/crates/core/src/ocr/registry.rs
Normal file
53
litellm-rust/crates/core/src/ocr/registry.rs
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
use super::adapters::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
|
||||
macro_rules! define_adapter_types {
|
||||
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum OcrAdapterKind {
|
||||
$( $variant, )+
|
||||
}
|
||||
|
||||
impl OcrAdapterKind {
|
||||
pub(crate) const fn provider(self) -> OcrProvider {
|
||||
match self {
|
||||
$( Self::$variant => <$adapter>::PROVIDER, )+
|
||||
}
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
super::adapters::for_each_ocr_adapter!(define_adapter_types);
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum OcrProvider {
|
||||
Mistral,
|
||||
}
|
||||
|
||||
impl OcrProvider {
|
||||
pub(crate) const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Mistral => "mistral",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn resolve_wire_adapter(
|
||||
model: &str,
|
||||
custom_llm_provider: Option<&str>,
|
||||
) -> Result<(String, OcrAdapterKind), Error> {
|
||||
let provider =
|
||||
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: OcrProvider::Mistral.as_str(),
|
||||
});
|
||||
let typed_provider = match provider.custom_llm_provider {
|
||||
"mistral" => OcrProvider::Mistral,
|
||||
value => return Err(Error::InvalidProvider(value.to_string())),
|
||||
};
|
||||
match typed_provider {
|
||||
OcrProvider::Mistral => Ok((provider.model.to_string(), OcrAdapterKind::Mistral)),
|
||||
}
|
||||
}
|
||||
|
|
@ -1,7 +1,7 @@
|
|||
use crate::Error;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::types::{OcrRequestData, OcrResponseData};
|
||||
use super::types::{LiteLLMOcrResponse, OcrRequestData};
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum OcrAuthStrategy {
|
||||
|
|
@ -49,14 +49,14 @@ pub trait OcrProviderConfig: Sync {
|
|||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<OcrResponseData, Error>;
|
||||
) -> Result<LiteLLMOcrResponse, Error>;
|
||||
|
||||
fn transform_ocr_response_with_params(
|
||||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
_optional_params: &Map<String, Value>,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
self.transform_ocr_response(model, response_json)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,14 @@
|
|||
use std::sync::Arc;
|
||||
use std::time::Duration;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::hooks::{NoopOcrHooks, OcrHooks};
|
||||
use super::registry::{OcrAdapterKind, resolve_wire_adapter};
|
||||
use crate::Error;
|
||||
use crate::constants::OCR_HTTP_TIMEOUT_SECS;
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct OcrRequestData {
|
||||
pub data: Value,
|
||||
|
|
@ -8,31 +16,145 @@ pub struct OcrRequestData {
|
|||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct OcrResponseData {
|
||||
#[serde(tag = "type")]
|
||||
pub enum OcrDocument {
|
||||
#[serde(rename = "document_url")]
|
||||
DocumentUrl {
|
||||
document_url: String,
|
||||
#[serde(flatten)]
|
||||
extra_fields: Map<String, Value>,
|
||||
},
|
||||
#[serde(rename = "image_url")]
|
||||
ImageUrl {
|
||||
image_url: String,
|
||||
#[serde(flatten)]
|
||||
extra_fields: Map<String, Value>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum OcrResponseFormat {
|
||||
#[default]
|
||||
Litellm,
|
||||
Native,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct OcrConnection {
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub extra_headers: Vec<(String, String)>,
|
||||
pub timeout: Duration,
|
||||
}
|
||||
|
||||
impl Default for OcrConnection {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
extra_headers: Vec::new(),
|
||||
timeout: Duration::from_secs(OCR_HTTP_TIMEOUT_SECS),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub struct LiteLLMOcrRequest {
|
||||
pub model: String,
|
||||
pub document: OcrDocument,
|
||||
pub connection: OcrConnection,
|
||||
pub hooks: Arc<dyn OcrHooks>,
|
||||
pub litellm_call_id: Option<String>,
|
||||
pub optional_params: Map<String, Value>,
|
||||
pub(crate) adapter: OcrAdapterKind,
|
||||
}
|
||||
|
||||
impl LiteLLMOcrRequest {
|
||||
pub fn new(
|
||||
model: String,
|
||||
document: OcrDocument,
|
||||
custom_llm_provider: Option<&str>,
|
||||
optional_params: Map<String, Value>,
|
||||
) -> Result<Self, Error> {
|
||||
let (model, adapter_kind) = resolve_wire_adapter(&model, custom_llm_provider)?;
|
||||
|
||||
Ok(Self {
|
||||
model,
|
||||
document,
|
||||
connection: OcrConnection::default(),
|
||||
hooks: Arc::new(NoopOcrHooks),
|
||||
litellm_call_id: None,
|
||||
optional_params,
|
||||
adapter: adapter_kind,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn response_format(
|
||||
&self,
|
||||
) -> Result<OcrResponseFormat, super::error::OcrRequestError> {
|
||||
self.optional_params
|
||||
.get("req_format")
|
||||
.map(|value| {
|
||||
serde_json::from_value(value.clone())
|
||||
.map_err(|_| super::error::OcrRequestError::RequestFormat)
|
||||
})
|
||||
.transpose()
|
||||
.map(|format| format.unwrap_or_default())
|
||||
}
|
||||
|
||||
pub fn with_host_hooks(
|
||||
self,
|
||||
hooks: Arc<dyn OcrHooks>,
|
||||
litellm_call_id: Option<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
hooks,
|
||||
litellm_call_id,
|
||||
..self
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
|
||||
pub struct LiteLLMOcrResponse {
|
||||
pub pages: Vec<Value>,
|
||||
pub model: String,
|
||||
pub document_annotation: Option<Value>,
|
||||
pub usage_info: Option<Value>,
|
||||
pub object: String,
|
||||
#[serde(flatten)]
|
||||
pub extra_fields: Map<String, Value>,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub provider_native_response: Option<Value>,
|
||||
}
|
||||
|
||||
impl OcrResponseData {
|
||||
impl LiteLLMOcrResponse {
|
||||
pub fn into_json(self) -> Value {
|
||||
let mut response = serde_json::json!({
|
||||
"pages": self.pages,
|
||||
"model": self.model,
|
||||
"document_annotation": self.document_annotation,
|
||||
"usage_info": self.usage_info,
|
||||
"object": self.object,
|
||||
});
|
||||
if let Value::Object(object) = &mut response {
|
||||
object.extend(self.extra_fields);
|
||||
if let Some(native_response) = self.provider_native_response {
|
||||
object.insert("provider_native_response".to_string(), native_response);
|
||||
}
|
||||
}
|
||||
response
|
||||
serde_json::to_value(self).expect("OCR response fields are JSON-compatible")
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn response_serialization_flattens_extra_fields_and_omits_absent_native_response() {
|
||||
let response = LiteLLMOcrResponse {
|
||||
pages: vec![],
|
||||
model: "model".into(),
|
||||
document_annotation: None,
|
||||
usage_info: None,
|
||||
object: "ocr".into(),
|
||||
extra_fields: json!({"provider_field":"kept"})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
.clone(),
|
||||
provider_native_response: None,
|
||||
};
|
||||
let serialized = response.into_json();
|
||||
assert_eq!(serialized["provider_field"], "kept");
|
||||
assert!(serialized.get("provider_native_response").is_none());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
154
litellm-rust/crates/core/src/ocr/wire.rs
Normal file
154
litellm-rust/crates/core/src/ocr/wire.rs
Normal file
|
|
@ -0,0 +1,154 @@
|
|||
use crate::ocr::error::OcrRequestError;
|
||||
use crate::ocr::error::OcrResponseError;
|
||||
use std::time::Duration;
|
||||
|
||||
use super::hooks::{OcrDuringCallRequest, OcrPreCallRequest};
|
||||
use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument};
|
||||
use crate::Error;
|
||||
use serde::{
|
||||
Deserialize,
|
||||
de::{DeserializeOwned, IntoDeserializer},
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct DecodedOcrResponse<T> {
|
||||
pub data: T,
|
||||
pub native: Option<Value>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct OcrWireRequest {
|
||||
pub model: String,
|
||||
pub document: Value,
|
||||
pub api_key: Option<String>,
|
||||
pub api_base: Option<String>,
|
||||
pub custom_llm_provider: Option<String>,
|
||||
pub extra_headers: Option<Map<String, Value>>,
|
||||
#[serde(default)]
|
||||
pub optional_params: Map<String, Value>,
|
||||
pub timeout_seconds: Option<f64>,
|
||||
}
|
||||
|
||||
pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> bool {
|
||||
super::registry::resolve_wire_adapter(model, custom_llm_provider).is_ok()
|
||||
}
|
||||
|
||||
pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, Error> {
|
||||
let document = decode_request_value(wire.document, "document")?;
|
||||
let headers = wire
|
||||
.extra_headers
|
||||
.unwrap_or_default()
|
||||
.into_iter()
|
||||
.map(|(name, value)| {
|
||||
let value = value
|
||||
.as_str()
|
||||
.ok_or_else(|| OcrRequestError::RequestField {
|
||||
path: format!("extra_headers.{name}"),
|
||||
})?;
|
||||
Ok((name, value.to_string()))
|
||||
})
|
||||
.collect::<Result<Vec<_>, OcrRequestError>>()?;
|
||||
let timeout = wire
|
||||
.timeout_seconds
|
||||
.map(|seconds| {
|
||||
Duration::try_from_secs_f64(seconds).map_err(|_| OcrRequestError::RequestField {
|
||||
path: "timeout_seconds".into(),
|
||||
})
|
||||
})
|
||||
.transpose()?;
|
||||
let defaults = OcrConnection::default();
|
||||
let request = LiteLLMOcrRequest::new(
|
||||
wire.model,
|
||||
document,
|
||||
wire.custom_llm_provider.as_deref(),
|
||||
wire.optional_params,
|
||||
)?;
|
||||
let connection = OcrConnection {
|
||||
api_key: nonblank(wire.api_key),
|
||||
api_base: nonblank(wire.api_base),
|
||||
extra_headers: headers,
|
||||
timeout: timeout.unwrap_or(defaults.timeout),
|
||||
};
|
||||
Ok(LiteLLMOcrRequest {
|
||||
connection,
|
||||
..request
|
||||
})
|
||||
}
|
||||
|
||||
fn nonblank(value: Option<String>) -> Option<String> {
|
||||
value
|
||||
.map(|s| s.trim().to_string())
|
||||
.filter(|s| !s.is_empty())
|
||||
}
|
||||
pub fn decode_request_value<T: DeserializeOwned>(
|
||||
value: Value,
|
||||
prefix: &str,
|
||||
) -> Result<T, OcrRequestError> {
|
||||
serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| {
|
||||
OcrRequestError::RequestField {
|
||||
path: format!("{prefix}.{}", error.path()),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
pub fn decode_response<T: DeserializeOwned>(
|
||||
bytes: &[u8],
|
||||
native: bool,
|
||||
) -> Result<DecodedOcrResponse<T>, OcrResponseError> {
|
||||
let mut deserializer = serde_json::Deserializer::from_slice(bytes);
|
||||
let data = serde_path_to_error::deserialize(&mut deserializer).map_err(|error| {
|
||||
OcrResponseError::ResponseField {
|
||||
path: error.path().to_string(),
|
||||
}
|
||||
})?;
|
||||
deserializer
|
||||
.end()
|
||||
.map_err(|_| OcrResponseError::ResponseField {
|
||||
path: "response".into(),
|
||||
})?;
|
||||
let native = if native {
|
||||
Some(
|
||||
serde_json::from_slice(bytes).map_err(|_| OcrResponseError::ResponseField {
|
||||
path: "response".into(),
|
||||
})?,
|
||||
)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
Ok(DecodedOcrResponse { data, native })
|
||||
}
|
||||
|
||||
pub fn decode_pre_call_result(
|
||||
original: OcrPreCallRequest,
|
||||
value: Value,
|
||||
) -> Result<OcrPreCallRequest, OcrRequestError> {
|
||||
#[derive(Deserialize)]
|
||||
struct Changed {
|
||||
document: OcrDocument,
|
||||
#[serde(default)]
|
||||
optional_params: Map<String, Value>,
|
||||
}
|
||||
let changed: Changed = decode_request_value(value, "guardrail")?;
|
||||
Ok(OcrPreCallRequest {
|
||||
document: changed.document,
|
||||
optional_params: Value::Object(changed.optional_params),
|
||||
..original
|
||||
})
|
||||
}
|
||||
|
||||
pub fn decode_during_call_result(
|
||||
original: OcrDuringCallRequest,
|
||||
value: Value,
|
||||
) -> Result<OcrDuringCallRequest, OcrRequestError> {
|
||||
#[derive(Deserialize)]
|
||||
struct Changed {
|
||||
body: Value,
|
||||
}
|
||||
let changed: Changed = decode_request_value(value, "guardrail")?;
|
||||
Ok(OcrDuringCallRequest {
|
||||
body: changed.body,
|
||||
..original
|
||||
})
|
||||
}
|
||||
|
|
@ -2,7 +2,7 @@ use std::collections::BTreeSet;
|
|||
|
||||
use crate::error::{Error, json_type_name};
|
||||
use crate::ocr::transformation::{OcrAuthStrategy, OcrProviderConfig, OcrResponseHandling};
|
||||
use crate::ocr::types::{OcrRequestData, OcrResponseData};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG;
|
||||
|
|
@ -440,7 +440,7 @@ fn transform_document_intelligence_response(
|
|||
model: &str,
|
||||
response_json: Value,
|
||||
preserve_native_response: bool,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let response = response_json
|
||||
.as_object()
|
||||
.ok_or_else(|| Error::InvalidType {
|
||||
|
|
@ -488,7 +488,7 @@ fn transform_document_intelligence_response(
|
|||
})
|
||||
.collect();
|
||||
|
||||
Ok(OcrResponseData {
|
||||
Ok(LiteLLMOcrResponse {
|
||||
usage_info: Some(json!({
|
||||
"pages_processed": pages.len(),
|
||||
"doc_size_bytes": null,
|
||||
|
|
@ -521,7 +521,7 @@ impl OcrProviderConfig for AzureAiOcrConfig {
|
|||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json)
|
||||
}
|
||||
|
||||
|
|
@ -599,7 +599,7 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
|
|||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
transform_document_intelligence_response(model, response_json, false)
|
||||
}
|
||||
|
||||
|
|
@ -608,7 +608,7 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
|
|||
model: &str,
|
||||
response_json: Value,
|
||||
optional_params: &Map<String, Value>,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
transform_document_intelligence_response(
|
||||
model,
|
||||
response_json,
|
||||
|
|
@ -718,7 +718,7 @@ mod tests {
|
|||
})
|
||||
}
|
||||
|
||||
fn assert_native_fields_preserved(response: &OcrResponseData, operation: &Value) {
|
||||
fn assert_native_fields_preserved(response: &LiteLLMOcrResponse, operation: &Value) {
|
||||
let analyze_result = &operation["analyzeResult"];
|
||||
|
||||
assert_eq!(response.extra_fields["content"], analyze_result["content"]);
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
use crate::error::{Error, json_type_name};
|
||||
use crate::ocr::transformation::OcrProviderConfig;
|
||||
use crate::ocr::types::{OcrRequestData, OcrResponseData};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
const SUPPORTED_OCR_PARAMS: &[&str] = &[
|
||||
|
|
@ -107,7 +107,7 @@ impl OcrProviderConfig for MistralOcrConfig {
|
|||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let response_object = response_json
|
||||
.as_object()
|
||||
.ok_or_else(|| Error::InvalidType {
|
||||
|
|
@ -128,7 +128,7 @@ impl OcrProviderConfig for MistralOcrConfig {
|
|||
let document_annotation = response_object.get("document_annotation").cloned();
|
||||
let usage_info = response_object.get("usage_info").cloned();
|
||||
|
||||
Ok(OcrResponseData {
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages,
|
||||
model,
|
||||
document_annotation,
|
||||
|
|
@ -178,7 +178,10 @@ pub fn transform_ocr_request(
|
|||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub fn transform_ocr_response(model: &str, response_json: Value) -> Result<OcrResponseData, Error> {
|
||||
pub fn transform_ocr_response(
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ use serde_json::{Map, Value, json};
|
|||
|
||||
use crate::error::{Error, json_type_name};
|
||||
use crate::ocr::transformation::OcrProviderConfig;
|
||||
use crate::ocr::types::{OcrRequestData, OcrResponseData};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData};
|
||||
|
||||
pub const REDUCTO_API_BASE: &str = "https://platform.reducto.ai";
|
||||
pub const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY";
|
||||
|
|
@ -283,7 +283,7 @@ fn build_pages(result: &Map<String, Value>) -> Vec<Value> {
|
|||
pub fn transform_reducto_response(
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let response = response_json
|
||||
.as_object()
|
||||
.ok_or_else(|| Error::InvalidType {
|
||||
|
|
@ -311,7 +311,7 @@ pub fn transform_reducto_response(
|
|||
"credits": usage.get("credits").cloned().unwrap_or(Value::Null),
|
||||
}));
|
||||
|
||||
Ok(OcrResponseData {
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages: build_pages(result),
|
||||
model: model.to_string(),
|
||||
document_annotation: None,
|
||||
|
|
@ -341,7 +341,7 @@ impl OcrProviderConfig for ReductoParseV3Config {
|
|||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
transform_reducto_response(model, response_json)
|
||||
}
|
||||
|
||||
|
|
@ -383,7 +383,7 @@ impl OcrProviderConfig for ReductoParseLegacyConfig {
|
|||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
transform_reducto_response(model, response_json)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
use crate::error::{Error, json_type_name};
|
||||
use crate::ocr::transformation::OcrProviderConfig;
|
||||
use crate::ocr::types::{OcrRequestData, OcrResponseData};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG;
|
||||
|
|
@ -226,7 +226,7 @@ impl OcrProviderConfig for VertexAiOcrConfig {
|
|||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json)
|
||||
}
|
||||
|
||||
|
|
@ -301,7 +301,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig {
|
|||
&self,
|
||||
model: &str,
|
||||
response_json: Value,
|
||||
) -> Result<OcrResponseData, Error> {
|
||||
) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let response = response_json
|
||||
.as_object()
|
||||
.ok_or_else(|| Error::InvalidType {
|
||||
|
|
@ -339,7 +339,7 @@ impl OcrProviderConfig for VertexAiDeepSeekOcrConfig {
|
|||
.get("usage_info")
|
||||
.cloned()
|
||||
.or_else(|| response.get("usage").cloned());
|
||||
Ok(OcrResponseData {
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages,
|
||||
model: object
|
||||
.get("model")
|
||||
|
|
|
|||
95
litellm-rust/crates/core/src/url_utils.rs
Normal file
95
litellm-rust/crates/core/src/url_utils.rs
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
use std::marker::PhantomData;
|
||||
|
||||
use thiserror::Error;
|
||||
use url::Url;
|
||||
|
||||
#[derive(Debug, Error)]
|
||||
pub(crate) enum ApiUrlError {
|
||||
#[error("invalid URL: {0}")]
|
||||
Parse(#[from] url::ParseError),
|
||||
#[error("URL cannot be used as a base")]
|
||||
CannotBeBase,
|
||||
}
|
||||
|
||||
pub(crate) struct Base;
|
||||
pub(crate) struct Complete;
|
||||
|
||||
pub(crate) struct ApiUrl<State> {
|
||||
url: Url,
|
||||
state: PhantomData<State>,
|
||||
}
|
||||
|
||||
impl ApiUrl<Base> {
|
||||
pub(crate) fn parse(value: &str) -> Result<Self, ApiUrlError> {
|
||||
Ok(Self {
|
||||
url: Url::parse(value.trim())?,
|
||||
state: PhantomData,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn complete_path(
|
||||
mut self,
|
||||
target: &[&str],
|
||||
) -> Result<ApiUrl<Complete>, ApiUrlError> {
|
||||
let existing: Vec<String> = self
|
||||
.url
|
||||
.path_segments()
|
||||
.ok_or(ApiUrlError::CannotBeBase)?
|
||||
.filter(|segment| !segment.is_empty())
|
||||
.map(str::to_string)
|
||||
.collect();
|
||||
let overlap = (0..=existing.len().min(target.len()))
|
||||
.rev()
|
||||
.find(|&length| {
|
||||
existing[existing.len() - length..]
|
||||
.iter()
|
||||
.map(String::as_str)
|
||||
.eq(target[..length].iter().copied())
|
||||
})
|
||||
.unwrap_or(0);
|
||||
self.url
|
||||
.path_segments_mut()
|
||||
.map_err(|()| ApiUrlError::CannotBeBase)?
|
||||
.pop_if_empty()
|
||||
.extend(target[overlap..].iter().copied());
|
||||
Ok(ApiUrl {
|
||||
url: self.url,
|
||||
state: PhantomData,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl ApiUrl<Complete> {
|
||||
pub(crate) fn into_string(self) -> String {
|
||||
self.url.into()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn completion_appends_only_the_missing_path_suffix() {
|
||||
for (base, expected) in [
|
||||
("https://example.test", "https://example.test/v1/ocr"),
|
||||
("https://example.test/v1", "https://example.test/v1/ocr"),
|
||||
("https://example.test/v1/ocr", "https://example.test/v1/ocr"),
|
||||
] {
|
||||
let actual = ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v1", "ocr"]))
|
||||
.map(|url| url.into_string())
|
||||
.expect("url builds");
|
||||
assert_eq!(actual, expected);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn completion_places_paths_before_queries() {
|
||||
let actual = ApiUrl::parse("https://example.test/v1?tenant=a")
|
||||
.and_then(|url| url.complete_path(&["v1", "ocr"]))
|
||||
.map(|url| url.into_string())
|
||||
.expect("url builds");
|
||||
assert_eq!(actual, "https://example.test/v1/ocr?tenant=a");
|
||||
}
|
||||
}
|
||||
227
litellm-rust/crates/core/tests/ocr.rs
Normal file
227
litellm-rust/crates/core/tests/ocr.rs
Normal file
|
|
@ -0,0 +1,227 @@
|
|||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use super::OcrClient;
|
||||
use super::hooks::{OcrHookFuture, OcrHooks, OcrLogFuture, OcrPreCallRequest};
|
||||
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
|
||||
use super::wire::{OcrWireRequest, decode_request};
|
||||
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming};
|
||||
|
||||
#[test]
|
||||
fn request_boundary_selects_mistral_and_rejects_unknown_providers() {
|
||||
let request = OcrWireRequest {
|
||||
model: "mistral/model".into(),
|
||||
document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}),
|
||||
api_key: Some("key".into()),
|
||||
api_base: None,
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
optional_params: json!({"extract_header":true,"unknown":42})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
.clone(),
|
||||
timeout_seconds: None,
|
||||
};
|
||||
assert!(decode_request(request).is_ok());
|
||||
assert!(
|
||||
decode_request(OcrWireRequest {
|
||||
model: "model".into(),
|
||||
document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}),
|
||||
api_key: Some("key".into()),
|
||||
api_base: None,
|
||||
custom_llm_provider: Some("unknown".into()),
|
||||
extra_headers: None,
|
||||
optional_params: serde_json::Map::new(),
|
||||
timeout_seconds: None,
|
||||
})
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn facade_executes_direct_mistral_once() {
|
||||
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
|
||||
"pages":[{"index":0,"markdown":"hello","custom":"preserved"}],
|
||||
"usage_info":{"pages_processed":1}
|
||||
}))])
|
||||
.await;
|
||||
let result = perform_ocr(wire_request(
|
||||
"mistral/model",
|
||||
&base,
|
||||
json!({"extract_header":true,"unknown":"ignored"}),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
server.await.unwrap();
|
||||
assert_eq!(result.pages[0]["markdown"], "hello");
|
||||
assert_eq!(result.pages[0]["custom"], "preserved");
|
||||
let requests = seen.lock().unwrap();
|
||||
assert_eq!(requests.len(), 1);
|
||||
assert!(requests[0].starts_with("POST /v1/ocr "));
|
||||
assert!(
|
||||
requests[0]
|
||||
.to_ascii_lowercase()
|
||||
.contains("authorization: bearer test-key\r\n")
|
||||
);
|
||||
let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
|
||||
assert_eq!(
|
||||
body,
|
||||
json!({
|
||||
"model":"model",
|
||||
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
|
||||
"extract_header":true
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn facade_retains_native_response_when_requested() {
|
||||
let provider_response = json!({
|
||||
"pages":[{"index":0,"markdown":"hello"}],
|
||||
"usage_info":{"pages_processed":1},
|
||||
"provider_only":"preserved"
|
||||
});
|
||||
let (base, _, server) = mock_server(vec![MockResponse::json(provider_response.clone())]).await;
|
||||
let response = perform_ocr(wire_request(
|
||||
"mistral/model",
|
||||
&base,
|
||||
json!({"req_format":"native"}),
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
server.await.unwrap();
|
||||
assert_eq!(response.provider_native_response, Some(provider_response));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn facade_uses_the_injected_http_client() {
|
||||
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
|
||||
let mut default_headers = reqwest::header::HeaderMap::new();
|
||||
default_headers.insert(
|
||||
"x-transport-owner",
|
||||
reqwest::header::HeaderValue::from_static("host"),
|
||||
);
|
||||
let provider_http = reqwest::Client::builder()
|
||||
.default_headers(default_headers)
|
||||
.build()
|
||||
.unwrap();
|
||||
OcrClient::new(provider_http)
|
||||
.unwrap()
|
||||
.perform(wire_request("mistral/model", &base, json!({})))
|
||||
.await
|
||||
.unwrap();
|
||||
server.await.unwrap();
|
||||
assert!(seen.lock().unwrap()[0].contains("x-transport-owner: host"));
|
||||
}
|
||||
|
||||
struct RecordingHooks {
|
||||
events: Arc<Mutex<Vec<&'static str>>>,
|
||||
block: bool,
|
||||
}
|
||||
|
||||
impl OcrHooks for RecordingHooks {
|
||||
fn has_guardrails(&self) -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
|
||||
Box::pin(async move {
|
||||
self.events.lock().unwrap().push("pre");
|
||||
if self.block {
|
||||
return Err(crate::Error::InvalidRequest("blocked".into()));
|
||||
}
|
||||
Ok(request)
|
||||
})
|
||||
}
|
||||
|
||||
fn during_call(
|
||||
&self,
|
||||
request: super::hooks::OcrDuringCallRequest,
|
||||
) -> OcrHookFuture<'_, super::hooks::OcrDuringCallRequest> {
|
||||
Box::pin(async move {
|
||||
self.events.lock().unwrap().push("during");
|
||||
Ok(request)
|
||||
})
|
||||
}
|
||||
|
||||
fn success<'a>(
|
||||
&'a self,
|
||||
_context: &'a CallLifecycleContext,
|
||||
_response: &'a super::LiteLLMOcrResponse,
|
||||
_timing: &'a CallLifecycleTiming,
|
||||
) -> OcrLogFuture<'a> {
|
||||
Box::pin(async move {
|
||||
self.events.lock().unwrap().push("success");
|
||||
})
|
||||
}
|
||||
|
||||
fn failure<'a>(
|
||||
&'a self,
|
||||
_context: &'a CallLifecycleContext,
|
||||
_error: &'a crate::Error,
|
||||
_timing: &'a CallLifecycleTiming,
|
||||
) -> OcrLogFuture<'a> {
|
||||
Box::pin(async move {
|
||||
self.events.lock().unwrap().push("failure");
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn lifecycle_orders_hooks_and_emits_one_success() {
|
||||
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
|
||||
let events = Arc::new(Mutex::new(Vec::new()));
|
||||
let request = wire_request("mistral/model", &base, json!({}));
|
||||
let request = super::LiteLLMOcrRequest {
|
||||
hooks: Arc::new(RecordingHooks {
|
||||
events: events.clone(),
|
||||
block: false,
|
||||
}),
|
||||
..request
|
||||
};
|
||||
perform_ocr(request).await.unwrap();
|
||||
server.await.unwrap();
|
||||
assert_eq!(*events.lock().unwrap(), ["pre", "during", "success"]);
|
||||
assert_eq!(seen.lock().unwrap().len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn lifecycle_blocking_prevents_execution_and_emits_one_failure() {
|
||||
let events = Arc::new(Mutex::new(Vec::new()));
|
||||
let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({}));
|
||||
let request = super::LiteLLMOcrRequest {
|
||||
hooks: Arc::new(RecordingHooks {
|
||||
events: events.clone(),
|
||||
block: true,
|
||||
}),
|
||||
..request
|
||||
};
|
||||
let error = perform_ocr(request).await.unwrap_err();
|
||||
assert!(matches!(error, crate::Error::InvalidRequest(_)));
|
||||
assert_eq!(*events.lock().unwrap(), ["pre", "failure"]);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upstream_failure_emits_one_terminal_failure() {
|
||||
let (base, seen, server) = mock_server(vec![MockResponse {
|
||||
status: 500,
|
||||
headers: vec![],
|
||||
body: json!({"error":"failed"}),
|
||||
}])
|
||||
.await;
|
||||
let events = Arc::new(Mutex::new(Vec::new()));
|
||||
let request = wire_request("mistral/model", &base, json!({}));
|
||||
let request = super::LiteLLMOcrRequest {
|
||||
hooks: Arc::new(RecordingHooks {
|
||||
events: events.clone(),
|
||||
block: false,
|
||||
}),
|
||||
..request
|
||||
};
|
||||
assert!(perform_ocr(request).await.is_err());
|
||||
server.await.unwrap();
|
||||
assert_eq!(*events.lock().unwrap(), ["pre", "during", "failure"]);
|
||||
assert_eq!(seen.lock().unwrap().len(), 1);
|
||||
}
|
||||
106
litellm-rust/crates/core/tests/ocr/support.rs
Normal file
106
litellm-rust/crates/core/tests/ocr/support.rs
Normal file
|
|
@ -0,0 +1,106 @@
|
|||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use serde_json::{Value, json};
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
use crate::ocr::wire::{OcrWireRequest, decode_request};
|
||||
use crate::ocr::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient};
|
||||
|
||||
pub(crate) fn ocr_client() -> OcrClient {
|
||||
OcrClient::for_test(reqwest::Client::new())
|
||||
}
|
||||
|
||||
pub(crate) async fn perform_ocr(
|
||||
request: LiteLLMOcrRequest,
|
||||
) -> Result<LiteLLMOcrResponse, crate::Error> {
|
||||
ocr_client().perform(request).await
|
||||
}
|
||||
|
||||
pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
|
||||
decode_request(OcrWireRequest {
|
||||
model: model.into(),
|
||||
document: json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
|
||||
api_key: Some("test-key".into()),
|
||||
api_base: Some(base.into()),
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
optional_params: options.as_object().unwrap().clone(),
|
||||
timeout_seconds: Some(2.0),
|
||||
})
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
pub(crate) struct MockResponse {
|
||||
pub status: u16,
|
||||
pub headers: Vec<(&'static str, String)>,
|
||||
pub body: Value,
|
||||
}
|
||||
|
||||
impl MockResponse {
|
||||
pub fn json(body: Value) -> Self {
|
||||
Self {
|
||||
status: 200,
|
||||
headers: vec![],
|
||||
body,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn mock_server(
|
||||
responses: Vec<MockResponse>,
|
||||
) -> (String, Arc<Mutex<Vec<String>>>, tokio::task::JoinHandle<()>) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let base = format!("http://{}", listener.local_addr().unwrap());
|
||||
let requests = Arc::new(Mutex::new(Vec::new()));
|
||||
let seen = requests.clone();
|
||||
let server_base = base.clone();
|
||||
let task = tokio::spawn(async move {
|
||||
for response in responses {
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let mut bytes = Vec::new();
|
||||
let mut buffer = [0u8; 4096];
|
||||
let header_end = loop {
|
||||
let n = socket.read(&mut buffer).await.unwrap();
|
||||
assert!(n > 0);
|
||||
bytes.extend_from_slice(&buffer[..n]);
|
||||
if let Some(index) = bytes.windows(4).position(|s| s == b"\r\n\r\n") {
|
||||
break index + 4;
|
||||
}
|
||||
};
|
||||
let length = String::from_utf8_lossy(&bytes[..header_end])
|
||||
.lines()
|
||||
.find_map(|line| {
|
||||
let (name, value) = line.split_once(':')?;
|
||||
name.eq_ignore_ascii_case("content-length")
|
||||
.then(|| value.trim().parse::<usize>().unwrap())
|
||||
})
|
||||
.unwrap_or(0);
|
||||
while bytes.len() < header_end + length {
|
||||
let n = socket.read(&mut buffer).await.unwrap();
|
||||
assert!(n > 0);
|
||||
bytes.extend_from_slice(&buffer[..n]);
|
||||
}
|
||||
seen.lock()
|
||||
.unwrap()
|
||||
.push(String::from_utf8_lossy(&bytes).into_owned());
|
||||
let body = serde_json::to_vec(&response.body).unwrap();
|
||||
let headers = response
|
||||
.headers
|
||||
.into_iter()
|
||||
.map(|(name, value)| {
|
||||
format!("{name}: {}\r\n", value.replace("{base}", &server_base))
|
||||
})
|
||||
.collect::<String>();
|
||||
let head = format!(
|
||||
"HTTP/1.1 {} OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n{}\r\n",
|
||||
response.status,
|
||||
body.len(),
|
||||
headers
|
||||
);
|
||||
socket.write_all(head.as_bytes()).await.unwrap();
|
||||
socket.write_all(&body).await.unwrap();
|
||||
}
|
||||
});
|
||||
(base, requests, task)
|
||||
}
|
||||
|
|
@ -41,6 +41,7 @@ pub(crate) fn chat_completions_error_to_pyerr(err: Error) -> PyErr {
|
|||
| Error::InvalidRequest(_)
|
||||
| Error::InvalidType { .. }
|
||||
| Error::MissingField(_)
|
||||
| Error::MissingApiKey { .. }
|
||||
| Error::Routing(_)
|
||||
// Nothing reached the provider, so serving it on Python cannot double
|
||||
// bill and is the only way the caller gets an answer at all.
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue