diff --git a/api/changelog.d/compliance-overviews-ingest-perf.changed.md b/api/changelog.d/compliance-overviews-ingest-perf.changed.md new file mode 100644 index 0000000000..6e833d184e --- /dev/null +++ b/api/changelog.d/compliance-overviews-ingest-perf.changed.md @@ -0,0 +1 @@ +Speed up compliance overview ingestion by reading ThreatScore mappings from the compliance template instead of each finding, generating time-ordered `uuid7` row ids and grouping inserted rows by framework and requirement diff --git a/api/src/backend/tasks/jobs/scan.py b/api/src/backend/tasks/jobs/scan.py index 0a8db605c0..8b57ebeecd 100644 --- a/api/src/backend/tasks/jobs/scan.py +++ b/api/src/backend/tasks/jobs/scan.py @@ -5,7 +5,6 @@ import json import random import re import time -import uuid from collections import defaultdict from collections.abc import Callable, Iterable from datetime import UTC, datetime @@ -73,6 +72,7 @@ from tasks.jobs.queries import ( COMPLIANCE_UPSERT_TENANT_SUMMARY_SQL, ) from tasks.utils import CustomEncoder, batched +from uuid6 import uuid7 logger = get_task_logger(__name__) @@ -1756,32 +1756,27 @@ def aggregate_findings(tenant_id: str, scan_id: str): def _aggregate_findings_by_region( - tenant_id: str, scan_id: str, modeled_threatscore_compliance_id: str + tenant_id: str, + scan_id: str, + normalized_threatscore_id: str, + threatscore_requirements_by_check: dict[str, list[str]], ) -> tuple[dict, dict]: """ Aggregate findings by region using streaming, column-scoped ORM reads. Reads only the consumed columns as tuples via ``values_list`` and streams them with ``.iterator()``, using the denormalized ``resource_regions`` array - instead of ``prefetch_related("resources")``. ``resource_regions`` mirrors the - regions of a finding's related resources, so it yields the same per-region - tally without joining the resource table. - - Args: - tenant_id: Tenant UUID - scan_id: Scan UUID - modeled_threatscore_compliance_id: ID for ThreatScore compliance framework + instead of ``prefetch_related("resources")``. ThreatScore requirement ids + are resolved per ``check_id`` from ``threatscore_requirements_by_check``. Returns: tuple: (check_status_by_region, findings_count_by_compliance) - check_status_by_region: {region: {check_id: status}} - - findings_count_by_compliance: {region: {normalized_id: {requirement_id: {total, pass}}}} + - findings_count_by_compliance: {region: {normalized_threatscore_id: {requirement_id: {total, pass}}}} """ check_status_by_region: dict = {} findings_count_by_compliance: dict = {} - normalized_id = re.sub(r"[^a-z0-9]", "", modeled_threatscore_compliance_id.lower()) - with rls_transaction(tenant_id, using=READ_REPLICA_ALIAS): findings = ( Finding.all_objects.filter( @@ -1790,14 +1785,12 @@ def _aggregate_findings_by_region( muted=False, status__in=["PASS", "FAIL"], ) - .values_list("check_id", "status", "resource_regions", "compliance") + .values_list("check_id", "status", "resource_regions") .iterator(chunk_size=DJANGO_FINDINGS_BATCH_SIZE) ) - for check_id, status, resource_regions, compliance in findings: - threatscore_requirements = (compliance or {}).get( - modeled_threatscore_compliance_id - ) + for check_id, status, resource_regions in findings: + threatscore_requirements = threatscore_requirements_by_check.get(check_id) for region in resource_regions or (): # Priority: FAIL > any other status @@ -1809,7 +1802,7 @@ def _aggregate_findings_by_region( if threatscore_requirements: compliance_key = findings_count_by_compliance.setdefault( region, {} - ).setdefault(normalized_id, {}) + ).setdefault(normalized_threatscore_id, {}) for requirement_id in threatscore_requirements: requirement_stats = compliance_key.setdefault( @@ -1848,15 +1841,28 @@ def create_compliance_requirements(tenant_id: str, scan_id: str): compliance_template = PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE[ provider_instance.provider ] - modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0" + normalized_threatscore_id = _normalized_compliance_key( + "ProwlerThreatScore", "1.0" + ) requirement_lookup: dict[str, list[tuple[str, str]]] = {} + threatscore_requirements_by_check: dict[str, list[str]] = {} for compliance_id, compliance in compliance_template.items(): + is_threatscore = ( + _normalized_compliance_key( + compliance["framework"], compliance["version"] + ) + == normalized_threatscore_id + ) for requirement_id, requirement in compliance["requirements"].items(): for check_id in requirement["checks"].keys(): requirement_lookup.setdefault(check_id, []).append( (compliance_id, requirement_id) ) + if is_threatscore: + threatscore_requirements_by_check.setdefault( + check_id, [] + ).append(requirement_id) regions = [] requirements_created = 0 @@ -1869,7 +1875,10 @@ def create_compliance_requirements(tenant_id: str, scan_id: str): # Aggregate findings by region using SQL for optimal performance check_status_by_region, findings_count_by_compliance = ( _aggregate_findings_by_region( - tenant_id, scan_id, modeled_threatscore_compliance_id + tenant_id, + scan_id, + normalized_threatscore_id, + threatscore_requirements_by_check, ) ) @@ -1934,23 +1943,35 @@ def create_compliance_requirements(tenant_id: str, scan_id: str): # Yield rows lazily (consumed batch-by-batch by COPY) so peak memory # stays bounded; tally requirement_statuses in the same pass. The # ORM fallback re-iterates from scratch, so the tally resets first. + # Region is the innermost loop so consecutive rows share the leading + # columns of the table's secondary indexes. def _iter_compliance_requirement_rows(): requirement_statuses.clear() - for region in regions: - region_stats = region_requirement_stats.get(region, {}) - region_findings = findings_count_by_compliance.get(region, {}) - for ( - compliance_id, - framework, - version, - modeled_compliance_id, - requirements, - ) in compliance_plan: - compliance_stats = region_stats.get(compliance_id, {}) - compliance_findings = region_findings.get( - modeled_compliance_id, {} + for ( + compliance_id, + framework, + version, + modeled_compliance_id, + requirements, + ) in compliance_plan: + stats_by_region = [ + ( + region, + region_requirement_stats.get(region, {}).get( + compliance_id, {} + ), + findings_count_by_compliance.get(region, {}).get( + modeled_compliance_id, {} + ), ) - for requirement_id, description, total_checks in requirements: + for region in regions + ] + for requirement_id, description, total_checks in requirements: + for ( + region, + compliance_stats, + compliance_findings, + ) in stats_by_region: stats = compliance_stats.get(requirement_id) if stats: passed_checks = stats["passed_checks"] @@ -1981,7 +2002,7 @@ def create_compliance_requirements(tenant_id: str, scan_id: str): requirement_statuses[key]["pass_count"] += 1 yield { - "id": uuid.uuid4(), + "id": uuid7(), "tenant_id": tenant_id_str, "inserted_at": utc_datetime_now, "compliance_id": compliance_id, diff --git a/api/src/backend/tasks/tests/test_scan.py b/api/src/backend/tasks/tests/test_scan.py index 8027588398..0cee24351e 100644 --- a/api/src/backend/tasks/tests/test_scan.py +++ b/api/src/backend/tasks/tests/test_scan.py @@ -1,6 +1,5 @@ import csv import json -import re import uuid from collections.abc import MutableMapping from contextlib import contextmanager @@ -10,6 +9,7 @@ from unittest.mock import MagicMock, patch import pytest from api.db_router import MainRouter +from api.db_utils import rls_transaction from api.exceptions import ProviderConnectionError, ProviderDeletedException from api.models import ( Finding, @@ -2795,6 +2795,167 @@ class TestCreateComplianceRequirements: assert count_after_first > 0 assert count_after_second == count_after_first + with rls_transaction(tenant_id): + row_versions = { + row_id.version + for row_id in ComplianceRequirementOverview.objects.filter( + scan_id=scan_id + ).values_list("id", flat=True) + } + assert row_versions == {7} + + def test_create_compliance_requirements_threatscore_counts_from_template( + self, + tenants_fixture, + scans_fixture, + aws_provider, + findings_fixture, + ): + """ThreatScore finding counts are derived from the template mapping, + not from each finding's stored ``compliance`` payload.""" + from api.models import ComplianceRequirementOverview + + with patch( + "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" + ) as mock_compliance_template: + tenant_id = str(tenants_fixture[0].id) + scan_id = str(scans_fixture[0].id) + + mock_compliance_template.__getitem__.return_value = { + "prowler_threatscore_aws": { + "framework": "ProwlerThreatScore", + "version": "1.0", + "requirements": { + "1.1.1": { + "description": "ThreatScore requirement", + "checks": {"test_check_id": None}, + }, + "1.1.2": { + "description": "Unrelated requirement", + "checks": {"other_check_id": None}, + }, + }, + }, + "other_framework": { + "framework": "Other", + "version": "2.0", + "requirements": { + "a": { + "description": "Same check, other framework", + "checks": {"test_check_id": None}, + }, + }, + }, + } + + create_compliance_requirements(tenant_id, scan_id) + + with rls_transaction(tenant_id): + counted = sum( + len(finding.resource_regions or []) + for finding in Finding.all_objects.filter( + scan_id=scan_id, + muted=False, + status__in=["PASS", "FAIL"], + check_id="test_check_id", + ) + ) + rows = list( + ComplianceRequirementOverview.objects.filter( + scan_id=scan_id + ).values_list("compliance_id", "requirement_id", "total_findings") + ) + assert counted > 0 + assert ( + sum( + total + for compliance_id, requirement_id, total in rows + if (compliance_id, requirement_id) + == ("prowler_threatscore_aws", "1.1.1") + ) + == counted + ) + assert all( + total == 0 + for compliance_id, requirement_id, total in rows + if (compliance_id, requirement_id) == ("prowler_threatscore_aws", "1.1.2") + ) + assert all( + total == 0 + for compliance_id, _, total in rows + if compliance_id == "other_framework" + ) + + def test_create_compliance_requirements_rows_across_regions_and_frameworks( + self, + tenants_fixture, + scans_fixture, + aws_provider, + ): + from api.models import ComplianceRequirementOverview + + tenant_id = str(tenants_fixture[0].id) + scan_id = str(scans_fixture[0].id) + check_status_by_region = { + "us-east-1": {"check_a": "FAIL", "check_b": "PASS"}, + "eu-west-1": {"check_a": "PASS"}, + } + template = { + "fw_one": { + "framework": "One", + "version": "1", + "requirements": { + "r1": {"description": "a", "checks": {"check_a": None}}, + "r2": { + "description": "a+b", + "checks": {"check_a": None, "check_b": None}, + }, + }, + }, + "fw_two": { + "framework": "Two", + "version": "2", + "requirements": { + "m1": {"description": "manual", "checks": {}}, + }, + }, + } + + with ( + patch( + "tasks.jobs.scan.PROWLER_COMPLIANCE_OVERVIEW_TEMPLATE" + ) as mock_compliance_template, + patch( + "tasks.jobs.scan._aggregate_findings_by_region", + return_value=(check_status_by_region, {}), + ), + ): + mock_compliance_template.__getitem__.return_value = template + result = create_compliance_requirements(tenant_id, scan_id) + + assert result["requirements_created"] == 6 + with rls_transaction(tenant_id): + rows = set( + ComplianceRequirementOverview.objects.filter( + scan_id=scan_id + ).values_list( + "compliance_id", + "requirement_id", + "region", + "requirement_status", + "passed_checks", + "failed_checks", + "total_checks", + ) + ) + assert rows == { + ("fw_one", "r1", "us-east-1", "FAIL", 0, 1, 1), + ("fw_one", "r1", "eu-west-1", "PASS", 1, 0, 1), + ("fw_one", "r2", "us-east-1", "FAIL", 1, 1, 2), + ("fw_one", "r2", "eu-west-1", "PASS", 1, 0, 2), + ("fw_two", "m1", "us-east-1", "MANUAL", 0, 0, 0), + ("fw_two", "m1", "eu-west-1", "MANUAL", 0, 0, 0), + } def test_create_compliance_requirements_kubernetes_provider( self, @@ -4723,17 +4884,11 @@ class TestAggregateFindingsByRegion: """Test function returns correct data structure.""" tenant_id = str(uuid.uuid4()) scan_id = str(uuid.uuid4()) - modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0" + normalized_id = "prowlerthreatscore10" - # (check_id, status, resource_regions, compliance) tuples - finding_rows = [ - ( - "check1", - "FAIL", - ["us-east-1"], - {modeled_threatscore_compliance_id: ["req1", "req2"]}, - ) - ] + # (check_id, status, resource_regions) tuples + finding_rows = [("check1", "FAIL", ["us-east-1"])] + threatscore_by_check = {"check1": ["req1", "req2"]} mock_queryset = MagicMock() mock_queryset.values_list.return_value = mock_queryset @@ -4747,13 +4902,16 @@ class TestAggregateFindingsByRegion: check_status_by_region, findings_count_by_compliance = ( _aggregate_findings_by_region( - tenant_id, scan_id, modeled_threatscore_compliance_id + tenant_id, + scan_id, + normalized_id, + threatscore_by_check, ) ) # Streaming query contract: column-scoped values_list + iterator mock_queryset.values_list.assert_called_once_with( - "check_id", "status", "resource_regions", "compliance" + "check_id", "status", "resource_regions" ) mock_queryset.iterator.assert_called_once() @@ -4774,13 +4932,14 @@ class TestAggregateFindingsByRegion: """Test that FAIL status takes priority over other statuses.""" tenant_id = str(uuid.uuid4()) scan_id = str(uuid.uuid4()) - modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0" + normalized_id = "prowlerthreatscore10" # Same check/region: PASS first, then FAIL — FAIL must win finding_rows = [ - ("check1", "PASS", ["us-east-1"], {}), - ("check1", "FAIL", ["us-east-1"], {}), + ("check1", "PASS", ["us-east-1"]), + ("check1", "FAIL", ["us-east-1"]), ] + threatscore_by_check = {} mock_queryset = MagicMock() mock_queryset.values_list.return_value = mock_queryset @@ -4793,12 +4952,15 @@ class TestAggregateFindingsByRegion: mock_findings_filter.return_value = mock_queryset check_status_by_region, _ = _aggregate_findings_by_region( - tenant_id, scan_id, modeled_threatscore_compliance_id + tenant_id, + scan_id, + normalized_id, + threatscore_by_check, ) # Streaming query contract: column-scoped values_list + iterator mock_queryset.values_list.assert_called_once_with( - "check_id", "status", "resource_regions", "compliance" + "check_id", "status", "resource_regions" ) mock_queryset.iterator.assert_called_once() @@ -4813,8 +4975,9 @@ class TestAggregateFindingsByRegion: """Test that muted findings are filtered out (muted=False in query).""" tenant_id = str(uuid.uuid4()) scan_id = str(uuid.uuid4()) - modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0" + normalized_id = "prowlerthreatscore10" + threatscore_by_check = {} mock_queryset = MagicMock() mock_queryset.values_list.return_value = mock_queryset mock_queryset.iterator.return_value = [] @@ -4826,12 +4989,15 @@ class TestAggregateFindingsByRegion: mock_findings_filter.return_value = mock_queryset _aggregate_findings_by_region( - tenant_id, scan_id, modeled_threatscore_compliance_id + tenant_id, + scan_id, + normalized_id, + threatscore_by_check, ) # Streaming query contract: column-scoped values_list + iterator mock_queryset.values_list.assert_called_once_with( - "check_id", "status", "resource_regions", "compliance" + "check_id", "status", "resource_regions" ) mock_queryset.iterator.assert_called_once() @@ -4851,23 +5017,14 @@ class TestAggregateFindingsByRegion: """Test that ThreatScore compliance counts are processed correctly.""" tenant_id = str(uuid.uuid4()) scan_id = str(uuid.uuid4()) - modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0" + normalized_id = "prowlerthreatscore10" # PASS and FAIL findings mapped to the same ThreatScore requirement finding_rows = [ - ( - "check1", - "PASS", - ["us-east-1"], - {modeled_threatscore_compliance_id: ["req1"]}, - ), - ( - "check2", - "FAIL", - ["us-east-1"], - {modeled_threatscore_compliance_id: ["req1"]}, - ), + ("check1", "PASS", ["us-east-1"]), + ("check2", "FAIL", ["us-east-1"]), ] + threatscore_by_check = {"check1": ["req1"], "check2": ["req1"]} mock_queryset = MagicMock() mock_queryset.values_list.return_value = mock_queryset @@ -4880,19 +5037,19 @@ class TestAggregateFindingsByRegion: mock_findings_filter.return_value = mock_queryset _, findings_count_by_compliance = _aggregate_findings_by_region( - tenant_id, scan_id, modeled_threatscore_compliance_id + tenant_id, + scan_id, + normalized_id, + threatscore_by_check, ) # Streaming query contract: column-scoped values_list + iterator mock_queryset.values_list.assert_called_once_with( - "check_id", "status", "resource_regions", "compliance" + "check_id", "status", "resource_regions" ) mock_queryset.iterator.assert_called_once() # Verify compliance counts - normalized_id = re.sub( - r"[^a-z0-9]", "", modeled_threatscore_compliance_id.lower() - ) assert "us-east-1" in findings_count_by_compliance assert normalized_id in findings_count_by_compliance["us-east-1"] assert "req1" in findings_count_by_compliance["us-east-1"][normalized_id] @@ -4909,13 +5066,14 @@ class TestAggregateFindingsByRegion: """Test aggregation across multiple regions.""" tenant_id = str(uuid.uuid4()) scan_id = str(uuid.uuid4()) - modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0" + normalized_id = "prowlerthreatscore10" # One finding per region finding_rows = [ - ("check1", "FAIL", ["us-east-1"], {}), - ("check1", "PASS", ["us-west-2"], {}), + ("check1", "FAIL", ["us-east-1"]), + ("check1", "PASS", ["us-west-2"]), ] + threatscore_by_check = {} mock_queryset = MagicMock() mock_queryset.values_list.return_value = mock_queryset @@ -4928,12 +5086,15 @@ class TestAggregateFindingsByRegion: mock_findings_filter.return_value = mock_queryset check_status_by_region, _ = _aggregate_findings_by_region( - tenant_id, scan_id, modeled_threatscore_compliance_id + tenant_id, + scan_id, + normalized_id, + threatscore_by_check, ) # Streaming query contract: column-scoped values_list + iterator mock_queryset.values_list.assert_called_once_with( - "check_id", "status", "resource_regions", "compliance" + "check_id", "status", "resource_regions" ) mock_queryset.iterator.assert_called_once() @@ -4951,16 +5112,10 @@ class TestAggregateFindingsByRegion: """A finding with multiple resource_regions is tallied in every region.""" tenant_id = str(uuid.uuid4()) scan_id = str(uuid.uuid4()) - modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0" + normalized_id = "prowlerthreatscore10" - finding_rows = [ - ( - "check1", - "FAIL", - ["us-east-1", "eu-west-1"], - {modeled_threatscore_compliance_id: ["req1"]}, - ) - ] + finding_rows = [("check1", "FAIL", ["us-east-1", "eu-west-1"])] + threatscore_by_check = {"check1": ["req1"]} mock_queryset = MagicMock() mock_queryset.values_list.return_value = mock_queryset @@ -4974,19 +5129,19 @@ class TestAggregateFindingsByRegion: check_status_by_region, findings_count_by_compliance = ( _aggregate_findings_by_region( - tenant_id, scan_id, modeled_threatscore_compliance_id + tenant_id, + scan_id, + normalized_id, + threatscore_by_check, ) ) # Streaming query contract: column-scoped values_list + iterator mock_queryset.values_list.assert_called_once_with( - "check_id", "status", "resource_regions", "compliance" + "check_id", "status", "resource_regions" ) mock_queryset.iterator.assert_called_once() - normalized_id = re.sub( - r"[^a-z0-9]", "", modeled_threatscore_compliance_id.lower() - ) for region in ("us-east-1", "eu-west-1"): assert check_status_by_region[region]["check1"] == "FAIL" req_stats = findings_count_by_compliance[region][normalized_id]["req1"] @@ -5000,12 +5155,13 @@ class TestAggregateFindingsByRegion: """A finding with no denormalized regions contributes nothing.""" tenant_id = str(uuid.uuid4()) scan_id = str(uuid.uuid4()) - modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0" + normalized_id = "prowlerthreatscore10" finding_rows = [ - ("check1", "FAIL", [], {modeled_threatscore_compliance_id: ["req1"]}), - ("check2", "PASS", None, {}), + ("check1", "FAIL", []), + ("check2", "PASS", None), ] + threatscore_by_check = {"check1": ["req1"]} mock_queryset = MagicMock() mock_queryset.values_list.return_value = mock_queryset @@ -5019,13 +5175,16 @@ class TestAggregateFindingsByRegion: check_status_by_region, findings_count_by_compliance = ( _aggregate_findings_by_region( - tenant_id, scan_id, modeled_threatscore_compliance_id + tenant_id, + scan_id, + normalized_id, + threatscore_by_check, ) ) # Streaming query contract: column-scoped values_list + iterator mock_queryset.values_list.assert_called_once_with( - "check_id", "status", "resource_regions", "compliance" + "check_id", "status", "resource_regions" ) mock_queryset.iterator.assert_called_once() @@ -5040,8 +5199,9 @@ class TestAggregateFindingsByRegion: """Test with no findings - should return empty dicts.""" tenant_id = str(uuid.uuid4()) scan_id = str(uuid.uuid4()) - modeled_threatscore_compliance_id = "ProwlerThreatScore-1.0" + normalized_id = "prowlerthreatscore10" + threatscore_by_check = {} mock_queryset = MagicMock() mock_queryset.values_list.return_value = mock_queryset mock_queryset.iterator.return_value = [] @@ -5054,13 +5214,16 @@ class TestAggregateFindingsByRegion: check_status_by_region, findings_count_by_compliance = ( _aggregate_findings_by_region( - tenant_id, scan_id, modeled_threatscore_compliance_id + tenant_id, + scan_id, + normalized_id, + threatscore_by_check, ) ) # Streaming query contract: column-scoped values_list + iterator mock_queryset.values_list.assert_called_once_with( - "check_id", "status", "resource_regions", "compliance" + "check_id", "status", "resource_regions" ) mock_queryset.iterator.assert_called_once()