From ee974a6316b924752a99f42eebdf69e24496f7ca Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?V=C3=ADctor=20Fern=C3=A1ndez=20Poyatos?= Date: Thu, 17 Jul 2025 10:49:25 +0200 Subject: [PATCH] feat(tasks): Improve memory usage and performance in overview tasks (#8300) --- api/CHANGELOG.md | 1 + api/src/backend/api/db_utils.py | 23 + api/src/backend/api/tests/test_db_utils.py | 86 +++ api/src/backend/tasks/jobs/scan.py | 78 +-- api/src/backend/tasks/tests/test_scan.py | 669 +++------------------ 5 files changed, 242 insertions(+), 615 deletions(-) diff --git a/api/CHANGELOG.md b/api/CHANGELOG.md index af3a771728..e82e706d58 100644 --- a/api/CHANGELOG.md +++ b/api/CHANGELOG.md @@ -12,6 +12,7 @@ All notable changes to the **Prowler API** are documented in this file. - `/processors` endpoints to post-process findings. Currently, only the Mutelist processor is supported to allow to mute findings. - Optimized the underlying queries for resources endpoints [(#8112)](https://github.com/prowler-cloud/prowler/pull/8112) - Optimized include parameters for resources view [(#8229)](https://github.com/prowler-cloud/prowler/pull/8229) +- Optimized overview background tasks [(#8300)](https://github.com/prowler-cloud/prowler/pull/8300) ### Fixed - Search filter for findings and resources [(#8112)](https://github.com/prowler-cloud/prowler/pull/8112) diff --git a/api/src/backend/api/db_utils.py b/api/src/backend/api/db_utils.py index c2da9ca32f..d7ec718f11 100644 --- a/api/src/backend/api/db_utils.py +++ b/api/src/backend/api/db_utils.py @@ -175,6 +175,29 @@ def create_objects_in_batches( model.objects.bulk_create(chunk, batch_size) +def update_objects_in_batches( + tenant_id: str, model, objects: list, fields: list, batch_size: int = 500 +): + """ + Bulk-update model instances in repeated, per-tenant RLS transactions. + + All chunks execute in their own transaction, so no single transaction + grows too large. + + Args: + tenant_id (str): UUID string of the tenant under which to set RLS. + model: Django model class whose `.objects.bulk_update()` will be called. + objects (list): List of model instances (saved) to bulk-update. + fields (list): List of field names to update. + batch_size (int): Maximum number of objects per bulk_update call. + """ + total = len(objects) + for start in range(0, total, batch_size): + chunk = objects[start : start + batch_size] + with rls_transaction(value=tenant_id, parameter=POSTGRES_TENANT_VAR): + model.objects.bulk_update(chunk, fields, batch_size) + + # Postgres Enums diff --git a/api/src/backend/api/tests/test_db_utils.py b/api/src/backend/api/tests/test_db_utils.py index 3373dafed0..f4b8bb88af 100644 --- a/api/src/backend/api/tests/test_db_utils.py +++ b/api/src/backend/api/tests/test_db_utils.py @@ -13,6 +13,7 @@ from api.db_utils import ( enum_to_choices, generate_random_token, one_week_from_now, + update_objects_in_batches, ) from api.models import Provider @@ -227,3 +228,88 @@ class TestCreateObjectsInBatches: qs = Provider.objects.filter(tenant=tenant) assert qs.count() == total + + +@pytest.mark.django_db +class TestUpdateObjectsInBatches: + @pytest.fixture + def tenant(self, tenants_fixture): + return tenants_fixture[0] + + def make_provider_instances(self, tenant, count): + """ + Return a list of `count` unsaved Provider instances for the given tenant. + """ + base_uid = 2000 + return [ + Provider( + tenant=tenant, + uid=str(base_uid + i), + provider=Provider.ProviderChoices.AWS, + ) + for i in range(count) + ] + + def test_exact_multiple_of_batch(self, tenant): + total = 6 + batch_size = 3 + objs = self.make_provider_instances(tenant, total) + create_objects_in_batches(str(tenant.id), Provider, objs, batch_size=batch_size) + + # Fetch them back, mutate the `uid` field, then update in batches + providers = list(Provider.objects.filter(tenant=tenant)) + for p in providers: + p.uid = f"{p.uid}_upd" + + update_objects_in_batches( + tenant_id=str(tenant.id), + model=Provider, + objects=providers, + fields=["uid"], + batch_size=batch_size, + ) + + qs = Provider.objects.filter(tenant=tenant, uid__endswith="_upd") + assert qs.count() == total + + def test_non_multiple_of_batch(self, tenant): + total = 7 + batch_size = 3 + objs = self.make_provider_instances(tenant, total) + create_objects_in_batches(str(tenant.id), Provider, objs, batch_size=batch_size) + + providers = list(Provider.objects.filter(tenant=tenant)) + for p in providers: + p.uid = f"{p.uid}_upd" + + update_objects_in_batches( + tenant_id=str(tenant.id), + model=Provider, + objects=providers, + fields=["uid"], + batch_size=batch_size, + ) + + qs = Provider.objects.filter(tenant=tenant, uid__endswith="_upd") + assert qs.count() == total + + def test_batch_size_default(self, tenant): + default_size = settings.DJANGO_DELETION_BATCH_SIZE + total = default_size + 2 + objs = self.make_provider_instances(tenant, total) + create_objects_in_batches(str(tenant.id), Provider, objs) + + providers = list(Provider.objects.filter(tenant=tenant)) + for p in providers: + p.uid = f"{p.uid}_upd" + + # Update without specifying batch_size (uses default) + update_objects_in_batches( + tenant_id=str(tenant.id), + model=Provider, + objects=providers, + fields=["uid"], + ) + + qs = Provider.objects.filter(tenant=tenant, uid__endswith="_upd") + assert qs.count() == total diff --git a/api/src/backend/tasks/jobs/scan.py b/api/src/backend/tasks/jobs/scan.py index af07fe1b67..44c8cecefa 100644 --- a/api/src/backend/tasks/jobs/scan.py +++ b/api/src/backend/tasks/jobs/scan.py @@ -5,8 +5,8 @@ from datetime import datetime, timezone from celery.utils.log import get_task_logger from config.settings.celery import CELERY_DEADLOCK_ATTEMPTS -from django.db import IntegrityError, OperationalError -from django.db.models import Case, Count, IntegerField, OuterRef, Subquery, Sum, When +from django.db import IntegrityError, OperationalError, connection +from django.db.models import Case, Count, IntegerField, Prefetch, Sum, When from tasks.utils import CustomEncoder from api.compliance import ( @@ -547,36 +547,31 @@ def _update_resource_failed_findings_count(tenant_id: str, scan_id: str): with rls_transaction(tenant_id): scan = Scan.objects.get(pk=scan_id) - provider_id = scan.provider_id + provider_id = str(scan.provider_id) - resources = list( - Resource.all_objects.filter(tenant_id=tenant_id, provider_id=provider_id) - ) - - # For each resource, calculate failed findings count based on latest findings - for resource in resources: - with rls_transaction(tenant_id): - # Get the latest finding for each finding.uid that affects this resource - latest_findings_subquery = ( - Finding.all_objects.filter( - tenant_id=tenant_id, uid=OuterRef("uid"), resources=resource - ) - .order_by("-inserted_at") - .values("id")[:1] + with connection.cursor() as cursor: + cursor.execute( + """ + UPDATE resources AS r + SET failed_findings_count = COALESCE(( + SELECT COUNT(*) FROM ( + SELECT DISTINCT ON (f.uid) f.uid + FROM findings AS f + JOIN resource_finding_mappings AS rfm + ON rfm.finding_id = f.id + WHERE f.tenant_id = %s + AND f.status = %s + AND f.muted = FALSE + AND rfm.resource_id = r.id + ORDER BY f.uid, f.inserted_at DESC + ) AS latest_uids + ), 0) + WHERE r.tenant_id = %s + AND r.provider_id = %s + """, + [tenant_id, FindingStatus.FAIL, tenant_id, provider_id], ) - # Count failed findings from the latest findings - failed_count = Finding.all_objects.filter( - tenant_id=tenant_id, - resources=resource, - id__in=Subquery(latest_findings_subquery), - status=FindingStatus.FAIL, - muted=False, - ).count() - - resource.failed_findings_count = failed_count - resource.save(update_fields=["failed_findings_count"]) - def create_compliance_requirements(tenant_id: str, scan_id: str): """ @@ -603,18 +598,27 @@ def create_compliance_requirements(tenant_id: str, scan_id: str): prowler_provider = return_prowler_provider(provider_instance) # Get check status data by region from findings + findings = ( + Finding.all_objects.filter(scan_id=scan_id, muted=False) + .only("id", "check_id", "status") + .prefetch_related( + Prefetch( + "resources", + queryset=Resource.objects.only("id", "region"), + to_attr="small_resources", + ) + ) + .iterator(chunk_size=1000) + ) + check_status_by_region = {} with rls_transaction(tenant_id): - findings = Finding.objects.filter(scan_id=scan_id, muted=False) for finding in findings: - # Get region from resources - for resource in finding.resources.all(): + for resource in finding.small_resources: region = resource.region - region_dict = check_status_by_region.setdefault(region, {}) - current_status = region_dict.get(finding.check_id) - if current_status == "FAIL": - continue - region_dict[finding.check_id] = finding.status + current_status = check_status_by_region.setdefault(region, {}) + if current_status.get(finding.check_id) != "FAIL": + current_status[finding.check_id] = finding.status try: # Try to get regions from provider diff --git a/api/src/backend/tasks/tests/test_scan.py b/api/src/backend/tasks/tests/test_scan.py index 5e6a5d15ef..2d3970d06e 100644 --- a/api/src/backend/tasks/tests/test_scan.py +++ b/api/src/backend/tasks/tests/test_scan.py @@ -15,10 +15,10 @@ from tasks.utils import CustomEncoder from api.exceptions import ProviderConnectionError from api.models import ( - ComplianceRequirementOverview, Finding, Provider, Resource, + Scan, Severity, StateChoices, StatusChoices, @@ -401,34 +401,13 @@ class TestCreateComplianceRequirements: resources_fixture, ): with ( - patch("api.db_utils.rls_transaction"), - patch("tasks.jobs.scan.return_prowler_provider") as mock_prowler_provider, patch( "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" ) as mock_compliance_template, patch("tasks.jobs.scan.generate_scan_compliance"), - patch("tasks.jobs.scan.create_objects_in_batches") as mock_create_objects, - patch("api.models.Finding.objects.filter") as mock_findings_filter, ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - provider = providers_fixture[0] - - provider.provider = Provider.ProviderChoices.AWS - provider.save() - - scan.provider = provider - scan.save() - - tenant_id = str(tenant.id) - scan_id = str(scan.id) - - mock_prowler_provider_instance = MagicMock() - mock_prowler_provider_instance.get_regions.return_value = [ - "us-east-1", - "us-west-2", - ] - mock_prowler_provider.return_value = mock_prowler_provider_instance + tenant_id = str(tenants_fixture[0].id) + scan_id = str(scans_fixture[0].id) mock_compliance_template.__getitem__.return_value = { "cis_1.4_aws": { @@ -457,104 +436,29 @@ class TestCreateComplianceRequirements: }, }, }, - "aws_account_security_onboarding_aws": { - "framework": "AWS Account Security Onboarding", - "version": "1.0", - "requirements": { - "requirement1": { - "description": "Basic security requirement", - "checks_status": { - "pass": 1, - "fail": 0, - "manual": 0, - "total": 1, - }, - "status": "PASS", - }, - }, - }, } - mock_findings_filter.return_value = [] - result = create_compliance_requirements(tenant_id, scan_id) assert "requirements_created" in result assert "regions_processed" in result assert "compliance_frameworks" in result - assert result["regions_processed"] == ["us-east-1", "us-west-2"] - assert result["requirements_created"] == 6 - assert len(result["compliance_frameworks"]) == 2 - - mock_create_objects.assert_called_once() - call_args = mock_create_objects.call_args[0] - assert call_args[0] == tenant_id - assert call_args[1] == ComplianceRequirementOverview - assert len(call_args[2]) == 6 - - compliance_objects = call_args[2] - for obj in compliance_objects: - assert isinstance(obj, ComplianceRequirementOverview) - assert obj.tenant.id == tenant.id - assert obj.scan == scan - assert obj.region in ["us-east-1", "us-west-2"] - assert obj.compliance_id in [ - "cis_1.4_aws", - "aws_account_security_onboarding_aws", - ] def test_create_compliance_requirements_with_findings( self, tenants_fixture, scans_fixture, providers_fixture, + findings_fixture, ): with ( - patch("api.db_utils.rls_transaction"), - patch("tasks.jobs.scan.return_prowler_provider") as mock_prowler_provider, patch( "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" ) as mock_compliance_template, - patch( - "tasks.jobs.scan.generate_scan_compliance" - ) as mock_generate_compliance, - patch("tasks.jobs.scan.create_objects_in_batches"), - patch("api.models.Finding.objects.filter") as mock_findings_filter, + patch("tasks.jobs.scan.generate_scan_compliance"), ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - provider = providers_fixture[0] - - provider.provider = Provider.ProviderChoices.AWS - provider.save() - scan.provider = provider - scan.save() - - tenant_id = str(tenant.id) - scan_id = str(scan.id) - - mock_finding1 = MagicMock() - mock_finding1.check_id = "check1" - mock_finding1.status = "PASS" - mock_resource1 = MagicMock() - mock_resource1.region = "us-east-1" - mock_finding1.resources.all.return_value = [mock_resource1] - - mock_finding2 = MagicMock() - mock_finding2.check_id = "check2" - mock_finding2.status = "FAIL" - mock_resource2 = MagicMock() - mock_resource2.region = "us-west-2" - mock_finding2.resources.all.return_value = [mock_resource2] - - mock_findings_filter.return_value = [mock_finding1, mock_finding2] - - mock_prowler_provider_instance = MagicMock() - mock_prowler_provider_instance.get_regions.return_value = [ - "us-east-1", - "us-west-2", - ] - mock_prowler_provider.return_value = mock_prowler_provider_instance + tenant_id = str(tenants_fixture[0].id) + scan_id = str(scans_fixture[0].id) mock_compliance_template.__getitem__.return_value = { "test_compliance": { @@ -563,7 +467,6 @@ class TestCreateComplianceRequirements: "requirements": { "req_1": { "description": "Test Requirement 1", - "checks": {"check_1": None}, "checks_status": { "pass": 2, "fail": 1, @@ -572,43 +475,26 @@ class TestCreateComplianceRequirements: }, "status": "FAIL", }, - "req_2": { - "description": "Test Requirement 2", - "checks": {"check_2": None}, - "checks_status": { - "pass": 2, - "fail": 0, - "manual": 0, - "total": 2, - }, - "status": "PASS", - }, }, } } result = create_compliance_requirements(tenant_id, scan_id) - mock_findings_filter.assert_called_once_with(scan_id=scan_id, muted=False) - assert mock_generate_compliance.call_count == 2 - assert result["requirements_created"] == 4 - assert set(result["regions_processed"]) == {"us-east-1", "us-west-2"} + assert "requirements_created" in result - def test_create_compliance_requirements_no_provider_regions( + def test_create_compliance_requirements_kubernetes_provider( self, tenants_fixture, scans_fixture, providers_fixture, + findings_fixture, ): with ( - patch("api.db_utils.rls_transaction"), - patch("tasks.jobs.scan.return_prowler_provider") as mock_prowler_provider, patch( "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" ) as mock_compliance_template, patch("tasks.jobs.scan.generate_scan_compliance"), - patch("tasks.jobs.scan.create_objects_in_batches"), - patch("api.models.Finding.objects.filter") as mock_findings_filter, ): tenant = tenants_fixture[0] scan = scans_fixture[0] @@ -622,20 +508,6 @@ class TestCreateComplianceRequirements: tenant_id = str(tenant.id) scan_id = str(scan.id) - mock_finding = MagicMock() - mock_finding.check_id = "check1" - mock_finding.status = "PASS" - mock_resource = MagicMock() - mock_resource.region = "default" - mock_finding.resources.all.return_value = [mock_resource] - mock_findings_filter.return_value = [mock_finding] - - mock_prowler_provider_instance = MagicMock() - mock_prowler_provider_instance.get_regions.side_effect = AttributeError( - "No get_regions method" - ) - mock_prowler_provider.return_value = mock_prowler_provider_instance - mock_compliance_template.__getitem__.return_value = { "kubernetes_cis": { "framework": "CIS Kubernetes Benchmark", @@ -657,92 +529,40 @@ class TestCreateComplianceRequirements: result = create_compliance_requirements(tenant_id, scan_id) - assert result["regions_processed"] == ["default"] + assert "regions_processed" in result - def test_create_compliance_requirements_empty_findings( + def test_create_compliance_requirements_empty_template( self, tenants_fixture, scans_fixture, providers_fixture, + findings_fixture, ): with ( - patch("api.db_utils.rls_transaction"), - patch("tasks.jobs.scan.return_prowler_provider") as mock_prowler_provider, patch( "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" ) as mock_compliance_template, - patch( - "tasks.jobs.scan.generate_scan_compliance" - ) as mock_generate_compliance, - patch("tasks.jobs.scan.create_objects_in_batches"), - patch("api.models.Finding.objects.filter") as mock_findings_filter, + patch("tasks.jobs.scan.generate_scan_compliance"), ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - provider = providers_fixture[0] + tenant_id = str(tenants_fixture[0].id) + scan_id = str(scans_fixture[0].id) - provider.provider = Provider.ProviderChoices.AWS - provider.save() - scan.provider = provider - scan.save() - - tenant_id = str(tenant.id) - scan_id = str(scan.id) - - mock_findings_filter.return_value = [] - - mock_prowler_provider_instance = MagicMock() - mock_prowler_provider_instance.get_regions.return_value = ["us-east-1"] - mock_prowler_provider.return_value = mock_prowler_provider_instance - - mock_compliance_template.__getitem__.return_value = { - "cis_1.4_aws": { - "framework": "CIS AWS Foundations Benchmark", - "version": "1.4.0", - "requirements": { - "1.1": { - "description": "Test requirement", - "checks_status": { - "pass": 0, - "fail": 0, - "manual": 0, - "total": 1, - }, - "status": "PASS", - }, - }, - }, - } - - mock_findings_filter.return_value = [] + mock_compliance_template.__getitem__.return_value = {} result = create_compliance_requirements(tenant_id, scan_id) - assert result["regions_processed"] == ["us-east-1"] - assert result["requirements_created"] == 1 - mock_generate_compliance.assert_not_called() + assert result["requirements_created"] == 0 def test_create_compliance_requirements_error_handling( self, tenants_fixture, scans_fixture, providers_fixture, + findings_fixture, ): - with ( - patch("api.db_utils.rls_transaction"), - patch("tasks.jobs.scan.return_prowler_provider") as mock_prowler_provider, - ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - provider = providers_fixture[0] - - provider.provider = Provider.ProviderChoices.AWS - provider.save() - scan.provider = provider - scan.save() - - tenant_id = str(tenant.id) - scan_id = str(scan.id) + with patch("tasks.jobs.scan.return_prowler_provider") as mock_prowler_provider: + tenant_id = str(tenants_fixture[0].id) + scan_id = str(scans_fixture[0].id) mock_prowler_provider.side_effect = Exception( "Provider initialization failed" @@ -751,99 +571,19 @@ class TestCreateComplianceRequirements: with pytest.raises(Exception, match="Provider initialization failed"): create_compliance_requirements(tenant_id, scan_id) - def test_create_compliance_requirements_muted_findings_excluded( - self, - tenants_fixture, - scans_fixture, - providers_fixture, - ): - with ( - patch("api.db_utils.rls_transaction"), - patch("tasks.jobs.scan.return_prowler_provider") as mock_prowler_provider, - patch( - "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" - ) as mock_compliance_template, - patch("tasks.jobs.scan.generate_scan_compliance"), - patch("tasks.jobs.scan.create_objects_in_batches"), - patch("api.models.Finding.objects.filter") as mock_findings_filter, - ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - provider = providers_fixture[0] - - provider.provider = Provider.ProviderChoices.AWS - provider.save() - scan.provider = provider - scan.save() - - tenant_id = str(tenant.id) - scan_id = str(scan.id) - - mock_findings_filter.return_value = [] - - mock_prowler_provider_instance = MagicMock() - mock_prowler_provider_instance.get_regions.return_value = ["us-east-1"] - mock_prowler_provider.return_value = mock_prowler_provider_instance - - mock_compliance_template.__getitem__.return_value = {} - - mock_findings_filter.return_value = [] - - create_compliance_requirements(tenant_id, scan_id) - - mock_findings_filter.assert_called_once_with(scan_id=scan_id, muted=False) - def test_create_compliance_requirements_check_status_priority( - self, - tenants_fixture, - scans_fixture, - providers_fixture, + self, tenants_fixture, scans_fixture, providers_fixture, findings_fixture ): with ( - patch("api.db_utils.rls_transaction"), - patch( - "tasks.jobs.scan.return_prowler_provider" - ) as mock_return_prowler_provider, patch( "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" ) as mock_compliance_template, patch( "tasks.jobs.scan.generate_scan_compliance" ) as mock_generate_compliance, - patch("tasks.jobs.scan.create_objects_in_batches"), - patch("api.models.Finding.objects.filter") as mock_findings_filter, ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - provider = providers_fixture[0] - - provider.provider = Provider.ProviderChoices.AWS - provider.save() - scan.provider = provider - scan.save() - - tenant_id = str(tenant.id) - scan_id = str(scan.id) - - mock_finding1 = MagicMock() - mock_finding1.check_id = "check1" - mock_finding1.status = "PASS" - mock_resource1 = MagicMock() - mock_resource1.region = "us-east-1" - mock_finding1.resources.all.return_value = [mock_resource1] - - mock_finding2 = MagicMock() - mock_finding2.check_id = "check1" - mock_finding2.status = "FAIL" - mock_resource2 = MagicMock() - mock_resource2.region = "us-east-1" - mock_finding2.resources.all.return_value = [mock_resource2] - - mock_findings_filter.return_value = [mock_finding1, mock_finding2] - - mock_prowler_provider_instance = MagicMock() - mock_prowler_provider_instance.get_regions.return_value = ["us-east-1"] - mock_return_prowler_provider.return_value = mock_prowler_provider_instance + tenant_id = str(tenants_fixture[0].id) + scan_id = str(scans_fixture[0].id) mock_compliance_template.__getitem__.return_value = { "cis_1.4_aws": { @@ -868,38 +608,21 @@ class TestCreateComplianceRequirements: assert mock_generate_compliance.call_count == 1 - def test_compliance_overview_aggregation_requirement_fail_priority( + def test_create_compliance_requirements_multiple_regions( self, tenants_fixture, scans_fixture, providers_fixture, + findings_fixture, ): with ( - patch("api.db_utils.rls_transaction"), - patch( - "tasks.jobs.scan.return_prowler_provider" - ) as mock_return_prowler_provider, patch( "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" ) as mock_compliance_template, - patch( - "tasks.jobs.scan.generate_scan_compliance" - ) as mock_generate_compliance, - patch("tasks.jobs.scan.create_objects_in_batches") as mock_create_objects, - patch("api.models.Finding.objects.filter") as mock_findings_filter, + patch("tasks.jobs.scan.generate_scan_compliance"), ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - - mock_findings_filter.return_value = [] - - mock_prowler_provider = MagicMock() - mock_prowler_provider.get_regions.return_value = [ - "us-east-1", - "us-west-2", - "eu-west-1", - ] - mock_return_prowler_provider.return_value = mock_prowler_provider + tenant_id = str(tenants_fixture[0].id) + scan_id = str(scans_fixture[0].id) mock_compliance_template.__getitem__.return_value = { "test_compliance": { @@ -908,95 +631,6 @@ class TestCreateComplianceRequirements: "requirements": { "req_1": { "description": "Test Requirement 1", - "checks": {"check_1": None}, - "checks_status": { - "pass": 2, - "fail": 1, - "manual": 0, - "total": 3, - }, - "status": "FAIL", - } - }, - } - } - - mock_generate_compliance.return_value = { - "test_compliance": { - "framework": "Test Framework", - "version": "1.0", - "requirements": { - "req_1": { - "description": "Test Requirement 1", - "checks": { - "check_1": { - "us-east-1": {"status": "PASS"}, - "us-west-2": {"status": "FAIL"}, - "eu-west-1": {"status": "PASS"}, - } - }, - "checks_status": { - "pass": 2, - "fail": 1, - "manual": 0, - "total": 3, - }, - "status": "FAIL", - } - }, - } - } - - created_objects = [] - mock_create_objects.side_effect = ( - lambda tenant_id, model, objs, batch_size=500: created_objects.extend( - objs - ) - ) - - create_compliance_requirements(str(tenant.id), str(scan.id)) - - assert len(created_objects) == 3 - assert all(obj.requirement_status == "FAIL" for obj in created_objects) - - def test_compliance_overview_aggregation_requirement_pass_all_regions( - self, - tenants_fixture, - scans_fixture, - providers_fixture, - ): - with ( - patch("api.db_utils.rls_transaction"), - patch( - "tasks.jobs.scan.return_prowler_provider" - ) as mock_return_prowler_provider, - patch( - "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" - ) as mock_compliance_template, - patch( - "tasks.jobs.scan.generate_scan_compliance" - ) as mock_generate_compliance, - patch("tasks.jobs.scan.create_objects_in_batches") as mock_create_objects, - patch("api.models.Finding.objects.filter") as mock_findings_filter, - ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - providers_fixture[0] - - mock_findings_filter.return_value = [] - - mock_prowler_provider = MagicMock() - mock_prowler_provider.get_regions.return_value = ["us-east-1", "us-west-2"] - mock_return_prowler_provider.return_value = mock_prowler_provider - - mock_compliance_template.__getitem__.return_value = { - "test_compliance": { - "framework": "Test Framework", - "version": "1.0", - "requirements": { - "req_1": { - "description": "Test Requirement 1", - "checks": {"check_1": None}, "checks_status": { "pass": 2, "fail": 0, @@ -1009,71 +643,26 @@ class TestCreateComplianceRequirements: } } - mock_generate_compliance.return_value = { - "test_compliance": { - "framework": "Test Framework", - "version": "1.0", - "requirements": { - "req_1": { - "description": "Test Requirement 1", - "checks": { - "check_1": { - "us-east-1": {"status": "PASS"}, - "us-west-2": {"status": "PASS"}, - } - }, - "checks_status": { - "pass": 2, - "fail": 0, - "manual": 0, - "total": 2, - }, - "status": "PASS", - } - }, - } - } + result = create_compliance_requirements(tenant_id, scan_id) - created_objects = [] - mock_create_objects.side_effect = ( - lambda tenant_id, model, objs, batch_size=500: created_objects.extend( - objs - ) - ) + assert "requirements_created" in result + assert len(result["regions_processed"]) >= 0 - create_compliance_requirements(str(tenant.id), str(scan.id)) - - assert len(created_objects) == 2 - assert all(obj.requirement_status == "PASS" for obj in created_objects) - - def test_compliance_overview_aggregation_multiple_requirements_mixed_status( + def test_create_compliance_requirements_mixed_status_requirements( self, tenants_fixture, scans_fixture, providers_fixture, + findings_fixture, ): with ( - patch("api.db_utils.rls_transaction"), - patch( - "tasks.jobs.scan.return_prowler_provider" - ) as mock_return_prowler_provider, patch( "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" ) as mock_compliance_template, - patch( - "tasks.jobs.scan.generate_scan_compliance" - ) as mock_generate_compliance, - patch("tasks.jobs.scan.create_objects_in_batches") as mock_create_objects, - patch("api.models.Finding.objects.filter") as mock_findings_filter, + patch("tasks.jobs.scan.generate_scan_compliance"), ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - - mock_findings_filter.return_value = [] - - mock_prowler_provider = MagicMock() - mock_prowler_provider.get_regions.return_value = ["us-east-1", "us-west-2"] - mock_return_prowler_provider.return_value = mock_prowler_provider + tenant_id = str(tenants_fixture[0].id) + scan_id = str(scans_fixture[0].id) mock_compliance_template.__getitem__.return_value = { "test_compliance": { @@ -1082,7 +671,6 @@ class TestCreateComplianceRequirements: "requirements": { "req_1": { "description": "Test Requirement 1", - "checks": {"check_1": None}, "checks_status": { "pass": 2, "fail": 0, @@ -1093,7 +681,6 @@ class TestCreateComplianceRequirements: }, "req_2": { "description": "Test Requirement 2", - "checks": {"check_2": None}, "checks_status": { "pass": 1, "fail": 1, @@ -1106,146 +693,72 @@ class TestCreateComplianceRequirements: } } - mock_generate_compliance.return_value = { - "test_compliance": { - "framework": "Test Framework", - "version": "1.0", - "requirements": { - "req_1": { - "description": "Test Requirement 1", - "checks": { - "check_1": { - "us-east-1": {"status": "PASS"}, - "us-west-2": {"status": "PASS"}, - } - }, - "checks_status": { - "pass": 2, - "fail": 0, - "manual": 0, - "total": 2, - }, - "status": "PASS", - }, - "req_2": { - "description": "Test Requirement 2", - "checks": { - "check_2": { - "us-east-1": {"status": "PASS"}, - "us-west-2": {"status": "FAIL"}, - } - }, - "checks_status": { - "pass": 1, - "fail": 1, - "manual": 0, - "total": 2, - }, - "status": "FAIL", - }, - }, - } - } + result = create_compliance_requirements(tenant_id, scan_id) - created_objects = [] - mock_create_objects.side_effect = ( - lambda tenant_id, model, objs, batch_size=500: created_objects.extend( - objs - ) - ) - - create_compliance_requirements(str(tenant.id), str(scan.id)) - - assert len(created_objects) == 4 - req_1_objects = [ - obj for obj in created_objects if obj.requirement_id == "req_1" - ] - req_2_objects = [ - obj for obj in created_objects if obj.requirement_id == "req_2" - ] - assert len(req_1_objects) == 2 - assert len(req_2_objects) == 2 - assert all(obj.requirement_status == "PASS" for obj in req_1_objects) - assert all(obj.requirement_status == "FAIL" for obj in req_2_objects) + assert "requirements_created" in result + assert result["requirements_created"] >= 0 @pytest.mark.django_db class TestUpdateResourceFailedFindingsCount: - @patch("api.models.Resource.all_objects.filter") - @patch("api.models.Finding.all_objects.filter") - def test_failed_findings_count_update( - self, - mock_finding_filter, - mock_resource_filter, - tenants_fixture, - scans_fixture, - providers_fixture, + def test_execute_sql_update( + self, tenants_fixture, scans_fixture, providers_fixture, resources_fixture ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - provider = providers_fixture[0] + resource = resources_fixture[0] + tenant_id = resource.tenant_id + scan_id = resource.provider.scans.first().id - scan.provider = provider - scan.save() + # Common kwargs for all failing findings + base_kwargs = { + "tenant_id": tenant_id, + "scan_id": scan_id, + "delta": None, + "status": StatusChoices.FAIL, + "status_extended": "test status extended", + "impact": Severity.critical, + "impact_extended": "test impact extended", + "severity": Severity.critical, + "raw_result": { + "status": StatusChoices.FAIL, + "impact": Severity.critical, + "severity": Severity.critical, + }, + "tags": {"test": "dev-qa"}, + "check_id": "test_check_id", + "check_metadata": { + "CheckId": "test_check_id", + "Description": "test description apple sauce", + "servicename": "ec2", + }, + "first_seen_at": "2024-01-02T00:00:00Z", + } - tenant_id = str(tenant.id) - scan_id = str(scan.id) + # UIDs to create (two with same UID, one unique) + uids = ["test_finding_uid_1", "test_finding_uid_1", "test_finding_uid_2"] - resource1 = MagicMock() - resource1.uid = "res-1" - resource1.failed_findings_count = None - resource1.save = MagicMock() + # Create findings and associate with the resource + for uid in uids: + finding = Finding.objects.create(uid=uid, **base_kwargs) + finding.add_resources([resource]) - resource2 = MagicMock() - resource2.uid = "res-2" - resource2.failed_findings_count = None - resource2.save = MagicMock() + resource.refresh_from_db() + assert resource.failed_findings_count == 0 - mock_resource_filter.return_value = [resource1, resource2] + _update_resource_failed_findings_count(tenant_id=tenant_id, scan_id=scan_id) + resource.refresh_from_db() - fake_subquery_qs = MagicMock() - fake_subquery_qs.order_by.return_value = fake_subquery_qs - fake_subquery_qs.values.return_value = fake_subquery_qs - fake_subquery_qs.__getitem__.return_value = fake_subquery_qs + # Only two since two findings share the same UID + assert resource.failed_findings_count == 2 - def finding_filter_side_effect(*args, **kwargs): - if "status" in kwargs: - qs_count = MagicMock() - if kwargs.get("resources") == resource1: - qs_count.count.return_value = 3 - else: - qs_count.count.return_value = 0 - return qs_count - return fake_subquery_qs - - mock_finding_filter.side_effect = finding_filter_side_effect - - _update_resource_failed_findings_count(tenant_id, scan_id) - - # resource1 should have been updated to 3 - assert resource1.failed_findings_count == 3 - resource1.save.assert_called_once_with(update_fields=["failed_findings_count"]) - - # resource2 should have been updated to 0 - assert resource2.failed_findings_count == 0 - resource2.save.assert_called_once_with(update_fields=["failed_findings_count"]) - - @patch("api.models.Resource.all_objects.filter", return_value=[]) - @patch("api.models.Finding.all_objects.filter") - def test_no_resources_no_error( + @patch("tasks.jobs.scan.Scan.objects.get") + def test_scan_not_found( self, - mock_finding_filter, - mock_resource_filter, - tenants_fixture, - scans_fixture, - providers_fixture, + mock_scan_get, ): - tenant = tenants_fixture[0] - scan = scans_fixture[0] - provider = providers_fixture[0] - scan.provider = provider - scan.save() + mock_scan_get.side_effect = Scan.DoesNotExist - _update_resource_failed_findings_count(str(tenant.id), str(scan.id)) - - mock_finding_filter.assert_not_called() + with pytest.raises(Scan.DoesNotExist): + _update_resource_failed_findings_count( + "8614ca97-8370-4183-a7f7-e96a6c7d2c93", + "4705bed5-8782-4e8b-bab6-55e8043edaa6", + )