mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-10-04 02:04:06 +00:00
fix(api): return the new scan id when a scan is created (#12878)
This commit is contained in:
@@ -0,0 +1 @@
|
||||
`POST /api/v1/scans` again returns the new scan id in the response `task_args`, which had been empty since the scan broker publish moved to transaction commit
|
||||
@@ -3951,6 +3951,43 @@ class TestScanViewSet:
|
||||
mock_enqueue_scan_execution.assert_called_once()
|
||||
# assert scan.scanner_args == expected_scanner_args
|
||||
|
||||
@patch("api.v1.views.enqueue_scan_execution_on_commit")
|
||||
def test_scans_create_returns_the_scan_id_in_task_args(
|
||||
self,
|
||||
mock_enqueue_scan_execution,
|
||||
authenticated_client,
|
||||
okta_provider,
|
||||
):
|
||||
"""The 202 is a task, so `task_args` is the only place the scan id is.
|
||||
|
||||
It is serialized before the on_commit publish that would otherwise fill
|
||||
the kwargs, so the record has to carry them from the start.
|
||||
"""
|
||||
payload = {
|
||||
"data": {
|
||||
"type": "scans",
|
||||
"attributes": {"name": "New Scan"},
|
||||
"relationships": {
|
||||
"provider": {
|
||||
"data": {"type": "providers", "id": str(okta_provider.id)}
|
||||
}
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
response = authenticated_client.post(
|
||||
reverse("scan-list"),
|
||||
data=payload,
|
||||
content_type=API_JSON_CONTENT_TYPE,
|
||||
)
|
||||
|
||||
assert response.status_code == status.HTTP_202_ACCEPTED
|
||||
scan = Scan.objects.get()
|
||||
assert response.json()["data"]["attributes"]["task_args"] == {
|
||||
"scan_id": str(scan.id),
|
||||
"provider_id": str(okta_provider.id),
|
||||
}
|
||||
|
||||
@patch("tasks.tasks.perform_scan_task.apply_async")
|
||||
def test_scans_create_queues_scan_when_provider_has_active_scan(
|
||||
self,
|
||||
|
||||
@@ -2822,6 +2822,15 @@ class ScanViewSet(ProviderVisibilityMixin, BaseRLSViewSet):
|
||||
tenant_id=self.request.tenant_id,
|
||||
task_id=pre_task_id,
|
||||
task_status=(QUEUED_SCAN_TASK_STATE if active_scan else None),
|
||||
# This response is serialized before the on_commit publish,
|
||||
# so without these the caller gets a task id and no scan id.
|
||||
# Kept in step with what `enqueue_scan_execution_on_commit`
|
||||
# publishes below.
|
||||
task_kwargs={
|
||||
"tenant_id": str(self.request.tenant_id),
|
||||
"scan_id": str(scan.id),
|
||||
"provider_id": str(scan.provider_id),
|
||||
},
|
||||
)
|
||||
|
||||
if not active_scan:
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import os
|
||||
from datetime import UTC, datetime, timedelta
|
||||
from pathlib import Path
|
||||
@@ -163,13 +164,28 @@ def create_scan_task_record(
|
||||
task_id: str,
|
||||
task_name: str = "scan-perform",
|
||||
task_status: str | None = states.PENDING,
|
||||
task_kwargs: dict | None = None,
|
||||
) -> Task:
|
||||
"""Pre-create the TaskResult + Task rows for a pre-generated task id.
|
||||
|
||||
Pass ``task_kwargs`` when the response built from this record is serialized
|
||||
before the broker publish. ``task_kwargs`` is otherwise only written by the
|
||||
``before_task_publish`` signal (``api/signals.py``), and the scan publish is
|
||||
deferred to ``on_commit``, so the 202 would carry an empty ``task_args`` and
|
||||
the caller would have no way to learn the scan id it was just handed a task
|
||||
for. The publish later overwrites the field with the same kwargs as a Python
|
||||
repr; both forms decode to the same dict (``decode_celery_field``).
|
||||
"""
|
||||
if task_status is None:
|
||||
task_status = states.PENDING
|
||||
|
||||
defaults = {"status": task_status, "task_name": task_name}
|
||||
if task_kwargs is not None:
|
||||
defaults["task_kwargs"] = json.dumps(task_kwargs)
|
||||
|
||||
task_result, _ = TaskResult.objects.update_or_create(
|
||||
task_id=str(task_id),
|
||||
defaults={"status": task_status, "task_name": task_name},
|
||||
defaults=defaults,
|
||||
)
|
||||
prowler_task, _ = Task.objects.update_or_create(
|
||||
id=str(task_id),
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import json
|
||||
import uuid
|
||||
from contextlib import contextmanager
|
||||
from datetime import UTC, datetime
|
||||
@@ -14,6 +15,7 @@ from api.models import (
|
||||
StateChoices,
|
||||
Task,
|
||||
)
|
||||
from api.v1.serializers import TaskSerializer
|
||||
from botocore.exceptions import ClientError
|
||||
from celery import states
|
||||
from django_celery_beat.models import IntervalSchedule, PeriodicTask
|
||||
@@ -32,6 +34,7 @@ from tasks.tasks import (
|
||||
_scan_tmp_output_directory,
|
||||
check_integrations_task,
|
||||
check_lighthouse_provider_connection_task,
|
||||
create_scan_task_record,
|
||||
generate_outputs_task,
|
||||
mute_findings_in_latest_scans_task,
|
||||
perform_attack_paths_scan_task,
|
||||
@@ -3384,3 +3387,79 @@ class TestTaskTimeLimits:
|
||||
"lighthouse-provider-connection-check",
|
||||
):
|
||||
assert celery_app.tasks[name].time_limit < default
|
||||
|
||||
|
||||
@pytest.mark.django_db
|
||||
class TestCreateScanTaskRecord:
|
||||
"""`task_kwargs` is what a response built before the publish can report."""
|
||||
|
||||
def _scan(self, tenant, provider):
|
||||
"""A manual scan, like the one `POST /api/v1/scans` creates."""
|
||||
return Scan.objects.create(
|
||||
tenant_id=tenant.id,
|
||||
provider=provider,
|
||||
name="Manual scan",
|
||||
trigger=Scan.TriggerChoices.MANUAL,
|
||||
state=StateChoices.AVAILABLE,
|
||||
)
|
||||
|
||||
def _publish_kwargs(self, tenant, scan):
|
||||
"""What `enqueue_scan_execution_on_commit` publishes for this scan."""
|
||||
return {
|
||||
"tenant_id": str(tenant.id),
|
||||
"scan_id": str(scan.id),
|
||||
"provider_id": str(scan.provider_id),
|
||||
}
|
||||
|
||||
def _task_args(self, task):
|
||||
"""Read the record back the way `TaskSerializer` does."""
|
||||
return TaskSerializer(task).data["task_args"]
|
||||
|
||||
def test_the_stored_kwargs_are_the_ones_the_publish_would_send(
|
||||
self, tenants_fixture, aws_provider
|
||||
):
|
||||
"""The 202 reports what is stored here, so it has to be the dispatch kwargs."""
|
||||
tenant = tenants_fixture[0]
|
||||
scan = self._scan(tenant, aws_provider)
|
||||
|
||||
task = create_scan_task_record(
|
||||
tenant_id=str(tenant.id),
|
||||
task_id=str(uuid.uuid4()),
|
||||
task_kwargs=self._publish_kwargs(tenant, scan),
|
||||
)
|
||||
|
||||
assert self._task_args(task) == {
|
||||
"scan_id": str(scan.id),
|
||||
"provider_id": str(aws_provider.id),
|
||||
}
|
||||
|
||||
def test_a_record_created_without_kwargs_reports_none(self, tenants_fixture):
|
||||
"""The argument is optional, so the other callers keep their behaviour."""
|
||||
task = create_scan_task_record(
|
||||
tenant_id=str(tenants_fixture[0].id),
|
||||
task_id=str(uuid.uuid4()),
|
||||
)
|
||||
|
||||
assert self._task_args(task) == {}
|
||||
|
||||
def test_the_publish_can_overwrite_the_stored_kwargs(
|
||||
self, tenants_fixture, aws_provider
|
||||
):
|
||||
"""django-celery-results stores a Python repr; both must decode alike."""
|
||||
tenant = tenants_fixture[0]
|
||||
scan = self._scan(tenant, aws_provider)
|
||||
task_id = str(uuid.uuid4())
|
||||
kwargs = self._publish_kwargs(tenant, scan)
|
||||
|
||||
task = create_scan_task_record(
|
||||
tenant_id=str(tenant.id), task_id=task_id, task_kwargs=kwargs
|
||||
)
|
||||
before = self._task_args(task)
|
||||
|
||||
# What `before_task_publish` writes once the task reaches the broker.
|
||||
task_result = TaskResult.objects.get(task_id=task_id)
|
||||
task_result.task_kwargs = json.dumps(repr(kwargs))
|
||||
task_result.save(update_fields=["task_kwargs"])
|
||||
task.refresh_from_db()
|
||||
|
||||
assert self._task_args(task) == before
|
||||
|
||||
Reference in New Issue
Block a user