From 4cb2778bfa578f287b4067c87410e9d758d0801b Mon Sep 17 00:00:00 2001 From: Chandrapal Badshah <12944530+Chan9390@users.noreply.github.com> Date: Wed, 15 Oct 2025 18:01:29 +0530 Subject: [PATCH] feat: add amazon bedrock --- .../api/migrations/0051_lighthouse_bedrock.py | 21 ++ api/src/backend/api/models.py | 1 + .../api/v1/serializer_utils/lighthouse.py | 39 +++ api/src/backend/api/v1/serializers.py | 33 +- api/src/backend/api/v1/views.py | 24 +- .../tasks/jobs/lighthouse_providers.py | 315 +++++++++++++++--- 6 files changed, 357 insertions(+), 76 deletions(-) create mode 100644 api/src/backend/api/migrations/0051_lighthouse_bedrock.py diff --git a/api/src/backend/api/migrations/0051_lighthouse_bedrock.py b/api/src/backend/api/migrations/0051_lighthouse_bedrock.py new file mode 100644 index 0000000000..3107619c90 --- /dev/null +++ b/api/src/backend/api/migrations/0051_lighthouse_bedrock.py @@ -0,0 +1,21 @@ +# Generated by Django 5.1.12 on 2025-10-14 11:46 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("api", "0050_lighthouse_multi_llm"), + ] + + operations = [ + migrations.AlterField( + model_name="lighthouseproviderconfiguration", + name="provider_type", + field=models.CharField( + choices=[("openai", "OpenAI"), ("bedrock", "AWS Bedrock")], + help_text="LLM provider name", + max_length=50, + ), + ) + ] diff --git a/api/src/backend/api/models.py b/api/src/backend/api/models.py index c2c0dbcb47..def2d249f4 100644 --- a/api/src/backend/api/models.py +++ b/api/src/backend/api/models.py @@ -1970,6 +1970,7 @@ class LighthouseProviderConfiguration(RowLevelSecurityProtectedModel): class LLMProviderChoices(models.TextChoices): OPENAI = "openai", _("OpenAI") + BEDROCK = "bedrock", _("AWS Bedrock") id = models.UUIDField(primary_key=True, default=uuid4, editable=False) inserted_at = models.DateTimeField(auto_now_add=True, editable=False) diff --git a/api/src/backend/api/v1/serializer_utils/lighthouse.py b/api/src/backend/api/v1/serializer_utils/lighthouse.py index f7ecef32cc..664caec34a 100644 --- a/api/src/backend/api/v1/serializer_utils/lighthouse.py +++ b/api/src/backend/api/v1/serializer_utils/lighthouse.py @@ -11,3 +11,42 @@ class OpenAICredentialsSerializer(serializers.Serializer): if not re.match(pattern, value or ""): raise serializers.ValidationError("Invalid OpenAI API key format.") return value + + +class BedrockCredentialsSerializer(serializers.Serializer): + """ + Serializer for AWS Bedrock credentials validation. + + Validates long-term AWS credentials (AKIA) and region format. + """ + + access_key_id = serializers.CharField() + secret_access_key = serializers.CharField() + region = serializers.CharField() + + def validate_access_key_id(self, value: str) -> str: + """Validate AWS access key ID format (AKIA for long-term credentials).""" + pattern = r"^AKIA[0-9A-Z]{16}$" + if not re.match(pattern, value or ""): + raise serializers.ValidationError( + "Invalid AWS access key ID format. Must be AKIA followed by 16 alphanumeric characters." + ) + return value + + def validate_secret_access_key(self, value: str) -> str: + """Validate AWS secret access key format (40 base64 characters).""" + pattern = r"^[A-Za-z0-9/+=]{40}$" + if not re.match(pattern, value or ""): + raise serializers.ValidationError( + "Invalid AWS secret access key format. Must be 40 base64 characters." + ) + return value + + def validate_region(self, value: str) -> str: + """Validate AWS region format.""" + pattern = r"^[a-z]{2}-[a-z]+-\d+$" + if not re.match(pattern, value or ""): + raise serializers.ValidationError( + "Invalid AWS region format. Expected format like 'us-east-1' or 'eu-west-2'." + ) + return value diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 511b1f89e4..359b79dd09 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -59,7 +59,10 @@ from api.v1.serializer_utils.integrations import ( S3ConfigSerializer, SecurityHubConfigSerializer, ) -from api.v1.serializer_utils.lighthouse import OpenAICredentialsSerializer +from api.v1.serializer_utils.lighthouse import ( + BedrockCredentialsSerializer, + OpenAICredentialsSerializer, +) from api.v1.serializer_utils.processors import ProcessorConfigField from api.v1.serializer_utils.providers import ProviderSecretField from prowler.lib.mutelist.mutelist import Mutelist @@ -3096,6 +3099,19 @@ class LighthouseProviderConfigCreateSerializer(RLSSerializer, BaseWriteSerialize e.detail[f"credentials/{key}"] = value del e.detail[key] raise e + elif ( + provider_type == LighthouseProviderConfiguration.LLMProviderChoices.BEDROCK + ): + try: + BedrockCredentialsSerializer(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) @@ -3154,6 +3170,21 @@ class LighthouseProviderConfigUpdateSerializer(BaseWriteSerializer): e.detail[f"credentials/{key}"] = value del e.detail[key] raise e + elif ( + credentials is not None + and provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.BEDROCK + ): + try: + BedrockCredentialsSerializer(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) diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index 6746d7490c..115bf13ede 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -4350,23 +4350,13 @@ class LighthouseProviderConfigViewSet(BaseRLSViewSet): @extend_schema( tags=["Lighthouse AI"], summary="Check LLM provider connection", - description="Validate provider credentials asynchronously and toggle is_active.", + description="Validate provider credentials asynchronously and toggle is_active. Supports OpenAI and AWS Bedrock providers.", request=None, responses={202: OpenApiResponse(response=TaskSerializer)}, ) @action(detail=True, methods=["post"], url_name="connection") def connection(self, request, pk=None): instance = self.get_object() - if ( - instance.provider_type - != LighthouseProviderConfiguration.LLMProviderChoices.OPENAI - ): - return Response( - data={ - "errors": [{"detail": "Only 'openai' provider supported in MVP"}] - }, - status=status.HTTP_400_BAD_REQUEST, - ) with transaction.atomic(): task = check_lighthouse_provider_connection_task.delay( @@ -4388,7 +4378,7 @@ class LighthouseProviderConfigViewSet(BaseRLSViewSet): @extend_schema( tags=["Lighthouse AI"], summary="Refresh LLM models catalog", - description="Fetch available models for this provider configuration and upsert into catalog.", + description="Fetch available models for this provider configuration and upsert into catalog. Supports OpenAI and AWS Bedrock providers.", request=None, responses={202: OpenApiResponse(response=TaskSerializer)}, ) @@ -4400,16 +4390,6 @@ class LighthouseProviderConfigViewSet(BaseRLSViewSet): ) def refresh_models(self, request, pk=None): instance = self.get_object() - if ( - instance.provider_type - != LighthouseProviderConfiguration.LLMProviderChoices.OPENAI - ): - return Response( - data={ - "errors": [{"detail": "Only 'openai' provider supported in MVP"}] - }, - status=status.HTTP_400_BAD_REQUEST, - ) with transaction.atomic(): task = refresh_lighthouse_provider_models_task.delay( diff --git a/api/src/backend/tasks/jobs/lighthouse_providers.py b/api/src/backend/tasks/jobs/lighthouse_providers.py index 1a68a5c1e2..7a9af6ab1d 100644 --- a/api/src/backend/tasks/jobs/lighthouse_providers.py +++ b/api/src/backend/tasks/jobs/lighthouse_providers.py @@ -1,6 +1,8 @@ -from typing import Dict, Set +from typing import Dict +import boto3 import openai +from botocore.exceptions import BotoCoreError, ClientError from celery.utils.log import get_task_logger from api.models import LighthouseProviderConfiguration, LighthouseProviderModels @@ -30,16 +32,53 @@ def _extract_openai_api_key( return api_key +def _extract_bedrock_credentials( + provider_cfg: LighthouseProviderConfiguration, +) -> Dict[str, str] | None: + """ + Safely extract AWS Bedrock credentials from a provider configuration. + + Args: + provider_cfg (LighthouseProviderConfiguration): The provider configuration instance + containing the credentials. + + Returns: + Dict[str, str] | None: Dictionary with 'access_key_id', 'secret_access_key', and + 'region' if present and valid, otherwise None. + """ + creds = provider_cfg.credentials_decoded + if not isinstance(creds, dict): + return None + + access_key_id = creds.get("access_key_id") + secret_access_key = creds.get("secret_access_key") + region = creds.get("region") + + # Validate all required fields are present and are strings + if ( + not isinstance(access_key_id, str) + or not access_key_id + or not isinstance(secret_access_key, str) + or not secret_access_key + or not isinstance(region, str) + or not region + ): + return None + + return { + "access_key_id": access_key_id, + "secret_access_key": secret_access_key, + "region": region, + } + + def check_lighthouse_provider_connection(provider_config_id: str) -> Dict: """ Validate a Lighthouse provider configuration by calling the provider API and toggle its active state accordingly. - Currently supports the OpenAI provider by invoking `models.list` to verify that - the provided credentials are valid. - Args: - provider_config_id (str): The primary key of the `LighthouseProviderConfiguration` + provider_config_id: The primary key of the `LighthouseProviderConfiguration` to validate. Returns: @@ -55,42 +94,192 @@ def check_lighthouse_provider_connection(provider_config_id: str) -> Dict: """ provider_cfg = LighthouseProviderConfiguration.objects.get(pk=provider_config_id) - # TODO: Add support for other providers - if ( - provider_cfg.provider_type - != LighthouseProviderConfiguration.LLMProviderChoices.OPENAI - ): - return {"connected": False, "error": "Unsupported provider type"} - - api_key = _extract_openai_api_key(provider_cfg) - if not api_key: - provider_cfg.is_active = False - provider_cfg.save() - return {"connected": False, "error": "API key is invalid or missing"} - try: - client = openai.OpenAI(api_key=api_key) - _ = client.models.list() + if ( + provider_cfg.provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.OPENAI + ): + api_key = _extract_openai_api_key(provider_cfg) + if not api_key: + provider_cfg.is_active = False + provider_cfg.save() + return {"connected": False, "error": "API key is invalid or missing"} + + # Test connection by listing models + client = openai.OpenAI(api_key=api_key) + _ = client.models.list() + + elif ( + provider_cfg.provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.BEDROCK + ): + bedrock_creds = _extract_bedrock_credentials(provider_cfg) + if not bedrock_creds: + provider_cfg.is_active = False + provider_cfg.save() + return { + "connected": False, + "error": "AWS credentials are invalid or missing", + } + + # Test connection by listing foundation models + bedrock_client = boto3.client( + "bedrock", + aws_access_key_id=bedrock_creds["access_key_id"], + aws_secret_access_key=bedrock_creds["secret_access_key"], + region_name=bedrock_creds["region"], + ) + _ = bedrock_client.list_foundation_models() + + else: + return {"connected": False, "error": "Unsupported provider type"} + + # Connection successful provider_cfg.is_active = True provider_cfg.save() return {"connected": True, "error": None} + except Exception as e: - logger.warning("OpenAI connection check failed: %s", str(e)) + logger.warning( + "%s connection check failed: %s", provider_cfg.provider_type, str(e) + ) provider_cfg.is_active = False provider_cfg.save() return {"connected": False, "error": str(e)} +def _fetch_openai_models(api_key: str) -> Dict[str, str]: + """ + Fetch available models from OpenAI API. + + Args: + api_key: OpenAI API key for authentication. + + Returns: + Dict mapping model_id to model_name. For OpenAI, both are the same + as the API doesn't provide separate display names. + + Raises: + Exception: If the API call fails. + """ + client = openai.OpenAI(api_key=api_key) + models = client.models.list() + # OpenAI uses model.id for both ID and display name + return {m.id: m.id for m in getattr(models, "data", [])} + + +def _fetch_bedrock_models(bedrock_creds: Dict[str, str]) -> Dict[str, str]: + """ + Fetch available models from AWS Bedrock with entitlement verification. + + This function: + 1. Lists foundation models with TEXT modality support + 2. Lists inference profiles with TEXT modality support + 3. Verifies user has entitlement access to each model + + Args: + bedrock_creds: Dictionary with 'access_key_id', 'secret_access_key', and 'region'. + + Returns: + Dict mapping model_id to model_name for all accessible models. + + Raises: + BotoCoreError, ClientError: If AWS API calls fail. + """ + bedrock_client = boto3.client( + "bedrock", + aws_access_key_id=bedrock_creds["access_key_id"], + aws_secret_access_key=bedrock_creds["secret_access_key"], + region_name=bedrock_creds["region"], + ) + + models_to_check: Dict[str, str] = {} + + # Step 1: Get foundation models with TEXT modality + foundation_response = bedrock_client.list_foundation_models() + model_summaries = foundation_response.get("modelSummaries", []) + + for model in model_summaries: + # Check if model supports TEXT input and output modality + input_modalities = model.get("inputModalities", []) + output_modalities = model.get("outputModalities", []) + + if "TEXT" not in input_modalities or "TEXT" not in output_modalities: + continue + + model_id = model.get("modelId") + if not model_id: + continue + + inference_types = model.get("inferenceTypesSupported", []) + + # Only include models with ON_DEMAND inference support + if "ON_DEMAND" in inference_types: + models_to_check[model_id] = model["modelName"] + + # Step 2: Get inference profiles + try: + inference_profiles_response = bedrock_client.list_inference_profiles() + inference_profiles = inference_profiles_response.get( + "inferenceProfileSummaries", [] + ) + + for profile in inference_profiles: + # Check if profile supports TEXT modality + input_modalities = profile.get("inputModalities", []) + output_modalities = profile.get("outputModalities", []) + + if "TEXT" not in input_modalities or "TEXT" not in output_modalities: + continue + + profile_id = profile.get("inferenceProfileId") + if profile_id: + models_to_check[profile_id] = profile["inferenceProfileName"] + + except (BotoCoreError, ClientError) as e: + logger.info( + "Could not fetch inference profiles in %s: %s", + bedrock_creds["region"], + str(e), + ) + + # Step 3: Verify entitlement availability for each model + available_models: Dict[str, str] = {} + + for model_id, model_name in models_to_check.items(): + try: + availability = bedrock_client.get_foundation_model_availability( + modelId=model_id + ) + + entitlement = availability.get("entitlementAvailability") + + # Only include models user has access to + if entitlement == "AVAILABLE": + available_models[model_id] = model_name + else: + logger.debug( + "Skipping model %s - entitlement status: %s", model_id, entitlement + ) + + except (BotoCoreError, ClientError) as e: + logger.debug( + "Could not check availability for model %s: %s", model_id, str(e) + ) + continue + + return available_models + + def refresh_lighthouse_provider_models(provider_config_id: str) -> Dict: """ Refresh the catalog of models for a Lighthouse provider configuration. - For the OpenAI provider, this fetches the current list of models, upserts entries - into `LighthouseProviderModels`, and deletes stale entries no longer returned by - the provider. + Fetches the current list of models from the provider, upserts entries into + `LighthouseProviderModels`, and deletes stale entries no longer returned. Args: - provider_config_id (str): The primary key of the `LighthouseProviderConfiguration` + provider_config_id: The primary key of the `LighthouseProviderConfiguration` whose models should be refreshed. Returns: @@ -104,45 +293,65 @@ def refresh_lighthouse_provider_models(provider_config_id: str) -> Dict: LighthouseProviderConfiguration.DoesNotExist: If no configuration exists with the given ID. """ provider_cfg = LighthouseProviderConfiguration.objects.get(pk=provider_config_id) + fetched_models: Dict[str, str] = {} - if ( - provider_cfg.provider_type - != LighthouseProviderConfiguration.LLMProviderChoices.OPENAI - ): - return { - "created": 0, - "updated": 0, - "deleted": 0, - "error": "Unsupported provider type", - } - - api_key = _extract_openai_api_key(provider_cfg) - if not api_key: - return { - "created": 0, - "updated": 0, - "deleted": 0, - "error": "API key is invalid or missing", - } - + # Fetch models from the appropriate provider try: - client = openai.OpenAI(api_key=api_key) - models = client.models.list() - fetched_ids: Set[str] = {m.id for m in getattr(models, "data", [])} - except Exception as e: # noqa: BLE001 - logger.warning("OpenAI models refresh failed: %s", str(e)) + if ( + provider_cfg.provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.OPENAI + ): + api_key = _extract_openai_api_key(provider_cfg) + if not api_key: + return { + "created": 0, + "updated": 0, + "deleted": 0, + "error": "API key is invalid or missing", + } + fetched_models = _fetch_openai_models(api_key) + + elif ( + provider_cfg.provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.BEDROCK + ): + bedrock_creds = _extract_bedrock_credentials(provider_cfg) + if not bedrock_creds: + return { + "created": 0, + "updated": 0, + "deleted": 0, + "error": "AWS credentials are invalid or missing", + } + fetched_models = _fetch_bedrock_models(bedrock_creds) + + else: + return { + "created": 0, + "updated": 0, + "deleted": 0, + "error": "Unsupported provider type", + } + + except Exception as e: + logger.warning( + "Unexpected error refreshing %s models: %s", + provider_cfg.provider_type, + str(e), + ) return {"created": 0, "updated": 0, "deleted": 0, "error": str(e)} + # Upsert models into the catalog created = 0 updated = 0 - for model_id in fetched_ids: + for model_id, model_name in fetched_models.items(): obj, was_created = LighthouseProviderModels.objects.update_or_create( tenant_id=provider_cfg.tenant_id, provider_configuration=provider_cfg, model_id=model_id, defaults={ - "model_name": model_id, # OpenAI doesn't return a separate display name + "model_name": model_name, "default_parameters": {}, }, ) @@ -156,7 +365,7 @@ def refresh_lighthouse_provider_models(provider_config_id: str) -> Dict: LighthouseProviderModels.objects.filter( tenant_id=provider_cfg.tenant_id, provider_configuration=provider_cfg ) - .exclude(model_id__in=fetched_ids) + .exclude(model_id__in=fetched_models.keys()) .delete() )