fix(Providers, Resources, Scans): rename provider_id and filter on more provider fields (#42)

* fix(Providers, Resources, Scans): filter on more provider fields

* Apply suggestions from code review

more python-y

Co-authored-by: Víctor Fernández Poyatos <victor@prowler.com>

---------

Co-authored-by: Víctor Fernández Poyatos <victor@prowler.com>
This commit is contained in:
Jon Young
2024-09-13 10:09:09 -04:00
committed by GitHub
parent 1cef6f0db7
commit 6a341b88f0
10 changed files with 178 additions and 105 deletions
+83 -26
View File
@@ -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"],
@@ -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,
+4 -3
View File
@@ -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
+21 -16
View File
@@ -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",
@@ -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",
},
}
}
+50 -39
View File
@@ -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",
+2 -2
View File
@@ -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):
+3 -4
View File
@@ -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",
+5 -5
View File
@@ -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"}},
+1 -1
View File
@@ -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",
),
],