diff --git a/src/backend/api/db_utils.py b/src/backend/api/db_utils.py new file mode 100644 index 0000000000..fcb674a5f4 --- /dev/null +++ b/src/backend/api/db_utils.py @@ -0,0 +1,111 @@ +from contextlib import contextmanager + +from django.conf import settings +from django.db import models +from psycopg2 import connect as psycopg2_connect +from psycopg2.extensions import new_type, register_type, register_adapter, AsIs + + +@contextmanager +def psycopg_connection(database_alias: str): + psycopg2_connection = None + try: + admin_db = settings.DATABASES[database_alias] + + psycopg2_connection = psycopg2_connect( + dbname=admin_db["NAME"], + user=admin_db["USER"], + password=admin_db["PASSWORD"], + host=admin_db["HOST"], + ) + yield psycopg2_connection + finally: + if psycopg2_connection is not None: + psycopg2_connection.close() + + +# Postgres Enums + + +class PostgresEnumMigration: + def __init__(self, enum_name: str, enum_values: tuple): + self.enum_name = enum_name + self.enum_values = enum_values + + def create_enum_type(self, apps, schema_editor): # noqa: F841 + string_enum_values = ", ".join([f"'{value}'" for value in self.enum_values]) + with schema_editor.connection.cursor() as cursor: + cursor.execute( + f"CREATE TYPE {self.enum_name} AS ENUM ({string_enum_values});" + ) + + def drop_enum_type(self, apps, schema_editor): # noqa: F841 + with schema_editor.connection.cursor() as cursor: + cursor.execute(f"DROP TYPE {self.enum_name};") + + +class PostgresEnumField(models.Field): + def __init__(self, enum_type_name, *args, **kwargs): + self.enum_type_name = enum_type_name + super().__init__(*args, **kwargs) + + def db_type(self, connection): + return self.enum_type_name + + def from_db_value(self, value, expression, connection): # noqa: F841 + return value + + def to_python(self, value): + if isinstance(value, EnumType): + return value.value + return value + + def get_prep_value(self, value): + if isinstance(value, EnumType): + return value.value + return value + + +class EnumType: + def __init__(self, value): + self.value = value + + def __str__(self): + return self.value + + +def enum_adapter(enum_obj): + return AsIs(f"'{enum_obj.value}'::{enum_obj.__class__.enum_type_name}") + + +def get_enum_oid(connection, enum_type_name: str): + with connection.cursor() as cursor: + cursor.execute("SELECT oid FROM pg_type WHERE typname = %s;", (enum_type_name,)) + result = cursor.fetchone() + if result is None: + raise ValueError(f"Enum type '{enum_type_name}' not found") + return result[0] + + +def register_enum(apps, schema_editor, enum_class): # noqa: F841 + with psycopg_connection(schema_editor.connection.alias) as connection: + enum_oid = get_enum_oid(connection, enum_class.enum_type_name) + enum_instance = new_type( + (enum_oid,), + enum_class.enum_type_name, + lambda value, cur: value, # noqa: F841 + ) + register_type(enum_instance, connection) + register_adapter(enum_class, enum_adapter) + + +# Postgres enum definition for Provider.provider + + +class ProviderEnum(EnumType): + enum_type_name = "provider" + + +class ProviderEnumField(PostgresEnumField): + def __init__(self, *args, **kwargs): + super().__init__("provider", *args, **kwargs) diff --git a/src/backend/api/filters.py b/src/backend/api/filters.py index dd6e48b924..8c94c6d15c 100644 --- a/src/backend/api/filters.py +++ b/src/backend/api/filters.py @@ -1,6 +1,7 @@ -from django_filters.rest_framework import FilterSet, BooleanFilter +from django_filters.rest_framework import FilterSet, BooleanFilter, CharFilter from rest_framework_json_api.django_filters.backends import DjangoFilterBackend +from api.db_utils import ProviderEnumField from api.models import Provider from api.rls import Tenant @@ -36,3 +37,8 @@ class ProviderFilter(FilterSet): "inserted_at": ["exact", "gte", "lte"], "updated_at": ["exact", "gte", "lte"], } + filter_overrides = { + ProviderEnumField: { + "filter_class": CharFilter, + }, + } diff --git a/src/backend/api/migrations/0001_initial.py b/src/backend/api/migrations/0001_initial.py index e9cd32946d..b9af1b7d1e 100644 --- a/src/backend/api/migrations/0001_initial.py +++ b/src/backend/api/migrations/0001_initial.py @@ -1,4 +1,5 @@ import uuid +from functools import partial import django.core.validators import django.db.models.deletion @@ -6,14 +7,28 @@ from django.conf import settings from django.db import migrations, models import api.rls +from api.db_utils import ( + PostgresEnumMigration, + register_enum, + ProviderEnumField, + ProviderEnum, +) +from api.models import Provider DB_NAME = settings.DATABASES["default"]["NAME"] DB_USER_NAME = settings.DATABASES["default"]["USER"] DB_USER_PASSWORD = settings.DATABASES["default"]["PASSWORD"] +ProviderEnumMigration = PostgresEnumMigration( + enum_name="provider", + enum_values=tuple(provider[0] for provider in Provider.ProviderChoices.choices), +) + class Migration(migrations.Migration): initial = True + # Required for our kind of `RunPython` operations + atomic = False dependencies = [] @@ -56,15 +71,21 @@ class Migration(migrations.Migration): ("name", models.CharField(max_length=100)), ], options={ - "db_table": "tenant", + "db_table": "tenants", }, ), migrations.RunSQL( # Needed for now since we don't have users yet f""" - GRANT SELECT, INSERT, UPDATE, DELETE ON TABLE tenant TO {DB_USER_NAME}; + GRANT SELECT, INSERT, UPDATE, DELETE ON TABLE tenants TO {DB_USER_NAME}; """ ), + # Create and register ProviderEnum type + migrations.RunPython( + ProviderEnumMigration.create_enum_type, + reverse_code=ProviderEnumMigration.drop_enum_type, + ), + migrations.RunPython(partial(register_enum, enum_class=ProviderEnum)), migrations.CreateModel( name="Provider", fields=[ @@ -81,15 +102,14 @@ class Migration(migrations.Migration): ("updated_at", models.DateTimeField(auto_now=True)), ( "provider", - models.CharField( + ProviderEnumField( choices=[ - ("aws", "Aws"), + ("aws", "AWS"), ("azure", "Azure"), - ("gcp", "Gcp"), + ("gcp", "GCP"), ("kubernetes", "Kubernetes"), ], default="aws", - max_length=10, ), ), ( @@ -113,8 +133,8 @@ class Migration(migrations.Migration): "connection_last_checked_at", models.DateTimeField(blank=True, null=True), ), - ("metadata", models.JSONField(default=dict)), - ("scanner_args", models.JSONField(default=dict)), + ("metadata", models.JSONField(blank=True, default=dict)), + ("scanner_args", models.JSONField(blank=True, default=dict)), ( "tenant", models.ForeignKey( @@ -124,6 +144,7 @@ class Migration(migrations.Migration): ], options={ "abstract": False, + "db_table": "providers", }, ), migrations.AddConstraint( diff --git a/src/backend/api/models.py b/src/backend/api/models.py index 4aefca56e2..3737c6cb19 100644 --- a/src/backend/api/models.py +++ b/src/backend/api/models.py @@ -4,7 +4,7 @@ from uuid import uuid4, UUID from django.core.validators import MinLengthValidator from django.db import models from django.utils.translation import gettext_lazy as _ - +from api.db_utils import ProviderEnumField from api.exceptions import ModelValidationError from api.rls import RowLevelSecurityConstraint from api.rls import RowLevelSecurityProtectedModel @@ -62,8 +62,8 @@ class Provider(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.CharField( - max_length=10, choices=ProviderChoices.choices, default=ProviderChoices.AWS + provider = ProviderEnumField( + choices=ProviderChoices.choices, default=ProviderChoices.AWS ) provider_id = models.CharField(max_length=63, validators=[MinLengthValidator(3)]) alias = models.CharField( @@ -83,6 +83,8 @@ class Provider(RowLevelSecurityProtectedModel): super().save(*args, **kwargs) class Meta(RowLevelSecurityProtectedModel.Meta): + db_table = "providers" + constraints = [ models.UniqueConstraint( fields=("tenant_id", "provider", "provider_id"), diff --git a/src/backend/api/rls.py b/src/backend/api/rls.py index 50db74a3bb..ed5d5a83c9 100644 --- a/src/backend/api/rls.py +++ b/src/backend/api/rls.py @@ -20,7 +20,7 @@ class Tenant(models.Model): name = models.CharField(max_length=100) class Meta: - db_table = "tenant" + db_table = "tenants" # TODO Add abstract class for non-RLS models diff --git a/src/backend/api/v1/serializers.py b/src/backend/api/v1/serializers.py index 7906f3f2f8..bd04da63cf 100644 --- a/src/backend/api/v1/serializers.py +++ b/src/backend/api/v1/serializers.py @@ -45,6 +45,10 @@ class TenantSerializer(BaseSerializerV1): # Providers +class ProviderEnumField(serializers.ChoiceField): + def __init__(self, **kwargs): + kwargs["choices"] = Provider.ProviderChoices.choices + super().__init__(**kwargs) class ProviderSerializer(RLSSerializer): @@ -52,6 +56,7 @@ class ProviderSerializer(RLSSerializer): Serializer for the Provider model. """ + provider = ProviderEnumField(choices=Provider.ProviderChoices.choices) connection = serializers.SerializerMethodField(read_only=True) class Meta: diff --git a/src/backend/api/v1/views.py b/src/backend/api/v1/views.py index 79693aedd3..61ad33531b 100644 --- a/src/backend/api/v1/views.py +++ b/src/backend/api/v1/views.py @@ -11,14 +11,15 @@ from rest_framework.decorators import action from rest_framework.generics import get_object_or_404 from rest_framework_json_api.views import Response -from api.base_views import BaseRLSViewSet -from api.base_views import BaseViewSet -from api.filters import ProviderFilter -from api.filters import TenantFilter +from api.base_views import BaseRLSViewSet, BaseViewSet +from api.filters import ProviderFilter, TenantFilter from api.models import Provider from api.rls import Tenant -from api.v1.serializers import ProviderSerializer, ProviderUpdateSerializer -from api.v1.serializers import TenantSerializer +from api.v1.serializers import ( + ProviderSerializer, + ProviderUpdateSerializer, + TenantSerializer, +) CACHE_DECORATOR = cache_control( max_age=django_settings.CACHE_MAX_AGE,