fix(api): queue provider scans when one is active (#11848)

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