diff --git a/prowler/providers/azure/azure_provider.py b/prowler/providers/azure/azure_provider.py index 8b399cdc48..bffc72396b 100644 --- a/prowler/providers/azure/azure_provider.py +++ b/prowler/providers/azure/azure_provider.py @@ -350,7 +350,7 @@ class AzureProvider(Provider): @property def resource_groups(self) -> dict[str, list[str]]: - """Mapping of subscription name to the list of resource groups to scan within it.""" + """Mapping of subscription ID to the list of resource groups to scan within it.""" return self._resource_groups # TODO: this should be moved to the argparse, if not we need to enforce it from the Provider @@ -1120,6 +1120,19 @@ class AzureProvider(Provider): return set(chain.from_iterable(locations.values())) def validate_resource_groups(self, resource_groups: list) -> dict[str, list[str]]: + """Validate requested resource groups across Azure subscriptions. + + Args: + resource_groups: Resource group names requested for scanning. + + Returns: + A mapping of subscription IDs to the matching resource group names. + + The matching is case-insensitive and resolved independently for each + subscription. If a subscription's resource groups cannot be queried, a + warning is logged and that subscription keeps an empty resource group + list so the remaining subscriptions can still be validated. + """ resource_groups = [r.strip() for r in resource_groups if r and r.strip()] if not resource_groups: return {} @@ -1140,10 +1153,11 @@ class AzureProvider(Provider): existing_rgs = { rg.name.lower(): rg.name for rg in rg_client.resource_groups.list() } - except Exception as e: + except Exception as error: logger.warning( f"Could not list resource groups for subscription '{display_name}' " - f"({subscription_id}): {e}. Skipping resource group filtering for this subscription." + f"({subscription_id}): {error.__class__.__name__}[{error.__traceback__.tb_lineno}]: {error}. " + "Skipping resource group filtering for this subscription." ) continue diff --git a/prowler/providers/azure/lib/service/service.py b/prowler/providers/azure/lib/service/service.py index ae9b127647..9d63639e94 100644 --- a/prowler/providers/azure/lib/service/service.py +++ b/prowler/providers/azure/lib/service/service.py @@ -1,3 +1,4 @@ +from collections.abc import Callable, Iterable from concurrent.futures import ThreadPoolExecutor, as_completed from kiota_authentication_azure.azure_identity_authentication_provider import ( @@ -50,7 +51,25 @@ class AzureService: return results - def list_with_rg_scope(self, subscription_id, list_all_fn, list_by_rg_fn): + def list_with_rg_scope( + self, + subscription_id: str, + list_all_fn: Callable[[], Iterable[object]], + list_by_rg_fn: Callable[..., Iterable[object]], + ) -> list[object]: + """List Azure resources using the provider resource group scope. + + Args: + subscription_id: Subscription ID whose resource group scope should be used. + list_all_fn: Callable that lists all resources in the subscription when + no resource group filter is configured. + list_by_rg_fn: Callable that lists resources for a single resource + group. It must accept ``resource_group_name`` as a keyword argument. + + Returns: + A list containing the resources returned by the selected Azure SDK + list operation. + """ if not self.resource_groups: return list(list_all_fn()) resource_groups = self.resource_groups.get(subscription_id, []) diff --git a/prowler/providers/azure/services/defender/defender_service.py b/prowler/providers/azure/services/defender/defender_service.py index b4ce4239cc..d68d88dc22 100644 --- a/prowler/providers/azure/services/defender/defender_service.py +++ b/prowler/providers/azure/services/defender/defender_service.py @@ -230,12 +230,12 @@ class Defender(AzureService): iot_security_solutions = {} for subscription_id, client in self.clients.items(): try: + iot_security_solutions.update({subscription_id: {}}) iot_security_solutions_list = self.list_with_rg_scope( subscription_id, client.iot_security_solution.list_by_subscription, client.iot_security_solution.list_by_resource_group, ) - iot_security_solutions.update({subscription_id: {}}) for iot_security_solution in iot_security_solutions_list: iot_security_solutions[subscription_id].update( { diff --git a/prowler/providers/azure/services/postgresql/postgresql_service.py b/prowler/providers/azure/services/postgresql/postgresql_service.py index e7fe98fe79..3d4855f256 100644 --- a/prowler/providers/azure/services/postgresql/postgresql_service.py +++ b/prowler/providers/azure/services/postgresql/postgresql_service.py @@ -19,13 +19,13 @@ class PostgreSQL(AzureService): flexible_servers = {} for subscription, client in self.clients.items(): try: + flexible_servers.update({subscription: []}) flexible_servers_list = self.list_with_rg_scope( subscription, client.servers.list, client.servers.list_by_resource_group, ) - flexible_servers.update({subscription: []}) for postgresql_server in flexible_servers_list: # Isolate each server: a failure collecting one server must # not abort collection of the remaining servers in the diff --git a/tests/providers/azure/azure_provider_test.py b/tests/providers/azure/azure_provider_test.py index b4fba94691..5e493a90c7 100644 --- a/tests/providers/azure/azure_provider_test.py +++ b/tests/providers/azure/azure_provider_test.py @@ -647,6 +647,32 @@ class TestAzureProviderValidateResourceGroups: assert result[sub_name] == ["rg-prod"] + @patch("prowler.providers.azure.azure_provider.ResourceManagementClient") + def test_validate_resource_groups_skips_subscription_when_listing_fails( + self, mock_rm_client + ): + accessible_subscription = str(uuid4()) + failing_subscription = str(uuid4()) + provider = self._make_provider( + subscriptions={ + accessible_subscription: "Accessible Sub", + failing_subscription: "Failing Sub", + } + ) + + mock_rg = MagicMock() + mock_rg.name = "rg-prod" + accessible_client = MagicMock() + accessible_client.resource_groups.list.return_value = [mock_rg] + failing_client = MagicMock() + failing_client.resource_groups.list.side_effect = Exception("Forbidden") + mock_rm_client.side_effect = [accessible_client, failing_client] + + result = provider.validate_resource_groups(["rg-prod"]) + + assert result[accessible_subscription] == ["rg-prod"] + assert result[failing_subscription] == [] + class TestAzureProviderSetupIdentitySubscriptions: """Regression tests ensuring identity.subscriptions preserves every diff --git a/tests/providers/azure/services/app/app_service_test.py b/tests/providers/azure/services/app/app_service_test.py index 8cbc5ad54f..76f602431e 100644 --- a/tests/providers/azure/services/app/app_service_test.py +++ b/tests/providers/azure/services/app/app_service_test.py @@ -331,6 +331,59 @@ class Test_App_get_apps: mock_client.web_apps.list.assert_not_called() assert result[AZURE_SUBSCRIPTION_ID] == {} + def test_get_apps_with_multiple_resource_groups(self): + mock_client = MagicMock() + mock_client.web_apps.list_by_resource_group.return_value = [] + + with ( + patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=set_mocked_azure_provider(), + ), + patch( + "prowler.providers.azure.services.monitor.monitor_service.Monitor", + new=MagicMock(), + ), + ): + from prowler.providers.azure.services.app.app_service import App + + app = App(set_mocked_azure_provider()) + + app.clients = {AZURE_SUBSCRIPTION_ID: mock_client} + app.resource_groups = {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} + + result = app._get_apps() + + assert mock_client.web_apps.list_by_resource_group.call_count == 2 + assert AZURE_SUBSCRIPTION_ID in result + + def test_get_apps_with_mixed_case_resource_group(self): + mock_client = MagicMock() + mock_client.web_apps.list_by_resource_group.return_value = [] + + with ( + patch( + "prowler.providers.common.provider.Provider.get_global_provider", + return_value=set_mocked_azure_provider(), + ), + patch( + "prowler.providers.azure.services.monitor.monitor_service.Monitor", + new=MagicMock(), + ), + ): + from prowler.providers.azure.services.app.app_service import App + + app = App(set_mocked_azure_provider()) + + app.clients = {AZURE_SUBSCRIPTION_ID: mock_client} + app.resource_groups = {AZURE_SUBSCRIPTION_ID: ["RG"]} + + app._get_apps() + + mock_client.web_apps.list_by_resource_group.assert_called_once_with( + resource_group_name="RG" + ) + class Test_App_get_functions: def test_get_functions_no_resource_groups(self): @@ -415,61 +468,6 @@ class Test_App_get_functions: mock_client.web_apps.list.assert_not_called() assert result[AZURE_SUBSCRIPTION_ID] == {} - def test_get_apps_with_multiple_resource_groups(self): - mock_client = MagicMock() - mock_client.web_apps.list_by_resource_group.return_value = [] - - with ( - patch( - "prowler.providers.common.provider.Provider.get_global_provider", - return_value=set_mocked_azure_provider(), - ), - patch( - "prowler.providers.azure.services.monitor.monitor_service.Monitor", - new=MagicMock(), - ), - ): - from prowler.providers.azure.services.app.app_service import App - - app = App(set_mocked_azure_provider()) - - app.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - app.resource_groups = {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} - - result = app._get_apps() - - assert mock_client.web_apps.list_by_resource_group.call_count == 2 - assert AZURE_SUBSCRIPTION_ID in result - - def test_get_apps_with_mixed_case_resource_group(self): - mock_client = MagicMock() - mock_client.web_apps.list_by_resource_group.return_value = [] - - with ( - patch( - "prowler.providers.common.provider.Provider.get_global_provider", - return_value=set_mocked_azure_provider(), - ), - patch( - "prowler.providers.azure.services.monitor.monitor_service.Monitor", - new=MagicMock(), - ), - ): - from prowler.providers.azure.services.app.app_service import App - - app = App(set_mocked_azure_provider()) - - app.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - app.resource_groups = {AZURE_SUBSCRIPTION_ID: ["RG"]} - - app._get_apps() - - mock_client.web_apps.list_by_resource_group.assert_called_once_with( - resource_group_name="RG" - ) - - -class Test_App_get_functions_extra: def test_get_functions_with_multiple_resource_groups(self): mock_client = MagicMock() mock_client.web_apps.list_by_resource_group.return_value = [] diff --git a/tests/providers/azure/services/containerregistry/containerregistry_service_test.py b/tests/providers/azure/services/containerregistry/containerregistry_service_test.py index c3468ca086..36883f97d9 100644 --- a/tests/providers/azure/services/containerregistry/containerregistry_service_test.py +++ b/tests/providers/azure/services/containerregistry/containerregistry_service_test.py @@ -95,8 +95,6 @@ class TestContainerRegistryService: class Test_ContainerRegistry_get_registries: def test_get_container_registries_no_resource_groups(self): - from unittest.mock import MagicMock, patch - mock_client = MagicMock() mock_client.registries.list.return_value = [] @@ -135,8 +133,6 @@ class Test_ContainerRegistry_get_registries: assert AZURE_SUBSCRIPTION_ID in result def test_get_container_registries_with_resource_group(self): - from unittest.mock import MagicMock, patch - mock_client = MagicMock() mock_client.registries.list_by_resource_group.return_value = [] @@ -177,8 +173,6 @@ class Test_ContainerRegistry_get_registries: assert AZURE_SUBSCRIPTION_ID in result def test_get_container_registries_empty_resource_group_for_subscription(self): - from unittest.mock import MagicMock, patch - mock_client = MagicMock() mock_provider = MagicMock() @@ -216,8 +210,6 @@ class Test_ContainerRegistry_get_registries: assert result[AZURE_SUBSCRIPTION_ID] == {} def test_get_container_registries_with_multiple_resource_groups(self): - from unittest.mock import MagicMock, patch - mock_client = MagicMock() mock_client.registries.list_by_resource_group.return_value = [] @@ -258,8 +250,6 @@ class Test_ContainerRegistry_get_registries: assert AZURE_SUBSCRIPTION_ID in result def test_get_container_registries_with_mixed_case_resource_group(self): - from unittest.mock import MagicMock, patch - mock_client = MagicMock() mock_client.registries.list_by_resource_group.return_value = [] diff --git a/tests/providers/azure/services/defender/defender_service_test.py b/tests/providers/azure/services/defender/defender_service_test.py index 71457fc6ac..b50ddd26c9 100644 --- a/tests/providers/azure/services/defender/defender_service_test.py +++ b/tests/providers/azure/services/defender/defender_service_test.py @@ -58,7 +58,7 @@ def mock_defender_get_assessments(_): } -def mock_defender_get_security_contacts(*args, **kwargs): +def mock_defender_get_security_contacts(*_args, **_kwargs): from prowler.providers.azure.services.defender.defender_service import ( NotificationsByRole, ) @@ -447,6 +447,53 @@ class Test_Defender_get_iot_security_solutions: mock_client.iot_security_solution.list_by_subscription.assert_not_called() assert result[AZURE_SUBSCRIPTION_ID] == {} + def test_get_iot_security_solutions_with_multiple_resource_groups(self): + mock_client = MagicMock() + mock_client.iot_security_solution.list_by_resource_group.return_value = [] + + 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={}), + ): + defender = Defender(set_mocked_azure_provider()) + + defender.clients = {AZURE_SUBSCRIPTION_ID: mock_client} + defender.resource_groups = {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} + + result = defender._get_iot_security_solutions() + + assert mock_client.iot_security_solution.list_by_resource_group.call_count == 2 + assert AZURE_SUBSCRIPTION_ID in result + + def test_get_iot_security_solutions_with_mixed_case_resource_group(self): + mock_client = MagicMock() + mock_client.iot_security_solution.list_by_resource_group.return_value = [] + + 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={}), + ): + defender = Defender(set_mocked_azure_provider()) + + defender.clients = {AZURE_SUBSCRIPTION_ID: mock_client} + defender.resource_groups = {AZURE_SUBSCRIPTION_ID: ["RG"]} + + defender._get_iot_security_solutions() + + mock_client.iot_security_solution.list_by_resource_group.assert_called_once_with( + resource_group_name="RG" + ) + class Test_Defender_get_jit_policies: def test_get_jit_policies_no_resource_groups(self): @@ -522,55 +569,6 @@ class Test_Defender_get_jit_policies: mock_client.jit_network_access_policies.list.assert_not_called() assert result[AZURE_SUBSCRIPTION_ID] == {} - def test_get_iot_security_solutions_with_multiple_resource_groups(self): - mock_client = MagicMock() - mock_client.iot_security_solution.list_by_resource_group.return_value = [] - - 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={}), - ): - defender = Defender(set_mocked_azure_provider()) - - defender.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - defender.resource_groups = {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} - - result = defender._get_iot_security_solutions() - - assert mock_client.iot_security_solution.list_by_resource_group.call_count == 2 - assert AZURE_SUBSCRIPTION_ID in result - - def test_get_iot_security_solutions_with_mixed_case_resource_group(self): - mock_client = MagicMock() - mock_client.iot_security_solution.list_by_resource_group.return_value = [] - - 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={}), - ): - defender = Defender(set_mocked_azure_provider()) - - defender.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - defender.resource_groups = {AZURE_SUBSCRIPTION_ID: ["RG"]} - - defender._get_iot_security_solutions() - - mock_client.iot_security_solution.list_by_resource_group.assert_called_once_with( - resource_group_name="RG" - ) - - -class Test_Defender_get_jit_policies_extra: def test_get_jit_policies_with_multiple_resource_groups(self): mock_client = MagicMock() mock_client.jit_network_access_policies.list_by_resource_group.return_value = [] diff --git a/tests/providers/azure/services/iam/azure_iam_service_test.py b/tests/providers/azure/services/iam/azure_iam_service_test.py index 3f1dfec6fc..27004ae512 100644 --- a/tests/providers/azure/services/iam/azure_iam_service_test.py +++ b/tests/providers/azure/services/iam/azure_iam_service_test.py @@ -30,7 +30,9 @@ class Test_IAM_get_roles: builtin, custom = iam._get_roles() - mock_client.role_definitions.list.assert_called_once() + mock_client.role_definitions.list.assert_called_once_with( + scope=f"/subscriptions/{AZURE_SUBSCRIPTION_ID}" + ) assert AZURE_SUBSCRIPTION_ID in builtin assert AZURE_SUBSCRIPTION_ID in custom @@ -55,7 +57,9 @@ class Test_IAM_get_roles: builtin, custom = iam._get_roles() - mock_client.role_definitions.list.assert_called_once() + mock_client.role_definitions.list.assert_called_once_with( + scope=f"/subscriptions/{AZURE_SUBSCRIPTION_ID}" + ) assert AZURE_SUBSCRIPTION_ID in builtin assert AZURE_SUBSCRIPTION_ID in custom @@ -80,7 +84,9 @@ class Test_IAM_get_roles: builtin, custom = iam._get_roles() - mock_client.role_definitions.list.assert_called_once() + mock_client.role_definitions.list.assert_called_once_with( + scope=f"/subscriptions/{AZURE_SUBSCRIPTION_ID}" + ) assert AZURE_SUBSCRIPTION_ID in builtin assert AZURE_SUBSCRIPTION_ID in custom @@ -108,7 +114,9 @@ class Test_IAM_get_role_assignments: result = iam._get_role_assignments() - mock_client.role_assignments.list_for_subscription.assert_called_once() + mock_client.role_assignments.list_for_subscription.assert_called_once_with( + filter="atScope()" + ) assert AZURE_SUBSCRIPTION_ID in result def test_get_role_assignments_with_resource_group(self): @@ -133,7 +141,9 @@ class Test_IAM_get_role_assignments: result = iam._get_role_assignments() - mock_client.role_assignments.list_for_subscription.assert_called_once() + mock_client.role_assignments.list_for_subscription.assert_called_once_with( + filter="atScope()" + ) assert AZURE_SUBSCRIPTION_ID in result def test_get_role_assignments_empty_resource_group_for_subscription(self): @@ -158,5 +168,7 @@ class Test_IAM_get_role_assignments: result = iam._get_role_assignments() - mock_client.role_assignments.list_for_subscription.assert_called_once() + mock_client.role_assignments.list_for_subscription.assert_called_once_with( + filter="atScope()" + ) assert AZURE_SUBSCRIPTION_ID in result diff --git a/tests/providers/azure/services/network/network_service_test.py b/tests/providers/azure/services/network/network_service_test.py index 9a440b90cc..3239666b44 100644 --- a/tests/providers/azure/services/network/network_service_test.py +++ b/tests/providers/azure/services/network/network_service_test.py @@ -286,6 +286,73 @@ class Test_Network_get_security_groups: mock_client.network_security_groups.list_all.assert_not_called() assert result[AZURE_SUBSCRIPTION_ID] == [] + def test_get_security_groups_with_multiple_resource_groups(self): + mock_client = MagicMock() + mock_client.network_security_groups = MagicMock() + mock_client.network_security_groups.list.return_value = [] + + with ( + patch( + "prowler.providers.azure.services.network.network_service.Network._get_security_groups", + new=mock_network_get_security_groups, + ), + patch( + "prowler.providers.azure.services.network.network_service.Network._get_bastion_hosts", + new=mock_network_get_bastion_hosts, + ), + patch( + "prowler.providers.azure.services.network.network_service.Network._get_network_watchers", + new=mock_network_get_network_watchers, + ), + patch( + "prowler.providers.azure.services.network.network_service.Network._get_public_ip_addresses", + new=mock_network_get_public_ip_addresses, + ), + ): + network = Network(set_mocked_azure_provider()) + + network.clients = {AZURE_SUBSCRIPTION_ID: mock_client} + network.resource_groups = {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} + + result = network._get_security_groups() + + assert mock_client.network_security_groups.list.call_count == 2 + assert AZURE_SUBSCRIPTION_ID in result + + def test_get_security_groups_with_mixed_case_resource_group(self): + mock_client = MagicMock() + mock_client.network_security_groups = MagicMock() + mock_client.network_security_groups.list.return_value = [] + + with ( + patch( + "prowler.providers.azure.services.network.network_service.Network._get_security_groups", + new=mock_network_get_security_groups, + ), + patch( + "prowler.providers.azure.services.network.network_service.Network._get_bastion_hosts", + new=mock_network_get_bastion_hosts, + ), + patch( + "prowler.providers.azure.services.network.network_service.Network._get_network_watchers", + new=mock_network_get_network_watchers, + ), + patch( + "prowler.providers.azure.services.network.network_service.Network._get_public_ip_addresses", + new=mock_network_get_public_ip_addresses, + ), + ): + network = Network(set_mocked_azure_provider()) + + network.clients = {AZURE_SUBSCRIPTION_ID: mock_client} + network.resource_groups = {AZURE_SUBSCRIPTION_ID: ["RG"]} + + network._get_security_groups() + + mock_client.network_security_groups.list.assert_called_once_with( + resource_group_name="RG" + ) + class Test_Network_get_network_watchers: def test_get_network_watchers_no_resource_groups(self): @@ -600,73 +667,6 @@ class Test_Network_get_public_ip_addresses: mock_client.public_ip_addresses.list_all.assert_not_called() assert result[AZURE_SUBSCRIPTION_ID] == [] - def test_get_security_groups_with_multiple_resource_groups(self): - mock_client = MagicMock() - mock_client.network_security_groups = MagicMock() - mock_client.network_security_groups.list.return_value = [] - - with ( - patch( - "prowler.providers.azure.services.network.network_service.Network._get_security_groups", - new=mock_network_get_security_groups, - ), - patch( - "prowler.providers.azure.services.network.network_service.Network._get_bastion_hosts", - new=mock_network_get_bastion_hosts, - ), - patch( - "prowler.providers.azure.services.network.network_service.Network._get_network_watchers", - new=mock_network_get_network_watchers, - ), - patch( - "prowler.providers.azure.services.network.network_service.Network._get_public_ip_addresses", - new=mock_network_get_public_ip_addresses, - ), - ): - network = Network(set_mocked_azure_provider()) - - network.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - network.resource_groups = {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} - - result = network._get_security_groups() - - assert mock_client.network_security_groups.list.call_count == 2 - assert AZURE_SUBSCRIPTION_ID in result - - def test_get_security_groups_with_mixed_case_resource_group(self): - mock_client = MagicMock() - mock_client.network_security_groups = MagicMock() - mock_client.network_security_groups.list.return_value = [] - - with ( - patch( - "prowler.providers.azure.services.network.network_service.Network._get_security_groups", - new=mock_network_get_security_groups, - ), - patch( - "prowler.providers.azure.services.network.network_service.Network._get_bastion_hosts", - new=mock_network_get_bastion_hosts, - ), - patch( - "prowler.providers.azure.services.network.network_service.Network._get_network_watchers", - new=mock_network_get_network_watchers, - ), - patch( - "prowler.providers.azure.services.network.network_service.Network._get_public_ip_addresses", - new=mock_network_get_public_ip_addresses, - ), - ): - network = Network(set_mocked_azure_provider()) - - network.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - network.resource_groups = {AZURE_SUBSCRIPTION_ID: ["RG"]} - - network._get_security_groups() - - mock_client.network_security_groups.list.assert_called_once_with( - resource_group_name="RG" - ) - class Test_Network_get_network_watchers_extra: def test_get_network_watchers_with_multiple_resource_groups(self): diff --git a/tests/providers/azure/services/vm/vm_service_test.py b/tests/providers/azure/services/vm/vm_service_test.py index e25c6ffce3..b6bb583fbf 100644 --- a/tests/providers/azure/services/vm/vm_service_test.py +++ b/tests/providers/azure/services/vm/vm_service_test.py @@ -1,5 +1,7 @@ from unittest.mock import MagicMock, patch +import pytest + from prowler.providers.azure.services.vm.vm_service import ( Disk, LinuxConfiguration, @@ -112,6 +114,23 @@ def mock_vm_get_virtual_machines_with_linux(_): } +@pytest.fixture +def vm_service_factory(): + def _build(mock_client, resource_groups): + with ( + patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), + patch.object(VirtualMachines, "_get_disks", return_value={}), + patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), + ): + vm_service = VirtualMachines(set_mocked_azure_provider()) + + vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} + vm_service.resource_groups = resource_groups + return vm_service + + return _build + + @patch( "prowler.providers.azure.services.vm.vm_service.VirtualMachines._get_virtual_machines", new=mock_vm_get_virtual_machines, @@ -403,7 +422,7 @@ class Test_VirtualMachine_SecurityProfile_Validation: This tests the actual scenario where Azure SDK objects are converted """ - def mock_list_vms(*args, **kwargs): + def mock_list_vms(*_args, **_kwargs): # Simulate Azure SDK VM object with security_profile mock_vm = MagicMock() mock_vm.id = "/subscriptions/test/resourceGroups/test-rg/providers/Microsoft.Compute/virtualMachines/test-vm" @@ -470,20 +489,12 @@ class Test_VirtualMachine_SecurityProfile_Validation: class Test_VM_get_virtual_machines: - def test_get_virtual_machines_no_resource_groups(self): + def test_get_virtual_machines_no_resource_groups(self, vm_service_factory): mock_client = MagicMock() mock_client.virtual_machines = MagicMock() mock_client.virtual_machines.list_all.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = None + vm_service = vm_service_factory(mock_client, None) result = vm_service._get_virtual_machines() @@ -491,20 +502,14 @@ class Test_VM_get_virtual_machines: mock_client.virtual_machines.list.assert_not_called() assert AZURE_SUBSCRIPTION_ID in result - def test_get_virtual_machines_with_resource_group(self): + def test_get_virtual_machines_with_resource_group(self, vm_service_factory): mock_client = MagicMock() mock_client.virtual_machines = MagicMock() mock_client.virtual_machines.list.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: [RESOURCE_GROUP]} + vm_service = vm_service_factory( + mock_client, {AZURE_SUBSCRIPTION_ID: [RESOURCE_GROUP]} + ) result = vm_service._get_virtual_machines() @@ -514,19 +519,13 @@ class Test_VM_get_virtual_machines: mock_client.virtual_machines.list_all.assert_not_called() assert AZURE_SUBSCRIPTION_ID in result - def test_get_virtual_machines_empty_resource_group_for_subscription(self): + def test_get_virtual_machines_empty_resource_group_for_subscription( + self, vm_service_factory + ): mock_client = MagicMock() mock_client.virtual_machines = MagicMock() - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: []} + vm_service = vm_service_factory(mock_client, {AZURE_SUBSCRIPTION_ID: []}) result = vm_service._get_virtual_machines() @@ -534,22 +533,45 @@ class Test_VM_get_virtual_machines: mock_client.virtual_machines.list_all.assert_not_called() assert result[AZURE_SUBSCRIPTION_ID] == {} + def test_get_virtual_machines_with_multiple_resource_groups( + self, vm_service_factory + ): + mock_client = MagicMock() + mock_client.virtual_machines = MagicMock() + mock_client.virtual_machines.list.return_value = [] + + vm_service = vm_service_factory( + mock_client, {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} + ) + + result = vm_service._get_virtual_machines() + + assert mock_client.virtual_machines.list.call_count == 2 + assert AZURE_SUBSCRIPTION_ID in result + + def test_get_virtual_machines_with_mixed_case_resource_group( + self, vm_service_factory + ): + mock_client = MagicMock() + mock_client.virtual_machines = MagicMock() + mock_client.virtual_machines.list.return_value = [] + + vm_service = vm_service_factory(mock_client, {AZURE_SUBSCRIPTION_ID: ["RG"]}) + + vm_service._get_virtual_machines() + + mock_client.virtual_machines.list.assert_called_once_with( + resource_group_name="RG" + ) + class Test_VM_get_disks: - def test_get_disks_no_resource_groups(self): + def test_get_disks_no_resource_groups(self, vm_service_factory): mock_client = MagicMock() mock_client.disks = MagicMock() mock_client.disks.list.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = None + vm_service = vm_service_factory(mock_client, None) result = vm_service._get_disks() @@ -557,20 +579,14 @@ class Test_VM_get_disks: mock_client.disks.list_by_resource_group.assert_not_called() assert AZURE_SUBSCRIPTION_ID in result - def test_get_disks_with_resource_group(self): + def test_get_disks_with_resource_group(self, vm_service_factory): mock_client = MagicMock() mock_client.disks = MagicMock() mock_client.disks.list_by_resource_group.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: [RESOURCE_GROUP]} + vm_service = vm_service_factory( + mock_client, {AZURE_SUBSCRIPTION_ID: [RESOURCE_GROUP]} + ) result = vm_service._get_disks() @@ -580,19 +596,11 @@ class Test_VM_get_disks: mock_client.disks.list.assert_not_called() assert AZURE_SUBSCRIPTION_ID in result - def test_get_disks_empty_resource_group_for_subscription(self): + def test_get_disks_empty_resource_group_for_subscription(self, vm_service_factory): mock_client = MagicMock() mock_client.disks = MagicMock() - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: []} + vm_service = vm_service_factory(mock_client, {AZURE_SUBSCRIPTION_ID: []}) result = vm_service._get_disks() @@ -602,20 +610,12 @@ class Test_VM_get_disks: class Test_VM_get_vm_scale_sets: - def test_get_vm_scale_sets_no_resource_groups(self): + def test_get_vm_scale_sets_no_resource_groups(self, vm_service_factory): mock_client = MagicMock() mock_client.virtual_machine_scale_sets = MagicMock() mock_client.virtual_machine_scale_sets.list_all.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = None + vm_service = vm_service_factory(mock_client, None) result = vm_service._get_vm_scale_sets() @@ -623,20 +623,14 @@ class Test_VM_get_vm_scale_sets: mock_client.virtual_machine_scale_sets.list.assert_not_called() assert AZURE_SUBSCRIPTION_ID in result - def test_get_vm_scale_sets_with_resource_group(self): + def test_get_vm_scale_sets_with_resource_group(self, vm_service_factory): mock_client = MagicMock() mock_client.virtual_machine_scale_sets = MagicMock() mock_client.virtual_machine_scale_sets.list.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: [RESOURCE_GROUP]} + vm_service = vm_service_factory( + mock_client, {AZURE_SUBSCRIPTION_ID: [RESOURCE_GROUP]} + ) result = vm_service._get_vm_scale_sets() @@ -646,19 +640,13 @@ class Test_VM_get_vm_scale_sets: mock_client.virtual_machine_scale_sets.list_all.assert_not_called() assert AZURE_SUBSCRIPTION_ID in result - def test_get_vm_scale_sets_empty_resource_group_for_subscription(self): + def test_get_vm_scale_sets_empty_resource_group_for_subscription( + self, vm_service_factory + ): mock_client = MagicMock() mock_client.virtual_machine_scale_sets = MagicMock() - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: []} + vm_service = vm_service_factory(mock_client, {AZURE_SUBSCRIPTION_ID: []}) result = vm_service._get_vm_scale_sets() @@ -666,83 +654,28 @@ class Test_VM_get_vm_scale_sets: mock_client.virtual_machine_scale_sets.list_all.assert_not_called() assert result[AZURE_SUBSCRIPTION_ID] == {} - def test_get_virtual_machines_with_multiple_resource_groups(self): - mock_client = MagicMock() - mock_client.virtual_machines = MagicMock() - mock_client.virtual_machines.list.return_value = [] - - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} - - result = vm_service._get_virtual_machines() - - assert mock_client.virtual_machines.list.call_count == 2 - assert AZURE_SUBSCRIPTION_ID in result - - def test_get_virtual_machines_with_mixed_case_resource_group(self): - mock_client = MagicMock() - mock_client.virtual_machines = MagicMock() - mock_client.virtual_machines.list.return_value = [] - - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: ["RG"]} - - vm_service._get_virtual_machines() - - mock_client.virtual_machines.list.assert_called_once_with( - resource_group_name="RG" - ) - class Test_VM_get_disks_extra: - def test_get_disks_with_multiple_resource_groups(self): + def test_get_disks_with_multiple_resource_groups(self, vm_service_factory): mock_client = MagicMock() mock_client.disks = MagicMock() mock_client.disks.list_by_resource_group.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} + vm_service = vm_service_factory( + mock_client, {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} + ) result = vm_service._get_disks() assert mock_client.disks.list_by_resource_group.call_count == 2 assert AZURE_SUBSCRIPTION_ID in result - def test_get_disks_with_mixed_case_resource_group(self): + def test_get_disks_with_mixed_case_resource_group(self, vm_service_factory): mock_client = MagicMock() mock_client.disks = MagicMock() mock_client.disks.list_by_resource_group.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: ["RG"]} + vm_service = vm_service_factory(mock_client, {AZURE_SUBSCRIPTION_ID: ["RG"]}) vm_service._get_disks() @@ -752,40 +685,26 @@ class Test_VM_get_disks_extra: class Test_VM_get_vm_scale_sets_extra: - def test_get_vm_scale_sets_with_multiple_resource_groups(self): + def test_get_vm_scale_sets_with_multiple_resource_groups(self, vm_service_factory): mock_client = MagicMock() mock_client.virtual_machine_scale_sets = MagicMock() mock_client.virtual_machine_scale_sets.list.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} + vm_service = vm_service_factory( + mock_client, {AZURE_SUBSCRIPTION_ID: RESOURCE_GROUP_LIST} + ) result = vm_service._get_vm_scale_sets() assert mock_client.virtual_machine_scale_sets.list.call_count == 2 assert AZURE_SUBSCRIPTION_ID in result - def test_get_vm_scale_sets_with_mixed_case_resource_group(self): + def test_get_vm_scale_sets_with_mixed_case_resource_group(self, vm_service_factory): mock_client = MagicMock() mock_client.virtual_machine_scale_sets = MagicMock() mock_client.virtual_machine_scale_sets.list.return_value = [] - with ( - patch.object(VirtualMachines, "_get_virtual_machines", return_value={}), - patch.object(VirtualMachines, "_get_disks", return_value={}), - patch.object(VirtualMachines, "_get_vm_scale_sets", return_value={}), - ): - vm_service = VirtualMachines(set_mocked_azure_provider()) - - vm_service.clients = {AZURE_SUBSCRIPTION_ID: mock_client} - vm_service.resource_groups = {AZURE_SUBSCRIPTION_ID: ["RG"]} + vm_service = vm_service_factory(mock_client, {AZURE_SUBSCRIPTION_ID: ["RG"]}) vm_service._get_vm_scale_sets()