diff --git a/litellm/secret_managers/get_azure_ad_token_provider.py b/litellm/secret_managers/get_azure_ad_token_provider.py index 5403675b97..c982856b5e 100644 --- a/litellm/secret_managers/get_azure_ad_token_provider.py +++ b/litellm/secret_managers/get_azure_ad_token_provider.py @@ -1,5 +1,9 @@ import os -from typing import Callable +from typing import Any, Callable, Optional, Union + +from litellm.types.secret_managers.get_azure_ad_token_provider import ( + AzureCredentialType, +) def get_azure_ad_token_provider() -> Callable[[], str]: @@ -15,24 +19,45 @@ def get_azure_ad_token_provider() -> Callable[[], str]: Callable that returns a temporary authentication token. """ import azure.identity as identity - from azure.identity import get_bearer_token_provider + from azure.identity import ( + CertificateCredential, + ClientSecretCredential, + ManagedIdentityCredential, + get_bearer_token_provider, + ) azure_scope = os.environ.get( "AZURE_SCOPE", "https://cognitiveservices.azure.com/.default" ) - cred = os.environ.get("AZURE_CREDENTIAL", "ClientSecretCredential") - - cred_cls = getattr(identity, cred) - # ClientSecretCredential, DefaultAzureCredential, AzureCliCredential - if cred == "ClientSecretCredential": - credential = cred_cls( + cred: Union[AzureCredentialType, str] = AzureCredentialType( + os.environ.get("AZURE_CREDENTIAL", AzureCredentialType.ClientSecretCredential) + ) + credential: Optional[ + Union[ + ClientSecretCredential, + ManagedIdentityCredential, + CertificateCredential, + Any, + ] + ] = None + if cred == AzureCredentialType.ClientSecretCredential: + credential = ClientSecretCredential( client_id=os.environ["AZURE_CLIENT_ID"], client_secret=os.environ["AZURE_CLIENT_SECRET"], tenant_id=os.environ["AZURE_TENANT_ID"], ) - elif cred == "ManagedIdentityCredential": - credential = cred_cls(client_id=os.environ["AZURE_CLIENT_ID"]) + elif cred == AzureCredentialType.ManagedIdentityCredential: + credential = ManagedIdentityCredential(client_id=os.environ["AZURE_CLIENT_ID"]) + elif cred == AzureCredentialType.CertificateCredential: + credential = CertificateCredential( + client_id=os.environ["AZURE_CLIENT_ID"], + tenant_id=os.environ["AZURE_TENANT_ID"], + certificate_path=os.environ["AZURE_CERTIFICATE_PATH"], + ) 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) diff --git a/litellm/types/secret_managers/get_azure_ad_token_provider.py b/litellm/types/secret_managers/get_azure_ad_token_provider.py new file mode 100644 index 0000000000..f318b4333b --- /dev/null +++ b/litellm/types/secret_managers/get_azure_ad_token_provider.py @@ -0,0 +1,7 @@ +from enum import Enum + + +class AzureCredentialType(str, Enum): + ClientSecretCredential = "ClientSecretCredential" + ManagedIdentityCredential = "ManagedIdentityCredential" + CertificateCredential = "CertificateCredential" diff --git a/poetry.lock b/poetry.lock index 2bd91505cb..cadc7cd6d7 100644 --- a/poetry.lock +++ b/poetry.lock @@ -258,10 +258,9 @@ tests-mypy = ["mypy (>=1.11.1) ; platform_python_implementation == \"CPython\" a name = "azure-core" version = "1.33.0" description = "Microsoft Azure Core Library for Python" -optional = true +optional = false python-versions = ">=3.8" -groups = ["main"] -markers = "extra == \"extra-proxy\"" +groups = ["main", "proxy-dev"] files = [ {file = "azure_core-1.33.0-py3-none-any.whl", hash = "sha256:9b5b6d0223a1d38c37500e6971118c1e0f13f54951e6893968b38910bc9cda8f"}, {file = "azure_core-1.33.0.tar.gz", hash = "sha256:f367aa07b5e3005fec2c1e184b882b0b039910733907d001c20fb08ebb8c0eb9"}, @@ -280,10 +279,9 @@ tracing = ["opentelemetry-api (>=1.26,<2.0)"] name = "azure-identity" version = "1.21.0" description = "Microsoft Azure Identity Library for Python" -optional = true +optional = false python-versions = ">=3.8" -groups = ["main"] -markers = "extra == \"extra-proxy\"" +groups = ["main", "proxy-dev"] files = [ {file = "azure_identity-1.21.0-py3-none-any.whl", hash = "sha256:258ea6325537352440f71b35c3dffe9d240eae4a5126c1b7ce5efd5766bd9fd9"}, {file = "azure_identity-1.21.0.tar.gz", hash = "sha256:ea22ce6e6b0f429bc1b8d9212d5b9f9877bd4c82f1724bfa910760612c07a9a6"}, @@ -499,7 +497,7 @@ version = "1.17.1" description = "Foreign Function Interface for Python calling C code." optional = false python-versions = ">=3.8" -groups = ["main", "dev"] +groups = ["main", "dev", "proxy-dev"] files = [ {file = "cffi-1.17.1-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:df8b1c11f177bc2313ec4b2d46baec87a5f3e71fc8b45dab2ee7cae86d9aba14"}, {file = "cffi-1.17.1-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:8f2cdc858323644ab277e9bb925ad72ae0e67f69e804f4898c070998d50b1a67"}, @@ -569,7 +567,7 @@ files = [ {file = "cffi-1.17.1-cp39-cp39-win_amd64.whl", hash = "sha256:d016c76bdd850f3c626af19b0542c9677ba156e4ee4fccfdd7848803533ef662"}, {file = "cffi-1.17.1.tar.gz", hash = "sha256:1c39c6016c32bc48dd54561950ebd6836e1670f2ae46128f67cf49e789c52824"}, ] -markers = {main = "(extra == \"proxy\" or extra == \"extra-proxy\") and (platform_python_implementation != \"PyPy\" or extra == \"proxy\")", dev = "platform_python_implementation != \"PyPy\""} +markers = {main = "platform_python_implementation != \"PyPy\" or extra == \"proxy\"", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""} [package.dependencies] pycparser = "*" @@ -729,7 +727,7 @@ version = "43.0.3" description = "cryptography is a package which provides cryptographic recipes and primitives to Python developers." optional = false python-versions = ">=3.7" -groups = ["main", "dev"] +groups = ["main", "dev", "proxy-dev"] files = [ {file = "cryptography-43.0.3-cp37-abi3-macosx_10_9_universal2.whl", hash = "sha256:bf7a1932ac4176486eab36a19ed4c0492da5d97123f1406cf15e41b05e787d2e"}, {file = "cryptography-43.0.3-cp37-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:63efa177ff54aec6e1c0aefaa1a241232dcd37413835a9b674b6e3f0ae2bfd3e"}, @@ -759,7 +757,6 @@ files = [ {file = "cryptography-43.0.3-pp39-pypy39_pp73-win_amd64.whl", hash = "sha256:2ce6fae5bdad59577b44e4dfed356944fbf1d925269114c28be377692643b4ff"}, {file = "cryptography-43.0.3.tar.gz", hash = "sha256:315b9001266a492a6ff443b61238f956b214dbec9910a081ba5b6646a055a805"}, ] -markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""} [package.dependencies] cffi = {version = ">=1.12", markers = "platform_python_implementation != \"PyPy\""} @@ -1986,10 +1983,9 @@ dev = ["absl-py", "pyink", "pylint (>=2.6.0)", "pytest", "pytest-xdist"] name = "msal" version = "1.32.3" description = "The Microsoft Authentication Library (MSAL) for Python library enables your app to access the Microsoft Cloud by supporting authentication of users with Microsoft Azure Active Directory accounts (AAD) and Microsoft Accounts (MSA) using industry standard OAuth2 and OpenID Connect." -optional = true +optional = false python-versions = ">=3.7" -groups = ["main"] -markers = "extra == \"extra-proxy\"" +groups = ["main", "proxy-dev"] files = [ {file = "msal-1.32.3-py3-none-any.whl", hash = "sha256:b2798db57760b1961b142f027ffb7c8169536bf77316e99a0df5c4aaebb11569"}, {file = "msal-1.32.3.tar.gz", hash = "sha256:5eea038689c78a5a70ca8ecbe1245458b55a857bd096efb6989c69ba15985d35"}, @@ -2007,10 +2003,9 @@ broker = ["pymsalruntime (>=0.14,<0.18) ; python_version >= \"3.6\" and platform name = "msal-extensions" version = "1.3.0" description = "Microsoft Authentication Library extensions (MSAL EX) provides a persistence API that can save your data on disk, encrypted on Windows, macOS and Linux. Concurrent data access will be coordinated by a file lock mechanism." -optional = true +optional = false python-versions = ">=3.7" -groups = ["main"] -markers = "extra == \"extra-proxy\"" +groups = ["main", "proxy-dev"] files = [ {file = "msal_extensions-1.3.0-py3-none-any.whl", hash = "sha256:105328ddcbdd342016c9949d8f89e3917554740c8ab26669c0fa0e069e730a0e"}, {file = "msal_extensions-1.3.0.tar.gz", hash = "sha256:96918996642b38c78cd59b55efa0f06fd1373c90e0949be8615697c048fba62c"}, @@ -2937,12 +2932,12 @@ version = "2.22" description = "C parser in Python" optional = false python-versions = ">=3.8" -groups = ["main", "dev"] +groups = ["main", "dev", "proxy-dev"] files = [ {file = "pycparser-2.22-py3-none-any.whl", hash = "sha256:c3702b6d3dd8c7abc1afa565d7e63d53a1d0bd86cdc24edd75470f4de499cfcc"}, {file = "pycparser-2.22.tar.gz", hash = "sha256:491c8be9c040f5390f5bf44a5b07752bd07f56edf992381b05c701439eec10f6"}, ] -markers = {main = "(extra == \"proxy\" or extra == \"extra-proxy\") and (platform_python_implementation != \"PyPy\" or extra == \"proxy\")", dev = "platform_python_implementation != \"PyPy\""} +markers = {main = "platform_python_implementation != \"PyPy\" or extra == \"proxy\"", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""} [[package]] name = "pydantic" @@ -3136,10 +3131,9 @@ windows-terminal = ["colorama (>=0.4.6)"] name = "pyjwt" version = "2.9.0" description = "JSON Web Token implementation in Python" -optional = true +optional = false python-versions = ">=3.8" -groups = ["main"] -markers = "extra == \"extra-proxy\" or extra == \"proxy\"" +groups = ["main", "proxy-dev"] files = [ {file = "PyJWT-2.9.0-py3-none-any.whl", hash = "sha256:3b02fb0f44517787776cf48f2ae25d8e14f300e6d7545a4315cee571a415e850"}, {file = "pyjwt-2.9.0.tar.gz", hash = "sha256:7e1e5b56cc735432a7369cbfa0efe50fa113ebecdc04ae6922deba8b84582d0c"}, @@ -3879,10 +3873,9 @@ crt = ["botocore[crt] (>=1.33.2,<2.0a.0)"] name = "six" version = "1.17.0" description = "Python 2 and 3 compatibility utilities" -optional = true +optional = false python-versions = "!=3.0.*,!=3.1.*,!=3.2.*,>=2.7" -groups = ["main"] -markers = "extra == \"extra-proxy\" or extra == \"proxy\"" +groups = ["main", "proxy-dev"] files = [ {file = "six-1.17.0-py2.py3-none-any.whl", hash = "sha256:4721f391ed90541fddacab5acf947aa0d3dc7d27b2e1e8eda2be8970586c3274"}, {file = "six-1.17.0.tar.gz", hash = "sha256:ff70335d468e7eb6ec65b95b99d3a2836546063f63acc5171de367e834932a81"}, @@ -4935,4 +4928,4 @@ utils = ["numpydoc"] [metadata] lock-version = "2.1" python-versions = ">=3.8.1,<4.0, !=3.9.7" -content-hash = "15bad8ae37c1e7cf21555b0023150bfd4bd7d6d548828f6c65a66283d14a189b" +content-hash = "fe1bc122aaeae89043f8099bfe7fda082110b06ddee5147cbfcdd085f072ef72" diff --git a/pyproject.toml b/pyproject.toml index d3fd3ecf0c..0cc0c4a826 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -130,6 +130,7 @@ prometheus-client = "0.20.0" opentelemetry-api = "1.25.0" opentelemetry-sdk = "1.25.0" opentelemetry-exporter-otlp = "1.25.0" +azure-identity = "^1.15.0" [build-system] requires = ["poetry-core", "wheel"] diff --git a/tests/litellm/secret_managers/test_get_azure_ad_token_provider.py b/tests/litellm/secret_managers/test_get_azure_ad_token_provider.py new file mode 100644 index 0000000000..7fd427dae4 --- /dev/null +++ b/tests/litellm/secret_managers/test_get_azure_ad_token_provider.py @@ -0,0 +1,138 @@ +import json +import os +import sys +from typing import Optional +from unittest.mock import MagicMock, patch + +# Adds the grandparent directory to sys.path to allow importing project modules +sys.path.insert(0, os.path.abspath("../..")) + +import pytest + +from litellm.secret_managers.get_azure_ad_token_provider import ( + get_azure_ad_token_provider, +) + + +class TestGetAzureAdTokenProvider: + @patch.dict( + os.environ, + { + "AZURE_CLIENT_ID": "test-client-id", + "AZURE_CLIENT_SECRET": "test-client-secret", + "AZURE_TENANT_ID": "test-tenant-id", + "AZURE_SCOPE": "https://cognitiveservices.azure.com/.default", + "AZURE_CREDENTIAL": "ClientSecretCredential", + }, + ) + @patch("azure.identity.get_bearer_token_provider") + @patch("azure.identity.ClientSecretCredential") + def test_get_azure_ad_token_provider_client_secret_credential( + self, mock_client_secret_credential, mock_get_bearer_token_provider + ): + """Test get_azure_ad_token_provider with ClientSecretCredential.""" + # Mock the Azure identity credential instance + mock_credential_instance = MagicMock() + mock_client_secret_credential.return_value = mock_credential_instance + + # Mock the bearer token provider + mock_token_provider = MagicMock(return_value="mock-token") + mock_get_bearer_token_provider.return_value = mock_token_provider + + # Call the function + result = get_azure_ad_token_provider() + + # Assertions + assert callable(result) + mock_client_secret_credential.assert_called_once_with( + client_id="test-client-id", + client_secret="test-client-secret", + tenant_id="test-tenant-id", + ) + mock_get_bearer_token_provider.assert_called_once_with( + mock_credential_instance, "https://cognitiveservices.azure.com/.default" + ) + + # Test that the returned callable works + token = result() + assert token == "mock-token" + + @patch.dict( + os.environ, + { + "AZURE_CLIENT_ID": "test-client-id", + "AZURE_SCOPE": "https://cognitiveservices.azure.com/.default", + "AZURE_CREDENTIAL": "ManagedIdentityCredential", + }, + ) + @patch("azure.identity.get_bearer_token_provider") + @patch("azure.identity.ManagedIdentityCredential") + def test_get_azure_ad_token_provider_managed_identity_credential( + self, mock_managed_identity_credential, mock_get_bearer_token_provider + ): + """Test get_azure_ad_token_provider with ManagedIdentityCredential.""" + # Mock the Azure identity credential instance + mock_credential_instance = MagicMock() + mock_managed_identity_credential.return_value = mock_credential_instance + + # Mock the bearer token provider + mock_token_provider = MagicMock(return_value="mock-managed-identity-token") + mock_get_bearer_token_provider.return_value = mock_token_provider + + # Call the function + result = get_azure_ad_token_provider() + + # Assertions + assert callable(result) + mock_managed_identity_credential.assert_called_once_with( + client_id="test-client-id" + ) + mock_get_bearer_token_provider.assert_called_once_with( + mock_credential_instance, "https://cognitiveservices.azure.com/.default" + ) + + # Test that the returned callable works + token = result() + assert token == "mock-managed-identity-token" + + @patch.dict( + os.environ, + { + "AZURE_CLIENT_ID": "test-client-id", + "AZURE_TENANT_ID": "test-tenant-id", + "AZURE_CERTIFICATE_PATH": "/path/to/cert.pem", + "AZURE_SCOPE": "https://cognitiveservices.azure.com/.default", + "AZURE_CREDENTIAL": "CertificateCredential", + }, + ) + @patch("azure.identity.get_bearer_token_provider") + @patch("azure.identity.CertificateCredential") + def test_get_azure_ad_token_provider_certificate_credential( + self, mock_certificate_credential, mock_get_bearer_token_provider + ): + """Test get_azure_ad_token_provider with CertificateCredential.""" + # Mock the Azure identity credential instance + mock_credential_instance = MagicMock() + mock_certificate_credential.return_value = mock_credential_instance + + # Mock the bearer token provider + mock_token_provider = MagicMock(return_value="mock-certificate-token") + mock_get_bearer_token_provider.return_value = mock_token_provider + + # Call the function + result = get_azure_ad_token_provider() + + # Assertions + assert callable(result) + mock_certificate_credential.assert_called_once_with( + client_id="test-client-id", + tenant_id="test-tenant-id", + certificate_path="/path/to/cert.pem", + ) + mock_get_bearer_token_provider.assert_called_once_with( + mock_credential_instance, "https://cognitiveservices.azure.com/.default" + ) + + # Test that the returned callable works + token = result() + assert token == "mock-certificate-token"