mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-24 13:01:56 +00:00
feat/PRWLR-4014 Implement SDK integration for POST /providers/{provider_id}/connection (#30)
* chore(deps): PRWLR-4014 include prowler library in python deps * feat(Backend,API): PRWLR-4014 add AWS provider test_connection through celery tasks * fix(Backend,API): PRWLR-4014 fix model handling in celery tasks * test(Tasks): PRWLR-4014 add unit tests for celery tasks * docs(API): PRWLR-4014 update endpoint tag * feat(Backend): PRWLR-4014 add decorator for tenant dependant Celery tasks * chore(Backend): PRWLR-4014 remove TODOs and improve docstrings * feat(Backend): PRWLR-4014 replace timezone.now for datetime.now(timezone.utc) * feat(Backend): PRWLR-4014 use SET LOCAL for api.tenant_id setting * feat(Backend, Tasks): PRWLR-4014 refactor tasks module to abstract business logic * fix(Backend): PRWLR-4014 use set_config for RLS config and set transaction at request dispatch level * fix(Tasks): PRWLR-4014 fix tasks tenant decorator
This commit is contained in:
committed by
GitHub
parent
8f2bd45872
commit
8183207802
Generated
+2409
-125
File diff suppressed because it is too large
Load Diff
+3
-1
@@ -23,9 +23,11 @@ djangorestframework-jsonapi = "7.0.2"
|
||||
drf-spectacular = "0.27.2"
|
||||
drf-spectacular-jsonapi = "0.4.1"
|
||||
gunicorn = "23.0.0"
|
||||
prowler = {git = "https://github.com/prowler-cloud/prowler.git", branch = "master"}
|
||||
psycopg2-binary = "2.9.9"
|
||||
pytest-celery = {extras = ["redis"], version = "^1.0.1"}
|
||||
python = "^3.12"
|
||||
# Needed for prowler compatibility
|
||||
python = ">=3.11,<3.13"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
bandit = "1.7.9"
|
||||
|
||||
@@ -29,6 +29,10 @@ class BaseViewSet(ModelViewSet):
|
||||
|
||||
|
||||
class BaseRLSViewSet(BaseViewSet):
|
||||
def dispatch(self, request, *args, **kwargs):
|
||||
with transaction.atomic():
|
||||
return super().dispatch(request, *args, **kwargs)
|
||||
|
||||
def initial(self, request, *args, **kwargs):
|
||||
# Ideally, this logic would be in the `.setup()` method but DRF view sets don't call it
|
||||
# https://docs.djangoproject.com/en/5.1/ref/class-based-views/base/#django.views.generic.base.View.setup
|
||||
@@ -44,9 +48,8 @@ class BaseRLSViewSet(BaseViewSet):
|
||||
except ValueError:
|
||||
raise ValidationError("X-Tenant-ID header must be a valid UUID")
|
||||
|
||||
with transaction.atomic():
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute("SET api.tenant_id = %s", [tenant_id])
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(f"SELECT set_config('api.tenant_id', '{tenant_id}', TRUE);")
|
||||
return super().initial(request, *args, **kwargs)
|
||||
|
||||
def get_serializer_context(self):
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
from functools import wraps
|
||||
|
||||
from django.db import connection, transaction
|
||||
|
||||
|
||||
def set_tenant(func):
|
||||
"""
|
||||
Decorator to set the tenant context for a Celery task based on the provided tenant_id.
|
||||
|
||||
This decorator extracts the `tenant_id` from the task's keyword arguments,
|
||||
and uses it to set the tenant context for the current database session.
|
||||
The `tenant_id` is then removed from the kwargs before the task function
|
||||
is executed. If `tenant_id` is not provided, a KeyError is raised.
|
||||
|
||||
Args:
|
||||
func (function): The Celery task function to be decorated.
|
||||
|
||||
Raises:
|
||||
KeyError: If `tenant_id` is not found in the task's keyword arguments.
|
||||
|
||||
Returns:
|
||||
function: The wrapped function with tenant context set.
|
||||
|
||||
Example:
|
||||
# This decorator MUST be defined the last in the decorator chain
|
||||
|
||||
@shared_task
|
||||
@set_tenant
|
||||
def some_task(arg1, **kwargs):
|
||||
# Task logic here
|
||||
pass
|
||||
|
||||
# When calling the task
|
||||
some_task.delay(arg1, tenant_id="1234-abcd-5678")
|
||||
|
||||
# The tenant context will be set before the task logic executes.
|
||||
"""
|
||||
|
||||
@wraps(func)
|
||||
@transaction.atomic
|
||||
def wrapper(*args, **kwargs):
|
||||
try:
|
||||
tenant_id = kwargs.pop("tenant_id")
|
||||
except KeyError:
|
||||
raise KeyError("This task requires the tenant_id")
|
||||
|
||||
with connection.cursor() as cursor:
|
||||
cursor.execute(f"SELECT set_config('api.tenant_id', '{tenant_id}', TRUE);")
|
||||
|
||||
return func(*args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
@@ -0,0 +1,34 @@
|
||||
from unittest.mock import patch, call
|
||||
|
||||
import pytest
|
||||
|
||||
from api.decorators import set_tenant
|
||||
|
||||
|
||||
@pytest.mark.django_db
|
||||
class TestSetTenantDecorator:
|
||||
@patch("api.decorators.connection.cursor")
|
||||
def test_set_tenant(self, mock_cursor):
|
||||
mock_cursor.return_value.__enter__.return_value = mock_cursor
|
||||
|
||||
@set_tenant
|
||||
def random_func(arg):
|
||||
return arg
|
||||
|
||||
tenant_id = "1234-abcd-5678"
|
||||
|
||||
result = random_func("test_arg", tenant_id=tenant_id)
|
||||
|
||||
assert (
|
||||
call(f"SELECT set_config('api.tenant_id', '{tenant_id}', TRUE);")
|
||||
in mock_cursor.execute.mock_calls
|
||||
)
|
||||
assert result == "test_arg"
|
||||
|
||||
def test_set_tenant_exception(self):
|
||||
@set_tenant
|
||||
def random_func(arg):
|
||||
return arg
|
||||
|
||||
with pytest.raises(KeyError):
|
||||
random_func("test_arg")
|
||||
@@ -1,4 +1,4 @@
|
||||
from datetime import datetime
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
from django.urls import reverse
|
||||
@@ -433,7 +433,15 @@ class TestProviderViewSet:
|
||||
)
|
||||
assert response.status_code == status.HTTP_404_NOT_FOUND
|
||||
|
||||
def test_providers_connection(self, client, providers, tenant_header):
|
||||
@patch("tasks.tasks.check_provider_connection_task.delay")
|
||||
def test_providers_connection(
|
||||
self, mock_provider_connection, client, providers, tenant_header
|
||||
):
|
||||
task_mock = Mock()
|
||||
task_mock.id = "12345"
|
||||
task_mock.status = "PENDING"
|
||||
mock_provider_connection.return_value = task_mock
|
||||
|
||||
provider1, *_ = providers
|
||||
assert provider1.connected is None
|
||||
assert provider1.connection_last_checked_at is None
|
||||
@@ -443,12 +451,11 @@ class TestProviderViewSet:
|
||||
headers=tenant_header,
|
||||
)
|
||||
assert response.status_code == status.HTTP_202_ACCEPTED
|
||||
mock_provider_connection.assert_called_once_with(
|
||||
provider_id=str(provider1.id), tenant_id=tenant_header["X-Tenant-ID"]
|
||||
)
|
||||
assert "Content-Location" in response.headers
|
||||
# TODO Assert a task is returned when they are implemented
|
||||
|
||||
provider1.refresh_from_db()
|
||||
assert provider1.connected is True
|
||||
assert isinstance(provider1.connection_last_checked_at, datetime)
|
||||
assert response.headers["Content-Location"] == f"api/v1/tasks/{task_mock.id}"
|
||||
|
||||
def test_providers_connection_invalid_provider(
|
||||
self, client, providers, tenant_header
|
||||
|
||||
@@ -31,6 +31,20 @@ class RLSSerializer(BaseSerializerV1):
|
||||
return super().create(validated_data)
|
||||
|
||||
|
||||
# Tasks
|
||||
|
||||
|
||||
class DelayedTaskSerializer(serializers.Serializer):
|
||||
id = serializers.CharField()
|
||||
status = serializers.CharField()
|
||||
|
||||
class JSONAPIMeta:
|
||||
resource_name = "Task"
|
||||
|
||||
def to_representation(self, obj):
|
||||
return {"id": obj.id, "status": obj.status}
|
||||
|
||||
|
||||
# Tenants
|
||||
|
||||
|
||||
|
||||
+19
-15
@@ -1,6 +1,4 @@
|
||||
from django.conf import settings as django_settings
|
||||
from django.urls import reverse
|
||||
from django.utils import timezone
|
||||
from django.utils.decorators import method_decorator
|
||||
from django.views.decorators.cache import cache_control
|
||||
from drf_spectacular.settings import spectacular_settings
|
||||
@@ -19,7 +17,9 @@ from api.v1.serializers import (
|
||||
ProviderSerializer,
|
||||
ProviderUpdateSerializer,
|
||||
TenantSerializer,
|
||||
DelayedTaskSerializer,
|
||||
)
|
||||
from tasks.tasks import check_provider_connection_task
|
||||
|
||||
CACHE_DECORATOR = cache_control(
|
||||
max_age=django_settings.CACHE_MAX_AGE,
|
||||
@@ -99,15 +99,13 @@ class TenantViewSet(BaseViewSet):
|
||||
destroy=extend_schema(
|
||||
summary="Delete a provider",
|
||||
description="Remove a provider from the system by their ID.",
|
||||
# TODO Update with task response when implemented
|
||||
responses={202: ProviderSerializer},
|
||||
responses={202: DelayedTaskSerializer},
|
||||
),
|
||||
)
|
||||
@method_decorator(CACHE_DECORATOR, name="list")
|
||||
@method_decorator(CACHE_DECORATOR, name="retrieve")
|
||||
class ProviderViewSet(BaseRLSViewSet):
|
||||
queryset = Provider.objects.all()
|
||||
serializer_class = ProviderSerializer
|
||||
http_method_names = ["get", "post", "patch", "delete"]
|
||||
filterset_class = ProviderFilter
|
||||
search_fields = ["provider", "provider_id", "alias"]
|
||||
@@ -124,6 +122,12 @@ class ProviderViewSet(BaseRLSViewSet):
|
||||
def get_queryset(self):
|
||||
return Provider.objects.all()
|
||||
|
||||
def get_serializer_class(self):
|
||||
# TODO Add `destroy` when refactored
|
||||
if self.action in {"connection"}:
|
||||
return DelayedTaskSerializer
|
||||
return ProviderSerializer
|
||||
|
||||
def partial_update(self, request, *args, **kwargs):
|
||||
instance = self.get_object()
|
||||
serializer = ProviderUpdateSerializer(instance, data=request.data, partial=True)
|
||||
@@ -133,25 +137,25 @@ class ProviderViewSet(BaseRLSViewSet):
|
||||
return Response(data=read_serializer.data, status=status.HTTP_200_OK)
|
||||
|
||||
@extend_schema(
|
||||
tags=["Provider"],
|
||||
summary="Check connection",
|
||||
description="Try to verify connection. For instance, Role & Credentials are set correctly",
|
||||
request=None,
|
||||
# TODO Update with task response when implemented
|
||||
responses={202: ProviderSerializer},
|
||||
responses={202: DelayedTaskSerializer},
|
||||
)
|
||||
@action(detail=True, methods=["post"], url_name="connection")
|
||||
def connection(self, request, pk=None):
|
||||
# TODO Connection check to third parties
|
||||
# TODO Returning provider data as a placeholder for tasks. We will do something similar for the created task
|
||||
provider = get_object_or_404(Provider, pk=pk)
|
||||
provider.connected = True
|
||||
provider.connection_last_checked_at = timezone.now()
|
||||
provider.save()
|
||||
serializer = ProviderSerializer(provider)
|
||||
get_object_or_404(Provider, pk=pk)
|
||||
task = check_provider_connection_task.delay(
|
||||
provider_id=pk, tenant_id=request.headers.get("X-Tenant-ID")
|
||||
)
|
||||
serializer = DelayedTaskSerializer(task)
|
||||
return Response(
|
||||
data=serializer.data,
|
||||
status=status.HTTP_202_ACCEPTED,
|
||||
headers={"Content-Location": reverse("provider-detail", kwargs={"pk": pk})},
|
||||
# TODO Use /tasks view name when implemented
|
||||
# headers={"Content-Location": reverse("task-detail", kwargs={"pk": task.id})},
|
||||
headers={"Content-Location": f"api/v1/tasks/{task.id}"},
|
||||
)
|
||||
|
||||
def destroy(self, request, *args, **kwargs):
|
||||
|
||||
@@ -4,5 +4,6 @@ from celery import Celery
|
||||
celery_app = Celery("tasks")
|
||||
|
||||
celery_app.config_from_object("django.conf:settings", namespace="CELERY")
|
||||
celery_app.conf.update(result_extended=True)
|
||||
|
||||
celery_app.autodiscover_tasks(["tasks"])
|
||||
celery_app.autodiscover_tasks(["api"])
|
||||
|
||||
@@ -6,5 +6,6 @@ VALKEY_DB = env("VALKEY_DB", default="0")
|
||||
|
||||
CELERY_BROKER_URL = f"redis://{VALKEY_HOST}:{VALKEY_PORT}/{VALKEY_DB}"
|
||||
CELERY_RESULT_BACKEND = "django-db"
|
||||
CELERY_TASK_TRACK_STARTED = True
|
||||
|
||||
CELERY_BROKER_CONNECTION_RETRY_ON_STARTUP = True
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
from prowler.providers.aws.aws_provider import AwsProvider
|
||||
|
||||
from api.models import Provider
|
||||
|
||||
logger = get_task_logger(__name__)
|
||||
|
||||
|
||||
def check_provider_connection(provider_id: str):
|
||||
"""
|
||||
Business logic to check the connection status of a provider.
|
||||
|
||||
Args:
|
||||
provider_id (str): The primary key of the Provider instance to check.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing:
|
||||
- 'connected' (bool): Indicates whether the provider is successfully connected.
|
||||
- 'error' (str or None): The error message if the connection failed, otherwise `None`.
|
||||
|
||||
Raises:
|
||||
ValueError: If the provider type is not supported.
|
||||
Model.DoesNotExist: If the provider does not exist.
|
||||
"""
|
||||
provider_instance = Provider.objects.get(pk=provider_id)
|
||||
match provider_instance.provider:
|
||||
case Provider.ProviderChoices.AWS.value:
|
||||
# TODO Refactor when proper credentials are implemented
|
||||
try:
|
||||
connection_result = AwsProvider.test_connection(
|
||||
raise_on_exception=False
|
||||
)
|
||||
# TODO: Improve this exception handling when SDK exceptions are implemented
|
||||
except Exception as e:
|
||||
logger.warning(str(e))
|
||||
raise e
|
||||
case _:
|
||||
raise ValueError(
|
||||
f"Provider type {provider_instance.provider} not supported"
|
||||
)
|
||||
|
||||
provider_instance.connected = connection_result.is_connected
|
||||
provider_instance.connection_last_checked_at = datetime.now(tz=timezone.utc)
|
||||
provider_instance.save()
|
||||
|
||||
connection_error = (
|
||||
f"{connection_result.error.__class__.__name__}: {connection_result.error}"
|
||||
if connection_result.error
|
||||
else None
|
||||
)
|
||||
return {"connected": connection_result.is_connected, "error": connection_error}
|
||||
+16
-11
@@ -1,16 +1,21 @@
|
||||
from celery import shared_task
|
||||
|
||||
from celery.utils.log import get_task_logger
|
||||
|
||||
# TODO: add celery metadata to CustomLogger
|
||||
logger = get_task_logger(__name__)
|
||||
from api.decorators import set_tenant
|
||||
from tasks.jobs.connection import check_provider_connection
|
||||
|
||||
|
||||
@shared_task
|
||||
def debug_task(x, y):
|
||||
try:
|
||||
logger.info("Hello world")
|
||||
except Exception as e:
|
||||
logger.error(e)
|
||||
@shared_task(name="provider-connection-check")
|
||||
@set_tenant
|
||||
def check_provider_connection_task(provider_id: str):
|
||||
"""
|
||||
Task to check the connection status of a provider.
|
||||
|
||||
return x + y
|
||||
Args:
|
||||
provider_id (str): The primary key of the Provider instance to check.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing:
|
||||
- 'connected' (bool): Indicates whether the provider is successfully connected.
|
||||
- 'error' (str or None): The error message if the connection failed, otherwise `None`.
|
||||
"""
|
||||
return check_provider_connection(provider_id)
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import patch, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from api.models import Provider
|
||||
from tasks.jobs.connection import check_provider_connection
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"provider_data, provider_class",
|
||||
[
|
||||
(
|
||||
{"provider": "aws", "provider_id": "123456789012", "alias": "aws"},
|
||||
"AwsProvider",
|
||||
),
|
||||
],
|
||||
)
|
||||
@pytest.mark.django_db
|
||||
def test_check_provider_connection(get_tenant, provider_data, provider_class):
|
||||
provider = Provider.objects.create(**provider_data, tenant_id=get_tenant.id)
|
||||
|
||||
mock_test_connection_result = MagicMock()
|
||||
mock_test_connection_result.is_connected = True
|
||||
|
||||
with patch(
|
||||
f"tasks.jobs.connection.{provider_class}.test_connection"
|
||||
) as mock_test_connection:
|
||||
mock_test_connection.return_value = mock_test_connection_result
|
||||
|
||||
check_provider_connection(
|
||||
provider_id=str(provider.id),
|
||||
)
|
||||
provider.refresh_from_db()
|
||||
|
||||
mock_test_connection.assert_called_once()
|
||||
assert provider.connected is True
|
||||
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")
|
||||
@pytest.mark.django_db
|
||||
def test_check_provider_connection_unsupported_provider(mock_provider_get):
|
||||
mock_provider_instance = MagicMock()
|
||||
mock_provider_instance.provider = "UNSUPPORTED_PROVIDER"
|
||||
mock_provider_get.return_value = mock_provider_instance
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match="Provider type UNSUPPORTED_PROVIDER not supported"
|
||||
):
|
||||
check_provider_connection("provider_id")
|
||||
|
||||
|
||||
@patch("tasks.jobs.connection.Provider.objects.get")
|
||||
@patch("tasks.jobs.connection.AwsProvider.test_connection")
|
||||
@pytest.mark.django_db
|
||||
def test_check_provider_connection_exception(mock_test_connection, mock_provider_get):
|
||||
mock_provider_instance = MagicMock()
|
||||
mock_provider_instance.provider = Provider.ProviderChoices.AWS.value
|
||||
mock_provider_get.return_value = mock_provider_instance
|
||||
|
||||
mock_test_connection.return_value = MagicMock()
|
||||
mock_test_connection.return_value.is_connected = False
|
||||
mock_test_connection.return_value.error = Exception()
|
||||
|
||||
result = check_provider_connection(provider_id="provider_id")
|
||||
|
||||
assert result["connected"] is False
|
||||
assert result["error"] is not None
|
||||
|
||||
mock_provider_instance.save.assert_called_once()
|
||||
assert mock_provider_instance.connected is False
|
||||
@@ -1,3 +0,0 @@
|
||||
# @pytest.mark.celery
|
||||
# def test_debug_task(celery_worker):
|
||||
# assert debug_task.delay(1, 2).get() == 4
|
||||
Reference in New Issue
Block a user