diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index dfdda6391d..fa732edba8 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -1,12 +1,9 @@ import glob -import json import logging import os -import tempfile from datetime import datetime, timedelta, timezone from urllib.parse import urljoin -import boto3 import sentry_sdk from allauth.socialaccount.models import SocialAccount, SocialApp from allauth.socialaccount.providers.github.views import GitHubOAuth2Adapter @@ -14,6 +11,7 @@ from allauth.socialaccount.providers.google.views import GoogleOAuth2Adapter from allauth.socialaccount.providers.saml.views import FinishACSView, LoginView from botocore.exceptions import ClientError, NoCredentialsError, ParamValidationError from celery.result import AsyncResult +from config.custom_logging import BackendLogger from config.env import env from config.settings.social_login import ( GITHUB_OAUTH_CALLBACK_URL, @@ -24,7 +22,7 @@ from django.conf import settings as django_settings from django.contrib.postgres.aggregates import ArrayAgg from django.contrib.postgres.search import SearchQuery from django.db import transaction -from django.db.models import Count, Exists, F, OuterRef, Prefetch, Q, Sum +from django.db.models import Count, F, Prefetch, Q, Sum from django.db.models.functions import Coalesce from django.http import HttpResponse from django.shortcuts import redirect @@ -52,8 +50,7 @@ from rest_framework.exceptions import ( ValidationError, ) from rest_framework.generics import GenericAPIView, get_object_or_404 -from rest_framework.permissions import SAFE_METHODS, IsAuthenticated -from rest_framework.viewsets import ViewSet +from rest_framework.permissions import SAFE_METHODS from rest_framework_json_api.views import RelationshipView, Response from rest_framework_simplejwt.exceptions import InvalidToken, TokenError from tasks.beat import schedule_provider_scan @@ -64,7 +61,6 @@ from tasks.tasks import ( check_provider_connection_task, delete_provider_task, delete_tenant_task, - generate_threatscore_report_task, perform_scan_task, ) @@ -82,6 +78,7 @@ from api.filters import ( IntegrationFilter, InvitationFilter, LatestFindingFilter, + LatestResourceFilter, MembershipFilter, ProcessorFilter, ProviderFilter, @@ -97,7 +94,6 @@ from api.filters import ( UserFilter, ) from api.models import ( - ComplianceOverview, ComplianceRequirementOverview, Finding, Integration, @@ -112,6 +108,7 @@ from api.models import ( Resource, ResourceFindingMapping, ResourceScanSummary, + ResourceTag, Role, RoleProviderGroupRelationship, SAMLConfiguration, @@ -171,6 +168,7 @@ from api.v1.serializers import ( ProviderSecretUpdateSerializer, ProviderSerializer, ProviderUpdateSerializer, + ResourceMetadataSerializer, ResourceSerializer, RoleCreateSerializer, RoleProviderGroupRelationshipSerializer, @@ -196,6 +194,8 @@ from api.v1.serializers import ( UserUpdateSerializer, ) +logger = logging.getLogger(BackendLogger.API) + CACHE_DECORATOR = cache_control( max_age=django_settings.CACHE_MAX_AGE, stale_while_revalidate=django_settings.CACHE_STALE_WHILE_REVALIDATE, @@ -292,7 +292,7 @@ class SchemaView(SpectacularAPIView): def get(self, request, *args, **kwargs): spectacular_settings.TITLE = "Prowler API" - spectacular_settings.VERSION = "1.9.0" + spectacular_settings.VERSION = "1.10.0" spectacular_settings.DESCRIPTION = ( "Prowler API specification.\n\nThis file is auto-generated." ) @@ -565,10 +565,25 @@ class SAMLConfigurationViewSet(BaseRLSViewSet): class TenantFinishACSView(FinishACSView): + def _rollback_saml_user(self, request): + """Helper function to rollback SAML user if it was just created and validation fails""" + saml_user_id = request.session.get("saml_user_created") + if saml_user_id: + User.objects.using(MainRouter.admin_db).filter(id=saml_user_id).delete() + request.session.pop("saml_user_created", None) + def dispatch(self, request, organization_slug): - super().dispatch(request, organization_slug) + try: + super().dispatch(request, organization_slug) + except Exception as e: + logger.error(f"SAML dispatch failed: {e}") + self._rollback_saml_user(request) + callback_url = env.str("AUTH_URL") + return redirect(f"{callback_url}?sso_saml_failed=true") + user = getattr(request, "user", None) if not user or not user.is_authenticated: + self._rollback_saml_user(request) callback_url = env.str("AUTH_URL") return redirect(f"{callback_url}?sso_saml_failed=true") @@ -591,7 +606,9 @@ class TenantFinishACSView(FinishACSView): SocialApp.DoesNotExist, SocialAccount.DoesNotExist, User.DoesNotExist, - ): + ) as e: + logger.error(f"SAML user is not authenticated: {e}") + self._rollback_saml_user(request) callback_url = env.str("AUTH_URL") return redirect(f"{callback_url}?sso_saml_failed=true") @@ -665,6 +682,7 @@ class TenantFinishACSView(FinishACSView): ) callback_url = env.str("SAML_SSO_CALLBACK_URL") redirect_url = f"{callback_url}?id={saml_token.id}" + request.session.pop("saml_user_created", None) return redirect(redirect_url) @@ -1478,15 +1496,6 @@ class ProviderViewSet(BaseRLSViewSet): }, request=None, ), - threatscore_report=extend_schema( - tags=["Scan"], - summary="Generate ThreatScore report", - description="Generate a ThreatScore report for a specific scan", - request=None, - responses={ - 202: OpenApiResponse(description="ThreatScore report generation started") - }, - ), ) @method_decorator(CACHE_DECORATOR, name="list") @method_decorator(CACHE_DECORATOR, name="retrieve") @@ -1543,8 +1552,6 @@ class ScanViewSet(BaseRLSViewSet): if hasattr(self, "response_serializer_class"): return self.response_serializer_class return ScanComplianceReportSerializer - elif self.action == "threatscore_report": - return None return super().get_serializer_class() def partial_update(self, request, *args, **kwargs): @@ -1809,39 +1816,6 @@ class ScanViewSet(BaseRLSViewSet): }, ) - @action(detail=True, methods=["post"], url_path="threatscore-report") - def threatscore_report(self, request, pk=None): - scan = self.get_object() - tenant_id = self.request.tenant_id - provider_id = str(scan.provider_id) - scan_id = str(scan.id) - compliance_id = request.data.get("compliance_id") - output_path = request.data.get( - "output_path", f"/tmp/threatscore_report_{scan_id}.pdf" - ) - only_failed = request.data.get("only_failed", True) - min_risk_level = request.data.get("min_risk_level", 4) - if not compliance_id: - return Response( - {"detail": "compliance_id is required"}, - status=status.HTTP_400_BAD_REQUEST, - ) - generate_threatscore_report_task.apply_async( - kwargs={ - "tenant_id": tenant_id, - "scan_id": scan_id, - "provider_id": provider_id, - "compliance_id": compliance_id, - "output_path": output_path, - "only_failed": only_failed, - "min_risk_level": min_risk_level, - } - ) - return Response( - {"detail": "ThreatScore report generation started."}, - status=status.HTTP_202_ACCEPTED, - ) - @extend_schema_view( list=extend_schema( @@ -1911,6 +1885,14 @@ class TaskViewSet(BaseRLSViewSet): summary="List all resources", description="Retrieve a list of all resources with options for filtering by various criteria. Resources are " "objects that are discovered by Prowler. They can be anything from a single host to a whole VPC.", + parameters=[ + OpenApiParameter( + name="filter[updated_at]", + description="At least one of the variations of the `filter[updated_at]` filter must be provided.", + required=True, + type=OpenApiTypes.DATE, + ) + ], ), retrieve=extend_schema( tags=["Resource"], @@ -1918,15 +1900,43 @@ class TaskViewSet(BaseRLSViewSet): description="Fetch detailed information about a specific resource by their ID. A Resource is an object that " "is discovered by Prowler. It can be anything from a single host to a whole VPC.", ), + metadata=extend_schema( + tags=["Resource"], + summary="Retrieve metadata values from resources", + description="Fetch unique metadata values from a set of resources. This is useful for dynamic filtering.", + parameters=[ + OpenApiParameter( + name="filter[updated_at]", + description="At least one of the variations of the `filter[updated_at]` filter must be provided.", + required=True, + type=OpenApiTypes.DATE, + ) + ], + filters=True, + ), + latest=extend_schema( + tags=["Resource"], + summary="List the latest resources", + description="Retrieve a list of the latest resources from the latest scans for each provider with options for " + "filtering by various criteria.", + filters=True, + ), + metadata_latest=extend_schema( + tags=["Resource"], + summary="Retrieve metadata values from the latest resources", + description="Fetch unique metadata values from a set of resources from the latest scans for each provider. " + "This is useful for dynamic filtering.", + filters=True, + ), ) @method_decorator(CACHE_DECORATOR, name="list") @method_decorator(CACHE_DECORATOR, name="retrieve") -class ResourceViewSet(BaseRLSViewSet): - queryset = Resource.objects.all() +class ResourceViewSet(PaginateByPkMixin, BaseRLSViewSet): + queryset = Resource.all_objects.all() serializer_class = ResourceSerializer http_method_names = ["get"] filterset_class = ResourceFilter - ordering = ["-inserted_at"] + ordering = ["-failed_findings_count", "-updated_at"] ordering_fields = [ "provider_uid", "uid", @@ -1937,6 +1947,14 @@ class ResourceViewSet(BaseRLSViewSet): "inserted_at", "updated_at", ] + prefetch_for_includes = { + "__all__": [], + "provider": [ + Prefetch( + "provider", queryset=Provider.all_objects.select_related("resources") + ) + ], + } # RBAC required permissions (implicit -> MANAGE_PROVIDERS enable unlimited visibility or check the visibility of # the provider through the provider group) required_permissions = [] @@ -1945,41 +1963,257 @@ class ResourceViewSet(BaseRLSViewSet): user_roles = get_role(self.request.user) if user_roles.unlimited_visibility: # User has unlimited visibility, return all scans - queryset = Resource.objects.filter(tenant_id=self.request.tenant_id) + queryset = Resource.all_objects.filter(tenant_id=self.request.tenant_id) else: # User lacks permission, filter providers based on provider groups associated with the role - queryset = Resource.objects.filter( + queryset = Resource.all_objects.filter( tenant_id=self.request.tenant_id, provider__in=get_providers(user_roles) ) search_value = self.request.query_params.get("filter[search]", None) if search_value: - # Django's ORM will build a LEFT JOIN and OUTER JOIN on the "through" table, resulting in duplicates - # The duplicates then require a `distinct` query search_query = SearchQuery( search_value, config="simple", search_type="plain" ) queryset = queryset.filter( - Q(tags__key=search_value) - | Q(tags__value=search_value) - | Q(tags__text_search=search_query) - | Q(tags__key__contains=search_value) - | Q(tags__value__contains=search_value) - | Q(uid=search_value) - | Q(name=search_value) - | Q(region=search_value) - | Q(service=search_value) - | Q(type=search_value) - | Q(text_search=search_query) - | Q(uid__contains=search_value) - | Q(name__contains=search_value) - | Q(region__contains=search_value) - | Q(service__contains=search_value) - | Q(type__contains=search_value) + Q(text_search=search_query) | Q(tags__text_search=search_query) ).distinct() return queryset + def _optimize_tags_loading(self, queryset): + """Optimize tags loading with prefetch_related to avoid N+1 queries""" + # Use prefetch_related to load all tags in a single query + return queryset.prefetch_related( + Prefetch( + "tags", + queryset=ResourceTag.objects.filter( + tenant_id=self.request.tenant_id + ).select_related(), + to_attr="prefetched_tags", + ) + ) + + def get_serializer_class(self): + if self.action in ["metadata", "metadata_latest"]: + return ResourceMetadataSerializer + return super().get_serializer_class() + + def get_filterset_class(self): + if self.action in ["latest", "metadata_latest"]: + return LatestResourceFilter + return ResourceFilter + + def filter_queryset(self, queryset): + # Do not apply filters when retrieving specific resource + if self.action == "retrieve": + return queryset + return super().filter_queryset(queryset) + + def list(self, request, *args, **kwargs): + filtered_queryset = self.filter_queryset(self.get_queryset()) + return self.paginate_by_pk( + request, + filtered_queryset, + manager=Resource.all_objects, + select_related=["provider"], + prefetch_related=["findings"], + ) + + def retrieve(self, request, *args, **kwargs): + queryset = self._optimize_tags_loading(self.get_queryset()) + instance = get_object_or_404(queryset, pk=kwargs.get("pk")) + mapping_ids = list( + ResourceFindingMapping.objects.filter( + resource=instance, tenant_id=request.tenant_id + ).values_list("finding_id", flat=True) + ) + latest_findings = ( + Finding.all_objects.filter(id__in=mapping_ids, tenant_id=request.tenant_id) + .order_by("uid", "-inserted_at") + .distinct("uid") + ) + setattr(instance, "latest_findings", latest_findings) + serializer = self.get_serializer(instance) + return Response(serializer.data, status=status.HTTP_200_OK) + + @action(detail=False, methods=["get"], url_name="latest") + def latest(self, request): + tenant_id = request.tenant_id + filtered_queryset = self.filter_queryset(self.get_queryset()) + + latest_scan_ids = ( + Scan.all_objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) + .order_by("provider_id", "-inserted_at") + .distinct("provider_id") + .values_list("id", flat=True) + ) + filtered_queryset = filtered_queryset.filter( + tenant_id=tenant_id, provider__scan__in=latest_scan_ids + ) + + return self.paginate_by_pk( + request, + filtered_queryset, + manager=Resource.all_objects, + select_related=["provider"], + prefetch_related=["findings"], + ) + + @action(detail=False, methods=["get"], url_name="metadata") + def metadata(self, request): + # Force filter validation + self.filter_queryset(self.get_queryset()) + + tenant_id = request.tenant_id + query_params = request.query_params + + queryset = ResourceScanSummary.objects.filter(tenant_id=tenant_id) + + if scans := query_params.get("filter[scan__in]") or query_params.get( + "filter[scan]" + ): + queryset = queryset.filter(scan_id__in=scans.split(",")) + else: + exact = query_params.get("filter[inserted_at]") + gte = query_params.get("filter[inserted_at__gte]") + lte = query_params.get("filter[inserted_at__lte]") + + date_filters = {} + if exact: + date = parse_date(exact) + datetime_start = datetime.combine( + date, datetime.min.time(), tzinfo=timezone.utc + ) + datetime_end = datetime_start + timedelta(days=1) + date_filters["scan_id__gte"] = uuid7_start( + datetime_to_uuid7(datetime_start) + ) + date_filters["scan_id__lt"] = uuid7_start( + datetime_to_uuid7(datetime_end) + ) + else: + if gte: + date_start = parse_date(gte) + datetime_start = datetime.combine( + date_start, datetime.min.time(), tzinfo=timezone.utc + ) + date_filters["scan_id__gte"] = uuid7_start( + datetime_to_uuid7(datetime_start) + ) + if lte: + date_end = parse_date(lte) + datetime_end = datetime.combine( + date_end + timedelta(days=1), + datetime.min.time(), + tzinfo=timezone.utc, + ) + date_filters["scan_id__lt"] = uuid7_start( + datetime_to_uuid7(datetime_end) + ) + + if date_filters: + queryset = queryset.filter(**date_filters) + + if service_filter := query_params.get("filter[service]") or query_params.get( + "filter[service__in]" + ): + queryset = queryset.filter(service__in=service_filter.split(",")) + if region_filter := query_params.get("filter[region]") or query_params.get( + "filter[region__in]" + ): + queryset = queryset.filter(region__in=region_filter.split(",")) + if resource_type_filter := query_params.get("filter[type]") or query_params.get( + "filter[type__in]" + ): + queryset = queryset.filter( + resource_type__in=resource_type_filter.split(",") + ) + + services = list( + queryset.values_list("service", flat=True).distinct().order_by("service") + ) + regions = list( + queryset.values_list("region", flat=True).distinct().order_by("region") + ) + resource_types = list( + queryset.values_list("resource_type", flat=True) + .exclude(resource_type__isnull=True) + .exclude(resource_type__exact="") + .distinct() + .order_by("resource_type") + ) + + result = { + "services": services, + "regions": regions, + "types": resource_types, + } + + serializer = self.get_serializer(data=result) + serializer.is_valid(raise_exception=True) + return Response(serializer.data) + + @action( + detail=False, + methods=["get"], + url_name="metadata_latest", + url_path="metadata/latest", + ) + def metadata_latest(self, request): + tenant_id = request.tenant_id + query_params = request.query_params + + latest_scans_queryset = ( + Scan.all_objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) + .order_by("provider_id", "-inserted_at") + .distinct("provider_id") + ) + + queryset = ResourceScanSummary.objects.filter( + tenant_id=tenant_id, + scan_id__in=latest_scans_queryset.values_list("id", flat=True), + ) + + if service_filter := query_params.get("filter[service]") or query_params.get( + "filter[service__in]" + ): + queryset = queryset.filter(service__in=service_filter.split(",")) + if region_filter := query_params.get("filter[region]") or query_params.get( + "filter[region__in]" + ): + queryset = queryset.filter(region__in=region_filter.split(",")) + if resource_type_filter := query_params.get("filter[type]") or query_params.get( + "filter[type__in]" + ): + queryset = queryset.filter( + resource_type__in=resource_type_filter.split(",") + ) + + services = list( + queryset.values_list("service", flat=True).distinct().order_by("service") + ) + regions = list( + queryset.values_list("region", flat=True).distinct().order_by("region") + ) + resource_types = list( + queryset.values_list("resource_type", flat=True) + .exclude(resource_type__isnull=True) + .exclude(resource_type__exact="") + .distinct() + .order_by("resource_type") + ) + + result = { + "services": services, + "regions": regions, + "types": resource_types, + } + + serializer = self.get_serializer(data=result) + serializer.is_valid(raise_exception=True) + return Response(serializer.data) + @extend_schema_view( list=extend_schema( @@ -2098,17 +2332,7 @@ class FindingViewSet(PaginateByPkMixin, BaseRLSViewSet): search_value, config="simple", search_type="plain" ) - resource_match = Resource.all_objects.filter( - text_search=search_query, - id__in=ResourceFindingMapping.objects.filter( - resource_id=OuterRef("pk"), - tenant_id=tenant_id, - ).values("resource_id"), - ) - - queryset = queryset.filter( - Q(text_search=search_query) | Q(Exists(resource_match)) - ) + queryset = queryset.filter(text_search=search_query) return queryset @@ -3244,7 +3468,7 @@ class ComplianceOverviewViewSet(BaseRLSViewSet, TaskManagementMixin): ) @method_decorator(CACHE_DECORATOR, name="list") class OverviewViewSet(BaseRLSViewSet): - queryset = ComplianceOverview.objects.all() + queryset = ScanSummary.objects.all() http_method_names = ["get"] ordering = ["-inserted_at"] # RBAC required permissions (implicit -> MANAGE_PROVIDERS enable unlimited visibility or check the visibility of @@ -3255,19 +3479,10 @@ class OverviewViewSet(BaseRLSViewSet): role = get_role(self.request.user) providers = get_providers(role) - def _get_filtered_queryset(model): - if role.unlimited_visibility: - return model.all_objects.filter(tenant_id=self.request.tenant_id) - return model.all_objects.filter( - tenant_id=self.request.tenant_id, scan__provider__in=providers - ) + if not role.unlimited_visibility: + self.allowed_providers = providers - if self.action == "providers": - return _get_filtered_queryset(Finding) - elif self.action in ("findings", "findings_severity", "services"): - return _get_filtered_queryset(ScanSummary) - else: - return super().get_queryset() + return ScanSummary.all_objects.filter(tenant_id=self.request.tenant_id) def get_serializer_class(self): if self.action == "providers": @@ -3300,18 +3515,24 @@ class OverviewViewSet(BaseRLSViewSet): @action(detail=False, methods=["get"], url_name="providers") def providers(self, request): tenant_id = self.request.tenant_id + queryset = self.get_queryset() + provider_filter = ( + {"provider__in": self.allowed_providers} + if hasattr(self, "allowed_providers") + else {} + ) latest_scan_ids = ( - Scan.all_objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) + Scan.all_objects.filter( + tenant_id=tenant_id, state=StateChoices.COMPLETED, **provider_filter + ) .order_by("provider_id", "-inserted_at") .distinct("provider_id") .values_list("id", flat=True) ) findings_aggregated = ( - ScanSummary.all_objects.filter( - tenant_id=tenant_id, scan_id__in=latest_scan_ids - ) + queryset.filter(scan_id__in=latest_scan_ids) .values( "scan__provider_id", provider=F("scan__provider__provider"), @@ -3347,7 +3568,7 @@ class OverviewViewSet(BaseRLSViewSet): ) return Response( - OverviewProviderSerializer(overview, many=True).data, + self.get_serializer(overview, many=True).data, status=status.HTTP_200_OK, ) @@ -3356,9 +3577,16 @@ class OverviewViewSet(BaseRLSViewSet): tenant_id = self.request.tenant_id queryset = self.get_queryset() filtered_queryset = self.filter_queryset(queryset) + provider_filter = ( + {"provider__in": self.allowed_providers} + if hasattr(self, "allowed_providers") + else {} + ) latest_scan_ids = ( - Scan.all_objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) + Scan.all_objects.filter( + tenant_id=tenant_id, state=StateChoices.COMPLETED, **provider_filter + ) .order_by("provider_id", "-inserted_at") .distinct("provider_id") .values_list("id", flat=True) @@ -3395,9 +3623,16 @@ class OverviewViewSet(BaseRLSViewSet): tenant_id = self.request.tenant_id queryset = self.get_queryset() filtered_queryset = self.filter_queryset(queryset) + provider_filter = ( + {"provider__in": self.allowed_providers} + if hasattr(self, "allowed_providers") + else {} + ) latest_scan_ids = ( - Scan.all_objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) + Scan.all_objects.filter( + tenant_id=tenant_id, state=StateChoices.COMPLETED, **provider_filter + ) .order_by("provider_id", "-inserted_at") .distinct("provider_id") .values_list("id", flat=True) @@ -3417,7 +3652,7 @@ class OverviewViewSet(BaseRLSViewSet): for item in severity_counts: severity_data[item["severity"]] = item["count"] - serializer = OverviewSeveritySerializer(severity_data) + serializer = self.get_serializer(severity_data) return Response(serializer.data, status=status.HTTP_200_OK) @action(detail=False, methods=["get"], url_name="services") @@ -3425,9 +3660,16 @@ class OverviewViewSet(BaseRLSViewSet): tenant_id = self.request.tenant_id queryset = self.get_queryset() filtered_queryset = self.filter_queryset(queryset) + provider_filter = ( + {"provider__in": self.allowed_providers} + if hasattr(self, "allowed_providers") + else {} + ) latest_scan_ids = ( - Scan.all_objects.filter(tenant_id=tenant_id, state=StateChoices.COMPLETED) + Scan.all_objects.filter( + tenant_id=tenant_id, state=StateChoices.COMPLETED, **provider_filter + ) .order_by("provider_id", "-inserted_at") .distinct("provider_id") .values_list("id", flat=True) @@ -3445,7 +3687,7 @@ class OverviewViewSet(BaseRLSViewSet): .order_by("service") ) - serializer = OverviewServiceSerializer(services_data, many=True) + serializer = self.get_serializer(services_data, many=True) return Response(serializer.data, status=status.HTTP_200_OK) @@ -3695,174 +3937,4 @@ class ProcessorViewSet(BaseRLSViewSet): return ProcessorCreateSerializer elif self.action == "partial_update": return ProcessorUpdateSerializer - return super().get_serializer_class() - - -class ReportsViewSet(ViewSet): - permission_classes = [IsAuthenticated] - logger = logging.getLogger("prowler.reports") - - def _get_task_and_validate(self, report_id, user): - try: - task = Task.objects.get(id=report_id) - except Task.DoesNotExist: - return None, Response( - {"detail": "Report not found"}, status=status.HTTP_404_NOT_FOUND - ) - - scan_id = None - try: - kwargs = json.loads(task.task_runner_task.task_kwargs.replace("'", '"')) - scan_id = kwargs.get("scan_id") - except Exception: - pass - if scan_id: - try: - scan = Scan.objects.get(id=scan_id) - if scan.tenant_id != user.tenant_id: - return None, Response( - {"detail": "Forbidden"}, status=status.HTTP_403_FORBIDDEN - ) - except Scan.DoesNotExist: - return None, Response( - {"detail": "Scan not found"}, status=status.HTTP_404_NOT_FOUND - ) - return task, None - - @action(detail=False, methods=["post"], url_path="pdf/generate") - def pdf_generate(self, request): - scan_id = request.data.get("scan_id") - compliance_id = request.data.get("compliance_id") - options = request.data.get("options", {}) - only_failed = options.get("only_failed", True) - min_risk_level = options.get("min_risk_level", 4) - if not scan_id or not compliance_id: - return Response( - {"detail": "scan_id and compliance_id are required"}, - status=status.HTTP_400_BAD_REQUEST, - ) - try: - scan = Scan.objects.get(id=scan_id) - if scan.tenant_id != request.user.tenant_id: - return Response( - {"detail": "Forbidden"}, status=status.HTTP_403_FORBIDDEN - ) - except Scan.DoesNotExist: - return Response( - {"detail": "Scan not found"}, status=status.HTTP_404_NOT_FOUND - ) - tenant_id = scan.tenant_id - provider_id = scan.provider_id - output_path = f"/tmp/threatscore_report_{scan_id}.pdf" - celery_task = generate_threatscore_report_task.apply_async( - kwargs={ - "tenant_id": tenant_id, - "scan_id": scan_id, - "provider_id": provider_id, - "compliance_id": compliance_id, - "output_path": output_path, - "only_failed": only_failed, - "min_risk_level": min_risk_level, - } - ) - self.logger.info( - f"User {request.user.id} initiated ThreatScore report for scan {scan_id} (task {celery_task.id})" - ) - return Response( - {"report_id": celery_task.id, "status": "initiated"}, - status=status.HTTP_202_ACCEPTED, - ) - - @action(detail=False, methods=["get"], url_path="pdf/status/(?P[^/]+)") - def pdf_status(self, request, report_id=None): - task, error = self._get_task_and_validate(report_id, request.user) - if error: - return error - state = task.state - output_path = None - try: - kwargs = json.loads(task.task_runner_task.task_kwargs.replace("'", '"')) - output_path = kwargs.get("output_path") - except Exception: - pass - if state == "SUCCESS": - status_str = "ready" - elif state in ("PENDING", "RECEIVED", "STARTED", "RETRY"): - status_str = "processing" - elif state == "FAILURE": - status_str = "failed" - else: - status_str = "initiated" - resp = {"status": status_str} - if status_str == "ready" and output_path: - resp["output_path"] = output_path - return Response(resp, status=status.HTTP_200_OK) - - @action(detail=False, methods=["get"], url_path="pdf/download/(?P[^/]+)") - def pdf_download(self, request, report_id=None): - task, error = self._get_task_and_validate(report_id, request.user) - if error: - return error - if task.state != "SUCCESS": - return Response( - {"detail": "Report not ready"}, status=status.HTTP_202_ACCEPTED - ) - output_path = getattr(task, "result", None) - if not output_path: - try: - kwargs = json.loads(task.task_runner_task.task_kwargs.replace("'", '"')) - output_path = kwargs.get("output_path") - except Exception: - pass - if not output_path: - self.logger.warning( - f"User {request.user.id} tried to download missing report {report_id} (no path)" - ) - return Response( - {"detail": "Report file not found"}, status=status.HTTP_404_NOT_FOUND - ) - if output_path.startswith("s3://"): - import re - - match = re.match(r"s3://([^/]+)/(.+)", output_path) - if not match: - return Response( - {"detail": "Invalid S3 URI"}, - status=status.HTTP_500_INTERNAL_SERVER_ERROR, - ) - bucket, key = match.groups() - s3 = boto3.client("s3") - try: - with tempfile.NamedTemporaryFile(delete=False) as tmp: - s3.download_fileobj(bucket, key, tmp) - tmp_path = tmp.name - with open(tmp_path, "rb") as f: - pdf_data = f.read() - os.unlink(tmp_path) - except Exception as e: - self.logger.warning( - f"User {request.user.id} failed S3 download for report {report_id}: {e}" - ) - return Response( - {"detail": "Report file not found in S3"}, - status=status.HTTP_404_NOT_FOUND, - ) - filename = os.path.basename(key) - else: - if not os.path.exists(output_path): - self.logger.warning( - f"User {request.user.id} tried to download missing report {report_id} at {output_path}" - ) - return Response( - {"detail": "Report file not found"}, - status=status.HTTP_404_NOT_FOUND, - ) - with open(output_path, "rb") as f: - pdf_data = f.read() - filename = os.path.basename(output_path) - self.logger.info( - f"User {request.user.id} downloaded ThreatScore report {report_id} from {output_path}" - ) - response = Response(pdf_data, content_type="application/pdf") - response["Content-Disposition"] = f'attachment; filename="{filename}"' - return response + return super().get_serializer_class() \ No newline at end of file diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index fe0ba11a46..544103a1e3 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -1,9 +1,8 @@ -import os from datetime import datetime, timedelta, timezone from pathlib import Path from shutil import rmtree -from celery import chain, current_task, shared_task +from celery import chain, shared_task from celery.utils.log import get_task_logger from config.celery import RLSTask from config.django.base import DJANGO_FINDINGS_BATCH_SIZE, DJANGO_TMP_OUTPUT_DIRECTORY @@ -34,9 +33,6 @@ from api.v1.serializers import ScanTaskSerializer from prowler.lib.check.compliance_models import Compliance from prowler.lib.outputs.compliance.generic.generic import GenericCompliance from prowler.lib.outputs.finding import Finding as FindingOutput -from util.compliance_report.threatscore_report_generator import ( - generate_threatscore_report, -) logger = get_task_logger(__name__) @@ -423,50 +419,4 @@ def check_lighthouse_connection_task(lighthouse_config_id: str, tenant_id: str = - 'error' (str or None): The error message if the connection failed, otherwise `None`. - 'available_models' (list): List of available models if connection is successful. """ - return check_lighthouse_connection(lighthouse_config_id=lighthouse_config_id) - - -@shared_task( - base=RLSTask, - name="scan-threatscore-report", - queue="scan-reports", -) -@set_tenant(keep_tenant=True) -def generate_threatscore_report_task( - scan_id: str, - compliance_id: str, - output_path: str, - provider_id: str, - tenant_id: str, - only_failed: bool = True, - min_risk_level: int = 4, -): - generate_threatscore_report( - scan_id=scan_id, - compliance_id=compliance_id, - output_path=output_path, - provider_id=provider_id, - tenant_id=tenant_id, - only_failed=only_failed, - min_risk_level=min_risk_level, - ) - - s3_uri = None - if os.path.exists(output_path): - try: - s3_uri = _upload_to_s3(tenant_id, output_path, scan_id) - except Exception as e: - logger.error(f"Error uploading PDF to S3: {e}") - - from api.models import Task - - task_id = current_task.request.id if hasattr(current_task, "request") else None - if task_id: - try: - task = Task.objects.get(id=task_id) - result_path = s3_uri if s3_uri else output_path - task.result = result_path - task.save(update_fields=["result"]) - except Exception as e: - logger.error(f"Error saving PDF location to Task: {e}") - return {"pdf_location": s3_uri if s3_uri else output_path} + return check_lighthouse_connection(lighthouse_config_id=lighthouse_config_id) \ No newline at end of file