revert(api): release providers blocked by scans whose worker died (#12915)

This commit is contained in:
César Arroba
2026-09-30 13:00:50 +02:00
committed by GitHub
parent a44a725507
commit f0da33f451
9 changed files with 18 additions and 918 deletions
@@ -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
@@ -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),
]
@@ -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()
-2
View File
@@ -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
)
-6
View File
@@ -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
-69
View File
@@ -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
@@ -86,7 +86,6 @@ _SKIP_RECOVERY = {
"attack-paths-cleanup-stale-scans",
"attack-paths-reap-orphaned-tmp-databases",
"reconcile-orphan-tasks",
"scan-release-stale",
}
+18 -124
View File
@@ -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)."""
@@ -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