diff --git a/prowler/CHANGELOG.md b/prowler/CHANGELOG.md index d75e189be8..f180827796 100644 --- a/prowler/CHANGELOG.md +++ b/prowler/CHANGELOG.md @@ -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) --- diff --git a/prowler/providers/azure/azure_provider.py b/prowler/providers/azure/azure_provider.py index be8b7e60a3..c9496ac0a5 100644 --- a/prowler/providers/azure/azure_provider.py +++ b/prowler/providers/azure/azure_provider.py @@ -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(): diff --git a/prowler/providers/azure/lib/regions/regions.py b/prowler/providers/azure/lib/regions/regions.py index 6b88ab5561..b72a7c3af4 100644 --- a/prowler/providers/azure/lib/regions/regions.py +++ b/prowler/providers/azure/lib/regions/regions.py @@ -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] diff --git a/prowler/providers/azure/lib/service/service.py b/prowler/providers/azure/lib/service/service.py index f8cfd417c9..a0a832ca01 100644 --- a/prowler/providers/azure/lib/service/service.py +++ b/prowler/providers/azure/lib/service/service.py @@ -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( diff --git a/prowler/providers/azure/models.py b/prowler/providers/azure/models.py index 752d1372c6..62d03db365 100644 --- a/prowler/providers/azure/models.py +++ b/prowler/providers/azure/models.py @@ -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): diff --git a/tests/providers/azure/azure_provider_test.py b/tests/providers/azure/azure_provider_test.py index 0846fdc540..1d9aa97e6b 100644 --- a/tests/providers/azure/azure_provider_test.py +++ b/tests/providers/azure/azure_provider_test.py @@ -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 diff --git a/tests/providers/azure/lib/regions/regions_test.py b/tests/providers/azure/lib/regions/regions_test.py index 2f8fdd6053..674897725d 100644 --- a/tests/providers/azure/lib/regions/regions_test.py +++ b/tests/providers/azure/lib/regions/regions_test.py @@ -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, }, } diff --git a/tests/providers/azure/lib/service/azure_service_test.py b/tests/providers/azure/lib/service/azure_service_test.py new file mode 100644 index 0000000000..9360be85ea --- /dev/null +++ b/tests/providers/azure/lib/service/azure_service_test.py @@ -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)