From 4ab119d6c9be0620f15d9d41756314176e654e7a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?V=C3=ADctor=20Fern=C3=A1ndez=20Poyatos?= Date: Tue, 5 Nov 2024 15:30:53 +0100 Subject: [PATCH] feat(Invitation): PRWLR-4722 Add invitations endpoints (#71) * feat(Invitation): PRWLR-4722 add model and enum * feat(Invitation): PRWLR-4722 add serializers * feat(Invitation): PRWLR-4722 add filters * feat(Invitation): PRWLR-4722 update token field constraints * feat(Invitation): PRWLR-4722 add serializers * feat(Invitation): PRWLR-4722 add views, url and custom logic * feat(Invitation): PRWLR-4722 update unique constraint in model * feat(Invitation): PRWLR-4722 update serializer validation error messages * fix(Invitation): PRWLR-4722 fix view logic * feat(User): PRWLR-4722 add invitation_code query param and logic to create user view * fix(Invitation): PRWLR-4722 fix invitation creation tenant filter * chore: PRWLR-4722 add comments * feat(Invitation): PRWLR-4722 add email filter to view * fix(Utils): PRWLR-4722 fix datetime functions * fix(User): PRWLR-4722 fix bug when creating users * fix(Tests): PRWLR-4722 adapt unit and integration tests * test(db-utils): PRWLR-4722 add new unit tests * test(Invitation): PRWLR-4722 add unit tests * test(Invitation): PRWLR-4722 add unit tests * fix(Invitation): PRWLR-4722 fix views and serializers * feat(Invitation): PRWLR-4722 refactor invitation validation and tests * chore: PRWLR-4722 update API spec * test(Invitation): PRWLR-4722 add more unit tests * feat(Invitation): PRWLR-4722 refactor invitation urls * chore: PRWLR-4722 update API spec --- src/backend/api/db_utils.py | 29 + src/backend/api/exceptions.py | 8 + src/backend/api/filters.py | 25 + src/backend/api/migrations/0001_initial.py | 91 +++ src/backend/api/models.py | 52 +- src/backend/api/specs/v1.yaml | 735 +++++++++++++++++- .../tests/integration/test_authentication.py | 2 + src/backend/api/tests/test_db_utils.py | 108 +++ src/backend/api/tests/test_utils.py | 109 ++- src/backend/api/tests/test_views.py | 614 ++++++++++++++- src/backend/api/utils.py | 70 +- src/backend/api/v1/serializers.py | 101 ++- src/backend/api/v1/urls.py | 20 +- src/backend/api/v1/views.py | 207 ++++- src/backend/conftest.py | 35 +- src/backend/tasks/tests/test_scan.py | 4 +- 16 files changed, 2156 insertions(+), 54 deletions(-) create mode 100644 src/backend/api/tests/test_db_utils.py diff --git a/src/backend/api/db_utils.py b/src/backend/api/db_utils.py index 282277322d..d90c2c2340 100644 --- a/src/backend/api/db_utils.py +++ b/src/backend/api/db_utils.py @@ -1,4 +1,6 @@ +import secrets from contextlib import contextmanager +from datetime import datetime, timezone, timedelta from django.conf import settings from django.contrib.auth.models import BaseUserManager @@ -71,6 +73,21 @@ def enum_to_choices(enum_class): return [(item.value, item.name.replace("_", " ").title()) for item in enum_class] +def one_week_from_now(): + """ + Return a datetime object with a date one week from now. + """ + return datetime.now(timezone.utc) + timedelta(days=7) + + +def generate_random_token(length: int = 14, symbols: str | None = None) -> str: + """ + Generate a random token with the specified length. + """ + _symbols = "23456789ABCDEFGHJKMNPQRSTVWXYZ" + return "".join(secrets.choice(symbols or _symbols) for _ in range(length)) + + # Postgres Enums @@ -240,3 +257,15 @@ class ProviderSecretTypeEnum(EnumType): class ProviderSecretTypeEnumField(PostgresEnumField): def __init__(self, *args, **kwargs): super().__init__("provider_secret_type", *args, **kwargs) + + +# Postgres enum definition for Provider secrets type + + +class InvitationStateEnum(EnumType): + enum_type_name = "invitation_state" + + +class InvitationStateEnumField(PostgresEnumField): + def __init__(self, *args, **kwargs): + super().__init__("invitation_state", *args, **kwargs) diff --git a/src/backend/api/exceptions.py b/src/backend/api/exceptions.py index 604ca0405c..12bc788d68 100644 --- a/src/backend/api/exceptions.py +++ b/src/backend/api/exceptions.py @@ -1,4 +1,6 @@ from django.core.exceptions import ValidationError as django_validation_error +from rest_framework import status +from rest_framework.exceptions import APIException from rest_framework_json_api.exceptions import exception_handler from rest_framework_json_api.serializers import ValidationError from rest_framework_simplejwt.exceptions import TokenError, InvalidToken @@ -24,6 +26,12 @@ class ModelValidationError(ValidationError): ) +class InvitationTokenExpiredException(APIException): + status_code = status.HTTP_410_GONE + default_detail = "The invitation token has expired and is no longer valid." + default_code = "token_expired" + + def custom_exception_handler(exc, context): if isinstance(exc, django_validation_error): if hasattr(exc, "error_dict"): diff --git a/src/backend/api/filters.py b/src/backend/api/filters.py index 756fb6276a..bc95a91a38 100644 --- a/src/backend/api/filters.py +++ b/src/backend/api/filters.py @@ -20,6 +20,7 @@ from api.db_utils import ( FindingDeltaEnumField, StatusEnumField, SeverityEnumField, + InvitationStateEnumField, ) from api.models import ( Membership, @@ -33,6 +34,7 @@ from api.models import ( SeverityChoices, StatusChoices, ProviderSecret, + Invitation, ) from api.rls import Tenant from api.uuid_utils import ( @@ -378,3 +380,26 @@ class ProviderSecretFilter(FilterSet): fields = { "name": ["exact", "icontains"], } + + +class InvitationFilter(FilterSet): + inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") + updated_at = DateFilter(field_name="updated_at", lookup_expr="date") + expires_at = DateFilter(field_name="expires_at", lookup_expr="date") + state = ChoiceFilter(choices=Invitation.State.choices) + state__in = ChoiceInFilter(choices=Invitation.State.choices, lookup_expr="in") + + class Meta: + model = Invitation + fields = { + "email": ["exact", "icontains"], + "inserted_at": ["date", "gte", "lte"], + "updated_at": ["date", "gte", "lte"], + "expires_at": ["date", "gte", "lte"], + "inviter": ["exact"], + } + filter_overrides = { + InvitationStateEnumField: { + "filter_class": CharFilter, + } + } diff --git a/src/backend/api/migrations/0001_initial.py b/src/backend/api/migrations/0001_initial.py index 6a6e389d3d..7069270823 100644 --- a/src/backend/api/migrations/0001_initial.py +++ b/src/backend/api/migrations/0001_initial.py @@ -33,6 +33,8 @@ from api.db_utils import ( StateEnumField, StateEnum, ScanTriggerEnumField, + InvitationStateEnum, + InvitationStateEnumField, register_enum, DB_PROWLER_USER, DB_PROWLER_PASSWORD, @@ -49,6 +51,7 @@ from api.models import ( SeverityChoices, Membership, ProviderSecret, + Invitation, ) DB_NAME = settings.DATABASES["default"]["NAME"] @@ -97,6 +100,11 @@ ProviderSecretTypeEnumMigration = PostgresEnumMigration( ), ) +InvitationStateEnumMigration = PostgresEnumMigration( + enum_name="invitation_state", + enum_values=tuple(state[0] for state in Invitation.State.choices), +) + class Migration(migrations.Migration): initial = True @@ -1207,4 +1215,87 @@ class Migration(migrations.Migration): statements=["SELECT", "INSERT", "UPDATE", "DELETE"], ), ), + migrations.RunPython( + InvitationStateEnumMigration.create_enum_type, + reverse_code=InvitationStateEnumMigration.drop_enum_type, + ), + migrations.RunPython(partial(register_enum, enum_class=InvitationStateEnum)), + migrations.CreateModel( + name="Invitation", + fields=[ + ( + "id", + models.UUIDField( + default=uuid.uuid4, + editable=False, + primary_key=True, + serialize=False, + ), + ), + ("inserted_at", models.DateTimeField(auto_now_add=True)), + ("updated_at", models.DateTimeField(auto_now=True)), + ("email", models.EmailField(max_length=254)), + ( + "state", + InvitationStateEnumField( + choices=[ + ("pending", "Invitation is pending"), + ("accepted", "Invitation was accepted by a user"), + ("expired", "Invitation expired after the configured time"), + ("revoked", "Invitation was revoked by a user"), + ], + default="pending", + ), + ), + ( + "token", + models.CharField( + unique=True, + default=api.db_utils.generate_random_token, + editable=False, + max_length=14, + validators=[django.core.validators.MinLengthValidator(14)], + ), + ), + ( + "expires_at", + models.DateTimeField(default=api.db_utils.one_week_from_now), + ), + ( + "inviter", + models.ForeignKey( + null=True, + on_delete=django.db.models.deletion.SET_NULL, + related_name="invitations", + related_query_name="invitation", + to=settings.AUTH_USER_MODEL, + ), + ), + ( + "tenant", + models.ForeignKey( + on_delete=django.db.models.deletion.CASCADE, to="api.tenant" + ), + ), + ], + options={ + "db_table": "invitations", + "abstract": False, + }, + ), + migrations.AddConstraint( + model_name="invitation", + constraint=models.UniqueConstraint( + fields=("tenant", "token", "email"), + name="unique_tenant_token_email_by_invitation", + ), + ), + migrations.AddConstraint( + model_name="invitation", + constraint=api.rls.RowLevelSecurityConstraint( + "tenant_id", + name="rls_on_invitation", + statements=["SELECT", "INSERT", "UPDATE", "DELETE"], + ), + ), ] diff --git a/src/backend/api/models.py b/src/backend/api/models.py index e71b2b57cc..ed7a2c5c91 100644 --- a/src/backend/api/models.py +++ b/src/backend/api/models.py @@ -27,6 +27,9 @@ from api.db_utils import ( StatusEnumField, CustomUserManager, ProviderSecretTypeEnumField, + InvitationStateEnumField, + one_week_from_now, + generate_random_token, ) from api.exceptions import ModelValidationError from api.rls import ( @@ -509,11 +512,8 @@ class Finding(PostgresPartitionedModel, RowLevelSecurityProtectedModel): impact_extended = models.TextField(blank=True, null=True) raw_result = models.JSONField(default=dict) - # TODO: review usability tags = models.JSONField(default=dict, null=True, blank=True) - check_id = models.CharField(max_length=100, blank=False, null=False) - # TODO: review usability check_metadata = models.JSONField(default=dict, null=False) # Relationships @@ -669,3 +669,49 @@ class ProviderSecret(RowLevelSecurityProtectedModel): def secret(self, value): encrypted_data = fernet.encrypt(json.dumps(value).encode()) self._secret = encrypted_data + + +class Invitation(RowLevelSecurityProtectedModel): + class State(models.TextChoices): + PENDING = "pending", _("Invitation is pending") + ACCEPTED = "accepted", _("Invitation was accepted by a user") + EXPIRED = "expired", _("Invitation expired after the configured time") + REVOKED = "revoked", _("Invitation was revoked by a user") + + id = models.UUIDField(primary_key=True, default=uuid4, editable=False) + inserted_at = models.DateTimeField(auto_now_add=True, editable=False) + updated_at = models.DateTimeField(auto_now=True, editable=False) + email = models.EmailField(max_length=254, blank=False, null=False) + state = InvitationStateEnumField(choices=State.choices, default=State.PENDING) + token = models.CharField( + max_length=14, + unique=True, + default=generate_random_token, + editable=False, + blank=False, + null=False, + validators=[MinLengthValidator(14)], + ) + expires_at = models.DateTimeField(default=one_week_from_now) + inviter = models.ForeignKey( + User, + on_delete=models.SET_NULL, + related_name="invitations", + related_query_name="invitation", + null=True, + ) + + class Meta(RowLevelSecurityProtectedModel.Meta): + db_table = "invitations" + + constraints = [ + models.UniqueConstraint( + fields=("tenant", "token", "email"), + name="unique_tenant_token_email_by_invitation", + ), + RowLevelSecurityConstraint( + field="tenant_id", + name="rls_on_%(class)s", + statements=["SELECT", "INSERT", "UPDATE", "DELETE"], + ), + ] diff --git a/src/backend/api/specs/v1.yaml b/src/backend/api/specs/v1.yaml index b148d9ed20..90b10f1534 100644 --- a/src/backend/api/specs/v1.yaml +++ b/src/backend/api/specs/v1.yaml @@ -503,6 +503,35 @@ paths: schema: $ref: '#/components/schemas/FindingResponse' description: '' + /api/v1/invitations/accept: + post: + operationId: invitations_accept_create + description: Accept an invitation to an existing tenant. This invitation cannot + be expired and the emails must match. + summary: Accept an invitation + tags: + - Invitation + requestBody: + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/InvitationAcceptRequest' + application/x-www-form-urlencoded: + schema: + $ref: '#/components/schemas/InvitationAcceptRequest' + multipart/form-data: + schema: + $ref: '#/components/schemas/InvitationAcceptRequest' + required: true + security: + - jwtAuth: [] + responses: + '201': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/OpenApiResponseResponse' + description: '' /api/v1/providers: get: operationId: providers_list @@ -2221,6 +2250,310 @@ paths: responses: '204': description: No response body + /api/v1/tenants/invitations: + get: + operationId: tenants_invitations_list + description: Retrieve a list of all tenant invitations with options for filtering + by various criteria. + summary: List all invitations + parameters: + - in: query + name: fields[Invitation] + schema: + type: array + items: + type: string + enum: + - inserted_at + - updated_at + - email + - state + - token + - expires_at + - inviter + - 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[email] + schema: + type: string + - in: query + name: filter[email__icontains] + schema: + type: string + - in: query + name: filter[expires_at] + schema: + type: string + format: date + - in: query + name: filter[expires_at__date] + schema: + type: string + format: date + - in: query + name: filter[expires_at__gte] + schema: + type: string + format: date-time + - in: query + name: filter[expires_at__lte] + schema: + type: string + format: date-time + - in: query + name: filter[inserted_at] + schema: + type: string + format: date + - in: query + name: filter[inserted_at__date] + 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[inviter] + schema: + type: string + format: uuid + - name: filter[search] + required: false + in: query + description: A search term. + schema: + type: string + - in: query + name: filter[state] + schema: + type: string + enum: + - accepted + - expired + - pending + - revoked + description: |- + * `pending` - Invitation is pending + * `accepted` - Invitation was accepted by a user + * `expired` - Invitation expired after the configured time + * `revoked` - Invitation was revoked by a user + - in: query + name: filter[state__in] + schema: + type: array + items: + type: string + enum: + - accepted + - expired + - pending + - revoked + description: |- + Multiple values may be separated by commas. + + * `pending` - Invitation is pending + * `accepted` - Invitation was accepted by a user + * `expired` - Invitation expired after the configured time + * `revoked` - Invitation was revoked by a user + explode: false + style: form + - in: query + name: filter[updated_at] + schema: + type: string + format: date + - in: query + name: filter[updated_at__date] + 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: + - inserted_at + - -inserted_at + - updated_at + - -updated_at + - expires_at + - -expires_at + - state + - -state + - inviter + - -inviter + explode: false + tags: + - Invitation + security: + - jwtAuth: [] + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/PaginatedInvitationList' + description: '' + post: + operationId: tenants_invitations_create + description: Add a new tenant invitation to the system by providing the required + invitation details. The invited user will have to accept the invitations or + create an account using the given code. + summary: Invite a user to a tenant + tags: + - Invitation + requestBody: + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/InvitationCreateRequest' + application/x-www-form-urlencoded: + schema: + $ref: '#/components/schemas/InvitationCreateRequest' + multipart/form-data: + schema: + $ref: '#/components/schemas/InvitationCreateRequest' + required: true + security: + - jwtAuth: [] + responses: + '201': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/InvitationCreateResponse' + description: '' + /api/v1/tenants/invitations/{id}: + get: + operationId: tenants_invitations_retrieve + description: Fetch detailed information about a specific invitation by its ID. + summary: Retrieve data from a tenant invitation + parameters: + - in: query + name: fields[Invitation] + schema: + type: array + items: + type: string + enum: + - inserted_at + - updated_at + - email + - state + - token + - expires_at + - inviter + - 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 + required: true + tags: + - Invitation + security: + - jwtAuth: [] + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/InvitationResponse' + description: '' + patch: + operationId: tenants_invitations_partial_update + description: Update certain fields of an existing tenant invitation's information + without affecting other fields. + summary: Partially update a tenant invitation + parameters: + - in: path + name: id + schema: + type: string + format: uuid + required: true + tags: + - Invitation + requestBody: + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/PatchedInvitationUpdateRequest' + application/x-www-form-urlencoded: + schema: + $ref: '#/components/schemas/PatchedInvitationUpdateRequest' + multipart/form-data: + schema: + $ref: '#/components/schemas/PatchedInvitationUpdateRequest' + required: true + security: + - jwtAuth: [] + responses: + '200': + content: + application/vnd.api+json: + schema: + $ref: '#/components/schemas/InvitationUpdateResponse' + description: '' + delete: + operationId: tenants_invitations_destroy + description: Revoke a tenant invitation from the system by their ID. + summary: Revoke a tenant invitation + parameters: + - in: path + name: id + schema: + type: string + format: uuid + required: true + tags: + - Invitation + security: + - jwtAuth: [] + responses: + '204': + description: No response body /api/v1/tokens: post: operationId: tokens_create @@ -2287,6 +2620,13 @@ paths: description: Create a new user account by providing the necessary registration details. summary: Register a new user + parameters: + - in: query + name: invitation_token + schema: + type: string + example: F3NMFPNDZHR4Z9 + description: Optional invitation code for joining an existing tenant. tags: - User requestBody: @@ -2739,7 +3079,7 @@ components: type: string enum: - Finding - Membership: + Invitation: type: object required: - type @@ -2748,13 +3088,328 @@ components: properties: type: allOf: - - $ref: '#/components/schemas/MembershipTypeEnum' + - $ref: '#/components/schemas/Type4b9Enum' 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: + inserted_at: + type: string + format: date-time + readOnly: true + updated_at: + type: string + format: date-time + readOnly: true + email: + type: string + format: email + maxLength: 254 + state: + enum: + - pending + - accepted + - expired + - revoked + type: string + description: |- + * `pending` - Invitation is pending + * `accepted` - Invitation was accepted by a user + * `expired` - Invitation expired after the configured time + * `revoked` - Invitation was revoked by a user + token: + type: string + readOnly: true + expires_at: + type: string + format: date-time + required: + - email + relationships: + type: object + properties: + inviter: + type: object + properties: + data: + type: object + properties: + id: + type: string + format: uuid + type: + type: string + enum: + - User + 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: The identifier of the related object. + title: Resource Identifier + nullable: true + InvitationAcceptRequest: + 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: + - Invitation + attributes: + type: object + properties: + invitation_token: + type: string + writeOnly: true + minLength: 1 + required: + - invitation_token + required: + - data + InvitationCreate: + type: object + required: + - type + additionalProperties: false + properties: + type: + allOf: + - $ref: '#/components/schemas/Type4b9Enum' + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. + attributes: + type: object + properties: + email: + type: string + format: email + maxLength: 254 + expires_at: + type: string + format: date-time + description: UTC. Default 7 days. If this attribute is provided, it + must be at least 24 hours in the future. + state: + enum: + - pending + - accepted + - expired + - revoked + type: string + description: |- + * `pending` - Invitation is pending + * `accepted` - Invitation was accepted by a user + * `expired` - Invitation expired after the configured time + * `revoked` - Invitation was revoked by a user + readOnly: true + token: + type: string + readOnly: true + required: + - email + relationships: + type: object + properties: + inviter: + type: object + properties: + data: + type: object + properties: + id: + type: string + format: uuid + type: + type: string + enum: + - User + 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: The identifier of the related object. + title: Resource Identifier + readOnly: true + nullable: true + InvitationCreateRequest: + 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: + - Invitation + attributes: + type: object + properties: + email: + type: string + format: email + minLength: 1 + maxLength: 254 + expires_at: + type: string + format: date-time + description: UTC. Default 7 days. If this attribute is provided, + it must be at least 24 hours in the future. + state: + enum: + - pending + - accepted + - expired + - revoked + type: string + description: |- + * `pending` - Invitation is pending + * `accepted` - Invitation was accepted by a user + * `expired` - Invitation expired after the configured time + * `revoked` - Invitation was revoked by a user + readOnly: true + token: + type: string + readOnly: true + minLength: 1 + required: + - email + relationships: + type: object + properties: + inviter: + type: object + properties: + data: + type: object + properties: + id: + type: string + format: uuid + type: + type: string + enum: + - User + 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: The identifier of the related object. + title: Resource Identifier + readOnly: true + nullable: true + required: + - data + InvitationCreateResponse: + type: object + properties: + data: + $ref: '#/components/schemas/InvitationCreate' + required: + - data + InvitationResponse: + type: object + properties: + data: + $ref: '#/components/schemas/Invitation' + required: + - data + InvitationUpdate: + type: object + required: + - type + - id + additionalProperties: false + properties: + type: + allOf: + - $ref: '#/components/schemas/Type4b9Enum' + 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: + email: + type: string + format: email + maxLength: 254 + expires_at: + type: string + format: date-time + state: + enum: + - pending + - accepted + - expired + - revoked + type: string + description: |- + * `pending` - Invitation is pending + * `accepted` - Invitation was accepted by a user + * `expired` - Invitation expired after the configured time + * `revoked` - Invitation was revoked by a user + readOnly: true + token: + type: string + readOnly: true + InvitationUpdateResponse: + type: object + properties: + data: + $ref: '#/components/schemas/InvitationUpdate' + required: + - data + Membership: + type: object + required: + - type + additionalProperties: false + properties: + type: + allOf: + - $ref: '#/components/schemas/MembershipTypeEnum' + description: The [type](https://jsonapi.org/format/#document-resource-object-identification) + member is used to describe resource objects that share common attributes + and relationships. attributes: type: object properties: @@ -2840,7 +3495,7 @@ components: type: object properties: data: - $ref: '#/components/schemas/Task' + $ref: '#/components/schemas/Membership' required: - data PaginatedFindingList: @@ -2852,6 +3507,15 @@ components: $ref: '#/components/schemas/Finding' required: - data + PaginatedInvitationList: + type: object + properties: + data: + type: array + items: + $ref: '#/components/schemas/Invitation' + required: + - data PaginatedMembershipList: type: object properties: @@ -2915,6 +3579,56 @@ components: $ref: '#/components/schemas/Tenant' required: - data + PatchedInvitationUpdateRequest: + 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: + - Invitation + id: + type: string + format: uuid + attributes: + type: object + properties: + email: + type: string + format: email + minLength: 1 + maxLength: 254 + expires_at: + type: string + format: date-time + state: + enum: + - pending + - accepted + - expired + - revoked + type: string + description: |- + * `pending` - Invitation is pending + * `accepted` - Invitation was accepted by a user + * `expired` - Invitation expired after the configured time + * `revoked` - Invitation was revoked by a user + readOnly: true + token: + type: string + readOnly: true + minLength: 1 + required: + - data PatchedProviderSecretUpdateRequest: type: object properties: @@ -2932,7 +3646,9 @@ components: and relationships. enum: - ProviderSecret - id: {} + id: + type: string + format: uuid attributes: type: object properties: @@ -3917,7 +4633,9 @@ components: 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: {} + id: + type: string + format: uuid attributes: type: object properties: @@ -4786,6 +5504,10 @@ components: type: string enum: - ProviderSecret + Type4b9Enum: + type: string + enum: + - Invitation Type4e8Enum: type: string enum: @@ -5023,3 +5745,6 @@ tags: - name: Task description: Endpoints for task management, allowing retrieval of task status and revoking tasks that have not started. +- name: Invitation + description: Endpoints for tenant invitations management, allowing retrieval and + filtering of invitations, creating new invitations, accepting and revoking them. diff --git a/src/backend/api/tests/integration/test_authentication.py b/src/backend/api/tests/integration/test_authentication.py index 1000f76555..c306c44cef 100644 --- a/src/backend/api/tests/integration/test_authentication.py +++ b/src/backend/api/tests/integration/test_authentication.py @@ -1,10 +1,12 @@ import pytest from django.urls import reverse +from unittest.mock import patch from rest_framework.test import APIClient from conftest import TEST_PASSWORD, get_api_tokens, get_authorization_header +@patch("api.v1.views.MainRouter.admin_db", new="default") @pytest.mark.django_db def test_basic_authentication(): client = APIClient() diff --git a/src/backend/api/tests/test_db_utils.py b/src/backend/api/tests/test_db_utils.py new file mode 100644 index 0000000000..15cbf88399 --- /dev/null +++ b/src/backend/api/tests/test_db_utils.py @@ -0,0 +1,108 @@ +from datetime import datetime, timezone +from enum import Enum +from unittest.mock import patch + +from api.db_utils import enum_to_choices, one_week_from_now, generate_random_token + + +class TestEnumToChoices: + def test_enum_to_choices_simple(self): + class Color(Enum): + RED = 1 + GREEN = 2 + BLUE = 3 + + expected_result = [ + (1, "Red"), + (2, "Green"), + (3, "Blue"), + ] + + result = enum_to_choices(Color) + assert result == expected_result + + def test_enum_to_choices_with_underscores(self): + class Status(Enum): + PENDING_APPROVAL = "pending" + IN_PROGRESS = "in_progress" + COMPLETED_SUCCESSFULLY = "completed" + + expected_result = [ + ("pending", "Pending Approval"), + ("in_progress", "In Progress"), + ("completed", "Completed Successfully"), + ] + + result = enum_to_choices(Status) + assert result == expected_result + + def test_enum_to_choices_empty_enum(self): + class EmptyEnum(Enum): + pass + + expected_result = [] + + result = enum_to_choices(EmptyEnum) + assert result == expected_result + + def test_enum_to_choices_numeric_values(self): + class Numbers(Enum): + ONE = 1 + TWO = 2 + THREE = 3 + + expected_result = [ + (1, "One"), + (2, "Two"), + (3, "Three"), + ] + + result = enum_to_choices(Numbers) + assert result == expected_result + + +class TestOneWeekFromNow: + def test_one_week_from_now(self): + with patch("api.db_utils.datetime") as mock_datetime: + mock_datetime.now.return_value = datetime(2023, 1, 1, tzinfo=timezone.utc) + expected_result = datetime(2023, 1, 8, tzinfo=timezone.utc) + + result = one_week_from_now() + assert result == expected_result + + def test_one_week_from_now_with_timezone(self): + with patch("api.db_utils.datetime") as mock_datetime: + mock_datetime.now.return_value = datetime( + 2023, 6, 15, 12, 0, tzinfo=timezone.utc + ) + expected_result = datetime(2023, 6, 22, 12, 0, tzinfo=timezone.utc) + + result = one_week_from_now() + assert result == expected_result + + +class TestGenerateRandomToken: + def test_generate_random_token_default_length(self): + token = generate_random_token() + assert len(token) == 14 + + def test_generate_random_token_custom_length(self): + length = 20 + token = generate_random_token(length=length) + assert len(token) == length + + def test_generate_random_token_with_symbols(self): + symbols = "ABC123" + token = generate_random_token(length=10, symbols=symbols) + assert len(token) == 10 + assert all(char in symbols for char in token) + + def test_generate_random_token_unique(self): + tokens = {generate_random_token() for _ in range(1000)} + # Assuming that generating 1000 tokens should result in unique values + assert len(tokens) == 1000 + + def test_generate_random_token_no_symbols_provided(self): + token = generate_random_token(length=5, symbols="") + # Default symbols + assert len(token) == 5 diff --git a/src/backend/api/tests/test_utils.py b/src/backend/api/tests/test_utils.py index 5d75b0820d..fcde1a9999 100644 --- a/src/backend/api/tests/test_utils.py +++ b/src/backend/api/tests/test_utils.py @@ -1,11 +1,16 @@ -from unittest.mock import MagicMock, patch +from datetime import datetime, timedelta, timezone +from unittest.mock import patch, MagicMock import pytest from prowler.providers.aws.aws_provider import AwsProvider from prowler.providers.azure.azure_provider import AzureProvider from prowler.providers.gcp.gcp_provider import GcpProvider from prowler.providers.kubernetes.kubernetes_provider import KubernetesProvider +from rest_framework.exceptions import ValidationError, NotFound +from api.db_router import MainRouter +from api.exceptions import InvitationTokenExpiredException +from api.models import Invitation from api.models import Provider from api.utils import ( merge_dicts, @@ -14,6 +19,7 @@ from api.utils import ( prowler_provider_connection_test, get_prowler_provider_kwargs, ) +from api.utils import validate_invitation class TestMergeDicts: @@ -209,3 +215,104 @@ class TestGetProwlerProviderKwargs: expected_result = {} assert result == expected_result + + +class TestValidateInvitation: + @pytest.fixture + def invitation(self): + invitation = MagicMock(spec=Invitation) + invitation.token = "VALID_TOKEN" + invitation.email = "user@example.com" + invitation.expires_at = datetime.now(timezone.utc) + timedelta(days=1) + invitation.state = Invitation.State.PENDING + invitation.tenant = MagicMock() + return invitation + + def test_valid_invitation(self, invitation): + with patch("api.utils.Invitation.objects.using") as mock_using: + mock_db = mock_using.return_value + mock_db.get.return_value = invitation + + result = validate_invitation("VALID_TOKEN", "user@example.com") + + assert result == invitation + mock_db.get.assert_called_once_with( + token="VALID_TOKEN", email="user@example.com" + ) + + def test_invitation_not_found_raises_validation_error(self): + with patch("api.utils.Invitation.objects.using") as mock_using: + mock_db = mock_using.return_value + mock_db.get.side_effect = Invitation.DoesNotExist + + with pytest.raises(ValidationError) as exc_info: + validate_invitation("INVALID_TOKEN", "user@example.com") + + assert exc_info.value.detail == { + "invitation_token": "Invalid invitation code." + } + mock_db.get.assert_called_once_with( + token="INVALID_TOKEN", email="user@example.com" + ) + + def test_invitation_not_found_raises_not_found(self): + with patch("api.utils.Invitation.objects.using") as mock_using: + mock_db = mock_using.return_value + mock_db.get.side_effect = Invitation.DoesNotExist + + with pytest.raises(NotFound) as exc_info: + validate_invitation( + "INVALID_TOKEN", "user@example.com", raise_not_found=True + ) + + assert exc_info.value.detail == "Invitation is not valid." + mock_db.get.assert_called_once_with( + token="INVALID_TOKEN", email="user@example.com" + ) + + def test_invitation_expired(self, invitation): + expired_time = datetime.now(timezone.utc) - timedelta(days=1) + invitation.expires_at = expired_time + + with patch("api.utils.Invitation.objects.using") as mock_using, patch( + "api.utils.datetime" + ) as mock_datetime: + mock_db = mock_using.return_value + mock_db.get.return_value = invitation + mock_datetime.now.return_value = datetime.now(timezone.utc) + + with pytest.raises(InvitationTokenExpiredException): + validate_invitation("VALID_TOKEN", "user@example.com") + + # Ensure the invitation state was updated to EXPIRED + assert invitation.state == Invitation.State.EXPIRED + invitation.save.assert_called_once_with(using=MainRouter.admin_db) + + def test_invitation_not_pending(self, invitation): + invitation.state = Invitation.State.ACCEPTED + + with patch("api.utils.Invitation.objects.using") as mock_using: + mock_db = mock_using.return_value + mock_db.get.return_value = invitation + + with pytest.raises(ValidationError) as exc_info: + validate_invitation("VALID_TOKEN", "user@example.com") + + assert exc_info.value.detail == { + "invitation_token": "This invitation is no longer valid." + } + + def test_invitation_with_different_email(self): + with patch("api.utils.Invitation.objects.using") as mock_using: + mock_db = mock_using.return_value + mock_db.get.side_effect = Invitation.DoesNotExist + + with pytest.raises(ValidationError) as exc_info: + validate_invitation("VALID_TOKEN", "different@example.com") + + assert exc_info.value.detail == { + "invitation_token": "Invalid invitation code." + } + mock_db.get.assert_called_once_with( + token="VALID_TOKEN", email="different@example.com" + ) diff --git a/src/backend/api/tests/test_views.py b/src/backend/api/tests/test_views.py index 1c132f5142..87b20d5bf6 100644 --- a/src/backend/api/tests/test_views.py +++ b/src/backend/api/tests/test_views.py @@ -1,18 +1,20 @@ import json from datetime import datetime +from datetime import timezone, timedelta from unittest.mock import ANY, Mock, patch import jwt import pytest -from api.models import Membership, Provider, ProviderSecret, Scan, User +from django.urls import reverse +from rest_framework import status + +from api.models import User, Membership, Provider, Scan, ProviderSecret, Invitation from api.rls import Tenant from conftest import ( API_JSON_CONTENT_TYPE, TEST_PASSWORD, TEST_USER, ) -from django.urls import reverse -from rest_framework import status TODAY = str(datetime.today().date()) @@ -34,6 +36,7 @@ class TestUserViewSet: assert response.status_code == status.HTTP_200_OK assert response.json()["data"]["attributes"]["email"] == create_test_user.email + @patch("api.db_router.MainRouter.admin_db", new="default") def test_users_create(self, client): valid_user_payload = { "name": "test", @@ -50,6 +53,7 @@ class TestUserViewSet: == valid_user_payload["email"].lower() ) + @patch("api.db_router.MainRouter.admin_db", new="default") def test_users_create_duplicated_email(self, client): # Create a user self.test_users_create(client) @@ -104,6 +108,7 @@ class TestUserViewSet: "NonExistentEmail@prowler.com", ], ) + @patch("api.db_router.MainRouter.admin_db", new="default") def test_users_create_used_email(self, authenticated_client, email): # First user created; no errors should occur user_payload = { @@ -295,7 +300,7 @@ class TestTenantViewSet: @pytest.fixture def extra_users(self, tenants_fixture): - _, tenant2 = tenants_fixture + _, tenant2, _ = tenants_fixture user2 = User.objects.create_user( name="testing2", password=TEST_PASSWORD, @@ -324,7 +329,7 @@ class TestTenantViewSet: assert len(response.json()["data"]) == len(tenants_fixture) def test_tenants_retrieve(self, authenticated_client, tenants_fixture): - tenant1, _ = tenants_fixture + tenant1, *_ = tenants_fixture response = authenticated_client.get( reverse("tenant-detail", kwargs={"pk": tenant1.id}) ) @@ -343,7 +348,7 @@ class TestTenantViewSet: ) assert response.status_code == status.HTTP_201_CREATED # Two tenants from the fixture + the new one - assert Tenant.objects.count() == 3 + assert Tenant.objects.count() == 4 assert ( response.json()["data"]["attributes"]["name"] == valid_tenant_payload["name"] @@ -359,7 +364,7 @@ class TestTenantViewSet: assert response.status_code == status.HTTP_400_BAD_REQUEST def test_tenants_partial_update(self, authenticated_client, tenants_fixture): - tenant1, _ = tenants_fixture + tenant1, *_ = tenants_fixture new_name = "This is the new name" payload = { "data": { @@ -380,7 +385,7 @@ class TestTenantViewSet: def test_tenants_partial_update_invalid_content_type( self, authenticated_client, tenants_fixture ): - tenant1, _ = tenants_fixture + tenant1, *_ = tenants_fixture response = authenticated_client.patch( reverse("tenant-detail", kwargs={"pk": tenant1.id}), data={} ) @@ -389,7 +394,7 @@ class TestTenantViewSet: def test_tenants_partial_update_invalid_content( self, authenticated_client, tenants_fixture ): - tenant1, _ = tenants_fixture + tenant1, *_ = tenants_fixture new_name = "This is the new name" payload = {"name": new_name} response = authenticated_client.patch( @@ -400,12 +405,12 @@ class TestTenantViewSet: assert response.status_code == status.HTTP_400_BAD_REQUEST def test_tenants_delete(self, authenticated_client, tenants_fixture): - tenant1, _ = tenants_fixture + tenant1, *_ = tenants_fixture response = authenticated_client.delete( reverse("tenant-detail", kwargs={"pk": tenant1.id}) ) assert response.status_code == status.HTTP_204_NO_CONTENT - assert Tenant.objects.count() == 1 + assert Tenant.objects.count() == len(tenants_fixture) - 1 def test_tenants_delete_invalid(self, authenticated_client): response = authenticated_client.delete( @@ -417,7 +422,7 @@ class TestTenantViewSet: def test_tenants_list_filter_search(self, authenticated_client, tenants_fixture): """Search is applied to tenants_fixture name.""" - tenant1, _ = tenants_fixture + tenant1, *_ = tenants_fixture response = authenticated_client.get( reverse("tenant-list"), {"filter[search]": tenant1.name} ) @@ -426,7 +431,7 @@ class TestTenantViewSet: assert response.json()["data"][0]["attributes"]["name"] == tenant1.name def test_tenants_list_query_param_name(self, authenticated_client, tenants_fixture): - tenant1, _ = tenants_fixture + tenant1, *_ = tenants_fixture response = authenticated_client.get( reverse("tenant-list"), {"name": tenant1.name} ) @@ -441,11 +446,11 @@ class TestTenantViewSet: ( [ ("name", "Tenant One", 1), - ("name.icontains", "Tenant", 2), - ("inserted_at", TODAY, 2), - ("inserted_at.gte", "2024-01-01", 2), + ("name.icontains", "Tenant", 3), + ("inserted_at", TODAY, 3), + ("inserted_at.gte", "2024-01-01", 3), ("inserted_at.lte", "2024-01-01", 0), - ("updated_at.gte", "2024-01-01", 2), + ("updated_at.gte", "2024-01-01", 3), ("updated_at.lte", "2024-01-01", 0), ] ), @@ -497,16 +502,16 @@ class TestTenantViewSet: assert response.json()["meta"]["pagination"]["pages"] == len(tenants_fixture) def test_tenants_list_sort_name(self, authenticated_client, tenants_fixture): - _, tenant2 = tenants_fixture + _, tenant2, _ = tenants_fixture response = authenticated_client.get(reverse("tenant-list"), {"sort": "-name"}) assert response.status_code == status.HTTP_200_OK - assert len(response.json()["data"]) == 2 + assert len(response.json()["data"]) == 3 assert response.json()["data"][0]["attributes"]["name"] == tenant2.name def test_tenants_list_memberships_as_owner( self, authenticated_client, tenants_fixture, extra_users ): - _, tenant2 = tenants_fixture + _, tenant2, _ = tenants_fixture response = authenticated_client.get( reverse("tenant-membership-list", kwargs={"tenant_pk": tenant2.id}) ) @@ -517,7 +522,7 @@ class TestTenantViewSet: def test_tenants_list_memberships_as_member( self, authenticated_client, tenants_fixture, extra_users ): - _, tenant2 = tenants_fixture + _, tenant2, _ = tenants_fixture _, user3_membership = extra_users user3, membership3 = user3_membership token_response = authenticated_client.post( @@ -539,7 +544,7 @@ class TestTenantViewSet: def test_tenants_delete_own_membership_as_member( self, authenticated_client, tenants_fixture, extra_users ): - tenant1, _ = tenants_fixture + tenant1, *_ = tenants_fixture membership = Membership.objects.get(tenant=tenant1, user__email=TEST_USER) response = authenticated_client.delete( @@ -555,7 +560,7 @@ class TestTenantViewSet: self, authenticated_client, tenants_fixture, extra_users ): # With extra_users, tenant2 has 2 owners - _, tenant2 = tenants_fixture + _, tenant2, _ = tenants_fixture user_membership = Membership.objects.get(tenant=tenant2, user__email=TEST_USER) response = authenticated_client.delete( reverse( @@ -569,7 +574,7 @@ class TestTenantViewSet: def test_tenants_delete_own_membership_as_last_owner( self, authenticated_client, tenants_fixture ): - _, tenant2 = tenants_fixture + _, tenant2, _ = tenants_fixture user_membership = Membership.objects.get(tenant=tenant2, user__email=TEST_USER) response = authenticated_client.delete( reverse( @@ -583,7 +588,7 @@ class TestTenantViewSet: def test_tenants_delete_another_membership_as_owner( self, authenticated_client, tenants_fixture, extra_users ): - _, tenant2 = tenants_fixture + _, tenant2, _ = tenants_fixture _, user3_membership = extra_users user3, membership3 = user3_membership @@ -599,7 +604,7 @@ class TestTenantViewSet: def test_tenants_delete_another_membership_as_member( self, authenticated_client, tenants_fixture, extra_users ): - _, tenant2 = tenants_fixture + _, tenant2, _ = tenants_fixture _, user3_membership = extra_users user3, membership3 = user3_membership @@ -619,10 +624,10 @@ class TestTenantViewSet: def test_tenants_list_memberships_not_member_of_tenant(self, authenticated_client): # Create a tenant the user is not a member of - tenant3 = Tenant.objects.create(name="Tenant Three") + tenant4 = Tenant.objects.create(name="Tenant Four") response = authenticated_client.get( - reverse("tenant-membership-list", kwargs={"tenant_pk": tenant3.id}) + reverse("tenant-membership-list", kwargs={"tenant_pk": tenant4.id}) ) assert response.status_code == status.HTTP_404_NOT_FOUND @@ -635,7 +640,7 @@ class TestMembershipViewSet: reverse("user-membership-list", kwargs={"user_pk": user_id}), ) assert response.status_code == status.HTTP_200_OK - assert len(response.json()["data"]) == len(tenants_fixture) + assert len(response.json()["data"]) == 2 def test_memberships_retrieve(self, authenticated_client, tenants_fixture): user_id = authenticated_client.user.pk @@ -2227,3 +2232,554 @@ class TestJWTFields: assert ( isinstance(payload[id_field], str) and payload[id_field] ), f"The field '{id_field}' is not a valid string" + + +@pytest.mark.django_db +class TestInvitationViewSet: + TOMORROW = datetime.now(timezone.utc) + timedelta(days=1, hours=1) + TOMORROW_ISO = TOMORROW.isoformat() + + def test_invitations_list(self, authenticated_client, invitations_fixture): + response = authenticated_client.get(reverse("invitation-list")) + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == len(invitations_fixture) + + def test_invitations_retrieve(self, authenticated_client, invitations_fixture): + invitation1, _ = invitations_fixture + response = authenticated_client.get( + reverse( + "invitation-detail", + kwargs={"pk": invitation1.id}, + ), + ) + assert response.status_code == status.HTTP_200_OK + assert response.json()["data"]["attributes"]["email"] == invitation1.email + assert response.json()["data"]["attributes"]["state"] == invitation1.state + assert response.json()["data"]["attributes"]["token"] == invitation1.token + assert response.json()["data"]["relationships"]["inviter"]["data"]["id"] == str( + invitation1.inviter.id + ) + + def test_invitations_invalid_retrieve(self, authenticated_client): + response = authenticated_client.get( + reverse( + "invitation-detail", + kwargs={ + "pk": "f498b103-c760-4785-9a3e-e23fafbb7b02", + }, + ), + ) + assert response.status_code == status.HTTP_404_NOT_FOUND + + def test_invitations_create_valid(self, authenticated_client, create_test_user): + user = create_test_user + data = { + "data": { + "type": "Invitation", + "attributes": { + "email": "any_email@prowler.com", + "expires_at": self.TOMORROW_ISO, + }, + } + } + response = authenticated_client.post( + reverse("invitation-list"), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_201_CREATED + assert Invitation.objects.count() == 1 + assert ( + response.json()["data"]["attributes"]["email"] + == data["data"]["attributes"]["email"] + ) + assert response.json()["data"]["attributes"]["expires_at"] == data["data"][ + "attributes" + ]["expires_at"].replace("+00:00", "Z") + assert ( + response.json()["data"]["attributes"]["state"] + == Invitation.State.PENDING.value + ) + assert response.json()["data"]["relationships"]["inviter"]["data"]["id"] == str( + user.id + ) + + @pytest.mark.parametrize( + "email", + [ + "invalid_email", + "invalid_email@", + # There is a pending invitation with this email + "testing@prowler.com", + # User is already a member of the tenant + TEST_USER, + ], + ) + def test_invitations_create_invalid_email( + self, email, authenticated_client, invitations_fixture + ): + data = { + "data": { + "type": "Invitation", + "attributes": { + "email": email, + "expires_at": self.TOMORROW_ISO, + }, + } + } + response = authenticated_client.post( + reverse("invitation-list"), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert response.json()["errors"][0]["code"] == "invalid" + assert ( + response.json()["errors"][0]["source"]["pointer"] + == "/data/attributes/email" + ) + + def test_invitations_create_invalid_expires_at( + self, authenticated_client, invitations_fixture + ): + data = { + "data": { + "type": "Invitation", + "attributes": { + "email": "thisisarandomemail@prowler.com", + "expires_at": ( + datetime.now(timezone.utc) + timedelta(hours=23) + ).isoformat(), + }, + } + } + response = authenticated_client.post( + reverse("invitation-list"), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert response.json()["errors"][0]["code"] == "invalid" + assert ( + response.json()["errors"][0]["source"]["pointer"] + == "/data/attributes/expires_at" + ) + + def test_invitations_partial_update_valid( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + new_email = "new_email@prowler.com" + new_expires_at = datetime.now(timezone.utc) + timedelta(days=7) + new_expires_at_iso = new_expires_at.isoformat() + data = { + "data": { + "id": str(invitation.id), + "type": "Invitation", + "attributes": { + "email": new_email, + "expires_at": new_expires_at_iso, + }, + } + } + assert invitation.email != new_email + assert invitation.expires_at != new_expires_at + + response = authenticated_client.patch( + reverse( + "invitation-detail", + kwargs={"pk": str(invitation.id)}, + ), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_200_OK + invitation.refresh_from_db() + + assert invitation.email == new_email + assert invitation.expires_at == new_expires_at + + @pytest.mark.parametrize( + "email", + [ + "invalid_email", + "invalid_email@", + # There is a pending invitation with this email + "testing@prowler.com", + # User is already a member of the tenant + TEST_USER, + ], + ) + def test_invitations_partial_update_invalid_email( + self, email, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + data = { + "data": { + "id": str(invitation.id), + "type": "Invitation", + "attributes": { + "email": email, + "expires_at": self.TOMORROW_ISO, + }, + } + } + response = authenticated_client.patch( + reverse( + "invitation-detail", + kwargs={"pk": str(invitation.id)}, + ), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert response.json()["errors"][0]["code"] == "invalid" + assert ( + response.json()["errors"][0]["source"]["pointer"] + == "/data/attributes/email" + ) + + def test_invitations_partial_update_invalid_expires_at( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + data = { + "data": { + "id": str(invitation.id), + "type": "Invitation", + "attributes": { + "expires_at": ( + datetime.now(timezone.utc) + timedelta(hours=23) + ).isoformat(), + }, + } + } + response = authenticated_client.patch( + reverse( + "invitation-detail", + kwargs={"pk": str(invitation.id)}, + ), + data=json.dumps(data), + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert response.json()["errors"][0]["code"] == "invalid" + assert ( + response.json()["errors"][0]["source"]["pointer"] + == "/data/attributes/expires_at" + ) + + def test_invitations_partial_update_invalid_content_type( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + response = authenticated_client.patch( + reverse( + "invitation-detail", + kwargs={"pk": str(invitation.id)}, + ), + data={}, + ) + assert response.status_code == status.HTTP_415_UNSUPPORTED_MEDIA_TYPE + + def test_invitations_partial_update_invalid_content( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + response = authenticated_client.patch( + reverse( + "invitation-detail", + kwargs={"pk": str(invitation.id)}, + ), + data={"email": "invalid_email"}, + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + + def test_invitations_partial_update_invalid_invitation(self, authenticated_client): + response = authenticated_client.patch( + reverse( + "invitation-detail", + kwargs={"pk": "54611fc8-b02e-4cc1-aaaa-34acae625629"}, + ), + data={}, + content_type="application/vnd.api+json", + ) + assert response.status_code == status.HTTP_404_NOT_FOUND + + def test_invitations_delete(self, authenticated_client, invitations_fixture): + invitation, *_ = invitations_fixture + assert invitation.state == Invitation.State.PENDING.value + + response = authenticated_client.delete( + reverse( + "invitation-detail", + kwargs={"pk": str(invitation.id)}, + ) + ) + invitation.refresh_from_db() + assert response.status_code == status.HTTP_204_NO_CONTENT + assert invitation.state == Invitation.State.REVOKED.value + + def test_invitations_invalid_delete(self, authenticated_client): + response = authenticated_client.delete( + reverse( + "invitation-detail", + kwargs={"pk": "54611fc8-b02e-4cc1-aaaa-34acae625629"}, + ) + ) + assert response.status_code == status.HTTP_404_NOT_FOUND + + def test_invitations_invalid_delete_invalid_state( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + invitation.state = Invitation.State.ACCEPTED.value + invitation.save() + + response = authenticated_client.delete( + reverse( + "invitation-detail", + kwargs={"pk": str(invitation.id)}, + ) + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert response.json()["errors"][0]["code"] == "invalid" + assert response.json()["errors"][0]["source"]["pointer"] == "/data" + assert ( + response.json()["errors"][0]["detail"] + == "This invitation cannot be revoked." + ) + + @patch("api.db_router.MainRouter.admin_db", new="default") + def test_invitations_accept_invitation_new_user(self, client, invitations_fixture): + invitation, *_ = invitations_fixture + + data = { + "name": "test", + "password": "newpassword123", + "email": invitation.email, + } + assert invitation.state == Invitation.State.PENDING.value + assert not User.objects.filter(email__iexact=invitation.email).exists() + + response = client.post( + reverse("user-list") + f"?invitation_token={invitation.token}", + data=data, + format="json", + ) + + invitation.refresh_from_db() + assert response.status_code == status.HTTP_201_CREATED + assert User.objects.filter(email__iexact=invitation.email).exists() + assert invitation.state == Invitation.State.ACCEPTED.value + assert Membership.objects.filter( + user__email__iexact=invitation.email, tenant=invitation.tenant + ).exists() + + @patch("api.db_router.MainRouter.admin_db", new="default") + def test_invitations_accept_invitation_existing_user( + self, authenticated_client, create_test_user, tenants_fixture + ): + *_, tenant = tenants_fixture + user = create_test_user + + invitation = Invitation.objects.create( + tenant=tenant, + email=TEST_USER, + inviter=user, + expires_at=self.TOMORROW, + ) + + data = { + "invitation_token": invitation.token, + } + + assert not Membership.objects.filter( + user__email__iexact=user.email, tenant=tenant + ).exists() + + response = authenticated_client.post( + reverse("invitation-accept"), data=data, format="json" + ) + + assert response.status_code == status.HTTP_201_CREATED + invitation.refresh_from_db() + assert Membership.objects.filter( + user__email__iexact=user.email, tenant=tenant + ).exists() + assert invitation.state == Invitation.State.ACCEPTED.value + + @patch("api.db_router.MainRouter.admin_db", new="default") + def test_invitations_accept_invitation_invalid_token(self, authenticated_client): + data = { + "invitation_token": "invalid_token", + } + + response = authenticated_client.post( + reverse("invitation-accept"), data=data, format="json" + ) + + assert response.status_code == status.HTTP_404_NOT_FOUND + assert response.json()["errors"][0]["code"] == "not_found" + + @patch("api.db_router.MainRouter.admin_db", new="default") + def test_invitations_accept_invitation_invalid_token_expired( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + invitation.expires_at = datetime.now(timezone.utc) - timedelta(days=1) + invitation.email = TEST_USER + invitation.save() + + data = { + "invitation_token": invitation.token, + } + + response = authenticated_client.post( + reverse("invitation-accept"), data=data, format="json" + ) + + assert response.status_code == status.HTTP_410_GONE + + @patch("api.db_router.MainRouter.admin_db", new="default") + def test_invitations_accept_invitation_invalid_token_expired_new_user( + self, client, invitations_fixture + ): + new_email = "new_email@prowler.com" + invitation, *_ = invitations_fixture + invitation.expires_at = datetime.now(timezone.utc) - timedelta(days=1) + invitation.email = new_email + invitation.save() + + data = { + "name": "test", + "password": "newpassword123", + "email": new_email, + } + + response = client.post( + reverse("user-list") + f"?invitation_token={invitation.token}", + data=data, + format="json", + ) + + assert response.status_code == status.HTTP_410_GONE + + @patch("api.db_router.MainRouter.admin_db", new="default") + def test_invitations_accept_invitation_invalid_token_accepted( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + invitation.state = Invitation.State.ACCEPTED.value + invitation.email = TEST_USER + invitation.save() + + data = { + "invitation_token": invitation.token, + } + + response = authenticated_client.post( + reverse("invitation-accept"), data=data, format="json" + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert response.json()["errors"][0]["code"] == "invalid" + assert ( + response.json()["errors"][0]["detail"] + == "This invitation is no longer valid." + ) + + @patch("api.db_router.MainRouter.admin_db", new="default") + def test_invitations_accept_invitation_invalid_token_revoked( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + invitation.state = Invitation.State.REVOKED.value + invitation.email = TEST_USER + invitation.save() + + data = { + "invitation_token": invitation.token, + } + + response = authenticated_client.post( + reverse("invitation-accept"), data=data, format="json" + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert ( + response.json()["errors"][0]["detail"] + == "This invitation is no longer valid." + ) + + @pytest.mark.parametrize( + "filter_name, filter_value, expected_count", + ( + [ + ("inserted_at", TODAY, 2), + ("inserted_at.gte", "2024-01-01", 2), + ("inserted_at.lte", "2024-01-01", 0), + ("updated_at.gte", "2024-01-01", 2), + ("updated_at.lte", "2024-01-01", 0), + ("expires_at.gte", TODAY, 1), + ("expires_at.lte", TODAY, 1), + ("expires_at", TODAY, 0), + ("email", "testing@prowler.com", 2), + ("email.icontains", "testing", 2), + ("inviter", "", 2), + ] + ), + ) + def test_invitations_filters( + self, + authenticated_client, + create_test_user, + invitations_fixture, + filter_name, + filter_value, + expected_count, + ): + user = create_test_user + response = authenticated_client.get( + reverse("invitation-list"), + { + f"filter[{filter_name}]": filter_value + if filter_name != "inviter" + else str(user.id) + }, + ) + + assert response.status_code == status.HTTP_200_OK + assert len(response.json()["data"]) == expected_count + + def test_invitations_list_filter_invalid(self, authenticated_client): + response = authenticated_client.get( + reverse("invitation-list"), + {"filter[invalid]": "whatever"}, + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + + @pytest.mark.parametrize( + "sort_field", + [ + "inserted_at", + "updated_at", + "expires_at", + "state", + "inviter", + ], + ) + def test_invitations_sort(self, authenticated_client, sort_field): + response = authenticated_client.get( + reverse("invitation-list"), + {"sort": sort_field}, + ) + assert response.status_code == status.HTTP_200_OK + + def test_invitations_sort_invalid(self, authenticated_client): + response = authenticated_client.get( + reverse("invitation-list"), + {"sort": "invalid"}, + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST diff --git a/src/backend/api/utils.py b/src/backend/api/utils.py index 8565f619b6..1f60669039 100644 --- a/src/backend/api/utils.py +++ b/src/backend/api/utils.py @@ -1,10 +1,15 @@ +from datetime import datetime, timezone + from prowler.providers.aws.aws_provider import AwsProvider from prowler.providers.azure.azure_provider import AzureProvider from prowler.providers.common.models import Connection from prowler.providers.gcp.gcp_provider import GcpProvider from prowler.providers.kubernetes.kubernetes_provider import KubernetesProvider +from rest_framework.exceptions import ValidationError, NotFound -from api.models import Provider +from api.db_router import MainRouter +from api.exceptions import InvitationTokenExpiredException +from api.models import Provider, Invitation def merge_dicts(default_dict: dict, replacement_dict: dict) -> dict: @@ -119,3 +124,66 @@ def prowler_provider_connection_test(provider: Provider) -> Connection: return prowler_provider.test_connection( **prowler_provider_kwargs, provider_id=provider.uid, raise_on_exception=False ) + + +def validate_invitation( + invitation_token: str, email: str, raise_not_found=False +) -> Invitation: + """ + Validates an invitation based on the provided token and email. + + This function attempts to retrieve an Invitation object using the given + `invitation_token` and `email`. It performs several checks to ensure that + the invitation is valid, not expired, and in the correct state for acceptance. + + Args: + invitation_token (str): The token associated with the invitation. + email (str): The email address associated with the invitation. + raise_not_found (bool, optional): If True, raises a `NotFound` exception + when the invitation is not found. If False, raises a `ValidationError`. + Defaults to False. + + Returns: + Invitation: The validated Invitation object. + + Raises: + NotFound: If `raise_not_found` is True and the invitation does not exist. + ValidationError: If the invitation does not exist and `raise_not_found` + is False, or if the invitation is invalid or in an incorrect state. + InvitationTokenExpiredException: If the invitation has expired. + + Notes: + - This function uses the admin database connector to bypass RLS protection + since the invitation may belong to a tenant the user is not a member of yet. + - If the invitation has expired, its state is updated to EXPIRED, and an + `InvitationTokenExpiredException` is raised. + - Only invitations in the PENDING state can be accepted. + + Examples: + invitation = validate_invitation("TOKEN123", "user@example.com") + """ + try: + # Admin DB connector is used to bypass RLS protection since the invitation belongs to a tenant the user + # is not a member of yet + invitation = Invitation.objects.using(MainRouter.admin_db).get( + token=invitation_token, email=email + ) + except Invitation.DoesNotExist: + if raise_not_found: + raise NotFound(detail="Invitation is not valid.") + else: + raise ValidationError({"invitation_token": "Invalid invitation code."}) + + # Check if the invitation has expired + if invitation.expires_at < datetime.now(timezone.utc): + invitation.state = Invitation.State.EXPIRED + invitation.save(using=MainRouter.admin_db) + raise InvitationTokenExpiredException() + + # Check the state of the invitation + if invitation.state != Invitation.State.PENDING: + raise ValidationError( + {"invitation_token": "This invitation is no longer valid."} + ) + + return invitation diff --git a/src/backend/api/v1/serializers.py b/src/backend/api/v1/serializers.py index 827bee410c..c28fb2e01b 100644 --- a/src/backend/api/v1/serializers.py +++ b/src/backend/api/v1/serializers.py @@ -1,4 +1,5 @@ import json +from datetime import datetime, timezone, timedelta from django.conf import settings from django.contrib.auth import authenticate @@ -9,8 +10,8 @@ from jwt.exceptions import InvalidKeyError from rest_framework_json_api import serializers from rest_framework_json_api.serializers import ValidationError from rest_framework_simplejwt.exceptions import TokenError -from rest_framework_simplejwt.tokens import RefreshToken from rest_framework_simplejwt.serializers import TokenObtainPairSerializer +from rest_framework_simplejwt.tokens import RefreshToken from api.models import ( StateChoices, @@ -23,6 +24,7 @@ from api.models import ( ResourceTag, Finding, ProviderSecret, + Invitation, ) from api.rls import Tenant from api.utils import merge_dicts @@ -881,6 +883,7 @@ class ProviderSecretUpdateSerializer(BaseWriteProviderSecretSerializer): class Meta: model = ProviderSecret fields = [ + "id", "inserted_at", "updated_at", "name", @@ -903,3 +906,99 @@ class ProviderSecretUpdateSerializer(BaseWriteProviderSecretSerializer): validated_attrs = super().validate(attrs) self.validate_secret_based_on_provider(provider.provider, secret_type, secret) return validated_attrs + + +# Invitations + + +class InvitationSerializer(RLSSerializer): + """ + Serializer for the Invitation model. + """ + + class Meta: + model = Invitation + fields = [ + "id", + "inserted_at", + "updated_at", + "email", + "state", + "token", + "expires_at", + "inviter", + "url", + ] + + +class InvitationBaseWriteSerializer(BaseWriteSerializer): + def validate_email(self, value): + user = User.objects.filter(email=value).first() + tenant_id = self.context["tenant_id"] + if user and Membership.objects.filter(user=user, tenant=tenant_id).exists(): + raise ValidationError( + "The user may already be a member of the tenant or there was an issue with the " + "email provided." + ) + if Invitation.objects.filter( + email=value, state=Invitation.State.PENDING + ).exists(): + raise ValidationError( + "Unable to process your request. Please check the information provided and " + "try again." + ) + return value + + def validate_expires_at(self, value): + now = datetime.now(timezone.utc) + if value and value < now + timedelta(hours=24): + raise ValidationError( + "Expiry date must be at least 24 hours in the future." + ) + return value + + +class InvitationCreateSerializer(InvitationBaseWriteSerializer, RLSSerializer): + expires_at = serializers.DateTimeField( + required=False, + help_text="UTC. Default 7 days. If this attribute is " + "provided, it must be at least 24 hours in the " + "future.", + ) + + class Meta: + model = Invitation + fields = ["email", "expires_at", "state", "token", "inviter"] + extra_kwargs = { + "token": {"read_only": True}, + "state": {"read_only": True}, + "inviter": {"read_only": True}, + "expires_at": {"required": False}, + } + + def create(self, validated_data): + inviter = self.context.get("request").user + validated_data["inviter"] = inviter + return super().create(validated_data) + + +class InvitationUpdateSerializer(InvitationBaseWriteSerializer): + class Meta: + model = Invitation + fields = ["id", "email", "expires_at", "state", "token"] + extra_kwargs = { + "token": {"read_only": True}, + "state": {"read_only": True}, + "expires_at": {"required": False}, + "email": {"required": False}, + } + + +class InvitationAcceptSerializer(RLSSerializer): + """Serializer for accepting an invitation.""" + + invitation_token = serializers.CharField(write_only=True) + + class Meta: + model = Invitation + fields = ["invitation_token"] diff --git a/src/backend/api/v1/urls.py b/src/backend/api/v1/urls.py index 448668149e..dc8f5f07a6 100644 --- a/src/backend/api/v1/urls.py +++ b/src/backend/api/v1/urls.py @@ -16,6 +16,8 @@ from api.v1.views import ( ResourceViewSet, FindingViewSet, ProviderSecretViewSet, + InvitationViewSet, + InvitationAcceptViewSet, ) router = routers.DefaultRouter(trailing_slash=False) @@ -23,7 +25,6 @@ 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"providers/secrets", ProviderSecretViewSet, basename="provider-secret") router.register(r"scans", ScanViewSet, basename="scan") router.register(r"tasks", TaskViewSet, basename="task") router.register(r"resources", ResourceViewSet, basename="resource") @@ -52,6 +53,23 @@ urlpatterns = [ ), name="providersecret-detail", ), + path( + "tenants/invitations", + InvitationViewSet.as_view({"get": "list", "post": "create"}), + name="invitation-list", + ), + path( + "tenants/invitations/", + InvitationViewSet.as_view( + {"get": "retrieve", "patch": "partial_update", "delete": "destroy"} + ), + name="invitation-detail", + ), + path( + "invitations/accept", + InvitationAcceptViewSet.as_view({"post": "accept"}), + name="invitation-accept", + ), path("", include(router.urls)), path("", include(tenants_router.urls)), path("", include(users_router.urls)), diff --git a/src/backend/api/v1/views.py b/src/backend/api/v1/views.py index c1db0b04c6..3fbb0c2796 100644 --- a/src/backend/api/v1/views.py +++ b/src/backend/api/v1/views.py @@ -18,13 +18,19 @@ from drf_spectacular.utils import ( from drf_spectacular.views import SpectacularAPIView from rest_framework import status, permissions from rest_framework.decorators import action -from rest_framework.exceptions import MethodNotAllowed, NotFound, PermissionDenied +from rest_framework.exceptions import ( + MethodNotAllowed, + NotFound, + PermissionDenied, + ValidationError, +) from rest_framework.generics import get_object_or_404, GenericAPIView from rest_framework_json_api.views import Response from rest_framework_simplejwt.exceptions import InvalidToken from rest_framework_simplejwt.exceptions import TokenError from api.base_views import BaseTenantViewset, BaseRLSViewSet, BaseViewSet +from api.db_router import MainRouter from api.filters import ( ProviderFilter, TenantFilter, @@ -34,6 +40,7 @@ from api.filters import ( ResourceFilter, FindingFilter, ProviderSecretFilter, + InvitationFilter, ) from api.models import ( User, @@ -44,8 +51,10 @@ from api.models import ( Resource, Finding, ProviderSecret, + Invitation, ) from api.rls import Tenant +from api.utils import validate_invitation from api.uuid_utils import datetime_to_uuid7 from api.v1.serializers import ( TokenSerializer, @@ -67,6 +76,10 @@ from api.v1.serializers import ( ProviderSecretSerializer, ProviderSecretUpdateSerializer, ProviderSecretCreateSerializer, + InvitationSerializer, + InvitationCreateSerializer, + InvitationUpdateSerializer, + InvitationAcceptSerializer, ) from tasks.tasks import ( check_provider_connection_task, @@ -173,6 +186,11 @@ class SchemaView(SpectacularAPIView): "description": "Endpoints for task management, allowing retrieval of task status and " "revoking tasks that have not started.", }, + { + "name": "Invitation", + "description": "Endpoints for tenant invitations management, allowing retrieval and filtering of " + "invitations, creating new invitations, accepting and revoking them.", + }, ] return super().get(request, *args, **kwargs) @@ -242,18 +260,53 @@ class UserViewSet(BaseViewSet): status=status.HTTP_200_OK, ) + @extend_schema( + parameters=[ + OpenApiParameter( + name="invitation_token", + description="Optional invitation code for joining an existing tenant.", + required=False, + type={"type": "string", "example": "F3NMFPNDZHR4Z9"}, + location=OpenApiParameter.QUERY, + ), + ] + ) def create(self, request, *args, **kwargs): + invitation_token = request.query_params.get("invitation_token", None) + invitation = None + serializer = self.get_serializer( data=request.data, context=self.get_serializer_context() ) serializer.is_valid(raise_exception=True) - user = serializer.save() - tenant = Tenant.objects.create( - name=f"{user.email.split('@')[0]} default tenant" + + if invitation_token: + invitation = validate_invitation( + invitation_token, serializer.validated_data["email"] + ) + + # Proceed with creating the user and membership + user = User.objects.db_manager(MainRouter.admin_db).create_user( + **serializer.validated_data ) - Membership.objects.create( - user=user, tenant=tenant, role=Membership.RoleChoices.OWNER + tenant = ( + invitation.tenant + if invitation_token + else Tenant.objects.using(MainRouter.admin_db).create( + name=f"{user.email.split('@')[0]} default tenant" + ) ) + role = ( + Membership.RoleChoices.MEMBER + if invitation_token + else Membership.RoleChoices.OWNER + ) + Membership.objects.using(MainRouter.admin_db).create( + user=user, tenant=tenant, role=role + ) + if invitation: + invitation.state = Invitation.State.ACCEPTED + invitation.save(using=MainRouter.admin_db) return Response(data=UserSerializer(user).data, status=status.HTTP_201_CREATED) def partial_update(self, request, *args, **kwargs): @@ -621,7 +674,7 @@ class ScanViewSet(BaseRLSViewSet): tenant_id=request.tenant_id, scan_id=str(scan.id), provider_id=str(scan.provider_id), - checks_to_execute=scan.scanner_args.get("checks_to_execute", []), + checks_to_execute=scan.scanner_args.get("checks_to_execute"), ) scan.task_id = task.id @@ -886,3 +939,143 @@ class ProviderSecretViewSet(BaseRLSViewSet): elif self.action == "partial_update": return ProviderSecretUpdateSerializer return super().get_serializer_class() + + +@extend_schema_view( + list=extend_schema( + tags=["Invitation"], + summary="List all invitations", + description="Retrieve a list of all tenant invitations with options for filtering by various criteria.", + ), + retrieve=extend_schema( + tags=["Invitation"], + summary="Retrieve data from a tenant invitation", + description="Fetch detailed information about a specific invitation by its ID.", + ), + create=extend_schema( + tags=["Invitation"], + summary="Invite a user to a tenant", + description="Add a new tenant invitation to the system by providing the required invitation details. The " + "invited user will have to accept the invitations or create an account using the given code.", + ), + partial_update=extend_schema( + tags=["Invitation"], + summary="Partially update a tenant invitation", + description="Update certain fields of an existing tenant invitation's information without affecting other " + "fields.", + ), + destroy=extend_schema( + tags=["Invitation"], + summary="Revoke a tenant invitation", + description="Revoke a tenant invitation from the system by their ID.", + ), +) +@method_decorator(CACHE_DECORATOR, name="list") +@method_decorator(CACHE_DECORATOR, name="retrieve") +class InvitationViewSet(BaseRLSViewSet): + queryset = Invitation.objects.all() + serializer_class = InvitationSerializer + filterset_class = InvitationFilter + http_method_names = ["get", "post", "patch", "delete"] + search_fields = ["email"] + ordering = ["-inserted_at"] + ordering_fields = [ + "inserted_at", + "updated_at", + "expires_at", + "state", + "inviter", + ] + + def get_queryset(self): + return Invitation.objects.all() + + def get_serializer_class(self): + if self.action == "create": + return InvitationCreateSerializer + elif self.action == "partial_update": + return InvitationUpdateSerializer + return super().get_serializer_class() + + def create(self, request, *args, **kwargs): + serializer = self.get_serializer( + data=request.data, + context={"tenant_id": self.request.tenant_id, "request": request}, + ) + serializer.is_valid(raise_exception=True) + serializer.save() + return Response(data=serializer.data, status=status.HTTP_201_CREATED) + + def partial_update(self, request, *args, **kwargs): + instance = self.get_object() + if instance.state != Invitation.State.PENDING: + raise ValidationError(detail="This invitation cannot be updated.") + serializer = self.get_serializer( + instance, + data=request.data, + partial=True, + context={"tenant_id": self.request.tenant_id, "request": request}, + ) + serializer.is_valid(raise_exception=True) + serializer.save() + return Response(data=serializer.data, status=status.HTTP_200_OK) + + def destroy(self, request, *args, **kwargs): + instance = self.get_object() + if instance.state != Invitation.State.PENDING: + raise ValidationError(detail="This invitation cannot be revoked.") + instance.state = Invitation.State.REVOKED + instance.save() + return Response(status=status.HTTP_204_NO_CONTENT) + + +class InvitationAcceptViewSet(BaseRLSViewSet): + queryset = Invitation.objects.all() + serializer_class = InvitationAcceptSerializer + http_method_names = ["post"] + + def get_queryset(self): + return Invitation.objects.all() + + def get_serializer_class(self): + if hasattr(self, "response_serializer_class"): + return self.response_serializer_class + return InvitationAcceptSerializer + + @extend_schema(exclude=True) + def create(self, request, *args, **kwargs): + raise MethodNotAllowed(method="POST") + + @extend_schema( + tags=["Invitation"], + summary="Accept an invitation", + description="Accept an invitation to an existing tenant. This invitation cannot be expired and the emails must " + "match.", + responses={201: OpenApiResponse(response=MembershipSerializer)}, + ) + @action(detail=False, methods=["post"], url_name="accept") + def accept(self, request): + serializer = self.get_serializer( + data=request.data, + context=self.get_serializer_context(), + ) + serializer.is_valid(raise_exception=True) + invitation_token = serializer.validated_data["invitation_token"] + user_email = request.user.email + + invitation = validate_invitation( + invitation_token, user_email, raise_not_found=True + ) + + # Proceed with accepting the invitation + user = User.objects.using(MainRouter.admin_db).get(email=user_email) + membership = Membership.objects.using(MainRouter.admin_db).create( + user=user, + tenant=invitation.tenant, + ) + invitation.state = Invitation.State.ACCEPTED + invitation.save(using=MainRouter.admin_db) + + self.response_serializer_class = MembershipSerializer + membership_serializer = self.get_serializer(membership) + return Response(data=membership_serializer.data, status=status.HTTP_201_CREATED) diff --git a/src/backend/conftest.py b/src/backend/conftest.py index 976b75ede4..d62ea275b4 100644 --- a/src/backend/conftest.py +++ b/src/backend/conftest.py @@ -2,6 +2,7 @@ import logging import pytest from django.conf import settings +from datetime import datetime, timezone, timedelta from django.db import connections as django_connections, connection as django_connection from django.urls import reverse from django_celery_results.models import TaskResult @@ -23,6 +24,7 @@ from api.models import ( Task, Membership, ProviderSecret, + Invitation, ) from api.rls import Tenant from api.v1.serializers import TokenSerializer @@ -121,12 +123,37 @@ def tenants_fixture(create_test_user): tenant=tenant2, role=Membership.RoleChoices.OWNER, ) - return tenant1, tenant2 + tenant3 = Tenant.objects.create( + name="Tenant Three", + ) + return tenant1, tenant2, tenant3 + + +@pytest.fixture +def invitations_fixture(create_test_user, tenants_fixture): + user = create_test_user + *_, tenant = tenants_fixture + valid_invitation = Invitation.objects.create( + email="testing@prowler.com", + state=Invitation.State.PENDING, + token="TESTING1234567", + inviter=user, + tenant=tenant, + ) + expired_invitation = Invitation.objects.create( + email="testing@prowler.com", + state=Invitation.State.EXPIRED, + token="TESTING1234568", + expires_at=datetime.now(timezone.utc) - timedelta(days=1), + inviter=user, + tenant=tenant, + ) + return valid_invitation, expired_invitation @pytest.fixture def providers_fixture(tenants_fixture): - tenant, _ = tenants_fixture + tenant, *_ = tenants_fixture provider1 = Provider.objects.create( provider="aws", uid="123456789012", @@ -178,7 +205,7 @@ def provider_secret_fixture(providers_fixture): @pytest.fixture def scans_fixture(tenants_fixture, providers_fixture): - tenant, _ = tenants_fixture + tenant, *_ = tenants_fixture provider, provider2, *_ = providers_fixture scan1 = Scan.objects.create( @@ -210,7 +237,7 @@ def scans_fixture(tenants_fixture, providers_fixture): @pytest.fixture def tasks_fixture(tenants_fixture): - tenant, _ = tenants_fixture + tenant, *_ = tenants_fixture task_runner_task1 = TaskResult.objects.create( task_id="81a1b34b-ff6e-498e-979c-d6a83260167f", diff --git a/src/backend/tasks/tests/test_scan.py b/src/backend/tasks/tests/test_scan.py index 23509233f8..c3eae5e40d 100644 --- a/src/backend/tasks/tests/test_scan.py +++ b/src/backend/tasks/tests/test_scan.py @@ -27,7 +27,7 @@ class TestPerformScan: assert len(Finding.objects.all()) == 0 assert len(Resource.objects.all()) == 0 - tenant, _ = tenants_fixture + tenant, *_ = tenants_fixture scan, *_ = scans_fixture provider, *_ = providers_fixture @@ -89,7 +89,7 @@ class TestPerformScan: scans_fixture, providers_fixture, ): - tenant, _ = tenants_fixture + tenant, *_ = tenants_fixture scan, *_ = scans_fixture provider, *_ = providers_fixture