diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index 002af353da..99fc46892a 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -31,6 +31,7 @@ from api.models import ( UserRoleRelationship, ) from api.rls import Tenant +from prowler.config.config import get_available_compliance_frameworks TODAY = str(datetime.today().date()) @@ -2277,7 +2278,8 @@ class TestScanViewSet: scan.save() monkeypatch.setattr( - "api.v1.views.env", type("env", (), {"str": lambda self, key: bucket})() + "api.v1.views.env", + type("env", (), {"str": lambda self, *args, **kwargs: "test-bucket"})(), ) class FakeS3Client: @@ -2346,6 +2348,165 @@ class TestScanViewSet: assert content_disposition.startswith('attachment; filename="') assert f'filename="{file_path.name}"' in content_disposition + def test_compliance_invalid_framework(self, authenticated_client, scans_fixture): + scan = scans_fixture[0] + scan.state = StateChoices.COMPLETED + scan.output_location = "dummy" + scan.save() + + url = reverse("scan-compliance", kwargs={"pk": scan.id, "name": "invalid"}) + resp = authenticated_client.get(url) + assert resp.status_code == status.HTTP_404_NOT_FOUND + assert resp.json()["errors"]["detail"] == "Compliance 'invalid' not found." + + def test_compliance_executing( + self, authenticated_client, scans_fixture, monkeypatch + ): + scan = scans_fixture[0] + scan.state = StateChoices.EXECUTING + scan.save() + task = Task.objects.create(tenant_id=scan.tenant_id) + scan.task = task + scan.save() + dummy = {"id": str(task.id), "state": StateChoices.EXECUTING} + + monkeypatch.setattr( + "api.v1.views.TaskSerializer", + lambda *args, **kwargs: type("S", (), {"data": dummy}), + ) + + framework = get_available_compliance_frameworks(scan.provider.provider)[0] + url = reverse("scan-compliance", kwargs={"pk": scan.id, "name": framework}) + resp = authenticated_client.get(url) + assert resp.status_code == status.HTTP_202_ACCEPTED + assert "Content-Location" in resp + assert dummy["id"] in resp["Content-Location"] + + def test_compliance_no_output(self, authenticated_client, scans_fixture): + scan = scans_fixture[0] + scan.state = StateChoices.COMPLETED + scan.output_location = "" + scan.save() + + framework = get_available_compliance_frameworks(scan.provider.provider)[0] + url = reverse("scan-compliance", kwargs={"pk": scan.id, "name": framework}) + resp = authenticated_client.get(url) + assert resp.status_code == status.HTTP_404_NOT_FOUND + assert resp.json()["errors"]["detail"] == "The scan has no reports." + + def test_compliance_s3_no_credentials( + self, authenticated_client, scans_fixture, monkeypatch + ): + scan = scans_fixture[0] + bucket = "bucket" + key = "file.zip" + scan.output_location = f"s3://{bucket}/{key}" + scan.state = StateChoices.COMPLETED + scan.save() + + monkeypatch.setattr( + "api.v1.views.get_s3_client", + lambda: (_ for _ in ()).throw(NoCredentialsError()), + ) + + framework = get_available_compliance_frameworks(scan.provider.provider)[0] + url = reverse("scan-compliance", kwargs={"pk": scan.id, "name": framework}) + resp = authenticated_client.get(url) + assert resp.status_code == status.HTTP_403_FORBIDDEN + assert resp.json()["errors"]["detail"] == "There is a problem with credentials." + + def test_compliance_s3_success( + self, authenticated_client, scans_fixture, monkeypatch + ): + scan = scans_fixture[0] + bucket = "bucket" + prefix = "path/scan.zip" + scan.output_location = f"s3://{bucket}/{prefix}" + scan.state = StateChoices.COMPLETED + scan.save() + + monkeypatch.setattr( + "api.v1.views.env", + type("env", (), {"str": lambda self, *args, **kwargs: "test-bucket"})(), + ) + + match_key = "path/compliance/mitre_attack_aws.csv" + + class FakeS3Client: + def list_objects_v2(self, Bucket, Prefix): + return {"Contents": [{"Key": match_key}]} + + def get_object(self, Bucket, Key): + return {"Body": io.BytesIO(b"ignored")} + + monkeypatch.setattr("api.v1.views.get_s3_client", lambda: FakeS3Client()) + + framework = match_key.split("/")[-1].split(".")[0] + url = reverse("scan-compliance", kwargs={"pk": scan.id, "name": framework}) + resp = authenticated_client.get(url) + assert resp.status_code == status.HTTP_200_OK + cd = resp["Content-Disposition"] + assert cd.startswith('attachment; filename="') + assert cd.endswith('filename="mitre_attack_aws.csv"') + + def test_compliance_s3_not_found( + self, authenticated_client, scans_fixture, monkeypatch + ): + scan = scans_fixture[0] + bucket = "bucket" + scan.output_location = f"s3://{bucket}/x/scan.zip" + scan.state = StateChoices.COMPLETED + scan.save() + + monkeypatch.setattr( + "api.v1.views.env", + type("env", (), {"str": lambda self, *args, **kwargs: "test-bucket"})(), + ) + + class FakeS3Client: + def list_objects_v2(self, Bucket, Prefix): + return {"Contents": []} + + def get_object(self, Bucket, Key): + return {"Body": io.BytesIO(b"ignored")} + + monkeypatch.setattr("api.v1.views.get_s3_client", lambda: FakeS3Client()) + + url = reverse("scan-compliance", kwargs={"pk": scan.id, "name": "cis_1.4_aws"}) + resp = authenticated_client.get(url) + assert resp.status_code == status.HTTP_404_NOT_FOUND + assert ( + resp.json()["errors"]["detail"] + == "No compliance file found for name 'cis_1.4_aws'." + ) + + def test_compliance_local_file( + self, authenticated_client, scans_fixture, tmp_path, monkeypatch + ): + scan = scans_fixture[0] + scan.state = StateChoices.COMPLETED + base = tmp_path / "reports" + comp_dir = base / "compliance" + comp_dir.mkdir(parents=True) + fname = comp_dir / "scan_cis.csv" + fname.write_bytes(b"ignored") + + scan.output_location = str(base / "scan.zip") + scan.save() + + monkeypatch.setattr( + glob, + "glob", + lambda p: [str(fname)] if p.endswith("*_cis_1.4_aws.csv") else [], + ) + + url = reverse("scan-compliance", kwargs={"pk": scan.id, "name": "cis_1.4_aws"}) + resp = authenticated_client.get(url) + assert resp.status_code == status.HTTP_200_OK + cd = resp["Content-Disposition"] + assert cd.startswith('attachment; filename="') + assert cd.endswith(f'filename="{fname.name}"') + @pytest.mark.django_db class TestTaskViewSet: diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index 105cc8dea1..e88beea750 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -1,7 +1,6 @@ import glob import os -import sentry_sdk from allauth.socialaccount.providers.github.views import GitHubOAuth2Adapter from allauth.socialaccount.providers.google.views import GoogleOAuth2Adapter from botocore.exceptions import ClientError, NoCredentialsError, ParamValidationError @@ -1223,8 +1222,12 @@ class ScanViewSet(BaseRLSViewSet): elif self.action == "partial_update": return ScanUpdateSerializer elif self.action == "report": + if hasattr(self, "response_serializer_class"): + return self.response_serializer_class return ScanReportSerializer elif self.action == "compliance": + if hasattr(self, "response_serializer_class"): + return self.response_serializer_class return ScanComplianceReportSerializer return super().get_serializer_class() @@ -1243,70 +1246,86 @@ class ScanViewSet(BaseRLSViewSet): ) return Response(data=read_serializer.data, status=status.HTTP_200_OK) - @action(detail=True, methods=["get"], url_name="report") - def report(self, request, pk=None): - scan_instance = self.get_object() + def _get_running_task_response(self, scan_instance): + """ + If the scan or its report-generation task is still executing, + return a 202 Response with the task payload and Content-Location. + """ + task = None - if scan_instance.state == StateChoices.EXECUTING: - # If the scan is still running, return the task - prowler_task = Task.objects.get(id=scan_instance.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": output_serializer.data["id"]} - ) - }, - ) - - try: - output_celery_task = Task.objects.get( - task_runner_task__task_name="scan-report", - task_runner_task__task_args__contains=pk, - ) - self.response_serializer_class = TaskSerializer - output_serializer = self.get_serializer(output_celery_task) - if output_serializer.data["state"] == StateChoices.EXECUTING: - # If the task is still running, return the task - return Response( - data=output_serializer.data, - status=status.HTTP_202_ACCEPTED, - headers={ - "Content-Location": reverse( - "task-detail", kwargs={"pk": output_serializer.data["id"]} - ) - }, - ) - except Task.DoesNotExist: - # If the task does not exist, it means that the task is removed from the database - pass - - output_location = scan_instance.output_location - if not output_location: - return Response( - {"detail": "The scan has no reports."}, - status=status.HTTP_404_NOT_FOUND, - ) - - if scan_instance.output_location.startswith("s3://"): + if scan_instance.state == StateChoices.EXECUTING and scan_instance.task: + task = scan_instance.task + else: try: - s3_client = get_s3_client() + task = Task.objects.get( + task_runner_task__task_name="scan-report", + task_runner_task__task_args__contains=str(scan_instance.id), + ) + except Task.DoesNotExist: + return None + + self.response_serializer_class = TaskSerializer + serializer = self.get_serializer(task) + + if serializer.data.get("state") != StateChoices.EXECUTING: + return None + + return Response( + data=serializer.data, + status=status.HTTP_202_ACCEPTED, + headers={ + "Content-Location": reverse( + "task-detail", kwargs={"pk": serializer.data["id"]} + ) + }, + ) + + def _load_file(self, path_pattern, s3=False, bucket=None, list_objects=False): + """ + Load binary content and filename. + If s3=True and list_objects=False: treat path_pattern as exact key. + If s3=True and list_objects=True: list by prefix, then pick first matching key. + Else: treat path_pattern as glob pattern on local FS. + Returns (content, filename) or Response on error. + """ + if s3: + try: + client = get_s3_client() except (ClientError, NoCredentialsError, ParamValidationError): return Response( {"detail": "There is a problem with credentials."}, status=status.HTTP_403_FORBIDDEN, ) - - bucket_name = env.str("DJANGO_OUTPUT_S3_AWS_OUTPUT_BUCKET") - key = output_location[len(f"s3://{bucket_name}/") :] + if list_objects: + # list keys under prefix then match suffix + prefix = os.path.dirname(path_pattern) + suffix = os.path.basename(path_pattern) + try: + resp = client.list_objects_v2(Bucket=bucket, Prefix=prefix) + except ClientError: + return Response( + {"detail": "Failed to list compliance files in S3."}, + status=status.HTTP_500_INTERNAL_SERVER_ERROR, + ) + contents = resp.get("Contents", []) + keys = [obj["Key"] for obj in contents if obj["Key"].endswith(suffix)] + if not keys: + return Response( + { + "detail": f"No compliance file found for name '{os.path.splitext(suffix)[0]}'." + }, + status=status.HTTP_404_NOT_FOUND, + ) + # path_pattern here is prefix, but in compliance we build correct suffix check before + key = keys[0] + else: + # path_pattern is exact key + key = path_pattern try: - s3_object = s3_client.get_object(Bucket=bucket_name, Key=key) + s3_obj = client.get_object(Bucket=bucket, Key=key) except ClientError as e: - error_code = e.response.get("Error", {}).get("Code") - if error_code == "NoSuchKey": + code = e.response.get("Error", {}).get("Code") + if code == "NoSuchKey": return Response( {"detail": "The scan has no reports."}, status=status.HTTP_404_NOT_FOUND, @@ -1315,28 +1334,56 @@ class ScanViewSet(BaseRLSViewSet): {"detail": "There is a problem with credentials."}, status=status.HTTP_403_FORBIDDEN, ) - file_content = s3_object["Body"].read() - filename = os.path.basename(output_location.split("/")[-1]) + content = s3_obj["Body"].read() + filename = os.path.basename(key) else: - zip_files = glob.glob(output_location) - try: - file_path = zip_files[0] - except IndexError as e: - sentry_sdk.capture_exception(e) + files = glob.glob(path_pattern) + if not files: return Response( {"detail": "The scan has no reports."}, status=status.HTTP_404_NOT_FOUND, ) - with open(file_path, "rb") as f: - file_content = f.read() - filename = os.path.basename(file_path) + filepath = files[0] + with open(filepath, "rb") as f: + content = f.read() + filename = os.path.basename(filepath) - response = HttpResponse( - file_content, content_type="application/x-zip-compressed" - ) + return content, filename + + def _serve_file(self, content, filename, content_type): + response = HttpResponse(content, content_type=content_type) response["Content-Disposition"] = f'attachment; filename="{filename}"' + return response + @action(detail=True, methods=["get"], url_name="report") + def report(self, request, pk=None): + scan = self.get_object() + # Check for executing tasks + running_resp = self._get_running_task_response(scan) + if running_resp: + return running_resp + + if not scan.output_location: + return Response( + {"detail": "The scan has no reports."}, status=status.HTTP_404_NOT_FOUND + ) + + if scan.output_location.startswith("s3://"): + bucket = env.str("DJANGO_OUTPUT_S3_AWS_OUTPUT_BUCKET", "") + key_prefix = scan.output_location.removeprefix(f"s3://{bucket}/") + loader = self._load_file( + key_prefix, s3=True, bucket=bucket, list_objects=False + ) + else: + loader = self._load_file(scan.output_location, s3=False) + + if isinstance(loader, Response): + return loader + + content, filename = loader + return self._serve_file(content, filename, "application/x-zip-compressed") + @action( detail=True, methods=["get"], @@ -1344,81 +1391,39 @@ class ScanViewSet(BaseRLSViewSet): url_name="compliance", ) def compliance(self, request, pk=None, name=None): - scan_instance = self.get_object() - if name not in get_available_compliance_frameworks( - scan_instance.provider.provider - ): + scan = self.get_object() + if name not in get_available_compliance_frameworks(scan.provider.provider): return Response( {"detail": f"Compliance '{name}' not found."}, status=status.HTTP_404_NOT_FOUND, ) - if scan_instance.output_location.startswith("s3://"): - try: - s3 = get_s3_client() - bucket = env.str("DJANGO_OUTPUT_S3_AWS_OUTPUT_BUCKET") - except (ClientError, NoCredentialsError, ParamValidationError): - return Response( - {"detail": "Problem with S3 credentials."}, - status=status.HTTP_403_FORBIDDEN, - ) + running_resp = self._get_running_task_response(scan) + if running_resp: + return running_resp - key_prefix = scan_instance.output_location[len(f"s3://{bucket}/") :] - compliance_prefix = f"{os.path.dirname(key_prefix)}/compliance/" - try: - resp = s3.list_objects_v2(Bucket=bucket, Prefix=compliance_prefix) - except ClientError: - return Response( - {"detail": "Failed to list compliance files in S3."}, - status=status.HTTP_500_INTERNAL_SERVER_ERROR, - ) - - contents = resp.get("Contents", []) - matches = [ - obj["Key"] for obj in contents if obj["Key"].endswith(f"_{name}.csv") - ] - if not matches: - return Response( - {"detail": f"No compliance file found for '{name}' in S3."}, - status=status.HTTP_404_NOT_FOUND, - ) - - file_path = matches[0] - try: - s3_obj = s3.get_object(Bucket=bucket, Key=file_path) - except ClientError as e: - error_code = e.response.get("Error", {}).get("Code") - if error_code == "NoSuchKey": - return Response( - {"detail": "The compliance file does not exist in S3."}, - status=status.HTTP_404_NOT_FOUND, - ) - - return Response( - {"detail": "Error downloading the compliance file from S3."}, - status=status.HTTP_500_INTERNAL_SERVER_ERROR, - ) - file_content = s3_obj["Body"].read() - else: - output_location = os.path.join( - os.path.dirname(scan_instance.output_location), "compliance" + if not scan.output_location: + return Response( + {"detail": "The scan has no reports."}, status=status.HTTP_404_NOT_FOUND ) - search_pattern = os.path.join(output_location, f"*_{name}.csv") - matching_files = glob.glob(search_pattern) - if not matching_files: - return Response( - {"detail": f"No compliance file found for name '{name}'."}, - status=status.HTTP_404_NOT_FOUND, - ) - file_path = matching_files[0] - with open(file_path, "rb") as f: - file_content = f.read() - filename = os.path.basename(file_path) - response = HttpResponse(file_content, content_type="text/csv") - response["Content-Disposition"] = f'attachment; filename="{filename}"' + if scan.output_location.startswith("s3://"): + bucket = env.str("DJANGO_OUTPUT_S3_AWS_OUTPUT_BUCKET", "") + key_prefix = scan.output_location.removeprefix(f"s3://{bucket}/") + prefix = os.path.join( + os.path.dirname(key_prefix), "compliance", f"{name}.csv" + ) + loader = self._load_file(prefix, s3=True, bucket=bucket, list_objects=True) + else: + base = os.path.dirname(scan.output_location) + pattern = os.path.join(base, "compliance", f"*_{name}.csv") + loader = self._load_file(pattern, s3=False) - return response + if isinstance(loader, Response): + return loader + + content, filename = loader + return self._serve_file(content, filename, "text/csv") def create(self, request, *args, **kwargs): input_serializer = self.get_serializer(data=request.data)