mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-23 12:31:54 +00:00
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:
committed by
GitHub
parent
60c75b4814
commit
6d69a192f3
@@ -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)
|
||||
@@ -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",
|
||||
(
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user