mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(ocr): preserve Azure Document Intelligence authentication
This commit is contained in:
parent
9132a5343b
commit
1b88828229
5 changed files with 134 additions and 19 deletions
|
|
@ -92,6 +92,18 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
def get_api_key_env_var(self) -> str | None:
|
||||
return AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV_VAR
|
||||
|
||||
def resolve_connection_params(
|
||||
self,
|
||||
*,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
dynamic_api_key: str | None,
|
||||
dynamic_api_base: str | None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
explicit_api_key: Final = None if api_key is None else dynamic_api_key or api_key
|
||||
explicit_api_base: Final = None if api_base is None else dynamic_api_base or api_base
|
||||
return explicit_api_key, explicit_api_base
|
||||
|
||||
def get_supported_ocr_params(self, model: str) -> list:
|
||||
"""
|
||||
Get supported OCR parameters for Azure Document Intelligence.
|
||||
|
|
@ -618,7 +630,11 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
|
|||
except SSRFError as ssrf_err:
|
||||
raise ValueError(f"Azure Document Intelligence: rejected polling URL ({ssrf_err})")
|
||||
|
||||
poll_headers = {"Ocp-Apim-Subscription-Key": raw_response.request.headers.get("Ocp-Apim-Subscription-Key", "")}
|
||||
poll_headers: Final = {
|
||||
header: raw_response.request.headers[header]
|
||||
for header in ("Ocp-Apim-Subscription-Key", "Authorization")
|
||||
if header in raw_response.request.headers
|
||||
}
|
||||
return operation_url, poll_headers
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -144,6 +144,16 @@ class BaseOCRConfig:
|
|||
"""
|
||||
return None
|
||||
|
||||
def resolve_connection_params(
|
||||
self,
|
||||
*,
|
||||
api_key: str | None,
|
||||
api_base: str | None,
|
||||
dynamic_api_key: str | None,
|
||||
dynamic_api_base: str | None,
|
||||
) -> tuple[str | None, str | None]:
|
||||
return dynamic_api_key or api_key, dynamic_api_base or api_base
|
||||
|
||||
def get_health_check_document(self) -> DocumentType:
|
||||
return { # mutable-ok: litellm.aocr rejects any document that is not a dict
|
||||
"type": "document_url",
|
||||
|
|
|
|||
|
|
@ -19,9 +19,6 @@ import litellm
|
|||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import request_timeout
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.azure_ai.ocr.common_utils import (
|
||||
is_azure_document_intelligence_model,
|
||||
)
|
||||
from litellm.llms.base_llm.ocr.transformation import (
|
||||
OCR_REQUEST_FORMAT_PARAM,
|
||||
BaseOCRConfig,
|
||||
|
|
@ -80,8 +77,6 @@ def _prepare_ocr_request(
|
|||
if doc_type not in ["document_url", "image_url"]:
|
||||
raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url', 'image_url', or 'file'")
|
||||
|
||||
caller_supplied_api_base: Final = api_base is not None
|
||||
|
||||
(
|
||||
model,
|
||||
custom_llm_provider,
|
||||
|
|
@ -94,16 +89,6 @@ def _prepare_ocr_request(
|
|||
api_key=api_key,
|
||||
)
|
||||
|
||||
suppress_dynamic_api_base: Final = (
|
||||
not caller_supplied_api_base
|
||||
and custom_llm_provider == "azure_ai"
|
||||
and is_azure_document_intelligence_model(model)
|
||||
)
|
||||
if dynamic_api_key:
|
||||
api_key = dynamic_api_key
|
||||
if dynamic_api_base and not suppress_dynamic_api_base:
|
||||
api_base = dynamic_api_base
|
||||
|
||||
ocr_provider_config: Final = ProviderConfigManager.get_provider_ocr_config(
|
||||
model=model,
|
||||
provider=litellm.LlmProviders(custom_llm_provider),
|
||||
|
|
@ -112,6 +97,13 @@ def _prepare_ocr_request(
|
|||
if ocr_provider_config is None:
|
||||
raise ValueError(f"OCR is not supported for provider: {custom_llm_provider}")
|
||||
|
||||
resolved_api_key, resolved_api_base = ocr_provider_config.resolve_connection_params(
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
dynamic_api_key=dynamic_api_key,
|
||||
dynamic_api_base=dynamic_api_base,
|
||||
)
|
||||
|
||||
verbose_logger.debug("OCR call - model: %s, provider: %s", model, custom_llm_provider)
|
||||
|
||||
litellm_params: Final = GenericLiteLLMParams.model_validate(kwargs)
|
||||
|
|
@ -156,7 +148,7 @@ def _prepare_ocr_request(
|
|||
optional_params=optional_params,
|
||||
litellm_params={
|
||||
"litellm_call_id": litellm_call_id,
|
||||
"api_base": api_base,
|
||||
"api_base": resolved_api_base,
|
||||
},
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
|
|
@ -164,8 +156,8 @@ def _prepare_ocr_request(
|
|||
return _PreparedOCRRequest(
|
||||
model=model,
|
||||
document=document,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
api_key=resolved_api_key,
|
||||
api_base=resolved_api_base,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
provider_config=ocr_provider_config,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from unittest.mock import MagicMock
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
|
@ -371,3 +372,35 @@ def test_validate_environment_falls_back_to_entra_token(monkeypatch):
|
|||
|
||||
assert headers["Authorization"] == "Bearer entra-token"
|
||||
assert "Ocp-Apim-Subscription-Key" not in headers
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("request_headers", "expected_poll_headers"),
|
||||
(
|
||||
(
|
||||
{"Ocp-Apim-Subscription-Key": "subscription-key"},
|
||||
{"Ocp-Apim-Subscription-Key": "subscription-key"},
|
||||
),
|
||||
(
|
||||
{"Authorization": "Bearer entra-token"},
|
||||
{"Authorization": "Bearer entra-token"},
|
||||
),
|
||||
),
|
||||
)
|
||||
def test_get_polling_target_preserves_request_authentication(
|
||||
request_headers: dict[str, str], expected_poll_headers: dict[str, str]
|
||||
) -> None:
|
||||
response: Final = httpx.Response(
|
||||
status_code=202,
|
||||
headers={"Operation-Location": "https://example.cognitiveservices.azure.com/operations/123"},
|
||||
request=httpx.Request(
|
||||
"POST",
|
||||
"https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze",
|
||||
headers=request_headers,
|
||||
),
|
||||
)
|
||||
|
||||
operation_url, poll_headers = AzureDocumentIntelligenceOCRConfig()._get_polling_target(response)
|
||||
|
||||
assert operation_url == "https://example.cognitiveservices.azure.com/operations/123"
|
||||
assert poll_headers == expected_poll_headers
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.llms.custom_httpx import llm_http_handler
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.ocr.legacy import _prepare_ocr_request
|
||||
from litellm.rust_bridge import bindings, configuration
|
||||
from litellm.rust_bridge.ocr_lifecycle import NATIVE_OCR_LIFECYCLE
|
||||
|
||||
|
|
@ -133,3 +134,66 @@ async def test_python_provider_errors_keep_public_exception(provider: Mock, asyn
|
|||
assert error.value.model == "mistral-ocr-latest"
|
||||
assert error.value.llm_provider == "mistral"
|
||||
assert provider.call_count == 1
|
||||
|
||||
|
||||
def test_document_intelligence_environment_key_is_not_replaced_by_generic_azure_key(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
monkeypatch.setenv("AZURE_AI_API_KEY", "generic-key")
|
||||
monkeypatch.setenv("AZURE_DOCUMENT_INTELLIGENCE_API_KEY", "document-key")
|
||||
prepared: Final = _prepare_ocr_request(
|
||||
model="azure_ai/doc-intelligence/prebuilt-layout",
|
||||
document={"type": "document_url", "document_url": "https://example.com/file.pdf"},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
timeout=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
kwargs={"litellm_logging_obj": Mock()},
|
||||
)
|
||||
|
||||
assert prepared.api_key is None
|
||||
headers: Final = prepared.provider_config.validate_environment(
|
||||
headers={},
|
||||
model=prepared.model,
|
||||
api_key=prepared.api_key,
|
||||
api_base=prepared.api_base,
|
||||
litellm_params=prepared.litellm_params,
|
||||
)
|
||||
assert headers["Ocp-Apim-Subscription-Key"] == "document-key"
|
||||
|
||||
|
||||
def test_document_intelligence_explicit_connection_is_preserved(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("AZURE_AI_API_KEY", "generic-key")
|
||||
monkeypatch.setenv("AZURE_AI_API_BASE", "https://generic.example.com")
|
||||
prepared: Final = _prepare_ocr_request(
|
||||
model="azure_ai/doc-intelligence/prebuilt-layout",
|
||||
document={"type": "document_url", "document_url": "https://example.com/file.pdf"},
|
||||
api_key="explicit-key",
|
||||
api_base="https://document.example.com",
|
||||
timeout=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
kwargs={"litellm_logging_obj": Mock()},
|
||||
)
|
||||
|
||||
assert prepared.api_key == "explicit-key"
|
||||
assert prepared.api_base == "https://document.example.com"
|
||||
|
||||
|
||||
def test_generic_azure_connection_still_applies_to_foundry_ocr(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("AZURE_AI_API_KEY", "generic-key")
|
||||
monkeypatch.setenv("AZURE_AI_API_BASE", "https://generic.example.com")
|
||||
prepared: Final = _prepare_ocr_request(
|
||||
model="azure_ai/mistral-document-ai-2505",
|
||||
document={"type": "document_url", "document_url": "https://example.com/file.pdf"},
|
||||
api_key=None,
|
||||
api_base=None,
|
||||
timeout=None,
|
||||
custom_llm_provider=None,
|
||||
extra_headers=None,
|
||||
kwargs={"litellm_logging_obj": Mock()},
|
||||
)
|
||||
|
||||
assert prepared.api_key == "generic-key"
|
||||
assert prepared.api_base == "https://generic.example.com"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue