mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
wip
This commit is contained in:
parent
3b9fc16f6b
commit
43a19d81ab
40 changed files with 1096 additions and 1122 deletions
3
.gitignore
vendored
3
.gitignore
vendored
|
|
@ -147,3 +147,6 @@ crash.*.log
|
|||
|
||||
ui/litellm-dashboard/out/
|
||||
litellm.log
|
||||
|
||||
.coverage-rust
|
||||
coverage-rust.xml
|
||||
|
|
|
|||
|
|
@ -2,6 +2,6 @@ litellm-core is the LiteLLM SDK in Rust — it makes the LLM call. Each top-leve
|
|||
|
||||
A route module owns everything the call needs: types, the provider template trait, provider transforms (under `providers/`), provider/auth/URL resolution, and the handler that performs the HTTP call. Handlers belong here, never in a host crate.
|
||||
|
||||
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or callback dispatch. Env reads are limited to credential fallback in a route's `prepare.rs`.
|
||||
Not here: serving HTTP (axum routes, extractors), config file reading, rollout state, databases, or host-specific callback execution. Core owns lifecycle sequencing and callback payload construction; hosts execute the selected integrations. Env reads are limited to credential fallback in a route's `prepare.rs`.
|
||||
|
||||
Routes (messages, ocr, realtime) and providers (anthropic, mistral, openai) are modules, not crates.
|
||||
|
|
|
|||
|
|
@ -1,66 +0,0 @@
|
|||
# CLAUDE.md
|
||||
|
||||
Rules for `litellm-rust/crates/core`.
|
||||
|
||||
## Responsibility
|
||||
|
||||
`core` is the LiteLLM SDK in Rust: it makes the LLM call. Every top-level
|
||||
LiteLLM call has a public entrypoint here, named after the route
|
||||
(`messages::messages()` is the Rust equivalent of `litellm.messages()`), and
|
||||
calling it returns a typed non-streaming response.
|
||||
|
||||
Allowed:
|
||||
- The public entrypoint for a route, plus its `<route>_stream` variant when the
|
||||
route supports streaming.
|
||||
- Provider resolution, auth header construction, URL building, and the provider
|
||||
HTTP call (shared reused client, connect + request timeouts).
|
||||
- Shared request/response structs.
|
||||
- Typed errors with stable, non-sensitive messages.
|
||||
- Deterministic validation helpers.
|
||||
- Serialization helpers that intentionally mirror Python output shape.
|
||||
- Route templates that match Python base config responsibilities, such as
|
||||
`messages::transformation::AnthropicMessagesProviderConfig`.
|
||||
|
||||
Not allowed:
|
||||
- Serving HTTP: axum routers, extractors, and other transport concerns.
|
||||
- Filesystem, database, or cache access.
|
||||
- Config file reading or rollout state; the host resolves those and passes them
|
||||
in. Env reads are limited to credential fallback in a route's `prepare.rs`.
|
||||
- Logging callbacks, tracing spans, spend writes, or customer callbacks.
|
||||
- Provider-specific branching that belongs in `providers`.
|
||||
- Panics for user/provider-controlled input.
|
||||
|
||||
## Typed Contracts (core rule)
|
||||
|
||||
Trait and function boundaries MUST be strongly typed. No stringly-typed JSON
|
||||
(`&str` / `String` / `Vec<String>` / bare `serde_json::Value`) as a transform
|
||||
input or output. Parse wire bytes into typed structs/enums at the host edge;
|
||||
`core` and `providers` operate only on those types (e.g. `RealtimeEvent`,
|
||||
`RealtimeTransformResult`, `OcrRequestData`). A `type`-style discriminator is a
|
||||
typed field on a struct, not a raw string threaded through the API.
|
||||
|
||||
## Structure
|
||||
|
||||
Use route names directly under `src/`: `messages`, `ocr`, future
|
||||
`chat_completions`, `embeddings`, and similar top-level LiteLLM calls. Do not
|
||||
invent broad names like `engine` for route contracts.
|
||||
|
||||
`src/messages` is the reference shape for a route module:
|
||||
|
||||
```
|
||||
mod.rs pub async fn messages(..) (+ messages_stream)
|
||||
types.rs request/response types
|
||||
transformation.rs the provider template trait
|
||||
prepare.rs provider resolution, auth headers, URL
|
||||
handler.rs the provider call
|
||||
client.rs the shared reqwest client
|
||||
```
|
||||
|
||||
## Parity Rules
|
||||
|
||||
- Every shared type used by a provider transform needs unit tests for
|
||||
serialization shape.
|
||||
- If Python parity requires always emitting a `null` field instead of omitting
|
||||
it, document that in code and pin it with a test.
|
||||
- Error enums should preserve enough detail for Python/HTTP hosts to map errors
|
||||
consistently without exposing document contents or upstream bodies.
|
||||
|
|
@ -63,3 +63,6 @@ pub(crate) const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY";
|
|||
pub(crate) const REDUCTO_ID_PREFIX: &str = "reducto://";
|
||||
pub(crate) const AZURE_AI_OCR_PATH: &str = "/providers/mistral/azure/ocr";
|
||||
pub(crate) const MISTRAL_OCR_API_BASE: &str = "https://api.mistral.ai/v1";
|
||||
|
||||
pub(crate) const COHERE_PARSE_API_BASE: &str = "https://api.cohere.com";
|
||||
pub(crate) const COHERE_API_KEY_ENV: &str = "COHERE_API_KEY";
|
||||
|
|
|
|||
|
|
@ -52,6 +52,15 @@ pub enum Error {
|
|||
Unsupported(&'static str),
|
||||
}
|
||||
|
||||
impl Error {
|
||||
pub const fn http_status_code(&self) -> Option<u16> {
|
||||
match self {
|
||||
Self::InvalidRequest(_) => Some(400),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, ThisError)]
|
||||
pub(crate) enum MediaError {
|
||||
#[error("media URL rejected by network policy")]
|
||||
|
|
|
|||
131
litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs
Normal file
131
litellm-rust/crates/core/src/ocr/adapters/azure/cohere.rs
Normal file
|
|
@ -0,0 +1,131 @@
|
|||
use super::super::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::cohere::{
|
||||
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
|
||||
};
|
||||
use crate::ocr::document::{inline_remote_document, validate_inline_document};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
use crate::providers::azure_ai::auth::AzureAuthInputs;
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
|
||||
|
||||
pub(crate) struct AzureCohereAdapter;
|
||||
|
||||
impl OcrAdapter for AzureCohereAdapter {
|
||||
type ProviderResponse = CohereResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::AzureAi;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let params = super::super::super::wire::decode_request_value::<CohereParams>(
|
||||
serde_json::Value::Object(request.optional_params.clone()),
|
||||
"optional_params",
|
||||
)?;
|
||||
let mut config = AzureAuthInputs::from_sourced_optional_params(
|
||||
&request.optional_params,
|
||||
&request.input_sources,
|
||||
)
|
||||
.map_err(Error::from)?;
|
||||
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
|
||||
let base = request
|
||||
.connection
|
||||
.api_base
|
||||
.clone()
|
||||
.or_else(|| credential_env(AZURE_AI_API_BASE_ENV))
|
||||
.filter(|base| !base.trim().is_empty())
|
||||
.ok_or_else(|| {
|
||||
Error::Auth(
|
||||
"Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(),
|
||||
)
|
||||
})?;
|
||||
let headers =
|
||||
super::validate_ai_environment(&request.connection, &config, &credential_env).await?;
|
||||
validate_document(&request.document)?;
|
||||
let remote = request.document.source().starts_with("http://")
|
||||
|| request.document.source().starts_with("https://");
|
||||
let document = inline_remote_document(
|
||||
client.document_fetcher(),
|
||||
request.document.clone(),
|
||||
&request.connection,
|
||||
)
|
||||
.await?;
|
||||
let body = transform_request(&request.model, document, params)?;
|
||||
transform_request_body(
|
||||
client,
|
||||
request,
|
||||
&complete_url(&base)?,
|
||||
&headers,
|
||||
!remote,
|
||||
body,
|
||||
|body| {
|
||||
validate_document(&body.document)?;
|
||||
validate_inline_document(&body.document)
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
transform_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
fn complete_url(base: &str) -> Result<String, OcrError> {
|
||||
let mut url = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
|
||||
if !matches!(url.scheme(), "http" | "https") {
|
||||
return Err(invalid_api_base().into());
|
||||
}
|
||||
let path = url.path().trim_end_matches('/').to_string();
|
||||
if path.ends_with("/v2/parse") {
|
||||
url.set_path(&path);
|
||||
return Ok(url.into());
|
||||
}
|
||||
url.set_path(path.strip_suffix("/models").unwrap_or(&path));
|
||||
ApiUrl::parse(url.as_str())
|
||||
.and_then(|url| url.complete_path(&["providers", "cohere", "v2", "parse"]))
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| invalid_api_base().into())
|
||||
}
|
||||
|
||||
fn invalid_api_base() -> OcrRequestError {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".into(),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn completes_foundry_urls_without_duplicate_paths_and_preserves_queries() {
|
||||
for suffix in [
|
||||
"",
|
||||
"/models",
|
||||
"/providers/cohere/v2",
|
||||
"/providers/cohere/v2/parse",
|
||||
] {
|
||||
assert_eq!(
|
||||
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
|
||||
"https://example.com/providers/cohere/v2/parse?tenant=a"
|
||||
);
|
||||
}
|
||||
assert_eq!(
|
||||
complete_url("https://example.com/v2/parse?tenant=a").unwrap(),
|
||||
"https://example.com/v2/parse?tenant=a"
|
||||
);
|
||||
assert!(complete_url("relative/path").is_err());
|
||||
}
|
||||
}
|
||||
|
|
@ -61,7 +61,7 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
|
|||
url: &str,
|
||||
headers: &[(String, String)],
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> Result<Vec<u8>, OcrError> {
|
||||
) -> Result<crate::ocr::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError> {
|
||||
polling::read_operation_response(
|
||||
client.polling_http(),
|
||||
response,
|
||||
|
|
@ -72,7 +72,6 @@ impl OcrAdapter for AzureDocumentIntelligenceAdapter {
|
|||
&request.hooks,
|
||||
)
|
||||
.await
|
||||
.map(|decoded| decoded.text.into_bytes())
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -44,7 +44,7 @@ pub(super) async fn read_operation_response(
|
|||
}
|
||||
let bytes = crate::ocr::client::read_response_bytes(response).await?;
|
||||
crate::ocr::handler::post_call(hooks, &bytes).await?;
|
||||
poll_operation(http_client, operation, headers, connection, native).await
|
||||
poll_operation(http_client, operation, headers, connection, native, hooks).await
|
||||
}
|
||||
|
||||
async fn poll_operation(
|
||||
|
|
@ -53,6 +53,7 @@ async fn poll_operation(
|
|||
headers: &[(String, String)],
|
||||
connection: &OcrConnection,
|
||||
native: bool,
|
||||
hooks: &Arc<dyn OcrHooks>,
|
||||
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
|
||||
let deadline = Instant::now()
|
||||
.checked_add(connection.poll_timeout)
|
||||
|
|
@ -88,7 +89,10 @@ async fn poll_operation(
|
|||
.await
|
||||
.map_err(|_| OcrPollingError::PollTimeout)??;
|
||||
match &decoded.data.status {
|
||||
Some(OperationStatus::Succeeded) => return Ok(decoded),
|
||||
Some(OperationStatus::Succeeded) => {
|
||||
crate::ocr::handler::post_call(hooks, decoded.text.as_bytes()).await?;
|
||||
return Ok(decoded);
|
||||
}
|
||||
Some(OperationStatus::Running | OperationStatus::NotStarted) => {
|
||||
tokio::time::timeout_at(deadline, tokio::time::sleep(Duration::from_secs(retry)))
|
||||
.await
|
||||
|
|
|
|||
|
|
@ -92,7 +92,7 @@ fn get_complete_url(
|
|||
})
|
||||
}
|
||||
|
||||
async fn validate_environment(
|
||||
pub(in crate::ocr::adapters) async fn validate_environment(
|
||||
connection: &OcrConnection,
|
||||
config: &AzureAuthInputs,
|
||||
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
mod cohere;
|
||||
mod document_intelligence;
|
||||
mod mistral;
|
||||
|
||||
|
|
@ -10,8 +11,10 @@ use crate::ocr::error::OcrError;
|
|||
use crate::ocr::types::OcrConnection;
|
||||
use crate::providers::azure_ai::auth::{AzureAuthInputs, AzureAuthService};
|
||||
|
||||
pub(crate) use cohere::AzureCohereAdapter;
|
||||
pub(crate) use document_intelligence::AzureDocumentIntelligenceAdapter;
|
||||
pub(crate) use mistral::AzureMistralAdapter;
|
||||
pub(super) use mistral::validate_environment as validate_ai_environment;
|
||||
|
||||
async fn resolve_entra(
|
||||
config: &AzureAuthInputs,
|
||||
|
|
@ -22,6 +25,10 @@ async fn resolve_entra(
|
|||
.get_or_init(AzureAuthService::default)
|
||||
.get_azure_ad_token(config, env_lookup)
|
||||
.await
|
||||
.or_else(|error| match error {
|
||||
crate::AuthError::EmptyAzureToken => Ok(None),
|
||||
other => Err(other),
|
||||
})
|
||||
.map(|credential| {
|
||||
credential.map(|credential| {
|
||||
let source = credential.source();
|
||||
|
|
|
|||
123
litellm-rust/crates/core/src/ocr/adapters/cohere.rs
Normal file
123
litellm-rust/crates/core/src/ocr/adapters/cohere.rs
Normal file
|
|
@ -0,0 +1,123 @@
|
|||
use super::OcrAdapter;
|
||||
use crate::Error;
|
||||
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
|
||||
use crate::ocr::OcrClient;
|
||||
use crate::ocr::codecs::cohere::{
|
||||
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
|
||||
};
|
||||
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::prepare::{credential_env, transform_request_body};
|
||||
use crate::ocr::registry::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
|
||||
use crate::url_utils::ApiUrl;
|
||||
|
||||
pub(crate) struct CohereAdapter;
|
||||
|
||||
impl OcrAdapter for CohereAdapter {
|
||||
type ProviderResponse = CohereResponse;
|
||||
const PROVIDER: OcrProvider = OcrProvider::Cohere;
|
||||
|
||||
async fn prepare_request(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
client: &OcrClient,
|
||||
) -> Result<reqwest::Request, OcrError> {
|
||||
let params = super::super::wire::decode_request_value::<CohereParams>(
|
||||
serde_json::Value::Object(request.optional_params.clone()),
|
||||
"optional_params",
|
||||
)?;
|
||||
let headers = validate_environment(&request.connection, &credential_env)?;
|
||||
let url = complete_url(
|
||||
request
|
||||
.connection
|
||||
.api_base
|
||||
.as_deref()
|
||||
.unwrap_or(COHERE_PARSE_API_BASE),
|
||||
)?;
|
||||
let body = transform_request(&request.model, request.document.clone(), params)?;
|
||||
transform_request_body(client, request, &url, &headers, true, body, |body| {
|
||||
validate_document(&body.document)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
||||
fn transform_ocr_response(
|
||||
&self,
|
||||
request: &LiteLLMOcrRequest,
|
||||
response: Self::ProviderResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
transform_response(&request.model, response)
|
||||
}
|
||||
}
|
||||
|
||||
fn complete_url(base: &str) -> Result<String, OcrError> {
|
||||
let parsed = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
|
||||
if !matches!(parsed.scheme(), "http" | "https") {
|
||||
return Err(invalid_api_base().into());
|
||||
}
|
||||
ApiUrl::parse(base)
|
||||
.and_then(|url| url.complete_path(&["v2", "parse"]))
|
||||
.map(|url| url.into_string())
|
||||
.map_err(|_| invalid_api_base().into())
|
||||
}
|
||||
|
||||
fn invalid_api_base() -> OcrRequestError {
|
||||
OcrRequestError::RequestField {
|
||||
path: "api_base".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 key = connection
|
||||
.api_key
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|key| !key.is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| env_lookup(COHERE_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
|
||||
.ok_or_else(|| {
|
||||
Error::Auth("Missing COHERE_API_KEY - set it in the environment or pass api_key".into())
|
||||
})?;
|
||||
Ok(
|
||||
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
|
||||
.chain(connection.extra_headers.clone())
|
||||
.collect(),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn completes_provider_urls_without_duplicate_paths_and_preserves_queries() {
|
||||
for suffix in ["", "/v2", "/v2/parse"] {
|
||||
assert_eq!(
|
||||
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
|
||||
"https://example.com/v2/parse?tenant=a"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_urls_and_blank_keys() {
|
||||
assert!(complete_url("relative/path").is_err());
|
||||
assert!(complete_url("ftp://example.com").is_err());
|
||||
assert!(matches!(
|
||||
validate_environment(
|
||||
&OcrConnection {
|
||||
api_key: Some(" ".into()),
|
||||
..Default::default()
|
||||
},
|
||||
&|_| None,
|
||||
),
|
||||
Err(OcrError::Public(Error::Auth(_)))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
@ -8,11 +8,13 @@ use super::registry::OcrProvider;
|
|||
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
|
||||
|
||||
mod azure;
|
||||
mod cohere;
|
||||
mod mistral;
|
||||
mod reducto;
|
||||
mod vertex;
|
||||
|
||||
pub(crate) use azure::{AzureDocumentIntelligenceAdapter, AzureMistralAdapter};
|
||||
pub(crate) use azure::{AzureCohereAdapter, AzureDocumentIntelligenceAdapter, AzureMistralAdapter};
|
||||
pub(crate) use cohere::CohereAdapter;
|
||||
pub(crate) use mistral::MistralAdapter;
|
||||
pub(crate) use reducto::{ReductoLegacyAdapter, ReductoV3Adapter};
|
||||
pub(crate) use vertex::{VertexDeepSeekAdapter, VertexMistralAdapter};
|
||||
|
|
@ -54,11 +56,16 @@ pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
|
|||
_url: &str,
|
||||
_headers: &[(String, String)],
|
||||
request: &LiteLLMOcrRequest,
|
||||
) -> impl Future<Output = Result<Vec<u8>, OcrError>> + Send {
|
||||
) -> impl Future<
|
||||
Output = Result<super::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError>,
|
||||
> + Send {
|
||||
async move {
|
||||
let bytes = super::client::read_response_bytes(response).await?;
|
||||
super::handler::post_call(&request.hooks, &bytes).await?;
|
||||
Ok(bytes)
|
||||
Ok(super::wire::decode_response(
|
||||
&bytes,
|
||||
request.response_format()? == super::types::OcrResponseFormat::Native,
|
||||
)?)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -66,6 +73,8 @@ pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
|
|||
macro_rules! for_each_ocr_adapter {
|
||||
($callback:ident) => {
|
||||
$callback! {
|
||||
Cohere, $crate::ocr::adapters::CohereAdapter, $crate::ocr::adapters::CohereAdapter, Cohere;
|
||||
AzureCohere, $crate::ocr::adapters::AzureCohereAdapter, $crate::ocr::adapters::AzureCohereAdapter, AzureAi;
|
||||
Mistral, $crate::ocr::adapters::MistralAdapter, $crate::ocr::adapters::MistralAdapter, Mistral;
|
||||
AzureMistral, $crate::ocr::adapters::AzureMistralAdapter, $crate::ocr::adapters::AzureMistralAdapter, AzureAi;
|
||||
AzureDocumentIntelligence, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, AzureAi;
|
||||
|
|
|
|||
235
litellm-rust/crates/core/src/ocr/codecs/cohere.rs
Normal file
235
litellm-rust/crates/core/src/ocr/codecs/cohere.rs
Normal file
|
|
@ -0,0 +1,235 @@
|
|||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::{Map, Value, json};
|
||||
|
||||
use crate::ocr::document::InlineDocument;
|
||||
use crate::ocr::error::{OcrRequestError, OcrResponseError};
|
||||
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub(crate) enum OutputFormat {
|
||||
#[default]
|
||||
Markdown,
|
||||
Blocks,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct CohereParams {
|
||||
#[serde(default)]
|
||||
pub output_format: OutputFormat,
|
||||
}
|
||||
|
||||
#[derive(Deserialize, Serialize)]
|
||||
pub(crate) struct CohereRequest {
|
||||
pub model: String,
|
||||
pub document: OcrDocument,
|
||||
pub output_format: OutputFormat,
|
||||
}
|
||||
|
||||
pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), OcrRequestError> {
|
||||
let OcrDocument::ImageUrl { image_url, .. } = document else {
|
||||
return Err(OcrRequestError::CohereImageOnly);
|
||||
};
|
||||
if image_url.is_empty() {
|
||||
return Err(OcrRequestError::CohereImageOnly);
|
||||
}
|
||||
if let Some(inline) = InlineDocument::parse(image_url)? {
|
||||
if !inline.mime_type().type_.eq_ignore_ascii_case("image") {
|
||||
return Err(OcrRequestError::CohereImageOnly);
|
||||
}
|
||||
inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
pub(crate) struct CohereResponse {
|
||||
#[serde(default)]
|
||||
pages: Vec<CoherePage>,
|
||||
meta: Option<CohereMeta>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CoherePage {
|
||||
index: Option<i64>,
|
||||
markdown: Option<CohereMarkdown>,
|
||||
blocks: Option<Vec<Map<String, Value>>>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CohereMarkdown {
|
||||
#[serde(default)]
|
||||
content: String,
|
||||
images: Option<Vec<Map<String, Value>>>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CohereMeta {
|
||||
billed_units: Option<CohereBilledUnits>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct CohereBilledUnits {
|
||||
pages: Option<i64>,
|
||||
}
|
||||
|
||||
pub(crate) fn transform_response(
|
||||
model: &str,
|
||||
response: CohereResponse,
|
||||
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
|
||||
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(|_| OcrResponseError::NumericRange("pages"))
|
||||
})?;
|
||||
let pages = response
|
||||
.pages
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(position, page)| {
|
||||
let index = page.index.map(Ok).unwrap_or_else(|| {
|
||||
i64::try_from(position).map_err(|_| OcrResponseError::NumericRange("page index"))
|
||||
})?;
|
||||
let (content, images) = page
|
||||
.markdown
|
||||
.map(|markdown| {
|
||||
let images =
|
||||
markdown
|
||||
.images
|
||||
.filter(|images| !images.is_empty())
|
||||
.map(|images| {
|
||||
images
|
||||
.into_iter()
|
||||
.map(|mut image| {
|
||||
if let Some(Value::Object(bbox)) =
|
||||
image.get("bounding_box").cloned()
|
||||
{
|
||||
image.insert("bbox".into(), Value::Object(bbox));
|
||||
}
|
||||
Value::Object(image)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
});
|
||||
(markdown.content, images)
|
||||
})
|
||||
.unwrap_or_default();
|
||||
let mut normalized = json!({"index": index, "markdown": content, "images": images});
|
||||
if let Some(blocks) = page.blocks {
|
||||
normalized["blocks"] = json!(blocks);
|
||||
}
|
||||
Ok(normalized)
|
||||
})
|
||||
.collect::<Result<Vec<_>, OcrResponseError>>()?;
|
||||
Ok(LiteLLMOcrResponse {
|
||||
pages,
|
||||
model: model.into(),
|
||||
document_annotation: None,
|
||||
usage_info: Some(json!({"pages_processed": pages_processed})),
|
||||
object: "ocr".into(),
|
||||
extra_fields: Map::new(),
|
||||
provider_native_response: None,
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn transform_request(
|
||||
model: &str,
|
||||
document: OcrDocument,
|
||||
params: CohereParams,
|
||||
) -> Result<CohereRequest, OcrRequestError> {
|
||||
validate_document(&document)?;
|
||||
Ok(CohereRequest {
|
||||
model: model.into(),
|
||||
document,
|
||||
output_format: params.output_format,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn response_normalizes_markdown_images_blocks_and_billed_pages() {
|
||||
let response = serde_json::from_value(json!({
|
||||
"pages": [
|
||||
{"index": 4, "markdown": {"content": "receipt", "images": [{"id":"image", "bounding_box":{"top_left_x":1}, "description":"scan"}]}},
|
||||
{"blocks": [{"type":"text","text":"total"}]}
|
||||
],
|
||||
"meta": {"billed_units":{"pages":3}}
|
||||
})).unwrap();
|
||||
let normalized = transform_response("parse-v5.0", response).unwrap();
|
||||
assert_eq!(normalized.pages[0]["index"], 4);
|
||||
assert_eq!(normalized.pages[0]["markdown"], "receipt");
|
||||
assert_eq!(normalized.pages[0]["images"][0]["bbox"]["top_left_x"], 1);
|
||||
assert_eq!(normalized.pages[0]["images"][0]["description"], "scan");
|
||||
assert_eq!(normalized.pages[1]["index"], 1);
|
||||
assert_eq!(normalized.pages[1]["markdown"], "");
|
||||
assert_eq!(normalized.pages[1]["blocks"][0]["text"], "total");
|
||||
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 3);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_defaults_and_invalid_fields() {
|
||||
for value in [
|
||||
json!({}),
|
||||
json!({"meta":null}),
|
||||
json!({"pages":[],"meta":{"billed_units":null}}),
|
||||
] {
|
||||
let normalized =
|
||||
transform_response("parse", serde_json::from_value(value).unwrap()).unwrap();
|
||||
assert!(normalized.pages.is_empty());
|
||||
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 0);
|
||||
}
|
||||
for value in [
|
||||
json!({"pages":null}),
|
||||
json!({"pages":[{"markdown":"text"}]}),
|
||||
json!({"pages":[{"index":"bad"}]}),
|
||||
] {
|
||||
assert!(serde_json::from_value::<CohereResponse>(value).is_err());
|
||||
}
|
||||
let normalized = transform_response(
|
||||
"parse",
|
||||
serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 1);
|
||||
assert!(normalized.pages[0]["images"].is_null());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_requires_image_and_supported_output_format() {
|
||||
for value in [
|
||||
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
|
||||
json!({"type":"image_url","image_url":""}),
|
||||
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
|
||||
] {
|
||||
assert_eq!(
|
||||
validate_document(&serde_json::from_value(value).unwrap()),
|
||||
Err(OcrRequestError::CohereImageOnly)
|
||||
);
|
||||
}
|
||||
assert!(serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err());
|
||||
for format in ["markdown", "blocks"] {
|
||||
assert!(
|
||||
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
|
||||
);
|
||||
}
|
||||
let request = transform_request(
|
||||
"parse-v5.0",
|
||||
serde_json::from_value(json!({
|
||||
"type":"image_url",
|
||||
"image_url":"https://example.com/image.png"
|
||||
}))
|
||||
.unwrap(),
|
||||
serde_json::from_value(json!({})).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
serde_json::to_value(request).unwrap()["output_format"],
|
||||
"markdown"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
pub(crate) mod cohere;
|
||||
pub(crate) mod deepseek;
|
||||
pub(crate) mod document_intelligence;
|
||||
pub(crate) mod mistral;
|
||||
|
|
|
|||
|
|
@ -4,6 +4,10 @@ use crate::error::TransportError;
|
|||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Error)]
|
||||
pub enum OcrRequestError {
|
||||
#[error(
|
||||
"Cohere Parse only accepts `image_url` documents; document_url and PDF inputs are not supported"
|
||||
)]
|
||||
CohereImageOnly,
|
||||
#[error("Invalid `req_format`. Expected 'native' or 'litellm'.")]
|
||||
RequestFormat,
|
||||
#[error("invalid OCR request field: {path}")]
|
||||
|
|
|
|||
|
|
@ -75,10 +75,10 @@ impl PreparedOcrCall {
|
|||
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
|
||||
match self.request.adapter {
|
||||
$( OcrAdapterKind::$variant => {
|
||||
let bytes = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?;
|
||||
let decoded = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?;
|
||||
Ok(OcrProviderResponse {
|
||||
request: self.request,
|
||||
bytes,
|
||||
data: OcrProviderData::$variant(decoded),
|
||||
})
|
||||
}, )+
|
||||
}
|
||||
|
|
@ -106,12 +106,14 @@ fn request_headers(request: &reqwest::Request) -> Result<Vec<(String, String)>,
|
|||
|
||||
macro_rules! provider_data {
|
||||
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
|
||||
enum OcrProviderData {
|
||||
$( $variant(super::wire::DecodedOcrResponse<<$adapter as OcrAdapter>::ProviderResponse>), )+
|
||||
}
|
||||
|
||||
impl OcrProviderResponse {
|
||||
pub(crate) fn normalize(self) -> Result<LiteLLMOcrResponse, Error> {
|
||||
let native = self.request.response_format()? == super::types::OcrResponseFormat::Native;
|
||||
match self.request.adapter {
|
||||
$( OcrAdapterKind::$variant => {
|
||||
let decoded = super::wire::decode_response::<<$adapter as OcrAdapter>::ProviderResponse>(&self.bytes, native)?;
|
||||
match self.data {
|
||||
$( OcrProviderData::$variant(decoded) => {
|
||||
let response = $instance.transform_ocr_response(&self.request, decoded.data)?;
|
||||
Ok(LiteLLMOcrResponse { provider_native_response: decoded.native, ..response })
|
||||
}, )+
|
||||
|
|
@ -123,7 +125,7 @@ macro_rules! provider_data {
|
|||
|
||||
pub(crate) struct OcrProviderResponse {
|
||||
request: LiteLLMOcrRequest,
|
||||
bytes: Vec<u8>,
|
||||
data: OcrProviderData,
|
||||
}
|
||||
|
||||
pub(crate) async fn post_call(hooks: &Arc<dyn OcrHooks>, bytes: &[u8]) -> Result<(), Error> {
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ super::adapters::for_each_ocr_adapter!(define_adapter_types);
|
|||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) enum OcrProvider {
|
||||
Cohere,
|
||||
Mistral,
|
||||
AzureAi,
|
||||
Reducto,
|
||||
|
|
@ -32,6 +33,7 @@ pub(crate) enum OcrProvider {
|
|||
impl OcrProvider {
|
||||
pub(crate) const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Cohere => "cohere",
|
||||
Self::Mistral => "mistral",
|
||||
Self::AzureAi => "azure_ai",
|
||||
Self::Reducto => "reducto",
|
||||
|
|
@ -50,6 +52,7 @@ pub(crate) fn resolve_wire_adapter(
|
|||
custom_llm_provider: OcrProvider::Mistral.as_str(),
|
||||
});
|
||||
let typed_provider = match provider.custom_llm_provider {
|
||||
"cohere" => OcrProvider::Cohere,
|
||||
"mistral" => OcrProvider::Mistral,
|
||||
"azure_ai" => OcrProvider::AzureAi,
|
||||
"reducto" => OcrProvider::Reducto,
|
||||
|
|
@ -57,10 +60,17 @@ pub(crate) fn resolve_wire_adapter(
|
|||
value => return Err(Error::InvalidProvider(value.to_string())),
|
||||
};
|
||||
let adapter = match typed_provider {
|
||||
OcrProvider::Cohere => OcrAdapterKind::Cohere,
|
||||
OcrProvider::Mistral => OcrAdapterKind::Mistral,
|
||||
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
|
||||
OcrAdapterKind::AzureDocumentIntelligence
|
||||
}
|
||||
OcrProvider::AzureAi
|
||||
if provider.model.to_ascii_lowercase().contains("cohere")
|
||||
&& provider.model.to_ascii_lowercase().contains("parse") =>
|
||||
{
|
||||
OcrAdapterKind::AzureCohere
|
||||
}
|
||||
OcrProvider::AzureAi => OcrAdapterKind::AzureMistral,
|
||||
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
|
||||
OcrAdapterKind::ReductoLegacy
|
||||
|
|
|
|||
|
|
@ -89,6 +89,7 @@ pub fn consumed_optional_param_names(
|
|||
|
||||
let (_, adapter) = super::registry::resolve_wire_adapter(model, custom_llm_provider)?;
|
||||
let provider_fields: &[&str] = match adapter {
|
||||
OcrAdapterKind::Cohere | OcrAdapterKind::AzureCohere => &["output_format"],
|
||||
OcrAdapterKind::Mistral | OcrAdapterKind::AzureMistral | OcrAdapterKind::VertexMistral => {
|
||||
MISTRAL_OPTION_FIELDS
|
||||
}
|
||||
|
|
@ -98,9 +99,9 @@ pub fn consumed_optional_param_names(
|
|||
OcrAdapterKind::VertexDeepSeek => DEEPSEEK_OPTION_FIELDS,
|
||||
};
|
||||
let auth_fields: &[&str] = match adapter {
|
||||
OcrAdapterKind::AzureMistral | OcrAdapterKind::AzureDocumentIntelligence => {
|
||||
AZURE_AUTH_OPTION_FIELDS
|
||||
}
|
||||
OcrAdapterKind::AzureMistral
|
||||
| OcrAdapterKind::AzureDocumentIntelligence
|
||||
| OcrAdapterKind::AzureCohere => AZURE_AUTH_OPTION_FIELDS,
|
||||
OcrAdapterKind::VertexMistral | OcrAdapterKind::VertexDeepSeek => VERTEX_AUTH_OPTION_FIELDS,
|
||||
_ => &[],
|
||||
};
|
||||
|
|
|
|||
|
|
@ -81,7 +81,7 @@ impl AzureAuthService {
|
|||
AzureCredentialPlan::Caller(caller) => {
|
||||
let credential = caller.acquire().await?;
|
||||
if credential.secret().expose().is_empty() {
|
||||
return Ok(None);
|
||||
return Err(AuthError::EmptyAzureToken);
|
||||
}
|
||||
Ok(Some(Sourced::new(credential, InputSource::Deployment)))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -180,8 +180,17 @@ impl super::hooks::OcrHooks for SubmissionBoundary {
|
|||
request: super::hooks::OcrPostCallRequest,
|
||||
) -> super::hooks::OcrHookFuture<'_, super::hooks::OcrPostCallRequest> {
|
||||
Box::pin(async move {
|
||||
assert_eq!(self.request_count.lock().unwrap().len(), 1);
|
||||
assert_eq!(request.original_response, json!(r#"{"submitted":true}"#));
|
||||
match self.request_count.lock().unwrap().len() {
|
||||
1 => assert_eq!(request.original_response, json!(r#"{"submitted":true}"#)),
|
||||
2 => assert!(
|
||||
request
|
||||
.original_response
|
||||
.as_str()
|
||||
.unwrap()
|
||||
.contains("succeeded")
|
||||
),
|
||||
count => panic!("unexpected callback after {count} requests"),
|
||||
}
|
||||
Ok(request)
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -73,7 +73,18 @@ pub(crate) fn ocr_error_to_pyerr(err: Error) -> PyErr {
|
|||
Error::Network(message) if message.contains("timed out") => {
|
||||
ocr_upstream_error(408, message)
|
||||
}
|
||||
other => core_error_to_pyerr(other),
|
||||
other => {
|
||||
let status = other.http_status_code();
|
||||
let error = core_error_to_pyerr(other);
|
||||
if let Some(status) = status {
|
||||
Python::attach(|py| {
|
||||
let value = error.value(py);
|
||||
value.setattr("status_code", status).ok();
|
||||
value.setattr("message", value.to_string()).ok();
|
||||
});
|
||||
}
|
||||
error
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -111,6 +122,18 @@ mod ocr_error_tests {
|
|||
.and_then(|args| args.extract())
|
||||
.expect("OCR failures retain status and unprefixed provider message");
|
||||
assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string()));
|
||||
|
||||
let mapped = ocr_error_to_pyerr(Error::InvalidRequest("invalid format".into()));
|
||||
assert!(mapped.is_instance_of::<PyValueError>(py));
|
||||
assert_eq!(
|
||||
mapped
|
||||
.value(py)
|
||||
.getattr("status_code")
|
||||
.unwrap()
|
||||
.extract::<u16>()
|
||||
.unwrap(),
|
||||
400
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -60,9 +60,14 @@ where
|
|||
E: Send + 'static,
|
||||
F: Future<Output = Result<T, E>> + Send + 'static,
|
||||
{
|
||||
let result = run_sync_value_on(py, runtime, async move {
|
||||
map_core_result(future.await, map_error)
|
||||
})?;
|
||||
if Handle::try_current().is_ok() {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
"synchronous native routes cannot run from a Tokio context; use the async route",
|
||||
));
|
||||
}
|
||||
|
||||
let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?;
|
||||
let result = map_core_result(result, map_error)?;
|
||||
Pythonized(result).into_pyobject(py).map(Bound::unbind)
|
||||
}
|
||||
|
||||
|
|
@ -76,8 +81,9 @@ where
|
|||
E: Send + 'static,
|
||||
F: Future<Output = Result<T, E>> + Send + 'static,
|
||||
{
|
||||
run_async_value(py, async move {
|
||||
let result = map_core_result(future.await, map_error)?;
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
let result = catch_future_panic(future).await?;
|
||||
let result = map_core_result(result, map_error)?;
|
||||
Ok(Pythonized(result))
|
||||
})
|
||||
}
|
||||
|
|
|
|||
|
|
@ -65,7 +65,7 @@ impl ResponsesWebSocketConnection {
|
|||
}
|
||||
}
|
||||
|
||||
#[pymodule(gil_used = true)]
|
||||
#[pymodule(gil_used = false)]
|
||||
mod _native {
|
||||
use pyo3::prelude::*;
|
||||
|
||||
|
|
|
|||
|
|
@ -10,7 +10,9 @@ use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_param_names, dec
|
|||
use litellm_core::ocr::{
|
||||
NativeOutcome, OcrAdmission, OcrCall, OcrClient, OcrHostOperation, OcrHostResult,
|
||||
};
|
||||
use litellm_python_interop::{from_py, to_py};
|
||||
use litellm_python_interop::{
|
||||
from_py_preserving_errors as from_py, to_py_preserving_errors as to_py,
|
||||
};
|
||||
|
||||
use crate::errors::{RustBridgeDeclined, ocr_error_to_pyerr};
|
||||
use crate::lifecycle::{PythonCallState, PythonRoute, missing_state, now, run_call};
|
||||
|
|
@ -82,9 +84,14 @@ impl PythonOcrHost {
|
|||
fn python_pre_call(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
request: OcrDuringCallRequest,
|
||||
mut request: OcrDuringCallRequest,
|
||||
) -> PyResult<OcrDuringCallRequest> {
|
||||
let pre_call = self.pre_call.as_ref().ok_or_else(missing_state)?;
|
||||
if let Some(body) = request.body.as_object_mut() {
|
||||
for name in &request.retained_fields {
|
||||
body.remove(name);
|
||||
}
|
||||
}
|
||||
let body = to_py(py, &request.body)?
|
||||
.into_bound(py)
|
||||
.cast_into::<PyDict>()?;
|
||||
|
|
@ -134,11 +141,9 @@ impl PythonOcrHost {
|
|||
.iter()
|
||||
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
Ok(OcrDuringCallRequest {
|
||||
body: from_py(&body)?,
|
||||
headers,
|
||||
..request
|
||||
})
|
||||
request.body = from_py(&body)?;
|
||||
request.headers = headers;
|
||||
Ok(request)
|
||||
}
|
||||
|
||||
fn python_post_call(
|
||||
|
|
@ -413,6 +418,13 @@ fn _ocr_lifecycle(
|
|||
kwargs: Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
if let Ok(gil_enabled) = py.import("sys")?.getattr("_is_gil_enabled")
|
||||
&& !gil_enabled.call0()?.is_truthy()?
|
||||
{
|
||||
return Err(pyo3::exceptions::PyRuntimeError::new_err(
|
||||
"native OCR requires the Python GIL",
|
||||
));
|
||||
}
|
||||
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
|
||||
let call = admitted_call(OcrCall::admit(
|
||||
client,
|
||||
|
|
|
|||
|
|
@ -2,4 +2,6 @@ mod gil;
|
|||
mod marshal;
|
||||
|
||||
pub use gil::{release_count, release_gil};
|
||||
pub use marshal::{Pythonized, from_py, panic_to_pyerr, to_py};
|
||||
pub use marshal::{
|
||||
Pythonized, from_py, from_py_preserving_errors, panic_to_pyerr, to_py, to_py_preserving_errors,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -1,12 +1,20 @@
|
|||
use std::any::Any;
|
||||
use std::panic::{AssertUnwindSafe, catch_unwind};
|
||||
|
||||
use pyo3::exceptions::PyValueError;
|
||||
use pyo3::panic::PanicException;
|
||||
use pyo3::prelude::*;
|
||||
use serde::Serialize;
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
pub fn from_py<T>(value: &Bound<'_, PyAny>) -> PyResult<T>
|
||||
where
|
||||
T: DeserializeOwned,
|
||||
{
|
||||
pythonize::depythonize(value).map_err(|error| PyValueError::new_err(error.to_string()))
|
||||
}
|
||||
|
||||
pub fn from_py_preserving_errors<T>(value: &Bound<'_, PyAny>) -> PyResult<T>
|
||||
where
|
||||
T: DeserializeOwned,
|
||||
{
|
||||
|
|
@ -17,14 +25,18 @@ pub fn to_py<T>(py: Python<'_>, value: &T) -> PyResult<Py<PyAny>>
|
|||
where
|
||||
T: Serialize + ?Sized,
|
||||
{
|
||||
pythonize_bound(py, value).map(Bound::unbind)
|
||||
pythonize::pythonize(py, value)
|
||||
.map(Bound::unbind)
|
||||
.map_err(|error| PyValueError::new_err(error.to_string()))
|
||||
}
|
||||
|
||||
fn pythonize_bound<'py, T>(py: Python<'py>, value: &T) -> PyResult<Bound<'py, PyAny>>
|
||||
pub fn to_py_preserving_errors<T>(py: Python<'_>, value: &T) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
T: Serialize + ?Sized,
|
||||
{
|
||||
pythonize::pythonize(py, value).map_err(PyErr::from)
|
||||
pythonize::pythonize(py, value)
|
||||
.map(Bound::unbind)
|
||||
.map_err(PyErr::from)
|
||||
}
|
||||
|
||||
pub struct Pythonized<T>(pub T);
|
||||
|
|
@ -38,7 +50,9 @@ where
|
|||
type Error = PyErr;
|
||||
|
||||
fn into_pyobject(self, py: Python<'py>) -> PyResult<Self::Output> {
|
||||
catch_unwind(AssertUnwindSafe(|| pythonize_bound(py, &self.0))).map_err(panic_to_pyerr)?
|
||||
catch_unwind(AssertUnwindSafe(|| pythonize::pythonize(py, &self.0)))
|
||||
.map_err(panic_to_pyerr)?
|
||||
.map_err(|error| PyValueError::new_err(error.to_string()))
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -112,7 +126,15 @@ value = Broken()
|
|||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let error = from_py::<i64>(&locals.get_item("value").unwrap().unwrap()).unwrap_err();
|
||||
let value = locals.get_item("value").unwrap().unwrap();
|
||||
let legacy_error = from_py::<i64>(&value).unwrap_err();
|
||||
assert!(legacy_error.is_instance_of::<PyValueError>(py));
|
||||
assert!(
|
||||
!legacy_error
|
||||
.value(py)
|
||||
.is(locals.get_item("failure").unwrap().unwrap())
|
||||
);
|
||||
let error = from_py_preserving_errors::<i64>(&value).unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.value(py)
|
||||
|
|
|
|||
|
|
@ -1,241 +0,0 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextvars
|
||||
import datetime
|
||||
from collections.abc import Awaitable, Callable, Coroutine
|
||||
from concurrent.futures import Future
|
||||
from typing import Protocol
|
||||
|
||||
|
||||
class CompletionLogging(Protocol):
|
||||
def success_handler(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
def failure_handler(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
def async_success_handler(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> Coroutine[object, object, None]: ...
|
||||
|
||||
def async_failure_handler(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> Awaitable[None]: ...
|
||||
|
||||
def handle_sync_success_callbacks_for_async_calls(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
|
||||
class CompletionExecutor(Protocol):
|
||||
def submit(
|
||||
self,
|
||||
function: Callable[..., object],
|
||||
/,
|
||||
*args: object,
|
||||
) -> Future[object]: ...
|
||||
|
||||
|
||||
class Completion(Protocol):
|
||||
def success(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
def failure(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
async def async_failure(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
|
||||
class PythonCompletion:
|
||||
def __init__(
|
||||
self,
|
||||
logging_obj: CompletionLogging,
|
||||
executor: CompletionExecutor | None,
|
||||
*,
|
||||
async_call: bool,
|
||||
internal_call: bool,
|
||||
completion_with_fallbacks: bool,
|
||||
) -> None:
|
||||
self._logging_obj = logging_obj
|
||||
self._executor = executor
|
||||
self._async_call = async_call
|
||||
self._internal_call = internal_call
|
||||
self._completion_with_fallbacks = completion_with_fallbacks
|
||||
|
||||
def success(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
if not self._async_call:
|
||||
assert self._executor is not None
|
||||
context = contextvars.copy_context()
|
||||
self._executor.submit(
|
||||
context.run,
|
||||
self._logging_obj.success_handler,
|
||||
result,
|
||||
start_time,
|
||||
end_time,
|
||||
)
|
||||
return
|
||||
|
||||
if not self._internal_call:
|
||||
if getattr(self._logging_obj, "_defer_async_logging", False):
|
||||
|
||||
def enqueue_deferred_logging() -> None:
|
||||
asyncio.create_task(self._dispatch_async_success(result, start_time, end_time))
|
||||
|
||||
setattr( # noqa: B010 # optional legacy logger field is absent from narrow test doubles
|
||||
self._logging_obj,
|
||||
"_enqueue_deferred_logging",
|
||||
enqueue_deferred_logging,
|
||||
)
|
||||
else:
|
||||
asyncio.create_task(self._dispatch_async_success(result, start_time, end_time))
|
||||
|
||||
self._logging_obj.handle_sync_success_callbacks_for_async_calls(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
def failure(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
if self._async_call and self._internal_call:
|
||||
return
|
||||
self._logging_obj.failure_handler(exception, traceback_exception, start_time, end_time)
|
||||
|
||||
async def async_failure(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
if not self._async_call or self._internal_call:
|
||||
return
|
||||
await self._logging_obj.async_failure_handler(exception, traceback_exception, start_time, end_time)
|
||||
|
||||
async def _dispatch_async_success(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
if self._completion_with_fallbacks:
|
||||
return
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( # pyright: ignore[reportUnknownMemberType] # legacy worker lacks generic coroutine annotations
|
||||
async_coroutine=self._logging_obj.async_success_handler(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
)
|
||||
self._logging_obj.handle_sync_success_callbacks_for_async_calls(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
|
||||
class CallCompletion:
|
||||
def __init__(self, implementation: Completion) -> None:
|
||||
self._python_implementation: Completion | None = implementation
|
||||
self._implementation: Completion | None = implementation
|
||||
self._attached = False
|
||||
|
||||
@property
|
||||
def python_implementation(self) -> Completion:
|
||||
assert self._python_implementation is not None
|
||||
return self._python_implementation
|
||||
|
||||
def attach(self, implementation: Completion) -> bool:
|
||||
if self._attached:
|
||||
return False
|
||||
self._implementation = implementation
|
||||
self._attached = True
|
||||
return True
|
||||
|
||||
def release(self) -> None:
|
||||
self._python_implementation = None
|
||||
self._implementation = None
|
||||
|
||||
def success(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
implementation = self._implementation
|
||||
assert implementation is not None
|
||||
implementation.success(result, start_time, end_time)
|
||||
|
||||
def failure(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
implementation = self._implementation
|
||||
assert implementation is not None
|
||||
implementation.failure(exception, traceback_exception, start_time, end_time)
|
||||
|
||||
async def async_failure(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
implementation = self._implementation
|
||||
assert implementation is not None
|
||||
await implementation.async_failure(
|
||||
exception,
|
||||
traceback_exception,
|
||||
start_time,
|
||||
end_time,
|
||||
)
|
||||
|
|
@ -3,7 +3,9 @@ import mimetypes
|
|||
import os
|
||||
import re
|
||||
from io import IOBase
|
||||
from typing import Any, Final
|
||||
from typing import Final, Literal, Protocol
|
||||
|
||||
from typing_extensions import NotRequired, ReadOnly, TypedDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
|
|
@ -31,7 +33,17 @@ def get_mime_type(file_path: str) -> str:
|
|||
return guessed or "application/octet-stream"
|
||||
|
||||
|
||||
def convert_file_document_to_url_document(document: dict[str, Any]) -> dict[str, str]:
|
||||
class FileReader(Protocol):
|
||||
def read(self) -> bytes | str: ...
|
||||
|
||||
|
||||
class FileDocument(TypedDict):
|
||||
type: ReadOnly[Literal["file"]]
|
||||
file: ReadOnly[bytes | os.PathLike[str] | FileReader]
|
||||
mime_type: ReadOnly[NotRequired[str]]
|
||||
|
||||
|
||||
def convert_file_document_to_url_document(document: FileDocument) -> dict[str, str]:
|
||||
file_input: Final = document.get("file")
|
||||
if file_input is None:
|
||||
raise ValueError(
|
||||
|
|
|
|||
|
|
@ -3181,10 +3181,11 @@ class ProxyBaseLLMRequestProcessing:
|
|||
Extracted as a static method so tests can exercise the production
|
||||
gating logic directly rather than reimplementing the finally block.
|
||||
"""
|
||||
pending: Final = getattr(logging_obj, "_native_pending_logging", None)
|
||||
if pending is not None:
|
||||
logging_obj._native_pending_logging = None # rebind-ok: consume the native release signal once
|
||||
pending.release(not exception_raised)
|
||||
if getattr(logging_obj, "call_type", None) in ("ocr", "aocr"):
|
||||
pending: Final = getattr(logging_obj, "_native_pending_logging", None)
|
||||
if pending is not None:
|
||||
logging_obj._native_pending_logging = None # rebind-ok: consume the native OCR release signal once
|
||||
pending.release(not exception_raised)
|
||||
_enqueue_fn: Final = getattr(logging_obj, "_enqueue_deferred_logging", None)
|
||||
if _enqueue_fn is None:
|
||||
return
|
||||
|
|
|
|||
|
|
@ -8,7 +8,6 @@ from typing import Final, Protocol, cast # noqa: TID251 # validates dynamicall
|
|||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.ocr import LiteLLMOcrRequest
|
||||
|
||||
|
||||
|
|
@ -44,8 +43,6 @@ NATIVE_OCR_LIFECYCLE: Final = NativeBinding("_ocr_lifecycle", validate=_binding)
|
|||
|
||||
|
||||
def select(request: LiteLLMOcrRequest) -> NativeOcrLifecycle | None:
|
||||
if not rust_enabled():
|
||||
return None
|
||||
if litellm.cache is not None or request.kwargs.get("caching") or request.kwargs.get("aocr"):
|
||||
return None
|
||||
return NATIVE_OCR_LIFECYCLE.load()
|
||||
|
|
@ -99,12 +96,6 @@ def call_azure_ad_token_provider(provider: object) -> str:
|
|||
|
||||
|
||||
def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception:
|
||||
if isinstance(error, ValueError) and "Invalid `req_format`" in str(error):
|
||||
return litellm.BadRequestError(
|
||||
message=str(error),
|
||||
model=request.model,
|
||||
llm_provider=request_provider,
|
||||
)
|
||||
mapper: Final = cast( # cast-ok: bounded adapter for the legacy public exception mapper
|
||||
ExceptionMapper, litellm.exception_type
|
||||
)
|
||||
|
|
|
|||
178
litellm/utils.py
178
litellm/utils.py
|
|
@ -7,6 +7,7 @@ import ast
|
|||
import asyncio
|
||||
import base64
|
||||
import binascii
|
||||
import contextvars
|
||||
import copy
|
||||
import datetime
|
||||
import hashlib
|
||||
|
|
@ -80,7 +81,6 @@ from litellm.constants import (
|
|||
PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO,
|
||||
TOOL_CHOICE_OBJECT_TOKEN_COUNT,
|
||||
)
|
||||
from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion
|
||||
from litellm.litellm_core_utils.core_helpers import normalize_drop_params
|
||||
from litellm.litellm_core_utils.fallback_generalizations import (
|
||||
match_capability_generalizations,
|
||||
|
|
@ -1196,6 +1196,79 @@ def function_setup(
|
|||
raise e
|
||||
|
||||
|
||||
def _dispatch_success_logging(
|
||||
logging_obj: LiteLLMLoggingObject,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
is_completion_with_fallbacks: bool,
|
||||
is_litellm_internal_call: bool,
|
||||
) -> None:
|
||||
if not is_litellm_internal_call:
|
||||
if getattr(logging_obj, "_defer_async_logging", False):
|
||||
|
||||
def _enqueue_deferred_logging() -> None:
|
||||
asyncio.create_task(
|
||||
_client_async_logging_helper(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
)
|
||||
)
|
||||
|
||||
logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging
|
||||
else:
|
||||
asyncio.create_task(
|
||||
_client_async_logging_helper(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
)
|
||||
)
|
||||
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
|
||||
async def _client_async_logging_helper(
|
||||
logging_obj: LiteLLMLoggingObject,
|
||||
result,
|
||||
start_time,
|
||||
end_time,
|
||||
is_completion_with_fallbacks: bool,
|
||||
):
|
||||
if (
|
||||
is_completion_with_fallbacks is False
|
||||
): # don't log the parent event litellm.completion_with_fallbacks as a 'log_success_event', this will lead to double logging the same call - https://github.com/BerriAI/litellm/issues/7477
|
||||
print_verbose(
|
||||
f"Async Wrapper: Completed Call, calling async_success_handler: {logging_obj.async_success_handler}"
|
||||
)
|
||||
################################################
|
||||
# Async Logging Worker
|
||||
################################################
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(
|
||||
async_coroutine=logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time)
|
||||
)
|
||||
|
||||
################################################
|
||||
# Sync Logging Worker
|
||||
################################################
|
||||
logging_obj.handle_sync_success_callbacks_for_async_calls(
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
|
||||
def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tuple[int | None, dict[str, Any]]:
|
||||
"""
|
||||
Get the number of retries from the kwargs and the retry policy.
|
||||
|
|
@ -1464,7 +1537,6 @@ def client(original_function):
|
|||
start_time: Final = datetime.datetime.now()
|
||||
result = None
|
||||
logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None)
|
||||
completion: CallCompletion | None = None
|
||||
|
||||
# only set litellm_call_id if its not in kwargs
|
||||
if "litellm_call_id" not in kwargs:
|
||||
|
|
@ -1480,17 +1552,7 @@ def client(original_function):
|
|||
|
||||
# Type assertion: logging_obj is guaranteed to be non-None after function_setup
|
||||
assert logging_obj is not None, "logging_obj should not be None after function_setup"
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor
|
||||
|
||||
completion = CallCompletion(
|
||||
PythonCompletion(
|
||||
logging_obj,
|
||||
logging_executor,
|
||||
async_call=False,
|
||||
internal_call=False,
|
||||
completion_with_fallbacks=False,
|
||||
)
|
||||
)
|
||||
## LOAD CREDENTIALS
|
||||
load_credentials_from_list(kwargs)
|
||||
kwargs["litellm_logging_obj"] = logging_obj
|
||||
|
|
@ -1588,12 +1650,7 @@ def client(original_function):
|
|||
except Exception as e:
|
||||
print_verbose(f"Error while checking max token limit: {e}")
|
||||
# MODEL CALL
|
||||
invocation_kwargs: Final = (
|
||||
{**kwargs, "_litellm_call_completion": completion}
|
||||
if original_function.__name__ == CallTypes.ocr.value
|
||||
else kwargs
|
||||
)
|
||||
result = original_function(*args, **invocation_kwargs)
|
||||
result = original_function(*args, **kwargs)
|
||||
end_time = datetime.datetime.now()
|
||||
if _is_streaming_request(
|
||||
kwargs=kwargs,
|
||||
|
|
@ -1647,8 +1704,8 @@ def client(original_function):
|
|||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata")
|
||||
update_response_metadata(
|
||||
_update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata")
|
||||
_update_response_metadata(
|
||||
result=result,
|
||||
logging_obj=logging_obj,
|
||||
model=model,
|
||||
|
|
@ -1656,8 +1713,21 @@ def client(original_function):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
# LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated
|
||||
verbose_logger.info("Wrapper: Completed Call, calling success_handler")
|
||||
completion.success(result, start_time, end_time)
|
||||
# Copy the current context to propagate it to the background thread
|
||||
# This is essential for OpenTelemetry span context propagation
|
||||
ctx: Final = contextvars.copy_context()
|
||||
executor: Final = getattr(sys.modules[__name__], "executor")
|
||||
executor.submit(
|
||||
ctx.run,
|
||||
logging_obj.success_handler,
|
||||
result,
|
||||
start_time,
|
||||
end_time,
|
||||
)
|
||||
# RETURN RESULT
|
||||
return result
|
||||
except Exception as e:
|
||||
call_type = original_function.__name__
|
||||
|
|
@ -1731,14 +1801,11 @@ def client(original_function):
|
|||
end_time = datetime.datetime.now()
|
||||
|
||||
# LOG FAILURE - handle streaming failure logging in the _next_ object, remove `handle_failure` once it's deprecated
|
||||
if completion is not None:
|
||||
completion.failure(e, traceback_exception, start_time, end_time)
|
||||
elif logging_obj:
|
||||
logging_obj.failure_handler(e, traceback_exception, start_time, end_time)
|
||||
if logging_obj:
|
||||
logging_obj.failure_handler(
|
||||
e, traceback_exception, start_time, end_time
|
||||
) # DO NOT MAKE THREADED - router retry fallback relies on this!
|
||||
raise e
|
||||
finally:
|
||||
if completion is not None:
|
||||
completion.release()
|
||||
|
||||
@wraps(original_function)
|
||||
async def wrapper_async(*args, **kwargs):
|
||||
|
|
@ -1747,7 +1814,6 @@ def client(original_function):
|
|||
result = None
|
||||
_update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata")
|
||||
logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None)
|
||||
completion: CallCompletion | None = None
|
||||
LLMCachingHandler: Final = _get_cached_llm_caching_handler()
|
||||
_llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler(
|
||||
original_function=original_function,
|
||||
|
|
@ -1773,15 +1839,7 @@ def client(original_function):
|
|||
|
||||
# Type assertion: logging_obj is guaranteed to be non-None after function_setup
|
||||
assert logging_obj is not None, "logging_obj should not be None after function_setup"
|
||||
completion = CallCompletion(
|
||||
PythonCompletion(
|
||||
logging_obj,
|
||||
None,
|
||||
async_call=True,
|
||||
internal_call=_is_litellm_internal_call,
|
||||
completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
)
|
||||
)
|
||||
|
||||
modified_kwargs: Final = await async_pre_call_deployment_hook(kwargs, call_type)
|
||||
if modified_kwargs is not None:
|
||||
kwargs = modified_kwargs
|
||||
|
|
@ -1866,12 +1924,7 @@ def client(original_function):
|
|||
|
||||
# MODEL CALL
|
||||
try:
|
||||
invocation_kwargs: Final = (
|
||||
{**kwargs, "_litellm_call_completion": completion}
|
||||
if original_function.__name__ == CallTypes.aocr.value
|
||||
else kwargs
|
||||
)
|
||||
result = await original_function(*args, **invocation_kwargs)
|
||||
result = await original_function(*args, **kwargs)
|
||||
except Exception as deployment_error:
|
||||
_deployment_call_end_time = datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with
|
||||
try:
|
||||
|
|
@ -1940,7 +1993,14 @@ def client(original_function):
|
|||
and _caching_handler_response is not None
|
||||
and _caching_handler_response.final_embedding_cached_response is not None
|
||||
):
|
||||
completion.success(result, start_time, end_time)
|
||||
_dispatch_success_logging(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
is_litellm_internal_call=_is_litellm_internal_call,
|
||||
)
|
||||
return _llm_caching_handler._combine_cached_embedding_response_with_api_result(
|
||||
_caching_handler_response=_caching_handler_response,
|
||||
embedding_response=result,
|
||||
|
|
@ -1956,7 +2016,14 @@ def client(original_function):
|
|||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
completion.success(result, start_time, end_time)
|
||||
_dispatch_success_logging(
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
is_completion_with_fallbacks=is_completion_with_fallbacks,
|
||||
is_litellm_internal_call=_is_litellm_internal_call,
|
||||
)
|
||||
|
||||
return result
|
||||
except Exception as e:
|
||||
|
|
@ -1964,12 +2031,17 @@ def client(original_function):
|
|||
# Reuse the timestamp taken right when the deployment call itself failed, before
|
||||
# the failure hook ran, so a slow callback doesn't inflate the reported duration.
|
||||
end_time = _deployment_call_end_time if _deployment_call_end_time is not None else datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with
|
||||
if completion is not None:
|
||||
completion.failure(e, traceback_exception, start_time, end_time)
|
||||
await completion.async_failure(e, traceback_exception, start_time, end_time)
|
||||
elif logging_obj and not _is_litellm_internal_call:
|
||||
logging_obj.failure_handler(e, traceback_exception, start_time, end_time)
|
||||
await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time)
|
||||
if logging_obj and not _is_litellm_internal_call:
|
||||
try:
|
||||
logging_obj.failure_handler(
|
||||
e, traceback_exception, start_time, end_time
|
||||
) # DO NOT MAKE THREADED - router retry fallback relies on this!
|
||||
except Exception as e:
|
||||
raise e
|
||||
try:
|
||||
await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time)
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
call_type = original_function.__name__
|
||||
num_retries, kwargs = _get_wrapper_num_retries(kwargs=kwargs, exception=e)
|
||||
|
|
@ -2039,8 +2111,6 @@ def client(original_function):
|
|||
raise e
|
||||
|
||||
finally:
|
||||
if completion is not None:
|
||||
completion.release()
|
||||
# Restore trace_id/session_id contextvars to their pre-call value once
|
||||
# this call (in this asyncio Task) is fully done - see
|
||||
# request_correlation_in_logs. Unlike wrapper()'s sync path, it's safe to
|
||||
|
|
|
|||
|
|
@ -1,402 +0,0 @@
|
|||
import asyncio
|
||||
import contextvars
|
||||
import datetime
|
||||
import weakref
|
||||
from collections.abc import Callable, Coroutine
|
||||
from concurrent.futures import Future, ThreadPoolExecutor
|
||||
from threading import get_ident
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils import thread_pool_executor
|
||||
from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.utils import client
|
||||
|
||||
|
||||
class RecordingExecutor:
|
||||
def __init__(self) -> None:
|
||||
self.submissions: list[tuple[Callable[..., object], tuple[object, ...]]] = []
|
||||
|
||||
def submit(self, function: Callable[..., object], *args: object) -> Future[object]:
|
||||
self.submissions.append((function, args))
|
||||
future: Final[Future[object]] = Future()
|
||||
future.set_result(function(*args))
|
||||
return future
|
||||
|
||||
|
||||
class RecordingCompletion:
|
||||
def __init__(self) -> None:
|
||||
self.successes: list[object] = []
|
||||
self.failures: list[Exception] = []
|
||||
|
||||
def success(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
self.successes.append(result)
|
||||
|
||||
def failure(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
self.failures.append(exception)
|
||||
|
||||
async def async_failure(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
self.failures.append(exception)
|
||||
|
||||
|
||||
class RecordingLogging:
|
||||
def __init__(self, marker: contextvars.ContextVar[str], observed: list[tuple[object, str]]) -> None:
|
||||
self._marker = marker
|
||||
self._observed = observed
|
||||
self._defer_async_logging = False
|
||||
self._enqueue_deferred_logging: Callable[[], None] | None = None
|
||||
|
||||
def success_handler(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
self._observed.append((result, self._marker.get()))
|
||||
|
||||
def failure_handler(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
async def async_success_handler(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
async def async_failure_handler(
|
||||
self,
|
||||
exception: Exception,
|
||||
traceback_exception: str,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
def handle_sync_success_callbacks_for_async_calls(
|
||||
self,
|
||||
result: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None: ...
|
||||
|
||||
|
||||
def test_python_completion_preserves_sync_context_and_response_identity() -> None:
|
||||
marker: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("marker")
|
||||
marker.set("request-context")
|
||||
executor: Final = RecordingExecutor()
|
||||
response: Final = object()
|
||||
observed: Final[list[tuple[object, str]]] = []
|
||||
logging_obj: Final = RecordingLogging(marker, observed)
|
||||
completion: Final = PythonCompletion(
|
||||
logging_obj,
|
||||
executor,
|
||||
async_call=False,
|
||||
internal_call=False,
|
||||
completion_with_fallbacks=False,
|
||||
)
|
||||
now: Final = datetime.datetime.now(datetime.timezone.utc)
|
||||
|
||||
completion.success(response, now, now)
|
||||
|
||||
assert observed == [(response, "request-context")]
|
||||
assert len(executor.submissions) == 1
|
||||
|
||||
|
||||
def test_sync_wrapper_dispatches_with_logging_executor_and_caller_context(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
marker: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("wrapper-context", default="missing")
|
||||
marker.set("request-context")
|
||||
caller_thread: Final = get_ident()
|
||||
response: Final = object()
|
||||
observed: Final[list[tuple[object, str, int]]] = []
|
||||
logging_obj: Final = MagicMock()
|
||||
|
||||
def record_success(result: object, start_time: datetime.datetime, end_time: datetime.datetime) -> None:
|
||||
observed.append((result, marker.get(), get_ident()))
|
||||
|
||||
def ocr(**kwargs: object) -> object:
|
||||
return response
|
||||
|
||||
logging_obj.success_handler.side_effect = record_success
|
||||
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(logging_obj, {})))
|
||||
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
|
||||
wrapped: Final = client(ocr)
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
monkeypatch.setattr(thread_pool_executor, "executor", executor)
|
||||
result: Final = wrapped()
|
||||
|
||||
assert result is response
|
||||
assert len(observed) == 1
|
||||
assert observed[0][0] is response
|
||||
assert observed[0][1] == "request-context"
|
||||
assert observed[0][2] != caller_thread
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_completion_attaches_once_and_forwards_final_objects() -> None:
|
||||
python_completion: Final = RecordingCompletion()
|
||||
native_completion: Final = RecordingCompletion()
|
||||
ignored_completion: Final = RecordingCompletion()
|
||||
completion: Final = CallCompletion(python_completion)
|
||||
response: Final = object()
|
||||
error: Final = ValueError("mapped failure")
|
||||
now: Final = datetime.datetime.now(datetime.timezone.utc)
|
||||
|
||||
assert completion.python_implementation is python_completion
|
||||
assert completion.attach(native_completion)
|
||||
assert not completion.attach(ignored_completion)
|
||||
|
||||
completion.success(response, now, now)
|
||||
completion.failure(error, "traceback", now, now)
|
||||
await completion.async_failure(error, "traceback", now, now)
|
||||
|
||||
assert native_completion.successes == [response]
|
||||
assert native_completion.failures == [error, error]
|
||||
assert python_completion.successes == []
|
||||
assert ignored_completion.successes == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_python_completion_retains_deferred_success_arguments(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
response: Final = object()
|
||||
logging_obj: Final = MagicMock()
|
||||
logging_obj._defer_async_logging = True
|
||||
logging_obj.async_success_handler = AsyncMock()
|
||||
completion: Final = CallCompletion(
|
||||
PythonCompletion(
|
||||
logging_obj,
|
||||
RecordingExecutor(),
|
||||
async_call=True,
|
||||
internal_call=False,
|
||||
completion_with_fallbacks=False,
|
||||
)
|
||||
)
|
||||
worker: Final = MagicMock()
|
||||
monkeypatch.setattr("litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER", worker)
|
||||
scheduled: Final[list[Coroutine[object, object, None]]] = []
|
||||
monkeypatch.setattr("asyncio.create_task", scheduled.append)
|
||||
now: Final = datetime.datetime.now(datetime.timezone.utc)
|
||||
|
||||
completion.success(response, now, now)
|
||||
completion.release()
|
||||
|
||||
logging_obj._enqueue_deferred_logging()
|
||||
assert len(scheduled) == 1
|
||||
await scheduled[0]
|
||||
await worker.ensure_initialized_and_enqueue.call_args.kwargs["async_coroutine"]
|
||||
logging_obj.async_success_handler.assert_awaited_once_with(result=response, start_time=now, end_time=now)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_ocr_wrapper_injects_completion_after_fresh_deployment_kwargs(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
native_completion: Final = RecordingCompletion()
|
||||
original_response: Final = object()
|
||||
replacement_response: Final = object()
|
||||
shared_metadata: Final = {"request": "shared"}
|
||||
replacement_kwargs: Final[dict[str, object]] = {"metadata": shared_metadata}
|
||||
hook_input: dict[str, object] | None = None
|
||||
|
||||
async def fresh_kwargs(kwargs: dict[str, object], call_type: str) -> dict[str, object]:
|
||||
nonlocal hook_input
|
||||
hook_input = kwargs
|
||||
return replacement_kwargs
|
||||
|
||||
async def aocr(**kwargs: object) -> object:
|
||||
completion = kwargs.get("_litellm_call_completion")
|
||||
assert isinstance(completion, CallCompletion)
|
||||
assert completion.attach(native_completion)
|
||||
assert kwargs["metadata"] is shared_metadata
|
||||
return original_response
|
||||
|
||||
async def replace_response(request_data: dict[str, object], response: object, call_type: object) -> object:
|
||||
assert response is original_response
|
||||
return replacement_response
|
||||
|
||||
monkeypatch.setattr("litellm.utils.async_pre_call_deployment_hook", fresh_kwargs)
|
||||
monkeypatch.setattr("litellm.utils.async_post_call_success_deployment_hook", replace_response)
|
||||
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {})))
|
||||
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
|
||||
wrapped: Final = client(aocr)
|
||||
|
||||
result: Final = await wrapped()
|
||||
|
||||
assert result is replacement_response
|
||||
assert native_completion.successes == [replacement_response]
|
||||
assert hook_input is not None
|
||||
assert "_litellm_call_completion" not in hook_input
|
||||
assert "_litellm_call_completion" not in replacement_kwargs
|
||||
assert replacement_kwargs["metadata"] is shared_metadata
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_ocr_wrapper_sends_final_failure_to_attached_completion(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
native_completion: Final = RecordingCompletion()
|
||||
mapped_error: Final = ValueError("mapped OCR failure")
|
||||
|
||||
async def aocr(**kwargs: object) -> object:
|
||||
completion = kwargs.get("_litellm_call_completion")
|
||||
assert isinstance(completion, CallCompletion)
|
||||
assert completion.attach(native_completion)
|
||||
raise mapped_error
|
||||
|
||||
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {})))
|
||||
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
|
||||
wrapped: Final = client(aocr)
|
||||
|
||||
with pytest.raises(ValueError, match="mapped OCR failure") as caught:
|
||||
await wrapped()
|
||||
|
||||
assert caught.value is mapped_error
|
||||
assert native_completion.failures == [mapped_error, mapped_error]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_ocr_wrapper_reports_metadata_failure_without_success(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
native_completion: Final = RecordingCompletion()
|
||||
response: Final = object()
|
||||
metadata_error: Final = ValueError("metadata failure")
|
||||
|
||||
async def aocr(**kwargs: object) -> object:
|
||||
completion = kwargs.get("_litellm_call_completion")
|
||||
assert isinstance(completion, CallCompletion)
|
||||
assert completion.attach(native_completion)
|
||||
return response
|
||||
|
||||
def fail_metadata(**kwargs: object) -> None:
|
||||
raise metadata_error
|
||||
|
||||
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {})))
|
||||
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
|
||||
monkeypatch.setattr("litellm.utils.update_response_metadata", fail_metadata)
|
||||
wrapped: Final = client(aocr)
|
||||
|
||||
with pytest.raises(ValueError, match="metadata failure") as caught:
|
||||
await wrapped()
|
||||
|
||||
assert caught.value is metadata_error
|
||||
assert native_completion.successes == []
|
||||
assert native_completion.failures == [metadata_error, metadata_error]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapper_completion_stays_separate_from_provider_options(
|
||||
monkeypatch: pytest.MonkeyPatch, asynchronous: bool
|
||||
) -> None:
|
||||
native_completion: Final = RecordingCompletion()
|
||||
response: Final = OCRResponse(model="mistral-ocr-latest", pages=[])
|
||||
metadata: Final = {"request": "shared"}
|
||||
pages: Final = [0, 2]
|
||||
|
||||
def ocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> OCRResponse:
|
||||
assert kwargs["metadata"] is metadata
|
||||
assert kwargs["pages"] is pages
|
||||
assert "_litellm_call_completion" not in kwargs
|
||||
assert _litellm_call_completion.attach(native_completion)
|
||||
return response
|
||||
|
||||
async def aocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> OCRResponse:
|
||||
return ocr(_litellm_call_completion=_litellm_call_completion, **kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.utils.function_setup",
|
||||
MagicMock(return_value=(MagicMock(), {"metadata": metadata, "pages": pages})),
|
||||
)
|
||||
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
|
||||
wrapped: Final = client(aocr if asynchronous else ocr)
|
||||
arguments: Final = {
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"document": {"type": "document_url", "document_url": "https://example.com/doc.pdf"},
|
||||
"api_key": "test-key",
|
||||
"metadata": metadata,
|
||||
"pages": pages,
|
||||
}
|
||||
|
||||
result: Final = await wrapped(**arguments) if asynchronous else wrapped(**arguments)
|
||||
|
||||
assert result is response
|
||||
assert native_completion.successes == [response]
|
||||
assert native_completion.failures == []
|
||||
assert arguments["metadata"] is metadata
|
||||
assert "_litellm_call_completion" not in arguments
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("asynchronous", "exit_path"),
|
||||
[(False, "success"), (True, "success"), (False, "callback_error"), (True, "callback_error"), (True, "cancelled")],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_wrapper_releases_completion_resources_on_every_exit(
|
||||
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, exit_path: str
|
||||
) -> None:
|
||||
retained: Final[list[tuple[CallCompletion, weakref.ReferenceType[object]]]] = []
|
||||
callback_error: Final = RuntimeError("failure callback failed")
|
||||
implementation: Final = MagicMock(spec=RecordingCompletion)
|
||||
implementation.failure.side_effect = callback_error
|
||||
implementation.async_failure = AsyncMock()
|
||||
response: Final = object()
|
||||
|
||||
def ocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> object:
|
||||
retained.append((_litellm_call_completion, weakref.ref(_litellm_call_completion.python_implementation)))
|
||||
assert _litellm_call_completion.attach(implementation)
|
||||
if exit_path == "cancelled":
|
||||
raise asyncio.CancelledError
|
||||
if exit_path == "callback_error":
|
||||
raise ValueError("provider failed")
|
||||
return response
|
||||
|
||||
async def aocr(*, _litellm_call_completion: CallCompletion, **kwargs: object) -> object:
|
||||
return ocr(_litellm_call_completion=_litellm_call_completion, **kwargs)
|
||||
|
||||
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {})))
|
||||
monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock())
|
||||
wrapped: Final = client(aocr if asynchronous else ocr)
|
||||
|
||||
if exit_path == "success":
|
||||
result: Final = await wrapped() if asynchronous else wrapped()
|
||||
assert result is response
|
||||
assert implementation.success.call_args.args[0] is response
|
||||
elif exit_path == "cancelled":
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await wrapped()
|
||||
implementation.success.assert_not_called()
|
||||
implementation.failure.assert_not_called()
|
||||
else:
|
||||
with pytest.raises(RuntimeError, match="failure callback failed") as caught:
|
||||
await wrapped() if asynchronous else wrapped()
|
||||
assert caught.value is callback_error
|
||||
implementation.async_failure.assert_not_called()
|
||||
|
||||
assert len(retained) == 1
|
||||
assert retained[0][1]() is None
|
||||
|
|
@ -1,9 +1,5 @@
|
|||
import base64
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import AzureDocumentIntelligenceOCRConfig
|
||||
|
|
@ -12,27 +8,6 @@ from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
|
|||
MODEL = "azure_ai/Cohere-parse-v5"
|
||||
API_BASE = "https://resource.services.ai.azure.com"
|
||||
PARSE_URL = f"{API_BASE}/providers/cohere/v2/parse"
|
||||
IMAGE_URL = "https://example.com/receipt.png"
|
||||
PNG_BYTES = base64.b64decode(
|
||||
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
|
||||
)
|
||||
PNG_DATA_URI = f"data:image/png;base64,{base64.b64encode(PNG_BYTES).decode()}"
|
||||
|
||||
|
||||
def _parse_response() -> dict:
|
||||
return {
|
||||
"id": "882bf973-9dfa-4d02-9d30-709247008efd",
|
||||
"pages": [{"index": 0, "type": "markdown", "markdown": {"content": "# Receipt\n\nTotal Due: $4.00"}}],
|
||||
"meta": {"api_version": {"version": "2"}, "billed_units": {"pages": 1}},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def disable_aiohttp_transport(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
yield
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -95,90 +70,3 @@ def test_validate_environment_requires_api_base(monkeypatch) -> None:
|
|||
|
||||
with pytest.raises(ValueError, match="AZURE_AI_API_BASE"):
|
||||
AzureAICohereParseConfig().validate_environment(headers={}, model="Cohere-parse-v5", api_key="key")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_inlines_remote_image_and_posts_to_foundry(disable_aiohttp_transport, respx_mock):
|
||||
respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"})
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
response = await litellm.aocr(
|
||||
model=MODEL,
|
||||
document={"type": "image_url", "image_url": IMAGE_URL},
|
||||
api_base=API_BASE,
|
||||
api_key="azure-key",
|
||||
)
|
||||
|
||||
request = route.calls.last.request
|
||||
assert request.headers["Authorization"] == "Bearer azure-key"
|
||||
assert json.loads(request.content) == {
|
||||
"model": "Cohere-parse-v5",
|
||||
"document": {"type": "image_url", "image_url": PNG_DATA_URI},
|
||||
"output_format": "markdown",
|
||||
}
|
||||
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
|
||||
assert response.usage_info.pages_processed == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_passes_data_uri_through_without_fetching(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
await litellm.aocr(
|
||||
model=MODEL,
|
||||
document={"type": "image_url", "image_url": PNG_DATA_URI},
|
||||
api_base=API_BASE,
|
||||
api_key="azure-key",
|
||||
output_format="blocks",
|
||||
)
|
||||
|
||||
body = json.loads(route.calls.last.request.content)
|
||||
assert body["document"]["image_url"] == PNG_DATA_URI
|
||||
assert body["output_format"] == "blocks"
|
||||
|
||||
|
||||
def test_ocr_sync_inlines_remote_image(respx_mock):
|
||||
respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"})
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
response = litellm.ocr(
|
||||
model=MODEL,
|
||||
document={"type": "image_url", "image_url": IMAGE_URL},
|
||||
api_base=API_BASE,
|
||||
api_key="azure-key",
|
||||
)
|
||||
|
||||
assert json.loads(route.calls.last.request.content)["document"]["image_url"] == PNG_DATA_URI
|
||||
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_rejects_pdf_before_calling_foundry(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info:
|
||||
await litellm.aocr(
|
||||
model=MODEL,
|
||||
document={"type": "document_url", "document_url": "https://example.com/doc.pdf"},
|
||||
api_base=API_BASE,
|
||||
api_key="azure-key",
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == "azure_ai"
|
||||
assert not route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ahealth_check_ocr_sends_an_image_to_the_foundry_cohere_parse_deployment(
|
||||
disable_aiohttp_transport, respx_mock
|
||||
):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
result = await litellm.ahealth_check(
|
||||
model_params={"model": MODEL, "api_base": API_BASE, "api_key": "test-key"}, mode="ocr"
|
||||
)
|
||||
|
||||
document = json.loads(route.calls.last.request.content)["document"]
|
||||
assert document["type"] == "image_url"
|
||||
assert document["image_url"].startswith("data:image/png;base64,")
|
||||
assert "error" not in result
|
||||
|
|
|
|||
|
|
@ -1,8 +1,11 @@
|
|||
import json
|
||||
from typing import Final
|
||||
from unittest.mock import Mock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
|
||||
|
||||
PARSE_URL = "https://api.cohere.com/v2/parse"
|
||||
MODEL = "cohere/parse-v5.0"
|
||||
|
|
@ -57,173 +60,38 @@ def _blocks_response() -> dict:
|
|||
}
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def disable_aiohttp_transport(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
yield
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_sends_markdown_parse_request_and_normalizes_pages(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
|
||||
|
||||
request = route.calls.last.request
|
||||
assert request.headers["Authorization"] == "Bearer test-key"
|
||||
assert json.loads(request.content) == {
|
||||
"model": "parse-v5.0",
|
||||
"document": IMAGE_DOCUMENT,
|
||||
"output_format": "markdown",
|
||||
}
|
||||
assert response.object == "ocr"
|
||||
assert [page.index for page in response.pages] == [0, 1]
|
||||
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
|
||||
assert response.pages[1].markdown == "Page two"
|
||||
assert response.pages[1].images is None
|
||||
image = response.pages[0].images[0]
|
||||
assert image.bbox == BOUNDING_BOX
|
||||
assert image.model_extra["description"] == "A parking receipt"
|
||||
assert image.model_extra["bounding_box_normalized"]["bottom_right_x"] == 1
|
||||
assert response.usage_info.pages_processed == 2
|
||||
assert response.get_provider_native_response() is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_usage_prefers_billed_units_over_page_count(disable_aiohttp_transport, respx_mock):
|
||||
respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=3))
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
|
||||
|
||||
assert response.usage_info.pages_processed == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_usage_falls_back_to_page_count_without_meta(disable_aiohttp_transport, respx_mock):
|
||||
respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=None))
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
|
||||
|
||||
assert response.usage_info.pages_processed == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_blocks_output_format_forwards_param_and_keeps_blocks(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_blocks_response())
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="blocks")
|
||||
|
||||
assert json.loads(route.calls.last.request.content)["output_format"] == "blocks"
|
||||
assert response.pages[0].markdown == ""
|
||||
assert response.pages[0].model_extra["blocks"] == [{"type": "text", "text": "Total Due: $4.00"}]
|
||||
assert response.usage_info.pages_processed == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_native_format_carries_provider_payload(disable_aiohttp_transport, respx_mock):
|
||||
payload = _markdown_response()
|
||||
route = respx_mock.post(PARSE_URL).respond(json=payload)
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", req_format="native")
|
||||
|
||||
assert "req_format" not in json.loads(route.calls.last.request.content)
|
||||
assert response.get_provider_native_response() == payload
|
||||
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_rejects_unknown_output_format_before_calling_provider(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="Invalid `output_format`: 'html'") as exc_info:
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="html")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert not route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"document",
|
||||
[
|
||||
{"type": "document_url", "document_url": "https://example.com/doc.pdf"},
|
||||
{"type": "image_url", "image_url": "data:application/pdf;base64,JVBERi0="},
|
||||
{"type": "image_url", "image_url": ""},
|
||||
],
|
||||
)
|
||||
async def test_aocr_rejects_non_image_documents_before_calling_provider(
|
||||
disable_aiohttp_transport, respx_mock, document
|
||||
):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info:
|
||||
await litellm.aocr(model=MODEL, document=document, api_key="test-key")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert not route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"api_base, expected_url",
|
||||
[
|
||||
("https://gateway.example.com", "https://gateway.example.com/v2/parse"),
|
||||
("https://gateway.example.com/cohere/", "https://gateway.example.com/cohere/v2/parse"),
|
||||
("https://gateway.example.com/v2", "https://gateway.example.com/v2/parse"),
|
||||
("https://gateway.example.com/v2/parse", "https://gateway.example.com/v2/parse"),
|
||||
],
|
||||
)
|
||||
async def test_aocr_posts_to_api_base_variants(disable_aiohttp_transport, respx_mock, api_base, expected_url):
|
||||
route = respx_mock.post(expected_url).respond(json=_markdown_response())
|
||||
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", api_base=api_base)
|
||||
|
||||
assert route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_surfaces_provider_error_with_its_status_and_message(disable_aiohttp_transport, respx_mock):
|
||||
respx_mock.post(PARSE_URL).respond(
|
||||
status_code=400, json={"id": "83b0d95e", "message": "output_format must be `blocks` or `markdown`"}
|
||||
@pytest.mark.parametrize("output_format", ["markdown", "blocks"])
|
||||
def test_transform_cohere_request_filters_options(output_format: str) -> None:
|
||||
config: Final = CohereParseConfig()
|
||||
params: Final = config.map_ocr_params(
|
||||
{"output_format": output_format, "req_format": "native", "unknown": True}, {}, "parse-v5.0"
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="output_format must be") as exc_info:
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
request: Final = config.transform_ocr_request("parse-v5.0", IMAGE_DOCUMENT, params, {})
|
||||
assert request.data == {"model": "parse-v5.0", "document": IMAGE_DOCUMENT, "output_format": output_format}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_reads_api_key_from_environment(disable_aiohttp_transport, respx_mock, monkeypatch):
|
||||
monkeypatch.setenv("COHERE_API_KEY", "env-key")
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT)
|
||||
|
||||
assert route.calls.last.request.headers["Authorization"] == "Bearer env-key"
|
||||
@pytest.mark.parametrize("native", [False, True])
|
||||
def test_transform_cohere_response_keeps_images_and_native_payload(native: bool) -> None:
|
||||
payload: Final = _markdown_response(3)
|
||||
response: Final = CohereParseConfig().transform_ocr_response(
|
||||
"parse-v5.0", httpx.Response(200, json=payload), Mock(), {"req_format": "native" if native else "litellm"}
|
||||
)
|
||||
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
|
||||
assert response.pages[0].images[0].bbox == BOUNDING_BOX
|
||||
assert response.pages[0].images[0].model_extra["description"] == "A parking receipt"
|
||||
assert response.pages[1].images is None
|
||||
assert response.usage_info.pages_processed == 3
|
||||
assert response.get_provider_native_response() == (payload if native else None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_without_api_key_names_the_env_var(disable_aiohttp_transport, respx_mock, monkeypatch):
|
||||
monkeypatch.delenv("COHERE_API_KEY", raising=False)
|
||||
monkeypatch.setattr(litellm, "cohere_key", None)
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
with pytest.raises(Exception, match="Missing COHERE_API_KEY"):
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT)
|
||||
|
||||
assert not route.called
|
||||
def test_transform_cohere_blocks() -> None:
|
||||
response: Final = CohereParseConfig().transform_ocr_response(
|
||||
"parse-v5.0", httpx.Response(200, json=_blocks_response()), Mock()
|
||||
)
|
||||
assert response.pages[0].model_extra["blocks"] == [{"type": "text", "text": "Total Due: $4.00"}]
|
||||
assert response.pages[0].markdown == ""
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ahealth_check_ocr_sends_an_image_cohere_parse_accepts(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
result = await litellm.ahealth_check(model_params={"model": MODEL, "api_key": "test-key"}, mode="ocr")
|
||||
|
||||
document = json.loads(route.calls.last.request.content)["document"]
|
||||
assert document["type"] == "image_url"
|
||||
assert document["image_url"].startswith("data:image/png;base64,")
|
||||
assert "error" not in result
|
||||
def test_transform_cohere_rejects_unsupported_output_format() -> None:
|
||||
with pytest.raises(litellm.UnsupportedParamsError, match="output_format"):
|
||||
CohereParseConfig().map_ocr_params({"output_format": "html"}, {}, "parse-v5.0")
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ Streaming: CSW.__anext__ stores args on logging_obj at stream end.
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Any
|
||||
from typing import Any, Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -297,6 +297,38 @@ async def test_no_flag_fires_create_task_normally():
|
|||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.parametrize("call_type", ["ocr", "aocr", "completion", "acompletion", "embedding", "responses"])
|
||||
@pytest.mark.parametrize("exception_raised", [False, True])
|
||||
def test_native_pending_logging_is_released_only_for_ocr(call_type: str, exception_raised: bool) -> None:
|
||||
pending: Final = MagicMock()
|
||||
enqueue: Final = MagicMock()
|
||||
logger: Final = MagicMock(
|
||||
call_type=call_type,
|
||||
_native_pending_logging=pending,
|
||||
_enqueue_deferred_logging=enqueue,
|
||||
)
|
||||
|
||||
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(
|
||||
logging_obj=logger,
|
||||
exception_raised=exception_raised,
|
||||
)
|
||||
ProxyBaseLLMRequestProcessing._flush_deferred_async_logging(
|
||||
logging_obj=logger,
|
||||
exception_raised=exception_raised,
|
||||
)
|
||||
|
||||
if call_type in ("ocr", "aocr"):
|
||||
pending.release.assert_called_once_with(not exception_raised)
|
||||
assert logger._native_pending_logging is None
|
||||
else:
|
||||
pending.release.assert_not_called()
|
||||
assert logger._native_pending_logging is pending
|
||||
if exception_raised:
|
||||
enqueue.assert_not_called()
|
||||
else:
|
||||
enqueue.assert_called_once_with()
|
||||
|
||||
|
||||
def test_flush_deferred_async_logging_fires_on_success():
|
||||
"""
|
||||
Happy path: with no exception, the production flush helper invokes the
|
||||
|
|
|
|||
|
|
@ -10,11 +10,11 @@ from litellm.rust_bridge.ocr import LiteLLMOcrRequest
|
|||
from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled,available", [(True, False), (False, True), (False, False)])
|
||||
def test_public_selection_requires_supported_native_ocr(enabled: bool, available: bool) -> None:
|
||||
@pytest.mark.parametrize("enabled", [True, False])
|
||||
def test_public_selection_requires_available_native_ocr(enabled: bool) -> None:
|
||||
native: Final = Mock(side_effect=AssertionError("must not admit"))
|
||||
litellm.rust(enabled)
|
||||
NATIVE_OCR_LIFECYCLE.override(native if available else None)
|
||||
NATIVE_OCR_LIFECYCLE.override(None)
|
||||
try:
|
||||
with pytest.raises(RuntimeError, match="Rust OCR is unavailable or does not support this request"):
|
||||
litellm.ocr("mistral/mistral-ocr-latest", {"type": "document_url", "document_url": "https://example.com"})
|
||||
|
|
@ -96,7 +96,7 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs()
|
|||
assert "timeout" not in captured[0]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled", [False, True], ids=["legacy", "native"])
|
||||
@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"])
|
||||
def test_public_duplicate_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None:
|
||||
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
|
||||
document: Final = {"type": "document_url", "document_url": "https://example.com"}
|
||||
|
|
@ -111,7 +111,7 @@ def test_public_duplicate_argument_error_does_not_depend_on_native_selection(ena
|
|||
assert native.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled", [False, True], ids=["legacy", "native"])
|
||||
@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"])
|
||||
def test_public_missing_required_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None:
|
||||
native: Final = Mock(side_effect=AssertionError("binding errors precede admission"))
|
||||
litellm.rust(enabled)
|
||||
|
|
@ -123,3 +123,27 @@ def test_public_missing_required_argument_error_does_not_depend_on_native_select
|
|||
NATIVE_OCR_LIFECYCLE.reset()
|
||||
litellm.rust(None)
|
||||
assert native.call_count == 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize("enabled", [False, True, None])
|
||||
async def test_public_ocr_ignores_rust_flag(
|
||||
monkeypatch: pytest.MonkeyPatch, asynchronous: bool, enabled: bool | None
|
||||
) -> None:
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
monkeypatch.setenv("LITELLM_RUST", "0")
|
||||
response: Final = OCRResponse(pages=[], model="mistral-ocr-latest")
|
||||
native: Final = AsyncMock(return_value=response) if asynchronous else Mock(return_value=response)
|
||||
litellm.rust(enabled)
|
||||
NATIVE_OCR_LIFECYCLE.override(native)
|
||||
try:
|
||||
if asynchronous:
|
||||
assert await litellm.aocr("mistral/mistral-ocr-latest", {}) is response
|
||||
else:
|
||||
assert litellm.ocr("mistral/mistral-ocr-latest", {}) is response
|
||||
assert native.call_count == 1
|
||||
finally:
|
||||
NATIVE_OCR_LIFECYCLE.reset()
|
||||
litellm.rust(None)
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import contextlib
|
||||
import contextvars
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
|
@ -7,6 +8,7 @@ import queue
|
|||
import threading
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from collections.abc import Iterator
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
|
|
@ -60,6 +62,36 @@ from litellm.utils import (
|
|||
# Adds the parent directory to the system path
|
||||
|
||||
|
||||
def test_non_ocr_wrapper_preserves_logging_executor_and_context(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
marker: Final = contextvars.ContextVar("non-ocr-logging-context", default="missing")
|
||||
token: Final = marker.set("caller-context")
|
||||
caller_thread: Final = threading.get_ident()
|
||||
response: Final = object()
|
||||
logger: Final = MagicMock()
|
||||
observed: Final = queue.Queue[tuple[object, str, int]]()
|
||||
|
||||
def record_success(result: object, start_time: datetime, end_time: datetime) -> None:
|
||||
observed.put((result, marker.get(), threading.get_ident()))
|
||||
|
||||
def embedding(**kwargs: object) -> object:
|
||||
return response
|
||||
|
||||
logger.success_handler.side_effect = record_success
|
||||
monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(logger, {})))
|
||||
try:
|
||||
with ThreadPoolExecutor(max_workers=1) as executor:
|
||||
monkeypatch.setattr("litellm.utils.executor", executor)
|
||||
result: Final = client(embedding)()
|
||||
logged_response, context, worker_thread = observed.get_nowait()
|
||||
assert result is response
|
||||
assert logged_response is response
|
||||
assert context == "caller-context"
|
||||
assert worker_thread != caller_thread
|
||||
assert observed.empty()
|
||||
finally:
|
||||
marker.reset(token)
|
||||
|
||||
|
||||
def test_cloudflare_model_info_includes_rpm(local_model_cost_map: None) -> None:
|
||||
assert litellm.get_model_info("cloudflare/@cf/meta/llama-3.1-8b-instruct-fp8")["rpm"] == 300
|
||||
assert litellm.get_model_info("cloudflare/@cf/moonshotai/kimi-k2.6")["rpm"] == 20
|
||||
|
|
|
|||
141
tests/test_litellm_rust/ocr/test_cohere.py
Normal file
141
tests/test_litellm_rust/ocr/test_cohere.py
Normal file
|
|
@ -0,0 +1,141 @@
|
|||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec
|
||||
|
||||
pytestmark = pytest.mark.requires_rust_extension
|
||||
MODELS: Final = ("cohere/parse-v5.0", "azure_ai/Cohere-parse-v5.0")
|
||||
IMAGE: Final = {"type": "image_url", "image_url": "data:image/png;base64,YWJj"}
|
||||
BOX: Final = {"top_left_x": 0, "top_left_y": 0, "bottom_right_x": 32, "bottom_right_y": 32}
|
||||
PAYLOAD: Final = {
|
||||
"pages": [
|
||||
{
|
||||
"index": 4,
|
||||
"markdown": {"content": "receipt", "images": [{"id": "image", "bounding_box": BOX, "description": "scan"}]},
|
||||
},
|
||||
{"markdown": {"content": "page two"}},
|
||||
],
|
||||
"meta": {"billed_units": {"pages": 3}},
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
async def test_public_cohere_request_and_normalization(
|
||||
recording_server: RecordingServer, model: str, asynchronous: bool
|
||||
) -> None:
|
||||
recording_server.enqueue(ResponseSpec(body=PAYLOAD))
|
||||
args: Final = {
|
||||
"model": model,
|
||||
"document": IMAGE,
|
||||
"api_base": recording_server.base_url,
|
||||
"api_key": "test-key",
|
||||
"req_format": "native",
|
||||
"unrecognized": True,
|
||||
}
|
||||
response: Final = await litellm.aocr(**args) if asynchronous else litellm.ocr(**args)
|
||||
request: Final = recording_server.requests[0]
|
||||
assert request.path == ("/providers/cohere/v2/parse" if model.startswith("azure_ai/") else "/v2/parse")
|
||||
assert request.headers["authorization"] == "Bearer test-key"
|
||||
assert request.body == {"model": model.split("/", 1)[1], "document": IMAGE, "output_format": "markdown"}
|
||||
assert [page.index for page in response.pages] == [4, 1]
|
||||
assert response.pages[0].markdown == "receipt"
|
||||
assert response.pages[0].images[0].bbox == BOX
|
||||
assert response.pages[0].images[0].model_extra["description"] == "scan"
|
||||
assert response.pages[1].images is None
|
||||
assert response.usage_info.pages_processed == 3
|
||||
assert response.get_provider_native_response() == PAYLOAD
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
async def test_public_cohere_blocks_and_usage_fallback(recording_server: RecordingServer, model: str) -> None:
|
||||
blocks: Final = [{"type": "text", "text": "total"}]
|
||||
recording_server.enqueue(ResponseSpec(body={"pages": [{"blocks": blocks}]}))
|
||||
response: Final = await litellm.aocr(
|
||||
model=model, document=IMAGE, api_base=recording_server.base_url, api_key="test-key", output_format="blocks"
|
||||
)
|
||||
assert recording_server.requests[0].body["output_format"] == "blocks"
|
||||
assert response.pages[0].model_extra["blocks"] == blocks
|
||||
assert response.pages[0].markdown == ""
|
||||
assert response.usage_info.pages_processed == 1
|
||||
assert response.get_provider_native_response() is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
@pytest.mark.parametrize(
|
||||
"document",
|
||||
[
|
||||
{"type": "document_url", "document_url": "https://example.com/file.pdf"},
|
||||
{"type": "image_url", "image_url": "data:application/pdf;base64,YQ=="},
|
||||
{"type": "image_url", "image_url": ""},
|
||||
],
|
||||
)
|
||||
async def test_public_cohere_rejects_non_images_before_network(
|
||||
recording_server: RecordingServer, model: str, document: dict[str, str]
|
||||
) -> None:
|
||||
recording_server.expected_requests = 0
|
||||
with pytest.raises(litellm.BadRequestError, match="only accepts `image_url`"):
|
||||
await litellm.aocr(model=model, document=document, api_base=recording_server.base_url, api_key="test-key")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
async def test_public_cohere_rejects_unknown_format(recording_server: RecordingServer, model: str) -> None:
|
||||
recording_server.expected_requests = 0
|
||||
with pytest.raises(litellm.BadRequestError, match="output_format"):
|
||||
await litellm.aocr(
|
||||
model=model, document=IMAGE, api_base=recording_server.base_url, api_key="test-key", output_format="html"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
async def test_public_cohere_provider_failure(recording_server: RecordingServer, model: str) -> None:
|
||||
recording_server.enqueue(ResponseSpec(status=400, body={"message": "output_format must be blocks or markdown"}))
|
||||
with pytest.raises(litellm.BadRequestError, match="output_format must be") as caught:
|
||||
await litellm.aocr(model=model, document=IMAGE, api_base=recording_server.base_url, api_key="test-key")
|
||||
assert caught.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("model", MODELS)
|
||||
async def test_public_cohere_health_check(recording_server: RecordingServer, model: str) -> None:
|
||||
recording_server.enqueue(ResponseSpec(body=PAYLOAD))
|
||||
response: Final = await litellm.ahealth_check(
|
||||
model_params={"model": model, "api_key": "test-key", "api_base": recording_server.base_url}, mode="ocr"
|
||||
)
|
||||
assert "error" not in response
|
||||
assert recording_server.requests[0].body["document"]["image_url"].startswith("data:image/png;base64,")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("suffix", ["", "/cohere/", "/v2", "/v2/parse"])
|
||||
async def test_public_cohere_url_variants(recording_server: RecordingServer, suffix: str) -> None:
|
||||
recording_server.enqueue(ResponseSpec(body=PAYLOAD))
|
||||
await litellm.aocr(model=MODELS[0], document=IMAGE, api_base=recording_server.base_url + suffix, api_key="test-key")
|
||||
assert recording_server.requests[0].path == ("/cohere/v2/parse" if suffix == "/cohere/" else "/v2/parse")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_public_cohere_environment_key_and_remote_url(
|
||||
recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setenv("COHERE_API_KEY", "env-key")
|
||||
recording_server.enqueue(ResponseSpec(body=PAYLOAD))
|
||||
document: Final = {"type": "image_url", "image_url": "https://example.com/receipt.png"}
|
||||
await litellm.aocr(model=MODELS[0], document=document, api_base=recording_server.base_url)
|
||||
assert recording_server.requests[0].headers["authorization"] == "Bearer env-key"
|
||||
assert recording_server.requests[0].body["document"] == document
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_public_cohere_missing_key(recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("COHERE_API_KEY", raising=False)
|
||||
recording_server.expected_requests = 0
|
||||
with pytest.raises(Exception, match="Missing COHERE_API_KEY"):
|
||||
await litellm.aocr(model=MODELS[0], document=IMAGE, api_base=recording_server.base_url)
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
import sys
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
|
@ -16,8 +17,26 @@ def ocr_server(recording_server: RecordingServer) -> RecordingServer:
|
|||
return recording_server
|
||||
|
||||
|
||||
def test_public_ocr_uses_native_route_when_enabled(ocr_server: RecordingServer) -> None:
|
||||
litellm.rust(True)
|
||||
def test_native_ocr_rejects_disabled_gil_before_provider_call(
|
||||
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
ocr_server.expected_requests = 0
|
||||
monkeypatch.setattr(sys, "_is_gil_enabled", lambda: False, raising=False)
|
||||
|
||||
with pytest.raises(RuntimeError, match="native OCR requires the Python GIL"):
|
||||
litellm.ocr(
|
||||
model=OCR_MODEL,
|
||||
document=OCR_DOCUMENT,
|
||||
api_key="test-key",
|
||||
api_base=ocr_server.base_url,
|
||||
)
|
||||
|
||||
assert not ocr_server.requests
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled", [False, True, None])
|
||||
def test_public_ocr_uses_native_route_independently_of_flag(ocr_server: RecordingServer, enabled: bool | None) -> None:
|
||||
litellm.rust(enabled)
|
||||
response: Final = litellm.ocr(
|
||||
model=OCR_MODEL,
|
||||
document=OCR_DOCUMENT,
|
||||
|
|
@ -29,18 +48,3 @@ def test_public_ocr_uses_native_route_when_enabled(ocr_server: RecordingServer)
|
|||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert len(ocr_server.requests) == 1
|
||||
assert not ocr_server.requests[0].headers.get("user-agent", "").startswith("python-httpx")
|
||||
|
||||
|
||||
def test_public_ocr_fails_before_network_when_native_is_disabled(ocr_server: RecordingServer) -> None:
|
||||
litellm.rust(False)
|
||||
ocr_server.expected_requests = 0
|
||||
|
||||
with pytest.raises(RuntimeError, match="Rust OCR is unavailable"):
|
||||
litellm.ocr(
|
||||
model=OCR_MODEL,
|
||||
document=OCR_DOCUMENT,
|
||||
api_key="test-key",
|
||||
api_base=ocr_server.base_url,
|
||||
)
|
||||
|
||||
assert ocr_server.requests == []
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
import datetime
|
||||
import gc
|
||||
import json
|
||||
import sys
|
||||
import threading
|
||||
import weakref
|
||||
|
|
@ -711,7 +712,9 @@ async def test_reducto_lifecycle_retains_upload_parse_and_post_call_boundaries(
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_document_intelligence_post_call_runs_before_polling(ocr_server: RecordingServer) -> None:
|
||||
async def test_document_intelligence_post_call_observes_submission_and_final_result(
|
||||
ocr_server: RecordingServer,
|
||||
) -> None:
|
||||
ocr_server.expected_requests = 2
|
||||
ocr_server.enqueue(
|
||||
ResponseSpec(
|
||||
|
|
@ -725,7 +728,7 @@ async def test_document_intelligence_post_call_runs_before_polling(ocr_server: R
|
|||
|
||||
class Observe(Logging):
|
||||
def post_call(self, *args, **kwargs):
|
||||
boundaries.append(tuple(request.method for request in ocr_server.requests))
|
||||
boundaries.append((tuple(request.method for request in ocr_server.requests), kwargs["original_response"]))
|
||||
return super().post_call(*args, **kwargs)
|
||||
|
||||
logger: Final = Observe(
|
||||
|
|
@ -740,7 +743,9 @@ async def test_document_intelligence_post_call_runs_before_polling(ocr_server: R
|
|||
response: Final = await call_aocr(
|
||||
ocr_server, model="azure_ai/doc-intelligence/prebuilt-read", litellm_logging_obj=logger
|
||||
)
|
||||
assert boundaries == [("POST",)]
|
||||
assert [methods for methods, _ in boundaries] == [("POST",), ("POST", "GET")]
|
||||
assert json.loads(boundaries[0][1])["status"] == "running"
|
||||
assert json.loads(boundaries[1][1])["status"] == "succeeded"
|
||||
assert [request.method for request in ocr_server.requests] == ["POST", "GET"]
|
||||
assert ocr_server.requests[1].path == "/operations/1"
|
||||
assert response.pages == []
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue