feat/PRWLR-4177 Add /tasks endpoints and data model (#35)

* feat(Backend): PRWLR-4177 add Task model and migrations

* feat(Tasks): PRWLR-4177 add RLSTask class

* feat(API): PRWLR-4177 add Task serializers

* feat(Backend, DB): PRWLR-4177 refactor db variables and add policy on task runner tasks

* feat(API): PRWLR-4177 add Tasks filters and sort fields

* feat(API, Tasks): PRWLR-4177 add deletion tasks and revoke logic to /tasks

* test(Task): PRWLR-4177 add deletion tasks unit tests

* test(Views): PRWLR-4177 add Tasks views unit tests and update outdated ones

* chore(API): PRWLR-4177 improve drf-spectacular annotations

* chore(API): PRWLR-4177 add PROGRESS task state

* chore(API): PRWLR-4177 update spec

* chore(API): PRWLR-4177 remove force query parameter from DELETE /tasks

* feat(Backend): PRWLR-4177 add APITimeoutError and raise when TaskResult is not created

* feat(Backend): PRWLR-4177 add specific error class for task timeouts
This commit is contained in:
Víctor Fernández Poyatos
2024-09-07 02:47:51 +02:00
committed by GitHub
parent f5462c9b27
commit ec67fc12e0
15 changed files with 1426 additions and 161 deletions
+13
View File
@@ -5,6 +5,19 @@ from django.db import models
from psycopg2 import connect as psycopg2_connect
from psycopg2.extensions import new_type, register_type, register_adapter, AsIs
DB_USER = settings.DATABASES["default"]["USER"] if not settings.TESTING else "test"
DB_PASSWORD = (
settings.DATABASES["default"]["PASSWORD"] if not settings.TESTING else "test"
)
DB_PROWLER_USER = (
settings.DATABASES["prowler_user"]["USER"] if not settings.TESTING else "test"
)
DB_PROWLER_PASSWORD = (
settings.DATABASES["prowler_user"]["PASSWORD"] if not settings.TESTING else "test"
)
TASK_RUNNER_DB_TABLE = "django_celery_results_taskresult"
POSTGRES_TENANT_VAR = "api.tenant_id"
@contextmanager
def psycopg_connection(database_alias: str):
+53 -12
View File
@@ -3,37 +3,41 @@ from rest_framework_json_api.django_filters.backends import DjangoFilterBackend
from rest_framework_json_api.serializers import ValidationError
from api.db_utils import ProviderEnumField
from api.models import Provider, Scan
from api.models import Provider, Scan, Task, StateChoices
from api.rls import Tenant
from api.v1.serializers import TaskBase
def provider_enum_filter(queryset, value, lookup_field: str = "provider"):
def enum_filter(queryset, value, enum_choices, lookup_field: str):
"""
Filter a queryset based on a provider value, using a specified lookup field.
Filter a queryset based on a provided value, using a specified lookup field
and validating against a given enumeration.
This function filters a given queryset by checking if the provided `value`
matches a valid choice in the `Provider.ProviderChoices` enum. If the `value`
is valid, the queryset is filtered using the specified `lookup_field`.
matches a valid choice in the specified `enum_choices` (a Django `TextChoices` enum).
If the `value` is valid, the queryset is filtered using the specified `lookup_field`.
Otherwise, a `ValidationError` is raised.
Args:
queryset (QuerySet): The Django queryset to be filtered.
value (str): The value to filter the queryset by, which must be a valid
option in `Provider.ProviderChoices`.
option in the specified `enum_choices`.
enum_choices: The enumeration class that defines the valid
choices for the `value`.
lookup_field (str): The field or lookup path within the model used for
filtering the queryset. Defaults to "provider".
filtering the queryset.
Returns:
QuerySet: A filtered queryset based on the provided `value` and `lookup_field`.
Raises:
ValidationError: If the provided `value` is not a valid choice in
`Provider.ProviderChoices`.
the specified `enum_choices`.
"""
if value not in Provider.ProviderChoices:
if value not in enum_choices:
raise ValidationError(
f"Invalid provider value: '{value}'. Valid values are: "
f"{', '.join(Provider.ProviderChoices)}"
f"{', '.join(enum_choices)}"
)
return queryset.filter(**{lookup_field: value})
@@ -63,7 +67,12 @@ class ProviderFilter(FilterSet):
provider = CharFilter(method="filter_provider")
def filter_provider(self, queryset, name, value):
return provider_enum_filter(queryset, value)
return enum_filter(
queryset,
value,
enum_choices=Provider.ProviderChoices,
lookup_field="provider",
)
class Meta:
model = Provider
@@ -86,7 +95,12 @@ class ScanFilter(FilterSet):
trigger = CharFilter(method="filter_trigger")
def filter_provider(self, queryset, name, value):
return provider_enum_filter(queryset, value, lookup_field="provider__provider")
return enum_filter(
queryset,
value,
enum_choices=Provider.ProviderChoices,
lookup_field="provider__provider",
)
def filter_trigger(self, queryset, name, value):
if value not in Scan.TriggerChoices:
@@ -106,3 +120,30 @@ class ScanFilter(FilterSet):
"started_at": ["exact", "gte", "lte"],
"trigger": ["exact"],
}
class TaskFilter(FilterSet):
name = CharFilter(field_name="task_runner_task__task_name", lookup_expr="exact")
name__icontains = CharFilter(
field_name="task_runner_task__task_name", lookup_expr="icontains"
)
state = CharFilter(method="filter_state", lookup_expr="exact")
task_state_inverse_mapping_values = {
v: k for k, v in TaskBase.state_mapping.items()
}
def filter_state(self, queryset, name, value):
if value not in StateChoices:
raise ValidationError(
f"Invalid provider value: '{value}'. Valid values are: "
f"{', '.join(StateChoices)}"
)
return queryset.filter(
task_runner_task__status=self.task_state_inverse_mapping_values[value]
)
class Meta:
model = Task
fields = []
+75 -12
View File
@@ -18,16 +18,15 @@ from api.db_utils import (
StateEnum,
ScanTriggerEnumField,
register_enum,
DB_PROWLER_USER,
DB_PROWLER_PASSWORD,
TASK_RUNNER_DB_TABLE,
POSTGRES_TENANT_VAR,
)
from api.models import Provider, Scan, StateChoices
DB_NAME = settings.DATABASES["default"]["NAME"]
DB_USER_NAME = (
settings.DATABASES["prowler_user"]["USER"] if not settings.TESTING else "test"
)
DB_USER_PASSWORD = (
settings.DATABASES["prowler_user"]["PASSWORD"] if not settings.TESTING else "test"
)
ProviderEnumMigration = PostgresEnumMigration(
enum_name="provider",
@@ -50,7 +49,7 @@ class Migration(migrations.Migration):
# Required for our kind of `RunPython` operations
atomic = False
dependencies = []
dependencies = [("django_celery_results", "0011_taskresult_periodic_task_name")]
operations = [
migrations.RunSQL(
@@ -60,8 +59,8 @@ class Migration(migrations.Migration):
IF NOT EXISTS (
SELECT
FROM pg_catalog.pg_roles
WHERE rolname = '{DB_USER_NAME}') THEN
CREATE ROLE {DB_USER_NAME} LOGIN PASSWORD '{DB_USER_PASSWORD}';
WHERE rolname = '{DB_PROWLER_USER}') THEN
CREATE ROLE {DB_PROWLER_USER} LOGIN PASSWORD '{DB_PROWLER_PASSWORD}';
END IF;
END
$$;
@@ -70,8 +69,8 @@ class Migration(migrations.Migration):
migrations.RunSQL(
# `runserver` command for dev tools requires read access to migrations
f"""
GRANT CONNECT ON DATABASE "{DB_NAME}" TO {DB_USER_NAME};
GRANT SELECT ON django_migrations TO {DB_USER_NAME};
GRANT CONNECT ON DATABASE "{DB_NAME}" TO {DB_PROWLER_USER};
GRANT SELECT ON django_migrations TO {DB_PROWLER_USER};
"""
),
# Create and register State type
@@ -103,7 +102,7 @@ class Migration(migrations.Migration):
migrations.RunSQL(
# Needed for now since we don't have users yet
f"""
GRANT SELECT, INSERT, UPDATE, DELETE ON TABLE tenants TO {DB_USER_NAME};
GRANT SELECT, INSERT, UPDATE, DELETE ON TABLE tenants TO {DB_PROWLER_USER};
"""
),
# Create and register ProviderEnum type
@@ -283,4 +282,68 @@ class Migration(migrations.Migration):
name="scans_prov_state_trig_sche_idx",
),
),
migrations.CreateModel(
name="Task",
fields=[
(
"id",
models.UUIDField(
default=uuid.uuid4,
editable=False,
primary_key=True,
serialize=False,
),
),
("inserted_at", models.DateTimeField(auto_now_add=True)),
(
"task_runner_task",
models.OneToOneField(
blank=True,
null=True,
on_delete=django.db.models.deletion.CASCADE,
related_name="task",
related_query_name="task",
to="django_celery_results.taskresult",
),
),
(
"tenant",
models.ForeignKey(
on_delete=django.db.models.deletion.CASCADE, to="api.tenant"
),
),
],
options={
"db_table": "tasks",
"abstract": False,
},
),
migrations.AddConstraint(
model_name="task",
constraint=api.rls.RowLevelSecurityConstraint(
"tenant_id",
name="rls_on_task",
statements=["SELECT", "INSERT", "UPDATE", "DELETE"],
),
),
migrations.AddIndex(
model_name="task",
index=models.Index(
fields=["id", "task_runner_task"],
name="tasks_id_trt_id_idx",
),
),
migrations.RunSQL(
f"""
ALTER TABLE {TASK_RUNNER_DB_TABLE} ENABLE ROW LEVEL SECURITY;
CREATE POLICY "{DB_PROWLER_USER}_{TASK_RUNNER_DB_TABLE}_select"
ON {TASK_RUNNER_DB_TABLE}
FOR SELECT
TO {DB_PROWLER_USER}
USING (
task_id::uuid in (SELECT id FROM tasks WHERE tenant_id = (NULLIF(current_setting('{POSTGRES_TENANT_VAR}', true), ''))::uuid)
);
GRANT SELECT ON TABLE {TASK_RUNNER_DB_TABLE} TO {DB_PROWLER_USER};
"""
),
]
+32
View File
@@ -4,6 +4,7 @@ from uuid import uuid4, UUID
from django.core.validators import MinLengthValidator
from django.db import models
from django.utils.translation import gettext_lazy as _
from django_celery_results.models import TaskResult
from uuid6 import uuid7
from api.db_utils import ProviderEnumField, StateEnumField, ScanTriggerEnumField
@@ -157,3 +158,34 @@ class Scan(RowLevelSecurityProtectedModel):
name="scans_prov_state_type_sche_idx",
),
]
class Task(RowLevelSecurityProtectedModel):
id = models.UUIDField(primary_key=True, default=uuid4, editable=False)
inserted_at = models.DateTimeField(auto_now_add=True, editable=False)
task_runner_task = models.OneToOneField(
TaskResult,
on_delete=models.CASCADE,
related_name="task",
related_query_name="task",
null=True,
blank=True,
)
class Meta(RowLevelSecurityProtectedModel.Meta):
db_table = "tasks"
constraints = [
RowLevelSecurityConstraint(
field="tenant_id",
name="rls_on_%(class)s",
statements=["SELECT", "INSERT", "UPDATE", "DELETE"],
),
]
indexes = [
models.Index(
fields=["id", "task_runner_task"],
name="tasks_id_trt_id_idx",
),
]
+4 -9
View File
@@ -1,15 +1,12 @@
from typing import Any
from uuid import uuid4
from django.conf import settings
from django.core.exceptions import ValidationError
from django.db import DEFAULT_DB_ALIAS
from django.db import models
from django.db.backends.ddl_references import Statement, Table
DB_PROWLER_USER = (
settings.DATABASES["default"]["USER"] if not settings.TESTING else "test"
)
from api.db_utils import DB_USER, POSTGRES_TENANT_VAR
class Tenant(models.Model):
@@ -29,8 +26,6 @@ class Tenant(models.Model):
# TODO Add abstract class for non-RLS models
class RowLevelSecurityConstraint(models.BaseConstraint):
TENANT_SETTING = "api.tenant_id"
rls_sql_query = """
ALTER TABLE %(table_name)s ENABLE ROW LEVEL SECURITY;
ALTER TABLE %(table_name)s FORCE ROW LEVEL SECURITY;
@@ -88,8 +83,8 @@ class RowLevelSecurityConstraint(models.BaseConstraint):
full_create_sql_query,
table_name=model._meta.db_table,
field_column=field_column,
db_user=DB_PROWLER_USER,
tenant_setting=self.TENANT_SETTING,
db_user=DB_USER,
tenant_setting=POSTGRES_TENANT_VAR,
)
def remove_sql(self, model: Any, schema_editor: Any) -> Any:
@@ -102,7 +97,7 @@ class RowLevelSecurityConstraint(models.BaseConstraint):
full_drop_sql_query,
table_name=Table(model._meta.db_table, schema_editor.quote_name),
field_column=field_column,
db_user=DB_PROWLER_USER,
db_user=DB_USER,
)
def __eq__(self, other: object) -> bool:
File diff suppressed because it is too large Load Diff
+66 -5
View File
@@ -377,16 +377,25 @@ class TestProviderViewSet:
)
assert response.status_code == status.HTTP_400_BAD_REQUEST
def test_providers_delete(self, client, providers_fixture, tenant_header):
@patch("api.v1.views.delete_provider_task.delay")
def test_providers_delete(
self, mock_delete_task, client, providers_fixture, tenant_header
):
task_mock = Mock()
task_mock.id = "12345"
mock_delete_task.return_value = task_mock
provider1, *_ = providers_fixture
response = client.delete(
reverse("provider-detail", kwargs={"pk": provider1.id}),
headers=tenant_header,
)
assert response.status_code == status.HTTP_202_ACCEPTED
mock_delete_task.assert_called_once_with(
provider_id=str(provider1.id), tenant_id=tenant_header["X-Tenant-ID"]
)
assert "Content-Location" in response.headers
assert Provider.objects.count() == len(providers_fixture) - 1
# TODO Assert a task is returned when they are implemented
assert response.headers["Content-Location"] == f"/api/v1/tasks/{task_mock.id}"
def test_providers_delete_invalid(self, client, tenant_header):
response = client.delete(
@@ -395,7 +404,7 @@ class TestProviderViewSet:
)
assert response.status_code == status.HTTP_404_NOT_FOUND
@patch("tasks.tasks.check_provider_connection_task.delay")
@patch("api.v1.views.check_provider_connection_task.delay")
def test_providers_connection(
self, mock_provider_connection, client, providers_fixture, tenant_header
):
@@ -417,7 +426,7 @@ class TestProviderViewSet:
provider_id=str(provider1.id), tenant_id=tenant_header["X-Tenant-ID"]
)
assert "Content-Location" in response.headers
assert response.headers["Content-Location"] == f"api/v1/tasks/{task_mock.id}"
assert response.headers["Content-Location"] == f"/api/v1/tasks/{task_mock.id}"
def test_providers_connection_invalid_provider(
self, client, providers_fixture, tenant_header
@@ -731,3 +740,55 @@ class TestScanViewSet:
reverse("scan-list"), {"sort": "invalid"}, headers=tenant_header
)
assert response.status_code == status.HTTP_400_BAD_REQUEST
@pytest.mark.django_db
class TestTaskViewSet:
def test_tasks_list(self, client, tasks_fixture, tenant_header):
response = client.get(reverse("task-list"), headers=tenant_header)
assert response.status_code == status.HTTP_200_OK
assert len(response.json()["data"]) == len(tasks_fixture)
def test_tasks_retrieve(self, client, tasks_fixture, tenant_header):
task1, *_ = tasks_fixture
response = client.get(
reverse("task-detail", kwargs={"pk": task1.id}),
headers=tenant_header,
)
assert response.status_code == status.HTTP_200_OK
assert (
response.json()["data"]["attributes"]["name"]
== task1.task_runner_task.task_name
)
def test_tasks_invalid_retrieve(self, client, tenant_header):
response = client.get(
reverse("task-detail", kwargs={"pk": "invalid_id"}), headers=tenant_header
)
assert response.status_code == status.HTTP_404_NOT_FOUND
@patch("api.v1.views.AsyncResult", return_value=Mock())
def test_tasks_revoke(
self, mock_async_result, client, tasks_fixture, tenant_header
):
_, task2 = tasks_fixture
response = client.delete(
reverse("task-detail", kwargs={"pk": task2.id}), headers=tenant_header
)
assert response.status_code == status.HTTP_202_ACCEPTED
assert response.headers["Content-Location"] == f"/api/v1/tasks/{task2.id}"
mock_async_result.return_value.revoke.assert_called_once()
def test_tasks_invalid_revoke(self, client, tenant_header):
response = client.delete(
reverse("task-detail", kwargs={"pk": "invalid_id"}), headers=tenant_header
)
assert response.status_code == status.HTTP_404_NOT_FOUND
def test_tasks_revoke_invalid_status(self, client, tasks_fixture, tenant_header):
task1, _ = tasks_fixture
response = client.delete(
reverse("task-detail", kwargs={"pk": task1.id}), headers=tenant_header
)
# Task status is SUCCESS
assert response.status_code == status.HTTP_400_BAD_REQUEST
+88 -8
View File
@@ -1,8 +1,10 @@
import json
from drf_spectacular.utils import extend_schema_field
from rest_framework_json_api import serializers
from rest_framework_json_api.serializers import ValidationError
from api.models import StateChoices, Provider, Scan
from api.models import StateChoices, Provider, Scan, Task
from api.rls import Tenant
from api.utils import merge_dicts
@@ -39,17 +41,95 @@ class StateEnumSerializerField(serializers.ChoiceField):
# Tasks
class DelayedTaskSerializer(serializers.Serializer):
id = serializers.CharField()
status = serializers.CharField()
class TaskBase(serializers.Serializer):
state_mapping = {
"PENDING": StateChoices.SCHEDULED,
"STARTED": StateChoices.EXECUTING,
"PROGRESS": StateChoices.EXECUTING,
"SUCCESS": StateChoices.COMPLETED,
"FAILURE": StateChoices.FAILED,
"REVOKED": StateChoices.CANCELLED,
}
class JSONAPIMeta:
resource_name = "Task"
def to_representation(self, obj):
return {"id": obj.id, "status": obj.status}
def map_state(self, task_result_state):
return self.state_mapping.get(task_result_state, StateChoices.AVAILABLE)
@extend_schema_field(
{
"type": "string",
"enum": StateChoices.values,
}
)
def get_state(self, obj):
task_result_state = (
obj.task_runner_task.status if obj.task_runner_task else None
)
return self.map_state(task_result_state)
class DelayedTaskSerializer(TaskBase):
id = serializers.CharField()
state = serializers.SerializerMethodField(read_only=True)
class Meta:
fields = [
"id",
"state",
]
@extend_schema_field(
{
"type": "string",
"enum": StateChoices.values,
}
)
def get_state(self, obj):
task_result_state = obj.status if obj else None
return self.map_state(task_result_state)
class TaskSerializer(RLSSerializer, TaskBase):
state = serializers.SerializerMethodField(read_only=True)
metadata = serializers.SerializerMethodField(read_only=True)
result = serializers.SerializerMethodField(read_only=True)
completed_at = serializers.DateTimeField(
source="task_runner_task.date_done", read_only=True
)
name = serializers.CharField(source="task_runner_task.task_name", read_only=True)
class Meta:
model = Task
fields = [
"id",
"inserted_at",
"completed_at",
"name",
"state",
"result",
"metadata",
]
@extend_schema_field(serializers.JSONField())
def get_metadata(self, obj):
return self.get_json_field(obj, "metadata")
@extend_schema_field(serializers.JSONField())
def get_result(self, obj):
return self.get_json_field(obj, "result")
@staticmethod
def get_json_field(obj, field_name):
"""Helper method to DRY the logic for loading JSON fields from task_runner_task."""
task_result_field = (
getattr(obj.task_runner_task, field_name, None)
if obj.task_runner_task
else None
)
return json.loads(task_result_field) if task_result_field else {}
# Tenants
+8 -1
View File
@@ -2,13 +2,20 @@ from django.urls import path, include
from drf_spectacular.views import SpectacularRedocView
from rest_framework import routers
from api.v1.views import SchemaView, TenantViewSet, ProviderViewSet, ScanViewSet
from api.v1.views import (
SchemaView,
TenantViewSet,
ProviderViewSet,
ScanViewSet,
TaskViewSet,
)
router = routers.DefaultRouter(trailing_slash=False)
router.register(r"tenants", TenantViewSet, basename="tenant")
router.register(r"providers", ProviderViewSet, basename="provider")
router.register(r"scans", ScanViewSet, basename="scan")
router.register(r"tasks", TaskViewSet, basename="task")
urlpatterns = [
path("", include(router.urls)),
+81 -16
View File
@@ -1,4 +1,5 @@
from django.conf import settings as django_settings
from django.db.models import F
from django.urls import reverse
from django.utils.decorators import method_decorator
from django.views.decorators.cache import cache_control
@@ -9,22 +10,24 @@ from rest_framework import status
from rest_framework.decorators import action
from rest_framework.generics import get_object_or_404
from rest_framework_json_api.views import Response
from celery.result import AsyncResult
from api.base_views import BaseRLSViewSet, BaseViewSet
from api.filters import ProviderFilter, TenantFilter, ScanFilter
from api.models import Provider, Scan
from api.filters import ProviderFilter, TenantFilter, ScanFilter, TaskFilter
from api.models import Provider, Scan, Task
from api.rls import Tenant
from api.v1.serializers import (
ProviderSerializer,
ProviderCreateSerializer,
ProviderUpdateSerializer,
TenantSerializer,
TaskSerializer,
DelayedTaskSerializer,
ScanSerializer,
ScanCreateSerializer,
ScanUpdateSerializer,
)
from tasks.tasks import check_provider_connection_task
from tasks.tasks import check_provider_connection_task, delete_provider_task
CACHE_DECORATOR = cache_control(
max_age=django_settings.CACHE_MAX_AGE,
@@ -102,6 +105,7 @@ class TenantViewSet(BaseViewSet):
responses={200: ProviderSerializer},
),
destroy=extend_schema(
tags=["Provider"],
summary="Delete a provider",
description="Remove a provider from the system by their ID.",
responses={202: DelayedTaskSerializer},
@@ -133,6 +137,8 @@ class ProviderViewSet(BaseRLSViewSet):
return ProviderCreateSerializer
elif self.action == "partial_update":
return ProviderUpdateSerializer
elif self.action in ["connection", "destroy"]:
return DelayedTaskSerializer
return super().get_serializer_class()
def partial_update(self, request, *args, **kwargs):
@@ -167,21 +173,24 @@ class ProviderViewSet(BaseRLSViewSet):
return Response(
data=serializer.data,
status=status.HTTP_202_ACCEPTED,
# TODO Use /tasks view name when implemented
# headers={"Content-Location": reverse("task-detail", kwargs={"pk": task.id})},
headers={"Content-Location": f"api/v1/tasks/{task.id}"},
headers={
"Content-Location": reverse("task-detail", kwargs={"pk": task.id})
},
)
def destroy(self, request, *args, **kwargs):
response = super().destroy(request, *args, **kwargs)
# TODO Background task to delete provider. For now, it will delete the provider from the system
# Same as /connection endpoint
response.status_code = status.HTTP_202_ACCEPTED
response.headers = {
"Content-Location": "/api/v1/tasks/5234",
**response.headers,
}
return response
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.headers.get("X-Tenant-ID")
)
serializer = DelayedTaskSerializer(task)
return Response(
data=serializer.data,
status=status.HTTP_202_ACCEPTED,
headers={
"Content-Location": reverse("task-detail", kwargs={"pk": task.id})
},
)
@extend_schema_view(
@@ -266,3 +275,59 @@ class ScanViewSet(BaseRLSViewSet):
instance, context=self.get_serializer_context()
)
return Response(data=read_serializer.data, status=status.HTTP_200_OK)
@extend_schema_view(
list=extend_schema(
summary="List all tasks",
description="Retrieve a list of all tasks with options for filtering by name, state, and other criteria.",
),
retrieve=extend_schema(
summary="Retrieve data from a specific task",
description="Fetch detailed information about a specific task by its ID.",
),
destroy=extend_schema(
tags=["Task"],
summary="Revoke a task",
description="Try to revoke a task using its ID. Only tasks that are not yet in progress can be revoked.",
responses={202: DelayedTaskSerializer},
),
)
class TaskViewSet(BaseRLSViewSet):
queryset = Task.objects.all()
serializer_class = TaskSerializer
http_method_names = ["get", "delete"]
filterset_class = TaskFilter
search_fields = ["name"]
ordering = ["inserted_at"]
ordering_fields = ["inserted_at", "completed_at", "name", "state"]
def get_queryset(self):
return Task.objects.annotate(
name=F("task_runner_task__task_name"), state=F("task_runner_task__status")
)
def destroy(self, request, *args, pk=None, **kwargs):
task = get_object_or_404(Task, pk=pk)
if task.task_runner_task.status not in ["PENDING", "RECEIVED"]:
serializer = TaskSerializer(task)
return Response(
data={
"detail": f"Task cannot be revoked. Status: '{serializer.data.get('state')}'"
},
status=status.HTTP_400_BAD_REQUEST,
headers={
"Content-Location": reverse("task-detail", kwargs={"pk": task.id})
},
)
task_instance = AsyncResult(pk)
task_instance.revoke()
serializer = DelayedTaskSerializer(task_instance)
return Response(
data=serializer.data,
status=status.HTTP_202_ACCEPTED,
headers={
"Content-Location": reverse("task-detail", kwargs={"pk": task.id})
},
)
+70 -1
View File
@@ -1,5 +1,9 @@
from celery import Celery
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")
@@ -7,3 +11,68 @@ celery_app.config_from_object("django.conf:settings", namespace="CELERY")
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,
kwargs=None,
task_id=None,
producer=None,
link=None,
link_error=None,
shadow=None,
**options,
):
from api.models import Task as APITask
from django_celery_results.models import TaskResult
result = super().apply_async(
args=args,
kwargs=kwargs,
task_id=task_id,
producer=producer,
link=link,
link_error=link_error,
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,
tenant_id=kwargs.get("tenant_id"),
task_runner_task=task_result_instance,
)
return result
+30 -2
View File
@@ -4,8 +4,8 @@ import pytest
from django.conf import settings
from django.db import connections as django_connections
from rest_framework import status
from api.models import Provider, Scan, StateChoices
from django_celery_results.models import TaskResult
from api.models import Provider, Scan, StateChoices, Task
from api.rls import Tenant
API_JSON_CONTENT_TYPE = "application/vnd.api+json"
@@ -125,6 +125,34 @@ def scans_fixture(tenants_fixture, providers_fixture):
return scan1, scan2, scan3
@pytest.fixture
def tasks_fixture(tenants_fixture):
tenant, _ = tenants_fixture
task_runner_task1 = TaskResult.objects.create(
task_id="81a1b34b-ff6e-498e-979c-d6a83260167f",
task_name="task_runner_task1",
status="SUCCESS",
)
task_runner_task2 = TaskResult.objects.create(
task_id="4d0260a5-2e1f-4a34-a976-8c5acb9f5499",
task_name="task_runner_task1",
status="PENDING",
)
task1 = Task.objects.create(
id=task_runner_task1.task_id,
task_runner_task=task_runner_task1,
tenant_id=tenant.id,
)
task2 = Task.objects.create(
id=task_runner_task2.task_id,
task_runner_task=task_runner_task2,
tenant_id=tenant.id,
)
return task1, task2
@pytest.fixture
def tenant_header(tenants_fixture):
return {"X-Tenant-ID": str(tenants_fixture[0].id)}
+25
View File
@@ -0,0 +1,25 @@
from celery.utils.log import get_task_logger
logger = get_task_logger(__name__)
def delete_instance(model, pk: str):
"""
Deletes an instance of the specified model.
This function retrieves an instance of the provided model using its primary key
and deletes it from the database.
Args:
model (Model): The Django model class from which to delete an instance.
pk (str): The primary key of the instance to delete.
Returns:
tuple: A tuple containing the number of objects deleted and a dictionary
with the count of deleted objects per model,
including related models if applicable.
Raises:
model.DoesNotExist: If no instance with the provided primary key exists.
"""
return model.objects.get(pk=pk).delete()
+22 -1
View File
@@ -1,10 +1,13 @@
from celery import shared_task
from api.decorators import set_tenant
from api.models import Provider
from config.celery import RLSTask
from tasks.jobs.connection import check_provider_connection
from tasks.jobs.deletion import delete_instance
@shared_task(name="provider-connection-check")
@shared_task(base=RLSTask, name="provider-connection-check")
@set_tenant
def check_provider_connection_task(provider_id: str):
"""
@@ -19,3 +22,21 @@ def check_provider_connection_task(provider_id: str):
- 'error' (str or None): The error message if the connection failed, otherwise `None`.
"""
return check_provider_connection(provider_id)
@shared_task(base=RLSTask, name="provider-deletion")
@set_tenant
def delete_provider_task(provider_id: str):
"""
Task to delete a specific Provider instance.
Args:
provider_id (str): The primary key of the `Provider` instance to be deleted.
Returns:
tuple: A tuple containing:
- The number of instances deleted.
- A dictionary with the count of deleted instances per model,
including related models if cascading deletes were triggered.
"""
return delete_instance(model=Provider, pk=provider_id)
+22
View File
@@ -0,0 +1,22 @@
import pytest
from django.core.exceptions import ObjectDoesNotExist
from api.models import Provider
from tasks.jobs.deletion import delete_instance
@pytest.mark.django_db
class TestDeleteInstance:
def test_delete_instance_success(self, providers_fixture):
instance = providers_fixture[0]
result = delete_instance(Provider, instance.id)
assert result
with pytest.raises(ObjectDoesNotExist):
Provider.objects.get(pk=instance.id)
def test_delete_instance_does_not_exist(self):
non_existent_pk = "babf6796-cfcc-4fd3-9dcf-88d012247645"
with pytest.raises(ObjectDoesNotExist):
delete_instance(Provider, non_existent_pk)