Azure LLM: fix passing through of azure_ad_token_provider parameter (#10694)

* Azure LLM: fix passing through of azure_ad_token_provider parameter

* add test

---------

Co-authored-by: Clara Luise Pohland <clara-luise.pohland@telekom.de>
This commit is contained in:
Clara Pohland 2025-05-10 08:38:13 +02:00 • committed by GitHub
parent 6f32189093
commit 6084d6c437
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 27 additions and 3 deletions

View file

@ -320,7 +320,7 @@ class BaseAzureLLM(BaseOpenAILLM):
api_version: Optional[str],
is_async: bool,
) -> dict:
azure_ad_token_provider: Optional[Callable[[], str]] = None
azure_ad_token_provider = litellm_params.get("azure_ad_token_provider")
# If we have api_key, then we have higher priority
azure_ad_token = litellm_params.get("azure_ad_token")
tenant_id = litellm_params.get("tenant_id", os.getenv("AZURE_TENANT_ID"))
@ -336,7 +336,11 @@ class BaseAzureLLM(BaseOpenAILLM):
)
max_retries = litellm_params.get("max_retries")
timeout = litellm_params.get("timeout")
if not api_key and tenant_id and client_id and client_secret:
if (
not api_key
and azure_ad_token_provider is None
and tenant_id and client_id and client_secret
):
verbose_logger.debug(
"Using Azure AD Token Provider from Entra ID for Azure Auth"
)
@ -345,7 +349,7 @@ class BaseAzureLLM(BaseOpenAILLM):
client_id=client_id,
client_secret=client_secret,
)
if azure_username and azure_password and client_id:
if azure_ad_token_provider is None and azure_username and azure_password and client_id:
verbose_logger.debug("Using Azure Username and Password for Azure Auth")
azure_ad_token_provider = get_azure_ad_token_from_username_password(
azure_username=azure_username,

View file

@ -260,6 +260,26 @@ def test_initialize_with_oidc_token_no_credentials(setup_mocks, monkeypatch):
# Verify expected result
assert result["azure_ad_token"] == "mock-oidc-token"
def test_initialize_with_ad_token_provider(setup_mocks, monkeypatch):
# Clear environment variables
monkeypatch.delenv("AZURE_CLIENT_ID", raising=False)
monkeypatch.delenv("AZURE_TENANT_ID", raising=False)
# Test with custom azure_ad_token_provider
result = BaseAzureLLM().initialize_azure_sdk_client(
litellm_params={
"azure_ad_token_provider": lambda: "mock-custom-token",
},
api_key=None,
api_base="https://test.openai.azure.com",
model_name="gpt-4",
api_version=None,
is_async=False,
)
# Verify expected result
assert result["azure_ad_token_provider"]() == "mock-custom-token"
def test_initialize_with_enable_token_refresh(setup_mocks, monkeypatch):
litellm._turn_on_debug()