diff --git a/litellm/__init__.py b/litellm/__init__.py index 64c60ca3374..850885b700b 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -167,12 +167,12 @@ prometheus_initialize_budget_metrics: Optional[bool] = False require_auth_for_metrics_endpoint: Optional[bool] = False argilla_batch_size: Optional[int] = None datadog_use_v1: Optional[bool] = False # if you want to use v1 datadog logged payload. -gcs_pub_sub_use_v1: Optional[ - bool -] = False # if you want to use v1 gcs pubsub logged payload -generic_api_use_v1: Optional[ - bool -] = False # if you want to use v1 generic api logged payload +gcs_pub_sub_use_v1: Optional[bool] = ( + False # if you want to use v1 gcs pubsub logged payload +) +generic_api_use_v1: Optional[bool] = ( + False # if you want to use v1 generic api logged payload +) argilla_transformation_object: Optional[Dict[str, Any]] = None _async_input_callback: List[ Union[str, Callable, "CustomLogger"] @@ -192,25 +192,25 @@ _async_failure_callback: List[ pre_call_rules: List[Callable] = [] post_call_rules: List[Callable] = [] turn_off_message_logging: Optional[bool] = False -standard_logging_payload_excluded_fields: Optional[ - List[str] -] = None # Fields to exclude from StandardLoggingPayload before callbacks receive it +standard_logging_payload_excluded_fields: Optional[List[str]] = ( + None # Fields to exclude from StandardLoggingPayload before callbacks receive it +) log_raw_request_response: bool = False redact_messages_in_exceptions: Optional[bool] = False redact_user_api_key_info: Optional[bool] = False filter_invalid_headers: Optional[bool] = False -add_user_information_to_llm_headers: Optional[ - bool -] = None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers +add_user_information_to_llm_headers: Optional[bool] = ( + None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers +) store_audit_logs = False # Enterprise feature, allow users to see audit logs ### end of callbacks ############# -email: Optional[ - str -] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -token: Optional[ - str -] = None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +email: Optional[str] = ( + None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) +token: Optional[str] = ( + None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) telemetry = True max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults drop_params = bool(os.getenv("LITELLM_DROP_PARAMS", False)) @@ -272,9 +272,9 @@ use_client: bool = False ssl_verify: Union[str, bool] = True ssl_security_level: Optional[str] = None ssl_certificate: Optional[str] = None -ssl_ecdh_curve: Optional[ - str -] = None # Set to 'X25519' to disable PQC and improve performance +ssl_ecdh_curve: Optional[str] = ( + None # Set to 'X25519' to disable PQC and improve performance +) disable_streaming_logging: bool = False disable_token_counter: bool = False disable_add_transform_inline_image_block: bool = False @@ -328,20 +328,24 @@ enable_loadbalancing_on_batch_endpoints: Optional[bool] = None enable_caching_on_provider_specific_optional_params: bool = ( False # feature-flag for caching on optional params - e.g. 'top_k' ) -caching: bool = False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -caching_with_models: bool = False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 -cache: Optional[ - "Cache" -] = None # cache object <- use this - https://docs.litellm.ai/docs/caching +caching: bool = ( + False # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) +caching_with_models: bool = ( + False # # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648 +) +cache: Optional["Cache"] = ( + None # cache object <- use this - https://docs.litellm.ai/docs/caching +) default_in_memory_ttl: Optional[float] = None default_redis_ttl: Optional[float] = None default_redis_batch_cache_expiry: Optional[float] = None model_alias_map: Dict[str, str] = {} model_group_settings: Optional["ModelGroupSettings"] = None max_budget: float = 0.0 # set the max budget across all providers -budget_duration: Optional[ - str -] = None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). +budget_duration: Optional[str] = ( + None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). +) default_soft_budget: float = ( DEFAULT_SOFT_BUDGET # by default all litellm proxy keys have a soft budget of 50.0 ) @@ -350,7 +354,9 @@ forward_traceparent_to_llm_provider: bool = False _current_cost = 0.0 # private variable, used if max budget is set error_logs: Dict = {} -add_function_to_prompt: bool = False # if function calling not supported by api, append function call details to system prompt +add_function_to_prompt: bool = ( + False # if function calling not supported by api, append function call details to system prompt +) client_session: Optional[httpx.Client] = None aclient_session: Optional[httpx.AsyncClient] = None model_fallbacks: Optional[List] = None # Deprecated for 'litellm.fallbacks' @@ -397,7 +403,9 @@ prometheus_emit_stream_label: bool = False disable_add_prefix_to_prompt: bool = ( False # used by anthropic, to disable adding prefix to prompt ) -disable_copilot_system_to_assistant: bool = False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. +disable_copilot_system_to_assistant: bool = ( + False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. +) public_mcp_servers: Optional[List[str]] = None public_model_groups: Optional[List[str]] = None public_agent_groups: Optional[List[str]] = None @@ -406,9 +414,9 @@ public_agent_groups: Optional[List[str]] = None # Old format: { "displayName": "url" } (for backward compatibility) public_model_groups_links: Dict[str, Union[str, Dict[str, Any]]] = {} #### REQUEST PRIORITIZATION ####### -priority_reservation: Optional[ - Dict[str, Union[float, "PriorityReservationDict"]] -] = None +priority_reservation: Optional[Dict[str, Union[float, "PriorityReservationDict"]]] = ( + None +) # priority_reservation_settings is lazy-loaded via __getattr__ # Only declare for type checking - at runtime __getattr__ handles it if TYPE_CHECKING: @@ -416,13 +424,17 @@ if TYPE_CHECKING: ######## Networking Settings ######## -use_aiohttp_transport: bool = True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead. +use_aiohttp_transport: bool = ( + True # Older variable, aiohttp is now the default. use disable_aiohttp_transport instead. +) aiohttp_trust_env: bool = False # set to true to use HTTP_ Proxy settings disable_aiohttp_transport: bool = False # Set this to true to use httpx instead disable_aiohttp_trust_env: bool = ( False # When False, aiohttp will respect HTTP(S)_PROXY env vars ) -force_ipv4: bool = False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6. +force_ipv4: bool = ( + False # when True, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6. +) network_mock: bool = False # When True, use mock transport — no real network calls ####### STOP SEQUENCE LIMIT ####### @@ -437,13 +449,13 @@ context_window_fallbacks: Optional[List] = None content_policy_fallbacks: Optional[List] = None allowed_fails: int = 3 allow_dynamic_callback_disabling: bool = True -num_retries_per_request: Optional[ - int -] = None # for the request overall (incl. fallbacks + model retries) +num_retries_per_request: Optional[int] = ( + None # for the request overall (incl. fallbacks + model retries) +) ####### SECRET MANAGERS ##################### -secret_manager_client: Optional[ - Any -] = None # list of instantiated key management clients - e.g. azure kv, infisical, etc. +secret_manager_client: Optional[Any] = ( + None # list of instantiated key management clients - e.g. azure kv, infisical, etc. +) _google_kms_resource_name: Optional[str] = None _key_management_system: Optional["KeyManagementSystem"] = None # Note: KeyManagementSettings must be eagerly imported because _key_management_settings @@ -456,12 +468,12 @@ output_parse_pii: bool = False from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map model_cost = get_model_cost_map(url=model_cost_map_url) -cost_discount_config: Dict[ - str, float -] = {} # Provider-specific cost discounts {"vertex_ai": 0.05} = 5% discount -cost_margin_config: Dict[ - str, Union[float, Dict[str, float]] -] = {} # Provider-specific or global cost margins. Examples: +cost_discount_config: Dict[str, float] = ( + {} +) # Provider-specific cost discounts {"vertex_ai": 0.05} = 5% discount +cost_margin_config: Dict[str, Union[float, Dict[str, float]]] = ( + {} +) # Provider-specific or global cost margins. Examples: # Percentage: {"openai": 0.10} = 10% margin # Fixed: {"openai": {"fixed_amount": 0.001}} = $0.001 per request # Global: {"global": 0.05} = 5% global margin on all providers @@ -1310,12 +1322,12 @@ from . import rag from .types.llms.custom_llm import CustomLLMItem custom_provider_map: List[CustomLLMItem] = [] -_custom_providers: List[ - str -] = [] # internal helper util, used to track names of custom providers -disable_hf_tokenizer_download: Optional[ - bool -] = None # disable huggingface tokenizer download. Defaults to openai clk100 +_custom_providers: List[str] = ( + [] +) # internal helper util, used to track names of custom providers +disable_hf_tokenizer_download: Optional[bool] = ( + None # disable huggingface tokenizer download. Defaults to openai clk100 +) global_disable_no_log_param: bool = False ### CLI UTILITIES ### @@ -1809,6 +1821,9 @@ if TYPE_CHECKING: from .llms.hosted_vllm.responses.transformation import ( HostedVLLMResponsesAPIConfig as HostedVLLMResponsesAPIConfig, ) + from .llms.azure_ai.responses.transformation import ( + AzureAIResponsesAPIConfig as AzureAIResponsesAPIConfig, + ) from .llms.github_copilot.chat.transformation import ( GithubCopilotConfig as GithubCopilotConfig, ) diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index 9164a3c8ae4..d3a11d279b7 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -230,6 +230,7 @@ LLM_CONFIG_NAMES = ( "XAIResponsesAPIConfig", "LiteLLMProxyResponsesAPIConfig", "HostedVLLMResponsesAPIConfig", + "AzureAIResponsesAPIConfig", "VolcEngineResponsesAPIConfig", "PerplexityResponsesConfig", "DatabricksResponsesAPIConfig", @@ -921,6 +922,10 @@ _LLM_CONFIGS_IMPORT_MAP = { ".llms.hosted_vllm.responses.transformation", "HostedVLLMResponsesAPIConfig", ), + "AzureAIResponsesAPIConfig": ( + ".llms.azure_ai.responses.transformation", + "AzureAIResponsesAPIConfig", + ), "VolcEngineResponsesAPIConfig": ( ".llms.volcengine.responses.transformation", "VolcEngineResponsesAPIConfig", diff --git a/litellm/llms/azure_ai/responses/__init__.py b/litellm/llms/azure_ai/responses/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/azure_ai/responses/transformation.py b/litellm/llms/azure_ai/responses/transformation.py new file mode 100644 index 00000000000..d66a7a9d7af --- /dev/null +++ b/litellm/llms/azure_ai/responses/transformation.py @@ -0,0 +1,133 @@ +""" +Responses API transformation for Azure AI provider. + +When api_base includes /projects/, route to the real upstream /responses +endpoint instead of falling back to the completions-style bridge. + +Ref: https://learn.microsoft.com/en-us/azure/foundry/foundry-models/how-to/generate-responses +""" + +from typing import Optional +from urllib.parse import urlparse + +import httpx + +import litellm +from litellm.llms.azure.common_utils import BaseAzureLLM +from litellm.llms.azure_ai.common_utils import AzureFoundryModelInfo +from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders +from litellm.utils import _add_path_to_api_base + + +class AzureAIResponsesAPIConfig(OpenAIResponsesAPIConfig): + """ + Configuration for Azure AI Foundry Responses API. + + Extends OpenAI's responses config because Azure AI Foundry project-based + endpoints follow the OpenAI /responses spec. Uses Azure-specific auth + (api-key header for *.services.ai.azure.com hosts) and constructs the + correct URL path for project-based endpoints. + """ + + @property + def custom_llm_provider(self) -> LlmProviders: + return LlmProviders.AZURE_AI + + def validate_environment( + self, + headers: dict, + model: str, + litellm_params: Optional[GenericLiteLLMParams], + ) -> dict: + litellm_params = litellm_params or GenericLiteLLMParams() + api_key = AzureFoundryModelInfo.get_api_key( + api_key=litellm_params.api_key, + ) + api_base = AzureFoundryModelInfo.get_api_base( + api_base=litellm_params.api_base, + ) + + if api_key: + if api_base and self._should_use_api_key_header(api_base): + headers["api-key"] = api_key + else: + headers["Authorization"] = f"Bearer {api_key}" + else: + # Fall back to Azure AD token-based auth + headers = BaseAzureLLM._base_validate_azure_environment( + headers=headers, litellm_params=litellm_params + ) + + headers["Content-Type"] = "application/json" + return headers + + @staticmethod + def _should_use_api_key_header(api_base: str) -> bool: + """ + Returns True if the request should use the ``api-key`` header. + + Azure AI Foundry endpoints under *.services.ai.azure.com and + *.openai.azure.com expect the ``api-key`` header instead of + ``Authorization: Bearer ...``. + """ + parsed_url = urlparse(api_base) + host = parsed_url.hostname + if host and ( + host.endswith(".services.ai.azure.com") + or host.endswith(".openai.azure.com") + ): + return True + return False + + def get_complete_url( + self, + api_base: Optional[str], + litellm_params: dict, + ) -> str: + """ + Build the full URL for the Azure AI Foundry Responses API. + + For project-based endpoints (api_base contains ``/projects/``), + appends ``/openai/v1/responses`` to the base. + + Example: + api_base = "https://.services.ai.azure.com/api/projects/" + -> "https://.services.ai.azure.com/api/projects//openai/v1/responses" + """ + api_base = AzureFoundryModelInfo.get_api_base(api_base=api_base) + + if api_base is None: + raise ValueError( + "api_base is required for Azure AI Responses API. " + "Set via api_base parameter or AZURE_AI_API_BASE environment variable." + ) + + # Extract api_version + api_version = litellm_params.get("api_version") + + # Parse query params from existing URL + original_url = httpx.URL(api_base) + query_params = dict(original_url.params) + + if "api-version" not in query_params and api_version: + query_params["api-version"] = api_version + + # Build the responses endpoint path + if "/projects/" in api_base: + new_url = _add_path_to_api_base( + api_base=api_base, ending_path="/openai/v1/responses" + ) + elif "services.ai.azure.com" in api_base: + new_url = _add_path_to_api_base( + api_base=api_base, ending_path="/models/responses" + ) + else: + new_url = _add_path_to_api_base( + api_base=api_base, ending_path="/v1/responses" + ) + + final_url = httpx.URL(new_url).copy_with(params=query_params) + return str(final_url) diff --git a/litellm/utils.py b/litellm/utils.py index f902644e760..590e1d694ff 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -783,9 +783,9 @@ def function_setup( # noqa: PLR0915 coroutine_checker = get_coroutine_checker_fn() ## DYNAMIC CALLBACKS ## - dynamic_callbacks: Optional[ - List[Union[str, Callable, "CustomLogger"]] - ] = kwargs.pop("callbacks", None) + dynamic_callbacks: Optional[List[Union[str, Callable, "CustomLogger"]]] = ( + kwargs.pop("callbacks", None) + ) all_callbacks = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks) if len(all_callbacks) > 0: @@ -1691,9 +1691,9 @@ def client(original_function): # noqa: PLR0915 exception=e, retry_policy=kwargs.get("retry_policy"), ) - kwargs[ - "retry_policy" - ] = reset_retry_policy() # prevent infinite loops + kwargs["retry_policy"] = ( + reset_retry_policy() + ) # prevent infinite loops litellm.num_retries = ( None # set retries to None to prevent infinite loops ) @@ -1740,9 +1740,9 @@ def client(original_function): # noqa: PLR0915 exception=e, retry_policy=kwargs.get("retry_policy"), ) - kwargs[ - "retry_policy" - ] = reset_retry_policy() # prevent infinite loops + kwargs["retry_policy"] = ( + reset_retry_policy() + ) # prevent infinite loops litellm.num_retries = ( None # set retries to None to prevent infinite loops ) @@ -3771,10 +3771,10 @@ def pre_process_non_default_params( if "response_format" in non_default_params: if provider_config is not None: - non_default_params[ - "response_format" - ] = provider_config.get_json_schema_from_pydantic_object( - response_format=non_default_params["response_format"] + non_default_params["response_format"] = ( + provider_config.get_json_schema_from_pydantic_object( + response_format=non_default_params["response_format"] + ) ) else: non_default_params["response_format"] = type_to_response_format_param( @@ -3903,16 +3903,16 @@ def pre_process_optional_params( True # so that main.py adds the function call to the prompt ) if "tools" in non_default_params: - optional_params[ - "functions_unsupported_model" - ] = non_default_params.pop("tools") + optional_params["functions_unsupported_model"] = ( + non_default_params.pop("tools") + ) non_default_params.pop( "tool_choice", None ) # causes ollama requests to hang elif "functions" in non_default_params: - optional_params[ - "functions_unsupported_model" - ] = non_default_params.pop("functions") + optional_params["functions_unsupported_model"] = ( + non_default_params.pop("functions") + ) elif ( litellm.add_function_to_prompt ): # if user opts to add it to prompt instead @@ -4893,9 +4893,7 @@ def _get_order_filtered_deployments( ) -> List: if target_order is not None: filtered = [ - d - for d in healthy_deployments - if _get_deployment_order(d) == target_order + d for d in healthy_deployments if _get_deployment_order(d) == target_order ] if filtered: return filtered @@ -7549,9 +7547,9 @@ class ModelResponseIterator: if convert_to_delta is True: _stream_response = ModelResponseStream() _stream_response.choices[0].delta.content = model_response.choices[0].message.content # type: ignore - self.model_response: Union[ - ModelResponse, ModelResponseStream - ] = _stream_response + self.model_response: Union[ModelResponse, ModelResponseStream] = ( + _stream_response + ) else: self.model_response = model_response self.is_done = False @@ -8545,6 +8543,8 @@ class ProviderConfigManager: return None elif litellm.LlmProviders.OPENROUTER == provider: return litellm.OpenRouterResponsesAPIConfig() + elif litellm.LlmProviders.AZURE_AI == provider: + return litellm.AzureAIResponsesAPIConfig() elif litellm.LlmProviders.HOSTED_VLLM == provider: return litellm.HostedVLLMResponsesAPIConfig() return None diff --git a/tests/test_litellm/llms/azure_ai/responses/test_azure_ai_responses.py b/tests/test_litellm/llms/azure_ai/responses/test_azure_ai_responses.py new file mode 100644 index 00000000000..312b35e1815 --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/responses/test_azure_ai_responses.py @@ -0,0 +1,218 @@ +""" +Tests for azure_ai Responses API support. + +Verifies that when api_base contains /projects/, litellm.responses() +routes to the real upstream /responses endpoint instead of falling +back to the completions-style bridge. + +Ref: https://github.com/BerriAI/litellm/issues/25407 +""" + +import json +import os +import sys +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +import litellm +from litellm.llms.azure_ai.responses.transformation import ( + AzureAIResponsesAPIConfig, +) +from litellm.types.router import GenericLiteLLMParams +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager + + +def _make_mock_responses_api_response(content: str = "Hello!") -> dict: + return { + "id": "resp-test-azure-ai", + "object": "response", + "created_at": 1234567890, + "model": "DeepSeek-R1", + "output": [ + { + "type": "message", + "id": "msg-test-azure-ai", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": content, + "annotations": [], + } + ], + } + ], + "status": "completed", + "usage": { + "input_tokens": 10, + "output_tokens": 20, + "total_tokens": 30, + }, + } + + +def _make_mock_http_client(response_body: dict) -> MagicMock: + mock_client = MagicMock() + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.headers = {"content-type": "application/json"} + mock_response.json.return_value = response_body + mock_response.text = json.dumps(response_body) + mock_client.post.return_value = mock_response + return mock_client + + +class TestAzureAIResponsesProviderConfig: + """Test that ProviderConfigManager returns AzureAIResponsesAPIConfig for azure_ai.""" + + def test_provider_config_registration(self): + config = ProviderConfigManager.get_provider_responses_api_config( + model="DeepSeek-R1", + provider=LlmProviders.AZURE_AI, + ) + assert config is not None + assert isinstance(config, AzureAIResponsesAPIConfig) + assert config.custom_llm_provider == LlmProviders.AZURE_AI + + +class TestAzureAIResponsesURL: + """Test get_complete_url() constructs the correct URL for various api_base patterns.""" + + def test_project_based_url(self): + """When api_base contains /projects/, append /openai/v1/responses.""" + config = AzureAIResponsesAPIConfig() + url = config.get_complete_url( + api_base="https://myresource.services.ai.azure.com/api/projects/my-project", + litellm_params={}, + ) + assert url == ( + "https://myresource.services.ai.azure.com/api/projects/my-project" + "/openai/v1/responses" + ) + + def test_project_based_url_with_api_version(self): + """api-version query param should be appended.""" + config = AzureAIResponsesAPIConfig() + url = config.get_complete_url( + api_base="https://myresource.services.ai.azure.com/api/projects/my-project", + litellm_params={"api_version": "2025-03-01-preview"}, + ) + assert "api-version=2025-03-01-preview" in url + assert "/openai/v1/responses" in url + + def test_services_ai_azure_url_without_projects(self): + """Standard services.ai.azure.com base (no /projects/) uses /models/responses.""" + config = AzureAIResponsesAPIConfig() + url = config.get_complete_url( + api_base="https://myresource.services.ai.azure.com", + litellm_params={}, + ) + assert "/models/responses" in url + + def test_generic_api_base(self): + """Non-Azure-Foundry base uses /v1/responses.""" + config = AzureAIResponsesAPIConfig() + url = config.get_complete_url( + api_base="https://custom-proxy.example.com", + litellm_params={}, + ) + assert url == "https://custom-proxy.example.com/v1/responses" + + def test_missing_api_base_raises(self): + """ValueError raised when api_base is None and no env var set.""" + config = AzureAIResponsesAPIConfig() + with patch.dict(os.environ, {}, clear=False): + # Ensure env var is not set + os.environ.pop("AZURE_AI_API_BASE", None) + with pytest.raises(ValueError, match="api_base is required"): + config.get_complete_url(api_base=None, litellm_params={}) + + +class TestAzureAIResponsesAuth: + """Test validate_environment() sets correct auth headers.""" + + def test_api_key_header_for_azure_foundry_host(self): + """*.services.ai.azure.com should use api-key header.""" + config = AzureAIResponsesAPIConfig() + headers = config.validate_environment( + headers={}, + model="DeepSeek-R1", + litellm_params=GenericLiteLLMParams( + api_key="test-key-123", + api_base="https://myresource.services.ai.azure.com/api/projects/proj", + ), + ) + assert headers.get("api-key") == "test-key-123" + assert "Authorization" not in headers + + def test_bearer_auth_for_non_azure_host(self): + """Non-Azure hosts should use Bearer auth.""" + config = AzureAIResponsesAPIConfig() + headers = config.validate_environment( + headers={}, + model="DeepSeek-R1", + litellm_params=GenericLiteLLMParams( + api_key="test-key-456", + api_base="https://custom-proxy.example.com", + ), + ) + assert headers.get("Authorization") == "Bearer test-key-456" + assert "api-key" not in headers + + def test_api_key_header_for_openai_azure_host(self): + """*.openai.azure.com should also use api-key header.""" + config = AzureAIResponsesAPIConfig() + headers = config.validate_environment( + headers={}, + model="gpt-4", + litellm_params=GenericLiteLLMParams( + api_key="test-key-789", + api_base="https://myresource.openai.azure.com", + ), + ) + assert headers.get("api-key") == "test-key-789" + + +class TestAzureAIResponsesE2E: + """End-to-end test with mocked HTTP client.""" + + def test_responses_create_routes_to_native_endpoint(self): + """ + Verify litellm.responses() uses native /responses routing for azure_ai + with a project-based api_base, rather than the completions bridge. + """ + mock_client = _make_mock_http_client( + _make_mock_responses_api_response("Hello from Azure AI Foundry!") + ) + + with patch( + "litellm.llms.custom_httpx.llm_http_handler._get_httpx_client", + return_value=mock_client, + ): + response = litellm.responses( + model="azure_ai/DeepSeek-R1", + input="Hello, how are you?", + api_base="https://myresource.services.ai.azure.com/api/projects/my-project", + api_key="test-key", + ) + + from litellm.types.llms.openai import ResponsesAPIResponse + + assert response is not None + assert isinstance(response, ResponsesAPIResponse) + assert len(response.output) > 0 + output_message = response.output[0] + assert output_message.role == "assistant" + assert "Azure AI Foundry" in output_message.content[0].text + + # Verify the URL used contains /responses (native endpoint) + call_args = mock_client.post.call_args + called_url = call_args[1].get("url", call_args[0][0] if call_args[0] else "") + assert "/responses" in str(called_url)