feat/PRWLR-4413 Add Postgres Enums for Django and update Provider.provider field (#28)

* feat(db): PRWLR-4413 add Provider Postgres Enum type for Django

* fix(Backend): PRWLR-4413 Fix initial migration for Providers

* feat(Backend): PRWLR-4413 add provider enum to Provider model

* fix(Backend, API): PRWLR-4413 fix ProviderEnum representation

* chore(Backend): PRWLR-4413 remove max_length constraint from provider enum

* chore(Backend): PRWLR-4413 refactor postgres enum creation to avoid boilerplate

* chore(Backend): PRWLR-4413 improve comments
This commit is contained in:
Víctor Fernández Poyatos
2024-08-21 18:02:46 +02:00
committed by GitHub
parent 8a2cfea677
commit 8f2bd45872
7 changed files with 165 additions and 19 deletions
+111
View File
@@ -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)
+7 -1
View File
@@ -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,
},
}
+29 -8
View File
@@ -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(
+5 -3
View File
@@ -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"),
+1 -1
View File
@@ -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
+5
View File
@@ -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:
+7 -6
View File
@@ -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,