From f8c840f283eb4973d2d8798a21d17d0175812fa4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adri=C3=A1n=20Jes=C3=BAs=20Pe=C3=B1a=20Rodr=C3=ADguez?= Date: Wed, 14 May 2025 10:02:41 +0200 Subject: [PATCH] fix: ensure proper folder creation (#7729) --- api/src/backend/api/tests/test_views.py | 94 ++++++++++++---------- api/src/backend/tasks/tests/test_export.py | 39 +++++---- 2 files changed, 72 insertions(+), 61 deletions(-) diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index 91de5775c9..0608c4a378 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -2,7 +2,9 @@ import glob import io import json import os +import tempfile from datetime import datetime, timedelta, timezone +from pathlib import Path from unittest.mock import ANY, MagicMock, Mock, patch import jwt @@ -2318,35 +2320,34 @@ class TestScanViewSet: assert response.status_code == 404 assert response.json()["errors"]["detail"] == "The scan has no reports." - def test_report_local_file( - self, authenticated_client, scans_fixture, tmp_path, monkeypatch - ): - """ - When output_location is a local file path, the view should read the file from disk - and return it with proper headers. - """ + def test_report_local_file(self, authenticated_client, scans_fixture, monkeypatch): scan = scans_fixture[0] - file_content = b"local zip file content" - file_path = tmp_path / "report.zip" - file_path.write_bytes(file_content) + with tempfile.TemporaryDirectory() as tmp: + tmp_path = Path(tmp) + base_tmp = tmp_path / "report_local_file" + base_tmp.mkdir(parents=True, exist_ok=True) - scan.output_location = str(file_path) - scan.state = StateChoices.COMPLETED - scan.save() + file_content = b"local zip file content" + file_path = base_tmp / "report.zip" + file_path.write_bytes(file_content) - monkeypatch.setattr( - glob, - "glob", - lambda pattern: [str(file_path)] if pattern == str(file_path) else [], - ) + scan.output_location = str(file_path) + scan.state = StateChoices.COMPLETED + scan.save() - url = reverse("scan-report", kwargs={"pk": scan.id}) - response = authenticated_client.get(url) - assert response.status_code == 200 - assert response.content == file_content - content_disposition = response.get("Content-Disposition") - assert content_disposition.startswith('attachment; filename="') - assert f'filename="{file_path.name}"' in content_disposition + monkeypatch.setattr( + glob, + "glob", + lambda pattern: [str(file_path)] if pattern == str(file_path) else [], + ) + + url = reverse("scan-report", kwargs={"pk": scan.id}) + response = authenticated_client.get(url) + assert response.status_code == 200 + assert response.content == file_content + content_disposition = response.get("Content-Disposition") + 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] @@ -2481,31 +2482,36 @@ class TestScanViewSet: ) def test_compliance_local_file( - self, authenticated_client, scans_fixture, tmp_path, monkeypatch + self, authenticated_client, scans_fixture, 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() + with tempfile.TemporaryDirectory() as tmp: + tmp_path = Path(tmp) + base = tmp_path / "reports" + comp_dir = base / "compliance" + comp_dir.mkdir(parents=True, exist_ok=True) + fname = comp_dir / "scan_cis.csv" + fname.write_bytes(b"ignored") - monkeypatch.setattr( - glob, - "glob", - lambda p: [str(fname)] if p.endswith("*_cis_1.4_aws.csv") else [], - ) + scan.output_location = str(base / "scan.zip") + scan.save() - 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}"') + 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}"') @patch("api.v1.views.Task.objects.get") @patch("api.v1.views.TaskSerializer") diff --git a/api/src/backend/tasks/tests/test_export.py b/api/src/backend/tasks/tests/test_export.py index 2aefbce9f9..6811fe7449 100644 --- a/api/src/backend/tasks/tests/test_export.py +++ b/api/src/backend/tasks/tests/test_export.py @@ -1,5 +1,6 @@ import os import zipfile +from pathlib import Path from unittest.mock import MagicMock, patch import pytest @@ -14,8 +15,9 @@ from tasks.jobs.export import ( @pytest.mark.django_db class TestOutputs: - def test_compress_output_files_creates_zip(self, tmp_path): - output_dir = tmp_path / "output" + def test_compress_output_files_creates_zip(self, tmpdir): + base_tmp = Path(str(tmpdir.mkdir("compress_output"))) + output_dir = base_tmp / "output" output_dir.mkdir() file_path = output_dir / "result.csv" file_path.write_text("data") @@ -54,16 +56,16 @@ class TestOutputs: @patch("tasks.jobs.export.get_s3_client") @patch("tasks.jobs.export.base") - def test_upload_to_s3_success(self, mock_base, mock_get_client, tmp_path): + def test_upload_to_s3_success(self, mock_base, mock_get_client, tmpdir): mock_base.DJANGO_OUTPUT_S3_AWS_OUTPUT_BUCKET = "test-bucket" - zip_path = tmp_path / "outputs.zip" + base_tmp = Path(str(tmpdir.mkdir("upload_success"))) + zip_path = base_tmp / "outputs.zip" zip_path.write_bytes(b"dummy") - compliance_dir = tmp_path / "compliance" + compliance_dir = base_tmp / "compliance" compliance_dir.mkdir() - compliance_file = compliance_dir / "report.csv" - compliance_file.write_text("ok") + (compliance_dir / "report.csv").write_text("ok") client_mock = MagicMock() mock_get_client.return_value = client_mock @@ -72,7 +74,6 @@ class TestOutputs: expected_uri = "s3://test-bucket/tenant-id/scan-id/outputs.zip" assert result == expected_uri - assert client_mock.upload_file.call_count == 2 @patch("tasks.jobs.export.get_s3_client") @@ -84,12 +85,14 @@ class TestOutputs: @patch("tasks.jobs.export.get_s3_client") @patch("tasks.jobs.export.base") - def test_upload_to_s3_skips_non_files(self, mock_base, mock_get_client, tmp_path): + def test_upload_to_s3_skips_non_files(self, mock_base, mock_get_client, tmpdir): mock_base.DJANGO_OUTPUT_S3_AWS_OUTPUT_BUCKET = "test-bucket" - zip_path = tmp_path / "results.zip" + base_tmp = Path(str(tmpdir.mkdir("upload_skips_non_files"))) + + zip_path = base_tmp / "results.zip" zip_path.write_bytes(b"zip") - compliance_dir = tmp_path / "compliance" + compliance_dir = base_tmp / "compliance" compliance_dir.mkdir() (compliance_dir / "subdir").mkdir() @@ -100,7 +103,6 @@ class TestOutputs: expected_uri = "s3://test-bucket/tenant/scan/results.zip" assert result == expected_uri - client_mock.upload_file.assert_called_once() @patch( @@ -110,23 +112,26 @@ class TestOutputs: @patch("tasks.jobs.export.base") @patch("tasks.jobs.export.logger.error") def test_upload_to_s3_failure_logs_error( - self, mock_logger, mock_base, mock_get_client, tmp_path + self, mock_logger, mock_base, mock_get_client, tmpdir ): mock_base.DJANGO_OUTPUT_S3_AWS_OUTPUT_BUCKET = "bucket" - zip_path = tmp_path / "zipfile.zip" + + base_tmp = Path(str(tmpdir.mkdir("upload_failure_logs"))) + zip_path = base_tmp / "zipfile.zip" zip_path.write_bytes(b"zip") - compliance_dir = tmp_path / "compliance" + compliance_dir = base_tmp / "compliance" compliance_dir.mkdir() (compliance_dir / "report.csv").write_text("csv") _upload_to_s3("tenant", str(zip_path), "scan") mock_logger.assert_called() - def test_generate_output_directory_creates_paths(self, tmp_path): + def test_generate_output_directory_creates_paths(self, tmpdir): from prowler.config.config import output_file_timestamp - base_dir = str(tmp_path) + base_tmp = Path(str(tmpdir.mkdir("generate_output"))) + base_dir = str(base_tmp) tenant_id = "t1" scan_id = "s1" provider = "aws"