mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
refactor(ocr): dispatch through Rust before Python preparation
This commit is contained in:
parent
0d16842e67
commit
c13f0aeedc
9 changed files with 491 additions and 358 deletions
|
|
@ -12,7 +12,7 @@ use litellm_core::chat_completions::{
|
|||
use litellm_core::error::CoreError;
|
||||
use litellm_core::messages::messages as run_messages;
|
||||
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
|
||||
use litellm_runtime::ocr::{OcrRequest, ocr as run_ocr};
|
||||
use litellm_runtime::ocr::{OcrRequest, ocr as run_ocr, ocr_decline_reason};
|
||||
use pyo3::exceptions::{PyRuntimeError, PyValueError};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyAny, PyDict};
|
||||
|
|
@ -208,6 +208,29 @@ fn marshal_inputs(
|
|||
Ok((document, extra_headers, optional_params, timeout))
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (model, custom_llm_provider=None, optional_params=None))]
|
||||
fn ocr_decline(
|
||||
py: Python<'_>,
|
||||
model: String,
|
||||
custom_llm_provider: Option<String>,
|
||||
optional_params: Option<Py<PyAny>>,
|
||||
) -> PyResult<Option<String>> {
|
||||
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
|
||||
Ok(ocr_decline_reason(OcrRequest {
|
||||
model: &model,
|
||||
document: Value::Null,
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
extra_headers: None,
|
||||
optional_params,
|
||||
timeout: None,
|
||||
litellm_call_id: Some("ocr-decline"),
|
||||
})
|
||||
.map(|reason| reason.to_string()))
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
|
|
@ -613,6 +636,7 @@ fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
|||
let py = module.py();
|
||||
module.add_function(wrap_pyfunction!(ocr, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(aocr, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(ocr_decline, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(transcription, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(atranscription, module)?)?;
|
||||
module.add_function(wrap_pyfunction!(messages, module)?)?;
|
||||
|
|
|
|||
|
|
@ -14,8 +14,8 @@ mod prepare;
|
|||
mod provider;
|
||||
mod types;
|
||||
|
||||
pub use prepare::{prepare_ocr_request, prepare_provider_request};
|
||||
pub use types::{OcrRequest, PreparedOcrRequest, ProviderOcrRequest};
|
||||
pub use prepare::{ocr_decline_reason, prepare_ocr_request, prepare_provider_request};
|
||||
pub use types::{OcrDeclineReason, OcrRequest, PreparedOcrRequest, ProviderOcrRequest};
|
||||
|
||||
use handler::execute_ocr_provider_call;
|
||||
|
||||
|
|
|
|||
|
|
@ -4,21 +4,35 @@ use std::time::{SystemTime, UNIX_EPOCH};
|
|||
use litellm_core::CoreResult;
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::common_utils::{convert_document_url_to_data_uri, has_header, string_headers};
|
||||
use super::provider::ocr_provider_config;
|
||||
use super::types::{OcrAuthStrategy, OcrRequest, PreparedOcrRequest, ProviderOcrRequest};
|
||||
use super::types::{
|
||||
OcrAuthStrategy, OcrDeclineReason, OcrRequest, PreparedOcrRequest, ProviderOcrRequest,
|
||||
};
|
||||
|
||||
pub fn ocr_decline_reason(request: OcrRequest<'_>) -> Option<OcrDeclineReason> {
|
||||
if request
|
||||
.optional_params
|
||||
.get("req_format")
|
||||
.and_then(Value::as_str)
|
||||
== Some("native")
|
||||
{
|
||||
return Some(OcrDeclineReason::NativeRequestFormat);
|
||||
}
|
||||
let provider = select_ocr_provider(request.model, request.custom_llm_provider);
|
||||
ocr_provider_config(provider.custom_llm_provider, provider.model)
|
||||
.is_none()
|
||||
.then_some(OcrDeclineReason::UnsupportedProvider)
|
||||
}
|
||||
|
||||
pub fn prepare_ocr_request(request: OcrRequest<'_>) -> PreparedOcrRequest {
|
||||
let litellm_call_id = request
|
||||
.litellm_call_id
|
||||
.map(str::to_string)
|
||||
.unwrap_or_else(new_ocr_call_id);
|
||||
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
|
||||
.unwrap_or(CustomLlmProvider {
|
||||
model: request.model,
|
||||
custom_llm_provider: "mistral",
|
||||
});
|
||||
let provider_info = select_ocr_provider(request.model, request.custom_llm_provider);
|
||||
|
||||
PreparedOcrRequest {
|
||||
model: provider_info.model.to_string(),
|
||||
|
|
@ -33,6 +47,16 @@ pub fn prepare_ocr_request(request: OcrRequest<'_>) -> PreparedOcrRequest {
|
|||
}
|
||||
}
|
||||
|
||||
fn select_ocr_provider<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
) -> CustomLlmProvider<'a> {
|
||||
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
|
||||
model,
|
||||
custom_llm_provider: "mistral",
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn prepare_provider_request(
|
||||
request: PreparedOcrRequest,
|
||||
) -> CoreResult<ProviderOcrRequest> {
|
||||
|
|
@ -102,3 +126,53 @@ fn new_ocr_call_id() -> String {
|
|||
.unwrap_or(0);
|
||||
format!("ocr-{timestamp}-{sequence}")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::{Map, json};
|
||||
|
||||
use super::*;
|
||||
|
||||
fn request<'a>(
|
||||
model: &'a str,
|
||||
custom_llm_provider: Option<&'a str>,
|
||||
optional_params: Map<String, Value>,
|
||||
) -> OcrRequest<'a> {
|
||||
OcrRequest {
|
||||
model,
|
||||
document: json!({}),
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider,
|
||||
extra_headers: None,
|
||||
optional_params,
|
||||
timeout: None,
|
||||
litellm_call_id: Some("decline-test"),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decline_uses_runtime_provider_selection_without_credentials() {
|
||||
assert_eq!(
|
||||
ocr_decline_reason(request("mistral/mistral-ocr-latest", None, Map::new())),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
ocr_decline_reason(request("gpt-4o", Some("openai"), Map::new())),
|
||||
Some(OcrDeclineReason::UnsupportedProvider)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decline_rejects_python_only_native_format() {
|
||||
let optional_params = Map::from_iter([("req_format".to_string(), json!("native"))]);
|
||||
assert_eq!(
|
||||
ocr_decline_reason(request(
|
||||
"azure_ai/doc-intelligence/prebuilt-layout",
|
||||
None,
|
||||
optional_params,
|
||||
)),
|
||||
Some(OcrDeclineReason::NativeRequestFormat)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,6 +4,25 @@ use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest};
|
|||
use litellm_core::ocr::transformation::OcrProviderTransformation;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum OcrDeclineReason {
|
||||
NativeRequestFormat,
|
||||
UnsupportedProvider,
|
||||
}
|
||||
|
||||
impl std::fmt::Display for OcrDeclineReason {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::NativeRequestFormat => {
|
||||
formatter.write_str("native OCR response format requires Python")
|
||||
}
|
||||
Self::UnsupportedProvider => {
|
||||
formatter.write_str("OCR provider is not supported by the Rust runtime")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum OcrAuthStrategy {
|
||||
Bearer,
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ import base64
|
|||
import mimetypes
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Callable, Coroutine, Mapping
|
||||
from collections.abc import Coroutine, Mapping
|
||||
from dataclasses import dataclass
|
||||
from io import IOBase
|
||||
from typing import Any, Final, cast
|
||||
|
|
@ -52,21 +52,6 @@ class _PreparedOCRRequest:
|
|||
litellm_logging_obj: LiteLLMLoggingObj
|
||||
|
||||
|
||||
@dataclass
|
||||
class _PreparedRustOCRCall:
|
||||
api_key: str | None
|
||||
api_base: str | None
|
||||
headers: dict[str, object]
|
||||
optional_params: dict[str, object]
|
||||
|
||||
|
||||
_RUST_OCR_PROVIDERS: Final = {
|
||||
"mistral",
|
||||
"azure_ai",
|
||||
"vertex_ai",
|
||||
}
|
||||
|
||||
|
||||
def _prepare_ocr_request(
|
||||
model: str,
|
||||
document: Mapping[str, object],
|
||||
|
|
@ -188,115 +173,59 @@ def _prepare_ocr_request(
|
|||
)
|
||||
|
||||
|
||||
def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
|
||||
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
|
||||
return False
|
||||
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
|
||||
|
||||
|
||||
def _rust_bridge_optional_params(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> dict[str, object]:
|
||||
optional_params: Final = dict(prepared_request.optional_params)
|
||||
if prepared_request.custom_llm_provider == "vertex_ai":
|
||||
vertex_project: Final = (
|
||||
prepared_request.litellm_params.get("vertex_project")
|
||||
or prepared_request.litellm_params.get("vertex_ai_project")
|
||||
or litellm.vertex_project
|
||||
or resolve_secret("VERTEXAI_PROJECT")
|
||||
def _materialize_ocr_document(document: Mapping[str, object]) -> dict[str, object]:
|
||||
if not isinstance(document, dict):
|
||||
raise ValueError(f"document must be a dict with 'type' and URL/file field, got {type(document)}")
|
||||
if document.get("type") == "file":
|
||||
return cast(dict[str, object], convert_file_document_to_url_document(document))
|
||||
if document.get("type") not in ("document_url", "image_url"):
|
||||
raise ValueError(
|
||||
f"Invalid document type: {document.get('type')}. Must be 'document_url', 'image_url', or 'file'"
|
||||
)
|
||||
vertex_location: Final = (
|
||||
prepared_request.litellm_params.get("vertex_location")
|
||||
or prepared_request.litellm_params.get("vertex_ai_location")
|
||||
or litellm.vertex_location
|
||||
or resolve_secret("VERTEXAI_LOCATION")
|
||||
or resolve_secret("VERTEX_LOCATION")
|
||||
)
|
||||
if vertex_project is not None:
|
||||
optional_params["vertex_project"] = vertex_project
|
||||
if vertex_location is not None:
|
||||
optional_params["vertex_location"] = vertex_location
|
||||
return optional_params
|
||||
return document
|
||||
|
||||
|
||||
def _rust_bridge_api_base(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_secret: Callable[[str], str | None],
|
||||
) -> str | None:
|
||||
if prepared_request.api_base is not None:
|
||||
return prepared_request.api_base
|
||||
if prepared_request.custom_llm_provider == "azure_ai":
|
||||
if is_azure_document_intelligence_model(prepared_request.model):
|
||||
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
|
||||
return resolve_secret("AZURE_AI_API_BASE")
|
||||
return None
|
||||
def _native_ocr_params(kwargs: Mapping[str, object]) -> dict[str, object]:
|
||||
return {key: value for key, value in kwargs.items() if key != "litellm_logging_obj"}
|
||||
|
||||
|
||||
def _prepare_rust_ocr_call(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
) -> _PreparedRustOCRCall:
|
||||
provider_config: Final = prepared_request.provider_config
|
||||
api_key_env_var: Final = provider_config.get_api_key_env_var()
|
||||
resolved_api_key: Final = prepared_request.api_key or (
|
||||
resolve_api_key(api_key_env_var) if api_key_env_var is not None else None
|
||||
)
|
||||
resolved_headers: Final = provider_config.validate_environment(
|
||||
headers=prepared_request.extra_headers or {},
|
||||
model=prepared_request.model,
|
||||
api_key=resolved_api_key,
|
||||
api_base=prepared_request.api_base,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
resolved_complete_url: Final = provider_config.get_complete_url(
|
||||
api_base=prepared_request.api_base,
|
||||
model=prepared_request.model,
|
||||
optional_params=prepared_request.optional_params,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
|
||||
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
|
||||
prepared_request.litellm_logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=resolved_api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": {
|
||||
"model": prepared_request.model,
|
||||
"document": prepared_request.document,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": resolved_complete_url,
|
||||
"headers": resolved_headers,
|
||||
},
|
||||
)
|
||||
return _PreparedRustOCRCall(
|
||||
api_key=resolved_api_key,
|
||||
api_base=rust_api_base,
|
||||
headers=cast(dict[str, object], resolved_headers),
|
||||
optional_params=rust_optional_params,
|
||||
)
|
||||
def _validate_ocr_request_format(
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
kwargs: Mapping[str, object],
|
||||
) -> None:
|
||||
requested_format: Final = kwargs.get(OCR_REQUEST_FORMAT_PARAM)
|
||||
if requested_format is not None:
|
||||
try:
|
||||
parse_ocr_request_format(requested_format)
|
||||
except ValueError as error:
|
||||
raise litellm.exceptions.UnsupportedParamsError(
|
||||
message=str(error),
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider or "",
|
||||
) from error
|
||||
|
||||
|
||||
def _run_rust_ocr(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
*,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> OCRResponse | None:
|
||||
if rust_ocr_bridge.load_rust_ocr() is None:
|
||||
return None
|
||||
prepared: Final = _prepare_rust_ocr_call(
|
||||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
)
|
||||
rust_response: Final = rust_ocr_bridge.ocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_response is None:
|
||||
return None
|
||||
|
|
@ -304,24 +233,25 @@ def _run_rust_ocr(
|
|||
|
||||
|
||||
async def _run_rust_aocr(
|
||||
prepared_request: _PreparedOCRRequest,
|
||||
resolve_api_key: Callable[[str], str | None],
|
||||
*,
|
||||
model: str,
|
||||
document: dict[str, object],
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
custom_llm_provider: str | None,
|
||||
extra_headers: dict[str, object] | None,
|
||||
optional_params: dict[str, object],
|
||||
timeout: float | httpx.Timeout | None,
|
||||
) -> OCRResponse | None:
|
||||
if rust_ocr_bridge.load_rust_aocr() is None:
|
||||
return None
|
||||
prepared: Final = _prepare_rust_ocr_call(
|
||||
prepared_request=prepared_request,
|
||||
resolve_api_key=resolve_api_key,
|
||||
)
|
||||
rust_response: Final = await rust_ocr_bridge.aocr(
|
||||
model=prepared_request.model,
|
||||
document=prepared_request.document,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
custom_llm_provider=prepared_request.custom_llm_provider,
|
||||
extra_headers=prepared.headers,
|
||||
optional_params=prepared.optional_params,
|
||||
timeout=prepared_request.effective_timeout,
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_response is None:
|
||||
return None
|
||||
|
|
@ -408,9 +338,40 @@ async def aocr(
|
|||
"kwargs": kwargs,
|
||||
}
|
||||
try:
|
||||
_validate_ocr_request_format(model, custom_llm_provider, kwargs)
|
||||
native_document: Final = _materialize_ocr_document(document)
|
||||
native_params: Final = _native_ocr_params(kwargs)
|
||||
native_available: Final = (
|
||||
rust_ocr_bridge.rust_ocr_enabled()
|
||||
and rust_ocr_bridge.load_rust_aocr() is not None
|
||||
and rust_ocr_bridge.load_rust_ocr_decline() is not None
|
||||
)
|
||||
if native_available:
|
||||
decline_reason: Final = rust_ocr_bridge.ocr_decline(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
optional_params=native_params,
|
||||
)
|
||||
if decline_reason is None:
|
||||
rust_response: Final = await _run_rust_aocr(
|
||||
model=model,
|
||||
document=native_document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=native_params,
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path")
|
||||
else:
|
||||
verbose_logger.debug("Async Rust OCR declined request: %s", decline_reason)
|
||||
|
||||
prepared: Final = _prepare_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
document=native_document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
timeout=timeout,
|
||||
|
|
@ -422,18 +383,6 @@ async def aocr(
|
|||
custom_llm_provider = prepared.custom_llm_provider
|
||||
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
|
||||
|
||||
if _rust_ocr_supported(prepared) and rust_ocr_bridge.rust_ocr_enabled():
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
rust_response: Final = await _run_rust_aocr(
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
)
|
||||
if rust_response is None:
|
||||
verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path")
|
||||
else:
|
||||
return rust_response
|
||||
|
||||
response = base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
|
|
@ -680,9 +629,40 @@ def ocr(
|
|||
try:
|
||||
_is_async: Final = kwargs.pop("aocr", False) is True
|
||||
completion_kwargs["aocr"] = _is_async
|
||||
_validate_ocr_request_format(model, custom_llm_provider, kwargs)
|
||||
native_document: Final = _materialize_ocr_document(document)
|
||||
native_params: Final = _native_ocr_params(kwargs)
|
||||
native_available: Final = (
|
||||
rust_ocr_bridge.rust_ocr_enabled()
|
||||
and rust_ocr_bridge.load_rust_ocr() is not None
|
||||
and rust_ocr_bridge.load_rust_ocr_decline() is not None
|
||||
)
|
||||
if native_available:
|
||||
decline_reason: Final = rust_ocr_bridge.ocr_decline(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
optional_params=native_params,
|
||||
)
|
||||
if decline_reason is None:
|
||||
rust_response: Final = _run_rust_ocr(
|
||||
model=model,
|
||||
document=native_document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
optional_params=native_params,
|
||||
timeout=timeout,
|
||||
)
|
||||
if rust_response is not None:
|
||||
return rust_response
|
||||
verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path")
|
||||
else:
|
||||
verbose_logger.debug("Rust OCR declined request: %s", decline_reason)
|
||||
|
||||
prepared: Final = _prepare_ocr_request(
|
||||
model=model,
|
||||
document=document,
|
||||
document=native_document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
kwargs=kwargs,
|
||||
|
|
@ -694,18 +674,6 @@ def ocr(
|
|||
custom_llm_provider = prepared.custom_llm_provider
|
||||
completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider})
|
||||
|
||||
if _rust_ocr_supported(prepared) and rust_ocr_bridge.rust_ocr_enabled():
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
rust_response: Final = _run_rust_ocr(
|
||||
prepared_request=prepared,
|
||||
resolve_api_key=get_secret_str,
|
||||
)
|
||||
if rust_response is None:
|
||||
verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path")
|
||||
else:
|
||||
return rust_response
|
||||
|
||||
response: Final = base_llm_http_handler.ocr(
|
||||
model=prepared.model,
|
||||
document=prepared.document,
|
||||
|
|
|
|||
|
|
@ -44,6 +44,16 @@ class RustAocr(Protocol):
|
|||
raise NotImplementedError
|
||||
|
||||
|
||||
class RustOcrDecline(Protocol):
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
optional_params: dict[str, object],
|
||||
) -> str | None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class _Unset:
|
||||
pass
|
||||
|
||||
|
|
@ -63,6 +73,7 @@ def _env_enables_rust_ocr() -> bool:
|
|||
_rust_ocr_enabled = _env_enables_rust_ocr()
|
||||
_rust_ocr_impl: RustOcr | None = None
|
||||
_rust_aocr_impl: RustAocr | None = None
|
||||
_rust_ocr_decline_impl: RustOcrDecline | None = None
|
||||
|
||||
|
||||
def use_litellm_rust(
|
||||
|
|
@ -70,14 +81,17 @@ def use_litellm_rust(
|
|||
*,
|
||||
ocr: RustOcr | None | _Unset = _UNSET,
|
||||
aocr: RustAocr | None | _Unset = _UNSET,
|
||||
ocr_decline: RustOcrDecline | None | _Unset = _UNSET,
|
||||
messages: RustMessages | None | _Unset = _UNSET,
|
||||
amessages: RustAmessages | None | _Unset = _UNSET,
|
||||
responses_websocket: Any | None | _Unset = _UNSET,
|
||||
transcription: Any | None | _Unset = _UNSET,
|
||||
atranscription: Any | None | _Unset = _UNSET,
|
||||
) -> None:
|
||||
global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl
|
||||
configuring_ocr: Final = not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset)
|
||||
global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl, _rust_ocr_decline_impl
|
||||
configuring_ocr: Final = (
|
||||
not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset) or not isinstance(ocr_decline, _Unset)
|
||||
)
|
||||
configuring_messages: Final = not isinstance(messages, _Unset) or not isinstance(amessages, _Unset)
|
||||
configuring_responses_websocket: Final = not isinstance(responses_websocket, _Unset)
|
||||
configuring_transcription: Final = not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset)
|
||||
|
|
@ -87,6 +101,8 @@ def use_litellm_rust(
|
|||
_rust_ocr_impl = ocr
|
||||
if not isinstance(aocr, _Unset):
|
||||
_rust_aocr_impl = aocr
|
||||
if not isinstance(ocr_decline, _Unset):
|
||||
_rust_ocr_decline_impl = ocr_decline
|
||||
if configuring_transcription:
|
||||
from litellm.rust_bridge.transcription import configure_rust_transcription
|
||||
|
||||
|
|
@ -138,6 +154,33 @@ def load_rust_aocr() -> RustAocr | None:
|
|||
return cast(RustAocr, getattr(native_bridge, "aocr", None))
|
||||
|
||||
|
||||
def load_rust_ocr_decline() -> RustOcrDecline | None:
|
||||
if _rust_ocr_decline_impl is not None:
|
||||
return _rust_ocr_decline_impl
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
native_bridge: Final = get_native_bridge()
|
||||
if native_bridge is None:
|
||||
return None
|
||||
return cast(RustOcrDecline, getattr(native_bridge, "ocr_decline", None))
|
||||
|
||||
|
||||
def ocr_decline(
|
||||
*,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
optional_params: dict[str, object],
|
||||
) -> str | None:
|
||||
decline: Final = load_rust_ocr_decline()
|
||||
if decline is None:
|
||||
return "Rust OCR decline bridge is unavailable"
|
||||
return decline(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
|
||||
def ocr(
|
||||
*,
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -11,10 +11,9 @@ supplied api_base is always honoured.
|
|||
from litellm.llms.azure_ai.ocr.common_utils import (
|
||||
is_azure_document_intelligence_model,
|
||||
)
|
||||
from litellm.ocr.main import _prepare_ocr_request, _rust_bridge_api_base
|
||||
from litellm.ocr.main import _prepare_ocr_request
|
||||
|
||||
_DOC = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
|
||||
_DOC_INTELLIGENCE_ENDPOINT = "https://di.cognitiveservices.azure.com"
|
||||
_AZURE_AI_API_BASE = "https://generic-azure-ai.example.com"
|
||||
|
||||
|
||||
|
|
@ -23,13 +22,6 @@ class _FakeLogging:
|
|||
return None
|
||||
|
||||
|
||||
def _resolve_secret(name: str) -> str | None:
|
||||
return {
|
||||
"AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT": _DOC_INTELLIGENCE_ENDPOINT,
|
||||
"AZURE_AI_API_BASE": _AZURE_AI_API_BASE,
|
||||
}.get(name)
|
||||
|
||||
|
||||
def _prepare(model: str, api_base: str | None):
|
||||
return _prepare_ocr_request(
|
||||
model=model,
|
||||
|
|
@ -64,7 +56,6 @@ class TestDocIntelligenceApiBaseResolution:
|
|||
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", None)
|
||||
|
||||
assert prepared.api_base is None
|
||||
assert _rust_bridge_api_base(prepared, _resolve_secret) == _DOC_INTELLIGENCE_ENDPOINT
|
||||
|
||||
def test_explicit_api_base_is_honoured_for_doc_intelligence(self, monkeypatch):
|
||||
"""A caller-supplied api_base must always win, even for doc-intelligence."""
|
||||
|
|
@ -74,7 +65,6 @@ class TestDocIntelligenceApiBaseResolution:
|
|||
prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", custom)
|
||||
|
||||
assert prepared.api_base == custom
|
||||
assert _rust_bridge_api_base(prepared, _resolve_secret) == custom
|
||||
|
||||
def test_generic_azure_ai_base_still_applies_to_mistral_ocr(self, monkeypatch):
|
||||
"""Non doc-intelligence azure_ai models keep using AZURE_AI_API_BASE."""
|
||||
|
|
|
|||
|
|
@ -4,39 +4,32 @@ providers that don't support a native response must reject it, and the Rust
|
|||
bridge (which only returns the normalized shape) must not serve native requests.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.ocr.main import _PreparedOCRRequest, _rust_ocr_supported
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
|
||||
DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
|
||||
|
||||
|
||||
def _prepared(optional_params: dict[str, object]) -> _PreparedOCRRequest:
|
||||
return _PreparedOCRRequest(
|
||||
model="doc-intelligence/prebuilt-layout",
|
||||
document=dict(DOCUMENT),
|
||||
api_key="fake-key",
|
||||
api_base="https://example.cognitiveservices.azure.com",
|
||||
custom_llm_provider="azure_ai",
|
||||
extra_headers=None,
|
||||
provider_config=MagicMock(),
|
||||
optional_params=optional_params,
|
||||
litellm_params={},
|
||||
effective_timeout=60.0,
|
||||
litellm_logging_obj=MagicMock(),
|
||||
def test_native_decline_wrapper_uses_runtime_result():
|
||||
rust_ocr_bridge.use_litellm_rust(
|
||||
True,
|
||||
ocr_decline=lambda model, custom_llm_provider, optional_params: (
|
||||
"native OCR response format requires Python" if optional_params.get("req_format") == "native" else None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("optional_params", [{}, {"req_format": "litellm"}])
|
||||
def test_rust_ocr_serves_default_format(optional_params):
|
||||
assert _rust_ocr_supported(_prepared(optional_params)) is True
|
||||
|
||||
|
||||
def test_rust_ocr_skipped_for_native_format():
|
||||
assert _rust_ocr_supported(_prepared({"req_format": "native"})) is False
|
||||
try:
|
||||
assert (
|
||||
rust_ocr_bridge.ocr_decline(
|
||||
model="azure_ai/doc-intelligence/prebuilt-layout",
|
||||
custom_llm_provider=None,
|
||||
optional_params={"req_format": "native"},
|
||||
)
|
||||
== "native OCR response format requires Python"
|
||||
)
|
||||
finally:
|
||||
rust_ocr_bridge.use_litellm_rust(False, ocr_decline=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
"""Tests for the optional Rust-backed OCR path."""
|
||||
|
||||
import importlib
|
||||
import builtins
|
||||
import importlib
|
||||
import types
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -101,6 +102,27 @@ class RecordingAsyncBridge:
|
|||
return dict(FAKE_OCR_RESPONSE)
|
||||
|
||||
|
||||
class RecordingDecline:
|
||||
def __init__(self, reason: str | None = None) -> None:
|
||||
self.reason = reason
|
||||
self.calls: list[dict[str, object]] = []
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
model: str,
|
||||
custom_llm_provider: str | None,
|
||||
optional_params: dict[str, object],
|
||||
) -> str | None:
|
||||
self.calls.append(
|
||||
{
|
||||
"model": model,
|
||||
"custom_llm_provider": custom_llm_provider,
|
||||
"optional_params": optional_params,
|
||||
}
|
||||
)
|
||||
return self.reason
|
||||
|
||||
|
||||
class RaisingBridge:
|
||||
def __call__(
|
||||
self,
|
||||
|
|
@ -214,10 +236,10 @@ def build_prepared_request(
|
|||
@pytest.fixture(autouse=True)
|
||||
def _reset_rust_flag():
|
||||
"""Keep the global toggle isolated between tests."""
|
||||
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None)
|
||||
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None, ocr_decline=None)
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
yield
|
||||
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None)
|
||||
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None, ocr_decline=None)
|
||||
rust_bridge_loader._cached_bridge = rust_bridge_loader._BRIDGE_SENTINEL
|
||||
|
||||
|
||||
|
|
@ -225,7 +247,7 @@ def _reset_rust_flag():
|
|||
def fake_bridge():
|
||||
"""Enable the Rust path with an injected recording bridge (no native wheel)."""
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
litellm.use_litellm_rust(True, ocr=bridge, ocr_decline=RecordingDecline())
|
||||
return bridge
|
||||
|
||||
|
||||
|
|
@ -233,7 +255,7 @@ def fake_bridge():
|
|||
def fake_async_bridge():
|
||||
"""Enable the async Rust path with an injected recording bridge."""
|
||||
bridge = RecordingAsyncBridge()
|
||||
litellm.use_litellm_rust(True, aocr=bridge)
|
||||
litellm.use_litellm_rust(True, aocr=bridge, ocr_decline=RecordingDecline())
|
||||
return bridge
|
||||
|
||||
|
||||
|
|
@ -434,90 +456,85 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response():
|
|||
|
||||
def test_run_rust_ocr_prepares_request_and_wraps_response():
|
||||
bridge = RecordingBridge()
|
||||
logging_obj = RecordingLogging()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
response = ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
logging_obj=logging_obj,
|
||||
api_base="https://proxy.internal",
|
||||
extra_headers={"x-trace-id": "trace-1"},
|
||||
optional_params={"include_image_base64": True},
|
||||
timeout=12.5,
|
||||
),
|
||||
resolve_api_key=lambda _name: None,
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
api_key="sk-test",
|
||||
api_base="https://proxy.internal",
|
||||
custom_llm_provider=None,
|
||||
extra_headers={"x-trace-id": "trace-1"},
|
||||
optional_params={"include_image_base64": True},
|
||||
timeout=12.5,
|
||||
)
|
||||
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert response.pages[0].markdown == "hello world"
|
||||
assert bridge.calls[0] == {
|
||||
"model": "mistral-ocr-latest",
|
||||
"model": MODEL,
|
||||
"document": DOCUMENT,
|
||||
"api_key": "sk-test",
|
||||
"api_base": "https://proxy.internal",
|
||||
"custom_llm_provider": "mistral",
|
||||
"extra_headers": {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
},
|
||||
"custom_llm_provider": None,
|
||||
"extra_headers": {"x-trace-id": "trace-1"},
|
||||
"optional_params": {"include_image_base64": True},
|
||||
"timeout_seconds": 12.5,
|
||||
}
|
||||
|
||||
|
||||
def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
|
||||
def test_run_rust_ocr_leaves_missing_key_for_runtime_resolution():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(api_key=None, timeout=None),
|
||||
resolve_api_key=lambda name: (
|
||||
"sk-from-vault" if name == "MISTRAL_API_KEY" else None
|
||||
),
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["api_key"] == "sk-from-vault"
|
||||
assert bridge.calls[0]["api_key"] is None
|
||||
|
||||
|
||||
def test_run_rust_ocr_prefers_explicit_key_over_resolver():
|
||||
def test_run_rust_ocr_forwards_explicit_key():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
def _resolver(name: str) -> str | None:
|
||||
raise AssertionError(f"resolver should not be called for {name}")
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
api_key="sk-explicit",
|
||||
timeout=None,
|
||||
),
|
||||
resolve_api_key=_resolver,
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
api_key="sk-explicit",
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["api_key"] == "sk-explicit"
|
||||
|
||||
|
||||
def test_run_rust_ocr_uses_provider_api_key_env_var():
|
||||
def test_run_rust_ocr_does_not_resolve_provider_api_key_env_var():
|
||||
bridge = RecordingBridge()
|
||||
resolver_calls = []
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
def _resolver(name):
|
||||
resolver_calls.append(name)
|
||||
return "sk-provider-env"
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"),
|
||||
model="provider-ocr-model",
|
||||
api_key=None,
|
||||
timeout=None,
|
||||
),
|
||||
resolve_api_key=_resolver,
|
||||
model="provider/provider-ocr-model",
|
||||
document=DOCUMENT,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert resolver_calls == ["PROVIDER_OCR_API_KEY"]
|
||||
assert bridge.calls[0]["api_key"] == "sk-provider-env"
|
||||
assert bridge.calls[0]["api_key"] is None
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
||||
|
|
@ -525,18 +542,18 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
|||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="vertex_ai",
|
||||
model="mistral-ocr-maas",
|
||||
litellm_params={
|
||||
"vertex_project": "project-1",
|
||||
"vertex_location": "us-central1",
|
||||
"vertex_credentials": "redacted",
|
||||
},
|
||||
optional_params={"include_image_base64": True},
|
||||
timeout=None,
|
||||
),
|
||||
resolve_api_key=lambda _name: None,
|
||||
model="vertex_ai/mistral-ocr-maas",
|
||||
document=DOCUMENT,
|
||||
api_key="sk-test",
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
optional_params={
|
||||
"include_image_base64": True,
|
||||
"vertex_project": "project-1",
|
||||
"vertex_location": "us-central1",
|
||||
},
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["optional_params"] == {
|
||||
|
|
@ -546,99 +563,85 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
|
|||
}
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
def _resolver(name: str) -> str | None:
|
||||
return {
|
||||
"VERTEXAI_PROJECT": "project-from-secret",
|
||||
"VERTEXAI_LOCATION": "us-east5",
|
||||
}.get(name)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="vertex_ai",
|
||||
model="mistral-ocr-maas",
|
||||
timeout=None,
|
||||
),
|
||||
resolve_api_key=_resolver,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["optional_params"]["vertex_project"] == "project-from-secret"
|
||||
assert bridge.calls[0]["optional_params"]["vertex_location"] == "us-east5"
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager():
|
||||
def test_prepare_rust_ocr_call_leaves_vertex_env_resolution_to_runtime():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="azure_ai",
|
||||
model="pixtral-12b-2409",
|
||||
api_base=None,
|
||||
timeout=None,
|
||||
),
|
||||
resolve_api_key=lambda name: (
|
||||
"https://azure.example.com" if name == "AZURE_AI_API_BASE" else None
|
||||
),
|
||||
model="vertex_ai/mistral-ocr-maas",
|
||||
document=DOCUMENT,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["api_base"] == "https://azure.example.com"
|
||||
assert bridge.calls[0]["optional_params"] == {}
|
||||
|
||||
|
||||
def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
|
||||
def test_prepare_rust_ocr_call_leaves_azure_ai_api_base_for_runtime_resolution():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
custom_llm_provider="azure_ai",
|
||||
model="doc-intelligence/prebuilt-layout",
|
||||
api_base=None,
|
||||
timeout=None,
|
||||
),
|
||||
resolve_api_key=lambda name: (
|
||||
"https://document-intelligence.example.com"
|
||||
if name == "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT"
|
||||
else None
|
||||
),
|
||||
model="azure_ai/pixtral-12b-2409",
|
||||
document=DOCUMENT,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["api_base"] == "https://document-intelligence.example.com"
|
||||
assert bridge.calls[0]["api_base"] is None
|
||||
|
||||
|
||||
def test_run_rust_ocr_runs_pre_call_logging():
|
||||
def test_prepare_rust_ocr_call_leaves_document_intelligence_endpoint_for_runtime():
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
model="azure_ai/doc-intelligence/prebuilt-layout",
|
||||
document=DOCUMENT,
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
optional_params={},
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert bridge.calls[0]["api_base"] is None
|
||||
|
||||
|
||||
def test_run_rust_ocr_does_not_run_python_provider_logging_preparation():
|
||||
logging_obj = RecordingLogging()
|
||||
bridge = RecordingBridge()
|
||||
litellm.use_litellm_rust(True, ocr=bridge)
|
||||
|
||||
ocr_main._run_rust_ocr(
|
||||
prepared_request=build_prepared_request(
|
||||
logging_obj=logging_obj,
|
||||
api_base="https://api.mistral.ai/v1",
|
||||
extra_headers={"x-trace-id": "trace-1"},
|
||||
optional_params={"include_image_base64": True},
|
||||
timeout=None,
|
||||
),
|
||||
resolve_api_key=lambda _name: None,
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
api_key="sk-test",
|
||||
api_base=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers={"x-trace-id": "trace-1"},
|
||||
optional_params={"include_image_base64": True},
|
||||
timeout=None,
|
||||
)
|
||||
|
||||
assert logging_obj.pre_call_kwargs is not None
|
||||
assert logging_obj.pre_call_kwargs["input"] == "OCR document processing"
|
||||
additional_args = logging_obj.pre_call_kwargs["additional_args"]
|
||||
complete_input = additional_args["complete_input_dict"]
|
||||
assert complete_input["document"] == DOCUMENT
|
||||
assert complete_input["include_image_base64"] is True
|
||||
assert additional_args["api_base"] == "https://api.mistral.ai/v1/ocr"
|
||||
assert additional_args["headers"] == {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
}
|
||||
assert logging_obj.pre_call_kwargs is None
|
||||
|
||||
|
||||
def test_ocr_routes_to_rust_when_enabled(fake_bridge):
|
||||
def test_ocr_routes_to_rust_before_python_preparation(fake_bridge, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
ocr_main,
|
||||
"_prepare_ocr_request",
|
||||
lambda **kwargs: (_ for _ in ()).throw(AssertionError("Python OCR preparation must not run")),
|
||||
)
|
||||
response = litellm.ocr(
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
|
|
@ -651,14 +654,11 @@ def test_ocr_routes_to_rust_when_enabled(fake_bridge):
|
|||
assert response.pages[0].markdown == "hello world"
|
||||
assert len(fake_bridge.calls) == 1
|
||||
call = fake_bridge.calls[0]
|
||||
assert call["model"] == "mistral-ocr-latest"
|
||||
assert call["model"] == MODEL
|
||||
assert call["document"] == DOCUMENT
|
||||
assert call["api_key"] == "sk-test"
|
||||
assert call["custom_llm_provider"] == "mistral"
|
||||
assert call["extra_headers"] == {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
}
|
||||
assert call["custom_llm_provider"] is None
|
||||
assert call["extra_headers"] == {"x-trace-id": "trace-1"}
|
||||
assert call["optional_params"].get("include_image_base64") is True
|
||||
|
||||
|
||||
|
|
@ -672,8 +672,8 @@ def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge):
|
|||
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert len(fake_bridge.calls) == 1
|
||||
assert fake_bridge.calls[0]["model"] == "pixtral-12b-2409"
|
||||
assert fake_bridge.calls[0]["custom_llm_provider"] == "azure_ai"
|
||||
assert fake_bridge.calls[0]["model"] == "azure_ai/pixtral-12b-2409"
|
||||
assert fake_bridge.calls[0]["custom_llm_provider"] is None
|
||||
|
||||
|
||||
def test_ocr_rust_path_converts_file_document_before_bridge(fake_bridge):
|
||||
|
|
@ -699,17 +699,22 @@ def test_ocr_exception_type_uses_resolved_provider_context(
|
|||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.use_litellm_rust(True, ocr=RaisingBridge())
|
||||
litellm.use_litellm_rust(True, ocr=RaisingBridge(), ocr_decline=RecordingDecline())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
||||
assert captured["model"] == "mistral-ocr-latest"
|
||||
assert captured["custom_llm_provider"] == "mistral"
|
||||
assert captured["model"] == MODEL
|
||||
assert captured["custom_llm_provider"] is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge):
|
||||
async def test_aocr_routes_to_async_rust_before_python_preparation(fake_async_bridge, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
ocr_main,
|
||||
"_prepare_ocr_request",
|
||||
lambda **kwargs: (_ for _ in ()).throw(AssertionError("Python OCR preparation must not run")),
|
||||
)
|
||||
response = await litellm.aocr(
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
|
|
@ -722,14 +727,11 @@ async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge):
|
|||
assert response.pages[0].markdown == "hello world"
|
||||
assert len(fake_async_bridge.calls) == 1
|
||||
call = fake_async_bridge.calls[0]
|
||||
assert call["model"] == "mistral-ocr-latest"
|
||||
assert call["model"] == MODEL
|
||||
assert call["document"] == DOCUMENT
|
||||
assert call["api_key"] == "sk-test"
|
||||
assert call["custom_llm_provider"] == "mistral"
|
||||
assert call["extra_headers"] == {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
}
|
||||
assert call["custom_llm_provider"] is None
|
||||
assert call["extra_headers"] == {"x-trace-id": "trace-1"}
|
||||
assert call["optional_params"].get("include_image_base64") is True
|
||||
|
||||
|
||||
|
|
@ -744,13 +746,13 @@ async def test_aocr_exception_type_uses_resolved_provider_context(
|
|||
return CapturedException("wrapped")
|
||||
|
||||
monkeypatch.setattr(ocr_main.litellm, "exception_type", fake_exception_type)
|
||||
litellm.use_litellm_rust(True, aocr=RaisingAsyncBridge())
|
||||
litellm.use_litellm_rust(True, aocr=RaisingAsyncBridge(), ocr_decline=RecordingDecline())
|
||||
|
||||
with pytest.raises(CapturedException):
|
||||
await litellm.aocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
||||
assert captured["model"] == "mistral-ocr-latest"
|
||||
assert captured["custom_llm_provider"] == "mistral"
|
||||
assert captured["model"] == MODEL
|
||||
assert captured["custom_llm_provider"] is None
|
||||
|
||||
|
||||
def test_ocr_forwards_timeout_to_rust(fake_bridge):
|
||||
|
|
@ -764,9 +766,7 @@ def test_ocr_forwards_timeout_to_rust(fake_bridge):
|
|||
def test_ocr_passes_default_request_timeout_to_rust(fake_bridge):
|
||||
litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
||||
from litellm.constants import request_timeout
|
||||
|
||||
assert fake_bridge.calls[0]["timeout_seconds"] == float(request_timeout)
|
||||
assert fake_bridge.calls[0]["timeout_seconds"] is None
|
||||
|
||||
|
||||
def test_ocr_does_not_route_to_rust_when_disabled():
|
||||
|
|
@ -800,6 +800,28 @@ def test_ocr_falls_back_to_python_when_bridge_unavailable(monkeypatch):
|
|||
assert isinstance(response, OCRResponse)
|
||||
|
||||
|
||||
def test_ocr_falls_back_to_python_when_runtime_declines(monkeypatch):
|
||||
bridge = RecordingBridge()
|
||||
decline = RecordingDecline("OCR provider is not supported by the Rust runtime")
|
||||
litellm.use_litellm_rust(True, ocr=bridge, ocr_decline=decline)
|
||||
prepared = build_prepared_request()
|
||||
fake_prepare = MagicMock(return_value=prepared)
|
||||
|
||||
monkeypatch.setattr(ocr_main, "_prepare_ocr_request", fake_prepare)
|
||||
monkeypatch.setattr(
|
||||
ocr_main.base_llm_http_handler,
|
||||
"ocr",
|
||||
lambda **kwargs: OCRResponse(pages=[], model="mistral-ocr-latest", object="ocr"),
|
||||
)
|
||||
|
||||
response = litellm.ocr(model=MODEL, document=DOCUMENT, api_key="sk-test")
|
||||
|
||||
assert isinstance(response, OCRResponse)
|
||||
assert len(decline.calls) == 1
|
||||
assert bridge.calls == []
|
||||
fake_prepare.assert_called_once()
|
||||
|
||||
|
||||
def test_ocr_provider_configs_expose_api_key_env_vars():
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue