mirror of
https://github.com/prowler-cloud/prowler.git
synced 2026-07-23 12:31:54 +00:00
chore: refactor code to deduplicate validation of llm provider creds
This commit is contained in:
@@ -6,6 +6,7 @@ from django.conf import settings
|
||||
from django.contrib.auth import authenticate
|
||||
from django.contrib.auth.models import update_last_login
|
||||
from django.contrib.auth.password_validation import validate_password
|
||||
from django.db import IntegrityError
|
||||
from drf_spectacular.utils import extend_schema_field
|
||||
from jwt.exceptions import InvalidKeyError
|
||||
from rest_framework.reverse import reverse
|
||||
@@ -18,7 +19,7 @@ from rest_framework_simplejwt.serializers import TokenObtainPairSerializer
|
||||
from rest_framework_simplejwt.tokens import RefreshToken
|
||||
|
||||
from api.db_router import MainRouter
|
||||
from api.exceptions import ConflictException, ModelValidationError
|
||||
from api.exceptions import ConflictException
|
||||
from api.models import (
|
||||
Finding,
|
||||
Integration,
|
||||
@@ -2964,9 +2965,12 @@ class LighthouseProviderConfigSerializer(RLSSerializer):
|
||||
"inserted_at": {"read_only": True},
|
||||
"updated_at": {"read_only": True},
|
||||
"is_active": {"read_only": True},
|
||||
"url": {"read_only": True},
|
||||
"url": {"read_only": True, "view_name": "lighthouse-providers-detail"},
|
||||
}
|
||||
|
||||
class JSONAPIMeta:
|
||||
resource_name = "lighthouse-providers"
|
||||
|
||||
def to_representation(self, instance):
|
||||
data = super().to_representation(instance)
|
||||
# Support JSON:API fields filter: fields[lighthouse-providers]=credentials
|
||||
@@ -2990,7 +2994,12 @@ class LighthouseProviderConfigSerializer(RLSSerializer):
|
||||
return [mask_value(v) for v in value]
|
||||
return value
|
||||
|
||||
data["credentials"] = mask_value(creds) if creds is not None else None
|
||||
# Always return masked credentials, even if creds is None
|
||||
if creds is not None:
|
||||
data["credentials"] = mask_value(creds)
|
||||
else:
|
||||
# If credentials_decoded returns None, return None for credentials field
|
||||
data["credentials"] = None
|
||||
|
||||
return data
|
||||
|
||||
@@ -3001,7 +3010,7 @@ class LighthouseProviderConfigCreateSerializer(RLSSerializer, BaseWriteSerialize
|
||||
Accepts credentials as JSON; stored encrypted via credentials_decoded.
|
||||
"""
|
||||
|
||||
credentials = serializers.JSONField(write_only=True)
|
||||
credentials = serializers.JSONField(write_only=True, required=True)
|
||||
|
||||
class Meta:
|
||||
model = LighthouseProviderConfiguration
|
||||
@@ -3016,25 +3025,22 @@ class LighthouseProviderConfigCreateSerializer(RLSSerializer, BaseWriteSerialize
|
||||
"base_url": {"required": False, "allow_null": True},
|
||||
}
|
||||
|
||||
def validate(self, attrs):
|
||||
tenant_id = self.context.get("tenant_id")
|
||||
provider_type = attrs.get("provider_type")
|
||||
if LighthouseProviderConfiguration.objects.filter(
|
||||
tenant_id=tenant_id, provider_type=provider_type
|
||||
).exists():
|
||||
raise ValidationError(
|
||||
{
|
||||
"provider_type": "Configuration for this provider already exists for the tenant.",
|
||||
}
|
||||
)
|
||||
return super().validate(attrs)
|
||||
|
||||
def create(self, validated_data):
|
||||
credentials = validated_data.pop("credentials")
|
||||
instance = super().create(validated_data)
|
||||
|
||||
instance = LighthouseProviderConfiguration(**validated_data)
|
||||
instance.tenant_id = self.context.get("tenant_id")
|
||||
instance.credentials_decoded = credentials
|
||||
instance.save()
|
||||
return instance
|
||||
|
||||
try:
|
||||
instance.save()
|
||||
return instance
|
||||
except IntegrityError:
|
||||
raise ValidationError(
|
||||
{
|
||||
"provider_type": "Configuration for this provider already exists for the tenant."
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class LighthouseProviderConfigUpdateSerializer(BaseWriteSerializer):
|
||||
@@ -3062,10 +3068,14 @@ class LighthouseProviderConfigUpdateSerializer(BaseWriteSerializer):
|
||||
|
||||
def update(self, instance, validated_data):
|
||||
credentials = validated_data.pop("credentials", None)
|
||||
instance = super().update(instance, validated_data)
|
||||
|
||||
for attr, value in validated_data.items():
|
||||
setattr(instance, attr, value)
|
||||
|
||||
if credentials is not None:
|
||||
instance.credentials_decoded = credentials
|
||||
instance.save()
|
||||
|
||||
instance.save()
|
||||
return instance
|
||||
|
||||
|
||||
@@ -3106,7 +3116,8 @@ class LighthouseTenantConfigSerializer(RLSSerializer):
|
||||
class LighthouseTenantConfigCreateSerializer(RLSSerializer, BaseWriteSerializer):
|
||||
"""
|
||||
Create serializer for LighthouseTenantConfiguration.
|
||||
Ensures one record per tenant.
|
||||
|
||||
Note: One record per tenant constraint is enforced by database and model.clean().
|
||||
"""
|
||||
|
||||
class Meta:
|
||||
@@ -3117,13 +3128,13 @@ class LighthouseTenantConfigCreateSerializer(RLSSerializer, BaseWriteSerializer)
|
||||
"default_models",
|
||||
]
|
||||
|
||||
def validate(self, attrs):
|
||||
tenant_id = self.context.get("tenant_id")
|
||||
if LighthouseTenantConfiguration.objects.filter(tenant_id=tenant_id).exists():
|
||||
def create(self, validated_data):
|
||||
try:
|
||||
return super().create(validated_data)
|
||||
except IntegrityError:
|
||||
raise ValidationError(
|
||||
{"tenant_id": "Tenant Lighthouse configuration already exists."}
|
||||
)
|
||||
return super().validate(attrs)
|
||||
|
||||
|
||||
class LighthouseTenantConfigUpdateSerializer(BaseWriteSerializer):
|
||||
@@ -3139,47 +3150,6 @@ class LighthouseTenantConfigUpdateSerializer(BaseWriteSerializer):
|
||||
"id": {"read_only": True},
|
||||
}
|
||||
|
||||
def validate(self, attrs):
|
||||
"""
|
||||
Validate only that default_provider (if supplied):
|
||||
- is a supported provider type, and
|
||||
- has an active provider configuration for this tenant.
|
||||
"""
|
||||
validated_attrs = super().validate(attrs)
|
||||
|
||||
default_provider = validated_attrs.get("default_provider", None)
|
||||
|
||||
# To set a default provider, it must be one of the supported provider types and have an active configuration for this tenant
|
||||
if default_provider is not None and default_provider != "":
|
||||
supported_providers = [
|
||||
choice[0]
|
||||
for choice in LighthouseProviderConfiguration.ProviderChoices.choices
|
||||
]
|
||||
if default_provider not in supported_providers:
|
||||
raise ModelValidationError(
|
||||
detail=f"Unsupported provider '{default_provider}'",
|
||||
code="unsupported_default_provider",
|
||||
pointer="/data/attributes/default_provider",
|
||||
)
|
||||
|
||||
tenant_id = (
|
||||
self.instance.tenant_id
|
||||
if self.instance
|
||||
else self.context.get("tenant_id")
|
||||
)
|
||||
if not LighthouseProviderConfiguration.objects.filter(
|
||||
tenant_id=tenant_id,
|
||||
provider_type=default_provider,
|
||||
is_active=True,
|
||||
).exists():
|
||||
raise ModelValidationError(
|
||||
detail=f"No active configuration found for provider '{default_provider}'",
|
||||
code="invalid_default_provider",
|
||||
pointer="/data/attributes/default_provider",
|
||||
)
|
||||
|
||||
return validated_attrs
|
||||
|
||||
|
||||
# Lighthouse: Provider models
|
||||
|
||||
|
||||
Reference in New Issue
Block a user