From 4b0801faf9157f1e596949e71b7fb249ecde3116 Mon Sep 17 00:00:00 2001 From: pedrooot Date: Thu, 15 Jan 2026 17:43:18 +0100 Subject: [PATCH] feat(api): add provider status --- api/src/backend/api/filters.py | 14 +++++ .../api/migrations/0068_provider_status.py | 62 +++++++++++++++++++ api/src/backend/api/models.py | 18 ++++++ api/src/backend/api/specs/v1.yaml | 35 +++++++++++ api/src/backend/api/v1/serializers.py | 4 ++ api/src/backend/tasks/jobs/connection.py | 22 ++++++- api/src/backend/tasks/jobs/scan.py | 3 + .../backend/tasks/tests/test_connection.py | 60 +++++++++++++++++- api/src/backend/tasks/tests/test_scan.py | 2 + 9 files changed, 217 insertions(+), 3 deletions(-) create mode 100644 api/src/backend/api/migrations/0068_provider_status.py diff --git a/api/src/backend/api/filters.py b/api/src/backend/api/filters.py index fdb2282e12..faa6fc1707 100644 --- a/api/src/backend/api/filters.py +++ b/api/src/backend/api/filters.py @@ -40,6 +40,7 @@ from api.models import ( ProviderComplianceScore, ProviderGroup, ProviderSecret, + ProviderStatusChoices, Resource, ResourceTag, Role, @@ -319,6 +320,18 @@ class ProviderFilter(FilterSet): choices=Provider.ProviderChoices.choices, lookup_expr="in", ) + status = ChoiceFilter( + choices=ProviderStatusChoices.choices, + help_text="""Filter by provider connection status. + Valid values: pending, checking, connected, error.""", + ) + status__in = ChoiceInFilter( + field_name="status", + choices=ProviderStatusChoices.choices, + lookup_expr="in", + help_text="""Filter by multiple provider connection statuses. + Accepts comma-separated values: pending,connected,error""", + ) class Meta: model = Provider @@ -329,6 +342,7 @@ class ProviderFilter(FilterSet): "alias": ["exact", "icontains", "in"], "inserted_at": ["gte", "lte"], "updated_at": ["gte", "lte"], + "status": ["exact", "in"], } filter_overrides = { ProviderEnumField: { diff --git a/api/src/backend/api/migrations/0068_provider_status.py b/api/src/backend/api/migrations/0068_provider_status.py new file mode 100644 index 0000000000..173b83ea61 --- /dev/null +++ b/api/src/backend/api/migrations/0068_provider_status.py @@ -0,0 +1,62 @@ +# Migration to add status field to Provider model +from django.db import migrations, models + +from api.db_router import MainRouter + + +def populate_provider_status(apps, schema_editor): + """ + Populate status field based on existing connected field values. + + Migration logic for production: + - connected=True → status="connected" (successful connection) + - connected=False → status="error" (connection was attempted but failed) + - connected=NULL → status="pending" (never attempted to connect) + """ + Provider = apps.get_model("api", "Provider") + db_alias = MainRouter.admin_db + + Provider.objects.using(db_alias).filter(connected=True).update(status="connected") + Provider.objects.using(db_alias).filter(connected=False).update(status="error") + Provider.objects.using(db_alias).filter(connected__isnull=True).update( + status="pending" + ) + + +def reverse_populate(apps, schema_editor): + """ + Reverse the status population. + """ + Provider = apps.get_model("api", "Provider") + db_alias = MainRouter.admin_db + + Provider.objects.using(db_alias).all().update(status=None) + + +class Migration(migrations.Migration): + dependencies = [ + ("api", "0067_tenant_compliance_summary"), + ] + + operations = [ + migrations.AddField( + model_name="provider", + name="status", + field=models.CharField( + blank=True, + choices=[ + ("pending", "Pending"), + ("checking", "Checking"), + ("connected", "Connected"), + ("error", "Error"), + ], + default="pending", + max_length=20, + null=True, + ), + ), + migrations.RunPython( + populate_provider_status, + reverse_code=reverse_populate, + ), + ] diff --git a/api/src/backend/api/models.py b/api/src/backend/api/models.py index aa568bd17c..ee5e9091a8 100644 --- a/api/src/backend/api/models.py +++ b/api/src/backend/api/models.py @@ -110,6 +110,17 @@ class PermissionChoices(models.TextChoices): NONE = "none", _("No permissions") +class ProviderStatusChoices(models.TextChoices): + """ + Represents the connection status states for a Provider. + """ + + PENDING = "pending", _("Pending") + CHECKING = "checking", _("Checking") + CONNECTED = "connected", _("Connected") + ERROR = "error", _("Error") + + class ActiveProviderManager(models.Manager): def get_queryset(self): return super().get_queryset().filter(self.active_provider_filter()) @@ -419,6 +430,13 @@ class Provider(RowLevelSecurityProtectedModel): ) connected = models.BooleanField(null=True, blank=True) connection_last_checked_at = models.DateTimeField(null=True, blank=True) + status = models.CharField( + max_length=20, + choices=ProviderStatusChoices.choices, + default=ProviderStatusChoices.PENDING, + null=True, + blank=True, + ) metadata = models.JSONField(default=dict, blank=True) scanner_args = models.JSONField(default=dict, blank=True) diff --git a/api/src/backend/api/specs/v1.yaml b/api/src/backend/api/specs/v1.yaml index 2e238a5ad4..5524e253e4 100644 --- a/api/src/backend/api/specs/v1.yaml +++ b/api/src/backend/api/specs/v1.yaml @@ -6645,6 +6645,34 @@ paths: connections. If not specified, both connected and failed providers are included. Providers with no connection attempt (status is null) are excluded from this filter. + - in: query + name: filter[status] + schema: + type: string + enum: + - pending + - checking + - connected + - error + description: |- + Filter by provider connection status. + Valid values: pending, checking, connected, error. + - in: query + name: filter[status__in] + schema: + type: array + items: + type: string + enum: + - pending + - checking + - connected + - error + description: |- + Filter by multiple provider connection statuses. + Accepts comma-separated values: pending,connected,error + explode: false + style: form - in: query name: filter[id] schema: @@ -16867,6 +16895,13 @@ components: last_checked_at: type: string format: date-time + status: + type: string + enum: + - pending + - checking + - connected + - error readOnly: true required: - provider diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 00c8c37dfb..42b60f48b5 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -890,6 +890,7 @@ class ProviderSerializer(RLSSerializer): "properties": { "connected": {"type": "boolean"}, "last_checked_at": {"type": "string", "format": "date-time"}, + "status": {"type": "string"}, }, } ) @@ -897,6 +898,7 @@ class ProviderSerializer(RLSSerializer): return { "connected": obj.connected, "last_checked_at": obj.connection_last_checked_at, + "status": obj.status, } @@ -927,6 +929,7 @@ class ProviderIncludeSerializer(RLSSerializer): "properties": { "connected": {"type": "boolean"}, "last_checked_at": {"type": "string", "format": "date-time"}, + "status": {"type": "string"}, }, } ) @@ -934,6 +937,7 @@ class ProviderIncludeSerializer(RLSSerializer): return { "connected": obj.connected, "last_checked_at": obj.connection_last_checked_at, + "status": obj.status, } diff --git a/api/src/backend/tasks/jobs/connection.py b/api/src/backend/tasks/jobs/connection.py index d7068ebf3b..6350e5d338 100644 --- a/api/src/backend/tasks/jobs/connection.py +++ b/api/src/backend/tasks/jobs/connection.py @@ -3,7 +3,12 @@ from datetime import datetime, timezone import openai from celery.utils.log import get_task_logger -from api.models import Integration, LighthouseConfiguration, Provider +from api.models import ( + Integration, + LighthouseConfiguration, + Provider, + ProviderStatusChoices, +) from api.utils import ( prowler_integration_connection_test, prowler_provider_connection_test, @@ -29,16 +34,31 @@ def check_provider_connection(provider_id: str): Model.DoesNotExist: If the provider does not exist. """ provider_instance = Provider.objects.get(pk=provider_id) + + # Set status to CHECKING before the connection test + provider_instance.status = ProviderStatusChoices.CHECKING + provider_instance.save(update_fields=["status"]) + try: connection_result = prowler_provider_connection_test(provider_instance) except Exception as e: logger.warning( f"Unexpected exception checking {provider_instance.provider} provider connection: {str(e)}" ) + # Set status to ERROR on exception + provider_instance.status = ProviderStatusChoices.ERROR + provider_instance.connected = False + provider_instance.connection_last_checked_at = datetime.now(tz=timezone.utc) + provider_instance.save() raise e provider_instance.connected = connection_result.is_connected provider_instance.connection_last_checked_at = datetime.now(tz=timezone.utc) + provider_instance.status = ( + ProviderStatusChoices.CONNECTED + if connection_result.is_connected + else ProviderStatusChoices.ERROR + ) provider_instance.save() connection_error = f"{connection_result.error}" if connection_result.error else None diff --git a/api/src/backend/tasks/jobs/scan.py b/api/src/backend/tasks/jobs/scan.py index 9e40e2df03..904a8eabde 100644 --- a/api/src/backend/tasks/jobs/scan.py +++ b/api/src/backend/tasks/jobs/scan.py @@ -39,6 +39,7 @@ from api.models import ( MuteRule, Processor, Provider, + ProviderStatusChoices, Resource, ResourceFindingMapping, ResourceScanSummary, @@ -801,8 +802,10 @@ def perform_prowler_scan( provider_instance, mutelist_processor ) provider_instance.connected = True + provider_instance.status = ProviderStatusChoices.CONNECTED except Exception as e: provider_instance.connected = False + provider_instance.status = ProviderStatusChoices.ERROR exc = ProviderConnectionError( f"Provider {provider_instance.provider} is not connected: {e}" ) diff --git a/api/src/backend/tasks/tests/test_connection.py b/api/src/backend/tasks/tests/test_connection.py index 30973f98bf..16b6fc9f37 100644 --- a/api/src/backend/tasks/tests/test_connection.py +++ b/api/src/backend/tasks/tests/test_connection.py @@ -9,7 +9,27 @@ from tasks.jobs.connection import ( check_provider_connection, ) -from api.models import Integration, LighthouseConfiguration, Provider +from api.models import ( + Integration, + LighthouseConfiguration, + Provider, + ProviderStatusChoices, +) + + +@pytest.mark.django_db +def test_provider_created_with_pending_status(tenants_fixture): + """Test that a newly created provider has PENDING status.""" + provider = Provider.objects.create( + provider="aws", + uid="123456789012", + alias="aws-test", + tenant_id=tenants_fixture[0].id, + ) + + assert provider.status == ProviderStatusChoices.PENDING + assert provider.connected is None + assert provider.connection_last_checked_at is None @pytest.mark.parametrize( @@ -25,6 +45,9 @@ def test_check_provider_connection( ): provider = Provider.objects.create(**provider_data, tenant_id=tenants_fixture[0].id) + # Verify initial state is PENDING + assert provider.status == ProviderStatusChoices.PENDING + mock_test_connection_result = MagicMock() mock_test_connection_result.is_connected = True @@ -37,10 +60,40 @@ def test_check_provider_connection( mock_provider_connection_test.assert_called_once() assert provider.connected is True + assert provider.status == ProviderStatusChoices.CONNECTED assert provider.connection_last_checked_at is not None assert provider.connection_last_checked_at <= datetime.now(tz=timezone.utc) +@patch("tasks.jobs.connection.Provider.objects.get") +@patch("tasks.jobs.connection.prowler_provider_connection_test") +@pytest.mark.django_db +def test_check_provider_connection_sets_checking_status( + mock_provider_connection_test, mock_provider_get +): + """Test that status is set to CHECKING before connection test runs.""" + mock_provider_instance = MagicMock() + mock_provider_instance.provider = Provider.ProviderChoices.AWS.value + mock_provider_get.return_value = mock_provider_instance + + captured_status = None + + def capture_status_during_test(provider): + nonlocal captured_status + captured_status = provider.status + result = MagicMock() + result.is_connected = True + result.error = None + return result + + mock_provider_connection_test.side_effect = capture_status_during_test + + check_provider_connection(provider_id="provider_id") + + # Verify status was CHECKING when prowler_provider_connection_test was called + assert captured_status == ProviderStatusChoices.CHECKING + + @patch("tasks.jobs.connection.Provider.objects.get") @pytest.mark.django_db def test_check_provider_connection_unsupported_provider(mock_provider_get): @@ -73,8 +126,11 @@ def test_check_provider_connection_exception( assert result["connected"] is False assert result["error"] is not None - mock_provider_instance.save.assert_called_once() + assert ( + mock_provider_instance.save.call_count == 2 + ) # Once for CHECKING, once for ERROR assert mock_provider_instance.connected is False + assert mock_provider_instance.status == ProviderStatusChoices.ERROR @pytest.mark.parametrize( diff --git a/api/src/backend/tasks/tests/test_scan.py b/api/src/backend/tasks/tests/test_scan.py index 8902b17b54..cbb21a32de 100644 --- a/api/src/backend/tasks/tests/test_scan.py +++ b/api/src/backend/tasks/tests/test_scan.py @@ -34,6 +34,7 @@ from api.models import ( Finding, MuteRule, Provider, + ProviderStatusChoices, Resource, Scan, StateChoices, @@ -258,6 +259,7 @@ class TestPerformScan: provider.refresh_from_db() assert provider.connected is False + assert provider.status == ProviderStatusChoices.ERROR assert isinstance(provider.connection_last_checked_at, datetime) @pytest.mark.parametrize(