fix(azure): use the selected cloud endpoints in Defender and Key Vault (#12813)

This commit is contained in:
César Arroba
2026-09-16 09:25:08 +02:00
committed by GitHub
parent 757cd44ecb
commit 75c22df63b
8 changed files with 190 additions and 7 deletions
@@ -106,3 +106,19 @@ class TestAzureServiceSovereignClouds:
service.__set_clients__(identity, session, logs_service, region_config)
logs_service.assert_called_once_with(credential=session, endpoint=logs_endpoint)
class TestAzureServiceRegionConfig:
def test_init_keeps_provider_region_config(self):
region_config = AzureRegionConfig(
name="AzureUSGovernment",
base_url="https://management.usgovcloudapi.net",
credential_scopes=["https://management.usgovcloudapi.net/.default"],
)
provider = MagicMock()
provider.region_config = region_config
with patch.object(AzureService, "__set_clients__", return_value={}):
service = AzureService(MagicMock(), provider)
assert service.region_config is region_config
@@ -1,6 +1,7 @@
from datetime import timedelta
from unittest.mock import MagicMock, patch
from prowler.providers.azure.models import AzureRegionConfig
from prowler.providers.azure.services.defender.defender_service import (
Assesment,
AutoProvisioningSetting,
@@ -618,3 +619,53 @@ class Test_Defender_get_jit_policies:
mock_client.jit_network_access_policies.list_by_resource_group.assert_called_once_with(
resource_group_name="RG"
)
US_GOV_REGION_CONFIG = AzureRegionConfig(
name="AzureUSGovernment",
base_url="https://management.usgovcloudapi.net",
credential_scopes=["https://management.usgovcloudapi.net/.default"],
)
class Test_Defender_get_security_contacts_sovereign_cloud:
def _defender(self, provider):
with (
patch(DEFENDER_INIT_PATCHES[0], return_value={}),
patch(DEFENDER_INIT_PATCHES[1], return_value={}),
patch(DEFENDER_INIT_PATCHES[2], return_value={}),
patch(DEFENDER_INIT_PATCHES[3], return_value={}),
patch(DEFENDER_INIT_PATCHES[4], return_value={}),
patch(DEFENDER_INIT_PATCHES[5], return_value={}),
patch(DEFENDER_INIT_PATCHES[6], return_value={}),
):
return Defender(provider)
def test_init_requests_token_for_cloud_scope(self):
provider = set_mocked_azure_provider(azure_region_config=US_GOV_REGION_CONFIG)
self._defender(provider)
provider.session.get_token.assert_called_once_with(
"https://management.usgovcloudapi.net/.default"
)
def test_get_security_contacts_uses_cloud_management_host(self):
provider = set_mocked_azure_provider(azure_region_config=US_GOV_REGION_CONFIG)
defender = self._defender(provider)
response = MagicMock()
response.json.return_value = {"value": []}
with patch(
"prowler.providers.azure.services.defender.defender_service.requests.get",
return_value=response,
) as mock_get:
result = defender._get_security_contacts(token="token")
assert result == {AZURE_SUBSCRIPTION_ID: {}}
mock_get.assert_called_once()
url = mock_get.call_args.args[0]
assert url.startswith(
f"https://management.usgovcloudapi.net/subscriptions/{AZURE_SUBSCRIPTION_ID}/"
)
assert "management.azure.com" not in url
@@ -470,3 +470,87 @@ class Test_KeyVault_get_key_vaults:
mock_client.vaults.list_by_resource_group.assert_called_once_with(
resource_group_name="MyRG"
)
class Test_KeyVault_get_keys:
def test_get_keys_builds_key_client_from_vault_uri(self):
mock_client = MagicMock()
mock_client.keys.list.return_value = []
mock_provider = MagicMock()
mock_provider.identity = MagicMock()
with (
patch(
"prowler.providers.common.provider.Provider.get_global_provider",
return_value=mock_provider,
),
patch(
"prowler.providers.azure.services.monitor.monitor_service.Monitor",
new=MagicMock(),
),
patch(
"prowler.providers.azure.services.keyvault.keyvault_service.KeyVault._get_key_vaults",
return_value={},
),
):
from prowler.providers.azure.services.keyvault.keyvault_service import (
KeyVault,
)
keyvault = KeyVault(set_mocked_azure_provider())
keyvault.clients = {AZURE_SUBSCRIPTION_ID: mock_client}
provider = set_mocked_azure_provider()
vault_uri = "https://my-vault.vault.usgovcloudapi.net/"
with patch(
"prowler.providers.azure.services.keyvault.keyvault_service.KeyClient"
) as mock_key_client_cls:
mock_key_client_cls.return_value.list_properties_of_keys.return_value = []
keys = keyvault._get_keys(
AZURE_SUBSCRIPTION_ID, RESOURCE_GROUP, "my-vault", vault_uri, provider
)
assert keys == []
mock_key_client_cls.assert_called_once_with(
vault_url=vault_uri, credential=provider.session
)
def test_get_keys_without_vault_uri_skips_rotation_policies(self):
mock_client = MagicMock()
mock_client.keys.list.return_value = []
mock_provider = MagicMock()
mock_provider.identity = MagicMock()
with (
patch(
"prowler.providers.common.provider.Provider.get_global_provider",
return_value=mock_provider,
),
patch(
"prowler.providers.azure.services.monitor.monitor_service.Monitor",
new=MagicMock(),
),
patch(
"prowler.providers.azure.services.keyvault.keyvault_service.KeyVault._get_key_vaults",
return_value={},
),
):
from prowler.providers.azure.services.keyvault.keyvault_service import (
KeyVault,
)
keyvault = KeyVault(set_mocked_azure_provider())
keyvault.clients = {AZURE_SUBSCRIPTION_ID: mock_client}
provider = set_mocked_azure_provider()
with patch(
"prowler.providers.azure.services.keyvault.keyvault_service.KeyClient"
) as mock_key_client_cls:
keys = keyvault._get_keys(
AZURE_SUBSCRIPTION_ID, RESOURCE_GROUP, "my-vault", "", provider
)
assert keys == []
mock_key_client_cls.assert_not_called()