From d239d299e251a43f3827c818829d6132ab4d9d23 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adri=C3=A1n=20Jes=C3=BAs=20Pe=C3=B1a=20Rodr=C3=ADguez?= Date: Thu, 31 Jul 2025 10:00:05 +0200 Subject: [PATCH] fix(s3): use enabled to filter (#8409) --- api/src/backend/tasks/jobs/connection.py | 7 +- api/src/backend/tasks/jobs/integrations.py | 1 + api/src/backend/tasks/tasks.py | 4 +- .../backend/tasks/tests/test_connection.py | 134 +++++++++++++++++- .../backend/tasks/tests/test_integrations.py | 24 ++++ api/src/backend/tasks/tests/test_tasks.py | 78 +++++++++- 6 files changed, 242 insertions(+), 6 deletions(-) diff --git a/api/src/backend/tasks/jobs/connection.py b/api/src/backend/tasks/jobs/connection.py index bbb9d5c9e5..d7068ebf3b 100644 --- a/api/src/backend/tasks/jobs/connection.py +++ b/api/src/backend/tasks/jobs/connection.py @@ -95,7 +95,12 @@ def check_integration_connection(integration_id: str): Args: integration_id (str): The primary key of the Integration instance to check. """ - integration = Integration.objects.get(pk=integration_id) + integration = Integration.objects.filter(pk=integration_id, enabled=True).first() + + if not integration: + logger.info(f"Integration {integration_id} is not enabled") + return {"connected": False, "error": "Integration is not enabled"} + try: result = prowler_integration_connection_test(integration) except Exception as e: diff --git a/api/src/backend/tasks/jobs/integrations.py b/api/src/backend/tasks/jobs/integrations.py index 7fd9980ab3..f4b6b6cf58 100644 --- a/api/src/backend/tasks/jobs/integrations.py +++ b/api/src/backend/tasks/jobs/integrations.py @@ -71,6 +71,7 @@ def upload_s3_integration( Integration.objects.filter( integrationproviderrelationship__provider_id=provider_id, integration_type=Integration.IntegrationChoices.AMAZON_S3, + enabled=True, ) ) diff --git a/api/src/backend/tasks/tasks.py b/api/src/backend/tasks/tasks.py index 5a6aa5cfc5..b0dff077aa 100644 --- a/api/src/backend/tasks/tasks.py +++ b/api/src/backend/tasks/tasks.py @@ -387,6 +387,7 @@ def generate_outputs_task(scan_id: str, provider_id: str, tenant_id: str): s3_integrations = Integration.objects.filter( integrationproviderrelationship__provider_id=provider_id, integration_type=Integration.IntegrationChoices.AMAZON_S3, + enabled=True, ) if s3_integrations: @@ -486,7 +487,8 @@ def check_integrations_task(tenant_id: str, provider_id: str): try: with rls_transaction(tenant_id): integrations = Integration.objects.filter( - integrationproviderrelationship__provider_id=provider_id + integrationproviderrelationship__provider_id=provider_id, + enabled=True, ) if not integrations.exists(): diff --git a/api/src/backend/tasks/tests/test_connection.py b/api/src/backend/tasks/tests/test_connection.py index 30c8b147bc..30973f98bf 100644 --- a/api/src/backend/tasks/tests/test_connection.py +++ b/api/src/backend/tasks/tests/test_connection.py @@ -1,10 +1,15 @@ +import uuid from datetime import datetime, timezone from unittest.mock import MagicMock, patch import pytest -from tasks.jobs.connection import check_lighthouse_connection, check_provider_connection +from tasks.jobs.connection import ( + check_integration_connection, + check_lighthouse_connection, + check_provider_connection, +) -from api.models import LighthouseConfiguration, Provider +from api.models import Integration, LighthouseConfiguration, Provider @pytest.mark.parametrize( @@ -127,3 +132,128 @@ def test_check_lighthouse_connection_missing_api_key(mock_lighthouse_get): assert result["available_models"] == [] assert mock_lighthouse_instance.is_active is False mock_lighthouse_instance.save.assert_called_once() + + +@pytest.mark.django_db +class TestCheckIntegrationConnection: + def setup_method(self): + self.integration_id = str(uuid.uuid4()) + + @patch("tasks.jobs.connection.Integration.objects.filter") + @patch("tasks.jobs.connection.prowler_integration_connection_test") + def test_check_integration_connection_success( + self, mock_prowler_test, mock_integration_filter + ): + """Test successful integration connection check with enabled=True filter.""" + mock_integration = MagicMock() + mock_integration.id = self.integration_id + mock_integration.integration_type = Integration.IntegrationChoices.AMAZON_S3 + + mock_queryset = MagicMock() + mock_queryset.first.return_value = mock_integration + mock_integration_filter.return_value = mock_queryset + + mock_connection_result = MagicMock() + mock_connection_result.is_connected = True + mock_connection_result.error = None + mock_prowler_test.return_value = mock_connection_result + + result = check_integration_connection(integration_id=self.integration_id) + + # Verify that Integration.objects.filter was called with enabled=True filter + mock_integration_filter.assert_called_once_with( + pk=self.integration_id, enabled=True + ) + mock_queryset.first.assert_called_once() + mock_prowler_test.assert_called_once_with(mock_integration) + + # Verify the integration properties were updated + assert mock_integration.connected is True + assert mock_integration.connection_last_checked_at is not None + mock_integration.save.assert_called_once() + + # Verify the return value + assert result["connected"] is True + assert result["error"] is None + + @patch("tasks.jobs.connection.Integration.objects.filter") + @patch("tasks.jobs.connection.prowler_integration_connection_test") + def test_check_integration_connection_failure( + self, mock_prowler_test, mock_integration_filter + ): + """Test failed integration connection check.""" + mock_integration = MagicMock() + mock_integration.id = self.integration_id + + mock_queryset = MagicMock() + mock_queryset.first.return_value = mock_integration + mock_integration_filter.return_value = mock_queryset + + test_error = Exception("Connection failed") + mock_connection_result = MagicMock() + mock_connection_result.is_connected = False + mock_connection_result.error = test_error + mock_prowler_test.return_value = mock_connection_result + + result = check_integration_connection(integration_id=self.integration_id) + + # Verify that Integration.objects.filter was called with enabled=True filter + mock_integration_filter.assert_called_once_with( + pk=self.integration_id, enabled=True + ) + mock_queryset.first.assert_called_once() + + # Verify the integration properties were updated + assert mock_integration.connected is False + assert mock_integration.connection_last_checked_at is not None + mock_integration.save.assert_called_once() + + # Verify the return value + assert result["connected"] is False + assert result["error"] == str(test_error) + + @patch("tasks.jobs.connection.Integration.objects.filter") + def test_check_integration_connection_not_enabled(self, mock_integration_filter): + """Test that disabled integrations return proper error response.""" + # Mock that no enabled integration is found + mock_queryset = MagicMock() + mock_queryset.first.return_value = None + mock_integration_filter.return_value = mock_queryset + + result = check_integration_connection(integration_id=self.integration_id) + + # Verify the filter was called with enabled=True + mock_integration_filter.assert_called_once_with( + pk=self.integration_id, enabled=True + ) + mock_queryset.first.assert_called_once() + + # Verify the return value matches the expected error response + assert result["connected"] is False + assert result["error"] == "Integration is not enabled" + + @patch("tasks.jobs.connection.Integration.objects.filter") + @patch("tasks.jobs.connection.prowler_integration_connection_test") + def test_check_integration_connection_exception( + self, mock_prowler_test, mock_integration_filter + ): + """Test integration connection check when prowler test raises exception.""" + mock_integration = MagicMock() + mock_integration.id = self.integration_id + + mock_queryset = MagicMock() + mock_queryset.first.return_value = mock_integration + mock_integration_filter.return_value = mock_queryset + + test_exception = Exception("Unexpected error during connection test") + mock_prowler_test.side_effect = test_exception + + with pytest.raises(Exception, match="Unexpected error during connection test"): + check_integration_connection(integration_id=self.integration_id) + + # Verify that Integration.objects.filter was called with enabled=True filter + mock_integration_filter.assert_called_once_with( + pk=self.integration_id, enabled=True + ) + mock_queryset.first.assert_called_once() + mock_prowler_test.assert_called_once_with(mock_integration) diff --git a/api/src/backend/tasks/tests/test_integrations.py b/api/src/backend/tasks/tests/test_integrations.py index 2413e3f0b4..72f6d88365 100644 --- a/api/src/backend/tasks/tests/test_integrations.py +++ b/api/src/backend/tasks/tests/test_integrations.py @@ -208,6 +208,30 @@ class TestS3IntegrationUploads: "S3 connection failed for integration i-1: failed" ) + @patch("tasks.jobs.integrations.rls_transaction") + @patch("tasks.jobs.integrations.Integration.objects.filter") + def test_upload_s3_integration_filters_enabled_only( + self, mock_integration_filter, mock_rls + ): + """Test that upload_s3_integration only processes enabled integrations.""" + tenant_id = "tenant-id" + provider_id = "provider-id" + output_directory = "/tmp/prowler_output/scan123" + + # Mock that no enabled integrations are found + mock_integration_filter.return_value = [] + mock_rls.return_value.__enter__.return_value = None + + result = upload_s3_integration(tenant_id, provider_id, output_directory) + + assert result is False + # Verify the filter includes the correct parameters including enabled=True + mock_integration_filter.assert_called_once_with( + integrationproviderrelationship__provider_id=provider_id, + integration_type=Integration.IntegrationChoices.AMAZON_S3, + enabled=True, + ) + def test_s3_integration_validates_and_normalizes_output_directory(self): """Test that S3 integration validation normalizes output_directory paths.""" from api.models import Integration diff --git a/api/src/backend/tasks/tests/test_tasks.py b/api/src/backend/tasks/tests/test_tasks.py index d74f13fd6e..37dbc5c83a 100644 --- a/api/src/backend/tasks/tests/test_tasks.py +++ b/api/src/backend/tasks/tests/test_tasks.py @@ -9,6 +9,8 @@ from tasks.tasks import ( s3_integration_task, ) +from api.models import Integration + # TODO Move this to outputs/reports jobs @pytest.mark.django_db @@ -418,6 +420,56 @@ class TestGenerateOutputs: ) assert "Error deleting output files" in caplog.text + @patch("tasks.tasks.rls_transaction") + @patch("tasks.tasks.Integration.objects.filter") + def test_generate_outputs_filters_enabled_s3_integrations( + self, mock_integration_filter, mock_rls + ): + """Test that generate_outputs_task only processes enabled S3 integrations.""" + with ( + patch("tasks.tasks.ScanSummary.objects.filter") as mock_summary, + patch("tasks.tasks.Provider.objects.get"), + patch("tasks.tasks.initialize_prowler_provider"), + patch("tasks.tasks.Compliance.get_bulk"), + patch("tasks.tasks.get_compliance_frameworks", return_value=[]), + patch("tasks.tasks.Finding.all_objects.filter") as mock_findings, + patch( + "tasks.tasks._generate_output_directory", return_value=("out", "comp") + ), + patch("tasks.tasks.FindingOutput._transform_findings_stats"), + patch("tasks.tasks.FindingOutput.transform_api_finding"), + patch("tasks.tasks._compress_output_files", return_value="/tmp/compressed"), + patch("tasks.tasks._upload_to_s3", return_value="s3://bucket/file.zip"), + patch("tasks.tasks.Scan.all_objects.filter"), + patch("tasks.tasks.rmtree"), + patch("tasks.tasks.s3_integration_task.apply_async") as mock_s3_task, + ): + mock_summary.return_value.exists.return_value = True + mock_findings.return_value.order_by.return_value.iterator.return_value = [ + [MagicMock()], + True, + ] + mock_integration_filter.return_value = [MagicMock()] + mock_rls.return_value.__enter__.return_value = None + + with ( + patch("tasks.tasks.OUTPUT_FORMATS_MAPPING", {}), + patch("tasks.tasks.COMPLIANCE_CLASS_MAP", {"aws": []}), + ): + generate_outputs_task( + scan_id=self.scan_id, + provider_id=self.provider_id, + tenant_id=self.tenant_id, + ) + + # Verify the S3 integrations filters + mock_integration_filter.assert_called_once_with( + integrationproviderrelationship__provider_id=self.provider_id, + integration_type=Integration.IntegrationChoices.AMAZON_S3, + enabled=True, + ) + mock_s3_task.assert_called_once() + class TestScanCompleteTasks: @patch("tasks.tasks.create_compliance_requirements_task.apply_async") @@ -465,7 +517,8 @@ class TestCheckIntegrationsTask: assert result == {"integrations_processed": 0} mock_integration_filter.assert_called_once_with( - integrationproviderrelationship__provider_id=self.provider_id + integrationproviderrelationship__provider_id=self.provider_id, + enabled=True, ) @patch("tasks.tasks.group") @@ -488,11 +541,32 @@ class TestCheckIntegrationsTask: assert result == {"integrations_processed": 0} mock_integration_filter.assert_called_once_with( - integrationproviderrelationship__provider_id=self.provider_id + integrationproviderrelationship__provider_id=self.provider_id, + enabled=True, ) # group should not be called since no integration tasks are created yet mock_group.assert_not_called() + @patch("tasks.tasks.rls_transaction") + @patch("tasks.tasks.Integration.objects.filter") + def test_check_integrations_disabled_integrations_ignored( + self, mock_integration_filter, mock_rls + ): + """Test that disabled integrations are not processed.""" + mock_integration_filter.return_value.exists.return_value = False + mock_rls.return_value.__enter__.return_value = None + + result = check_integrations_task( + tenant_id=self.tenant_id, + provider_id=self.provider_id, + ) + + assert result == {"integrations_processed": 0} + mock_integration_filter.assert_called_once_with( + integrationproviderrelationship__provider_id=self.provider_id, + enabled=True, + ) + @patch("tasks.tasks.upload_s3_integration") def test_s3_integration_task_success(self, mock_upload): mock_upload.return_value = True