mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-24 04:51:51 +00:00
fix(api): queue provider scans when one is active (#11848)
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
[
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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")
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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")
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user