diff --git a/api/changelog.d/scan-release-stale.fixed.md b/api/changelog.d/scan-release-stale.fixed.md deleted file mode 100644 index c7a986d4ee..0000000000 --- a/api/changelog.d/scan-release-stale.fixed.md +++ /dev/null @@ -1 +0,0 @@ -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 deleted file mode 100644 index 24f67e1c2b..0000000000 --- a/api/src/backend/api/migrations/0101_scan_release_stale_periodic_task.py +++ /dev/null @@ -1,48 +0,0 @@ -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 deleted file mode 100644 index a55fb41e33..0000000000 --- a/api/src/backend/api/tests/test_scan_dead_release_view.py +++ /dev/null @@ -1,148 +0,0 @@ -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 dab9be431c..5dc175dd42 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -333,7 +333,6 @@ 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, @@ -2809,7 +2808,6 @@ 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 f5ec7356a8..664fcd1c4a 100644 --- a/api/src/backend/config/django/base.py +++ b/api/src/backend/config/django/base.py @@ -321,12 +321,6 @@ 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 deleted file mode 100644 index 9b6f785975..0000000000 --- a/api/src/backend/tasks/jobs/dead_scans.py +++ /dev/null @@ -1,69 +0,0 @@ -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 f184317bcf..7290fb313f 100644 --- a/api/src/backend/tasks/jobs/orphan_recovery.py +++ b/api/src/backend/tasks/jobs/orphan_recovery.py @@ -86,7 +86,6 @@ _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 e5393a347a..9ec6526963 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, MainRouter +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.exceptions import ProviderDeletedException @@ -29,7 +29,6 @@ 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 @@ -59,11 +58,6 @@ 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, @@ -109,19 +103,18 @@ 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 live scan that has already been dispatched for a provider.""" - dead = dead_scan_q(datetime.now(UTC)) + """Return a scan that has already been dispatched for a provider.""" executing_scan = ( - Scan.objects.select_for_update(of=("self",)) + Scan.objects.select_for_update() .filter( tenant_id=tenant_id, provider_id=provider_id, state=StateChoices.EXECUTING, ) - .exclude(dead) .order_by("-inserted_at") .first() ) @@ -138,7 +131,6 @@ 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() ) @@ -266,64 +258,28 @@ 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 - return _dispatch_queued_provider_scan_locked(tenant_id, provider_id) + 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 -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, + 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), ) - 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) + return queued_scan def _dispatch_next_queued_provider_scan_best_effort( @@ -337,51 +293,6 @@ 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, @@ -710,17 +621,6 @@ 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( @@ -835,12 +735,6 @@ 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 deleted file mode 100644 index ec1be61372..0000000000 --- a/api/src/backend/tasks/tests/test_scan_release_stale.py +++ /dev/null @@ -1,519 +0,0 @@ -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