feat(api): revert changes

This commit is contained in:
pedrooot
2025-07-15 18:14:14 +02:00
parent 6874d318fa
commit 289de9376f
2 changed files with 357 additions and 335 deletions
+355 -283
View File
@@ -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<report_id>[^/]+)")
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<report_id>[^/]+)")
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()
+2 -52
View File
@@ -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)