From bf04261af6c455839f5d63d6d39d58974efeb1ff Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adri=C3=A1n=20Jes=C3=BAs=20Pe=C3=B1a=20Rodr=C3=ADguez?= Date: Wed, 13 Nov 2024 18:17:08 +0100 Subject: [PATCH] feat(provider-groups): PRWLR-4725 add provider-groups system (#82) * feat(provider-groups): PRWLR-4725 add provider-groups system * feat(provider-groups): PRWLR-4725 add provider-groups migrations * feat(provider-groups): PRWLR-4725 improve provider-groups models --- src/backend/api/filters.py | 15 + src/backend/api/fixtures/6_dev_rbac.json | 62 +++ src/backend/api/migrations/0001_initial.py | 105 +++++ src/backend/api/models.py | 57 +++ src/backend/api/specs/v1.yaml | 505 ++++++++++++++++++++- src/backend/api/tests/test_views.py | 255 ++++++++++- src/backend/api/v1/serializers.py | 83 ++++ src/backend/api/v1/urls.py | 2 + src/backend/api/v1/views.py | 102 +++++ src/backend/conftest.py | 20 + 10 files changed, 1204 insertions(+), 2 deletions(-) create mode 100644 src/backend/api/fixtures/6_dev_rbac.json diff --git a/src/backend/api/filters.py b/src/backend/api/filters.py index 94d9f0ee3d..946a0df1e0 100644 --- a/src/backend/api/filters.py +++ b/src/backend/api/filters.py @@ -25,6 +25,7 @@ from api.models import ( User, Membership, Provider, + ProviderGroup, Resource, ResourceTag, Scan, @@ -136,6 +137,20 @@ class ProviderRelationshipFilterSet(FilterSet): ) +class ProviderGroupFilter(FilterSet): + inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") + updated_at = DateFilter(field_name="updated_at", lookup_expr="date") + + class Meta: + model = ProviderGroup + fields = { + "id": ["exact", "in"], + "name": ["exact", "in"], + "inserted_at": ["gte", "lte"], + "updated_at": ["gte", "lte"], + } + + class ScanFilter(ProviderRelationshipFilterSet): inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") completed_at = DateFilter(field_name="completed_at", lookup_expr="date") diff --git a/src/backend/api/fixtures/6_dev_rbac.json b/src/backend/api/fixtures/6_dev_rbac.json new file mode 100644 index 0000000000..38917e7546 --- /dev/null +++ b/src/backend/api/fixtures/6_dev_rbac.json @@ -0,0 +1,62 @@ +[ + { + "model": "api.providergroup", + "pk": "3fe28fb8-e545-424c-9b8f-69aff638f430", + "fields": { + "name": "first_group", + "inserted_at": "2024-11-13T11:36:19.503Z", + "updated_at": "2024-11-13T11:36:19.503Z", + "tenant": "12646005-9067-4d2a-a098-8bb378604362" + } + }, + { + "model": "api.providergroup", + "pk": "525e91e7-f3f3-4254-bbc3-27ce1ade86b1", + "fields": { + "name": "second_group", + "inserted_at": "2024-11-13T11:36:25.421Z", + "updated_at": "2024-11-13T11:36:25.421Z", + "tenant": "12646005-9067-4d2a-a098-8bb378604362" + } + }, + { + "model": "api.providergroup", + "pk": "481769f5-db2b-447b-8b00-1dee18db90ec", + "fields": { + "name": "third_group", + "inserted_at": "2024-11-13T11:36:37.603Z", + "updated_at": "2024-11-13T11:36:37.603Z", + "tenant": "12646005-9067-4d2a-a098-8bb378604362" + } + }, + { + "model": "api.providergroupmembership", + "pk": "13625bd3-f428-4021-ac1b-b0bd41b6e02f", + "fields": { + "tenant": "12646005-9067-4d2a-a098-8bb378604362", + "provider": "1b59e032-3eb6-4694-93a5-df84cd9b3ce2", + "provider_group": "3fe28fb8-e545-424c-9b8f-69aff638f430", + "inserted_at": "2024-11-13T11:55:17.138Z" + } + }, + { + "model": "api.providergroupmembership", + "pk": "54784ebe-42d2-4937-aa6a-e21c62879567", + "fields": { + "tenant": "12646005-9067-4d2a-a098-8bb378604362", + "provider": "15fce1fa-ecaa-433f-a9dc-62553f3a2555", + "provider_group": "3fe28fb8-e545-424c-9b8f-69aff638f430", + "inserted_at": "2024-11-13T11:55:17.138Z" + } + }, + { + "model": "api.providergroupmembership", + "pk": "c8bd52d5-42a5-48fe-8e0a-3eef154b8ebe", + "fields": { + "tenant": "12646005-9067-4d2a-a098-8bb378604362", + "provider": "15fce1fa-ecaa-433f-a9dc-62553f3a2555", + "provider_group": "525e91e7-f3f3-4254-bbc3-27ce1ade86b1", + "inserted_at": "2024-11-13T11:55:41.237Z" + } + } +] diff --git a/src/backend/api/migrations/0001_initial.py b/src/backend/api/migrations/0001_initial.py index 03def2143b..011c33a00d 100644 --- a/src/backend/api/migrations/0001_initial.py +++ b/src/backend/api/migrations/0001_initial.py @@ -448,6 +448,111 @@ class Migration(migrations.Migration): name="unique_provider_uids", ), ), + migrations.CreateModel( + name="ProviderGroup", + fields=[ + ( + "id", + models.UUIDField( + default=uuid.uuid4, + editable=False, + primary_key=True, + serialize=False, + ), + ), + ("name", models.CharField(max_length=255)), + ("inserted_at", models.DateTimeField(auto_now_add=True)), + ("updated_at", models.DateTimeField(auto_now=True)), + ], + options={ + "db_table": "provider_groups", + }, + ), + migrations.CreateModel( + name="ProviderGroupMembership", + fields=[ + ( + "id", + models.UUIDField( + default=uuid.uuid4, + editable=False, + primary_key=True, + serialize=False, + ), + ), + ("inserted_at", models.DateTimeField(auto_now_add=True)), + ], + options={ + "db_table": "provider_group_memberships", + }, + ), + migrations.AddField( + model_name="providergroup", + name="tenant", + field=models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, + to="api.tenant", + ), + ), + migrations.AddField( + model_name="providergroup", + name="providers", + field=models.ManyToManyField( + related_name="provider_groups", + through="api.ProviderGroupMembership", + to="api.provider", + ), + ), + migrations.AddField( + model_name="providergroupmembership", + name="tenant", + field=models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, to="api.tenant" + ), + ), + migrations.AddField( + model_name="providergroupmembership", + name="provider", + field=models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, to="api.provider" + ), + ), + migrations.AddField( + model_name="providergroupmembership", + name="provider_group", + field=models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, to="api.providergroup" + ), + ), + migrations.AddConstraint( + model_name="providergroup", + constraint=api.rls.RowLevelSecurityConstraint( + "tenant_id", + name="rls_on_providergroup", + statements=["SELECT", "INSERT", "UPDATE", "DELETE"], + ), + ), + migrations.AddConstraint( + model_name="providergroup", + constraint=models.UniqueConstraint( + fields=("tenant_id", "name"), name="unique_group_name_per_tenant" + ), + ), + migrations.AddConstraint( + model_name="providergroupmembership", + constraint=api.rls.RowLevelSecurityConstraint( + "tenant_id", + name="rls_on_providergroupmembership", + statements=["SELECT", "INSERT", "UPDATE", "DELETE"], + ), + ), + migrations.AddConstraint( + model_name="providergroupmembership", + constraint=models.UniqueConstraint( + fields=("provider_id", "provider_group"), + name="unique_provider_group_membership", + ), + ), migrations.CreateModel( name="Task", fields=[ diff --git a/src/backend/api/models.py b/src/backend/api/models.py index 4692193dc1..a189fa601a 100644 --- a/src/backend/api/models.py +++ b/src/backend/api/models.py @@ -235,6 +235,63 @@ class Provider(RowLevelSecurityProtectedModel): ] +class ProviderGroup(RowLevelSecurityProtectedModel): + id = models.UUIDField(primary_key=True, default=uuid4, editable=False) + name = models.CharField(max_length=255) + inserted_at = models.DateTimeField(auto_now_add=True, editable=False) + updated_at = models.DateTimeField(auto_now=True, editable=False) + providers = models.ManyToManyField( + Provider, through="ProviderGroupMembership", related_name="provider_groups" + ) + + class Meta: + db_table = "provider_groups" + constraints = [ + models.UniqueConstraint( + fields=["tenant_id", "name"], + name="unique_group_name_per_tenant", + ), + RowLevelSecurityConstraint( + field="tenant_id", + name="rls_on_%(class)s", + statements=["SELECT", "INSERT", "UPDATE", "DELETE"], + ), + ] + + class JSONAPIMeta: + resource_name = "provider-groups" + + +class ProviderGroupMembership(RowLevelSecurityProtectedModel): + id = models.UUIDField(primary_key=True, default=uuid4, editable=False) + provider = models.ForeignKey( + Provider, + on_delete=models.CASCADE, + ) + provider_group = models.ForeignKey( + ProviderGroup, + on_delete=models.CASCADE, + ) + inserted_at = models.DateTimeField(auto_now_add=True, editable=False) + + class Meta: + db_table = "provider_group_memberships" + constraints = [ + models.UniqueConstraint( + fields=["provider_id", "provider_group"], + name="unique_provider_group_membership", + ), + RowLevelSecurityConstraint( + field="tenant_id", + name="rls_on_%(class)s", + statements=["SELECT", "INSERT", "UPDATE", "DELETE"], + ), + ] + + class JSONAPIMeta: + resource_name = "provider-group-memberships" + + class Task(RowLevelSecurityProtectedModel): id = models.UUIDField(primary_key=True, default=uuid4, editable=False) inserted_at = models.DateTimeField(auto_now_add=True, editable=False) diff --git a/src/backend/api/specs/v1.yaml b/src/backend/api/specs/v1.yaml index 4570c25939..7e83d19137 100644 --- a/src/backend/api/specs/v1.yaml +++ b/src/backend/api/specs/v1.yaml @@ -539,6 +539,296 @@ paths: schema: $ref: '#/components/schemas/OpenApiResponseResponse' description: '' + /api/v1/provider_groups: + get: + operationId: provider_groups_list + description: Retrieve a list of all provider groups with options for filtering + by various criteria. + summary: List all provider groups + parameters: + - in: query + name: fields[provider-groups] + schema: + type: array + items: + type: string + enum: + - name + - inserted_at + - updated_at + - providers + - url + description: endpoint return only specific fields in the response on a per-type + basis by including a fields[TYPE] query parameter. + explode: false + - in: query + name: filter[id] + schema: + type: string + format: uuid + - in: query + name: filter[id__in] + schema: + type: array + items: + type: string + format: uuid + description: Multiple values may be separated by commas. + explode: false + style: form + - in: query + name: filter[inserted_at] + schema: + type: string + format: date + - in: query + name: filter[inserted_at__gte] + schema: + type: string + format: date-time + - in: query + name: filter[inserted_at__lte] + schema: + type: string + format: date-time + - in: query + name: filter[name] + schema: + type: string + - in: query + name: filter[name__in] + schema: + type: array + items: + type: string + description: Multiple values may be separated by commas. + explode: false + style: form + - name: filter[search] + required: false + in: query + description: A search term. + schema: + type: string + - in: query + name: filter[updated_at] + schema: + type: string + format: date + - in: query + name: filter[updated_at__gte] + schema: + type: string + format: date-time + - in: query + name: filter[updated_at__lte] + schema: + type: string + format: date-time + - name: page[number] + required: false + in: query + description: A page number within the paginated result set. + schema: + type: integer + - name: page[size] + required: false + in: query + description: Number of results to return per page. + schema: + type: integer + - name: sort + required: false + in: query + description: '[list of fields to sort by](https://jsonapi.org/format/#fetching-sorting)' + schema: + type: array + items: + type: string + enum: + - id + - -id + - name + - -name + - inserted_at + - -inserted_at + - updated_at + - -updated_at + - providers + - -providers + - url + - -url + explode: false + tags: + - Provider Group + security: + - jwtAuth: [] + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/PaginatedProviderGroupList' + description: '' + post: + operationId: provider_groups_create + description: Add a new provider group to the system by providing the required + provider group details. + summary: Create a new provider group + tags: + - Provider Group + requestBody: + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/ProviderGroupRequest' + application/x-www-form-urlencoded: + schema: + $ref: '#/components/schemas/ProviderGroupRequest' + multipart/form-data: + schema: + $ref: '#/components/schemas/ProviderGroupRequest' + required: true + security: + - jwtAuth: [] + responses: + '201': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/ProviderGroupResponse' + description: '' + /api/v1/provider_groups/{id}: + get: + operationId: provider_groups_retrieve + description: Fetch detailed information about a specific provider group by their + ID. + summary: Retrieve data from a provider group + parameters: + - in: query + name: fields[provider-groups] + schema: + type: array + items: + type: string + enum: + - name + - inserted_at + - updated_at + - providers + - url + description: endpoint return only specific fields in the response on a per-type + basis by including a fields[TYPE] query parameter. + explode: false + - in: path + name: id + schema: + type: string + format: uuid + description: A UUID string identifying this provider group. + required: true + tags: + - Provider Group + security: + - jwtAuth: [] + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/ProviderGroupResponse' + description: '' + patch: + operationId: provider_groups_partial_update + description: Update certain fields of an existing provider group's information + without affecting other fields. + summary: Partially update a provider group + parameters: + - in: path + name: id + schema: + type: string + format: uuid + description: A UUID string identifying this provider group. + required: true + tags: + - Provider Group + requestBody: + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/PatchedProviderGroupUpdateRequest' + application/x-www-form-urlencoded: + schema: + $ref: '#/components/schemas/PatchedProviderGroupUpdateRequest' + multipart/form-data: + schema: + $ref: '#/components/schemas/PatchedProviderGroupUpdateRequest' + required: true + security: + - jwtAuth: [] + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/SerializerMetaclassResponse' + description: '' + delete: + operationId: provider_groups_destroy + description: Remove a provider group from the system by their ID. + summary: Delete a provider group + parameters: + - in: path + name: id + schema: + type: string + format: uuid + description: A UUID string identifying this provider group. + required: true + tags: + - Provider Group + security: + - jwtAuth: [] + responses: + '204': + description: No response body + /api/v1/provider_groups/{id}/providers: + put: + operationId: provider_groups_providers_update + description: Add one or more providers to an existing provider group. + summary: Add providers to a provider group + parameters: + - in: path + name: id + schema: + type: string + format: uuid + description: A UUID string identifying this provider group. + required: true + tags: + - Provider Group + requestBody: + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/ProviderGroupMembershipUpdateRequest' + application/x-www-form-urlencoded: + schema: + $ref: '#/components/schemas/ProviderGroupMembershipUpdateRequest' + multipart/form-data: + schema: + $ref: '#/components/schemas/ProviderGroupMembershipUpdateRequest' + required: true + security: + - jwtAuth: [] + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/OpenApiResponseResponse' + description: '' /api/v1/providers: get: operationId: providers_list @@ -3649,6 +3939,15 @@ components: $ref: '#/components/schemas/Membership' required: - data + PaginatedProviderGroupList: + type: object + properties: + data: + type: array + items: + $ref: '#/components/schemas/ProviderGroup' + required: + - data PaginatedProviderList: type: object properties: @@ -3762,6 +4061,37 @@ components: minLength: 1 required: - data + PatchedProviderGroupUpdateRequest: + type: object + properties: + data: + type: object + required: + - type + - id + additionalProperties: false + properties: + type: + type: string + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + enum: + - provider-groups + id: + type: string + format: uuid + attributes: + type: object + properties: + name: + type: string + minLength: 1 + maxLength: 255 + required: + - name + required: + - data PatchedProviderSecretUpdateRequest: type: object properties: @@ -4300,6 +4630,177 @@ components: $ref: '#/components/schemas/ProviderCreate' required: - data + ProviderGroup: + type: object + required: + - type + - id + additionalProperties: false + properties: + type: + allOf: + - $ref: '#/components/schemas/ProviderGroupTypeEnum' + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + id: + type: string + format: uuid + attributes: + type: object + properties: + name: + type: string + maxLength: 255 + inserted_at: + type: string + format: date-time + readOnly: true + updated_at: + type: string + format: date-time + readOnly: true + required: + - name + relationships: + type: object + properties: + providers: + type: object + properties: + data: + type: array + items: + type: object + properties: + id: + type: string + format: uuid + title: Resource Identifier + description: The identifier of the related object. + type: + type: string + enum: + - Provider + title: Resource Type Name + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common + attributes and relationships. + required: + - id + - type + required: + - data + description: A related resource object from type Provider + title: Provider + readOnly: true + ProviderGroupMembershipUpdateRequest: + type: object + properties: + data: + type: object + required: + - type + additionalProperties: false + properties: + type: + type: string + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + enum: + - provider-group-memberships + attributes: + type: object + properties: + provider_ids: + type: array + items: + type: string + format: uuid + description: List of provider UUIDs to add to the group + required: + - provider_ids + required: + - data + ProviderGroupRequest: + type: object + properties: + data: + type: object + required: + - type + additionalProperties: false + properties: + type: + type: string + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + enum: + - provider-groups + attributes: + type: object + properties: + name: + type: string + minLength: 1 + maxLength: 255 + inserted_at: + type: string + format: date-time + readOnly: true + updated_at: + type: string + format: date-time + readOnly: true + required: + - name + relationships: + type: object + properties: + providers: + type: object + properties: + data: + type: array + items: + type: object + properties: + id: + type: string + format: uuid + title: Resource Identifier + description: The identifier of the related object. + type: + type: string + enum: + - Provider + title: Resource Type Name + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share + common attributes and relationships. + required: + - id + - type + required: + - data + description: A related resource object from type Provider + title: Provider + readOnly: true + required: + - data + ProviderGroupResponse: + type: object + properties: + data: + $ref: '#/components/schemas/ProviderGroup' + required: + - data + ProviderGroupTypeEnum: + type: string + enum: + - provider-groups ProviderResponse: type: object properties: @@ -5291,7 +5792,7 @@ components: type: object properties: data: - $ref: '#/components/schemas/Provider' + $ref: '#/components/schemas/ProviderGroup' required: - data Task: @@ -5867,6 +6368,8 @@ tags: description: Endpoints for managing tenants, along with their memberships. - name: Provider description: Endpoints for managing providers (AWS, GCP, Azure, etc...). +- name: Provider Group + description: Endpoints for managing provider groups. - name: Scan description: Endpoints for triggering manual scans and viewing scan results. - name: Resource diff --git a/src/backend/api/tests/test_views.py b/src/backend/api/tests/test_views.py index 4492c32ed7..3994f417e2 100644 --- a/src/backend/api/tests/test_views.py +++ b/src/backend/api/tests/test_views.py @@ -8,7 +8,16 @@ import pytest from django.urls import reverse from rest_framework import status -from api.models import User, Membership, Provider, Scan, ProviderSecret, Invitation +from api.models import ( + User, + Membership, + Provider, + ProviderGroup, + ProviderGroupMembership, + Scan, + ProviderSecret, + Invitation, +) from api.rls import Tenant from conftest import ( API_JSON_CONTENT_TYPE, @@ -1123,6 +1132,250 @@ class TestProviderViewSet: assert response.status_code == status.HTTP_400_BAD_REQUEST +@pytest.mark.django_db +class TestProviderGroupViewSet: + def test_provider_group_list(self, authenticated_client, provider_groups_fixture): + response = authenticated_client.get(reverse("providergroup-list")) + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == len(provider_groups_fixture) + + def test_provider_group_retrieve( + self, authenticated_client, provider_groups_fixture + ): + provider_group = provider_groups_fixture[0] + response = authenticated_client.get( + reverse("providergroup-detail", kwargs={"pk": provider_group.id}) + ) + assert response.status_code == status.HTTP_200_OK + data = response.json()["data"] + assert data["id"] == str(provider_group.id) + assert data["attributes"]["name"] == provider_group.name + + def test_provider_group_create(self, authenticated_client): + data = { + "data": { + "type": "provider-groups", + "attributes": { + "name": "Test Provider Group", + }, + } + } + response = authenticated_client.post( + reverse("providergroup-list"), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_201_CREATED + response_data = response.json()["data"] + assert response_data["attributes"]["name"] == "Test Provider Group" + assert ProviderGroup.objects.filter(name="Test Provider Group").exists() + + def test_provider_group_create_invalid(self, authenticated_client): + data = { + "data": { + "type": "provider-groups", + "attributes": { + # Name is missing + }, + } + } + response = authenticated_client.post( + reverse("providergroup-list"), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + errors = response.json()["errors"] + assert errors[0]["source"]["pointer"] == "/data/attributes/name" + + def test_provider_group_partial_update( + self, authenticated_client, provider_groups_fixture + ): + provider_group = provider_groups_fixture[1] + data = { + "data": { + "id": str(provider_group.id), + "type": "provider-groups", + "attributes": { + "name": "Updated Provider Group Name", + }, + } + } + response = authenticated_client.patch( + reverse("providergroup-detail", kwargs={"pk": provider_group.id}), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_200_OK + provider_group.refresh_from_db() + assert provider_group.name == "Updated Provider Group Name" + + def test_provider_group_partial_update_invalid( + self, authenticated_client, provider_groups_fixture + ): + provider_group = provider_groups_fixture[2] + data = { + "data": { + "id": str(provider_group.id), + "type": "provider-groups", + "attributes": { + "name": "", # Invalid name + }, + } + } + response = authenticated_client.patch( + reverse("providergroup-detail", kwargs={"pk": provider_group.id}), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + errors = response.json()["errors"] + assert errors[0]["source"]["pointer"] == "/data/attributes/name" + + def test_provider_group_destroy( + self, authenticated_client, provider_groups_fixture + ): + provider_group = provider_groups_fixture[2] + response = authenticated_client.delete( + reverse("providergroup-detail", kwargs={"pk": provider_group.id}) + ) + assert response.status_code == status.HTTP_204_NO_CONTENT + assert not ProviderGroup.objects.filter(id=provider_group.id).exists() + + def test_provider_group_destroy_invalid(self, authenticated_client): + response = authenticated_client.delete( + reverse("providergroup-detail", kwargs={"pk": "non-existent-id"}) + ) + assert response.status_code == status.HTTP_404_NOT_FOUND + + def test_provider_group_providers_update( + self, authenticated_client, provider_groups_fixture, providers_fixture + ): + provider_group = provider_groups_fixture[0] + provider_ids = [str(provider.id) for provider in providers_fixture] + + data = { + "data": { + "type": "provider-group-memberships", + "id": str(provider_group.id), + "attributes": {"provider_ids": provider_ids}, + } + } + + response = authenticated_client.put( + reverse("providergroup-providers", kwargs={"pk": provider_group.id}), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_200_OK + memberships = ProviderGroupMembership.objects.filter( + provider_group=provider_group + ) + assert memberships.count() == len(provider_ids) + for membership in memberships: + assert str(membership.provider_id) in provider_ids + + def test_provider_group_providers_update_non_existent_provider( + self, authenticated_client, provider_groups_fixture, providers_fixture + ): + provider_group = provider_groups_fixture[0] + provider_ids = [str(provider.id) for provider in providers_fixture] + provider_ids[-1] = "1b59e032-3eb6-4694-93a5-df84cd9b3ce2" + + data = { + "data": { + "type": "provider-group-memberships", + "id": str(provider_group.id), + "attributes": {"provider_ids": provider_ids}, + } + } + + response = authenticated_client.put( + reverse("providergroup-providers", kwargs={"pk": provider_group.id}), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + errors = response.json()["errors"] + assert ( + errors[0]["detail"] + == f"The following provider IDs do not exist: {provider_ids[-1]}" + ) + + def test_provider_group_providers_update_invalid_provider( + self, authenticated_client, provider_groups_fixture + ): + provider_group = provider_groups_fixture[1] + invalid_provider_id = "non-existent-id" + data = { + "data": { + "type": "provider-group-memberships", + "id": str(provider_group.id), + "attributes": {"provider_ids": [invalid_provider_id]}, + } + } + + response = authenticated_client.put( + reverse("providergroup-providers", kwargs={"pk": provider_group.id}), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + errors = response.json()["errors"] + assert errors[0]["detail"] == "Must be a valid UUID." + + def test_provider_group_providers_update_invalid_payload( + self, authenticated_client, provider_groups_fixture + ): + provider_group = provider_groups_fixture[2] + data = { + # Missing "provider_ids" + } + + response = authenticated_client.put( + reverse("providergroup-providers", kwargs={"pk": provider_group.id}), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + errors = response.json()["errors"] + assert errors[0]["detail"] == "Received document does not contain primary data" + + def test_provider_group_retrieve_not_found(self, authenticated_client): + response = authenticated_client.get( + reverse("providergroup-detail", kwargs={"pk": "non-existent-id"}) + ) + assert response.status_code == status.HTTP_404_NOT_FOUND + + def test_provider_group_list_filters( + self, authenticated_client, provider_groups_fixture + ): + provider_group = provider_groups_fixture[0] + response = authenticated_client.get( + reverse("providergroup-list"), {"filter[name]": provider_group.name} + ) + assert response.status_code == status.HTTP_200_OK + data = response.json()["data"] + assert len(data) == 1 + assert data[0]["attributes"]["name"] == provider_group.name + + def test_provider_group_list_sorting( + self, authenticated_client, provider_groups_fixture + ): + response = authenticated_client.get( + reverse("providergroup-list"), {"sort": "name"} + ) + assert response.status_code == status.HTTP_200_OK + data = response.json()["data"] + names = [item["attributes"]["name"] for item in data] + assert names == sorted(names) + + def test_provider_group_invalid_method(self, authenticated_client): + response = authenticated_client.put(reverse("providergroup-list")) + assert response.status_code == status.HTTP_405_METHOD_NOT_ALLOWED + + @pytest.mark.django_db class TestProviderSecretViewSet: def test_provider_secrets_list(self, authenticated_client, provider_secret_fixture): diff --git a/src/backend/api/v1/serializers.py b/src/backend/api/v1/serializers.py index c28fb2e01b..2a4f96f04e 100644 --- a/src/backend/api/v1/serializers.py +++ b/src/backend/api/v1/serializers.py @@ -18,6 +18,8 @@ from api.models import ( User, Membership, Provider, + ProviderGroup, + ProviderGroupMembership, Scan, Task, Resource, @@ -354,6 +356,87 @@ class MembershipSerializer(serializers.ModelSerializer): fields = ["id", "user", "tenant", "role", "date_joined"] +# Provider Groups +class ProviderGroupSerializer(RLSSerializer, BaseWriteSerializer): + providers = serializers.ResourceRelatedField(many=True, read_only=True) + + def validate(self, attrs): + tenant = self.context["tenant_id"] + name = attrs.get("name", self.instance.name if self.instance else None) + + # Exclude the current instance when checking for uniqueness during updates + queryset = ProviderGroup.objects.filter(tenant=tenant, name=name) + if self.instance: + queryset = queryset.exclude(pk=self.instance.pk) + + if queryset.exists(): + raise serializers.ValidationError( + { + "name": "A provider group with this name already exists for this tenant." + } + ) + + return super().validate(attrs) + + class Meta: + model = ProviderGroup + fields = ["id", "name", "inserted_at", "updated_at", "providers", "url"] + read_only_fields = ["id", "inserted_at", "updated_at"] + extra_kwargs = { + "id": {"read_only": True}, + "inserted_at": {"read_only": True}, + "updated_at": {"read_only": True}, + } + + +class ProviderGroupUpdateSerializer(RLSSerializer, BaseWriteSerializer): + """ + Serializer for updating the ProviderGroup model. + Only allows "name" field to be updated. + """ + + class Meta: + model = ProviderGroup + fields = ["id", "name"] + + +class ProviderGroupMembershipUpdateSerializer(RLSSerializer, BaseWriteSerializer): + """ + Serializer for modifying provider group memberships + """ + + provider_ids = serializers.ListField( + child=serializers.UUIDField(), + help_text="List of provider UUIDs to add to the group", + ) + + def validate(self, attrs): + tenant_id = self.context["tenant_id"] + provider_ids = attrs.get("provider_ids", []) + + existing_provider_ids = set( + Provider.objects.filter( + id__in=provider_ids, tenant_id=tenant_id + ).values_list("id", flat=True) + ) + provided_provider_ids = set(provider_ids) + + missing_provider_ids = provided_provider_ids - existing_provider_ids + + if missing_provider_ids: + raise serializers.ValidationError( + { + "provider_ids": f"The following provider IDs do not exist: {', '.join(str(id) for id in missing_provider_ids)}" + } + ) + + return super().validate(attrs) + + class Meta: + model = ProviderGroupMembership + fields = ["id", "provider_ids"] + + # Providers class ProviderEnumSerializerField(serializers.ChoiceField): def __init__(self, **kwargs): diff --git a/src/backend/api/v1/urls.py b/src/backend/api/v1/urls.py index dc8f5f07a6..4f45d38f25 100644 --- a/src/backend/api/v1/urls.py +++ b/src/backend/api/v1/urls.py @@ -15,6 +15,7 @@ from api.v1.views import ( TaskViewSet, ResourceViewSet, FindingViewSet, + ProviderGroupViewSet, ProviderSecretViewSet, InvitationViewSet, InvitationAcceptViewSet, @@ -25,6 +26,7 @@ router = routers.DefaultRouter(trailing_slash=False) router.register(r"users", UserViewSet, basename="user") router.register(r"tenants", TenantViewSet, basename="tenant") router.register(r"providers", ProviderViewSet, basename="provider") +router.register(r"provider_groups", ProviderGroupViewSet, basename="providergroup") router.register(r"scans", ScanViewSet, basename="scan") router.register(r"tasks", TaskViewSet, basename="task") router.register(r"resources", ResourceViewSet, basename="resource") diff --git a/src/backend/api/v1/views.py b/src/backend/api/v1/views.py index a6fdcd6df7..c310a74d0f 100644 --- a/src/backend/api/v1/views.py +++ b/src/backend/api/v1/views.py @@ -33,6 +33,7 @@ from api.base_views import BaseTenantViewset, BaseRLSViewSet, BaseUserViewset from api.db_router import MainRouter from api.filters import ( ProviderFilter, + ProviderGroupFilter, TenantFilter, MembershipFilter, ScanFilter, @@ -47,6 +48,8 @@ from api.models import ( User, Membership, Provider, + ProviderGroup, + ProviderGroupMembership, Scan, Task, Resource, @@ -64,6 +67,9 @@ from api.v1.serializers import ( UserCreateSerializer, UserUpdateSerializer, MembershipSerializer, + ProviderGroupSerializer, + ProviderGroupUpdateSerializer, + ProviderGroupMembershipUpdateSerializer, ProviderSerializer, ProviderCreateSerializer, ProviderUpdateSerializer, @@ -168,6 +174,10 @@ class SchemaView(SpectacularAPIView): "name": "Provider", "description": "Endpoints for managing providers (AWS, GCP, Azure, etc...).", }, + { + "name": "Provider Group", + "description": "Endpoints for managing provider groups.", + }, { "name": "Scan", "description": "Endpoints for triggering manual scans and viewing scan results.", @@ -466,6 +476,98 @@ class TenantMembersViewSet(BaseTenantViewset): return Response(status=status.HTTP_204_NO_CONTENT) +@extend_schema(tags=["Provider Group"]) +@extend_schema_view( + list=extend_schema( + summary="List all provider groups", + description="Retrieve a list of all provider groups with options for filtering by various criteria.", + ), + retrieve=extend_schema( + summary="Retrieve data from a provider group", + description="Fetch detailed information about a specific provider group by their ID.", + ), + create=extend_schema( + summary="Create a new provider group", + description="Add a new provider group to the system by providing the required provider group details.", + ), + partial_update=extend_schema( + summary="Partially update a provider group", + description="Update certain fields of an existing provider group's information without affecting other fields.", + request=ProviderGroupUpdateSerializer, + responses={200: ProviderGroupSerializer}, + ), + destroy=extend_schema( + summary="Delete a provider group", + description="Remove a provider group from the system by their ID.", + ), + update=extend_schema(exclude=True), +) +class ProviderGroupViewSet(BaseRLSViewSet): + queryset = ProviderGroup.objects.all() + serializer_class = ProviderGroupSerializer + filterset_class = ProviderGroupFilter + http_method_names = ["get", "post", "patch", "put", "delete"] + ordering = ["inserted_at"] + + def get_queryset(self): + return ProviderGroup.objects.prefetch_related("providers") + + def get_serializer_class(self): + if self.action == "partial_update": + return ProviderGroupUpdateSerializer + elif self.action == "providers": + if hasattr(self, "response_serializer_class"): + return self.response_serializer_class + return ProviderGroupMembershipUpdateSerializer + return super().get_serializer_class() + + @extend_schema( + tags=["Provider Group"], + summary="Add providers to a provider group", + description="Add one or more providers to an existing provider group.", + request=ProviderGroupMembershipUpdateSerializer, + responses={200: OpenApiResponse(response=ProviderGroupSerializer)}, + ) + @action(detail=True, methods=["put"], url_name="providers") + def providers(self, request, pk=None): + provider_group = self.get_object() + + # Validate input data + serializer = self.get_serializer_class()( + data=request.data, + context=self.get_serializer_context(), + ) + serializer.is_valid(raise_exception=True) + + provider_ids = serializer.validated_data["provider_ids"] + + # Update memberships + ProviderGroupMembership.objects.filter( + provider_group=provider_group, tenant_id=request.tenant_id + ).delete() + + provider_group_memberships = [ + ProviderGroupMembership( + tenant_id=self.request.tenant_id, + provider_group=provider_group, + provider_id=provider_id, + ) + for provider_id in provider_ids + ] + + ProviderGroupMembership.objects.bulk_create( + provider_group_memberships, ignore_conflicts=True + ) + + # Return the updated provider group with providers + provider_group.refresh_from_db() + self.response_serializer_class = ProviderGroupSerializer + response_serializer = ProviderGroupSerializer( + provider_group, context=self.get_serializer_context() + ) + return Response(data=response_serializer.data, status=status.HTTP_200_OK) + + @extend_schema_view( list=extend_schema( summary="List all providers", diff --git a/src/backend/conftest.py b/src/backend/conftest.py index d62ea275b4..430f0aacea 100644 --- a/src/backend/conftest.py +++ b/src/backend/conftest.py @@ -17,6 +17,7 @@ from api.models import ( from api.models import ( User, Provider, + ProviderGroup, Resource, ResourceTag, Scan, @@ -189,6 +190,25 @@ def providers_fixture(tenants_fixture): return provider1, provider2, provider3, provider4, provider5 +@pytest.fixture +def provider_groups_fixture(tenants_fixture): + tenant, *_ = tenants_fixture + pgroup1 = ProviderGroup.objects.create( + name="Group One", + tenant_id=tenant.id, + ) + pgroup2 = ProviderGroup.objects.create( + name="Group Two", + tenant_id=tenant.id, + ) + pgroup3 = ProviderGroup.objects.create( + name="Group Three", + tenant_id=tenant.id, + ) + + return pgroup1, pgroup2, pgroup3 + + @pytest.fixture def provider_secret_fixture(providers_fixture): return tuple(