diff --git a/.env b/.env index 1fa497f3e6..1f53b81153 100644 --- a/.env +++ b/.env @@ -24,6 +24,10 @@ POSTGRES_USER=prowler POSTGRES_PASSWORD=postgres POSTGRES_DB=prowler_db +# Celery-Prowler task settings +TASK_RETRY_DELAY_SECONDS=0.1 +TASK_RETRY_ATTEMPTS=5 + # Valkey settings # If running Valkey and celery on host, use localhost, else use 'valkey' VALKEY_HOST=valkey diff --git a/api/CHANGELOG.md b/api/CHANGELOG.md index e7ad901fc4..937c090f85 100644 --- a/api/CHANGELOG.md +++ b/api/CHANGELOG.md @@ -20,6 +20,7 @@ All notable changes to the **Prowler API** are documented in this file. ### Fixed - Fixed task lookup to use task_kwargs instead of task_args for scan report resolution. [(#7830)](https://github.com/prowler-cloud/prowler/pull/7830) - Fixed Kubernetes UID validation to allow valid context names [(#7871)](https://github.com/prowler-cloud/prowler/pull/7871) +- Fixed a race condition when creating background tasks [(#7876)](https://github.com/prowler-cloud/prowler/pull/7876). --- diff --git a/api/src/backend/api/models.py b/api/src/backend/api/models.py index 2b4fec00cf..9564399868 100644 --- a/api/src/backend/api/models.py +++ b/api/src/backend/api/models.py @@ -1,7 +1,9 @@ import json import re +import time from uuid import UUID, uuid4 +from config.env import env from cryptography.fernet import Fernet from django.conf import settings from django.contrib.auth.models import AbstractBaseUser @@ -352,6 +354,42 @@ class ProviderGroupMembership(RowLevelSecurityProtectedModel): resource_name = "provider_groups-provider" +class TaskManager(models.Manager): + def get_with_retry( + self, + id: str, + max_retries: int = None, + delay_seconds: float = None, + ): + """ + Retry fetching a Task by ID in case it hasn't been created yet. + + Args: + id (str): The Celery task ID (expected to match Task model PK). + max_retries (int, optional): Number of retry attempts. Defaults to env TASK_RETRY_ATTEMPTS or 5. + delay_seconds (float, optional): Delay between retries in seconds. Defaults to env TASK_RETRY_DELAY_SECONDS or 0.1. + + Returns: + Task: The retrieved Task instance. + + Raises: + Task.DoesNotExist: If the task is not found after all retries. + """ + max_retries = max_retries or env.int("TASK_RETRY_ATTEMPTS", default=5) + delay_seconds = delay_seconds or env.float( + "TASK_RETRY_DELAY_SECONDS", default=0.1 + ) + + for _attempt in range(max_retries): + try: + return self.get(id=id) + except self.model.DoesNotExist: + time.sleep(delay_seconds) + raise self.model.DoesNotExist( + f"Task with ID {id} not found after {max_retries} retries." + ) + + class Task(RowLevelSecurityProtectedModel): id = models.UUIDField(primary_key=True, default=uuid4, editable=False) inserted_at = models.DateTimeField(auto_now_add=True, editable=False) @@ -364,6 +402,8 @@ class Task(RowLevelSecurityProtectedModel): blank=True, ) + objects = TaskManager() + class Meta(RowLevelSecurityProtectedModel.Meta): db_table = "tasks" diff --git a/api/src/backend/api/tests/test_models.py b/api/src/backend/api/tests/test_models.py index c2beeb3583..de2ec7c59e 100644 --- a/api/src/backend/api/tests/test_models.py +++ b/api/src/backend/api/tests/test_models.py @@ -1,6 +1,9 @@ +import uuid +from unittest import mock + import pytest -from api.models import Resource, ResourceTag +from api.models import Resource, ResourceTag, Task @pytest.mark.django_db @@ -120,3 +123,35 @@ class TestResourceModel: # compliance={}, # ) # assert Finding.objects.filter(uid=long_uid).exists() + + +@pytest.mark.django_db +class TestTaskManager: + def test_get_with_retry_success(self): + task_id = uuid.uuid4() + call_counter = {"count": 0} + + def side_effect(*args, **kwargs): + if call_counter["count"] < 2: + call_counter["count"] += 1 + raise Task.DoesNotExist() + return Task(id=task_id) + + with mock.patch.object(Task.objects, "get", side_effect=side_effect): + task = Task.objects.get_with_retry( + task_id, max_retries=5, delay_seconds=0.01 + ) + + assert task.id == task_id + assert call_counter["count"] == 2 + + def test_get_with_retry_fail(self): + non_existent_id = uuid.uuid4() + + with mock.patch.object(Task.objects, "get", side_effect=Task.DoesNotExist): + with pytest.raises(Task.DoesNotExist) as excinfo: + Task.objects.get_with_retry( + non_existent_id, max_retries=3, delay_seconds=0.01 + ) + + assert str(non_existent_id) in str(excinfo.value) diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index b8679217f9..8fc87d2808 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -1086,7 +1086,7 @@ class ProviderViewSet(BaseRLSViewSet): task = check_provider_connection_task.delay( provider_id=pk, tenant_id=self.request.tenant_id ) - prowler_task = Task.objects.get(id=task.id) + prowler_task = Task.objects.get_with_retry(id=task.id) serializer = TaskSerializer(prowler_task) return Response( data=serializer.data, @@ -1109,7 +1109,7 @@ class ProviderViewSet(BaseRLSViewSet): task = delete_provider_task.delay( provider_id=pk, tenant_id=self.request.tenant_id ) - prowler_task = Task.objects.get(id=task.id) + prowler_task = Task.objects.get_with_retry(id=task.id) serializer = TaskSerializer(prowler_task) return Response( data=serializer.data, @@ -1489,10 +1489,10 @@ class ScanViewSet(BaseRLSViewSet): }, ) + prowler_task = Task.objects.get_with_retry(id=task.id) scan.task_id = task.id scan.save(update_fields=["task_id"]) - prowler_task = Task.objects.get(id=task.id) self.response_serializer_class = TaskSerializer output_serializer = self.get_serializer(prowler_task) @@ -2823,7 +2823,7 @@ class ScheduleViewSet(BaseRLSViewSet): with transaction.atomic(): task = schedule_provider_scan(provider_instance) - prowler_task = Task.objects.get(id=task.id) + prowler_task = Task.objects.get_with_retry(id=task.id) self.response_serializer_class = TaskSerializer output_serializer = self.get_serializer(prowler_task)