fix(azure): pass authority to credentials for sovereign clouds (#10284)

Co-authored-by: Hugo P.Brito <hugopbrit@gmail.com>
Co-authored-by: Hugo Pereira Brito <101209179+HugoPBrito@users.noreply.github.com>
This commit is contained in:
authored and GitHub committed 2026-05-29 15:17:41 +02:00
1 parent 2a641b39c8
commit e3c4368d32
8 files changed
+525 -11

No files matched your search

+1
View File
@@ -18,6 +18,7 @@ All notable changes to the **Prowler SDK** are documented in this file.
- ENS RD 311/2022 (AWS) compliance mapping: `vpc_different_regions` was uncorrectly mapped under the `mp.com.4` family (Network segregation). That check is now mapped to a new `op.cont.2.aws.vpc.1` requirement under the Continuity of Service control [(#11372)](https://github.com/prowler-cloud/prowler/pull/11372)
- Compliance CSV row count now matches the UI per requirement by sourcing rows from the framework JSON's `requirement.Checks` instead of the stale `finding.compliance` snapshot [(#11370)](https://github.com/prowler-cloud/prowler/pull/11370)
- OpenStack provider exception codes moved from the `10000-10999` range, shared with the AlibabaCloud provider, to the free `17000-17999` range to keep error codes unambiguous [(#11382)](https://github.com/prowler-cloud/prowler/pull/11382)
- Azure provider now supports authentication against sovereign clouds (`AzureChinaCloud`, `AzureUSGovernment`) [(#10284)](https://github.com/prowler-cloud/prowler/pull/10284)
- Deprecate `s3_bucket_default_encryption` check for AWS provider since SSE-S3 is automatically applied to all S3 buckets by AWS as of January 5, 2023 and can no longer be disabled [(#11230)](https://github.com/prowler-cloud/prowler/pull/11230)
---
+50 -9
View File
@@ -241,7 +241,10 @@ class AzureProvider(Provider):
azure_credentials = None
if tenant_id and client_id and client_secret:
azure_credentials = self.validate_static_credentials(
tenant_id=tenant_id, client_id=client_id, client_secret=client_secret
tenant_id=tenant_id,
client_id=client_id,
client_secret=client_secret,
region_config=self._region_config,
)
# Set up the Azure session
@@ -410,6 +413,9 @@ class AzureProvider(Provider):
authority=config["authority"],
base_url=config["base_url"],
credential_scopes=config["credential_scopes"],
graph_host=config["graph_host"],
graph_scope=config["graph_scope"],
logs_endpoint=config["logs_endpoint"],
)
except ArgumentTypeError as validation_error:
logger.error(
@@ -507,6 +513,7 @@ class AzureProvider(Provider):
tenant_id=azure_credentials["tenant_id"],
client_id=azure_credentials["client_id"],
client_secret=azure_credentials["client_secret"],
authority=region_config.authority,
)
return credentials
except ClientAuthenticationError as error:
@@ -579,7 +586,10 @@ class AzureProvider(Provider):
)
else:
try:
credentials = InteractiveBrowserCredential(tenant_id=tenant_id)
credentials = InteractiveBrowserCredential(
tenant_id=tenant_id,
authority=region_config.authority,
)
except Exception as error:
logger.critical(
"Failed to retrieve azure credentials using browser authentication"
@@ -662,6 +672,7 @@ class AzureProvider(Provider):
tenant_id=tenant_id,
client_id=client_id,
client_secret=client_secret,
region_config=region_config,
)
# Set up the Azure session
@@ -675,7 +686,11 @@ class AzureProvider(Provider):
region_config,
)
# Create a SubscriptionClient
subscription_client = SubscriptionClient(credentials)
subscription_client = SubscriptionClient(
credentials,
base_url=region_config.base_url,
credential_scopes=region_config.credential_scopes,
)
# Get info from the subscriptions
available_subscriptions = []
@@ -1039,7 +1054,11 @@ class AzureProvider(Provider):
}
"""
credentials = self.session
subscription_client = SubscriptionClient(credentials)
subscription_client = SubscriptionClient(
credentials,
base_url=self.region_config.base_url,
credential_scopes=self.region_config.credential_scopes,
)
locations = {}
for subscription_id, display_name in self._identity.subscriptions.items():
@@ -1084,7 +1103,10 @@ class AzureProvider(Provider):
@staticmethod
def validate_static_credentials(
tenant_id: str = None, client_id: str = None, client_secret: str = None
tenant_id: str = None,
client_id: str = None,
client_secret: str = None,
region_config: AzureRegionConfig = None,
) -> dict:
"""
Validates the static credentials for the Azure provider.
@@ -1093,6 +1115,9 @@ class AzureProvider(Provider):
tenant_id (str): The Azure Active Directory tenant ID.
client_id (str): The Azure client ID.
client_secret (str): The Azure client secret.
region_config (AzureRegionConfig): The region configuration used to
build the per-cloud login endpoint and Graph scope. Defaults to
the public-cloud configuration when not provided.
Raises:
AzureNotValidTenantIdError: If the provided Azure Tenant ID is not valid.
@@ -1129,8 +1154,13 @@ class AzureProvider(Provider):
message="The provided Azure Client Secret is not valid.",
)
if region_config is None:
region_config = AzureProvider.setup_region_config("AzureCloud")
try:
AzureProvider.verify_client(tenant_id, client_id, client_secret)
AzureProvider.verify_client(
tenant_id, client_id, client_secret, region_config
)
return {
"tenant_id": tenant_id,
"client_id": client_id,
@@ -1162,7 +1192,9 @@ class AzureProvider(Provider):
)
@staticmethod
def verify_client(tenant_id, client_id, client_secret) -> None:
def verify_client(
tenant_id, client_id, client_secret, region_config: AzureRegionConfig = None
) -> None:
"""
Verifies the Azure client credentials using the specified tenant ID, client ID, and client secret.
@@ -1170,6 +1202,9 @@ class AzureProvider(Provider):
tenant_id (str): The Azure Active Directory tenant ID.
client_id (str): The Azure client ID.
client_secret (str): The Azure client secret.
region_config (AzureRegionConfig): The region configuration used to
build the per-cloud login endpoint and Graph scope. Defaults to
the public-cloud configuration when not provided.
Raises:
AzureNotValidTenantIdError: If the provided Azure Tenant ID is not valid.
@@ -1179,7 +1214,13 @@ class AzureProvider(Provider):
Returns:
None
"""
url = f"https://login.microsoftonline.com/{tenant_id}/oauth2/v2.0/token"
if region_config is None:
region_config = AzureProvider.setup_region_config("AzureCloud")
# `authority` is None for the public cloud and a bare host (e.g.
# `login.chinacloudapi.cn`) for sovereign clouds, mirroring the
# `AzureAuthorityHosts` constants used by azure-identity.
login_endpoint = region_config.authority or "login.microsoftonline.com"
url = f"https://{login_endpoint}/{tenant_id}/oauth2/v2.0/token"
headers = {
"Content-Type": "application/x-www-form-urlencoded",
"Accept": "application/json",
@@ -1188,7 +1229,7 @@ class AzureProvider(Provider):
"grant_type": "client_credentials",
"client_id": client_id,
"client_secret": client_secret,
"scope": "https://graph.microsoft.com/.default",
"scope": region_config.graph_scope,
}
response = requests.post(url, headers=headers, data=data).json()
if "access_token" not in response.keys() and "error_codes" in response.keys():
@@ -4,6 +4,18 @@ AZURE_CHINA_CLOUD = "https://management.chinacloudapi.cn"
AZURE_US_GOV_CLOUD = "https://management.usgovcloudapi.net"
AZURE_GENERIC_CLOUD = "https://management.azure.com"
AZURE_GENERIC_GRAPH_HOST = "https://graph.microsoft.com"
AZURE_CHINA_GRAPH_HOST = "https://microsoftgraph.chinacloudapi.cn"
AZURE_US_GOV_GRAPH_HOST = "https://graph.microsoft.us"
AZURE_GENERIC_GRAPH_SCOPE = f"{AZURE_GENERIC_GRAPH_HOST}/.default"
AZURE_CHINA_GRAPH_SCOPE = f"{AZURE_CHINA_GRAPH_HOST}/.default"
AZURE_US_GOV_GRAPH_SCOPE = f"{AZURE_US_GOV_GRAPH_HOST}/.default"
AZURE_GENERIC_LOGS_ENDPOINT = "https://api.loganalytics.io"
AZURE_CHINA_LOGS_ENDPOINT = "https://api.loganalytics.azure.cn"
AZURE_US_GOV_LOGS_ENDPOINT = "https://api.loganalytics.us"
def get_regions_config(region):
allowed_regions = {
@@ -11,16 +23,25 @@ def get_regions_config(region):
"authority": None,
"base_url": AZURE_GENERIC_CLOUD,
"credential_scopes": [AZURE_GENERIC_CLOUD + "/.default"],
"graph_host": AZURE_GENERIC_GRAPH_HOST,
"graph_scope": AZURE_GENERIC_GRAPH_SCOPE,
"logs_endpoint": AZURE_GENERIC_LOGS_ENDPOINT,
},
"AzureChinaCloud": {
"authority": AzureAuthorityHosts.AZURE_CHINA,
"base_url": AZURE_CHINA_CLOUD,
"credential_scopes": [AZURE_CHINA_CLOUD + "/.default"],
"graph_host": AZURE_CHINA_GRAPH_HOST,
"graph_scope": AZURE_CHINA_GRAPH_SCOPE,
"logs_endpoint": AZURE_CHINA_LOGS_ENDPOINT,
},
"AzureUSGovernment": {
"authority": AzureAuthorityHosts.AZURE_GOVERNMENT,
"base_url": AZURE_US_GOV_CLOUD,
"credential_scopes": [AZURE_US_GOV_CLOUD + "/.default"],
"graph_host": AZURE_US_GOV_GRAPH_HOST,
"graph_scope": AZURE_US_GOV_GRAPH_SCOPE,
"logs_endpoint": AZURE_US_GOV_LOGS_ENDPOINT,
},
}
return allowed_regions[region]
+30 -2
View File
@@ -1,5 +1,11 @@
from concurrent.futures import ThreadPoolExecutor, as_completed
from kiota_authentication_azure.azure_identity_authentication_provider import (
AzureIdentityAuthenticationProvider,
)
from msgraph.graph_request_adapter import GraphRequestAdapter
from msgraph_core import GraphClientFactory
from prowler.lib.logger import logger
from prowler.providers.azure.azure_provider import AzureProvider
@@ -47,10 +53,32 @@ class AzureService:
clients = {}
try:
if "GraphServiceClient" in str(service):
clients.update({identity.tenant_domain: service(credentials=session)})
# GraphServiceClient(credentials, scopes=...) only customises the
# OAuth scope; the underlying httpx client's base URL stays at
# graph.microsoft.com. For sovereign clouds we must also point
# the HTTP transport at the per-cloud host, which is done by
# building a custom GraphRequestAdapter with a NationalClouds
# base URL.
auth_provider = AzureIdentityAuthenticationProvider(
session, scopes=[region_config.graph_scope]
)
http_client = GraphClientFactory.create_with_default_middleware(
host=region_config.graph_host
)
request_adapter = GraphRequestAdapter(auth_provider, client=http_client)
clients.update(
{identity.tenant_domain: service(request_adapter=request_adapter)}
)
elif "LogsQueryClient" in str(service):
for subscription_id, display_name in identity.subscriptions.items():
clients.update({subscription_id: service(credential=session)})
clients.update(
{
subscription_id: service(
credential=session,
endpoint=region_config.logs_endpoint,
)
}
)
else:
for subscription_id, display_name in identity.subscriptions.items():
clients.update(
+3
View File
@@ -20,6 +20,9 @@ class AzureRegionConfig(BaseModel):
authority: Optional[str] = None
base_url: str = ""
credential_scopes: list = []
graph_host: str = "https://graph.microsoft.com"
graph_scope: str = "https://graph.microsoft.com/.default"
logs_endpoint: str = "https://api.loganalytics.io"
class AzureSubscription(BaseModel):
@@ -725,6 +725,300 @@ class TestAzureProviderSetupIdentitySubscriptions:
}
class TestAzureProviderSovereignCloudSupport:
"""Sovereign-cloud authentication coverage across AzureCloud,
AzureChinaCloud and AzureUSGovernment for every authentication code path
Prowler exposes. Pinned to issue #8425."""
REGION_CASES = [
(
"AzureCloud",
None,
"https://management.azure.com",
["https://management.azure.com/.default"],
"https://graph.microsoft.com/.default",
"https://api.loganalytics.io",
"login.microsoftonline.com",
),
(
"AzureChinaCloud",
"login.chinacloudapi.cn",
"https://management.chinacloudapi.cn",
["https://management.chinacloudapi.cn/.default"],
"https://microsoftgraph.chinacloudapi.cn/.default",
"https://api.loganalytics.azure.cn",
"login.chinacloudapi.cn",
),
(
"AzureUSGovernment",
"login.microsoftonline.us",
"https://management.usgovcloudapi.net",
["https://management.usgovcloudapi.net/.default"],
"https://graph.microsoft.us/.default",
"https://api.loganalytics.us",
"login.microsoftonline.us",
),
]
@pytest.mark.parametrize(
"region,authority,base_url,credential_scopes,graph_scope,logs_endpoint,_login_endpoint",
REGION_CASES,
)
def test_setup_region_config_per_cloud(
self,
region,
authority,
base_url,
credential_scopes,
graph_scope,
logs_endpoint,
_login_endpoint,
):
config = AzureProvider.setup_region_config(region)
# graph_host mirrors graph_scope without the `/.default` suffix; we
# derive it here to avoid threading a separate parameter through every
# parametrized test in this class.
expected_graph_host = graph_scope.removesuffix("/.default")
assert config == AzureRegionConfig(
name=region,
authority=authority,
base_url=base_url,
credential_scopes=credential_scopes,
graph_host=expected_graph_host,
graph_scope=graph_scope,
logs_endpoint=logs_endpoint,
)
@pytest.mark.parametrize(
"region,authority,_base_url,_credential_scopes,_graph_scope,_logs_endpoint,_login_endpoint",
REGION_CASES,
)
def test_setup_session_static_credentials_passes_authority(
self,
region,
authority,
_base_url,
_credential_scopes,
_graph_scope,
_logs_endpoint,
_login_endpoint,
):
with patch(
"prowler.providers.azure.azure_provider.ClientSecretCredential"
) as mock_client_secret_credential:
azure_credentials = {
"tenant_id": str(uuid4()),
"client_id": str(uuid4()),
"client_secret": "fake-secret-value",
}
region_config = AzureProvider.setup_region_config(region)
AzureProvider.setup_session(
az_cli_auth=False,
sp_env_auth=False,
browser_auth=False,
managed_identity_auth=False,
tenant_id=azure_credentials["tenant_id"],
azure_credentials=azure_credentials,
region_config=region_config,
)
mock_client_secret_credential.assert_called_once_with(
tenant_id=azure_credentials["tenant_id"],
client_id=azure_credentials["client_id"],
client_secret=azure_credentials["client_secret"],
authority=authority,
)
@pytest.mark.parametrize(
"region,authority,_base_url,_credential_scopes,_graph_scope,_logs_endpoint,_login_endpoint",
REGION_CASES,
)
def test_setup_session_browser_auth_passes_authority(
self,
region,
authority,
_base_url,
_credential_scopes,
_graph_scope,
_logs_endpoint,
_login_endpoint,
):
with patch(
"prowler.providers.azure.azure_provider.InteractiveBrowserCredential"
) as mock_interactive_browser_credential:
tenant_id = str(uuid4())
region_config = AzureProvider.setup_region_config(region)
AzureProvider.setup_session(
az_cli_auth=False,
sp_env_auth=False,
browser_auth=True,
managed_identity_auth=False,
tenant_id=tenant_id,
azure_credentials=None,
region_config=region_config,
)
mock_interactive_browser_credential.assert_called_once_with(
tenant_id=tenant_id,
authority=authority,
)
@pytest.mark.parametrize(
"region,authority,_base_url,_credential_scopes,_graph_scope,_logs_endpoint,_login_endpoint",
REGION_CASES,
)
def test_setup_session_default_credential_passes_authority(
self,
region,
authority,
_base_url,
_credential_scopes,
_graph_scope,
_logs_endpoint,
_login_endpoint,
):
with patch(
"prowler.providers.azure.azure_provider.DefaultAzureCredential"
) as mock_default_credential:
region_config = AzureProvider.setup_region_config(region)
AzureProvider.setup_session(
az_cli_auth=True,
sp_env_auth=False,
browser_auth=False,
managed_identity_auth=False,
tenant_id=None,
azure_credentials=None,
region_config=region_config,
)
_, called_kwargs = mock_default_credential.call_args
assert called_kwargs["authority"] == authority
assert called_kwargs["exclude_cli_credential"] is False
assert called_kwargs["exclude_environment_credential"] is True
assert called_kwargs["exclude_managed_identity_credential"] is True
@pytest.mark.parametrize(
"region,_authority,_base_url,_credential_scopes,graph_scope,_logs_endpoint,login_endpoint",
REGION_CASES,
)
def test_verify_client_uses_per_cloud_endpoints(
self,
region,
_authority,
_base_url,
_credential_scopes,
graph_scope,
_logs_endpoint,
login_endpoint,
):
tenant_id = str(uuid4())
client_id = str(uuid4())
client_secret = "fake-secret"
region_config = AzureProvider.setup_region_config(region)
with patch("prowler.providers.azure.azure_provider.requests.post") as mock_post:
mock_post.return_value = MagicMock()
mock_post.return_value.json.return_value = {"access_token": "fake-token"}
AzureProvider.verify_client(
tenant_id, client_id, client_secret, region_config
)
mock_post.assert_called_once()
args, kwargs = mock_post.call_args
assert args[0] == (
f"https://{login_endpoint}/{tenant_id}/oauth2/v2.0/token"
)
assert kwargs["data"]["scope"] == graph_scope
assert kwargs["data"]["client_id"] == client_id
assert kwargs["data"]["client_secret"] == client_secret
@pytest.mark.parametrize(
"region,_authority,base_url,credential_scopes,_graph_scope,_logs_endpoint,_login_endpoint",
REGION_CASES,
)
def test_test_connection_passes_base_url_to_subscription_client(
self,
region,
_authority,
base_url,
credential_scopes,
_graph_scope,
_logs_endpoint,
_login_endpoint,
):
subscription_client_instance = MagicMock()
subscription_client_instance.subscriptions = MagicMock()
subscription_client_instance.subscriptions.list = MagicMock(return_value=[])
subscription_client_class = MagicMock(return_value=subscription_client_instance)
with (
patch(
"prowler.providers.azure.azure_provider.AzureProvider.setup_session"
) as mock_setup_session,
patch(
"prowler.providers.azure.azure_provider.SubscriptionClient",
subscription_client_class,
),
):
mock_setup_session.return_value = MagicMock()
AzureProvider.test_connection(
az_cli_auth=True,
region=region,
raise_on_exception=False,
)
subscription_client_class.assert_called_once()
_, kwargs = subscription_client_class.call_args
assert kwargs["base_url"] == base_url
assert kwargs["credential_scopes"] == credential_scopes
@pytest.mark.parametrize(
"region,_authority,base_url,credential_scopes,_graph_scope,_logs_endpoint,_login_endpoint",
REGION_CASES,
)
def test_get_locations_passes_base_url_to_subscription_client(
self,
region,
_authority,
base_url,
credential_scopes,
_graph_scope,
_logs_endpoint,
_login_endpoint,
):
subscription_client_instance = MagicMock()
subscription_client_instance.subscriptions = MagicMock()
subscription_client_instance.subscriptions.list_locations = MagicMock(
return_value=[]
)
subscription_client_class = MagicMock(return_value=subscription_client_instance)
with (
patch.object(AzureProvider, "__init__", return_value=None),
patch(
"prowler.providers.azure.azure_provider.SubscriptionClient",
subscription_client_class,
),
):
azure_provider = AzureProvider()
azure_provider._session = MagicMock()
azure_provider._region_config = AzureProvider.setup_region_config(region)
azure_provider._identity = AzureIdentityInfo(subscriptions={})
azure_provider.get_locations()
subscription_client_class.assert_called_once()
_, kwargs = subscription_client_class.call_args
assert kwargs["base_url"] == base_url
assert kwargs["credential_scopes"] == credential_scopes
class TestAzureProviderSetupIdentityEventLoop:
"""Regression for the Celery worker scenario where
asyncio.get_event_loop() raised "There is no current event loop in
@@ -2,8 +2,17 @@ from azure.identity import AzureAuthorityHosts
from prowler.providers.azure.lib.regions.regions import (
AZURE_CHINA_CLOUD,
AZURE_CHINA_GRAPH_HOST,
AZURE_CHINA_GRAPH_SCOPE,
AZURE_CHINA_LOGS_ENDPOINT,
AZURE_GENERIC_CLOUD,
AZURE_GENERIC_GRAPH_HOST,
AZURE_GENERIC_GRAPH_SCOPE,
AZURE_GENERIC_LOGS_ENDPOINT,
AZURE_US_GOV_CLOUD,
AZURE_US_GOV_GRAPH_HOST,
AZURE_US_GOV_GRAPH_SCOPE,
AZURE_US_GOV_LOGS_ENDPOINT,
get_regions_config,
)
@@ -20,16 +29,25 @@ class Test_azure_regions:
"authority": None,
"base_url": AZURE_GENERIC_CLOUD,
"credential_scopes": [AZURE_GENERIC_CLOUD + "/.default"],
"graph_host": AZURE_GENERIC_GRAPH_HOST,
"graph_scope": AZURE_GENERIC_GRAPH_SCOPE,
"logs_endpoint": AZURE_GENERIC_LOGS_ENDPOINT,
},
"AzureChinaCloud": {
"authority": AzureAuthorityHosts.AZURE_CHINA,
"base_url": AZURE_CHINA_CLOUD,
"credential_scopes": [AZURE_CHINA_CLOUD + "/.default"],
"graph_host": AZURE_CHINA_GRAPH_HOST,
"graph_scope": AZURE_CHINA_GRAPH_SCOPE,
"logs_endpoint": AZURE_CHINA_LOGS_ENDPOINT,
},
"AzureUSGovernment": {
"authority": AzureAuthorityHosts.AZURE_GOVERNMENT,
"base_url": AZURE_US_GOV_CLOUD,
"credential_scopes": [AZURE_US_GOV_CLOUD + "/.default"],
"graph_host": AZURE_US_GOV_GRAPH_HOST,
"graph_scope": AZURE_US_GOV_GRAPH_SCOPE,
"logs_endpoint": AZURE_US_GOV_LOGS_ENDPOINT,
},
}
@@ -0,0 +1,108 @@
from unittest.mock import MagicMock, patch
import pytest
from prowler.providers.azure.lib.service.service import AzureService
from prowler.providers.azure.models import AzureIdentityInfo, AzureRegionConfig
REGION_CASES = [
(
"AzureCloud",
"https://graph.microsoft.com",
"https://graph.microsoft.com/.default",
"https://api.loganalytics.io",
),
(
"AzureChinaCloud",
"https://microsoftgraph.chinacloudapi.cn",
"https://microsoftgraph.chinacloudapi.cn/.default",
"https://api.loganalytics.azure.cn",
),
(
"AzureUSGovernment",
"https://graph.microsoft.us",
"https://graph.microsoft.us/.default",
"https://api.loganalytics.us",
),
]
def _identity_and_session():
identity = AzureIdentityInfo(
tenant_domain="tenant.onmicrosoft.com",
subscriptions={"sub-1": "Subscription 1"},
)
session = MagicMock()
return identity, session
class TestAzureServiceSovereignClouds:
"""Cover __set_clients__ kwargs for the Graph and Logs clients across the
three sovereign clouds — these are the two service slots in service.py
that historically defaulted to public-cloud endpoints."""
@pytest.mark.parametrize(
"_region,graph_host,graph_scope,_logs_endpoint",
REGION_CASES,
)
def test_set_clients_graph_uses_per_cloud_host_scope_and_adapter(
self, _region, graph_host, graph_scope, _logs_endpoint
):
graph_service = MagicMock()
graph_service.__str__ = MagicMock(return_value="GraphServiceClient")
region_config = AzureRegionConfig(
graph_host=graph_host,
graph_scope=graph_scope,
logs_endpoint=_logs_endpoint,
)
identity, session = _identity_and_session()
with (
patch.object(AzureService, "__init__", return_value=None),
patch(
"prowler.providers.azure.lib.service.service.AzureIdentityAuthenticationProvider"
) as mock_auth_provider_cls,
patch(
"prowler.providers.azure.lib.service.service.GraphClientFactory"
) as mock_factory,
patch(
"prowler.providers.azure.lib.service.service.GraphRequestAdapter"
) as mock_adapter_cls,
):
service = AzureService.__new__(AzureService)
service.__set_clients__(identity, session, graph_service, region_config)
mock_auth_provider_cls.assert_called_once_with(session, scopes=[graph_scope])
mock_factory.create_with_default_middleware.assert_called_once_with(
host=graph_host
)
mock_adapter_cls.assert_called_once_with(
mock_auth_provider_cls.return_value,
client=mock_factory.create_with_default_middleware.return_value,
)
graph_service.assert_called_once_with(
request_adapter=mock_adapter_cls.return_value
)
@pytest.mark.parametrize(
"_region,_graph_host,_graph_scope,logs_endpoint",
REGION_CASES,
)
def test_set_clients_logs_passes_per_cloud_endpoint(
self, _region, _graph_host, _graph_scope, logs_endpoint
):
logs_service = MagicMock()
logs_service.__str__ = MagicMock(return_value="LogsQueryClient")
region_config = AzureRegionConfig(
graph_host=_graph_host,
graph_scope=_graph_scope,
logs_endpoint=logs_endpoint,
)
identity, session = _identity_and_session()
with patch.object(AzureService, "__init__", return_value=None):
service = AzureService.__new__(AzureService)
service.__set_clients__(identity, session, logs_service, region_config)
logs_service.assert_called_once_with(credential=session, endpoint=logs_endpoint)