fix: move lighthouse validation from model to serializer

This commit is contained in:
Chandrapal Badshah
2025-10-09 11:54:50 +05:30
parent fb870711d9
commit a56686bf1a
2 changed files with 182 additions and 131 deletions
+1 -131
View File
@@ -1872,22 +1872,6 @@ class LighthouseConfiguration(RowLevelSecurityProtectedModel):
def clean(self):
super().clean()
# Validate temperature
if not 0 <= self.temperature <= 1:
raise ModelValidationError(
detail="Temperature must be between 0 and 1",
code="invalid_temperature",
pointer="/data/attributes/temperature",
)
# Validate max_tokens
if not 500 <= self.max_tokens <= 5000:
raise ModelValidationError(
detail="Max tokens must be between 500 and 5000",
code="invalid_max_tokens",
pointer="/data/attributes/max_tokens",
)
@property
def api_key_decoded(self):
"""Return the decrypted API key, or None if unavailable or invalid."""
@@ -1912,15 +1896,6 @@ class LighthouseConfiguration(RowLevelSecurityProtectedModel):
code="invalid_api_key",
pointer="/data/attributes/api_key",
)
# Validate OpenAI API key format
openai_key_pattern = r"^sk-[\w-]+T3BlbkFJ[\w-]+$"
if not re.match(openai_key_pattern, value):
raise ModelValidationError(
detail="Invalid OpenAI API key format.",
code="invalid_api_key",
pointer="/data/attributes/api_key",
)
self.api_key = fernet.encrypt(value.encode())
def save(self, *args, **kwargs):
@@ -1995,35 +1970,6 @@ class LighthouseProviderConfiguration(RowLevelSecurityProtectedModel):
class ProviderChoices(models.TextChoices):
OPENAI = "openai", _("OpenAI")
@staticmethod
def validate_openai_credentials(credentials: dict):
"""
Validate OpenAI provider credentials.
Checks:
- api_key is present and is a string
- api_key matches expected OpenAI format
Raises:
ModelValidationError: If credentials are invalid
"""
api_key = credentials.get("api_key")
if not isinstance(api_key, str) or not api_key:
raise ModelValidationError(
detail="OpenAI credentials must include 'api_key'",
code="invalid_openai_credentials",
pointer="/data/attributes/credentials/api_key",
)
# Validate OpenAI API key format
openai_key_pattern = r"^sk-[\w-]+T3BlbkFJ[\w-]+$"
if not re.match(openai_key_pattern, api_key):
raise ModelValidationError(
detail="Invalid OpenAI API key format.",
code="invalid_openai_api_key_format",
pointer="/data/attributes/credentials/api_key",
)
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)
@@ -2069,16 +2015,7 @@ class LighthouseProviderConfiguration(RowLevelSecurityProtectedModel):
@credentials_decoded.setter
def credentials_decoded(self, value):
"""
Set and encrypt credentials.
This is the single source of truth for credential validation.
Validates based on provider_type and encrypts before storage.
Args:
value: Dict containing provider-specific credentials
Raises:
ModelValidationError: If credentials are invalid or missing
Set and encrypt credentials (assumes serializer performed validation).
"""
if not value:
raise ModelValidationError(
@@ -2086,11 +2023,6 @@ class LighthouseProviderConfiguration(RowLevelSecurityProtectedModel):
code="invalid_credentials",
pointer="/data/attributes/credentials",
)
# Validate credentials based on provider type
if self.provider_type == self.ProviderChoices.OPENAI:
self.validate_openai_credentials(value)
self.credentials = fernet.encrypt(json.dumps(value).encode())
def delete(self, *args, **kwargs):
@@ -2173,68 +2105,6 @@ class LighthouseTenantConfiguration(RowLevelSecurityProtectedModel):
def clean(self):
super().clean()
# Validate default_provider is supported, configured and active (if provided)
if self.default_provider:
# Check if provider type is supported
supported_providers = set(
LighthouseProviderConfiguration.ProviderChoices.values
)
if self.default_provider not in supported_providers:
raise ModelValidationError(
detail=f"Unsupported provider: '{self.default_provider}'. Supported providers: {', '.join(supported_providers)}",
code="unsupported_provider",
pointer="/data/attributes/default_provider",
)
# Check if provider is configured and active
if not LighthouseProviderConfiguration.objects.filter(
tenant_id=self.tenant_id,
provider_type=self.default_provider,
is_active=True,
).exists():
raise ModelValidationError(
detail=f"No active configuration found for provider '{self.default_provider}'",
code="invalid_default_provider",
pointer="/data/attributes/default_provider",
)
# Validate default_models mapping
if self.default_models is not None and not isinstance(
self.default_models, dict
):
raise ModelValidationError(
detail="default_models must be an object mapping provider->model",
code="invalid_default_models",
pointer="/data/attributes/default_models",
)
for provider_type, model_id in (self.default_models or {}).items():
# Provider must exist and be active
provider_cfg = LighthouseProviderConfiguration.objects.filter(
tenant_id=self.tenant_id,
provider_type=provider_type,
is_active=True,
).first()
if not provider_cfg:
raise ModelValidationError(
detail=f"No active configuration found for provider '{provider_type}'",
code="invalid_default_models_provider",
pointer="/data/attributes/default_models",
)
# Model must exist under that provider configuration
if not LighthouseProviderModels.objects.filter(
tenant_id=self.tenant_id,
provider_configuration=provider_cfg,
model_id=model_id,
).exists():
raise ModelValidationError(
detail=f"Invalid model '{model_id}' for provider '{provider_type}'",
code="invalid_default_models_model",
pointer="/data/attributes/default_models",
)
class Meta(RowLevelSecurityProtectedModel.Meta):
db_table = "lighthouse_tenant_config"
+181
View File
@@ -59,6 +59,7 @@ from api.v1.serializer_utils.integrations import (
S3ConfigSerializer,
SecurityHubConfigSerializer,
)
from api.v1.serializer_utils.lighthouse import OpenAICredentialsSerializer
from api.v1.serializer_utils.processors import ProcessorConfigField
from api.v1.serializer_utils.providers import ProviderSecretField
from prowler.lib.mutelist.mutelist import Mutelist
@@ -2755,6 +2756,16 @@ class LighthouseConfigCreateSerializer(RLSSerializer, BaseWriteSerializer):
"updated_at": {"read_only": True},
}
def validate_temperature(self, value):
if not 0 <= value <= 1:
raise ValidationError("Temperature must be between 0 and 1.")
return value
def validate_max_tokens(self, value):
if not 500 <= value <= 5000:
raise ValidationError("Max tokens must be between 500 and 5000.")
return value
def validate(self, attrs):
tenant_id = self.context.get("request").tenant_id
if LighthouseConfiguration.objects.filter(tenant_id=tenant_id).exists():
@@ -2763,6 +2774,11 @@ class LighthouseConfigCreateSerializer(RLSSerializer, BaseWriteSerializer):
"tenant_id": "Lighthouse configuration already exists for this tenant."
}
)
api_key = attrs.get("api_key")
if api_key is not None:
OpenAICredentialsSerializer(data={"api_key": api_key}).is_valid(
raise_exception=True
)
return super().validate(attrs)
def create(self, validated_data):
@@ -2807,6 +2823,24 @@ class LighthouseConfigUpdateSerializer(BaseWriteSerializer):
"max_tokens": {"required": False},
}
def validate_temperature(self, value):
if not 0 <= value <= 1:
raise ValidationError("Temperature must be between 0 and 1.")
return value
def validate_max_tokens(self, value):
if not 500 <= value <= 5000:
raise ValidationError("Max tokens must be between 500 and 5000.")
return value
def validate(self, attrs):
api_key = attrs.get("api_key", None)
if api_key is not None:
OpenAICredentialsSerializer(data={"api_key": api_key}).is_valid(
raise_exception=True
)
return super().validate(attrs)
def update(self, instance, validated_data):
api_key = validated_data.pop("api_key", None)
instance = super().update(instance, validated_data)
@@ -3042,6 +3076,24 @@ class LighthouseProviderConfigCreateSerializer(RLSSerializer, BaseWriteSerialize
}
)
def validate(self, attrs):
provider_type = attrs.get("provider_type")
credentials = attrs.get("credentials") or {}
if provider_type == LighthouseProviderConfiguration.ProviderChoices.OPENAI:
try:
OpenAICredentialsSerializer(data=credentials).is_valid(
raise_exception=True
)
except ValidationError as e:
details = e.detail.copy()
for key, value in details.items():
e.detail[f"credentials/{key}"] = value
del e.detail[key]
raise e
return super().validate(attrs)
class LighthouseProviderConfigUpdateSerializer(BaseWriteSerializer):
"""
@@ -3078,6 +3130,27 @@ class LighthouseProviderConfigUpdateSerializer(BaseWriteSerializer):
instance.save()
return instance
def validate(self, attrs):
provider_type = getattr(self.instance, "provider_type", None)
credentials = attrs.get("credentials", None)
if (
credentials is not None
and provider_type == LighthouseProviderConfiguration.ProviderChoices.OPENAI
):
try:
OpenAICredentialsSerializer(data=credentials).is_valid(
raise_exception=True
)
except ValidationError as e:
details = e.detail.copy()
for key, value in details.items():
e.detail[f"credentials/{key}"] = value
del e.detail[key]
raise e
return super().validate(attrs)
# Lighthouse: Tenant configuration
@@ -3128,6 +3201,58 @@ class LighthouseTenantConfigCreateSerializer(RLSSerializer, BaseWriteSerializer)
"default_models",
]
def validate(self, attrs):
request = self.context.get("request")
tenant_id = self.context.get("tenant_id") or (
getattr(request, "tenant_id", None) if request else None
)
default_provider = attrs.get("default_provider", "")
default_models = attrs.get("default_models", {})
if default_provider:
supported = set(LighthouseProviderConfiguration.ProviderChoices.values)
if default_provider not in supported:
raise ValidationError(
{"default_provider": f"Unsupported provider '{default_provider}'."}
)
if not LighthouseProviderConfiguration.objects.filter(
tenant_id=tenant_id, provider_type=default_provider, is_active=True
).exists():
raise ValidationError(
{
"default_provider": f"No active configuration found for '{default_provider}'."
}
)
if default_models is not None and not isinstance(default_models, dict):
raise ValidationError(
{"default_models": "Must be an object mapping provider -> model_id."}
)
for provider_type, model_id in (default_models or {}).items():
provider_cfg = LighthouseProviderConfiguration.objects.filter(
tenant_id=tenant_id, provider_type=provider_type, is_active=True
).first()
if not provider_cfg:
raise ValidationError(
{
"default_models": f"No active configuration for provider '{provider_type}'."
}
)
if not LighthouseProviderModels.objects.filter(
tenant_id=tenant_id,
provider_configuration=provider_cfg,
model_id=model_id,
).exists():
raise ValidationError(
{
"default_models": f"Invalid model '{model_id}' for provider '{provider_type}'."
}
)
return super().validate(attrs)
def create(self, validated_data):
try:
return super().create(validated_data)
@@ -3150,6 +3275,62 @@ class LighthouseTenantConfigUpdateSerializer(BaseWriteSerializer):
"id": {"read_only": True},
}
def validate(self, attrs):
request = self.context.get("request")
tenant_id = self.context.get("tenant_id") or (
getattr(request, "tenant_id", None) if request else None
)
default_provider = attrs.get(
"default_provider", getattr(self.instance, "default_provider", "")
)
default_models = attrs.get(
"default_models", getattr(self.instance, "default_models", {})
)
if default_provider:
supported = set(LighthouseProviderConfiguration.ProviderChoices.values)
if default_provider not in supported:
raise ValidationError(
{"default_provider": f"Unsupported provider '{default_provider}'."}
)
if not LighthouseProviderConfiguration.objects.filter(
tenant_id=tenant_id, provider_type=default_provider, is_active=True
).exists():
raise ValidationError(
{
"default_provider": f"No active configuration found for '{default_provider}'."
}
)
if default_models is not None and not isinstance(default_models, dict):
raise ValidationError(
{"default_models": "Must be an object mapping provider -> model_id."}
)
for provider_type, model_id in (default_models or {}).items():
provider_cfg = LighthouseProviderConfiguration.objects.filter(
tenant_id=tenant_id, provider_type=provider_type, is_active=True
).first()
if not provider_cfg:
raise ValidationError(
{
"default_models": f"No active configuration for provider '{provider_type}'."
}
)
if not LighthouseProviderModels.objects.filter(
tenant_id=tenant_id,
provider_configuration=provider_cfg,
model_id=model_id,
).exists():
raise ValidationError(
{
"default_models": f"Invalid model '{model_id}' for provider '{provider_type}'."
}
)
return super().validate(attrs)
# Lighthouse: Provider models