From a56686bf1a83768acddd29e2af99690cbc4d0e33 Mon Sep 17 00:00:00 2001 From: Chandrapal Badshah <12944530+Chan9390@users.noreply.github.com> Date: Thu, 9 Oct 2025 11:54:50 +0530 Subject: [PATCH] fix: move lighthouse validation from model to serializer --- api/src/backend/api/models.py | 132 +------------------ api/src/backend/api/v1/serializers.py | 181 ++++++++++++++++++++++++++ 2 files changed, 182 insertions(+), 131 deletions(-) diff --git a/api/src/backend/api/models.py b/api/src/backend/api/models.py index 68f901996f..4631e3f422 100644 --- a/api/src/backend/api/models.py +++ b/api/src/backend/api/models.py @@ -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" diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 2d33b64aec..c00d43aa41 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -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