diff --git a/src/backend/api/db_utils.py b/src/backend/api/db_utils.py index da1920e1db..87f3081479 100644 --- a/src/backend/api/db_utils.py +++ b/src/backend/api/db_utils.py @@ -114,13 +114,13 @@ class ProviderEnumField(PostgresEnumField): # Postgres enum definition for Scan.type -class ScanTypeEnum(EnumType): - enum_type_name = "scan_type" +class ScanTriggerEnum(EnumType): + enum_type_name = "scan_trigger" -class ScanTypeEnumField(PostgresEnumField): +class ScanTriggerEnumField(PostgresEnumField): def __init__(self, *args, **kwargs): - super().__init__("scan_type", *args, **kwargs) + super().__init__("scan_trigger", *args, **kwargs) # Postgres enum definition for state diff --git a/src/backend/api/filters.py b/src/backend/api/filters.py index b2a92c8940..9b77f53f02 100644 --- a/src/backend/api/filters.py +++ b/src/backend/api/filters.py @@ -83,19 +83,19 @@ class ProviderFilter(FilterSet): class ScanFilter(FilterSet): provider = CharFilter(method="filter_provider") - type = CharFilter(method="filter_type") + trigger = CharFilter(method="filter_trigger") def filter_provider(self, queryset, name, value): return provider_enum_filter(queryset, value, lookup_field="provider__provider") - def filter_type(self, queryset, name, value): - if value not in Scan.TypeChoices: + def filter_trigger(self, queryset, name, value): + if value not in Scan.TriggerChoices: raise ValidationError( - f"Invalid scan type value: '{value}'. Valid values are: " - f"{', '.join(Scan.TypeChoices)}" + f"Invalid scan trigger value: '{value}'. Valid values are: " + f"{', '.join(Scan.TriggerChoices)}" ) - return queryset.filter(type=value) + return queryset.filter(trigger=value) class Meta: model = Scan @@ -104,5 +104,5 @@ class ScanFilter(FilterSet): "provider_id": ["exact"], "name": ["exact", "icontains"], "started_at": ["exact", "gte", "lte"], - "type": ["exact"], + "trigger": ["exact"], } diff --git a/src/backend/api/migrations/0001_initial.py b/src/backend/api/migrations/0001_initial.py index df1701da4c..a2bb746224 100644 --- a/src/backend/api/migrations/0001_initial.py +++ b/src/backend/api/migrations/0001_initial.py @@ -13,10 +13,10 @@ from api.db_utils import ( PostgresEnumMigration, ProviderEnum, ProviderEnumField, - ScanTypeEnum, + ScanTriggerEnum, StateEnumField, StateEnum, - ScanTypeEnumField, + ScanTriggerEnumField, register_enum, ) from api.models import Provider, Scan, StateChoices @@ -34,9 +34,9 @@ ProviderEnumMigration = PostgresEnumMigration( enum_values=tuple(provider[0] for provider in Provider.ProviderChoices.choices), ) -ScanTypeEnumMigration = PostgresEnumMigration( - enum_name="scan_type", - enum_values=tuple(scan_type[0] for scan_type in Scan.TypeChoices.choices), +ScanTriggerEnumMigration = PostgresEnumMigration( + enum_name="scan_trigger", + enum_values=tuple(scan_trigger[0] for scan_trigger in Scan.TriggerChoices.choices), ) StateEnumMigration = PostgresEnumMigration( @@ -188,12 +188,12 @@ class Migration(migrations.Migration): name="unique_provider_ids", ), ), - # Create and register ScanTypeEnum type + # Create and register ScanTriggerEnum type migrations.RunPython( - ScanTypeEnumMigration.create_enum_type, - reverse_code=ScanTypeEnumMigration.drop_enum_type, + ScanTriggerEnumMigration.create_enum_type, + reverse_code=ScanTriggerEnumMigration.drop_enum_type, ), - migrations.RunPython(partial(register_enum, enum_class=ScanTypeEnum)), + migrations.RunPython(partial(register_enum, enum_class=ScanTriggerEnum)), migrations.CreateModel( name="Scan", fields=[ @@ -216,8 +216,8 @@ class Migration(migrations.Migration): ), ), ( - "type", - ScanTypeEnumField( + "trigger", + ScanTriggerEnumField( choices=[("scheduled", "Scheduled"), ("manual", "Manual")] ), ), @@ -279,8 +279,8 @@ class Migration(migrations.Migration): migrations.AddIndex( model_name="scan", index=models.Index( - fields=["provider", "state", "type", "scheduled_at"], - name="scans_prov_state_type_sche_idx", + fields=["provider", "state", "trigger", "scheduled_at"], + name="scans_prov_state_trig_sche_idx", ), ), ] diff --git a/src/backend/api/models.py b/src/backend/api/models.py index cfd545ca30..d313ea8e7f 100644 --- a/src/backend/api/models.py +++ b/src/backend/api/models.py @@ -6,7 +6,7 @@ from django.db import models from django.utils.translation import gettext_lazy as _ from uuid6 import uuid7 -from api.db_utils import ProviderEnumField, StateEnumField, ScanTypeEnumField +from api.db_utils import ProviderEnumField, StateEnumField, ScanTriggerEnumField from api.exceptions import ModelValidationError from api.rls import RowLevelSecurityConstraint from api.rls import RowLevelSecurityProtectedModel @@ -110,7 +110,7 @@ class Provider(RowLevelSecurityProtectedModel): class Scan(RowLevelSecurityProtectedModel): - class TypeChoices(models.TextChoices): + class TriggerChoices(models.TextChoices): SCHEDULED = "scheduled", _("Scheduled") MANUAL = "manual", _("Manual") @@ -124,8 +124,8 @@ class Scan(RowLevelSecurityProtectedModel): related_name="scans", related_query_name="scan", ) - type = ScanTypeEnumField( - choices=TypeChoices.choices, + trigger = ScanTriggerEnumField( + choices=TriggerChoices.choices, ) state = StateEnumField(choices=StateChoices.choices, default=StateChoices.AVAILABLE) unique_resource_count = models.IntegerField(default=0) @@ -153,7 +153,7 @@ class Scan(RowLevelSecurityProtectedModel): indexes = [ models.Index( - fields=["provider", "state", "type", "scheduled_at"], + fields=["provider", "state", "trigger", "scheduled_at"], name="scans_prov_state_type_sche_idx", ), ] diff --git a/src/backend/api/tests/test_views.py b/src/backend/api/tests/test_views.py index 174034e5a2..799cff0e8e 100644 --- a/src/backend/api/tests/test_views.py +++ b/src/backend/api/tests/test_views.py @@ -589,7 +589,7 @@ class TestScanViewSet: scan = Scan.objects.get() assert scan.name == scan_json_payload["data"]["attributes"]["name"] assert scan.provider == provider5 - assert scan.type == Scan.TypeChoices.MANUAL + assert scan.trigger == Scan.TriggerChoices.MANUAL assert scan.scanner_args == expected_scanner_args @pytest.mark.parametrize( @@ -601,7 +601,7 @@ class TestScanViewSet: "type": "Scan", "attributes": { "name": "a", - "type": Scan.TypeChoices.MANUAL, + "trigger": Scan.TriggerChoices.MANUAL, }, "relationships": { "provider": { @@ -682,7 +682,7 @@ class TestScanViewSet: "filter_name, filter_value", [ ("provider", "aws"), - ("type", Scan.TypeChoices.MANUAL), + ("trigger", Scan.TriggerChoices.MANUAL), ("name", "Scan 1"), ("started_at", "2024-01-01 00:00:00"), ], @@ -715,7 +715,7 @@ class TestScanViewSet: [ "provider_id", "name", - "type", + "trigger", "inserted_at", "updated_at", ], diff --git a/src/backend/api/v1/serializers.py b/src/backend/api/v1/serializers.py index 93df31068d..72e0a5a669 100644 --- a/src/backend/api/v1/serializers.py +++ b/src/backend/api/v1/serializers.py @@ -130,14 +130,16 @@ class ProviderUpdateSerializer(BaseWriteSerializer): # Scans -class ScanTypeEnumSerializerField(serializers.ChoiceField): +class ScanTriggerEnumSerializerField(serializers.ChoiceField): def __init__(self, **kwargs): - kwargs["choices"] = Scan.TypeChoices.choices + kwargs["choices"] = Scan.TriggerChoices.choices super().__init__(**kwargs) class ScanSerializer(RLSSerializer): - type_ = serializers.ChoiceField(choices=Scan.TypeChoices.choices, read_only=True) + trigger = serializers.ChoiceField( + choices=Scan.TriggerChoices.choices, read_only=True + ) state = StateEnumSerializerField(read_only=True) class Meta: @@ -145,7 +147,7 @@ class ScanSerializer(RLSSerializer): fields = [ "id", "name", - "type_", + "trigger", "state", "unique_resource_count", "progress", @@ -158,13 +160,6 @@ class ScanSerializer(RLSSerializer): "url", ] - def get_fields(self): - """`type` is a Python reserved keyword.""" - fields = super().get_fields() - type_ = fields.pop("type_") - fields["type"] = type_ - return fields - class ScanCreateSerializer(RLSSerializer, BaseWriteSerializer): class Meta: @@ -182,8 +177,8 @@ class ScanCreateSerializer(RLSSerializer, BaseWriteSerializer): provider.scanner_args, validated_data["scanner_args"] ) - if not validated_data.get("type"): - validated_data["type"] = Scan.TypeChoices.MANUAL.value + if not validated_data.get("trigger"): + validated_data["trigger"] = Scan.TriggerChoices.MANUAL.value return RLSSerializer.create(self, validated_data) diff --git a/src/backend/api/v1/views.py b/src/backend/api/v1/views.py index 65d536ee6b..fb3eb68700 100644 --- a/src/backend/api/v1/views.py +++ b/src/backend/api/v1/views.py @@ -221,7 +221,7 @@ class ScanViewSet(BaseRLSViewSet): ordering_fields = [ "provider_id", "name", - "type", + "trigger", "attempted_at", "scheduled_at", "inserted_at", diff --git a/src/backend/conftest.py b/src/backend/conftest.py index 91c26557bd..f5e318b512 100644 --- a/src/backend/conftest.py +++ b/src/backend/conftest.py @@ -104,21 +104,21 @@ def scans_fixture(tenants_fixture, providers_fixture): scan1 = Scan.objects.create( name="Scan 1", provider=provider, - type=Scan.TypeChoices.MANUAL, + trigger=Scan.TriggerChoices.MANUAL, state=StateChoices.AVAILABLE, tenant_id=tenant.id, ) scan2 = Scan.objects.create( name="Scan 2", provider=provider, - type=Scan.TypeChoices.SCHEDULED, + trigger=Scan.TriggerChoices.SCHEDULED, state=StateChoices.FAILED, tenant_id=tenant.id, ) scan3 = Scan.objects.create( name="Scan 3", provider=provider, - type=Scan.TypeChoices.SCHEDULED, + trigger=Scan.TriggerChoices.SCHEDULED, state=StateChoices.AVAILABLE, tenant_id=tenant.id, )