From ffaa267b5ee5511bad675d07d688d930c6aa01b1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?V=C3=ADctor=20Fern=C3=A1ndez=20Poyatos?= Date: Sun, 1 Dec 2024 09:12:19 +0100 Subject: [PATCH] feat(scan, schedule): add next_scan_at field to scans and POST /schedules/daily (#5978) --- api/src/backend/api/filters.py | 2 + .../backend/api/migrations/0001_initial.py | 1 + api/src/backend/api/models.py | 1 + api/src/backend/api/specs/v1.yaml | 81 ++++++++++++++++++- api/src/backend/api/tests/test_views.py | 39 +++++++++ api/src/backend/api/v1/serializers.py | 21 +++++ api/src/backend/api/v1/urls.py | 30 +++---- api/src/backend/api/v1/views.py | 63 +++++++++++++-- api/src/backend/tasks/beat.py | 26 +++++- api/src/backend/tasks/tasks.py | 19 ++++- api/src/backend/tasks/tests/test_beat.py | 53 ++++++++++++ 11 files changed, 306 insertions(+), 30 deletions(-) create mode 100644 api/src/backend/tasks/tests/test_beat.py diff --git a/api/src/backend/api/filters.py b/api/src/backend/api/filters.py index 7a3e850430..9b1156a89a 100644 --- a/api/src/backend/api/filters.py +++ b/api/src/backend/api/filters.py @@ -164,6 +164,7 @@ class ScanFilter(ProviderRelationshipFilterSet): inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") completed_at = DateFilter(field_name="completed_at", lookup_expr="date") started_at = DateFilter(field_name="started_at", lookup_expr="date") + next_scan_at = DateFilter(field_name="next_scan_at", lookup_expr="date") trigger = ChoiceFilter(choices=Scan.TriggerChoices.choices) class Meta: @@ -172,6 +173,7 @@ class ScanFilter(ProviderRelationshipFilterSet): "provider": ["exact", "in"], "name": ["exact", "icontains"], "started_at": ["gte", "lte"], + "next_scan_at": ["gte", "lte"], "trigger": ["exact"], } diff --git a/api/src/backend/api/migrations/0001_initial.py b/api/src/backend/api/migrations/0001_initial.py index f2ec03ed5b..6288cf4093 100644 --- a/api/src/backend/api/migrations/0001_initial.py +++ b/api/src/backend/api/migrations/0001_initial.py @@ -679,6 +679,7 @@ class Migration(migrations.Migration): ("updated_at", models.DateTimeField(auto_now=True)), ("started_at", models.DateTimeField(null=True, blank=True)), ("completed_at", models.DateTimeField(null=True, blank=True)), + ("next_scan_at", models.DateTimeField(null=True, blank=True)), ( "provider", models.ForeignKey( diff --git a/api/src/backend/api/models.py b/api/src/backend/api/models.py index 4c106cd8c7..8b2dc54072 100644 --- a/api/src/backend/api/models.py +++ b/api/src/backend/api/models.py @@ -403,6 +403,7 @@ class Scan(RowLevelSecurityProtectedModel): updated_at = models.DateTimeField(auto_now=True, editable=False) started_at = models.DateTimeField(null=True, blank=True) completed_at = models.DateTimeField(null=True, blank=True) + next_scan_at = models.DateTimeField(null=True, blank=True) # TODO: mutelist foreign key class Meta(RowLevelSecurityProtectedModel.Meta): diff --git a/api/src/backend/api/specs/v1.yaml b/api/src/backend/api/specs/v1.yaml index 047d1920e7..3a72c33e9a 100644 --- a/api/src/backend/api/specs/v1.yaml +++ b/api/src/backend/api/specs/v1.yaml @@ -754,9 +754,8 @@ paths: /api/v1/findings/findings_services_regions: get: operationId: findings_findings_services_regions_retrieve - description: Fetch services and regions affected in findings by date. - summary: Retrieve the services and regions that are impacted by findings in - a specific date + description: Fetch services and regions affected in findings. + summary: Retrieve the services and regions that are impacted by findings parameters: - in: query name: fields[finding-dynamic-filters] @@ -2765,6 +2764,7 @@ paths: - started_at - completed_at - scheduled_at + - next_scan_at - url description: endpoint return only specific fields in the response on a per-type basis by including a fields[TYPE] query parameter. @@ -2787,6 +2787,21 @@ paths: name: filter[name__icontains] schema: type: string + - in: query + name: filter[next_scan_at] + schema: + type: string + format: date + - in: query + name: filter[next_scan_at__gte] + schema: + type: string + format: date-time + - in: query + name: filter[next_scan_at__lte] + schema: + type: string + format: date-time - in: query name: filter[provider] schema: @@ -3002,6 +3017,7 @@ paths: - started_at - completed_at - scheduled_at + - next_scan_at - url description: endpoint return only specific fields in the response on a per-type basis by including a fields[TYPE] query parameter. @@ -3060,6 +3076,35 @@ paths: schema: $ref: '#/components/schemas/ScanUpdateResponse' description: '' + /api/v1/schedules/daily: + post: + operationId: schedules_daily_create + description: Schedules a daily scan for the specified provider. This endpoint + creates a periodic task that will execute a scan every 24 hours. + summary: Create a daily schedule scan for a given provider + tags: + - Schedule + requestBody: + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/ScheduleDailyCreateRequest' + application/x-www-form-urlencoded: + schema: + $ref: '#/components/schemas/ScheduleDailyCreateRequest' + multipart/form-data: + schema: + $ref: '#/components/schemas/ScheduleDailyCreateRequest' + required: true + security: + - jwtAuth: [] + responses: + '202': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/OpenApiResponseResponse' + description: '' /api/v1/tasks: get: operationId: tasks_list @@ -7058,6 +7103,10 @@ components: type: string format: date-time nullable: true + next_scan_at: + type: string + format: date-time + nullable: true relationships: type: object properties: @@ -7205,6 +7254,32 @@ components: $ref: '#/components/schemas/ScanUpdate' required: - data + ScheduleDailyCreateRequest: + 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: + - daily-schedules + attributes: + type: object + properties: + provider_id: + type: string + format: uuid + required: + - provider_id + required: + - data SerializerMetaclassResponse: type: object properties: diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index e98ef25fef..e064452a83 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -3325,3 +3325,42 @@ class TestOverviewViewSet: ) # TODO Add more tests for the rest of overviews + + +@pytest.mark.django_db +class TestScheduleViewSet: + @pytest.mark.parametrize("method", ["get", "post"]) + def test_schedule_invalid_method_list(self, method, authenticated_client): + response = getattr(authenticated_client, method)(reverse("schedule-list")) + assert response.status_code == status.HTTP_405_METHOD_NOT_ALLOWED + + @patch("api.v1.views.Task.objects.get") + @patch("api.v1.views.schedule_provider_scan") + def test_schedule_daily( + self, + mock_schedule_scan, + mock_task_get, + authenticated_client, + providers_fixture, + tasks_fixture, + ): + provider, *_ = providers_fixture + prowler_task = tasks_fixture[0] + mock_schedule_scan.return_value.id = prowler_task.id + mock_task_get.return_value = prowler_task + json_payload = { + "provider_id": str(provider.id), + } + response = authenticated_client.post( + reverse("schedule-daily"), data=json_payload, format="json" + ) + assert response.status_code == status.HTTP_202_ACCEPTED + + def test_schedule_daily_provider_does_not_exist(self, authenticated_client): + json_payload = { + "provider_id": "4846c2f9-84b2-442b-94dd-3082e8eb9584", + } + response = authenticated_client.post( + reverse("schedule-daily"), data=json_payload, format="json" + ) + assert response.status_code == status.HTTP_404_NOT_FOUND diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 5c14221299..3c9d42bdef 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -536,9 +536,11 @@ class ScanSerializer(RLSSerializer): "duration", "provider", "task", + "inserted_at", "started_at", "completed_at", "scheduled_at", + "next_scan_at", "url", ] @@ -1319,3 +1321,22 @@ class OverviewSeveritySerializer(serializers.Serializer): def get_root_meta(self, _resource, _many): return {"version": "v1"} + + +# Schedules + + +class ScheduleDailyCreateSerializer(serializers.Serializer): + provider_id = serializers.UUIDField(required=True) + + class JSONAPIMeta: + resource_name = "daily-schedules" + + # TODO: DRY this when we have more time + def validate(self, data): + if hasattr(self, "initial_data"): + initial_data = set(self.initial_data.keys()) - {"id", "type"} + unknown_keys = initial_data - set(self.fields.keys()) + if unknown_keys: + raise ValidationError(f"Invalid fields: {unknown_keys}") + return data diff --git a/api/src/backend/api/v1/urls.py b/api/src/backend/api/v1/urls.py index c212c95e06..eabc98640a 100644 --- a/api/src/backend/api/v1/urls.py +++ b/api/src/backend/api/v1/urls.py @@ -1,26 +1,27 @@ -from django.urls import path, include +from django.urls import include, path from drf_spectacular.views import SpectacularRedocView from rest_framework_nested import routers from api.v1.views import ( + ComplianceOverviewViewSet, CustomTokenObtainView, CustomTokenRefreshView, - SchemaView, - UserViewSet, - TenantViewSet, - TenantMembersViewSet, - MembershipViewSet, - ProviderViewSet, - ScanViewSet, - TaskViewSet, - ResourceViewSet, FindingViewSet, + InvitationAcceptViewSet, + InvitationViewSet, + MembershipViewSet, + OverviewViewSet, ProviderGroupViewSet, ProviderSecretViewSet, - InvitationViewSet, - InvitationAcceptViewSet, - OverviewViewSet, - ComplianceOverviewViewSet, + ProviderViewSet, + ResourceViewSet, + ScanViewSet, + ScheduleViewSet, + SchemaView, + TaskViewSet, + TenantMembersViewSet, + TenantViewSet, + UserViewSet, ) router = routers.DefaultRouter(trailing_slash=False) @@ -37,6 +38,7 @@ router.register( r"compliance-overviews", ComplianceOverviewViewSet, basename="complianceoverview" ) router.register(r"overviews", OverviewViewSet, basename="overview") +router.register(r"schedules", ScheduleViewSet, basename="schedule") tenants_router = routers.NestedSimpleRouter(router, r"tenants", lookup="tenant") tenants_router.register( diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index d9b36af945..69e0a375f4 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -100,6 +100,7 @@ from api.v1.serializers import ( ScanCreateSerializer, ScanSerializer, ScanUpdateSerializer, + ScheduleDailyCreateSerializer, TaskSerializer, TenantSerializer, TokenRefreshSerializer, @@ -725,14 +726,6 @@ class ProviderViewSet(BaseRLSViewSet): }, ) - def create(self, request, *args, **kwargs): - serializer = self.get_serializer(data=request.data) - serializer.is_valid(raise_exception=True) - provider = serializer.save() - # Schedule a daily scan for the new provider - schedule_provider_scan(provider) - return Response(data=serializer.data, status=status.HTTP_201_CREATED) - @extend_schema_view( list=extend_schema( @@ -1549,3 +1542,57 @@ class OverviewViewSet(BaseRLSViewSet): serializer = OverviewSeveritySerializer(severity_data) return Response(serializer.data, status=status.HTTP_200_OK) + + +@extend_schema(tags=["Schedule"]) +@extend_schema_view( + daily=extend_schema( + summary="Create a daily schedule scan for a given provider", + description="Schedules a daily scan for the specified provider. This endpoint creates a periodic task " + "that will execute a scan every 24 hours.", + request=ScheduleDailyCreateSerializer, + responses={202: OpenApiResponse(response=TaskSerializer)}, + ) +) +class ScheduleViewSet(BaseRLSViewSet): + # TODO: change to Schedule when implemented + queryset = Task.objects.none() + http_method_names = ["post"] + + def get_queryset(self): + return super().get_queryset() + + def get_serializer_class(self): + if self.action == "daily": + if hasattr(self, "response_serializer_class"): + return self.response_serializer_class + return ScheduleDailyCreateSerializer + return super().get_serializer_class() + + @extend_schema(exclude=True) + def create(self, request, *args, **kwargs): + raise MethodNotAllowed(method="POST") + + @action(detail=False, methods=["post"], url_name="daily") + def daily(self, request): + serializer = self.get_serializer(data=request.data) + serializer.is_valid(raise_exception=True) + provider_id = serializer.validated_data["provider_id"] + + provider_instance = get_object_or_404(Provider, pk=provider_id) + with transaction.atomic(): + task = schedule_provider_scan(provider_instance) + + prowler_task = Task.objects.get(id=task.id) + self.response_serializer_class = TaskSerializer + output_serializer = self.get_serializer(prowler_task) + + return Response( + data=output_serializer.data, + status=status.HTTP_202_ACCEPTED, + headers={ + "Content-Location": reverse( + "task-detail", kwargs={"pk": prowler_task.id} + ) + }, + ) diff --git a/api/src/backend/tasks/beat.py b/api/src/backend/tasks/beat.py index b53fb700c9..d3446bf85a 100644 --- a/api/src/backend/tasks/beat.py +++ b/api/src/backend/tasks/beat.py @@ -1,7 +1,8 @@ import json -from django.utils import timezone -from django_celery_beat.models import PeriodicTask, IntervalSchedule +from django_celery_beat.models import IntervalSchedule, PeriodicTask +from rest_framework_json_api.serializers import ValidationError +from tasks.tasks import perform_scheduled_scan_task from api.models import Provider @@ -16,7 +17,7 @@ def schedule_provider_scan(provider_instance: Provider): task_name = f"scan-perform-scheduled-{provider_instance.id}" # Schedule the task - PeriodicTask.objects.create( + _, created = PeriodicTask.objects.get_or_create( interval=schedule, name=task_name, task="scan-perform-scheduled", @@ -26,6 +27,23 @@ def schedule_provider_scan(provider_instance: Provider): "provider_id": str(provider_instance.id), } ), - start_time=provider_instance.inserted_at + timezone.timedelta(hours=24), one_off=False, ) + if not created: + raise ValidationError( + [ + { + "detail": "There is already a scheduled scan for this provider.", + "status": 400, + "source": {"pointer": "/data/attributes/provider_id"}, + "code": "invalid", + } + ] + ) + + return perform_scheduled_scan_task.apply_async( + kwargs={ + "tenant_id": str(provider_instance.tenant_id), + "provider_id": str(provider_instance.id), + }, + ) diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index dc5f6ad158..79bc005ae8 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -1,5 +1,8 @@ +from datetime import datetime, timedelta, timezone + from celery import shared_task from config.celery import RLSTask +from django_celery_beat.models import PeriodicTask from tasks.jobs.connection import check_provider_connection from tasks.jobs.deletion import delete_provider from tasks.jobs.scan import aggregate_findings, perform_prowler_scan @@ -98,20 +101,34 @@ def perform_scheduled_scan_task(self, tenant_id: str, provider_id: str): with tenant_transaction(tenant_id): provider_instance = Provider.objects.get(pk=provider_id) + periodic_task_instance = PeriodicTask.objects.get( + name=f"scan-perform-scheduled-{provider_id}" + ) + next_scan_date = datetime.combine( + datetime.now(timezone.utc), periodic_task_instance.date_changed.time() + ) + timedelta(hours=24) scan_instance = Scan.objects.create( tenant_id=tenant_id, name="Daily scheduled scan", provider=provider_instance, trigger=Scan.TriggerChoices.SCHEDULED, + next_scan_at=next_scan_date, task_id=task_id, ) - return perform_prowler_scan( + result = perform_prowler_scan( tenant_id=tenant_id, scan_id=str(scan_instance.id), provider_id=provider_id, ) + perform_scan_summary_task.apply_async( + kwargs={ + "tenant_id": tenant_id, + "scan_id": str(scan_instance.id), + } + ), + return result @shared_task(name="scan-summary") diff --git a/api/src/backend/tasks/tests/test_beat.py b/api/src/backend/tasks/tests/test_beat.py new file mode 100644 index 0000000000..78b5acb039 --- /dev/null +++ b/api/src/backend/tasks/tests/test_beat.py @@ -0,0 +1,53 @@ +import json +from unittest.mock import patch + +import pytest +from django_celery_beat.models import IntervalSchedule, PeriodicTask +from rest_framework_json_api.serializers import ValidationError +from tasks.beat import schedule_provider_scan + + +@pytest.mark.django_db +class TestScheduleProviderScan: + def test_schedule_provider_scan_success(self, providers_fixture): + provider_instance, *_ = providers_fixture + + with patch( + "tasks.tasks.perform_scheduled_scan_task.apply_async" + ) as mock_apply_async: + result = schedule_provider_scan(provider_instance) + + assert result is not None + + mock_apply_async.assert_called_once_with( + kwargs={ + "tenant_id": str(provider_instance.tenant_id), + "provider_id": str(provider_instance.id), + }, + ) + + task_name = f"scan-perform-scheduled-{provider_instance.id}" + periodic_task = PeriodicTask.objects.get(name=task_name) + assert periodic_task is not None + assert periodic_task.interval.every == 24 + assert periodic_task.interval.period == IntervalSchedule.HOURS + assert periodic_task.task == "scan-perform-scheduled" + assert json.loads(periodic_task.kwargs) == { + "tenant_id": str(provider_instance.tenant_id), + "provider_id": str(provider_instance.id), + } + + def test_schedule_provider_scan_already_exists(self, providers_fixture): + provider_instance, *_ = providers_fixture + + # First, schedule the scan + with patch("tasks.tasks.perform_scheduled_scan_task.apply_async"): + schedule_provider_scan(provider_instance) + + # Now, try scheduling again, should raise ValidationError + with pytest.raises(ValidationError) as exc_info: + schedule_provider_scan(provider_instance) + + assert "There is already a scheduled scan for this provider." in str( + exc_info.value + )