From 6d69a192f3ff18e545d89e0fc0822779c4b84166 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?V=C3=ADctor=20Fern=C3=A1ndez=20Poyatos?= Date: Thu, 17 Oct 2024 18:07:06 +0200 Subject: [PATCH] fix(Finding, Resource): PRWLR-5057 Fix include query parameter for /findings and /resources (#55) * fix(Finding, Resource): PRWLR-5057 fix include query parameter * fix(Finding, Resource): PRWLR-5057 optimize requests * test(Finding, Resource): PRWLR-5057 add unit tests for include --- src/backend/api/renderers.py | 23 ++++++++++++ src/backend/api/tests/test_views.py | 53 ++++++++++++++++++++++++++++ src/backend/api/v1/serializers.py | 54 +++-------------------------- src/backend/api/v1/views.py | 11 ++++-- src/backend/config/django/base.py | 2 +- 5 files changed, 91 insertions(+), 52 deletions(-) create mode 100644 src/backend/api/renderers.py diff --git a/src/backend/api/renderers.py b/src/backend/api/renderers.py new file mode 100644 index 0000000000..ccca52c26f --- /dev/null +++ b/src/backend/api/renderers.py @@ -0,0 +1,23 @@ +from contextlib import nullcontext + +from rest_framework_json_api.renderers import JSONRenderer + +from api.db_utils import tenant_transaction + + +class APIJSONRenderer(JSONRenderer): + """JSONRenderer override to apply tenant RLS when there are included resources in the request.""" + + def render(self, data, accepted_media_type=None, renderer_context=None): + request = renderer_context.get("request") + tenant_id = getattr(request, "tenant_id", None) if request else None + include_param_present = "include" in request.query_params if request else False + + # Use tenant_transaction if needed for included resources, otherwise do nothing + context_manager = ( + tenant_transaction(tenant_id) + if tenant_id and include_param_present + else nullcontext() + ) + with context_manager: + return super().render(data, accepted_media_type, renderer_context) diff --git a/src/backend/api/tests/test_views.py b/src/backend/api/tests/test_views.py index 13c9f970bc..0312d6aaea 100644 --- a/src/backend/api/tests/test_views.py +++ b/src/backend/api/tests/test_views.py @@ -1710,6 +1710,35 @@ class TestResourceViewSet: assert response.status_code == status.HTTP_200_OK assert len(response.json()["data"]) == len(resources_fixture) + @pytest.mark.parametrize( + "include_values, expected_resources", + [ + ("provider", ["Provider"]), + ("findings", ["Finding"]), + ("provider,findings", ["Provider", "Finding"]), + ], + ) + def test_resources_list_include( + self, + include_values, + expected_resources, + authenticated_client, + resources_fixture, + findings_fixture, + ): + response = authenticated_client.get( + reverse("resource-list"), {"include": include_values} + ) + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == len(resources_fixture) + assert "included" in response.json() + + included_data = response.json()["included"] + for expected_type in expected_resources: + assert any( + d.get("type") == expected_type for d in included_data + ), f"Expected type '{expected_type}' not found in included data" + @pytest.mark.parametrize( "filter_name, filter_value, expected_count", ( @@ -1867,6 +1896,30 @@ class TestFindingViewSet: == findings_fixture[0].status ) + @pytest.mark.parametrize( + "include_values, expected_resources", + [ + ("resources", ["Resource"]), + ("scan", ["Scan"]), + ("resources.provider,scan", ["Resource", "Scan", "Provider"]), + ], + ) + def test_findings_list_include( + self, include_values, expected_resources, authenticated_client, findings_fixture + ): + response = authenticated_client.get( + reverse("finding-list"), {"include": include_values} + ) + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == len(findings_fixture) + assert "included" in response.json() + + included_data = response.json()["included"] + for expected_type in expected_resources: + assert any( + d.get("type") == expected_type for d in included_data + ), f"Expected type '{expected_type}' not found in included data" + @pytest.mark.parametrize( "filter_name, filter_value, expected_count", ( diff --git a/src/backend/api/v1/serializers.py b/src/backend/api/v1/serializers.py index 66bf0784eb..2423f9d5e1 100644 --- a/src/backend/api/v1/serializers.py +++ b/src/backend/api/v1/serializers.py @@ -6,7 +6,6 @@ from django.contrib.auth.password_validation import validate_password from drf_spectacular.utils import extend_schema_field from jwt.exceptions import InvalidKeyError from rest_framework_json_api import serializers -from rest_framework_json_api.relations import SerializerMethodResourceRelatedField from rest_framework_json_api.serializers import ValidationError from rest_framework_simplejwt.exceptions import TokenError from rest_framework_simplejwt.tokens import RefreshToken @@ -20,10 +19,7 @@ from api.models import ( Task, Resource, ResourceTag, - SeverityChoices, - StatusChoices, Finding, - ResourceFindingMapping, ProviderSecret, ) from api.rls import Tenant @@ -510,9 +506,7 @@ class ResourceSerializer(RLSSerializer): tags = serializers.SerializerMethodField() type_ = serializers.CharField(read_only=True) - findings = SerializerMethodResourceRelatedField( - method_name="get_findings", many=True, read_only=True - ) + findings = serializers.ResourceRelatedField(many=True, read_only=True) class Meta: model = Resource @@ -538,12 +532,9 @@ class ResourceSerializer(RLSSerializer): included_serializers = { "findings": "api.v1.serializers.FindingSerializer", + "provider": "api.v1.serializers.ProviderSerializer", } - def get_findings(self, obj): - mappings = ResourceFindingMapping.objects.filter(resource=obj) - return Finding.objects.filter(id__in=[m.finding_id for m in mappings]) - @extend_schema_field( { "type": "object", @@ -562,40 +553,12 @@ class ResourceSerializer(RLSSerializer): return fields -class FindingDeltaEnumSerializerField(serializers.ChoiceField): - def __init__(self, **kwargs): - kwargs["choices"] = Finding.DeltaChoices.choices - super().__init__(**kwargs) - - -class SeverityEnumSerializerField(serializers.ChoiceField): - def __init__(self, **kwargs): - kwargs["choices"] = SeverityChoices.choices - super().__init__(**kwargs) - - -class StatusEnumSerializerField(serializers.ChoiceField): - def __init__(self, **kwargs): - kwargs["choices"] = StatusChoices.choices - super().__init__(**kwargs) - - -class ResourceFindingMappingSerializer(serializers.ModelSerializer): - resource = ResourceSerializer(read_only=True) - - class Meta: - model = ResourceFindingMapping - fields = ["resource"] - - class FindingSerializer(RLSSerializer): """ Serializer for the Finding model. """ - resources = SerializerMethodResourceRelatedField( - method_name="get_resources", many=True, read_only=True - ) + resources = serializers.ResourceRelatedField(many=True, read_only=True) class Meta: model = Finding @@ -618,17 +581,10 @@ class FindingSerializer(RLSSerializer): ] included_serializers = { - "scan": "api.v1.serializers.ScanSerializer", - "resources": "api.v1.serializers.ResourceSerializer", + "scan": ScanSerializer, + "resources": ResourceSerializer, } - class JSONAPIMeta: - resource_name = "Findings" - - def get_resources(self, obj): - mappings = ResourceFindingMapping.objects.filter(finding=obj) - return Resource.objects.filter(id__in={m.resource_id for m in mappings}) - # Provider secrets class BaseWriteProviderSecretSerializer(BaseWriteSerializer): diff --git a/src/backend/api/v1/views.py b/src/backend/api/v1/views.py index a05e7b42f1..a921bbfcdd 100644 --- a/src/backend/api/v1/views.py +++ b/src/backend/api/v1/views.py @@ -3,6 +3,7 @@ from django.conf import settings as django_settings from django.contrib.postgres.search import SearchQuery from django.db import transaction from django.db.models import F, Q +from django.db.models import Prefetch from django.urls import reverse from django.utils.decorators import method_decorator from django.views.decorators.cache import cache_control @@ -22,6 +23,7 @@ from rest_framework.generics import get_object_or_404, GenericAPIView from rest_framework_json_api.views import Response from rest_framework_simplejwt.exceptions import InvalidToken from rest_framework_simplejwt.exceptions import TokenError + from api.base_views import BaseTenantViewset, BaseRLSViewSet, BaseViewSet from api.filters import ( ProviderFilter, @@ -742,6 +744,13 @@ class ResourceViewSet(BaseRLSViewSet): class FindingViewSet(BaseRLSViewSet): queryset = Finding.objects.all() serializer_class = FindingSerializer + prefetch_for_includes = { + "__all__": [], + "resources": [ + Prefetch("resources", queryset=Resource.objects.select_related("findings")) + ], + "scan": [Prefetch("scan", queryset=Scan.objects.select_related("findings"))], + } http_method_names = ["get"] filterset_class = FindingFilter ordering = ["-id"] @@ -760,8 +769,6 @@ class FindingViewSet(BaseRLSViewSet): return datetime_to_uuid7(inserted_at) def get_queryset(self): - # TODO: require scan_id filter, or if none provided, inject today - queryset = Finding.objects.all() search_value = self.request.query_params.get("filter[search]", None) diff --git a/src/backend/config/django/base.py b/src/backend/config/django/base.py index 8866f93c49..28ffa76dbc 100644 --- a/src/backend/config/django/base.py +++ b/src/backend/config/django/base.py @@ -76,7 +76,7 @@ REST_FRAMEWORK = { "rest_framework.parsers.FormParser", "rest_framework.parsers.MultiPartParser", ), - "DEFAULT_RENDERER_CLASSES": ("rest_framework_json_api.renderers.JSONRenderer",), + "DEFAULT_RENDERER_CLASSES": ("api.renderers.APIJSONRenderer",), "DEFAULT_METADATA_CLASS": "rest_framework_json_api.metadata.JSONAPIMetadata", "DEFAULT_FILTER_BACKENDS": ( "rest_framework_json_api.filters.QueryParameterValidationFilter",