diff --git a/src/backend/api/filters.py b/src/backend/api/filters.py index 75fe085996..f445e54965 100644 --- a/src/backend/api/filters.py +++ b/src/backend/api/filters.py @@ -40,6 +40,14 @@ def enum_filter(queryset, value, enum_choices, lookup_field: str): ValidationError: If the provided `value` is not a valid choice in the specified `enum_choices`. """ + if "__in" in lookup_field: + values = value + if isinstance(values, str): + values = value.split(",") + + values = [v for v in values if v in enum_choices] + return queryset.filter(**{lookup_field: values}) + if value not in enum_choices: raise ValidationError( f"Invalid provider value: '{value}'. Valid values are: " @@ -49,6 +57,12 @@ def enum_filter(queryset, value, enum_choices, lookup_field: str): return queryset.filter(**{lookup_field: value}) +def extract_lookup_expr(name): + """Extract the lookup expression from the filter name.""" + parts = name.split("__") + return parts[-1] if len(parts) > 1 else "exact" + + class CustomDjangoFilterBackend(DjangoFilterBackend): def to_html(self, _request, _queryset, _view): """Override this method to use the Browsable API in dev environments. @@ -75,9 +89,10 @@ class ProviderFilter(FilterSet): inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") updated_at = DateFilter(field_name="updated_at", lookup_expr="date") connected = BooleanFilter() - provider = CharFilter(method="filter_provider") + provider = CharFilter(method="filter_provider_type") + provider__in = CharFilter(method="filter_provider_type_in") - def filter_provider(self, queryset, name, value): + def filter_provider_type(self, queryset, name, value): return enum_filter( queryset, value, @@ -85,12 +100,21 @@ class ProviderFilter(FilterSet): lookup_field="provider", ) + def filter_provider_type_in(self, queryset, name, value): + return enum_filter( + queryset, + value, + enum_choices=Provider.ProviderChoices, + lookup_field="provider__in", + ) + class Meta: model = Provider fields = { - "provider": ["exact"], - "provider_id": ["exact", "icontains"], - "alias": ["exact", "icontains"], + "provider": ["exact", "in"], + "id": ["exact", "in"], + "uid": ["exact", "icontains", "in"], + "alias": ["exact", "icontains", "in"], "inserted_at": ["gte", "lte"], "updated_at": ["gte", "lte"], } @@ -101,21 +125,64 @@ class ProviderFilter(FilterSet): } -class ScanFilter(FilterSet): - inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") - completed_at = DateFilter(field_name="completed_at", lookup_expr="date") - started_at = DateFilter(field_name="started_at", lookup_expr="date") - provider = CharFilter(method="filter_provider") - trigger = CharFilter(method="filter_trigger") +class ProviderRelationshipFilterSet(FilterSet): + provider_type = CharFilter(method="filter_provider_type") + provider_type__in = CharFilter(method="filter_provider_type_in") + provider_uid = CharFilter(method="filter_provider_uid") + provider_uid__in = CharFilter(method="filter_provider_uid_in") + provider_uid__icontains = CharFilter(method="filter_provider_uid_icontains") + provider_alias = CharFilter(method="filter_provider_alias") + provider_alias__in = CharFilter(method="filter_provider_alias_in") + provider_alias__icontains = CharFilter(method="filter_provider_alias_icontains") - def filter_provider(self, queryset, name, value): + def filter_provider_type(self, queryset, name, value): return enum_filter( queryset, value, enum_choices=Provider.ProviderChoices, - lookup_field="provider__provider", + lookup_field="provider__provider__exact", ) + def filter_provider_type_in(self, queryset, name, value): + if not isinstance(value, list): + value = value.split(",") + + return enum_filter( + queryset, + value, + enum_choices=Provider.ProviderChoices, + lookup_field="provider__provider__in", + ) + + def filter_provider_uid(self, queryset, name, value): + return queryset.filter(provider__uid=value) + + def filter_provider_uid_in(self, queryset, name, value): + if isinstance(value, str): + value = value.split(",") + return queryset.filter(provider__uid__in=value) + + def filter_provider_uid_icontains(self, queryset, name, value): + return queryset.filter(provider__uid__icontains=value) + + def filter_provider_alias(self, queryset, name, value): + return queryset.filter(provider__alias=value) + + def filter_provider_alias_in(self, queryset, name, value): + if isinstance(value, str): + value = value.split(",") + return queryset.filter(provider__alias__in=value) + + def filter_provider_alias_icontains(self, queryset, name, value): + return queryset.filter(provider__alias__icontains=value) + + +class ScanFilter(ProviderRelationshipFilterSet): + inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") + completed_at = DateFilter(field_name="completed_at", lookup_expr="date") + started_at = DateFilter(field_name="started_at", lookup_expr="date") + trigger = CharFilter(method="filter_trigger") + def filter_trigger(self, queryset, name, value): if value not in Scan.TriggerChoices: raise ValidationError( @@ -128,8 +195,7 @@ class ScanFilter(FilterSet): class Meta: model = Scan fields = { - "provider": ["exact"], - "provider_id": ["exact", "in"], + "provider": ["exact", "in"], "name": ["exact", "icontains"], "started_at": ["gte", "lte"], "trigger": ["exact"], @@ -173,8 +239,7 @@ class ResourceTagFilter(FilterSet): search = ["text_search"] -class ResourceFilter(FilterSet): - provider = CharFilter(method="filter_provider") +class ResourceFilter(ProviderRelationshipFilterSet): tag_key = CharFilter(method="filter_tag_key") tag_value = CharFilter(method="filter_tag_value") tag = CharFilter(method="filter_tag") @@ -182,18 +247,10 @@ class ResourceFilter(FilterSet): 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"], + "provider": ["exact", "in"], "uid": ["exact", "icontains"], "name": ["exact", "icontains"], "region": ["exact", "icontains", "in"], diff --git a/src/backend/api/fixtures/1_dev_providers.json b/src/backend/api/fixtures/1_dev_providers.json index 6323b516f1..22292f199b 100644 --- a/src/backend/api/fixtures/1_dev_providers.json +++ b/src/backend/api/fixtures/1_dev_providers.json @@ -7,7 +7,7 @@ "inserted_at": "2024-08-01T17:20:27.050Z", "updated_at": "2024-08-01T17:20:27.050Z", "provider": "gcp", - "provider_id": "a12322-test321", + "uid": "a12322-test321", "alias": "gcp_testing_2", "connected": null, "connection_last_checked_at": null, @@ -22,7 +22,7 @@ "inserted_at": "2024-08-01T17:19:42.453Z", "updated_at": "2024-08-01T17:19:42.453Z", "provider": "gcp", - "provider_id": "a12345-test123", + "uid": "a12345-test123", "alias": "gcp_testing_1", "connected": null, "connection_last_checked_at": null, @@ -37,7 +37,7 @@ "inserted_at": "2024-08-01T17:19:09.556Z", "updated_at": "2024-08-01T17:19:09.556Z", "provider": "aws", - "provider_id": "123456789020", + "uid": "123456789020", "alias": "aws_testing_2", "connected": null, "connection_last_checked_at": null, @@ -52,7 +52,7 @@ "inserted_at": "2024-08-01T17:20:16.962Z", "updated_at": "2024-08-01T17:20:16.962Z", "provider": "gcp", - "provider_id": "a12322-test123", + "uid": "a12322-test123", "alias": "gcp_testing_3", "connected": null, "connection_last_checked_at": null, @@ -67,7 +67,7 @@ "inserted_at": "2024-08-01T17:18:58.132Z", "updated_at": "2024-08-01T17:18:58.132Z", "provider": "aws", - "provider_id": "123456789015", + "uid": "123456789015", "alias": "aws_testing_1", "connected": null, "connection_last_checked_at": null, @@ -82,7 +82,7 @@ "inserted_at": "2024-08-06T16:03:26.176Z", "updated_at": "2024-08-06T16:03:26.176Z", "provider": "azure", - "provider_id": "8851db6b-42e5-4533-aa9e-30a32d67e875", + "uid": "8851db6b-42e5-4533-aa9e-30a32d67e875", "alias": "azure_testing", "connected": null, "connection_last_checked_at": null, @@ -98,7 +98,7 @@ "inserted_at": "2024-08-06T16:03:07.037Z", "updated_at": "2024-08-06T16:03:07.037Z", "provider": "kubernetes", - "provider_id": "kubernetes-test-12345", + "uid": "kubernetes-test-12345", "alias": "k8s_testing", "connected": null, "connection_last_checked_at": null, diff --git a/src/backend/api/migrations/0001_initial.py b/src/backend/api/migrations/0001_initial.py index 8fb82ee65e..29ae972f68 100644 --- a/src/backend/api/migrations/0001_initial.py +++ b/src/backend/api/migrations/0001_initial.py @@ -138,10 +138,11 @@ class Migration(migrations.Migration): ), ), ( - "provider_id", + "uid", models.CharField( max_length=63, validators=[django.core.validators.MinLengthValidator(3)], + verbose_name="Unique identifier for the provider, set by the provider", ), ), ( @@ -183,8 +184,8 @@ class Migration(migrations.Migration): migrations.AddConstraint( model_name="provider", constraint=models.UniqueConstraint( - fields=("tenant_id", "provider", "provider_id"), - name="unique_provider_ids", + fields=("tenant_id", "provider", "uid"), + name="unique_provider_uids", ), ), # Create and register ScanTriggerEnum type diff --git a/src/backend/api/models.py b/src/backend/api/models.py index 6e029b28fd..85c59a4341 100644 --- a/src/backend/api/models.py +++ b/src/backend/api/models.py @@ -32,16 +32,16 @@ class Provider(RowLevelSecurityProtectedModel): KUBERNETES = "kubernetes", _("Kubernetes") @staticmethod - def validate_aws_provider_id(value): + def validate_aws_uid(value): if not re.match(r"^\d{12}$", value): raise ModelValidationError( detail="AWS provider ID must be exactly 12 digits.", - code="aws-provider-id", - pointer="/data/attributes/provider_id", + code="aws-uid", + pointer="/data/attributes/uid", ) @staticmethod - def validate_azure_provider_id(value): + def validate_azure_uid(value): try: val = UUID(value, version=4) if str(val) != value: @@ -49,28 +49,28 @@ class Provider(RowLevelSecurityProtectedModel): except ValueError: raise ModelValidationError( detail="Azure provider ID must be a valid UUID.", - code="azure-provider-id", - pointer="/data/attributes/provider_id", + code="azure-uid", + pointer="/data/attributes/uid", ) @staticmethod - def validate_gcp_provider_id(value): + def validate_gcp_uid(value): if not re.match(r"^[a-z][a-z0-9-]{5,29}$", value): raise ModelValidationError( detail="GCP provider ID must be 6 to 30 characters, start with a letter, and contain only lowercase " "letters, numbers, and hyphens.", - code="gcp-provider-id", - pointer="/data/attributes/provider_id", + code="gcp-uid", + pointer="/data/attributes/uid", ) @staticmethod - def validate_kubernetes_provider_id(value): + def validate_kubernetes_uid(value): if not re.match(r"^[a-z0-9]([-a-z0-9]{1,61}[a-z0-9])?$", value): raise ModelValidationError( detail="K8s provider ID must be up to 63 characters, start and end with a lowercase letter or number, " "and contain only lowercase alphanumeric characters and hyphens.", - code="kubernetes-provider-id", - pointer="/data/attributes/provider_id", + code="kubernetes-uid", + pointer="/data/attributes/uid", ) id = models.UUIDField(primary_key=True, default=uuid4, editable=False) @@ -79,7 +79,12 @@ class Provider(RowLevelSecurityProtectedModel): provider = ProviderEnumField( choices=ProviderChoices.choices, default=ProviderChoices.AWS ) - provider_id = models.CharField(max_length=63, validators=[MinLengthValidator(3)]) + uid = models.CharField( + "Unique identifier for the provider, set by the provider", + max_length=63, + blank=False, + validators=[MinLengthValidator(3)], + ) alias = models.CharField( blank=True, null=True, max_length=100, validators=[MinLengthValidator(3)] ) @@ -90,7 +95,7 @@ class Provider(RowLevelSecurityProtectedModel): def clean(self): super().clean() - getattr(self, f"validate_{self.provider}_provider_id")(self.provider_id) + getattr(self, f"validate_{self.provider}_uid")(self.uid) def save(self, *args, **kwargs): self.full_clean() @@ -101,8 +106,8 @@ class Provider(RowLevelSecurityProtectedModel): constraints = [ models.UniqueConstraint( - fields=("tenant_id", "provider", "provider_id"), - name="unique_provider_ids", + fields=("tenant_id", "provider", "uid"), + name="unique_provider_uids", ), RowLevelSecurityConstraint( field="tenant_id", diff --git a/src/backend/api/tests/integration/test_tenants.py b/src/backend/api/tests/integration/test_tenants.py index d1790f7134..cf3afcf8e6 100644 --- a/src/backend/api/tests/integration/test_tenants.py +++ b/src/backend/api/tests/integration/test_tenants.py @@ -19,7 +19,7 @@ def test_check_resources_between_different_tenants( "attributes": { "alias": "test_provider_tenant_1", "provider": "aws", - "provider_id": "123456789012", + "uid": "123456789012", }, } } @@ -39,7 +39,7 @@ def test_check_resources_between_different_tenants( "attributes": { "alias": "test_provider_tenant_2", "provider": "aws", - "provider_id": "123456789013", + "uid": "123456789013", }, } } diff --git a/src/backend/api/tests/test_views.py b/src/backend/api/tests/test_views.py index ee50c87f3f..debc28ba6e 100644 --- a/src/backend/api/tests/test_views.py +++ b/src/backend/api/tests/test_views.py @@ -213,10 +213,7 @@ class TestProviderViewSet: ) assert response.status_code == status.HTTP_200_OK assert response.json()["data"]["attributes"]["provider"] == provider1.provider - assert ( - response.json()["data"]["attributes"]["provider_id"] - == provider1.provider_id - ) + assert response.json()["data"]["attributes"]["uid"] == provider1.uid assert response.json()["data"]["attributes"]["alias"] == provider1.alias def test_providers_invalid_retrieve(self, client, tenant_header): @@ -230,16 +227,16 @@ class TestProviderViewSet: "provider_json_payload", ( [ - {"provider": "aws", "provider_id": "111111111111", "alias": "test"}, - {"provider": "gcp", "provider_id": "a12322-test54321", "alias": "test"}, + {"provider": "aws", "uid": "111111111111", "alias": "test"}, + {"provider": "gcp", "uid": "a12322-test54321", "alias": "test"}, { "provider": "kubernetes", - "provider_id": "kubernetes-test-123456789", + "uid": "kubernetes-test-123456789", "alias": "test", }, { "provider": "azure", - "provider_id": "8851db6b-42e5-4533-aa9e-30a32d67e875", + "uid": "8851db6b-42e5-4533-aa9e-30a32d67e875", "alias": "test", }, ] @@ -255,9 +252,7 @@ class TestProviderViewSet: assert response.status_code == status.HTTP_201_CREATED assert Provider.objects.count() == 1 assert Provider.objects.get().provider == provider_json_payload["provider"] - assert ( - Provider.objects.get().provider_id == provider_json_payload["provider_id"] - ) + assert Provider.objects.get().uid == provider_json_payload["uid"] assert Provider.objects.get().alias == provider_json_payload["alias"] @pytest.mark.parametrize( @@ -265,51 +260,51 @@ class TestProviderViewSet: ( [ ( - {"provider": "aws", "provider_id": "1", "alias": "test"}, + {"provider": "aws", "uid": "1", "alias": "test"}, "min_length", - "provider_id", + "uid", ), ( { "provider": "aws", - "provider_id": "1111111111111", + "uid": "1111111111111", "alias": "test", }, - "aws-provider-id", - "provider_id", + "aws-uid", + "uid", ), ( - {"provider": "aws", "provider_id": "aaaaaaaaaaaa", "alias": "test"}, - "aws-provider-id", - "provider_id", + {"provider": "aws", "uid": "aaaaaaaaaaaa", "alias": "test"}, + "aws-uid", + "uid", ), ( - {"provider": "gcp", "provider_id": "1234asdf", "alias": "test"}, - "gcp-provider-id", - "provider_id", + {"provider": "gcp", "uid": "1234asdf", "alias": "test"}, + "gcp-uid", + "uid", ), ( { "provider": "kubernetes", - "provider_id": "-1234asdf", + "uid": "-1234asdf", "alias": "test", }, - "kubernetes-provider-id", - "provider_id", + "kubernetes-uid", + "uid", ), ( { "provider": "azure", - "provider_id": "8851db6b-42e5-4533-aa9e-30a32d67e87", + "uid": "8851db6b-42e5-4533-aa9e-30a32d67e87", "alias": "test", }, - "azure-provider-id", - "provider_id", + "azure-uid", + "uid", ), ( { "provider": "does-not-exist", - "provider_id": "8851db6b-42e5-4533-aa9e-30a32d67e87", + "uid": "8851db6b-42e5-4533-aa9e-30a32d67e87", "alias": "test", }, "invalid_choice", @@ -383,7 +378,7 @@ class TestProviderViewSet: "attribute_key, attribute_value", [ ("provider", "aws"), - ("provider_id", "123456789012"), + ("uid", "123456789012"), ], ) def test_providers_partial_update_invalid_fields( @@ -470,8 +465,9 @@ class TestProviderViewSet: ( [ ("provider", "aws", 2), - ("provider_id", "123456789012", 1), - ("provider_id.icontains", "1", 5), + ("provider.in", "azure,gcp", 2), + ("uid", "123456789012", 1), + ("uid.icontains", "1", 5), ("alias", "aws_testing_1", 1), ("alias.icontains", "aws", 2), ("inserted_at", TODAY, 5), @@ -522,7 +518,7 @@ class TestProviderViewSet: ( [ "provider", - "provider_id", + "uid", "alias", "connected", "inserted_at", @@ -733,7 +729,14 @@ class TestScanViewSet: "filter_name, filter_value, expected_count", ( [ - ("provider", "aws", 3), + ("provider_type", "aws", 3), + ("provider_type.in", "gcp,azure", 0), + ("provider_uid", "123456789012", 2), + ("provider_uid.icontains", "1", 3), + ("provider_uid.in", "123456789012,123456789013", 3), + ("provider_alias", "aws_testing_1", 2), + ("provider_alias.icontains", "aws", 3), + ("provider_alias.in", "aws_testing_1,aws_testing_2", 3), ("name", "Scan 1", 1), ("name.icontains", "Scan", 3), ("started_at", "2024-01-02", 3), @@ -781,7 +784,7 @@ class TestScanViewSet: ): response = client.get( reverse("scan-list"), - {"filter[provider_id]": scans_fixture[0].provider.id}, + {"filter[provider]": scans_fixture[0].provider.id}, headers=tenant_header, ) assert response.status_code == status.HTTP_200_OK @@ -791,7 +794,7 @@ class TestScanViewSet: response = client.get( reverse("scan-list"), { - "filter[provider_id.in]": [ + "filter[provider.in]": [ scans_fixture[0].provider.id, scans_fixture[1].provider.id, ] @@ -804,7 +807,6 @@ class TestScanViewSet: @pytest.mark.parametrize( "sort_field", [ - "provider_id", "name", "trigger", "inserted_at", @@ -910,6 +912,15 @@ class TestResourceViewSet: ("inserted_at.gte", "2024-01-01 00:00:00", 3), ("updated_at.lte", "2024-01-01 00:00:00", 0), ("type.icontains", "prowler", 2), + # provider filters + ("provider_type", "aws", 3), + ("provider_type.in", "azure,gcp", 0), + ("provider_uid", "123456789012", 2), + ("provider_uid.in", "123456789012", 2), + ("provider_uid.in", "123456789012,123456789012", 2), + ("provider_uid.icontains", "1", 3), + ("provider_alias", "aws_testing_1", 2), + ("provider_alias.icontains", "aws", 3), # tags searching ("tag", "key3:value:value", 0), ("tag_key", "key3", 1), @@ -950,7 +961,7 @@ class TestResourceViewSet: response = client.get( reverse("resource-list"), { - "filter[provider_id.in]": [ + "filter[provider.in]": [ resources_fixture[0].provider.id, resources_fixture[1].provider.id, ] @@ -980,7 +991,7 @@ class TestResourceViewSet: @pytest.mark.parametrize( "sort_field", [ - "provider_id", + "uid", "uid", "name", "region", diff --git a/src/backend/api/v1/serializers.py b/src/backend/api/v1/serializers.py index be0d0731f9..e47e0cb0e1 100644 --- a/src/backend/api/v1/serializers.py +++ b/src/backend/api/v1/serializers.py @@ -180,7 +180,7 @@ class ProviderSerializer(RLSSerializer): "inserted_at", "updated_at", "provider", - "provider_id", + "uid", "alias", "connection", "scanner_args", @@ -206,7 +206,7 @@ class ProviderSerializer(RLSSerializer): class ProviderCreateSerializer(RLSSerializer, BaseWriteSerializer): class Meta: model = Provider - fields = ["alias", "provider", "provider_id", "scanner_args"] + fields = ["alias", "provider", "uid", "scanner_args"] class ProviderUpdateSerializer(BaseWriteSerializer): diff --git a/src/backend/api/v1/views.py b/src/backend/api/v1/views.py index 90c8eb2ab9..16c4b3f0d5 100644 --- a/src/backend/api/v1/views.py +++ b/src/backend/api/v1/views.py @@ -128,11 +128,11 @@ class ProviderViewSet(BaseRLSViewSet): serializer_class = ProviderSerializer http_method_names = ["get", "post", "patch", "delete"] filterset_class = ProviderFilter - search_fields = ["provider", "provider_id", "alias"] + search_fields = ["provider", "uid", "alias"] ordering = ["inserted_at"] ordering_fields = [ "provider", - "provider_id", + "uid", "alias", "connected", "inserted_at", @@ -238,7 +238,6 @@ class ScanViewSet(BaseRLSViewSet): filterset_class = ScanFilter ordering = ["inserted_at"] ordering_fields = [ - "provider_id", "name", "trigger", "attempted_at", @@ -362,7 +361,7 @@ class ResourceViewSet(BaseRLSViewSet): filterset_class = ResourceFilter ordering = ["inserted_at"] ordering_fields = [ - "provider_id", + "provider_uid", "uid", "name", "region", diff --git a/src/backend/conftest.py b/src/backend/conftest.py index 19bc1babf5..f95394fe2f 100644 --- a/src/backend/conftest.py +++ b/src/backend/conftest.py @@ -59,31 +59,31 @@ def providers_fixture(tenants_fixture): tenant, _ = tenants_fixture provider1 = Provider.objects.create( provider="aws", - provider_id="123456789012", + uid="123456789012", alias="aws_testing_1", tenant_id=tenant.id, ) provider2 = Provider.objects.create( provider="aws", - provider_id="123456789013", + uid="123456789013", alias="aws_testing_2", tenant_id=tenant.id, ) provider3 = Provider.objects.create( provider="gcp", - provider_id="a12322-test321", + uid="a12322-test321", alias="gcp_testing", tenant_id=tenant.id, ) provider4 = Provider.objects.create( provider="kubernetes", - provider_id="kubernetes-test-12345", + uid="kubernetes-test-12345", alias="k8s_testing", tenant_id=tenant.id, ) provider5 = Provider.objects.create( provider="azure", - provider_id="37b065f8-26b0-4218-a665-0b23d07b27d9", + uid="37b065f8-26b0-4218-a665-0b23d07b27d9", alias="azure_testing", tenant_id=tenant.id, scanner_args={"key1": "value1", "key2": {"key21": "value21"}}, diff --git a/src/backend/tasks/tests/test_connection.py b/src/backend/tasks/tests/test_connection.py index b2466cd025..dad47eff13 100644 --- a/src/backend/tasks/tests/test_connection.py +++ b/src/backend/tasks/tests/test_connection.py @@ -11,7 +11,7 @@ from tasks.jobs.connection import check_provider_connection "provider_data, provider_class", [ ( - {"provider": "aws", "provider_id": "123456789012", "alias": "aws"}, + {"provider": "aws", "uid": "123456789012", "alias": "aws"}, "AwsProvider", ), ],