diff --git a/.gitignore b/.gitignore index c294aa24c0..a215af5677 100644 --- a/.gitignore +++ b/.gitignore @@ -125,6 +125,7 @@ celerybeat.pid # Environments .env +*.env .venv env/ venv/ diff --git a/src/backend/api/filters.py b/src/backend/api/filters.py index 8e8c548228..75fe085996 100644 --- a/src/backend/api/filters.py +++ b/src/backend/api/filters.py @@ -1,3 +1,4 @@ +from django.db.models import Q from django_filters.rest_framework import ( FilterSet, BooleanFilter, @@ -8,7 +9,7 @@ from rest_framework_json_api.django_filters.backends import DjangoFilterBackend from rest_framework_json_api.serializers import ValidationError from api.db_utils import ProviderEnumField -from api.models import Provider, Scan, Task, StateChoices +from api.models import Provider, Resource, ResourceTag, Scan, Task, StateChoices from api.rls import Tenant from api.v1.serializers import TaskBase @@ -128,7 +129,7 @@ class ScanFilter(FilterSet): model = Scan fields = { "provider": ["exact"], - "provider_id": ["exact"], + "provider_id": ["exact", "in"], "name": ["exact", "icontains"], "started_at": ["gte", "lte"], "trigger": ["exact"], @@ -160,3 +161,56 @@ class TaskFilter(FilterSet): class Meta: model = Task fields = [] + + +class ResourceTagFilter(FilterSet): + class Meta: + model = ResourceTag + fields = { + "key": ["exact", "icontains"], + "value": ["exact", "icontains"], + } + search = ["text_search"] + + +class ResourceFilter(FilterSet): + provider = CharFilter(method="filter_provider") + tag_key = CharFilter(method="filter_tag_key") + tag_value = CharFilter(method="filter_tag_value") + tag = CharFilter(method="filter_tag") + tags = CharFilter(method="filter_tag") + inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") + updated_at = DateFilter(field_name="updated_at", lookup_expr="date") + + def filter_provider(self, queryset, name, value): + return enum_filter( + queryset, + value, + enum_choices=Provider.ProviderChoices, + lookup_field="provider__provider", + ) + + class Meta: + model = Resource + fields = { + "provider_id": ["exact", "in"], + "uid": ["exact", "icontains"], + "name": ["exact", "icontains"], + "region": ["exact", "icontains", "in"], + "service": ["exact", "icontains", "in"], + "type": ["exact", "icontains", "in"], + "inserted_at": ["gte", "lte"], + "updated_at": ["gte", "lte"], + } + + def filter_tag_key(self, queryset, name, value): + return queryset.filter(Q(tags__key=value) | Q(tags__key__icontains=value)) + + def filter_tag_value(self, queryset, name, value): + return queryset.filter(Q(tags__value=value) | Q(tags__value__icontains=value)) + + def filter_tag(self, queryset, name, value): + # we won't know what the user wants to filter on just based on the value + # and we don't want to build special filtering logic for every possible + # provider tag spec, so we'll just do a full text search + return queryset.filter(tags__text_search=value) diff --git a/src/backend/api/fixtures/2_dev_resources.json b/src/backend/api/fixtures/2_dev_resources.json new file mode 100644 index 0000000000..23dc85ed93 --- /dev/null +++ b/src/backend/api/fixtures/2_dev_resources.json @@ -0,0 +1,45 @@ +[ + { + "model": "api.resource", + "pk": "a3ba9470-a240-49a6-8196-9230a267a220", + "fields": { + "tenant": "12646005-9067-4d2a-a098-8bb378604362", + "provider": "37b065f8-26b0-4218-a665-0b23d07b27d9", + "uid": "unique-1", + "name": "testing 1", + "inserted_at": "2024-08-01T17:20:27.050Z", + "updated_at": "2024-08-01T17:20:27.050Z" + } + }, + { + "model": "api.resource", + "pk": "85f18c25-4deb-460e-87e2-12548f2508ed", + "fields": { + "tenant": "12646005-9067-4d2a-a098-8bb378604362", + "provider": "37b065f8-26b0-4218-a665-0b23d07b27d9", + "uid": "unique-2", + "name": "testing 2", + "inserted_at": "2024-08-01T17:20:27.050Z", + "updated_at": "2024-08-01T17:20:27.050Z" + } + }, + { + "model": "api.resourcetag", + "pk": "057c38c5-94aa-46ee-98bb-9ec5b0886bbf", + "fields": { + "tenant": "12646005-9067-4d2a-a098-8bb378604362", + "key": "key", + "value": "tag value", + "inserted_at": "2024-08-01T17:20:27.050Z", + "updated_at": "2024-08-01T17:20:27.050Z" + } + }, + { + "model": "api.resourcetagmapping", + "fields": { + "tag": "057c38c5-94aa-46ee-98bb-9ec5b0886bbf", + "resource": "85f18c25-4deb-460e-87e2-12548f2508ed", + "tenant": "12646005-9067-4d2a-a098-8bb378604362" + } + } +] diff --git a/src/backend/api/migrations/0001_initial.py b/src/backend/api/migrations/0001_initial.py index 65124f869a..8fb82ee65e 100644 --- a/src/backend/api/migrations/0001_initial.py +++ b/src/backend/api/migrations/0001_initial.py @@ -346,4 +346,270 @@ class Migration(migrations.Migration): GRANT SELECT ON TABLE {TASK_RUNNER_DB_TABLE} TO {DB_PROWLER_USER}; """ ), + # Resources + migrations.RunSQL( + sql=""" + CREATE EXTENSION IF NOT EXISTS pg_trgm; + """, + reverse_sql=""" + DROP EXTENSION IF EXISTS pg_trgm; + """, + ), + migrations.CreateModel( + name="Resource", + fields=[ + ( + "id", + models.UUIDField( + default=uuid.uuid4, + editable=False, + primary_key=True, + serialize=False, + ), + ), + ("inserted_at", models.DateTimeField(auto_now_add=True)), + ("updated_at", models.DateTimeField(auto_now=True)), + ( + "uid", + models.TextField( + verbose_name="Unique identifier for the resource, set by the provider" + ), + ), + ( + "name", + models.TextField( + verbose_name="Name of the resource, as set in the provider" + ), + ), + ( + "region", + models.TextField( + verbose_name="Location of the resource, as set by the provider" + ), + ), + ( + "service", + models.TextField( + verbose_name="Service of the resource, as set by the provider" + ), + ), + ( + "type", + models.TextField( + verbose_name="Type of the resource, as set by the provider" + ), + ), + ( + "text_search", + models.GeneratedField( + db_persist=True, + expression=django.contrib.postgres.search.CombinedSearchVector( + django.contrib.postgres.search.CombinedSearchVector( + django.contrib.postgres.search.CombinedSearchVector( + django.contrib.postgres.search.SearchVector( + "uid", config="simple", weight="A" + ), + "||", + django.contrib.postgres.search.SearchVector( + "name", config="simple", weight="B" + ), + django.contrib.postgres.search.SearchConfig( + "simple" + ), + ), + "||", + django.contrib.postgres.search.SearchVector( + "region", config="simple", weight="C" + ), + django.contrib.postgres.search.SearchConfig("simple"), + ), + "||", + django.contrib.postgres.search.SearchVector( + "service", "type", config="simple", weight="D" + ), + django.contrib.postgres.search.SearchConfig("simple"), + ), + null=True, + output_field=django.contrib.postgres.search.SearchVectorField(), + ), + ), + ( + "provider", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + related_name="resources", + related_query_name="resource", + to="api.provider", + ), + ), + ( + "tenant", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, to="api.tenant" + ), + ), + ], + options={ + "db_table": "resources", + "abstract": False, + }, + ), + migrations.CreateModel( + name="ResourceTag", + fields=[ + ( + "id", + models.UUIDField( + default=uuid.uuid4, + editable=False, + primary_key=True, + serialize=False, + ), + ), + ("inserted_at", models.DateTimeField(auto_now_add=True)), + ("updated_at", models.DateTimeField(auto_now=True)), + ("key", models.TextField()), + ("value", models.TextField()), + ( + "text_search", + models.GeneratedField( + db_persist=True, + expression=django.contrib.postgres.search.CombinedSearchVector( + django.contrib.postgres.search.SearchVector( + "key", config="simple", weight="A" + ), + "||", + django.contrib.postgres.search.SearchVector( + "value", config="simple", weight="B" + ), + django.contrib.postgres.search.SearchConfig("simple"), + ), + null=True, + output_field=django.contrib.postgres.search.SearchVectorField(), + ), + ), + ( + "tenant", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, to="api.tenant" + ), + ), + ], + options={ + "db_table": "resource_tags", + "abstract": False, + }, + ), + migrations.CreateModel( + name="ResourceTagMapping", + fields=[ + ( + "id", + models.UUIDField( + default=uuid.uuid4, + editable=False, + primary_key=True, + serialize=False, + ), + ), + ( + "resource", + models.ForeignKey( + on_delete=django.db.models.deletion.DO_NOTHING, + to="api.resource", + ), + ), + ( + "tag", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + to="api.resourcetag", + ), + ), + ( + "tenant", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + to="api.tenant", + ), + ), + ], + options={ + "db_table": "resource_tag_mappings", + "abstract": False, + }, + ), + migrations.AddField( + model_name="resource", + name="tags", + field=models.ManyToManyField( + through="api.ResourceTagMapping", + to="api.resourcetag", + verbose_name="Tags associated with the resource, by provider", + ), + ), + migrations.AddIndex( + model_name="resourcetag", + index=django.contrib.postgres.indexes.GinIndex( + fields=["text_search"], name="gin_resource_tags_search_idx" + ), + ), + migrations.AddIndex( + model_name="resource", + index=django.contrib.postgres.indexes.GinIndex( + fields=["text_search"], name="gin_resources_search_idx" + ), + ), + migrations.AddConstraint( + model_name="resourcetag", + constraint=models.UniqueConstraint( + fields=("tenant_id", "key", "value"), + name="unique_resource_tags_by_tenant_key_value", + ), + ), + migrations.AddConstraint( + model_name="resourcetag", + constraint=api.rls.RowLevelSecurityConstraint( + "tenant_id", + name="rls_on_resourcetag", + statements=["SELECT"], + ), + ), + migrations.AddConstraint( + model_name="resourcetagmapping", + constraint=models.UniqueConstraint( + fields=("tenant_id", "resource_id", "tag_id"), + name="unique_resource_tag_mappings_by_tenant_resource_tag", + ), + ), + migrations.AddConstraint( + model_name="resourcetagmapping", + constraint=api.rls.RowLevelSecurityConstraint( + "tenant_id", + name="rls_on_resourcetagmapping", + statements=["SELECT"], + ), + ), + migrations.AddIndex( + model_name="resource", + index=models.Index( + fields=["uid", "region", "service", "name"], + name="idx_resource_uid_reg_serv_name", + ), + ), + migrations.AddConstraint( + model_name="resource", + constraint=models.UniqueConstraint( + fields=("tenant_id", "provider_id", "uid"), + name="unique_resources_by_provider", + ), + ), + migrations.AddConstraint( + model_name="resource", + constraint=api.rls.RowLevelSecurityConstraint( + "tenant_id", + name="rls_on_resource", + statements=["SELECT"], + ), + ), ] diff --git a/src/backend/api/models.py b/src/backend/api/models.py index 2984adb399..6e029b28fd 100644 --- a/src/backend/api/models.py +++ b/src/backend/api/models.py @@ -1,6 +1,8 @@ import re from uuid import uuid4, UUID +from django.contrib.postgres.indexes import GinIndex +from django.contrib.postgres.search import SearchVector, SearchVectorField from django.core.validators import MinLengthValidator from django.db import models from django.utils.translation import gettext_lazy as _ @@ -155,7 +157,7 @@ class Scan(RowLevelSecurityProtectedModel): indexes = [ models.Index( fields=["provider", "state", "trigger", "scheduled_at"], - name="scans_prov_state_type_sche_idx", + name="scans_prov_state_trig_sche_idx", ), ] @@ -189,3 +191,155 @@ class Task(RowLevelSecurityProtectedModel): name="tasks_id_trt_id_idx", ), ] + + +class ResourceTag(RowLevelSecurityProtectedModel): + id = models.UUIDField(primary_key=True, default=uuid4, editable=False) + inserted_at = models.DateTimeField(auto_now_add=True, editable=False) + updated_at = models.DateTimeField(auto_now=True, editable=False) + + key = models.TextField(blank=False) + value = models.TextField(blank=False) + + text_search = models.GeneratedField( + expression=SearchVector("key", weight="A", config="simple") + + SearchVector("value", weight="B", config="simple"), + output_field=SearchVectorField(), + db_persist=True, + null=True, + editable=False, + ) + + class Meta(RowLevelSecurityProtectedModel.Meta): + db_table = "resource_tags" + + indexes = [ + GinIndex(fields=["text_search"], name="gin_resource_tags_search_idx"), + ] + + constraints = [ + models.UniqueConstraint( + fields=("tenant_id", "key", "value"), + name="unique_resource_tags_by_tenant_key_value", + ), + RowLevelSecurityConstraint( + field="tenant_id", + name="rls_on_%(class)s", + statements=["SELECT"], + ), + ] + + +class Resource(RowLevelSecurityProtectedModel): + id = models.UUIDField(primary_key=True, default=uuid4, editable=False) + inserted_at = models.DateTimeField(auto_now_add=True, editable=False) + updated_at = models.DateTimeField(auto_now=True, editable=False) + + provider = models.ForeignKey( + Provider, + on_delete=models.CASCADE, + related_name="resources", + related_query_name="resource", + ) + + uid = models.TextField( + "Unique identifier for the resource, set by the provider", blank=False + ) + name = models.TextField("Name of the resource, as set in the provider", blank=False) + region = models.TextField( + "Location of the resource, as set by the provider", blank=False + ) + service = models.TextField( + "Service of the resource, as set by the provider", blank=False + ) + type = models.TextField("Type of the resource, as set by the provider", blank=False) + + text_search = models.GeneratedField( + expression=SearchVector("uid", weight="A", config="simple") + + SearchVector("name", weight="B", config="simple") + + SearchVector("region", weight="C", config="simple") + + SearchVector("service", "type", weight="D", config="simple"), + output_field=SearchVectorField(), + db_persist=True, + null=True, + editable=False, + ) + + tags = models.ManyToManyField( + ResourceTag, + verbose_name="Tags associated with the resource, by provider", + through="ResourceTagMapping", + ) + + def get_tags(self) -> dict: + return {tag.key: tag.value for tag in self.tags.all()} + + def clear_tags(self): + self.tags.clear() + self.save() + + def upsert_or_delete_tags(self, tags: list[ResourceTag] | None): + if tags is None: + self.clear_tags() + return + + # Add new relationships with the tenant_id field + for tag in tags: + ResourceTagMapping.objects.update_or_create( + tag=tag, resource=self, tenant_id=self.tenant_id + ) + + # Save the instance + self.save() + + class Meta(RowLevelSecurityProtectedModel.Meta): + db_table = "resources" + + indexes = [ + models.Index( + fields=["uid", "region", "service", "name"], + name="idx_resource_uid_reg_serv_name", + ), + GinIndex(fields=["text_search"], name="gin_resources_search_idx"), + ] + + constraints = [ + models.UniqueConstraint( + fields=("tenant_id", "provider_id", "uid"), + name="unique_resources_by_provider", + ), + RowLevelSecurityConstraint( + field="tenant_id", + name="rls_on_%(class)s", + statements=["SELECT"], + ), + ] + + +class ResourceTagMapping(RowLevelSecurityProtectedModel): + # NOTE that we don't really need a primary key here, + # but everything is easier with django if we do + id = models.UUIDField(primary_key=True, default=uuid4, editable=False) + resource = models.ForeignKey(Resource, on_delete=models.DO_NOTHING) + tag = models.ForeignKey(ResourceTag, on_delete=models.CASCADE) + + class Meta(RowLevelSecurityProtectedModel.Meta): + db_table = "resource_tag_mappings" + + # django will automatically create indexes for: + # - resource_id + # - tag_id + # - tenant_id + # - id + + constraints = [ + models.UniqueConstraint( + fields=("tenant_id", "resource_id", "tag_id"), + name="unique_resource_tag_mappings_by_tenant_resource_tag", + ), + RowLevelSecurityConstraint( + field="tenant_id", + name="rls_on_%(class)s", + statements=["SELECT"], + ), + ] diff --git a/src/backend/api/tests/test_models.py b/src/backend/api/tests/test_models.py new file mode 100644 index 0000000000..c7fdf9deb1 --- /dev/null +++ b/src/backend/api/tests/test_models.py @@ -0,0 +1,89 @@ +import pytest + +from api.models import Resource, ResourceTag + + +@pytest.mark.django_db +class TestResourceModel: + def test_setting_tags(self, providers_fixture): + provider, *_ = providers_fixture + + resource = Resource.objects.create( + tenant_id=provider.tenant_id, + provider=provider, + uid="arn:aws:ec2:us-east-1:123456789012:instance/i-1234567890abcdef0", + name="My Instance 1", + region="us-east-1", + service="ec2", + type="prowler-test", + ) + + tags = [ + ResourceTag.objects.create( + tenant_id=provider.tenant_id, + key="key", + value="value", + ), + ResourceTag.objects.create( + tenant_id=provider.tenant_id, + key="key2", + value="value2", + ), + ] + + resource.upsert_or_delete_tags(tags) + + assert len(tags) == len(resource.tags.all()) + + tags_dict = resource.get_tags() + + for tag in tags: + assert tag.key in tags_dict + assert tag.value == tags_dict[tag.key] + + def test_adding_tags(self, resources_fixture): + resource, *_ = resources_fixture + + tags = [ + ResourceTag.objects.create( + tenant_id=resource.tenant_id, + key="env", + value="test", + ), + ] + before_count = len(resource.tags.all()) + + resource.upsert_or_delete_tags(tags) + + assert before_count + 1 == len(resource.tags.all()) + + tags_dict = resource.get_tags() + + assert "env" in tags_dict + assert tags_dict["env"] == "test" + + def test_adding_duplicate_tags(self, resources_fixture): + resource, *_ = resources_fixture + + tags = resource.tags.all() + + before_count = len(resource.tags.all()) + + resource.upsert_or_delete_tags(tags) + + # should be the same number of tags + assert before_count == len(resource.tags.all()) + + def test_add_tags_none(self, resources_fixture): + resource, *_ = resources_fixture + resource.upsert_or_delete_tags(None) + + assert len(resource.tags.all()) == 0 + assert resource.get_tags() == {} + + def test_clear_tags(self, resources_fixture): + resource, *_ = resources_fixture + resource.clear_tags() + + assert len(resource.tags.all()) == 0 + assert resource.get_tags() == {} diff --git a/src/backend/api/tests/test_views.py b/src/backend/api/tests/test_views.py index 49780eee01..ee50c87f3f 100644 --- a/src/backend/api/tests/test_views.py +++ b/src/backend/api/tests/test_views.py @@ -776,6 +776,31 @@ class TestScanViewSet: ) assert response.status_code == status.HTTP_400_BAD_REQUEST + def test_scan_filter_by_provider_id_exact( + self, client, scans_fixture, tenant_header + ): + response = client.get( + reverse("scan-list"), + {"filter[provider_id]": scans_fixture[0].provider.id}, + headers=tenant_header, + ) + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == 2 + + def test_scan_filter_by_provider_id_in(self, client, scans_fixture, tenant_header): + response = client.get( + reverse("scan-list"), + { + "filter[provider_id.in]": [ + scans_fixture[0].provider.id, + scans_fixture[1].provider.id, + ] + }, + headers=tenant_header, + ) + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == 2 + @pytest.mark.parametrize( "sort_field", [ @@ -849,3 +874,156 @@ class TestTaskViewSet: ) # Task status is SUCCESS assert response.status_code == status.HTTP_400_BAD_REQUEST + + +@pytest.mark.django_db +class TestResourceViewSet: + def test_resources_list_none(self, client, tenant_header): + response = client.get(reverse("resource-list"), headers=tenant_header) + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == 0 + + def test_resources_list(self, client, resources_fixture, tenant_header): + response = client.get(reverse("resource-list"), headers=tenant_header) + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == len(resources_fixture) + assert ( + response.json()["data"][0]["attributes"]["uid"] == resources_fixture[0].uid + ) + + @pytest.mark.parametrize( + "filter_name, filter_value, expected_count", + ( + [ + ( + "uid", + "arn:aws:ec2:us-east-1:123456789012:instance/i-1234567890abcdef0", + 1, + ), + ("uid.icontains", "i-1234567890abcdef", 3), + ("name", "My Instance 2", 1), + ("name.icontains", "ce 2", 1), + ("region", "eu-west-1", 1), + ("region.icontains", "west", 1), + ("service", "ec2", 2), + ("service.icontains", "ec", 2), + ("inserted_at.gte", "2024-01-01 00:00:00", 3), + ("updated_at.lte", "2024-01-01 00:00:00", 0), + ("type.icontains", "prowler", 2), + # tags searching + ("tag", "key3:value:value", 0), + ("tag_key", "key3", 1), + ("tag_value", "value2", 2), + ("tag", "key3:multi word value3", 1), + ("tags", "key3:multi word value3", 1), + ("tags", "multi word", 1), + # full text search on resource + ("search", "arn", 3), + ("search", "def1", 1), + # full text search on resource tags + ("search", "multi word", 1), + ("search", "key2", 2), + ] + ), + ) + def test_resource_filters( + self, + client, + resources_fixture, + tenant_header, + filter_name, + filter_value, + expected_count, + ): + response = client.get( + reverse("resource-list"), + {f"filter[{filter_name}]": filter_value}, + headers=tenant_header, + ) + + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == expected_count + + def test_resource_filter_by_provider_id_in( + self, client, resources_fixture, tenant_header + ): + response = client.get( + reverse("resource-list"), + { + "filter[provider_id.in]": [ + resources_fixture[0].provider.id, + resources_fixture[1].provider.id, + ] + }, + headers=tenant_header, + ) + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == 2 + + @pytest.mark.parametrize( + "filter_name", + ( + [ + "resource", # Invalid filter name + "invalid", + ] + ), + ) + def test_resources_filters_invalid(self, client, tenant_header, filter_name): + response = client.get( + reverse("resource-list"), + {f"filter[{filter_name}]": "whatever"}, + headers=tenant_header, + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + + @pytest.mark.parametrize( + "sort_field", + [ + "provider_id", + "uid", + "name", + "region", + "service", + "type", + "inserted_at", + "updated_at", + ], + ) + def test_resources_sort(self, client, tenant_header, sort_field): + response = client.get( + reverse("resource-list"), {"sort": sort_field}, headers=tenant_header + ) + assert response.status_code == status.HTTP_200_OK + + def test_resources_sort_invalid(self, client, tenant_header): + response = client.get( + reverse("resource-list"), {"sort": "invalid"}, headers=tenant_header + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert response.json()["errors"][0]["code"] == "invalid" + assert response.json()["errors"][0]["source"]["pointer"] == "/data" + assert ( + response.json()["errors"][0]["detail"] == "invalid sort parameter: invalid" + ) + + def test_resources_retrieve(self, client, resources_fixture, tenant_header): + resource_1, *_ = resources_fixture + response = client.get( + reverse("resource-detail", kwargs={"pk": resource_1.id}), + headers=tenant_header, + ) + assert response.status_code == status.HTTP_200_OK + assert response.json()["data"]["attributes"]["uid"] == resource_1.uid + assert response.json()["data"]["attributes"]["name"] == resource_1.name + assert response.json()["data"]["attributes"]["region"] == resource_1.region + assert response.json()["data"]["attributes"]["service"] == resource_1.service + assert response.json()["data"]["attributes"]["type"] == resource_1.type + assert response.json()["data"]["attributes"]["tags"] == resource_1.get_tags() + + def test_resources_invalid_retrieve(self, client, tenant_header): + response = client.get( + reverse("resource-detail", kwargs={"pk": "random_id"}), + headers=tenant_header, + ) + assert response.status_code == status.HTTP_404_NOT_FOUND diff --git a/src/backend/api/v1/serializers.py b/src/backend/api/v1/serializers.py index f38e830d60..c8ab81bb3c 100644 --- a/src/backend/api/v1/serializers.py +++ b/src/backend/api/v1/serializers.py @@ -4,7 +4,7 @@ from drf_spectacular.utils import extend_schema_field from rest_framework_json_api import serializers from rest_framework_json_api.serializers import ValidationError -from api.models import StateChoices, Provider, Scan, Task +from api.models import StateChoices, Provider, Scan, Task, Resource, ResourceTag from api.rls import Tenant from api.utils import merge_dicts @@ -276,3 +276,59 @@ class ScanUpdateSerializer(BaseWriteSerializer): extra_kwargs = { "id": {"read_only": True}, } + + +class ResourceTagSerializer(RLSSerializer): + """ + Serializer fore the ResourceTag model + """ + + class Meta: + model = ResourceTag + fields = ["key", "value"] + + +class ResourceSerializer(RLSSerializer): + """ + Serializer for the Resource model. + """ + + tags = serializers.SerializerMethodField() + type_ = serializers.CharField(read_only=True) + + class Meta: + model = Resource + fields = [ + "id", + "inserted_at", + "updated_at", + "uid", + "name", + "region", + "service", + "type_", + "tags", + "provider", + ] + extra_kwargs = { + "id": {"read_only": True}, + "inserted_at": {"read_only": True}, + "updated_at": {"read_only": True}, + } + + @extend_schema_field( + { + "type": "object", + "description": "Tags associated with the resource", + "example": {"env": "prod", "owner": "johndoe"}, + } + ) + def get_tags(self, obj): + return obj.get_tags() + + def get_fields(self): + """`type` is a Python reserved keyword.""" + fields = super().get_fields() + type_ = fields.pop("type_") + fields["type"] = type_ + return fields diff --git a/src/backend/api/v1/urls.py b/src/backend/api/v1/urls.py index 7e7d1e9c07..36168b87e3 100644 --- a/src/backend/api/v1/urls.py +++ b/src/backend/api/v1/urls.py @@ -8,6 +8,7 @@ from api.v1.views import ( ProviderViewSet, ScanViewSet, TaskViewSet, + ResourceViewSet, ) router = routers.DefaultRouter(trailing_slash=False) @@ -16,6 +17,7 @@ router.register(r"tenants", TenantViewSet, basename="tenant") router.register(r"providers", ProviderViewSet, basename="provider") router.register(r"scans", ScanViewSet, basename="scan") router.register(r"tasks", TaskViewSet, basename="task") +router.register(r"resources", ResourceViewSet, basename="resource") urlpatterns = [ path("", include(router.urls)), diff --git a/src/backend/api/v1/views.py b/src/backend/api/v1/views.py index fd2f1e5aad..90c8eb2ab9 100644 --- a/src/backend/api/v1/views.py +++ b/src/backend/api/v1/views.py @@ -3,6 +3,9 @@ from django.db.models import F from django.urls import reverse from django.utils.decorators import method_decorator from django.views.decorators.cache import cache_control +from django.contrib.postgres.search import SearchQuery +from django.db.models import Q + from drf_spectacular.settings import spectacular_settings from drf_spectacular.utils import extend_schema, extend_schema_view from drf_spectacular.views import SpectacularAPIView @@ -13,8 +16,14 @@ from rest_framework_json_api.views import Response from celery.result import AsyncResult from api.base_views import BaseRLSViewSet, BaseViewSet -from api.filters import ProviderFilter, TenantFilter, ScanFilter, TaskFilter -from api.models import Provider, Scan, Task +from api.filters import ( + ProviderFilter, + TenantFilter, + ScanFilter, + TaskFilter, + ResourceFilter, +) +from api.models import Provider, Scan, Task, Resource from api.rls import Tenant from api.v1.serializers import ( ProviderSerializer, @@ -26,6 +35,7 @@ from api.v1.serializers import ( ScanSerializer, ScanCreateSerializer, ScanUpdateSerializer, + ResourceSerializer, ) from tasks.tasks import check_provider_connection_task, delete_provider_task @@ -331,3 +341,62 @@ class TaskViewSet(BaseRLSViewSet): "Content-Location": reverse("task-detail", kwargs={"pk": task.id}) }, ) + + +@extend_schema_view( + list=extend_schema( + summary="List all resources", + description="Retrieve a list of all resources with options for filtering by various criteria. Resources are objects that are discovered by Prowler. They can be anything from a single host to a whole VPC.", + ), + retrieve=extend_schema( + summary="Retrieve data for a resource", + description="Fetch detailed information about a specific resource by their ID. A Resource is an object that is discovered by Prowler. It can be anything from a single host to a whole VPC.", + ), +) +@method_decorator(CACHE_DECORATOR, name="list") +@method_decorator(CACHE_DECORATOR, name="retrieve") +class ResourceViewSet(BaseRLSViewSet): + queryset = Resource.objects.all() + serializer_class = ResourceSerializer + http_method_names = ["get"] + filterset_class = ResourceFilter + ordering = ["inserted_at"] + ordering_fields = [ + "provider_id", + "uid", + "name", + "region", + "service", + "type", + "inserted_at", + "updated_at", + ] + + def get_queryset(self): + queryset = Resource.objects.all() + search_value = self.request.query_params.get("filter[search]", None) + + if search_value: + # Django's ORM will build a LEFT JOIN and OUTER JOIN on the "through" table, resulting in duplicates + # The duplicates then require a `distinct` query + search_query = SearchQuery( + search_value, config="simple", search_type="plain" + ) + queryset = queryset.filter( + Q(tags__key=search_value) + | Q(tags__value=search_value) + | Q(tags__text_search=search_query) + | Q(tags__key__contains=search_value) + | Q(tags__value__contains=search_value) + | Q(uid=search_value) + | Q(name=search_value) + | Q(region=search_value) + | Q(service=search_value) + | Q(text_search=search_query) + | Q(uid__contains=search_value) + | Q(name__contains=search_value) + | Q(region__contains=search_value) + | Q(service__contains=search_value) + ).distinct() + + return queryset diff --git a/src/backend/config/django/base.py b/src/backend/config/django/base.py index d1e833fbdf..d2dc736239 100644 --- a/src/backend/config/django/base.py +++ b/src/backend/config/django/base.py @@ -15,6 +15,7 @@ INSTALLED_APPS = [ "django.contrib.sessions", "django.contrib.messages", "django.contrib.staticfiles", + "django.contrib.postgres", "api", "rest_framework", "corsheaders", diff --git a/src/backend/conftest.py b/src/backend/conftest.py index a113eed7b1..d71a164496 100644 --- a/src/backend/conftest.py +++ b/src/backend/conftest.py @@ -5,7 +5,7 @@ from django.conf import settings from django.db import connections as django_connections from rest_framework import status from django_celery_results.models import TaskResult -from api.models import Provider, Scan, StateChoices, Task +from api.models import Provider, Resource, ResourceTag, Scan, StateChoices, Task from api.rls import Tenant API_JSON_CONTENT_TYPE = "application/vnd.api+json" @@ -95,7 +95,7 @@ def providers_fixture(tenants_fixture): @pytest.fixture def scans_fixture(tenants_fixture, providers_fixture): tenant, _ = tenants_fixture - provider, *_ = providers_fixture + provider, provider2, *_ = providers_fixture scan1 = Scan.objects.create( name="Scan 1", @@ -115,7 +115,7 @@ def scans_fixture(tenants_fixture, providers_fixture): ) scan3 = Scan.objects.create( name="Scan 3", - provider=provider, + provider=provider2, trigger=Scan.TriggerChoices.SCHEDULED, state=StateChoices.AVAILABLE, tenant_id=tenant.id, @@ -152,6 +152,68 @@ def tasks_fixture(tenants_fixture): return task1, task2 +@pytest.fixture +def resources_fixture(providers_fixture): + provider, *_ = providers_fixture + + tags = [ + ResourceTag.objects.create( + tenant_id=provider.tenant_id, + key="key", + value="value", + ), + ResourceTag.objects.create( + tenant_id=provider.tenant_id, + key="key2", + value="value2", + ), + ] + + resource1 = Resource.objects.create( + tenant_id=provider.tenant_id, + provider=provider, + uid="arn:aws:ec2:us-east-1:123456789012:instance/i-1234567890abcdef0", + name="My Instance 1", + region="us-east-1", + service="ec2", + type="prowler-test", + ) + + resource1.upsert_or_delete_tags(tags) + + resource2 = Resource.objects.create( + tenant_id=provider.tenant_id, + provider=provider, + uid="arn:aws:ec2:us-east-1:123456789012:instance/i-1234567890abcdef1", + name="My Instance 2", + region="eu-west-1", + service="ec2", + type="prowler-test", + ) + resource2.upsert_or_delete_tags(tags) + + resource3 = Resource.objects.create( + tenant_id=providers_fixture[1].tenant_id, + provider=providers_fixture[1], + uid="arn:aws:ec2:us-east-1:123456789012:bucket/i-1234567890abcdef2", + name="My Bucket 3", + region="us-east-1", + service="s3", + type="test", + ) + + tags = [ + ResourceTag.objects.create( + tenant_id=provider.tenant_id, + key="key3", + value="multi word value3", + ), + ] + resource3.upsert_or_delete_tags(tags) + + return resource1, resource2, resource3 + + @pytest.fixture def tenant_header(tenants_fixture): return {"X-Tenant-ID": str(tenants_fixture[0].id)}