diff --git a/api/changelog.d/scan-release-stale.fixed.md b/api/changelog.d/scan-release-stale.fixed.md new file mode 100644 index 0000000000..c7a986d4ee --- /dev/null +++ b/api/changelog.d/scan-release-stale.fixed.md @@ -0,0 +1 @@ +Scans whose worker was killed mid-run no longer block their provider: a scan whose task already failed, whose worker no longer answers, or that shows no progress for 12 hours is marked failed and the next queued scan runs diff --git a/api/src/backend/api/migrations/0101_scan_release_stale_periodic_task.py b/api/src/backend/api/migrations/0101_scan_release_stale_periodic_task.py new file mode 100644 index 0000000000..24f67e1c2b --- /dev/null +++ b/api/src/backend/api/migrations/0101_scan_release_stale_periodic_task.py @@ -0,0 +1,48 @@ +from django.db import migrations + +TASK_NAME = "scan-release-stale" +INTERVAL_MINUTES = 5 + + +def create_periodic_task(apps, _schema_editor): + IntervalSchedule = apps.get_model("django_celery_beat", "IntervalSchedule") + PeriodicTask = apps.get_model("django_celery_beat", "PeriodicTask") + + schedule, _ = IntervalSchedule.objects.get_or_create( + every=INTERVAL_MINUTES, + period="minutes", + ) + + PeriodicTask.objects.update_or_create( + name=TASK_NAME, + defaults={ + "task": TASK_NAME, + "interval": schedule, + "enabled": True, + }, + ) + + +def delete_periodic_task(apps, _schema_editor): + IntervalSchedule = apps.get_model("django_celery_beat", "IntervalSchedule") + PeriodicTask = apps.get_model("django_celery_beat", "PeriodicTask") + + PeriodicTask.objects.filter(name=TASK_NAME).delete() + + # Clean up the schedule if no other task references it + IntervalSchedule.objects.filter( + every=INTERVAL_MINUTES, + period="minutes", + periodictask__isnull=True, + ).delete() + + +class Migration(migrations.Migration): + dependencies = [ + ("api", "0100_attack_paths_tmp_db_reap_periodic_task"), + ("django_celery_beat", "0019_alter_periodictasks_options"), + ] + + operations = [ + migrations.RunPython(create_periodic_task, delete_periodic_task), + ] diff --git a/api/src/backend/api/tests/test_scan_dead_release_view.py b/api/src/backend/api/tests/test_scan_dead_release_view.py new file mode 100644 index 0000000000..a55fb41e33 --- /dev/null +++ b/api/src/backend/api/tests/test_scan_dead_release_view.py @@ -0,0 +1,148 @@ +import uuid +from datetime import UTC, datetime, timedelta +from unittest.mock import patch + +import pytest +from api.models import Scan, StateChoices, Task +from celery import states +from django.urls import reverse +from django_celery_results.models import TaskResult +from rest_framework import status + +API_JSON_CONTENT_TYPE = "application/vnd.api+json" + + +def _task(tenant_id, task_status): + task_result = TaskResult.objects.create( + task_id=str(uuid.uuid4()), task_name="scan-perform", status=task_status + ) + return Task.objects.create( + id=task_result.task_id, task_runner_task=task_result, tenant_id=tenant_id + ) + + +def _dead_executing_scan(tenant, provider): + return Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Killed scan", + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.EXECUTING, + started_at=datetime.now(UTC) - timedelta(hours=2), + task=_task(tenant.id, states.FAILURE), + ) + + +def _post_scan(client, provider): + return client.post( + reverse("scan-list"), + data={ + "data": { + "type": "scans", + "attributes": {"name": "New Scan"}, + "relationships": { + "provider": {"data": {"type": "providers", "id": str(provider.id)}} + }, + } + }, + content_type=API_JSON_CONTENT_TYPE, + ) + + +@pytest.mark.django_db +class TestScanCreateReleasesDeadScan: + @patch("tasks.tasks.perform_scan_task.apply_async") + def test_dead_scan_does_not_block_new_scan( + self, + mock_apply_async, + authenticated_client, + tenants_fixture, + aws_provider, + django_capture_on_commit_callbacks, + ): + dead = _dead_executing_scan(tenants_fixture[0], aws_provider) + + with ( + patch("tasks.jobs.dead_scans.ping_workers") as ping, + django_capture_on_commit_callbacks(execute=True), + ): + response = _post_scan(authenticated_client, aws_provider) + + ping.assert_not_called() + + assert response.status_code == status.HTTP_202_ACCEPTED + dead.refresh_from_db() + assert dead.state == StateChoices.FAILED + new_scan = Scan.objects.exclude(id=dead.id).get() + assert new_scan.task.task_runner_task.status == states.PENDING + mock_apply_async.assert_called_once() + assert mock_apply_async.call_args.kwargs["kwargs"]["scan_id"] == str( + new_scan.id + ) + + @patch("tasks.tasks.perform_scan_task.apply_async") + def test_queued_scan_behind_dead_scan_runs_first_and_new_scan_queues( + self, + mock_apply_async, + authenticated_client, + tenants_fixture, + aws_provider, + django_capture_on_commit_callbacks, + ): + tenant = tenants_fixture[0] + dead = _dead_executing_scan(tenant, aws_provider) + queued = Scan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + name="Queued scan", + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.AVAILABLE, + task=_task(tenant.id, "QUEUED"), + ) + + with ( + patch("tasks.jobs.dead_scans.ping_workers") as ping, + django_capture_on_commit_callbacks(execute=True), + ): + response = _post_scan(authenticated_client, aws_provider) + + ping.assert_not_called() + + assert response.status_code == status.HTTP_202_ACCEPTED + dead.refresh_from_db() + queued.task.task_runner_task.refresh_from_db() + assert dead.state == StateChoices.FAILED + assert queued.task.task_runner_task.status == states.PENDING + mock_apply_async.assert_called_once() + assert mock_apply_async.call_args.kwargs["kwargs"]["scan_id"] == str(queued.id) + new_scan = Scan.objects.exclude(id__in=(dead.id, queued.id)).get() + assert new_scan.task.task_runner_task.status == "QUEUED" + + @patch("tasks.tasks.perform_scan_task.apply_async") + def test_live_scan_still_queues_new_scan( + self, + mock_apply_async, + authenticated_client, + tenants_fixture, + aws_provider, + django_capture_on_commit_callbacks, + ): + live = _dead_executing_scan(tenants_fixture[0], aws_provider) + TaskResult.objects.filter(pk=live.task.task_runner_task.pk).update( + status=states.STARTED + ) + + with ( + patch("tasks.jobs.dead_scans.ping_workers") as ping, + django_capture_on_commit_callbacks(execute=True), + ): + response = _post_scan(authenticated_client, aws_provider) + + ping.assert_not_called() + + assert response.status_code == status.HTTP_202_ACCEPTED + live.refresh_from_db() + assert live.state == StateChoices.EXECUTING + new_scan = Scan.objects.exclude(id=live.id).get() + assert new_scan.task.task_runner_task.status == "QUEUED" + mock_apply_async.assert_not_called() diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index 5dc175dd42..dab9be431c 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -333,6 +333,7 @@ from tasks.jobs.attack_paths import db_utils as attack_paths_db_utils from tasks.jobs.export import get_s3_client, get_s3_presign_client from tasks.tasks import ( QUEUED_SCAN_TASK_STATE, + _release_provider_scan_slot, backfill_compliance_summaries_task, backfill_scan_resource_summaries_task, check_integration_connection_task, @@ -2808,6 +2809,7 @@ class ScanViewSet(ProviderVisibilityMixin, BaseRLSViewSet): tenant_id=self.request.tenant_id, id__in=self.get_provider_queryset().values("id"), ) + _release_provider_scan_slot(self.request.tenant_id, provider.id) active_scan = get_active_provider_scan( self.request.tenant_id, provider.id ) diff --git a/api/src/backend/config/django/base.py b/api/src/backend/config/django/base.py index 664fcd1c4a..f5ec7356a8 100644 --- a/api/src/backend/config/django/base.py +++ b/api/src/backend/config/django/base.py @@ -321,6 +321,12 @@ UI_BASE_URL = env.str("DJANGO_UI_BASE_URL", "").rstrip("/") CSRF_COOKIE_SECURE = True SESSION_COOKIE_SECURE = True +# A scan is dead when its task already finished, its worker stopped answering, or it +# shows no progress (`updated_at`) for the backstop; a dispatched scan that never +# started is dead after the dispatch age. +SCAN_STALE_BACKSTOP_HOURS = env.int("SCAN_STALE_BACKSTOP_HOURS", 12) +SCAN_DISPATCH_STALE_HOURS = env.int("SCAN_DISPATCH_STALE_HOURS", 24) + # Attack Paths ATTACK_PATHS_SCAN_INACTIVITY_THRESHOLD_MINUTES = env.int( "ATTACK_PATHS_SCAN_INACTIVITY_THRESHOLD_MINUTES", 30 diff --git a/api/src/backend/tasks/jobs/dead_scans.py b/api/src/backend/tasks/jobs/dead_scans.py new file mode 100644 index 0000000000..9b6f785975 --- /dev/null +++ b/api/src/backend/tasks/jobs/dead_scans.py @@ -0,0 +1,69 @@ +from datetime import UTC, datetime, timedelta + +from api.db_router import MainRouter +from api.models import Scan, StateChoices +from celery import states +from celery.utils.log import get_task_logger +from config.django.base import SCAN_DISPATCH_STALE_HOURS, SCAN_STALE_BACKSTOP_HOURS +from django.db.models import Q +from django_celery_results.models import TaskResult +from tasks.jobs.attack_paths.cleanup import _ping_workers as ping_workers + +logger = get_task_logger(__name__) + +DISPATCHED_SCAN_TASK_STATES = (states.PENDING, states.STARTED, "PROGRESS") + +_TASK_STATUS = "task__task_runner_task__status" + + +def dead_scan_q(now: datetime) -> Q: + """Single DB-only definition of a dead scan: it will never finish on its own.""" + backstop = now - timedelta(hours=SCAN_STALE_BACKSTOP_HOURS) + dispatch_cutoff = now - timedelta(hours=SCAN_DISPATCH_STALE_HOURS) + + executing = Q(state=StateChoices.EXECUTING) & ( + Q(**{f"{_TASK_STATUS}__in": states.READY_STATES}) | Q(updated_at__lt=backstop) + ) + never_started = Q( + state__in=(StateChoices.AVAILABLE, StateChoices.SCHEDULED), + task__isnull=False, + **{f"{_TASK_STATUS}__in": DISPATCHED_SCAN_TASK_STATES}, + task__task_runner_task__date_created__lt=dispatch_cutoff, + ) + return executing | never_started + + +def fail_unresponsive_scan_tasks() -> int: + """Mark the task of every executing scan whose worker no longer answers as failed. + + Workers with unknown liveness (ping error), responsive workers and scans with no + recorded worker are left alone; the latter fall to the staleness backstop. + """ + rows = list( + Scan.all_objects.using(MainRouter.admin_db) + .filter(state=StateChoices.EXECUTING, task__task_runner_task__isnull=False) + .exclude(**{f"{_TASK_STATUS}__in": states.READY_STATES}) + .exclude(task__task_runner_task__worker__isnull=True) + .exclude(task__task_runner_task__worker="") + .values_list("task__task_runner_task_id", "task__task_runner_task__worker") + ) + workers = {worker for _, worker in rows} + if not workers: + return 0 + + _, unresponsive = ping_workers(workers) + if not unresponsive: + return 0 + + updated = ( + TaskResult.objects.using(MainRouter.admin_db) + .filter(id__in=[task_id for task_id, worker in rows if worker in unresponsive]) + .exclude(status__in=states.READY_STATES) + .update(status=states.FAILURE, date_done=datetime.now(UTC)) + ) + logger.warning( + "Marked %s task(s) failed for %s unresponsive worker(s)", + updated, + len(unresponsive), + ) + return updated diff --git a/api/src/backend/tasks/jobs/orphan_recovery.py b/api/src/backend/tasks/jobs/orphan_recovery.py index 7290fb313f..f184317bcf 100644 --- a/api/src/backend/tasks/jobs/orphan_recovery.py +++ b/api/src/backend/tasks/jobs/orphan_recovery.py @@ -86,6 +86,7 @@ _SKIP_RECOVERY = { "attack-paths-cleanup-stale-scans", "attack-paths-reap-orphaned-tmp-databases", "reconcile-orphan-tasks", + "scan-release-stale", } diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index 9ec6526963..e5393a347a 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -9,7 +9,7 @@ from api.compliance import ( get_compliance_frameworks, get_prowler_provider_compliance, ) -from api.db_router import READ_REPLICA_ALIAS +from api.db_router import READ_REPLICA_ALIAS, MainRouter from api.db_utils import delete_related_daily_task, rls_transaction from api.decorators import handle_provider_deletion, set_tenant from api.exceptions import ProviderDeletedException @@ -29,6 +29,7 @@ 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.db.models import Q from django_celery_beat.models import PeriodicTask from django_celery_results.models import TaskResult from prowler.lib.check.compliance_models import Compliance @@ -58,6 +59,11 @@ from tasks.jobs.connection import ( check_lighthouse_connection, check_provider_connection, ) +from tasks.jobs.dead_scans import ( + DISPATCHED_SCAN_TASK_STATES, + dead_scan_q, + fail_unresponsive_scan_tasks, +) from tasks.jobs.deletion import delete_provider, delete_tenant from tasks.jobs.export import ( COMPLIANCE_CLASS_MAP, @@ -103,18 +109,19 @@ 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.""" + """Return a live scan that has already been dispatched for a provider.""" + dead = dead_scan_q(datetime.now(UTC)) executing_scan = ( - Scan.objects.select_for_update() + Scan.objects.select_for_update(of=("self",)) .filter( tenant_id=tenant_id, provider_id=provider_id, state=StateChoices.EXECUTING, ) + .exclude(dead) .order_by("-inserted_at") .first() ) @@ -131,6 +138,7 @@ def _get_dispatched_provider_scan(tenant_id: str, provider_id: str): task__isnull=False, task__task_runner_task__status__in=DISPATCHED_SCAN_TASK_STATES, ) + .exclude(dead) .order_by("-inserted_at") .first() ) @@ -258,28 +266,64 @@ def _get_or_create_queued_scheduled_scan( ) +def _dispatch_queued_provider_scan_locked(tenant_id: str, provider_id: str): + """Dispatch the oldest queued scan; the caller holds the provider lock.""" + 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(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 + return _dispatch_queued_provider_scan_locked(tenant_id, provider_id) - 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), +def _release_provider_scan_slot(tenant_id: str, provider_id: str): + """Fail the provider's dead scans and dispatch the next queued one. + + Must run inside a transaction that already holds the provider lock. + Returns the dispatched scan, or None. + """ + now = datetime.now(UTC) + dead_scans = list( + Scan.objects.select_for_update(of=("self",)) + .select_related("task__task_runner_task") + .filter(tenant_id=tenant_id, provider_id=provider_id) + .filter(dead_scan_q(now)) + ) + for scan in dead_scans: + logger.warning( + "Scan %s of provider %s has no live worker; marking it failed", + scan.id, + provider_id, ) - return queued_scan + scan.state = StateChoices.FAILED + scan.completed_at = now + scan.save(update_fields=["state", "completed_at", "updated_at"]) + task_result = scan.task.task_runner_task if scan.task else None + if task_result and task_result.status not in states.READY_STATES: + task_result.status = states.FAILURE + task_result.date_done = now + task_result.save(update_fields=["status", "date_done"]) + + return _dispatch_queued_provider_scan_locked(tenant_id, provider_id) def _dispatch_next_queued_provider_scan_best_effort( @@ -293,6 +337,51 @@ def _dispatch_next_queued_provider_scan_best_effort( ) +def release_stale_scans() -> dict: + """Run the per-provider healer for every provider with a dead or queued scan.""" + dispatched = failed = 0 + try: + unresponsive_tasks = fail_unresponsive_scan_tasks() + except Exception: + unresponsive_tasks = 0 + logger.exception("Failed to check scan workers for liveness") + + now = datetime.now(UTC) + queued = Q( + state=StateChoices.AVAILABLE, + task__isnull=False, + task__task_runner_task__status=QUEUED_SCAN_TASK_STATE, + ) + candidates = list( + Scan.all_objects.using(MainRouter.admin_db) + .filter(dead_scan_q(now) | queued) + .values_list("tenant_id", "provider_id") + .distinct() + ) + for tenant_id, provider_id in candidates: + try: + with rls_transaction(str(tenant_id)): + if ( + not Provider.objects.select_for_update() + .filter(pk=provider_id) + .exists() + ): + continue + if _release_provider_scan_slot(str(tenant_id), str(provider_id)): + dispatched += 1 + except Exception: + failed += 1 + logger.exception( + "Failed to release stale scans for provider %s", provider_id + ) + return { + "providers_checked": len(candidates), + "unresponsive_tasks": unresponsive_tasks, + "dispatched": dispatched, + "failed": failed, + } + + def _get_or_create_next_scheduled_scan( tenant_id: str, provider_id: str, @@ -621,6 +710,17 @@ def perform_scheduled_scan_task(self, tenant_id: str, provider_id: str): scheduler_task_id=periodic_task_instance.id, ) + released_scan = _release_provider_scan_slot(tenant_id, provider_id) + if released_scan and released_scan.trigger == Scan.TriggerChoices.SCHEDULED: + # The released queued scan is this tick's run. + _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=released_scan).data + active_scan = get_active_provider_scan(tenant_id, provider_id) if active_scan: logger.warning( @@ -735,6 +835,12 @@ def reap_orphaned_attack_paths_tmp_databases_task(): return reap_orphaned_tmp_databases() +@shared_task(name="scan-release-stale", queue="celery") +def release_stale_scans_task(): + """Periodic watchdog: fail dead scans and unblock providers nobody requests.""" + return release_stale_scans() + + @shared_task(name="reconcile-orphan-tasks", queue="celery") def reconcile_orphan_tasks_task(): """Periodic watchdog: recover tasks whose worker is gone (deploys, crashes).""" diff --git a/api/src/backend/tasks/tests/test_scan_release_stale.py b/api/src/backend/tasks/tests/test_scan_release_stale.py new file mode 100644 index 0000000000..ec1be61372 --- /dev/null +++ b/api/src/backend/tasks/tests/test_scan_release_stale.py @@ -0,0 +1,519 @@ +import uuid +from datetime import UTC, datetime, timedelta +from unittest.mock import patch + +import pytest +from api.models import Scan, StateChoices, Task +from celery import states +from django_celery_beat.models import IntervalSchedule, PeriodicTask +from django_celery_results.models import TaskResult +from tasks.jobs.dead_scans import dead_scan_q +from tasks.tasks import ( + _release_provider_scan_slot, + perform_scheduled_scan_task, + release_stale_scans, +) + +PING = "tasks.jobs.dead_scans.ping_workers" +BACKSTOP_OVER = timedelta(hours=12, minutes=5) +BACKSTOP_UNDER = timedelta(hours=11, minutes=55) +DISPATCH_OVER = timedelta(hours=24, minutes=5) +DISPATCH_UNDER = timedelta(hours=23, minutes=55) + + +def _task(tenant_id, status, worker=None, date_created=None): + task_result = TaskResult.objects.create( + task_id=str(uuid.uuid4()), + task_name="scan-perform", + status=status, + worker=worker, + ) + if date_created: + TaskResult.objects.filter(pk=task_result.pk).update(date_created=date_created) + return Task.objects.create( + id=task_result.task_id, task_runner_task=task_result, tenant_id=tenant_id + ) + + +def _executing( + tenant, + provider, + task_status=states.FAILURE, + worker=None, + idle=None, + trigger=Scan.TriggerChoices.MANUAL, +): + scan = Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Executing scan", + trigger=trigger, + state=StateChoices.EXECUTING, + started_at=datetime.now(UTC) - timedelta(hours=1), + task=_task(tenant.id, task_status, worker=worker), + ) + if idle: + Scan.objects.filter(pk=scan.pk).update(updated_at=datetime.now(UTC) - idle) + return scan + + +def _queued(tenant, provider, trigger=Scan.TriggerChoices.MANUAL): + return Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + name="Queued scan", + trigger=trigger, + state=StateChoices.AVAILABLE, + task=_task(tenant.id, "QUEUED"), + ) + + +def _dispatched(tenant, provider, task_status, task_age): + return Scan.objects.create( + tenant_id=tenant.id, + provider=provider, + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.AVAILABLE, + task=_task(tenant.id, task_status, date_created=datetime.now(UTC) - task_age), + ) + + +def _dead_ids(tenant_id): + return set( + Scan.objects.filter(tenant_id=tenant_id) + .filter(dead_scan_q(datetime.now(UTC))) + .values_list("id", flat=True) + ) + + +@pytest.mark.django_db +class TestDeadScanQ: + @pytest.mark.parametrize("task_status", sorted(states.READY_STATES)) + def test_executing_with_finished_task_is_dead( + self, task_status, tenants_fixture, aws_provider + ): + scan = _executing(tenants_fixture[0], aws_provider, task_status=task_status) + assert _dead_ids(tenants_fixture[0].id) == {scan.id} + + def test_executing_with_running_task_is_alive(self, tenants_fixture, aws_provider): + _executing(tenants_fixture[0], aws_provider, task_status=states.STARTED) + assert _dead_ids(tenants_fixture[0].id) == set() + + def test_backstop_over_and_under_twelve_hours(self, tenants_fixture, aws_provider): + tenant = tenants_fixture[0] + over = _executing( + tenant, aws_provider, task_status=states.STARTED, idle=BACKSTOP_OVER + ) + _executing( + tenant, aws_provider, task_status=states.STARTED, idle=BACKSTOP_UNDER + ) + assert _dead_ids(tenant.id) == {over.id} + + def test_dispatched_never_started_over_and_under_twenty_four_hours( + self, tenants_fixture, aws_provider + ): + tenant = tenants_fixture[0] + over = _dispatched(tenant, aws_provider, states.PENDING, DISPATCH_OVER) + _dispatched(tenant, aws_provider, states.PENDING, DISPATCH_UNDER) + assert _dead_ids(tenant.id) == {over.id} + + def test_queued_and_completed_scans_are_not_dead( + self, tenants_fixture, aws_provider + ): + tenant = tenants_fixture[0] + _queued(tenant, aws_provider) + Scan.objects.create( + tenant_id=tenant.id, + provider=aws_provider, + trigger=Scan.TriggerChoices.MANUAL, + state=StateChoices.COMPLETED, + ) + assert _dead_ids(tenant.id) == set() + + +@pytest.mark.django_db +class TestReleaseProviderScanSlot: + def test_fails_dead_scan_and_dispatches_queued( + self, tenants_fixture, aws_provider, django_capture_on_commit_callbacks + ): + tenant = tenants_fixture[0] + dead = _executing(tenant, aws_provider, task_status=states.STARTED) + TaskResult.objects.filter(pk=dead.task.task_runner_task.pk).update( + status=states.FAILURE + ) + queued = _queued(tenant, aws_provider) + + with patch("tasks.tasks.perform_scan_task.apply_async") as publish: + with django_capture_on_commit_callbacks(execute=True): + released = _release_provider_scan_slot( + str(tenant.id), str(aws_provider.id) + ) + + assert released.id == queued.id + publish.assert_called_once() + dead.refresh_from_db() + assert dead.state == StateChoices.FAILED + assert dead.completed_at is not None + queued.task.task_runner_task.refresh_from_db() + assert queued.task.task_runner_task.status == states.PENDING + + def test_fails_dead_scan_without_queue(self, tenants_fixture, aws_provider): + tenant = tenants_fixture[0] + dead = _executing(tenant, aws_provider) + + assert _release_provider_scan_slot(str(tenant.id), str(aws_provider.id)) is None + dead.refresh_from_db() + assert dead.state == StateChoices.FAILED + + def test_backstop_scan_gets_task_marked_failed(self, tenants_fixture, aws_provider): + tenant = tenants_fixture[0] + dead = _executing( + tenant, aws_provider, task_status=states.STARTED, idle=BACKSTOP_OVER + ) + + _release_provider_scan_slot(str(tenant.id), str(aws_provider.id)) + + dead.refresh_from_db() + dead.task.task_runner_task.refresh_from_db() + assert dead.state == StateChoices.FAILED + assert dead.task.task_runner_task.status == states.FAILURE + assert dead.task.task_runner_task.date_done is not None + + def test_live_scan_is_kept_and_queue_stays(self, tenants_fixture, aws_provider): + tenant = tenants_fixture[0] + live = _executing( + tenant, aws_provider, task_status=states.STARTED, idle=BACKSTOP_UNDER + ) + queued = _queued(tenant, aws_provider) + + assert _release_provider_scan_slot(str(tenant.id), str(aws_provider.id)) is None + live.refresh_from_db() + queued.task.task_runner_task.refresh_from_db() + assert live.state == StateChoices.EXECUTING + assert queued.task.task_runner_task.status == "QUEUED" + + def test_stale_dispatched_scan_is_failed_and_recent_one_kept( + self, tenants_fixture, aws_provider + ): + tenant = tenants_fixture[0] + old = _dispatched(tenant, aws_provider, states.PENDING, DISPATCH_OVER) + recent = _dispatched(tenant, aws_provider, states.PENDING, DISPATCH_UNDER) + + _release_provider_scan_slot(str(tenant.id), str(aws_provider.id)) + + old.refresh_from_db() + recent.refresh_from_db() + assert old.state == StateChoices.FAILED + assert recent.state == StateChoices.AVAILABLE + + +@pytest.mark.django_db +class TestScheduledScanHealsDeadScan: + def _periodic_task(self, provider_id, tenant_id): + interval, _ = IntervalSchedule.objects.get_or_create(every=24, period="hours") + return PeriodicTask.objects.create( + name=f"scan-perform-scheduled-{provider_id}", + task="scan-perform-scheduled", + interval=interval, + kwargs=f'{{"tenant_id": "{tenant_id}", "provider_id": "{provider_id}"}}', + enabled=True, + ) + + def _run(self, tenant, provider, task_id): + task_result = TaskResult.objects.create( + task_id=task_id, + task_name="scan-perform-scheduled", + status="STARTED", + date_created=datetime.now(UTC), + ) + Task.objects.create( + id=task_id, task_runner_task=task_result, tenant_id=tenant.id + ) + request = perform_scheduled_scan_task.request + previous = getattr(request, "id", None) + request.id = task_id + try: + with ( + patch("tasks.tasks.perform_prowler_scan") as scan, + patch("tasks.tasks.reconcile_scan_mute_rules"), + patch("tasks.tasks._perform_scan_complete_tasks"), + patch(PING) as ping, + ): + result = perform_scheduled_scan_task.run( + tenant_id=str(tenant.id), provider_id=str(provider.id) + ) + ping.assert_not_called() + return result, scan + finally: + request.id = previous + + def test_runs_scheduled_scan_when_dead_scan_blocks_provider( + self, tenants_fixture, aws_provider + ): + tenant = tenants_fixture[0] + self._periodic_task(aws_provider.id, tenant.id) + dead = _executing(tenant, aws_provider) + + result, scan = self._run(tenant, aws_provider, str(uuid.uuid4())) + + dead.refresh_from_db() + assert dead.state == StateChoices.FAILED + scan.assert_called_once() + assert result is not None + + def test_dead_scan_with_queued_scheduled_scan_returns_that_scan( + self, tenants_fixture, aws_provider, django_capture_on_commit_callbacks + ): + tenant = tenants_fixture[0] + self._periodic_task(aws_provider.id, tenant.id) + _executing(tenant, aws_provider) + queued = _queued(tenant, aws_provider, trigger=Scan.TriggerChoices.SCHEDULED) + + with patch("tasks.tasks.perform_scan_task.apply_async") as publish: + with django_capture_on_commit_callbacks(execute=True): + result, scan = self._run(tenant, aws_provider, str(uuid.uuid4())) + + assert result["id"] == str(queued.id) + publish.assert_called_once() + scan.assert_not_called() + assert ( + Scan.objects.filter( + provider=aws_provider, + trigger=Scan.TriggerChoices.SCHEDULED, + state=StateChoices.AVAILABLE, + ).count() + == 1 + ) + assert Scan.objects.filter( + provider=aws_provider, state=StateChoices.SCHEDULED + ).exists() + + def test_dead_scan_with_queued_manual_scan_queues_scheduled_run( + self, tenants_fixture, aws_provider, django_capture_on_commit_callbacks + ): + tenant = tenants_fixture[0] + self._periodic_task(aws_provider.id, tenant.id) + _executing(tenant, aws_provider) + manual = _queued(tenant, aws_provider) + + with patch("tasks.tasks.perform_scan_task.apply_async") as publish: + with django_capture_on_commit_callbacks(execute=True): + result, scan = self._run(tenant, aws_provider, str(uuid.uuid4())) + + publish.assert_called_once() + scan.assert_not_called() + assert result["id"] != str(manual.id) + queued_scheduled = Scan.objects.get(id=result["id"]) + assert queued_scheduled.task.task_runner_task.status == "QUEUED" + + +@pytest.mark.django_db +class TestReleaseStaleScansSweeper: + def test_fails_dead_scan_without_queue(self, tenants_fixture, aws_provider): + tenant = tenants_fixture[0] + dead = _executing(tenant, aws_provider) + + with patch(PING) as ping: + counts = release_stale_scans() + + ping.assert_not_called() + dead.refresh_from_db() + assert dead.state == StateChoices.FAILED + assert counts["dispatched"] == 0 + assert counts["failed"] == 0 + + def test_dispatches_queued_scan_behind_dead_scan( + self, tenants_fixture, aws_provider, django_capture_on_commit_callbacks + ): + tenant = tenants_fixture[0] + dead = _executing(tenant, aws_provider) + queued = _queued(tenant, aws_provider) + + with patch("tasks.tasks.perform_scan_task.apply_async") as publish: + with django_capture_on_commit_callbacks(execute=True): + counts = release_stale_scans() + + dead.refresh_from_db() + queued.task.task_runner_task.refresh_from_db() + assert dead.state == StateChoices.FAILED + assert queued.task.task_runner_task.status == states.PENDING + assert counts["dispatched"] == 1 + publish.assert_called_once() + + def test_unresponsive_worker_scan_is_released_in_the_same_run( + self, tenants_fixture, aws_provider, django_capture_on_commit_callbacks + ): + tenant = tenants_fixture[0] + dead = _executing( + tenant, aws_provider, task_status=states.STARTED, worker="w-dead@host" + ) + queued = _queued(tenant, aws_provider) + + with ( + patch(PING, return_value=(set(), {"w-dead@host"})) as ping, + patch("tasks.tasks.perform_scan_task.apply_async") as publish, + ): + with django_capture_on_commit_callbacks(execute=True): + counts = release_stale_scans() + + ping.assert_called_once_with({"w-dead@host"}) + dead.refresh_from_db() + dead.task.task_runner_task.refresh_from_db() + assert dead.state == StateChoices.FAILED + assert dead.task.task_runner_task.status == states.FAILURE + assert dead.task.task_runner_task.date_done is not None + assert counts["unresponsive_tasks"] == 1 + assert counts["dispatched"] == 1 + publish.assert_called_once() + queued.task.task_runner_task.refresh_from_db() + assert queued.task.task_runner_task.status == states.PENDING + + def test_other_tasks_of_an_unresponsive_worker_are_left_to_orphan_recovery( + self, tenants_fixture, aws_provider + ): + tenant = tenants_fixture[0] + _executing( + tenant, aws_provider, task_status=states.STARTED, worker="w-dead@host" + ) + summary = TaskResult.objects.create( + task_id=str(uuid.uuid4()), + task_name="scan-summary", + status=states.STARTED, + worker="w-dead@host", + ) + + with patch(PING, return_value=(set(), {"w-dead@host"})): + counts = release_stale_scans() + + summary.refresh_from_db() + assert summary.status == states.STARTED + assert counts["unresponsive_tasks"] == 1 + + def test_responsive_worker_scan_is_preserved(self, tenants_fixture, aws_provider): + live = _executing( + tenants_fixture[0], + aws_provider, + task_status=states.STARTED, + worker="w-live@host", + ) + + with patch(PING, return_value=({"w-live@host"}, set())): + release_stale_scans() + + live.refresh_from_db() + live.task.task_runner_task.refresh_from_db() + assert live.state == StateChoices.EXECUTING + assert live.task.task_runner_task.status == states.STARTED + + def test_unknown_liveness_is_preserved(self, tenants_fixture, aws_provider): + scan = _executing( + tenants_fixture[0], + aws_provider, + task_status=states.STARTED, + worker="w-unknown@host", + ) + + with patch(PING, return_value=(set(), None)): + release_stale_scans() + + scan.refresh_from_db() + assert scan.state == StateChoices.EXECUTING + + def test_scan_without_worker_is_left_to_the_backstop( + self, tenants_fixture, aws_provider + ): + tenant = tenants_fixture[0] + no_worker = _executing(tenant, aws_provider, task_status=states.STARTED) + + with patch(PING) as ping: + release_stale_scans() + + ping.assert_not_called() + no_worker.refresh_from_db() + assert no_worker.state == StateChoices.EXECUTING + + Scan.objects.filter(pk=no_worker.pk).update( + updated_at=datetime.now(UTC) - BACKSTOP_OVER + ) + with patch(PING) as ping: + release_stale_scans() + + ping.assert_not_called() + no_worker.refresh_from_db() + assert no_worker.state == StateChoices.FAILED + + def test_ping_failure_does_not_stop_the_release( + self, tenants_fixture, aws_provider + ): + dead = _executing(tenants_fixture[0], aws_provider) + alive = _executing( + tenants_fixture[0], + aws_provider, + task_status=states.STARTED, + worker="w@host", + ) + + with patch(PING, side_effect=RuntimeError("broker down")): + counts = release_stale_scans() + + dead.refresh_from_db() + alive.refresh_from_db() + assert dead.state == StateChoices.FAILED + assert alive.state == StateChoices.EXECUTING + assert counts["unresponsive_tasks"] == 0 + + def test_backstop_scan_is_reaped_and_recent_one_kept( + self, tenants_fixture, aws_provider + ): + tenant = tenants_fixture[0] + over = _executing( + tenant, aws_provider, task_status=states.STARTED, idle=BACKSTOP_OVER + ) + under = _executing( + tenant, aws_provider, task_status=states.STARTED, idle=BACKSTOP_UNDER + ) + + release_stale_scans() + + over.refresh_from_db() + under.refresh_from_db() + assert over.state == StateChoices.FAILED + assert under.state == StateChoices.EXECUTING + + def test_handles_two_tenants(self, tenants_fixture, aws_provider, provider_factory): + tenant_a, tenant_b = tenants_fixture[0], tenants_fixture[1] + provider_b = provider_factory(tenant=tenant_b) + dead_a = _executing(tenant_a, aws_provider) + dead_b = _executing( + tenant_b, provider_b, task_status=states.STARTED, worker="w-b@host" + ) + + with patch(PING, return_value=(set(), {"w-b@host"})): + counts = release_stale_scans() + + dead_a.refresh_from_db() + dead_b.refresh_from_db() + assert dead_a.state == StateChoices.FAILED + assert dead_b.state == StateChoices.FAILED + assert counts["providers_checked"] == 2 + + def test_one_provider_failure_does_not_stop_the_rest( + self, tenants_fixture, aws_provider, provider_factory + ): + tenant_a, tenant_b = tenants_fixture[0], tenants_fixture[1] + provider_b = provider_factory(tenant=tenant_b) + _executing(tenant_a, aws_provider) + dead_b = _executing(tenant_b, provider_b) + real = _release_provider_scan_slot + + def flaky(tenant_id, provider_id): + if provider_id == str(aws_provider.id): + raise RuntimeError("boom") + return real(tenant_id, provider_id) + + with patch("tasks.tasks._release_provider_scan_slot", side_effect=flaky): + counts = release_stale_scans() + + dead_b.refresh_from_db() + assert dead_b.state == StateChoices.FAILED + assert counts["failed"] == 1