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:
Víctor Fernández Poyatos
2024-08-29 15:53:07 +02:00
committed by GitHub
parent 8f2bd45872
commit 8183207802
15 changed files with 2696 additions and 166 deletions
Generated
+2409 -125
View File
File diff suppressed because it is too large Load Diff
+3 -1
View File
@@ -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"
+6 -3
View File
@@ -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):
+52
View File
@@ -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
+34
View File
@@ -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")
+14 -7
View File
@@ -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
+14
View File
@@ -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
View File
@@ -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):
+2 -1
View File
@@ -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"])
+1
View File
@@ -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
View File
+53
View File
@@ -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
View File
@@ -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
-3
View File
@@ -1,3 +0,0 @@
# @pytest.mark.celery
# def test_debug_task(celery_worker):
# assert debug_task.delay(1, 2).get() == 4