refactor(ocr): remove python provider transforms

This commit is contained in:
Ishaan Jaff 2026-06-25 14:30:06 -07:00
parent 1528c9b9d1
commit b2e7c48ebd
No known key found for this signature in database
27 changed files with 174 additions and 4457 deletions

View file

@ -371,12 +371,10 @@ UTILS_MODULE_NAMES = (
"redact_message_input_output_from_logging",
"CustomStreamWrapper",
"BaseGoogleGenAIGenerateContentConfig",
"BaseOCRConfig",
"BaseSearchConfig",
"BaseTextToSpeechConfig",
"BedrockModelInfo",
"CohereModelInfo",
"MistralOCRConfig",
"Rules",
"AsyncHTTPHandler",
"HTTPHandler",
@ -1290,7 +1288,6 @@ _UTILS_MODULE_IMPORT_MAP = {
"litellm.llms.base_llm.google_genai.transformation",
"BaseGoogleGenAIGenerateContentConfig",
),
"BaseOCRConfig": ("litellm.llms.base_llm.ocr.transformation", "BaseOCRConfig"),
"BaseSearchConfig": (
"litellm.llms.base_llm.search.transformation",
"BaseSearchConfig",
@ -1301,7 +1298,6 @@ _UTILS_MODULE_IMPORT_MAP = {
),
"BedrockModelInfo": ("litellm.llms.bedrock.common_utils", "BedrockModelInfo"),
"CohereModelInfo": ("litellm.llms.cohere.common_utils", "CohereModelInfo"),
"MistralOCRConfig": ("litellm.llms.mistral.ocr.transformation", "MistralOCRConfig"),
"Rules": ("litellm.litellm_core_utils.rules", "Rules"),
"AsyncHTTPHandler": ("litellm.llms.custom_httpx.http_handler", "AsyncHTTPHandler"),
"HTTPHandler": ("litellm.llms.custom_httpx.http_handler", "HTTPHandler"),

View file

@ -1,13 +1,3 @@
"""Azure AI OCR module."""
from .common_utils import get_azure_ai_ocr_config
from .document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
from .transformation import AzureAIOCRConfig
__all__ = [
"AzureAIOCRConfig",
"AzureDocumentIntelligenceOCRConfig",
"get_azure_ai_ocr_config",
]
__all__: list[str] = []

View file

@ -1,52 +0,0 @@
"""
Common utilities for Azure AI OCR providers.
This module provides routing logic to determine which OCR configuration to use
based on the model name.
"""
from typing import TYPE_CHECKING, Optional
from litellm._logging import verbose_logger
if TYPE_CHECKING:
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
"""
Determine which Azure AI OCR configuration to use based on the model name.
Azure AI supports multiple OCR services:
- Azure Document Intelligence: azure_ai/doc-intelligence/<model>
- Mistral OCR (via Azure AI): azure_ai/<model>
Args:
model: The model name (e.g., "azure_ai/doc-intelligence/prebuilt-read",
"azure_ai/pixtral-12b-2409")
Returns:
OCR configuration instance for the specified model
Examples:
>>> get_azure_ai_ocr_config("azure_ai/doc-intelligence/prebuilt-read")
<AzureDocumentIntelligenceOCRConfig object>
>>> get_azure_ai_ocr_config("azure_ai/pixtral-12b-2409")
<AzureAIOCRConfig object>
"""
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
# Check for Azure Document Intelligence models
if "doc-intelligence" in model or "documentintelligence" in model:
verbose_logger.debug(
f"Routing {model} to Azure Document Intelligence OCR config"
)
return AzureDocumentIntelligenceOCRConfig()
# Default to Mistral-based OCR for other azure_ai models
verbose_logger.debug(f"Routing {model} to Azure AI (Mistral) OCR config")
return AzureAIOCRConfig()

View file

@ -1,5 +1,3 @@
"""Azure Document Intelligence OCR module."""
from .transformation import AzureDocumentIntelligenceOCRConfig
__all__ = ["AzureDocumentIntelligenceOCRConfig"]
__all__: list[str] = []

View file

@ -1,815 +0,0 @@
"""
Azure Document Intelligence OCR transformation implementation.
Azure Document Intelligence (formerly Form Recognizer) provides advanced document analysis capabilities.
This implementation transforms between Mistral OCR format and Azure Document Intelligence API v4.0.
Note: Azure Document Intelligence API is async - POST returns 202 Accepted with Operation-Location header.
The operation location must be polled until the analysis completes.
"""
import asyncio
import re
import time
from typing import Any, Dict
from urllib.parse import quote
import httpx
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.url_utils import SSRFError, assert_same_origin
from litellm.constants import (
AZURE_DOCUMENT_INTELLIGENCE_API_VERSION,
AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI,
AZURE_OPERATION_POLLING_TIMEOUT,
)
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.ocr.transformation import (
BaseOCRConfig,
DocumentType,
OCRPage,
OCRPageDimensions,
OCRRequestData,
OCRResponse,
OCRUsageInfo,
)
from litellm.secret_managers.main import get_secret_str
AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"
class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
"""
Azure Document Intelligence OCR transformation configuration.
Supports Azure Document Intelligence v4.0 (2024-11-30) API.
Model route: azure_ai/doc-intelligence/<model>
Supported models:
- prebuilt-layout: Extracts text with markdown, tables, and structure (closest to Mistral OCR)
- prebuilt-read: Basic text extraction optimized for reading
- prebuilt-document: General document analysis
Reference: https://learn.microsoft.com/en-us/azure/ai-services/document-intelligence/
"""
def __init__(self) -> None:
super().__init__()
def get_api_key_env_var(self) -> str | None:
return AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR
def get_supported_ocr_params(self, model: str) -> list:
"""
Get supported OCR parameters for Azure Document Intelligence.
Azure DI exposes a `pages` query parameter on the analyze endpoint
(1-based, e.g. "1-3,5,7-9"). To keep the public request shape
aligned with Mistral OCR, callers pass `pages` using Mistral
semantics — a list of 0-based integers — or a pre-formatted
Azure-style string. Other Mistral-specific params (e.g.
`include_image_base64`) are not supported by Azure DI and are
ignored during transformation.
"""
return ["pages"]
def map_ocr_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
) -> dict:
"""
Map OCR params to Azure DI format.
Translates Mistral-style `pages` (list[int], 0-based) into Azure's
`pages` query string (1-based, e.g. "1,2,3" or "1-3,5"). A raw
string that already matches Azure's format is passed through
unchanged.
"""
pages = non_default_params.get("pages")
if pages is None:
return optional_params
normalized = self._normalize_pages_param(pages)
if normalized:
optional_params["pages"] = normalized
return optional_params
@staticmethod
def _normalize_pages_param(pages: Any) -> str:
"""
Convert a caller-provided `pages` value to Azure DI's query-string
form. Azure expects 1-based page numbers, grammar: `^(\\d+(-\\d+)?)(,\\s*(\\d+(-\\d+)?))*$`.
Accepted inputs:
- list[int]: Mistral-style 0-based indices. Converted to 1-based
and joined (e.g. [0,1,2] -> "1,2,3").
- list[str]: tokens like "1" or "3-5". Validated, joined as-is
(treated as Azure-native, i.e. 1-based).
- str: already in Azure format. Validated and whitespace-stripped.
"""
pages_pattern = re.compile(r"^\s*\d+(-\d+)?(\s*,\s*\d+(-\d+)?)*\s*$")
if isinstance(pages, str):
if not pages_pattern.match(pages):
raise ValueError(
f"Invalid `pages` string for Azure Document Intelligence: "
f"{pages!r}. Expected format like '1-3,5,7-9'."
)
return pages.replace(" ", "")
if isinstance(pages, list):
if len(pages) == 0:
return ""
if any(isinstance(p, bool) for p in pages):
raise ValueError("`pages` must be integers, not booleans")
if all(isinstance(p, int) for p in pages):
if any(p < 0 for p in pages):
raise ValueError(
"`pages` integers must be >= 0 (Mistral 0-based indices)"
)
# Mistral 0-based -> Azure 1-based.
return ",".join(str(p + 1) for p in sorted(set(pages)))
if all(isinstance(p, str) for p in pages):
joined = ",".join(p.strip() for p in pages)
if not pages_pattern.match(joined):
raise ValueError(
f"Invalid `pages` list for Azure Document Intelligence: "
f"{pages!r}. Expected tokens like '1' or '3-5'."
)
return joined
raise ValueError(
"`pages` must be a list[int] (0-based, Mistral-style) or a "
"string like '1-3,5,7-9'."
)
def validate_environment(
self,
headers: Dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> Dict:
"""
Validate environment and return headers for Azure Document Intelligence.
Authentication uses Ocp-Apim-Subscription-Key header.
"""
# Get API key from environment if not provided
if api_key is None:
api_key = get_secret_str(AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR)
if api_key is None:
raise ValueError(
"Missing Azure Document Intelligence API Key - Set AZURE_DOCUMENT_INTELLIGENCE_API_KEY environment variable or pass api_key parameter"
)
# Validate API base/endpoint is provided
if api_base is None:
api_base = get_secret_str("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
if api_base is None:
raise ValueError(
"Missing Azure Document Intelligence Endpoint - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT environment variable or pass api_base parameter"
)
headers = {
"Ocp-Apim-Subscription-Key": api_key,
"Content-Type": "application/json",
**headers,
}
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Azure Document Intelligence endpoint.
Format: {endpoint}/documentintelligence/documentModels/{modelId}:analyze?api-version=2024-11-30
Note: API version 2024-11-30 uses /documentintelligence/ path (not /formrecognizer/)
Args:
api_base: Azure Document Intelligence endpoint (e.g., https://your-resource.cognitiveservices.azure.com)
model: Model ID (e.g., "prebuilt-layout", "prebuilt-read")
optional_params: Optional parameters
Returns: Complete URL for Azure DI analyze endpoint
"""
if api_base is None:
api_base = get_secret_str("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
if api_base is None:
raise ValueError(
"Missing Azure Document Intelligence Endpoint - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT environment variable or pass api_base parameter"
)
# Ensure no trailing slash
api_base = api_base.rstrip("/")
# Extract model ID from full model path if needed
# Model can be "prebuilt-layout" or "azure_ai/doc-intelligence/prebuilt-layout"
model_id = model
if "/" in model:
# Extract the last part after the last slash
model_id = model.split("/")[-1]
encoded_model_id = encode_url_path_segment(model_id, field_name="model_id")
# Azure Document Intelligence analyze endpoint
# Note: API version 2024-11-30+ uses /documentintelligence/ (not /formrecognizer/)
url = (
f"{api_base}/documentintelligence/documentModels/{encoded_model_id}:analyze"
f"?api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}"
)
# Azure DI accepts `pages` as a query param (1-based, e.g. "1-3,5").
# `optional_params` has already been normalized in `map_ocr_params`.
pages = optional_params.get("pages") if optional_params else None
if pages:
url += f"&pages={quote(str(pages), safe=',-')}"
return url
def _extract_base64_from_data_uri(self, data_uri: str) -> str:
"""
Extract base64 content from a data URI.
Args:
data_uri: Data URI like "data:application/pdf;base64,..."
Returns:
Base64 string without the data URI prefix
"""
# Match pattern: data:[<mediatype>][;base64],<data>
match = re.match(r"data:([^;]+)(?:;base64)?,(.+)", data_uri)
if match:
return match.group(2)
return data_uri
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request to Azure Document Intelligence format.
Mistral OCR format:
{
"document": {
"type": "document_url",
"document_url": "https://example.com/doc.pdf"
}
}
Azure DI format:
{
"urlSource": "https://example.com/doc.pdf"
}
OR
{
"base64Source": "base64_encoded_content"
}
Args:
model: Model name
document: Document dict from user (Mistral format)
optional_params: Already mapped optional parameters
headers: Request headers
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug(
f"Azure Document Intelligence transform_ocr_request - model: {model}"
)
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Extract document URL from Mistral format
doc_type = document.get("type")
document_url = None
if doc_type == "document_url":
document_url = document.get("document_url", "")
elif doc_type == "image_url":
document_url = document.get("image_url", "")
else:
raise ValueError(
f"Invalid document type: {doc_type}. Must be 'document_url' or 'image_url'"
)
if not document_url:
raise ValueError("Document URL is required")
# Build Azure DI request
data: Dict[str, Any] = {}
# Check if it's a data URI (base64)
if document_url.startswith("data:"):
# Extract base64 content
base64_content = self._extract_base64_from_data_uri(document_url)
data["base64Source"] = base64_content
verbose_logger.debug("Using base64Source for Azure Document Intelligence")
else:
# Regular URL
data["urlSource"] = document_url
verbose_logger.debug("Using urlSource for Azure Document Intelligence")
# Azure DI: `pages` is a query param (wired in get_complete_url),
# not a body field. Other Mistral-specific params (e.g.
# include_image_base64, image_limit) are unsupported and ignored.
return OCRRequestData(data=data, files=None)
def _extract_page_markdown(self, page_data: Dict[str, Any]) -> str:
"""
Extract text from Azure DI page and format as markdown.
Azure DI provides text in 'lines' array. We concatenate them with newlines.
Args:
page_data: Azure DI page object
Returns:
Markdown-formatted text
"""
lines = page_data.get("lines", [])
if not lines:
return ""
# Extract text content from each line
text_lines = [line.get("content", "") for line in lines]
# Join with newlines to preserve structure
return "\n".join(text_lines)
def _convert_dimensions(
self, width: float, height: float, unit: str
) -> OCRPageDimensions:
"""
Convert Azure DI dimensions to pixels.
Azure DI provides dimensions in inches. We convert to pixels using configured DPI.
Args:
width: Width in specified unit
height: Height in specified unit
unit: Unit of measurement (e.g., "inch")
Returns:
OCRPageDimensions with pixel values
"""
# Convert to pixels using configured DPI
dpi = AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI
if unit == "inch":
width_px = int(width * dpi)
height_px = int(height * dpi)
else:
# If unit is not inches, assume it's already in pixels
width_px = int(width)
height_px = int(height)
return OCRPageDimensions(width=width_px, height=height_px, dpi=dpi)
@staticmethod
def _check_timeout(start_time: float, timeout_secs: int) -> None:
"""
Check if operation has timed out.
Args:
start_time: Start time of the operation
timeout_secs: Timeout duration in seconds
Raises:
TimeoutError: If operation has exceeded timeout
"""
if time.time() - start_time > timeout_secs:
raise TimeoutError(
f"Azure Document Intelligence operation polling timed out after {timeout_secs} seconds"
)
@staticmethod
def _get_retry_after(response: httpx.Response) -> int:
"""
Get retry-after duration from response headers.
Args:
response: HTTP response
Returns:
Retry-after duration in seconds (default: 2)
"""
retry_after = int(response.headers.get("retry-after", "2"))
verbose_logger.debug(f"Retry polling after: {retry_after} seconds")
return retry_after
@staticmethod
def _check_operation_status(response: httpx.Response) -> str:
"""
Check Azure DI operation status from response.
Args:
response: HTTP response from operation endpoint
Returns:
Operation status string
Raises:
ValueError: If operation failed or status is unknown
"""
try:
result = response.json()
status = result.get("status")
verbose_logger.debug(f"Azure DI operation status: {status}")
if status == "succeeded":
return "succeeded"
elif status == "failed":
error_msg = result.get("error", {}).get("message", "Unknown error")
raise ValueError(
f"Azure Document Intelligence analysis failed: {error_msg}"
)
elif status in ["running", "notStarted"]:
return "running"
else:
raise ValueError(f"Unknown operation status: {status}")
except Exception as e:
if "succeeded" in str(e) or "failed" in str(e):
raise
# If we can't parse JSON, something went wrong
raise ValueError(f"Failed to parse Azure DI operation response: {e}")
def _poll_operation_sync(
self,
operation_url: str,
headers: Dict[str, str],
timeout_secs: int,
) -> httpx.Response:
"""
Poll Azure Document Intelligence operation until completion (sync).
Azure DI POST returns 202 with Operation-Location header.
We need to poll that URL until status is "succeeded" or "failed".
Args:
operation_url: The Operation-Location URL to poll
headers: Request headers (including auth)
timeout_secs: Total timeout in seconds
Returns:
Final response with completed analysis
"""
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
client = _get_httpx_client()
start_time = time.time()
verbose_logger.debug(f"Polling Azure DI operation: {operation_url}")
while True:
self._check_timeout(start_time=start_time, timeout_secs=timeout_secs)
# Poll the operation status
response = client.get(url=operation_url, headers=headers)
# Check operation status
status = self._check_operation_status(response=response)
if status == "succeeded":
return response
elif status == "running":
# Wait before polling again
retry_after = self._get_retry_after(response=response)
time.sleep(retry_after)
async def _poll_operation_async(
self,
operation_url: str,
headers: Dict[str, str],
timeout_secs: int,
) -> httpx.Response:
"""
Poll Azure Document Intelligence operation until completion (async).
Args:
operation_url: The Operation-Location URL to poll
headers: Request headers (including auth)
timeout_secs: Total timeout in seconds
Returns:
Final response with completed analysis
"""
import litellm
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
client = get_async_httpx_client(llm_provider=litellm.LlmProviders.AZURE_AI)
start_time = time.time()
verbose_logger.debug(f"Polling Azure DI operation (async): {operation_url}")
while True:
self._check_timeout(start_time=start_time, timeout_secs=timeout_secs)
# Poll the operation status
response = await client.get(url=operation_url, headers=headers)
# Check operation status
status = self._check_operation_status(response=response)
if status == "succeeded":
return response
elif status == "running":
# Wait before polling again
retry_after = self._get_retry_after(response=response)
await asyncio.sleep(retry_after)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: Any,
**kwargs,
) -> OCRResponse:
"""
Transform Azure Document Intelligence response to Mistral OCR format.
Handles async operation polling: If response is 202 Accepted, polls Operation-Location
until analysis completes.
Azure DI response (after polling):
{
"status": "succeeded",
"analyzeResult": {
"content": "Full document text...",
"pages": [
{
"pageNumber": 1,
"width": 8.5,
"height": 11,
"unit": "inch",
"lines": [{"content": "text", "boundingBox": [...]}]
}
]
}
}
Mistral OCR format:
{
"pages": [
{
"index": 0,
"markdown": "extracted text",
"dimensions": {"width": 816, "height": 1056, "dpi": 96}
}
],
"model": "azure_ai/doc-intelligence/prebuilt-layout",
"usage_info": {"pages_processed": 1},
"object": "ocr"
}
Args:
model: Model name
raw_response: Raw HTTP response from Azure DI (may be 202 Accepted)
logging_obj: Logging object
Returns:
OCRResponse in Mistral format
"""
try:
# Check if we got 202 Accepted (async operation started)
if raw_response.status_code == 202:
verbose_logger.debug(
"Azure DI returned 202 Accepted, polling operation..."
)
# Get Operation-Location header
operation_url = raw_response.headers.get("Operation-Location")
if not operation_url:
raise ValueError(
"Azure Document Intelligence returned 202 but no Operation-Location header found"
)
# Reject cross-origin polling URLs — the auth headers
# below would otherwise leak to whatever URL the upstream
# (or an attacker-controlled upstream) returns. VERIA-51.
try:
assert_same_origin(operation_url, str(raw_response.request.url))
except SSRFError as ssrf_err:
raise ValueError(
f"Azure Document Intelligence: rejected polling URL ({ssrf_err})"
)
# Get headers for polling (need auth)
poll_headers = {
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get(
"Ocp-Apim-Subscription-Key", ""
)
}
# Get timeout from kwargs or use default
timeout_secs = AZURE_OPERATION_POLLING_TIMEOUT
# Poll until operation completes
raw_response = self._poll_operation_sync(
operation_url=operation_url,
headers=poll_headers,
timeout_secs=timeout_secs,
)
# Now parse the completed response
response_json = raw_response.json()
verbose_logger.debug(
f"Azure Document Intelligence response status: {response_json.get('status')}"
)
# Check if request succeeded
status = response_json.get("status")
if status != "succeeded":
raise ValueError(
f"Azure Document Intelligence analysis failed with status: {status}"
)
# Extract analyze result
analyze_result = response_json.get("analyzeResult", {})
azure_pages = analyze_result.get("pages", [])
# Transform pages to Mistral format
mistral_pages = []
for azure_page in azure_pages:
page_number = azure_page.get("pageNumber", 1)
index = page_number - 1 # Convert to 0-based index
# Extract markdown text
markdown = self._extract_page_markdown(azure_page)
# Convert dimensions
width = azure_page.get("width", 8.5)
height = azure_page.get("height", 11)
unit = azure_page.get("unit", "inch")
dimensions = self._convert_dimensions(
width=width, height=height, unit=unit
)
# Build OCR page
ocr_page = OCRPage(
index=index, markdown=markdown, dimensions=dimensions
)
mistral_pages.append(ocr_page)
# Build usage info
usage_info = OCRUsageInfo(
pages_processed=len(mistral_pages), doc_size_bytes=None
)
# Return Mistral OCR response
return OCRResponse(
pages=mistral_pages,
model=model,
usage_info=usage_info,
object="ocr",
)
except Exception as e:
verbose_logger.error(
f"Error parsing Azure Document Intelligence response: {e}"
)
raise e
async def async_transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: Any,
**kwargs,
) -> OCRResponse:
"""
Async transform Azure Document Intelligence response to Mistral OCR format.
Handles async operation polling: If response is 202 Accepted, polls Operation-Location
until analysis completes using async polling.
Args:
model: Model name
raw_response: Raw HTTP response from Azure DI (may be 202 Accepted)
logging_obj: Logging object
Returns:
OCRResponse in Mistral format
"""
try:
# Check if we got 202 Accepted (async operation started)
if raw_response.status_code == 202:
verbose_logger.debug(
"Azure DI returned 202 Accepted, polling operation (async)..."
)
# Get Operation-Location header
operation_url = raw_response.headers.get("Operation-Location")
if not operation_url:
raise ValueError(
"Azure Document Intelligence returned 202 but no Operation-Location header found"
)
# Reject cross-origin polling URLs (see sync path). VERIA-51.
try:
assert_same_origin(operation_url, str(raw_response.request.url))
except SSRFError as ssrf_err:
raise ValueError(
f"Azure Document Intelligence: rejected polling URL ({ssrf_err})"
)
# Get headers for polling (need auth)
poll_headers = {
"Ocp-Apim-Subscription-Key": raw_response.request.headers.get(
"Ocp-Apim-Subscription-Key", ""
)
}
# Get timeout from kwargs or use default
timeout_secs = AZURE_OPERATION_POLLING_TIMEOUT
# Poll until operation completes (async)
raw_response = await self._poll_operation_async(
operation_url=operation_url,
headers=poll_headers,
timeout_secs=timeout_secs,
)
# Now parse the completed response
response_json = raw_response.json()
verbose_logger.debug(
f"Azure Document Intelligence response status: {response_json.get('status')}"
)
# Check if request succeeded
status = response_json.get("status")
if status != "succeeded":
raise ValueError(
f"Azure Document Intelligence analysis failed with status: {status}"
)
# Extract analyze result
analyze_result = response_json.get("analyzeResult", {})
azure_pages = analyze_result.get("pages", [])
# Transform pages to Mistral format
mistral_pages = []
for azure_page in azure_pages:
page_number = azure_page.get("pageNumber", 1)
index = page_number - 1 # Convert to 0-based index
# Extract markdown text
markdown = self._extract_page_markdown(azure_page)
# Convert dimensions
width = azure_page.get("width", 8.5)
height = azure_page.get("height", 11)
unit = azure_page.get("unit", "inch")
dimensions = self._convert_dimensions(
width=width, height=height, unit=unit
)
# Build OCR page
ocr_page = OCRPage(
index=index, markdown=markdown, dimensions=dimensions
)
mistral_pages.append(ocr_page)
# Build usage info
usage_info = OCRUsageInfo(
pages_processed=len(mistral_pages), doc_size_bytes=None
)
# Return Mistral OCR response
return OCRResponse(
pages=mistral_pages,
model=model,
usage_info=usage_info,
object="ocr",
)
except Exception as e:
verbose_logger.error(
f"Error parsing Azure Document Intelligence response (async): {e}"
)
raise e

View file

@ -1,287 +0,0 @@
"""
Azure AI OCR transformation implementation.
"""
from typing import Dict
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.image_handling import (
async_convert_url_to_base64,
convert_url_to_base64,
)
from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRRequestData
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.secret_managers.main import get_secret_str
AZURE_AI_OCR_API_KEY_ENV_VAR = "AZURE_AI_API_KEY"
class AzureAIOCRConfig(MistralOCRConfig):
"""
Azure AI OCR transformation configuration.
Azure AI uses Mistral's OCR API but with a different endpoint format.
Inherits transformation logic from MistralOCRConfig since they use the same format.
Reference: Azure AI Foundry OCR documentation
Important: Azure AI only supports base64 data URIs (data:image/..., data:application/pdf;base64,...).
Regular URLs are not supported.
"""
def __init__(self) -> None:
super().__init__()
def get_api_key_env_var(self) -> str | None:
return AZURE_AI_OCR_API_KEY_ENV_VAR
def validate_environment(
self,
headers: Dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> Dict:
"""
Validate environment and return headers for Azure AI OCR.
Azure AI uses Bearer token authentication with AZURE_AI_API_KEY.
"""
# Get API key from environment if not provided
if api_key is None:
api_key = get_secret_str(AZURE_AI_OCR_API_KEY_ENV_VAR)
if api_key is None:
raise ValueError(
"Missing Azure AI API Key - A call is being made to Azure AI but no key is set either in the environment variables or via params"
)
# Validate API base is provided
if api_base is None:
api_base = get_secret_str("AZURE_AI_API_BASE")
if api_base is None:
raise ValueError(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter"
)
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
**headers,
}
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Azure AI OCR endpoint.
Azure AI endpoint format: https://<api_base>/providers/mistral/azure/ocr
Args:
api_base: Azure AI API base URL
model: Model name (not used in URL construction)
optional_params: Optional parameters
Returns: Complete URL for Azure AI OCR endpoint
"""
if api_base is None:
raise ValueError(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter"
)
# Ensure no trailing slash
api_base = api_base.rstrip("/")
# Azure AI OCR endpoint format
return f"{api_base}/providers/mistral/azure/ocr"
def _convert_url_to_data_uri_sync(self, url: str) -> str:
"""
Synchronously convert a URL to a base64 data URI.
Azure AI OCR doesn't have internet access, so we need to fetch URLs
and convert them to base64 data URIs.
Args:
url: The URL to convert
Returns:
Base64 data URI string
"""
verbose_logger.debug(
f"Azure AI OCR: Converting URL to base64 data URI (sync): {url}"
)
# Fetch and convert to base64 data URI
# convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
data_uri = convert_url_to_base64(url=url)
verbose_logger.debug(
f"Azure AI OCR: Converted URL to data URI (length: {len(data_uri)})"
)
return data_uri
async def _convert_url_to_data_uri_async(self, url: str) -> str:
"""
Asynchronously convert a URL to a base64 data URI.
Azure AI OCR doesn't have internet access, so we need to fetch URLs
and convert them to base64 data URIs.
Args:
url: The URL to convert
Returns:
Base64 data URI string
"""
verbose_logger.debug(
f"Azure AI OCR: Converting URL to base64 data URI (async): {url}"
)
# Fetch and convert to base64 data URI asynchronously
# async_convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
data_uri = await async_convert_url_to_base64(url=url)
verbose_logger.debug(
f"Azure AI OCR: Converted URL to data URI (length: {len(data_uri)})"
)
return data_uri
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Azure AI, converting URLs to base64 data URIs (sync).
Azure AI OCR doesn't have internet access, so we automatically fetch
any URLs and convert them to base64 data URIs synchronously.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug(
f"Azure AI OCR transform_ocr_request (sync) - model: {model}"
)
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Check if we need to convert URL to base64
doc_type = document.get("type")
transformed_document = document.copy()
if doc_type == "document_url":
document_url = document.get("document_url", "")
# If it's not already a data URI, convert it
if document_url and not document_url.startswith("data:"):
verbose_logger.debug(
"Azure AI OCR: Converting document URL to base64 data URI (sync)"
)
data_uri = self._convert_url_to_data_uri_sync(url=document_url)
transformed_document["document_url"] = data_uri
elif doc_type == "image_url":
image_url = document.get("image_url", "")
# If it's not already a data URI, convert it
if image_url and not image_url.startswith("data:"):
verbose_logger.debug(
"Azure AI OCR: Converting image URL to base64 data URI (sync)"
)
data_uri = self._convert_url_to_data_uri_sync(url=image_url)
transformed_document["image_url"] = data_uri
# Call parent's transform to build the request
return super().transform_ocr_request(
model=model,
document=transformed_document,
optional_params=optional_params,
headers=headers,
**kwargs,
)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Azure AI, converting URLs to base64 data URIs (async).
Azure AI OCR doesn't have internet access, so we automatically fetch
any URLs and convert them to base64 data URIs asynchronously.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug(
f"Azure AI OCR async_transform_ocr_request - model: {model}"
)
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Check if we need to convert URL to base64
doc_type = document.get("type")
transformed_document = document.copy()
if doc_type == "document_url":
document_url = document.get("document_url", "")
# If it's not already a data URI, convert it
if document_url and not document_url.startswith("data:"):
verbose_logger.debug(
"Azure AI OCR: Converting document URL to base64 data URI (async)"
)
data_uri = await self._convert_url_to_data_uri_async(url=document_url)
transformed_document["document_url"] = data_uri
elif doc_type == "image_url":
image_url = document.get("image_url", "")
# If it's not already a data URI, convert it
if image_url and not image_url.startswith("data:"):
verbose_logger.debug(
"Azure AI OCR: Converting image URL to base64 data URI (async)"
)
data_uri = await self._convert_url_to_data_uri_async(url=image_url)
transformed_document["image_url"] = data_uri
# Call parent's transform to build the request
return super().transform_ocr_request(
model=model,
document=transformed_document,
optional_params=optional_params,
headers=headers,
**kwargs,
)

View file

@ -1,23 +1,17 @@
"""Base OCR transformation module."""
from .transformation import (
BaseOCRConfig,
DocumentType,
OCRPage,
OCRPageDimensions,
OCRPageImage,
OCRRequestData,
OCRResponse,
OCRUsageInfo,
)
__all__ = [
"BaseOCRConfig",
"DocumentType",
"OCRResponse",
"OCRPage",
"OCRPageDimensions",
"OCRPageImage",
"OCRUsageInfo",
"OCRRequestData",
]

View file

@ -1,26 +1,11 @@
"""
Base OCR transformation configuration.
"""
"""OCR response models shared by the Python public API and proxy."""
from typing import TYPE_CHECKING, Any, Dict, List, Union
from typing import Any, Dict, List
import httpx
from pydantic import PrivateAttr
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.base import LiteLLMPydanticObjectBase
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
# DocumentType for OCR - providers always receive a dict with
# type="document_url" or type="image_url" (str values only).
# File-type inputs are preprocessed to this format in litellm/ocr/main.py.
DocumentType = Dict[str, str]
class OCRPageDimensions(LiteLLMPydanticObjectBase):
"""Page dimensions from OCR response."""
@ -76,191 +61,3 @@ class OCRResponse(LiteLLMPydanticObjectBase):
# Define private attributes using PrivateAttr
_hidden_params: dict = PrivateAttr(default_factory=dict)
class OCRRequestData(LiteLLMPydanticObjectBase):
"""OCR request data structure."""
data: Union[Dict, bytes] | None = None
files: Dict[str, Any] | None = None
class BaseOCRConfig:
"""
Base configuration for OCR transformations.
Handles provider-agnostic OCR operations.
"""
def __init__(self) -> None:
pass
def get_supported_ocr_params(self, model: str) -> list:
"""
Get supported OCR parameters for this provider.
Override this method in provider-specific implementations.
"""
return []
def get_api_key_env_var(self) -> str | None:
"""
Return the provider-specific API key environment variable name, if any.
"""
return None
def map_ocr_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
) -> dict:
"""Map OCR parameters to provider-specific parameters."""
return optional_params
def validate_environment(
self,
headers: Dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> Dict:
"""
Validate environment and return headers.
Override in provider-specific implementations.
"""
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for OCR endpoint.
Override in provider-specific implementations.
"""
raise NotImplementedError("get_complete_url must be implemented by provider")
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request to provider-specific format.
Override in provider-specific implementations.
Note: By the time this method is called, any file-type documents have already
been converted to document_url/image_url format with base64 data URIs by
the preprocessing in litellm/ocr/main.py.
Args:
model: Model name
document: Document to process - always a dict with type="document_url" or type="image_url"
optional_params: Optional parameters for the request
headers: Request headers
Returns:
OCRRequestData with data and files fields
"""
raise NotImplementedError(
"transform_ocr_request must be implemented by provider"
)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Async transform OCR request to provider-specific format.
Optional method - providers can override if they need async transformations
(e.g., Azure AI for URL-to-base64 conversion).
Default implementation falls back to sync transform_ocr_request.
Args:
model: Model name
document: Document to process (Mistral format dict, or file path, bytes, etc.)
optional_params: Optional parameters for the request
headers: Request headers
Returns:
OCRRequestData with data and files fields
"""
# Default implementation: call sync version
return self.transform_ocr_request(
model=model,
document=document,
optional_params=optional_params,
headers=headers,
**kwargs,
)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs,
) -> OCRResponse:
"""
Transform provider-specific OCR response to standard format.
Override in provider-specific implementations.
"""
raise NotImplementedError(
"transform_ocr_response must be implemented by provider"
)
async def async_transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs,
) -> OCRResponse:
"""
Async transform provider-specific OCR response to standard format.
Optional method - providers can override if they need async transformations
(e.g., Azure Document Intelligence for async operation polling).
Default implementation falls back to sync transform_ocr_response.
Args:
model: Model name
raw_response: Raw HTTP response
logging_obj: Logging object
Returns:
OCRResponse in standard format
"""
# Default implementation: call sync version
return self.transform_ocr_response(
model=model,
raw_response=raw_response,
logging_obj=logging_obj,
**kwargs,
)
def get_error_class(
self,
error_message: str,
status_code: int,
headers: dict,
) -> Exception:
"""Get appropriate error class for the provider."""
return BaseLLMException(
status_code=status_code,
message=error_message,
headers=headers,
)

View file

@ -55,7 +55,6 @@ from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, OCRResponse
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
@ -1456,301 +1455,6 @@ class BaseLLMHTTPHandler:
api_key=api_key,
)
def _prepare_ocr_request(
self,
model: str,
document: Dict[str, str],
optional_params: dict,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_base: Optional[str],
headers: Optional[Dict[str, Any]],
provider_config: BaseOCRConfig,
litellm_params: dict,
) -> Tuple[Dict[str, Any], str, Dict[str, Any], None]:
"""
Shared logic for preparing OCR requests.
Returns: (headers, complete_url, data, files)
"""
from litellm.llms.base_llm.ocr.transformation import OCRRequestData
headers = provider_config.validate_environment(
api_key=api_key,
api_base=api_base,
headers=headers or {},
model=model,
litellm_params=litellm_params,
)
complete_url = provider_config.get_complete_url(
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
)
# Transform the request to get data and files
transformed_result = provider_config.transform_ocr_request(
model=model,
document=document,
optional_params=optional_params,
headers=headers,
api_key=api_key,
api_base=api_base,
)
# All providers return OCRRequestData
if not isinstance(transformed_result, OCRRequestData):
raise ValueError(
f"Provider {provider_config.__class__.__name__} must return OCRRequestData"
)
# Data is always a dict for Mistral OCR format
if not isinstance(transformed_result.data, dict):
raise ValueError(
f"Expected dict data for OCR request, got {type(transformed_result.data)}"
)
data = transformed_result.data
## LOGGING
logging_obj.pre_call(
input="OCR document processing",
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": complete_url,
"headers": headers,
},
)
return headers, complete_url, data, None
async def _async_prepare_ocr_request(
self,
model: str,
document: Dict[str, str],
optional_params: dict,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_base: Optional[str],
headers: Optional[Dict[str, Any]],
provider_config: BaseOCRConfig,
litellm_params: dict,
) -> Tuple[Dict[str, Any], str, Dict[str, Any], None]:
"""
Async version of _prepare_ocr_request for providers that need async transforms.
Returns: (headers, complete_url, data, files)
"""
from litellm.llms.base_llm.ocr.transformation import OCRRequestData
headers = provider_config.validate_environment(
api_key=api_key,
api_base=api_base,
headers=headers or {},
model=model,
litellm_params=litellm_params,
)
complete_url = provider_config.get_complete_url(
api_base=api_base,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
)
# Use async transform (providers can override this method if they need async operations)
transformed_result = await provider_config.async_transform_ocr_request(
model=model,
document=document,
optional_params=optional_params,
headers=headers,
api_key=api_key,
api_base=api_base,
)
# All providers return OCRRequestData
if not isinstance(transformed_result, OCRRequestData):
raise ValueError(
f"Provider {provider_config.__class__.__name__} must return OCRRequestData"
)
# Data is always a dict for Mistral OCR format
if not isinstance(transformed_result.data, dict):
raise ValueError(
f"Expected dict data for OCR request, got {type(transformed_result.data)}"
)
data = transformed_result.data
## LOGGING
logging_obj.pre_call(
input="OCR document processing",
api_key=api_key,
additional_args={
"complete_input_dict": data,
"api_base": complete_url,
"headers": headers,
},
)
return headers, complete_url, data, None
def _transform_ocr_response(
self,
provider_config: BaseOCRConfig,
model: str,
response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
) -> OCRResponse:
"""Shared logic for transforming OCR responses."""
return provider_config.transform_ocr_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
def ocr(
self,
model: str,
document: Dict[str, str],
optional_params: dict,
timeout: Union[float, httpx.Timeout],
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_base: Optional[str],
custom_llm_provider: str,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
aocr: bool = False,
headers: Optional[Dict[str, Any]] = None,
provider_config: Optional[BaseOCRConfig] = None,
litellm_params: Optional[dict] = None,
) -> Union[OCRResponse, Coroutine[Any, Any, OCRResponse]]:
"""
Sync OCR handler.
"""
if provider_config is None:
raise ValueError(
f"No provider config found for model: {model} and provider: {custom_llm_provider}"
)
if litellm_params is None:
litellm_params = {}
if aocr is True:
return self.async_ocr(
model=model,
document=document,
optional_params=optional_params,
timeout=timeout,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
client=client,
headers=headers,
provider_config=provider_config,
litellm_params=litellm_params,
)
# Prepare the request
headers, complete_url, data, files = self._prepare_ocr_request(
model=model,
document=document,
optional_params=optional_params,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
headers=headers,
provider_config=provider_config,
litellm_params=litellm_params,
)
if client is None or not isinstance(client, HTTPHandler):
client = _get_httpx_client()
try:
# Make the POST request with JSON data (Mistral format)
response = client.post(
url=complete_url,
headers=headers,
json=data,
timeout=timeout,
)
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
return self._transform_ocr_response(
provider_config=provider_config,
model=model,
response=response,
logging_obj=logging_obj,
)
async def async_ocr(
self,
model: str,
document: Dict[str, str],
optional_params: dict,
timeout: Union[float, httpx.Timeout],
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
api_base: Optional[str],
custom_llm_provider: str,
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
headers: Optional[Dict[str, Any]] = None,
provider_config: Optional[BaseOCRConfig] = None,
litellm_params: Optional[dict] = None,
) -> OCRResponse:
"""
Async OCR handler.
"""
if provider_config is None:
raise ValueError(
f"No provider config found for model: {model} and provider: {custom_llm_provider}"
)
if litellm_params is None:
litellm_params = {}
# Prepare the request using async prepare method
headers, complete_url, data, files = await self._async_prepare_ocr_request(
model=model,
document=document,
optional_params=optional_params,
logging_obj=logging_obj,
api_key=api_key,
api_base=api_base,
headers=headers,
provider_config=provider_config,
litellm_params=litellm_params,
)
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders(custom_llm_provider),
)
else:
async_httpx_client = client
try:
# Make the async POST request with JSON data (Mistral format)
response = await async_httpx_client.post(
url=complete_url,
headers=headers,
json=data,
timeout=timeout,
)
except Exception as e:
raise self._handle_error(e=e, provider_config=provider_config)
# Use async response transform for async operations
return await provider_config.async_transform_ocr_response(
model=model,
raw_response=response,
logging_obj=logging_obj,
)
def search(
self,
query: Union[str, List[str]],
@ -5871,7 +5575,6 @@ class BaseLLMHTTPHandler:
BaseGoogleGenAIGenerateContentConfig,
BaseAnthropicMessagesConfig,
BaseBatchesConfig,
BaseOCRConfig,
BaseVideoConfig,
BaseSearchConfig,
BaseTextToSpeechConfig,

View file

@ -1,239 +0,0 @@
"""
Mistral OCR transformation implementation.
"""
from typing import Any, Dict
import httpx
from litellm._logging import verbose_logger
from litellm.llms.base_llm.ocr.transformation import (
BaseOCRConfig,
DocumentType,
OCRRequestData,
OCRResponse,
)
from litellm.secret_managers.main import get_secret_str
MISTRAL_OCR_API_KEY_ENV_VAR = "MISTRAL_API_KEY"
class MistralOCRConfig(BaseOCRConfig):
"""
Mistral OCR transformation configuration.
Reference: https://docs.mistral.ai/api/#tag/ocr
"""
def __init__(self) -> None:
super().__init__()
def get_supported_ocr_params(self, model: str) -> list:
"""
Get supported OCR parameters for Mistral OCR.
Mistral OCR supports:
- pages: List of page numbers to process
- include_image_base64: Whether to include base64 encoded images
- image_limit: Maximum number of images to return
- image_min_size: Minimum size of images to include
- bbox_annotation_format: Format for bounding box annotations
- document_annotation_format: Format for document annotations
- document_annotation_prompt: Prompt for document annotation extraction
- extract_header: Whether to extract document header
- extract_footer: Whether to extract document footer
- table_format: Table output format ("markdown" or "html")
- confidence_scores_granularity: Confidence score level ("word" or "page")
- id: Request identifier
"""
return [
"pages",
"include_image_base64",
"image_limit",
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
"document_annotation_prompt",
"extract_header",
"extract_footer",
"table_format",
"confidence_scores_granularity",
"id",
]
def get_api_key_env_var(self) -> str | None:
return MISTRAL_OCR_API_KEY_ENV_VAR
def map_ocr_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
) -> dict:
"""
Map OCR parameters to Mistral-specific format.
Mistral accepts these parameters directly, so no transformation needed.
Just filter out unsupported params.
"""
supported_params = self.get_supported_ocr_params(model=model)
# Only include params that are in the supported list
mapped_params = {}
for param, value in non_default_params.items():
if param in supported_params:
mapped_params[param] = value
return mapped_params
def validate_environment(
self,
headers: Dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> Dict:
"""
Validate environment and return headers for Mistral OCR.
"""
# Get API key from environment if not provided
if api_key is None:
api_key = get_secret_str(MISTRAL_OCR_API_KEY_ENV_VAR)
if api_key is None:
raise ValueError(
"Missing Mistral API Key - A call is being made to Mistral but no key is set either in the environment variables or via params"
)
headers = {
"Authorization": f"Bearer {api_key}",
**headers,
}
# Don't set Content-Type for multipart/form-data - httpx will handle it
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Mistral OCR endpoint.
Returns: https://api.mistral.ai/v1/ocr
"""
if api_base is None:
api_base = "https://api.mistral.ai/v1"
# Ensure no trailing slash
api_base = api_base.rstrip("/")
# Remove /v1 if it's already in the base to avoid duplication
if api_base.endswith("/v1"):
return f"{api_base}/ocr"
return f"{api_base}/v1/ocr"
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request to Mistral-specific format.
Mistral OCR API accepts:
{
"model": "mistral-ocr-latest",
"document": {
"type": "document_url",
"document_url": "<https-url or data-uri>"
},
"pages": [0], # optional
"include_image_base64": false, # optional
...
}
Args:
model: Model name (e.g., "mistral-ocr-latest")
document: Document dict from user (Mistral format) - already validated in main.py
optional_params: Already mapped optional parameters
headers: Request headers
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug(f"Mistral OCR transform_ocr_request - model: {model}")
# Document parameter is the Mistral-format dict from the user
# Just pass it through as-is to the Mistral API
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Build request data - use document dict directly
data = {
"model": model,
"document": document, # Pass through the Mistral-format document dict
}
# Add all optional parameters from the already-mapped optional_params
data.update(optional_params)
# No multipart files - using JSON
return OCRRequestData(data=data, files=None)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: Any,
**kwargs,
) -> OCRResponse:
"""
Return Mistral OCR response in native format.
Mistral OCR is the standard format for LiteLLM OCR responses.
No transformation needed - return native response.
Mistral OCR returns:
{
"pages": [
{
"index": 0,
"markdown": "extracted text content",
"images": [...],
"dimensions": {...}
},
...
],
"model": "mistral-ocr-2505-completion",
"document_annotation": null,
"usage_info": {...}
}
"""
try:
response_json = raw_response.json()
verbose_logger.debug(f"Mistral OCR response keys: {response_json.keys()}")
# Return native Mistral format - no transformation
return OCRResponse(
pages=response_json.get("pages", []),
model=response_json.get("model", model),
document_annotation=response_json.get("document_annotation"),
usage_info=response_json.get("usage_info"),
object="ocr",
)
except Exception as e:
verbose_logger.error(f"Error parsing Mistral OCR response: {e}")
raise e

View file

@ -1,159 +0,0 @@
import base64
import binascii
from collections import defaultdict
from typing import TYPE_CHECKING, Any, Dict, List, NoReturn, Optional, Tuple
from litellm.constants import request_timeout
REDUCTO_API_BASE = "https://platform.reducto.ai"
REDUCTO_ID_PREFIX = "reducto://"
if TYPE_CHECKING:
from litellm.llms.base_llm.ocr.transformation import OCRPage
def _normalize_api_base(api_base: Optional[str]) -> str:
return (api_base or REDUCTO_API_BASE).rstrip("/")
def _raise_bad_request(message: str, model: str) -> NoReturn:
import litellm
raise litellm.BadRequestError(
message=message,
model=model,
llm_provider="reducto",
)
def extract_file_id_or_bytes(
source_url: str,
model: str,
) -> Tuple[Optional[str], Optional[bytes], Optional[str]]:
if source_url.startswith(REDUCTO_ID_PREFIX):
return source_url, None, None
if source_url.startswith("http://") or source_url.startswith("https://"):
_raise_bad_request(
"Reducto requires type='file' (auto-uploaded) or a reducto:// id. Plain http(s) URLs are not supported; upload the file first.",
model=model,
)
if not source_url.startswith("data:"):
_raise_bad_request(
"Reducto requires a reducto:// id or a base64 data URI after OCR preprocessing.",
model=model,
)
try:
header, encoded = source_url.split(",", 1)
except ValueError:
_raise_bad_request("Invalid Reducto data URI provided.", model=model)
if ";base64" not in header:
_raise_bad_request(
"Reducto only supports base64-encoded data URIs.", model=model
)
mime = header.removeprefix("data:").split(";")[0] or "application/octet-stream"
try:
raw_bytes = base64.b64decode(encoded, validate=True)
except (binascii.Error, ValueError):
_raise_bad_request("Invalid Reducto base64 payload provided.", model=model)
return None, raw_bytes, mime
def _extract_file_id_from_upload_response(response: Any) -> str:
try:
payload = response.json()
except ValueError as exc:
raise ValueError(
"Reducto /upload returned a non-JSON 200 response: {}".format(response.text)
) from exc
file_id = (payload or {}).get("file_id") if isinstance(payload, dict) else None
if not isinstance(file_id, str) or not file_id:
raise ValueError(
"Reducto /upload returned 200 without a file_id; got payload={}".format(
payload
)
)
return file_id
def upload_bytes_sync(
raw_bytes: bytes,
mime: Optional[str],
api_key: str,
api_base: Optional[str],
) -> str:
import litellm
response = litellm.module_level_client.post(
url="{}{}".format(_normalize_api_base(api_base), "/upload"),
headers={"Authorization": f"Bearer {api_key}"},
files={"file": ("document", raw_bytes, mime or "application/octet-stream")},
timeout=request_timeout,
)
response.raise_for_status()
return _extract_file_id_from_upload_response(response)
async def upload_bytes_async(
raw_bytes: bytes,
mime: Optional[str],
api_key: str,
api_base: Optional[str],
) -> str:
import litellm
response = await litellm.module_level_aclient.post(
url="{}{}".format(_normalize_api_base(api_base), "/upload"),
headers={"Authorization": f"Bearer {api_key}"},
files={"file": ("document", raw_bytes, mime or "application/octet-stream")},
timeout=request_timeout,
)
response.raise_for_status()
return _extract_file_id_from_upload_response(response)
def build_pages_from_reducto(result: Dict[str, Any]) -> List["OCRPage"]:
from litellm.llms.base_llm.ocr.transformation import OCRPage
chunks = result.get("chunks", []) or []
blocks_by_page: Dict[int, List[Dict[str, Any]]] = defaultdict(list)
for chunk in chunks:
for block in chunk.get("blocks", []) or []:
page_no = (block.get("bbox") or {}).get("page")
if page_no is None:
continue
try:
normalized_page = int(page_no)
except (TypeError, ValueError):
continue
blocks_by_page[normalized_page].append(block)
if not blocks_by_page:
fallback_markdown = "\n\n".join(
chunk.get("content", "") for chunk in chunks if chunk.get("content")
)
if fallback_markdown == "":
return []
return [OCRPage(index=0, markdown=fallback_markdown)]
pages: List["OCRPage"] = []
for page_no, blocks in sorted(blocks_by_page.items()):
markdown = "\n\n".join(
block.get("content", "") for block in blocks if block.get("content")
)
page_index = max(page_no - 1, 0)
page = OCRPage(
index=page_index,
markdown=markdown,
)
# OCRPage accepts extra keys at runtime; assign blocks after construction
# so static typing does not reject provider-specific metadata.
setattr(page, "blocks", blocks)
pages.append(page)
return pages

View file

@ -1,241 +0,0 @@
from typing import Any, Dict, Optional, Tuple
import httpx
from litellm.llms.base_llm.ocr.transformation import (
BaseOCRConfig,
DocumentType,
OCRRequestData,
OCRResponse,
OCRUsageInfo,
)
from litellm.llms.reducto.common import (
REDUCTO_API_BASE,
build_pages_from_reducto,
extract_file_id_or_bytes,
upload_bytes_async,
upload_bytes_sync,
)
class _BaseReductoOCRConfig(BaseOCRConfig):
def map_ocr_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
) -> dict:
mapped_params = dict(optional_params)
supported_params = self.get_supported_ocr_params(model=model)
for param, value in non_default_params.items():
if param in supported_params:
mapped_params[param] = value
return mapped_params
def validate_environment(
self,
headers: Dict,
model: str,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
litellm_params: Optional[dict] = None,
**kwargs,
) -> Dict:
from litellm.secret_managers.main import get_secret_str
resolved_key = api_key or get_secret_str("REDUCTO_API_KEY")
if resolved_key is None:
raise ValueError(
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
)
return {
"Authorization": f"Bearer {resolved_key}",
"Content-Type": "application/json",
**headers,
}
def get_complete_url(
self,
api_base: Optional[str],
model: str,
optional_params: dict,
litellm_params: Optional[dict] = None,
**kwargs,
) -> str:
return "{}/parse".format((api_base or REDUCTO_API_BASE).rstrip("/"))
def _get_source_url(self, document: DocumentType, model: str) -> str:
source_url = document.get("document_url") or document.get("image_url")
if source_url is None:
raise ValueError(
"Reducto expected OCR preprocessing to produce document_url or image_url for model={}".format(
model
)
)
return source_url
@staticmethod
def _resolve_credentials(
api_key: Optional[str], api_base: Optional[str]
) -> Tuple[str, str]:
from litellm.secret_managers.main import get_secret_str
resolved_key = api_key or get_secret_str("REDUCTO_API_KEY")
if resolved_key is None:
raise ValueError(
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
)
resolved_base = (api_base or REDUCTO_API_BASE).rstrip("/")
return resolved_key, resolved_base
def _ensure_file_id_sync(
self,
model: str,
document: DocumentType,
api_key: Optional[str],
api_base: Optional[str],
) -> str:
source_url = self._get_source_url(document=document, model=model)
file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model)
if file_id is not None:
return file_id
resolved_key, resolved_base = self._resolve_credentials(api_key, api_base)
return upload_bytes_sync(
raw_bytes=raw_bytes or b"",
mime=mime,
api_key=resolved_key,
api_base=resolved_base,
)
async def _ensure_file_id_async(
self,
model: str,
document: DocumentType,
api_key: Optional[str],
api_base: Optional[str],
) -> str:
source_url = self._get_source_url(document=document, model=model)
file_id, raw_bytes, mime = extract_file_id_or_bytes(source_url, model=model)
if file_id is not None:
return file_id
resolved_key, resolved_base = self._resolve_credentials(api_key, api_base)
return await upload_bytes_async(
raw_bytes=raw_bytes or b"",
mime=mime,
api_key=resolved_key,
api_base=resolved_base,
)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: Any,
**kwargs,
) -> OCRResponse:
response_json = raw_response.json()
result = response_json.get("result", response_json) or {}
usage = response_json.get("usage", {}) or {}
response = OCRResponse(
pages=build_pages_from_reducto(result),
model=model,
usage_info=OCRUsageInfo(
pages_processed=usage.get("num_pages"),
credits=usage.get("credits"),
),
object="ocr",
)
response._hidden_params["reducto_raw"] = response_json
return response
class ReductoParseV3Config(_BaseReductoOCRConfig):
def get_supported_ocr_params(self, model: str) -> list:
return ["formatting", "retrieval", "settings"]
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id = self._ensure_file_id_sync(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(data={"input": file_id, **optional_params}, files=None)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id = await self._ensure_file_id_async(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(data={"input": file_id, **optional_params}, files=None)
class ReductoParseLegacyConfig(_BaseReductoOCRConfig):
def get_supported_ocr_params(self, model: str) -> list:
return ["enhance"]
def _build_legacy_body(self, file_id: str, optional_params: dict) -> Dict[str, Any]:
body: Dict[str, Any] = {"document_url": file_id}
enhance = optional_params.get("enhance")
if enhance is not None:
body["options"] = {"enhance": enhance}
return body
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id = self._ensure_file_id_sync(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(
data=self._build_legacy_body(
file_id=file_id, optional_params=optional_params
),
files=None,
)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
file_id = await self._ensure_file_id_async(
model=model,
document=document,
api_key=kwargs.get("api_key"),
api_base=kwargs.get("api_base"),
)
return OCRRequestData(
data=self._build_legacy_body(
file_id=file_id, optional_params=optional_params
),
files=None,
)

View file

@ -1,5 +1,3 @@
"""Vertex AI OCR module."""
from .transformation import VertexAIOCRConfig
__all__ = ["VertexAIOCRConfig"]
__all__: list[str] = []

View file

@ -1,41 +0,0 @@
"""
Common utilities for Vertex AI OCR providers.
This module provides routing logic to determine which OCR configuration to use
based on the model name.
"""
from typing import TYPE_CHECKING, Optional
if TYPE_CHECKING:
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
def get_vertex_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
"""
Determine which Vertex AI OCR configuration to use based on the model name.
Vertex AI supports multiple OCR services:
- Vertex AI OCR: vertex_ai/<model>
Args:
model: The model name (e.g., "vertex_ai/ocr/<model>")
Returns:
OCR configuration instance for the specified model
Examples:
>>> get_vertex_ai_ocr_config("vertex_ai/deepseek-ai/deepseek-ocr-maas")
<VertexAIDeepSeekOCRConfig object>
>>> get_vertex_ai_ocr_config("vertex_ai/ocr/mistral-ocr-maas")
<VertexAIOCRConfig object>
"""
from litellm.llms.vertex_ai.ocr.deepseek_transformation import (
VertexAIDeepSeekOCRConfig,
)
from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig
if "deepseek" in model:
return VertexAIDeepSeekOCRConfig()
return VertexAIOCRConfig()

View file

@ -1,399 +0,0 @@
"""
Vertex AI DeepSeek OCR transformation implementation.
"""
import json
from typing import TYPE_CHECKING, Any, Dict
import httpx
from litellm._logging import verbose_logger
from litellm.llms.base_llm.ocr.transformation import (
BaseOCRConfig,
DocumentType,
OCRPage,
OCRRequestData,
OCRResponse,
OCRUsageInfo,
)
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
VERTEX_AI_DEEPSEEK_OCR_API_KEY_ENV_VAR = "VERTEX_AI_API_KEY"
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class VertexAIDeepSeekOCRConfig(BaseOCRConfig):
"""
Vertex AI DeepSeek OCR transformation configuration.
This transformation converts standard LiteLLM OCR requests to the
Vertex AI DeepSeek OCR OpenAPI endpoint shape and normalizes the response.
"""
def __init__(self) -> None:
super().__init__()
self.vertex_base = VertexBase()
def get_api_key_env_var(self) -> str | None:
return VERTEX_AI_DEEPSEEK_OCR_API_KEY_ENV_VAR
def validate_environment(
self,
headers: Dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> Dict:
"""
Validate environment and return headers for Vertex AI OCR.
Vertex AI uses Bearer token authentication with access token from credentials.
"""
if api_key is not None:
return {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
**headers,
}
# Extract Vertex AI parameters using safe helpers from VertexBase
# Use safe_get_* methods that don't mutate litellm_params dict
litellm_params = litellm_params or {}
vertex_project = VertexBase.safe_get_vertex_ai_project(
litellm_params=litellm_params
)
vertex_credentials = VertexBase.safe_get_vertex_ai_credentials(
litellm_params=litellm_params
)
# Get access token from Vertex credentials
access_token, project_id = self.vertex_base.get_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
)
headers = {
"Authorization": f"Bearer {access_token}",
"Content-Type": "application/json",
**headers,
}
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Vertex AI DeepSeek OCR endpoint.
Args:
api_base: Vertex AI API base URL (optional)
model: Model name (e.g., "deepseek-ai/deepseek-ocr-maas")
optional_params: Optional parameters
litellm_params: LiteLLM parameters containing vertex_project, vertex_location
Returns: Complete URL for Vertex AI OCR endpoint
"""
# Extract Vertex AI parameters using safe helpers from VertexBase
# Use safe_get_* methods that don't mutate litellm_params dict
litellm_params = litellm_params or {}
vertex_project = VertexBase.safe_get_vertex_ai_project(
litellm_params=litellm_params
)
vertex_location = VertexBase.safe_get_vertex_ai_location(
litellm_params=litellm_params
)
if vertex_project is None:
raise ValueError(
"Missing vertex_project - Set VERTEXAI_PROJECT environment variable or pass vertex_project parameter"
)
if vertex_location is None:
vertex_location = "us-central1"
# Get API base URL
if api_base is None:
api_base = "https://aiplatform.googleapis.com"
# Ensure no trailing slash
api_base = api_base.rstrip("/")
return f"{api_base}/v1/projects/{vertex_project}/locations/{vertex_location}/endpoints/openapi/chat/completions"
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Vertex AI DeepSeek OCR.
Converts OCR document format to the Vertex AI DeepSeek OCR payload:
- Input: {"type": "image_url", "image_url": "gs://..."}
- Output: {"model": "deepseek-ai/deepseek-ocr-maas", "messages": [{"role": "user", "content": [{"type": "image_url", "image_url": "gs://..."}]}]}
Args:
model: Model name (e.g., "deepseek-ai/deepseek-ocr-maas")
document: Document dict from user (Mistral OCR format)
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data for the DeepSeek OCR endpoint
"""
verbose_logger.debug(
"Vertex AI DeepSeek OCR transform_ocr_request (sync) called"
)
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Extract document type and URL
doc_type = document.get("type")
image_url = None
document_url = None
if doc_type == "image_url":
image_url = document.get("image_url", "")
elif doc_type == "document_url":
document_url = document.get("document_url", "")
else:
raise ValueError(
f"Unsupported document type: {doc_type}. Expected 'image_url' or 'document_url'"
)
# Build DeepSeek OCR message content
content_item = {}
if image_url:
content_item = {"type": "image_url", "image_url": image_url}
elif document_url:
# For document URLs, we use image_url type as well (Vertex AI supports both)
content_item = {"type": "image_url", "image_url": document_url}
# Build DeepSeek OCR request
data = {
"model": "deepseek-ai/" + model,
"messages": [{"role": "user", "content": [content_item]}],
}
# Add optional parameters (stream, temperature, etc.)
deepseek_ocr_params = {}
for key, value in optional_params.items():
if key in ["stream", "temperature", "max_tokens", "top_p", "n", "stop"]:
deepseek_ocr_params[key] = value
data.update(deepseek_ocr_params)
verbose_logger.debug("Vertex AI DeepSeek OCR: Transformed request")
return OCRRequestData(data=data, files=None)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Vertex AI DeepSeek OCR (async).
Same as sync version - no async-specific logic needed.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data for the DeepSeek OCR endpoint
"""
return self.transform_ocr_request(
model=model,
document=document,
optional_params=optional_params,
headers=headers,
**kwargs,
)
def transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs,
) -> OCRResponse:
"""
Transform Vertex AI DeepSeek OCR response to OCR format.
Vertex AI DeepSeek OCR returns an OpenAPI response:
{
"id": "...",
"choices": [{
"message": {
"role": "assistant",
"content": "<OCR result as JSON string or markdown>"
}
}],
"usage": {...}
}
We need to extract the content and convert it to OCRResponse format.
Args:
model: Model name
raw_response: Raw HTTP response from Vertex AI
logging_obj: Logging object
**kwargs: Additional arguments
Returns:
OCRResponse in standard format
"""
verbose_logger.debug("Vertex AI DeepSeek OCR transform_ocr_response called")
verbose_logger.debug(f"Raw response: {raw_response.text}")
try:
response_json = raw_response.json()
# Extract OCR content from provider response
choices = response_json.get("choices", [])
if not choices:
raise ValueError("No choices in DeepSeek OCR response")
message = choices[0].get("message", {})
content = message.get("content", "")
if not content:
raise ValueError("No content in DeepSeek OCR response")
# Try to parse content as JSON (OCR result might be JSON string)
ocr_data = None
try:
# If content is a JSON string, parse it
if isinstance(content, str) and content.strip().startswith("{"):
ocr_data = json.loads(content)
elif isinstance(content, dict):
ocr_data = content
else:
# If content is markdown text, create a single page with the markdown
ocr_data = {
"pages": [{"index": 0, "markdown": content}],
"model": model,
"usage_info": response_json.get("usage", {}),
}
except json.JSONDecodeError:
# If JSON parsing fails, treat content as markdown
ocr_data = {
"pages": [{"index": 0, "markdown": content}],
"model": model,
"usage_info": response_json.get("usage", {}),
}
# Ensure we have the expected structure
if "pages" not in ocr_data:
# If OCR data doesn't have pages, wrap the content in a page
ocr_data = {
"pages": [
{
"index": 0,
"markdown": (
content
if isinstance(content, str)
else json.dumps(content)
),
}
],
"model": ocr_data.get("model", model),
"usage_info": ocr_data.get(
"usage_info", response_json.get("usage", {})
),
}
# Convert usage info if present
usage_info = None
if "usage_info" in ocr_data:
usage_dict = ocr_data["usage_info"]
if isinstance(usage_dict, dict):
usage_info = OCRUsageInfo(**usage_dict)
# Build OCRResponse
pages = []
for page_data in ocr_data.get("pages", []):
# Ensure page has required fields
if isinstance(page_data, dict):
page = OCRPage(
index=page_data.get("index", 0),
markdown=page_data.get("markdown", ""),
images=page_data.get("images"),
dimensions=page_data.get("dimensions"),
)
pages.append(page)
if not pages:
# Create a default page if none exist
pages = [
OCRPage(
index=0, markdown=content if isinstance(content, str) else ""
)
]
return OCRResponse(
pages=pages,
model=ocr_data.get("model", model),
document_annotation=ocr_data.get("document_annotation"),
usage_info=usage_info,
object="ocr",
)
except Exception as e:
verbose_logger.error(f"Error parsing Vertex AI DeepSeek OCR response: {e}")
raise e
async def async_transform_ocr_response(
self,
model: str,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs,
) -> OCRResponse:
"""
Async transform Vertex AI DeepSeek OCR response to OCR format.
Same as sync version - no async-specific logic needed.
Args:
model: Model name
raw_response: Raw HTTP response
logging_obj: Logging object
**kwargs: Additional arguments
Returns:
OCRResponse in standard format
"""
return self.transform_ocr_response(
model=model,
raw_response=raw_response,
logging_obj=logging_obj,
**kwargs,
)

View file

@ -1,314 +0,0 @@
"""
Vertex AI Mistral OCR transformation implementation.
"""
from typing import Dict
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.image_handling import (
async_convert_url_to_base64,
convert_url_to_base64,
)
from litellm.llms.base_llm.ocr.transformation import DocumentType, OCRRequestData
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
VERTEX_AI_OCR_API_KEY_ENV_VAR = "VERTEX_AI_API_KEY"
class VertexAIOCRConfig(MistralOCRConfig):
"""
Vertex AI Mistral OCR transformation configuration.
Vertex AI uses Mistral's OCR API format through the Mistral publisher endpoint.
Inherits transformation logic from MistralOCRConfig since they use the same format.
Reference: Vertex AI Mistral OCR documentation
Important: Vertex AI OCR only supports base64 data URIs (data:image/..., data:application/pdf;base64,...).
Regular URLs are not supported.
"""
def __init__(self) -> None:
super().__init__()
self.vertex_base = VertexBase()
def get_api_key_env_var(self) -> str | None:
return VERTEX_AI_OCR_API_KEY_ENV_VAR
def validate_environment(
self,
headers: Dict,
model: str,
api_key: str | None = None,
api_base: str | None = None,
litellm_params: dict | None = None,
**kwargs,
) -> Dict:
"""
Validate environment and return headers for Vertex AI OCR.
Vertex AI uses Bearer token authentication with access token from credentials.
"""
if api_key is not None:
return {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
**headers,
}
# Extract Vertex AI parameters using safe helpers from VertexBase
# Use safe_get_* methods that don't mutate litellm_params dict
litellm_params = litellm_params or {}
vertex_project = VertexBase.safe_get_vertex_ai_project(
litellm_params=litellm_params
)
vertex_credentials = VertexBase.safe_get_vertex_ai_credentials(
litellm_params=litellm_params
)
# Get access token from Vertex credentials
access_token, project_id = self.vertex_base.get_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
)
headers = {
"Authorization": f"Bearer {access_token}",
"Content-Type": "application/json",
**headers,
}
return headers
def get_complete_url(
self,
api_base: str | None,
model: str,
optional_params: dict,
litellm_params: dict | None = None,
**kwargs,
) -> str:
"""
Get complete URL for Vertex AI OCR endpoint.
Vertex AI endpoint format:
https://{location}-aiplatform.googleapis.com/v1/projects/{project}/locations/{location}/publishers/mistralai/ocr
Args:
api_base: Vertex AI API base URL (optional)
model: Model name (not used in URL construction)
optional_params: Optional parameters
litellm_params: LiteLLM parameters containing vertex_project, vertex_location
Returns: Complete URL for Vertex AI OCR endpoint
"""
# Extract Vertex AI parameters using safe helpers from VertexBase
# Use safe_get_* methods that don't mutate litellm_params dict
litellm_params = litellm_params or {}
vertex_project = VertexBase.safe_get_vertex_ai_project(
litellm_params=litellm_params
)
vertex_location = VertexBase.safe_get_vertex_ai_location(
litellm_params=litellm_params
)
if vertex_project is None:
raise ValueError(
"Missing vertex_project - Set VERTEXAI_PROJECT environment variable or pass vertex_project parameter"
)
if vertex_location is None:
vertex_location = "us-central1"
# Get API base URL
if api_base is None:
api_base = get_vertex_base_url(vertex_location)
# Ensure no trailing slash
api_base = api_base.rstrip("/")
# Vertex AI OCR endpoint format for Mistral publisher
# Format: https://{region}-aiplatform.googleapis.com/v1/projects/{project}/locations/{region}/publishers/mistralai/models/{model}:rawPredict
return f"{api_base}/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/mistralai/models/{model}:rawPredict"
def _convert_url_to_data_uri_sync(self, url: str) -> str:
"""
Synchronously convert a URL to a base64 data URI.
Vertex AI OCR doesn't have internet access, so we need to fetch URLs
and convert them to base64 data URIs.
Args:
url: The URL to convert
Returns:
Base64 data URI string
"""
verbose_logger.debug(
f"Vertex AI OCR: Converting URL to base64 data URI (sync): {url}"
)
# Fetch and convert to base64 data URI
# convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
data_uri = convert_url_to_base64(url=url)
verbose_logger.debug(
f"Vertex AI OCR: Converted URL to data URI (length: {len(data_uri)})"
)
return data_uri
async def _convert_url_to_data_uri_async(self, url: str) -> str:
"""
Asynchronously convert a URL to a base64 data URI.
Vertex AI OCR doesn't have internet access, so we need to fetch URLs
and convert them to base64 data URIs.
Args:
url: The URL to convert
Returns:
Base64 data URI string
"""
verbose_logger.debug(
f"Vertex AI OCR: Converting URL to base64 data URI (async): {url}"
)
# Fetch and convert to base64 data URI asynchronously
# async_convert_url_to_base64 already returns a full data URI like "data:image/jpeg;base64,..."
data_uri = await async_convert_url_to_base64(url=url)
verbose_logger.debug(
f"Vertex AI OCR: Converted URL to data URI (length: {len(data_uri)})"
)
return data_uri
def transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Vertex AI, converting URLs to base64 data URIs (sync).
Vertex AI OCR doesn't have internet access, so we automatically fetch
any URLs and convert them to base64 data URIs synchronously.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug("Vertex AI OCR transform_ocr_request (sync) called")
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Check if we need to convert URL to base64
doc_type = document.get("type")
transformed_document = document.copy()
if doc_type == "document_url":
document_url = document.get("document_url", "")
# If it's not already a data URI, convert it
if document_url and not document_url.startswith("data:"):
verbose_logger.debug(
"Vertex AI OCR: Converting document URL to base64 data URI (sync)"
)
data_uri = self._convert_url_to_data_uri_sync(url=document_url)
transformed_document["document_url"] = data_uri
elif doc_type == "image_url":
image_url = document.get("image_url", "")
# If it's not already a data URI, convert it
if image_url and not image_url.startswith("data:"):
verbose_logger.debug(
"Vertex AI OCR: Converting image URL to base64 data URI (sync)"
)
data_uri = self._convert_url_to_data_uri_sync(url=image_url)
transformed_document["image_url"] = data_uri
# Call parent's transform to build the request
return super().transform_ocr_request(
model=model,
document=transformed_document,
optional_params=optional_params,
headers=headers,
**kwargs,
)
async def async_transform_ocr_request(
self,
model: str,
document: DocumentType,
optional_params: dict,
headers: dict,
**kwargs,
) -> OCRRequestData:
"""
Transform OCR request for Vertex AI, converting URLs to base64 data URIs (async).
Vertex AI OCR doesn't have internet access, so we automatically fetch
any URLs and convert them to base64 data URIs asynchronously.
Args:
model: Model name
document: Document dict from user
optional_params: Already mapped optional parameters
headers: Request headers
**kwargs: Additional arguments
Returns:
OCRRequestData with JSON data
"""
verbose_logger.debug(
f"Vertex AI OCR async_transform_ocr_request - model: {model}"
)
if not isinstance(document, dict):
raise ValueError(f"Expected document dict, got {type(document)}")
# Check if we need to convert URL to base64
doc_type = document.get("type")
transformed_document = document.copy()
if doc_type == "document_url":
document_url = document.get("document_url", "")
# If it's not already a data URI, convert it
if document_url and not document_url.startswith("data:"):
verbose_logger.debug(
"Vertex AI OCR: Converting document URL to base64 data URI (async)"
)
data_uri = await self._convert_url_to_data_uri_async(url=document_url)
transformed_document["document_url"] = data_uri
elif doc_type == "image_url":
image_url = document.get("image_url", "")
# If it's not already a data URI, convert it
if image_url and not image_url.startswith("data:"):
verbose_logger.debug(
"Vertex AI OCR: Converting image URL to base64 data URI (async)"
)
data_uri = await self._convert_url_to_data_uri_async(url=image_url)
transformed_document["image_url"] = data_uri
# Call parent's transform to build the request
return super().transform_ocr_request(
model=model,
document=transformed_document,
optional_params=optional_params,
headers=headers,
**kwargs,
)

View file

@ -6,9 +6,8 @@ import base64
import mimetypes
import os
import re
from dataclasses import dataclass
from io import IOBase
from typing import Any, Callable, Coroutine, Union, cast
from typing import Any, Coroutine, Union, cast
import httpx
@ -23,30 +22,7 @@ from litellm.ocr.rust_bridge import (
load_rust_aocr,
load_rust_ocr,
)
from litellm.types.router import GenericLiteLLMParams
from litellm.utils import client
@dataclass
class _PreparedOCRRequest:
model: str
document: dict[str, Any]
api_key: str | None
api_base: str | None
custom_llm_provider: str
extra_headers: dict[str, object] | None
optional_params: dict[str, object]
litellm_params: dict[str, object]
effective_timeout: Union[float, httpx.Timeout]
litellm_logging_obj: LiteLLMLoggingObj
@dataclass
class _PreparedRustOCRCall:
api_key: str | None
api_base: str | None
headers: dict[str, object]
optional_params: dict[str, object]
from litellm.utils import client, filter_out_litellm_params
def _timeout_to_seconds(
@ -65,7 +41,7 @@ def _timeout_to_seconds(
return float(timeout)
def _prepare_ocr_request(
def _resolve_ocr_call_context(
model: str,
document: dict[str, Any],
api_key: str | None,
@ -74,7 +50,17 @@ def _prepare_ocr_request(
custom_llm_provider: str | None,
extra_headers: dict[str, Any] | None,
kwargs: dict[str, Any],
) -> _PreparedOCRRequest:
) -> tuple[
str,
dict[str, Any],
str | None,
str | None,
str,
dict[str, object] | None,
dict[str, object],
Union[float, httpx.Timeout],
LiteLLMLoggingObj,
]:
litellm_logging_obj = cast(LiteLLMLoggingObj, kwargs.pop("litellm_logging_obj"))
litellm_call_id = cast(str | None, kwargs.get("litellm_call_id", None))
@ -114,10 +100,13 @@ def _prepare_ocr_request(
verbose_logger.debug(f"OCR call - model: {model}, provider: {custom_llm_provider}")
litellm_params = GenericLiteLLMParams(**kwargs)
optional_params = dict(kwargs)
optional_params = {
key: value
for key, value in filter_out_litellm_params(kwargs=kwargs).items()
if key not in _RUST_BRIDGE_INTERNAL_PARAMS
}
verbose_logger.debug(f"OCR optional_params after mapping: {optional_params}")
verbose_logger.debug(f"OCR optional_params forwarded to Rust: {optional_params}")
effective_timeout = timeout or request_timeout
@ -132,142 +121,74 @@ def _prepare_ocr_request(
custom_llm_provider=custom_llm_provider,
)
return _PreparedOCRRequest(
model=model,
document=document,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=cast(dict[str, object] | None, extra_headers),
optional_params=cast(dict[str, object], optional_params),
litellm_params=dict(litellm_params),
effective_timeout=effective_timeout,
litellm_logging_obj=litellm_logging_obj,
return (
model,
document,
api_key,
api_base,
custom_llm_provider,
cast(dict[str, object] | None, extra_headers),
cast(dict[str, object], optional_params),
effective_timeout,
litellm_logging_obj,
)
def _rust_bridge_optional_params(
prepared_request: _PreparedOCRRequest,
resolve_secret: Callable[[str], str | None],
) -> dict[str, object]:
optional_params = dict(prepared_request.optional_params)
if prepared_request.custom_llm_provider == "vertex_ai":
vertex_project = (
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")
)
vertex_location = (
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
def _is_azure_document_intelligence(prepared_request: _PreparedOCRRequest) -> bool:
return prepared_request.custom_llm_provider == "azure_ai/doc-intelligence" or (
prepared_request.custom_llm_provider == "azure_ai"
and (
"doc-intelligence" in prepared_request.model
or "documentintelligence" in prepared_request.model
)
)
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 _is_azure_document_intelligence(prepared_request):
return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT")
if prepared_request.custom_llm_provider == "azure_ai":
return resolve_secret("AZURE_AI_API_BASE")
return None
def _rust_bridge_api_key(
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> str | None:
if prepared_request.api_key is not None:
return prepared_request.api_key
provider = prepared_request.custom_llm_provider
if provider == "mistral":
return resolve_api_key("MISTRAL_API_KEY")
if provider == "azure_ai/doc-intelligence" or _is_azure_document_intelligence(
prepared_request
):
return resolve_api_key("AZURE_DOCUMENT_INTELLIGENCE_API_KEY")
if provider == "azure_ai":
return resolve_api_key("AZURE_AI_API_KEY")
if provider == "vertex_ai":
return resolve_api_key("VERTEX_AI_API_KEY") or resolve_api_key("VERTEXAI_API_KEY")
if provider == "reducto":
return resolve_api_key("REDUCTO_API_KEY")
return None
def _prepare_rust_ocr_call(
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
) -> _PreparedRustOCRCall:
resolved_api_key = _rust_bridge_api_key(prepared_request, resolve_api_key)
rust_api_base = _rust_bridge_api_base(prepared_request, resolve_api_key)
rust_optional_params = _rust_bridge_optional_params(
prepared_request, resolve_api_key
)
resolved_headers = prepared_request.extra_headers or {}
prepared_request.litellm_logging_obj.pre_call(
def _run_pre_call_logging(
litellm_logging_obj: LiteLLMLoggingObj,
model: str,
document: dict[str, Any],
api_key: str | None,
api_base: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
) -> None:
litellm_logging_obj.pre_call(
input="OCR document processing",
api_key=resolved_api_key,
api_key=api_key,
additional_args={
"complete_input_dict": {
"model": prepared_request.model,
"document": prepared_request.document,
**rust_optional_params,
"model": model,
"document": document,
**optional_params,
},
"api_base": rust_api_base or prepared_request.api_base,
"headers": resolved_headers,
"api_base": api_base,
"headers": extra_headers or {},
},
)
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 _run_rust_ocr(
rust_ocr: RustOcr,
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
model: str,
document: dict[str, Any],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: Union[float, httpx.Timeout],
litellm_logging_obj: LiteLLMLoggingObj,
) -> OCRResponse:
prepared = _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
_run_pre_call_logging(
litellm_logging_obj=litellm_logging_obj,
model=model,
document=document,
api_key=api_key,
api_base=api_base,
extra_headers=extra_headers,
optional_params=optional_params,
)
return OCRResponse.model_validate(
rust_ocr(
model=prepared_request.model,
document=cast(dict[str, object], 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_seconds=_timeout_to_seconds(prepared_request.effective_timeout),
model=model,
document=cast(dict[str, object], document),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=_timeout_to_seconds(timeout),
)
)
@ -280,23 +201,35 @@ def _missing_rust_bridge_error() -> RuntimeError:
async def _run_rust_aocr(
rust_aocr: RustAocr,
prepared_request: _PreparedOCRRequest,
resolve_api_key: Callable[[str], str | None],
model: str,
document: dict[str, Any],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: Union[float, httpx.Timeout],
litellm_logging_obj: LiteLLMLoggingObj,
) -> OCRResponse:
prepared = _prepare_rust_ocr_call(
prepared_request=prepared_request,
resolve_api_key=resolve_api_key,
_run_pre_call_logging(
litellm_logging_obj=litellm_logging_obj,
model=model,
document=document,
api_key=api_key,
api_base=api_base,
extra_headers=extra_headers,
optional_params=optional_params,
)
return OCRResponse.model_validate(
await rust_aocr(
model=prepared_request.model,
document=cast(dict[str, object], 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_seconds=_timeout_to_seconds(prepared_request.effective_timeout),
model=model,
document=cast(dict[str, object], document),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=_timeout_to_seconds(timeout),
)
)
@ -381,7 +314,17 @@ async def aocr(
"kwargs": kwargs,
}
try:
prepared = _prepare_ocr_request(
(
model,
document,
api_key,
api_base,
custom_llm_provider,
extra_headers,
optional_params,
effective_timeout,
litellm_logging_obj,
) = _resolve_ocr_call_context(
model=model,
document=document,
api_key=api_key,
@ -391,8 +334,6 @@ async def aocr(
extra_headers=extra_headers,
kwargs=kwargs,
)
model = prepared.model
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update(
{"model": model, "custom_llm_provider": custom_llm_provider}
)
@ -401,12 +342,17 @@ async def aocr(
if rust_aocr is None:
raise _missing_rust_bridge_error()
from litellm.secret_managers.main import get_secret_str
return await _run_rust_aocr(
rust_aocr=rust_aocr,
prepared_request=prepared,
resolve_api_key=get_secret_str,
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=effective_timeout,
litellm_logging_obj=litellm_logging_obj,
)
except Exception as e:
raise litellm.exception_type(
@ -424,6 +370,8 @@ async def aocr(
_MIME_PATTERN = re.compile(r"^[\w.+-]+/[\w.+-]+$")
_RUST_BRIDGE_INTERNAL_PARAMS = {"original_generic_function"}
_MIME_TYPE_MAP = {
".pdf": "application/pdf",
".png": "image/png",
@ -629,7 +577,17 @@ def ocr(
try:
_is_async = kwargs.pop("aocr", False) is True
completion_kwargs["aocr"] = _is_async
prepared = _prepare_ocr_request(
(
model,
document,
api_key,
api_base,
custom_llm_provider,
extra_headers,
optional_params,
effective_timeout,
litellm_logging_obj,
) = _resolve_ocr_call_context(
model=model,
document=document,
api_key=api_key,
@ -639,8 +597,6 @@ def ocr(
extra_headers=extra_headers,
timeout=timeout,
)
model = prepared.model
custom_llm_provider = prepared.custom_llm_provider
completion_kwargs.update(
{"model": model, "custom_llm_provider": custom_llm_provider}
)
@ -649,12 +605,17 @@ def ocr(
if rust_ocr is None:
raise _missing_rust_bridge_error()
from litellm.secret_managers.main import get_secret_str
return _run_rust_ocr(
rust_ocr=rust_ocr,
prepared_request=prepared,
resolve_api_key=get_secret_str,
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=effective_timeout,
litellm_logging_obj=litellm_logging_obj,
)
except Exception as e:
raise litellm.exception_type(

View file

@ -301,7 +301,6 @@ if TYPE_CHECKING:
from litellm.llms.base_llm.google_genai.transformation import (
BaseGoogleGenAIGenerateContentConfig,
)
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig
from litellm.llms.base_llm.sandbox.transformation import BaseSandboxConfig
from litellm.llms.base_llm.search.transformation import BaseSearchConfig
from litellm.llms.base_llm.text_to_speech.transformation import (
@ -309,7 +308,6 @@ if TYPE_CHECKING:
)
from litellm.llms.bedrock.common_utils import BedrockModelInfo
from litellm.llms.cohere.common_utils import CohereModelInfo
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
# Type stubs for lazy-loaded functions and classes
from litellm.litellm_core_utils.cached_imports import (
@ -9640,48 +9638,6 @@ class ProviderConfigManager:
return get_openrouter_image_edit_config(model)
return None
@staticmethod
def get_provider_ocr_config(
model: str,
provider: LlmProviders,
) -> Optional["BaseOCRConfig"]:
"""
Get OCR configuration for a given provider.
"""
from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig
# Special handling for Azure AI - distinguish between Mistral OCR and Document Intelligence
if provider == litellm.LlmProviders.AZURE_AI:
from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config
return get_azure_ai_ocr_config(model=model)
if provider == litellm.LlmProviders.VERTEX_AI:
from litellm.llms.vertex_ai.ocr.common_utils import get_vertex_ai_ocr_config
return get_vertex_ai_ocr_config(model=model)
if provider == litellm.LlmProviders.REDUCTO:
from litellm.llms.reducto.ocr.transformation import (
ReductoParseLegacyConfig,
ReductoParseV3Config,
)
if model == "parse-v3":
return ReductoParseV3Config()
if model == "parse-legacy":
return ReductoParseLegacyConfig()
return None
MistralOCRConfig = getattr(sys.modules[__name__], "MistralOCRConfig")
PROVIDER_TO_CONFIG_MAP = {
litellm.LlmProviders.MISTRAL: MistralOCRConfig,
}
config_class = PROVIDER_TO_CONFIG_MAP.get(provider, None)
if config_class is None:
return None
return config_class()
@staticmethod
def get_provider_search_config(
provider: "SearchProviders",

View file

@ -10,10 +10,6 @@ import os
import pytest
from base_ocr_unit_tests import BaseOCRTest
from litellm.constants import AZURE_DOCUMENT_INTELLIGENCE_API_VERSION
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
class TestAzureDocumentIntelligenceOCR(BaseOCRTest):
@ -47,125 +43,3 @@ class TestAzureDocumentIntelligenceOCR(BaseOCRTest):
"api_base": endpoint,
}
class TestAzureDocumentIntelligencePagesParam:
"""
Unit tests for the Mistral-compatible `pages` parameter translation to
Azure Document Intelligence's `pages` query string.
These tests exercise the transformation layer directly and do not
require Azure credentials or a network call.
"""
@pytest.fixture
def cfg(self) -> AzureDocumentIntelligenceOCRConfig:
return AzureDocumentIntelligenceOCRConfig()
def test_get_supported_ocr_params_includes_pages(self, cfg):
assert cfg.get_supported_ocr_params("prebuilt-layout") == ["pages"]
def test_map_ocr_params_mistral_zero_based_int_list(self, cfg):
mapped = cfg.map_ocr_params({"pages": [0, 1, 2]}, {}, "prebuilt-layout")
assert mapped == {"pages": "1,2,3"}
def test_map_ocr_params_dedupes_and_sorts(self, cfg):
mapped = cfg.map_ocr_params({"pages": [2, 0, 0, 1]}, {}, "prebuilt-layout")
assert mapped == {"pages": "1,2,3"}
def test_map_ocr_params_empty_list_omits_pages(self, cfg):
mapped = cfg.map_ocr_params({"pages": []}, {}, "prebuilt-layout")
assert mapped == {}
def test_map_ocr_params_azure_native_string_range(self, cfg):
mapped = cfg.map_ocr_params({"pages": "3-9"}, {}, "prebuilt-layout")
assert mapped == {"pages": "3-9"}
def test_map_ocr_params_azure_native_string_with_spaces_stripped(self, cfg):
mapped = cfg.map_ocr_params({"pages": "1-3, 5"}, {}, "prebuilt-layout")
assert mapped == {"pages": "1-3,5"}
def test_map_ocr_params_list_of_string_tokens(self, cfg):
mapped = cfg.map_ocr_params({"pages": ["1", "3-5"]}, {}, "prebuilt-layout")
assert mapped == {"pages": "1,3-5"}
def test_map_ocr_params_invalid_string_raises(self, cfg):
with pytest.raises(ValueError, match="Invalid `pages` string"):
cfg.map_ocr_params({"pages": "a,b"}, {}, "prebuilt-layout")
def test_map_ocr_params_negative_index_raises(self, cfg):
with pytest.raises(ValueError, match="must be >= 0"):
cfg.map_ocr_params({"pages": [-1]}, {}, "prebuilt-layout")
def test_map_ocr_params_bool_list_raises(self, cfg):
with pytest.raises(ValueError, match="must be integers, not booleans"):
cfg.map_ocr_params({"pages": [True, False]}, {}, "prebuilt-layout")
def test_map_ocr_params_unsupported_type_raises(self, cfg):
with pytest.raises(ValueError):
cfg.map_ocr_params({"pages": 5}, {}, "prebuilt-layout")
def test_get_complete_url_appends_pages_query(self, cfg):
url = cfg.get_complete_url(
api_base="https://example.cognitiveservices.azure.com/",
model="azure_ai/doc-intelligence/prebuilt-layout",
optional_params={"pages": "1-3,5"},
)
assert (
f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url
), url
assert "pages=1-3,5" in url, url
assert "/documentintelligence/documentModels/prebuilt-layout:analyze" in url
def test_get_complete_url_no_pages_when_optional_params_empty(self, cfg):
url = cfg.get_complete_url(
api_base="https://example.cognitiveservices.azure.com",
model="prebuilt-layout",
optional_params={},
)
assert "pages=" not in url
def test_transform_ocr_request_does_not_put_pages_in_body(self, cfg):
req = cfg.transform_ocr_request(
model="prebuilt-layout",
document={
"type": "document_url",
"document_url": "https://example.com/x.pdf",
},
optional_params={"pages": "1,2,3"},
headers={},
)
assert req.data is not None
assert "pages" not in req.data
assert req.data.get("urlSource") == "https://example.com/x.pdf"
def test_end_to_end_mistral_shape_to_azure_query(self, cfg):
"""
Caller sends Mistral-style `pages: [2,3,4,5,6,7,8]` (0-based,
meaning human pages 3-9). LiteLLM should turn that into Azure's
`&pages=3,4,5,6,7,8,9` on the analyze URL, and the body should
still only contain urlSource.
"""
non_default_params = {"pages": [2, 3, 4, 5, 6, 7, 8]}
optional_params = cfg.map_ocr_params(
non_default_params=non_default_params,
optional_params={},
model="prebuilt-layout",
)
url = cfg.get_complete_url(
api_base="https://example.cognitiveservices.azure.com",
model="prebuilt-layout",
optional_params=optional_params,
)
req = cfg.transform_ocr_request(
model="prebuilt-layout",
document={
"type": "document_url",
"document_url": "https://example.com/x.pdf",
},
optional_params=optional_params,
headers={},
)
assert "pages=3,4,5,6,7,8,9" in url
assert req.data == {"urlSource": "https://example.com/x.pdf"}

View file

@ -15,7 +15,6 @@ from base_ocr_unit_tests import BaseOCRTest
def load_vertex_ai_credentials():
"""Load Vertex AI credentials for tests"""
# Define the path to the vertex_key.json file
print("loading vertex ai credentials")
filepath = os.path.dirname(os.path.abspath(__file__))
vertex_key_path = filepath + "/vertex_key.json"
@ -23,7 +22,6 @@ def load_vertex_ai_credentials():
try:
with open(vertex_key_path, "r") as file:
# Read the file content
print("Read vertexai file path")
content = file.read()
# If the file is empty or not valid JSON, create an empty dictionary
@ -110,32 +108,3 @@ class TestVertexAIDeepSeekOCR(BaseOCRTest):
def test_ocr_response_structure(self):
"""Skip this test for DeepSeek OCR - PDF URLs not supported"""
pass
def test_vertex_ai_ocr_routing():
"""
Test that Vertex AI OCR routing correctly selects the right config based on model name.
"""
from litellm.llms.vertex_ai.ocr.common_utils import get_vertex_ai_ocr_config
from litellm.llms.vertex_ai.ocr.deepseek_transformation import (
VertexAIDeepSeekOCRConfig,
)
from litellm.llms.vertex_ai.ocr.transformation import VertexAIOCRConfig
# Test DeepSeek OCR routing
deepseek_config = get_vertex_ai_ocr_config("vertex_ai/deepseek-ocr-maas")
assert isinstance(
deepseek_config, VertexAIDeepSeekOCRConfig
), "DeepSeek model should route to VertexAIDeepSeekOCRConfig"
# Test Mistral OCR routing (should use default VertexAIOCRConfig)
mistral_config = get_vertex_ai_ocr_config("vertex_ai/mistral-ocr-2505")
assert isinstance(
mistral_config, VertexAIOCRConfig
), "Mistral model should route to VertexAIOCRConfig"
# Test other DeepSeek variants
deepseek_variant = get_vertex_ai_ocr_config("vertex_ai/deepseek-ocr-maas")
assert isinstance(
deepseek_variant, VertexAIDeepSeekOCRConfig
), "DeepSeek variant should route to VertexAIDeepSeekOCRConfig"

View file

@ -1,33 +0,0 @@
import pytest
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
def test_should_encode_azure_document_intelligence_model_id():
config = AzureDocumentIntelligenceOCRConfig()
url = config.get_complete_url(
api_base="https://example.cognitiveservices.azure.com",
model="prebuilt-layout?x=1#frag",
optional_params={},
litellm_params={},
)
assert (
url
== "https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout%3Fx%3D1%23frag:analyze?api-version=2024-11-30"
)
def test_should_reject_dot_segment_azure_document_intelligence_model_id():
config = AzureDocumentIntelligenceOCRConfig()
with pytest.raises(ValueError, match="model_id cannot be a dot path segment"):
config.get_complete_url(
api_base="https://example.cognitiveservices.azure.com",
model="azure_ai/doc-intelligence/..",
optional_params={},
litellm_params={},
)

View file

@ -1,178 +0,0 @@
"""
Unit tests for MistralOCRConfig transformation.
Tests the supported OCR parameters and their mapping behaviour.
No real API calls are made — all tests are fully mocked/local.
"""
import pytest
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
@pytest.fixture
def config() -> MistralOCRConfig:
return MistralOCRConfig()
MODEL = "mistral-ocr-latest"
class TestGetSupportedOcrParams:
def test_extract_header_in_supported_params(self, config: MistralOCRConfig) -> None:
"""extract_header must be in the Mistral OCR supported params list."""
supported = config.get_supported_ocr_params(model=MODEL)
assert "extract_header" in supported
def test_extract_footer_in_supported_params(self, config: MistralOCRConfig) -> None:
"""extract_footer must be in the Mistral OCR supported params list."""
supported = config.get_supported_ocr_params(model=MODEL)
assert "extract_footer" in supported
def test_existing_params_still_present(self, config: MistralOCRConfig) -> None:
"""Ensure the previously supported params were not accidentally removed."""
supported = config.get_supported_ocr_params(model=MODEL)
for param in [
"pages",
"include_image_base64",
"image_limit",
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
]:
assert (
param in supported
), f"Previously supported param '{param}' is missing"
class TestMapOcrParams:
def test_extract_header_passed_through(self, config: MistralOCRConfig) -> None:
"""extract_header=True must survive the map_ocr_params filter."""
result = config.map_ocr_params(
non_default_params={"extract_header": True},
optional_params={},
model=MODEL,
)
assert result == {"extract_header": True}
def test_extract_footer_passed_through(self, config: MistralOCRConfig) -> None:
"""extract_footer=True must survive the map_ocr_params filter."""
result = config.map_ocr_params(
non_default_params={"extract_footer": True},
optional_params={},
model=MODEL,
)
assert result == {"extract_footer": True}
def test_extract_header_and_footer_together(self, config: MistralOCRConfig) -> None:
"""Both params can be passed together and are both forwarded."""
result = config.map_ocr_params(
non_default_params={"extract_header": True, "extract_footer": False},
optional_params={},
model=MODEL,
)
assert result == {"extract_header": True, "extract_footer": False}
def test_unknown_param_is_dropped(self, config: MistralOCRConfig) -> None:
"""Parameters not in the supported list must be silently dropped."""
result = config.map_ocr_params(
non_default_params={"extract_header": True, "unsupported_param": "value"},
optional_params={},
model=MODEL,
)
assert "extract_header" in result
assert "unsupported_param" not in result
class TestNewSupportedParams:
"""Verify the newly added params are in the supported list."""
@pytest.mark.parametrize(
"param_name",
[
"table_format",
"confidence_scores_granularity",
"document_annotation_prompt",
"id",
],
)
def test_new_param_in_supported_list(
self, config: MistralOCRConfig, param_name: str
) -> None:
supported = config.get_supported_ocr_params(model=MODEL)
assert param_name in supported
class TestNewParamsMapOcr:
"""Verify the newly added params survive map_ocr_params."""
@pytest.mark.parametrize(
"param_name,param_value",
[
("table_format", "html"),
("table_format", "markdown"),
("confidence_scores_granularity", "word"),
("confidence_scores_granularity", "page"),
("document_annotation_prompt", "Extract all invoice line items"),
("id", "req-123"),
],
)
def test_new_param_passed_through(
self, config: MistralOCRConfig, param_name: str, param_value: str
) -> None:
result = config.map_ocr_params(
non_default_params={param_name: param_value},
optional_params={},
model=MODEL,
)
assert result == {param_name: param_value}
class TestTransformOcrRequest:
"""Verify params end up in the final request body via transform_ocr_request."""
SAMPLE_DOCUMENT = {
"type": "document_url",
"document_url": "https://example.com/doc.pdf",
}
@pytest.mark.parametrize(
"param_name,param_value",
[
("table_format", "html"),
("confidence_scores_granularity", "word"),
("document_annotation_prompt", "Extract all invoice line items"),
("id", "req-123"),
("extract_header", True),
("pages", [0, 1]),
],
)
def test_param_included_in_request_body(
self, config: MistralOCRConfig, param_name: str, param_value
) -> None:
result = config.transform_ocr_request(
model=MODEL,
document=self.SAMPLE_DOCUMENT,
optional_params={param_name: param_value},
headers={},
)
assert result.data[param_name] == param_value
assert result.data["model"] == MODEL
assert result.data["document"] == self.SAMPLE_DOCUMENT
assert result.files is None
def test_multiple_new_params_together(self, config: MistralOCRConfig) -> None:
"""Multiple new params can be passed together in a single request."""
optional_params = {
"table_format": "html",
"confidence_scores_granularity": "page",
"extract_header": True,
}
result = config.transform_ocr_request(
model=MODEL,
document=self.SAMPLE_DOCUMENT,
optional_params=optional_params,
headers={},
)
for key, value in optional_params.items():
assert result.data[key] == value

View file

@ -1,59 +0,0 @@
import json
import litellm
import pytest
@pytest.fixture()
def disable_aiohttp_transport():
original_disable_aiohttp = litellm.disable_aiohttp_transport
litellm.disable_aiohttp_transport = True
litellm.in_memory_llm_clients_cache.flush_cache()
try:
yield
finally:
litellm.disable_aiohttp_transport = original_disable_aiohttp
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.mark.asyncio
async def test_parse_legacy_wraps_enhance_under_options(
disable_aiohttp_transport, respx_mock
):
upload_route = respx_mock.post("https://platform.reducto.ai/upload").respond(
json={"file_id": "reducto://legacy.pdf"}
)
parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(
json={
"usage": {"num_pages": 1, "credits": 1},
"result": {
"chunks": [
{
"content": "Legacy parse",
"blocks": [{"content": "Legacy parse", "bbox": {"page": 1}}],
}
]
},
}
)
response = await litellm.aocr(
model="reducto/parse-legacy",
document={
"type": "file",
"file": b"%PDF-1.4 legacy",
"mime_type": "application/pdf",
},
api_key="legacy-key",
api_base="https://platform.reducto.ai",
enhance={"agentic": [{"type": "table"}]},
)
assert upload_route.called
assert parse_route.called
request_body = json.loads(parse_route.calls[0].request.read())
assert request_body == {
"document_url": "reducto://legacy.pdf",
"options": {"enhance": {"agentic": [{"type": "table"}]}},
}
assert response.pages[0].markdown == "Legacy parse"

View file

@ -1,152 +0,0 @@
import json
import litellm
import pytest
def _reducto_parse_response() -> dict:
return {
"job_id": "job_123",
"usage": {"num_pages": 3, "credits": 3},
"result": {
"chunks": [
{
"content": "Page 1 block A",
"blocks": [
{
"content": "Page 1 block A",
"bbox": {"page": 1},
"kind": "text",
}
],
},
{
"content": "Page 2 block A",
"blocks": [
{
"content": "Page 2 block A",
"bbox": {"page": 2},
"kind": "table",
}
],
},
{
"content": "Page 1 block B",
"blocks": [
{
"content": "Page 1 block B",
"bbox": {"page": 1},
"kind": "text",
}
],
},
{
"content": "Page 3 block A",
"blocks": [
{
"content": "Page 3 block A",
"bbox": {"page": 3},
"kind": "figure",
}
],
},
]
},
}
@pytest.fixture()
def disable_aiohttp_transport():
original_disable_aiohttp = litellm.disable_aiohttp_transport
litellm.disable_aiohttp_transport = True
litellm.in_memory_llm_clients_cache.flush_cache()
try:
yield
finally:
litellm.disable_aiohttp_transport = original_disable_aiohttp
litellm.in_memory_llm_clients_cache.flush_cache()
@pytest.mark.asyncio
async def test_parse_v3_file_upload_and_response_mapping(
disable_aiohttp_transport, respx_mock
):
upload_route = respx_mock.post("https://platform.reducto.ai/upload").respond(
json={"file_id": "reducto://uploaded.pdf"}
)
parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(
json=_reducto_parse_response()
)
response = await litellm.aocr(
model="reducto/parse-v3",
document={
"type": "file",
"file": b"%PDF-1.4 reducto",
"mime_type": "application/pdf",
},
api_key="test-key",
api_base="https://platform.reducto.ai",
formatting={"table_output_format": "html"},
retrieval={"chunk_mode": "section"},
settings={"ocr_system": "standard"},
)
assert upload_route.called
assert parse_route.called
assert len(upload_route.calls) == 1
assert len(parse_route.calls) == 1
upload_request = upload_route.calls[0].request
assert upload_request.headers["authorization"] == "Bearer test-key"
assert "application/json" not in upload_request.headers["content-type"]
upload_body = upload_request.read()
assert b'filename="document"' in upload_body
assert b"application/pdf" in upload_body
parse_request_body = json.loads(parse_route.calls[0].request.read())
assert parse_request_body["input"] == "reducto://uploaded.pdf"
assert parse_request_body["formatting"] == {"table_output_format": "html"}
assert parse_request_body["retrieval"] == {"chunk_mode": "section"}
assert parse_request_body["settings"] == {"ocr_system": "standard"}
assert response.usage_info is not None
assert response.usage_info.credits == 3
assert response.usage_info.pages_processed == 3
assert len(response.pages) == 3
assert response.pages[0].index == 0
assert response.pages[0].markdown == "Page 1 block A\n\nPage 1 block B"
assert getattr(response.pages[0], "blocks")[0]["bbox"]["page"] == 1
assert response.pages[1].markdown == "Page 2 block A"
assert response.pages[2].markdown == "Page 3 block A"
assert response._hidden_params["reducto_raw"]["usage"]["credits"] == 3
@pytest.mark.asyncio
async def test_parse_v3_reducto_id_passthrough_skips_upload(
disable_aiohttp_transport, respx_mock
):
upload_route = respx_mock.post("https://platform.reducto.ai/upload").respond(
json={"file_id": "reducto://should-not-upload.pdf"}
)
parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(
json=_reducto_parse_response()
)
response = await litellm.aocr(
model="reducto/parse-v3",
document={
"type": "document_url",
"document_url": "reducto://already-uploaded.pdf",
},
api_key="test-key",
api_base="https://platform.reducto.ai",
retrieval={"chunk_mode": "section"},
)
assert not upload_route.called
assert parse_route.called
parse_request_body = json.loads(parse_route.calls[0].request.read())
assert parse_request_body["input"] == "reducto://already-uploaded.pdf"
assert parse_request_body["retrieval"]["chunk_mode"] == "section"
assert response.pages[0].markdown.startswith("Page 1 block A")

View file

@ -1,213 +0,0 @@
import json
import os
from unittest.mock import AsyncMock, Mock
import httpx
import litellm
import pytest
from litellm.llms.reducto.common import (
extract_file_id_or_bytes,
upload_bytes_async,
upload_bytes_sync,
)
@pytest.fixture()
def disable_aiohttp_transport(monkeypatch):
original_disable_aiohttp = litellm.disable_aiohttp_transport
litellm.disable_aiohttp_transport = True
litellm.in_memory_llm_clients_cache.flush_cache()
monkeypatch.setenv("REDUCTO_API_KEY", "env-reducto-key")
try:
yield
finally:
litellm.disable_aiohttp_transport = original_disable_aiohttp
litellm.in_memory_llm_clients_cache.flush_cache()
os.environ.pop("REDUCTO_API_KEY", None)
@pytest.mark.asyncio
async def test_parse_v3_rejects_plain_http_urls(disable_aiohttp_transport):
with pytest.raises(litellm.BadRequestError, match="upload the file first"):
await litellm.aocr(
model="reducto/parse-v3",
document={
"type": "document_url",
"document_url": "https://example.com/document.pdf",
},
api_key="test-key",
api_base="https://platform.reducto.ai",
)
@pytest.mark.asyncio
async def test_parse_v3_image_data_uri_upload_uses_image_mime(
disable_aiohttp_transport, respx_mock
):
upload_route = respx_mock.post("https://custom.reducto.test/upload").respond(
json={"file_id": "reducto://uploaded-image.png"}
)
parse_route = respx_mock.post("https://custom.reducto.test/parse").respond(
json={
"usage": {"num_pages": 1, "credits": 1},
"result": {
"chunks": [
{
"content": "Image OCR",
"blocks": [{"content": "Image OCR", "bbox": {"page": 1}}],
}
]
},
}
)
response = await litellm.aocr(
model="reducto/parse-v3",
document={
"type": "file",
"file": b"\x89PNG\r\n\x1a\npng",
"mime_type": "image/png",
},
api_key="programmatic-key",
api_base="https://custom.reducto.test/",
)
assert upload_route.called
assert parse_route.called
upload_request = upload_route.calls[0].request
assert upload_request.headers["authorization"] == "Bearer programmatic-key"
assert b"image/png" in upload_request.read()
parse_request_body = json.loads(parse_route.calls[0].request.read())
assert parse_request_body["input"] == "reducto://uploaded-image.png"
assert response.pages[0].markdown == "Image OCR"
@pytest.mark.asyncio
async def test_parse_v3_uses_programmatic_api_key_over_env(
disable_aiohttp_transport, respx_mock
):
upload_route = respx_mock.post("https://platform.reducto.ai/upload").respond(
json={"file_id": "reducto://uploaded.pdf"}
)
parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(
json={
"usage": {"num_pages": 1, "credits": 1},
"result": {
"chunks": [
{
"content": "Programmatic auth",
"blocks": [
{"content": "Programmatic auth", "bbox": {"page": 1}}
],
}
]
},
}
)
await litellm.aocr(
model="reducto/parse-v3",
document={
"type": "file",
"file": b"%PDF-1.4 auth",
"mime_type": "application/pdf",
},
api_key="passed-key",
api_base="https://platform.reducto.ai",
)
assert upload_route.calls[0].request.headers["authorization"] == "Bearer passed-key"
assert parse_route.calls[0].request.headers["authorization"] == "Bearer passed-key"
def test_upload_bytes_sync_uses_shared_client(monkeypatch):
captured = {}
def fake_post(*, url, headers, files, timeout):
captured["url"] = url
captured["headers"] = headers
captured["files"] = files
captured["timeout"] = timeout
return httpx.Response(
200,
json={"file_id": "reducto://sync-upload"},
request=httpx.Request("POST", url),
)
sync_post = Mock(side_effect=fake_post)
monkeypatch.setattr(litellm.module_level_client, "post", sync_post)
class ForbiddenSyncClient:
def __init__(self, *args, **kwargs):
raise AssertionError("should not construct")
monkeypatch.setattr(httpx, "Client", ForbiddenSyncClient)
file_id = upload_bytes_sync(
raw_bytes=b"%PDF-1.4 sync",
mime="application/pdf",
api_key="sync-key",
api_base="https://sync.reducto.test/",
)
assert file_id == "reducto://sync-upload"
sync_post.assert_called_once()
assert captured["url"] == "https://sync.reducto.test/upload"
assert captured["headers"] == {"Authorization": "Bearer sync-key"}
assert captured["files"]["file"] == (
"document",
b"%PDF-1.4 sync",
"application/pdf",
)
@pytest.mark.asyncio
async def test_upload_bytes_async_uses_shared_aclient(monkeypatch):
captured = {}
async def fake_post(*, url, headers, files, timeout):
captured["url"] = url
captured["headers"] = headers
captured["files"] = files
captured["timeout"] = timeout
return httpx.Response(
200,
json={"file_id": "reducto://async-upload"},
request=httpx.Request("POST", url),
)
async_post = AsyncMock(side_effect=fake_post)
monkeypatch.setattr(litellm.module_level_aclient, "post", async_post)
class ForbiddenAsyncClient:
def __init__(self, *args, **kwargs):
raise AssertionError("should not construct")
monkeypatch.setattr(httpx, "AsyncClient", ForbiddenAsyncClient)
file_id = await upload_bytes_async(
raw_bytes=b"%PDF-1.4 async",
mime="application/pdf",
api_key="async-key",
api_base="https://async.reducto.test/",
)
assert file_id == "reducto://async-upload"
async_post.assert_awaited_once()
assert captured["url"] == "https://async.reducto.test/upload"
assert captured["headers"] == {"Authorization": "Bearer async-key"}
assert captured["files"]["file"] == (
"document",
b"%PDF-1.4 async",
"application/pdf",
)
def test_extract_file_id_or_bytes_raises_on_malformed_data_uri():
with pytest.raises(litellm.BadRequestError, match="Invalid Reducto data URI"):
extract_file_id_or_bytes("data:application/pdf", model="reducto/parse-v3")
with pytest.raises(litellm.BadRequestError, match="Invalid Reducto base64 payload"):
extract_file_id_or_bytes("data:;base64,!!!not-base64", model="reducto/parse-v3")

View file

@ -1,177 +0,0 @@
"""
VERIA-51: polling URLs returned by upstream APIs (Azure DALL-E,
Azure Document Intelligence, Black Forest Labs) used to be followed
without origin validation. The handlers attached the operator's API
key to the polling request, so an attacker who could influence the
upstream response (or a compromised upstream) could redirect the proxy
to send credentials anywhere.
These tests assert each handler now rejects polling URLs that don't
share an origin with the original request URL.
"""
from unittest.mock import MagicMock, patch
import httpx
import pytest
# Azure DALL-E sync + async paths route through ``assert_same_origin``
# the same way as the cases below. The helper itself is unit-tested in
# ``tests/test_litellm/litellm_core_utils/test_url_utils.py``; the
# tests here exercise the wiring at sites with simpler signatures.
# ── Azure Document Intelligence polling ───────────────────────────────────────
def test_azure_di_sync_rejects_cross_origin_polling():
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
config = AzureDocumentIntelligenceOCRConfig()
raw_response = MagicMock()
raw_response.status_code = 202
raw_response.headers = {
"Operation-Location": "https://attacker.example.com/results/xyz",
}
raw_response.request = MagicMock()
raw_response.request.url = (
"https://eastus.cognitiveservices.azure.com/documentintelligence/.../analyze"
)
raw_response.request.headers = {"Ocp-Apim-Subscription-Key": "leak-me"}
with pytest.raises(ValueError, match="rejected polling URL"):
config.transform_ocr_response(
model="azure-doc-intel",
raw_response=raw_response,
logging_obj=MagicMock(),
request_data={},
optional_params={},
litellm_params={},
encoding=None,
response={},
)
# ── Black Forest Labs polling ─────────────────────────────────────────────────
def test_bfl_image_generation_sync_rejects_cross_origin_polling():
from litellm.llms.black_forest_labs.image_generation.handler import (
BlackForestLabsImageGeneration,
)
handler = BlackForestLabsImageGeneration()
initial_response = MagicMock()
initial_response.status_code = 200
initial_response.json = MagicMock(
return_value={"polling_url": "https://attacker.example.com/get_result"}
)
initial_response.request = MagicMock()
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro"
sync_client = MagicMock()
sync_client.get = MagicMock()
with pytest.raises(Exception, match="Rejected polling URL"):
handler._poll_for_result_sync(
initial_response=initial_response,
headers={"x-key": "secret"},
sync_client=sync_client,
)
sync_client.get.assert_not_called()
@pytest.mark.asyncio
async def test_bfl_image_generation_async_rejects_cross_origin_polling():
from litellm.llms.black_forest_labs.image_generation.handler import (
BlackForestLabsImageGeneration,
)
handler = BlackForestLabsImageGeneration()
initial_response = MagicMock()
initial_response.status_code = 200
initial_response.json = MagicMock(
return_value={"polling_url": "https://attacker.example.com/get_result"}
)
initial_response.request = MagicMock()
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro"
async_client = MagicMock()
async_client.get = MagicMock()
with pytest.raises(Exception, match="Rejected polling URL"):
await handler._poll_for_result_async(
initial_response=initial_response,
headers={"x-key": "secret"},
async_client=async_client,
)
async_client.get.assert_not_called()
def test_bfl_image_edit_sync_rejects_cross_origin_polling():
from litellm.llms.black_forest_labs.image_edit.handler import (
BlackForestLabsImageEdit,
)
handler = BlackForestLabsImageEdit()
initial_response = MagicMock()
initial_response.status_code = 200
initial_response.json = MagicMock(
return_value={"polling_url": "https://attacker.example.com/get_result"}
)
initial_response.request = MagicMock()
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro/edit"
sync_client = MagicMock()
sync_client.get = MagicMock()
with pytest.raises(Exception, match="Rejected polling URL"):
handler._poll_for_result_sync(
initial_response=initial_response,
headers={"x-key": "secret"},
sync_client=sync_client,
)
sync_client.get.assert_not_called()
def test_bfl_image_generation_same_origin_polling_passes():
"""Sanity check: when the polling URL shares origin with the original
request, the origin check passes and polling proceeds."""
from litellm.llms.black_forest_labs.image_generation.handler import (
BlackForestLabsImageGeneration,
)
handler = BlackForestLabsImageGeneration()
initial_response = MagicMock()
initial_response.status_code = 200
initial_response.json = MagicMock(
return_value={"polling_url": "https://api.bfl.ai/v1/get_result?id=abc"}
)
initial_response.request = MagicMock()
initial_response.request.url = "https://api.bfl.ai/v1/flux-pro"
sync_client = MagicMock()
poll_response = MagicMock()
poll_response.status_code = 200
poll_response.json = MagicMock(return_value={"status": "Ready"})
sync_client.get = MagicMock(return_value=poll_response)
result = handler._poll_for_result_sync(
initial_response=initial_response,
headers={"x-key": "secret"},
sync_client=sync_client,
)
sync_client.get.assert_called_once()
assert result is poll_response

View file

@ -3,7 +3,6 @@
import importlib
import builtins
import types
from typing import Any
import httpx
import pytest
@ -151,33 +150,6 @@ class RecordingLogging:
}
def build_prepared_request(
*,
logging_obj: RecordingLogging | None = None,
model: str = "mistral-ocr-latest",
document: dict[str, object] = DOCUMENT,
api_key: str | None = "sk-test",
api_base: str | None = None,
custom_llm_provider: str = "mistral",
extra_headers: dict[str, object] | None = None,
optional_params: dict[str, object] | None = None,
litellm_params: dict[str, object] | None = None,
timeout: float | httpx.Timeout | None = 12.5,
) -> Any:
return ocr_main._PreparedOCRRequest(
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 or {},
litellm_params=litellm_params or {},
effective_timeout=timeout,
litellm_logging_obj=logging_obj or RecordingLogging(),
)
@pytest.fixture(autouse=True)
def _reset_rust_bridge():
"""Keep the global bridge state isolated between tests."""
@ -320,14 +292,15 @@ def test_run_rust_ocr_forwards_args_and_wraps_response():
response = ocr_main._run_rust_ocr(
rust_ocr=bridge,
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="mistral-ocr-latest",
document=DOCUMENT,
api_key="sk-test",
api_base="https://proxy.internal",
custom_llm_provider="mistral",
extra_headers={"x-trace-id": "trace-1"},
optional_params={"include_image_base64": True},
timeout=12.5,
litellm_logging_obj=logging_obj,
)
assert isinstance(response, OCRResponse)
@ -344,167 +317,21 @@ def test_run_rust_ocr_forwards_args_and_wraps_response():
"timeout_seconds": 12.5,
}
def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing():
"""No explicit api_key: the resolver (get_secret_str in production) supplies it,
so secret-manager backends (AWS/Azure/GCP/Vault) work like the Python path."""
bridge = RecordingBridge()
ocr_main._run_rust_ocr(
rust_ocr=bridge,
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
),
)
assert bridge.calls[0]["api_key"] == "sk-from-vault"
def test_run_rust_ocr_resolves_reducto_key():
bridge = RecordingBridge()
resolver_calls = []
def _resolver(name):
resolver_calls.append(name)
return "sk-provider-env"
ocr_main._run_rust_ocr(
rust_ocr=bridge,
prepared_request=build_prepared_request(
custom_llm_provider="reducto",
model="parse-v3",
api_key=None,
timeout=None,
),
resolve_api_key=_resolver,
)
assert resolver_calls == ["REDUCTO_API_KEY"]
assert bridge.calls[0]["api_key"] == "sk-provider-env"
def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata():
bridge = RecordingBridge()
ocr_main._run_rust_ocr(
rust_ocr=bridge,
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,
)
assert bridge.calls[0]["optional_params"] == {
"include_image_base64": True,
"vertex_project": "project-1",
"vertex_location": "us-central1",
}
def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_manager():
bridge = RecordingBridge()
def _resolver(name: str) -> str | None:
return {
"VERTEXAI_PROJECT": "project-from-secret",
"VERTEXAI_LOCATION": "us-east5",
}.get(name)
ocr_main._run_rust_ocr(
rust_ocr=bridge,
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():
bridge = RecordingBridge()
ocr_main._run_rust_ocr(
rust_ocr=bridge,
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
),
)
assert bridge.calls[0]["api_base"] == "https://azure.example.com"
def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
bridge = RecordingBridge()
ocr_main._run_rust_ocr(
rust_ocr=bridge,
prepared_request=build_prepared_request(
custom_llm_provider="azure_ai/doc-intelligence",
model="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
),
)
assert bridge.calls[0]["api_base"] == "https://document-intelligence.example.com"
def test_run_rust_ocr_prefers_explicit_key_over_resolver():
bridge = RecordingBridge()
resolver_calls = []
def _resolver(name):
resolver_calls.append(name)
return "sk-from-vault"
ocr_main._run_rust_ocr(
rust_ocr=bridge,
prepared_request=build_prepared_request(api_key="sk-explicit", timeout=None),
resolve_api_key=_resolver,
)
assert bridge.calls[0]["api_key"] == "sk-explicit"
assert resolver_calls == [] # resolver never consulted when a key is supplied
def test_run_rust_ocr_runs_pre_call_logging():
"""The Rust shortcut must run pre_call so callbacks and spend tracking fire."""
logging_obj = RecordingLogging()
ocr_main._run_rust_ocr(
rust_ocr=RecordingBridge(),
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="mistral-ocr-latest",
document=DOCUMENT,
api_key="sk-test",
api_base="https://api.mistral.ai/v1",
custom_llm_provider="mistral",
extra_headers={"x-trace-id": "trace-1"},
optional_params={"include_image_base64": True},
timeout=12.5,
litellm_logging_obj=logging_obj,
)
assert logging_obj.pre_call_kwargs is not None
@ -540,6 +367,19 @@ def test_ocr_routes_to_rust_by_default(fake_bridge):
assert call["optional_params"].get("include_image_base64") is True
def test_ocr_filters_internal_litellm_params_before_rust(fake_bridge):
litellm.ocr(
model=MODEL,
document=DOCUMENT,
api_key="sk-test",
include_image_base64=True,
original_generic_function=lambda: None,
litellm_metadata={"trace": "internal"},
)
assert fake_bridge.calls[0]["optional_params"] == {"include_image_base64": True}
def test_ocr_routes_azure_ai_to_rust_by_default(fake_bridge):
response = litellm.ocr(
model="azure_ai/pixtral-12b-2409",