mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
refactor(ocr): remove python provider transforms
This commit is contained in:
parent
1528c9b9d1
commit
b2e7c48ebd
27 changed files with 174 additions and 4457 deletions
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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] = []
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -1,5 +1,3 @@
|
|||
"""Azure Document Intelligence OCR module."""
|
||||
|
||||
from .transformation import AzureDocumentIntelligenceOCRConfig
|
||||
|
||||
__all__ = ["AzureDocumentIntelligenceOCRConfig"]
|
||||
__all__: list[str] = []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -1,5 +1,3 @@
|
|||
"""Vertex AI OCR module."""
|
||||
|
||||
from .transformation import VertexAIOCRConfig
|
||||
|
||||
__all__ = ["VertexAIOCRConfig"]
|
||||
__all__: list[str] = []
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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={},
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
@ -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")
|
||||
|
|
@ -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")
|
||||
|
|
@ -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
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue