diff --git a/api/CHANGELOG.md b/api/CHANGELOG.md index 9450ef87ba..8b36cb94a1 100644 --- a/api/CHANGELOG.md +++ b/api/CHANGELOG.md @@ -12,6 +12,7 @@ All notable changes to the **Prowler API** are documented in this file. - Support for Oracle Cloud Infrastructure (OCI) provider [(#8927)](https://github.com/prowler-cloud/prowler/pull/8927) - Support muting findings based on simple rules with custom reason [(#9051)](https://github.com/prowler-cloud/prowler/pull/9051) - Support C5 compliance framework for the GCP provider [(#9097)](https://github.com/prowler-cloud/prowler/pull/9097) +- Support for Amazon Bedrock and OpenAI compatible providers in Lighthouse AI [(#8957)](https://github.com/prowler-cloud/prowler/pull/8957) --- diff --git a/api/src/backend/api/migrations/0053_lighthouse_bedrock_openai_compatible.py b/api/src/backend/api/migrations/0053_lighthouse_bedrock_openai_compatible.py new file mode 100644 index 0000000000..d7054d6443 --- /dev/null +++ b/api/src/backend/api/migrations/0053_lighthouse_bedrock_openai_compatible.py @@ -0,0 +1,25 @@ +# Generated by Django 5.1.12 on 2025-10-14 11:46 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + dependencies = [ + ("api", "0052_mute_rules"), + ] + + operations = [ + migrations.AlterField( + model_name="lighthouseproviderconfiguration", + name="provider_type", + field=models.CharField( + choices=[ + ("openai", "OpenAI"), + ("bedrock", "AWS Bedrock"), + ("openai_compatible", "OpenAI Compatible"), + ], + 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 8e3cbd2282..3a0f1e9df2 100644 --- a/api/src/backend/api/models.py +++ b/api/src/backend/api/models.py @@ -2039,6 +2039,8 @@ class LighthouseProviderConfiguration(RowLevelSecurityProtectedModel): class LLMProviderChoices(models.TextChoices): OPENAI = "openai", _("OpenAI") + BEDROCK = "bedrock", _("AWS Bedrock") + OPENAI_COMPATIBLE = "openai_compatible", _("OpenAI Compatible") 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/specs/v1.yaml b/api/src/backend/api/specs/v1.yaml index efc99e9b2f..958e176472 100644 --- a/api/src/backend/api/specs/v1.yaml +++ b/api/src/backend/api/specs/v1.yaml @@ -3585,7 +3585,7 @@ paths: summary: Get Lighthouse AI Tenant config parameters: - in: query - name: fields[lighthouse-config] + name: fields[lighthouse-configurations] schema: type: array items: @@ -3690,7 +3690,7 @@ paths: summary: List all LLM models parameters: - in: query - name: fields[LighthouseProviderModels] + name: fields[lighthouse-models] schema: type: array items: @@ -3727,26 +3727,34 @@ paths: name: filter[provider_type] schema: type: string - x-spec-enum-id: 80350f3ab09becd9 + x-spec-enum-id: 30bc523cb9ec0e8b enum: + - bedrock - openai + - openai_compatible description: |- LLM provider name * `openai` - OpenAI + * `bedrock` - AWS Bedrock + * `openai_compatible` - OpenAI Compatible - in: query name: filter[provider_type__in] schema: type: array items: type: string - x-spec-enum-id: 80350f3ab09becd9 + x-spec-enum-id: 30bc523cb9ec0e8b enum: + - bedrock - openai + - openai_compatible description: |- Multiple values may be separated by commas. * `openai` - OpenAI + * `bedrock` - AWS Bedrock + * `openai_compatible` - OpenAI Compatible explode: false style: form - name: filter[search] @@ -3811,7 +3819,7 @@ paths: summary: Retrieve LLM model details parameters: - in: query - name: fields[LighthouseProviderModels] + name: fields[lighthouse-models] schema: type: array items: @@ -3849,7 +3857,7 @@ paths: get: operationId: lighthouse_providers_list description: Retrieve all LLM provider configurations for the current tenant - summary: List all LLM provider configs + summary: List all LLM provider configurations parameters: - in: query name: fields[lighthouse-providers] @@ -3876,26 +3884,34 @@ paths: name: filter[provider_type] schema: type: string - x-spec-enum-id: 80350f3ab09becd9 + x-spec-enum-id: 30bc523cb9ec0e8b enum: + - bedrock - openai + - openai_compatible description: |- LLM provider name * `openai` - OpenAI + * `bedrock` - AWS Bedrock + * `openai_compatible` - OpenAI Compatible - in: query name: filter[provider_type__in] schema: type: array items: type: string - x-spec-enum-id: 80350f3ab09becd9 + x-spec-enum-id: 30bc523cb9ec0e8b enum: + - bedrock - openai + - openai_compatible description: |- Multiple values may be separated by commas. * `openai` - OpenAI + * `bedrock` - AWS Bedrock + * `openai_compatible` - OpenAI Compatible explode: false style: form - name: filter[search] @@ -3957,7 +3973,7 @@ paths: operationId: lighthouse_providers_create description: Create a per-tenant configuration for an LLM provider. Only one configuration per provider type is allowed per tenant. - summary: Create LLM provider config + summary: Create LLM provider configuration tags: - Lighthouse AI requestBody: @@ -3986,7 +4002,7 @@ paths: operationId: lighthouse_providers_retrieve description: Get details for a specific provider configuration in the current tenant. - summary: Retrieve LLM provider config + summary: Retrieve LLM provider configuration parameters: - in: query name: fields[lighthouse-providers] @@ -4026,7 +4042,7 @@ paths: patch: operationId: lighthouse_providers_partial_update description: Partially update a provider configuration (e.g., base_url, is_active). - summary: Update LLM provider config + summary: Update LLM provider configuration parameters: - in: path name: id @@ -4062,7 +4078,7 @@ paths: operationId: lighthouse_providers_destroy description: Delete a provider configuration. Any tenant defaults that reference this provider are cleared during deletion. - summary: Delete LLM provider config + summary: Delete LLM provider configuration parameters: - in: path name: id @@ -4121,7 +4137,7 @@ paths: post: operationId: lighthouse_providers_refresh_models_create description: Fetch available models for this provider configuration and upsert - into catalog. + into catalog. Supports OpenAI, OpenAI-compatible, and AWS Bedrock providers. summary: Refresh LLM models catalog parameters: - in: path @@ -11954,12 +11970,16 @@ components: provider_type: enum: - openai + - bedrock + - openai_compatible type: string - x-spec-enum-id: 80350f3ab09becd9 + x-spec-enum-id: 30bc523cb9ec0e8b description: |- LLM provider name * `openai` - OpenAI + * `bedrock` - AWS Bedrock + * `openai_compatible` - OpenAI Compatible base_url: type: string format: uri @@ -11991,18 +12011,64 @@ components: provider_type: enum: - openai + - bedrock + - openai_compatible type: string - x-spec-enum-id: 80350f3ab09becd9 + x-spec-enum-id: 30bc523cb9ec0e8b description: |- - LLM provider name + LLM provider type. Determines which credential format to use. See 'credentials' field documentation for provider-specific requirements. * `openai` - OpenAI + * `bedrock` - AWS Bedrock + * `openai_compatible` - OpenAI Compatible base_url: type: string format: uri nullable: true - maxLength: 200 + description: Base URL for the LLM provider API. Required for 'openai_compatible' + provider type. credentials: + oneOf: + - type: object + title: OpenAI Credentials + properties: + api_key: + type: string + description: OpenAI API key. Must start with 'sk-' followed by + alphanumeric characters, hyphens, or underscores. + pattern: ^sk-[\w-]+$ + required: + - api_key + - type: object + title: AWS Bedrock Credentials + properties: + access_key_id: + type: string + description: AWS access key ID. + pattern: ^AKIA[0-9A-Z]{16}$ + secret_access_key: + type: string + description: AWS secret access key. + pattern: ^[A-Za-z0-9/+=]{40}$ + region: + type: string + description: 'AWS region identifier where Bedrock is available. + Examples: us-east-1, us-west-2, eu-west-1, ap-northeast-1.' + pattern: ^[a-z]{2}-[a-z]+-\d+$ + required: + - access_key_id + - secret_access_key + - region + - type: object + title: OpenAI Compatible Credentials + properties: + api_key: + type: string + description: 'API key for OpenAI-compatible provider. The format + varies by provider. Note: The ''base_url'' field (separate from + credentials) is required when using this provider type.' + required: + - api_key writeOnly: true is_active: type: boolean @@ -12031,18 +12097,66 @@ components: provider_type: enum: - openai + - bedrock + - openai_compatible type: string - x-spec-enum-id: 80350f3ab09becd9 + x-spec-enum-id: 30bc523cb9ec0e8b description: |- - LLM provider name + LLM provider type. Determines which credential format to use. See 'credentials' field documentation for provider-specific requirements. * `openai` - OpenAI + * `bedrock` - AWS Bedrock + * `openai_compatible` - OpenAI Compatible base_url: type: string format: uri nullable: true - maxLength: 200 + minLength: 1 + description: Base URL for the LLM provider API. Required for 'openai_compatible' + provider type. credentials: + oneOf: + - type: object + title: OpenAI Credentials + properties: + api_key: + type: string + description: OpenAI API key. Must start with 'sk-' followed + by alphanumeric characters, hyphens, or underscores. + pattern: ^sk-[\w-]+$ + required: + - api_key + - type: object + title: AWS Bedrock Credentials + properties: + access_key_id: + type: string + description: AWS access key ID. + pattern: ^AKIA[0-9A-Z]{16}$ + secret_access_key: + type: string + description: AWS secret access key. + pattern: ^[A-Za-z0-9/+=]{40}$ + region: + type: string + description: 'AWS region identifier where Bedrock is available. + Examples: us-east-1, us-west-2, eu-west-1, ap-northeast-1.' + pattern: ^[a-z]{2}-[a-z]+-\d+$ + required: + - access_key_id + - secret_access_key + - region + - type: object + title: OpenAI Compatible Credentials + properties: + api_key: + type: string + description: 'API key for OpenAI-compatible provider. The + format varies by provider. Note: The ''base_url'' field + (separate from credentials) is required when using this + provider type.' + required: + - api_key writeOnly: true is_active: type: boolean @@ -12088,19 +12202,65 @@ components: provider_type: enum: - openai + - bedrock + - openai_compatible type: string - x-spec-enum-id: 80350f3ab09becd9 + x-spec-enum-id: 30bc523cb9ec0e8b readOnly: true description: |- LLM provider name * `openai` - OpenAI + * `bedrock` - AWS Bedrock + * `openai_compatible` - OpenAI Compatible base_url: type: string format: uri nullable: true - maxLength: 200 + description: Base URL for the LLM provider API. Required for 'openai_compatible' + provider type. credentials: + oneOf: + - type: object + title: OpenAI Credentials + properties: + api_key: + type: string + description: OpenAI API key. Must start with 'sk-' followed by + alphanumeric characters, hyphens, or underscores. + pattern: ^sk-[\w-]+$ + required: + - api_key + - type: object + title: AWS Bedrock Credentials + properties: + access_key_id: + type: string + description: AWS access key ID. + pattern: ^AKIA[0-9A-Z]{16}$ + secret_access_key: + type: string + description: AWS secret access key. + pattern: ^[A-Za-z0-9/+=]{40}$ + region: + type: string + description: 'AWS region identifier where Bedrock is available. + Examples: us-east-1, us-west-2, eu-west-1, ap-northeast-1.' + pattern: ^[a-z]{2}-[a-z]+-\d+$ + required: + - access_key_id + - secret_access_key + - region + - type: object + title: OpenAI Compatible Credentials + properties: + api_key: + type: string + description: 'API key for OpenAI-compatible provider. The format + varies by provider. Note: The ''base_url'' field (separate from + credentials) is required when using this provider type.' + required: + - api_key writeOnly: true is_active: type: boolean @@ -12124,7 +12284,7 @@ components: member is used to describe resource objects that share common attributes and relationships. enum: - - LighthouseProviderModels + - lighthouse-models id: type: string format: uuid @@ -12197,7 +12357,7 @@ components: member is used to describe resource objects that share common attributes and relationships. enum: - - lighthouse-config + - lighthouse-configurations id: type: string format: uuid @@ -12234,7 +12394,7 @@ components: member is used to describe resource objects that share common attributes and relationships. enum: - - lighthouse-config + - lighthouse-configurations id: type: string format: uuid @@ -13436,19 +13596,67 @@ components: provider_type: enum: - openai + - bedrock + - openai_compatible type: string - x-spec-enum-id: 80350f3ab09becd9 + x-spec-enum-id: 30bc523cb9ec0e8b readOnly: true description: |- LLM provider name * `openai` - OpenAI + * `bedrock` - AWS Bedrock + * `openai_compatible` - OpenAI Compatible base_url: type: string format: uri nullable: true - maxLength: 200 + minLength: 1 + description: Base URL for the LLM provider API. Required for 'openai_compatible' + provider type. credentials: + oneOf: + - type: object + title: OpenAI Credentials + properties: + api_key: + type: string + description: OpenAI API key. Must start with 'sk-' followed + by alphanumeric characters, hyphens, or underscores. + pattern: ^sk-[\w-]+$ + required: + - api_key + - type: object + title: AWS Bedrock Credentials + properties: + access_key_id: + type: string + description: AWS access key ID. + pattern: ^AKIA[0-9A-Z]{16}$ + secret_access_key: + type: string + description: AWS secret access key. + pattern: ^[A-Za-z0-9/+=]{40}$ + region: + type: string + description: 'AWS region identifier where Bedrock is available. + Examples: us-east-1, us-west-2, eu-west-1, ap-northeast-1.' + pattern: ^[a-z]{2}-[a-z]+-\d+$ + required: + - access_key_id + - secret_access_key + - region + - type: object + title: OpenAI Compatible Credentials + properties: + api_key: + type: string + description: 'API key for OpenAI-compatible provider. The + format varies by provider. Note: The ''base_url'' field + (separate from credentials) is required when using this + provider type.' + required: + - api_key writeOnly: true is_active: type: boolean @@ -13470,7 +13678,7 @@ components: member is used to describe resource objects that share common attributes and relationships. enum: - - lighthouse-config + - lighthouse-configurations id: type: string format: uuid diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index f6e169ea2b..97e6002e00 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -9950,3 +9950,380 @@ class TestMuteRuleViewSet: assert len(data) == len(mute_rules_fixture) for rule_data in data: assert rule_data["id"] != str(other_rule.id) + + @pytest.mark.parametrize( + "credentials", + [ + {}, # empty credentials + { + "access_key_id": "AKIAIOSFODNN7EXAMPLE" + }, # missing secret_access_key and region + { + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY" + }, # missing access_key_id and region + { + "access_key_id": "AKIAIOSFODNN7EXAMPLE", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + }, # missing region + { # invalid access_key_id format (not starting with AKIA) + "access_key_id": "ABCD0123456789ABCDEF", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + "region": "us-east-1", + }, + { # invalid access_key_id format (wrong length) + "access_key_id": "AKIAIOSFODNN7EXAMPL", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + "region": "us-east-1", + }, + { # invalid secret_access_key format (wrong length) + "access_key_id": "AKIAIOSFODNN7EXAMPLE", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEK", + "region": "us-east-1", + }, + { # invalid region format + "access_key_id": "AKIAIOSFODNN7EXAMPLE", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + "region": "invalid-region", + }, + { # invalid region format (uppercase) + "access_key_id": "AKIAIOSFODNN7EXAMPLE", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + "region": "US-EAST-1", + }, + ], + ) + def test_bedrock_invalid_credentials(self, authenticated_client, credentials): + """Bedrock provider with invalid credentials should error""" + payload = { + "data": { + "type": "lighthouse-providers", + "attributes": { + "provider_type": "bedrock", + "credentials": credentials, + }, + } + } + resp = authenticated_client.post( + reverse("lighthouse-providers-list"), + data=payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert resp.status_code == status.HTTP_400_BAD_REQUEST + + def test_bedrock_valid_credentials_success(self, authenticated_client): + """Bedrock provider with valid AWS credentials should succeed and mask credentials""" + valid_credentials = { + "access_key_id": "AKIAIOSFODNN7EXAMPLE", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + "region": "us-east-1", + } + payload = { + "data": { + "type": "lighthouse-providers", + "attributes": { + "provider_type": "bedrock", + "credentials": valid_credentials, + }, + } + } + resp = authenticated_client.post( + reverse("lighthouse-providers-list"), + data=payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert resp.status_code == status.HTTP_201_CREATED + data = resp.json()["data"] + + # Verify credentials are returned masked + masked_creds = data["attributes"].get("credentials") + assert masked_creds is not None + assert "access_key_id" in masked_creds + assert "secret_access_key" in masked_creds + assert "region" in masked_creds + # Verify all characters are masked with asterisks + assert all(c == "*" for c in masked_creds["access_key_id"]) + assert all(c == "*" for c in masked_creds["secret_access_key"]) + + def test_bedrock_provider_duplicate_per_tenant(self, authenticated_client): + """Creating a second Bedrock provider for same tenant should fail""" + valid_credentials = { + "access_key_id": "AKIAIOSFODNN7EXAMPLE", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + "region": "us-west-2", + } + payload = { + "data": { + "type": "lighthouse-providers", + "attributes": { + "provider_type": "bedrock", + "credentials": valid_credentials, + }, + } + } + # First creation succeeds + resp1 = authenticated_client.post( + reverse("lighthouse-providers-list"), + data=payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert resp1.status_code == status.HTTP_201_CREATED + + # Second creation should fail with validation error + resp2 = authenticated_client.post( + reverse("lighthouse-providers-list"), + data=payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert resp2.status_code == status.HTTP_400_BAD_REQUEST + assert "already exists" in str(resp2.json()).lower() + + def test_bedrock_patch_credentials_and_fields_filter(self, authenticated_client): + """PATCH credentials and verify fields filter returns decrypted values""" + valid_credentials = { + "access_key_id": "AKIAIOSFODNN7EXAMPLE", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + "region": "eu-west-1", + } + create_payload = { + "data": { + "type": "lighthouse-providers", + "attributes": { + "provider_type": "bedrock", + "credentials": valid_credentials, + }, + } + } + create_resp = authenticated_client.post( + reverse("lighthouse-providers-list"), + data=create_payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert create_resp.status_code == status.HTTP_201_CREATED + provider_id = create_resp.json()["data"]["id"] + + # Update credentials with new valid ones + new_credentials = { + "access_key_id": "AKIAZZZZZZZZZZZZZZZZ", + "secret_access_key": "aBcDeFgHiJkLmNoPqRsTuVwXyZ0123456789+/==", + "region": "ap-south-1", + } + patch_payload = { + "data": { + "type": "lighthouse-providers", + "id": provider_id, + "attributes": { + "credentials": new_credentials, + "is_active": False, + }, + } + } + patch_resp = authenticated_client.patch( + reverse("lighthouse-providers-detail", kwargs={"pk": provider_id}), + data=patch_payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert patch_resp.status_code == status.HTTP_200_OK + updated = patch_resp.json()["data"]["attributes"] + assert updated["is_active"] is False + + # Default GET should return masked credentials + get_resp = authenticated_client.get( + reverse("lighthouse-providers-detail", kwargs={"pk": provider_id}) + ) + assert get_resp.status_code == status.HTTP_200_OK + masked = get_resp.json()["data"]["attributes"]["credentials"] + assert all(c == "*" for c in masked["access_key_id"]) + assert all(c == "*" for c in masked["secret_access_key"]) + + # Fields filter should return decrypted credentials + get_full = authenticated_client.get( + reverse("lighthouse-providers-detail", kwargs={"pk": provider_id}) + + "?fields[lighthouse-providers]=credentials" + ) + assert get_full.status_code == status.HTTP_200_OK + creds = get_full.json()["data"]["attributes"]["credentials"] + assert creds["access_key_id"] == new_credentials["access_key_id"] + assert creds["secret_access_key"] == new_credentials["secret_access_key"] + assert creds["region"] == new_credentials["region"] + + def test_bedrock_partial_credential_update(self, authenticated_client): + """Test partial update of Bedrock credentials (e.g., only region)""" + # Create provider with full credentials + initial_credentials = { + "access_key_id": "AKIAIOSFODNN7EXAMPLE", + "secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + "region": "us-east-1", + } + create_payload = { + "data": { + "type": "lighthouse-providers", + "attributes": { + "provider_type": "bedrock", + "credentials": initial_credentials, + }, + } + } + create_resp = authenticated_client.post( + reverse("lighthouse-providers-list"), + data=create_payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert create_resp.status_code == status.HTTP_201_CREATED + provider_id = create_resp.json()["data"]["id"] + + # Update only the region field + partial_update = { + "region": "eu-west-1", + } + patch_payload = { + "data": { + "type": "lighthouse-providers", + "id": provider_id, + "attributes": { + "credentials": partial_update, + }, + } + } + patch_resp = authenticated_client.patch( + reverse("lighthouse-providers-detail", kwargs={"pk": provider_id}), + data=patch_payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert patch_resp.status_code == status.HTTP_200_OK + + # Verify credentials with fields filter - region should be updated, keys preserved + get_full = authenticated_client.get( + reverse("lighthouse-providers-detail", kwargs={"pk": provider_id}) + + "?fields[lighthouse-providers]=credentials" + ) + assert get_full.status_code == status.HTTP_200_OK + creds = get_full.json()["data"]["attributes"]["credentials"] + + # Original keys should be preserved + assert creds["access_key_id"] == initial_credentials["access_key_id"] + assert creds["secret_access_key"] == initial_credentials["secret_access_key"] + # Region should be updated + assert creds["region"] == "eu-west-1" + + @pytest.mark.parametrize( + "attributes", + [ + pytest.param( + { + "provider_type": "openai_compatible", + "credentials": {"api_key": "compat-key"}, + }, + id="missing", + ), + pytest.param( + { + "provider_type": "openai_compatible", + "credentials": {"api_key": "compat-key"}, + "base_url": "", + }, + id="empty", + ), + ], + ) + def test_openai_compatible_missing_base_url(self, authenticated_client, attributes): + payload = { + "data": { + "type": "lighthouse-providers", + "attributes": attributes, + } + } + + resp = authenticated_client.post( + reverse("lighthouse-providers-list"), + data=payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert resp.status_code == status.HTTP_400_BAD_REQUEST + error_detail = str(resp.json()).lower() + assert "base_url" in error_detail + + def test_openai_compatible_invalid_credentials(self, authenticated_client): + payload = { + "data": { + "type": "lighthouse-providers", + "attributes": { + "provider_type": "openai_compatible", + "base_url": "https://compat.example/v1", + "credentials": {"api_key": ""}, + }, + } + } + + resp = authenticated_client.post( + reverse("lighthouse-providers-list"), + data=payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert resp.status_code == status.HTTP_400_BAD_REQUEST + errors = resp.json().get("errors", []) + assert any( + error.get("source", {}).get("pointer") + == "/data/attributes/credentials/api_key" + for error in errors + ) + assert any( + "may not be blank" in error.get("detail", "").lower() for error in errors + ) + + def test_openai_compatible_patch_credentials_and_fields(self, authenticated_client): + create_payload = { + "data": { + "type": "lighthouse-providers", + "attributes": { + "provider_type": "openai_compatible", + "base_url": "https://compat.example/v1", + "credentials": {"api_key": "compat-key-123"}, + }, + } + } + + create_resp = authenticated_client.post( + reverse("lighthouse-providers-list"), + data=create_payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert create_resp.status_code == status.HTTP_201_CREATED + provider_id = create_resp.json()["data"]["id"] + + updated_base_url = "https://compat.example/v2" + updated_api_key = "compat-key-456" + patch_payload = { + "data": { + "type": "lighthouse-providers", + "id": provider_id, + "attributes": { + "base_url": updated_base_url, + "credentials": {"api_key": updated_api_key}, + }, + } + } + + patch_resp = authenticated_client.patch( + reverse("lighthouse-providers-detail", kwargs={"pk": provider_id}), + data=patch_payload, + content_type=API_JSON_CONTENT_TYPE, + ) + assert patch_resp.status_code == status.HTTP_200_OK + updated_attrs = patch_resp.json()["data"]["attributes"] + assert updated_attrs["base_url"] == updated_base_url + assert updated_attrs["credentials"]["api_key"] == "*" * len(updated_api_key) + + get_resp = authenticated_client.get( + reverse("lighthouse-providers-detail", kwargs={"pk": provider_id}) + ) + assert get_resp.status_code == status.HTTP_200_OK + masked = get_resp.json()["data"]["attributes"]["credentials"]["api_key"] + assert masked == "*" * len(updated_api_key) + + get_full = authenticated_client.get( + reverse("lighthouse-providers-detail", kwargs={"pk": provider_id}) + + "?fields[lighthouse-providers]=credentials" + ) + assert get_full.status_code == status.HTTP_200_OK + creds = get_full.json()["data"]["attributes"]["credentials"] + assert creds["api_key"] == updated_api_key diff --git a/api/src/backend/api/v1/serializer_utils/lighthouse.py b/api/src/backend/api/v1/serializer_utils/lighthouse.py index f7ecef32cc..e29c679ce6 100644 --- a/api/src/backend/api/v1/serializer_utils/lighthouse.py +++ b/api/src/backend/api/v1/serializer_utils/lighthouse.py @@ -1,5 +1,6 @@ import re +from drf_spectacular.utils import extend_schema_field from rest_framework_json_api import serializers @@ -11,3 +12,198 @@ class OpenAICredentialsSerializer(serializers.Serializer): if not re.match(pattern, value or ""): raise serializers.ValidationError("Invalid OpenAI API key format.") return value + + def to_internal_value(self, data): + """Check for unknown fields before DRF filters them out.""" + if not isinstance(data, dict): + raise serializers.ValidationError( + {"non_field_errors": ["Credentials must be an object"]} + ) + + allowed_fields = set(self.fields.keys()) + provided_fields = set(data.keys()) + extra_fields = provided_fields - allowed_fields + + if extra_fields: + raise serializers.ValidationError( + { + "non_field_errors": [ + f"Unknown fields in credentials: {', '.join(sorted(extra_fields))}" + ] + } + ) + + return super().to_internal_value(data) + + +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 + + def to_internal_value(self, data): + """Check for unknown fields before DRF filters them out.""" + if not isinstance(data, dict): + raise serializers.ValidationError( + {"non_field_errors": ["Credentials must be an object"]} + ) + + allowed_fields = set(self.fields.keys()) + provided_fields = set(data.keys()) + extra_fields = provided_fields - allowed_fields + + if extra_fields: + raise serializers.ValidationError( + { + "non_field_errors": [ + f"Unknown fields in credentials: {', '.join(sorted(extra_fields))}" + ] + } + ) + + return super().to_internal_value(data) + + +class BedrockCredentialsUpdateSerializer(BedrockCredentialsSerializer): + """ + Serializer for AWS Bedrock credentials during UPDATE operations. + + Inherits all validation logic from BedrockCredentialsSerializer but makes + all fields optional to support partial updates. + """ + + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + # Make all fields optional for updates + for field in self.fields.values(): + field.required = False + + +class OpenAICompatibleCredentialsSerializer(serializers.Serializer): + """ + Minimal serializer for OpenAI-compatible credentials. + + Many OpenAI-compatible providers do not use the same key format as OpenAI. + We only require a non-empty API key string. Additional fields can be added later + without breaking existing configurations. + """ + + api_key = serializers.CharField() + + def validate_api_key(self, value: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise serializers.ValidationError("API key is required.") + return value.strip() + + def to_internal_value(self, data): + """Check for unknown fields before DRF filters them out.""" + if not isinstance(data, dict): + raise serializers.ValidationError( + {"non_field_errors": ["Credentials must be an object"]} + ) + + allowed_fields = set(self.fields.keys()) + provided_fields = set(data.keys()) + extra_fields = provided_fields - allowed_fields + + if extra_fields: + raise serializers.ValidationError( + { + "non_field_errors": [ + f"Unknown fields in credentials: {', '.join(sorted(extra_fields))}" + ] + } + ) + + return super().to_internal_value(data) + + +@extend_schema_field( + { + "oneOf": [ + { + "type": "object", + "title": "OpenAI Credentials", + "properties": { + "api_key": { + "type": "string", + "description": "OpenAI API key. Must start with 'sk-' followed by alphanumeric characters, " + "hyphens, or underscores.", + "pattern": "^sk-[\\w-]+$", + } + }, + "required": ["api_key"], + }, + { + "type": "object", + "title": "AWS Bedrock Credentials", + "properties": { + "access_key_id": { + "type": "string", + "description": "AWS access key ID.", + "pattern": "^AKIA[0-9A-Z]{16}$", + }, + "secret_access_key": { + "type": "string", + "description": "AWS secret access key.", + "pattern": "^[A-Za-z0-9/+=]{40}$", + }, + "region": { + "type": "string", + "description": "AWS region identifier where Bedrock is available. Examples: us-east-1, " + "us-west-2, eu-west-1, ap-northeast-1.", + "pattern": "^[a-z]{2}-[a-z]+-\\d+$", + }, + }, + "required": ["access_key_id", "secret_access_key", "region"], + }, + { + "type": "object", + "title": "OpenAI Compatible Credentials", + "properties": { + "api_key": { + "type": "string", + "description": "API key for OpenAI-compatible provider. The format varies by provider. " + "Note: The 'base_url' field (separate from credentials) is required when using this provider type.", + } + }, + "required": ["api_key"], + }, + ] + } +) +class LighthouseCredentialsField(serializers.JSONField): + pass diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 10a004e3cf..1af474e6de 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -60,7 +60,13 @@ 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, + BedrockCredentialsUpdateSerializer, + LighthouseCredentialsField, + OpenAICompatibleCredentialsSerializer, + OpenAICredentialsSerializer, +) from api.v1.serializer_utils.processors import ProcessorConfigField from api.v1.serializer_utils.providers import ProviderSecretField from prowler.lib.mutelist.mutelist import Mutelist @@ -3076,7 +3082,12 @@ class LighthouseProviderConfigCreateSerializer(RLSSerializer, BaseWriteSerialize Accepts credentials as JSON; stored encrypted via credentials_decoded. """ - credentials = serializers.JSONField(write_only=True, required=True) + credentials = LighthouseCredentialsField(write_only=True, required=True) + base_url = serializers.URLField( + required=False, + allow_null=True, + help_text="Base URL for the LLM provider API. Required for 'openai_compatible' provider type.", + ) class Meta: model = LighthouseProviderConfiguration @@ -3088,7 +3099,10 @@ class LighthouseProviderConfigCreateSerializer(RLSSerializer, BaseWriteSerialize ] extra_kwargs = { "is_active": {"required": False}, - "base_url": {"required": False, "allow_null": True}, + "provider_type": { + "help_text": "LLM provider type. Determines which credential format to use. " + "See 'credentials' field documentation for provider-specific requirements." + }, } def create(self, validated_data): @@ -3111,6 +3125,7 @@ class LighthouseProviderConfigCreateSerializer(RLSSerializer, BaseWriteSerialize def validate(self, attrs): provider_type = attrs.get("provider_type") credentials = attrs.get("credentials") or {} + base_url = attrs.get("base_url") if provider_type == LighthouseProviderConfiguration.LLMProviderChoices.OPENAI: try: @@ -3123,6 +3138,35 @@ 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 + elif ( + provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.OPENAI_COMPATIBLE + ): + if not base_url: + raise ValidationError({"base_url": "Base URL is required."}) + try: + OpenAICompatibleCredentialsSerializer(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) @@ -3132,7 +3176,12 @@ class LighthouseProviderConfigUpdateSerializer(BaseWriteSerializer): Update serializer for LighthouseProviderConfiguration. """ - credentials = serializers.JSONField(write_only=True, required=False) + credentials = LighthouseCredentialsField(write_only=True, required=False) + base_url = serializers.URLField( + required=False, + allow_null=True, + help_text="Base URL for the LLM provider API. Required for 'openai_compatible' provider type.", + ) class Meta: model = LighthouseProviderConfiguration @@ -3146,7 +3195,6 @@ class LighthouseProviderConfigUpdateSerializer(BaseWriteSerializer): extra_kwargs = { "id": {"read_only": True}, "provider_type": {"read_only": True}, - "base_url": {"required": False, "allow_null": True}, "is_active": {"required": False}, } @@ -3157,7 +3205,11 @@ class LighthouseProviderConfigUpdateSerializer(BaseWriteSerializer): setattr(instance, attr, value) if credentials is not None: - instance.credentials_decoded = credentials + # Merge partial credentials with existing ones + # New values overwrite existing ones, but unspecified fields are preserved + existing_credentials = instance.credentials_decoded or {} + merged_credentials = {**existing_credentials, **credentials} + instance.credentials_decoded = merged_credentials instance.save() return instance @@ -3165,6 +3217,7 @@ class LighthouseProviderConfigUpdateSerializer(BaseWriteSerializer): def validate(self, attrs): provider_type = getattr(self.instance, "provider_type", None) credentials = attrs.get("credentials", None) + base_url = attrs.get("base_url", None) if ( credentials is not None @@ -3181,6 +3234,40 @@ 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: + BedrockCredentialsUpdateSerializer(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 + elif ( + credentials is not None + and provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.OPENAI_COMPATIBLE + ): + if base_url is None: + pass + elif not base_url: + raise ValidationError({"base_url": "Base URL cannot be empty."}) + try: + OpenAICompatibleCredentialsSerializer(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 b33b42e461..2cdbf02c8b 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -4323,28 +4323,30 @@ class LighthouseConfigViewSet(BaseRLSViewSet): @extend_schema_view( list=extend_schema( tags=["Lighthouse AI"], - summary="List all LLM provider configs", + summary="List all LLM provider configurations", description="Retrieve all LLM provider configurations for the current tenant", ), retrieve=extend_schema( tags=["Lighthouse AI"], - summary="Retrieve LLM provider config", + summary="Retrieve LLM provider configuration", description="Get details for a specific provider configuration in the current tenant.", ), create=extend_schema( tags=["Lighthouse AI"], - summary="Create LLM provider config", - description="Create a per-tenant configuration for an LLM provider. Only one configuration per provider type is allowed per tenant.", + summary="Create LLM provider configuration", + description="Create a per-tenant configuration for an LLM provider. Only one configuration per provider type " + "is allowed per tenant.", ), partial_update=extend_schema( tags=["Lighthouse AI"], - summary="Update LLM provider config", + summary="Update LLM provider configuration", description="Partially update a provider configuration (e.g., base_url, is_active).", ), destroy=extend_schema( tags=["Lighthouse AI"], - summary="Delete LLM provider config", - description="Delete a provider configuration. Any tenant defaults that reference this provider are cleared during deletion.", + summary="Delete LLM provider configuration", + description="Delete a provider configuration. Any tenant defaults that reference this provider are cleared " + "during deletion.", ), ) class LighthouseProviderConfigViewSet(BaseRLSViewSet): @@ -4409,16 +4411,6 @@ class LighthouseProviderConfigViewSet(BaseRLSViewSet): @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( @@ -4440,7 +4432,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, OpenAI-compatible, and AWS Bedrock providers.", request=None, responses={202: OpenApiResponse(response=TaskSerializer)}, ) @@ -4452,16 +4444,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..53f9e4dd97 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,71 @@ def _extract_openai_api_key( return api_key +def _extract_openai_compatible_params( + provider_cfg: LighthouseProviderConfiguration, +) -> Dict[str, str] | None: + """ + Extract base_url and api_key for OpenAI-compatible providers. + """ + creds = provider_cfg.credentials_decoded + base_url = provider_cfg.base_url + if not isinstance(creds, dict): + return None + api_key = creds.get("api_key") + if not isinstance(api_key, str) or not api_key: + return None + if not isinstance(base_url, str) or not base_url: + return None + return {"base_url": base_url, "api_key": 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 +112,238 @@ 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() + + elif ( + provider_cfg.provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.OPENAI_COMPATIBLE + ): + params = _extract_openai_compatible_params(provider_cfg) + if not params: + provider_cfg.is_active = False + provider_cfg.save() + return { + "connected": False, + "error": "Base URL or API key is invalid or missing", + } + + # Test connection using OpenAI SDK with custom base_url + # Note: base_url should include version (e.g., https://openrouter.ai/api/v1) + client = openai.OpenAI( + api_key=params["api_key"], + base_url=params["base_url"], + ) + _ = client.models.list() + + 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_openai_compatible_models(base_url: str, api_key: str) -> Dict[str, str]: + """ + Fetch available models from an OpenAI-compatible API using the OpenAI SDK. + + Returns a mapping of model_id -> model_name. Prefers the 'name' attribute + if available (e.g., from OpenRouter), otherwise falls back to 'id'. + + Note: base_url should include version (e.g., https://openrouter.ai/api/v1) + """ + client = openai.OpenAI(api_key=api_key, base_url=base_url) + models = client.models.list() + + available_models: Dict[str, str] = {} + for model in models.data: + model_id = model.id + # Prefer provider-supplied human-friendly name when available + name = getattr(model, "name", None) + if name: + available_models[model_id] = name + else: + available_models[model_id] = model_id + + return available_models + + +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 +357,81 @@ 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) + + elif ( + provider_cfg.provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.OPENAI_COMPATIBLE + ): + params = _extract_openai_compatible_params(provider_cfg) + if not params: + return { + "created": 0, + "updated": 0, + "deleted": 0, + "error": "Base URL or API key is invalid or missing", + } + fetched_models = _fetch_openai_compatible_models( + params["base_url"], params["api_key"] + ) + + 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 +445,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() ) diff --git a/api/src/backend/tasks/tests/test_tasks.py b/api/src/backend/tasks/tests/test_tasks.py index a98341ef35..b2d8b7bc2b 100644 --- a/api/src/backend/tasks/tests/test_tasks.py +++ b/api/src/backend/tasks/tests/test_tasks.py @@ -1,16 +1,24 @@ import uuid from unittest.mock import MagicMock, patch +import openai import pytest +from botocore.exceptions import ClientError from tasks.tasks import ( _perform_scan_complete_tasks, check_integrations_task, + check_lighthouse_provider_connection_task, generate_outputs_task, + refresh_lighthouse_provider_models_task, s3_integration_task, security_hub_integration_task, ) -from api.models import Integration +from api.models import ( + Integration, + LighthouseProviderConfiguration, + LighthouseProviderModels, +) # TODO Move this to outputs/reports jobs @@ -1097,3 +1105,363 @@ class TestCheckIntegrationsTask: assert result is False mock_upload.assert_called_once_with(self.tenant_id, self.provider_id, scan_id) + + +@pytest.mark.django_db +class TestCheckLighthouseProviderConnectionTask: + def setup_method(self): + self.tenant_id = str(uuid.uuid4()) + + @pytest.mark.parametrize( + "provider_type,credentials,base_url,expected_result", + [ + ( + LighthouseProviderConfiguration.LLMProviderChoices.OPENAI, + {"api_key": "sk-test123"}, + None, + {"connected": True, "error": None}, + ), + ( + LighthouseProviderConfiguration.LLMProviderChoices.OPENAI_COMPATIBLE, + {"api_key": "sk-test123"}, + "https://openrouter.ai/api/v1", + {"connected": True, "error": None}, + ), + ( + LighthouseProviderConfiguration.LLMProviderChoices.BEDROCK, + { + "access_key_id": "AKIA123", + "secret_access_key": "secret", + "region": "us-east-1", + }, + None, + {"connected": True, "error": None}, + ), + ], + ) + def test_check_connection_success_all_providers( + self, tenants_fixture, provider_type, credentials, base_url, expected_result + ): + """Test successful connection check for all provider types.""" + # Create provider configuration + provider_cfg = LighthouseProviderConfiguration( + tenant_id=tenants_fixture[0].id, + provider_type=provider_type, + base_url=base_url, + is_active=False, + ) + provider_cfg.credentials_decoded = credentials + provider_cfg.save() + + # Mock the appropriate API calls + with ( + patch("tasks.jobs.lighthouse_providers.openai.OpenAI") as mock_openai, + patch("tasks.jobs.lighthouse_providers.boto3.client") as mock_boto3, + ): + mock_client = MagicMock() + mock_client.models.list.return_value = MagicMock() + mock_client.list_foundation_models.return_value = {} + mock_openai.return_value = mock_client + mock_boto3.return_value = mock_client + + # Execute + result = check_lighthouse_provider_connection_task( + provider_config_id=str(provider_cfg.id), + tenant_id=str(tenants_fixture[0].id), + ) + + # Assert + assert result == expected_result + provider_cfg.refresh_from_db() + assert provider_cfg.is_active is True + + @pytest.mark.parametrize( + "provider_type,credentials,base_url,exception_to_raise", + [ + ( + LighthouseProviderConfiguration.LLMProviderChoices.OPENAI, + {"api_key": "sk-invalid"}, + None, + openai.AuthenticationError( + "Invalid API key", response=MagicMock(), body=None + ), + ), + ( + LighthouseProviderConfiguration.LLMProviderChoices.OPENAI_COMPATIBLE, + {"api_key": "sk-invalid"}, + "https://openrouter.ai/api/v1", + openai.APIConnectionError(request=MagicMock()), + ), + ( + LighthouseProviderConfiguration.LLMProviderChoices.BEDROCK, + { + "access_key_id": "AKIA123", + "secret_access_key": "secret", + "region": "us-east-1", + }, + None, + ClientError( + {"Error": {"Code": "AccessDenied", "Message": "Access Denied"}}, + "list_foundation_models", + ), + ), + ], + ) + def test_check_connection_api_failure( + self, + tenants_fixture, + provider_type, + credentials, + base_url, + exception_to_raise, + ): + """Test connection check when API calls fail.""" + # Create provider configuration + provider_cfg = LighthouseProviderConfiguration( + tenant_id=tenants_fixture[0].id, + provider_type=provider_type, + base_url=base_url, + is_active=True, + ) + provider_cfg.credentials_decoded = credentials + provider_cfg.save() + + # Mock the API to raise exception + with ( + patch("tasks.jobs.lighthouse_providers.openai.OpenAI") as mock_openai, + patch("tasks.jobs.lighthouse_providers.boto3.client") as mock_boto3, + ): + mock_client = MagicMock() + if ( + provider_type + == LighthouseProviderConfiguration.LLMProviderChoices.BEDROCK + ): + mock_client.list_foundation_models.side_effect = exception_to_raise + mock_boto3.return_value = mock_client + else: + mock_client.models.list.side_effect = exception_to_raise + mock_openai.return_value = mock_client + + # Execute + result = check_lighthouse_provider_connection_task( + provider_config_id=str(provider_cfg.id), + tenant_id=str(tenants_fixture[0].id), + ) + + # Assert + assert result["connected"] is False + assert result["error"] is not None + provider_cfg.refresh_from_db() + assert provider_cfg.is_active is False + + def test_check_connection_updates_active_status(self, tenants_fixture): + """Test that connection check toggles is_active from True to False on failure.""" + # Create provider with is_active=True + provider_cfg = LighthouseProviderConfiguration( + tenant_id=tenants_fixture[0].id, + provider_type=LighthouseProviderConfiguration.LLMProviderChoices.OPENAI, + base_url=None, + is_active=True, + ) + provider_cfg.credentials_decoded = {"api_key": "sk-test123"} + provider_cfg.save() + + # Mock API to fail + with patch("tasks.jobs.lighthouse_providers.openai.OpenAI") as mock_openai: + mock_client = MagicMock() + mock_client.models.list.side_effect = openai.AuthenticationError( + "Invalid", response=MagicMock(), body=None + ) + mock_openai.return_value = mock_client + + # Execute + result = check_lighthouse_provider_connection_task( + provider_config_id=str(provider_cfg.id), + tenant_id=str(tenants_fixture[0].id), + ) + + # Assert status changed + assert result["connected"] is False + provider_cfg.refresh_from_db() + assert provider_cfg.is_active is False + + def test_check_connection_provider_does_not_exist(self, tenants_fixture): + """Test that checking non-existent provider raises DoesNotExist.""" + non_existent_id = str(uuid.uuid4()) + + with pytest.raises(LighthouseProviderConfiguration.DoesNotExist): + check_lighthouse_provider_connection_task( + provider_config_id=non_existent_id, + tenant_id=str(tenants_fixture[0].id), + ) + + +@pytest.mark.django_db +class TestRefreshLighthouseProviderModelsTask: + def setup_method(self): + self.tenant_id = str(uuid.uuid4()) + + @pytest.mark.parametrize( + "provider_type,credentials,base_url,mock_models,expected_count", + [ + ( + LighthouseProviderConfiguration.LLMProviderChoices.OPENAI, + {"api_key": "sk-test123"}, + None, + {"gpt-5": "gpt-5", "gpt-4o": "gpt-4o"}, + 2, + ), + ( + LighthouseProviderConfiguration.LLMProviderChoices.OPENAI_COMPATIBLE, + {"api_key": "sk-test123"}, + "https://openrouter.ai/api/v1", + {"model-1": "Model One", "model-2": "Model Two"}, + 2, + ), + ( + LighthouseProviderConfiguration.LLMProviderChoices.BEDROCK, + { + "access_key_id": "AKIA123", + "secret_access_key": "secret", + "region": "us-east-1", + }, + None, + {"openai.gpt-oss-120b-1:0": "gpt-oss-120b"}, + 1, + ), + ], + ) + def test_refresh_models_create_new( + self, + tenants_fixture, + provider_type, + credentials, + base_url, + mock_models, + expected_count, + ): + """Test creating new models for all provider types.""" + # Create provider configuration + provider_cfg = LighthouseProviderConfiguration( + tenant_id=tenants_fixture[0].id, + provider_type=provider_type, + base_url=base_url, + is_active=True, + ) + provider_cfg.credentials_decoded = credentials + provider_cfg.save() + + # Mock the fetch functions + with ( + patch( + "tasks.jobs.lighthouse_providers._fetch_openai_models", + return_value=mock_models, + ), + patch( + "tasks.jobs.lighthouse_providers._fetch_openai_compatible_models", + return_value=mock_models, + ), + patch( + "tasks.jobs.lighthouse_providers._fetch_bedrock_models", + return_value=mock_models, + ), + ): + # Execute + result = refresh_lighthouse_provider_models_task( + provider_config_id=str(provider_cfg.id), + tenant_id=str(tenants_fixture[0].id), + ) + + # Assert + assert result["created"] == expected_count + assert result["updated"] == 0 + assert result["deleted"] == 0 + assert ( + LighthouseProviderModels.objects.filter( + provider_configuration=provider_cfg + ).count() + == expected_count + ) + + def test_refresh_models_mixed_operations(self, tenants_fixture): + """Test mixed create, update, and delete operations.""" + # Create provider configuration + provider_cfg = LighthouseProviderConfiguration( + tenant_id=tenants_fixture[0].id, + provider_type=LighthouseProviderConfiguration.LLMProviderChoices.OPENAI, + base_url=None, + is_active=True, + ) + provider_cfg.credentials_decoded = {"api_key": "sk-test123"} + provider_cfg.save() + + # Create 2 existing models (A, B) + LighthouseProviderModels.objects.create( + tenant_id=tenants_fixture[0].id, + provider_configuration=provider_cfg, + model_id="model-a", + model_name="Model A", + ) + LighthouseProviderModels.objects.create( + tenant_id=tenants_fixture[0].id, + provider_configuration=provider_cfg, + model_id="model-b", + model_name="Model B", + ) + + # Mock API to return models B (existing), C (new) - A will be deleted + mock_models = {"model-b": "Model B", "model-c": "Model C"} + with patch( + "tasks.jobs.lighthouse_providers._fetch_openai_models", + return_value=mock_models, + ): + # Execute + result = refresh_lighthouse_provider_models_task( + provider_config_id=str(provider_cfg.id), + tenant_id=str(tenants_fixture[0].id), + ) + + # Assert + assert result["created"] == 1 # model-c created + assert result["updated"] == 1 # model-b updated + assert result["deleted"] == 1 # model-a deleted + + # Verify only B and C exist + remaining_models = LighthouseProviderModels.objects.filter( + provider_configuration=provider_cfg + ) + assert remaining_models.count() == 2 + assert set(remaining_models.values_list("model_id", flat=True)) == { + "model-b", + "model-c", + } + + def test_refresh_models_api_exception(self, tenants_fixture): + """Test refresh when API raises an exception.""" + # Create provider configuration + provider_cfg = LighthouseProviderConfiguration( + tenant_id=tenants_fixture[0].id, + provider_type=LighthouseProviderConfiguration.LLMProviderChoices.OPENAI, + base_url=None, + is_active=True, + ) + provider_cfg.credentials_decoded = {"api_key": "sk-test123"} + provider_cfg.save() + + # Mock fetch to raise exception + with patch( + "tasks.jobs.lighthouse_providers._fetch_openai_models", + side_effect=openai.APIError("API Error", request=MagicMock(), body=None), + ): + # Execute + result = refresh_lighthouse_provider_models_task( + provider_config_id=str(provider_cfg.id), + tenant_id=str(tenants_fixture[0].id), + ) + + # Assert + assert result["created"] == 0 + assert result["updated"] == 0 + assert result["deleted"] == 0 + assert "error" in result + assert result["error"] is not None