From bb5397d9b2582e29c4b4f2ec613a25ef0c926ddf Mon Sep 17 00:00:00 2001 From: Emerson Gomes Date: Wed, 17 Dec 2025 18:13:41 -0600 Subject: [PATCH] fix: enforce scheme for Azure AI rerank api_base --- .../llms/azure_ai/rerank/transformation.py | 6 ++ .../test_azure_ai_rerank_transformation.py | 100 ++++++++++++++++++ 2 files changed, 106 insertions(+) create mode 100644 tests/test_litellm/llms/azure_ai/rerank/test_azure_ai_rerank_transformation.py diff --git a/litellm/llms/azure_ai/rerank/transformation.py b/litellm/llms/azure_ai/rerank/transformation.py index 376f460826..f577a42ed5 100644 --- a/litellm/llms/azure_ai/rerank/transformation.py +++ b/litellm/llms/azure_ai/rerank/transformation.py @@ -30,6 +30,12 @@ class AzureAIRerankConfig(CohereRerankConfig): "Azure AI API Base is required. api_base=None. Set in call or via `AZURE_AI_API_BASE` env var." ) original_url = httpx.URL(api_base) + if not original_url.is_absolute_url: + raise ValueError( + "Azure AI API Base must be an absolute URL including scheme (e.g. " + "'https://.services.ai.azure.com'). " + f"Got api_base={api_base!r}." + ) normalized_path = original_url.path.rstrip("/") # Allow callers to pass either full v1/v2 rerank endpoints: diff --git a/tests/test_litellm/llms/azure_ai/rerank/test_azure_ai_rerank_transformation.py b/tests/test_litellm/llms/azure_ai/rerank/test_azure_ai_rerank_transformation.py new file mode 100644 index 0000000000..1f42511343 --- /dev/null +++ b/tests/test_litellm/llms/azure_ai/rerank/test_azure_ai_rerank_transformation.py @@ -0,0 +1,100 @@ +import os +import sys + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +from litellm.llms.azure_ai.rerank.transformation import AzureAIRerankConfig + + +class TestAzureAIRerankConfigGetCompleteUrl: + def setup_method(self): + self.config = AzureAIRerankConfig() + self.model = "azure_ai/cohere-rerank-v3-english" + + def test_api_base_required(self): + with pytest.raises(ValueError) as exc_info: + self.config.get_complete_url(api_base=None, model=self.model) + + assert "api_base=None" in str(exc_info.value) + + @pytest.mark.parametrize( + "api_base", + [ + "example.com", + "example.com/v1", + "//example.com/v1", + "/v1/rerank", + ], + ) + def test_api_base_requires_scheme(self, api_base): + with pytest.raises(ValueError) as exc_info: + self.config.get_complete_url(api_base=api_base, model=self.model) + + error_message = str(exc_info.value).lower() + assert "absolute url" in error_message + assert "scheme" in error_message + + @pytest.mark.parametrize( + "api_base, expected_url", + [ + ( + "https://my-resource.services.ai.azure.com/v1/rerank/", + "https://my-resource.services.ai.azure.com/v1/rerank", + ), + ( + "https://my-resource.services.ai.azure.com/providers/cohere/v2/rerank/", + "https://my-resource.services.ai.azure.com/providers/cohere/v2/rerank", + ), + ], + ) + def test_preserves_full_rerank_endpoint(self, api_base, expected_url): + url = self.config.get_complete_url(api_base=api_base, model=self.model) + assert url == expected_url + + @pytest.mark.parametrize( + "api_base, expected_url", + [ + ( + "https://my-resource.services.ai.azure.com/v1", + "https://my-resource.services.ai.azure.com/v1/rerank", + ), + ( + "https://my-resource.services.ai.azure.com/v2/", + "https://my-resource.services.ai.azure.com/v2/rerank", + ), + ( + "https://my-resource.services.ai.azure.com/providers/cohere/v2", + "https://my-resource.services.ai.azure.com/providers/cohere/v2/rerank", + ), + ( + "https://my-resource.services.ai.azure.com/providers/cohere/v2/", + "https://my-resource.services.ai.azure.com/providers/cohere/v2/rerank", + ), + ], + ) + def test_appends_rerank_for_version_paths(self, api_base, expected_url): + url = self.config.get_complete_url(api_base=api_base, model=self.model) + assert url == expected_url + + @pytest.mark.parametrize( + "api_base", + [ + "https://my-resource.services.ai.azure.com", + "https://my-resource.services.ai.azure.com/", + ], + ) + def test_defaults_to_v1_rerank_when_base_has_no_path(self, api_base): + url = self.config.get_complete_url(api_base=api_base, model=self.model) + assert url == "https://my-resource.services.ai.azure.com/v1/rerank" + + def test_preserves_query_params(self): + url = self.config.get_complete_url( + api_base="https://my-resource.services.ai.azure.com/v1?r=1", + model=self.model, + ) + assert url == "https://my-resource.services.ai.azure.com/v1/rerank?r=1" +