mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-24 04:51:51 +00:00
feat(api): add provider status
This commit is contained in:
@@ -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,
|
||||
),
|
||||
]
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}"
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user