mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
Merge pull request #39862 from BerriAI/litellm_lit_6992_cohere_parse
feat(ocr): add Cohere Parse support for cohere and azure_ai
This commit is contained in:
commit
bf51dea36b
18 changed files with 987 additions and 7 deletions
|
|
@ -6,14 +6,13 @@ import base64
|
|||
from collections.abc import Awaitable, Callable
|
||||
from typing import TYPE_CHECKING, Final, Literal
|
||||
|
||||
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS
|
||||
from litellm.llms.base_llm.ocr.transformation import BaseOCRConfig, DocumentType
|
||||
from litellm.types.utils import LIST_BATCHES_SUPPORTED_PROVIDERS, LlmProviders
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
# Minimal PDF for health checks - base64 encoded 1-page PDF with just "test"
|
||||
TEST_PDF_URL = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y="
|
||||
|
||||
# Minimal image for health checks - base64 encoded 512x512 blue circle on a white background PNG
|
||||
TEST_IMAGE_BASE64 = "iVBORw0KGgoAAAANSUhEUgAAAgAAAAIACAIAAAB7GkOtAAAJk0lEQVR42u3VQREAIRADwVWCOmTjBVzwSLorCri6nbkAVBpPACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAAAgAAAIAZdY+HgEBgIRr/meeGgGA8EMvDAgAuPh6gACAi68HCAA4+mKAAICjLwYIALj7SoAAgLuvBAgA7r4pAQKAu29KgADg7psSIAA4/SYDCADuvikBAoDTbzKAAOD0mwwgADj9JgMIAE6/yQACgNNvMoAA4PSbDCAAOP0mAwgATr/JAAKA668BIAA4/TKAAOD0mwwgALj+pgEIAE6/yQACgOtvGoAA4PSbDCAAuP6mAQgATr/JAAKA628agADg9JsMIAC4/qYBCACuv2kAAoDrbxqAAOD0mwwgALj+pgEIAK6/aQACgOtvGoAA4PqbBiAArr+ZBiAATr+ZDCAArr+ZBiAArr+ZBiAArr+ZBiAArr+ZBggArr+ZBggArr+ZBggArr+ZBggArr+ZBggArr+ZBggArr+ZBggAAmAmAAKA62+mAQKA62+mAQKA62/m1xYAXH/TAAQA1980AAHA9TcNQAAEwEwAEADX30wDEADX30wDEADX30wDEADX30wDEAABMBMABMD1N9MABMD1N9MABEAAzAQAAXD9zTQAAXD9zTRAABAAMwEQAFx/Mw0QAFx/Mw0QAATATAAEANffTAMEANffTAMEAAEwEwABwPU30wABQADMBEAAXH8z0wABcP3NTAMEQADMTAAEwPU3Mw0QAAEwMwEQANffTAMQAAEwEwAEwPU30wAEQADMBAABcP3NNAABEAAzAUAAXH8zDUAABMBMABAA199MAwQAATATAAHA9TfTAAFAAMwEQABw/c00QAAQADMBEAAEwEwABMD1NzMNEAABMDMBEADX38w0QAAEwMwEQAAEwMwEQABcfzPTAAEQADMTAAEQADMTAAFw/c1MAwRAAMxMAARAAMxMAATA9TczDRAAATATAARAAMwEAAFw/c00AAEQADMBQAAEwEwAEADX30wDBAABMBMAAUAAzARAABAAMwEQAFx/Mw0QAAEwMwEQAAEwMwEQAAEwMwEQANffzDRAAATAzARAAATAzARAAATAzARAAATAzARAAFx/M9MAARAAMxMAARAAMxMAARAAMxMAARAAMxMAAXD9zUwDBEAAzEwABEAAzAQAARAAMwFAAATATAAQAAEwEwABQADMBEAAEAAzARAAXH8zDRAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAAATAzARAA19/MNEAANMDM9UcABMBMABAAATATAAHwBAJgJgACgACYCYAAIABmAiAA+IvMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAABMDMBEAANMDMXH8BEAAzEwABEAAzEwABEAAzEwABEAAzEwABEAAzEwABEAAzAUAABMBMABAADTBz/REAATATAAFAAMwEQAAQADMBEAAEwEwABAANMHP9BUAAzEwABEAAzEwABEAAzEwABEAAzEwABEADzMz1FwABMDMBEAABMDMBEAABMDMBEAANMDPXXwAEwMwEQAAEwMwEQAAEwMwEQAA0wMxcfwEQADMTAAEQADMBQAA0wMz1RwAEwEwAEAABMBMABEADzFx/AUAAzARAABAAMwEQADTAzPUXAATATAAEAAEwEwABQAPMXH8BEAAzEwABEAAzEwAB0AAzc/0FQADMTAAEQAPMzPUXAAEwMwEQAAEwMwEQAA0wM9dfAATAzARAADTAzFx/ARAAMwFAADTAzPVHAATATAAQAA0wc/0RAAEwEwAEQAPMXH8EQADMBAAB0AAz118AEAAzARAANMDM9RcABMBMAAQADTBz/QUAATATAAFAA8xcfwFAA8xcfwFAAMwEQADQADPXXwAEwMwEQAA0wMxcfwHQADNz/QVAAMxMAARAA8xcfwRAA8xcfwRAAMwEAAHQADPXHwHQADPXHwEQADMBQAA0wMz1RwA0wMz1RwAEwEwAEAANMHP9BQANMHP9BQANMHP9BQANMHP9BQABMBMAAUADzFx/AUADzFx/AUADzFx/AUADzFx/AUADzFx/AUADzPVHABAAEwAEAA0w1x8BQAPM9UcA0ABz/REANMBcfwQADTDXHwHQADPXHwHQADPXHwHQADPXHwHQADPXHwHQADPXHwGQATOnHwHQADPXHwHQADPXXwDQADPXXwDQADPXXwDQADPXXwCQAXP6EQA0wFx/BAANMNcfAUADzPVHAJABc/oRADTAXH8EABkwpx8BQAPM9UcAkAFz+hEANMBcfwQAGTCnHwFAA8z1RwCQAXP6EQBkwJx+BAANMNcfAUAGzOlHAJABc/oRAGTA6QcBQAacfhAAZMDpRwBABpx+BABkwOlHAEAGnH4EAJTA3UcAQAacfgQAlMDdRwBACdx9BACUwN1HAEAJ3H0EAJTA3UcAQAwcfQQAmmLgsyIA0NIDHw4BgJYe+DQIAISHwVMjAJDQDI+AAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACACAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAgAAAIAAACAIAAACAAAAgAAAIAgAAAIAAACAAAAgCAAAAIAABNHpialFcmLajuAAAAAElFTkSuQmCC"
|
||||
|
|
@ -29,6 +28,14 @@ def get_image_file_for_health_check() -> bytes:
|
|||
return base64.b64decode(TEST_IMAGE_BASE64)
|
||||
|
||||
|
||||
def _ocr_health_check_document(model: str, custom_llm_provider: str) -> DocumentType:
|
||||
from litellm.utils import ProviderConfigManager
|
||||
|
||||
provider: Final = next((known for known in LlmProviders if known.value == custom_llm_provider), None)
|
||||
config: Final = ProviderConfigManager.get_provider_ocr_config(model=model, provider=provider) if provider else None
|
||||
return (config or BaseOCRConfig()).get_health_check_document()
|
||||
|
||||
|
||||
class HealthCheckHelpers:
|
||||
@staticmethod
|
||||
async def ahealth_check_wildcard_models(
|
||||
|
|
@ -247,9 +254,6 @@ class HealthCheckHelpers:
|
|||
),
|
||||
"ocr": lambda: litellm.aocr(
|
||||
**_filter_model_params(model_params=model_params),
|
||||
document={
|
||||
"type": "document_url",
|
||||
"document_url": TEST_PDF_URL,
|
||||
},
|
||||
document=_ocr_health_check_document(model=model, custom_llm_provider=custom_llm_provider),
|
||||
),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Azure AI OCR module."""
|
||||
|
||||
from .cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from .common_utils import get_azure_ai_ocr_config
|
||||
from .document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
|
|
@ -7,6 +8,7 @@ from .document_intelligence.transformation import (
|
|||
from .transformation import AzureAIOCRConfig
|
||||
|
||||
__all__ = [
|
||||
"AzureAICohereParseConfig",
|
||||
"AzureAIOCRConfig",
|
||||
"AzureDocumentIntelligenceOCRConfig",
|
||||
"get_azure_ai_ocr_config",
|
||||
|
|
|
|||
91
litellm/llms/azure_ai/ocr/cohere_parse_transformation.py
Normal file
91
litellm/llms/azure_ai/ocr/cohere_parse_transformation.py
Normal file
|
|
@ -0,0 +1,91 @@
|
|||
"""Cohere Parse served from Azure AI Foundry (`/providers/cohere/v2/parse`)."""
|
||||
|
||||
from collections.abc import Mapping
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
||||
async_convert_url_to_base64,
|
||||
convert_url_to_base64,
|
||||
)
|
||||
from litellm.llms.azure_ai.common_utils import get_azure_ai_auth_headers
|
||||
from litellm.llms.cohere.ocr.transformation import COHERE_PARSE_PATH, CohereParseConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
AZURE_AI_API_KEY_ENV_VAR: Final = "AZURE_AI_API_KEY"
|
||||
AZURE_AI_API_BASE_ENV_VAR: Final = "AZURE_AI_API_BASE"
|
||||
AZURE_AI_COHERE_PROVIDER_PATH: Final = "/providers/cohere"
|
||||
AZURE_AI_MODELS_PATH_SUFFIX: Final = "/models"
|
||||
|
||||
|
||||
class AzureAICohereParseConfig(CohereParseConfig):
|
||||
"""Same request and response shape as Cohere Parse, behind Azure AI auth and URL layout.
|
||||
|
||||
Foundry cannot fetch external URLs, so remote images are inlined as base64 data URIs.
|
||||
"""
|
||||
|
||||
def get_api_key_env_var(self) -> str | None:
|
||||
return AZURE_AI_API_KEY_ENV_VAR
|
||||
|
||||
def _llm_provider(self) -> str:
|
||||
return "azure_ai"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.validate_environment signature
|
||||
) -> dict[str, str]: # mutable-ok: BaseOCRConfig signature
|
||||
resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR)
|
||||
if resolved_base is None:
|
||||
raise ValueError(
|
||||
f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable "
|
||||
"or pass api_base parameter"
|
||||
)
|
||||
resolved_key: Final = api_key or get_secret_str(AZURE_AI_API_KEY_ENV_VAR)
|
||||
return { # mutable-ok: BaseOCRConfig signature
|
||||
**get_azure_ai_auth_headers(api_key=resolved_key, litellm_params=litellm_params),
|
||||
"Content-Type": "application/json",
|
||||
**headers,
|
||||
}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.get_complete_url signature
|
||||
) -> str:
|
||||
resolved_base: Final = api_base or get_secret_str(AZURE_AI_API_BASE_ENV_VAR)
|
||||
if resolved_base is None:
|
||||
raise ValueError(
|
||||
f"Missing Azure AI API Base - Set {AZURE_AI_API_BASE_ENV_VAR} environment variable "
|
||||
"or pass api_base parameter"
|
||||
)
|
||||
url: Final = httpx.URL(resolved_base)
|
||||
if not url.is_absolute_url:
|
||||
raise ValueError(
|
||||
"Azure AI API Base must be an absolute URL including scheme (e.g. "
|
||||
f"'https://<resource>.services.ai.azure.com'). Got api_base={resolved_base!r}."
|
||||
)
|
||||
path: Final = url.path.rstrip("/")
|
||||
if path.endswith(COHERE_PARSE_PATH):
|
||||
return str(url.copy_with(path=path))
|
||||
if path.endswith(f"{AZURE_AI_COHERE_PROVIDER_PATH}/v2"):
|
||||
return str(url.copy_with(path=f"{path}/parse"))
|
||||
return str(
|
||||
url.copy_with(
|
||||
path=f"{path.removesuffix(AZURE_AI_MODELS_PATH_SUFFIX)}{AZURE_AI_COHERE_PROVIDER_PATH}{COHERE_PARSE_PATH}"
|
||||
)
|
||||
)
|
||||
|
||||
def _resolve_image_url_sync(self, image_url: str) -> str:
|
||||
return convert_url_to_base64(image_url)
|
||||
|
||||
async def _resolve_image_url_async(self, image_url: str) -> str:
|
||||
return await async_convert_url_to_base64(image_url)
|
||||
|
|
@ -24,6 +24,11 @@ def is_azure_document_intelligence_model(model: str) -> bool:
|
|||
return "doc-intelligence" in lowered or "documentintelligence" in lowered
|
||||
|
||||
|
||||
def is_azure_cohere_parse_model(model: str) -> bool:
|
||||
lowered: Final = model.lower()
|
||||
return "cohere" in lowered and "parse" in lowered
|
||||
|
||||
|
||||
def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
|
||||
"""
|
||||
Determine which Azure AI OCR configuration to use based on the model name.
|
||||
|
|
@ -46,6 +51,7 @@ def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
|
|||
>>> get_azure_ai_ocr_config("azure_ai/pixtral-12b-2409")
|
||||
<AzureAIOCRConfig object>
|
||||
"""
|
||||
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
|
||||
AzureDocumentIntelligenceOCRConfig,
|
||||
)
|
||||
|
|
@ -56,6 +62,10 @@ def get_azure_ai_ocr_config(model: str) -> Optional["BaseOCRConfig"]:
|
|||
verbose_logger.debug("Routing %s to Azure Document Intelligence OCR config", model)
|
||||
return AzureDocumentIntelligenceOCRConfig()
|
||||
|
||||
if is_azure_cohere_parse_model(model):
|
||||
verbose_logger.debug("Routing %s to Azure AI Cohere Parse config", model)
|
||||
return AzureAICohereParseConfig()
|
||||
|
||||
# Default to Mistral-based OCR for other azure_ai models
|
||||
verbose_logger.debug("Routing %s to Azure AI (Mistral) OCR config", model)
|
||||
return AzureAIOCRConfig()
|
||||
|
|
|
|||
|
|
@ -33,6 +33,8 @@ OCR_REQUEST_FORMAT_HEADER: Final = "x-req-format"
|
|||
|
||||
PROVIDER_NATIVE_RESPONSE_KEY: Final = "provider_native_response"
|
||||
|
||||
HEALTH_CHECK_PDF_DATA_URI: Final = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y="
|
||||
|
||||
|
||||
def parse_ocr_request_format(value: object) -> OCRRequestFormat:
|
||||
if value == "litellm":
|
||||
|
|
@ -142,6 +144,16 @@ class BaseOCRConfig:
|
|||
"""
|
||||
return None
|
||||
|
||||
def supports_rust_bridge(self) -> bool:
|
||||
"""Whether the Rust OCR bridge may serve this config when it is enabled for the provider."""
|
||||
return True
|
||||
|
||||
def get_health_check_document(self) -> DocumentType:
|
||||
return { # mutable-ok: litellm.aocr rejects any document that is not a dict
|
||||
"type": "document_url",
|
||||
"document_url": HEALTH_CHECK_PDF_DATA_URI,
|
||||
}
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
|
|
|
|||
3
litellm/llms/cohere/ocr/__init__.py
Normal file
3
litellm/llms/cohere/ocr/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
|
||||
|
||||
__all__ = ("CohereParseConfig",)
|
||||
301
litellm/llms/cohere/ocr/transformation.py
Normal file
301
litellm/llms/cohere/ocr/transformation.py
Normal file
|
|
@ -0,0 +1,301 @@
|
|||
"""Cohere Parse (`POST /v2/parse`) exposed through LiteLLM's OCR interface."""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Final, Literal
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from litellm.exceptions import BadRequestError, UnsupportedParamsError
|
||||
from litellm.llms.base_llm.ocr.transformation import (
|
||||
OCR_REQUEST_FORMAT_PARAM,
|
||||
BaseOCRConfig,
|
||||
DocumentType,
|
||||
OCRPage,
|
||||
OCRPageImage,
|
||||
OCRRequestData,
|
||||
OCRRequestFormat,
|
||||
OCRResponse,
|
||||
OCRUsageInfo,
|
||||
parse_ocr_request_format,
|
||||
)
|
||||
from litellm.llms.cohere.common_utils import CohereError
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
|
||||
COHERE_API_KEY_ENV_VAR: Final = "COHERE_API_KEY"
|
||||
COHERE_PARSE_API_BASE: Final = "https://api.cohere.com"
|
||||
COHERE_PARSE_PATH: Final = "/v2/parse"
|
||||
COHERE_PARSE_OUTPUT_FORMAT_PARAM: Final = "output_format"
|
||||
COHERE_PARSE_OUTPUT_FORMATS: Final = ("markdown", "blocks")
|
||||
COHERE_PARSE_DEFAULT_OUTPUT_FORMAT: Final = "markdown"
|
||||
COHERE_PARSE_SUPPORTED_PARAMS: Final = (COHERE_PARSE_OUTPUT_FORMAT_PARAM, OCR_REQUEST_FORMAT_PARAM)
|
||||
COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI: Final = (
|
||||
"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC"
|
||||
)
|
||||
COHERE_PARSE_IMAGE_ONLY_MESSAGE: Final = (
|
||||
"Cohere Parse only accepts `image_url` documents (an image URL or a base64 image data URI); "
|
||||
"`document_url` and PDF inputs are not supported."
|
||||
)
|
||||
|
||||
_NATIVE_RESPONSE_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_BOUNDING_BOX_ADAPTER: Final = TypeAdapter(Mapping[str, object])
|
||||
|
||||
|
||||
class _CohereParseDocument(TypedDict):
|
||||
type: ReadOnly[Literal["image_url"]]
|
||||
image_url: ReadOnly[str]
|
||||
|
||||
|
||||
class _CohereParseRequestBody(TypedDict):
|
||||
model: ReadOnly[str]
|
||||
document: ReadOnly[_CohereParseDocument]
|
||||
output_format: ReadOnly[str]
|
||||
|
||||
|
||||
class _MarkdownPage(TypedDict):
|
||||
index: ReadOnly[int]
|
||||
markdown: ReadOnly[str]
|
||||
images: ReadOnly[Sequence[OCRPageImage] | None]
|
||||
|
||||
|
||||
class _BlocksPage(_MarkdownPage):
|
||||
blocks: ReadOnly[Sequence[Mapping[str, object]]]
|
||||
|
||||
|
||||
class _CohereParseMarkdown(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
content: str = ""
|
||||
images: Sequence[Mapping[str, object]] | None = None
|
||||
|
||||
|
||||
class _CohereParsePage(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
index: int | None = None
|
||||
markdown: _CohereParseMarkdown | None = None
|
||||
blocks: Sequence[Mapping[str, object]] | None = None
|
||||
|
||||
|
||||
class _CohereParseBilledUnits(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
pages: int | None = None
|
||||
|
||||
|
||||
class _CohereParseMeta(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
billed_units: _CohereParseBilledUnits | None = None
|
||||
|
||||
|
||||
class _CohereParseResponse(BaseModel):
|
||||
model_config = ConfigDict(frozen=True, extra="allow")
|
||||
|
||||
pages: Sequence[_CohereParsePage] = ()
|
||||
meta: _CohereParseMeta | None = None
|
||||
|
||||
|
||||
def _requested_format(optional_params: Mapping[str, object] | None) -> OCRRequestFormat:
|
||||
if optional_params is None:
|
||||
return "litellm"
|
||||
return "native" if optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native" else "litellm"
|
||||
|
||||
|
||||
def _page_image(image: Mapping[str, object]) -> OCRPageImage:
|
||||
bounding_box: Final = image.get("bounding_box")
|
||||
if not isinstance(bounding_box, Mapping):
|
||||
return OCRPageImage.model_validate(image)
|
||||
bbox: Final = _BOUNDING_BOX_ADAPTER.validate_python(bounding_box)
|
||||
return OCRPageImage.model_validate(MappingProxyType({**image, "bbox": bbox}))
|
||||
|
||||
|
||||
def _normalize_page(page: _CohereParsePage, position: int) -> OCRPage:
|
||||
markdown: Final = page.markdown
|
||||
images: Final = tuple(_page_image(image) for image in markdown.images) if markdown and markdown.images else None
|
||||
normalized: Final[_MarkdownPage] = {
|
||||
"index": page.index if page.index is not None else position,
|
||||
"markdown": markdown.content if markdown else "",
|
||||
"images": images,
|
||||
}
|
||||
if page.blocks is None:
|
||||
return OCRPage.model_validate(normalized)
|
||||
with_blocks: Final[_BlocksPage] = {**normalized, "blocks": page.blocks}
|
||||
return OCRPage.model_validate(with_blocks)
|
||||
|
||||
|
||||
def _billed_pages(parsed: _CohereParseResponse) -> int | None:
|
||||
if parsed.meta is None or parsed.meta.billed_units is None:
|
||||
return None
|
||||
return parsed.meta.billed_units.pages
|
||||
|
||||
|
||||
class CohereParseConfig(BaseOCRConfig):
|
||||
"""Cohere Parse, an image-only document understanding endpoint returning markdown or blocks."""
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list[str]: # mutable-ok: BaseOCRConfig signature
|
||||
return list(COHERE_PARSE_SUPPORTED_PARAMS) # mutable-ok: BaseOCRConfig signature
|
||||
|
||||
def get_api_key_env_var(self) -> str | None:
|
||||
return COHERE_API_KEY_ENV_VAR
|
||||
|
||||
def supports_rust_bridge(self) -> bool:
|
||||
return False
|
||||
|
||||
def get_health_check_document(self) -> DocumentType:
|
||||
return { # mutable-ok: litellm.aocr rejects any document that is not a dict
|
||||
"type": "image_url",
|
||||
"image_url": COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI,
|
||||
}
|
||||
|
||||
def _llm_provider(self) -> str:
|
||||
return "cohere"
|
||||
|
||||
def map_ocr_params(
|
||||
self,
|
||||
non_default_params: Mapping[str, object],
|
||||
optional_params: Mapping[str, object],
|
||||
model: str,
|
||||
) -> dict[str, object]: # mutable-ok: BaseOCRConfig signature
|
||||
output_format: Final = non_default_params.get(COHERE_PARSE_OUTPUT_FORMAT_PARAM)
|
||||
if output_format is not None and output_format not in COHERE_PARSE_OUTPUT_FORMATS:
|
||||
raise UnsupportedParamsError(
|
||||
message=(
|
||||
f"Invalid `{COHERE_PARSE_OUTPUT_FORMAT_PARAM}`: {output_format!r}. "
|
||||
f"Expected one of {', '.join(COHERE_PARSE_OUTPUT_FORMATS)}."
|
||||
),
|
||||
model=model,
|
||||
llm_provider=self._llm_provider(),
|
||||
)
|
||||
requested_format: Final = non_default_params.get(OCR_REQUEST_FORMAT_PARAM)
|
||||
request_format: Final = parse_ocr_request_format(requested_format) if requested_format is not None else None
|
||||
overrides: Final = tuple(
|
||||
(key, value)
|
||||
for key, value in (
|
||||
(COHERE_PARSE_OUTPUT_FORMAT_PARAM, output_format),
|
||||
(OCR_REQUEST_FORMAT_PARAM, request_format),
|
||||
)
|
||||
if value is not None
|
||||
)
|
||||
return {**optional_params, **dict(overrides)} # mutable-ok: BaseOCRConfig signature
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: Mapping[str, str],
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.validate_environment signature
|
||||
) -> dict[str, str]: # mutable-ok: BaseOCRConfig signature
|
||||
resolved_key: Final = api_key or get_secret_str(COHERE_API_KEY_ENV_VAR)
|
||||
if resolved_key is None:
|
||||
raise ValueError(
|
||||
f"Missing {COHERE_API_KEY_ENV_VAR} - set it in the environment or pass api_key to "
|
||||
"litellm.ocr()/litellm.aocr()"
|
||||
)
|
||||
return { # mutable-ok: BaseOCRConfig signature
|
||||
"Authorization": f"Bearer {resolved_key}",
|
||||
"Content-Type": "application/json",
|
||||
**headers,
|
||||
}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
model: str,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.get_complete_url signature
|
||||
) -> str:
|
||||
url: Final = httpx.URL(api_base or COHERE_PARSE_API_BASE)
|
||||
path: Final = url.path.rstrip("/")
|
||||
if path.endswith(COHERE_PARSE_PATH):
|
||||
return str(url.copy_with(path=path))
|
||||
if path.endswith("/v2"):
|
||||
return str(url.copy_with(path=f"{path}/parse"))
|
||||
return str(url.copy_with(path=f"{path}{COHERE_PARSE_PATH}"))
|
||||
|
||||
def _image_url(self, document: DocumentType, model: str) -> str:
|
||||
image_url: Final = document.get("image_url", "")
|
||||
if document.get("type") != "image_url" or not image_url or image_url.startswith("data:application/pdf"):
|
||||
raise BadRequestError(
|
||||
message=COHERE_PARSE_IMAGE_ONLY_MESSAGE,
|
||||
model=model,
|
||||
llm_provider=self._llm_provider(),
|
||||
)
|
||||
return image_url
|
||||
|
||||
def _resolve_image_url_sync(self, image_url: str) -> str:
|
||||
return image_url
|
||||
|
||||
async def _resolve_image_url_async(self, image_url: str) -> str:
|
||||
return image_url
|
||||
|
||||
def _build_request(self, model: str, image_url: str, optional_params: Mapping[str, object]) -> OCRRequestData:
|
||||
body: Final[_CohereParseRequestBody] = {
|
||||
"model": model,
|
||||
"document": {"type": "image_url", "image_url": image_url},
|
||||
"output_format": str(
|
||||
optional_params.get(COHERE_PARSE_OUTPUT_FORMAT_PARAM, COHERE_PARSE_DEFAULT_OUTPUT_FORMAT)
|
||||
),
|
||||
}
|
||||
return OCRRequestData(data=dict(body), files=None) # mutable-ok: OCRRequestData.data is a dict
|
||||
|
||||
def transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: Mapping[str, object],
|
||||
headers: Mapping[str, str],
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.transform_ocr_request signature
|
||||
) -> OCRRequestData:
|
||||
image_url: Final = self._resolve_image_url_sync(self._image_url(document, model))
|
||||
return self._build_request(model=model, image_url=image_url, optional_params=optional_params)
|
||||
|
||||
async def async_transform_ocr_request(
|
||||
self,
|
||||
model: str,
|
||||
document: DocumentType,
|
||||
optional_params: Mapping[str, object],
|
||||
headers: Mapping[str, str],
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.async_transform_ocr_request signature
|
||||
) -> OCRRequestData:
|
||||
image_url: Final = await self._resolve_image_url_async(self._image_url(document, model))
|
||||
return self._build_request(model=model, image_url=image_url, optional_params=optional_params)
|
||||
|
||||
def transform_ocr_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: "LiteLLMLoggingObj",
|
||||
optional_params: Mapping[str, object] | None = None,
|
||||
**kwargs: object, # kwargs-ok: BaseOCRConfig.transform_ocr_response signature
|
||||
) -> OCRResponse:
|
||||
native: Final = _NATIVE_RESPONSE_ADAPTER.validate_python(raw_response.json())
|
||||
parsed: Final = _CohereParseResponse.model_validate(native)
|
||||
pages: Final = [ # mutable-ok: OCRResponse.pages is a list
|
||||
_normalize_page(page, position) for position, page in enumerate(parsed.pages)
|
||||
]
|
||||
billed_pages: Final = _billed_pages(parsed)
|
||||
response: Final = OCRResponse(
|
||||
pages=pages,
|
||||
model=model,
|
||||
usage_info=OCRUsageInfo(pages_processed=billed_pages if billed_pages is not None else len(pages)),
|
||||
)
|
||||
if _requested_format(optional_params) == "native":
|
||||
response.set_provider_native_response(native)
|
||||
return response
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: Mapping[str, str],
|
||||
) -> Exception:
|
||||
return CohereError(status_code=status_code, message=error_message)
|
||||
|
|
@ -10243,6 +10243,16 @@
|
|||
],
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/mistral/"
|
||||
},
|
||||
"azure_ai/Cohere-parse-v5": {
|
||||
"deprecation_date": "2026-12-15",
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "ocr",
|
||||
"ocr_cost_per_page": 0.0015,
|
||||
"source": "https://cohere.com/blog/parse",
|
||||
"supported_endpoints": [
|
||||
"/v1/ocr"
|
||||
]
|
||||
},
|
||||
"azure_ai/doc-intelligence/prebuilt-read": {
|
||||
"litellm_provider": "azure_ai",
|
||||
"ocr_cost_per_page": 0.0015,
|
||||
|
|
@ -14116,6 +14126,15 @@
|
|||
"output_vector_size": 1536,
|
||||
"supports_embedding_image_input": true
|
||||
},
|
||||
"cohere/parse-v5.0": {
|
||||
"litellm_provider": "cohere",
|
||||
"mode": "ocr",
|
||||
"ocr_cost_per_page": 0.0015,
|
||||
"source": "https://cohere.com/blog/parse",
|
||||
"supported_endpoints": [
|
||||
"/v1/ocr"
|
||||
]
|
||||
},
|
||||
"cohere.rerank-v3-5:0": {
|
||||
"input_cost_per_query": 0.002,
|
||||
"input_cost_per_token": 0.0,
|
||||
|
|
|
|||
|
|
@ -191,6 +191,8 @@ def _prepare_ocr_request(
|
|||
def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool:
|
||||
if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native":
|
||||
return False
|
||||
if not prepared_request.provider_config.supports_rust_bridge():
|
||||
return False
|
||||
return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -559,6 +559,7 @@
|
|||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": true,
|
||||
"ocr": true,
|
||||
"a2a": true,
|
||||
"interactions": true
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9294,6 +9294,11 @@ class ProviderConfigManager:
|
|||
|
||||
return get_vertex_ai_ocr_config(model=model)
|
||||
|
||||
if provider == litellm.LlmProviders.COHERE:
|
||||
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
|
||||
|
||||
return CohereParseConfig()
|
||||
|
||||
if provider == litellm.LlmProviders.REDUCTO:
|
||||
from litellm.llms.reducto.ocr.transformation import (
|
||||
ReductoParseLegacyConfig,
|
||||
|
|
|
|||
|
|
@ -10243,6 +10243,16 @@
|
|||
],
|
||||
"source": "https://azure.microsoft.com/en-us/pricing/details/ai-foundry-models/mistral/"
|
||||
},
|
||||
"azure_ai/Cohere-parse-v5": {
|
||||
"deprecation_date": "2026-12-15",
|
||||
"litellm_provider": "azure_ai",
|
||||
"mode": "ocr",
|
||||
"ocr_cost_per_page": 0.0015,
|
||||
"source": "https://cohere.com/blog/parse",
|
||||
"supported_endpoints": [
|
||||
"/v1/ocr"
|
||||
]
|
||||
},
|
||||
"azure_ai/doc-intelligence/prebuilt-read": {
|
||||
"litellm_provider": "azure_ai",
|
||||
"ocr_cost_per_page": 0.0015,
|
||||
|
|
@ -14116,6 +14126,15 @@
|
|||
"output_vector_size": 1536,
|
||||
"supports_embedding_image_input": true
|
||||
},
|
||||
"cohere/parse-v5.0": {
|
||||
"litellm_provider": "cohere",
|
||||
"mode": "ocr",
|
||||
"ocr_cost_per_page": 0.0015,
|
||||
"source": "https://cohere.com/blog/parse",
|
||||
"supported_endpoints": [
|
||||
"/v1/ocr"
|
||||
]
|
||||
},
|
||||
"cohere.rerank-v3-5:0": {
|
||||
"input_cost_per_query": 0.002,
|
||||
"input_cost_per_token": 0.0,
|
||||
|
|
|
|||
|
|
@ -594,6 +594,7 @@
|
|||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": true,
|
||||
"ocr": true,
|
||||
"a2a": true,
|
||||
"interactions": true
|
||||
}
|
||||
|
|
|
|||
|
|
@ -453,3 +453,32 @@ async def test_realtime_health_check_uses_model_level_vertex_params():
|
|||
"Authorization": "Bearer model-level-token",
|
||||
"x-goog-user-project": "model-level-project",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"model, custom_llm_provider, expected_document_type, expected_uri_prefix",
|
||||
[
|
||||
("mistral/mistral-ocr-latest", "mistral", "document_url", "data:application/pdf;base64,"),
|
||||
("azure_ai/mistral-document-ai-2512", "azure_ai", "document_url", "data:application/pdf;base64,"),
|
||||
("cohere/parse-v5.0", "cohere", "image_url", "data:image/png;base64,"),
|
||||
("azure_ai/Cohere-parse-v5", "azure_ai", "image_url", "data:image/png;base64,"),
|
||||
],
|
||||
)
|
||||
async def test_ocr_health_check_sends_the_document_kind_the_provider_config_accepts(
|
||||
model, custom_llm_provider, expected_document_type, expected_uri_prefix
|
||||
):
|
||||
handlers = HealthCheckHelpers.get_mode_handlers(
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
model_params={"model": model, "api_key": "sk-test"},
|
||||
)
|
||||
|
||||
with patch( # test-quality-ok: the public health-check path has no dependency injection seam
|
||||
"litellm.aocr", new_callable=AsyncMock, return_value={}
|
||||
) as mock_aocr:
|
||||
await handlers["ocr"]()
|
||||
|
||||
document = mock_aocr.call_args.kwargs["document"]
|
||||
assert document["type"] == expected_document_type
|
||||
assert document[expected_document_type].startswith(expected_uri_prefix)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,184 @@
|
|||
import base64
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from litellm.llms.azure_ai.ocr.common_utils import get_azure_ai_ocr_config
|
||||
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import AzureDocumentIntelligenceOCRConfig
|
||||
from litellm.llms.azure_ai.ocr.transformation import AzureAIOCRConfig
|
||||
|
||||
MODEL = "azure_ai/Cohere-parse-v5"
|
||||
API_BASE = "https://resource.services.ai.azure.com"
|
||||
PARSE_URL = f"{API_BASE}/providers/cohere/v2/parse"
|
||||
IMAGE_URL = "https://example.com/receipt.png"
|
||||
PNG_BYTES = base64.b64decode(
|
||||
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
|
||||
)
|
||||
PNG_DATA_URI = f"data:image/png;base64,{base64.b64encode(PNG_BYTES).decode()}"
|
||||
|
||||
|
||||
def _parse_response() -> dict:
|
||||
return {
|
||||
"id": "882bf973-9dfa-4d02-9d30-709247008efd",
|
||||
"pages": [{"index": 0, "type": "markdown", "markdown": {"content": "# Receipt\n\nTotal Due: $4.00"}}],
|
||||
"meta": {"api_version": {"version": "2"}, "billed_units": {"pages": 1}},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def disable_aiohttp_transport(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
yield
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model, expected_config",
|
||||
[
|
||||
("Cohere-parse-v5", AzureAICohereParseConfig),
|
||||
("cohere-parse-v5", AzureAICohereParseConfig),
|
||||
("cohere/parse-v5", AzureAICohereParseConfig),
|
||||
("invoice-parser", AzureAIOCRConfig),
|
||||
("parse-v5", AzureAIOCRConfig),
|
||||
("mistral-ocr-4-0", AzureAIOCRConfig),
|
||||
("mistral-document-ai-2512", AzureAIOCRConfig),
|
||||
("doc-intelligence/prebuilt-read", AzureDocumentIntelligenceOCRConfig),
|
||||
],
|
||||
)
|
||||
def test_azure_ai_ocr_routing(model: str, expected_config: type) -> None:
|
||||
assert type(get_azure_ai_ocr_config(model)) is expected_config
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"api_base, expected_url",
|
||||
[
|
||||
(API_BASE, PARSE_URL),
|
||||
(f"{API_BASE}/", PARSE_URL),
|
||||
(f"{API_BASE}/models", PARSE_URL),
|
||||
(f"{API_BASE}/providers/cohere/v2", PARSE_URL),
|
||||
(f"{API_BASE}/providers/cohere/v2/parse", PARSE_URL),
|
||||
],
|
||||
)
|
||||
def test_get_complete_url_targets_the_cohere_provider_route(api_base: str, expected_url: str) -> None:
|
||||
url = AzureAICohereParseConfig().get_complete_url(api_base=api_base, model="Cohere-parse-v5", optional_params={})
|
||||
|
||||
assert url == expected_url
|
||||
|
||||
|
||||
def test_get_complete_url_falls_back_to_env_api_base(monkeypatch) -> None:
|
||||
monkeypatch.setenv("AZURE_AI_API_BASE", API_BASE)
|
||||
|
||||
url = AzureAICohereParseConfig().get_complete_url(api_base=None, model="Cohere-parse-v5", optional_params={})
|
||||
|
||||
assert url == PARSE_URL
|
||||
|
||||
|
||||
def test_get_complete_url_requires_api_base(monkeypatch) -> None:
|
||||
monkeypatch.delenv("AZURE_AI_API_BASE", raising=False)
|
||||
|
||||
with pytest.raises(ValueError, match="AZURE_AI_API_BASE"):
|
||||
AzureAICohereParseConfig().get_complete_url(api_base=None, model="Cohere-parse-v5", optional_params={})
|
||||
|
||||
|
||||
def test_get_complete_url_rejects_relative_api_base() -> None:
|
||||
with pytest.raises(ValueError, match="absolute URL"):
|
||||
AzureAICohereParseConfig().get_complete_url(
|
||||
api_base="resource.services.ai.azure.com", model="Cohere-parse-v5", optional_params={}
|
||||
)
|
||||
|
||||
|
||||
def test_validate_environment_requires_api_base(monkeypatch) -> None:
|
||||
monkeypatch.delenv("AZURE_AI_API_BASE", raising=False)
|
||||
|
||||
with pytest.raises(ValueError, match="AZURE_AI_API_BASE"):
|
||||
AzureAICohereParseConfig().validate_environment(headers={}, model="Cohere-parse-v5", api_key="key")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_inlines_remote_image_and_posts_to_foundry(disable_aiohttp_transport, respx_mock):
|
||||
respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"})
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
response = await litellm.aocr(
|
||||
model=MODEL,
|
||||
document={"type": "image_url", "image_url": IMAGE_URL},
|
||||
api_base=API_BASE,
|
||||
api_key="azure-key",
|
||||
)
|
||||
|
||||
request = route.calls.last.request
|
||||
assert request.headers["Authorization"] == "Bearer azure-key"
|
||||
assert json.loads(request.content) == {
|
||||
"model": "Cohere-parse-v5",
|
||||
"document": {"type": "image_url", "image_url": PNG_DATA_URI},
|
||||
"output_format": "markdown",
|
||||
}
|
||||
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
|
||||
assert response.usage_info.pages_processed == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_passes_data_uri_through_without_fetching(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
await litellm.aocr(
|
||||
model=MODEL,
|
||||
document={"type": "image_url", "image_url": PNG_DATA_URI},
|
||||
api_base=API_BASE,
|
||||
api_key="azure-key",
|
||||
output_format="blocks",
|
||||
)
|
||||
|
||||
body = json.loads(route.calls.last.request.content)
|
||||
assert body["document"]["image_url"] == PNG_DATA_URI
|
||||
assert body["output_format"] == "blocks"
|
||||
|
||||
|
||||
def test_ocr_sync_inlines_remote_image(respx_mock):
|
||||
respx_mock.get(IMAGE_URL).respond(content=PNG_BYTES, headers={"Content-Type": "image/png"})
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
response = litellm.ocr(
|
||||
model=MODEL,
|
||||
document={"type": "image_url", "image_url": IMAGE_URL},
|
||||
api_base=API_BASE,
|
||||
api_key="azure-key",
|
||||
)
|
||||
|
||||
assert json.loads(route.calls.last.request.content)["document"]["image_url"] == PNG_DATA_URI
|
||||
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_rejects_pdf_before_calling_foundry(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info:
|
||||
await litellm.aocr(
|
||||
model=MODEL,
|
||||
document={"type": "document_url", "document_url": "https://example.com/doc.pdf"},
|
||||
api_base=API_BASE,
|
||||
api_key="azure-key",
|
||||
)
|
||||
|
||||
assert exc_info.value.llm_provider == "azure_ai"
|
||||
assert not route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ahealth_check_ocr_sends_an_image_to_the_foundry_cohere_parse_deployment(
|
||||
disable_aiohttp_transport, respx_mock
|
||||
):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_parse_response())
|
||||
|
||||
result = await litellm.ahealth_check(
|
||||
model_params={"model": MODEL, "api_base": API_BASE, "api_key": "test-key"}, mode="ocr"
|
||||
)
|
||||
|
||||
document = json.loads(route.calls.last.request.content)["document"]
|
||||
assert document["type"] == "image_url"
|
||||
assert document["image_url"].startswith("data:image/png;base64,")
|
||||
assert "error" not in result
|
||||
58
tests/test_litellm/llms/cohere/ocr/test_cohere_parse_cost.py
Normal file
58
tests/test_litellm/llms/cohere/ocr/test_cohere_parse_cost.py
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.cost_calculator import completion_cost
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRPage, OCRResponse, OCRUsageInfo
|
||||
|
||||
COST_PER_PAGE = 0.0015
|
||||
REPO_ROOT = Path(__file__).parents[5]
|
||||
COST_MAPS = [
|
||||
REPO_ROOT / "model_prices_and_context_window.json",
|
||||
REPO_ROOT / "litellm" / "model_prices_and_context_window_backup.json",
|
||||
]
|
||||
MODELS = [("cohere/parse-v5.0", "cohere"), ("azure_ai/Cohere-parse-v5", "azure_ai")]
|
||||
|
||||
|
||||
def _ocr_response(model: str, pages_processed: int) -> OCRResponse:
|
||||
return OCRResponse(
|
||||
pages=[OCRPage(index=i, markdown=f"page {i}") for i in range(pages_processed)],
|
||||
model=model,
|
||||
usage_info=OCRUsageInfo(pages_processed=pages_processed),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("cost_map_path", COST_MAPS, ids=lambda path: path.name)
|
||||
@pytest.mark.parametrize("model, provider", MODELS)
|
||||
def test_pricing_entry(cost_map_path: Path, model: str, provider: str) -> None:
|
||||
with open(cost_map_path) as f:
|
||||
info = json.load(f).get(model)
|
||||
|
||||
assert info is not None, f"{model} missing from {cost_map_path.name}"
|
||||
assert info["litellm_provider"] == provider
|
||||
assert info["mode"] == "ocr"
|
||||
assert info["supported_endpoints"] == ["/v1/ocr"]
|
||||
assert info["ocr_cost_per_page"] == COST_PER_PAGE
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model, provider", MODELS)
|
||||
def test_model_info_resolves_ocr_mode_and_price(local_model_cost_map, model: str, provider: str) -> None:
|
||||
info = litellm.get_model_info(model=model, custom_llm_provider=provider)
|
||||
|
||||
assert info["mode"] == "ocr"
|
||||
assert info["ocr_cost_per_page"] == COST_PER_PAGE
|
||||
|
||||
|
||||
@pytest.mark.parametrize("model, provider", MODELS)
|
||||
@pytest.mark.parametrize("pages_processed", [1, 3])
|
||||
def test_cost_scales_with_billed_pages(local_model_cost_map, model: str, provider: str, pages_processed: int) -> None:
|
||||
cost = completion_cost(
|
||||
completion_response=_ocr_response(model.split("/", 1)[1], pages_processed),
|
||||
model=model,
|
||||
custom_llm_provider=provider,
|
||||
call_type="ocr",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(COST_PER_PAGE * pages_processed)
|
||||
|
|
@ -0,0 +1,229 @@
|
|||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
PARSE_URL = "https://api.cohere.com/v2/parse"
|
||||
MODEL = "cohere/parse-v5.0"
|
||||
IMAGE_DOCUMENT = {"type": "image_url", "image_url": "https://example.com/receipt.png"}
|
||||
BOUNDING_BOX = {"top_left_x": 0, "top_left_y": 0, "bottom_right_x": 32, "bottom_right_y": 32}
|
||||
|
||||
|
||||
def _markdown_response(billed_pages: int | None = 2) -> dict:
|
||||
return {
|
||||
"id": "272900cc-04c0-4da2-a505-2cea58d231bf",
|
||||
"pages": [
|
||||
{
|
||||
"index": 0,
|
||||
"type": "markdown",
|
||||
"markdown": {
|
||||
"content": "# Receipt\n\nTotal Due: $4.00",
|
||||
"images": [
|
||||
{
|
||||
"id": "img-0",
|
||||
"description": "A parking receipt",
|
||||
"category": "other",
|
||||
"bounding_box": BOUNDING_BOX,
|
||||
"bounding_box_normalized": {
|
||||
"top_left_x": 0,
|
||||
"top_left_y": 0,
|
||||
"bottom_right_x": 1,
|
||||
"bottom_right_y": 1,
|
||||
},
|
||||
}
|
||||
],
|
||||
},
|
||||
},
|
||||
{"index": 1, "type": "markdown", "markdown": {"content": "Page two"}},
|
||||
],
|
||||
**(
|
||||
{"meta": {"api_version": {"version": "2"}, "billed_units": {"pages": billed_pages}}} if billed_pages else {}
|
||||
),
|
||||
}
|
||||
|
||||
|
||||
def _blocks_response() -> dict:
|
||||
return {
|
||||
"id": "94474f83-e30d-4763-b4bc-52af6e12c4f7",
|
||||
"pages": [
|
||||
{
|
||||
"index": 0,
|
||||
"type": "blocks",
|
||||
"blocks": [{"type": "text", "text": "Total Due: $4.00"}],
|
||||
}
|
||||
],
|
||||
"meta": {"api_version": {"version": "2"}, "billed_units": {"pages": 1}},
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def disable_aiohttp_transport(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
yield
|
||||
litellm.in_memory_llm_clients_cache.flush_cache()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_sends_markdown_parse_request_and_normalizes_pages(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
|
||||
|
||||
request = route.calls.last.request
|
||||
assert request.headers["Authorization"] == "Bearer test-key"
|
||||
assert json.loads(request.content) == {
|
||||
"model": "parse-v5.0",
|
||||
"document": IMAGE_DOCUMENT,
|
||||
"output_format": "markdown",
|
||||
}
|
||||
assert response.object == "ocr"
|
||||
assert [page.index for page in response.pages] == [0, 1]
|
||||
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
|
||||
assert response.pages[1].markdown == "Page two"
|
||||
assert response.pages[1].images is None
|
||||
image = response.pages[0].images[0]
|
||||
assert image.bbox == BOUNDING_BOX
|
||||
assert image.model_extra["description"] == "A parking receipt"
|
||||
assert image.model_extra["bounding_box_normalized"]["bottom_right_x"] == 1
|
||||
assert response.usage_info.pages_processed == 2
|
||||
assert response.get_provider_native_response() is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_usage_prefers_billed_units_over_page_count(disable_aiohttp_transport, respx_mock):
|
||||
respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=3))
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
|
||||
|
||||
assert response.usage_info.pages_processed == 3
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_usage_falls_back_to_page_count_without_meta(disable_aiohttp_transport, respx_mock):
|
||||
respx_mock.post(PARSE_URL).respond(json=_markdown_response(billed_pages=None))
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
|
||||
|
||||
assert response.usage_info.pages_processed == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_blocks_output_format_forwards_param_and_keeps_blocks(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_blocks_response())
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="blocks")
|
||||
|
||||
assert json.loads(route.calls.last.request.content)["output_format"] == "blocks"
|
||||
assert response.pages[0].markdown == ""
|
||||
assert response.pages[0].model_extra["blocks"] == [{"type": "text", "text": "Total Due: $4.00"}]
|
||||
assert response.usage_info.pages_processed == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_native_format_carries_provider_payload(disable_aiohttp_transport, respx_mock):
|
||||
payload = _markdown_response()
|
||||
route = respx_mock.post(PARSE_URL).respond(json=payload)
|
||||
|
||||
response = await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", req_format="native")
|
||||
|
||||
assert "req_format" not in json.loads(route.calls.last.request.content)
|
||||
assert response.get_provider_native_response() == payload
|
||||
assert response.pages[0].markdown == "# Receipt\n\nTotal Due: $4.00"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_rejects_unknown_output_format_before_calling_provider(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="Invalid `output_format`: 'html'") as exc_info:
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", output_format="html")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert not route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"document",
|
||||
[
|
||||
{"type": "document_url", "document_url": "https://example.com/doc.pdf"},
|
||||
{"type": "image_url", "image_url": "data:application/pdf;base64,JVBERi0="},
|
||||
{"type": "image_url", "image_url": ""},
|
||||
],
|
||||
)
|
||||
async def test_aocr_rejects_non_image_documents_before_calling_provider(
|
||||
disable_aiohttp_transport, respx_mock, document
|
||||
):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="only accepts `image_url` documents") as exc_info:
|
||||
await litellm.aocr(model=MODEL, document=document, api_key="test-key")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert not route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"api_base, expected_url",
|
||||
[
|
||||
("https://gateway.example.com", "https://gateway.example.com/v2/parse"),
|
||||
("https://gateway.example.com/cohere/", "https://gateway.example.com/cohere/v2/parse"),
|
||||
("https://gateway.example.com/v2", "https://gateway.example.com/v2/parse"),
|
||||
("https://gateway.example.com/v2/parse", "https://gateway.example.com/v2/parse"),
|
||||
],
|
||||
)
|
||||
async def test_aocr_posts_to_api_base_variants(disable_aiohttp_transport, respx_mock, api_base, expected_url):
|
||||
route = respx_mock.post(expected_url).respond(json=_markdown_response())
|
||||
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key", api_base=api_base)
|
||||
|
||||
assert route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_surfaces_provider_error_with_its_status_and_message(disable_aiohttp_transport, respx_mock):
|
||||
respx_mock.post(PARSE_URL).respond(
|
||||
status_code=400, json={"id": "83b0d95e", "message": "output_format must be `blocks` or `markdown`"}
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.BadRequestError, match="output_format must be") as exc_info:
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT, api_key="test-key")
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_reads_api_key_from_environment(disable_aiohttp_transport, respx_mock, monkeypatch):
|
||||
monkeypatch.setenv("COHERE_API_KEY", "env-key")
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT)
|
||||
|
||||
assert route.calls.last.request.headers["Authorization"] == "Bearer env-key"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aocr_without_api_key_names_the_env_var(disable_aiohttp_transport, respx_mock, monkeypatch):
|
||||
monkeypatch.delenv("COHERE_API_KEY", raising=False)
|
||||
monkeypatch.setattr(litellm, "cohere_key", None)
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
with pytest.raises(Exception, match="Missing COHERE_API_KEY"):
|
||||
await litellm.aocr(model=MODEL, document=IMAGE_DOCUMENT)
|
||||
|
||||
assert not route.called
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ahealth_check_ocr_sends_an_image_cohere_parse_accepts(disable_aiohttp_transport, respx_mock):
|
||||
route = respx_mock.post(PARSE_URL).respond(json=_markdown_response())
|
||||
|
||||
result = await litellm.ahealth_check(model_params={"model": MODEL, "api_key": "test-key"}, mode="ocr")
|
||||
|
||||
document = json.loads(route.calls.last.request.content)["document"]
|
||||
assert document["type"] == "image_url"
|
||||
assert document["image_url"].startswith("data:image/png;base64,")
|
||||
assert "error" not in result
|
||||
|
|
@ -4,11 +4,14 @@ providers that don't support a native response must reject it, and the Rust
|
|||
bridge (which only returns the normalized shape) must not serve native requests.
|
||||
"""
|
||||
|
||||
import dataclasses
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig
|
||||
from litellm.llms.cohere.ocr.transformation import CohereParseConfig
|
||||
from litellm.ocr.main import _PreparedOCRRequest, _rust_ocr_supported
|
||||
|
||||
DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"}
|
||||
|
|
@ -39,6 +42,13 @@ def test_rust_ocr_skipped_for_native_format():
|
|||
assert _rust_ocr_supported(_prepared({"req_format": "native"})) is False
|
||||
|
||||
|
||||
@pytest.mark.parametrize("provider_config", [CohereParseConfig(), AzureAICohereParseConfig()])
|
||||
def test_rust_ocr_skipped_for_configs_without_bridge_support(provider_config):
|
||||
prepared = dataclasses.replace(_prepared({}), provider_config=provider_config)
|
||||
|
||||
assert _rust_ocr_supported(prepared) is False
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_native_format_rejected_for_provider_without_support_as_bad_request():
|
||||
with pytest.raises(litellm.BadRequestError, match="not supported for provider") as exc_info:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue