feat(api): add provider status

This commit is contained in:
pedrooot
2026-01-15 17:43:18 +01:00
parent 1bf49747ad
commit 4b0801faf9
9 changed files with 217 additions and 3 deletions
+14
View File
@@ -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: {
@@ -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,
),
]
+18
View File
@@ -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)
+35
View File
@@ -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
+4
View File
@@ -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,
}
+21 -1
View File
@@ -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
+3
View File
@@ -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}"
)
+58 -2
View File
@@ -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(
+2
View File
@@ -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(