diff --git a/src/backend/api/filters.py b/src/backend/api/filters.py index f374b87359..8e8c548228 100644 --- a/src/backend/api/filters.py +++ b/src/backend/api/filters.py @@ -1,4 +1,9 @@ -from django_filters.rest_framework import FilterSet, BooleanFilter, CharFilter +from django_filters.rest_framework import ( + FilterSet, + BooleanFilter, + CharFilter, + DateFilter, +) from rest_framework_json_api.django_filters.backends import DjangoFilterBackend from rest_framework_json_api.serializers import ValidationError @@ -53,16 +58,21 @@ class CustomDjangoFilterBackend(DjangoFilterBackend): class TenantFilter(FilterSet): + inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") + updated_at = DateFilter(field_name="updated_at", lookup_expr="date") + class Meta: model = Tenant fields = { "name": ["exact", "icontains"], - "inserted_at": ["exact", "gte", "lte"], - "updated_at": ["exact", "gte", "lte"], + "inserted_at": ["date", "gte", "lte"], + "updated_at": ["gte", "lte"], } 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") @@ -80,8 +90,8 @@ class ProviderFilter(FilterSet): "provider": ["exact"], "provider_id": ["exact", "icontains"], "alias": ["exact", "icontains"], - "inserted_at": ["exact", "gte", "lte"], - "updated_at": ["exact", "gte", "lte"], + "inserted_at": ["gte", "lte"], + "updated_at": ["gte", "lte"], } filter_overrides = { ProviderEnumField: { @@ -91,6 +101,9 @@ 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") @@ -117,7 +130,7 @@ class ScanFilter(FilterSet): "provider": ["exact"], "provider_id": ["exact"], "name": ["exact", "icontains"], - "started_at": ["exact", "gte", "lte"], + "started_at": ["gte", "lte"], "trigger": ["exact"], } diff --git a/src/backend/api/tests/test_views.py b/src/backend/api/tests/test_views.py index 6122ab2a03..49780eee01 100644 --- a/src/backend/api/tests/test_views.py +++ b/src/backend/api/tests/test_views.py @@ -3,28 +3,31 @@ from unittest.mock import Mock, patch import pytest from django.urls import reverse from rest_framework import status - +from datetime import datetime from api.models import Provider, Scan from api.rls import Tenant from conftest import API_JSON_CONTENT_TYPE, NO_TENANT_HTTP_STATUS +TODAY = str(datetime.today().date()) + + @pytest.mark.django_db class TestTenantViewSet: @pytest.fixture def valid_tenant_payload(self): return { "name": "Tenant Three", - "inserted_at": "2023-01-05T00:00:00Z", - "updated_at": "2023-01-06T00:00:00Z", + "inserted_at": "2023-01-05", + "updated_at": "2023-01-06", } @pytest.fixture def invalid_tenant_payload(self): return { "name": "", - "inserted_at": "2023-01-05T00:00:00Z", - "updated_at": "2023-01-06T00:00:00Z", + "inserted_at": "2023-01-05", + "updated_at": "2023-01-06", } def test_tenants_list(self, client, tenants_fixture): @@ -125,12 +128,37 @@ class TestTenantViewSet: response = client.get(reverse("tenant-list"), {"random": "value"}) assert response.status_code == status.HTTP_400_BAD_REQUEST - def test_tenants_list_filter_name(self, client, tenants_fixture): - tenant1, _ = tenants_fixture - response = client.get(reverse("tenant-list"), {"filter[name]": tenant1.name}) + @pytest.mark.parametrize( + "filter_name, filter_value, expected_count", + ( + [ + ("name", "Tenant One", 1), + ("name.icontains", "Tenant", 2), + ("inserted_at", TODAY, 2), + ("inserted_at.gte", "2024-01-01", 2), + ("inserted_at.lte", "2024-01-01", 0), + ("updated_at.gte", "2024-01-01", 2), + ("updated_at.lte", "2024-01-01", 0), + ] + ), + ) + def test_tenants_filters( + self, + client, + tenants_fixture, + tenant_header, + filter_name, + filter_value, + expected_count, + ): + response = client.get( + reverse("tenant-list"), + {f"filter[{filter_name}]": filter_value}, + headers=tenant_header, + ) + assert response.status_code == status.HTTP_200_OK - assert len(response.json()["data"]) == 1 - assert response.json()["data"][0]["attributes"]["name"] == tenant1.name + assert len(response.json()["data"]) == expected_count def test_tenants_list_filter_invalid(self, client): response = client.get(reverse("tenant-list"), {"filter[invalid]": "whatever"}) @@ -438,25 +466,39 @@ class TestProviderViewSet: assert response.status_code == status.HTTP_404_NOT_FOUND @pytest.mark.parametrize( - "filter_name, filter_value", + "filter_name, filter_value, expected_count", ( [ - ("provider", "aws"), - ("provider_id", "12345"), - ("alias", "test"), - ("search", "test"), - ("inserted_at", "2024-01-01 00:00:00"), - ("updated_at", "2024-01-01 00:00:00"), + ("provider", "aws", 2), + ("provider_id", "123456789012", 1), + ("provider_id.icontains", "1", 5), + ("alias", "aws_testing_1", 1), + ("alias.icontains", "aws", 2), + ("inserted_at", TODAY, 5), + ("inserted_at.gte", "2024-01-01", 5), + ("inserted_at.lte", "2024-01-01", 0), + ("updated_at.gte", "2024-01-01", 5), + ("updated_at.lte", "2024-01-01", 0), ] ), ) - def test_providers_filters(self, client, tenant_header, filter_name, filter_value): + def test_providers_filters( + self, + client, + providers_fixture, + tenant_header, + filter_name, + filter_value, + expected_count, + ): response = client.get( reverse("provider-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 @pytest.mark.parametrize( "filter_name", @@ -688,21 +730,36 @@ class TestScanViewSet: assert response.status_code == status.HTTP_400_BAD_REQUEST @pytest.mark.parametrize( - "filter_name, filter_value", - [ - ("provider", "aws"), - ("trigger", Scan.TriggerChoices.MANUAL), - ("name", "Scan 1"), - ("started_at", "2024-01-01 00:00:00"), - ], + "filter_name, filter_value, expected_count", + ( + [ + ("provider", "aws", 3), + ("name", "Scan 1", 1), + ("name.icontains", "Scan", 3), + ("started_at", "2024-01-02", 3), + ("started_at.gte", "2024-01-01", 3), + ("started_at.lte", "2024-01-01", 0), + ("trigger", Scan.TriggerChoices.MANUAL, 1), + ] + ), ) - def test_scans_filters(self, client, tenant_header, filter_name, filter_value): + def test_scans_filters( + self, + client, + scans_fixture, + tenant_header, + filter_name, + filter_value, + expected_count, + ): response = client.get( reverse("scan-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 @pytest.mark.parametrize( "filter_name", diff --git a/src/backend/conftest.py b/src/backend/conftest.py index 30c953028a..a113eed7b1 100644 --- a/src/backend/conftest.py +++ b/src/backend/conftest.py @@ -47,13 +47,9 @@ def disable_logging(): def tenants_fixture(): tenant1 = Tenant.objects.create( name="Tenant One", - inserted_at="2023-01-01T00:00:00Z", - updated_at="2023-01-02T00:00:00Z", ) tenant2 = Tenant.objects.create( name="Tenant Two", - inserted_at="2023-01-03T00:00:00Z", - updated_at="2023-01-04T00:00:00Z", ) return tenant1, tenant2 @@ -107,6 +103,7 @@ def scans_fixture(tenants_fixture, providers_fixture): trigger=Scan.TriggerChoices.MANUAL, state=StateChoices.AVAILABLE, tenant_id=tenant.id, + started_at="2024-01-02T00:00:00Z", ) scan2 = Scan.objects.create( name="Scan 2", @@ -114,6 +111,7 @@ def scans_fixture(tenants_fixture, providers_fixture): trigger=Scan.TriggerChoices.SCHEDULED, state=StateChoices.FAILED, tenant_id=tenant.id, + started_at="2024-01-02T00:00:00Z", ) scan3 = Scan.objects.create( name="Scan 3", @@ -121,6 +119,7 @@ def scans_fixture(tenants_fixture, providers_fixture): trigger=Scan.TriggerChoices.SCHEDULED, state=StateChoices.AVAILABLE, tenant_id=tenant.id, + started_at="2024-01-02T00:00:00Z", ) return scan1, scan2, scan3