diff --git a/api/CHANGELOG.md b/api/CHANGELOG.md index 85387bf9d7..b917037e21 100644 --- a/api/CHANGELOG.md +++ b/api/CHANGELOG.md @@ -8,6 +8,10 @@ All notable changes to the **Prowler API** are documented in this file. - Compliance PDF reports no longer require provider credentials: findings are enriched from the provider metadata stored in the database, so reports generate even after the provider secret is deleted or its credentials become invalid [(#11845)](https://github.com/prowler-cloud/prowler/pull/11845) +### 🐞 Fixed + +- Provider scans now queue behind active provider scans instead of dispatching concurrently, and resource failed-finding counters retry database conflicts with stable row locking [(#11848)](https://github.com/prowler-cloud/prowler/pull/11848) + --- ## [1.33.1] (Prowler v5.32.1) diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index a4951d936e..b3a0f47e87 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -65,6 +65,7 @@ from api.v1.views import ( TenantFinishACSView, ) from botocore.exceptions import ClientError, NoCredentialsError +from celery import states from conftest import ( API_JSON_CONTENT_TYPE, TEST_PASSWORD, @@ -3644,21 +3645,15 @@ class TestScanViewSet: ), ], ) - @patch("api.v1.views.Task.objects.get") - @patch("api.v1.views.perform_scan_task.apply_async") + @patch("api.v1.views.enqueue_scan_execution_on_commit") def test_scans_create_valid( self, - mock_perform_scan_task, - mock_task_get, + mock_enqueue_scan_execution, authenticated_client, scan_json_payload, _expected_scanner_args, providers_fixture, - tasks_fixture, ): - prowler_task = tasks_fixture[0] - mock_perform_scan_task.return_value.id = prowler_task.id - mock_task_get.return_value = prowler_task *_, provider5 = providers_fixture # Provider5 has these scanner_args # scanner_args={"key1": "value1", "key2": {"key21": "value21"}} @@ -3683,8 +3678,121 @@ class TestScanViewSet: assert scan.name == scan_json_payload["data"]["attributes"]["name"] assert scan.provider == provider5 assert scan.trigger == Scan.TriggerChoices.MANUAL + mock_enqueue_scan_execution.assert_called_once() # assert scan.scanner_args == expected_scanner_args + @patch("tasks.tasks.perform_scan_task.apply_async") + def test_scans_create_queues_scan_when_provider_has_active_scan( + self, + mock_perform_scan_task, + authenticated_client, + providers_fixture, + tenants_fixture, + django_capture_on_commit_callbacks, + ): + tenant, *_ = tenants_fixture + provider, *_ = providers_fixture + task_result = TaskResult.objects.create( + task_id=str(uuid4()), + task_name="scan-perform", + status=states.PENDING, + ) + prowler_task = Task.objects.create( + id=task_result.task_id, + tenant_id=tenant.id, + task_runner_task=task_result, + ) + Scan.objects.create( + name="Active scan", + provider=provider, + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.AVAILABLE, + tenant_id=tenant.id, + task=prowler_task, + ) + + with django_capture_on_commit_callbacks(execute=True): + response = authenticated_client.post( + reverse("scan-list"), + data={ + "data": { + "type": "scans", + "attributes": {"name": "Duplicate Scan"}, + "relationships": { + "provider": { + "data": {"type": "providers", "id": str(provider.id)} + } + }, + } + }, + content_type=API_JSON_CONTENT_TYPE, + ) + + assert response.status_code == status.HTTP_202_ACCEPTED + assert response.json()["data"]["id"] != str(prowler_task.id) + assert Scan.objects.count() == 2 + queued_scan = Scan.objects.exclude(task=prowler_task).get() + assert queued_scan.trigger == Scan.TriggerChoices.MANUAL + assert queued_scan.state == StateChoices.AVAILABLE + assert queued_scan.task.task_runner_task.status == "QUEUED" + mock_perform_scan_task.assert_not_called() + + @patch("tasks.tasks.perform_scan_task.apply_async") + def test_scans_create_queues_scan_when_scheduled_scan_is_claimed( + self, + mock_perform_scan_task, + authenticated_client, + providers_fixture, + tenants_fixture, + django_capture_on_commit_callbacks, + ): + tenant, *_ = tenants_fixture + provider, *_ = providers_fixture + task_result = TaskResult.objects.create( + task_id=str(uuid4()), + task_name="scan-perform-scheduled", + status=states.STARTED, + ) + prowler_task = Task.objects.create( + id=task_result.task_id, + tenant_id=tenant.id, + task_runner_task=task_result, + ) + Scan.objects.create( + name="Claimed scheduled scan", + provider=provider, + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.SCHEDULED, + tenant_id=tenant.id, + task=prowler_task, + ) + + with django_capture_on_commit_callbacks(execute=True): + response = authenticated_client.post( + reverse("scan-list"), + data={ + "data": { + "type": "scans", + "attributes": {"name": "Manual Scan"}, + "relationships": { + "provider": { + "data": {"type": "providers", "id": str(provider.id)} + } + }, + } + }, + content_type=API_JSON_CONTENT_TYPE, + ) + + assert response.status_code == status.HTTP_202_ACCEPTED + assert response.json()["data"]["id"] != str(prowler_task.id) + assert Scan.objects.count() == 2 + queued_scan = Scan.objects.exclude(task=prowler_task).get() + assert queued_scan.trigger == Scan.TriggerChoices.MANUAL + assert queued_scan.state == StateChoices.AVAILABLE + assert queued_scan.task.task_runner_task.status == "QUEUED" + mock_perform_scan_task.assert_not_called() + @pytest.mark.parametrize( "scan_json_payload, error_code", [ diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index e7907cf6f8..e392505818 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -237,7 +237,7 @@ from api.v1.serializers import ( UserUpdateSerializer, ) from botocore.exceptions import ClientError, NoCredentialsError, ParamValidationError -from celery import chain, states +from celery import chain from celery.result import AsyncResult from config.custom_logging import BackendLogger from config.env import env @@ -283,7 +283,6 @@ from django.utils.dateparse import parse_date from django.utils.decorators import method_decorator from django.views.decorators.cache import cache_control from django_celery_beat.models import PeriodicTask -from django_celery_results.models import TaskResult from drf_spectacular.settings import spectacular_settings from drf_spectacular.types import OpenApiTypes from drf_spectacular.utils import ( @@ -322,17 +321,20 @@ from tasks.beat import schedule_provider_scan from tasks.jobs.attack_paths import db_utils as attack_paths_db_utils from tasks.jobs.export import get_s3_client from tasks.tasks import ( + QUEUED_SCAN_TASK_STATE, backfill_compliance_summaries_task, backfill_scan_resource_summaries_task, check_integration_connection_task, check_lighthouse_connection_task, check_lighthouse_provider_connection_task, check_provider_connection_task, + create_scan_task_record, delete_provider_task, delete_tenant_task, + enqueue_scan_execution_on_commit, + get_active_provider_scan, jira_integration_task, mute_historical_findings_task, - perform_scan_task, reaggregate_all_finding_group_summaries_task, refresh_lighthouse_provider_models_task, ) @@ -2717,12 +2719,23 @@ class ScanViewSet(BaseRLSViewSet): def create(self, request, *args, **kwargs): input_serializer = self.get_serializer(data=request.data) input_serializer.is_valid(raise_exception=True) + provider = input_serializer.validated_data.get("provider") + active_scan = None # Broker publish is deferred to on_commit so the worker cannot read # Scan before BaseRLSViewSet's dispatch-wide atomic commits. pre_task_id = str(uuid.uuid4()) with transaction.atomic(): + if provider: + provider = Provider.objects.select_for_update().get( + id=provider.id, + tenant_id=self.request.tenant_id, + ) + active_scan = get_active_provider_scan( + self.request.tenant_id, provider.id + ) + scan = input_serializer.save() scan.task_id = pre_task_id scan.save(update_fields=["task_id"]) @@ -2733,29 +2746,18 @@ class ScanViewSet(BaseRLSViewSet): provider_id=str(scan.provider_id), ) - task_result, _ = TaskResult.objects.get_or_create( - task_id=pre_task_id, - defaults={"status": states.PENDING, "task_name": "scan-perform"}, - ) - prowler_task, _ = Task.objects.update_or_create( - id=pre_task_id, + prowler_task = create_scan_task_record( tenant_id=self.request.tenant_id, - defaults={"task_runner_task": task_result}, + task_id=pre_task_id, + task_status=(QUEUED_SCAN_TASK_STATE if active_scan else None), ) - scan_kwargs = { - "tenant_id": self.request.tenant_id, - "scan_id": str(scan.id), - "provider_id": str(scan.provider_id), - # Disabled for now - # checks_to_execute=scan.scanner_args.get("checks_to_execute") - } - - transaction.on_commit( - lambda: perform_scan_task.apply_async( - kwargs=scan_kwargs, task_id=pre_task_id + if not active_scan: + enqueue_scan_execution_on_commit( + tenant_id=self.request.tenant_id, + scan=scan, + task_id=pre_task_id, ) - ) self.response_serializer_class = TaskSerializer output_serializer = self.get_serializer(prowler_task) diff --git a/api/src/backend/tasks/jobs/scan.py b/api/src/backend/tasks/jobs/scan.py index d69e0c8941..8290c45a7a 100644 --- a/api/src/backend/tasks/jobs/scan.py +++ b/api/src/backend/tasks/jobs/scan.py @@ -1,6 +1,7 @@ import csv import io import json +import random import re import time import uuid @@ -306,6 +307,55 @@ def _store_resources( return resource_instance, (resource_instance.uid, resource_instance.region) +def _bulk_update_resource_failed_findings_counts( + tenant_id: str, + scan_id: str, + resources_to_update: list[Resource], +) -> None: + """Persist failed finding counters with stable row locking and retry.""" + if not resources_to_update: + return + + sorted_resources = sorted( + resources_to_update, key=lambda resource: str(resource.id) + ) + for start in range(0, len(sorted_resources), SCAN_DB_BATCH_SIZE): + chunk = sorted_resources[start : start + SCAN_DB_BATCH_SIZE] + chunk_ids = [resource.id for resource in chunk] + + for attempt in range(CELERY_DEADLOCK_ATTEMPTS): + try: + with rls_transaction(tenant_id): + list( + Resource.objects.select_for_update() + .filter(id__in=chunk_ids) + .order_by("id") + .values_list("id", flat=True) + ) + Resource.objects.bulk_update( + chunk, + ["failed_findings_count"], + batch_size=SCAN_DB_BATCH_SIZE, + ) + break + except OperationalError: + if attempt < CELERY_DEADLOCK_ATTEMPTS - 1: + logger.warning( + "Resource failed findings count update hit a database " + "conflict on scan %s. Retrying chunk %s/%s " + "(attempt %s/%s).", + scan_id, + start // SCAN_DB_BATCH_SIZE + 1, + (len(sorted_resources) + SCAN_DB_BATCH_SIZE - 1) + // SCAN_DB_BATCH_SIZE, + attempt + 1, + CELERY_DEADLOCK_ATTEMPTS, + ) + time.sleep((0.1 * (2**attempt)) + random.uniform(0, 0.1)) + continue + raise + + def _copy_compliance_requirement_rows( tenant_id: str, rows: list[dict[str, Any]] ) -> None: @@ -1182,16 +1232,11 @@ def perform_prowler_scan( resources_to_update.append(resource_instance) if resources_to_update: - # Single rls_transaction wrapping the bulk_update (previously - # `update_objects_in_batches` opened one rls_transaction per - # chunk; for tenants with many resources this collapsed N - # BEGINs/COMMITs into 1). - with rls_transaction(tenant_id): - Resource.objects.bulk_update( - resources_to_update, - ["failed_findings_count"], - batch_size=SCAN_DB_BATCH_SIZE, - ) + _bulk_update_resource_failed_findings_counts( + tenant_id=tenant_id, + scan_id=scan_id, + resources_to_update=resources_to_update, + ) except ProviderDeletedException as e: logger.warning(str(e)) diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index e7bb0982cd..7a1ff54131 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -2,6 +2,7 @@ import os from datetime import UTC, datetime, timedelta from pathlib import Path from shutil import rmtree +from uuid import uuid4 from api.compliance import ( get_compliance_frameworks, @@ -10,14 +11,24 @@ from api.compliance import ( from api.db_router import READ_REPLICA_ALIAS from api.db_utils import delete_related_daily_task, rls_transaction from api.decorators import handle_provider_deletion, set_tenant -from api.models import Finding, Integration, Provider, Scan, ScanSummary, StateChoices +from api.models import ( + Finding, + Integration, + Provider, + Scan, + ScanSummary, + StateChoices, + Task, +) from api.utils import initialize_prowler_provider from api.v1.serializers import ScanTaskSerializer -from celery import chain, group, shared_task +from celery import chain, group, shared_task, states from celery.utils.log import get_task_logger from config.celery import RLSTask from config.django.base import DJANGO_FINDINGS_BATCH_SIZE, DJANGO_TMP_OUTPUT_DIRECTORY +from django.db import transaction from django_celery_beat.models import PeriodicTask +from django_celery_results.models import TaskResult from prowler.lib.check.compliance_models import Compliance from prowler.lib.outputs.compliance.compliance import ( process_universal_compliance_frameworks, @@ -85,6 +96,220 @@ from tasks.utils import ( ) logger = get_task_logger(__name__) +QUEUED_SCAN_TASK_STATE = "QUEUED" +DISPATCHED_SCAN_TASK_STATES = (states.PENDING, states.STARTED, "PROGRESS") + + +def _get_dispatched_provider_scan(tenant_id: str, provider_id: str): + """Return a scan that has already been dispatched for a provider.""" + executing_scan = ( + Scan.objects.select_for_update() + .filter( + tenant_id=tenant_id, + provider_id=provider_id, + state=StateChoices.EXECUTING, + ) + .order_by("-inserted_at") + .first() + ) + if executing_scan: + return executing_scan + + return ( + Scan.objects.select_for_update(of=("self",)) + .select_related("task__task_runner_task") + .filter( + tenant_id=tenant_id, + provider_id=provider_id, + state__in=(StateChoices.AVAILABLE, StateChoices.SCHEDULED), + task__isnull=False, + task__task_runner_task__status__in=DISPATCHED_SCAN_TASK_STATES, + ) + .order_by("-inserted_at") + .first() + ) + + +def _get_queued_provider_scan(tenant_id: str, provider_id: str): + """Return the next DB-queued scan for a provider.""" + return ( + Scan.objects.select_for_update(of=("self",)) + .select_related("task__task_runner_task") + .filter( + tenant_id=tenant_id, + provider_id=provider_id, + state=StateChoices.AVAILABLE, + task__isnull=False, + task__task_runner_task__status=QUEUED_SCAN_TASK_STATE, + ) + .order_by("inserted_at", "id") + .first() + ) + + +def get_active_provider_scan(tenant_id: str, provider_id: str): + """Return a dispatched or DB-queued scan for a provider.""" + return _get_dispatched_provider_scan( + tenant_id, provider_id + ) or _get_queued_provider_scan(tenant_id, provider_id) + + +def create_scan_task_record( + tenant_id: str, + task_id: str, + task_name: str = "scan-perform", + task_status: str | None = states.PENDING, +) -> Task: + if task_status is None: + task_status = states.PENDING + + task_result, _ = TaskResult.objects.update_or_create( + task_id=str(task_id), + defaults={"status": task_status, "task_name": task_name}, + ) + prowler_task, _ = Task.objects.update_or_create( + id=str(task_id), + tenant_id=tenant_id, + defaults={"task_runner_task": task_result}, + ) + return prowler_task + + +def enqueue_scan_execution_on_commit( + tenant_id: str, + scan: Scan, + task_id: str, +) -> None: + transaction.on_commit( + lambda: perform_scan_task.apply_async( + kwargs={ + "tenant_id": str(tenant_id), + "scan_id": str(scan.id), + "provider_id": str(scan.provider_id), + }, + task_id=str(task_id), + ) + ) + + +def _get_queued_scheduled_scan(tenant_id: str, provider_id: str): + return ( + Scan.objects.select_for_update(of=("self",)) + .select_related("task__task_runner_task") + .filter( + tenant_id=tenant_id, + provider_id=provider_id, + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.AVAILABLE, + task__isnull=False, + task__task_runner_task__status=QUEUED_SCAN_TASK_STATE, + ) + .order_by("inserted_at", "id") + .first() + ) + + +def _get_or_create_queued_scheduled_scan( + tenant_id: str, + provider_id: str, + periodic_task_instance: PeriodicTask, + scheduled_at: datetime, +) -> Scan: + queued_scan = _get_queued_scheduled_scan(tenant_id, provider_id) + if queued_scan: + return queued_scan + + task_id = str(uuid4()) + queued_task = create_scan_task_record( + tenant_id=tenant_id, + task_id=task_id, + task_status=QUEUED_SCAN_TASK_STATE, + ) + return Scan.objects.create( + tenant_id=tenant_id, + name="Daily scheduled scan", + provider_id=provider_id, + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.AVAILABLE, + scheduled_at=scheduled_at, + scheduler_task_id=periodic_task_instance.id, + task=queued_task, + ) + + +def _dispatch_next_queued_provider_scan(tenant_id: str, provider_id: str): + with rls_transaction(tenant_id): + if not Provider.objects.select_for_update().filter(pk=provider_id).exists(): + return None + + if _get_dispatched_provider_scan(tenant_id, provider_id): + return None + + queued_scan = _get_queued_provider_scan(tenant_id, provider_id) + if not queued_scan or not queued_scan.task: + return None + + task_result = queued_scan.task.task_runner_task + task_result.status = states.PENDING + task_result.task_name = "scan-perform" + task_result.save(update_fields=["status", "task_name"]) + enqueue_scan_execution_on_commit( + tenant_id=tenant_id, + scan=queued_scan, + task_id=str(queued_scan.task_id), + ) + return queued_scan + + +def _dispatch_next_queued_provider_scan_best_effort( + tenant_id: str, provider_id: str +) -> None: + try: + _dispatch_next_queued_provider_scan(tenant_id, provider_id) + except Exception: + logger.exception( + "Failed to dispatch next queued scan for provider %s", provider_id + ) + + +def _get_or_create_next_scheduled_scan( + tenant_id: str, + provider_id: str, + periodic_task_instance: PeriodicTask, + next_scan_datetime: datetime, +) -> Scan: + interval = periodic_task_instance.interval + now = datetime.now(UTC) + while next_scan_datetime <= now: + next_scan_datetime += timedelta(**{interval.period: interval.every}) + + return _get_or_create_scheduled_scan( + tenant_id=tenant_id, + provider_id=provider_id, + scheduler_task_id=periodic_task_instance.id, + scheduled_at=next_scan_datetime, + update_state=True, + ) + + +def _ensure_next_scheduled_scan_best_effort( + tenant_id: str, + provider_id: str, + periodic_task_instance: PeriodicTask, + next_scan_datetime: datetime, +) -> None: + try: + with rls_transaction(tenant_id): + _get_or_create_next_scheduled_scan( + tenant_id=tenant_id, + provider_id=provider_id, + periodic_task_instance=periodic_task_instance, + next_scan_datetime=next_scan_datetime, + ) + except Exception: + logger.exception( + "Failed to ensure next scheduled scan for provider %s", provider_id + ) def _cleanup_orphan_scheduled_scans( @@ -117,6 +342,7 @@ def _cleanup_orphan_scheduled_scans( trigger=Scan.TriggerChoices.SCHEDULED, state=StateChoices.AVAILABLE, scheduler_task_id=scheduler_task_id, + task__isnull=True, ) scheduled_scan_exists = Scan.objects.filter( @@ -292,16 +518,17 @@ def perform_scan_task( ) return None - result = perform_prowler_scan( - tenant_id=tenant_id, - scan_id=scan_id, - provider_id=provider_id, - checks_to_execute=checks_to_execute, - ) - - _perform_scan_complete_tasks(tenant_id, scan_id, provider_id) - - return result + try: + result = perform_prowler_scan( + tenant_id=tenant_id, + scan_id=scan_id, + provider_id=provider_id, + checks_to_execute=checks_to_execute, + ) + _perform_scan_complete_tasks(tenant_id, scan_id, provider_id) + return result + finally: + _dispatch_next_queued_provider_scan_best_effort(tenant_id, provider_id) # acks_late=False: like scan-perform; a dropped run is re-fired by Beat on the next tick. @@ -335,7 +562,7 @@ def perform_scheduled_scan_task(self, tenant_id: str, provider_id: str): task_id = self.request.id with rls_transaction(tenant_id): - if not Provider.objects.filter(pk=provider_id).exists(): + if not Provider.objects.select_for_update().filter(pk=provider_id).exists(): logger.warning( "scheduled scan-perform skipped: provider %s no longer exists " "(tenant=%s)", @@ -348,22 +575,6 @@ def perform_scheduled_scan_task(self, tenant_id: str, provider_id: str): periodic_task_instance = PeriodicTask.objects.get( name=f"scan-perform-scheduled-{provider_id}" ) - executing_scan = ( - Scan.objects.filter( - tenant_id=tenant_id, - provider_id=provider_id, - trigger=Scan.TriggerChoices.SCHEDULED, - state=StateChoices.EXECUTING, - ) - .order_by("-started_at") - .first() - ) - if executing_scan: - logger.warning( - f"Scheduled scan already executing for provider {provider_id}. Skipping." - ) - return ScanTaskSerializer(instance=executing_scan).data - executed_scan = Scan.objects.filter( tenant_id=tenant_id, provider_id=provider_id, @@ -388,6 +599,26 @@ def perform_scheduled_scan_task(self, tenant_id: str, provider_id: str): scheduler_task_id=periodic_task_instance.id, ) + active_scan = get_active_provider_scan(tenant_id, provider_id) + if active_scan: + logger.warning( + "Scan already queued or executing for provider %s. Queueing scheduled run.", + provider_id, + ) + queued_scheduled_scan = _get_or_create_queued_scheduled_scan( + tenant_id=tenant_id, + provider_id=provider_id, + periodic_task_instance=periodic_task_instance, + scheduled_at=current_scan_datetime, + ) + _get_or_create_next_scheduled_scan( + tenant_id=tenant_id, + provider_id=provider_id, + periodic_task_instance=periodic_task_instance, + next_scan_datetime=next_scan_datetime, + ) + return ScanTaskSerializer(instance=queued_scheduled_scan).data + scan_instance = _get_or_create_scheduled_scan( tenant_id=tenant_id, provider_id=provider_id, @@ -403,24 +634,16 @@ def perform_scheduled_scan_task(self, tenant_id: str, provider_id: str): scan_id=str(scan_instance.id), provider_id=provider_id, ) + _perform_scan_complete_tasks(tenant_id, str(scan_instance.id), provider_id) + return result finally: - with rls_transaction(tenant_id): - now = datetime.now(UTC) - if next_scan_datetime <= now: - interval_delta = timedelta(**{interval.period: interval.every}) - while next_scan_datetime <= now: - next_scan_datetime += interval_delta - _get_or_create_scheduled_scan( - tenant_id=tenant_id, - provider_id=provider_id, - scheduler_task_id=periodic_task_instance.id, - scheduled_at=next_scan_datetime, - update_state=True, - ) - - _perform_scan_complete_tasks(tenant_id, str(scan_instance.id), provider_id) - - return result + _ensure_next_scheduled_scan_best_effort( + tenant_id=tenant_id, + provider_id=provider_id, + periodic_task_instance=periodic_task_instance, + next_scan_datetime=next_scan_datetime, + ) + _dispatch_next_queued_provider_scan_best_effort(tenant_id, provider_id) @shared_task(name="scan-summary", queue="overview") diff --git a/api/src/backend/tasks/tests/test_scan.py b/api/src/backend/tasks/tests/test_scan.py index 2d251b247f..2fd2dde05f 100644 --- a/api/src/backend/tasks/tests/test_scan.py +++ b/api/src/backend/tasks/tests/test_scan.py @@ -21,11 +21,13 @@ from api.models import ( StateChoices, StatusChoices, ) +from django.db import IntegrityError, OperationalError from prowler.lib.check.models import Severity from prowler.lib.outputs.finding import Status from tasks.jobs.scan import ( _ATTACK_SURFACE_MAPPING_CACHE, _aggregate_findings_by_region, + _bulk_update_resource_failed_findings_counts, _copy_compliance_requirement_rows, _create_compliance_summaries, _create_finding_delta, @@ -858,6 +860,98 @@ class TestPerformScan: # Assert that failed_findings_count was reset to 0 during the scan assert resource.failed_findings_count == 0 + def test_failed_findings_count_update_retries_deadlock_in_stable_order( + self, resources_fixture, monkeypatch + ): + resource1, resource2, _ = resources_fixture + tenant_id = str(resource1.tenant_id) + resource1.failed_findings_count = 2 + resource2.failed_findings_count = 3 + resources_to_update = [resource2, resource1] + expected_order = [ + str(resource.id) + for resource in sorted(resources_to_update, key=lambda item: str(item.id)) + ] + original_bulk_update = Resource.objects.bulk_update + bulk_update_calls = [] + + def flaky_bulk_update(objects, fields, batch_size=None): + bulk_update_calls.append([str(obj.id) for obj in objects]) + if len(bulk_update_calls) == 1: + raise OperationalError("deadlock detected") + return original_bulk_update(objects, fields, batch_size=batch_size) + + monkeypatch.setattr("tasks.jobs.scan.SCAN_DB_BATCH_SIZE", 10) + monkeypatch.setattr(Resource.objects, "bulk_update", flaky_bulk_update) + + _bulk_update_resource_failed_findings_counts( + tenant_id=tenant_id, + scan_id="scan-id", + resources_to_update=resources_to_update, + ) + + resource1.refresh_from_db() + resource2.refresh_from_db() + assert resource1.failed_findings_count == 2 + assert resource2.failed_findings_count == 3 + assert bulk_update_calls == [expected_order, expected_order] + + def test_failed_findings_count_update_does_not_retry_integrity_error( + self, resources_fixture, monkeypatch + ): + resource, *_ = resources_fixture + resource.failed_findings_count = 2 + bulk_update_calls = [] + sleep_calls = [] + + def failing_bulk_update(objects, fields, batch_size=None): + bulk_update_calls.append([str(obj.id) for obj in objects]) + raise IntegrityError("constraint violation") + + monkeypatch.setattr(Resource.objects, "bulk_update", failing_bulk_update) + monkeypatch.setattr("tasks.jobs.scan.time.sleep", sleep_calls.append) + + with pytest.raises(IntegrityError, match="constraint violation"): + _bulk_update_resource_failed_findings_counts( + tenant_id=str(resource.tenant_id), + scan_id="scan-id", + resources_to_update=[resource], + ) + + assert len(bulk_update_calls) == 1 + assert sleep_calls == [] + + def test_failed_findings_count_update_adds_jitter_to_retry_backoff( + self, resources_fixture, monkeypatch + ): + from tasks.jobs import scan as scan_jobs + + resource, *_ = resources_fixture + resource.failed_findings_count = 2 + bulk_update_calls = [] + sleep_calls = [] + original_bulk_update = Resource.objects.bulk_update + + def flaky_bulk_update(objects, fields, batch_size=None): + bulk_update_calls.append([str(obj.id) for obj in objects]) + if len(bulk_update_calls) == 1: + raise OperationalError("deadlock detected") + return original_bulk_update(objects, fields, batch_size=batch_size) + + monkeypatch.setattr(Resource.objects, "bulk_update", flaky_bulk_update) + monkeypatch.setattr(scan_jobs, "random", MagicMock()) + scan_jobs.random.uniform.return_value = 0.037 + monkeypatch.setattr("tasks.jobs.scan.time.sleep", sleep_calls.append) + + _bulk_update_resource_failed_findings_counts( + tenant_id=str(resource.tenant_id), + scan_id="scan-id", + resources_to_update=[resource], + ) + + scan_jobs.random.uniform.assert_called_once_with(0, 0.1) + assert sleep_calls == [0.137] + def test_perform_prowler_scan_with_active_mute_rules( self, tenants_fixture, diff --git a/api/src/backend/tasks/tests/test_tasks.py b/api/src/backend/tasks/tests/test_tasks.py index 95634e8a95..6475e525d6 100644 --- a/api/src/backend/tasks/tests/test_tasks.py +++ b/api/src/backend/tasks/tests/test_tasks.py @@ -14,6 +14,7 @@ from api.models import ( Task, ) from botocore.exceptions import ClientError +from celery import states from django_celery_beat.models import IntervalSchedule, PeriodicTask from django_celery_results.models import TaskResult from tasks.jobs.lighthouse_providers import ( @@ -2286,6 +2287,51 @@ class TestCleanupOrphanScheduledScans: assert Scan.objects.filter(id=scheduled_scan.id).exists() assert Scan.objects.filter(id=available_scan_other_task.id).exists() + def test_cleanup_keeps_db_queued_scheduled_scans( + self, tenants_fixture, providers_fixture + ): + """DB-queued scheduled scans have a task and must not be deleted as orphans.""" + tenant = tenants_fixture[0] + provider = providers_fixture[0] + periodic_task = self._create_periodic_task(provider.id, tenant.id) + task_result = TaskResult.objects.create( + task_id=str(uuid.uuid4()), + task_name="scan-perform", + status="QUEUED", + ) + queued_task = Task.objects.create( + id=task_result.task_id, + task_runner_task=task_result, + tenant_id=tenant.id, + ) + queued_scan = Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Queued scheduled scan", + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.AVAILABLE, + scheduler_task_id=periodic_task.id, + task=queued_task, + ) + scheduled_scan = Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Daily scheduled scan", + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.SCHEDULED, + scheduler_task_id=periodic_task.id, + ) + + deleted_count = _cleanup_orphan_scheduled_scans( + tenant_id=str(tenant.id), + provider_id=str(provider.id), + scheduler_task_id=periodic_task.id, + ) + + assert deleted_count == 0 + assert Scan.objects.filter(id=queued_scan.id).exists() + assert Scan.objects.filter(id=scheduled_scan.id).exists() + @pytest.mark.django_db class TestPerformScheduledScanTask: @@ -2334,10 +2380,10 @@ class TestPerformScheduledScanTask: ) return task_result - def test_skip_when_scheduled_scan_executing( + def test_queues_scheduled_scan_when_scheduled_scan_is_executing( self, tenants_fixture, providers_fixture ): - """Skip a scheduled run when another scheduled scan is already executing.""" + """Queue a scheduled run when another scheduled scan is executing.""" tenant = tenants_fixture[0] provider = providers_fixture[0] periodic_task = self._create_periodic_task(provider.id, tenant.id) @@ -2364,8 +2410,16 @@ class TestPerformScheduledScanTask: mock_scan.assert_not_called() mock_complete_tasks.assert_not_called() - assert result["id"] == str(executing_scan.id) - assert result["state"] == StateChoices.EXECUTING + assert result["id"] != str(executing_scan.id) + assert result["state"] == StateChoices.AVAILABLE + queued_scheduled_scan = Scan.objects.get( + tenant_id=tenant.id, + provider=provider, + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.AVAILABLE, + ) + assert result["id"] == str(queued_scheduled_scan.id) + assert queued_scheduled_scan.task.task_runner_task.status == "QUEUED" assert ( Scan.objects.filter( tenant_id=tenant.id, @@ -2373,7 +2427,133 @@ class TestPerformScheduledScanTask: trigger=Scan.TriggerChoices.SCHEDULED, state=StateChoices.SCHEDULED, ).count() - == 0 + == 1 + ) + + def test_queues_scheduled_scan_when_manual_scan_is_pending( + self, tenants_fixture, providers_fixture + ): + """Queue one scheduled run when a manual scan is already dispatched.""" + tenant = tenants_fixture[0] + provider = providers_fixture[0] + self._create_periodic_task(provider.id, tenant.id) + task_id = str(uuid.uuid4()) + self._create_task_result(tenant.id, task_id) + manual_task_result = TaskResult.objects.create( + task_id=str(uuid.uuid4()), + task_name="scan-perform", + status=states.PENDING, + ) + manual_task = Task.objects.create( + id=manual_task_result.task_id, + task_runner_task=manual_task_result, + tenant_id=tenant.id, + ) + manual_scan = Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Manual scan", + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.AVAILABLE, + task=manual_task, + ) + + with ( + patch("tasks.tasks.perform_prowler_scan") as mock_scan, + patch("tasks.tasks._perform_scan_complete_tasks") as mock_complete_tasks, + self._override_task_request(perform_scheduled_scan_task, id=task_id), + ): + result = perform_scheduled_scan_task.run( + tenant_id=str(tenant.id), provider_id=str(provider.id) + ) + + mock_scan.assert_not_called() + mock_complete_tasks.assert_not_called() + assert result["id"] != str(manual_scan.id) + assert result["state"] == StateChoices.AVAILABLE + queued_scheduled_scan = Scan.objects.get( + tenant_id=tenant.id, + provider=provider, + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.AVAILABLE, + ) + assert result["id"] == str(queued_scheduled_scan.id) + assert queued_scheduled_scan.task.task_runner_task.status == "QUEUED" + scheduled_scan = Scan.objects.get( + tenant_id=tenant.id, + provider=provider, + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.SCHEDULED, + ) + assert scheduled_scan.scheduled_at > datetime.now(UTC) + + def test_coalesces_scheduled_scan_when_one_is_already_queued( + self, tenants_fixture, providers_fixture + ): + """Reuse the existing queued scheduled scan instead of adding another.""" + tenant = tenants_fixture[0] + provider = providers_fixture[0] + periodic_task = self._create_periodic_task(provider.id, tenant.id) + task_id = str(uuid.uuid4()) + self._create_task_result(tenant.id, task_id) + manual_task_result = TaskResult.objects.create( + task_id=str(uuid.uuid4()), + task_name="scan-perform", + status=states.PENDING, + ) + manual_task = Task.objects.create( + id=manual_task_result.task_id, + task_runner_task=manual_task_result, + tenant_id=tenant.id, + ) + Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Manual scan", + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.AVAILABLE, + task=manual_task, + ) + queued_task_result = TaskResult.objects.create( + task_id=str(uuid.uuid4()), + task_name="scan-perform", + status="QUEUED", + ) + queued_task = Task.objects.create( + id=queued_task_result.task_id, + task_runner_task=queued_task_result, + tenant_id=tenant.id, + ) + queued_scheduled_scan = Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Daily scheduled scan", + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.AVAILABLE, + scheduler_task_id=periodic_task.id, + task=queued_task, + ) + + with ( + patch("tasks.tasks.perform_prowler_scan") as mock_scan, + patch("tasks.tasks._perform_scan_complete_tasks") as mock_complete_tasks, + self._override_task_request(perform_scheduled_scan_task, id=task_id), + ): + result = perform_scheduled_scan_task.run( + tenant_id=str(tenant.id), provider_id=str(provider.id) + ) + + mock_scan.assert_not_called() + mock_complete_tasks.assert_not_called() + assert result["id"] == str(queued_scheduled_scan.id) + assert ( + Scan.objects.filter( + tenant_id=tenant.id, + provider=provider, + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.AVAILABLE, + ).count() + == 1 ) def test_creates_next_scheduled_scan_after_completion( @@ -2435,6 +2615,41 @@ class TestPerformScheduledScanTask: == 1 ) + def test_next_scheduled_scan_failure_does_not_mask_completed_scan( + self, tenants_fixture, providers_fixture, caplog + ): + """Keep scheduled scan success when next-run creation fails.""" + tenant = tenants_fixture[0] + provider = providers_fixture[0] + self._create_periodic_task(provider.id, tenant.id) + task_id = str(uuid.uuid4()) + self._create_task_result(tenant.id, task_id) + + def _complete_scan(tenant_id, scan_id, provider_id): + scan_instance = Scan.objects.get(id=scan_id) + scan_instance.state = StateChoices.COMPLETED + scan_instance.save() + return {"status": "ok"} + + with ( + patch("tasks.tasks.perform_prowler_scan", side_effect=_complete_scan), + patch("tasks.tasks._perform_scan_complete_tasks"), + patch( + "tasks.tasks._get_or_create_next_scheduled_scan", + side_effect=RuntimeError("scheduler unavailable"), + ), + patch("tasks.tasks._dispatch_next_queued_provider_scan") as mock_dispatch, + self._override_task_request(perform_scheduled_scan_task, id=task_id), + caplog.at_level("ERROR"), + ): + result = perform_scheduled_scan_task.run( + tenant_id=str(tenant.id), provider_id=str(provider.id) + ) + + assert result == {"status": "ok"} + mock_dispatch.assert_called_once_with(str(tenant.id), str(provider.id)) + assert "Failed to ensure next scheduled scan" in caplog.text + def test_dedupes_multiple_scheduled_scans_before_run( self, tenants_fixture, providers_fixture ): @@ -2549,6 +2764,104 @@ class TestPerformScanTask: mock_scan.assert_not_called() mock_complete_tasks.assert_not_called() + def test_dispatches_next_queued_scan_after_completion( + self, + tenants_fixture, + providers_fixture, + django_capture_on_commit_callbacks, + ): + """Dispatch the next queued scan for the provider after completion.""" + tenant = tenants_fixture[0] + provider = providers_fixture[0] + current_scan = Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Running scan", + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.AVAILABLE, + ) + queued_task_result = TaskResult.objects.create( + task_id=str(uuid.uuid4()), + task_name="scan-perform", + status="QUEUED", + ) + queued_task = Task.objects.create( + id=queued_task_result.task_id, + task_runner_task=queued_task_result, + tenant_id=tenant.id, + ) + queued_scan = Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Queued scan", + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.AVAILABLE, + task=queued_task, + ) + + def _complete_scan(tenant_id, scan_id, provider_id, checks_to_execute=None): + scan_instance = Scan.objects.get(id=scan_id) + scan_instance.state = StateChoices.COMPLETED + scan_instance.save() + return {"status": "ok"} + + with ( + patch("tasks.tasks.perform_prowler_scan", side_effect=_complete_scan), + patch("tasks.tasks._perform_scan_complete_tasks"), + patch("tasks.tasks.perform_scan_task.apply_async") as mock_apply_async, + ): + with django_capture_on_commit_callbacks(execute=True): + result = perform_scan_task.run( + tenant_id=str(tenant.id), + scan_id=str(current_scan.id), + provider_id=str(provider.id), + ) + + queued_task_result.refresh_from_db() + assert result == {"status": "ok"} + assert queued_task_result.status == states.PENDING + mock_apply_async.assert_called_once_with( + kwargs={ + "tenant_id": str(tenant.id), + "scan_id": str(queued_scan.id), + "provider_id": str(provider.id), + }, + task_id=str(queued_task.id), + ) + + def test_dispatch_failure_does_not_mask_completed_scan( + self, tenants_fixture, providers_fixture, caplog + ): + """Keep scan success when queued dispatch fails after completion.""" + tenant = tenants_fixture[0] + provider = providers_fixture[0] + current_scan = Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Running scan", + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.AVAILABLE, + ) + + with ( + patch("tasks.tasks.perform_prowler_scan", return_value={"status": "ok"}), + patch("tasks.tasks._perform_scan_complete_tasks"), + patch( + "tasks.tasks._dispatch_next_queued_provider_scan", + side_effect=RuntimeError("dispatch unavailable"), + ) as mock_dispatch, + caplog.at_level("ERROR"), + ): + result = perform_scan_task.run( + tenant_id=str(tenant.id), + scan_id=str(current_scan.id), + provider_id=str(provider.id), + ) + + assert result == {"status": "ok"} + mock_dispatch.assert_called_once_with(str(tenant.id), str(provider.id)) + assert "Failed to dispatch next queued scan" in caplog.text + @pytest.mark.django_db class TestReaggregateAllFindingGroupSummaries: diff --git a/api/src/backend/tasks/utils.py b/api/src/backend/tasks/utils.py index 26bc031d7b..cac9a10f0d 100644 --- a/api/src/backend/tasks/utils.py +++ b/api/src/backend/tasks/utils.py @@ -103,6 +103,7 @@ def _get_or_create_scheduled_scan( trigger=Scan.TriggerChoices.SCHEDULED, state__in=(StateChoices.SCHEDULED, StateChoices.AVAILABLE), scheduler_task_id=scheduler_task_id, + task__isnull=True, ).order_by("scheduled_at", "inserted_at") )