mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-24 21:11:53 +00:00
fix(Task): PRWLR-4970 Fix Celery task issues when status is pending and race conditions (#49)
* fix(Task): PRWLR-4970 add TaskResult entry to database when task reaches broker * fix(Task, Scan): PRWLR-4970 remove race conditions using atomic transactions * chore(Django): PRWLR-4970 bump Django version to 5.1.1
This commit is contained in:
committed by
GitHub
parent
ded28baa2f
commit
6bd8a17a5f
@@ -28,7 +28,7 @@ start_prod_server() {
|
||||
|
||||
start_worker() {
|
||||
echo "Starting the worker..."
|
||||
poetry run python -m celery -A config.celery worker -l "${DJANGO_LOGGING_LEVEL:-info}" -E
|
||||
poetry run python -m celery -A config.celery worker -l "${DJANGO_LOGGING_LEVEL:-info}" -Q default,scans -E
|
||||
}
|
||||
|
||||
case "$1" in
|
||||
|
||||
Generated
+505
-496
File diff suppressed because it is too large
Load Diff
+1
-1
@@ -12,7 +12,7 @@ version = "1.0.0"
|
||||
|
||||
[tool.poetry.dependencies]
|
||||
celery = {extras = ["pytest"], version = "^5.4.0"}
|
||||
django = "5.0.8"
|
||||
django = "5.1.1"
|
||||
django-celery-results = "^2.5.1"
|
||||
django-cors-headers = "4.4.0"
|
||||
django-environ = "0.11.2"
|
||||
|
||||
@@ -4,3 +4,6 @@ from django.apps import AppConfig
|
||||
class ApiConfig(AppConfig):
|
||||
default_auto_field = "django.db.models.BigAutoField"
|
||||
name = "api"
|
||||
|
||||
def ready(self):
|
||||
from api import signals # noqa: F401
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
from celery import states
|
||||
from celery.signals import before_task_publish
|
||||
from django_celery_results.backends.database import DatabaseBackend
|
||||
|
||||
from config.celery import celery_app
|
||||
|
||||
|
||||
def create_task_result_on_publish(sender=None, headers=None, **kwargs): # noqa: F841
|
||||
"""Celery signal to store TaskResult entries when tasks reach the broker."""
|
||||
db_result_backend = DatabaseBackend(celery_app)
|
||||
request = type("request", (object,), headers)
|
||||
|
||||
db_result_backend.store_result(
|
||||
headers["id"],
|
||||
None,
|
||||
states.PENDING,
|
||||
traceback=None,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
before_task_publish.connect(
|
||||
create_task_result_on_publish, dispatch_uid="create_task_result_on_publish"
|
||||
)
|
||||
+18
-12
@@ -1,6 +1,7 @@
|
||||
from celery.result import AsyncResult
|
||||
from django.conf import settings as django_settings
|
||||
from django.contrib.postgres.search import SearchQuery
|
||||
from django.db import transaction
|
||||
from django.db.models import F, Q
|
||||
from django.urls import reverse
|
||||
from django.utils.decorators import method_decorator
|
||||
@@ -456,9 +457,10 @@ class ProviderViewSet(BaseRLSViewSet):
|
||||
@action(detail=True, methods=["post"], url_name="connection")
|
||||
def connection(self, request, pk=None):
|
||||
get_object_or_404(Provider, pk=pk)
|
||||
task = check_provider_connection_task.delay(
|
||||
provider_id=pk, tenant_id=request.tenant_id
|
||||
)
|
||||
with transaction.atomic():
|
||||
task = check_provider_connection_task.delay(
|
||||
provider_id=pk, tenant_id=request.tenant_id
|
||||
)
|
||||
prowler_task = Task.objects.get(id=task.id)
|
||||
serializer = TaskSerializer(prowler_task)
|
||||
return Response(
|
||||
@@ -473,7 +475,10 @@ class ProviderViewSet(BaseRLSViewSet):
|
||||
|
||||
def destroy(self, request, *args, pk=None, **kwargs):
|
||||
get_object_or_404(Provider, pk=pk)
|
||||
task = delete_provider_task.delay(provider_id=pk, tenant_id=request.tenant_id)
|
||||
with transaction.atomic():
|
||||
task = delete_provider_task.delay(
|
||||
provider_id=pk, tenant_id=request.tenant_id
|
||||
)
|
||||
prowler_task = Task.objects.get(id=task.id)
|
||||
serializer = TaskSerializer(prowler_task)
|
||||
return Response(
|
||||
@@ -560,14 +565,15 @@ class ScanViewSet(BaseRLSViewSet):
|
||||
def create(self, request, *args, **kwargs):
|
||||
input_serializer = self.get_serializer(data=request.data)
|
||||
input_serializer.is_valid(raise_exception=True)
|
||||
scan = input_serializer.save()
|
||||
|
||||
task = perform_scan_task.delay(
|
||||
tenant_id=request.tenant_id,
|
||||
scan_id=str(scan.id),
|
||||
provider_id=str(scan.provider_id),
|
||||
checks_to_execute=scan.scanner_args.get("checks_to_execute", []),
|
||||
)
|
||||
with transaction.atomic():
|
||||
scan = input_serializer.save()
|
||||
with transaction.atomic():
|
||||
task = perform_scan_task.delay(
|
||||
tenant_id=request.tenant_id,
|
||||
scan_id=str(scan.id),
|
||||
provider_id=str(scan.provider_id),
|
||||
checks_to_execute=scan.scanner_args.get("checks_to_execute", []),
|
||||
)
|
||||
|
||||
scan.task_id = task.id
|
||||
scan.save(update_fields=["task_id"])
|
||||
|
||||
@@ -1,9 +1,4 @@
|
||||
import time
|
||||
|
||||
from celery import Celery, Task
|
||||
from django.utils.translation import gettext_lazy as _
|
||||
from rest_framework import status
|
||||
from rest_framework.exceptions import APIException
|
||||
|
||||
celery_app = Celery("tasks")
|
||||
|
||||
@@ -13,36 +8,7 @@ celery_app.conf.update(result_extended=True)
|
||||
celery_app.autodiscover_tasks(["api"])
|
||||
|
||||
|
||||
class TaskTimeoutError(APIException):
|
||||
status_code = status.HTTP_504_GATEWAY_TIMEOUT
|
||||
default_detail = _("The request timed out")
|
||||
default_code = "service_unavailable_timeout"
|
||||
|
||||
|
||||
class RLSTask(Task):
|
||||
def wait_for_task_result(self, result, timeout=10, poll_interval=0.1):
|
||||
"""
|
||||
Wait for the Task runner task to be created, with a timeout.
|
||||
|
||||
Args:
|
||||
result: The result object that contains the task_id.
|
||||
timeout: Maximum time to wait for the TaskResult to be created (in seconds).
|
||||
poll_interval: Time between each check (in seconds).
|
||||
|
||||
Raises:
|
||||
TimeoutError: If the TaskResult is not created within the specified timeout.
|
||||
"""
|
||||
from django_celery_results.models import TaskResult
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
while not TaskResult.objects.filter(task_id=result.task_id).exists():
|
||||
if time.time() - start_time > timeout:
|
||||
raise TaskTimeoutError(
|
||||
f"Task runner task was not created within {timeout} seconds"
|
||||
)
|
||||
time.sleep(poll_interval)
|
||||
|
||||
def apply_async(
|
||||
self,
|
||||
args=None,
|
||||
@@ -67,8 +33,6 @@ class RLSTask(Task):
|
||||
shadow=shadow,
|
||||
**options,
|
||||
)
|
||||
# The TaskResult row is delayed a bit, so we need to wait for it to be created
|
||||
self.wait_for_task_result(result, timeout=10, poll_interval=0.05)
|
||||
task_result_instance = TaskResult.objects.get(task_id=result.task_id)
|
||||
APITask.objects.create(
|
||||
id=task_result_instance.task_id,
|
||||
|
||||
@@ -21,7 +21,6 @@ from api.models import (
|
||||
StateChoices,
|
||||
)
|
||||
from api.v1.serializers import ScanTaskSerializer
|
||||
from config.celery import TaskTimeoutError
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
@@ -79,13 +78,7 @@ def perform_prowler_scan(
|
||||
provider_instance = Provider.objects.get(pk=provider_id)
|
||||
start_time = time.time()
|
||||
unique_resources = set()
|
||||
# Prevent race conditions
|
||||
while not Scan.objects.filter(id=scan_id).exists():
|
||||
if time.time() - start_time > 10:
|
||||
raise TaskTimeoutError(
|
||||
f"Could not find scan with given id {scan_id} within 10 seconds"
|
||||
)
|
||||
time.sleep(0.1)
|
||||
|
||||
scan_instance = Scan.objects.get(pk=scan_id)
|
||||
scan_instance.state = StateChoices.EXECUTING
|
||||
scan_instance.started_at = datetime.now(tz=timezone.utc)
|
||||
|
||||
@@ -43,7 +43,7 @@ def delete_provider_task(provider_id: str):
|
||||
return delete_instance(model=Provider, pk=provider_id)
|
||||
|
||||
|
||||
@shared_task(base=RLSTask, name="scan-perform")
|
||||
@shared_task(base=RLSTask, name="scan-perform", queue="scans")
|
||||
def perform_scan_task(
|
||||
tenant_id: str, scan_id: str, provider_id: str, checks_to_execute: list[str] = None
|
||||
):
|
||||
|
||||
Reference in New Issue
Block a user