fix(ocr): preserve Azure Document Intelligence authentication

This commit is contained in:
Yujong Lee 2026-09-12 11:13:39 -07:00
parent 9132a5343b
commit 1b88828229
5 changed files with 134 additions and 19 deletions

View file

@ -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

View file

@ -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",

View file

@ -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,

View file

@ -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

View file

@ -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"