[Azure OpenAI Feature] - Support DefaultAzureCredential without hard-coded environment variables (#12841)
* DefaultAzureCredential * update get_azure_ad_token_provider * fixes for get_azure_ad_token_provider * test_get_azure_ad_token_provider_with_default_azure_credential * test_get_azure_ad_token_fallback_to_default_azure_credential * docs DefaultAzureCredential * fix linting
This commit is contained in:
parent
7bdb5593bf
commit
133c26c015
@ -618,23 +618,43 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
|
||||
### Azure AD Token Refresh - `DefaultAzureCredential`
|
||||
|
||||
Use this if you want to use Azure `DefaultAzureCredential` for Authentication on your requests
|
||||
Use this if you want to use Azure `DefaultAzureCredential` for Authentication on your requests. `DefaultAzureCredential` automatically discovers and uses available Azure credentials from multiple sources.
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
**Option 1: Explicit DefaultAzureCredential (Recommended)**
|
||||
```python
|
||||
from litellm import completion
|
||||
from azure.identity import DefaultAzureCredential, get_bearer_token_provider
|
||||
|
||||
# DefaultAzureCredential automatically discovers credentials from:
|
||||
# - Environment variables (AZURE_CLIENT_ID, AZURE_CLIENT_SECRET, AZURE_TENANT_ID)
|
||||
# - Managed Identity (AKS, Azure VMs, etc.)
|
||||
# - Azure CLI credentials
|
||||
# - And other Azure identity sources
|
||||
token_provider = get_bearer_token_provider(DefaultAzureCredential(), "https://cognitiveservices.azure.com/.default")
|
||||
|
||||
|
||||
response = completion(
|
||||
model = "azure/<your deployment name>", # model = azure/<your deployment name>
|
||||
api_base = "", # azure api base
|
||||
api_version = "", # azure api version
|
||||
azure_ad_token_provider=token_provider
|
||||
azure_ad_token_provider=token_provider,
|
||||
messages = [{"role": "user", "content": "good morning"}],
|
||||
)
|
||||
```
|
||||
|
||||
**Option 2: LiteLLM Auto-Fallback to DefaultAzureCredential**
|
||||
```python
|
||||
import litellm
|
||||
|
||||
# Enable automatic fallback to DefaultAzureCredential
|
||||
litellm.enable_azure_ad_token_refresh = True
|
||||
|
||||
response = litellm.completion(
|
||||
model = "azure/<your deployment name>",
|
||||
api_base = "",
|
||||
api_version = "",
|
||||
messages = [{"role": "user", "content": "good morning"}],
|
||||
)
|
||||
```
|
||||
@ -642,6 +662,8 @@ response = completion(
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="PROXY config.yaml">
|
||||
|
||||
**Scenario 1: With Environment Variables (Traditional)**
|
||||
|
||||
1. Add relevant env vars
|
||||
|
||||
```bash
|
||||
@ -663,12 +685,48 @@ litellm_settings:
|
||||
enable_azure_ad_token_refresh: true # 👈 KEY CHANGE
|
||||
```
|
||||
|
||||
**Scenario 2: Managed Identity (AKS, Azure VMs) - No Hard-coded Credentials Required**
|
||||
|
||||
Perfect for AKS clusters, Azure VMs, or other managed environments where Azure automatically injects credentials.
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gpt-3.5-turbo
|
||||
litellm_params:
|
||||
model: azure/your-deployment-name
|
||||
api_base: https://openai-gpt-4-test-v-1.openai.azure.com/
|
||||
|
||||
litellm_settings:
|
||||
enable_azure_ad_token_refresh: true # 👈 KEY CHANGE
|
||||
```
|
||||
|
||||
**Scenario 3: Azure CLI Authentication**
|
||||
|
||||
If you're authenticated via `az login`, no additional configuration needed:
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: gpt-3.5-turbo
|
||||
litellm_params:
|
||||
model: azure/your-deployment-name
|
||||
api_base: https://openai-gpt-4-test-v-1.openai.azure.com/
|
||||
|
||||
litellm_settings:
|
||||
enable_azure_ad_token_refresh: true # 👈 KEY CHANGE
|
||||
```
|
||||
|
||||
3. Start proxy
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
**How it works**:
|
||||
- LiteLLM first tries Service Principal authentication (if environment variables are available)
|
||||
- If that fails, it automatically falls back to `DefaultAzureCredential`
|
||||
- `DefaultAzureCredential` will use Managed Identity, Azure CLI credentials, or other available Azure identity sources
|
||||
- This eliminates the need for hard-coded credentials in managed environments like AKS
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
|
||||
@ -278,6 +278,7 @@ def get_azure_ad_token(
|
||||
3. From username and password
|
||||
4. From OIDC token
|
||||
5. From a service principal with secret workflow
|
||||
6. From DefaultAzureCredential
|
||||
|
||||
Args:
|
||||
litellm_params: Dictionary containing authentication parameters
|
||||
@ -352,18 +353,27 @@ def get_azure_ad_token(
|
||||
azure_tenant_id=tenant_id,
|
||||
scope=scope,
|
||||
)
|
||||
# Try to get token provider from service principal
|
||||
# Try to get token provider from service principal or DefaultAzureCredential
|
||||
elif (
|
||||
azure_ad_token_provider is None
|
||||
and litellm.enable_azure_ad_token_refresh is True
|
||||
):
|
||||
verbose_logger.debug(
|
||||
"Using Azure AD token provider based on Service Principal with Secret workflow for Azure Auth"
|
||||
"Using Azure AD token provider based on Service Principal with Secret workflow or DefaultAzureCredential for Azure Auth"
|
||||
)
|
||||
try:
|
||||
azure_ad_token_provider = get_azure_ad_token_provider(azure_scope=scope)
|
||||
except ValueError:
|
||||
verbose_logger.debug("Azure AD Token Provider could not be used.")
|
||||
|
||||
#########################################################
|
||||
# If litellm.enable_azure_ad_token_refresh is True and no other token provider is available,
|
||||
# try to get DefaultAzureCredential provider
|
||||
#########################################################
|
||||
if azure_ad_token_provider is None and azure_ad_token is None:
|
||||
azure_ad_token_provider = BaseAzureLLM._try_get_default_azure_credential_provider(
|
||||
scope=scope,
|
||||
)
|
||||
|
||||
# Execute the token provider to get the token if available
|
||||
if azure_ad_token_provider and callable(azure_ad_token_provider):
|
||||
@ -387,6 +397,38 @@ def get_azure_ad_token(
|
||||
|
||||
|
||||
class BaseAzureLLM(BaseOpenAILLM):
|
||||
@staticmethod
|
||||
def _try_get_default_azure_credential_provider(
|
||||
scope: str,
|
||||
) -> Optional[Callable[[], str]]:
|
||||
"""
|
||||
Try to get DefaultAzureCredential provider
|
||||
|
||||
Args:
|
||||
scope: Azure scope for the token
|
||||
|
||||
Returns:
|
||||
Token provider callable if DefaultAzureCredential is enabled and available, None otherwise
|
||||
"""
|
||||
from litellm.types.secret_managers.get_azure_ad_token_provider import (
|
||||
AzureCredentialType,
|
||||
)
|
||||
|
||||
verbose_logger.debug(
|
||||
"Attempting to use DefaultAzureCredential for Azure Auth"
|
||||
)
|
||||
|
||||
try:
|
||||
azure_ad_token_provider = get_azure_ad_token_provider(
|
||||
azure_scope=scope,
|
||||
azure_credential=AzureCredentialType.DefaultAzureCredential,
|
||||
)
|
||||
verbose_logger.debug("Successfully obtained Azure AD token provider using DefaultAzureCredential")
|
||||
return azure_ad_token_provider
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"DefaultAzureCredential failed: {str(e)}")
|
||||
return None
|
||||
|
||||
def get_azure_openai_client(
|
||||
self,
|
||||
api_key: Optional[str],
|
||||
|
||||
@ -6,7 +6,10 @@ from litellm.types.secret_managers.get_azure_ad_token_provider import (
|
||||
)
|
||||
|
||||
|
||||
def get_azure_ad_token_provider(azure_scope: Optional[str] = None) -> Callable[[], str]:
|
||||
def get_azure_ad_token_provider(
|
||||
azure_scope: Optional[str] = None,
|
||||
azure_credential: Optional[AzureCredentialType] = None,
|
||||
) -> Callable[[], str]:
|
||||
"""
|
||||
Get Azure AD token provider based on Service Principal with Secret workflow.
|
||||
|
||||
@ -27,6 +30,7 @@ def get_azure_ad_token_provider(azure_scope: Optional[str] = None) -> Callable[[
|
||||
from azure.identity import (
|
||||
CertificateCredential,
|
||||
ClientSecretCredential,
|
||||
DefaultAzureCredential,
|
||||
ManagedIdentityCredential,
|
||||
get_bearer_token_provider,
|
||||
)
|
||||
@ -37,14 +41,17 @@ def get_azure_ad_token_provider(azure_scope: Optional[str] = None) -> Callable[[
|
||||
or "https://cognitiveservices.azure.com/.default"
|
||||
)
|
||||
|
||||
cred: str = os.environ.get(
|
||||
"AZURE_CREDENTIAL", AzureCredentialType.ClientSecretCredential
|
||||
cred: str = (
|
||||
azure_credential.value if azure_credential else None
|
||||
or os.environ.get("AZURE_CREDENTIAL", AzureCredentialType.ClientSecretCredential)
|
||||
or AzureCredentialType.ClientSecretCredential
|
||||
)
|
||||
credential: Optional[
|
||||
Union[
|
||||
ClientSecretCredential,
|
||||
ManagedIdentityCredential,
|
||||
CertificateCredential,
|
||||
DefaultAzureCredential,
|
||||
Any,
|
||||
]
|
||||
] = None
|
||||
@ -62,10 +69,15 @@ def get_azure_ad_token_provider(azure_scope: Optional[str] = None) -> Callable[[
|
||||
tenant_id=os.environ["AZURE_TENANT_ID"],
|
||||
certificate_path=os.environ["AZURE_CERTIFICATE_PATH"],
|
||||
)
|
||||
elif cred == AzureCredentialType.DefaultAzureCredential:
|
||||
# DefaultAzureCredential doesn't require explicit environment variables
|
||||
# It automatically discovers credentials from the environment (managed identity, CLI, etc.)
|
||||
credential = DefaultAzureCredential()
|
||||
else:
|
||||
cred_cls = getattr(identity, cred)
|
||||
credential = cred_cls()
|
||||
|
||||
if credential is None:
|
||||
raise ValueError("No credential provided")
|
||||
|
||||
return get_bearer_token_provider(credential, azure_scope)
|
||||
|
||||
@ -5,3 +5,4 @@ class AzureCredentialType(str, Enum):
|
||||
ClientSecretCredential = "ClientSecretCredential"
|
||||
ManagedIdentityCredential = "ManagedIdentityCredential"
|
||||
CertificateCredential = "CertificateCredential"
|
||||
DefaultAzureCredential = "DefaultAzureCredential"
|
||||
|
||||
@ -12,7 +12,13 @@ sys.path.insert(
|
||||
) # Adds the parent directory to the system path
|
||||
import litellm
|
||||
from litellm.llms.azure.common_utils import BaseAzureLLM, get_azure_ad_token
|
||||
from litellm.secret_managers.get_azure_ad_token_provider import (
|
||||
get_azure_ad_token_provider,
|
||||
)
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.secret_managers.get_azure_ad_token_provider import (
|
||||
AzureCredentialType,
|
||||
)
|
||||
from litellm.types.utils import CallTypes
|
||||
|
||||
|
||||
@ -1298,7 +1304,7 @@ def test_get_azure_ad_token_with_token_refresh(setup_mocks, monkeypatch):
|
||||
|
||||
# Verify the debug message was logged
|
||||
setup_mocks["logger"].debug.assert_any_call(
|
||||
"Using Azure AD token provider based on Service Principal with Secret workflow for Azure Auth"
|
||||
"Using Azure AD token provider based on Service Principal with Secret workflow or DefaultAzureCredential for Azure Auth"
|
||||
)
|
||||
|
||||
# Verify get_azure_ad_token_provider was called
|
||||
@ -1325,7 +1331,7 @@ def test_get_azure_ad_token_with_token_refresh_error(setup_mocks):
|
||||
|
||||
# Verify the debug message was logged
|
||||
setup_mocks["logger"].debug.assert_any_call(
|
||||
"Using Azure AD token provider based on Service Principal with Secret workflow for Azure Auth"
|
||||
"Using Azure AD token provider based on Service Principal with Secret workflow or DefaultAzureCredential for Azure Auth"
|
||||
)
|
||||
|
||||
# Verify error was logged
|
||||
@ -1333,8 +1339,8 @@ def test_get_azure_ad_token_with_token_refresh_error(setup_mocks):
|
||||
"Azure AD Token Provider could not be used."
|
||||
)
|
||||
|
||||
# Verify get_azure_ad_token_provider was called
|
||||
setup_mocks["token_provider"].assert_called_once()
|
||||
# Verify get_azure_ad_token_provider was called twice (once for service principal, once for DefaultAzureCredential)
|
||||
assert setup_mocks["token_provider"].call_count == 2
|
||||
|
||||
# Verify the token is None since the provider raised an error
|
||||
assert token is None
|
||||
@ -1380,3 +1386,101 @@ def test_token_provider_raises_exception(setup_mocks):
|
||||
|
||||
# Verify the error was logged
|
||||
setup_mocks["logger"].error.assert_called()
|
||||
|
||||
|
||||
def test_get_azure_ad_token_provider_with_default_azure_credential():
|
||||
"""
|
||||
Test that get_azure_ad_token_provider correctly uses DefaultAzureCredential
|
||||
when explicitly specified as the credential type. This verifies that the function
|
||||
can dynamically instantiate DefaultAzureCredential and return a working token provider.
|
||||
"""
|
||||
# Mock Azure identity classes
|
||||
with patch('azure.identity.DefaultAzureCredential') as mock_default_cred, \
|
||||
patch('azure.identity.get_bearer_token_provider') as mock_token_provider:
|
||||
|
||||
# Configure mocks
|
||||
mock_credential_instance = MagicMock()
|
||||
mock_default_cred.return_value = mock_credential_instance
|
||||
mock_token_provider.return_value = lambda: "test-default-azure-token"
|
||||
|
||||
# Test with DefaultAzureCredential specified explicitly
|
||||
token_provider = get_azure_ad_token_provider(
|
||||
azure_scope="https://cognitiveservices.azure.com/.default",
|
||||
azure_credential=AzureCredentialType.DefaultAzureCredential
|
||||
)
|
||||
|
||||
# Verify DefaultAzureCredential was instantiated
|
||||
mock_default_cred.assert_called_once_with()
|
||||
|
||||
# Verify get_bearer_token_provider was called with the right parameters
|
||||
mock_token_provider.assert_called_once_with(
|
||||
mock_credential_instance,
|
||||
"https://cognitiveservices.azure.com/.default"
|
||||
)
|
||||
|
||||
# Verify the returned token provider works
|
||||
token = token_provider()
|
||||
assert token == "test-default-azure-token"
|
||||
|
||||
|
||||
def test_get_azure_ad_token_fallback_to_default_azure_credential(setup_mocks, monkeypatch):
|
||||
"""
|
||||
Test that get_azure_ad_token falls back to DefaultAzureCredential when the
|
||||
service principal method fails but token refresh is enabled. This tests the
|
||||
complete fallback flow from service principal to DefaultAzureCredential.
|
||||
"""
|
||||
# Clear environment variables that might interfere
|
||||
monkeypatch.delenv("AZURE_USERNAME", raising=False)
|
||||
monkeypatch.delenv("AZURE_PASSWORD", raising=False)
|
||||
monkeypatch.delenv("AZURE_CLIENT_SECRET", raising=False)
|
||||
monkeypatch.delenv("AZURE_CLIENT_ID", raising=False)
|
||||
monkeypatch.delenv("AZURE_TENANT_ID", raising=False)
|
||||
|
||||
# Reset mocks to ensure clean state
|
||||
setup_mocks["token_provider"].reset_mock()
|
||||
|
||||
# Enable token refresh
|
||||
setup_mocks["litellm"].enable_azure_ad_token_refresh = True
|
||||
|
||||
# Configure get_azure_ad_token_provider to fail first (service principal)
|
||||
# but succeed on second call (DefaultAzureCredential)
|
||||
def mock_token_provider_side_effect(*args, **kwargs):
|
||||
# If called with azure_credential=DefaultAzureCredential, return a working provider
|
||||
if kwargs.get("azure_credential") == AzureCredentialType.DefaultAzureCredential:
|
||||
return lambda: "mock-default-azure-credential-token"
|
||||
# Otherwise (service principal call), return None to simulate failure
|
||||
return None
|
||||
|
||||
setup_mocks["token_provider"].side_effect = mock_token_provider_side_effect
|
||||
|
||||
# Create test parameters with no other auth methods available
|
||||
litellm_params = GenericLiteLLMParams()
|
||||
|
||||
# Call the function
|
||||
token = get_azure_ad_token(litellm_params)
|
||||
|
||||
# Verify the success debug message was logged
|
||||
setup_mocks["logger"].debug.assert_any_call(
|
||||
"Successfully obtained Azure AD token provider using DefaultAzureCredential"
|
||||
)
|
||||
|
||||
# Verify get_azure_ad_token_provider was called twice:
|
||||
# 1. First with just azure_scope (service principal attempt)
|
||||
# 2. Second with azure_credential=DefaultAzureCredential (fallback)
|
||||
assert setup_mocks["token_provider"].call_count == 2
|
||||
|
||||
# Verify the calls were made with expected parameters
|
||||
calls = setup_mocks["token_provider"].call_args_list
|
||||
|
||||
# First call should be service principal attempt (no azure_credential)
|
||||
first_call_kwargs = calls[0][1]
|
||||
assert "azure_scope" in first_call_kwargs
|
||||
assert first_call_kwargs.get("azure_credential") is None
|
||||
|
||||
# Second call should be DefaultAzureCredential attempt
|
||||
second_call_kwargs = calls[1][1]
|
||||
assert "azure_scope" in second_call_kwargs
|
||||
assert second_call_kwargs.get("azure_credential") == AzureCredentialType.DefaultAzureCredential
|
||||
|
||||
# Verify the token is what we expect from our DefaultAzureCredential mock
|
||||
assert token == "mock-default-azure-credential-token"
|
||||
|
||||
Loading…
Reference in New Issue
Block a user