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:
Víctor Fernández Poyatos
2024-10-04 11:54:15 +02:00
committed by GitHub
parent ded28baa2f
commit 6bd8a17a5f
9 changed files with 554 additions and 555 deletions
+1 -1
View File
@@ -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
View File
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -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"
+3
View File
@@ -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
+24
View File
@@ -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
View File
@@ -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"])
-36
View File
@@ -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,
+1 -8
View File
@@ -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)
+1 -1
View File
@@ -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
):