mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-10-03 17:54:05 +00:00
fix(api): release providers blocked by scans whose worker died (#12899)
This commit is contained in:
@@ -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
|
||||
@@ -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),
|
||||
]
|
||||
@@ -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()
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -86,6 +86,7 @@ _SKIP_RECOVERY = {
|
||||
"attack-paths-cleanup-stale-scans",
|
||||
"attack-paths-reap-orphaned-tmp-databases",
|
||||
"reconcile-orphan-tasks",
|
||||
"scan-release-stale",
|
||||
}
|
||||
|
||||
|
||||
|
||||
+124
-18
@@ -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)."""
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user