From bf179212a5dc4b771b93f45edc34de4071df5606 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Rub=C3=A9n=20De=20la=20Torre=20Vico?= <98956809+puchy22@users.noreply.github.com> Date: Thu, 24 Sep 2026 13:12:44 +0200 Subject: [PATCH] fix(api): return the new scan id when a scan is created (#12878) --- .../scan-create-task-args.fixed.md | 1 + api/src/backend/api/tests/test_views.py | 37 +++++++++ api/src/backend/api/v1/views.py | 9 +++ api/src/backend/tasks/tasks.py | 18 ++++- api/src/backend/tasks/tests/test_tasks.py | 79 +++++++++++++++++++ 5 files changed, 143 insertions(+), 1 deletion(-) create mode 100644 api/changelog.d/scan-create-task-args.fixed.md diff --git a/api/changelog.d/scan-create-task-args.fixed.md b/api/changelog.d/scan-create-task-args.fixed.md new file mode 100644 index 0000000000..26d42cf0c4 --- /dev/null +++ b/api/changelog.d/scan-create-task-args.fixed.md @@ -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 diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index 43a13e4d4c..64a2250082 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -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, diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index 659a9218d1..9642b36d9d 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -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: diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index d160a2b6c7..332f893aa2 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -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), diff --git a/api/src/backend/tasks/tests/test_tasks.py b/api/src/backend/tasks/tests/test_tasks.py index b4854fe196..5ab1cdcf54 100644 --- a/api/src/backend/tasks/tests/test_tasks.py +++ b/api/src/backend/tasks/tests/test_tasks.py @@ -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