mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-24 04:51:51 +00:00
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:
committed by
GitHub
parent
8a2cfea677
commit
8f2bd45872
@@ -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)
|
||||
@@ -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,
|
||||
},
|
||||
}
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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"),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user