From ec67fc12e040b5c20fdc886a76f4d8fc9313faa4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?V=C3=ADctor=20Fern=C3=A1ndez=20Poyatos?= Date: Sat, 7 Sep 2024 02:47:51 +0200 Subject: [PATCH] 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 --- src/backend/api/db_utils.py | 13 + src/backend/api/filters.py | 65 +- src/backend/api/migrations/0001_initial.py | 87 +- src/backend/api/models.py | 32 + src/backend/api/rls.py | 13 +- src/backend/api/specs/v1.yaml | 931 ++++++++++++++++++--- src/backend/api/tests/test_views.py | 71 +- src/backend/api/v1/serializers.py | 96 ++- src/backend/api/v1/urls.py | 9 +- src/backend/api/v1/views.py | 97 ++- src/backend/config/celery.py | 71 +- src/backend/conftest.py | 32 +- src/backend/tasks/jobs/deletion.py | 25 + src/backend/tasks/tasks.py | 23 +- src/backend/tasks/tests/test_deletion.py | 22 + 15 files changed, 1426 insertions(+), 161 deletions(-) create mode 100644 src/backend/tasks/jobs/deletion.py create mode 100644 src/backend/tasks/tests/test_deletion.py diff --git a/src/backend/api/db_utils.py b/src/backend/api/db_utils.py index 87f3081479..ca35350308 100644 --- a/src/backend/api/db_utils.py +++ b/src/backend/api/db_utils.py @@ -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): diff --git a/src/backend/api/filters.py b/src/backend/api/filters.py index 9b77f53f02..f374b87359 100644 --- a/src/backend/api/filters.py +++ b/src/backend/api/filters.py @@ -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 = [] diff --git a/src/backend/api/migrations/0001_initial.py b/src/backend/api/migrations/0001_initial.py index a2bb746224..65124f869a 100644 --- a/src/backend/api/migrations/0001_initial.py +++ b/src/backend/api/migrations/0001_initial.py @@ -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}; + """ + ), ] diff --git a/src/backend/api/models.py b/src/backend/api/models.py index d313ea8e7f..2984adb399 100644 --- a/src/backend/api/models.py +++ b/src/backend/api/models.py @@ -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", + ), + ] diff --git a/src/backend/api/rls.py b/src/backend/api/rls.py index 965b092b59..1d1c6d2a4e 100644 --- a/src/backend/api/rls.py +++ b/src/backend/api/rls.py @@ -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: diff --git a/src/backend/api/specs/v1.yaml b/src/backend/api/specs/v1.yaml index b0c82f694a..e2d355e0a7 100644 --- a/src/backend/api/specs/v1.yaml +++ b/src/backend/api/specs/v1.yaml @@ -27,7 +27,8 @@ paths: - provider_id - alias - connection - - metadata + - scanner_args + - url description: endpoint return only specific fields in the response on a per-type basis by including a fields[TYPE] query parameter. explode: false @@ -62,20 +63,14 @@ paths: name: filter[provider] schema: type: string - enum: - - aws - - azure - - gcp - - kubernetes - description: |- - * `aws` - Aws - * `azure` - Azure - * `gcp` - Gcp - * `kubernetes` - Kubernetes - in: query name: filter[provider_id] schema: type: string + - in: query + name: filter[provider_id__icontains] + schema: + type: string - name: filter[search] required: false in: query @@ -155,13 +150,13 @@ paths: content: application/vnd.api+json: schema: - $ref: '#/components/schemas/ProviderRequest' + $ref: '#/components/schemas/ProviderCreateRequest' application/x-www-form-urlencoded: schema: - $ref: '#/components/schemas/ProviderRequest' + $ref: '#/components/schemas/ProviderCreateRequest' multipart/form-data: schema: - $ref: '#/components/schemas/ProviderRequest' + $ref: '#/components/schemas/ProviderCreateRequest' required: true security: - cookieAuth: [] @@ -172,7 +167,7 @@ paths: content: application/vnd.api+json: schema: - $ref: '#/components/schemas/ProviderResponse' + $ref: '#/components/schemas/ProviderCreateResponse' description: '' /api/v1/providers/{id}: get: @@ -193,7 +188,8 @@ paths: - provider_id - alias - connection - - metadata + - scanner_args + - url description: endpoint return only specific fields in the response on a per-type basis by including a fields[TYPE] query parameter. explode: false @@ -236,13 +232,13 @@ paths: content: application/vnd.api+json: schema: - $ref: '#/components/schemas/PatchedProviderRequest' + $ref: '#/components/schemas/PatchedProviderUpdateRequest' application/x-www-form-urlencoded: schema: - $ref: '#/components/schemas/PatchedProviderRequest' + $ref: '#/components/schemas/PatchedProviderUpdateRequest' multipart/form-data: schema: - $ref: '#/components/schemas/PatchedProviderRequest' + $ref: '#/components/schemas/PatchedProviderUpdateRequest' required: true security: - cookieAuth: [] @@ -253,7 +249,7 @@ paths: content: application/vnd.api+json: schema: - $ref: '#/components/schemas/ProviderResponse' + $ref: '#/components/schemas/SerializerMetaclassResponse' description: '' delete: operationId: providers_destroy @@ -307,6 +303,399 @@ paths: schema: $ref: '#/components/schemas/SerializerMetaclassResponse' description: '' + /api/v1/scans: + get: + operationId: scans_list + description: Retrieve a list of all scans with options for filtering by various + criteria. + summary: List all scans + parameters: + - in: query + name: fields[Scan] + schema: + type: array + items: + type: string + enum: + - name + - state + - unique_resource_count + - progress + - scanner_args + - duration + - provider + - started_at + - completed_at + - scheduled_at + - url + - type + description: endpoint return only specific fields in the response on a per-type + basis by including a fields[TYPE] query parameter. + explode: false + - in: query + name: filter[name] + schema: + type: string + - in: query + name: filter[name__icontains] + schema: + type: string + - in: query + name: filter[provider] + schema: + type: string + - in: query + name: filter[provider_id] + schema: + type: string + format: uuid + - name: filter[search] + required: false + in: query + description: A search term. + schema: + type: string + - in: query + name: filter[started_at] + schema: + type: string + format: date-time + - in: query + name: filter[started_at__gte] + schema: + type: string + format: date-time + - in: query + name: filter[started_at__lte] + schema: + type: string + format: date-time + - in: query + name: filter[type] + schema: + type: string + - name: page[number] + required: false + in: query + description: A page number within the paginated result set. + schema: + type: integer + - name: page[size] + required: false + in: query + description: Number of results to return per page. + schema: + type: integer + - name: sort + required: false + in: query + description: '[list of fields to sort by](https://jsonapi.org/format/#fetching-sorting)' + schema: + type: array + items: + type: string + enum: + - provider_id + - -provider_id + - name + - -name + - type + - -type + - attempted_at + - -attempted_at + - scheduled_at + - -scheduled_at + - inserted_at + - -inserted_at + - updated_at + - -updated_at + explode: false + tags: + - Scan + security: + - cookieAuth: [] + - basicAuth: [] + - {} + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/PaginatedScanList' + description: '' + post: + operationId: scans_create + description: Trigger a manual scan by providing the required scan details. If + `scanner_args` are not provided, the system will automatically use the default + settings from the associated provider. If you do provide `scanner_args`, these + settings will be merged with the provider's defaults. This means that your + provided settings will override the defaults only where they conflict, while + the rest of the default settings will remain intact. + summary: Trigger a manual scan + tags: + - Scan + requestBody: + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/ScanCreateRequest' + application/x-www-form-urlencoded: + schema: + $ref: '#/components/schemas/ScanCreateRequest' + multipart/form-data: + schema: + $ref: '#/components/schemas/ScanCreateRequest' + required: true + security: + - cookieAuth: [] + - basicAuth: [] + - {} + responses: + '201': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/ScanCreateResponse' + description: '' + /api/v1/scans/{id}: + get: + operationId: scans_retrieve + description: Fetch detailed information about a specific scan by its ID. + summary: Retrieve data from a specific scan + parameters: + - in: query + name: fields[Scan] + schema: + type: array + items: + type: string + enum: + - name + - state + - unique_resource_count + - progress + - scanner_args + - duration + - provider + - started_at + - completed_at + - scheduled_at + - url + - type + description: endpoint return only specific fields in the response on a per-type + basis by including a fields[TYPE] query parameter. + explode: false + - in: path + name: id + schema: + type: string + format: uuid + description: A UUID string identifying this scan. + required: true + tags: + - Scan + security: + - cookieAuth: [] + - basicAuth: [] + - {} + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/ScanResponse' + description: '' + patch: + operationId: scans_partial_update + description: Update certain fields of an existing scan without affecting other + fields. + summary: Partially update a scan + parameters: + - in: path + name: id + schema: + type: string + format: uuid + description: A UUID string identifying this scan. + required: true + tags: + - Scan + requestBody: + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/PatchedScanUpdateRequest' + application/x-www-form-urlencoded: + schema: + $ref: '#/components/schemas/PatchedScanUpdateRequest' + multipart/form-data: + schema: + $ref: '#/components/schemas/PatchedScanUpdateRequest' + required: true + security: + - cookieAuth: [] + - basicAuth: [] + - {} + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/SerializerMetaclassResponse' + description: '' + /api/v1/tasks: + get: + operationId: tasks_list + description: Retrieve a list of all tasks with options for filtering by name, + state, and other criteria. + summary: List all tasks + parameters: + - in: query + name: fields[Task] + schema: + type: array + items: + type: string + enum: + - inserted_at + - completed_at + - name + - state + - result + - metadata + description: endpoint return only specific fields in the response on a per-type + basis by including a fields[TYPE] query parameter. + explode: false + - in: query + name: filter[name] + schema: + type: string + - in: query + name: filter[name__icontains] + schema: + type: string + - name: filter[search] + required: false + in: query + description: A search term. + schema: + type: string + - in: query + name: filter[state] + schema: + type: string + - name: page[number] + required: false + in: query + description: A page number within the paginated result set. + schema: + type: integer + - name: page[size] + required: false + in: query + description: Number of results to return per page. + schema: + type: integer + - name: sort + required: false + in: query + description: '[list of fields to sort by](https://jsonapi.org/format/#fetching-sorting)' + schema: + type: array + items: + type: string + enum: + - inserted_at + - -inserted_at + - completed_at + - -completed_at + - name + - -name + - state + - -state + explode: false + tags: + - Task + security: + - cookieAuth: [] + - basicAuth: [] + - {} + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/PaginatedTaskList' + description: '' + /api/v1/tasks/{id}: + get: + operationId: tasks_retrieve + description: Fetch detailed information about a specific task by its ID. + summary: Retrieve data from a specific task + parameters: + - in: query + name: fields[Task] + schema: + type: array + items: + type: string + enum: + - inserted_at + - completed_at + - name + - state + - result + - metadata + description: endpoint return only specific fields in the response on a per-type + basis by including a fields[TYPE] query parameter. + explode: false + - in: path + name: id + schema: + type: string + format: uuid + description: A UUID string identifying this task. + required: true + tags: + - Task + security: + - cookieAuth: [] + - basicAuth: [] + - {} + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/TaskResponse' + description: '' + delete: + operationId: tasks_destroy + description: Try to revoke a task using its ID. If the task is being already + executed, its result will be ignored but it will finish. To prevent this, + use the `force` query parameter under your own risk. + summary: Revoke a task + parameters: + - in: path + name: id + schema: + type: string + format: uuid + description: A UUID string identifying this task. + required: true + tags: + - Task + security: + - cookieAuth: [] + - basicAuth: [] + - {} + responses: + '202': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/SerializerMetaclassResponse' + description: '' /api/v1/tenants: get: operationId: tenants_list @@ -542,6 +931,35 @@ paths: description: No response body components: schemas: + DelayedTask: + type: object + required: + - type + additionalProperties: false + properties: + type: + allOf: + - $ref: '#/components/schemas/TypeB52Enum' + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + attributes: + type: object + properties: + id: + type: string + state: + type: string + enum: + - available + - scheduled + - executing + - completed + - failed + - cancelled + readOnly: true + required: + - id PaginatedProviderList: type: object properties: @@ -551,6 +969,24 @@ components: $ref: '#/components/schemas/Provider' required: - data + PaginatedScanList: + type: object + properties: + data: + type: array + items: + $ref: '#/components/schemas/Scan' + required: + - data + PaginatedTaskList: + type: object + properties: + data: + type: array + items: + $ref: '#/components/schemas/Task' + required: + - data PaginatedTenantList: type: object properties: @@ -560,7 +996,7 @@ components: $ref: '#/components/schemas/Tenant' required: - data - PatchedProviderRequest: + PatchedProviderUpdateRequest: type: object properties: data: @@ -577,53 +1013,46 @@ components: and relationships. enum: - Provider + id: {} + attributes: + type: object + properties: + alias: + type: string + nullable: true + maxLength: 100 + minLength: 3 + scanner_args: {} + required: + - data + PatchedScanUpdateRequest: + type: object + properties: + data: + type: object + required: + - type + - id + additionalProperties: false + properties: + type: + type: string + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + enum: + - Scan id: type: string format: uuid attributes: type: object properties: - inserted_at: + name: type: string - format: date-time - readOnly: true - updated_at: - type: string - format: date-time - readOnly: true - provider: - enum: - - aws - - azure - - gcp - - kubernetes - type: string - description: |- - * `aws` - Aws - * `azure` - Azure - * `gcp` - Gcp - * `kubernetes` - Kubernetes - provider_id: - type: string - minLength: 3 - maxLength: 63 - alias: - type: string - minLength: 3 + nullable: true maxLength: 100 - connection: - type: object - properties: - connected: - type: boolean - last_checked_at: - type: string - format: date-time - readOnly: true - metadata: {} - required: - - provider_id - - alias + minLength: 3 required: - data PatchedTenantRequest: @@ -674,7 +1103,7 @@ components: properties: type: allOf: - - $ref: '#/components/schemas/ProviderTypeEnum' + - $ref: '#/components/schemas/Type4e8Enum' description: The [type](https://jsonapi.org/format/#document-resource-object-identification) member is used to describe resource objects that share common attributes and relationships. @@ -700,9 +1129,9 @@ components: - kubernetes type: string description: |- - * `aws` - Aws + * `aws` - AWS * `azure` - Azure - * `gcp` - Gcp + * `gcp` - GCP * `kubernetes` - Kubernetes provider_id: type: string @@ -710,6 +1139,7 @@ components: minLength: 3 alias: type: string + nullable: true maxLength: 100 minLength: 3 connection: @@ -721,11 +1151,50 @@ components: type: string format: date-time readOnly: true - metadata: {} + scanner_args: {} + required: + - provider + - provider_id + ProviderCreate: + type: object + required: + - type + additionalProperties: false + properties: + type: + allOf: + - $ref: '#/components/schemas/Type4e8Enum' + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + attributes: + type: object + properties: + alias: + type: string + nullable: true + maxLength: 100 + minLength: 3 + provider: + enum: + - aws + - azure + - gcp + - kubernetes + type: string + description: |- + * `aws` - AWS + * `azure` - Azure + * `gcp` - GCP + * `kubernetes` - Kubernetes + provider_id: + type: string + maxLength: 63 + minLength: 3 + scanner_args: {} required: - provider_id - - alias - ProviderRequest: + ProviderCreateRequest: type: object properties: data: @@ -744,14 +1213,11 @@ components: attributes: type: object properties: - inserted_at: + alias: type: string - format: date-time - readOnly: true - updated_at: - type: string - format: date-time - readOnly: true + nullable: true + maxLength: 100 + minLength: 3 provider: enum: - aws @@ -760,31 +1226,24 @@ components: - kubernetes type: string description: |- - * `aws` - Aws + * `aws` - AWS * `azure` - Azure - * `gcp` - Gcp + * `gcp` - GCP * `kubernetes` - Kubernetes provider_id: type: string minLength: 3 maxLength: 63 - alias: - type: string - minLength: 3 - maxLength: 100 - connection: - type: object - properties: - connected: - type: boolean - last_checked_at: - type: string - format: date-time - readOnly: true - metadata: {} + scanner_args: {} required: - provider_id - - alias + required: + - data + ProviderCreateResponse: + type: object + properties: + data: + $ref: '#/components/schemas/ProviderCreate' required: - data ProviderResponse: @@ -794,10 +1253,231 @@ components: $ref: '#/components/schemas/Provider' required: - data - ProviderTypeEnum: - type: string - enum: - - Provider + Scan: + type: object + required: + - type + - id + additionalProperties: false + properties: + type: + allOf: + - $ref: '#/components/schemas/Type2bbEnum' + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + id: + type: string + format: uuid + attributes: + type: object + properties: + name: + type: string + nullable: true + maxLength: 100 + minLength: 3 + state: + enum: + - available + - scheduled + - executing + - completed + - failed + - cancelled + type: string + description: |- + * `available` - Available + * `scheduled` - Scheduled + * `executing` - Executing + * `completed` - Completed + * `failed` - Failed + * `cancelled` - Cancelled + readOnly: true + unique_resource_count: + type: integer + maximum: 2147483647 + minimum: -2147483648 + progress: + type: integer + maximum: 2147483647 + minimum: -2147483648 + scanner_args: {} + duration: + type: integer + maximum: 2147483647 + minimum: -2147483648 + nullable: true + started_at: + type: string + format: date-time + nullable: true + completed_at: + type: string + format: date-time + nullable: true + scheduled_at: + type: string + format: date-time + nullable: true + type: + enum: + - scheduled + - manual + type: string + description: |- + * `scheduled` - Scheduled + * `manual` - Manual + readOnly: true + relationships: + type: object + properties: + provider: + type: object + properties: + data: + type: object + properties: + id: + type: string + format: uuid + type: + type: string + enum: + - Provider + title: Resource Type Name + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common + attributes and relationships. + required: + - id + - type + required: + - data + description: The identifier of the related object. + title: Resource Identifier + required: + - provider + ScanCreate: + type: object + required: + - type + additionalProperties: false + properties: + type: + allOf: + - $ref: '#/components/schemas/Type2bbEnum' + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + attributes: + type: object + properties: + scanner_args: {} + name: + type: string + nullable: true + maxLength: 100 + minLength: 3 + relationships: + type: object + properties: + provider: + type: object + properties: + data: + type: object + properties: + id: + type: string + format: uuid + type: + type: string + enum: + - Provider + title: Resource Type Name + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common + attributes and relationships. + required: + - id + - type + required: + - data + description: The identifier of the related object. + title: Resource Identifier + required: + - provider + ScanCreateRequest: + type: object + properties: + data: + type: object + required: + - type + additionalProperties: false + properties: + type: + type: string + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + enum: + - Scan + attributes: + type: object + properties: + scanner_args: {} + name: + type: string + nullable: true + maxLength: 100 + minLength: 3 + relationships: + type: object + properties: + provider: + type: object + properties: + data: + type: object + properties: + id: + type: string + format: uuid + type: + type: string + enum: + - Provider + title: Resource Type Name + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share + common attributes and relationships. + required: + - id + - type + required: + - data + description: The identifier of the related object. + title: Resource Identifier + required: + - provider + required: + - data + ScanCreateResponse: + type: object + properties: + data: + $ref: '#/components/schemas/ScanCreate' + required: + - data + ScanResponse: + type: object + properties: + data: + $ref: '#/components/schemas/Scan' + required: + - data SerializerMetaclassResponse: type: object properties: @@ -805,6 +1485,57 @@ components: $ref: '#/components/schemas/Provider' required: - data + Task: + type: object + required: + - type + - id + additionalProperties: false + properties: + type: + allOf: + - $ref: '#/components/schemas/TypeB52Enum' + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + id: + type: string + format: uuid + attributes: + type: object + properties: + inserted_at: + type: string + format: date-time + readOnly: true + completed_at: + type: string + format: date-time + readOnly: true + name: + type: string + readOnly: true + state: + type: string + enum: + - available + - scheduled + - executing + - completed + - failed + - cancelled + readOnly: true + result: + readOnly: true + metadata: + readOnly: true + TaskResponse: + type: object + properties: + data: + $ref: '#/components/schemas/Task' + required: + - data Tenant: type: object required: @@ -883,6 +1614,18 @@ components: type: string enum: - Tenant + Type2bbEnum: + type: string + enum: + - Scan + Type4e8Enum: + type: string + enum: + - Provider + TypeB52Enum: + type: string + enum: + - Task securitySchemes: basicAuth: type: http diff --git a/src/backend/api/tests/test_views.py b/src/backend/api/tests/test_views.py index 799cff0e8e..6122ab2a03 100644 --- a/src/backend/api/tests/test_views.py +++ b/src/backend/api/tests/test_views.py @@ -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 diff --git a/src/backend/api/v1/serializers.py b/src/backend/api/v1/serializers.py index 72e0a5a669..f38e830d60 100644 --- a/src/backend/api/v1/serializers.py +++ b/src/backend/api/v1/serializers.py @@ -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 diff --git a/src/backend/api/v1/urls.py b/src/backend/api/v1/urls.py index 15fd1b4685..7e7d1e9c07 100644 --- a/src/backend/api/v1/urls.py +++ b/src/backend/api/v1/urls.py @@ -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)), diff --git a/src/backend/api/v1/views.py b/src/backend/api/v1/views.py index fb3eb68700..fd2f1e5aad 100644 --- a/src/backend/api/v1/views.py +++ b/src/backend/api/v1/views.py @@ -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}) + }, + ) diff --git a/src/backend/config/celery.py b/src/backend/config/celery.py index ce508b3406..409f0003f2 100644 --- a/src/backend/config/celery.py +++ b/src/backend/config/celery.py @@ -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 diff --git a/src/backend/conftest.py b/src/backend/conftest.py index f5e318b512..30c953028a 100644 --- a/src/backend/conftest.py +++ b/src/backend/conftest.py @@ -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)} diff --git a/src/backend/tasks/jobs/deletion.py b/src/backend/tasks/jobs/deletion.py new file mode 100644 index 0000000000..b203cf113e --- /dev/null +++ b/src/backend/tasks/jobs/deletion.py @@ -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() diff --git a/src/backend/tasks/tasks.py b/src/backend/tasks/tasks.py index daeb19ec69..8fb8f1a199 100644 --- a/src/backend/tasks/tasks.py +++ b/src/backend/tasks/tasks.py @@ -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) diff --git a/src/backend/tasks/tests/test_deletion.py b/src/backend/tasks/tests/test_deletion.py new file mode 100644 index 0000000000..630d1d1fa1 --- /dev/null +++ b/src/backend/tasks/tests/test_deletion.py @@ -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)