fix(api): return the new scan id when a scan is created (#12878)

This commit is contained in:
Rubén De la Torre Vico
2026-09-24 13:12:44 +02:00
committed by GitHub
parent 60f936a10b
commit bf179212a5
5 changed files with 143 additions and 1 deletions
@@ -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
+37
View File
@@ -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,
+9
View File
@@ -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:
+17 -1
View File
@@ -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),
+79
View File
@@ -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