From b2e7c48ebd80d608379c65f8ed61f41e6fe69f91 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Thu, 25 Jun 2026 14:30:06 -0700 Subject: [PATCH] refactor(ocr): remove python provider transforms --- litellm/_lazy_imports_registry.py | 4 - litellm/llms/azure_ai/ocr/__init__.py | 12 +- litellm/llms/azure_ai/ocr/common_utils.py | 52 -- .../ocr/document_intelligence/__init__.py | 4 +- .../document_intelligence/transformation.py | 815 ------------------ litellm/llms/azure_ai/ocr/transformation.py | 287 ------ litellm/llms/base_llm/ocr/__init__.py | 6 - litellm/llms/base_llm/ocr/transformation.py | 207 +---- litellm/llms/custom_httpx/llm_http_handler.py | 297 ------- litellm/llms/mistral/ocr/transformation.py | 239 ----- litellm/llms/reducto/common.py | 159 ---- litellm/llms/reducto/ocr/transformation.py | 241 ------ litellm/llms/vertex_ai/ocr/__init__.py | 4 +- litellm/llms/vertex_ai/ocr/common_utils.py | 41 - .../vertex_ai/ocr/deepseek_transformation.py | 399 --------- litellm/llms/vertex_ai/ocr/transformation.py | 314 ------- litellm/ocr/main.py | 315 +++---- litellm/utils.py | 44 - .../test_ocr_azure_document_intelligence.py | 126 --- tests/ocr_tests/test_ocr_vertex_ai.py | 31 - ...ocument_intelligence_ocr_transformation.py | 33 - .../ocr/test_mistral_ocr_transformation.py | 178 ---- .../llms/reducto/test_parse_legacy.py | 59 -- .../llms/reducto/test_parse_v3.py | 152 ---- .../test_litellm/llms/reducto/test_upload.py | 213 ----- .../llms/test_polling_url_origin_match.py | 177 ---- tests/test_litellm/ocr/test_rust_bridge.py | 222 +---- 27 files changed, 174 insertions(+), 4457 deletions(-) delete mode 100644 litellm/llms/azure_ai/ocr/common_utils.py delete mode 100644 litellm/llms/azure_ai/ocr/document_intelligence/transformation.py delete mode 100644 litellm/llms/azure_ai/ocr/transformation.py delete mode 100644 litellm/llms/mistral/ocr/transformation.py delete mode 100644 litellm/llms/reducto/common.py delete mode 100644 litellm/llms/reducto/ocr/transformation.py delete mode 100644 litellm/llms/vertex_ai/ocr/common_utils.py delete mode 100644 litellm/llms/vertex_ai/ocr/deepseek_transformation.py delete mode 100644 litellm/llms/vertex_ai/ocr/transformation.py delete mode 100644 tests/test_litellm/llms/azure_ai/test_azure_document_intelligence_ocr_transformation.py delete mode 100644 tests/test_litellm/llms/mistral/ocr/test_mistral_ocr_transformation.py delete mode 100644 tests/test_litellm/llms/reducto/test_parse_legacy.py delete mode 100644 tests/test_litellm/llms/reducto/test_parse_v3.py delete mode 100644 tests/test_litellm/llms/reducto/test_upload.py delete mode 100644 tests/test_litellm/llms/test_polling_url_origin_match.py diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 4f131354d2e..7eebadd8c33 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -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"), diff --git a/litellm/llms/azure_ai/ocr/__init__.py b/litellm/llms/azure_ai/ocr/__init__.py index ade1165b848..c84d5f0bf59 100644 --- a/litellm/llms/azure_ai/ocr/__init__.py +++ b/litellm/llms/azure_ai/ocr/__init__.py @@ -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] = [] diff --git a/litellm/llms/azure_ai/ocr/common_utils.py b/litellm/llms/azure_ai/ocr/common_utils.py deleted file mode 100644 index d736b891532..00000000000 --- a/litellm/llms/azure_ai/ocr/common_utils.py +++ /dev/null @@ -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/ - - Mistral OCR (via Azure AI): azure_ai/ - - 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") - - - >>> get_azure_ai_ocr_config("azure_ai/pixtral-12b-2409") - - """ - 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() diff --git a/litellm/llms/azure_ai/ocr/document_intelligence/__init__.py b/litellm/llms/azure_ai/ocr/document_intelligence/__init__.py index 32d700fd195..149f0fab624 100644 --- a/litellm/llms/azure_ai/ocr/document_intelligence/__init__.py +++ b/litellm/llms/azure_ai/ocr/document_intelligence/__init__.py @@ -1,5 +1,3 @@ """Azure Document Intelligence OCR module.""" -from .transformation import AzureDocumentIntelligenceOCRConfig - -__all__ = ["AzureDocumentIntelligenceOCRConfig"] +__all__: list[str] = [] diff --git a/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py b/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py deleted file mode 100644 index cc65ad706ab..00000000000 --- a/litellm/llms/azure_ai/ocr/document_intelligence/transformation.py +++ /dev/null @@ -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/ - - 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:[][;base64], - 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 diff --git a/litellm/llms/azure_ai/ocr/transformation.py b/litellm/llms/azure_ai/ocr/transformation.py deleted file mode 100644 index ee35fc28994..00000000000 --- a/litellm/llms/azure_ai/ocr/transformation.py +++ /dev/null @@ -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:///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, - ) diff --git a/litellm/llms/base_llm/ocr/__init__.py b/litellm/llms/base_llm/ocr/__init__.py index 2aea2d67807..58d619b0566 100644 --- a/litellm/llms/base_llm/ocr/__init__.py +++ b/litellm/llms/base_llm/ocr/__init__.py @@ -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", ] diff --git a/litellm/llms/base_llm/ocr/transformation.py b/litellm/llms/base_llm/ocr/transformation.py index a2946c62506..f0b373b423a 100644 --- a/litellm/llms/base_llm/ocr/transformation.py +++ b/litellm/llms/base_llm/ocr/transformation.py @@ -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, - ) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index d33ec295e94..4e613b33c9e 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -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, diff --git a/litellm/llms/mistral/ocr/transformation.py b/litellm/llms/mistral/ocr/transformation.py deleted file mode 100644 index 3c0460cd51e..00000000000 --- a/litellm/llms/mistral/ocr/transformation.py +++ /dev/null @@ -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": "" - }, - "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 diff --git a/litellm/llms/reducto/common.py b/litellm/llms/reducto/common.py deleted file mode 100644 index 4e7d96dbe87..00000000000 --- a/litellm/llms/reducto/common.py +++ /dev/null @@ -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 diff --git a/litellm/llms/reducto/ocr/transformation.py b/litellm/llms/reducto/ocr/transformation.py deleted file mode 100644 index cc338ecc484..00000000000 --- a/litellm/llms/reducto/ocr/transformation.py +++ /dev/null @@ -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, - ) diff --git a/litellm/llms/vertex_ai/ocr/__init__.py b/litellm/llms/vertex_ai/ocr/__init__.py index 915fbd49030..c79a6befb64 100644 --- a/litellm/llms/vertex_ai/ocr/__init__.py +++ b/litellm/llms/vertex_ai/ocr/__init__.py @@ -1,5 +1,3 @@ """Vertex AI OCR module.""" -from .transformation import VertexAIOCRConfig - -__all__ = ["VertexAIOCRConfig"] +__all__: list[str] = [] diff --git a/litellm/llms/vertex_ai/ocr/common_utils.py b/litellm/llms/vertex_ai/ocr/common_utils.py deleted file mode 100644 index 3e5fbe23447..00000000000 --- a/litellm/llms/vertex_ai/ocr/common_utils.py +++ /dev/null @@ -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/ - - Args: - model: The model name (e.g., "vertex_ai/ocr/") - - Returns: - OCR configuration instance for the specified model - - Examples: - >>> get_vertex_ai_ocr_config("vertex_ai/deepseek-ai/deepseek-ocr-maas") - - - >>> get_vertex_ai_ocr_config("vertex_ai/ocr/mistral-ocr-maas") - - """ - 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() diff --git a/litellm/llms/vertex_ai/ocr/deepseek_transformation.py b/litellm/llms/vertex_ai/ocr/deepseek_transformation.py deleted file mode 100644 index a98311d04eb..00000000000 --- a/litellm/llms/vertex_ai/ocr/deepseek_transformation.py +++ /dev/null @@ -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": "" - } - }], - "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, - ) diff --git a/litellm/llms/vertex_ai/ocr/transformation.py b/litellm/llms/vertex_ai/ocr/transformation.py deleted file mode 100644 index a725762b3c5..00000000000 --- a/litellm/llms/vertex_ai/ocr/transformation.py +++ /dev/null @@ -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, - ) diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 52244b40ab8..734ad3a7544 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -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( diff --git a/litellm/utils.py b/litellm/utils.py index e0aa8575473..1deb489a507 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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", diff --git a/tests/ocr_tests/test_ocr_azure_document_intelligence.py b/tests/ocr_tests/test_ocr_azure_document_intelligence.py index 7269890b7b6..097a4a34633 100644 --- a/tests/ocr_tests/test_ocr_azure_document_intelligence.py +++ b/tests/ocr_tests/test_ocr_azure_document_intelligence.py @@ -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"} - diff --git a/tests/ocr_tests/test_ocr_vertex_ai.py b/tests/ocr_tests/test_ocr_vertex_ai.py index 1ba5b9d0883..b911435eea7 100644 --- a/tests/ocr_tests/test_ocr_vertex_ai.py +++ b/tests/ocr_tests/test_ocr_vertex_ai.py @@ -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" diff --git a/tests/test_litellm/llms/azure_ai/test_azure_document_intelligence_ocr_transformation.py b/tests/test_litellm/llms/azure_ai/test_azure_document_intelligence_ocr_transformation.py deleted file mode 100644 index e638be68ec0..00000000000 --- a/tests/test_litellm/llms/azure_ai/test_azure_document_intelligence_ocr_transformation.py +++ /dev/null @@ -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={}, - ) diff --git a/tests/test_litellm/llms/mistral/ocr/test_mistral_ocr_transformation.py b/tests/test_litellm/llms/mistral/ocr/test_mistral_ocr_transformation.py deleted file mode 100644 index 97461561a05..00000000000 --- a/tests/test_litellm/llms/mistral/ocr/test_mistral_ocr_transformation.py +++ /dev/null @@ -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 diff --git a/tests/test_litellm/llms/reducto/test_parse_legacy.py b/tests/test_litellm/llms/reducto/test_parse_legacy.py deleted file mode 100644 index db19460baa3..00000000000 --- a/tests/test_litellm/llms/reducto/test_parse_legacy.py +++ /dev/null @@ -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" diff --git a/tests/test_litellm/llms/reducto/test_parse_v3.py b/tests/test_litellm/llms/reducto/test_parse_v3.py deleted file mode 100644 index 140b9737dc0..00000000000 --- a/tests/test_litellm/llms/reducto/test_parse_v3.py +++ /dev/null @@ -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") diff --git a/tests/test_litellm/llms/reducto/test_upload.py b/tests/test_litellm/llms/reducto/test_upload.py deleted file mode 100644 index 4fae90436bb..00000000000 --- a/tests/test_litellm/llms/reducto/test_upload.py +++ /dev/null @@ -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") diff --git a/tests/test_litellm/llms/test_polling_url_origin_match.py b/tests/test_litellm/llms/test_polling_url_origin_match.py deleted file mode 100644 index f1f910bc731..00000000000 --- a/tests/test_litellm/llms/test_polling_url_origin_match.py +++ /dev/null @@ -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 diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index 7136cac4362..5647b72667e 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -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",