fix(azure): refine resource group scoped scans

This commit is contained in:
Hugo P.Brito
2026-07-02 11:23:58 +01:00
parent cd90a91158
commit 8dbb97849e
11 changed files with 342 additions and 366 deletions
+17 -3
View File
@@ -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
+20 -1
View File
@@ -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, [])
@@ -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(
{
@@ -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
@@ -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
@@ -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 = []
@@ -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 = []
@@ -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 = []
@@ -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
@@ -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):
@@ -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()