mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-10-09 21:14:22 +00:00
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:
8 files changed
+525
-11
No files matched your search
@@ -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)
|
||||
|
||||
---
|
||||
|
||||
@@ -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]
|
||||
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
Reference in new issue
Block a user