mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-23 04:21:52 +00:00
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:
committed by
GitHub
parent
f5462c9b27
commit
ec67fc12e0
@@ -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
@@ -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 = []
|
||||
|
||||
@@ -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};
|
||||
"""
|
||||
),
|
||||
]
|
||||
|
||||
@@ -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",
|
||||
),
|
||||
]
|
||||
|
||||
@@ -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:
|
||||
|
||||
+837
-94
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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})
|
||||
},
|
||||
)
|
||||
|
||||
@@ -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
@@ -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)}
|
||||
|
||||
@@ -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()
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user