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
This commit is contained in:
Víctor Fernández Poyatos
2024-10-17 18:07:06 +02:00
committed by GitHub
parent 60c75b4814
commit 6d69a192f3
5 changed files with 91 additions and 52 deletions
+23
View File
@@ -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)
+53
View File
@@ -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",
(
+5 -49
View File
@@ -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):
+9 -2
View File
@@ -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)
+1 -1
View File
@@ -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",